엔진 게이트웨이 계약 고정
This commit is contained in:
parent
ebef20560e
commit
f0771db919
8 changed files with 1123 additions and 72 deletions
|
|
@ -7,7 +7,8 @@ gateway, and a future Node.js gateway must preserve these shapes.
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Literal, Optional
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal, Optional, cast
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
|
@ -74,9 +75,128 @@ class StreamErrorEvent(BaseModel):
|
|||
detail: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EngineGatewaySsePacket:
|
||||
event: EngineGatewaySseEvent
|
||||
payload: StreamTokenEvent | StreamDoneEvent | StreamErrorEvent
|
||||
|
||||
|
||||
class EngineGatewaySseDecodeError(ValueError):
|
||||
"""Gateway SSE line does not match the token/done/error wire contract."""
|
||||
|
||||
|
||||
def sse_frame(event: EngineGatewaySseEvent, payload: BaseModel | dict[str, Any]) -> str:
|
||||
if isinstance(payload, BaseModel):
|
||||
body = payload.model_dump()
|
||||
else:
|
||||
body = payload
|
||||
return f"event: {event}\ndata: {json.dumps(body, ensure_ascii=False)}\n\n"
|
||||
|
||||
|
||||
def parse_sse_event(value: str) -> EngineGatewaySseEvent:
|
||||
event = value.strip()
|
||||
if event in ENGINE_GATEWAY_SSE_EVENTS:
|
||||
return cast(EngineGatewaySseEvent, event)
|
||||
raise EngineGatewaySseDecodeError(f"unknown engine gateway SSE event: {event or '<empty>'}")
|
||||
|
||||
|
||||
def parse_sse_payload(
|
||||
event: EngineGatewaySseEvent,
|
||||
data: str,
|
||||
*,
|
||||
allow_legacy_token_text: bool = True,
|
||||
) -> EngineGatewaySsePacket | None:
|
||||
payload_text = data.strip()
|
||||
if not payload_text or payload_text == "[DONE]":
|
||||
return None
|
||||
|
||||
payload: Any
|
||||
try:
|
||||
payload = json.loads(payload_text)
|
||||
except json.JSONDecodeError as exc:
|
||||
if event == ENGINE_GATEWAY_SSE_TOKEN and allow_legacy_token_text:
|
||||
payload = payload_text
|
||||
elif event == ENGINE_GATEWAY_SSE_ERROR:
|
||||
payload = {"detail": payload_text}
|
||||
else:
|
||||
raise EngineGatewaySseDecodeError(
|
||||
f"invalid engine gateway SSE {event} JSON payload"
|
||||
) from exc
|
||||
|
||||
if event == ENGINE_GATEWAY_SSE_TOKEN:
|
||||
token = _to_stream_token_event(payload, allow_legacy_token_text)
|
||||
if token is None:
|
||||
return None
|
||||
return EngineGatewaySsePacket(event=event, payload=token)
|
||||
if event == ENGINE_GATEWAY_SSE_DONE:
|
||||
return EngineGatewaySsePacket(event=event, payload=_to_stream_done_event(payload))
|
||||
if event == ENGINE_GATEWAY_SSE_ERROR:
|
||||
return EngineGatewaySsePacket(event=event, payload=_to_stream_error_event(payload))
|
||||
raise EngineGatewaySseDecodeError(f"unknown engine gateway SSE event: {event}")
|
||||
|
||||
|
||||
class EngineGatewaySseLineDecoder:
|
||||
"""Stateful decoder for raw gateway SSE lines."""
|
||||
|
||||
def __init__(self, default_event: EngineGatewaySseEvent = ENGINE_GATEWAY_SSE_TOKEN) -> None:
|
||||
self.current_event = default_event
|
||||
|
||||
def feed_line(self, raw_line: str) -> EngineGatewaySsePacket | None:
|
||||
line = raw_line.strip()
|
||||
if not line:
|
||||
return None
|
||||
if line.startswith("event:"):
|
||||
self.current_event = parse_sse_event(line[len("event:"):])
|
||||
return None
|
||||
if not line.startswith("data:"):
|
||||
return None
|
||||
return parse_sse_payload(
|
||||
self.current_event,
|
||||
line[len("data:"):],
|
||||
)
|
||||
|
||||
|
||||
def _to_stream_token_event(
|
||||
payload: Any,
|
||||
allow_legacy_token_text: bool,
|
||||
) -> StreamTokenEvent | None:
|
||||
if isinstance(payload, dict) and "text" in payload:
|
||||
return StreamTokenEvent(text=str(payload["text"]))
|
||||
if isinstance(payload, str) and allow_legacy_token_text:
|
||||
return StreamTokenEvent(text=payload)
|
||||
return None
|
||||
|
||||
|
||||
def _to_stream_done_event(payload: Any) -> StreamDoneEvent:
|
||||
if not isinstance(payload, dict):
|
||||
payload = {}
|
||||
return StreamDoneEvent(
|
||||
provider=str(payload.get("provider") or ""),
|
||||
model=str(payload.get("model") or ""),
|
||||
tokens_in=_safe_int(payload.get("tokens_in")),
|
||||
tokens_out=_safe_int(payload.get("tokens_out")),
|
||||
cost_usd=_safe_float(payload.get("cost_usd")),
|
||||
turns=_safe_int(payload.get("turns")),
|
||||
)
|
||||
|
||||
|
||||
def _to_stream_error_event(payload: Any) -> StreamErrorEvent:
|
||||
if isinstance(payload, dict) and payload.get("detail"):
|
||||
return StreamErrorEvent(detail=str(payload["detail"]))
|
||||
if isinstance(payload, str) and payload:
|
||||
return StreamErrorEvent(detail=payload)
|
||||
return StreamErrorEvent(detail="engine stream error")
|
||||
|
||||
|
||||
def _safe_int(value: Any) -> int:
|
||||
try:
|
||||
return int(value or 0)
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
|
||||
|
||||
def _safe_float(value: Any) -> float:
|
||||
try:
|
||||
return float(value or 0.0)
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue