Explorar el Código

refactor: unify benchmark case and mode ids

Problem: Benchmark modes and case IDs were duplicated across Literal declarations, ordered tuples, configuration fields, and catalog keys, allowing the accepted configuration surface to drift from executable benchmark coverage.

Risk: Replacing strings with enums could change JSON compatibility, catalog lookup behavior, runtime request values, or matrix cardinality; regression coverage preserves stable ordering, string lookup and serialization, plain DebugRunRequest mode values, complete catalog coverage, and all 12 mode-by-case runs.
zhenyu.hu hace 2 semanas
padre
commit
f7fd343f74

+ 49 - 45
src/agent_lab/application/benchmark.py

@@ -1,8 +1,9 @@
 import json
 import unicodedata
 from collections.abc import Mapping
+from enum import StrEnum
 from types import MappingProxyType
-from typing import Literal, TypeAlias
+from typing import Annotated, Literal, TypeAlias
 from urllib.parse import urlsplit
 
 import httpx
@@ -17,29 +18,24 @@ from agent_lab.domain.events import ToolCallEvent
 from agent_lab.domain.messages import StreamItem
 
 
-BenchmarkMode: TypeAlias = Literal["dual_agent", "chat_agent_tools"]
-BenchmarkCaseId: TypeAlias = Literal[
-    "ordinary_chat",
-    "device_volume_silent",
-    "calendar_schedule_template",
-    "web_search_two_answers",
-    "session_terminate",
-    "parallel_volume_schedule",
-]
+class BenchmarkMode(StrEnum):
+    DUAL_AGENT = "dual_agent"
+    CHAT_AGENT_TOOLS = "chat_agent_tools"
+
+
+class BenchmarkCaseId(StrEnum):
+    ORDINARY_CHAT = "ordinary_chat"
+    DEVICE_VOLUME_SILENT = "device_volume_silent"
+    CALENDAR_SCHEDULE_TEMPLATE = "calendar_schedule_template"
+    WEB_SEARCH_TWO_ANSWERS = "web_search_two_answers"
+    SESSION_TERMINATE = "session_terminate"
+    PARALLEL_VOLUME_SCHEDULE = "parallel_volume_schedule"
+
+
 JsonScalar: TypeAlias = str | int | float | bool | None
 
-BENCHMARK_MODES: tuple[BenchmarkMode, ...] = (
-    "dual_agent",
-    "chat_agent_tools",
-)
-BENCHMARK_CASE_IDS: tuple[BenchmarkCaseId, ...] = (
-    "ordinary_chat",
-    "device_volume_silent",
-    "calendar_schedule_template",
-    "web_search_two_answers",
-    "session_terminate",
-    "parallel_volume_schedule",
-)
+BENCHMARK_MODES: tuple[BenchmarkMode, ...] = tuple(BenchmarkMode)
+BENCHMARK_CASE_IDS: tuple[BenchmarkCaseId, ...] = tuple(BenchmarkCaseId)
 
 BENCHMARK_SYSTEM_PROMPT = "Exercise the requested business scenario."
 BENCHMARK_MAX_TOKENS = 128
@@ -109,11 +105,11 @@ class BenchmarkTarget(_StrictBenchmarkModel):
 class BenchmarkConfig(BenchmarkTarget):
     schema_version: Literal[1]
     runs_per_case: int = Field(default=1, ge=1)
