Summarize streams with real progress: adapters reassemble SSE/NDJSON deltas and report 0-99% as tokens arrive
This commit is contained in:
parent
cd1b08fda0
commit
7fa99e8291
8 changed files with 246 additions and 42 deletions
|
|
@ -113,7 +113,8 @@ class LlmProvider(Protocol):
|
||||||
name: str
|
name: str
|
||||||
|
|
||||||
async def summarize(self, transcript: str, *, title: str | None = None,
|
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:
|
def get_transcription_provider(settings: Settings) -> TranscriptionProvider | None:
|
||||||
|
|
|
||||||
|
|
@ -38,6 +38,30 @@ def build_user_message(transcript: str, title: str | None,
|
||||||
return msg
|
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:
|
def parse_summary(data: object, model: str) -> SummaryResult:
|
||||||
if not isinstance(data, dict):
|
if not isinstance(data, dict):
|
||||||
return SummaryResult(model=model)
|
return SummaryResult(model=model)
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,12 @@ from contextlib import asynccontextmanager
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
|
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:
|
class OllamaProvider:
|
||||||
|
|
@ -39,10 +44,13 @@ class OllamaProvider:
|
||||||
yield client
|
yield client
|
||||||
|
|
||||||
async def summarize(self, transcript: str, *, title: str | None = None,
|
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 = {
|
payload = {
|
||||||
"model": self.model,
|
"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",
|
"format": "json",
|
||||||
# qwen3-family models "think" by default: a long reasoning
|
# qwen3-family models "think" by default: a long reasoning
|
||||||
# chain before the JSON answer, brutally slow on CPU and it
|
# chain before the JSON answer, brutally slow on CPU and it
|
||||||
|
|
@ -57,20 +65,20 @@ class OllamaProvider:
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
async with self._client() as client:
|
async with self._client() as client, client.stream(
|
||||||
resp = await client.post(f"{self.base_url}/api/chat", json=payload)
|
"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:
|
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||||
raise ProviderTransientError(f"Ollama unreachable: {type(e).__name__}") from 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
|
import json as _json
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
@ -78,3 +86,44 @@ class OllamaProvider:
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
raise ProviderTransientError("Ollama reply was not JSON.") from e
|
raise ProviderTransientError("Ollama reply was not JSON.") from e
|
||||||
return parse_summary(data, self.model)
|
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
|
||||||
|
|
|
||||||
|
|
@ -10,7 +10,54 @@ from contextlib import asynccontextmanager
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
|
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:
|
class OpenAICompatProvider:
|
||||||
|
|
@ -43,7 +90,8 @@ class OpenAICompatProvider:
|
||||||
yield client
|
yield client
|
||||||
|
|
||||||
async def summarize(self, transcript: str, *, title: str | None = None,
|
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 = (
|
headers = (
|
||||||
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||||||
)
|
)
|
||||||
|
|
@ -51,6 +99,10 @@ class OpenAICompatProvider:
|
||||||
"model": self.model,
|
"model": self.model,
|
||||||
"temperature": 0.2,
|
"temperature": 0.2,
|
||||||
"response_format": {"type": "json_object"},
|
"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": [
|
"messages": [
|
||||||
{"role": "system", "content": SYSTEM_PROMPT},
|
{"role": "system", "content": SYSTEM_PROMPT},
|
||||||
{"role": "user",
|
{"role": "user",
|
||||||
|
|
@ -58,23 +110,22 @@ class OpenAICompatProvider:
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
async with self._client() as client:
|
async with self._client() as client, client.stream(
|
||||||
resp = await client.post(
|
"POST", f"{self.base_url}/v1/chat/completions",
|
||||||
f"{self.base_url}/v1/chat/completions", headers=headers, json=payload
|
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:
|
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||||
raise ProviderTransientError(f"LLM unreachable: {type(e).__name__}") from 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
|
import json as _json
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
|
||||||
|
|
@ -90,6 +90,49 @@ def _thread_progress_reporter(job_id):
|
||||||
|
|
||||||
return report
|
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
|
MAX_TRIES = 3
|
||||||
|
|
||||||
# A `running` job younger than this is treated as live work, not a crash
|
# 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
|
rec.processing_status = ProcessingStatus.processing
|
||||||
await session.commit() # visible before the long LLM call
|
await session.commit() # visible before the long LLM call
|
||||||
try:
|
try:
|
||||||
result = await provider.summarize(text, title=rec.title,
|
result = await provider.summarize(
|
||||||
tone=job.tone)
|
text, title=rec.title, tone=job.tone,
|
||||||
|
on_progress=_async_progress_writer(job.id))
|
||||||
except AIError as e:
|
except AIError as e:
|
||||||
await _fail(session, rec, job, str(e), ctx, e)
|
await _fail(session, rec, job, str(e), ctx, e)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
|
|
||||||
|
|
@ -129,6 +129,32 @@ async def test_openai_compat_parses_and_sanitizes():
|
||||||
assert res.model == "qwen"
|
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():
|
async def test_openai_compat_partial_json_gets_defaults():
|
||||||
import json as _json
|
import json as _json
|
||||||
|
|
||||||
|
|
@ -165,14 +191,21 @@ async def test_ollama_happy_path():
|
||||||
async def ok(request: httpx.Request) -> httpx.Response:
|
async def ok(request: httpx.Request) -> httpx.Response:
|
||||||
assert request.url.path == "/api/chat"
|
assert request.url.path == "/api/chat"
|
||||||
payload = _json.loads(request.content.decode())
|
payload = _json.loads(request.content.decode())
|
||||||
assert payload["format"] == "json" and payload["stream"] is False
|
assert payload["format"] == "json" and payload["stream"] is True
|
||||||
return httpx.Response(200, json={
|
body = "\n".join([
|
||||||
"message": {"content": _json.dumps({"short": "S.", "detailed": "D."})}
|
_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",
|
p = OllamaProvider("http://ollama:11434", model="llama3",
|
||||||
http_client=mock_client(ok))
|
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 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():
|
async def test_ollama_404_is_config_error():
|
||||||
|
|
|
||||||
|
|
@ -95,7 +95,8 @@ class FakeLlm:
|
||||||
self.fail = fail
|
self.fail = fail
|
||||||
self.seen = []
|
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)
|
self.seen.append(transcript)
|
||||||
if self.fail is not None:
|
if self.fail is not None:
|
||||||
raise self.fail
|
raise self.fail
|
||||||
|
|
|
||||||
|
|
@ -131,7 +131,8 @@ async def test_user_edit_survives_auto_pipeline(client, monkeypatch):
|
||||||
class FakeLlm:
|
class FakeLlm:
|
||||||
name = "fake-llm"
|
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=(),
|
return SummaryResult(short="s", detailed="d", key_points=(),
|
||||||
decisions=(), action_items=(), questions=(),
|
decisions=(), action_items=(), questions=(),
|
||||||
model="fake-llm-1")
|
model="fake-llm-1")
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue