feat(server): 新增 ProgressHub 进程内进度发布/订阅单例
This commit is contained in:
@@ -0,0 +1,44 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Any, Dict, List
|
||||
|
||||
_event = Dict[str, Any]
|
||||
|
||||
|
||||
class ProgressHub:
|
||||
"""进程内会话级进度发布/订阅(单例)。
|
||||
|
||||
- subscribe(sid) 返回专属 asyncio.Queue;emit(sid, event) 向该 sid 全部队列投递。
|
||||
- emit 从同步线程(FastAPI 线程池中的 sync 端点)调用,经由已注册事件循环
|
||||
run_coroutine_threadsafe 安全投递;未注册 loop 时降级为直接放入队列(同线程场景)。
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._loop: asyncio.AbstractEventLoop | None = None
|
||||
self._subs: Dict[str, List[asyncio.Queue]] = {}
|
||||
|
||||
def register_loop(self, loop: asyncio.AbstractEventLoop) -> None:
|
||||
self._loop = loop
|
||||
|
||||
def subscribe(self, sid: str) -> asyncio.Queue:
|
||||
q: asyncio.Queue = asyncio.Queue()
|
||||
self._subs.setdefault(sid, []).append(q)
|
||||
return q
|
||||
|
||||
def unsubscribe(self, sid: str, q: asyncio.Queue) -> None:
|
||||
qs = self._subs.get(sid)
|
||||
if qs and q in qs:
|
||||
qs.remove(q)
|
||||
if not qs:
|
||||
self._subs.pop(sid, None)
|
||||
|
||||
def emit(self, sid: str, event: _event) -> None:
|
||||
for q in list(self._subs.get(sid, [])):
|
||||
if self._loop is not None:
|
||||
asyncio.run_coroutine_threadsafe(q.put(event), self._loop)
|
||||
else:
|
||||
q.put_nowait(event)
|
||||
|
||||
|
||||
hub = ProgressHub()
|
||||
@@ -0,0 +1,35 @@
|
||||
import asyncio
|
||||
import pytest
|
||||
from genesis.server.hub import ProgressHub
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subscribe_receives_emitted_event():
|
||||
h = ProgressHub()
|
||||
h.register_loop(asyncio.get_running_loop())
|
||||
q = h.subscribe("s1")
|
||||
h.emit("s1", {"type": "progress", "step": "parse", "status": "ok"})
|
||||
event = await asyncio.wait_for(q.get(), 1.0)
|
||||
assert event["step"] == "parse"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unsubscribe_stops_delivery():
|
||||
h = ProgressHub()
|
||||
h.register_loop(asyncio.get_running_loop())
|
||||
q = h.subscribe("s1")
|
||||
h.unsubscribe("s1", q)
|
||||
h.emit("s1", {"type": "progress"})
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
await asyncio.wait_for(q.get(), 0.2)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_subscribers_all_receive():
|
||||
h = ProgressHub()
|
||||
h.register_loop(asyncio.get_running_loop())
|
||||
q1, q2 = h.subscribe("s1"), h.subscribe("s1")
|
||||
h.emit("s1", {"type": "progress", "step": "gen"})
|
||||
a = await asyncio.wait_for(q1.get(), 1.0)
|
||||
b = await asyncio.wait_for(q2.get(), 1.0)
|
||||
assert a["step"] == b["step"] == "gen"
|
||||
Reference in New Issue
Block a user