vignette/apps/api/app/contracts/engine_gateway.py
2026-06-29 08:12:14 +09:00

250 lines
7.5 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,
)
ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL = "gateway-default"
def normalize_engine_gateway_model(model: Optional[str]) -> Optional[str]:
"""Return an explicit model override, or None for gateway default routing."""
value = (model or "").strip()
if not value or value == ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL:
return None
return value
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
def structured_payload_from_response(resp: GenerateResponse) -> dict[str, Any] | None:
"""Return structured output, or a JSON object embedded in legacy text."""
if isinstance(resp.structured, dict):
return resp.structured
raw = (resp.text or "").strip()
if not raw:
return None
raw = _strip_json_fence(raw)
parsed = _json_object_or_none(raw)
if parsed is not None:
return parsed
start, end = raw.find("{"), raw.rfind("}")
if 0 <= start < end:
return _json_object_or_none(raw[start : end + 1])
return 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 _strip_json_fence(value: str) -> str:
if not value.startswith("```"):
return value
text = value.split("```", 2)[1] if value.count("```") >= 2 else value.strip("`")
if text.lstrip().lower().startswith("json"):
text = text.lstrip()[4:]
return text.strip()
def _json_object_or_none(value: str) -> dict[str, Any] | None:
try:
parsed = json.loads(value)
except (json.JSONDecodeError, ValueError):
return None
return parsed if isinstance(parsed, dict) else None
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