-    modes: list[BenchmarkMode] = Field(
+    modes: list[Annotated[BenchmarkMode, Field(strict=False)]] = Field(
         default_factory=lambda: list(BENCHMARK_MODES),
         min_length=1,
     )
-    cases: list[BenchmarkCaseId] = Field(
+    cases: list[Annotated[BenchmarkCaseId, Field(strict=False)]] = Field(
         default_factory=lambda: list(BENCHMARK_CASE_IDS),
         min_length=1,
     )
@@ -127,7 +123,10 @@ class BenchmarkConfig(BenchmarkTarget):
 
     @field_validator("modes", "cases")
     @classmethod
-    def reject_duplicates(cls, values: list[str]) -> list[str]:
+    def reject_duplicates(
+        cls,
+        values: list[BenchmarkMode | BenchmarkCaseId],
+    ) -> list[BenchmarkMode | BenchmarkCaseId]:
         if len(values) != len(set(values)):
             raise ValueError("values must not contain duplicates")
         return values
@@ -211,8 +210,8 @@ _SCHEDULE_CONFIRMATION = (
 
 BENCHMARK_CASE_CATALOG: Mapping[BenchmarkCaseId, BenchmarkCase] = MappingProxyType(
     {
-        "ordinary_chat": BenchmarkCase(
-            case_id="ordinary_chat",
+        BenchmarkCaseId.ORDINARY_CHAT: BenchmarkCase(
+            case_id=BenchmarkCaseId.ORDINARY_CHAT,
             user_message="Hello.",
             enabled_tools=(
                 "session.terminate",
@@ -230,8 +229,8 @@ BENCHMARK_CASE_CATALOG: Mapping[BenchmarkCaseId, BenchmarkCase] = MappingProxyTy
             mock_visible_messages=("Ordinary answer.",),
             mock_events=(),
         ),
-        "device_volume_silent": BenchmarkCase(
-            case_id="device_volume_silent",
+        BenchmarkCaseId.DEVICE_VOLUME_SILENT: BenchmarkCase(
+            case_id=BenchmarkCaseId.DEVICE_VOLUME_SILENT,
             user_message="set volume to 30",
             enabled_tools=("device.volume.adjust",),
             expectation=BenchmarkExpectation(
@@ -250,8 +249,8 @@ BENCHMARK_CASE_CATALOG: Mapping[BenchmarkCaseId, BenchmarkCase] = MappingProxyTy
                 ),
             ),
         ),
-        "calendar_schedule_template": BenchmarkCase(
-            case_id="calendar_schedule_template",
+        BenchmarkCaseId.CALENDAR_SCHEDULE_TEMPLATE: BenchmarkCase(
+            case_id=BenchmarkCaseId.CALENDAR_SCHEDULE_TEMPLATE,
             user_message=(
                 'schedule "Design review" at 2026-07-14T09:30:00+08:00 '
                 "timezone Asia/Shanghai"
@@ -273,8 +272,8 @@ BENCHMARK_CASE_CATALOG: Mapping[BenchmarkCaseId, BenchmarkCase] = MappingProxyTy
                 ),
             ),
         ),
-        "web_search_two_answers": BenchmarkCase(
-            case_id="web_search_two_answers",
+        BenchmarkCaseId.WEB_SEARCH_TWO_ANSWERS: BenchmarkCase(
+            case_id=BenchmarkCaseId.WEB_SEARCH_TWO_ANSWERS,
             user_message="event batch safety",
             enabled_tools=("knowledge.web.search",),
             expectation=BenchmarkExpectation(
@@ -293,8 +292,8 @@ BENCHMARK_CASE_CATALOG: Mapping[BenchmarkCaseId, BenchmarkCase] = MappingProxyTy
                 ),
             ),
         ),
-        "session_terminate": BenchmarkCase(
-            case_id="session_terminate",
+        BenchmarkCaseId.SESSION_TERMINATE: BenchmarkCase(
+            case_id=BenchmarkCaseId.SESSION_TERMINATE,
             user_message="end this session",
             enabled_tools=("session.terminate",),
             expectation=BenchmarkExpectation(
@@ -307,8 +306,8 @@ BENCHMARK_CASE_CATALOG: Mapping[BenchmarkCaseId, BenchmarkCase] = MappingProxyTy
             mock_visible_messages=("Goodbye.",),
             mock_events=(_mock_event("terminate-1", "session.terminate"),),
         ),
-        "parallel_volume_schedule": BenchmarkCase(
-            case_id="parallel_volume_schedule",
+        BenchmarkCaseId.PARALLEL_VOLUME_SCHEDULE: BenchmarkCase(
+            case_id=BenchmarkCaseId.PARALLEL_VOLUME_SCHEDULE,
             user_message=(
                 'set volume to 30 and schedule "Design review" at '
                 "2026-07-14T09:30:00+08:00 timezone Asia/Shanghai"
@@ -350,13 +349,19 @@ BENCHMARK_CASE_CATALOG: Mapping[BenchmarkCaseId, BenchmarkCase] = MappingProxyTy
 )
 
 
+def _normalize_benchmark_mode(mode: BenchmarkMode | str) -> BenchmarkMode:
+    try:
+        return BenchmarkMode(mode)
+    except ValueError as exc:
+        raise ValueError(f"unsupported benchmark mode: {mode}") from exc
+
+
 def build_benchmark_request(
     case: BenchmarkCase,
-    mode: BenchmarkMode,
+    mode: BenchmarkMode | str,
     model: str,
 ) -> DebugRunRequest:
-    if mode not in BENCHMARK_MODES:
-        raise ValueError(f"unsupported benchmark mode: {mode}")
+    normalized_mode = _normalize_benchmark_mode(mode)
     validated_model = _validate_model(model)
     return DebugRunRequest(
         user_message=case.user_message,
@@ -376,16 +381,15 @@ def build_benchmark_request(
             max_parallel_events=BENCHMARK_MAX_PARALLEL_EVENTS,
             batch_timeout_seconds=BENCHMARK_BATCH_TIMEOUT_SECONDS,
         ),
-        tool_invocation_mode=mode,
+        tool_invocation_mode=normalized_mode.value,
     )
 
 
 def build_mock_rounds(
     case: BenchmarkCase,
-    mode: BenchmarkMode,
+    mode: BenchmarkMode | str,
 ) -> list[list[StreamItem]]:
-    if mode not in BENCHMARK_MODES:
-        raise ValueError(f"unsupported benchmark mode: {mode}")
+    normalized_mode = _normalize_benchmark_mode(mode)
 
     first_round: list[StreamItem] = []
     if case.mock_visible_messages:
@@ -404,7 +408,7 @@ def build_mock_rounds(
         )
         event_item = (
             StreamItem.text_event(event)
-            if mode == "dual_agent"
+            if normalized_mode is BenchmarkMode.DUAL_AGENT
             else StreamItem.provider_tool_call(event)
         )
         first_round.append(event_item)

+ 55 - 3
tests/test_benchmark_config.py

@@ -1,4 +1,5 @@
 import json
+from enum import StrEnum
 
 import pytest
 from pydantic import ValidationError
@@ -7,7 +8,9 @@ from agent_lab.application.benchmark import (
     BENCHMARK_CASE_CATALOG,
     BENCHMARK_CASE_IDS,
     BENCHMARK_MODES,
+    BenchmarkCaseId,
     BenchmarkConfig,
+    BenchmarkMode,
     BenchmarkTarget,
     build_benchmark_request,
     build_mock_rounds,
@@ -31,8 +34,30 @@ VALID_CONFIG = {
 }
 
 
+def test_benchmark_ids_are_ordered_string_compatible_enums():
+    assert isinstance(BenchmarkMode, type)
+    assert issubclass(BenchmarkMode, StrEnum)
+    assert isinstance(BenchmarkCaseId, type)
+    assert issubclass(BenchmarkCaseId, StrEnum)
+    assert tuple(BenchmarkMode) == BENCHMARK_MODES
+    assert tuple(BenchmarkCaseId) == BENCHMARK_CASE_IDS
+    assert tuple(str(mode) for mode in BENCHMARK_MODES) == (
+        "dual_agent",
+        "chat_agent_tools",
+    )
+    assert tuple(str(case_id) for case_id in BENCHMARK_CASE_IDS) == (
+        "ordinary_chat",
+        "device_volume_silent",
+        "calendar_schedule_template",
+        "web_search_two_answers",
+        "session_terminate",
+        "parallel_volume_schedule",
+    )
+
+
 def test_benchmark_config_accepts_the_versioned_single_target_schema():
     config = BenchmarkConfig.model_validate_json(json.dumps(VALID_CONFIG))
+    python_config = BenchmarkConfig.model_validate(VALID_CONFIG)
 
     assert config.schema_version == 1
     assert config.base_url == "https://provider.example/v1"
@@ -40,6 +65,12 @@ def test_benchmark_config_accepts_the_versioned_single_target_schema():
     assert config.runs_per_case == 1
     assert tuple(config.modes) == BENCHMARK_MODES
     assert tuple(config.cases) == BENCHMARK_CASE_IDS
+    assert all(isinstance(mode, StrEnum) for mode in config.modes)
+    assert all(isinstance(case_id, StrEnum) for case_id in config.cases)
+    assert python_config.modes == config.modes
+    assert python_config.cases == config.cases
+    assert config.model_dump(mode="json")["modes"] == VALID_CONFIG["modes"]
+    assert config.model_dump(mode="json")["cases"] == VALID_CONFIG["cases"]
 
 
 def test_benchmark_target_uses_the_same_strict_url_and_model_validation():
@@ -188,7 +219,12 @@ def test_benchmark_config_rejects_secrets_and_unknown_json_fields(field: str):
 
 
 def test_builtin_case_catalog_is_complete_and_immutable():
+    assert all(isinstance(case_id, StrEnum) for case_id in BENCHMARK_CASE_CATALOG)
     assert tuple(BENCHMARK_CASE_CATALOG) == BENCHMARK_CASE_IDS
+    assert set(BENCHMARK_CASE_CATALOG) == set(BenchmarkCaseId)
+    assert tuple(
+        case.case_id for case in BENCHMARK_CASE_CATALOG.values()
+    ) == BENCHMARK_CASE_IDS
 
     with pytest.raises(TypeError):
         BENCHMARK_CASE_CATALOG["ordinary_chat"] = BENCHMARK_CASE_CATALOG[
@@ -357,7 +393,9 @@ def test_builtin_cases_own_semantics_and_deterministic_mock_data():
 
 
 @pytest.mark.parametrize("case_id", BENCHMARK_CASE_IDS)
-def test_build_benchmark_request_uses_fixed_fairness_controls(case_id: str):
+def test_build_benchmark_request_uses_fixed_fairness_controls(
+    case_id: BenchmarkCaseId,
+):
     case = BENCHMARK_CASE_CATALOG[case_id]
     requests = {
         mode: build_benchmark_request(case, mode, "comparison-model")
@@ -386,6 +424,20 @@ def test_build_benchmark_request_uses_fixed_fairness_controls(case_id: str):
     assert tools.tool_invocation_mode == "chat_agent_tools"
 
 
+@pytest.mark.parametrize("mode", BENCHMARK_MODES)
+def test_benchmark_helpers_accept_enum_and_string_modes(mode: BenchmarkMode):
+    assert isinstance(mode, StrEnum)
+    case = BENCHMARK_CASE_CATALOG["ordinary_chat"]
+
+    enum_request = build_benchmark_request(case, mode, "comparison-model")
+    string_request = build_benchmark_request(case, mode.value, "comparison-model")
+
+    assert enum_request == string_request
+    assert enum_request.tool_invocation_mode == mode.value
+    assert type(enum_request.tool_invocation_mode) is str
+    assert build_mock_rounds(case, mode) == build_mock_rounds(case, mode.value)
+
+
 def test_build_benchmark_request_rejects_an_empty_model():
     case = BENCHMARK_CASE_CATALOG["ordinary_chat"]
 
@@ -396,8 +448,8 @@ def test_build_benchmark_request_rejects_an_empty_model():
 @pytest.mark.parametrize("mode", BENCHMARK_MODES)
 @pytest.mark.parametrize("case_id", BENCHMARK_CASE_IDS)
 def test_mock_rounds_are_mode_correct_and_derived_from_catalog(
-    mode: str,
-    case_id: str,
+    mode: BenchmarkMode,
+    case_id: BenchmarkCaseId,
 ):
     case = BENCHMARK_CASE_CATALOG[case_id]
     rounds = build_mock_rounds(case, mode)

+ 21 - 3
tests/test_tool_invocation_comparison.py

@@ -2,6 +2,7 @@ import asyncio
 import json
 from collections.abc import AsyncIterator
 from dataclasses import dataclass
+from enum import StrEnum
 from typing import Any
 
 import pytest
@@ -10,6 +11,8 @@ from agent_lab.application.benchmark import (
     BENCHMARK_CASE_CATALOG,
     BENCHMARK_CASE_IDS,
     BENCHMARK_MODES,
+    BenchmarkCaseId,
+    BenchmarkMode,
     build_benchmark_request,
     build_mock_rounds,
 )
@@ -214,7 +217,7 @@ def _provider_tool_names(
     return tuple(tool["function"]["name"] for tool in tools)
 
 
-def _assert_port_calls(scenario: str, ports: RecordingPorts) -> None:
+def _assert_port_calls(scenario: BenchmarkCaseId, ports: RecordingPorts) -> None:
     expected_session: list[tuple[str, str | None]] = []
     expected_volume: list[tuple[str, str, int | None, int | None]] = []
     expected_calendar: list[
@@ -258,12 +261,27 @@ def _assert_port_calls(scenario: str, ports: RecordingPorts) -> None:
     assert ports.search_calls == expected_search
 
 
+def test_invocation_matrix_remains_twelve_enum_runs():
+    matrix = tuple(
+        (mode, scenario)
+        for scenario in BENCHMARK_CASE_IDS
+        for mode in BENCHMARK_MODES
+    )
+
+    assert len(matrix) == 12
+    assert len(set(matrix)) == 12
+    assert all(isinstance(mode, StrEnum) for mode, _ in matrix)
+    assert all(isinstance(scenario, StrEnum) for _, scenario in matrix)
+    assert {type(mode) for mode, _ in matrix} == {BenchmarkMode}
+    assert {type(scenario) for _, scenario in matrix} == {BenchmarkCaseId}
+
+
 @pytest.mark.asyncio
 @pytest.mark.parametrize("mode", BENCHMARK_MODES)
 @pytest.mark.parametrize("scenario", BENCHMARK_CASE_IDS)
 async def test_invocation_modes_share_business_semantics(
-    mode: str,
-    scenario: str,
+    mode: BenchmarkMode,
+    scenario: BenchmarkCaseId,
 ) -> None:
     case = BENCHMARK_CASE_CATALOG[scenario]
     request = build_benchmark_request(case, mode, "comparison-model")