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 class _FakeProcess: def __init__(self): self.returncode = None self.stdin = _FakeStdin() self.stdout = None self.stderr = None async def wait(self): self.returncode = 0 def kill(self): self.returncode = -9 def _capture_subprocess(): captured = [] async def fake_create_subprocess_exec(*args, **kwargs): captured.append(args) return _FakeProcess() return captured, patch.object( gateway.asyncio, "create_subprocess_exec", fake_create_subprocess_exec, ) def _request(model=None, session_id=None): return gateway.GwGenerateReq( ai_role="client", messages=[gateway.GwMessage(role="user", content="hello")], model=model, session_id=session_id, ) 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), "x-engine-gateway-default-model-sentinel": contract.ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL, } def _model_arg(args): if "--model" not in args: return None return args[args.index("--model") + 1] async def _read_streaming_response(response): chunks = [] async for chunk in response.body_iterator: if isinstance(chunk, bytes): chunks.append(chunk.decode("utf-8")) else: chunks.append(str(chunk)) return "".join(chunks) class _FakeStreamSession: def __init__(self, events, model="test-model"): self.events = events self.model = model self.closed = False async def turn_stream(self, content, timeout=600.0): self.content = content self.timeout = timeout for event in self.events: yield event async def close(self): self.closed = True class GatewayModelTest(unittest.TestCase): def test_contract_owns_gateway_default_model_sentinel(self): self.assertEqual(contract.ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL, "gateway-default") self.assertIsNone(contract.normalize_engine_gateway_model(None)) self.assertIsNone(contract.normalize_engine_gateway_model("")) self.assertIsNone(contract.normalize_engine_gateway_model(" gateway-default ")) self.assertEqual( contract.normalize_engine_gateway_model(" request-model "), "request-model", ) def setUp(self): gateway.SESSIONS.clear() def tearDown(self): gateway.SESSIONS.clear() def test_gateway_reuses_shared_engine_contract_models(self): self.assertIs(gateway.GwGenerateReq, contract.GenerateRequest) self.assertIs(gateway.GwMessage, contract.EngineMessage) self.assertIs(engine_client.GenerateRequest, contract.GenerateRequest) self.assertEqual(contract.ENGINE_GATEWAY_SSE_EVENTS, ("token", "done", "error")) def test_split_messages_returns_named_current_turn_prompt_parts(self): parts = gateway._split_messages( [ contract.EngineMessage(role="system", content="system one"), contract.EngineMessage(role="system", content=""), contract.EngineMessage(role="system", content="system two"), contract.EngineMessage(role="assistant", content="previous counselor"), contract.EngineMessage(role="user", content="previous client"), contract.EngineMessage(role="user", content="current client"), ] ) self.assertIsInstance(parts, gateway.GatewayPromptParts) self.assertEqual(parts.system_prompt, "system one\n\nsystem two") self.assertEqual(parts.user_payload, "current client") def test_split_messages_injects_client_history_before_current_counselor_turn(self): parts = gateway._split_messages( [ contract.EngineMessage(role="system", content="client persona system"), contract.EngineMessage(role="user", content="상담자 이전 질문"), contract.EngineMessage(role="assistant", content="내담자 이전 답변"), contract.EngineMessage(role="user", content="이번 상담자 발화"), ], ai_role="client", ) self.assertEqual(parts.system_prompt, "client persona system") self.assertIn("[직전 대화]", parts.user_payload) self.assertIn("상담자: 상담자 이전 질문", parts.user_payload) self.assertIn("내담자: 내담자 이전 답변", parts.user_payload) self.assertIn("[이번 상담자 발화]", parts.user_payload) self.assertTrue(parts.user_payload.rstrip().endswith("이번 상담자 발화")) def test_split_messages_does_not_inject_history_for_evaluator_requests(self): parts = gateway._split_messages( [ contract.EngineMessage(role="system", content="eval system"), contract.EngineMessage(role="user", content="이전 평가 입력"), contract.EngineMessage(role="assistant", content="이전 평가 출력"), contract.EngineMessage(role="user", content="이번 평가 입력"), ], ai_role="evaluator", ) self.assertEqual(parts.system_prompt, "eval system") self.assertEqual(parts.user_payload, "이번 평가 입력") self.assertNotIn("[직전 대화]", parts.user_payload) def test_split_messages_preserves_no_user_payload_boundary(self): parts = gateway._split_messages( [contract.EngineMessage(role="system", content="system only")] ) self.assertEqual(parts.system_prompt, "system only") self.assertEqual(parts.user_payload, "") def test_sse_frame_helper_preserves_gateway_wire_contract(self): self.assertEqual( contract.sse_frame("token", contract.StreamTokenEvent(text="hello")), 'event: token\ndata: {"text": "hello"}\n\n', ) self.assertEqual( contract.sse_frame("error", contract.StreamErrorEvent(detail="failed")), '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_generate_response_structured_payload_prefers_structured_field(self): response = contract.GenerateResponse( text='{"reply":"text"}', provider="test", model="test-model", structured={"reply": "structured"}, ) self.assertEqual( contract.structured_payload_from_response(response), {"reply": "structured"}, ) def test_generate_response_structured_payload_accepts_fenced_and_embedded_json(self): fenced = contract.GenerateResponse( text='```json\n{"reply":"fenced"}\n```', provider="test", model="test-model", ) embedded = contract.GenerateResponse( text='prefix {"reply":"embedded"} suffix', provider="test", model="test-model", ) self.assertEqual(contract.structured_payload_from_response(fenced), {"reply": "fenced"}) self.assertEqual( contract.structured_payload_from_response(embedded), {"reply": "embedded"}, ) def test_generate_response_structured_payload_rejects_non_object_json(self): response = contract.GenerateResponse( text='["not", "object"]', provider="test", model="test-model", ) self.assertIsNone(contract.structured_payload_from_response(response)) 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["defaultModelSentinel"], contract.ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL, ) 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 ( patch.object(gateway, "DEFAULT_MODEL", "env-default"), patch.object(gateway, "FALLBACK_MODEL", ""), process_patch, ): session, ephemeral = asyncio.run( gateway._resolve_session(_request(model=" request-model "), "system prompt") ) try: self.assertIs(ephemeral, True) self.assertEqual(session.model, "request-model") self.assertEqual(_model_arg(captured[0]), "request-model") self.assertIn("--system-prompt", captured[0]) finally: asyncio.run(session.close()) def test_resolve_session_preserves_default_model_for_gateway_default(self): captured, process_patch = _capture_subprocess() with ( patch.object(gateway, "DEFAULT_MODEL", "env-default"), patch.object(gateway, "FALLBACK_MODEL", ""), process_patch, ): session, ephemeral = asyncio.run( gateway._resolve_session(_request(model="gateway-default"), "") ) try: self.assertIs(ephemeral, True) self.assertIsNone(session.model) self.assertEqual(_model_arg(captured[0]), "env-default") finally: asyncio.run(session.close()) def test_resolve_session_does_not_reuse_session_with_different_model(self): captured, process_patch = _capture_subprocess() existing = gateway.EngineSession(model="old-model") existing.proc = _FakeProcess() gateway.SESSIONS["sid"] = existing with ( patch.object(gateway, "DEFAULT_MODEL", ""), patch.object(gateway, "FALLBACK_MODEL", ""), process_patch, ): session, ephemeral = asyncio.run( gateway._resolve_session(_request(model="new-model", session_id="sid"), "") ) try: self.assertIs(ephemeral, True) self.assertIsNot(session, existing) self.assertEqual(_model_arg(captured[0]), "new-model") self.assertIs(gateway.SESSIONS["sid"], existing) finally: asyncio.run(session.close()) def test_resolve_session_reuses_live_session_id_without_starting_claude(self): captured, process_patch = _capture_subprocess() existing = gateway.EngineSession(model=None) existing.proc = _FakeProcess() gateway.SESSIONS["sid"] = existing with ( patch.object(gateway, "DEFAULT_MODEL", "env-default"), patch.object(gateway, "FALLBACK_MODEL", ""), process_patch, ): session, ephemeral = asyncio.run( gateway._resolve_session(_request(session_id="sid"), "new system prompt") ) self.assertIs(session, existing) self.assertIs(ephemeral, False) self.assertEqual(captured, []) def test_resolve_session_creates_fresh_ephemeral_for_missing_session_id(self): captured, process_patch = _capture_subprocess() with ( patch.object(gateway, "DEFAULT_MODEL", "env-default"), patch.object(gateway, "FALLBACK_MODEL", ""), process_patch, ): session, ephemeral = asyncio.run( gateway._resolve_session(_request(session_id="missing"), "system prompt") ) try: self.assertIs(ephemeral, True) self.assertNotIn(session.id, gateway.SESSIONS) self.assertEqual(len(captured), 1) self.assertIn("--system-prompt", captured[0]) finally: asyncio.run(session.close()) def test_v1_generate_reuses_session_id_without_ephemeral_close(self): existing = gateway.EngineSession(model=None) existing.proc = _FakeProcess() gateway.SESSIONS["sid"] = existing calls = [] closes = [] async def fake_turn(content, timeout=120.0): calls.append((content, timeout)) return {"text": "reused response", "cost_usd": 0.01, "is_error": False} async def fake_close(): closes.append(True) existing.turn = fake_turn existing.close = fake_close response = asyncio.run(gateway.v1_generate(_request(session_id="sid"))) validated = contract.GenerateResponse.model_validate(response) self.assertEqual(validated.text, "reused response") self.assertEqual(validated.provider, "claude_cli") self.assertEqual(validated.cost_usd, 0.01) self.assertEqual(calls, [("hello", gateway.GENERATE_TURN_TIMEOUT_SECONDS)]) self.assertEqual(closes, []) def test_v1_generate_closes_fresh_ephemeral_session(self): started = [] turned = [] closed = [] async def fake_start(self): started.append(self) self.proc = _FakeProcess() async def fake_turn(self, content, timeout=120.0): turned.append((self, content, timeout)) return {"text": "fresh response", "cost_usd": 0.02, "is_error": False} async def fake_close(self): closed.append(self) with ( patch.object(gateway.EngineSession, "start", fake_start), patch.object(gateway.EngineSession, "turn", fake_turn), 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(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", gateway.GENERATE_TURN_TIMEOUT_SECONDS)]) self.assertEqual(closed, [started[0]]) self.assertNotIn(started[0].id, gateway.SESSIONS) def test_v1_generate_rejects_missing_user_before_session_resolution(self): req = contract.GenerateRequest( messages=[contract.EngineMessage(role="system", content="system only")] ) with ( patch.object(gateway, "_resolve_session") as resolve_session, self.assertRaises(gateway.HTTPException) as raised, ): asyncio.run(gateway.v1_generate(req)) self.assertEqual(raised.exception.status_code, 400) self.assertEqual(raised.exception.detail, "no user message in payload") resolve_session.assert_not_called() def test_v1_stream_frames_token_and_done_events(self): session = _FakeStreamSession( [ {"type": "delta", "text": "안녕"}, {"type": "done", "cost_usd": 0.03, "turns": 2}, ], model="stream-model", ) async def fake_resolve(req, system_prompt): return session, True with patch.object(gateway, "_resolve_session", fake_resolve): response = asyncio.run(gateway.v1_stream(_request())) body = asyncio.run(_read_streaming_response(response)) self.assertIn("event: token", body) self.assertIn('data: {"text": "안녕"}', body) self.assertIn("event: done", body) self.assertIn('"provider": "claude_cli"', body) self.assertIn('"model": "stream-model"', body) self.assertIn('"cost_usd": 0.03', body) self.assertEqual(session.content, "hello") self.assertEqual(session.timeout, 600.0) self.assertTrue(session.closed) def test_v1_stream_rejects_missing_user_before_session_resolution(self): req = contract.GenerateRequest( messages=[contract.EngineMessage(role="system", content="system only")] ) with ( patch.object(gateway, "_resolve_session") as resolve_session, self.assertRaises(gateway.HTTPException) as raised, ): asyncio.run(gateway.v1_stream(req)) self.assertEqual(raised.exception.status_code, 400) self.assertEqual(raised.exception.detail, "no user message in payload") resolve_session.assert_not_called() def test_v1_stream_frames_engine_error_event(self): session = _FakeStreamSession( [ {"type": "done", "is_error": True, "error": "engine failed"}, ] ) async def fake_resolve(req, system_prompt): return session, False with patch.object(gateway, "_resolve_session", fake_resolve): response = asyncio.run(gateway.v1_stream(_request())) body = asyncio.run(_read_streaming_response(response)) self.assertIn("event: error", body) self.assertIn('data: {"detail": "engine failed"}', body) self.assertFalse(session.closed) if __name__ == "__main__": unittest.main()