117 lines
3.9 KiB
Python
117 lines
3.9 KiB
Python
"""Regression tests for evaluator low-cost model routing."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import unittest
|
|
from typing import Any
|
|
|
|
from .config import settings
|
|
from .engine_client import GenerateResponse
|
|
from .services import evaluator, orchestrator, persona, state_machine
|
|
|
|
|
|
def _initial_state() -> state_machine.SessionState:
|
|
return state_machine.init_state(
|
|
params=persona.P1.openness_params(),
|
|
)
|
|
|
|
|
|
def _turn_context() -> orchestrator.TurnContext:
|
|
return orchestrator.prepare_turn(
|
|
session_id="evaluator-model-session",
|
|
case_id="evaluator-model-case",
|
|
card=persona.P1,
|
|
state=_initial_state(),
|
|
learner_text="요즘 학교 가는 게 너무 부담돼요.",
|
|
theory_mode="humanistic",
|
|
)
|
|
|
|
|
|
class CaptureEvaluatorEngine:
|
|
def __init__(self) -> None:
|
|
self.requests: list[Any] = []
|
|
|
|
async def generate(self, req: Any) -> GenerateResponse:
|
|
self.requests.append(req)
|
|
loop = req.metadata.get("loop")
|
|
if loop == "fast":
|
|
structured = {
|
|
"techniques": [],
|
|
"client_state_read": [],
|
|
"appropriateness": "neutral",
|
|
"rapport_signal": 0.0,
|
|
"intent_deviations": [],
|
|
}
|
|
else:
|
|
structured = {
|
|
"strengths": [],
|
|
"improvements": [],
|
|
"alternative_utterances": [],
|
|
"intent_deviations": [],
|
|
}
|
|
return GenerateResponse(
|
|
text="",
|
|
model=req.model or "gateway-default",
|
|
provider="fake-provider",
|
|
structured=structured,
|
|
)
|
|
|
|
|
|
class EvaluatorModelRoutingTest(unittest.IsolatedAsyncioTestCase):
|
|
async def asyncSetUp(self) -> None:
|
|
self._fast_model = settings.evaluator_fast_model
|
|
self._deep_model = settings.evaluator_deep_model
|
|
settings.evaluator_fast_model = ""
|
|
settings.evaluator_deep_model = ""
|
|
|
|
async def asyncTearDown(self) -> None:
|
|
settings.evaluator_fast_model = self._fast_model
|
|
settings.evaluator_deep_model = self._deep_model
|
|
|
|
async def test_fast_evaluator_uses_configured_model_override(self) -> None:
|
|
settings.evaluator_fast_model = "cheap-fast"
|
|
engine = CaptureEvaluatorEngine()
|
|
|
|
result = await evaluator.evaluate_turn(
|
|
_turn_context(),
|
|
"괜찮아요.",
|
|
engine=engine, # type: ignore[arg-type]
|
|
)
|
|
|
|
self.assertIsNone(result.error)
|
|
self.assertEqual(len(engine.requests), 1)
|
|
self.assertEqual(engine.requests[0].ai_role, "evaluator")
|
|
self.assertEqual(engine.requests[0].model, "cheap-fast")
|
|
|
|
async def test_deep_evaluator_uses_configured_model_override(self) -> None:
|
|
settings.evaluator_deep_model = "cheap-deep"
|
|
engine = CaptureEvaluatorEngine()
|
|
|
|
result = await evaluator.evaluate_session(
|
|
session_id="evaluator-model-session",
|
|
stage="라포",
|
|
masked_turns=[
|
|
{"speaker": "counselor", "text": "천천히 이야기해줘도 괜찮아요."},
|
|
{"speaker": "client", "text": "잘 모르겠어요."},
|
|
],
|
|
engine=engine, # type: ignore[arg-type]
|
|
)
|
|
|
|
self.assertIsNone(result.error)
|
|
self.assertEqual(len(engine.requests), 1)
|
|
self.assertEqual(engine.requests[0].ai_role, "evaluator")
|
|
self.assertEqual(engine.requests[0].model, "cheap-deep")
|
|
|
|
async def test_blank_model_settings_keep_gateway_default_routing(self) -> None:
|
|
settings.evaluator_fast_model = " "
|
|
settings.evaluator_deep_model = ""
|
|
engine = CaptureEvaluatorEngine()
|
|
|
|
await evaluator.evaluate_turn(
|
|
_turn_context(),
|
|
"괜찮아요.",
|
|
engine=engine, # type: ignore[arg-type]
|
|
)
|
|
|
|
self.assertEqual(len(engine.requests), 1)
|
|
self.assertIsNone(engine.requests[0].model)
|