Summarize streams with real progress: adapters reassemble SSE/NDJSON deltas and report 0-99% as tokens arrive

This commit is contained in:
avi 2026-09-15 19:45:02 -05:00
commit 7fa99e8291
8 changed files with 246 additions and 42 deletions

View file

@ -113,7 +113,8 @@ class LlmProvider(Protocol):
name: str
async def summarize(self, transcript: str, *, title: str | None = None,
tone: str | None = None) -> SummaryResult: ...
tone: str | None = None,
on_progress=None) -> SummaryResult: ... # Callable[[int], None] | None
def get_transcription_provider(settings: Settings) -> TranscriptionProvider | None:

View file

@ -38,6 +38,30 @@ def build_user_message(transcript: str, title: str | None,
return msg
# Rough typical length of a summary JSON reply. The stream cannot know the
# model's final size, so progress is an estimate that crawls toward 99 and
# the job sets 100 on success — honest "almost there", never a fake jump.
EST_OUTPUT_CHARS = 1200
def make_progress_ticker(on_progress):
"""Wrap an optional on_progress callback into a feed(n_chars) sink.
Reports a clamped 0..99 estimate from streamed output size. Silent
no-op when the caller has no callback; failures never disturb the
summary itself."""
if on_progress is None:
return lambda n_chars: None
def feed(n_chars: int) -> None:
try: # noqa: SIM105 — swallow deliberately: progress is display state
on_progress(min(99, n_chars * 100 // EST_OUTPUT_CHARS))
except Exception: # display state only
pass
return feed
def parse_summary(data: object, model: str) -> SummaryResult:
if not isinstance(data, dict):
return SummaryResult(model=model)

View file

@ -8,7 +8,12 @@ from contextlib import asynccontextmanager
import httpx
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
from shonar.services.ai._llm import SYSTEM_PROMPT, build_user_message, parse_summary
from shonar.services.ai._llm import (
SYSTEM_PROMPT,
build_user_message,
make_progress_ticker,
parse_summary,
)
class OllamaProvider:
@ -39,10 +44,13 @@ class OllamaProvider:
yield client
async def summarize(self, transcript: str, *, title: str | None = None,
tone: str | None = None) -> SummaryResult:
tone: str | None = None, on_progress=None) -> SummaryResult:
ticker = make_progress_ticker(on_progress)
payload = {
"model": self.model,
"stream": False,
# Stream (NDJSON lines) so the UI shows real progress while the
# model generates; the reply is reassembled from the chunks.
"stream": True,
"format": "json",
# qwen3-family models "think" by default: a long reasoning
# chain before the JSON answer, brutally slow on CPU and it
@ -57,20 +65,20 @@ class OllamaProvider:
],
}
try:
async with self._client() as client:
resp = await client.post(f"{self.base_url}/api/chat", json=payload)
async with self._client() as client, client.stream(
"POST", f"{self.base_url}/api/chat", json=payload
) as resp:
if resp.status_code == 404:
# Missing model and missing route both 404 here; both
# are configuration, not weather.
raise ProviderConfigError(
"Ollama has no such model or route (HTTP 404).")
if resp.status_code != 200:
raise ProviderTransientError(
f"Summarization failed (HTTP {resp.status_code}).")
content = await _collect_reply(resp, ticker)
except (httpx.TimeoutException, httpx.TransportError) as e:
raise ProviderTransientError(f"Ollama unreachable: {type(e).__name__}") from e
if resp.status_code == 404:
# Missing model and missing route both 404 here; both are
# configuration, not weather.
raise ProviderConfigError("Ollama has no such model or route (HTTP 404).")
if resp.status_code != 200:
raise ProviderTransientError(f"Summarization failed (HTTP {resp.status_code}).")
try:
content = resp.json()["message"]["content"]
except (ValueError, KeyError, TypeError) as e:
raise ProviderTransientError("Ollama sent an unreadable reply.") from e
import json as _json
try:
@ -78,3 +86,44 @@ class OllamaProvider:
except ValueError as e:
raise ProviderTransientError("Ollama reply was not JSON.") from e
return parse_summary(data, self.model)
async def _collect_reply(resp: httpx.Response, ticker) -> str:
"""Reassemble the assistant reply from Ollama's streamed chat response.
Streaming answers are NDJSON (one JSON object per line, ``done`` on the
last); a server honoring stream=false returns one JSON body — both work.
The running character count feeds [ticker] for UI progress."""
import json as _json
ctype = resp.headers.get("content-type", "")
if "ndjson" not in ctype and "event-stream" not in ctype:
body = await resp.aread()
try:
content = _json.loads(body)["message"]["content"]
except (ValueError, KeyError, TypeError) as e:
raise ProviderTransientError("Ollama sent an unreadable reply.") from e
ticker(len(content or ""))
return content or ""
parts: list[str] = []
total = 0
async for line in resp.aiter_lines():
line = line.strip()
if not line:
continue
try:
obj = _json.loads(line)
except ValueError:
continue
piece = (obj.get("message") or {}).get("content") or ""
if piece:
parts.append(piece)
total += len(piece)
ticker(total)
if obj.get("done"):
break
content = "".join(parts)
if not content:
raise ProviderTransientError("Ollama sent an unreadable reply.")
return content

View file

@ -10,7 +10,54 @@ from contextlib import asynccontextmanager
import httpx
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
from shonar.services.ai._llm import SYSTEM_PROMPT, build_user_message, parse_summary
from shonar.services.ai._llm import (
SYSTEM_PROMPT,
build_user_message,
make_progress_ticker,
parse_summary,
)
async def _collect_reply(resp: httpx.Response, ticker) -> str:
"""Reassemble the assistant reply from a (possibly streamed) response.
Feeds the running character count to [ticker] as deltas arrive so the
UI can show progress. Servers that ignored "stream": true answer with
a plain JSON body — that path is handled too."""
import json as _json
ctype = resp.headers.get("content-type", "")
if "text/event-stream" not in ctype:
body = await resp.aread()
try:
content = _json.loads(body)["choices"][0]["message"]["content"]
except (ValueError, KeyError, IndexError, TypeError) as e:
raise ProviderTransientError("LLM sent an unreadable reply.") from e
ticker(len(content or ""))
return content or ""
parts: list[str] = []
total = 0
async for line in resp.aiter_lines():
if not line.startswith("data:"):
continue
data = line[5:].strip()
if data == "[DONE]":
break
try:
chunk = _json.loads(data)
delta = chunk["choices"][0].get("delta") or {}
piece = delta.get("content") or ""
except (ValueError, KeyError, IndexError, TypeError):
continue # keep-alives / usage chunks / odd frames: not content
if piece:
parts.append(piece)
total += len(piece)
ticker(total)
content = "".join(parts)
if not content:
raise ProviderTransientError("LLM sent an unreadable reply.")
return content
class OpenAICompatProvider:
@ -43,7 +90,8 @@ class OpenAICompatProvider:
yield client
async def summarize(self, transcript: str, *, title: str | None = None,
tone: str | None = None) -> SummaryResult:
tone: str | None = None, on_progress=None) -> SummaryResult:
ticker = make_progress_ticker(on_progress)
headers = (
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
)
@ -51,6 +99,10 @@ class OpenAICompatProvider:
"model": self.model,
"temperature": 0.2,
"response_format": {"type": "json_object"},
# Stream so the UI gets real progress while the model works;
# the full reply is reassembled from the deltas. Servers that
# ignore "stream" still work — the non-SSE body is handled too.
"stream": True,
"messages": [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user",
@ -58,23 +110,22 @@ class OpenAICompatProvider:
],
}
try:
async with self._client() as client:
resp = await client.post(
f"{self.base_url}/v1/chat/completions", headers=headers, json=payload
)
async with self._client() as client, client.stream(
"POST", f"{self.base_url}/v1/chat/completions",
headers=headers, json=payload,
) as resp:
if resp.status_code in (401, 403, 404):
raise ProviderConfigError(
f"LLM refused the request (HTTP {resp.status_code}).")
if resp.status_code == 429 or resp.status_code >= 500:
raise ProviderTransientError(
f"LLM busy (HTTP {resp.status_code}).")
if resp.status_code != 200:
raise ProviderTransientError(
f"Summarization failed (HTTP {resp.status_code}).")
content = await _collect_reply(resp, ticker)
except (httpx.TimeoutException, httpx.TransportError) as e:
raise ProviderTransientError(f"LLM unreachable: {type(e).__name__}") from e
if resp.status_code in (401, 403, 404):
raise ProviderConfigError(f"LLM refused the request (HTTP {resp.status_code}).")
if resp.status_code == 429 or resp.status_code >= 500:
raise ProviderTransientError(f"LLM busy (HTTP {resp.status_code}).")
if resp.status_code != 200:
raise ProviderTransientError(f"Summarization failed (HTTP {resp.status_code}).")
try:
body = resp.json()
content = body["choices"][0]["message"]["content"]
except (ValueError, KeyError, IndexError, TypeError) as e:
raise ProviderTransientError("LLM sent an unreadable reply.") from e
import json as _json
try:

