Skip to content
Merged
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: 1 addition & 3 deletions python/packages/a2a/agent_framework_a2a/_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,9 +157,7 @@ def __init__(
self.client = factory.create(agent_card, interceptors=interceptors) # type: ignore
except Exception as transport_error:
# Transport negotiation failed - fall back to minimal agent card with JSONRPC
fallback_url = (
agent_card.supported_interfaces[0].url if agent_card.supported_interfaces else url
)
fallback_url = agent_card.supported_interfaces[0].url if agent_card.supported_interfaces else url
if not fallback_url:
raise ValueError(
"A2A transport negotiation failed and no fallback URL is available. "
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -520,8 +520,11 @@ async def _run_impl(

session_context = await self._run_before_providers(session=session, input_messages=input_messages, options=opts)

# NOTE: session is created after providers run so that future provider-contributed
# tools/config could be folded into runtime_options before session creation.
# Merge provider-contributed tools into runtime_options before session creation.
if session_context.tools:
existing = list(opts.get("tools") or [])
opts["tools"] = existing + list(session_context.tools)

copilot_session = await self._get_or_create_session(session, streaming=False, runtime_options=opts)

# Build the prompt from the full set of messages in the session context,
Expand Down Expand Up @@ -605,8 +608,11 @@ async def _stream_updates(

session_context = await self._run_before_providers(session=session, input_messages=input_messages, options=opts)

# NOTE: session is created after providers run so that future provider-contributed
# tools/config could be folded into runtime_options before session creation.
# Merge provider-contributed tools into runtime_options before session creation.
if session_context.tools:
existing = list(opts.get("tools") or [])
opts["tools"] = existing + list(session_context.tools)

copilot_session = await self._get_or_create_session(session, streaming=True, runtime_options=opts)

if _ctx_holder is not None:
Expand Down Expand Up @@ -891,7 +897,8 @@ async def _create_session(
mcp_servers = opts.get("mcp_servers") or self._mcp_servers or None
provider = opts.get("provider") or self._provider or None
instruction_directories = opts.get("instruction_directories", self._instruction_directories)
tools = self._prepare_tools(self._tools) if self._tools else None
all_tools = list(self._tools or []) + list(opts.get("tools") or [])
tools = self._prepare_tools(all_tools) if all_tools else None

return await self._client.create_session(
on_permission_request=permission_handler,
Expand Down Expand Up @@ -929,7 +936,8 @@ async def _resume_session(
mcp_servers = opts.get("mcp_servers") or self._mcp_servers or None
provider = opts.get("provider") or self._provider or None
instruction_directories = opts.get("instruction_directories", self._instruction_directories)
tools = self._prepare_tools(self._tools) if self._tools else None
all_tools = list(self._tools or []) + list(opts.get("tools") or [])
tools = self._prepare_tools(all_tools) if all_tools else None

return await self._client.resume_session(
session_id,
Expand Down
228 changes: 228 additions & 0 deletions python/packages/github_copilot/tests/test_github_copilot_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -2477,3 +2477,231 @@ async def test_runtime_on_function_approval_rejected_streaming(self, mock_client
with pytest.raises(ValueError, match="on_function_approval"):
async for _ in agent.run("hello", stream=True, options={"on_function_approval": lambda _c: True}):
pass

async def test_provider_tools_forwarded_to_session(
self,
mock_client: MagicMock,
mock_session: MagicMock,
assistant_message_event: SessionEvent,
) -> None:
"""Test that tools added by context providers are forwarded to session creation."""
mock_session.send_and_wait.return_value = assistant_message_event

class ToolInjectingProvider(ContextProvider):
def __init__(self) -> None:
super().__init__(source_id="tool-injector")

async def before_run(
self,
*,
agent: Any,
session: AgentSession,
context: Any,
state: dict[str, Any],
) -> None:
from agent_framework._tools import normalize_tools

def load_skill(skill_name: str) -> str:
"""Load a skill by name."""
return f"Loaded: {skill_name}"

context.extend_tools(self.source_id, normalize_tools([load_skill]))

provider = ToolInjectingProvider()
agent = GitHubCopilotAgent(client=mock_client, context_providers=[provider])
session = agent.create_session()
await agent.run("Hello", session=session)

call_kwargs = mock_client.create_session.call_args.kwargs
assert call_kwargs.get("tools") is not None
tool_names = [t.name for t in call_kwargs["tools"]]
assert "load_skill" in tool_names

async def test_provider_tools_merged_with_constructor_tools(
self,
mock_client: MagicMock,
mock_session: MagicMock,
assistant_message_event: SessionEvent,
) -> None:
"""Test that provider tools are merged with constructor tools, not replacing them."""
mock_session.send_and_wait.return_value = assistant_message_event

def my_tool(x: str) -> str:
"""A constructor tool."""
return x

class ToolInjectingProvider(ContextProvider):
def __init__(self) -> None:
super().__init__(source_id="tool-injector")

async def before_run(
self,
*,
agent: Any,
session: AgentSession,
context: Any,
state: dict[str, Any],
) -> None:
from agent_framework._tools import normalize_tools

def load_skill(skill_name: str) -> str:
"""Load a skill by name."""
return f"Loaded: {skill_name}"

context.extend_tools(self.source_id, normalize_tools([load_skill]))

provider = ToolInjectingProvider()
agent = GitHubCopilotAgent(
client=mock_client,
tools=[my_tool],
context_providers=[provider],
)
session = agent.create_session()
await agent.run("Hello", session=session)

call_kwargs = mock_client.create_session.call_args.kwargs
assert call_kwargs.get("tools") is not None
tool_names = [t.name for t in call_kwargs["tools"]]
assert "my_tool" in tool_names
assert "load_skill" in tool_names

async def test_provider_tools_forwarded_in_streaming(
self,
mock_client: MagicMock,
mock_session: MagicMock,
assistant_delta_event: SessionEvent,
session_idle_event: SessionEvent,
) -> None:
"""Test that provider tools are forwarded in the streaming path."""
events = [assistant_delta_event, session_idle_event]

def mock_on(handler: Any) -> Any:
for event in events:
handler(event)
return lambda: None

mock_session.on = mock_on

class ToolInjectingProvider(ContextProvider):
def __init__(self) -> None:
super().__init__(source_id="tool-injector")

async def before_run(
self,
*,
agent: Any,
session: AgentSession,
context: Any,
state: dict[str, Any],
) -> None:
from agent_framework._tools import normalize_tools

def load_skill(skill_name: str) -> str:
"""Load a skill by name."""
return f"Loaded: {skill_name}"

context.extend_tools(self.source_id, normalize_tools([load_skill]))

provider = ToolInjectingProvider()
agent = GitHubCopilotAgent(client=mock_client, context_providers=[provider])
session = agent.create_session()
async for _ in agent.run("Hello", stream=True, session=session):
pass

call_kwargs = mock_client.create_session.call_args.kwargs
assert call_kwargs.get("tools") is not None
tool_names = [t.name for t in call_kwargs["tools"]]
assert "load_skill" in tool_names
Comment thread
giles17 marked this conversation as resolved.

async def test_provider_tools_forwarded_to_resume_session(
self,
mock_client: MagicMock,
mock_session: MagicMock,
assistant_message_event: SessionEvent,
) -> None:
"""Test that provider tools are forwarded when resuming an existing session."""
mock_session.send_and_wait.return_value = assistant_message_event

class ToolInjectingProvider(ContextProvider):
def __init__(self) -> None:
super().__init__(source_id="tool-injector")

async def before_run(
self,
*,
agent: Any,
session: AgentSession,
context: Any,
state: dict[str, Any],
) -> None:
from agent_framework._tools import normalize_tools

def load_skill(skill_name: str) -> str:
"""Load a skill by name."""
return f"Loaded: {skill_name}"

context.extend_tools(self.source_id, normalize_tools([load_skill]))

provider = ToolInjectingProvider()
agent = GitHubCopilotAgent(client=mock_client, context_providers=[provider])
session = agent.create_session()
session.service_session_id = "existing-id"
await agent.run("Hello", session=session)

mock_client.create_session.assert_not_called()
mock_client.resume_session.assert_called_once()
call_kwargs = mock_client.resume_session.call_args.kwargs
assert call_kwargs.get("tools") is not None
tool_names = [t.name for t in call_kwargs["tools"]]
assert "load_skill" in tool_names

async def test_provider_tools_forwarded_to_resume_session_streaming(
self,
mock_client: MagicMock,
mock_session: MagicMock,
assistant_delta_event: SessionEvent,
session_idle_event: SessionEvent,
) -> None:
"""Test that provider tools are forwarded when resuming an existing session in streaming mode."""
events = [assistant_delta_event, session_idle_event]

def mock_on(handler: Any) -> Any:
for event in events:
handler(event)
return lambda: None

mock_session.on = mock_on

class ToolInjectingProvider(ContextProvider):
def __init__(self) -> None:
super().__init__(source_id="tool-injector")

async def before_run(
self,
*,
agent: Any,
session: AgentSession,
context: Any,
state: dict[str, Any],
) -> None:
from agent_framework._tools import normalize_tools

def load_skill(skill_name: str) -> str:
"""Load a skill by name."""
return f"Loaded: {skill_name}"

context.extend_tools(self.source_id, normalize_tools([load_skill]))

provider = ToolInjectingProvider()
agent = GitHubCopilotAgent(client=mock_client, context_providers=[provider])
session = agent.create_session()
session.service_session_id = "existing-id"
async for _ in agent.run("Hello", stream=True, session=session):
pass

mock_client.create_session.assert_not_called()
mock_client.resume_session.assert_called_once()
call_kwargs = mock_client.resume_session.call_args.kwargs
assert call_kwargs.get("tools") is not None
tool_names = [t.name for t in call_kwargs["tools"]]
assert "load_skill" in tool_names
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,7 @@ def convert_volume(value: float, factor: float) -> str:
# 2. Define a class-based skill for temperature conversion
# ---------------------------------------------------------------------------


class TemperatureConverterSkill(ClassSkill):
"""A temperature-converter skill defined as a Python class.

Expand Down Expand Up @@ -178,6 +179,7 @@ def convert_temperature(self, value: float, factor: float, offset: float = 0) ->
# 3. Wire everything together and run the agent
# ---------------------------------------------------------------------------


async def main() -> None:
"""Run the combined skills demo."""
endpoint = os.environ["FOUNDRY_PROJECT_ENDPOINT"]
Expand Down
2 changes: 1 addition & 1 deletion python/uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading