vignette/apps/api/engine_gateway/test_gateway_model.py
2026-06-28 12:18:20 +09:00

305 lines
10 KiB
Python

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()