엔진 게이트웨이 계약 고정

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)