174 lines
5.9 KiB
Python
174 lines
5.9 KiB
Python
"""In-process SSE patch broadcaster for live terminal updates."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
import threading
|
|
import time
|
|
from collections import defaultdict
|
|
from typing import Any, AsyncIterator, DefaultDict, Iterable, Optional, Set
|
|
|
|
|
|
HEARTBEAT_INTERVAL_SECONDS = 30
|
|
QUEUE_MAXSIZE = 256
|
|
|
|
|
|
class SseManager:
|
|
def __init__(self) -> None:
|
|
self._queues: DefaultDict[str, set[asyncio.Queue[dict[str, Any]]]] = defaultdict(set)
|
|
self._queue_cities: dict[int, frozenset[str]] = {}
|
|
self._queue_loops: dict[int, asyncio.AbstractEventLoop] = {}
|
|
self._lock = threading.RLock()
|
|
self._revision = 0
|
|
|
|
def _next_revision(self) -> int:
|
|
with self._lock:
|
|
self._revision += 1
|
|
return self._revision
|
|
|
|
@staticmethod
|
|
def _normalize_city(value: Any) -> str:
|
|
return str(value or "").strip().lower()
|
|
|
|
@classmethod
|
|
def _normalize_city_set(cls, cities: Optional[Iterable[str]]) -> Set[str]:
|
|
return {
|
|
cls._normalize_city(city)
|
|
for city in (cities or [])
|
|
if cls._normalize_city(city)
|
|
}
|
|
|
|
def _track_revision(self, event: dict[str, Any]) -> None:
|
|
try:
|
|
revision = int(event.get("revision") or 0)
|
|
except (TypeError, ValueError):
|
|
return
|
|
if revision <= 0:
|
|
return
|
|
with self._lock:
|
|
if revision > self._revision:
|
|
self._revision = revision
|
|
|
|
def connection_count(self) -> int:
|
|
with self._lock:
|
|
return sum(len(queue_set) for queue_set in self._queues.values())
|
|
|
|
def broadcast(self, city: str, changes: dict[str, Any]) -> dict[str, Any]:
|
|
event = {
|
|
"type": "city_patch",
|
|
"city": self._normalize_city(city),
|
|
"changes": changes or {},
|
|
"revision": self._next_revision(),
|
|
"ts": int(time.time() * 1000),
|
|
}
|
|
return self.broadcast_event(event)
|
|
|
|
def broadcast_event(self, event: dict[str, Any]) -> dict[str, Any]:
|
|
city = self._normalize_city(event.get("city"))
|
|
if city:
|
|
event = {**event, "city": city}
|
|
self._track_revision(event)
|
|
if not city:
|
|
return event
|
|
|
|
with self._lock:
|
|
queue_items = [
|
|
(
|
|
queue,
|
|
self._queue_cities.get(id(queue), frozenset()),
|
|
self._queue_loops.get(id(queue)),
|
|
)
|
|
for queue_set in self._queues.values()
|
|
for queue in queue_set
|
|
]
|
|
|
|
for queue, subscribed_cities, loop in queue_items:
|
|
if subscribed_cities and city not in subscribed_cities:
|
|
continue
|
|
if loop and loop.is_running():
|
|
loop.call_soon_threadsafe(self._put_queue_event, queue, event)
|
|
continue
|
|
self._put_queue_event(queue, event)
|
|
return event
|
|
|
|
@staticmethod
|
|
def _put_queue_event(queue: asyncio.Queue[dict[str, Any]], event: dict[str, Any]) -> None:
|
|
try:
|
|
queue.put_nowait(event)
|
|
except asyncio.QueueFull:
|
|
try:
|
|
queue.get_nowait()
|
|
except asyncio.QueueEmpty:
|
|
pass
|
|
try:
|
|
queue.put_nowait(event)
|
|
except asyncio.QueueFull:
|
|
pass
|
|
|
|
async def event_stream(
|
|
self,
|
|
user_id: str,
|
|
*,
|
|
cities: Optional[Iterable[str]] = None,
|
|
replay_events: Optional[Iterable[dict[str, Any]]] = None,
|
|
connected_revision: Optional[int] = None,
|
|
resync_event: Optional[dict[str, Any]] = None,
|
|
) -> AsyncIterator[str]:
|
|
user_key = str(user_id or "anon")
|
|
city_set = frozenset(self._normalize_city_set(cities))
|
|
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue(maxsize=QUEUE_MAXSIZE)
|
|
loop = asyncio.get_running_loop()
|
|
with self._lock:
|
|
self._queues[user_key].add(queue)
|
|
self._queue_cities[id(queue)] = city_set
|
|
self._queue_loops[id(queue)] = loop
|
|
if connected_revision is not None:
|
|
self._revision = max(self._revision, int(connected_revision or 0))
|
|
|
|
try:
|
|
yield self._format_event({
|
|
"type": "connected",
|
|
"revision": self._revision,
|
|
"cities": sorted(city_set),
|
|
"ts": int(time.time() * 1000),
|
|
})
|
|
for event in replay_events or []:
|
|
self._track_revision(event)
|
|
yield self._format_event(event)
|
|
if resync_event:
|
|
self._track_revision({"revision": resync_event.get("latest_revision")})
|
|
yield self._format_event(resync_event)
|
|
while True:
|
|
try:
|
|
event = await asyncio.wait_for(
|
|
queue.get(),
|
|
timeout=HEARTBEAT_INTERVAL_SECONDS,
|
|
)
|
|
except asyncio.TimeoutError:
|
|
event = {
|
|
"type": "heartbeat",
|
|
"revision": self._revision,
|
|
"ts": int(time.time() * 1000),
|
|
}
|
|
yield self._format_event(event)
|
|
finally:
|
|
with self._lock:
|
|
self._queues[user_key].discard(queue)
|
|
self._queue_cities.pop(id(queue), None)
|
|
self._queue_loops.pop(id(queue), None)
|
|
if not self._queues[user_key]:
|
|
self._queues.pop(user_key, None)
|
|
|
|
@staticmethod
|
|
def _format_event(event: dict[str, Any]) -> str:
|
|
if str(event.get("type") or "").startswith("city_observation_patch"):
|
|
event = {
|
|
**event,
|
|
"sse_emitted_at_ms": int(time.time() * 1000),
|
|
}
|
|
return f"data: {json.dumps(event, ensure_ascii=False, separators=(',', ':'))}\n\n"
|
|
|
|
|
|
sse_manager = SseManager()
|