Coverage for src / lexigram / admin / controllers / progress.py: 95%
73 statements
« prev ^ index » next coverage.py v7.13.5, created at 2026-08-13 22:14 +0800
« prev ^ index » next coverage.py v7.13.5, created at 2026-08-13 22:14 +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 super().__init__(renderer=renderer)
134 self.tracker = tracker
136 @get("/progress/{task_id}/stream")
137 async def stream_progress(self, request: Request) -> StreamingResponse:
138 """Stream progress updates via Server-Sent Events.
140 The stream uses :meth:`ProgressTrackerProtocol.subscribe` so
141 updates are pushed immediately rather than polled. The generator
142 closes automatically when the task reaches a terminal state.
144 Args:
145 request: Starlette request. The task id comes from the route.
147 Returns:
148 SSE stream of progress snapshots.
149 """
151 async def event_generator() -> Any:
152 try:
153 task_id = request.path_params["task_id"]
154 # Guard: emit an error event and stop if the task is unknown.
155 snap = await self.tracker.get(task_id)
156 if snap is None:
157 yield (
158 f"event: error\ndata: "
159 f"{dumps_str({'error': 'Task not found'})}\n\n"
160 )
161 return
163 # subscribe() yields until terminal state then closes.
164 async for current_snap in self.tracker.subscribe(task_id):
165 yield (
166 f"event: progress\ndata: "
167 f"{dumps_str(_snap_to_dict(current_snap))}\n\n"
168 )
170 except asyncio.CancelledError:
171 # Client disconnected — clean exit.
172 pass
173 except (RuntimeError, ValueError, TypeError, OSError) as exc:
174 yield f"event: error\ndata: {dumps_str({'error': str(exc)})}\n\n"
176 return StreamingResponse(
177 event_generator(),
178 media_type="text/event-stream",
179 headers={
180 "Cache-Control": "no-cache",
181 "X-Accel-Buffering": "no", # Disable nginx buffering
182 },
183 )
185 @get("/progress/{task_id}")
186 async def get_task_status(
187 self,
188 request: Request,
189 ) -> dict[str, Any] | tuple[dict[str, str], int]:
190 """Return current task status as a JSON dict.
192 Args:
193 request: Starlette request. The task id comes from the route.
195 Returns:
196 Progress dict or a 404 error dict.
197 """
198 task_id = request.path_params["task_id"]
199 snap = await self.tracker.get(task_id)
200 if snap is None:
201 return {"error": "Task not found"}, 404
202 return _snap_to_dict(snap)