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
« 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━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
5Production-grade streaming optimization for LLM responses:
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
14Usage:
15 optimizer = StreamingOptimizer()
16 async for chunk in optimizer.optimize(llm_stream):
17 yield chunk
18"""
20from __future__ import annotations
22import asyncio
23import time
24from dataclasses import dataclass, field
25from enum import Enum
26from typing import Any, AsyncIterable, AsyncIterator, Callable, Dict, List, Optional
29# ---------------------------------------------------------------------------
30# Stream Chunk
31# ---------------------------------------------------------------------------
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)
44# ---------------------------------------------------------------------------
45# Configuration
46# ---------------------------------------------------------------------------
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
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
72# ---------------------------------------------------------------------------
73# Performance Tracker
74# ---------------------------------------------------------------------------
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
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)
97class MetricsCollector:
98 """Collect streaming performance metrics."""
100 def __init__(self):
101 self._metrics = StreamMetrics()
102 self._latencies: List[float] = []
104 def record_received(self, tokens: int = 1) -> None:
105 self._metrics.total_chunks_received += 1
106 self._metrics.total_tokens_received += tokens
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)
113 def record_backpressure(self) -> None:
114 self._metrics.backpressure_events += 1
116 def record_buffer_size(self, size: int) -> None:
117 if size > self._metrics.peak_buffer_size:
118 self._metrics.peak_buffer_size = size
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
131# ---------------------------------------------------------------------------
132# Stream Transformer
133# ---------------------------------------------------------------------------
136StreamTransformer = Callable[[StreamChunk], Optional[StreamChunk]]
139class TransformerPipeline:
140 """Chain of stream transformers applied to each chunk."""
142 def __init__(self):
143 self._transformers: List[StreamTransformer] = []
145 def add(self, transformer: StreamTransformer) -> "TransformerPipeline":
146 self._transformers.append(transformer)
147 return self
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
157 @property
158 def size(self) -> int:
159 return len(self._transformers)
162# Built-in transformers
165def strip_leading_whitespace() -> StreamTransformer:
166 """Remove leading whitespace from first chunk."""
167 first = True
169 def transformer(chunk: StreamChunk) -> StreamChunk:
170 nonlocal first
171 if first:
172 chunk.content = chunk.content.lstrip()
173 first = False
174 return chunk
176 return transformer
179def normalize_newlines() -> StreamTransformer:
180 """Normalize all line endings to '\n'."""
182 def transformer(chunk: StreamChunk) -> StreamChunk:
183 chunk.content = chunk.content.replace("\r\n", "\n").replace("\r", "\n")
184 return chunk
186 return transformer
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
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
205# ---------------------------------------------------------------------------
206# Streaming Optimizer
207# ---------------------------------------------------------------------------
210class StreamingOptimizer:
211 """
212 Optimize LLM streaming with aggregation, backpressure, and metrics.
214 Usage:
215 opt = StreamingOptimizer()
216 async for chunk in opt.optimize(llm_stream):
217 yield chunk.content
219 With transformers:
220 opt = StreamingOptimizer()
221 opt.pipeline.add(strip_leading_whitespace())
222 async for chunk in opt.optimize(llm_stream):
223 ...
224 """
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
235 @property
236 def pipeline(self) -> TransformerPipeline:
237 return self._pipeline
239 @property
240 def config(self) -> StreamConfig:
241 return self._config
243 async def optimize(
244 self, stream: AsyncIterable[StreamChunk]
245 ) -> AsyncIterator[StreamChunk]:
246 """
247 Optimize a stream of chunks.
249 Applies aggregation and transformer pipeline.
250 """
251 self._start_time = time.time()
253 async for chunk in stream:
254 self._metrics.record_received()
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)
261 self._buffer.append(chunk)
262 self._metrics.record_buffer_size(len(self._buffer))
264 # Check if we should emit
265 result = await self._maybe_emit()
266 if result is not None:
267 yield result
269 # Flush remaining buffer
270 async for chunk in self._flush():
271 yield chunk
273 self._metrics.finalize((time.time() - self._start_time) * 1000)
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
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
292 tokens = sum(
293 chunk.metadata.get("approx_tokens", 1) for chunk in self._buffer
294 )
296 should_emit = False
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 )
315 if not should_emit:
316 return None
318 return self._emit_aggregated()
320 def _emit_aggregated(self) -> Optional[StreamChunk]:
321 """Aggregate buffered chunks into a single emission."""
322 if not self._buffer:
323 return None
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()
331 # Clear buffer
332 count = len(self._buffer)
333 self._buffer.clear()
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 )
344 # Apply transformer pipeline
345 chunk = self._pipeline.apply(chunk)
346 if chunk is None:
347 return None
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)
354 return chunk
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
363 def get_metrics(self) -> StreamMetrics:
364 return self._metrics._metrics
366 def reset_metrics(self) -> None:
367 self._metrics = MetricsCollector()
368 self._start_time = None