diff --git a/src/google/adk/evaluation/eval_case.py b/src/google/adk/evaluation/eval_case.py index 43b690734ff..6af51b6f21d 100644 --- a/src/google/adk/evaluation/eval_case.py +++ b/src/google/adk/evaluation/eval_case.py @@ -28,6 +28,7 @@ from .common import EvalBaseModel from .conversation_scenarios import ConversationScenario as ConversationScenario from .eval_rubrics import Rubric +from ..events.event import Event class IntermediateData(EvalBaseModel): @@ -140,6 +141,9 @@ class SessionInput(EvalBaseModel): state: SessionState = Field(default_factory=dict) """The state of the session.""" + events: Optional[list[Event]] = Field(default=None) + """Optional pre-populated events to initialize the session with.""" + StaticConversation: TypeAlias = list[Invocation] """A conversation where the user's queries for each invocation are already specified.""" diff --git a/src/google/adk/evaluation/evaluation_generator.py b/src/google/adk/evaluation/evaluation_generator.py index e7bc95a399f..2b78c2e8965 100644 --- a/src/google/adk/evaluation/evaluation_generator.py +++ b/src/google/adk/evaluation/evaluation_generator.py @@ -142,19 +142,24 @@ async def _get_or_create_eval_session( if pinned_session_id: # A pinned id may name a session the caller prepared, so reuse it instead - # of replacing it; `initial_session.state` then applies only on create. + # of replacing it; `initial_session.state` and `initial_session.events` + # then apply only on create. session = await session_service.get_session( app_name=app_name, user_id=user_id, session_id=pinned_session_id ) if session: return session - return await session_service.create_session( + session = await session_service.create_session( app_name=app_name, user_id=user_id, state=initial_session.state if initial_session else {}, session_id=pinned_session_id or fallback_session_id or str(uuid.uuid4()), ) + if initial_session and initial_session.events: + for event in initial_session.events: + await session_service.append_event(session=session, event=event) + return session # Keyword-argument names accepted by `Runner`, used when building the eval diff --git a/src/google/adk/evaluation/local_eval_sets_manager.py b/src/google/adk/evaluation/local_eval_sets_manager.py index 69e4878b630..44af8eebe4a 100644 --- a/src/google/adk/evaluation/local_eval_sets_manager.py +++ b/src/google/adk/evaluation/local_eval_sets_manager.py @@ -28,6 +28,7 @@ from typing_extensions import override from ..errors.not_found_error import NotFoundError +from ..events.event import Event from ._eval_sets_manager_utils import add_eval_case_to_eval_set from ._eval_sets_manager_utils import delete_eval_case_from_eval_set from ._eval_sets_manager_utils import get_eval_case_from_eval_set @@ -151,10 +152,20 @@ def convert_eval_set_to_pydantic_schema( "initial_session" in old_eval_case and len(old_eval_case["initial_session"]) > 0 ): + initial_session_data = old_eval_case["initial_session"] + raw_events = initial_session_data.get("events") + events = None + if raw_events: + events = [ + Event.model_validate(e) if isinstance(e, dict) else e + for e in raw_events + ] session_input = SessionInput( - app_name=old_eval_case["initial_session"].get("app_name", ""), - user_id=old_eval_case["initial_session"].get("user_id", ""), - state=old_eval_case["initial_session"].get("state", {}), + app_name=initial_session_data.get("app_name", ""), + user_id=initial_session_data.get("user_id", ""), + session_id=initial_session_data.get("session_id"), + state=initial_session_data.get("state", {}), + events=events, ) new_eval_case = EvalCase( diff --git a/tests/unittests/evaluation/test_eval_case.py b/tests/unittests/evaluation/test_eval_case.py index 6c532b13ad0..7ee7cb1561a 100644 --- a/tests/unittests/evaluation/test_eval_case.py +++ b/tests/unittests/evaluation/test_eval_case.py @@ -23,6 +23,7 @@ from google.adk.evaluation.eval_case import InvocationEvent from google.adk.evaluation.eval_case import InvocationEvents from google.adk.evaluation.eval_case import SessionInput +from google.adk.events.event import Event from google.genai import types as genai_types import pytest @@ -79,6 +80,33 @@ def test_session_input_session_id_defaults_to_none(): assert SessionInput(app_name='a', user_id='u').session_id is None +def test_session_input_accepts_events(): + """Tests that SessionInput accepts events and round-trips them.""" + event = Event( + content=genai_types.Content( + parts=[genai_types.Part.from_text(text='hello')] + ), + author='user', + ) + session_input = SessionInput(app_name='a', user_id='u', events=[event]) + + assert session_input.events is not None + assert len(session_input.events) == 1 + assert session_input.events[0].content.parts[0].text == 'hello' + + round_tripped = SessionInput.model_validate_json( + session_input.model_dump_json() + ) + assert round_tripped.events is not None + assert len(round_tripped.events) == 1 + assert round_tripped.events[0].content.parts[0].text == 'hello' + + +def test_session_input_events_defaults_to_none(): + """Tests that events is optional and defaults to None.""" + assert SessionInput(app_name='a', user_id='u').events is None + + def test_get_all_tool_calls_with_none_input(): """Tests that an empty list is returned when intermediate_data is None.""" assert get_all_tool_calls(None) == [] diff --git a/tests/unittests/evaluation/test_evaluation_generator.py b/tests/unittests/evaluation/test_evaluation_generator.py index 70cd5769a5f..06dfa776216 100644 --- a/tests/unittests/evaluation/test_evaluation_generator.py +++ b/tests/unittests/evaluation/test_evaluation_generator.py @@ -27,6 +27,7 @@ from google.adk.evaluation.eval_case import get_all_tool_calls from google.adk.evaluation.eval_case import SessionInput from google.adk.evaluation.eval_set import EvalSet +from google.adk.evaluation.evaluation_generator import _get_or_create_eval_session from google.adk.evaluation.evaluation_generator import _LiveSession from google.adk.evaluation.evaluation_generator import _send_audio_to_live from google.adk.evaluation.evaluation_generator import EvaluationGenerator @@ -831,6 +832,64 @@ async def test_generate_inferences_live_audio_with_text_streams_audio( assert sent_blob.data == b"fake-audio" +class TestGetOrCreateEvalSession: + """Test cases for _get_or_create_eval_session.""" + + @pytest.mark.asyncio + async def test_appends_initial_events_on_create(self): + session_service = InMemorySessionService() + seed_event = _build_event( + "user", [types.Part(text="prior seed event")], "inv_seed" + ) + initial_session = SessionInput( + app_name="test_app", + user_id="test_user", + session_id="s1", + events=[seed_event], + ) + + session = await _get_or_create_eval_session( + session_service=session_service, + initial_session=initial_session, + fallback_session_id=None, + ) + + assert len(session.events) == 1 + assert session.events[0].content.parts[0].text == "prior seed event" + + @pytest.mark.asyncio + async def test_pinned_session_skips_appending_events_if_already_exists(self): + session_service = InMemorySessionService() + existing = await session_service.create_session( + app_name="test_app", + user_id="test_user", + session_id="s1", + ) + existing_event = _build_event( + "user", [types.Part(text="existing")], "inv0" + ) + await session_service.append_event(existing, existing_event) + + seed_event = _build_event( + "user", [types.Part(text="new seed")], "inv_seed" + ) + initial_session = SessionInput( + app_name="test_app", + user_id="test_user", + session_id="s1", + events=[seed_event], + ) + + session = await _get_or_create_eval_session( + session_service=session_service, + initial_session=initial_session, + fallback_session_id=None, + ) + + assert len(session.events) == 1 + assert session.events[0].content.parts[0].text == "existing" + + class TestSendAudioToLive: """Test cases for _send_audio_to_live.""" @@ -1046,6 +1105,41 @@ async def test_pinned_session_id_preserves_existing_session( "earlier turn" ] + @pytest.mark.asyncio + async def test_initial_session_prepopulates_events( + self, mocker, mock_runner + ): + """SessionInput.events are appended to the session on creation.""" + session_service = InMemorySessionService() + seed_event = _build_event( + "user", [types.Part(text="prior context")], "inv_seed" + ) + mock_user_sim = mocker.MagicMock(spec=UserSimulator) + mock_user_sim.get_next_user_message = mocker.AsyncMock( + return_value=NextUserMessage( + status=UserSimulatorStatus.STOP_SIGNAL_DETECTED + ) + ) + + await EvaluationGenerator._generate_inferences_from_root_agent( + root_agent=mocker.MagicMock(), + user_simulator=mock_user_sim, + initial_session=SessionInput( + app_name="test_app", + user_id="u", + session_id="fixed", + events=[seed_event], + ), + session_service=session_service, + ) + + reloaded = await session_service.get_session( + app_name="test_app", user_id="u", session_id="fixed" + ) + assert reloaded is not None + assert len(reloaded.events) == 1 + assert reloaded.events[0].content.parts[0].text == "prior context" + @pytest.mark.asyncio async def test_generates_inferences_with_user_simulator_live( self, mocker, mock_runner, mock_session_service diff --git a/tests/unittests/evaluation/test_local_eval_sets_manager.py b/tests/unittests/evaluation/test_local_eval_sets_manager.py index 541d4dd8b29..4bacd8aa082 100644 --- a/tests/unittests/evaluation/test_local_eval_sets_manager.py +++ b/tests/unittests/evaluation/test_local_eval_sets_manager.py @@ -165,6 +165,39 @@ def test_convert_eval_set_to_pydantic_schema_empty_initial_session(self): assert eval_set.eval_set_id == eval_set_id assert eval_set.eval_cases[0].session_input is None + def test_convert_eval_set_to_pydantic_schema_with_initial_session_events(self): + eval_set_id = "test_eval_set" + eval_set_in_json_format = [{ + "name": "session_with_events", + "data": [{"query": "Test", "reference": "Test Ref"}], + "initial_session": { + "app_name": "my_app", + "user_id": "u1", + "session_id": "s1", + "state": {"k": "v"}, + "events": [{ + "author": "user", + "content": {"parts": [{"text": "Hello context"}]}, + "invocation_id": "inv0", + }], + }, + }] + + eval_set = convert_eval_set_to_pydantic_schema( + eval_set_id, eval_set_in_json_format + ) + + assert eval_set.eval_set_id == eval_set_id + session_input = eval_set.eval_cases[0].session_input + assert session_input is not None + assert session_input.app_name == "my_app" + assert session_input.user_id == "u1" + assert session_input.session_id == "s1" + assert session_input.state == {"k": "v"} + assert session_input.events is not None + assert len(session_input.events) == 1 + assert session_input.events[0].content.parts[0].text == "Hello context" + def test_convert_eval_set_to_pydantic_schema_invalid_data(self): # This test implicitly checks for potential validation errors during Pydantic # object creation