vignette/apps/api/app/contracts/engine_gateway.py
2026-06-28 20:12:05 +09:00

202 lines
6 KiB
Python

"""Engine gateway HTTP and SSE contract.
Keep this module adapter-neutral. The Python API client, the current Python
gateway, and a future Node.js gateway must preserve these shapes.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from typing import Any, Literal, Optional, cast
from pydantic import BaseModel, Field
AIRole = Literal["client", "counselor", "evaluator"]
EngineMessageRole = Literal["system", "user", "assistant"]
EngineGatewaySseEvent = Literal["token", "done", "error"]
ENGINE_GATEWAY_SSE_TOKEN: EngineGatewaySseEvent = "token"
ENGINE_GATEWAY_SSE_DONE: EngineGatewaySseEvent = "done"
ENGINE_GATEWAY_SSE_ERROR: EngineGatewaySseEvent = "error"
ENGINE_GATEWAY_SSE_EVENTS: tuple[EngineGatewaySseEvent, ...] = (
ENGINE_GATEWAY_SSE_TOKEN,
ENGINE_GATEWAY_SSE_DONE,
ENGINE_GATEWAY_SSE_ERROR,
)
class EngineMessage(BaseModel):
role: EngineMessageRole
content: str
cache: bool = False
class GenerateRequest(BaseModel):
ai_role: AIRole = "client"
messages: list[EngineMessage]
model: Optional[str] = None
max_tokens: int = 1024
temperature: float = 0.7
structured_schema: Optional[dict[str, Any]] = None
session_id: Optional[str] = None
metadata: dict[str, Any] = Field(default_factory=dict)
class StreamRequest(GenerateRequest):
"""SSE stream request for live client AI turns."""
class GenerateResponse(BaseModel):
text: str
model: str
provider: str
tokens_in: int = 0
tokens_out: int = 0
cost_usd: float = 0.0
inference_geo: Optional[str] = None
structured: Optional[dict[str, Any]] = None
class StreamTokenEvent(BaseModel):
text: str
class StreamDoneEvent(BaseModel):
provider: str
model: str
tokens_in: int = 0
tokens_out: int = 0
cost_usd: float = 0.0
turns: int = 0
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