"""Tests for OAuth Proxy consent page display, CSP policy, and consent binding cookie."""
import re
import secrets
import time
from unittest.mock import Mock
from urllib.parse import parse_qs, urlparse
import pytest
from key_value.aio.stores.memory import MemoryStore
from mcp.server.auth.provider import AuthorizationParams
from mcp.shared.auth import OAuthClientInformationFull
from mcp_types import Icon
from pydantic import AnyUrl
from starlette.applications import Starlette
from starlette.testclient import TestClient
from fastmcp import FastMCP
from fastmcp.server.auth.auth import AccessToken, TokenVerifier
from fastmcp.server.auth.oauth_proxy import OAuthProxy
from fastmcp.server.auth.oauth_proxy.models import OAuthTransaction
class _Verifier(TokenVerifier):
"""Minimal token verifier for security tests."""
def __init__(self):
self.required_scopes = ["read"]
async def verify_token(self, token: str):
return AccessToken(
token=token, client_id="c", scopes=self.required_scopes, expires_at=None
)
@pytest.fixture
def oauth_proxy_https():
"""OAuthProxy configured with HTTPS base_url for __Host- cookies."""
return OAuthProxy(
upstream_authorization_endpoint="https://github.com/login/oauth/authorize",
upstream_token_endpoint="https://github.com/login/oauth/access_token",
upstream_client_id="client-id",
upstream_client_secret="client-secret",
token_verifier=_Verifier(),
base_url="https://myserver.example",
client_storage=MemoryStore(),
jwt_signing_key="test-secret",
)
@pytest.fixture
def oauth_proxy_https_remember():
"""OAuthProxy in 'remember' mode for silent-consent cookie tests."""
return OAuthProxy(
upstream_authorization_endpoint="https://github.com/login/oauth/authorize",
upstream_token_endpoint="https://github.com/login/oauth/access_token",
upstream_client_id="client-id",
upstream_client_secret="client-secret",
token_verifier=_Verifier(),
base_url="https://myserver.example",
client_storage=MemoryStore(),
jwt_signing_key="test-secret",
require_authorization_consent="remember",
)
async def _start_flow(
proxy: OAuthProxy, client_id: str, redirect: str
) -> tuple[str, str]:
"""Register client and start auth; returns (txn_id, consent_url)."""
await proxy.register_client(
OAuthClientInformationFull(
client_id=client_id,
client_secret="s",
redirect_uris=[AnyUrl(redirect)],
)
)
params = AuthorizationParams(
redirect_uri=AnyUrl(redirect),
redirect_uri_provided_explicitly=True,
state="client-state-xyz",
code_challenge="challenge",
scopes=["read"],
)
consent_url = await proxy.authorize(
OAuthClientInformationFull(
client_id=client_id,
client_secret="s",
redirect_uris=[AnyUrl(redirect)],
),
params,
)
qs = parse_qs(urlparse(consent_url).query)
return qs["txn_id"][0], consent_url
def _extract_csrf(html: str) -> str | None:
"""Extract CSRF token from HTML form."""
m = re.search(r"name=\"csrf_token\"\s+value=\"([^\"]+)\"", html)
return m.group(1) if m else None
class TestConsentPageServerIcon:
"""Tests for server icon display in OAuth consent screen."""
async def test_consent_screen_displays_server_icon(self):
"""Test that consent screen shows server's custom icon when available."""
# Create mock JWT verifier
verifier = Mock(spec=TokenVerifier)
verifier.required_scopes = ["read"]
verifier.verify_token = Mock(return_value=None)
# Create OAuthProxy
proxy = OAuthProxy(
upstream_authorization_endpoint="https://oauth.example.com/authorize",
upstream_token_endpoint="https://oauth.example.com/token",
upstream_client_id="upstream-client",
upstream_client_secret="upstream-secret",
token_verifier=verifier,
base_url="https://proxy.example.com",
client_storage=MemoryStore(),
jwt_signing_key="test-secret",
)
# Create FastMCP server with custom icon
server = FastMCP(
name="My Custom Server",
auth=proxy,
icons=[Icon(src="https://example.com/custom-icon.png")],
website_url="https://example.com",
)
# Create HTTP app
app = server.http_app()
# Register a test client with the proxy
client_info = OAuthClientInformationFull(
client_id="test-client",
client_secret="test-secret",
redirect_uris=[AnyUrl("http://localhost:12345/callback")],
)
await proxy.register_client(client_info)
# Create a transaction manually
txn_id = "test-txn-id"
transaction = OAuthTransaction(
txn_id=txn_id,
client_id="test-client",
client_redirect_uri="http://localhost:12345/callback",
client_state="client-state",
code_challenge="challenge",
code_challenge_method="S256",
scopes=["read"],
created_at=time.time(),
)
await proxy._transaction_store.put(key=txn_id, value=transaction)
# Make request to consent page
with TestClient(app) as client:
response = client.get(f"/consent?txn_id={txn_id}")
# Check that response is successful
assert response.status_code == 200
# Check that HTML contains custom icon
assert "https://example.com/custom-icon.png" in response.text
# Check that server name is used as alt text
assert 'alt="My Custom Server"' in response.text
async def test_consent_screen_falls_back_to_fastmcp_logo(self):
"""Test that consent screen shows FastMCP logo when no server icon provided."""
# Create mock JWT verifier
verifier = Mock(spec=TokenVerifier)
verifier.required_scopes = ["read"]
verifier.verify_token = Mock(return_value=None)
# Create OAuthProxy
proxy = OAuthProxy(
upstream_authorization_endpoint="https://oauth.example.com/authorize",
upstream_token_endpoint="https://oauth.example.com/token",
upstream_client_id="upstream-client",
upstream_client_secret="upstream-secret",
token_verifier=verifier,
base_url="https://proxy.example.com",
client_storage=MemoryStore(),
jwt_signing_key="test-secret",
)
# Create FastMCP server without icon
server = FastMCP(name="Server Without Icon", auth=proxy)
# Create HTTP app
app = server.http_app()
# Register a test client
client_info = OAuthClientInformationFull(
client_id="test-client",
client_secret="test-secret",
redirect_uris=[AnyUrl("http://localhost:12345/callback")],
)
await proxy.register_client(client_info)
# Create a transaction
txn_id = "test-txn-id"
transaction = OAuthTransaction(
txn_id=txn_id,
client_id="test-client",
client_redirect_uri="http://localhost:12345/callback",
client_state="client-state",
code_challenge="challenge",
code_challenge_method="S256",
scopes=["read"],
created_at=time.time(),
)
await proxy._transaction_store.put(key=txn_id, value=transaction)
# Make request to consent page
with TestClient(app) as client:
response = client.get(f"/consent?txn_id={txn_id}")
# Check that response is successful
assert response.status_code == 200
# Check that HTML contains FastMCP logo
assert "gofastmcp.com/assets/brand/blue-logo.png" in response.text
# Check that alt text is still the server name
assert 'alt="Server Without Icon"' in response.text
async def test_consent_screen_escapes_server_name(self):
"""Test that server name is properly HTML-escaped."""
# Create mock JWT verifier
verifier = Mock(spec=TokenVerifier)
verifier.required_scopes = ["read"]
verifier.verify_token = Mock(return_value=None)
# Create OAuthProxy
proxy = OAuthProxy(
upstream_authorization_endpoint="https://oauth.example.com/authorize",
upstream_token_endpoint="https://oauth.example.com/token",
upstream_client_id="upstream-client",
upstream_client_secret="upstream-secret",
token_verifier=verifier,
base_url="https://proxy.example.com",
client_storage=MemoryStore(),
jwt_signing_key="test-secret",
)
# Create FastMCP server with special characters in name
server = FastMCP(
name='Server',
auth=proxy,
icons=[Icon(src="https://example.com/icon.png")],
)
# Create HTTP app
app = server.http_app()
# Register a test client
client_info = OAuthClientInformationFull(
client_id="test-client",
client_secret="test-secret",
redirect_uris=[AnyUrl("http://localhost:12345/callback")],
)
await proxy.register_client(client_info)
# Create a transaction
txn_id = "test-txn-id"
transaction = OAuthTransaction(
txn_id=txn_id,
client_id="test-client",
client_redirect_uri="http://localhost:12345/callback",
client_state="client-state",
code_challenge="challenge",
code_challenge_method="S256",
scopes=["read"],
created_at=time.time(),
)
await proxy._transaction_store.put(key=txn_id, value=transaction)
# Make request to consent page
with TestClient(app) as client:
response = client.get(f"/consent?txn_id={txn_id}")
# Check that response is successful
assert response.status_code == 200
# Check that script tag is escaped
assert "