p5jsai-api / upstream.py
li2895's picture
Fix upstream delta text parsing
671eec5
Raw
History Blame Contribute Delete
11.4 kB
import json
import logging
import time
import uuid
from typing import Any, AsyncIterator
import httpx
from fastapi import HTTPException
from config import (
UPSTREAM_URL, UPSTREAM_HEADERS, DEFAULT_MODEL,
UPSTREAM_TIMEOUT_SECS, UPSTREAM_CONNECT_TIMEOUT_SECS,
UPSTREAM_MAX_CONNECTIONS, UPSTREAM_MAX_KEEPALIVE_CONNECTIONS,
UPSTREAM_KEEPALIVE_EXPIRY_SECS, UPSTREAM_PROXY_URL,
)
from response_cache import CompletionArtifact
logger = logging.getLogger(__name__)
_UPSTREAM_CLIENT: httpx.AsyncClient | None = None
def create_upstream_client() -> httpx.AsyncClient:
transport_kwargs: dict[str, Any] = {}
if UPSTREAM_PROXY_URL:
transport_kwargs["proxy"] = UPSTREAM_PROXY_URL
transport = httpx.AsyncHTTPTransport(
retries=0,
limits=httpx.Limits(
max_connections=UPSTREAM_MAX_CONNECTIONS,
max_keepalive_connections=UPSTREAM_MAX_KEEPALIVE_CONNECTIONS,
keepalive_expiry=UPSTREAM_KEEPALIVE_EXPIRY_SECS,
),
**transport_kwargs,
)
return httpx.AsyncClient(
timeout=httpx.Timeout(UPSTREAM_TIMEOUT_SECS, connect=UPSTREAM_CONNECT_TIMEOUT_SECS),
transport=transport,
trust_env=True,
http2=False,
)
def get_upstream_client() -> httpx.AsyncClient:
global _UPSTREAM_CLIENT
if _UPSTREAM_CLIENT is None:
_UPSTREAM_CLIENT = create_upstream_client()
return _UPSTREAM_CLIENT
async def close_upstream_client() -> None:
global _UPSTREAM_CLIENT
if _UPSTREAM_CLIENT is not None:
await _UPSTREAM_CLIENT.aclose()
_UPSTREAM_CLIENT = None
async def iter_upstream_events(payload: dict) -> AsyncIterator[tuple[str, dict]]:
client = get_upstream_client()
request_started = time.time()
event_counts: dict[str, int] = {}
raw_line_samples: list[str] = []
parsed_samples: list[dict] = []
skipped_lines = 0
json_errors = 0
byte_count = 0
try:
async with client.stream("POST", UPSTREAM_URL, headers=UPSTREAM_HEADERS, json=payload) as r:
logger.info(
"upstream response status=%s content-type=%s model=%s messages=%s",
r.status_code,
r.headers.get("content-type"),
payload.get("model"),
len(payload.get("messages") or []),
)
if r.status_code != 200:
body = await r.aread()
err_body = body.decode("utf-8", errors="replace")[:2000]
logger.warning(
"upstream non-200 status=%s body=%r headers=%s",
r.status_code,
err_body[:1000],
dict(r.headers),
)
yield ("error", {"status_code": r.status_code, "body": err_body or f"upstream returned HTTP {r.status_code}"})
return
buffer = b""
async for chunk in r.aiter_bytes():
if not chunk:
continue
byte_count += len(chunk)
buffer += chunk
while b"\n" in buffer:
line, buffer = buffer.split(b"\n", 1)
s = line.decode("utf-8", errors="replace").rstrip("\r")
if s and len(raw_line_samples) < 8:
raw_line_samples.append(s[:500])
if not s or s.startswith("event:") or s.startswith("ddata:"):
skipped_lines += 1
continue
if s.startswith("ata:"):
s = "d" + s
if not s.startswith("data:"):
skipped_lines += 1
continue
data = s[5:].lstrip()
if data == "[DONE]":
event_counts["done"] = event_counts.get("done", 0) + 1
yield ("done", {})
continue
try:
obj = json.loads(data)
except Exception as exc:
json_errors += 1
if json_errors <= 3:
logger.warning("upstream JSON parse failed: %s raw=%r", exc, data[:500])
continue
if not isinstance(obj, dict):
skipped_lines += 1
continue
if len(parsed_samples) < 5:
parsed_samples.append(obj)
t = obj.get("type")
if not t or t == "job_id":
skipped_lines += 1
continue
if t == "delta" and "text" in obj and "index" not in obj:
obj = {
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": obj.get("text", "")},
}
t = "content_block_delta"
if t == "error":
yield ("error", {"status_code": 502, "body": obj.get("message") or json.dumps(obj, ensure_ascii=False)})
return
event_counts[t] = event_counts.get(t, 0) + 1
yield (t, obj)
if buffer.strip():
tail = buffer.decode("utf-8", errors="replace").strip()
if len(raw_line_samples) < 8:
raw_line_samples.append(tail[:500])
logger.info(
"upstream stream finished status=%s bytes=%s duration=%.2fs events=%s skipped=%s json_errors=%s samples=%s parsed_samples=%s",
r.status_code,
byte_count,
time.time() - request_started,
event_counts,
skipped_lines,
json_errors,
raw_line_samples,
parsed_samples,
)
except httpx.TimeoutException as exc:
logger.exception("upstream timeout after %.2fs model=%s error=%s", time.time() - request_started, payload.get("model"), exc)
yield ("error", {"status_code": 504, "body": f"upstream timeout: {type(exc).__name__}: {exc}"})
except httpx.HTTPError as exc:
logger.exception("upstream HTTP error after %.2fs model=%s error=%s", time.time() - request_started, payload.get("model"), exc)
yield ("error", {"status_code": 502, "body": f"upstream HTTP error: {type(exc).__name__}: {exc}"})
except Exception as exc:
logger.exception("upstream unexpected error after %.2fs model=%s error=%s", time.time() - request_started, payload.get("model"), exc)
yield ("error", {"status_code": 502, "body": f"upstream unexpected error: {type(exc).__name__}: {exc}"})
async def collect_upstream_text(payload: dict) -> tuple[str, str, dict, str, dict]:
"""Run through the whole upstream stream and collect:
(full_text, model_id, usage, stop_reason, message_start_obj)
"""
text_parts: list[str] = []
model_id = payload.get("model", DEFAULT_MODEL)
stop_reason = "end_turn"
usage = {"input_tokens": 0, "output_tokens": 0}
msg_start = {}
event_counts: dict[str, int] = {}
async for event_name, obj in iter_upstream_events(payload):
event_counts[event_name] = event_counts.get(event_name, 0) + 1
if event_name == "error":
logger.warning("upstream collection failed events=%s error=%s", event_counts, obj)
raise HTTPException(status_code=obj.get("status_code", 502), detail=obj.get("body", "upstream error"))
if event_name == "message_start":
msg_start = obj.get("message", {})
model_id = msg_start.get("model", model_id)
if "usage" in msg_start:
usage["input_tokens"] = msg_start["usage"].get("input_tokens", 0)
elif event_name == "content_block_delta":
delta = obj.get("delta", {})
if delta.get("type") == "text_delta":
text_parts.append(delta.get("text", ""))
elif event_name == "message_delta":
d = obj.get("delta", {})
if d.get("stop_reason"):
stop_reason = d["stop_reason"]
u = obj.get("usage")
if u and "output_tokens" in u:
usage["output_tokens"] = u["output_tokens"]
full_text = "".join(text_parts)
if not full_text:
logger.warning(
"upstream returned no text model=%s events=%s msg_start=%s usage=%s stop_reason=%s",
model_id,
event_counts,
msg_start,
usage,
stop_reason,
)
return (full_text, model_id, usage, stop_reason, msg_start)
def build_completion_artifact(raw_text: str, model_id: str, usage: dict, stop_reason: str) -> CompletionArtifact:
return CompletionArtifact(
schema_version=1,
raw_text=raw_text,
model_id=model_id,
usage_input_tokens=int(usage.get("input_tokens", 0)),
usage_output_tokens=int(usage.get("output_tokens", 0)),
stop_reason=stop_reason or "end_turn",
stored_at=time.time(),
)
class LiveArtifactCapture:
def __init__(self, fallback_model: str):
self.fallback_model = fallback_model
self.model_id = fallback_model
self.raw_text_parts: list[str] = []
self.usage = {"input_tokens": 0, "output_tokens": 0}
self.stop_reason = "end_turn"
self.failed = False
self.completed = False
def observe(self, event_name: str, obj: dict) -> None:
if event_name == "error":
self.failed = True
return
if event_name == "done":
self.completed = True
if not self.raw_text_parts:
logger.warning(
"live stream completed without text model=%s usage=%s stop_reason=%s",
self.model_id,
self.usage,
self.stop_reason,
)
return
if event_name == "message_start":
message = obj.get("message", {})
self.model_id = message.get("model", self.model_id)
if "usage" in message:
self.usage["input_tokens"] = message["usage"].get("input_tokens", 0)
return
if event_name == "content_block_delta":
delta = obj.get("delta", {})
if delta.get("type") == "text_delta":
self.raw_text_parts.append(delta.get("text", ""))
return
if event_name == "message_delta":
delta = obj.get("delta", {})
if delta.get("stop_reason"):
self.stop_reason = delta["stop_reason"]
usage = obj.get("usage")
if usage and "output_tokens" in usage:
self.usage["output_tokens"] = usage["output_tokens"]
def is_cacheable(self) -> bool:
return self.completed and not self.failed
def build(self) -> CompletionArtifact:
return build_completion_artifact("".join(self.raw_text_parts), self.model_id, self.usage, self.stop_reason)
async def fetch_completion_artifact(payload: dict) -> CompletionArtifact:
text, model_id, usage, stop_reason, _ = await collect_upstream_text(payload)
return build_completion_artifact(text, model_id, usage, stop_reason)