vignette/apps/api/engine_gateway/test_gateway_model.py

1214 lines
45 KiB
Python

import asyncio
import json
import secrets
import shutil
import subprocess
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from jsonschema import Draft202012Validator
from fastapi.testclient import TestClient
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
class _StreamStdin(_FakeStdin):
def __init__(self):
self.writes = []
def write(self, value):
self.writes.append(value)
async def drain(self):
return None
class _StreamStdout:
def __init__(self, objects):
self.lines = [
(json.dumps(obj, ensure_ascii=False) + "\n").encode("utf-8")
for obj in objects
]
async def readline(self):
return self.lines.pop(0) if self.lines else b""
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.turns = 0
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 GatewayAuthenticationTest(unittest.TestCase):
SECRET = "engine-gateway-test-secret-" + ("x" * 32)
def test_unset_secret_preserves_local_gateway_compatibility(self):
with patch.object(gateway, "ENGINE_GATEWAY_SHARED_SECRET", ""):
response = TestClient(gateway.app).delete("/session/not-running")
self.assertEqual(response.status_code, 200)
def test_configured_secret_rejects_missing_and_wrong_credentials(self):
with patch.object(gateway, "ENGINE_GATEWAY_SHARED_SECRET", self.SECRET):
client = TestClient(gateway.app)
missing = client.post("/v1/generate", json={})
wrong = client.post(
"/v1/generate",
json={},
headers={gateway.ENGINE_TOKEN_HEADER: "wrong"},
)
authenticated = client.post(
"/v1/generate",
json={},
headers={gateway.ENGINE_TOKEN_HEADER: self.SECRET},
)
self.assertEqual(missing.status_code, 401)
self.assertEqual(wrong.status_code, 401)
self.assertEqual(authenticated.status_code, 422)
def test_health_probe_remains_unauthenticated(self):
with patch.object(gateway, "ENGINE_GATEWAY_SHARED_SECRET", self.SECRET):
response = TestClient(gateway.app).get("/health")
self.assertEqual(response.status_code, 200)
def test_ready_probe_requires_credentials_because_it_runs_generation(self):
with patch.object(gateway, "ENGINE_GATEWAY_SHARED_SECRET", self.SECRET):
response = TestClient(gateway.app).get("/ready")
self.assertEqual(response.status_code, 401)
def test_gateway_rejects_weak_configured_secret_at_startup(self):
for value in ("too-short", "example-gateway-secret-with-32-characters"):
with (
self.subTest(value=value),
patch.dict(
gateway.os.environ,
{"ENGINE_GATEWAY_SHARED_SECRET": value},
),
self.assertRaises(RuntimeError),
):
gateway._load_gateway_shared_secret()
def test_invalid_token_uses_constant_time_comparison(self):
with (
patch.object(gateway, "ENGINE_GATEWAY_SHARED_SECRET", self.SECRET),
patch.object(
gateway.secrets,
"compare_digest",
wraps=secrets.compare_digest,
) as compare_digest,
):
response = TestClient(gateway.app).get(
"/v1/capabilities",
headers={gateway.ENGINE_TOKEN_HEADER: "wrong"},
)
self.assertEqual(response.status_code, 401)
compare_digest.assert_called_once_with("wrong", self.SECRET)
def test_openapi_schema_does_not_expose_secret_or_auth_header(self):
with patch.object(gateway, "ENGINE_GATEWAY_SHARED_SECRET", self.SECRET):
schema = json.dumps(gateway.app.openapi())
self.assertNotIn(self.SECRET, schema)
self.assertNotIn(gateway.ENGINE_TOKEN_HEADER, schema)
def test_engine_client_adds_gateway_token_to_default_headers(self):
with patch.object(engine_client.httpx, "AsyncClient") as async_client_cls:
client = engine_client.EngineClient(
"http://127.0.0.1:9099",
shared_secret=self.SECRET,
)
client._new_client()
self.assertEqual(
async_client_cls.call_args.kwargs["headers"],
{gateway.ENGINE_TOKEN_HEADER: self.SECRET},
)
def test_engine_client_omits_gateway_token_when_unset(self):
with patch.object(engine_client.httpx, "AsyncClient") as async_client_cls:
client = engine_client.EngineClient(
"http://127.0.0.1:9099",
shared_secret="",
)
client._new_client()
self.assertEqual(async_client_cls.call_args.kwargs["headers"], {})
class EngineClientReadinessTest(unittest.IsolatedAsyncioTestCase):
@staticmethod
def _response(status_code, payload):
return engine_client.httpx.Response(
status_code,
json=payload,
request=engine_client.httpx.Request("GET", "http://127.0.0.1:9099/ready"),
)
async def test_health_fails_when_dedicated_live_client_provider_is_unready(self):
client = engine_client.EngineClient("http://127.0.0.1:9099")
client.engine_mode = "agy_cli"
client.default_model = "gemini-3.6-flash-high"
client.default_reasoning_effort = "high"
client.live_client_provider = "claude_cli"
client._client = AsyncMock()
client._client.get = AsyncMock(
side_effect=[
self._response(200, {"ok": True, "detail": "agy ready"}),
self._response(503, {"ok": False, "detail": "claude login required"}),
]
)
detail = await client.health_detail()
self.assertFalse(detail["ok"])
self.assertEqual(detail["default_engine"]["provider"], "agy_cli")
self.assertTrue(detail["default_engine"]["ok"])
self.assertEqual(detail["live_client_engine"]["provider"], "claude_cli")
self.assertFalse(detail["live_client_engine"]["ok"])
self.assertIn("실시간 내담자 공급자 claude_cli", detail["detail"])
calls = client._client.get.await_args_list
self.assertEqual(calls[0].kwargs["params"]["provider"], "agy_cli")
self.assertEqual(calls[0].kwargs["params"]["model"], "gemini-3.6-flash-high")
self.assertEqual(calls[1].kwargs["params"], {"provider": "claude_cli"})
async def test_health_reuses_one_probe_when_live_and_default_provider_match(self):
client = engine_client.EngineClient("http://127.0.0.1:9099")
client.engine_mode = "claude_cli"
client.live_client_provider = "claude_cli"
client._client = AsyncMock()
client._client.get = AsyncMock(
return_value=self._response(200, {"ok": True, "detail": "ready"})
)
detail = await client.health_detail()
self.assertTrue(detail["ok"])
client._client.get.assert_awaited_once()
async def test_gateway_liveness_does_not_replace_provider_readiness(self):
client = engine_client.EngineClient("http://127.0.0.1:9099")
client.engine_mode = "claude_cli"
client.live_client_provider = "claude_cli"
client._client = AsyncMock()
client._client.get = AsyncMock(
side_effect=[
self._response(404, {"detail": "missing"}),
self._response(200, {"status": "ok"}),
]
)
detail = await client.health_detail()
self.assertFalse(detail["ok"])
self.assertIn("공급자 준비상태 엔드포인트", detail["detail"])
self.assertEqual(client._client.get.await_count, 2)
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_engine_client_payload_includes_provider_model_and_reasoning_defaults(self):
client = engine_client.EngineClient("http://127.0.0.1:9099")
client.engine_mode = "codex_cli"
client.live_client_provider = None
client.default_model = "gpt-5.6-terra"
client.default_reasoning_effort = "medium"
payload = client._payload(
contract.GenerateRequest(
messages=[contract.EngineMessage(role="user", content="hello")]
)
)
self.assertEqual(payload["provider"], "codex_cli")
self.assertEqual(payload["model"], "gpt-5.6-terra")
self.assertEqual(payload["reasoning_effort"], "medium")
def test_engine_client_uses_dedicated_live_provider_without_foreign_model_defaults(self):
client = engine_client.EngineClient("http://127.0.0.1:9099")
client.engine_mode = "agy_cli"
client.default_model = "gemini-3.6-flash-high"
client.default_reasoning_effort = "high"
client.live_client_provider = "claude_cli"
payload = client._payload(
contract.GenerateRequest(
ai_role="client",
session_id="session-id",
messages=[contract.EngineMessage(role="user", content="hello")],
)
)
self.assertEqual(payload["provider"], "claude_cli")
self.assertNotIn("model", payload)
self.assertNotIn("reasoning_effort", payload)
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")
self.assertEqual(parts.current_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", cache=True),
contract.EngineMessage(role="system", content="dynamic state", cache=False),
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("dynamic state", parts.user_payload)
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("이번 상담자 발화"))
self.assertNotIn("[직전 대화]", parts.current_user_payload)
self.assertIn("dynamic state", parts.current_user_payload)
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, "")
self.assertEqual(parts.current_user_payload, "")
def test_claude_process_enables_real_partial_streaming_without_disk_session_copy(self):
self.assertIn("--include-partial-messages", gateway.BASE_ARGS)
self.assertIn("--no-session-persistence", gateway.BASE_ARGS)
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_repairs_only_trailing_commas(self):
response = contract.GenerateResponse(
text=(
'```json\n'
'{"reply":"keep literal , } and escaped \\\" text",'
'"items":[{"value":1,},],}\n'
'```'
),
provider="test",
model="test-model",
)
self.assertEqual(
contract.structured_payload_from_response(response),
{
"reply": 'keep literal , } and escaped " text',
"items": [{"value": 1}],
},
)
def test_generate_response_structured_payload_does_not_repair_other_corruption(self):
response = contract.GenerateResponse(
text='{"reply":"missing separator" "items":[]}',
provider="test",
model="test-model",
)
self.assertIsNone(contract.structured_payload_from_response(response))
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_engine_session_passes_reasoning_effort_to_claude_cli(self):
captured, process_patch = _capture_subprocess()
with (
patch.object(gateway, "DEFAULT_MODEL", ""),
patch.object(gateway, "FALLBACK_MODEL", ""),
process_patch,
):
session = gateway.EngineSession(
model="opus",
reasoning_effort="high",
)
asyncio.run(session.start())
try:
self.assertIn("--effort", captured[0])
self.assertEqual(
captured[0][captured[0].index("--effort") + 1], "high"
)
finally:
asyncio.run(session.close())
def test_engine_session_emits_partial_stream_events_without_final_message_duplication(self):
process = _FakeProcess()
process.stdin = _StreamStdin()
process.stdout = _StreamStdout(
[
{
"type": "stream_event",
"event": {
"type": "content_block_delta",
"delta": {"type": "text_delta", "text": ""},
},
},
{
"type": "stream_event",
"event": {
"type": "content_block_delta",
"delta": {"type": "text_delta", "text": "녕!"},
},
},
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "안녕!"}]},
},
{
"type": "result",
"is_error": False,
"total_cost_usd": 0.01,
"modelUsage": {
"claude-opus-4-8": {
"inputTokens": 12,
"outputTokens": 7,
"cacheReadInputTokens": 101,
"cacheCreationInputTokens": 23,
},
"claude-haiku-4-5": {
"inputTokens": 3,
"outputTokens": 2,
"cacheReadInputTokens": 9,
"cacheCreationInputTokens": 0,
},
},
},
]
)
session = gateway.EngineSession()
session.proc = process
async def collect():
return [event async for event in session.turn_stream("질문")]
events = asyncio.run(collect())
self.assertEqual(
events,
[
{"type": "delta", "text": ""},
{"type": "delta", "text": "녕!"},
{
"type": "done",
"text": "안녕!",
"cost_usd": 0.01,
"tokens_in": 148,
"tokens_out": 9,
"turns": 1,
"is_error": False,
"error": "안녕!",
},
],
)
def test_engine_session_uses_terminal_result_when_assistant_event_is_absent(self):
process = _FakeProcess()
process.stdin = _StreamStdin()
process.stdout = _StreamStdout(
[
{
"type": "result",
"is_error": False,
"result": '{"goal":{"score":0.2}}',
"usage": {"input_tokens": 10, "output_tokens": 5},
}
]
)
session = gateway.EngineSession()
session.proc = process
result = asyncio.run(session.turn("평가"))
self.assertEqual(result["text"], '{"goal":{"score":0.2}}')
self.assertFalse(result["is_error"])
def test_engine_session_streams_terminal_result_when_assistant_event_is_absent(self):
process = _FakeProcess()
process.stdin = _StreamStdin()
process.stdout = _StreamStdout(
[
{
"type": "result",
"is_error": False,
"result": "terminal-only",
"usage": {"input_tokens": 10, "output_tokens": 2},
}
]
)
session = gateway.EngineSession()
session.proc = process
async def collect():
return [event async for event in session.turn_stream("질문")]
events = asyncio.run(collect())
self.assertEqual(events[0], {"type": "delta", "text": "terminal-only"})
self.assertEqual(events[-1]["text"], "terminal-only")
def test_claude_result_tokens_falls_back_to_top_level_usage(self):
self.assertEqual(
gateway._claude_result_tokens(
{
"usage": {
"input_tokens": 17,
"output_tokens": 5,
"cache_read_input_tokens": 200,
"cache_creation_input_tokens": 30,
}
}
),
(247, 5),
)
def test_v1_generate_routes_non_claude_provider_through_registry(self):
result = SimpleNamespace(
text="registry response",
model="gpt-5.6-terra",
provider="codex_cli",
tokens_in=12,
tokens_out=3,
cost_usd=0.0,
inference_geo=None,
structured=None,
)
request = contract.GenerateRequest(
provider="codex_cli",
model="gpt-5.6-terra",
reasoning_effort="medium",
messages=[contract.EngineMessage(role="user", content="hello")],
)
with patch.object(
gateway,
"generate_with_provider",
AsyncMock(return_value=result),
) as generate:
response = asyncio.run(gateway.v1_generate(request))
self.assertEqual(response["provider"], "codex_cli")
self.assertEqual(response["model"], "gpt-5.6-terra")
generate.assert_awaited_once()
def test_v1_stream_forwards_non_claude_provider_deltas(self):
result = SimpleNamespace(
text="안녕",
model="gemini-3.6-flash-high",
provider="agy_cli",
tokens_in=12,
tokens_out=2,
cost_usd=0.0,
)
request = contract.GenerateRequest(
provider="agy_cli",
model="gemini-3.6-flash-high",
reasoning_effort="high",
messages=[contract.EngineMessage(role="user", content="hello")],
)
async def fake_stream(*args, **kwargs):
yield SimpleNamespace(type="delta", text="", result=None)
yield SimpleNamespace(type="delta", text="", result=None)
yield SimpleNamespace(type="done", text="", result=result)
with patch.object(gateway, "stream_with_provider", fake_stream):
response = asyncio.run(gateway.v1_stream(request))
body = asyncio.run(_read_streaming_response(response))
self.assertEqual(body.count("event: token"), 2)
self.assertIn('data: {"text": ""}', body)
self.assertIn('data: {"text": ""}', body)
self.assertIn("event: done", body)
self.assertIn('"provider": "agy_cli"', body)
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, False)
self.assertIsNot(session, existing)
self.assertEqual(_model_arg(captured[0]), "new-model")
self.assertIs(gateway.SESSIONS["sid"], session)
self.assertEqual(existing.proc.returncode, 0)
finally:
gateway.SESSIONS.pop("sid", None)
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_binds_missing_client_session_id_to_resident_pool(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, False)
self.assertIs(gateway.SESSIONS["missing"], session)
self.assertEqual(len(captured), 1)
self.assertIn("--system-prompt", captured[0])
finally:
gateway.SESSIONS.pop("missing", None)
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,
"tokens_in": 321,
"tokens_out": 12,
"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(validated.tokens_in, 321)
self.assertEqual(validated.tokens_out, 12)
self.assertEqual(
calls,
[("[이번 상담자 발화]\nhello", 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,
"tokens_in": 654,
"tokens_out": 21,
"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()))
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(validated.tokens_in, 654)
self.assertEqual(validated.tokens_out, 21)
self.assertEqual(len(started), 1)
self.assertEqual(
turned,
[
(
started[0],
"[이번 상담자 발화]\nhello",
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,
"tokens_in": 456,
"tokens_out": 18,
"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.assertIn('"tokens_in": 456', body)
self.assertIn('"tokens_out": 18', body)
self.assertEqual(session.content, "[이번 상담자 발화]\nhello")
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()