feat: Add A2A endpoints to fast api server when --a2a option is specified (WIP)

PiperOrigin-RevId: 776211580
This commit is contained in:
Xiang (Sean) Zhou
2025-06-26 11:18:44 -07:00
committed by Copybara-Service
parent 22629a17bd
commit e79651cd86
2 changed files with 240 additions and 5 deletions
+159 -1
View File
@@ -13,10 +13,14 @@
# limitations under the License.
import asyncio
import json
import logging
import os
from pathlib import Path
import sys
import tempfile
import time
from typing import Any
from typing import Optional
from unittest.mock import MagicMock
from unittest.mock import patch
@@ -465,6 +469,9 @@ def test_app(
artifact_service_uri="",
memory_service_uri="",
allow_origins=["*"],
a2a=False, # Disable A2A for most tests
host="127.0.0.1",
port=8000,
)
# Create a TestClient that doesn't start a real server
@@ -520,6 +527,134 @@ async def create_test_eval_set(
return test_session_info
@pytest.fixture
@pytest.mark.skipif(
sys.version_info < (3, 10), reason="A2A requires Python 3.10+"
)
def temp_agents_dir_with_a2a():
"""Create a temporary agents directory with A2A agent configurations for testing."""
with tempfile.TemporaryDirectory() as temp_dir:
# Create test agent directory
agent_dir = Path(temp_dir) / "test_a2a_agent"
agent_dir.mkdir()
# Create agent.json file
agent_card = {
"name": "test_a2a_agent",
"description": "Test A2A agent",
"version": "1.0.0",
"author": "test",
"capabilities": ["text"],
}
with open(agent_dir / "agent.json", "w") as f:
json.dump(agent_card, f)
# Create a simple agent.py file
agent_py_content = """
from google.adk.agents.base_agent import BaseAgent
class TestA2AAgent(BaseAgent):
def __init__(self):
super().__init__(name="test_a2a_agent")
"""
with open(agent_dir / "agent.py", "w") as f:
f.write(agent_py_content)
yield temp_dir
@pytest.fixture
@pytest.mark.skipif(
sys.version_info < (3, 10), reason="A2A requires Python 3.10+"
)
def test_app_with_a2a(
mock_session_service,
mock_artifact_service,
mock_memory_service,
mock_agent_loader,
mock_eval_sets_manager,
mock_eval_set_results_manager,
temp_agents_dir_with_a2a,
):
"""Create a TestClient for the FastAPI app with A2A enabled."""
# Mock A2A related classes
with (
patch("signal.signal", return_value=None),
patch(
"google.adk.cli.fast_api.InMemorySessionService",
return_value=mock_session_service,
),
patch(
"google.adk.cli.fast_api.InMemoryArtifactService",
return_value=mock_artifact_service,
),
patch(
"google.adk.cli.fast_api.InMemoryMemoryService",
return_value=mock_memory_service,
),
patch(
"google.adk.cli.fast_api.AgentLoader",
return_value=mock_agent_loader,
),
patch(
"google.adk.cli.fast_api.LocalEvalSetsManager",
return_value=mock_eval_sets_manager,
),
patch(
"google.adk.cli.fast_api.LocalEvalSetResultsManager",
return_value=mock_eval_set_results_manager,
),
patch(
"google.adk.cli.cli_eval.run_evals",
new=mock_run_evals_for_fast_api,
),
patch("a2a.server.tasks.InMemoryTaskStore") as mock_task_store,
patch(
"google.adk.a2a.executor.a2a_agent_executor.A2aAgentExecutor"
) as mock_executor,
patch(
"a2a.server.request_handlers.DefaultRequestHandler"
) as mock_handler,
patch("a2a.server.apps.A2AStarletteApplication") as mock_a2a_app,
):
# Configure mocks
mock_task_store.return_value = MagicMock()
mock_executor.return_value = MagicMock()
mock_handler.return_value = MagicMock()
# Mock A2AStarletteApplication
mock_app_instance = MagicMock()
mock_app_instance.routes.return_value = (
[]
) # Return empty routes for testing
mock_a2a_app.return_value = mock_app_instance
# Change to temp directory
original_cwd = os.getcwd()
os.chdir(temp_agents_dir_with_a2a)
try:
app = get_fast_api_app(
agents_dir=".",
web=True,
session_service_uri="",
artifact_service_uri="",
memory_service_uri="",
allow_origins=["*"],
a2a=True,
host="127.0.0.1",
port=8000,
)
client = TestClient(app)
yield client
finally:
os.chdir(original_cwd)
#################################################
# Test Cases
#################################################
@@ -760,5 +895,28 @@ def test_debug_trace(test_app):
logger.info("Debug trace test completed successfully")
@pytest.mark.skipif(
sys.version_info < (3, 10), reason="A2A requires Python 3.10+"
)
def test_a2a_agent_discovery(test_app_with_a2a):
"""Test that A2A agents are properly discovered and configured."""
# This test mainly verifies that the A2A setup doesn't break the app
response = test_app_with_a2a.get("/list-apps")
assert response.status_code == 200
logger.info("A2A agent discovery test passed")
@pytest.mark.skipif(
sys.version_info < (3, 10), reason="A2A requires Python 3.10+"
)
def test_a2a_disabled_by_default(test_app):
"""Test that A2A functionality is disabled by default."""
# The regular test_app fixture has a2a=False
# This test ensures no A2A routes are added
response = test_app.get("/list-apps")
assert response.status_code == 200
logger.info("A2A disabled by default test passed")
if __name__ == "__main__":
pytest.main(["-xvs", __file__])