Python: Fix A2AAgent to invoke context providers before and after run (#4757)

* Fix A2AAgent to invoke context providers before and after run

A2AAgent.run() bypassed the context provider lifecycle (before_run/after_run)
that BaseAgent defines as a contract for all agents. This caused A2AAgent to
violate the semantic definition of BaseAgent, resulting in inconsistency with
other agent implementations.

The fix follows the same pattern used by WorkflowAgent:
- Create SessionContext and run before_run on all context providers before
  processing the A2A stream
- Collect response updates and run after_run on all context providers after
  the stream is fully consumed
- Auto-create a session when context providers are configured but no session
  is explicitly passed

Fixes #4754

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Apply pre-commit auto-fixes

* Remove reproduction report from repository

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Address PR review feedback for #4754

- Validate messages when no continuation_token: raise ValueError if
  normalized_messages is empty, preventing IndexError on messages[-1]
- Import BaseContextProvider/SessionContext from public agent_framework
  package instead of internal agent_framework._sessions module
- Add test for ValueError on run(None) without continuation_token

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Improve test coverage for empty-messages guard in A2AAgent.run (#4754)

- Parameterize test to cover both messages=None and messages=[] inputs
- Add test verifying run(None, continuation_token=...) does not raise

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

---------

Co-authored-by: Copilot <copilot@github.com>
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
Eduard van Valkenburg
2026-03-19 11:45:42 +01:00
committed by GitHub
Unverified
parent bf8d9672e1
commit 4b21f38650
2 changed files with 243 additions and 4 deletions
+189 -1
View File
@@ -23,11 +23,14 @@ from a2a.types import Role as A2ARole
from agent_framework import (
AgentResponse,
AgentResponseUpdate,
AgentSession,
BaseContextProvider,
Content,
Message,
SessionContext,
)
from agent_framework.a2a import A2AAgent
from pytest import fixture, raises
from pytest import fixture, mark, raises
from agent_framework_a2a import A2AContinuationToken
from agent_framework_a2a._agent import _get_uri_data # type: ignore
@@ -851,3 +854,188 @@ async def test_poll_task_completed(a2a_agent: A2AAgent, mock_a2a_client: MockA2A
# endregion
# region Context Provider Tests
class TrackingContextProvider(BaseContextProvider):
"""A context provider that records when before_run and after_run are called."""
def __init__(self) -> None:
super().__init__(source_id="tracking-provider")
self.before_run_called = False
self.after_run_called = False
self.before_run_context: SessionContext | None = None
self.after_run_context: SessionContext | None = None
async def before_run(
self,
*,
agent: Any,
session: AgentSession,
context: SessionContext,
state: dict[str, Any],
) -> None:
self.before_run_called = True
self.before_run_context = context
async def after_run(
self,
*,
agent: Any,
session: AgentSession,
context: SessionContext,
state: dict[str, Any],
) -> None:
self.after_run_called = True
self.after_run_context = context
async def test_run_invokes_context_providers(mock_a2a_client: MockA2AClient) -> None:
"""Test that context providers are invoked during non-streaming run."""
provider = TrackingContextProvider()
agent = A2AAgent(
name="Test Agent",
client=mock_a2a_client,
context_providers=[provider],
http_client=None,
)
mock_a2a_client.add_message_response("msg-1", "Hello from A2A")
session = agent.create_session()
response = await agent.run("Hello", session=session)
assert provider.before_run_called
assert provider.after_run_called
assert response.text == "Hello from A2A"
async def test_run_streaming_invokes_context_providers(mock_a2a_client: MockA2AClient) -> None:
"""Test that context providers are invoked during streaming run."""
provider = TrackingContextProvider()
agent = A2AAgent(
name="Test Agent",
client=mock_a2a_client,
context_providers=[provider],
http_client=None,
)
mock_a2a_client.add_message_response("msg-1", "Streamed response")
session = agent.create_session()
stream = agent.run("Hello", stream=True, session=session)
updates = []
async for update in stream:
updates.append(update)
assert provider.before_run_called
assert provider.after_run_called
assert len(updates) == 1
assert updates[0].text == "Streamed response"
async def test_context_providers_receive_response(mock_a2a_client: MockA2AClient) -> None:
"""Test that after_run providers can access the response via session context."""
provider = TrackingContextProvider()
agent = A2AAgent(
name="Test Agent",
client=mock_a2a_client,
context_providers=[provider],
http_client=None,
)
mock_a2a_client.add_message_response("msg-1", "Response text")
session = agent.create_session()
await agent.run("Hello", session=session)
assert provider.after_run_context is not None
assert provider.after_run_context.response is not None
assert provider.after_run_context.response.text == "Response text"
async def test_context_providers_receive_input_messages(mock_a2a_client: MockA2AClient) -> None:
"""Test that before_run providers can access input messages via session context."""
provider = TrackingContextProvider()
agent = A2AAgent(
name="Test Agent",
client=mock_a2a_client,
context_providers=[provider],
http_client=None,
)
mock_a2a_client.add_message_response("msg-1", "Reply")
session = agent.create_session()
await agent.run("Hello world", session=session)
assert provider.before_run_context is not None
assert len(provider.before_run_context.input_messages) > 0
assert provider.before_run_context.input_messages[-1].text == "Hello world"
async def test_run_without_context_providers(mock_a2a_client: MockA2AClient) -> None:
"""Test that run works normally when no context providers are configured."""
agent = A2AAgent(
name="Test Agent",
client=mock_a2a_client,
http_client=None,
)
mock_a2a_client.add_message_response("msg-1", "Hello")
response = await agent.run("Hello")
assert response.text == "Hello"
async def test_run_creates_session_for_providers_when_none_provided(mock_a2a_client: MockA2AClient) -> None:
"""Test that a session is auto-created when context providers are configured but no session is passed."""
provider = TrackingContextProvider()
agent = A2AAgent(
name="Test Agent",
client=mock_a2a_client,
context_providers=[provider],
http_client=None,
)
mock_a2a_client.add_message_response("msg-1", "Hello")
await agent.run("Hello")
assert provider.before_run_called
assert provider.after_run_called
@mark.parametrize("messages", [None, []])
async def test_run_raises_when_no_messages_and_no_continuation_token(
mock_a2a_client: MockA2AClient, messages: list[str] | None
) -> None:
"""Test that run() raises ValueError when messages is None/empty and no continuation_token is provided."""
agent = A2AAgent(
name="Test Agent",
client=mock_a2a_client,
http_client=None,
)
with raises(ValueError, match="At least one message is required"):
await agent.run(messages)
async def test_run_with_continuation_token_does_not_require_messages(mock_a2a_client: MockA2AClient) -> None:
"""Test that run() does not raise when messages is None but a continuation_token is provided."""
task = Task(
id="task-cont",
context_id="ctx-cont",
status=TaskStatus(state=TaskState.completed, message=None),
)
mock_a2a_client.resubscribe_responses.append((task, None))
agent = A2AAgent(
name="Test Agent",
client=mock_a2a_client,
http_client=None,
)
token = A2AContinuationToken(task_id="task-cont", context_id="ctx-cont")
response = await agent.run(None, continuation_token=token)
assert response is not None
# endregion