306 lines
9.1 KiB
Python
306 lines
9.1 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"]
|
|
EngineProvider = Literal[
|
|
"claude_cli",
|
|
"claude_api",
|
|
"codex_cli",
|
|
"agy_cli",
|
|
"openai",
|
|
"solar",
|
|
]
|
|
ReasoningEffort = Literal["low", "medium", "high", "xhigh", "max", "ultra"]
|
|
EngineCapabilitySource = Literal["live_cli", "live_api", "static_cli", "unavailable"]
|
|
|
|
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"
|
|
ENGINE_PROVIDERS: tuple[EngineProvider, ...] = (
|
|
"claude_cli",
|
|
"claude_api",
|
|
"codex_cli",
|
|
"agy_cli",
|
|
"openai",
|
|
"solar",
|
|
)
|
|
ENGINE_REASONING_EFFORTS: tuple[ReasoningEffort, ...] = (
|
|
"low",
|
|
"medium",
|
|
"high",
|
|
"xhigh",
|
|
"max",
|
|
"ultra",
|
|
)
|
|
ENGINE_PROVIDER_DEFAULTS: dict[EngineProvider, tuple[str, Optional[ReasoningEffort]]] = {
|
|
"claude_cli": (ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL, "high"),
|
|
"claude_api": (ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL, "high"),
|
|
"codex_cli": ("gpt-5.6-terra", "medium"),
|
|
"agy_cli": ("gemini-3.6-flash-high", "high"),
|
|
"openai": (ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL, None),
|
|
"solar": (ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL, None),
|
|
}
|
|
|
|
|
|
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]
|
|
provider: Optional[EngineProvider] = None
|
|
model: Optional[str] = None
|
|
reasoning_effort: Optional[ReasoningEffort] = 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 EngineModelOption(BaseModel):
|
|
id: str
|
|
label: str
|
|
description: str = ""
|
|
reasoning_efforts: list[ReasoningEffort] = Field(default_factory=list)
|
|
default_reasoning_effort: Optional[ReasoningEffort] = None
|
|
is_default: bool = False
|
|
|
|
|
|
class EngineCapabilitiesResponse(BaseModel):
|
|
provider: EngineProvider
|
|
available: bool
|
|
source: EngineCapabilitySource
|
|
models: list[EngineModelOption] = Field(default_factory=list)
|
|
default_model: Optional[str] = None
|
|
default_reasoning_effort: Optional[ReasoningEffort] = None
|
|
detail: str = ""
|
|
fetched_at: float
|
|
|
|
|
|
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
|