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
117 changes: 112 additions & 5 deletions src/agents/mcp/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,9 @@ class RequireApprovalObject(TypedDict, total=False):

T = TypeVar("T")

_SAFE_EXCEPTION_GROUP_MESSAGE = "MCP request failed with additional errors."
_SAFE_EXCEPTION_MESSAGE = "An additional error occurred during the MCP request."


def _safe_transport_cause(http_error: Exception) -> Exception | None:
"""Keep a transport exception only when its HTTPX URLs need no sanitization."""
Expand Down Expand Up @@ -136,6 +139,37 @@ def _first_unsafe_transport_error(http_errors: list[Exception]) -> Exception | N
return next((error for error in http_errors if _safe_transport_cause(error) is None), None)


def _is_http_transport_error(error: BaseException) -> bool:
"""Return whether an exception is an HTTPX transport error."""
return isinstance(error, httpx.HTTPStatusError | httpx.RequestError)


def _credential_safe_exception_group(error_group: BaseExceptionGroup) -> BaseExceptionGroup:
"""Replace an exception group with a fixed-data graph that retains control semantics."""
safe_exceptions = [
_credential_safe_exception_group(error)
if isinstance(error, BaseExceptionGroup)
else _credential_safe_exception_leaf(error)
for error in error_group.exceptions
]
return BaseExceptionGroup(_SAFE_EXCEPTION_GROUP_MESSAGE, safe_exceptions)


def _credential_safe_exception_leaf(error: BaseException) -> BaseException:
"""Create a fixed-data replacement for one retained exception leaf."""
if isinstance(error, asyncio.CancelledError):
return asyncio.CancelledError()
if isinstance(error, KeyboardInterrupt):
return KeyboardInterrupt()
if isinstance(error, SystemExit):
return SystemExit()
if isinstance(error, GeneratorExit):
return GeneratorExit()
if isinstance(error, Exception):
return RuntimeError(_SAFE_EXCEPTION_MESSAGE)
return BaseException(_SAFE_EXCEPTION_MESSAGE)


def _log_transport_warning(message: str, http_error: Exception) -> None:
"""Log a transport failure without attaching credential-bearing request URLs."""
if _debug.DONT_LOG_TOOL_DATA:
Expand Down Expand Up @@ -826,6 +860,64 @@ def _raise_mapped_transport_error(error: UserError, cause: Exception | None) ->
raise error from None
raise error from cause

def _user_error_for_request_operation(
self,
operation: str,
http_error: Exception,
) -> UserError:
"""Build a credential-safe error for an MCP request operation."""
error_message = f"Failed to {operation} on MCP server '{self._error_name}': "
if isinstance(http_error, httpx.HTTPStatusError):
error_message += f"HTTP error {http_error.response.status_code}"
elif isinstance(http_error, httpx.ConnectError):
error_message += "Connection lost. The server may have disconnected."
elif isinstance(http_error, httpx.TimeoutException):
error_message += "Connection timeout."
else:
error_message += "Request failed."
return UserError(error_message)

async def _run_request_with_transport_error_redaction(
self,
operation: str,
func: Callable[[], Awaitable[T]],
) -> T:
"""Run an MCP request without retaining credential-bearing HTTP errors."""
transport_error: UserError | None = None
base_error_group: BaseExceptionGroup | None = None
try:
return await func()
except (httpx.HTTPStatusError, httpx.RequestError) as http_error:
transport_error = self._user_error_for_request_operation(operation, http_error)
except BaseExceptionGroup as error_group:
http_errors = self._extract_http_errors_from_exception(error_group)
if not http_errors:
raise
Comment thread
seratch marked this conversation as resolved.
selected_http_error = http_errors[0]
http_group, remaining_group = error_group.split(_is_http_transport_error)
assert http_group is not None
mapped_transport_error = self._user_error_for_request_operation(
operation,
selected_http_error,
)
if remaining_group is None:
transport_error = mapped_transport_error
else:
safe_remaining_group = _credential_safe_exception_group(remaining_group)
base_error_group = BaseExceptionGroup(
_SAFE_EXCEPTION_GROUP_MESSAGE,
[mapped_transport_error, *safe_remaining_group.exceptions],
)
http_errors.clear()
del selected_http_error
del http_group
del remaining_group

if base_error_group is not None:
raise base_error_group
assert transport_error is not None
self._raise_mapped_transport_error(transport_error, None)

async def _run_with_retries(self, func: Callable[[], Awaitable[T]]) -> T:
attempts = 0
while True:
Expand Down Expand Up @@ -1079,7 +1171,10 @@ async def list_prompts(
raise UserError("Server not initialized. Make sure you call `connect()` first.")
session = self.session
assert session is not None
return await self._maybe_serialize_request(lambda: session.list_prompts())
return await self._run_request_with_transport_error_redaction(
"list prompts",
lambda: self._maybe_serialize_request(lambda: session.list_prompts()),
)

async def get_prompt(
self, name: str, arguments: dict[str, Any] | None = None
Expand All @@ -1089,15 +1184,21 @@ async def get_prompt(
raise UserError("Server not initialized. Make sure you call `connect()` first.")
session = self.session
assert session is not None
return await self._maybe_serialize_request(lambda: session.get_prompt(name, arguments))
return await self._run_request_with_transport_error_redaction(
"get prompt",
lambda: self._maybe_serialize_request(lambda: session.get_prompt(name, arguments)),
)

async def list_resources(self, cursor: str | None = None) -> ListResourcesResult:
"""List the resources available on the server."""
if not self.session:
raise UserError("Server not initialized. Make sure you call `connect()` first.")
session = self.session
assert session is not None
return await self._maybe_serialize_request(lambda: session.list_resources(cursor))
return await self._run_request_with_transport_error_redaction(
"list resources",
lambda: self._maybe_serialize_request(lambda: session.list_resources(cursor)),
)

async def list_resource_templates(
self, cursor: str | None = None
Expand All @@ -1107,7 +1208,10 @@ async def list_resource_templates(
raise UserError("Server not initialized. Make sure you call `connect()` first.")
session = self.session
assert session is not None
return await self._maybe_serialize_request(lambda: session.list_resource_templates(cursor))
return await self._run_request_with_transport_error_redaction(
"list resource templates",
lambda: self._maybe_serialize_request(lambda: session.list_resource_templates(cursor)),
)

async def read_resource(self, uri: str) -> ReadResourceResult:
"""Read the contents of a specific resource by URI.
Expand All @@ -1122,7 +1226,10 @@ async def read_resource(self, uri: str) -> ReadResourceResult:
assert session is not None
from pydantic import AnyUrl

return await self._maybe_serialize_request(lambda: session.read_resource(AnyUrl(uri)))
return await self._run_request_with_transport_error_redaction(
"read resource",
lambda: self._maybe_serialize_request(lambda: session.read_resource(AnyUrl(uri))),
)

async def cleanup(self):
"""Cleanup the server."""
Expand Down
45 changes: 27 additions & 18 deletions src/agents/voice/result.py
Original file line number Diff line number Diff line change
Expand Up @@ -289,18 +289,25 @@ async def _wait_for_completion(self):
tasks.append(self._dispatcher_task)
await asyncio.gather(*tasks)

def _cleanup_tasks(self):
self._finish_turn()
async def _cleanup_tasks(self):
current_task = asyncio.current_task()
tasks: list[asyncio.Task[Any]] = []
seen: set[asyncio.Task[Any]] = set()
for task in [*self._tasks, self._dispatcher_task, self.text_generation_task]:
if task is None or task is current_task or task in seen:
continue
seen.add(task)
tasks.append(task)

for task in self._tasks:
for task in tasks:
if not task.done():
task.cancel()

if self._dispatcher_task and not self._dispatcher_task.done():
self._dispatcher_task.cancel()

if self.text_generation_task and not self.text_generation_task.done():
self.text_generation_task.cancel()
try:
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
finally:
self._finish_turn()

def _check_errors(self):
for task in self._tasks:
Expand All @@ -316,7 +323,7 @@ async def stream(self) -> AsyncIterator[VoiceStreamEvent]:
try:
event = await self._queue.get()
except asyncio.CancelledError:
self._cleanup_tasks()
await self._cleanup_tasks()
raise
if isinstance(event, VoiceStreamEventError):
self._stored_exception = event.error
Expand All @@ -333,15 +340,17 @@ async def stream(self) -> AsyncIterator[VoiceStreamEvent]:

# On the normal completion path, let the producer task finish gracefully so any active
# trace context can emit `trace_end` before we run cleanup.
if (
saw_session_end
and self.text_generation_task is not None
and not self.text_generation_task.done()
):
await asyncio.shield(self.text_generation_task)

self._check_errors()
self._cleanup_tasks()
try:
if (
saw_session_end
and self.text_generation_task is not None
and not self.text_generation_task.done()
):
await asyncio.shield(self.text_generation_task)

self._check_errors()
finally:
await self._cleanup_tasks()

if self._stored_exception:
raise self._stored_exception
Loading