diff --git a/configs/prompts/judge.yaml b/configs/prompts/judge.yaml index e4b4ba14..8a56c32d 100644 --- a/configs/prompts/judge.yaml +++ b/configs/prompts/judge.yaml @@ -1071,3 +1071,137 @@ judge: "summary": "<1-2 sentence summary for this turn>" }} ] + + audio_native_prompt: | + You are an expert evaluator analyzing whether an AI voice agent correctly **understood** the key entities a user provided during a spoken conversation. + + This is a speech-to-speech / audio-native agent: there is no reliable speech-to-text transcript of the user to inspect. Instead, the evidence of what the agent actually understood is the **arguments it passed to its tool calls**. When a user says "my id is EM123" and the agent then calls a tool with "EMP124", the agent misheard the entity. + + You are given a Conversation containing user turns, tool calls, and tool responses. **The assistant's own spoken turns are redacted** (shown as `[Assistant speaks]`) and deliberately withheld: the agent's speech is itself imperfectly transcribed, so it must never be used as evidence. Judge only what the user said (user turns) against what the agent did (tool-call arguments). + + Your task — evaluate each of the agent's TOOL CALLS individually: + 1. Each tool call in the Conversation is labelled with a stable id: `Tool Call [N]`. Produce **one output entry per tool call** that uses at least one user-provided entity, keyed by that id (`tool_call_id: N`). Be exhaustive — a single turn often contains MORE THAN ONE tool call (e.g. a lookup followed by a search, or a rebook followed by a seat assignment); rate each labelled tool call separately, never merge two of them or skip one. + 2. Within each tool call, look at **every argument**. For any argument whose value traces back to something the USER said — a user-provided key entity (name, confirmation code, date, amount, seat preference, etc.) — work out what the user actually said for it. + 3. Rate whether the value in that tool call matches what the user said: correct or incorrect. + 4. A user often states an entity in one turn and the agent uses it in a tool call several turns later — always rate it at the tool call, never at the user turn. If the same user value is used in more than one tool call, rate it in EACH one (so it can be correct in one tool call and wrong in another). + 5. Only rate arguments that trace back to a value the USER provided. Ignore any argument sourced from a tool response or generated by the agent (e.g. a journey/record id returned by a prior search, or a passenger id). If a tool call has no user-provided arguments at all, omit it from your output entirely. + + ## What Counts as an Entity + An entity must have a **specific, concrete value** — something that could be passed as an input to a program or tool (not an AI, but a script or database lookup). Ask yourself: could this value be stored in a variable and used programmatically? + + - Names (people, places, organizations): e.g. "John Smith", "Austin", "Delta Airlines" + - Specific dates and times: e.g. "December 15th", "3:45 PM" — NOT vague references like "tomorrow morning" or "later today" + - Confirmation codes / reference numbers: e.g. "ABC123", "ZK3FFW" + - Flight numbers: e.g. "UA 204" + - Amounts and prices: use the specific value only, e.g. "$120" — for qualifier phrases like "under $120", only use the specific value + - Addresses: e.g. "123 Main Street" + - Phone numbers: e.g. "555-867-5309" + - Email addresses: e.g. "john@example.com" + - Other specific identifiers: seat numbers, loyalty numbers, booking IDs, employee IDs, etc. + + **Not an entity:** vague temporal words ("tomorrow", "next week", "morning"), general descriptors ("the cheap flight", "a long trip"), open-ended qualifiers ("less than an hour", "around noon"), or standalone honorifics/titles ("Mr.", "Ms."). + + ## IMPORTANT: What This Metric Measures + This measures whether the agent **understood** the entities the user provided, as evidenced by the values it passed to its tool calls. It is NOT a faithfulness or tool-correctness metric. Do NOT penalize: + - The agent choosing the wrong tool, or omitting a tool call it should have made + - Values in a tool call that came from a tool RESPONSE or that the agent generated itself (only user-provided values are evaluated) + - A user entity that was never used in any tool call — it simply is not rated (do not report it) + + Only mark a value incorrect when a tool call used a **different value** than the user provided (e.g. user "EM123" → tool call "EMP124", user "March 25th" → tool call "2026-03-26"). + + ## Non-Spoken Tags in the User Text + The user text may contain non-spoken tags — metadata that was never said aloud and must never be treated as entities: + + - **Audio-direction tags** like [slow], [firm], [annoyed] describe how the words were meant to be spoken. Ignore them entirely. + - **Interruption tags** mark text the user likely never finished saying. No special handling is needed: anything that was never spoken cannot appear in a tool call, so there is simply nothing to rate. + + {interruption_tags_reference} + + ## Conversation + {conversation} + + ## Correctness Criteria + - Entity must have been used correctly by the agent — the value passed to the tool call must match what the user said. + - Numbers: "150" and "one hundred fifty" are equivalent + - Dates: "December 15th" and "Dec 15" are equivalent + - Names: Case-insensitive match. Phonetically equivalent spelling variants are acceptable on the **first mention** of a name (e.g., German Meyer/Meier/Maier/Mayer — all sound the same when spoken). Once the user has explicitly disambiguated the spelling (e.g., spelled it out letter by letter, or corrected a prior spelling), subsequent usage must match that clarified spelling exactly. Flag if: + - The value the agent used clearly differs (different sound, not just different spelling) + - The agent used a different spelling variant **after** the user has explicitly clarified the correct spelling + - Non-Latin scripts: A name or identifier rendered in its native writing system (e.g. Kanji, Hangul, Arabic) is equivalent to its romanized/phonetic form if they represent the same spoken entity. For example, "山田" and "Yamada" should be treated as matching. + - Phonetic letter spellings: In languages that do not use the Latin alphabet, Roman/Latin letters are spoken using language-specific phonetic equivalents. These are correct, not errors (e.g., in Korean "에이" for "A" or "비" for "B"). + - Number surface variants: Many languages have systematic spoken-form differences that do not change the underlying value and must not be penalized: + - Reversed digit order in spoken form (e.g., German "vierundzwanzig" = four-and-twenty = 24) + - Grammatical gender/agreement variants (e.g., Spanish "veintiún" vs "veintiuna" for 21) + - Large-unit grouping differences (e.g., Japanese 万-system: "ichi-man go-sen" / 一万五千 = 15,000) + + Important note: The user text will often feature things formatted like "one two three" instead of "123". Your goal is to evaluate semantic equivalence — these are considered equivalent if the value matches what the agent used. + + ## Examples + + The user gives an id in turn 1 and corrects it in turn 2; nothing is rated at those user turns — each **tool call** is rated on its own. Tool Call [1] used the pre-correction value EM123 (wrong, because the user had already corrected it) and Tool Call [2] used the corrected value EM456 (right). Tool Call [3] uses only a `record_id` that came from a tool response, so it is omitted entirely. + + **Example Input (Conversation):** + Turn 1 - User: My employee id is E M one two three. + Turn 1 - [Assistant speaks] + Turn 2 - User: Actually, sorry — the id is E M four five six. + Turn 2 - [Assistant speaks] + Turn 4 - [Assistant speaks] + Turn 4 - Tool Call [1] (lookup_employee): {{'employee_id': 'EM123'}} + Turn 4 - Tool Response (lookup_employee): {{'error': 'not_found'}} + Turn 6 - [Assistant speaks] + Turn 6 - Tool Call [2] (lookup_employee): {{'employee_id': 'EM456'}} + Turn 6 - Tool Response (lookup_employee): {{'employee_id': 'EM456', 'record_id': 'R-9012'}} + Turn 6 - Tool Call [3] (send_summary): {{'record_id': 'R-9012'}} + + **Example Response:** + [ + {{ + "tool_call_id": 1, + "tool_name": "lookup_employee", + "entities": [ + {{ + "type": "employee_id", + "value": "E M four five six", + "captured_value": "EM123", + "analysis": "The user first said EM123 but corrected it to EM456 in turn 2, so EM456 is the value they provided. This tool call used the pre-correction value EM123 — incorrect.", + "correct": false + }} + ], + "summary": "Tool call used the pre-correction id EM123; the user's provided id is EM456 — incorrect." + }}, + {{ + "tool_call_id": 2, + "tool_name": "lookup_employee", + "entities": [ + {{ + "type": "employee_id", + "value": "E M four five six", + "captured_value": "EM456", + "analysis": "Matches the id the user provided (after their turn-2 correction).", + "correct": true + }} + ], + "summary": "Tool call used EM456 — correct." + }} + ] + + Note Tool Call [3] does not appear: its only argument (record_id) came from a tool response, not the user. + + ## Response Format + Respond with a JSON array. Output **one entry per tool call that uses at least one user-provided entity**, keyed by the tool call's `[N]` id. Tool calls with no user-provided entities, and turns where the user only speaks, must NOT appear. + [ + {{ + "tool_call_id": , + "tool_name": "", + "entities": [ + {{ + "type": "", + "value": "", + "captured_value": "", + "analysis": "", + "correct": + }} + ], + "summary": "<1-2 sentence summary for this tool call>" + }} + ] diff --git a/src/eva/__init__.py b/src/eva/__init__.py index eb2eaa6a..5a7690b0 100644 --- a/src/eva/__init__.py +++ b/src/eva/__init__.py @@ -11,4 +11,4 @@ # Bump metrics_version when changes affect metric computation (metrics code, # judge prompts, pricing tables, postprocessor). -metrics_version = "2.2.0" +metrics_version = "2.3.0" diff --git a/src/eva/metrics/diagnostic/transcription_accuracy_key_entities.py b/src/eva/metrics/diagnostic/transcription_accuracy_key_entities.py index 8647d7f4..387fb8ab 100644 --- a/src/eva/metrics/diagnostic/transcription_accuracy_key_entities.py +++ b/src/eva/metrics/diagnostic/transcription_accuracy_key_entities.py @@ -1,4 +1,16 @@ -"""Transcription accuracy key entities metric using LLM-as-judge (entire conversation). +"""Key entity accuracy metric using LLM-as-judge (entire conversation). + +Measures whether the user's key entities (names, dates, confirmation codes, +amounts, etc.) survived intact through the agent's pipeline. The comparison +target depends on the pipeline type: + +- **CASCADE**: compares what the user was instructed to say (``intended_user_turns``) + against what the agent's STT transcribed (``transcribed_user_turns``) — i.e. + transcription accuracy. +- **Audio-native (S2S / AUDIO_LLM)**: there is no reliable STT output to inspect, + so the agent's **tool-call arguments** (from ``conversation_trace``) become the + evidence of what it understood. If the user says "my id is EM123" and the agent + then calls a tool with "EMP124", the agent misheard the entity. Debug metric for diagnosing model performance issues, not directly used in final evaluation scores. @@ -14,19 +26,24 @@ parse_judge_response_list, resolve_turn_id, ) -from eva.models.config import PipelineType from eva.models.results import MetricScore @register_metric class TranscriptionAccuracyKeyEntitiesMetric(TextJudgeMetric): - """LLM-based transcription accuracy metric for key entities only (entire conversation). + """LLM-based key entity accuracy metric (entire conversation). - Evaluates STT transcription accuracy by comparing key entities (names, dates, - confirmation codes, amounts, etc.) between what the user was supposed to say - (intended_user_turns) and what STT transcribed (transcribed_user_turns). + Evaluates whether the user's key entities (names, dates, confirmation codes, + amounts, etc.) survived intact through the agent's pipeline. The comparison + target depends on the pipeline type: - Computes the ratio of correctly transcribed entities per turn: + - **CASCADE**: compares ``intended_user_turns`` against ``transcribed_user_turns`` + (STT output) — transcription accuracy. + - **Audio-native (S2S / AUDIO_LLM)**: no reliable STT output exists, so the agent's + tool-call arguments (from ``conversation_trace``) are used as the evidence of + what it understood. + + Computes the ratio of correctly captured entities per turn: score = correct_entities / total_entities Entity types evaluated: @@ -48,11 +65,12 @@ class TranscriptionAccuracyKeyEntitiesMetric(TextJudgeMetric): """ name = "transcription_accuracy_key_entities" - version = "v0.3" - description = "Debug metric: LLM judge evaluation of STT key entity transcription accuracy for entire conversation" + version = "v0.4" + description = ( + "Debug metric: LLM judge evaluation of user key entity accuracy (STT in cascade, tool calls in audio-native)" + ) category = "diagnostic" exclude_from_pass_at_k = True - supported_pipeline_types = frozenset({PipelineType.CASCADE}) rating_scale = None # Custom scoring (not 1-3 scale) default_aggregation = "mean" @@ -70,9 +88,7 @@ async def compute(self, context: MetricContext) -> MetricScore: MetricScore with aggregated score and per-turn details """ try: - turns_to_evaluate = self._get_turns_to_evaluate(context) - user_turns_text = self._format_user_turns(turns_to_evaluate, context) - prompt = self.get_judge_prompt(user_turns=user_turns_text) + prompt, turns_to_evaluate = self._build_prompt_and_turns(context) response_text = await self._call_judge_raw(prompt, context) turn_evaluations = parse_judge_response_list(response_text) @@ -98,13 +114,16 @@ async def compute(self, context: MetricContext) -> MetricScore: error=error, ) - # Compute scores for each turn, keyed by turn_id + # Compute scores for each unit (user turn in cascade, tool call in audio-native), keyed by turn_id per_turn_ratings: dict[int, float | None] = {} per_turn_normalized: dict[int, float | None] = {} per_turn_explanations: dict[int, str] = {} per_turn_entity_details: dict[int, dict] = {} for turn_eval in turn_evaluations: + # Cascade keys results on ``turn_id``, audio-native on ``tool_call_id`` — accept either. + if isinstance(turn_eval, dict): + turn_eval.setdefault("turn_id", turn_eval.get("tool_call_id")) turn_id = resolve_turn_id(turn_eval, turns_to_evaluate, self.name) if turn_id is None: continue @@ -206,14 +225,34 @@ def _build_per_entity_type_sub_metrics( ) return sub_metrics + def _build_prompt_and_turns(self, context: MetricContext) -> tuple[str, list[int]]: + """Select the comparison source and judge prompt based on pipeline type. + + CASCADE compares intended vs STT-transcribed user turns, keyed on user + turns. Audio-native pipelines (S2S / AUDIO_LLM) have no reliable STT, so the + judge instead rates the agent's tool-call arguments against what the user + said — keyed on the tool-call turns (see ``_get_audio_native_turns``). + + Returns: + Tuple of (rendered judge prompt, sorted turn IDs to evaluate). + """ + if context.is_audio_native: + conversation, tool_call_ids = self._build_audio_native_conversation(context) + prompt = self.get_judge_prompt(prompt_key="audio_native_prompt", conversation=conversation) + return prompt, tool_call_ids + + turn_ids = self._get_turns_to_evaluate(context) + prompt = self.get_judge_prompt(user_turns=self._format_user_turns(turn_ids, context)) + return prompt, turn_ids + @staticmethod def _get_turns_to_evaluate(context: MetricContext) -> list[int]: - """Return sorted turn IDs present in both tts_text_user and transcript_user.""" + """Return sorted turn IDs present in both intended and transcribed user turns.""" return sorted(context.intended_user_turns.keys() & context.transcribed_user_turns.keys()) @staticmethod def _format_user_turns(turn_ids: list[int], context: MetricContext) -> str: - """Format user turns for the prompt.""" + """Format expected-vs-transcribed user turn pairs for the cascade prompt.""" return "\n\n".join( f"Turn {tid}:\n" f'Expected: "{context.intended_user_turns[tid]}"\n' @@ -221,6 +260,48 @@ def _format_user_turns(turn_ids: list[int], context: MetricContext) -> str: for tid in turn_ids ) + @staticmethod + def _build_audio_native_conversation(context: MetricContext) -> tuple[str, list[int]]: + """Build the redacted conversation for the audio-native judge and the tool-call ids. + + Audio-native scoring is keyed per **tool call** — the point where the agent + commits to a value — not per user turn: a user may state an entity in one turn + and the agent may use it in a tool call several turns later, a single turn may + contain several tool calls, and the same entity can appear (correctly or not) in + more than one. Each tool call is therefore labelled with a stable 1-based id + (``Tool Call [N]``) that the judge rates individually. + + User turns and tool calls/responses are shown verbatim — the user turns are the + entity source, the tool calls are the evidence of what the agent understood. + Assistant turns are redacted to ``[Assistant speaks]`` so the judge cannot use + the agent's own (possibly mis-transcribed) speech to penalize an entity: e.g. if + the user says "INC462" and the agent asks "Did you say INC463?", that read-back + is an artifact of the agent's transcription, not evidence the user was misheard. + + Returns: + Tuple of (formatted conversation, ordered list of tool-call ids). + """ + lines: list[str] = [] + tool_call_ids: list[int] = [] + idx = 0 + for entry in context.conversation_trace or []: + role = entry.get("role") + entry_type = entry.get("type") + turn_id = entry.get("turn_id", "?") + if role == "user": + lines.append(f"Turn {turn_id} - User: {entry.get('content', '')}") + elif role == "assistant": + lines.append(f"Turn {turn_id} - [Assistant speaks]") + elif entry_type == "tool_call": + idx += 1 + tool_call_ids.append(idx) + params = entry.get("parameters", {}) + lines.append(f"Turn {turn_id} - Tool Call [{idx}] ({entry.get('tool_name', 'unknown')}): {params}") + elif entry_type == "tool_response": + resp = entry.get("tool_response", {}) + lines.append(f"Turn {turn_id} - Tool Response ({entry.get('tool_name', 'unknown')}): {resp}") + return "\n".join(lines), tool_call_ids + async def _call_judge_raw(self, prompt: str, context: MetricContext) -> str | None: """Call LLM judge and return raw response text. diff --git a/src/eva/metrics/signatures.py b/src/eva/metrics/signatures.py index c16e4692..bdcae19f 100644 --- a/src/eva/metrics/signatures.py +++ b/src/eva/metrics/signatures.py @@ -64,14 +64,19 @@ def _source_hash(cls: type) -> str: def _prompt_hash_for_metric(cls: type[BaseMetric]) -> str | None: """Return the prompt template hash for judge metrics, or None. - All judge metrics in this codebase use `judge.{name}.user_prompt`. - A judge metric without a corresponding template raises KeyError — - that's a configuration bug we want surfaced. + Hashes *every* prompt template under ``judge.{name}`` (sorted by key), not just + ``user_prompt``, so a metric that ships multiple prompts — e.g. a per-pipeline + variant like ``audio_native_prompt`` — has all of them covered by drift detection. + A judge metric without any templates raises KeyError — that's a configuration + bug we want surfaced. """ if not issubclass(cls, TextJudgeMetric | AudioJudgeMetric): return None - template = get_prompt_manager().get_template(f"judge.{cls.name}.user_prompt") - return hash_prompt_template(template) + templates = get_prompt_manager().prompts.get("judge", {}).get(cls.name) + if not isinstance(templates, dict): + raise KeyError(f"No judge prompts found for metric {cls.name!r}") + combined = "\n".join(templates[key] for key in sorted(templates) if isinstance(templates[key], str)) + return hash_prompt_template(combined) def compute_all_metric_signatures() -> dict[str, dict[str, str | None]]: diff --git a/tests/fixtures/metric_signatures.json b/tests/fixtures/metric_signatures.json index 4dc4639f..0543a4b1 100644 --- a/tests/fixtures/metric_signatures.json +++ b/tests/fixtures/metric_signatures.json @@ -85,9 +85,9 @@ }, "TranscriptionAccuracyKeyEntitiesMetric": { "name": "transcription_accuracy_key_entities", - "prompt_hash": "877dfa61f232", - "source_hash": "70e92b2f7221", - "version": "v0.3" + "prompt_hash": "8492d1f0d168", + "source_hash": "d3c32f2497b1", + "version": "v0.4" }, "TurnTakingMetric": { "name": "turn_taking", diff --git a/tests/unit/metrics/test_transcription_accuracy_key_entities.py b/tests/unit/metrics/test_transcription_accuracy_key_entities.py index 654c7ed7..01ea0112 100644 --- a/tests/unit/metrics/test_transcription_accuracy_key_entities.py +++ b/tests/unit/metrics/test_transcription_accuracy_key_entities.py @@ -9,6 +9,7 @@ TranscriptionAccuracyKeyEntitiesMetric, ) from eva.metrics.utils import aggregate_per_turn_scores +from eva.models.config import PipelineType from .conftest import make_judge_metric, make_metric_context @@ -375,3 +376,248 @@ async def test_no_response_from_judge(self, metric): assert result.error == "No response from judge" assert result.score == 0.0 + + +class TestPipelineSupport: + def test_supports_cascade_and_audio_native(self): + """Metric supports cascade and both audio-native pipeline types.""" + supported = TranscriptionAccuracyKeyEntitiesMetric.supported_pipeline_types + assert PipelineType.CASCADE in supported + assert PipelineType.S2S in supported + assert PipelineType.AUDIO_LLM in supported + + +class TestAudioNativeCompute: + """Audio-native pipelines rate the agent's tool-call arguments against what the user said.""" + + @pytest.mark.asyncio + async def test_uses_audio_native_prompt_with_tool_call_evidence(self, metric): + """Audio-native path builds the tool-call prompt from the redacted conversation trace.""" + context = make_metric_context( + pipeline_type=PipelineType.S2S, + intended_user_turns={1: "My employee id is E M one two three"}, + transcribed_user_turns={}, # unreliable / absent for audio-native + conversation_trace=[ + {"role": "user", "content": "My employee id is E M one two three", "turn_id": 1}, + { + "type": "tool_call", + "tool_name": "lookup_employee", + "parameters": {"employee_id": "EMP124"}, + "turn_id": 1, + }, + ], + ) + response = _make_judge_response( + [ + { + "tool_call_id": 1, + "tool_name": "lookup_employee", + "summary": "Misheard employee id", + "entities": [ + {"type": "employee_id", "value": "EM123", "correct": False, "skipped": False}, + ], + } + ] + ) + generate_text = AsyncMock(return_value=(response, None)) + metric.llm_client.generate_text = generate_text + + result = await metric.compute(context) + + assert result.error is None + assert result.score == 0.0 + assert result.normalized_score == 0.0 + # Prompt must be the audio-native variant fed with tool-call evidence, not STT. + sent_prompt = generate_text.call_args[0][0][0]["content"] + assert "tool call" in sent_prompt.lower() + assert "EMP124" in sent_prompt + assert "Transcribed:" not in sent_prompt + # User turn is shown inline in the single Conversation block, tool call is labelled [1]. + assert "User: My employee id is E M one two three" in sent_prompt + assert "Tool Call [1] (lookup_employee)" in sent_prompt + + @pytest.mark.asyncio + async def test_assistant_turns_are_redacted(self, metric): + """Assistant speech is redacted so a mis-transcribed read-back can't penalize an entity. + + If the user says "QX7710" and the agent's (imperfectly transcribed) reply reads + it back as "QX7711", the judge must not see that text and treat it as a mishear. + (Sentinel values chosen to not collide with any example in the prompt template.) + """ + context = make_metric_context( + pipeline_type=PipelineType.S2S, + intended_user_turns={1: "My incident is QX7710"}, + conversation_trace=[ + {"role": "user", "content": "My incident is QX7710", "turn_id": 1}, + {"role": "assistant", "content": "Did you say QX7711?", "turn_id": 1}, + { + "type": "tool_call", + "tool_name": "get_incident", + "parameters": {"incident_id": "QX7710"}, + "turn_id": 1, + }, + ], + ) + response = _make_judge_response( + [ + { + "tool_call_id": 1, + "summary": "ok", + "entities": [{"type": "incident", "value": "QX7710", "correct": True}], + } + ] + ) + generate_text = AsyncMock(return_value=(response, None)) + metric.llm_client.generate_text = generate_text + + await metric.compute(context) + + sent_prompt = generate_text.call_args[0][0][0]["content"] + assert "[Assistant speaks]" in sent_prompt # assistant turn present but redacted + assert "QX7711" not in sent_prompt # the mis-transcribed read-back is withheld + # Both the user turn AND the tool call must be visible to the judge. + assert "User: My incident is QX7710" in sent_prompt + assert "Tool Call [1] (get_incident): {'incident_id': 'QX7710'}" in sent_prompt + + @pytest.mark.asyncio + async def test_correct_entity_in_tool_call_scores_one(self, metric): + """Entity correctly captured in the tool call → score 1.0.""" + context = make_metric_context( + pipeline_type=PipelineType.AUDIO_LLM, + intended_user_turns={1: "Book it for December 15th"}, + conversation_trace=[ + {"role": "user", "content": "Book it for December 15th", "turn_id": 1}, + { + "type": "tool_call", + "tool_name": "create_booking", + "parameters": {"date": "2026-12-15"}, + "turn_id": 1, + }, + ], + ) + response = _make_judge_response( + [ + { + "tool_call_id": 1, + "summary": "Date captured correctly", + "entities": [{"type": "date", "value": "December 15th", "correct": True, "skipped": False}], + } + ] + ) + metric.llm_client.generate_text = AsyncMock(return_value=(response, None)) + + result = await metric.compute(context) + + assert result.error is None + assert result.score == 1.0 + assert result.normalized_score == 1.0 + + @pytest.mark.asyncio + async def test_rates_per_tool_call_across_turns_with_correction(self, metric): + """Each tool call is rated on its own; a corrected entity is wrong in one call, right in another. + + User states an id in turn 1, corrects it in turn 2; the agent then makes two + tool calls (turns 4 and 6). Nothing is rated at the user turns — the two tool + calls are keyed 1 and 2 and rated independently. + """ + context = make_metric_context( + pipeline_type=PipelineType.S2S, + intended_user_turns={1: "My employee id is E M one two three", 2: "Sorry, it's E M four five six"}, + conversation_trace=[ + {"role": "user", "content": "My employee id is E M one two three", "turn_id": 1}, + {"role": "user", "content": "Sorry, it's E M four five six", "turn_id": 2}, + { + "type": "tool_call", + "tool_name": "lookup_employee", + "parameters": {"employee_id": "EM455"}, + "turn_id": 4, + }, + { + "type": "tool_call", + "tool_name": "lookup_employee", + "parameters": {"employee_id": "EM456"}, + "turn_id": 6, + }, + ], + ) + response = _make_judge_response( + [ + { + "tool_call_id": 1, + "summary": "wrong", + "entities": [{"type": "employee_id", "value": "EM456", "correct": False}], + }, + { + "tool_call_id": 2, + "summary": "right", + "entities": [{"type": "employee_id", "value": "EM456", "correct": True}], + }, + ] + ) + generate_text = AsyncMock(return_value=(response, None)) + metric.llm_client.generate_text = generate_text + + result = await metric.compute(context) + + assert result.error is None + # Rated units are the two tool calls (ids 1, 2), never the user turns (1, 2 as turns). + assert sorted(result.details["per_turn_ratings"].keys()) == [1, 2] + assert result.score == 0.5 # one tool call wrong, one right + sent_prompt = generate_text.call_args[0][0][0]["content"] + assert "Turn 4 - Tool Call [1]" in sent_prompt + assert "Turn 6 - Tool Call [2]" in sent_prompt + + @pytest.mark.asyncio + async def test_multiple_tool_calls_in_one_turn_rated_separately(self, metric): + """Two tool calls in the SAME turn are each labelled and rated individually.""" + context = make_metric_context( + pipeline_type=PipelineType.S2S, + intended_user_turns={1: "Confirmation is six V O R J U, last name Thompson"}, + conversation_trace=[ + {"role": "user", "content": "Confirmation is six V O R J U, last name Thompson", "turn_id": 1}, + { + "type": "tool_call", + "tool_name": "get_reservation", + "parameters": {"confirmation_number": "6VORJU", "last_name": "Thompson"}, + "turn_id": 1, + }, + { + "type": "tool_call", + "tool_name": "search_options", + "parameters": {"confirmation_number": "6VORXX"}, + "turn_id": 1, + }, + ], + ) + response = _make_judge_response( + [ + { + "tool_call_id": 1, + "tool_name": "get_reservation", + "summary": "both correct", + "entities": [ + {"type": "confirmation_code", "value": "6VORJU", "correct": True}, + {"type": "name", "value": "Thompson", "correct": True}, + ], + }, + { + "tool_call_id": 2, + "tool_name": "search_options", + "summary": "code wrong", + "entities": [{"type": "confirmation_code", "value": "6VORJU", "correct": False}], + }, + ] + ) + generate_text = AsyncMock(return_value=(response, None)) + metric.llm_client.generate_text = generate_text + + result = await metric.compute(context) + + assert result.error is None + # Both tool calls in turn 1 are rated as separate units (ids 1 and 2). + assert sorted(result.details["per_turn_ratings"].keys()) == [1, 2] + # Tool call 1: 2/2 correct = 1.0; tool call 2: 0/1 = 0.0; mean = 0.5. + assert result.score == 0.5 + sent_prompt = generate_text.call_args[0][0][0]["content"] + assert "Tool Call [1] (get_reservation)" in sent_prompt + assert "Tool Call [2] (search_options)" in sent_prompt