diff --git a/tests/tools/test_mcp_oauth.py b/tests/tools/test_mcp_oauth.py index 55f92f070acd..8034969013b3 100644 --- a/tests/tools/test_mcp_oauth.py +++ b/tests/tools/test_mcp_oauth.py @@ -561,6 +561,54 @@ class TestCallbackPortReservation: # Reservation was consumed by adoption. assert port not in mod._reserved_sockets + def test_concurrent_flows_keep_their_own_callback_ports(self, monkeypatch): + """#34260: flow A's waiter listens on A's port even after flow B + overwrites the legacy module-level global. + + This is the callback-side sibling of the #44588 redirect-handler fix: + without a per-flow waiter, A's callback wait would bind B's port and + A's redirect (pointing at A's port) would never be received. + """ + import asyncio + import threading + import urllib.request + import tools.mcp_oauth as mod + + monkeypatch.setattr(mod, "_is_interactive", lambda: False) + monkeypatch.setattr(mod, "_raise_if_non_interactive", lambda lead: None) + + cfg_a: dict = {} + port_a = mod._configure_callback_port(cfg_a) + waiter_a = mod._make_callback_waiter(port_a) + # Flow B configures afterwards — overwrites mod._oauth_port. + cfg_b: dict = {} + port_b = mod._configure_callback_port(cfg_b) + assert mod._oauth_port == port_b != port_a + + async def drive(): + task = asyncio.create_task(waiter_a()) + await asyncio.sleep(0.2) + + def hit(): + # The redirect goes to flow A's port — where A's waiter + # must be listening despite the clobbered global. + urllib.request.urlopen( + f"http://127.0.0.1:{port_a}/callback?code=flowA&state=sA", + timeout=5, + ) + + threading.Thread(target=hit, daemon=True).start() + return await asyncio.wait_for(task, timeout=10) + + try: + code, state = asyncio.run(drive()) + finally: + leftover = mod._reserved_sockets.pop(port_b, None) + if leftover is not None: + leftover.close() + assert code == "flowA" + assert state == "sA" + # --------------------------------------------------------------------------- # remove_oauth_tokens diff --git a/tools/mcp_oauth.py b/tools/mcp_oauth.py index b905ca3de77c..cb8917b28e91 100644 --- a/tools/mcp_oauth.py +++ b/tools/mcp_oauth.py @@ -668,21 +668,14 @@ def _make_redirect_handler(port: int, redirect_uri: str | None = None): async def _wait_for_callback() -> tuple[str, str | None]: - """Wait for the OAuth callback to arrive on the local callback server. + """Wait for the OAuth callback on the legacy module-level port. - Uses the module-level ``_oauth_port`` which is set by ``build_oauth_auth`` - before this is ever called. Polls for the result without blocking the - event loop. - - On an interactive TTY, races the HTTP listener against a stdin paste - fallback so users without an SSH tunnel can copy the redirect URL (or - just the ``code=...&state=...`` query string) from a browser on another - machine and paste it back. The HTTP listener wins when the redirect - reaches it first; the paste fallback wins when it doesn't. + Kept for backwards compatibility with callers that never went through + :func:`build_oauth_auth`'s per-flow wiring. New code paths receive a + per-flow waiter from :func:`_make_callback_waiter` so concurrent OAuth + flows cannot cross ports (#34260). Raises: - OAuthNonInteractiveError: If the callback times out (no user present - to complete the browser auth). RuntimeError: If ``_oauth_port`` has not been set, which would indicate that ``build_oauth_auth`` was skipped — the asserting form below was a silent bug when running Python with ``-O``/``-OO``. @@ -692,106 +685,129 @@ async def _wait_for_callback() -> tuple[str, str | None]: "OAuth callback port not set — build_oauth_auth must be called " "before _wait_for_oauth_callback" ) + return await _make_callback_waiter(_oauth_port)() - # Reject before binding the callback listener in non-interactive contexts. - # Reaching here means the SDK entered the authorization-code flow (a valid - # or refreshable token would never call the callback handler), so a cached - # token file is present but unusable. Binding the listener here would block - # for the full 300s timeout and — on the next connection retry — collide - # with the still-bound/TIME_WAIT port, surfacing as - # ``OSError: [Errno 98] Address already in use``. Failing fast keeps - # gateway startup independent of an unusable optional MCP server. This - # guard holds "regardless of whether a token file exists" — the point the - # build_oauth_auth token-file guard cannot cover. See #57836. - _raise_if_non_interactive( - "OAuth callback requires an interactive session but none is " - "available (non-interactive/background context); skipping browser " - "authorization without binding a callback listener." - ) - # The callback server is already running (started in build_oauth_auth). - # We just need to poll for the result. - handler_cls, result = _make_callback_handler() +def _make_callback_waiter(port: int): + """Return a callback waiter bound to a single OAuth flow's port. - # Start a temporary server on the known port, adopting the socket - # reserved at port-selection time when one exists. Holding the bound - # socket from _reserve_callback_port() until here closes the TOCTOU - # window where another process could steal the port between selection - # and bind (#22161). allow_reuse_address is set BEFORE binding (setting - # it after the constructor has already bound is a no-op) so a lingering - # TIME_WAIT socket from a previous flow cannot block the next one - # (#44590). - try: - server = HTTPServer( - ("127.0.0.1", _oauth_port), handler_cls, bind_and_activate=False - ) - reserved = _reserved_sockets.pop(_oauth_port, None) - if reserved is not None: - # Adopt the reserved (already bound) socket and start listening. - server.socket.close() - server.socket = reserved - server.server_address = reserved.getsockname() - server.server_activate() - else: - server.allow_reuse_address = True - server.server_bind() - server.server_activate() - except OSError as exc: - # The loopback callback port is genuinely in use: a concurrent OAuth - # flow, a leftover listener, or a fixed `oauth.redirect_port` that - # collided. build_oauth_auth does not start its own callback server, - # so there is nothing to poll here; surface a clear, actionable error - # instead of a misleading "timed out". - raise OAuthNonInteractiveError( - f"OAuth callback port {_oauth_port} is already in use ({exc}). " - "Close any other in-progress login, or set a free `oauth.redirect_port` " - "in the server config, then retry." - ) from exc + Closing over the port (instead of reading the module-level + ``_oauth_port``) keeps concurrent OAuth flows isolated: flow A's waiter + listens on flow A's port even when flow B's ``_configure_callback_port`` + overwrites the legacy global afterwards (#34260, the callback-side + sibling of the #44588 redirect-handler fix). - server_thread = threading.Thread(target=server.handle_request, daemon=True) - server_thread.start() + The waiter polls for the redirect without blocking the event loop. On an + interactive TTY it races the HTTP listener against a stdin paste fallback + so users without an SSH tunnel can paste the redirect URL (or just the + ``code=...&state=...`` query string) from a browser on another machine. - # Optional paste-fallback thread: only on interactive TTYs. Reads one - # line from stdin and writes the parsed code/state into the shared - # result dict. The HTTP listener and this thread race for the result; - # whichever fills it first wins. - paste_thread: threading.Thread | None = None - if _is_interactive(): - print( - "\n Or paste the redirect URL here (or the ``?code=...&state=...`` " - "portion) and press Enter. Type ``skip`` + Enter to continue " - "without this server:", - file=sys.stderr, - flush=True, - ) - paste_thread = threading.Thread( - target=_paste_callback_reader, args=(result,), daemon=True - ) - paste_thread.start() + Raises (when awaited): + OAuthNonInteractiveError: If the callback times out (no user present + to complete the browser auth), or in non-interactive contexts. + """ - timeout = 300.0 - poll_interval = 0.5 - elapsed = 0.0 - try: - while elapsed < timeout: - if result["auth_code"] is not None or result["error"] is not None: - break - await asyncio.sleep(poll_interval) - elapsed += poll_interval - finally: - server.server_close() - - if result["error"] == _USER_SKIPPED_SENTINEL: - raise OAuthNonInteractiveError("user_skipped") - if result["error"]: - raise RuntimeError(f"OAuth authorization failed: {result['error']}") - if result["auth_code"] is None: - raise OAuthNonInteractiveError( - "OAuth callback timed out — no authorization code received. " - "Ensure you completed the browser authorization flow." + async def _wait() -> tuple[str, str | None]: + # Reject before binding the callback listener in non-interactive + # contexts. Reaching here means the SDK entered the authorization-code + # flow (a valid or refreshable token would never call the callback + # handler), so a cached token file is present but unusable. Binding the + # listener here would block for the full 300s timeout and — on the next + # connection retry — collide with the still-bound/TIME_WAIT port, + # surfacing as ``OSError: [Errno 98] Address already in use``. Failing + # fast keeps gateway startup independent of an unusable optional MCP + # server. This guard holds "regardless of whether a token file exists" + # — the point the build_oauth_auth token-file guard cannot cover. + # See #57836. + _raise_if_non_interactive( + "OAuth callback requires an interactive session but none is " + "available (non-interactive/background context); skipping browser " + "authorization without binding a callback listener." ) - return result["auth_code"], result["state"] + handler_cls, result = _make_callback_handler() + + # Start a temporary server on this flow's port, adopting the socket + # reserved at port-selection time when one exists. Holding the bound + # socket from _reserve_callback_port() until here closes the TOCTOU + # window where another process could steal the port between selection + # and bind (#22161). allow_reuse_address is set BEFORE binding (setting + # it after the constructor has already bound is a no-op) so a lingering + # TIME_WAIT socket from a previous flow cannot block the next one + # (#44590). + try: + server = HTTPServer( + ("127.0.0.1", port), handler_cls, bind_and_activate=False + ) + reserved = _reserved_sockets.pop(port, None) + if reserved is not None: + # Adopt the reserved (already bound) socket and start listening. + server.socket.close() + server.socket = reserved + server.server_address = reserved.getsockname() + server.server_activate() + else: + server.allow_reuse_address = True + server.server_bind() + server.server_activate() + except OSError as exc: + # The loopback callback port is genuinely in use: a concurrent OAuth + # flow, a leftover listener, or a fixed `oauth.redirect_port` that + # collided. build_oauth_auth does not start its own callback server, + # so there is nothing to poll here; surface a clear, actionable error + # instead of a misleading "timed out". + raise OAuthNonInteractiveError( + f"OAuth callback port {port} is already in use ({exc}). " + "Close any other in-progress login, or set a free `oauth.redirect_port` " + "in the server config, then retry." + ) from exc + + server_thread = threading.Thread(target=server.handle_request, daemon=True) + server_thread.start() + + # Optional paste-fallback thread: only on interactive TTYs. Reads one + # line from stdin and writes the parsed code/state into the shared + # result dict. The HTTP listener and this thread race for the result; + # whichever fills it first wins. + paste_thread: threading.Thread | None = None + if _is_interactive(): + print( + "\n Or paste the redirect URL here (or the ``?code=...&state=...`` " + "portion) and press Enter. Type ``skip`` + Enter to continue " + "without this server:", + file=sys.stderr, + flush=True, + ) + paste_thread = threading.Thread( + target=_paste_callback_reader, args=(result,), daemon=True + ) + paste_thread.start() + + timeout = 300.0 + poll_interval = 0.5 + elapsed = 0.0 + try: + while elapsed < timeout: + if result["auth_code"] is not None or result["error"] is not None: + break + await asyncio.sleep(poll_interval) + elapsed += poll_interval + finally: + server.server_close() + + if result["error"] == _USER_SKIPPED_SENTINEL: + raise OAuthNonInteractiveError("user_skipped") + if result["error"]: + raise RuntimeError(f"OAuth authorization failed: {result['error']}") + if result["auth_code"] is None: + raise OAuthNonInteractiveError( + "OAuth callback timed out — no authorization code received. " + "Ensure you completed the browser authorization flow." + ) + + return result["auth_code"], result["state"] + + return _wait def _paste_callback_reader(result: dict) -> None: @@ -1052,17 +1068,18 @@ def build_oauth_auth( client_metadata = _build_client_metadata(cfg) _maybe_preregister_client(storage, cfg, client_metadata) - # Use closure factories to avoid global state pollution (#44588). + # Use closure factories to avoid global state pollution (#44588, #34260). resolved_port = cfg.get("_resolved_port", _oauth_port) redirect_handler = _make_redirect_handler( resolved_port, redirect_uri=cfg.get("redirect_uri") or None ) + callback_handler = _make_callback_waiter(resolved_port) return OAuthClientProvider( server_url=server_url, client_metadata=client_metadata, storage=storage, redirect_handler=redirect_handler, - callback_handler=_wait_for_callback, + callback_handler=callback_handler, timeout=float(cfg.get("timeout", 300)), ) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index 889b7424634b..757dee1b1a8a 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -522,8 +522,8 @@ class MCPOAuthManager: _configure_callback_port, _is_interactive, _maybe_preregister_client, + _make_callback_waiter, _make_redirect_handler, - _wait_for_callback, ) if not _OAUTH_AVAILABLE: @@ -547,6 +547,7 @@ class MCPOAuthManager: resolved_port = cfg.get("_resolved_port", 0) redirect_handler = _make_redirect_handler(resolved_port) + callback_handler = _make_callback_waiter(resolved_port) return _HERMES_PROVIDER_CLS( server_name=server_name, @@ -555,7 +556,7 @@ class MCPOAuthManager: client_metadata=client_metadata, storage=storage, redirect_handler=redirect_handler, - callback_handler=_wait_for_callback, + callback_handler=callback_handler, timeout=float(cfg.get("timeout", 300)), )