perf: split streaming SSE events in linear time

This commit is contained in:
9qeklajc
2026-09-26 16:02:11 +02:00
parent d8793f7785
commit bf3bb2b747
4 changed files with 125 additions and 101 deletions
-81
View File
@@ -1,81 +0,0 @@
# Streaming latency patterns
Patterns taken from LiteLLM 1.93 (`litellm/proxy/pass_through_endpoints/`,
`litellm/proxy/common_request_processing.py`, `litellm/litellm_core_utils/logging_worker.py`)
and how they map onto Routstr's streaming hot path in `routstr/upstream/base.py`.
## Where time goes today
Every SSE event in `handle_streaming_chat_completion` and
`handle_streaming_responses_completion` is parsed, mutated (`model`, `id`,
`provider`, `provider_url`), observed for usage, and reserialized. Cost scales with
chunks per second times concurrent streams, all on one event-loop thread.
## 1. Fast JSON in the per-chunk path — done
LiteLLM parses request bodies with `orjson` (`common_utils/http_parsing_utils.py`).
Routstr: `routstr/upstream/json_codec.py` wraps `orjson` with a stdlib fallback
(orjson rejects `NaN` on load and non-string keys / >64-bit ints on dump). Used for
the per-event parse and reserialize in both streaming paths.
Wire change: emitted events are compact UTF-8 JSON (`{"a":1}`, raw `é`) instead of
stdlib's `{"a": 1}` with `\u00e9`. Both are valid JSON.
Measured on a typical chat chunk: 2.85µs → 0.54µs per parse+serialize (5.3x).
## 2. Resolve per-stream invariants once — done (partial)
LiteLLM computes `fast_path`, `cost_injection_active` and `debug_enabled` once per
stream, then runs a branch-free loop (`common_request_processing.py:2632`).
Routstr: `_apply_provider_field` (and the OpenRouter/generic overrides) ran
`public_provider_url(self.base_url)` — a `urlsplit` plus `ipaddress` parse — on every
chunk. It is now `lru_cache`d in `routstr/upstream/model_paths.py`, which keeps
subclass semantics. 1.67µs → 0.03µs per chunk.
Combined per-chunk saving from 1 and 2: 4.52µs → 0.57µs.
## 3. Linear buffer handling — next
`buffer = (buffer + chunk).replace(b"\r\n", b"\n")` re-copies and rescans the whole
unconsumed buffer on every network chunk, and `b"\n\n" in buffer` rescans it again.
This is quadratic in event size — it bites on large single events such as
Responses API `response.completed`, which carries the full output. Normalize only the
new chunk (holding back a trailing `\r`) and search for the delimiter from the
previous end offset.
## 4. Raw passthrough, parse usage at end — next
LiteLLM's pass-through hot path forwards `aiter_bytes()` chunks untouched and appends
them to `raw_bytes`; usage is reconstructed once after the stream
(`streaming_handler.py:chunk_processor`, `_convert_raw_bytes_to_str_lines`).
Routstr parses every event to rewrite `model`/`id` and feed
`MissingUsageEstimator.observe`. Candidates to skip parsing: events whose `model`
already equals `requested_model` and whose `id` is stable, with usage observed from a
cheap byte check (`b'"usage"'`) or at end of stream. LiteLLM makes byte-level mutation
safe by returning the original chunk on any failure
(`_process_chunk_with_cost_injection`). Needs a framing test suite before starting.
## 5. Settlement off the response path — next
LiteLLM enqueues end-of-stream work on a bounded `asyncio.Queue` with a semaphore
and per-task timeout (`logging_worker.py`, "+200 RPS").
Routstr runs `adjust_payment_for_tokens` inline after the last upstream chunk, so the
client waits on a DB session and writes before the stream closes. The reservation is
already persisted, so settlement can move to a worker and the stale-reservation sweep
stays the backstop. Unlike LiteLLM's logging queue, this queue must never drop work,
and the cost trailer the client receives must be computed before the stream closes or
be dropped from the contract.
## Related, outside the streaming loop
- `keys.db` runs WAL with default `synchronous=FULL`, so every reservation commit
fsyncs before the upstream request is sent. `synchronous=NORMAL` removes that and
cannot corrupt the database.
- `fastapi run` serves one worker. Multiple workers are blocked: the lifespan starts
payout, auto top-up and refund tasks per process, and there is no leader election.
- There is no inbound admission control. A concurrency gate that returns 429 with
`Retry-After` before `pay_for_request` would shed load before it reaches the DB.
+7 -20
View File
@@ -74,6 +74,7 @@ from .litellm_routing import detect_litellm_prefix
from .model_paths import public_provider_url
from .rate_limit import UPSTREAM_RATE_LIMIT, classify_rate_limit
from .reasoning_effort import apply_reasoning_effort
from .sse_splitter import SSEEventSplitter
from .stream_ownership import (
ClosingStreamingResponse,
OwnedUpstreamStream,
@@ -1346,21 +1347,14 @@ class BaseUpstreamProvider:
# byte boundaries, so a single event's JSON can span chunks and
# multiple events can arrive together; buffering makes parsing
# boundary-independent for every provider.
buffer = b""
splitter = SSEEventSplitter()
async for chunk in response.aiter_bytes():
# Normalize the *joined* buffer, not each chunk in
# isolation: a CRLF event delimiter can straddle two
# ``aiter_bytes`` chunks (``...\r`` then ``\n...``). A
# per-chunk replace would leave a stray ``\r`` and the
# ``\n\n`` split would miss the delimiter, merging two
# events into one frame and breaking SSE clients.
buffer = (buffer + chunk).replace(b"\r\n", b"\n")
while b"\n\n" in buffer:
raw_event, buffer = buffer.split(b"\n\n", 1)
for raw_event in splitter.feed(chunk):
for out in _process_event(raw_event):
yield out
# Flush any trailing event that lacked a final blank line.
buffer = splitter.flush()
if buffer.strip():
for out in _process_event(buffer, final=True):
yield out
@@ -1795,20 +1789,13 @@ class BaseUpstreamProvider:
try:
# Buffer across network chunks; dispatch only on the SSE event
# delimiter so parsing is independent of byte boundaries.
buffer = b""
splitter = SSEEventSplitter()
async for chunk in response.aiter_bytes():
# Normalize the *joined* buffer, not each chunk in
# isolation: a CRLF event delimiter can straddle two
# ``aiter_bytes`` chunks (``...\r`` then ``\n...``). A
# per-chunk replace would leave a stray ``\r`` and the
# ``\n\n`` split would miss the delimiter, merging two
# events into one frame and breaking SSE clients.
buffer = (buffer + chunk).replace(b"\r\n", b"\n")
while b"\n\n" in buffer:
raw_event, buffer = buffer.split(b"\n\n", 1)
for raw_event in splitter.feed(chunk):
for out in _process_event(raw_event):
yield out
buffer = splitter.flush()
if buffer.strip():
for out in _process_event(buffer, final=True):
yield out
+44
View File
@@ -0,0 +1,44 @@
"""Incremental SSE event splitting that stays linear in stream size."""
class SSEEventSplitter:
"""Split upstream bytes into SSE events delimited by a blank line.
CRLF is normalized to LF. Each call only scans newly received bytes, so a
large event arriving over many network chunks (e.g. a Responses API
``response.completed`` carrying the full output) costs O(n) rather than
rescanning the buffered prefix on every chunk.
"""
def __init__(self) -> None:
self._buffer = bytearray()
# A trailing CR may be the first half of a CRLF split across chunks.
self._pending_cr = False
def feed(self, chunk: bytes) -> list[bytes]:
"""Add ``chunk`` and return the events it completed, without delimiters."""
if self._pending_cr:
chunk = b"\r" + chunk
self._pending_cr = chunk.endswith(b"\r")
if self._pending_cr:
chunk = chunk[:-1]
# The delimiter may straddle the old tail and the new chunk.
scan_from = max(len(self._buffer) - 1, 0)
self._buffer += chunk.replace(b"\r\n", b"\n")
events: list[bytes] = []
start = 0
while (end := self._buffer.find(b"\n\n", scan_from)) != -1:
events.append(bytes(self._buffer[start:end]))
start = scan_from = end + 2
if start:
del self._buffer[:start]
return events
def flush(self) -> bytes:
"""Return any trailing bytes that never saw a closing blank line."""
tail = bytes(self._buffer) + (b"\r" if self._pending_cr else b"")
self._buffer.clear()
self._pending_cr = False
return tail
+74
View File
@@ -0,0 +1,74 @@
import random
import pytest
from routstr.upstream.sse_splitter import SSEEventSplitter
def _reference_split(chunks: list[bytes]) -> tuple[list[bytes], bytes]:
"""The original rescanning implementation the splitter replaces."""
events: list[bytes] = []
buffer = b""
for chunk in chunks:
buffer = (buffer + chunk).replace(b"\r\n", b"\n")
while b"\n\n" in buffer:
raw_event, buffer = buffer.split(b"\n\n", 1)
events.append(raw_event)
return events, buffer
def _split(chunks: list[bytes]) -> tuple[list[bytes], bytes]:
splitter = SSEEventSplitter()
events = [event for chunk in chunks for event in splitter.feed(chunk)]
return events, splitter.flush()
STREAMS = [
b'data: {"a":1}\n\ndata: {"b":2}\n\ndata: [DONE]\n\n',
b'data: {"a":1}\r\n\r\ndata: {"b":2}\r\n\r\ndata: [DONE]\r\n\r\n',
b': OPENROUTER PROCESSING\n\ndata: {"a":1}\n\n: keepalive\n\ndata: [DONE]\n\n',
b'event: response.created\ndata: {"type":"x"}\n\nevent: done\ndata: {"t":1}\n\n',
b'data: {"part":\ndata: "two"}\n\n\n\ndata: {"trailing":true}',
b'data: {"a":1}\r\n\r\ndata: {"tail":1}\r',
b"\n\n\n\n",
b"",
]
@pytest.mark.parametrize("stream", STREAMS)
def test_matches_reference_at_every_two_way_split(stream: bytes) -> None:
for cut in range(len(stream) + 1):
chunks = [stream[:cut], stream[cut:]]
assert _split(chunks) == _reference_split(chunks)
@pytest.mark.parametrize("stream", STREAMS)
def test_matches_reference_on_random_chunkings(stream: bytes) -> None:
rng = random.Random(0)
for _ in range(200):
cuts = sorted(rng.sample(range(len(stream) + 1), min(len(stream), 6)))
bounds = [0, *cuts, len(stream)]
chunks = [stream[a:b] for a, b in zip(bounds, bounds[1:])]
assert _split(chunks) == _reference_split(chunks)
def test_byte_at_a_time_crlf_stream() -> None:
stream = b'data: {"a":1}\r\n\r\ndata: {"b":2}\r\n\r\n'
events, tail = _split([bytes([b]) for b in stream])
assert events == [b'data: {"a":1}', b'data: {"b":2}']
assert tail == b""
def test_flush_returns_held_back_carriage_return() -> None:
splitter = SSEEventSplitter()
assert splitter.feed(b"data: x\r") == []
assert splitter.flush() == b"data: x\r"
assert splitter.flush() == b""
def test_large_event_over_many_chunks() -> None:
payload = b"data: " + b"x" * 200_000 + b"\n\n"
chunks = [payload[i : i + 64] for i in range(0, len(payload), 64)]
events, tail = _split(chunks)
assert events == [payload[:-2]]
assert tail == b""