한신대 피드백 개선팩 반영
This commit is contained in:
parent
5a9c110c11
commit
6b6241f468
25 changed files with 1247 additions and 94 deletions
|
|
@ -338,22 +338,39 @@ async def close_session(sid: str):
|
|||
# session_id 가 오면 풀을 재사용해 멀티턴 prompt caching 이점을 살린다.
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _split_messages(messages: list[GwMessage]) -> GatewayPromptParts:
|
||||
"""EngineMessage[] → named prompt parts for the current gateway turn.
|
||||
|
||||
- system 들은 합쳐서 --system-prompt 로 주입할 텍스트로.
|
||||
- 마지막 user 발화를 이번 턴 stdin content 로.
|
||||
- 상주 세션 재사용 시에는 풀이 이미 컨텍스트를 들고 있으므로 마지막 user 만 보냄.
|
||||
"""
|
||||
def _split_messages(messages: list[GwMessage], *, ai_role: AIRole | None = None) -> GatewayPromptParts:
|
||||
"""EngineMessage[] → named prompt parts for the current gateway turn."""
|
||||
system_parts: list[str] = []
|
||||
last_user = ""
|
||||
non_system: list[GwMessage] = []
|
||||
for m in messages:
|
||||
if m.role == "system":
|
||||
system_parts.append(m.content)
|
||||
elif m.role == "user":
|
||||
last_user = m.content
|
||||
else:
|
||||
non_system.append(m)
|
||||
|
||||
last_user_index: int | None = None
|
||||
for index, m in enumerate(non_system):
|
||||
if m.role == "user":
|
||||
last_user_index = index
|
||||
|
||||
last_user = ""
|
||||
if last_user_index is not None:
|
||||
last_user = non_system[last_user_index].content
|
||||
|
||||
user_payload = last_user
|
||||
if ai_role == "client" and last_user_index is not None:
|
||||
history_parts: list[str] = []
|
||||
for m in non_system[:last_user_index]:
|
||||
content = m.content.strip()
|
||||
if not content:
|
||||
continue
|
||||
speaker = "상담자" if m.role == "user" else "내담자"
|
||||
history_parts.append(f"{speaker}: {content}")
|
||||
if history_parts:
|
||||
user_payload = "[직전 대화]\n" + "\n".join(history_parts) + "\n\n[이번 상담자 발화]\n" + last_user
|
||||
|
||||
system_prompt = "\n\n".join(p for p in system_parts if p.strip())
|
||||
return GatewayPromptParts(system_prompt=system_prompt, user_payload=last_user)
|
||||
return GatewayPromptParts(system_prompt=system_prompt, user_payload=user_payload)
|
||||
|
||||
|
||||
def _inject_schema(system_prompt: str, schema: Optional[dict[str, Any]]) -> str:
|
||||
|
|
@ -400,7 +417,7 @@ 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 반환."""
|
||||
prompt_parts = _split_messages(req.messages)
|
||||
prompt_parts = _split_messages(req.messages, ai_role=req.ai_role)
|
||||
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")
|
||||
|
|
@ -436,7 +453,7 @@ async def v1_generate(req: GwGenerateReq):
|
|||
@app.post("/v1/stream")
|
||||
async def v1_stream(req: GwGenerateReq):
|
||||
"""SSE 토큰 스트림. data: 라인으로 텍스트 델타를 흘리고 done/error 프레이밍."""
|
||||
prompt_parts = _split_messages(req.messages)
|
||||
prompt_parts = _split_messages(req.messages, ai_role=req.ai_role)
|
||||
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")
|
||||
|
|
|
|||
|
|
@ -210,6 +210,39 @@ class GatewayModelTest(unittest.TestCase):
|
|||
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")]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue