Coverage for agentos/core/streaming_optimizer.py: 0%

186 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 10:59 +0800

1""" 

2AgentOS Streaming Optimizer — SSE Stream Processing & Backpressure 

3━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 

4 

5Production-grade streaming optimization for LLM responses: 

6 

7 - Chunk aggregation (avoid 1-token-at-a-time jitter) 

8 - Backpressure handling (pause/resume flow control) 

9 - Stream transformation pipeline (chains of transformers) 

10 - Token counting and rate estimation 

11 - Adaptive chunk sizing based on network conditions 

12 - Stream cancellation and cleanup 

13 

14Usage: 

15 optimizer = StreamingOptimizer() 

16 async for chunk in optimizer.optimize(llm_stream): 

17 yield chunk 

18""" 

19 

20from __future__ import annotations 

21 

22import asyncio 

23import time 

24from dataclasses import dataclass, field 

25from enum import Enum 

26from typing import Any, AsyncIterable, AsyncIterator, Callable, Dict, List, Optional 

27 

28 

29# --------------------------------------------------------------------------- 

30# Stream Chunk 

31# --------------------------------------------------------------------------- 

32 

33 

34@dataclass 

35class StreamChunk: 

36 """A single chunk from an LLM stream.""" 

37 content: str 

38 index: int 

39 timestamp: float = field(default_factory=time.time) 

40 is_final: bool = False 

41 metadata: Dict[str, Any] = field(default_factory=dict) 

42 

43 

44# --------------------------------------------------------------------------- 

45# Configuration 

46# --------------------------------------------------------------------------- 

47 

48 

49class AggregationStrategy(str, Enum): 

50 """How to aggregate small chunks.""" 

51 NONE = "none" # Pass through as-is 

52 FIXED_SIZE = "fixed" # Wait for N tokens before emitting 

53 TIME_WINDOW = "time" # Emit every T milliseconds 

54 ADAPTIVE = "adaptive" # Adjust based on latency 

55 

56 

57@dataclass 

58class StreamConfig: 

59 """Configuration for streaming optimization.""" 

60 strategy: AggregationStrategy = AggregationStrategy.ADAPTIVE 

61 min_chunk_size: int = 3 # Minimum tokens before emitting 

62 max_chunk_size: int = 50 # Maximum tokens in a single chunk 

63 time_window_ms: int = 50 # Max wait time between emissions 

64 buffer_max_chunks: int = 100 # Max unacknowledged chunks (backpressure) 

65 adaptive_latency_target_ms: int = 100 # Target round-trip latency 

66 adaptive_min_chunk_size: int = 1 

67 adaptive_max_chunk_size: int = 100 

68 enable_compression: bool = False 

69 track_performance: bool = True 

70 

71 

72# --------------------------------------------------------------------------- 

73# Performance Tracker 

74# --------------------------------------------------------------------------- 

75 

76 

77@dataclass 

78class StreamMetrics: 

79 """Stream performance metrics.""" 

80 total_chunks_received: int = 0 

81 total_chunks_emitted: int = 0 

82 total_tokens_received: int = 0 

83 total_tokens_emitted: int = 0 

84 avg_chunk_latency_ms: float = 0.0 

85 total_wall_time_ms: float = 0.0 

86 backpressure_events: int = 0 

87 peak_buffer_size: int = 0 

88 tokens_per_second: float = 0.0 

89 

90 @property 

91 def aggregation_ratio(self) -> float: 

92 if self.total_chunks_received == 0: 

93 return 1.0 

94 return self.total_chunks_received / max(1, self.total_chunks_emitted) 

95 

96 

97class MetricsCollector: 

98 """Collect streaming performance metrics.""" 

99 

100 def __init__(self): 

101 self._metrics = StreamMetrics() 

102 self._latencies: List[float] = [] 

103 

104 def record_received(self, tokens: int = 1) -> None: 

105 self._metrics.total_chunks_received += 1 

106 self._metrics.total_tokens_received += tokens 

107 

108 def record_emitted(self, tokens: int = 1, latency_ms: float = 0.0) -> None: 

109 self._metrics.total_chunks_emitted += 1 

110 self._metrics.total_tokens_emitted += tokens 

111 self._latencies.append(latency_ms) 

112 

113 def record_backpressure(self) -> None: 

114 self._metrics.backpressure_events += 1 

115 

