mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat(conformance): Implements adk conformance test cli with replay mode
PiperOrigin-RevId: 808633566
This commit is contained in:
committed by
Copybara-Service
parent
c9ea80af28
commit
e86647d446
@@ -179,6 +179,99 @@ def cli_conformance_create(
|
|||||||
asyncio.run(run_conformance_create(test_paths))
|
asyncio.run(run_conformance_create(test_paths))
|
||||||
|
|
||||||
|
|
||||||
|
@conformance.command("test", cls=HelpfulCommand)
|
||||||
|
@click.argument(
|
||||||
|
"paths",
|
||||||
|
nargs=-1,
|
||||||
|
type=click.Path(
|
||||||
|
exists=True, file_okay=False, dir_okay=True, resolve_path=True
|
||||||
|
),
|
||||||
|
)
|
||||||
|
@click.option(
|
||||||
|
"--mode",
|
||||||
|
type=click.Choice(["replay", "live"], case_sensitive=False),
|
||||||
|
default="replay",
|
||||||
|
show_default=True,
|
||||||
|
help=(
|
||||||
|
"Test mode: 'replay' verifies against recorded interactions, 'live'"
|
||||||
|
" runs evaluation-based verification."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
@click.pass_context
|
||||||
|
def cli_conformance_test(
|
||||||
|
ctx,
|
||||||
|
paths: tuple[str, ...],
|
||||||
|
mode: str,
|
||||||
|
):
|
||||||
|
"""Run conformance tests to verify agent behavior consistency.
|
||||||
|
|
||||||
|
Validates that agents produce consistent outputs by comparing against recorded
|
||||||
|
interactions or evaluating live execution results.
|
||||||
|
|
||||||
|
PATHS can be any number of folder paths. Each folder can either:
|
||||||
|
- Contain a spec.yaml file directly (single test case)
|
||||||
|
- Contain subdirectories with spec.yaml files (multiple test cases)
|
||||||
|
|
||||||
|
If no paths are provided, defaults to searching the 'tests' folder.
|
||||||
|
|
||||||
|
TEST MODES:
|
||||||
|
|
||||||
|
\b
|
||||||
|
replay : Verifies agent interactions match previously recorded behaviors
|
||||||
|
exactly. Compares LLM requests/responses and tool calls/results.
|
||||||
|
live : Runs evaluation-based verification (not yet implemented)
|
||||||
|
|
||||||
|
DIRECTORY STRUCTURE:
|
||||||
|
|
||||||
|
Test cases must follow this structure:
|
||||||
|
|
||||||
|
\b
|
||||||
|
category/
|
||||||
|
test_name/
|
||||||
|
spec.yaml # Test specification
|
||||||
|
generated-recordings.yaml # Recorded interactions (replay mode)
|
||||||
|
generated-session.yaml # Session data (replay mode)
|
||||||
|
|
||||||
|
EXAMPLES:
|
||||||
|
|
||||||
|
\b
|
||||||
|
# Run all tests in current directory's 'tests' folder
|
||||||
|
adk conformance test
|
||||||
|
|
||||||
|
\b
|
||||||
|
# Run tests from specific folders
|
||||||
|
adk conformance test tests/core tests/tools
|
||||||
|
|
||||||
|
\b
|
||||||
|
# Run a single test case
|
||||||
|
adk conformance test tests/core/description_001
|
||||||
|
|
||||||
|
\b
|
||||||
|
# Run in live mode (when available)
|
||||||
|
adk conformance test --mode=live tests/core
|
||||||
|
"""
|
||||||
|
|
||||||
|
try:
|
||||||
|
from .conformance.cli_test import run_conformance_test
|
||||||
|
except ImportError as e:
|
||||||
|
click.secho(
|
||||||
|
f"Error: Missing conformance testing dependencies: {e}",
|
||||||
|
fg="red",
|
||||||
|
err=True,
|
||||||
|
)
|
||||||
|
click.secho(
|
||||||
|
"Please install the required conformance testing package dependencies.",
|
||||||
|
fg="yellow",
|
||||||
|
err=True,
|
||||||
|
)
|
||||||
|
ctx.exit(1)
|
||||||
|
|
||||||
|
# Convert to Path objects, use default if empty (paths are already resolved by Click)
|
||||||
|
test_paths = [Path(p) for p in paths] if paths else [Path("tests").resolve()]
|
||||||
|
|
||||||
|
asyncio.run(run_conformance_test(test_paths=test_paths, mode=mode.lower()))
|
||||||
|
|
||||||
|
|
||||||
@main.command("create", cls=HelpfulCommand)
|
@main.command("create", cls=HelpfulCommand)
|
||||||
@click.option(
|
@click.option(
|
||||||
"--model",
|
"--model",
|
||||||
|
|||||||
@@ -0,0 +1,55 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Loading utilities for conformance testing."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import click
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
from ...sessions.session import Session
|
||||||
|
from .test_case import TestSpec
|
||||||
|
|
||||||
|
|
||||||
|
def load_test_case(test_case_dir: Path) -> TestSpec:
|
||||||
|
"""Load TestSpec from spec.yaml file."""
|
||||||
|
spec_file = test_case_dir / "spec.yaml"
|
||||||
|
with open(spec_file, "r", encoding="utf-8") as f:
|
||||||
|
data: dict[str, Any] = yaml.safe_load(f)
|
||||||
|
return TestSpec.model_validate(data)
|
||||||
|
|
||||||
|
|
||||||
|
def load_recorded_session(test_case_dir: Path) -> Optional[Session]:
|
||||||
|
"""Load recorded session data from generated-session.yaml file."""
|
||||||
|
session_file = test_case_dir / "generated-session.yaml"
|
||||||
|
if not session_file.exists():
|
||||||
|
return None
|
||||||
|
|
||||||
|
with open(session_file, "r", encoding="utf-8") as f:
|
||||||
|
session_data = yaml.safe_load(f)
|
||||||
|
if not session_data:
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
return Session.model_validate(session_data)
|
||||||
|
except Exception as e:
|
||||||
|
click.secho(
|
||||||
|
f"Warning: Failed to parse session data: {e}", fg="yellow", err=True
|
||||||
|
)
|
||||||
|
return None
|
||||||
@@ -0,0 +1,181 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Validation logic for conformance test replay mode."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
import difflib
|
||||||
|
import json
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from ...events.event import Event
|
||||||
|
from ...sessions.session import Session
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ComparisonResult:
|
||||||
|
"""Result of comparing two objects during conformance testing."""
|
||||||
|
|
||||||
|
success: bool
|
||||||
|
error_message: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_mismatch_message(
|
||||||
|
context: str, actual_value: str, recorded_value: str
|
||||||
|
) -> str:
|
||||||
|
"""Generate a generic mismatch error message."""
|
||||||
|
return (
|
||||||
|
f"{context} mismatch - \nActual: \n{actual_value} \nRecorded:"
|
||||||
|
f" \n{recorded_value}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_diff_message(
|
||||||
|
context: str, actual_dict: dict, recorded_dict: dict
|
||||||
|
) -> str:
|
||||||
|
"""Generate a diff-based error message for comparison failures."""
|
||||||
|
# Convert to pretty-printed JSON for better readability
|
||||||
|
actual_json = json.dumps(actual_dict, indent=2, sort_keys=True)
|
||||||
|
recorded_json = json.dumps(recorded_dict, indent=2, sort_keys=True)
|
||||||
|
|
||||||
|
# Generate unified diff
|
||||||
|
diff_lines = list(
|
||||||
|
difflib.unified_diff(
|
||||||
|
recorded_json.splitlines(keepends=True),
|
||||||
|
actual_json.splitlines(keepends=True),
|
||||||
|
fromfile=f"recorded {context}\n",
|
||||||
|
tofile=f"actual {context}\n",
|
||||||
|
lineterm="",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if diff_lines:
|
||||||
|
return f"{context} mismatch:\n" + "".join(diff_lines)
|
||||||
|
else:
|
||||||
|
# Fallback to generic format if diff doesn't work
|
||||||
|
return _generate_mismatch_message(context, actual_json, recorded_json)
|
||||||
|
|
||||||
|
|
||||||
|
def compare_event(
|
||||||
|
actual_event: Event, recorded_event: Event, index: int
|
||||||
|
) -> ComparisonResult:
|
||||||
|
"""Compare a single actual event with a recorded event."""
|
||||||
|
# Comprehensive exclude dict for all fields that can differ between runs
|
||||||
|
excluded_fields = {
|
||||||
|
# Event-level fields that vary per run
|
||||||
|
"id": True,
|
||||||
|
"timestamp": True,
|
||||||
|
"invocation_id": True,
|
||||||
|
"long_running_tool_ids": True,
|
||||||
|
# Content fields that vary per run
|
||||||
|
"content": {
|
||||||
|
"parts": {
|
||||||
|
"__all__": {
|
||||||
|
"thought_signature": True,
|
||||||
|
"function_call": {"id": True},
|
||||||
|
"function_response": {"id": True},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
# Action fields that vary per run
|
||||||
|
"actions": {
|
||||||
|
"state_delta": {
|
||||||
|
"_adk_recordings_config": True,
|
||||||
|
"_adk_replay_config": True,
|
||||||
|
},
|
||||||
|
"requested_auth_configs": True,
|
||||||
|
"requested_tool_confirmations": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Compare events using model dumps with comprehensive exclude dict
|
||||||
|
actual_dict = actual_event.model_dump(
|
||||||
|
exclude_none=True, exclude=excluded_fields
|
||||||
|
)
|
||||||
|
recorded_dict = recorded_event.model_dump(
|
||||||
|
exclude_none=True, exclude=excluded_fields
|
||||||
|
)
|
||||||
|
|
||||||
|
if actual_dict != recorded_dict:
|
||||||
|
return ComparisonResult(
|
||||||
|
success=False,
|
||||||
|
error_message=_generate_diff_message(
|
||||||
|
f"event {index}", actual_dict, recorded_dict
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
return ComparisonResult(success=True)
|
||||||
|
|
||||||
|
|
||||||
|
def compare_events(
|
||||||
|
actual_events: list[Event], recorded_events: list[Event]
|
||||||
|
) -> ComparisonResult:
|
||||||
|
"""Compare actual events with recorded events."""
|
||||||
|
if len(actual_events) != len(recorded_events):
|
||||||
|
return ComparisonResult(
|
||||||
|
success=False,
|
||||||
|
error_message=_generate_mismatch_message(
|
||||||
|
"Event count", str(len(actual_events)), str(len(recorded_events))
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
for i, (actual, recorded) in enumerate(zip(actual_events, recorded_events)):
|
||||||
|
result = compare_event(actual, recorded, i)
|
||||||
|
if not result.success:
|
||||||
|
return result
|
||||||
|
|
||||||
|
return ComparisonResult(success=True)
|
||||||
|
|
||||||
|
|
||||||
|
def compare_session(
|
||||||
|
actual_session: Session, recorded_session: Session
|
||||||
|
) -> ComparisonResult:
|
||||||
|
"""Compare actual session with recorded session using comprehensive exclude list.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
ComparisonResult with success status and optional error message
|
||||||
|
"""
|
||||||
|
# Comprehensive exclude dict for all fields that can differ between runs
|
||||||
|
excluded_fields = {
|
||||||
|
# Session-level fields that vary per run
|
||||||
|
"id": True,
|
||||||
|
"last_update_time": True,
|
||||||
|
# State fields that contain ADK internal configuration
|
||||||
|
"state": {
|
||||||
|
"_adk_recordings_config": True,
|
||||||
|
"_adk_replay_config": True,
|
||||||
|
},
|
||||||
|
# Events comparison handled separately
|
||||||
|
"events": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Compare sessions using model dumps with comprehensive exclude dict
|
||||||
|
actual_dict = actual_session.model_dump(
|
||||||
|
exclude_none=True, exclude=excluded_fields
|
||||||
|
)
|
||||||
|
recorded_dict = recorded_session.model_dump(
|
||||||
|
exclude_none=True, exclude=excluded_fields
|
||||||
|
)
|
||||||
|
|
||||||
|
if actual_dict != recorded_dict:
|
||||||
|
return ComparisonResult(
|
||||||
|
success=False,
|
||||||
|
error_message=_generate_diff_message(
|
||||||
|
"session", actual_dict, recorded_dict
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
return ComparisonResult(success=True)
|
||||||
@@ -0,0 +1,342 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""CLI implementation for ADK conformance testing."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
import textwrap
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import click
|
||||||
|
from google.genai import types
|
||||||
|
|
||||||
|
from ..adk_web_server import RunAgentRequest
|
||||||
|
from ._generated_file_utils import load_recorded_session
|
||||||
|
from ._generated_file_utils import load_test_case
|
||||||
|
from ._replay_validators import compare_events
|
||||||
|
from ._replay_validators import compare_session
|
||||||
|
from .adk_web_server_client import AdkWebServerClient
|
||||||
|
from .test_case import TestCase
|
||||||
|
from .test_case import TestSpec
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _TestResult:
|
||||||
|
"""Result of running a single conformance test."""
|
||||||
|
|
||||||
|
category: str
|
||||||
|
name: str
|
||||||
|
success: bool
|
||||||
|
error_message: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _ConformanceTestSummary:
|
||||||
|
"""Summary of all conformance test results."""
|
||||||
|
|
||||||
|
total_tests: int
|
||||||
|
passed_tests: int
|
||||||
|
failed_tests: int
|
||||||
|
results: list[_TestResult]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def success_rate(self) -> float:
|
||||||
|
"""Calculate the success rate as a percentage."""
|
||||||
|
if self.total_tests == 0:
|
||||||
|
return 0.0
|
||||||
|
return (self.passed_tests / self.total_tests) * 100
|
||||||
|
|
||||||
|
|
||||||
|
class ConformanceTestRunner:
|
||||||
|
"""Runs conformance tests in replay mode."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
test_paths: list[Path],
|
||||||
|
client: AdkWebServerClient,
|
||||||
|
mode: str = "replay",
|
||||||
|
user_id: str = "adk_conformance_test_user",
|
||||||
|
):
|
||||||
|
self.test_paths = test_paths
|
||||||
|
self.mode = mode
|
||||||
|
self.client = client
|
||||||
|
self.user_id = user_id
|
||||||
|
|
||||||
|
def _discover_test_cases(self) -> list[TestCase]:
|
||||||
|
"""Discover test cases from specified folder paths."""
|
||||||
|
test_cases = []
|
||||||
|
for test_path in self.test_paths:
|
||||||
|
if not test_path.exists() or not test_path.is_dir():
|
||||||
|
click.secho(f"Invalid path: {test_path}", fg="yellow", err=True)
|
||||||
|
continue
|
||||||
|
|
||||||
|
for spec_file in test_path.rglob("spec.yaml"):
|
||||||
|
test_case_dir = spec_file.parent
|
||||||
|
category = test_case_dir.parent.name
|
||||||
|
name = test_case_dir.name
|
||||||
|
|
||||||
|
# Skip if recordings missing in replay mode
|
||||||
|
if (
|
||||||
|
self.mode == "replay"
|
||||||
|
and not (test_case_dir / "generated-recordings.yaml").exists()
|
||||||
|
):
|
||||||
|
click.secho(
|
||||||
|
f"Skipping {category}/{name}: no recordings",
|
||||||
|
fg="yellow",
|
||||||
|
err=True,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
test_spec = load_test_case(test_case_dir)
|
||||||
|
test_cases.append(
|
||||||
|
TestCase(
|
||||||
|
category=category,
|
||||||
|
name=name,
|
||||||
|
dir=test_case_dir,
|
||||||
|
test_spec=test_spec,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return sorted(test_cases, key=lambda tc: (tc.category, tc.name))
|
||||||
|
|
||||||
|
async def _run_user_messages(
|
||||||
|
self, session_id: str, test_case: TestCase
|
||||||
|
) -> None:
|
||||||
|
"""Run all user messages for a test case."""
|
||||||
|
for user_message_index, user_message in enumerate(
|
||||||
|
test_case.test_spec.user_messages
|
||||||
|
):
|
||||||
|
content = types.UserContent(parts=[types.Part(text=user_message)])
|
||||||
|
|
||||||
|
request = RunAgentRequest(
|
||||||
|
app_name=test_case.test_spec.agent,
|
||||||
|
user_id=self.user_id,
|
||||||
|
session_id=session_id,
|
||||||
|
new_message=content,
|
||||||
|
streaming=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Run the agent but don't collect events here
|
||||||
|
async for _ in self.client.run_agent(
|
||||||
|
request,
|
||||||
|
mode="replay",
|
||||||
|
test_case_dir=str(test_case.dir),
|
||||||
|
user_message_index=user_message_index,
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def _validate_test_results(
|
||||||
|
self, session_id: str, test_case: TestCase
|
||||||
|
) -> _TestResult:
|
||||||
|
"""Validate test results by comparing with recorded data."""
|
||||||
|
# Get final session and use its events for comparison
|
||||||
|
final_session = await self.client.get_session(
|
||||||
|
app_name=test_case.test_spec.agent,
|
||||||
|
user_id=self.user_id,
|
||||||
|
session_id=session_id,
|
||||||
|
)
|
||||||
|
if not final_session:
|
||||||
|
return _TestResult(
|
||||||
|
category=test_case.category,
|
||||||
|
name=test_case.name,
|
||||||
|
success=False,
|
||||||
|
error_message="No final session available for comparison",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Load recorded session data for comparison
|
||||||
|
recorded_session = load_recorded_session(test_case.dir)
|
||||||
|
if not recorded_session:
|
||||||
|
return _TestResult(
|
||||||
|
category=test_case.category,
|
||||||
|
name=test_case.name,
|
||||||
|
success=False,
|
||||||
|
error_message="No recorded session found for replay comparison",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Compare events and session
|
||||||
|
events_result = compare_events(
|
||||||
|
final_session.events, recorded_session.events
|
||||||
|
)
|
||||||
|
session_result = compare_session(final_session, recorded_session)
|
||||||
|
|
||||||
|
# Determine overall success
|
||||||
|
success = events_result.success and session_result.success
|
||||||
|
error_messages = []
|
||||||
|
if not events_result.success and events_result.error_message:
|
||||||
|
error_messages.append(f"Event mismatch: {events_result.error_message}")
|
||||||
|
if not session_result.success and session_result.error_message:
|
||||||
|
error_messages.append(f"Session mismatch: {session_result.error_message}")
|
||||||
|
|
||||||
|
return _TestResult(
|
||||||
|
category=test_case.category,
|
||||||
|
name=test_case.name,
|
||||||
|
success=success,
|
||||||
|
error_message="\n\n".join(error_messages) if error_messages else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _run_test_case_replay(self, test_case: TestCase) -> _TestResult:
|
||||||
|
"""Run a single test case in replay mode."""
|
||||||
|
try:
|
||||||
|
# Create session
|
||||||
|
session = await self.client.create_session(
|
||||||
|
app_name=test_case.test_spec.agent, user_id=self.user_id, state={}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Run each user message
|
||||||
|
try:
|
||||||
|
await self._run_user_messages(session.id, test_case)
|
||||||
|
except Exception as e:
|
||||||
|
return _TestResult(
|
||||||
|
category=test_case.category,
|
||||||
|
name=test_case.name,
|
||||||
|
success=False,
|
||||||
|
error_message=f"Replay verification failed: {e}",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Validate results and return test result
|
||||||
|
result = await self._validate_test_results(session.id, test_case)
|
||||||
|
|
||||||
|
# Clean up session
|
||||||
|
await self.client.delete_session(
|
||||||
|
app_name=test_case.test_spec.agent,
|
||||||
|
user_id=self.user_id,
|
||||||
|
session_id=session.id,
|
||||||
|
)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
return _TestResult(
|
||||||
|
category=test_case.category,
|
||||||
|
name=test_case.name,
|
||||||
|
success=False,
|
||||||
|
error_message=f"Test setup failed: {e}",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def run_all_tests(self) -> _ConformanceTestSummary:
|
||||||
|
"""Run all discovered test cases."""
|
||||||
|
test_cases = self._discover_test_cases()
|
||||||
|
if not test_cases:
|
||||||
|
click.secho("No test cases found!", fg="yellow", err=True)
|
||||||
|
return _ConformanceTestSummary(
|
||||||
|
total_tests=0,
|
||||||
|
passed_tests=0,
|
||||||
|
failed_tests=0,
|
||||||
|
results=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
click.echo(f"""
|
||||||
|
Found {len(test_cases)} test cases to run in {self.mode} mode
|
||||||
|
""")
|
||||||
|
|
||||||
|
results: list[_TestResult] = []
|
||||||
|
for test_case in test_cases:
|
||||||
|
click.echo(f"Running {test_case.category}/{test_case.name}...", nl=False)
|
||||||
|
if self.mode == "replay":
|
||||||
|
result = await self._run_test_case_replay(test_case)
|
||||||
|
else:
|
||||||
|
# TODO: Implement live mode
|
||||||
|
result = _TestResult(
|
||||||
|
category=test_case.category,
|
||||||
|
name=test_case.name,
|
||||||
|
success=False,
|
||||||
|
error_message="Live mode not yet implemented",
|
||||||
|
)
|
||||||
|
results.append(result)
|
||||||
|
_print_test_case_result(result)
|
||||||
|
|
||||||
|
passed = sum(1 for r in results if r.success)
|
||||||
|
return _ConformanceTestSummary(
|
||||||
|
total_tests=len(results),
|
||||||
|
passed_tests=passed,
|
||||||
|
failed_tests=len(results) - passed,
|
||||||
|
results=results,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def run_conformance_test(
|
||||||
|
test_paths: list[Path],
|
||||||
|
mode: str = "replay",
|
||||||
|
) -> None:
|
||||||
|
"""Run conformance tests."""
|
||||||
|
_print_test_header(mode)
|
||||||
|
|
||||||
|
async with AdkWebServerClient() as client:
|
||||||
|
runner = ConformanceTestRunner(test_paths, client, mode)
|
||||||
|
summary = await runner.run_all_tests()
|
||||||
|
|
||||||
|
_print_test_summary(summary)
|
||||||
|
|
||||||
|
|
||||||
|
def _print_test_header(mode: str) -> None:
|
||||||
|
"""Print the conformance test header."""
|
||||||
|
click.echo("=" * 50)
|
||||||
|
click.echo(f"Running ADK conformance tests in {mode} mode...")
|
||||||
|
click.echo("=" * 50)
|
||||||
|
|
||||||
|
|
||||||
|
def _print_test_case_result(result: _TestResult) -> None:
|
||||||
|
"""Print the result of a single test case."""
|
||||||
|
if result.success:
|
||||||
|
click.secho(" ✓ PASS", fg="green")
|
||||||
|
else:
|
||||||
|
click.secho(" ✗ FAIL", fg="red")
|
||||||
|
if result.error_message:
|
||||||
|
click.secho(f"Error: {result.error_message}", fg="red", err=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _print_test_result_details(result: _TestResult) -> None:
|
||||||
|
"""Print detailed information about a failed test result."""
|
||||||
|
click.secho(f"\n✗ {result.category}/{result.name}\n", fg="red")
|
||||||
|
if result.error_message:
|
||||||
|
indented_message = textwrap.indent(result.error_message, " ")
|
||||||
|
click.secho(indented_message, fg="red", err=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _print_test_summary(summary: _ConformanceTestSummary) -> None:
|
||||||
|
"""Print the conformance test summary results."""
|
||||||
|
# Print summary
|
||||||
|
click.echo("\n" + "=" * 50)
|
||||||
|
click.echo("CONFORMANCE TEST SUMMARY")
|
||||||
|
click.echo("=" * 50)
|
||||||
|
|
||||||
|
if summary.total_tests == 0:
|
||||||
|
click.secho("No tests were run.", fg="yellow")
|
||||||
|
return
|
||||||
|
|
||||||
|
click.echo(f"Total tests: {summary.total_tests}")
|
||||||
|
click.secho(f"Passed: {summary.passed_tests}", fg="green")
|
||||||
|
|
||||||
|
if summary.failed_tests > 0:
|
||||||
|
click.secho(f"Failed: {summary.failed_tests}", fg="red")
|
||||||
|
else:
|
||||||
|
click.echo(f"Failed: {summary.failed_tests}")
|
||||||
|
|
||||||
|
click.echo(f"Success rate: {summary.success_rate:.1f}%")
|
||||||
|
|
||||||
|
# List failed tests
|
||||||
|
failed_tests = [r for r in summary.results if not r.success]
|
||||||
|
if failed_tests:
|
||||||
|
click.echo("\nFailed tests:")
|
||||||
|
for result in failed_tests:
|
||||||
|
_print_test_result_details(result)
|
||||||
|
|
||||||
|
# Exit with error code if any tests failed
|
||||||
|
if summary.failed_tests > 0:
|
||||||
|
raise click.ClickException(f"{summary.failed_tests} test(s) failed")
|
||||||
|
else:
|
||||||
|
click.secho("\nAll tests passed! 🎉", fg="green")
|
||||||
Reference in New Issue
Block a user