엔진 게이트웨이 계약 고정

This commit is contained in:
Yun Chan 2026-06-28 20:12:05 +09:00
parent ebef20560e
commit f0771db919
8 changed files with 1123 additions and 72 deletions

View file

@ -7,7 +7,8 @@ gateway, and a future Node.js gateway must preserve these shapes.
from __future__ import annotations from __future__ import annotations
import json import json
from typing import Any, Literal, Optional from dataclasses import dataclass
from typing import Any, Literal, Optional, cast
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
@ -74,9 +75,128 @@ class StreamErrorEvent(BaseModel):
detail: str 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: def sse_frame(event: EngineGatewaySseEvent, payload: BaseModel | dict[str, Any]) -> str:
if isinstance(payload, BaseModel): if isinstance(payload, BaseModel):
body = payload.model_dump() body = payload.model_dump()
else: else:
body = payload body = payload
return f"event: {event}\ndata: {json.dumps(body, ensure_ascii=False)}\n\n" 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

View file

@ -24,6 +24,8 @@ from .config import settings
from .contracts.engine_gateway import ( from .contracts.engine_gateway import (
AIRole, AIRole,
EngineMessage, EngineMessage,
EngineGatewaySseLineDecoder,
EngineGatewaySsePacket,
GenerateRequest, GenerateRequest,
GenerateResponse, GenerateResponse,
StreamRequest, StreamRequest,
@ -147,9 +149,8 @@ class EngineClient:
async def stream(self, req: StreamRequest) -> AsyncIterator[str]: async def stream(self, req: StreamRequest) -> AsyncIterator[str]:
"""SSE 토큰 스트림 프록시. """SSE 토큰 스트림 프록시.
게이트웨이 SSE(`text/event-stream`) data: 청크를 그대로 yield. 게이트웨이 SSE(`text/event-stream`) 원시 non-empty line을 yield한다.
라우트(sessions.py) 이걸 받아 자체 SSE 이벤트(heartbeat 포함) 재방출. 호출부는 raw line 대신 `stream_packets()` 사용한다.
TODO: 게이트웨이 이벤트 프레이밍 확정(토큰/usage/done 이벤트 구분).
""" """
try: try:
async with self.client.stream( async with self.client.stream(
@ -164,6 +165,19 @@ class EngineClient:
except httpx.HTTPError as e: except httpx.HTTPError as e:
raise EngineError(f"engine stream transport error: {e}") from e raise EngineError(f"engine stream transport error: {e}") from e
async def stream_packets(self, req: StreamRequest) -> AsyncIterator[EngineGatewaySsePacket]:
"""Decode gateway SSE into the shared token/done/error contract.
Keep wire-format parsing at the gateway client boundary so app services
do not depend on raw SSE line structure. A future Node.js gateway should
only need to preserve `app.contracts.engine_gateway`.
"""
decoder = EngineGatewaySseLineDecoder()
async for raw in self.stream(req):
packet = decoder.feed_line(raw)
if packet is not None:
yield packet
# 앱 전역 싱글톤 (main lifespan 에서 startup/shutdown) # 앱 전역 싱글톤 (main lifespan 에서 startup/shutdown)
engine_client = EngineClient() engine_client = EngineClient()

View file

@ -31,6 +31,14 @@ from ..engine_client import (
GenerateResponse, GenerateResponse,
StreamRequest, StreamRequest,
) )
from ..contracts.engine_gateway import (
ENGINE_GATEWAY_SSE_DONE,
ENGINE_GATEWAY_SSE_ERROR,
EngineGatewaySseDecodeError,
StreamDoneEvent,
StreamErrorEvent,
StreamTokenEvent,
)
from . import guardrail, persona, state_machine from . import guardrail, persona, state_machine
from .persona import PersonaCard, PersonaStateContext from .persona import PersonaCard, PersonaStateContext
from .state_machine import SessionState, Stage from .state_machine import SessionState, Stage
@ -345,31 +353,23 @@ async def run_turn_stream(
return return
try: try:
current_event = "message"
started = time.perf_counter() started = time.perf_counter()
async for raw in engine.stream(req): async for packet in engine.stream_packets(req):
# engine_client.stream 은 게이트웨이 SSE 의 *원시 라인*을 그대로 yield 한다. if packet.event == ENGINE_GATEWAY_SSE_ERROR:
# 게이트웨이 프레이밍: "event: token|done|error" + "data: {...}". payload = packet.payload
line = raw.strip() detail = payload.detail if isinstance(payload, StreamErrorEvent) else "engine stream error"
if line.startswith("event:"):
current_event = line[len("event:"):].strip() or "message"
continue
if not line.startswith("data:"):
continue
payload = _extract_sse_payload(line)
if current_event == "error":
detail = _payload_detail(payload, "engine stream error")
yield StreamEvent("error", {"detail": detail}) yield StreamEvent("error", {"detail": detail})
return return
if current_event == "done": if packet.event == ENGINE_GATEWAY_SSE_DONE:
if isinstance(payload, dict): payload = packet.payload
stream_meta = payload if isinstance(payload, StreamDoneEvent):
stream_meta = payload.model_dump()
break break
text_piece = _payload_text(payload) payload = packet.payload
if text_piece is None: if not isinstance(payload, StreamTokenEvent):
continue continue
text_piece = payload.text
accumulated += text_piece accumulated += text_piece
# 출력 가드레일(누적 스캔) — 수단정보 발견 시 차단·재생성 신호 # 출력 가드레일(누적 스캔) — 수단정보 발견 시 차단·재생성 신호
@ -411,6 +411,8 @@ async def run_turn_stream(
"cost_usd": _safe_float(stream_meta.get("cost_usd")), "cost_usd": _safe_float(stream_meta.get("cost_usd")),
}, },
) )
except EngineGatewaySseDecodeError as e:
yield StreamEvent("error", {"detail": str(e)})
except EngineError as e: except EngineError as e:
yield StreamEvent("error", {"detail": str(e)}) yield StreamEvent("error", {"detail": str(e)})
@ -427,34 +429,6 @@ async def _record_llm_audit(
return return
def _extract_sse_payload(raw_line: str) -> Any:
"""게이트웨이 SSE data 라인의 JSON payload를 추출.
token은 {"text": "..."}이고, done/error도 JSON 객체다. 구형/테스트 fixture가
plain text data를 보내면 문자열 그대로 반환한다.
"""
import json as _json
line = raw_line.strip()
if not line.startswith("data:"):
return None
payload = line[len("data:"):].strip()
if not payload or payload == "[DONE]":
return None
try:
return _json.loads(payload)
except _json.JSONDecodeError:
return payload
def _payload_text(payload: Any) -> Optional[str]:
if isinstance(payload, dict) and "text" in payload:
return str(payload["text"])
if isinstance(payload, str):
return payload
return None
def _optional_str(value: Any) -> Optional[str]: def _optional_str(value: Any) -> Optional[str]:
if value is None: if value is None:
return None return None
@ -462,14 +436,6 @@ def _optional_str(value: Any) -> Optional[str]:
return text or None return text or None
def _payload_detail(payload: Any, fallback: str) -> str:
if isinstance(payload, dict) and payload.get("detail"):
return str(payload["detail"])
if isinstance(payload, str) and payload:
return payload
return fallback
def _safe_int(value: Any) -> int: def _safe_int(value: Any) -> int:
try: try:
return int(value or 0) return int(value or 0)

View file

@ -25,6 +25,7 @@ from app.contracts.engine_gateway import (
ENGINE_GATEWAY_SSE_ERROR, ENGINE_GATEWAY_SSE_ERROR,
ENGINE_GATEWAY_SSE_TOKEN, ENGINE_GATEWAY_SSE_TOKEN,
EngineMessage as GwMessage, EngineMessage as GwMessage,
GenerateResponse,
GenerateRequest as GwGenerateReq, GenerateRequest as GwGenerateReq,
StreamDoneEvent, StreamDoneEvent,
StreamErrorEvent, StreamErrorEvent,
@ -419,16 +420,16 @@ async def v1_generate(req: GwGenerateReq):
structured = json.loads(text) structured = json.loads(text)
except (json.JSONDecodeError, TypeError): except (json.JSONDecodeError, TypeError):
structured = None # 파싱 실패는 호출부가 text 로 폴백 structured = None # 파싱 실패는 호출부가 text 로 폴백
return { return GenerateResponse(
"text": text, text=text,
"model": s.model or DEFAULT_MODEL or "claude-opus-4-8", model=s.model or DEFAULT_MODEL or "claude-opus-4-8",
"provider": "claude_cli", provider="claude_cli",
"tokens_in": 0, tokens_in=0,
"tokens_out": 0, tokens_out=0,
"cost_usd": result.get("cost_usd", 0.0), cost_usd=result.get("cost_usd", 0.0),
"inference_geo": "us", inference_geo="us",
"structured": structured, structured=structured,
} ).model_dump()
@app.post("/v1/stream") @app.post("/v1/stream")

View file

@ -0,0 +1,90 @@
{
"version": 1,
"generate_request": {
"ai_role": "client",
"messages": [
{
"role": "system",
"content": "Keep replies concise.",
"cache": true
},
{
"role": "user",
"content": "hello",
"cache": false
}
],
"model": "gateway-default",
"max_tokens": 256,
"temperature": 0.4,
"structured_schema": {
"type": "object",
"properties": {
"reply": {
"type": "string"
}
},
"required": [
"reply"
],
"additionalProperties": false
},
"session_id": "contract-session-1",
"metadata": {
"trace_id": "trace-001",
"route": "golden"
}
},
"generate_response": {
"text": "{\"reply\":\"hello\"}",
"model": "claude-opus-4-8",
"provider": "claude_cli",
"tokens_in": 11,
"tokens_out": 7,
"cost_usd": 0.0123,
"inference_geo": "us",
"structured": {
"reply": "hello"
}
},
"stream_frames": [
"event: token\ndata: {\"text\":\"hel\"}\n\n",
"event: token\ndata: {\"text\":\"lo\"}\n\n",
"event: done\ndata: {\"provider\":\"claude_cli\",\"model\":\"claude-opus-4-8\",\"tokens_in\":11,\"tokens_out\":7,\"cost_usd\":0.0123,\"turns\":2}\n\n",
"event: error\ndata: {\"detail\":\"engine failed\"}\n\n"
],
"stream_packets": [
{
"event": "token",
"payload": {
"text": "hel"
}
},
{
"event": "token",
"payload": {
"text": "lo"
}
},
{
"event": "done",
"payload": {
"provider": "claude_cli",
"model": "claude-opus-4-8",
"tokens_in": 11,
"tokens_out": 7,
"cost_usd": 0.0123,
"turns": 2
}
},
{
"event": "error",
"payload": {
"detail": "engine failed"
}
}
],
"compatibility_lines": [
"data: [DONE]"
]
}

View file

@ -0,0 +1,333 @@
{
"$defs": {
"EngineMessage": {
"properties": {
"cache": {
"default": false,
"title": "Cache",
"type": "boolean"
},
"content": {
"title": "Content",
"type": "string"
},
"role": {
"enum": [
"system",
"user",
"assistant"
],
"title": "Role",
"type": "string"
}
},
"required": [
"role",
"content"
],
"title": "EngineMessage",
"type": "object"
},
"GenerateRequest": {
"properties": {
"ai_role": {
"default": "client",
"enum": [
"client",
"counselor",
"evaluator"
],
"title": "Ai Role",
"type": "string"
},
"max_tokens": {
"default": 1024,
"title": "Max Tokens",
"type": "integer"
},
"messages": {
"items": {
"$ref": "#/$defs/EngineMessage"
},
"title": "Messages",
"type": "array"
},
"metadata": {
"title": "Metadata",
"type": "object"
},
"model": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"default": null,
"title": "Model"
},
"session_id": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"default": null,
"title": "Session Id"
},
"structured_schema": {
"anyOf": [
{
"type": "object"
},
{
"type": "null"
}
],
"default": null,
"title": "Structured Schema"
},
"temperature": {
"default": 0.7,
"title": "Temperature",
"type": "number"
}
},
"required": [
"messages"
],
"title": "GenerateRequest",
"type": "object"
},
"GenerateResponse": {
"properties": {
"cost_usd": {
"default": 0.0,
"title": "Cost Usd",
"type": "number"
},
"inference_geo": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"default": null,
"title": "Inference Geo"
},
"model": {
"title": "Model",
"type": "string"
},
"provider": {
"title": "Provider",
"type": "string"
},
"structured": {
"anyOf": [
{
"type": "object"
},
{
"type": "null"
}
],
"default": null,
"title": "Structured"
},
"text": {
"title": "Text",
"type": "string"
},
"tokens_in": {
"default": 0,
"title": "Tokens In",
"type": "integer"
},
"tokens_out": {
"default": 0,
"title": "Tokens Out",
"type": "integer"
}
},
"required": [
"text",
"model",
"provider"
],
"title": "GenerateResponse",
"type": "object"
},
"StreamDoneEvent": {
"properties": {
"cost_usd": {
"default": 0.0,
"title": "Cost Usd",
"type": "number"
},
"model": {
"title": "Model",
"type": "string"
},
"provider": {
"title": "Provider",
"type": "string"
},
"tokens_in": {
"default": 0,
"title": "Tokens In",
"type": "integer"
},
"tokens_out": {
"default": 0,
"title": "Tokens Out",
"type": "integer"
},
"turns": {
"default": 0,
"title": "Turns",
"type": "integer"
}
},
"required": [
"provider",
"model"
],
"title": "StreamDoneEvent",
"type": "object"
},
"StreamErrorEvent": {
"properties": {
"detail": {
"title": "Detail",
"type": "string"
}
},
"required": [
"detail"
],
"title": "StreamErrorEvent",
"type": "object"
},
"StreamPacket": {
"oneOf": [
{
"additionalProperties": false,
"properties": {
"event": {
"const": "token"
},
"payload": {
"$ref": "#/$defs/StreamTokenEvent"
}
},
"required": [
"event",
"payload"
],
"type": "object"
},
{
"additionalProperties": false,
"properties": {
"event": {
"const": "done"
},
"payload": {
"$ref": "#/$defs/StreamDoneEvent"
}
},
"required": [
"event",
"payload"
],
"type": "object"
},
{
"additionalProperties": false,
"properties": {
"event": {
"const": "error"
},
"payload": {
"$ref": "#/$defs/StreamErrorEvent"
}
},
"required": [
"event",
"payload"
],
"type": "object"
}
]
},
"StreamTokenEvent": {
"properties": {
"text": {
"title": "Text",
"type": "string"
}
},
"required": [
"text"
],
"title": "StreamTokenEvent",
"type": "object"
}
},
"$id": "https://vignette.local/schemas/engine_gateway_contract.v1.json",
"$schema": "https://json-schema.org/draft/2020-12/schema",
"additionalProperties": false,
"properties": {
"compatibility_lines": {
"items": {
"type": "string"
},
"type": "array"
},
"generate_request": {
"$ref": "#/$defs/GenerateRequest"
},
"generate_response": {
"$ref": "#/$defs/GenerateResponse"
},
"stream_frames": {
"items": {
"type": "string"
},
"type": "array"
},
"stream_packets": {
"items": {
"$ref": "#/$defs/StreamPacket"
},
"type": "array"
},
"version": {
"const": 1
}
},
"required": [
"version",
"generate_request",
"generate_response",
"stream_frames",
"stream_packets",
"compatibility_lines"
],
"title": "EngineGatewayGoldenContract",
"type": "object",
"x-engine-gateway-sse-events": [
"token",
"done",
"error"
]
}

View file

@ -1,12 +1,22 @@
import asyncio import asyncio
import json
import shutil
import subprocess
import unittest import unittest
from pathlib import Path
from unittest.mock import patch from unittest.mock import patch
from jsonschema import Draft202012Validator
from app import engine_client from app import engine_client
from app.contracts import engine_gateway as contract from app.contracts import engine_gateway as contract
from engine_gateway import gateway from engine_gateway import gateway
GOLDEN_CONTRACT_PATH = Path(__file__).with_name("golden") / "engine_gateway_contract.v1.json"
GOLDEN_SCHEMA_PATH = Path(__file__).with_name("golden") / "engine_gateway_schema.v1.json"
class _FakeStdin: class _FakeStdin:
def close(self): def close(self):
pass pass
@ -49,6 +59,85 @@ def _request(model=None, session_id=None):
) )
def _load_golden_contract():
return json.loads(GOLDEN_CONTRACT_PATH.read_text(encoding="utf-8"))
def _load_golden_schema():
return json.loads(GOLDEN_SCHEMA_PATH.read_text(encoding="utf-8"))
def _contract_schema_doc():
defs = {}
def add_model(name, model):
schema = model.model_json_schema(ref_template="#/$defs/{model}")
defs.update(schema.pop("$defs", {}))
defs[name] = schema
add_model("GenerateRequest", contract.GenerateRequest)
add_model("GenerateResponse", contract.GenerateResponse)
add_model("StreamTokenEvent", contract.StreamTokenEvent)
add_model("StreamDoneEvent", contract.StreamDoneEvent)
add_model("StreamErrorEvent", contract.StreamErrorEvent)
defs["StreamPacket"] = {
"oneOf": [
{
"type": "object",
"additionalProperties": False,
"required": ["event", "payload"],
"properties": {
"event": {"const": "token"},
"payload": {"$ref": "#/$defs/StreamTokenEvent"},
},
},
{
"type": "object",
"additionalProperties": False,
"required": ["event", "payload"],
"properties": {
"event": {"const": "done"},
"payload": {"$ref": "#/$defs/StreamDoneEvent"},
},
},
{
"type": "object",
"additionalProperties": False,
"required": ["event", "payload"],
"properties": {
"event": {"const": "error"},
"payload": {"$ref": "#/$defs/StreamErrorEvent"},
},
},
]
}
return {
"$schema": "https://json-schema.org/draft/2020-12/schema",
"$id": "https://vignette.local/schemas/engine_gateway_contract.v1.json",
"title": "EngineGatewayGoldenContract",
"type": "object",
"additionalProperties": False,
"required": [
"version",
"generate_request",
"generate_response",
"stream_frames",
"stream_packets",
"compatibility_lines",
],
"properties": {
"version": {"const": 1},
"generate_request": {"$ref": "#/$defs/GenerateRequest"},
"generate_response": {"$ref": "#/$defs/GenerateResponse"},
"stream_frames": {"type": "array", "items": {"type": "string"}},
"stream_packets": {"type": "array", "items": {"$ref": "#/$defs/StreamPacket"}},
"compatibility_lines": {"type": "array", "items": {"type": "string"}},
},
"$defs": defs,
"x-engine-gateway-sse-events": list(contract.ENGINE_GATEWAY_SSE_EVENTS),
}
def _model_arg(args): def _model_arg(args):
if "--model" not in args: if "--model" not in args:
return None return None
@ -104,6 +193,160 @@ class GatewayModelTest(unittest.TestCase):
'event: error\ndata: {"detail": "failed"}\n\n', 'event: error\ndata: {"detail": "failed"}\n\n',
) )
def test_sse_line_decoder_parses_token_done_error_contract(self):
body = (
contract.sse_frame("token", contract.StreamTokenEvent(text="안녕"))
+ contract.sse_frame(
"done",
contract.StreamDoneEvent(
provider="claude_cli",
model="stream-model",
tokens_in=3,
tokens_out=5,
cost_usd=0.012,
turns=2,
),
)
+ contract.sse_frame("error", contract.StreamErrorEvent(detail="failed"))
)
decoder = contract.EngineGatewaySseLineDecoder()
packets = []
for line in body.splitlines():
packet = decoder.feed_line(line)
if packet is not None:
packets.append(packet)
self.assertEqual([packet.event for packet in packets], ["token", "done", "error"])
self.assertIsInstance(packets[0].payload, contract.StreamTokenEvent)
self.assertEqual(packets[0].payload.text, "안녕")
self.assertIsInstance(packets[1].payload, contract.StreamDoneEvent)
self.assertEqual(packets[1].payload.provider, "claude_cli")
self.assertEqual(packets[1].payload.model, "stream-model")
self.assertEqual(packets[1].payload.tokens_in, 3)
self.assertEqual(packets[1].payload.tokens_out, 5)
self.assertEqual(packets[1].payload.cost_usd, 0.012)
self.assertIsInstance(packets[2].payload, contract.StreamErrorEvent)
self.assertEqual(packets[2].payload.detail, "failed")
self.assertIn('data: {"text": "안녕"}', body)
def test_sse_line_decoder_keeps_legacy_plain_text_token_compatibility(self):
decoder = contract.EngineGatewaySseLineDecoder()
self.assertIsNone(decoder.feed_line("event: token"))
packet = decoder.feed_line("data: plain token")
self.assertIsNotNone(packet)
assert packet is not None
self.assertEqual(packet.event, "token")
self.assertIsInstance(packet.payload, contract.StreamTokenEvent)
self.assertEqual(packet.payload.text, "plain token")
def test_sse_line_decoder_ignores_provider_done_sentinel(self):
decoder = contract.EngineGatewaySseLineDecoder()
self.assertIsNone(decoder.feed_line("data: [DONE]"))
def test_sse_line_decoder_rejects_invalid_event_and_invalid_done_json(self):
decoder = contract.EngineGatewaySseLineDecoder()
with self.assertRaises(contract.EngineGatewaySseDecodeError):
decoder.feed_line("event: progress")
with self.assertRaises(contract.EngineGatewaySseDecodeError):
contract.parse_sse_payload("done", "not-json")
def test_engine_gateway_golden_fixture_validates_shared_wire_shapes(self):
golden = _load_golden_contract()
request = contract.GenerateRequest.model_validate(golden["generate_request"])
response = contract.GenerateResponse.model_validate(golden["generate_response"])
self.assertEqual(golden["version"], 1)
self.assertEqual(request.ai_role, "client")
self.assertEqual(request.messages[0].role, "system")
self.assertTrue(request.messages[0].cache)
self.assertEqual(request.structured_schema["required"], ["reply"])
self.assertEqual(response.structured, {"reply": "hello"})
decoder = contract.EngineGatewaySseLineDecoder()
packets = []
for frame in golden["stream_frames"]:
for line in frame.splitlines():
packet = decoder.feed_line(line)
if packet is not None:
packets.append(
{
"event": packet.event,
"payload": packet.payload.model_dump(),
}
)
self.assertEqual(packets, golden["stream_packets"])
compatibility_decoder = contract.EngineGatewaySseLineDecoder()
for line in golden["compatibility_lines"]:
self.assertIsNone(compatibility_decoder.feed_line(line))
def test_engine_gateway_json_schema_artifact_validates_golden_fixture(self):
schema = _load_golden_schema()
golden = _load_golden_contract()
self.assertEqual(schema, _contract_schema_doc())
Draft202012Validator.check_schema(schema)
Draft202012Validator(schema).validate(golden)
self.assertEqual(schema["x-engine-gateway-sse-events"], ["token", "done", "error"])
def test_node_conformance_runner_validates_engine_gateway_artifacts(self):
node = shutil.which("node")
if node is None:
self.skipTest("node executable is not available")
repo_root = Path(__file__).resolve().parents[3]
runner = repo_root / "scripts" / "check-engine-gateway-contract.mjs"
completed = subprocess.run(
[
node,
str(runner),
"--fixture",
str(GOLDEN_CONTRACT_PATH),
"--schema",
str(GOLDEN_SCHEMA_PATH),
"--json",
],
cwd=repo_root,
check=True,
capture_output=True,
text=True,
encoding="utf-8",
)
payload = json.loads(completed.stdout)
self.assertTrue(payload["ok"])
self.assertEqual(payload["streamEvents"], ["token", "done", "error"])
self.assertEqual(payload["decodedPacketCount"], 4)
def test_engine_client_stream_packets_owns_sse_decode_boundary(self):
client = engine_client.EngineClient(base_url="http://engine.test")
req = contract.StreamRequest(messages=[contract.EngineMessage(role="user", content="hello")])
async def fake_stream(_req):
yield "event: token"
yield 'data: {"text":"hello"}'
yield "event: done"
yield 'data: {"provider":"fake","model":"m","tokens_in":1,"tokens_out":2}'
async def collect():
with patch.object(client, "stream", fake_stream):
return [packet async for packet in client.stream_packets(req)]
packets = asyncio.run(collect())
self.assertEqual([packet.event for packet in packets], ["token", "done"])
self.assertIsInstance(packets[0].payload, contract.StreamTokenEvent)
self.assertEqual(packets[0].payload.text, "hello")
self.assertIsInstance(packets[1].payload, contract.StreamDoneEvent)
self.assertEqual(packets[1].payload.provider, "fake")
self.assertEqual(packets[1].payload.tokens_out, 2)
def test_resolve_session_uses_request_model_for_claude_cli(self): def test_resolve_session_uses_request_model_for_claude_cli(self):
captured, process_patch = _capture_subprocess() captured, process_patch = _capture_subprocess()
with ( with (
@ -221,9 +464,11 @@ class GatewayModelTest(unittest.TestCase):
existing.close = fake_close existing.close = fake_close
response = asyncio.run(gateway.v1_generate(_request(session_id="sid"))) response = asyncio.run(gateway.v1_generate(_request(session_id="sid")))
validated = contract.GenerateResponse.model_validate(response)
self.assertEqual(response["text"], "reused response") self.assertEqual(validated.text, "reused response")
self.assertEqual(response["provider"], "claude_cli") self.assertEqual(validated.provider, "claude_cli")
self.assertEqual(validated.cost_usd, 0.01)
self.assertEqual(calls, [("hello", 120.0)]) self.assertEqual(calls, [("hello", 120.0)])
self.assertEqual(closes, []) self.assertEqual(closes, [])
@ -249,8 +494,11 @@ class GatewayModelTest(unittest.TestCase):
patch.object(gateway.EngineSession, "close", fake_close), patch.object(gateway.EngineSession, "close", fake_close),
): ):
response = asyncio.run(gateway.v1_generate(_request(session_id="missing"))) response = asyncio.run(gateway.v1_generate(_request(session_id="missing")))
validated = contract.GenerateResponse.model_validate(response)
self.assertEqual(response["text"], "fresh response") self.assertEqual(validated.text, "fresh response")
self.assertEqual(validated.provider, "claude_cli")
self.assertEqual(validated.cost_usd, 0.02)
self.assertEqual(len(started), 1) self.assertEqual(len(started), 1)
self.assertEqual(turned, [(started[0], "hello", 120.0)]) self.assertEqual(turned, [(started[0], "hello", 120.0)])
self.assertEqual(closed, [started[0]]) self.assertEqual(closed, [started[0]])

