diff --git a/backend/shonar/services/ai/__init__.py b/backend/shonar/services/ai/__init__.py index c40114b..0cd13d9 100644 --- a/backend/shonar/services/ai/__init__.py +++ b/backend/shonar/services/ai/__init__.py @@ -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: diff --git a/backend/shonar/services/ai/_llm.py b/backend/shonar/services/ai/_llm.py index 611dc65..7cb04b6 100644 --- a/backend/shonar/services/ai/_llm.py +++ b/backend/shonar/services/ai/_llm.py @@ -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) diff --git a/backend/shonar/services/ai/ollama.py b/backend/shonar/services/ai/ollama.py index 45da6cb..424f96a 100644 --- a/backend/shonar/services/ai/ollama.py +++ b/backend/shonar/services/ai/ollama.py @@ -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 diff --git a/backend/shonar/services/ai/openai_compat.py b/backend/shonar/services/ai/openai_compat.py index a27dc3b..0fe4d9a 100644 --- a/backend/shonar/services/ai/openai_compat.py +++ b/backend/shonar/services/ai/openai_compat.py @@ -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: diff --git a/backend/shonar/services/processing.py b/backend/shonar/services/processing.py index cbf5ec4..912ff1e 100644 --- a/backend/shonar/services/processing.py +++ b/backend/shonar/services/processing.py @@ -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() diff --git a/backend/tests/test_ai_adapters.py b/backend/tests/test_ai_adapters.py index e9a4599..5c7d260 100644 --- a/backend/tests/test_ai_adapters.py +++ b/backend/tests/test_ai_adapters.py @@ -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(): diff --git a/backend/tests/test_ai_pipeline.py b/backend/tests/test_ai_pipeline.py index 89f50a5..71dd24d 100644 --- a/backend/tests/test_ai_pipeline.py +++ b/backend/tests/test_ai_pipeline.py @@ -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 diff --git a/backend/tests/test_transcript_edits.py b/backend/tests/test_transcript_edits.py index 6502516..a5b193d 100644 --- a/backend/tests/test_transcript_edits.py +++ b/backend/tests/test_transcript_edits.py @@ -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")