116 def record_buffer_size(self, size: int) -> None: 

117 if size > self._metrics.peak_buffer_size: 

118 self._metrics.peak_buffer_size = size 

119 

120 def finalize(self, wall_time_ms: float) -> StreamMetrics: 

121 self._metrics.total_wall_time_ms = wall_time_ms 

122 if self._latencies: 

123 self._metrics.avg_chunk_latency_ms = sum(self._latencies) / len(self._latencies) 

124 if wall_time_ms > 0: 

125 self._metrics.tokens_per_second = ( 

126 self._metrics.total_tokens_emitted / (wall_time_ms / 1000) 

127 ) 

128 return self._metrics 

129 

130 

131# --------------------------------------------------------------------------- 

132# Stream Transformer 

133# --------------------------------------------------------------------------- 

134 

135 

136StreamTransformer = Callable[[StreamChunk], Optional[StreamChunk]] 

137 

138 

139class TransformerPipeline: 

140 """Chain of stream transformers applied to each chunk.""" 

141 

142 def __init__(self): 

143 self._transformers: List[StreamTransformer] = [] 

144 

145 def add(self, transformer: StreamTransformer) -> "TransformerPipeline": 

146 self._transformers.append(transformer) 

147 return self 

148 

149 def apply(self, chunk: StreamChunk) -> Optional[StreamChunk]: 

150 current = chunk 

151 for t in self._transformers: 

152 if current is None: 

153 return None 

154 current = t(current) 

155 return current 

156 

157 @property 

158 def size(self) -> int: 

159 return len(self._transformers) 

160 

161 

162# Built-in transformers 

163 

164 

165def strip_leading_whitespace() -> StreamTransformer: 

166 """Remove leading whitespace from first chunk.""" 

167 first = True 

168 

169 def transformer(chunk: StreamChunk) -> StreamChunk: 

170 nonlocal first 

171 if first: 

172 chunk.content = chunk.content.lstrip() 

173 first = False 

174 return chunk 

175 

176 return transformer 

177 

178 

179def normalize_newlines() -> StreamTransformer: 

180 """Normalize all line endings to '\n'.""" 

181 

182 def transformer(chunk: StreamChunk) -> StreamChunk: 

183 chunk.content = chunk.content.replace("\r\n", "\n").replace("\r", "\n") 

184 return chunk 

185 

186 return transformer 

187 

188 

189def filter_empty_chunks() -> StreamTransformer: 

190 """Drop chunks with empty content.""" 

191 def transformer(chunk: StreamChunk) -> Optional[StreamChunk]: 

192 return chunk if chunk.content else None 

193 return transformer 

194 

195 

196def add_token_count() -> StreamTransformer: 

197 """Add approximate token count to metadata.""" 

198 def transformer(chunk: StreamChunk) -> StreamChunk: 

199 # Rough estimate: ~1.3 chars per token 

200 chunk.metadata["approx_tokens"] = max(1, int(len(chunk.content) / 1.3)) 

201 return chunk 

202 return transformer 

203 

204 

205# --------------------------------------------------------------------------- 

206# Streaming Optimizer 

207# --------------------------------------------------------------------------- 

208 

209 

210class StreamingOptimizer: 

211 """ 

212 Optimize LLM streaming with aggregation, backpressure, and metrics. 

213 

214 Usage: 

215 opt = StreamingOptimizer() 

216 async for chunk in opt.optimize(llm_stream): 

217 yield chunk.content 

218 

219 With transformers: 

220 opt = StreamingOptimizer() 

221 opt.pipeline.add(strip_leading_whitespace()) 

222 async for chunk in opt.optimize(llm_stream): 

223 ... 

224 """ 

225 

226 def __init__(self, config: Optional[StreamConfig] = None): 

227 self._config = config or StreamConfig() 

228 self._pipeline = TransformerPipeline() 

229 self._buffer: List[StreamChunk] = [] 

230 self._token_accumulator: List[str] = [] 

231 self._metrics = MetricsCollector() 

232 self._start_time: Optional[float] = None 

233 self._paused = False 

234 

235 @property 

236 def pipeline(self) -> TransformerPipeline: 

237 return self._pipeline 

238 

239 @property 

240 def config(self) -> StreamConfig: 

241 return self._config 

242 

243 async def optimize( 

244 self, stream: AsyncIterable[StreamChunk] 

245 ) -> AsyncIterator[StreamChunk]: 

246 """ 

247 Optimize a stream of chunks. 

248 

249 Applies aggregation and transformer pipeline. 

250 """ 

251 self._start_time = time.time() 

252 

253 async for chunk in stream: 

254 self._metrics.record_received() 

255 

256 # Backpressure: wait if buffer is full 

257 while len(self._buffer) >= self._config.buffer_max_chunks: 

258 self._metrics.record_backpressure() 

259 await asyncio.sleep(0.01) 

260 

261 self._buffer.append(chunk) 

262 self._metrics.record_buffer_size(len(self._buffer)) 

263 

264 # Check if we should emit 

265 result = await self._maybe_emit() 

266 if result is not None: 

267 yield result 

268 

269 # Flush remaining buffer 

270 async for chunk in self._flush(): 

271 yield chunk 

272 

273 self._metrics.finalize((time.time() - self._start_time) * 1000) 

274 

275 async def optimize_simple( 

276 self, text_stream: AsyncIterable[str] 

277 ) -> AsyncIterator[str]: 

278 """ 

279 Optimize a simple text stream (strings instead of StreamChunk objects). 

280 """ 

281 async for chunk in self.optimize( 

282 StreamChunk(content=text, index=i) 

283 for i, text in enumerate(text_stream) 

284 ): 

285 yield chunk.content 

286 

287 async def _maybe_emit(self) -> Optional[StreamChunk]: 

288 """Check if we should emit an aggregated chunk.""" 

289 if not self._buffer: 

290 return None 

291 

292 tokens = sum( 

293 chunk.metadata.get("approx_tokens", 1) for chunk in self._buffer 

294 ) 

295 

296 should_emit = False 

297 

298 if self._config.strategy == AggregationStrategy.NONE: 

299 should_emit = True 

300 elif self._config.strategy == AggregationStrategy.FIXED_SIZE: 

301 should_emit = tokens >= self._config.min_chunk_size 

302 elif self._config.strategy == AggregationStrategy.TIME_WINDOW: 

303 if self._buffer: 

304 elapsed = (time.time() - self._buffer[0].timestamp) * 1000 

305 should_emit = elapsed >= self._config.time_window_ms 

306 elif self._config.strategy == AggregationStrategy.ADAPTIVE: 

307 should_emit = ( 

308 tokens >= self._config.adaptive_min_chunk_size 

309 and ( 

310 tokens >= self._config.max_chunk_size 

311 or self._buffer[-1].is_final 

312 ) 

313 ) 

314 

315 if not should_emit: 

316 return None 

317 

318 return self._emit_aggregated() 

319 

320 def _emit_aggregated(self) -> Optional[StreamChunk]: 

321 """Aggregate buffered chunks into a single emission.""" 

322 if not self._buffer: 

323 return None 

324 

325 # Aggregate content 

326 content = "".join(c.content for c in self._buffer) 

327 is_final = self._buffer[-1].is_final 

328 index = self._buffer[-1].index 

329 timestamp = time.time() 

330 

331 # Clear buffer 

332 count = len(self._buffer) 

333 self._buffer.clear() 

334 

335 # Build aggregated chunk 

336 chunk = StreamChunk( 

337 content=content, 

338 index=index, 

339 timestamp=timestamp, 

340 is_final=is_final, 

341 metadata={"aggregated_from": count}, 

342 ) 

343 

344 # Apply transformer pipeline 

345 chunk = self._pipeline.apply(chunk) 

346 if chunk is None: 

347 return None 

348 

349 # Record metrics 

350 latency_ms = (timestamp - self._start_time) * 1000 if self._start_time else 0 

351 approx_tokens = chunk.metadata.get("approx_tokens", 1) 

352 self._metrics.record_emitted(tokens=approx_tokens, latency_ms=latency_ms) 

353 

354 return chunk 

355 

356 async def _flush(self) -> AsyncIterator[StreamChunk]: 

357 """Flush remaining buffered chunks.""" 

358 while self._buffer: 

359 chunk = self._emit_aggregated() 

360 if chunk is not None: 

361 yield chunk 

362 

363 def get_metrics(self) -> StreamMetrics: 

364 return self._metrics._metrics 

365 

366 def reset_metrics(self) -> None: 

367 self._metrics = MetricsCollector() 

368 self._start_time = None