View file

@ -90,6 +90,49 @@ def _thread_progress_reporter(job_id):
return report
def _async_progress_writer(job_id):
"""Progress callback for providers that stream on our own event loop
(the LLM adapters). Unlike _thread_progress_reporter there is no
worker thread: schedule the tiny write as a task on the running loop.
Same contract as the thread version — best-effort display state — and
identical values are coalesced so a chatty stream does not hammer
SQLite."""
loop = asyncio.get_running_loop()
last = -1
def report(pct: int) -> None:
nonlocal last
if int(pct) == last:
return
last = int(pct)
async def _write() -> None:
from sqlalchemy import update
from shonar.db.session import session_factory
try:
async with session_factory()() as s:
await s.execute(
update(ProcessingJob)
.where(ProcessingJob.id == job_id)
.values(progress=last)
)
await s.commit()
except Exception: # pragma: no cover - display state only
logger.debug("summarize progress write failed for job %s",
job_id, exc_info=True)
try:
loop.create_task(_write())
except RuntimeError: # loop already gone (shutdown race)
logger.debug("summarize progress dropped for job %s", job_id)
return report
MAX_TRIES = 3
# A `running` job younger than this is treated as live work, not a crash
@ -437,8 +480,9 @@ async def run_summarize(ctx: dict, recording_id: str) -> None:
rec.processing_status = ProcessingStatus.processing
await session.commit() # visible before the long LLM call
try:
result = await provider.summarize(text, title=rec.title,
tone=job.tone)
result = await provider.summarize(
text, title=rec.title, tone=job.tone,
on_progress=_async_progress_writer(job.id))
except AIError as e:
await _fail(session, rec, job, str(e), ctx, e)
await session.commit()

