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

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 if renderer is None: 

134 from lexigram.admin.engine.renderer import AdminRenderer 

135 

136 renderer = AdminRenderer() 

137 super().__init__(renderer=renderer) 

138 self.tracker = tracker 

139 

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

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

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

143 

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. 

147 

148 Args: 

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

150 

151 Returns: 

152 SSE stream of progress snapshots. 

153 """ 

154 

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 

166 

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 ) 

173 

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" 

179 

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 ) 

188 

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. 

195 

196 Args: 

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

198 

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)