260 lines
9.4 KiB
Python
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)
|