View file

@ -129,6 +129,32 @@ async def test_openai_compat_parses_and_sanitizes():
assert res.model == "qwen"
async def test_openai_compat_streams_with_progress():
"""stream:true replies (SSE deltas) reassemble and tick progress."""
import json as _json
def sse(request: httpx.Request) -> httpx.Response:
payload = _json.loads(request.content.decode())
assert payload["stream"] is True
pieces = ['{"short": "Stand', 'up.", "detailed": "The team met."',
', "key_points": ["a"]', "}"]
lines = [f"data: {_json.dumps({'choices': [{'delta': {'content': p}}]})}"
for p in pieces]
lines.append("data: [DONE]")
return httpx.Response(
200, content=("\n\n".join(lines) + "\n\n").encode(),
headers={"content-type": "text/event-stream"})
p = OpenAICompatProvider("http://llm:8000", model="qwen",
http_client=mock_client(sse))
seen_pcts = []
res = await p.summarize("t", on_progress=seen_pcts.append)
assert res.short == "Standup."
assert res.key_points == ("a",)
assert seen_pcts and seen_pcts == sorted(seen_pcts)
assert all(0 <= p <= 99 for p in seen_pcts) # never claims 100 early
async def test_openai_compat_partial_json_gets_defaults():
import json as _json
@ -165,14 +191,21 @@ async def test_ollama_happy_path():
async def ok(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/api/chat"
payload = _json.loads(request.content.decode())
assert payload["format"] == "json" and payload["stream"] is False
return httpx.Response(200, json={
"message": {"content": _json.dumps({"short": "S.", "detailed": "D."})}
})
assert payload["format"] == "json" and payload["stream"] is True
body = "\n".join([
_json.dumps({"message": {"content": "{\"short\": \"S."}}),
_json.dumps({"message": {"content": "\", \"detailed\": \"D.\"}"}}),
_json.dumps({"done": True}),
])
return httpx.Response(
200, json=None, content=body.encode(),
headers={"content-type": "application/x-ndjson"})
p = OllamaProvider("http://ollama:11434", model="llama3",
http_client=mock_client(ok))
res = await p.summarize("meeting notes")
seen_pcts = []
res = await p.summarize("meeting notes", on_progress=seen_pcts.append)
assert res.short == "S." and res.model == "llama3"
assert seen_pcts and seen_pcts == sorted(seen_pcts) # monotonic progress
async def test_ollama_404_is_config_error():

View file

@ -95,7 +95,8 @@ class FakeLlm:
self.fail = fail
self.seen = []
async def summarize(self, transcript, *, title=None, tone=None):
async def summarize(self, transcript, *, title=None, tone=None,
on_progress=None):
self.seen.append(transcript)
if self.fail is not None:
raise self.fail

View file

@ -131,7 +131,8 @@ async def test_user_edit_survives_auto_pipeline(client, monkeypatch):
class FakeLlm:
name = "fake-llm"
async def summarize(self, transcript, *, title=None, tone=None):
async def summarize(self, transcript, *, title=None, tone=None,
on_progress=None):
return SummaryResult(short="s", detailed="d", key_points=(),
decisions=(), action_items=(), questions=(),
model="fake-llm-1")