vignette/apps/api/app/test_evaluator_model_routing.py

260 lines
9.4 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 EngineError, 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 FailingEvaluatorEngine:
def __init__(self) -> None:
self.requests: list[Any] = []
async def generate(self, req: Any) -> GenerateResponse:
self.requests.append(req)
raise EngineError("synthetic evaluator failure")
class EvaluatorModelRoutingTest(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self) -> None:
self._fast_model = settings.evaluator_fast_model
self._deep_model = settings.evaluator_deep_model
self._cache_enabled = settings.evaluator_semantic_cache_enabled
self._cache_ttl = settings.evaluator_semantic_cache_ttl_seconds
self._cache_max_entries = settings.evaluator_semantic_cache_max_entries
settings.evaluator_fast_model = ""
settings.evaluator_deep_model = ""
settings.evaluator_semantic_cache_enabled = True
settings.evaluator_semantic_cache_ttl_seconds = 900
settings.evaluator_semantic_cache_max_entries = 256
evaluator.clear_evaluator_semantic_cache()
async def asyncTearDown(self) -> None:
settings.evaluator_fast_model = self._fast_model
settings.evaluator_deep_model = self._deep_model
settings.evaluator_semantic_cache_enabled = self._cache_enabled
settings.evaluator_semantic_cache_ttl_seconds = self._cache_ttl
settings.evaluator_semantic_cache_max_entries = self._cache_max_entries
evaluator.clear_evaluator_semantic_cache()
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)
async def test_fast_evaluator_reuses_semantic_cache_for_identical_prompt(self) -> None:
engine = CaptureEvaluatorEngine()
ctx = _turn_context()
first = await evaluator.evaluate_turn(
ctx,
"괜찮아요.",
engine=engine, # type: ignore[arg-type]
)
second = await evaluator.evaluate_turn(
ctx,
"괜찮아요.",
engine=engine, # type: ignore[arg-type]
)
self.assertIsNone(first.error)
self.assertIsNone(second.error)
self.assertEqual(len(engine.requests), 1)
stats = evaluator.evaluator_semantic_cache_stats()
self.assertEqual(stats["misses"], 1)
self.assertEqual(stats["hits"], 1)
self.assertEqual(stats["stores"], 1)
self.assertEqual(stats["entries"], 1)
async def test_deep_evaluator_reuses_semantic_cache_for_identical_prompt(self) -> None:
engine = CaptureEvaluatorEngine()
masked_turns = [
{"speaker": "counselor", "text": "천천히 이야기해줘도 괜찮아요."},
{"speaker": "client", "text": "잘 모르겠어요."},
]
first = await evaluator.evaluate_session(
session_id="evaluator-model-session",
stage="라포",
masked_turns=masked_turns,
engine=engine, # type: ignore[arg-type]
)
second = await evaluator.evaluate_session(
session_id="evaluator-model-session",
stage="라포",
masked_turns=masked_turns,
engine=engine, # type: ignore[arg-type]
)
self.assertIsNone(first.error)
self.assertIsNone(second.error)
self.assertEqual(len(engine.requests), 1)
stats = evaluator.evaluator_semantic_cache_stats()
self.assertEqual(stats["misses"], 1)
self.assertEqual(stats["hits"], 1)
self.assertEqual(stats["stores"], 1)
async def test_semantic_cache_key_separates_model_override(self) -> None:
engine = CaptureEvaluatorEngine()
settings.evaluator_fast_model = "cheap-fast-a"
await evaluator.evaluate_turn(
_turn_context(),
"괜찮아요.",
engine=engine, # type: ignore[arg-type]
)
settings.evaluator_fast_model = "cheap-fast-b"
await evaluator.evaluate_turn(
_turn_context(),
"괜찮아요.",
engine=engine, # type: ignore[arg-type]
)
self.assertEqual(len(engine.requests), 2)
self.assertEqual(engine.requests[0].model, "cheap-fast-a")
self.assertEqual(engine.requests[1].model, "cheap-fast-b")
stats = evaluator.evaluator_semantic_cache_stats()
self.assertEqual(stats["misses"], 2)
self.assertEqual(stats["hits"], 0)
async def test_semantic_cache_hit_does_not_record_second_audit_event(self) -> None:
engine = CaptureEvaluatorEngine()
audit_payloads: list[dict[str, Any]] = []
async def audit_hook(payload: dict[str, Any]) -> None:
audit_payloads.append(payload)
await evaluator.evaluate_turn(
_turn_context(),
"괜찮아요.",
engine=engine, # type: ignore[arg-type]
audit_hook=audit_hook,
)
await evaluator.evaluate_turn(
_turn_context(),
"괜찮아요.",
engine=engine, # type: ignore[arg-type]
audit_hook=audit_hook,
)
self.assertEqual(len(engine.requests), 1)
self.assertEqual(len(audit_payloads), 1)
self.assertEqual(audit_payloads[0]["provider"], "fake-provider")
stats = evaluator.evaluator_semantic_cache_stats()
self.assertEqual(stats["hits"], 1)
async def test_engine_error_is_not_cached(self) -> None:
engine = FailingEvaluatorEngine()
first = await evaluator.evaluate_turn(
_turn_context(),
"괜찮아요.",
engine=engine, # type: ignore[arg-type]
)
second = await evaluator.evaluate_turn(
_turn_context(),
"괜찮아요.",
engine=engine, # type: ignore[arg-type]
)
self.assertIn("engine_error", first.error or "")
self.assertIn("engine_error", second.error or "")
self.assertEqual(len(engine.requests), 2)
stats = evaluator.evaluator_semantic_cache_stats()
self.assertEqual(stats["hits"], 0)
self.assertEqual(stats["stores"], 0)
self.assertEqual(stats["misses"], 2)