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

1"""Progress tracking controller for SSE-based real-time updates.""" 

2 

3from __future__ import annotations 

4 

5import asyncio 

6from collections.abc import AsyncIterator 

7from typing import TYPE_CHECKING, Any 

8 

9from starlette.requests import Request 

10from starlette.responses import StreamingResponse 

11 

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 

21 

22if TYPE_CHECKING: 

23 from lexigram.admin.engine.renderer import AdminRenderer 

24 

25 

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 } 

37 

38 

39class LocalProgressTracker: 

40 """In-process :class:`ProgressTrackerProtocol` implementation owned by lexigram-admin. 

41 

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 """ 

48 

49 def __init__(self) -> None: 

50 self._snapshots: dict[str, ProgressSnapshot] = {} 

51 self._subscribers: dict[str, list[asyncio.Queue[ProgressSnapshot]]] = {} 

52 

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 ) 

65 

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 ) 

77 

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 ) 

89 

90 async def get(self, task_id: str) -> ProgressSnapshot | None: 

91 return self._snapshots.get(task_id) 

92 

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 

101 

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) 

112 

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) 

117 

118 

119@inject 

120class ProgressController(AdminController): 

121 """Controller for progress tracking endpoints. 

122 

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 """ 

127 

128 def __init__( 

129 self, 

130 tracker: ProgressTrackerProtocol, 

131 renderer: AdminRenderer | None = None, 

132 ) -> None: 

133 super().__init__(renderer=renderer) 

134 self.tracker = tracker 

135 

136 @get("/progress/{task_id}/stream") 

137 async def stream_progress(self, request: Request) -> StreamingResponse: 

138 """Stream progress updates via Server-Sent Events. 

139 

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. 

143 

144 Args: 

145 request: Starlette request. The task id comes from the route. 

146 

147 Returns: 

148 SSE stream of progress snapshots. 

149 """ 

150 

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 

162 

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 ) 

169 

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" 

175 

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 ) 

184 

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. 

191 

192 Args: 

193 request: Starlette request. The task id comes from the route. 

194 

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)