mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-19 02:26:12 +03:00
feat(mcp): add browser OAuth for remote servers (#5316)
This commit is contained in:
@@ -0,0 +1,283 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.config.schema import MCPServerConfig
|
||||
from nanobot.webui.mcp_oauth_api import (
|
||||
McpOAuthError,
|
||||
McpOAuthManager,
|
||||
prepare_mcp_oauth_redirect_uri,
|
||||
validate_mcp_oauth_redirect_uri,
|
||||
)
|
||||
|
||||
|
||||
class _Connection:
|
||||
def __init__(self) -> None:
|
||||
self.closed = False
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
def _config() -> MCPServerConfig:
|
||||
return MCPServerConfig(
|
||||
type="streamableHttp",
|
||||
auth="oauth",
|
||||
url="https://mcp.example.com/mcp",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_browser_flow_retries_current_server_and_ignores_unrelated_reload_failure(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
manager = McpOAuthManager()
|
||||
connection = _Connection()
|
||||
received: dict[str, object] = {}
|
||||
reload_calls = 0
|
||||
|
||||
monkeypatch.setattr(
|
||||
"nanobot.webui.mcp_oauth_api.validate_url_target",
|
||||
lambda _url: (True, ""),
|
||||
)
|
||||
|
||||
async def connect(servers, _registry, *, oauth_handlers):
|
||||
assert set(servers) == {"xmind"}
|
||||
handlers = oauth_handlers["xmind"]
|
||||
await handlers.redirect_handler(
|
||||
"https://accounts.example.com/authorize?client_id=test&state=state-123"
|
||||
)
|
||||
received["callback"] = await handlers.callback_handler()
|
||||
return {"xmind": connection}
|
||||
|
||||
async def reload_mcp() -> dict[str, object]:
|
||||
nonlocal reload_calls
|
||||
reload_calls += 1
|
||||
if reload_calls == 1:
|
||||
return {
|
||||
"ok": False,
|
||||
"requires_restart": False,
|
||||
"failed": ["xmind"],
|
||||
}
|
||||
return {
|
||||
"ok": False,
|
||||
"requires_restart": False,
|
||||
"connected": ["xmind"],
|
||||
"failed": ["notion"],
|
||||
"message": "MCP config reloaded, but some servers did not connect: notion",
|
||||
}
|
||||
|
||||
monkeypatch.setattr("nanobot.webui.mcp_oauth_api.connect_mcp_servers", connect)
|
||||
|
||||
started = await manager.start(
|
||||
"xmind",
|
||||
_config(),
|
||||
"https://agent.example.com/auth/mcp/callback",
|
||||
reload_mcp=reload_mcp,
|
||||
)
|
||||
|
||||
assert started["status"] == "authorization_required"
|
||||
assert started["authorization_url"].startswith("https://accounts.example.com/authorize?")
|
||||
manager.submit_callback(state="state-123", code="oauth-code", error=None)
|
||||
with pytest.raises(McpOAuthError, match="expired"):
|
||||
manager.submit_callback(state="state-123", code="replayed-code", error=None)
|
||||
|
||||
for _ in range(10):
|
||||
await asyncio.sleep(0)
|
||||
if reload_calls == 2:
|
||||
break
|
||||
assert reload_calls == 2
|
||||
|
||||
for _ in range(10):
|
||||
first, second = await asyncio.gather(
|
||||
manager.status(started["flow_id"]),
|
||||
manager.status(started["flow_id"]),
|
||||
)
|
||||
if first["status"] == "connected":
|
||||
break
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert first["status"] == "connected"
|
||||
assert second["status"] == "connected"
|
||||
assert first["hot_reload"]["failed"] == ["notion"]
|
||||
assert received["callback"] == ("oauth-code", "state-123")
|
||||
assert reload_calls == 2
|
||||
assert connection.closed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remote_http_flow_accepts_a_pasted_loopback_callback(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
manager = McpOAuthManager()
|
||||
connection = _Connection()
|
||||
received: dict[str, object] = {}
|
||||
|
||||
monkeypatch.setattr(
|
||||
"nanobot.webui.mcp_oauth_api.validate_url_target",
|
||||
lambda _url: (True, ""),
|
||||
)
|
||||
|
||||
async def connect(_servers, _registry, *, oauth_handlers):
|
||||
handlers = oauth_handlers["linear"]
|
||||
received["redirect_uri"] = handlers.redirect_uri
|
||||
await handlers.redirect_handler(
|
||||
"https://accounts.example.com/authorize?client_id=test&state=manual-state"
|
||||
)
|
||||
received["callback"] = await handlers.callback_handler()
|
||||
return {"linear": connection}
|
||||
|
||||
monkeypatch.setattr("nanobot.webui.mcp_oauth_api.connect_mcp_servers", connect)
|
||||
|
||||
started = await manager.start(
|
||||
"linear",
|
||||
_config(),
|
||||
"http://192.0.2.10:8765/auth/mcp/callback",
|
||||
reload_mcp=lambda: asyncio.sleep(
|
||||
0,
|
||||
result={"ok": True, "requires_restart": False},
|
||||
),
|
||||
)
|
||||
|
||||
assert started["status"] == "authorization_required"
|
||||
assert started["completion_input"] == "callback_url"
|
||||
assert received["redirect_uri"] == "http://127.0.0.1:8765/auth/mcp/callback"
|
||||
|
||||
with pytest.raises(McpOAuthError, match="complete callback URL"):
|
||||
manager.submit_callback_url(
|
||||
flow_id=started["flow_id"],
|
||||
callback_url=(
|
||||
"http://127.0.0.1:8765/wrong?code=oauth-code&state=manual-state"
|
||||
),
|
||||
)
|
||||
with pytest.raises(McpOAuthError, match="different or expired"):
|
||||
manager.submit_callback_url(
|
||||
flow_id=started["flow_id"],
|
||||
callback_url=(
|
||||
"http://127.0.0.1:8765/auth/mcp/callback"
|
||||
"?code=oauth-code&state=other-state"
|
||||
),
|
||||
)
|
||||
|
||||
submitted = manager.submit_callback_url(
|
||||
flow_id=started["flow_id"],
|
||||
callback_url=(
|
||||
"http://127.0.0.1:8765/auth/mcp/callback"
|
||||
"?code=oauth-code&state=manual-state"
|
||||
),
|
||||
)
|
||||
assert submitted["status"] == "connecting"
|
||||
|
||||
for _ in range(20):
|
||||
await asyncio.sleep(0)
|
||||
result = await manager.status(started["flow_id"])
|
||||
if result["status"] == "connected":
|
||||
break
|
||||
|
||||
assert result["status"] == "connected"
|
||||
assert result["completion_input"] == "callback_url"
|
||||
assert received["callback"] == ("oauth-code", "manual-state")
|
||||
assert connection.closed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_browser_flow_surfaces_provider_denial_without_callback_description(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
manager = McpOAuthManager()
|
||||
monkeypatch.setattr(
|
||||
"nanobot.webui.mcp_oauth_api.validate_url_target",
|
||||
lambda _url: (True, ""),
|
||||
)
|
||||
|
||||
async def connect(_servers, _registry, *, oauth_handlers):
|
||||
handlers = oauth_handlers["notion"]
|
||||
await handlers.redirect_handler("https://accounts.example.com/auth?state=deny-state")
|
||||
await handlers.callback_handler()
|
||||
return {}
|
||||
|
||||
monkeypatch.setattr("nanobot.webui.mcp_oauth_api.connect_mcp_servers", connect)
|
||||
started = await manager.start(
|
||||
"notion",
|
||||
_config(),
|
||||
"https://agent.example.com/auth/mcp/callback",
|
||||
reload_mcp=lambda: asyncio.sleep(0, result={"ok": True}),
|
||||
)
|
||||
|
||||
with pytest.raises(McpOAuthError, match="access_denied"):
|
||||
manager.submit_callback(state="deny-state", code=None, error="access_denied")
|
||||
for _ in range(10):
|
||||
await asyncio.sleep(0)
|
||||
result = await manager.status(started["flow_id"])
|
||||
if result["status"] == "failed":
|
||||
break
|
||||
|
||||
assert result["status"] == "failed"
|
||||
assert result["error"] == "Authorization was not completed (access_denied)."
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("authorization_url", "url_is_safe", "state"),
|
||||
[
|
||||
("https://127.0.0.1/authorize?state=private-state", False, "private-state"),
|
||||
("http://accounts.example.com/authorize?state=http-state", True, "http-state"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_browser_flow_blocks_unsafe_authorization_url(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
authorization_url: str,
|
||||
url_is_safe: bool,
|
||||
state: str,
|
||||
) -> None:
|
||||
manager = McpOAuthManager()
|
||||
monkeypatch.setattr(
|
||||
"nanobot.webui.mcp_oauth_api.validate_url_target",
|
||||
lambda _url: (url_is_safe, "private address"),
|
||||
)
|
||||
|
||||
async def connect(_servers, _registry, *, oauth_handlers):
|
||||
await oauth_handlers["linear"].redirect_handler(authorization_url)
|
||||
return {}
|
||||
|
||||
monkeypatch.setattr("nanobot.webui.mcp_oauth_api.connect_mcp_servers", connect)
|
||||
|
||||
result = await manager.start(
|
||||
"linear",
|
||||
_config(),
|
||||
"https://agent.example.com/auth/mcp/callback",
|
||||
reload_mcp=lambda: asyncio.sleep(0, result={"ok": True}),
|
||||
)
|
||||
|
||||
assert result["status"] == "failed"
|
||||
assert result["error"] == "The MCP server returned an unsafe authorization URL."
|
||||
with pytest.raises(McpOAuthError, match="expired"):
|
||||
manager.submit_callback(state=state, code="code", error=None)
|
||||
|
||||
|
||||
def test_redirect_uri_requires_https_except_for_loopback() -> None:
|
||||
assert validate_mcp_oauth_redirect_uri(
|
||||
"https://agent.example.com/auth/mcp/callback"
|
||||
) == "https://agent.example.com/auth/mcp/callback"
|
||||
assert validate_mcp_oauth_redirect_uri(
|
||||
"http://127.0.0.1:8765/auth/mcp/callback"
|
||||
) == "http://127.0.0.1:8765/auth/mcp/callback"
|
||||
|
||||
with pytest.raises(McpOAuthError, match="HTTPS or localhost"):
|
||||
validate_mcp_oauth_redirect_uri("http://192.0.2.10/auth/mcp/callback")
|
||||
with pytest.raises(McpOAuthError, match="Invalid"):
|
||||
validate_mcp_oauth_redirect_uri("https://agent.example.com/wrong")
|
||||
|
||||
|
||||
def test_remote_http_redirect_prepares_a_manual_loopback_callback() -> None:
|
||||
assert prepare_mcp_oauth_redirect_uri(
|
||||
"https://agent.example.com/auth/mcp/callback"
|
||||
) == ("https://agent.example.com/auth/mcp/callback", False)
|
||||
assert prepare_mcp_oauth_redirect_uri(
|
||||
"http://127.0.0.1:8765/auth/mcp/callback"
|
||||
) == ("http://127.0.0.1:8765/auth/mcp/callback", False)
|
||||
assert prepare_mcp_oauth_redirect_uri(
|
||||
"http://agent.example.com:9443/auth/mcp/callback"
|
||||
) == ("http://127.0.0.1:9443/auth/mcp/callback", True)
|
||||
@@ -3,7 +3,9 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
from mcp.shared.auth import OAuthToken
|
||||
|
||||
from nanobot.agent.tools.mcp_oauth import MCPOAuthStorage, mcp_oauth_has_credentials
|
||||
from nanobot.config.loader import load_config
|
||||
from nanobot.webui.mcp_presets_api import (
|
||||
McpPresetError,
|
||||
@@ -38,6 +40,9 @@ def test_mcp_presets_payload_lists_supported_cards(tmp_path, monkeypatch: pytest
|
||||
"aws-docs",
|
||||
"brave-search",
|
||||
"postman",
|
||||
"xmind",
|
||||
"notion",
|
||||
"linear",
|
||||
}.issubset(names)
|
||||
browserbase = next(preset for preset in payload["presets"] if preset["name"] == "browserbase")
|
||||
assert browserbase["installed"] is False
|
||||
@@ -55,6 +60,37 @@ def test_mcp_presets_payload_lists_supported_cards(tmp_path, monkeypatch: pytest
|
||||
assert manifest["trust"]["review_status"] == "builtin_preset"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_preset_is_one_click_configured_after_token_storage(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_use_config(tmp_path, monkeypatch)
|
||||
|
||||
payload = mcp_presets_action("enable", {"name": ["xmind"]})
|
||||
|
||||
row = next(item for item in payload["presets"] if item["name"] == "xmind")
|
||||
assert row["installed"] is True
|
||||
assert row["configured"] is False
|
||||
assert row["status"] == "authorization_required"
|
||||
assert row["transport"] == "streamableHttp"
|
||||
assert row["auth"] == "oauth"
|
||||
config = load_config()
|
||||
cfg = config.tools.mcp_servers["xmind"]
|
||||
assert cfg.type == "streamableHttp"
|
||||
assert cfg.auth == "oauth"
|
||||
assert cfg.url == "https://app.xmind.com/api/mcp"
|
||||
|
||||
await MCPOAuthStorage("xmind", cfg.url).set_tokens(OAuthToken(access_token="secret"))
|
||||
connected = mcp_presets_payload()
|
||||
row = next(item for item in connected["presets"] if item["name"] == "xmind")
|
||||
assert row["configured"] is True
|
||||
assert row["status"] == "configured"
|
||||
|
||||
mcp_presets_action("remove", {"name": ["xmind"]})
|
||||
assert await MCPOAuthStorage("xmind", cfg.url).get_tokens() is None
|
||||
|
||||
|
||||
def test_enable_browserbase_writes_scrubbed_config_payload(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
@@ -296,11 +332,11 @@ def test_test_mcp_preset_scrubs_connection_errors(
|
||||
assert "<redacted>" in payload["last_action"]["error"]
|
||||
|
||||
|
||||
def test_unlisted_oauth_placeholder_is_not_enabled(tmp_path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_unknown_oauth_placeholder_is_not_enabled(tmp_path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_use_config(tmp_path, monkeypatch)
|
||||
|
||||
with pytest.raises(McpPresetError) as exc:
|
||||
mcp_presets_action("enable", {"name": ["linear"]})
|
||||
mcp_presets_action("enable", {"name": ["asana"]})
|
||||
|
||||
assert exc.value.status == 404
|
||||
|
||||
@@ -414,6 +450,72 @@ def test_import_mcp_config_and_tool_allowlist(
|
||||
assert load_config().tools.mcp_servers["docs"].enabled_tools == []
|
||||
|
||||
|
||||
def test_import_recognizes_known_and_explicit_oauth_servers(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_use_config(tmp_path, monkeypatch)
|
||||
|
||||
payload = custom_mcp_action(
|
||||
"import",
|
||||
{
|
||||
"config": [
|
||||
(
|
||||
'{"mcpServers":{'
|
||||
'"notion-work":{"url":"https://mcp.notion.com/mcp"},'
|
||||
'"company-mcp":{"url":"https://mcp.example.com/mcp","auth":"oauth"},'
|
||||
'"notion-pat":{"url":"https://mcp.notion.com/mcp",'
|
||||
'"headers":{"Authorization":"Bearer secret"}}'
|
||||
'}}'
|
||||
)
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
config = load_config()
|
||||
assert config.tools.mcp_servers["notion-work"].auth == "oauth"
|
||||
assert config.tools.mcp_servers["company-mcp"].auth == "oauth"
|
||||
assert config.tools.mcp_servers["notion-pat"].auth is None
|
||||
rows = {row["name"]: row for row in payload["presets"]}
|
||||
assert rows["notion-work"]["status"] == "authorization_required"
|
||||
assert rows["company-mcp"]["status"] == "authorization_required"
|
||||
assert rows["notion-pat"]["status"] == "configured"
|
||||
assert "Bearer secret" not in str(payload)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_replacing_oauth_config_removes_its_stored_credentials(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_use_config(tmp_path, monkeypatch)
|
||||
server_url = "https://mcp.example.com/mcp"
|
||||
custom_mcp_action(
|
||||
"custom",
|
||||
{
|
||||
"name": ["company-mcp"],
|
||||
"transport": ["streamableHttp"],
|
||||
"url": [server_url],
|
||||
"auth": ["oauth"],
|
||||
},
|
||||
)
|
||||
await MCPOAuthStorage("company-mcp", server_url).set_tokens(
|
||||
OAuthToken(access_token="secret")
|
||||
)
|
||||
assert mcp_oauth_has_credentials("company-mcp", server_url)
|
||||
|
||||
custom_mcp_action(
|
||||
"custom",
|
||||
{
|
||||
"name": ["company-mcp"],
|
||||
"transport": ["streamableHttp"],
|
||||
"url": [server_url],
|
||||
},
|
||||
)
|
||||
|
||||
assert not mcp_oauth_has_credentials("company-mcp", server_url)
|
||||
|
||||
|
||||
def test_normalize_mcp_preset_mentions_accepts_configured_custom_server(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
import pytest
|
||||
@@ -28,6 +28,7 @@ def _router(*, authorized: bool = True) -> WebUISettingsRouter:
|
||||
),
|
||||
runtime_surface="browser",
|
||||
runtime_capabilities={},
|
||||
mcp_oauth_redirect_uri=lambda _request: "https://gateway.example/auth/mcp/callback",
|
||||
)
|
||||
|
||||
|
||||
@@ -39,6 +40,121 @@ def _mutation_request(path: str, payload: dict[str, object]) -> SimpleNamespace:
|
||||
return request
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_oauth_start_uses_gateway_callback_and_requires_api_auth(monkeypatch) -> None:
|
||||
config = SimpleNamespace(
|
||||
type="streamableHttp",
|
||||
auth="oauth",
|
||||
url="https://app.xmind.com/api/mcp",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"nanobot.webui.settings_routes.ensure_mcp_oauth_server",
|
||||
lambda _query, *, config_path=None: ("xmind", config),
|
||||
)
|
||||
router = _router()
|
||||
start = AsyncMock(return_value={
|
||||
"status": "authorization_required",
|
||||
"flow_id": "flow-123",
|
||||
"name": "xmind",
|
||||
"authorization_url": "https://xmind.example/authorize?state=state-123",
|
||||
})
|
||||
router._mcp_oauth = SimpleNamespace(start=start)
|
||||
request = _mutation_request(
|
||||
"/api/settings/mcp-oauth/start",
|
||||
{"name": "xmind"},
|
||||
)
|
||||
|
||||
response = await router.dispatch(None, request, "/api/settings/mcp-oauth/start")
|
||||
|
||||
assert response is not None
|
||||
assert response.status_code == 200
|
||||
assert json.loads(response.body)["flow_id"] == "flow-123"
|
||||
start.assert_awaited_once_with(
|
||||
"xmind",
|
||||
config,
|
||||
"https://gateway.example/auth/mcp/callback",
|
||||
reload_mcp=ANY,
|
||||
reset_credentials=False,
|
||||
)
|
||||
|
||||
denied = _router(authorized=False)
|
||||
denied_response = await denied.dispatch(None, request, "/api/settings/mcp-oauth/start")
|
||||
assert denied_response is not None
|
||||
assert denied_response.status_code == 401
|
||||
|
||||
failed = _router()
|
||||
failed._mcp_oauth = SimpleNamespace(
|
||||
start=AsyncMock(side_effect=RuntimeError("upstream secret response"))
|
||||
)
|
||||
failed_response = await failed.dispatch(None, request, "/api/settings/mcp-oauth/start")
|
||||
assert failed_response is not None
|
||||
assert failed_response.status_code == 500
|
||||
assert json.loads(failed_response.body) == {"error": "MCP OAuth start failed"}
|
||||
assert b"upstream secret response" not in failed_response.body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_oauth_callback_is_state_authenticated_and_returns_close_page() -> None:
|
||||
router = _router(authorized=False)
|
||||
submit = MagicMock(return_value="xmind")
|
||||
router._mcp_oauth = SimpleNamespace(submit_callback=submit)
|
||||
request = SimpleNamespace(
|
||||
path="/auth/mcp/callback?code=oauth-code&state=state-123",
|
||||
headers=Headers(),
|
||||
)
|
||||
|
||||
response = await router.dispatch(None, request, "/auth/mcp/callback")
|
||||
|
||||
assert response is not None
|
||||
assert response.status_code == 200
|
||||
assert response.headers["Content-Type"] == "text/html; charset=utf-8"
|
||||
assert response.headers["Cache-Control"] == "no-store"
|
||||
assert "frame-ancestors 'none'" in response.headers["Content-Security-Policy"]
|
||||
assert b"window.close" in response.body
|
||||
assert b"Authorization received" in response.body
|
||||
assert b"oauth-code" not in response.body
|
||||
submit.assert_called_once_with(state="state-123", code="oauth-code", error=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_oauth_manual_completion_reads_websocket_payload() -> None:
|
||||
callback_url = (
|
||||
"http://127.0.0.1:8765/auth/mcp/callback?code=oauth-code&state=state-123"
|
||||
)
|
||||
router = _router()
|
||||
submit = MagicMock(
|
||||
return_value={
|
||||
"flow_id": "flow-123",
|
||||
"name": "linear",
|
||||
"status": "connecting",
|
||||
"expires_in": 299,
|
||||
"completion_input": "callback_url",
|
||||
}
|
||||
)
|
||||
router._mcp_oauth = SimpleNamespace(submit_callback_url=submit)
|
||||
request = _mutation_request(
|
||||
"/api/settings/mcp-oauth/complete",
|
||||
{"flow_id": "flow-123", "callback_url": callback_url},
|
||||
)
|
||||
|
||||
response = await router.dispatch(None, request, "/api/settings/mcp-oauth/complete")
|
||||
|
||||
assert response is not None
|
||||
assert response.status_code == 200
|
||||
assert json.loads(response.body)["status"] == "connecting"
|
||||
assert b"oauth-code" not in response.body
|
||||
submit.assert_called_once_with(flow_id="flow-123", callback_url=callback_url)
|
||||
|
||||
denied = _router(authorized=False)
|
||||
denied_response = await denied.dispatch(
|
||||
None,
|
||||
request,
|
||||
"/api/settings/mcp-oauth/complete",
|
||||
)
|
||||
assert denied_response is not None
|
||||
assert denied_response.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("provider", "authorization_response"),
|
||||
[
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import cast
|
||||
|
||||
from websockets.datastructures import Headers
|
||||
from websockets.http11 import Request as WsRequest
|
||||
|
||||
from nanobot.channels.websocket.runtime import WebSocketConfig
|
||||
from nanobot.webui.ws_http import GatewayHTTPHandler
|
||||
|
||||
|
||||
def _handler(config: WebSocketConfig) -> GatewayHTTPHandler:
|
||||
handler = object.__new__(GatewayHTTPHandler)
|
||||
handler.config = config
|
||||
return handler
|
||||
|
||||
|
||||
def _request(**headers: str) -> WsRequest:
|
||||
return cast(WsRequest, SimpleNamespace(headers=Headers(headers)))
|
||||
|
||||
|
||||
def test_mcp_oauth_callback_uses_configured_public_websocket_origin() -> None:
|
||||
handler = _handler(WebSocketConfig(path="/ws", public_ws_url="wss://agent.example/ws"))
|
||||
|
||||
redirect_uri = handler._mcp_oauth_redirect_uri(_request(Host="ignored.example"))
|
||||
|
||||
assert redirect_uri == "https://agent.example/auth/mcp/callback"
|
||||
|
||||
|
||||
def test_mcp_oauth_callback_uses_safe_forwarded_request_origin() -> None:
|
||||
handler = _handler(WebSocketConfig(path="/ws", host="127.0.0.1", port=8765))
|
||||
|
||||
redirect_uri = handler._mcp_oauth_redirect_uri(
|
||||
_request(Host="nanobot.example:9443", **{"X-Forwarded-Proto": "https"})
|
||||
)
|
||||
|
||||
assert redirect_uri == "https://nanobot.example:9443/auth/mcp/callback"
|
||||
Reference in New Issue
Block a user