mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-14 08:00:10 +00:00
* fix(mcp): reject credentials that cannot travel as HTTP header values
A request-scoped secret or user_auth credential with a trailing newline
(the usual result of reading a token from a file, or a CRLF env-file),
CR/LF, surrounding whitespace, or characters outside Latin-1 sailed
through the credential interceptors into the HTTP client, where httpx/h11
reject it with an exception that echoes the full value:
LocalProtocolError: Illegal header value b'Bearer sk-...\n'
ToolErrorHandlingMiddleware copies that message into a model-visible
ToolMessage, so the secret landed in the prompt, the checkpoint, and
traces - everywhere headers_from_context promises it never goes.
Add illegal_header_value_reason to mcp/headers.py, mirroring the
transport's own rules (Latin-1 encodable; h11's field_vchar is [^\x00\s]
with SP/HTAB legal only between visible characters), and fail closed in
both interceptors before the value can reach the client. The denial names
only the secret key (plus the reason) and never repeats the value.
Illegal values are denied regardless of on_missing: the key is present,
so a passthrough fallback would silently run the call under the shared
discovery credential - the exact authority confusion the deny default
exists to prevent.
Values the transport accepts are not rejected: embedded SP/HTAB
('Bearer <token>'), Latin-1 high bytes, and DEL all still pass, pinned
by tests against h11's observed behaviour.
* fix(mcp): tighten header value validation to httpx's ASCII boundary
The validator mirrored h11's Latin-1 boundary, but the transport rejects
more than h11 does: build_server_params hands dict[str, str] headers
through the MCP SDK's create_mcp_http_client into httpx.AsyncClient, and
httpx (pinned 0.28.1) encodes str header values as ASCII - so a Latin-1
high byte like 'Bearer caf\xe9' passed validation here only to raise
UnicodeEncodeError inside httpx before h11 ever ran, with the exception
message repeating the offending value.
Validate str values against ASCII instead, flip the tests that pinned
Latin-1 high bytes as transportable, and pin the boundary against the
real client: create_mcp_http_client must reject what the validator
flags and construct cleanly for what it accepts (embedded SP/HTAB and
DEL still pass).
Addresses review feedback on the ASCII vs Latin-1 boundary.
* fix(mcp): validate OAuth and static header values at the same boundary
The validator added for headers_from_context and user_auth left two paths
uncovered. A token endpoint returning an access_token or token_type with a
newline reached httpx/h11, which raise with the full token in the message, and
ToolErrorHandlingMiddleware copies that message into a model-visible
ToolMessage -- the leak this PR set out to close. The operator's static headers
had the same hole.
OAuthTokenManager.get_authorization_header now renders the Authorization value
through one checked helper, so the tool interceptor, the initial discovery
headers and the durable task path are all covered by a single guard. The
rendered value is what gets checked rather than the two fields separately,
because that is what the transport sees: an access_token with leading
whitespace is legal once it follows "Bearer ".
build_server_params applies the same check to statically configured headers.
build_servers_config already isolates a per-server failure, so a bad value
drops that one server and logs the reason instead of the value.
* docs(mcp): correct which transport echoes the full header value
The rationale claimed httpx and h11 both render the full value into their
exception message. Only h11 does, on the line break and surrounding whitespace
cases. httpx's ASCII failure is a UnicodeEncodeError naming the offending
character and its position, not the credential, so at most one character
escapes there; refusing the value up front buys an actionable error rather than
an encode failure raised from inside the client.
Corrected in headers.py and in every copy of the claim: context_headers.py,
user_scoped_auth.py, oauth.py, client.py, mcp/AGENTS.md, docs/MCP_SERVER.md,
the frontend mcp.mdx, and the test comments carrying the same wording. No
behavior change.
---------
Co-authored-by: Terminator666666 <Terminator666666@users.noreply.github.com>
741 lines
28 KiB
Python
741 lines
28 KiB
Python
"""Tests for MCP OAuth support."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
import threading
|
|
from typing import Any
|
|
|
|
import pytest
|
|
|
|
from deerflow.config.extensions_config import ExtensionsConfig
|
|
from deerflow.mcp.oauth import OAuthTokenManager, build_oauth_tool_interceptor, get_initial_oauth_headers
|
|
|
|
|
|
class _MockResponse:
|
|
def __init__(self, payload: dict[str, Any]):
|
|
self._payload = payload
|
|
|
|
def raise_for_status(self) -> None:
|
|
return None
|
|
|
|
def json(self) -> dict[str, Any]:
|
|
return self._payload
|
|
|
|
|
|
class _MockAsyncClient:
|
|
def __init__(self, payload: dict[str, Any], post_calls: list[dict[str, Any]], **kwargs):
|
|
self._payload = payload
|
|
self._post_calls = post_calls
|
|
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, exc_type, exc, tb):
|
|
return False
|
|
|
|
async def post(self, url: str, data: dict[str, Any]):
|
|
self._post_calls.append({"url": url, "data": data})
|
|
return _MockResponse(self._payload)
|
|
|
|
|
|
def test_oauth_token_manager_fetches_and_caches_token(monkeypatch):
|
|
post_calls: list[dict[str, Any]] = []
|
|
|
|
def _client_factory(*args, **kwargs):
|
|
return _MockAsyncClient(
|
|
payload={
|
|
"access_token": "token-123",
|
|
"token_type": "Bearer",
|
|
"expires_in": 3600,
|
|
},
|
|
post_calls=post_calls,
|
|
**kwargs,
|
|
)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _client_factory)
|
|
|
|
config = ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"secure-http": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id",
|
|
"client_secret": "client-secret",
|
|
},
|
|
}
|
|
}
|
|
}
|
|
)
|
|
|
|
manager = OAuthTokenManager.from_extensions_config(config)
|
|
|
|
first = asyncio.run(manager.get_authorization_header("secure-http"))
|
|
second = asyncio.run(manager.get_authorization_header("secure-http"))
|
|
|
|
assert first == "Bearer token-123"
|
|
assert second == "Bearer token-123"
|
|
assert len(post_calls) == 1
|
|
assert post_calls[0]["url"] == "https://auth.example.com/oauth/token"
|
|
assert post_calls[0]["data"]["grant_type"] == "client_credentials"
|
|
|
|
|
|
def test_oauth_extra_token_params_cannot_override_grant_type(monkeypatch):
|
|
post_calls: list[dict[str, Any]] = []
|
|
|
|
def _client_factory(*args, **kwargs):
|
|
return _MockAsyncClient(
|
|
payload={
|
|
"access_token": "token-123",
|
|
"token_type": "Bearer",
|
|
"expires_in": 3600,
|
|
},
|
|
post_calls=post_calls,
|
|
**kwargs,
|
|
)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _client_factory)
|
|
|
|
config = ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"secure-http": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id",
|
|
"client_secret": "client-secret",
|
|
# A careless copy-paste from another OAuth config.
|
|
"extra_token_params": {
|
|
"grant_type": "password",
|
|
"resource": "https://api.example.com",
|
|
},
|
|
},
|
|
}
|
|
}
|
|
}
|
|
)
|
|
|
|
manager = OAuthTokenManager.from_extensions_config(config)
|
|
|
|
asyncio.run(manager.get_authorization_header("secure-http"))
|
|
|
|
# The reserved grant_type must win over the operator-supplied param so
|
|
# the value sent to the token endpoint matches the branch logic that
|
|
# picked client_credentials below. Other extension params still pass
|
|
# through unchanged.
|
|
assert post_calls[0]["data"]["grant_type"] == "client_credentials"
|
|
assert post_calls[0]["data"]["resource"] == "https://api.example.com"
|
|
|
|
|
|
def test_build_oauth_interceptor_injects_authorization_header(monkeypatch):
|
|
post_calls: list[dict[str, Any]] = []
|
|
|
|
def _client_factory(*args, **kwargs):
|
|
return _MockAsyncClient(
|
|
payload={
|
|
"access_token": "token-abc",
|
|
"token_type": "Bearer",
|
|
"expires_in": 3600,
|
|
},
|
|
post_calls=post_calls,
|
|
**kwargs,
|
|
)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _client_factory)
|
|
|
|
config = ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"secure-sse": {
|
|
"enabled": True,
|
|
"type": "sse",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id",
|
|
"client_secret": "client-secret",
|
|
},
|
|
}
|
|
}
|
|
}
|
|
)
|
|
|
|
interceptor = build_oauth_tool_interceptor(config)
|
|
assert interceptor is not None
|
|
|
|
class _Request:
|
|
def __init__(self):
|
|
self.server_name = "secure-sse"
|
|
self.headers = {"X-Test": "1"}
|
|
|
|
def override(self, **kwargs):
|
|
updated = _Request()
|
|
updated.server_name = self.server_name
|
|
updated.headers = kwargs.get("headers")
|
|
return updated
|
|
|
|
captured: dict[str, Any] = {}
|
|
|
|
async def _handler(request):
|
|
captured["headers"] = request.headers
|
|
return "ok"
|
|
|
|
result = asyncio.run(interceptor(_Request(), _handler))
|
|
|
|
assert result == "ok"
|
|
assert captured["headers"]["Authorization"] == "Bearer token-abc"
|
|
assert captured["headers"]["X-Test"] == "1"
|
|
|
|
|
|
def test_get_initial_oauth_headers(monkeypatch):
|
|
post_calls: list[dict[str, Any]] = []
|
|
|
|
def _client_factory(*args, **kwargs):
|
|
return _MockAsyncClient(
|
|
payload={
|
|
"access_token": "token-initial",
|
|
"token_type": "Bearer",
|
|
"expires_in": 3600,
|
|
},
|
|
post_calls=post_calls,
|
|
**kwargs,
|
|
)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _client_factory)
|
|
|
|
config = ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"secure-http": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id",
|
|
"client_secret": "client-secret",
|
|
},
|
|
},
|
|
"no-oauth": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://example.com/mcp",
|
|
},
|
|
}
|
|
}
|
|
)
|
|
|
|
headers = asyncio.run(get_initial_oauth_headers(config))
|
|
|
|
assert headers == {"secure-http": "Bearer token-initial"}
|
|
assert len(post_calls) == 1
|
|
|
|
|
|
def test_get_initial_oauth_headers_one_failing_server_does_not_drop_others(monkeypatch):
|
|
"""A single OAuth server whose token endpoint fails must not drop headers
|
|
(and therefore tools) from healthy servers."""
|
|
|
|
class _FailingClient:
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, exc_type, exc, tb):
|
|
return False
|
|
|
|
async def post(self, url: str, data: dict[str, Any]):
|
|
raise RuntimeError("token endpoint unreachable")
|
|
|
|
class _OkClient:
|
|
def __init__(self, post_calls: list[dict[str, Any]], **kwargs):
|
|
self._post_calls = post_calls
|
|
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, exc_type, exc, tb):
|
|
return False
|
|
|
|
async def post(self, url: str, data: dict[str, Any]):
|
|
self._post_calls.append({"url": url, "data": data})
|
|
return _MockResponse(
|
|
payload={
|
|
"access_token": "token-ok",
|
|
"token_type": "Bearer",
|
|
"expires_in": 3600,
|
|
}
|
|
)
|
|
|
|
ok_post_calls: list[dict[str, Any]] = []
|
|
|
|
def _client_factory(**kwargs):
|
|
# The first call is for the failing server, second for the healthy one,
|
|
# because OAuthTokenManager iterates _oauth_by_server in dict order
|
|
# ('broken-http' < 'secure-http').
|
|
if not hasattr(_client_factory, "_count"):
|
|
_client_factory._count = 0 # type: ignore[attr-defined]
|
|
_client_factory._count += 1 # type: ignore[attr-defined]
|
|
if _client_factory._count == 1: # type: ignore[attr-defined]
|
|
return _FailingClient()
|
|
return _OkClient(post_calls=ok_post_calls)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _client_factory)
|
|
|
|
config = ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"broken-http": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://broken.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.broken.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id",
|
|
"client_secret": "client-secret",
|
|
},
|
|
},
|
|
"secure-http": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id-2",
|
|
"client_secret": "client-secret-2",
|
|
},
|
|
},
|
|
}
|
|
}
|
|
)
|
|
|
|
headers = asyncio.run(get_initial_oauth_headers(config))
|
|
|
|
# The healthy server's header must still be present.
|
|
assert headers == {"secure-http": "Bearer token-ok"}
|
|
assert len(ok_post_calls) == 1
|
|
|
|
|
|
def test_oauth_refresh_token_rotation_persists_rotated_value(monkeypatch):
|
|
"""When a provider rotates the refresh_token, _fetch_token must capture
|
|
the new value so the next refresh uses it instead of the stale original."""
|
|
post_calls: list[dict[str, Any]] = []
|
|
|
|
def _client_factory(*args, **kwargs):
|
|
return _MockAsyncClient(
|
|
payload={
|
|
"access_token": "at-1",
|
|
"token_type": "Bearer",
|
|
"expires_in": 3600,
|
|
"refresh_token": "rt-rotated-1",
|
|
},
|
|
post_calls=post_calls,
|
|
**kwargs,
|
|
)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _client_factory)
|
|
|
|
config = ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"rotating-srv": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "refresh_token",
|
|
"refresh_token": "rt-original-seed",
|
|
},
|
|
}
|
|
}
|
|
}
|
|
)
|
|
|
|
manager = OAuthTokenManager.from_extensions_config(config)
|
|
|
|
# Force the _is_expiring check to always return True so we hit _fetch_token.
|
|
monkeypatch.setattr(OAuthTokenManager, "_is_expiring", lambda self, token, oauth: True)
|
|
|
|
first = asyncio.run(manager.get_authorization_header("rotating-srv"))
|
|
assert first == "Bearer at-1"
|
|
assert len(post_calls) == 1
|
|
# First call posted the original seed token.
|
|
assert post_calls[0]["data"]["refresh_token"] == "rt-original-seed"
|
|
|
|
# On the second call, the rotated refresh_token from the first response
|
|
# must be used.
|
|
second = asyncio.run(manager.get_authorization_header("rotating-srv"))
|
|
assert second == "Bearer at-1"
|
|
assert len(post_calls) == 2
|
|
assert post_calls[1]["data"]["refresh_token"] == "rt-rotated-1"
|
|
|
|
|
|
def test_get_authorization_header_concurrent_threads_no_deadlock(monkeypatch):
|
|
"""Concurrent callers on different event loops/threads must not deadlock.
|
|
|
|
The embedded/TUI sync tool-call path (``DeerFlowClient.stream()`` ->
|
|
LangGraph's ``ToolNode._func`` -> a ``ThreadPoolExecutor`` ->
|
|
``deerflow.tools.sync.make_sync_tool_wrapper``'s per-call ``asyncio.run()``)
|
|
invokes ``get_authorization_header`` from a fresh event loop on a fresh OS
|
|
thread for every concurrent tool call. A per-server ``asyncio.Lock`` binds
|
|
to whichever loop first contends on it; when a caller on a *different*
|
|
loop later releases/wakes a waiter, it does so without
|
|
``call_soon_threadsafe``, so the waiting loop's selector is never woken
|
|
and that caller hangs forever with no exception (a silent hang). A third
|
|
concurrent caller instead hits a synchronous ``RuntimeError: ... is bound
|
|
to a different event loop``. Both failure modes are reproducible with the
|
|
old ``asyncio.Lock``-per-server implementation.
|
|
|
|
This test uses a bounded thread-join timeout so that a regression back to
|
|
the old behavior fails this test quickly instead of hanging the whole
|
|
suite.
|
|
"""
|
|
post_calls: list[dict[str, Any]] = []
|
|
post_calls_guard = threading.Lock()
|
|
holder_in_critical_section = threading.Event()
|
|
|
|
class _SlowMockAsyncClient:
|
|
def __init__(self, **kwargs):
|
|
pass
|
|
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, exc_type, exc, tb):
|
|
return False
|
|
|
|
async def post(self, url: str, data: dict[str, Any]):
|
|
with post_calls_guard:
|
|
post_calls.append({"url": url, "data": data})
|
|
# Signal that this call is inside the critical section (the lock
|
|
# is held) and stay there briefly so the other threads have time
|
|
# to reach their own acquire() and genuinely contend, rather than
|
|
# racing to also take an uncontended fast path.
|
|
holder_in_critical_section.set()
|
|
await asyncio.sleep(0.3)
|
|
return _MockResponse(
|
|
{
|
|
"access_token": "concurrent-token",
|
|
"token_type": "Bearer",
|
|
"expires_in": 3600,
|
|
}
|
|
)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _SlowMockAsyncClient)
|
|
|
|
config = ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"secure-http": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id",
|
|
"client_secret": "client-secret",
|
|
},
|
|
}
|
|
}
|
|
}
|
|
)
|
|
|
|
manager = OAuthTokenManager.from_extensions_config(config)
|
|
results: dict[str, Any] = {}
|
|
|
|
def run_in_own_loop(name: str, wait_for_holder: bool) -> None:
|
|
if wait_for_holder:
|
|
# Only start once another thread is confirmed to be holding the
|
|
# lock, guaranteeing this call contends instead of racing for
|
|
# the uncontended fast path itself.
|
|
assert holder_in_critical_section.wait(timeout=5), "holder thread never entered critical section"
|
|
try:
|
|
results[name] = asyncio.run(manager.get_authorization_header("secure-http"))
|
|
except BaseException as exc: # noqa: BLE001 - captured to assert absence below
|
|
results[name] = exc
|
|
|
|
threads = [
|
|
threading.Thread(target=run_in_own_loop, args=("holder", False), name="holder", daemon=True),
|
|
threading.Thread(target=run_in_own_loop, args=("waiter-1", True), name="waiter-1", daemon=True),
|
|
threading.Thread(target=run_in_own_loop, args=("waiter-2", True), name="waiter-2", daemon=True),
|
|
]
|
|
|
|
for t in threads:
|
|
t.start()
|
|
|
|
# Bounded timeout: under the old per-server asyncio.Lock, at least one of
|
|
# these threads would never return. Joining with a timeout keeps a
|
|
# regression from hanging the test suite forever; it fails fast instead.
|
|
for t in threads:
|
|
t.join(timeout=5)
|
|
|
|
still_alive = [t.name for t in threads if t.is_alive()]
|
|
assert not still_alive, f"deadlock: thread(s) still blocked after bounded timeout: {still_alive}"
|
|
|
|
for name, result in results.items():
|
|
assert not isinstance(result, BaseException), f"{name} raised instead of completing: {result!r}"
|
|
assert result == "Bearer concurrent-token"
|
|
|
|
# De-duplication must be preserved: three concurrent callers racing for
|
|
# the same (initially uncached) server must still only perform ONE real
|
|
# token fetch, not one per caller.
|
|
assert len(post_calls) == 1
|
|
|
|
|
|
def test_get_authorization_header_cancelled_while_waiting_does_not_leak_lock(monkeypatch):
|
|
"""A caller cancelled while waiting on the per-server lock must not leak it.
|
|
|
|
``get_authorization_header`` runs ``lock.acquire()`` on a real OS thread via
|
|
``asyncio.to_thread`` so a blocking wait never blocks the event loop. Once that
|
|
thread has actually started running ``lock.acquire()``, Python cannot interrupt
|
|
it: cancelling the *caller* only stops the caller from continuing, it does not
|
|
stop the thread. If cancellation at that await let the thread go on to acquire
|
|
the lock unobserved (nobody left holding a reference that will call
|
|
``release()`` for it), the lock would stay held forever and every subsequent
|
|
call for this server would block permanently at the same line -- the very
|
|
cross-thread deadlock this file's lock was introduced to fix, reintroduced via
|
|
a different path.
|
|
|
|
This test holds the per-server lock (simulating another in-flight caller),
|
|
starts a second caller that has to wait for it, cancels that waiter while it
|
|
is genuinely blocked in its executor thread, releases the original holder, and
|
|
then asserts a third caller completes within a bounded timeout and performs
|
|
exactly one token fetch. Every potentially-hanging await is wrapped in a
|
|
bounded timeout so a regression fails this test quickly instead of hanging the
|
|
suite.
|
|
"""
|
|
post_calls: list[dict[str, Any]] = []
|
|
|
|
def _client_factory(*args, **kwargs):
|
|
return _MockAsyncClient(
|
|
payload={
|
|
"access_token": "after-cancel-token",
|
|
"token_type": "Bearer",
|
|
"expires_in": 3600,
|
|
},
|
|
post_calls=post_calls,
|
|
**kwargs,
|
|
)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _client_factory)
|
|
|
|
config = ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"secure-http": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id",
|
|
"client_secret": "client-secret",
|
|
},
|
|
}
|
|
}
|
|
}
|
|
)
|
|
|
|
manager = OAuthTokenManager.from_extensions_config(config)
|
|
lock = manager._locks["secure-http"]
|
|
|
|
async def scenario() -> None:
|
|
# Simulate another in-flight caller already holding the per-server lock
|
|
# (uncontended, so this succeeds immediately without blocking).
|
|
lock.acquire()
|
|
try:
|
|
waiter = asyncio.create_task(manager.get_authorization_header("secure-http"))
|
|
|
|
# Let the waiter's asyncio.to_thread(lock.acquire) actually get
|
|
# scheduled onto an executor thread and start genuinely blocking on
|
|
# the real lock before cancelling it -- otherwise the cancellation
|
|
# could land before the thread even starts, which would not exercise
|
|
# the bug.
|
|
await asyncio.sleep(0.2)
|
|
|
|
waiter.cancel()
|
|
# The original holder finishes its own work and releases *before* we
|
|
# wait on the cancelled waiter: a correct fix must keep the lock's
|
|
# eventual acquisition shielded from this coroutine's cancellation and
|
|
# wait for it to actually land before releasing, so awaiting the
|
|
# cancelled waiter can legitimately block until the lock is free
|
|
# either way.
|
|
lock.release()
|
|
|
|
with pytest.raises(asyncio.CancelledError):
|
|
await asyncio.wait_for(waiter, timeout=5)
|
|
|
|
# The crux of the regression: under the bug, the waiter's abandoned
|
|
# executor thread went on to acquire the lock with nobody left to
|
|
# release it, so this third call would block forever. Bound it so a
|
|
# regression fails fast instead of hanging the test itself.
|
|
third = await asyncio.wait_for(manager.get_authorization_header("secure-http"), timeout=5)
|
|
assert third == "Bearer after-cancel-token"
|
|
finally:
|
|
# Test-only safety net, independent of the assertions above: under
|
|
# the bug, the lock is left permanently locked with a background
|
|
# thread (from whichever caller's orphaned acquisition landed last)
|
|
# still parked on a *subsequent* acquire() that will now never
|
|
# return. asyncio.run()'s own teardown joins every thread the
|
|
# default executor ever created before it returns, so leaving that
|
|
# thread stuck would hang this test process at interpreter/loop
|
|
# shutdown even after the failure above is already reported. Forcing
|
|
# the lock open here lets any such thread finish so the process can
|
|
# exit; it is a no-op once the fix keeps the lock correctly balanced.
|
|
if lock.locked():
|
|
lock.release()
|
|
|
|
asyncio.run(scenario())
|
|
|
|
# Exactly one real token fetch: the cancelled waiter must never reach
|
|
# _fetch_token, so the third call is the only one that performs it.
|
|
assert len(post_calls) == 1
|
|
|
|
|
|
# --- Illegal header values ---------------------------------------------------
|
|
#
|
|
# What the token endpoint returns is not this process's to control. An
|
|
# access_token or token_type carrying a newline reaches h11, which raises with
|
|
# the full value in the message, and ToolErrorHandlingMiddleware copies that
|
|
# message into a model-visible ToolMessage. Every one of these asserts the
|
|
# token never appears in what the caller sees.
|
|
|
|
|
|
def _oauth_server_config() -> ExtensionsConfig:
|
|
return ExtensionsConfig.model_validate(
|
|
{
|
|
"mcpServers": {
|
|
"secure-http": {
|
|
"enabled": True,
|
|
"type": "http",
|
|
"url": "https://api.example.com/mcp",
|
|
"oauth": {
|
|
"enabled": True,
|
|
"token_url": "https://auth.example.com/oauth/token",
|
|
"grant_type": "client_credentials",
|
|
"client_id": "client-id",
|
|
"client_secret": "client-secret",
|
|
},
|
|
}
|
|
}
|
|
}
|
|
)
|
|
|
|
|
|
def _token_endpoint_returns(monkeypatch, payload: dict[str, Any]) -> list[dict[str, Any]]:
|
|
post_calls: list[dict[str, Any]] = []
|
|
|
|
def _client_factory(*args, **kwargs):
|
|
return _MockAsyncClient(payload=payload, post_calls=post_calls, **kwargs)
|
|
|
|
monkeypatch.setattr("httpx.AsyncClient", _client_factory)
|
|
return post_calls
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("payload", "secret"),
|
|
[
|
|
(
|
|
{"access_token": "oauth-secret-123\n", "token_type": "Bearer", "expires_in": 3600},
|
|
"oauth-secret-123",
|
|
),
|
|
(
|
|
{"access_token": "oauth-secret-456", "token_type": "Bearer\r", "expires_in": 3600},
|
|
"oauth-secret-456",
|
|
),
|
|
(
|
|
{"access_token": "oauth-secret-caf\u00e9", "token_type": "Bearer", "expires_in": 3600},
|
|
"oauth-secret-caf\u00e9",
|
|
),
|
|
(
|
|
{"access_token": "oauth-secret-789 ", "token_type": "Bearer", "expires_in": 3600},
|
|
"oauth-secret-789",
|
|
),
|
|
],
|
|
ids=["trailing-newline", "cr-in-token-type", "non-ascii", "trailing-space"],
|
|
)
|
|
def test_illegal_oauth_token_is_denied_without_leaking(monkeypatch, payload, secret):
|
|
_token_endpoint_returns(monkeypatch, payload)
|
|
manager = OAuthTokenManager.from_extensions_config(_oauth_server_config())
|
|
|
|
with pytest.raises(ValueError) as excinfo:
|
|
asyncio.run(manager.get_authorization_header("secure-http"))
|
|
|
|
message = str(excinfo.value)
|
|
assert secret not in message
|
|
assert "secure-http" in message
|
|
|
|
|
|
def test_oauth_token_kept_legal_by_the_space_after_token_type_is_accepted(monkeypatch):
|
|
# The rendered value is the boundary, not the two fields on their own: this
|
|
# access_token carries leading whitespace, which the transport tolerates
|
|
# once it follows "Bearer ". Denying it would refuse a token the server
|
|
# would have accepted.
|
|
_token_endpoint_returns(monkeypatch, {"access_token": " leading-space-token", "token_type": "Bearer", "expires_in": 3600})
|
|
manager = OAuthTokenManager.from_extensions_config(_oauth_server_config())
|
|
|
|
assert asyncio.run(manager.get_authorization_header("secure-http")) == "Bearer leading-space-token"
|
|
|
|
|
|
def test_oauth_interceptor_denies_illegal_token_without_calling_the_handler(monkeypatch):
|
|
_token_endpoint_returns(monkeypatch, {"access_token": "oauth-secret-abc\n", "token_type": "Bearer", "expires_in": 3600})
|
|
config = _oauth_server_config()
|
|
interceptor = build_oauth_tool_interceptor(config)
|
|
assert interceptor is not None
|
|
|
|
class _Request:
|
|
server_name = "secure-http"
|
|
headers: dict[str, str] = {}
|
|
|
|
def override(self, **kwargs): # pragma: no cover - denied before reached
|
|
raise AssertionError("the request must never be forwarded with an illegal token")
|
|
|
|
handler_calls: list[Any] = []
|
|
|
|
async def _handler(request): # pragma: no cover - denied before reached
|
|
handler_calls.append(request)
|
|
return "ok"
|
|
|
|
with pytest.raises(ValueError) as excinfo:
|
|
asyncio.run(interceptor(_Request(), _handler))
|
|
|
|
assert "oauth-secret-abc" not in str(excinfo.value)
|
|
assert handler_calls == []
|
|
|
|
|
|
def test_initial_oauth_headers_skips_server_with_illegal_token(monkeypatch, caplog):
|
|
_token_endpoint_returns(monkeypatch, {"access_token": "oauth-secret-xyz\n", "token_type": "Bearer", "expires_in": 3600})
|
|
config = _oauth_server_config()
|
|
|
|
with caplog.at_level(logging.WARNING, logger="deerflow.mcp.oauth"):
|
|
headers = asyncio.run(get_initial_oauth_headers(config))
|
|
|
|
# No header at all rather than a broken one: the connection then fails
|
|
# authentication at the server, which says nothing about the token.
|
|
assert headers == {}
|
|
assert "oauth-secret-xyz" not in caplog.text
|