import asyncio import unittest from unittest.mock import patch 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] class GatewayModelTest(unittest.TestCase): def setUp(self): gateway.SESSIONS.clear() def tearDown(self): gateway.SESSIONS.clear() 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) if __name__ == "__main__": unittest.main()