vignette/apps/api/app/test_evaluator_model_routing.py

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)