View file

@ -0,0 +1,279 @@
#!/usr/bin/env node
import { readFileSync } from "node:fs";
import { resolve } from "node:path";
import assert from "node:assert/strict";
const DEFAULT_FIXTURE = "apps/api/engine_gateway/golden/engine_gateway_contract.v1.json";
const DEFAULT_SCHEMA = "apps/api/engine_gateway/golden/engine_gateway_schema.v1.json";
const EXPECTED_EVENTS = ["token", "done", "error"];
function usage() {
return [
"Usage: node scripts/check-engine-gateway-contract.mjs [--fixture path] [--schema path] [--json]",
"",
"Validates the engine gateway v1 golden fixture without importing Python code.",
].join("\n");
}
function parseArgs(argv) {
const args = {
fixture: DEFAULT_FIXTURE,
schema: DEFAULT_SCHEMA,
json: false,
};
for (let i = 0; i < argv.length; i += 1) {
const arg = argv[i];
if (arg === "--fixture") {
args.fixture = argv[i + 1];
i += 1;
} else if (arg === "--schema") {
args.schema = argv[i + 1];
i += 1;
} else if (arg === "--json") {
args.json = true;
} else if (arg === "--help" || arg === "-h") {
console.log(usage());
process.exit(0);
} else {
throw new Error(`Unknown argument: ${arg}\n${usage()}`);
}
}
if (!args.fixture || !args.schema) {
throw new Error(`Missing path argument.\n${usage()}`);
}
return args;
}
function readJson(path) {
return JSON.parse(readFileSync(resolve(path), "utf8"));
}
function assertPlainObject(value, label) {
assert.equal(typeof value, "object", `${label} must be an object`);
assert.notEqual(value, null, `${label} must not be null`);
assert.equal(Array.isArray(value), false, `${label} must not be an array`);
}
function assertOnlyKeys(value, allowedKeys, label) {
for (const key of Object.keys(value)) {
assert.ok(allowedKeys.includes(key), `${label}.${key} is not in the v1 contract`);
}
}
function assertRequired(value, requiredKeys, label) {
for (const key of requiredKeys) {
assert.ok(Object.hasOwn(value, key), `${label}.${key} is required`);
}
}
function assertNullableObject(value, label) {
if (value === null || value === undefined) return;
assertPlainObject(value, label);
}
function assertInteger(value, label) {
assert.equal(Number.isInteger(value), true, `${label} must be an integer`);
}
function assertNumber(value, label) {
assert.equal(typeof value, "number", `${label} must be a number`);
assert.equal(Number.isFinite(value), true, `${label} must be finite`);
}
function assertString(value, label) {
assert.equal(typeof value, "string", `${label} must be a string`);
}
function schemaDef(schema, name) {
const def = schema?.$defs?.[name];
assertPlainObject(def, `$defs.${name}`);
return def;
}
function assertRootShape(schema, fixture) {
assert.deepEqual(schema["x-engine-gateway-sse-events"], EXPECTED_EVENTS);
assertRequired(fixture, schema.required, "fixture");
assertOnlyKeys(fixture, Object.keys(schema.properties), "fixture");
assert.equal(fixture.version, schema.properties.version.const);
}
function assertGenerateRequest(schema, request) {
const requestDef = schemaDef(schema, "GenerateRequest");
const messageDef = schemaDef(schema, "EngineMessage");
const allowedRequestKeys = Object.keys(requestDef.properties);
const messageRoles = messageDef.properties.role.enum;
const aiRoles = requestDef.properties.ai_role.enum;
assertPlainObject(request, "generate_request");
assertRequired(request, requestDef.required, "generate_request");
assertOnlyKeys(request, allowedRequestKeys, "generate_request");
assert.ok(aiRoles.includes(request.ai_role), "generate_request.ai_role must be a known role");
assert.ok(Array.isArray(request.messages), "generate_request.messages must be an array");
assert.ok(request.messages.length > 0, "generate_request.messages must not be empty");
request.messages.forEach((message, index) => {
assertPlainObject(message, `generate_request.messages[${index}]`);
assertRequired(message, messageDef.required, `generate_request.messages[${index}]`);
assertOnlyKeys(message, Object.keys(messageDef.properties), `generate_request.messages[${index}]`);
assert.ok(messageRoles.includes(message.role), `generate_request.messages[${index}].role is invalid`);
assertString(message.content, `generate_request.messages[${index}].content`);
if (Object.hasOwn(message, "cache")) {
assert.equal(typeof message.cache, "boolean", `generate_request.messages[${index}].cache must be boolean`);
}
});
if (Object.hasOwn(request, "model")) {
assert.ok(request.model === null || typeof request.model === "string", "generate_request.model must be string or null");
}
if (Object.hasOwn(request, "session_id")) {
assert.ok(
request.session_id === null || typeof request.session_id === "string",
"generate_request.session_id must be string or null",
);
}
if (Object.hasOwn(request, "max_tokens")) assertInteger(request.max_tokens, "generate_request.max_tokens");
if (Object.hasOwn(request, "temperature")) assertNumber(request.temperature, "generate_request.temperature");
if (Object.hasOwn(request, "structured_schema")) assertNullableObject(request.structured_schema, "generate_request.structured_schema");
if (Object.hasOwn(request, "metadata")) assertPlainObject(request.metadata, "generate_request.metadata");
}
function assertGenerateResponse(schema, response) {
const responseDef = schemaDef(schema, "GenerateResponse");
assertPlainObject(response, "generate_response");
assertRequired(response, responseDef.required, "generate_response");
assertOnlyKeys(response, Object.keys(responseDef.properties), "generate_response");
assertString(response.text, "generate_response.text");
assertString(response.model, "generate_response.model");
assertString(response.provider, "generate_response.provider");
if (Object.hasOwn(response, "tokens_in")) assertInteger(response.tokens_in, "generate_response.tokens_in");
if (Object.hasOwn(response, "tokens_out")) assertInteger(response.tokens_out, "generate_response.tokens_out");
if (Object.hasOwn(response, "cost_usd")) assertNumber(response.cost_usd, "generate_response.cost_usd");
if (Object.hasOwn(response, "inference_geo")) {
assert.ok(
response.inference_geo === null || typeof response.inference_geo === "string",
"generate_response.inference_geo must be string or null",
);
}
if (Object.hasOwn(response, "structured")) assertNullableObject(response.structured, "generate_response.structured");
}
function parseSseFrame(frame) {
assertString(frame, "stream_frame");
let event = null;
const dataLines = [];
for (const line of frame.split(/\r?\n/)) {
if (line === "") continue;
if (line.startsWith("event:")) {
event = line.slice("event:".length).trim();
} else if (line.startsWith("data:")) {
dataLines.push(line.slice("data:".length).trimStart());
}
}
assert.ok(event, "stream frame must contain an event line");
assert.ok(EXPECTED_EVENTS.includes(event), `unknown stream event: ${event}`);
assert.ok(dataLines.length > 0, `stream event ${event} must contain data`);
return packetFromEventData(event, dataLines.join("\n"));
}
function packetFromEventData(event, data) {
if (data === "[DONE]") return null;
if (event === "token") {
let payload;
try {
payload = JSON.parse(data);
} catch {
payload = { text: data };
}
assertPlainObject(payload, "token payload");
assertString(payload.text, "token payload.text");
assertOnlyKeys(payload, ["text"], "token payload");
return { event, payload };
}
let payload;
try {
payload = JSON.parse(data);
} catch (error) {
throw new Error(`stream event ${event} must contain JSON data: ${error.message}`);
}
if (event === "done") {
assertPlainObject(payload, "done payload");
assertRequired(payload, ["provider", "model"], "done payload");
assertOnlyKeys(payload, ["provider", "model", "tokens_in", "tokens_out", "cost_usd", "turns"], "done payload");
assertString(payload.provider, "done payload.provider");
assertString(payload.model, "done payload.model");
if (Object.hasOwn(payload, "tokens_in")) assertInteger(payload.tokens_in, "done payload.tokens_in");
if (Object.hasOwn(payload, "tokens_out")) assertInteger(payload.tokens_out, "done payload.tokens_out");
if (Object.hasOwn(payload, "cost_usd")) assertNumber(payload.cost_usd, "done payload.cost_usd");
if (Object.hasOwn(payload, "turns")) assertInteger(payload.turns, "done payload.turns");
return { event, payload };
}
assert.equal(event, "error");
assertPlainObject(payload, "error payload");
assertRequired(payload, ["detail"], "error payload");
assertOnlyKeys(payload, ["detail"], "error payload");
assertString(payload.detail, "error payload.detail");
return { event, payload };
}
function assertStreamContract(fixture) {
assert.ok(Array.isArray(fixture.stream_frames), "stream_frames must be an array");
assert.ok(Array.isArray(fixture.stream_packets), "stream_packets must be an array");
assert.ok(Array.isArray(fixture.compatibility_lines), "compatibility_lines must be an array");
const decodedPackets = [];
for (const frame of fixture.stream_frames) {
const packet = parseSseFrame(frame);
if (packet) decodedPackets.push(packet);
}
assert.deepEqual(decodedPackets, fixture.stream_packets);
for (const line of fixture.compatibility_lines) {
assertString(line, "compatibility line");
if (line.trim() === "data: [DONE]") continue;
throw new Error(`unsupported compatibility line: ${line}`);
}
}
function validateContract(schema, fixture) {
assertRootShape(schema, fixture);
assertGenerateRequest(schema, fixture.generate_request);
assertGenerateResponse(schema, fixture.generate_response);
assertStreamContract(fixture);
return {
version: fixture.version,
streamEvents: schema["x-engine-gateway-sse-events"],
streamFrameCount: fixture.stream_frames.length,
decodedPacketCount: fixture.stream_packets.length,
compatibilityLineCount: fixture.compatibility_lines.length,
};
}
function main() {
const args = parseArgs(process.argv.slice(2));
const schema = readJson(args.schema);
const fixture = readJson(args.fixture);
const result = validateContract(schema, fixture);
if (args.json) {
console.log(JSON.stringify({ ok: true, ...result }, null, 2));
} else {
console.log(
`engine_gateway_contract.v${result.version} OK: ${result.decodedPacketCount} packets, ${result.compatibilityLineCount} compatibility lines`,
);
}
}
try {
main();
} catch (error) {
console.error(error instanceof Error ? error.message : String(error));
process.exit(1);
}