Coverage for src/lexigram/admin/controllers/progress.py: 95%
76 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-21 14:56 +0800
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-21 14:56 +0800
1"""Progress tracking controller for SSE-based real-time updates."""
3from __future__ import annotations
5import asyncio
6from collections.abc import AsyncIterator
7from typing import TYPE_CHECKING, Any
9from starlette.requests import Request
10from starlette.responses import StreamingResponse
12from lexigram.admin.controllers.base import AdminController
13from lexigram.contracts.infra.tasks.progress import (
14 ProgressSnapshot,
15 ProgressStatus,
16 ProgressTrackerProtocol,
17)
18from lexigram.contracts.web import get
19from lexigram.di.decorators import inject
20from lexigram.serialization import dumps_str
22if TYPE_CHECKING:
23 from lexigram.admin.engine.renderer import AdminRenderer
26def _snap_to_dict(snap: ProgressSnapshot) -> dict[str, Any]:
27 """Convert a :class:`ProgressSnapshot` to a JSON-safe response dict."""
28 return {
29 "id": snap.task_id,
30 "status": snap.status.value,
31 "progress": snap.percent,
32 "current": snap.current,
33 "total": snap.total,
34 "message": snap.message,
35 "error": snap.error or None,
36 }
39class LocalProgressTracker:
40 """In-process :class:`ProgressTrackerProtocol` implementation owned by lexigram-admin.
42 Used as the DI-resolution fallback when no integrator (e.g. lexigram-tasks)
43 has registered a real tracker — keeps :class:`ProgressController` mountable
44 without a direct import of any sibling package. State is process-local and
45 lost on restart, matching the previous `InMemoryProgressTracker` fallback's
46 behavior.
47 """
49 def __init__(self) -> None:
50 self._snapshots: dict[str, ProgressSnapshot] = {}
51 self._subscribers: dict[str, list[asyncio.Queue[ProgressSnapshot]]] = {}
53 async def update(
54 self, task_id: str, current: int, total: int, message: str = ""
55 ) -> None:
56 await self._publish(
57 ProgressSnapshot(
58 task_id=task_id,
59 current=current,
60 total=total,
61 status=ProgressStatus.RUNNING,
62 message=message,
63 )
64 )
66 async def complete(self, task_id: str, result: str = "") -> None:
67 prev = self._snapshots.get(task_id)
68 await self._publish(
69 ProgressSnapshot(
70 task_id=task_id,
71 current=prev.current if prev else 0,
72 total=prev.total if prev else 0,
73 status=ProgressStatus.COMPLETE,
74 message=result,
75 )
76 )
78 async def fail(self, task_id: str, error: str) -> None:
79 prev = self._snapshots.get(task_id)
80 await self._publish(
81 ProgressSnapshot(
82 task_id=task_id,
83 current=prev.current if prev else 0,
84 total=prev.total if prev else 0,
85 status=ProgressStatus.FAILED,
86 error=error,
87 )
88 )
90 async def get(self, task_id: str) -> ProgressSnapshot | None:
91 return self._snapshots.get(task_id)
93 async def subscribe(self, task_id: str) -> AsyncIterator[ProgressSnapshot]:
94 existing = self._snapshots.get(task_id)
95 if existing is not None and existing.status in (
96 ProgressStatus.COMPLETE,
97 ProgressStatus.FAILED,
98 ):
99 yield existing
100 return
102 queue: asyncio.Queue[ProgressSnapshot] = asyncio.Queue()
103 self._subscribers.setdefault(task_id, []).append(queue)
104 try:
105 while True:
106 snap = await queue.get()
107 yield snap
108 if snap.status in (ProgressStatus.COMPLETE, ProgressStatus.FAILED):
109 return
110 finally:
111 self._subscribers.get(task_id, []).remove(queue)
113 async def _publish(self, snap: ProgressSnapshot) -> None:
114 self._snapshots[snap.task_id] = snap
115 for queue in self._subscribers.get(snap.task_id, []):
116 await queue.put(snap)
119@inject
120class ProgressController(AdminController):
121 """Controller for progress tracking endpoints.
123 Exposes SSE streaming and point-in-time status queries backed by
124 :class:`ProgressTrackerProtocol`. Inject :class:`LocalProgressTracker`
125 (or any conforming implementation) via the DI container.
126 """
128 def __init__(
129 self,
130 tracker: ProgressTrackerProtocol,
131 renderer: AdminRenderer | None = None,
132 ) -> None:
133 if renderer is None:
134 from lexigram.admin.engine.renderer import AdminRenderer
136 renderer = AdminRenderer()
137 super().__init__(renderer=renderer)
138 self.tracker = tracker
140 @get("/progress/{task_id}/stream")
141 async def stream_progress(self, request: Request) -> StreamingResponse:
142 """Stream progress updates via Server-Sent Events.
144 The stream uses :meth:`ProgressTrackerProtocol.subscribe` so
145 updates are pushed immediately rather than polled. The generator
146 closes automatically when the task reaches a terminal state.
148 Args:
149 request: Starlette request. The task id comes from the route.
151 Returns:
152 SSE stream of progress snapshots.
153 """
155 async def event_generator() -> Any:
156 try:
157 task_id = request.path_params["task_id"]
158 # Guard: emit an error event and stop if the task is unknown.
159 snap = await self.tracker.get(task_id)
160 if snap is None:
161 yield (
162 f"event: error\ndata: "
163 f"{dumps_str({'error': 'Task not found'})}\n\n"
164 )
165 return
167 # subscribe() yields until terminal state then closes.
168 async for current_snap in self.tracker.subscribe(task_id):
169 yield (
170 f"event: progress\ndata: "
171 f"{dumps_str(_snap_to_dict(current_snap))}\n\n"
172 )
174 except asyncio.CancelledError:
175 # Client disconnected — clean exit.
176 pass
177 except (RuntimeError, ValueError, TypeError, OSError) as exc:
178 yield f"event: error\ndata: {dumps_str({'error': str(exc)})}\n\n"
180 return StreamingResponse(
181 event_generator(),
182 media_type="text/event-stream",
183 headers={
184 "Cache-Control": "no-cache",
185 "X-Accel-Buffering": "no", # Disable nginx buffering
186 },
187 )
189 @get("/progress/{task_id}")
190 async def get_task_status(
191 self,
192 request: Request,
193 ) -> dict[str, Any] | tuple[dict[str, str], int]:
194 """Return current task status as a JSON dict.
196 Args:
197 request: Starlette request. The task id comes from the route.
199 Returns:
200 Progress dict or a 404 error dict.
201 """
202 task_id = request.path_params["task_id"]
203 snap = await self.tracker.get(task_id)
204 if snap is None:
205 return {"error": "Task not found"}, 404
206 return _snap_to_dict(snap)