런타임 계약과 학습자 흐름 보강
This commit is contained in:
parent
f456b8997a
commit
206018b088
56 changed files with 4306 additions and 1008 deletions
|
|
@ -13,6 +13,7 @@ import json
|
|||
import os
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
|
||||
from fastapi import FastAPI, HTTPException
|
||||
|
|
@ -30,6 +31,7 @@ from app.contracts.engine_gateway import (
|
|||
StreamDoneEvent,
|
||||
StreamErrorEvent,
|
||||
StreamTokenEvent,
|
||||
normalize_engine_gateway_model,
|
||||
sse_frame,
|
||||
)
|
||||
|
||||
|
|
@ -40,6 +42,15 @@ DEFAULT_BUDGET = float(os.environ.get("SESSION_BUDGET_USD", "5.0"))
|
|||
READY_TTL_SECONDS = float(os.environ.get("ENGINE_READY_TTL_SECONDS", "30"))
|
||||
READY_TIMEOUT_SECONDS = float(os.environ.get("ENGINE_READY_TIMEOUT_SECONDS", "20"))
|
||||
READY_BUDGET_USD = float(os.environ.get("ENGINE_READY_BUDGET_USD", "0.5"))
|
||||
GATEWAY_PROVIDER = "claude_cli"
|
||||
GATEWAY_FALLBACK_MODEL_NAME = "claude-opus-4-8"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GatewayPromptParts:
|
||||
system_prompt: str
|
||||
user_payload: str
|
||||
|
||||
|
||||
BASE_ARGS = [
|
||||
"-p",
|
||||
|
|
@ -53,14 +64,6 @@ BASE_ARGS = [
|
|||
"--exclude-dynamic-system-prompt-sections",
|
||||
]
|
||||
|
||||
|
||||
def _model_override(model: Optional[str]) -> Optional[str]:
|
||||
value = (model or "").strip()
|
||||
if not value or value == "gateway-default":
|
||||
return None
|
||||
return value
|
||||
|
||||
|
||||
class EngineSession:
|
||||
"""claude -p 상주 프로세스 1개 = 상담 회기 1개."""
|
||||
|
||||
|
|
@ -73,7 +76,7 @@ class EngineSession:
|
|||
self.id = uuid.uuid4().hex
|
||||
self.system_prompt = system_prompt
|
||||
self.budget = budget
|
||||
self.model = _model_override(model)
|
||||
self.model = normalize_engine_gateway_model(model)
|
||||
self.proc: asyncio.subprocess.Process | None = None
|
||||
self.lock = asyncio.Lock() # 한 회기 안의 턴은 직렬(상담 왕복)
|
||||
self.cost_usd = 0.0
|
||||
|
|
@ -335,28 +338,22 @@ async def close_session(sid: str):
|
|||
# session_id 가 오면 풀을 재사용해 멀티턴 prompt caching 이점을 살린다.
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _split_messages(messages: list[GwMessage]) -> tuple[str, str]:
|
||||
"""EngineMessage[] → (system_prompt, user_payload).
|
||||
def _split_messages(messages: list[GwMessage]) -> GatewayPromptParts:
|
||||
"""EngineMessage[] → named prompt parts for the current gateway turn.
|
||||
|
||||
- system 들은 합쳐서 --append-system-prompt 로 주입할 텍스트로.
|
||||
- system 들은 합쳐서 --system-prompt 로 주입할 텍스트로.
|
||||
- 마지막 user 발화를 이번 턴 stdin content 로.
|
||||
- 직전 assistant/user 히스토리는 (단발 모드라) system 뒤에 맥락으로 직렬화.
|
||||
(상주 세션 재사용 시에는 풀이 이미 컨텍스트를 들고 있으므로 마지막 user 만 보냄.)
|
||||
- 상주 세션 재사용 시에는 풀이 이미 컨텍스트를 들고 있으므로 마지막 user 만 보냄.
|
||||
"""
|
||||
system_parts: list[str] = []
|
||||
history_parts: list[str] = []
|
||||
last_user = ""
|
||||
for m in messages:
|
||||
if m.role == "system":
|
||||
system_parts.append(m.content)
|
||||
elif m.role == "assistant":
|
||||
history_parts.append(f"[이전 상담자 발화]\n{m.content}")
|
||||
elif m.role == "user":
|
||||
if last_user:
|
||||
history_parts.append(f"[이전 내담자 발화]\n{last_user}")
|
||||
last_user = m.content
|
||||
system_prompt = "\n\n".join(p for p in system_parts if p.strip())
|
||||
return system_prompt, last_user
|
||||
return GatewayPromptParts(system_prompt=system_prompt, user_payload=last_user)
|
||||
|
||||
|
||||
def _inject_schema(system_prompt: str, schema: Optional[dict[str, Any]]) -> str:
|
||||
|
|
@ -375,12 +372,16 @@ def _inject_schema(system_prompt: str, schema: Optional[dict[str, Any]]) -> str:
|
|||
return (system_prompt + directive) if system_prompt else directive.lstrip()
|
||||
|
||||
|
||||
def _response_model_name(session: EngineSession) -> str:
|
||||
return session.model or DEFAULT_MODEL or GATEWAY_FALLBACK_MODEL_NAME
|
||||
|
||||
|
||||
async def _resolve_session(req: GwGenerateReq, system_prompt: str) -> tuple[EngineSession, bool]:
|
||||
"""session_id 가 있고 살아있으면 재사용, 아니면 단발용 임시 세션 생성.
|
||||
|
||||
반환: (session, ephemeral). ephemeral=True 면 호출부가 응답 후 close 한다.
|
||||
"""
|
||||
requested_model = _model_override(req.model)
|
||||
requested_model = normalize_engine_gateway_model(req.model)
|
||||
if req.session_id and req.session_id in SESSIONS:
|
||||
s = SESSIONS[req.session_id]
|
||||
if s.proc is not None and s.proc.returncode is None:
|
||||
|
|
@ -399,14 +400,14 @@ async def _resolve_session(req: GwGenerateReq, system_prompt: str) -> tuple[Engi
|
|||
@app.post("/v1/generate")
|
||||
async def v1_generate(req: GwGenerateReq):
|
||||
"""단발 생성 (평가 deep-loop, 회기종료 압축 등). GenerateResponse 호환 dict 반환."""
|
||||
system_prompt, user_payload = _split_messages(req.messages)
|
||||
system_prompt = _inject_schema(system_prompt, req.structured_schema)
|
||||
if not user_payload:
|
||||
prompt_parts = _split_messages(req.messages)
|
||||
system_prompt = _inject_schema(prompt_parts.system_prompt, req.structured_schema)
|
||||
if not prompt_parts.user_payload:
|
||||
raise HTTPException(400, "no user message in payload")
|
||||
|
||||
s, ephemeral = await _resolve_session(req, system_prompt)
|
||||
try:
|
||||
result = await s.turn(user_payload, timeout=120.0)
|
||||
result = await s.turn(prompt_parts.user_payload, timeout=120.0)
|
||||
finally:
|
||||
if ephemeral:
|
||||
await s.close()
|
||||
|
|
@ -422,8 +423,8 @@ async def v1_generate(req: GwGenerateReq):
|
|||
structured = None # 파싱 실패는 호출부가 text 로 폴백
|
||||
return GenerateResponse(
|
||||
text=text,
|
||||
model=s.model or DEFAULT_MODEL or "claude-opus-4-8",
|
||||
provider="claude_cli",
|
||||
model=_response_model_name(s),
|
||||
provider=GATEWAY_PROVIDER,
|
||||
tokens_in=0,
|
||||
tokens_out=0,
|
||||
cost_usd=result.get("cost_usd", 0.0),
|
||||
|
|
@ -435,16 +436,16 @@ async def v1_generate(req: GwGenerateReq):
|
|||
@app.post("/v1/stream")
|
||||
async def v1_stream(req: GwGenerateReq):
|
||||
"""SSE 토큰 스트림. data: 라인으로 텍스트 델타를 흘리고 done/error 프레이밍."""
|
||||
system_prompt, user_payload = _split_messages(req.messages)
|
||||
system_prompt = _inject_schema(system_prompt, req.structured_schema)
|
||||
if not user_payload:
|
||||
prompt_parts = _split_messages(req.messages)
|
||||
system_prompt = _inject_schema(prompt_parts.system_prompt, req.structured_schema)
|
||||
if not prompt_parts.user_payload:
|
||||
raise HTTPException(400, "no user message in payload")
|
||||
|
||||
s, ephemeral = await _resolve_session(req, system_prompt)
|
||||
|
||||
async def _sse():
|
||||
try:
|
||||
async for evt in s.turn_stream(user_payload, timeout=600.0):
|
||||
async for evt in s.turn_stream(prompt_parts.user_payload, timeout=600.0):
|
||||
if evt.get("type") == "delta":
|
||||
yield sse_frame(
|
||||
ENGINE_GATEWAY_SSE_TOKEN,
|
||||
|
|
@ -460,8 +461,8 @@ async def v1_stream(req: GwGenerateReq):
|
|||
yield sse_frame(
|
||||
ENGINE_GATEWAY_SSE_DONE,
|
||||
StreamDoneEvent(
|
||||
provider="claude_cli",
|
||||
model=s.model or DEFAULT_MODEL or "claude-opus-4-8",
|
||||
provider=GATEWAY_PROVIDER,
|
||||
model=_response_model_name(s),
|
||||
tokens_in=0,
|
||||
tokens_out=0,
|
||||
cost_usd=evt.get("cost_usd", 0.0),
|
||||
|
|
|
|||
|
|
@ -329,5 +329,6 @@
|
|||
"token",
|
||||
"done",
|
||||
"error"
|
||||
]
|
||||
],
|
||||
"x-engine-gateway-default-model-sentinel": "gateway-default"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -135,6 +135,7 @@ def _contract_schema_doc():
|
|||
},
|
||||
"$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,
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -171,6 +172,16 @@ class _FakeStreamSession:
|
|||
|
||||
|
||||
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()
|
||||
|
||||
|
|
@ -183,6 +194,30 @@ class GatewayModelTest(unittest.TestCase):
|
|||
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_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")),
|
||||
|
|
@ -286,6 +321,46 @@ class GatewayModelTest(unittest.TestCase):
|
|||
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()
|
||||
|
|
@ -321,6 +396,10 @@ class GatewayModelTest(unittest.TestCase):
|
|||
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)
|
||||
|
||||
|
|
@ -504,6 +583,21 @@ class GatewayModelTest(unittest.TestCase):
|
|||
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(
|
||||
[
|
||||
|
|
@ -530,6 +624,21 @@ class GatewayModelTest(unittest.TestCase):
|
|||
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(
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue