695 lines
26 KiB
Python
695 lines
26 KiB
Python
import asyncio
|
|
import json
|
|
import shutil
|
|
import subprocess
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest.mock import patch
|
|
|
|
from jsonschema import Draft202012Validator
|
|
|
|
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
|
|
|
|
|
|
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.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 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_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")
|
|
|
|
def test_split_messages_injects_client_history_before_current_counselor_turn(self):
|
|
parts = gateway._split_messages(
|
|
[
|
|
contract.EngineMessage(role="system", content="client persona system"),
|
|
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("상담자: 상담자 이전 질문", parts.user_payload)
|
|
self.assertIn("내담자: 내담자 이전 답변", parts.user_payload)
|
|
self.assertIn("[이번 상담자 발화]", parts.user_payload)
|
|
self.assertTrue(parts.user_payload.rstrip().endswith("이번 상담자 발화"))
|
|
|
|
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, "")
|
|
|
|
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_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_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")))
|
|
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(calls, [("hello", 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, "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")))
|
|
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(len(started), 1)
|
|
self.assertEqual(turned, [(started[0], "hello", 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, "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_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()
|