diff --git a/src/google/adk/evaluation/_efficiency_evaluators.py b/src/google/adk/evaluation/_efficiency_evaluators.py index 71115f07333..3ddc17c4a5c 100644 --- a/src/google/adk/evaluation/_efficiency_evaluators.py +++ b/src/google/adk/evaluation/_efficiency_evaluators.py @@ -199,6 +199,8 @@ def __init__( self._eval_metric = eval_metric def _compute_value(self, invocation: Invocation) -> Optional[float]: + if invocation.inference_call_count is not None: + return float(invocation.inference_call_count) if not isinstance(invocation.intermediate_data, InvocationEvents): return None model_events = [ diff --git a/src/google/adk/evaluation/eval_case.py b/src/google/adk/evaluation/eval_case.py index 1e424d5f77c..292e713443f 100644 --- a/src/google/adk/evaluation/eval_case.py +++ b/src/google/adk/evaluation/eval_case.py @@ -116,6 +116,14 @@ class Invocation(EvalBaseModel): creation_timestamp: float = 0.0 """Timestamp for the current invocation, primarily intended for debugging purposes.""" + inference_call_count: Optional[int] = Field(default=None, ge=0) + """Number of model requests reconstructed from streamed event boundaries. + + Set during session-event conversion, before turn completion markers are + removed. None keeps event-based counting for older eval files and unary + invocations. Usage and intermediate content remain on their original events. + """ + duration: Optional[float] = None """Wall-clock seconds this invocation took, measured while it ran. diff --git a/src/google/adk/evaluation/evaluation_generator.py b/src/google/adk/evaluation/evaluation_generator.py index 65e0187b384..1e27e4d9cab 100644 --- a/src/google/adk/evaluation/evaluation_generator.py +++ b/src/google/adk/evaluation/evaluation_generator.py @@ -1028,7 +1028,10 @@ def convert_events_to_eval_invocations( final_response = event.content final_event = event - should_add_event = event.grounding_metadata is not None + should_add_event = ( + event.grounding_metadata is not None + or event.usage_metadata is not None + ) for p in event.content.parts: if ( p.function_call @@ -1040,7 +1043,10 @@ def convert_events_to_eval_invocations( break if should_add_event: events_to_add.append(event) - elif event.grounding_metadata is not None: + elif ( + event.grounding_metadata is not None + or event.usage_metadata is not None + ): events_to_add.append(event) invocation_events = [] @@ -1062,15 +1068,54 @@ def convert_events_to_eval_invocations( model_version=e.model_version, ) ) + event_pairs = list(zip(events_to_add, invocation_events)) + merged_events = [] + merged_event_sources: list[Event] = [] + for index, (event, invocation_event) in enumerate(event_pairs): + if ( + not (event.content and event.content.parts) + and event.grounding_metadata is None + and event.usage_metadata is not None + ): + # Live usage can arrive before or after content. Merge into an event + # already counted as a model call, without changing the input events + # or overwriting usage reported by another call. + # TODO: Index by author, model and Live session if long invocations + # make this quadratic search expensive. + model_event = next( + ( + candidate + for candidate_source, candidate in ( + event_pairs[:index][::-1] + event_pairs[index + 1 :] + ) + if candidate_source.live_session_id == event.live_session_id + and candidate.model_version is not None + and candidate.author == event.author + and candidate.usage_metadata is None + and ( + event.model_version is None + or candidate.model_version == event.model_version + ) + ), + None, + ) + if model_event is not None: + model_event.usage_metadata = event.usage_metadata + continue + merged_events.append(invocation_event) + merged_event_sources.append(event) invocations.append( Invocation( invocation_id=invocation_id, user_content=user_content, final_response=final_response, intermediate_data=InvocationEvents( - invocation_events=invocation_events + invocation_events=merged_events ), creation_timestamp=invocation_timestamp, + inference_call_count=EvaluationGenerator._count_streamed_inference_calls( + events, merged_event_sources + ), duration=(durations_per_invocation or {}).get(invocation_id), app_details=app_details, ) @@ -1078,6 +1123,83 @@ def convert_events_to_eval_invocations( return invocations + @staticmethod + def _count_streamed_inference_calls( + events: list[Event], merged_event_sources: list[Event] + ) -> Optional[int]: + """Counts requests before stream completion markers are projected out.""" + stream_keys = { + (event.author, event.live_session_id) + for event in events + if event.live_session_id is not None + or event.turn_complete is not None + or event.partial + } + if not stream_keys: + # Old eval files and ordinary unary responses retain event-based counting. + return None + + # Standalone usage merged into a model event must not start another call. + # Keep the raw sequence for boundaries, including markers absent from evals. + model_event_ids = {id(event) for event in merged_event_sources} + count = 0 + active_sessions: set[tuple[str, Optional[str]]] = set() + for event in events: + if event.author.lower() == _USER_AUTHOR: + # Live input transcriptions can interleave with the model response. + # Their session ID survives normalization into user text content. + if event.input_transcription is None and event.live_session_id is None: + active_sessions.clear() + continue + + has_model_data = ( + event.model_version is not None or event.usage_metadata is not None + ) + # Completion/interrupt metadata does not start a new request. Text flushes + # carry model_version so sessions that omit partial chunks still count. + has_payload = ( + bool(event.content and event.content.parts) + or event.usage_metadata is not None + or ( + event.grounding_metadata is not None + and not event.turn_complete + and not event.interrupted + ) + ) + key = (event.author, event.live_session_id) + if id(event) in model_event_ids and has_model_data and has_payload: + if key not in stream_keys: + # Unary responses still count individually in mixed invocations, + # including usage-only responses with no content. + count += 1 + elif key not in active_sessions: + count += 1 + active_sessions.add(key) + + if event.get_function_responses() and event.live_session_id is None: + # Legacy tool responses omit connection identity. Preserve their + # author-level boundary without closing other authors' requests. + active_sessions = { + active_key + for active_key in active_sessions + if active_key[0] != event.author + } + elif ( + event.turn_complete + or event.interrupted + or event.get_function_responses() + or ( + event.live_session_id is None + and event.turn_complete is None + and not event.partial + ) + ): + # SSE has a final response, possibly only usage/error metadata. Live + # can flush text before tools, so it needs a real connection boundary. + active_sessions.discard(key) + + return count + @staticmethod def _get_app_details_by_invocation_id( events: list[Event], request_intercepter: _RequestIntercepterPlugin diff --git a/src/google/adk/live/_live_llm_flow.py b/src/google/adk/live/_live_llm_flow.py index 6aa973bffc9..93ba665b264 100644 --- a/src/google/adk/live/_live_llm_flow.py +++ b/src/google/adk/live/_live_llm_flow.py @@ -435,12 +435,14 @@ async def postprocess_live_flow( # Handle transcription events ONCE per llm_response, outside the event loop if llm_response.input_transcription: + model_response_event.live_session_id = llm_response.live_session_id model_response_event.input_transcription = llm_response.input_transcription model_response_event.partial = llm_response.partial yield model_response_event return if llm_response.output_transcription: + model_response_event.live_session_id = llm_response.live_session_id model_response_event.output_transcription = ( llm_response.output_transcription ) diff --git a/src/google/adk/models/gemini_llm_connection.py b/src/google/adk/models/gemini_llm_connection.py index c6bdba5ee13..6f7f638b2d4 100644 --- a/src/google/adk/models/gemini_llm_connection.py +++ b/src/google/adk/models/gemini_llm_connection.py @@ -280,6 +280,7 @@ def __build_full_text_response( grounding_metadata=grounding_metadata, interrupted=interrupted, partial=False, + model_version=self._model_version, live_session_id=self._gemini_session.session_id, ) diff --git a/tests/unittests/evaluation/test_evaluation_generator.py b/tests/unittests/evaluation/test_evaluation_generator.py index b45be459459..ec85e9e6737 100644 --- a/tests/unittests/evaluation/test_evaluation_generator.py +++ b/tests/unittests/evaluation/test_evaluation_generator.py @@ -18,8 +18,13 @@ import builtins from google.adk.agents.base_agent import BaseAgent +from google.adk.agents.invocation_context import InvocationContext +from google.adk.agents.llm_agent import LlmAgent +from google.adk.agents.run_config import RunConfig from google.adk.apps.app import App from google.adk.evaluation import evaluation_generator as evaluation_generator_module +from google.adk.evaluation._efficiency_evaluators import _InferenceCallCountV1Evaluator +from google.adk.evaluation._efficiency_evaluators import _TokenUsageV1Evaluator from google.adk.evaluation.app_details import AgentDetails from google.adk.evaluation.app_details import AppDetails from google.adk.evaluation.conversation_scenarios import ConversationScenario @@ -38,6 +43,9 @@ from google.adk.evaluation.simulation.user_simulator import UserSimulator from google.adk.events.event import Event from google.adk.events.event_actions import EventActions +from google.adk.flows.llm_flows.base_llm_flow import BaseLlmFlow +from google.adk.live.live_request_queue import LiveRequestQueue +from google.adk.models.gemini_llm_connection import GeminiLlmConnection from google.adk.models.llm_request import LlmRequest from google.adk.plugins.base_plugin import BasePlugin from google.adk.sessions.in_memory_session_service import InMemorySessionService @@ -2258,3 +2266,1107 @@ def test_convert_events_empty_invocation_events_when_no_agent_events(): assert len(invocations) == 1 assert invocations[0].intermediate_data.invocation_events == [] + + +@pytest.mark.parametrize( + "chunks", [["Hello"], ["Hel", "lo"], ["H", "el", "lo"]] +) +@pytest.mark.parametrize( + "usage_position", ["before", "after", "combined", "absent"] +) +async def test_live_usage_preserved_without_extra_model_calls( + chunks, usage_position +): + """Exercise real Live response conversion, including consolidated final text.""" + messages = [ + types.LiveServerMessage( + server_content=types.LiveServerContent( + model_turn=types.Content( + role="model", parts=[types.Part(text=text)] + ) + ) + ) + for text in chunks + ] + usage = types.UsageMetadata( + prompt_token_count=10, response_token_count=5, total_token_count=15 + ) + if usage_position == "before": + messages.insert(0, types.LiveServerMessage(usage_metadata=usage)) + elif usage_position == "after": + messages.append(types.LiveServerMessage(usage_metadata=usage)) + elif usage_position == "combined": + messages[0].usage_metadata = usage + messages.append( + types.LiveServerMessage( + server_content=types.LiveServerContent(turn_complete=True) + ) + ) + + class LocalTransport: + session_id = "test-session" + + async def receive(self): + for message in messages: + yield message + + connection = GeminiLlmConnection( + LocalTransport(), model_version="gemini-live-test" + ) + events = [_build_event("user", [types.Part(text="Hi")], "inv1")] + async for response in connection.receive(): + events.append( + Event( + author="agent", + invocation_id="inv1", + **response.model_dump(exclude_none=True), + ) + ) + original_events = [event.model_copy(deep=True) for event in events] + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + expected_tokens = None if usage_position == "absent" else 15 + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == expected_tokens + ) + # Transport chunks and standalone usage reports belong to one model request. + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + assert invocations[0].final_response.parts[0].text == "Hello" + assert events == original_events + + +@pytest.mark.parametrize("content", [None, types.Content(parts=[])]) +def test_standalone_usage_preserved_without_model_event(content): + events = [ + Event( + author="agent", + invocation_id="inv1", + content=content, + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=10, + candidates_token_count=5, + total_token_count=15, + ), + ) + ] + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == 15 + ) + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + + +@pytest.mark.parametrize("usage_first", [False, True]) +def test_standalone_usage_keeps_separate_model_calls_and_invocations( + usage_first, +): + events = [] + for invocation_id in ["inv1", "inv2"]: + for tokens in [15, 28]: + model_event = Event( + author="agent", + invocation_id=invocation_id, + model_version="gemini-test", + content=types.Content(parts=[types.Part(text="response")]), + ) + usage_event = Event( + author="agent", + invocation_id=invocation_id, + model_version="gemini-test", + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=tokens, total_token_count=tokens + ), + ) + events.extend( + [usage_event, model_event] + if usage_first + else [model_event, usage_event] + ) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + for invocation in invocations: + assert [ + e.usage_metadata.total_token_count + for e in invocation.intermediate_data.invocation_events + ] == [15, 28] + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == 43 + ) + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + + +@pytest.mark.parametrize("different_field", ["author", "model_version"]) +def test_standalone_usage_does_not_merge_into_different_model(different_field): + model_event = Event( + author="agent", + invocation_id="inv1", + model_version="gemini-test", + content=types.Content(parts=[types.Part(text="response")]), + ) + usage_event = Event( + author="agent", + invocation_id="inv1", + model_version="gemini-test", + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=15, total_token_count=15 + ), + ) + setattr(usage_event, different_field, "different") + + invocations = EvaluationGenerator.convert_events_to_eval_invocations( + [model_event, usage_event] + ) + + assert len(invocations[0].intermediate_data.invocation_events) == 2 + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == 15 + ) + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + + +async def _live_count_events( + messages, + *, + author="agent", + invocation_id="inv1", + model_version="gemini-live-test", +): + """Convert deterministic transport messages through the real Live adapter.""" + + class LocalTransport: + session_id = "count-test" + + async def receive(self): + for message in messages: + yield message + + connection = GeminiLlmConnection( + LocalTransport(), model_version=model_version + ) + return [ + Event( + author=author, + invocation_id=invocation_id, + **response.model_dump(exclude_none=True), + ) + async for response in connection.receive() + ] + + +def _live_text(text, *, thought=False): + return types.LiveServerMessage( + server_content=types.LiveServerContent( + model_turn=types.Content( + role="model", parts=[types.Part(text=text, thought=thought)] + ) + ) + ) + + +def _live_complete(): + return types.LiveServerMessage( + server_content=types.LiveServerContent(turn_complete=True) + ) + + +@pytest.mark.parametrize( + "shape, expected_tools, expected_text", + [ + ("thought_split", 0, "Answer"), + ("thought_combined", 0, "Answer"), + ("text_tool", 1, "Checking"), + ("tool_only", 1, None), + ("parallel_tools", 2, None), + ("interrupted", 0, "Hello world"), + ], +) +async def test_live_response_segments_count_as_one_inference( + shape, expected_tools, expected_text +): + """Thought, text, and tools within a Live response are one inference.""" + calls = [types.FunctionCall(id="call1", name="lookup", args={})] + if shape == "parallel_tools": + calls.append(types.FunctionCall(id="call2", name="other", args={})) + tool_message = types.LiveServerMessage( + tool_call=types.LiveServerToolCall(function_calls=calls) + ) + if shape == "thought_split": + messages = [_live_text("Reasoning", thought=True), _live_text("Answer")] + elif shape == "thought_combined": + messages = [ + types.LiveServerMessage( + server_content=types.LiveServerContent( + model_turn=types.Content( + role="model", + parts=[ + types.Part(text="Reasoning", thought=True), + types.Part(text="Answer", thought=False), + ], + ) + ) + ) + ] + elif shape == "text_tool": + messages = [_live_text("Checking"), tool_message] + elif shape in ("tool_only", "parallel_tools"): + messages = [tool_message] + else: + messages = [_live_text("Hello "), _live_text("world")] + if shape == "interrupted": + messages.append( + types.LiveServerMessage( + server_content=types.LiveServerContent(interrupted=True) + ) + ) + else: + messages.append(_live_complete()) + events = await _live_count_events(messages) + original = [event.model_copy(deep=True) for event in events] + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + assert ( + len(get_all_tool_calls(invocations[0].intermediate_data)) + == expected_tools + ) + if expected_text is not None: + assert invocations[0].final_response.parts[0].text == expected_text + assert events == original + + +async def test_live_count_keeps_turns_and_interleaved_authors_separate(): + """An author's completion does not close another author's in-flight call.""" + first = await _live_count_events( + [_live_text("First"), _live_complete()], author="child1" + ) + second = await _live_count_events( + [_live_text("Second "), _live_text("answer"), _live_complete()], + author="child2", + ) + third = await _live_count_events( + [_live_text("Third"), _live_complete()], author="child1" + ) + events = first[:-1] + second[:1] + first[-1:] + second[1:] + third + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 3 + ) + + +async def test_typed_user_message_separates_live_generations_mid_turn(): + """A new client message closes the generation that has not completed.""" + events = await _live_count_events([_live_text("First")]) + events.append( + Event( + author="user", + invocation_id="inv1", + content=types.Content( + role="user", parts=[types.Part(text="New question")] + ), + ) + ) + events.extend( + await _live_count_events([_live_text("Second"), _live_complete()]) + ) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + assert invocations[0].user_content.parts[0].text == "New question" + assert invocations[0].final_response.parts[0].text == "Second" + + +@pytest.mark.parametrize( + "chunks", [[b"abcdef"], [b"ab", b"cdef"], [b"ab", b"cd", b"ef"]] +) +@pytest.mark.parametrize("with_usage", [False, True]) +async def test_live_audio_chunks_count_as_one_inference(chunks, with_usage): + """Audio chunk boundaries do not change the count for one generation.""" + messages = [ + types.LiveServerMessage( + server_content=types.LiveServerContent( + model_turn=types.Content( + role="model", + parts=[ + types.Part( + inline_data=types.Blob( + mime_type="audio/pcm;rate=24000", data=chunk + ) + ) + ], + ) + ) + ) + for chunk in chunks + ] + if with_usage: + messages.append( + types.LiveServerMessage( + usage_metadata=types.UsageMetadata( + prompt_token_count=10, + response_token_count=5, + total_token_count=15, + ) + ) + ) + messages.append(_live_complete()) + events = await _live_count_events(messages) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + assert _TokenUsageV1Evaluator().evaluate_invocations( + invocations + ).overall_score == (15 if with_usage else None) + contents = [ + event.content + for event in invocations[0].intermediate_data.invocation_events + if event.content is not None + ] + [invocations[0].final_response] + assert ( + b"".join( + part.inline_data.data + for content in contents + for part in content.parts + if part.inline_data is not None + ) + == b"abcdef" + ) + + +def test_live_tool_response_closes_call_without_turn_complete(): + """Gemini 3.x tools separate two requests before a turn_complete arrives.""" + events = [ + Event( + author="agent", + invocation_id="inv1", + live_session_id="live", + model_version="model", + content=types.Content( + role="model", + parts=[ + types.Part(text="Checking"), + types.Part( + function_call=types.FunctionCall( + id="call1", name="lookup", args={} + ) + ), + ], + ), + ), + Event( + author="agent", + invocation_id="inv1", + content=types.Content( + role="user", + parts=[ + types.Part( + function_response=types.FunctionResponse( + id="call1", name="lookup", response={"ok": True} + ) + ), + ], + ), + ), + Event( + author="agent", + invocation_id="inv1", + live_session_id="live", + partial=True, + model_version="model", + content=types.Content( + role="model", parts=[types.Part(text="Answer")] + ), + ), + Event( + author="agent", + invocation_id="inv1", + live_session_id="live", + partial=False, + model_version="model", + content=types.Content( + role="model", parts=[types.Part(text="Answer")] + ), + ), + Event( + author="agent", + invocation_id="inv1", + live_session_id="live", + turn_complete=True, + model_version="model", + ), + ] + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + assert len(get_all_tool_calls(invocations[0].intermediate_data)) == 1 + assert invocations[0].final_response.parts[0].text == "Answer" + + +def test_partial_usage_is_preserved_and_count_survives_serialization(): + """Filtering or recounting projected chunks loses usage or call boundaries.""" + usage = types.GenerateContentResponseUsageMetadata( + prompt_token_count=10, candidates_token_count=5, total_token_count=15 + ) + events = [ + Event( + author="agent", + invocation_id="inv1", + partial=True, + model_version="model", + usage_metadata=usage, + content=types.Content(parts=[types.Part(text="Hel")]), + ), + Event( + author="agent", + invocation_id="inv1", + partial=False, + model_version="model", + content=types.Content(parts=[types.Part(text="Hello")]), + ), + ] + original = [event.model_copy(deep=True) for event in events] + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + restored = [ + type(invocations[0]).model_validate_json(invocations[0].model_dump_json()) + ] + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(restored) + .overall_score + == 1 + ) + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(restored).overall_score + == 15 + ) + assert events == original + + +def test_non_live_streamed_calls_close_on_their_final_response(): + """Two SSE requests keep separate counts without Live completion markers.""" + events = [] + for answer in ["First", "Second"]: + events.extend([ + Event( + author="agent", + invocation_id="inv1", + model_version="model", + partial=True, + content=types.Content(parts=[types.Part(text=answer[:2])]), + ), + Event( + author="agent", + invocation_id="inv1", + model_version="model", + partial=False, + content=types.Content(parts=[types.Part(text=answer)]), + ), + ]) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + + +async def test_live_persisted_text_without_usage_counts_a_call(): + """Persisted sessions omit partial chunks but retain one completed request.""" + events = await _live_count_events( + [_live_text("Hel"), _live_text("lo"), _live_complete()] + ) + persisted = [event for event in events if not event.partial] + + invocations = EvaluationGenerator.convert_events_to_eval_invocations( + persisted + ) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + assert invocations[0].final_response.parts[0].text == "Hello" + + +async def test_non_live_calls_remain_separate_from_a_live_author(): + """Grouping one author's Live stream must not deduplicate unary sub-agents.""" + events = await _live_count_events( + [_live_text("Hel"), _live_text("lo"), _live_complete()] + ) + events.extend([ + Event( + author="unary_agent", + invocation_id="inv1", + model_version="model", + content=types.Content(parts=[types.Part(text="First")]), + ), + Event( + author="unary_agent", + invocation_id="inv1", + model_version="model", + content=types.Content(parts=[types.Part(text="Second")]), + ), + ]) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 3 + ) + + +@pytest.mark.parametrize("streamed", [False, True]) +async def test_same_author_can_switch_from_live_to_other_requests(streamed): + """Live grouping must not collapse later unary or SSE model requests.""" + events = await _live_count_events([_live_text("Live"), _live_complete()]) + for answer in ["First", "Second"]: + if streamed: + events.append( + Event( + author="agent", + invocation_id="inv1", + model_version="model", + partial=True, + content=types.Content(parts=[types.Part(text=answer)]), + ) + ) + events.append( + Event( + author="agent", + invocation_id="inv1", + model_version="model", + content=types.Content(parts=[types.Part(text=answer)]), + ) + ) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 3 + ) + + +async def test_interrupted_completion_grounding_does_not_start_another_call(): + """An interrupted text flush and its grounding completion share one call.""" + events = await _live_count_events([ + _live_text("Answer"), + types.LiveServerMessage( + server_content=types.LiveServerContent( + turn_complete=True, + interrupted=True, + grounding_metadata=types.GroundingMetadata( + web_search_queries=["query"] + ), + ) + ), + ]) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + assert invocations[0].final_response.parts[0].text == "Answer" + + +def test_new_live_session_separates_unfinished_requests(): + """Reconnects cannot merge requests that share an author but not a session.""" + events = [ + Event( + author="agent", + invocation_id="inv1", + live_session_id=session_id, + partial=True, + model_version="model", + content=types.Content(parts=[types.Part(text="Answer")]), + ) + for session_id in ["first", "second"] + ] + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + + +@pytest.mark.parametrize("usage_first", [False, True]) +async def test_live_grouping_keeps_unary_usage_merging(usage_first): + """Reconstructing streamed boundaries must respect unary usage merging.""" + events = await _live_count_events([_live_text("Live"), _live_complete()]) + content = Event( + author="agent", + invocation_id="inv1", + model_version="model", + content=types.Content(parts=[types.Part(text="Unary")]), + ) + usage = Event( + author="agent", + invocation_id="inv1", + model_version="model", + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=10, candidates_token_count=5, total_token_count=15 + ), + ) + events.extend([usage, content] if usage_first else [content, usage]) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == 15 + ) + + +async def _live_receiver_count_events(messages): + """Run the real LLM Live receiver across successive adapter receives.""" + + class LocalTransport: + session_id = "receiver-test" + + def __init__(self): + self.messages = iter(messages) + + async def receive(self): + for message in self.messages: + yield message + + context = InvocationContext( + invocation_id="inv1", + agent=LlmAgent(name="agent", model="gemini-2.0-flash"), + session=Session(id="session", app_name="test", user_id="user"), + session_service=InMemorySessionService(), + live_request_queue=LiveRequestQueue(), + run_config=RunConfig(), + ) + connection = GeminiLlmConnection( + LocalTransport(), model_version="gemini-live-test" + ) + return [ + event + async for event in BaseLlmFlow()._receive_from_model( + connection, context, LlmRequest() + ) + ] + + +@pytest.mark.parametrize("normalized", [False, True]) +async def test_input_transcription_does_not_split_a_live_request(normalized): + """Recognized user audio can arrive during an existing model turn.""" + events = await _live_receiver_count_events([ + _live_text("Hello "), + types.LiveServerMessage( + server_content=types.LiveServerContent( + input_transcription=types.Transcription( + text="Question", finished=False + ) + ) + ), + _live_text("world"), + _live_complete(), + ]) + if normalized: + events = EvaluationGenerator._normalize_live_transcriptions(events) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + assert invocations[0].final_response.parts[0].text == "Hello world" + if normalized: + assert invocations[0].user_content.parts[0].text == "Question" + + +@pytest.mark.parametrize("normalized", [False, True]) +async def test_live_barge_in_counts_interrupted_and_new_generations(normalized): + """A transcribed interruption followed by a new answer counts twice.""" + events = await _live_receiver_count_events([ + _live_text("First answer"), + types.LiveServerMessage( + server_content=types.LiveServerContent(interrupted=True) + ), + types.LiveServerMessage( + server_content=types.LiveServerContent( + input_transcription=types.Transcription( + text="stop", finished=True + ) + ) + ), + _live_text("Stopped"), + _live_complete(), + ]) + if normalized: + events = EvaluationGenerator._normalize_live_transcriptions(events) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert len(invocations) == 1 + assert invocations[0].inference_call_count == 2 + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + assert invocations[0].final_response.parts[0].text == "Stopped" + if normalized: + assert invocations[0].user_content.parts[0].text == "stop" + + +async def test_late_live_usage_is_preserved_without_starting_another_call(): + """The receiver can surface usage on its next receive after completion.""" + events = await _live_receiver_count_events([ + _live_text("Answer"), + _live_complete(), + types.LiveServerMessage( + usage_metadata=types.UsageMetadata( + prompt_token_count=10, + response_token_count=5, + total_token_count=15, + ) + ), + ]) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == 15 + ) + + +async def test_gemini_3_interruption_then_completion_counts_one_call(): + """Gemini 3.x attaches empty grounding metadata to a completion marker.""" + events = await _live_count_events( + [ + _live_text("Answer"), + types.LiveServerMessage( + server_content=types.LiveServerContent(interrupted=True) + ), + _live_complete(), + ], + model_version="gemini-3.1-flash-live-preview", + ) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + + +@pytest.mark.parametrize( + "part", + [ + types.Part( + executable_code=types.ExecutableCode( + language="PYTHON", code="print(1)" + ) + ), + types.Part( + code_execution_result=types.CodeExecutionResult( + outcome="OUTCOME_OK", output="1" + ) + ), + types.Part(thought_signature=b"signature"), + ], +) +def test_usage_preserved_on_non_text_model_content(part): + event = Event( + author="agent", + invocation_id="inv1", + model_version="gemini-test", + content=types.Content(parts=[part]), + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=10, + candidates_token_count=5, + total_token_count=15, + ), + ) + original_event = event.model_copy(deep=True) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations([event]) + + assert len(invocations[0].intermediate_data.invocation_events) == 1 + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == 15 + ) + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 1 + ) + assert event == original_event + + +def test_standalone_usage_does_not_merge_across_live_connections(): + events = [ + Event( + author="agent", + invocation_id="inv1", + live_session_id="first", + model_version="gemini-test", + content=types.Content(parts=[types.Part(text="response")]), + ), + Event( + author="agent", + invocation_id="inv1", + live_session_id="first", + turn_complete=True, + ), + Event( + author="agent", + invocation_id="inv1", + live_session_id="second", + model_version="gemini-test", + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=15, total_token_count=15 + ), + ), + Event( + author="agent", + invocation_id="inv1", + live_session_id="second", + turn_complete=True, + ), + ] + original_events = [event.model_copy(deep=True) for event in events] + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + invocation_events = invocations[0].intermediate_data.invocation_events + assert len(invocation_events) == 2 + assert invocation_events[0].usage_metadata is None + assert invocation_events[1].usage_metadata.total_token_count == 15 + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == 15 + ) + assert ( + _InferenceCallCountV1Evaluator() + .evaluate_invocations(invocations) + .overall_score + == 2 + ) + assert events == original_events + + +@pytest.mark.parametrize("complete_first", [False, True]) +def test_interleaved_live_connections_keep_independent_call_state( + complete_first, +): + def chunk(session_id): + return Event( + author="agent", + invocation_id="inv1", + live_session_id=session_id, + model_version="model", + partial=True, + content=types.Content(parts=[types.Part(text="chunk")]), + ) + + events = [chunk("first"), chunk("second")] + if complete_first: + events.extend([ + Event( + author="agent", + invocation_id="inv1", + live_session_id="first", + turn_complete=True, + ), + chunk("second"), + ]) + else: + events.extend([chunk("first"), chunk("second")]) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert invocations[0].inference_call_count == 2 + + +@pytest.mark.parametrize("boundary", ["complete", "interrupted", "tool"]) +def test_live_boundary_only_closes_its_own_connection(boundary): + def chunk(): + return Event( + author="agent", + invocation_id="inv1", + live_session_id="active", + model_version="model", + partial=True, + content=types.Content(parts=[types.Part(text="chunk")]), + ) + + boundary_event = Event( + author="agent", invocation_id="inv1", live_session_id="other" + ) + if boundary == "complete": + boundary_event.turn_complete = True + elif boundary == "interrupted": + boundary_event.interrupted = True + else: + boundary_event.content = types.Content( + parts=[ + types.Part( + function_response=types.FunctionResponse( + name="lookup", response={"ok": True} + ) + ) + ] + ) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations( + [chunk(), boundary_event, chunk()] + ) + + assert invocations[0].inference_call_count == 1 + + +@pytest.mark.parametrize("author", ["agent", "unary_agent"]) +async def test_contentless_unary_calls_remain_separate_in_live_invocation( + author, +): + events = await _live_count_events([_live_text("Live"), _live_complete()]) + events.extend([ + Event( + author=author, + invocation_id="inv1", + model_version="model", + usage_metadata=types.GenerateContentResponseUsageMetadata( + prompt_token_count=10, + candidates_token_count=5, + total_token_count=15, + ), + ) + for _ in range(2) + ]) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert invocations[0].inference_call_count == 3 + assert ( + _TokenUsageV1Evaluator().evaluate_invocations(invocations).overall_score + == 30 + ) + + +def test_non_live_streamed_calls_close_on_contentless_final_usage(): + events = [] + for answer in ["First", "Second"]: + usage = types.GenerateContentResponseUsageMetadata(total_token_count=15) + events.extend([ + Event( + author="agent", + invocation_id="inv1", + model_version="model", + partial=True, + content=types.Content(parts=[types.Part(text=answer)]), + usage_metadata=usage, + ), + Event( + author="agent", + invocation_id="inv1", + model_version="model", + partial=False, + usage_metadata=usage, + ), + ]) + + invocations = EvaluationGenerator.convert_events_to_eval_invocations(events) + + assert invocations[0].inference_call_count == 2 diff --git a/tests/unittests/live/test_live_llm_flow.py b/tests/unittests/live/test_live_llm_flow.py index 6acfce02bb7..38b8d5148f5 100644 --- a/tests/unittests/live/test_live_llm_flow.py +++ b/tests/unittests/live/test_live_llm_flow.py @@ -129,12 +129,21 @@ async def test_postprocess_live_flow_yields_voice_activity(): assert events[0].voice_activity == vad -async def test_postprocess_live_flow_yields_input_and_output_transcriptions(): +@pytest.mark.parametrize( + 'transcription_field', ['input_transcription', 'output_transcription'] +) +async def test_postprocess_live_flow_yields_input_and_output_transcriptions( + transcription_field, +): """Input and output transcription updates yield events with partial flags preserved.""" flow = _TestBaseLlmFlow() context = _create_test_context(live_request_queue=LiveRequestQueue()) - input_transcription = types.Transcription(text='hello', finished=False) - response = LlmResponse(input_transcription=input_transcription, partial=True) + transcription = types.Transcription(text='hello', finished=False) + response = LlmResponse( + **{transcription_field: transcription}, + partial=True, + live_session_id='live', + ) event = Event( id='ev-1', invocation_id=context.invocation_id, @@ -149,7 +158,8 @@ async def test_postprocess_live_flow_yields_input_and_output_transcriptions(): ] assert len(events) == 1 - assert events[0].input_transcription == input_transcription + assert getattr(events[0], transcription_field) == transcription + assert events[0].live_session_id == 'live' assert events[0].partial is True