Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions src/google/adk/evaluation/eval_case.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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."""
Expand Down
9 changes: 7 additions & 2 deletions src/google/adk/evaluation/evaluation_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
17 changes: 14 additions & 3 deletions src/google/adk/evaluation/local_eval_sets_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
28 changes: 28 additions & 0 deletions tests/unittests/evaluation/test_eval_case.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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) == []
Expand Down
94 changes: 94 additions & 0 deletions tests/unittests/evaluation/test_evaluation_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""

Expand Down Expand Up @@ -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
Expand Down
33 changes: 33 additions & 0 deletions tests/unittests/evaluation/test_local_eval_sets_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down