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("--append-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()) if __name__ == "__main__": unittest.main()