엔진 게이트웨이 계약 고정

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
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

View file

@ -24,6 +24,8 @@ from .config import settings
from .contracts.engine_gateway import (
AIRole,
EngineMessage,
EngineGatewaySseLineDecoder,
EngineGatewaySsePacket,
GenerateRequest,
GenerateResponse,
StreamRequest,
@ -147,9 +149,8 @@ class EngineClient:
async def stream(self, req: StreamRequest) -> AsyncIterator[str]:
"""SSE 토큰 스트림 프록시.
게이트웨이 SSE(`text/event-stream`) data: 청크를 그대로 yield.
라우트(sessions.py) 이걸 받아 자체 SSE 이벤트(heartbeat 포함) 재방출.
TODO: 게이트웨이 이벤트 프레이밍 확정(토큰/usage/done 이벤트 구분).
게이트웨이 SSE(`text/event-stream`) 원시 non-empty line을 yield한다.
호출부는 raw line 대신 `stream_packets()` 사용한다.
"""
try:
async with self.client.stream(
@ -164,6 +165,19 @@ class EngineClient:
except httpx.HTTPError as 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)
engine_client = EngineClient()

View file

@ -31,6 +31,14 @@ from ..engine_client import (
GenerateResponse,
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 .persona import PersonaCard, PersonaStateContext
from .state_machine import SessionState, Stage
@ -345,31 +353,23 @@ async def run_turn_stream(
return
try:
current_event = "message"
started = time.perf_counter()
async for raw in engine.stream(req):
# engine_client.stream 은 게이트웨이 SSE 의 *원시 라인*을 그대로 yield 한다.
# 게이트웨이 프레이밍: "event: token|done|error" + "data: {...}".
line = raw.strip()
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")
async for packet in engine.stream_packets(req):
if packet.event == ENGINE_GATEWAY_SSE_ERROR:
payload = packet.payload
detail = payload.detail if isinstance(payload, StreamErrorEvent) else "engine stream error"
yield StreamEvent("error", {"detail": detail})
return
if current_event == "done":
if isinstance(payload, dict):
stream_meta = payload
if packet.event == ENGINE_GATEWAY_SSE_DONE:
payload = packet.payload
if isinstance(payload, StreamDoneEvent):
stream_meta = payload.model_dump()
break
text_piece = _payload_text(payload)
if text_piece is None:
payload = packet.payload
if not isinstance(payload, StreamTokenEvent):
continue
text_piece = payload.text
accumulated += text_piece
# 출력 가드레일(누적 스캔) — 수단정보 발견 시 차단·재생성 신호
@ -411,6 +411,8 @@ async def run_turn_stream(
"cost_usd": _safe_float(stream_meta.get("cost_usd")),
},
)
except EngineGatewaySseDecodeError as e:
yield StreamEvent("error", {"detail": str(e)})
except EngineError as e:
yield StreamEvent("error", {"detail": str(e)})
@ -427,34 +429,6 @@ async def _record_llm_audit(
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]:
if value is None:
return None
@ -462,14 +436,6 @@ def _optional_str(value: Any) -> Optional[str]:
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:
try:
return int(value or 0)

View file

@ -25,6 +25,7 @@ from app.contracts.engine_gateway import (
ENGINE_GATEWAY_SSE_ERROR,
ENGINE_GATEWAY_SSE_TOKEN,
EngineMessage as GwMessage,
GenerateResponse,
GenerateRequest as GwGenerateReq,
StreamDoneEvent,
StreamErrorEvent,
@ -419,16 +420,16 @@ async def v1_generate(req: GwGenerateReq):
structured = json.loads(text)
except (json.JSONDecodeError, TypeError):
structured = None # 파싱 실패는 호출부가 text 로 폴백
return {
"text": text,
"model": s.model or DEFAULT_MODEL or "claude-opus-4-8",
"provider": "claude_cli",
"tokens_in": 0,
"tokens_out": 0,
"cost_usd": result.get("cost_usd", 0.0),
"inference_geo": "us",
"structured": structured,
}
return GenerateResponse(
text=text,
model=s.model or DEFAULT_MODEL or "claude-opus-4-8",
provider="claude_cli",
tokens_in=0,
tokens_out=0,
cost_usd=result.get("cost_usd", 0.0),
inference_geo="us",
structured=structured,
).model_dump()
@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 json
import shutil
import subprocess
import unittest
from pathlib import Path
from unittest.mock import patch
from jsonschema import Draft202012Validator
from app import engine_client
from app.contracts import engine_gateway as contract
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:
def close(self):
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):
if "--model" not in args:
return None
@ -104,6 +193,160 @@ class GatewayModelTest(unittest.TestCase):
'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):
captured, process_patch = _capture_subprocess()
with (
@ -221,9 +464,11 @@ class GatewayModelTest(unittest.TestCase):
existing.close = fake_close
response = asyncio.run(gateway.v1_generate(_request(session_id="sid")))
validated = contract.GenerateResponse.model_validate(response)
self.assertEqual(response["text"], "reused response")
self.assertEqual(response["provider"], "claude_cli")
self.assertEqual(validated.text, "reused response")
self.assertEqual(validated.provider, "claude_cli")
self.assertEqual(validated.cost_usd, 0.01)
self.assertEqual(calls, [("hello", 120.0)])
self.assertEqual(closes, [])
@ -249,8 +494,11 @@ class GatewayModelTest(unittest.TestCase):
patch.object(gateway.EngineSession, "close", fake_close),
):
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(turned, [(started[0], "hello", 120.0)])
self.assertEqual(closed, [started[0]])