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
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue