import asyncio import unittest from unittest.mock import patch from app import engine_client from app.contracts import engine_gateway as contract from engine_gateway import gateway 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 _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 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_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_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"))) self.assertEqual(response["text"], "reused response") self.assertEqual(response["provider"], "claude_cli") self.assertEqual(calls, [("hello", 120.0)]) 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"))) self.assertEqual(response["text"], "fresh response") self.assertEqual(len(started), 1) self.assertEqual(turned, [(started[0], "hello", 120.0)]) self.assertEqual(closed, [started[0]]) self.assertNotIn(started[0].id, gateway.SESSIONS) 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_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()