feat(oauth): Phase B — claude.ai protocol compatibility

Add 3 features for claude.ai Connectors compatibility:

- FEATURE-3: client_secret_basic auth on token endpoint (RFC 6749 §2.3.1)
- FEATURE-2: Token revocation endpoint /oauth/revoke (RFC 7009)
- FEATURE-4: resource parameter support with JWT aud claim (RFC 8707)

23 new tests (481 total), all passing.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
airano
2026-02-24 20:49:23 +03:30
parent 1736779d69
commit d3bcb31053
8 changed files with 852 additions and 3 deletions

View File

@@ -0,0 +1,165 @@
"""Tests for client_secret_basic (HTTP Basic Auth) on the OAuth token endpoint."""
import base64
import os
import tempfile
from unittest.mock import AsyncMock, MagicMock
import pytest
def _make_request(headers=None, body=None, content_type="application/x-www-form-urlencoded"):
"""Create a mock Starlette Request for oauth_token."""
request = AsyncMock()
request.headers = MagicMock()
_headers = {"content-type": content_type}
if headers:
_headers.update(headers)
request.headers.get = lambda key, default="": _headers.get(key.lower(), default)
form_data = body or {}
form = AsyncMock(return_value=form_data)
request.form = form
return request
def _encode_basic(client_id: str, client_secret: str) -> str:
"""Encode client_id:client_secret as Basic Auth header value."""
raw = f"{client_id}:{client_secret}"
encoded = base64.b64encode(raw.encode("utf-8")).decode("utf-8")
return f"Basic {encoded}"
@pytest.fixture
def temp_storage():
"""Create temporary storage for OAuth tests."""
with tempfile.TemporaryDirectory() as tmpdir:
os.environ["OAUTH_STORAGE_PATH"] = tmpdir
os.environ["OAUTH_JWT_SECRET_KEY"] = "test_secret_key_for_basic_auth_tests"
from core.oauth import client_registry, server, storage, token_manager
client_registry._client_registry = None
token_manager._token_manager = None
storage._storage = None
server._oauth_server = None
yield tmpdir
os.environ.pop("OAUTH_STORAGE_PATH", None)
os.environ.pop("OAUTH_JWT_SECRET_KEY", None)
client_registry._client_registry = None
token_manager._token_manager = None
storage._storage = None
server._oauth_server = None
async def _call_oauth_token(request):
"""Import and call the oauth_token endpoint from server.py."""
from server import oauth_token
return await oauth_token(request)
@pytest.mark.unit
class TestClientSecretBasicAuth:
"""Tests for client_secret_basic authentication on the token endpoint."""
async def test_basic_auth_header_parsed(self, temp_storage):
"""Basic Auth header is correctly parsed (base64 encoded client_id:client_secret)."""
request = _make_request(
headers={"authorization": _encode_basic("my_client_id", "my_client_secret")},
body={"grant_type": "client_credentials"},
)
response = await _call_oauth_token(request)
# The request will likely fail on actual client validation,
# but we're testing that it gets past the credential parsing stage.
# If it returned invalid_request with "Missing client_id", Basic Auth parsing failed.
import json
data = json.loads(response.body)
assert data.get("error") != "invalid_request" or "Missing client_id" not in data.get(
"error_description", ""
)
async def test_body_params_take_priority_over_basic_auth(self, temp_storage):
"""Body params take priority over Basic Auth (setdefault behavior)."""
request = _make_request(
headers={"authorization": _encode_basic("basic_id", "basic_secret")},
body={
"grant_type": "client_credentials",
"client_id": "body_id",
"client_secret": "body_secret",
},
)
response = await _call_oauth_token(request)
import json
data = json.loads(response.body)
# The body params should be used, not the Basic Auth ones.
# If the error mentions "body_id", body params were used.
# If it mentions "basic_id", Basic Auth overrode body params (wrong).
# Since the client won't exist, we expect invalid_client error.
# The key check: it should NOT have replaced body_id with basic_id.
if data.get("error") == "invalid_client":
# The error came from actual client validation, meaning
# credentials were parsed. The body params were used.
pass
else:
# Should not get "Missing client_id" error
assert "Missing client_id" not in data.get("error_description", "")
async def test_malformed_basic_auth_returns_invalid_client(self, temp_storage):
"""Malformed Basic Auth header returns invalid_client with 401 status."""
request = _make_request(
headers={"authorization": "Basic !!!not-valid-base64!!!"},
body={"grant_type": "client_credentials"},
)
response = await _call_oauth_token(request)
import json
data = json.loads(response.body)
assert data["error"] == "invalid_client"
assert "Invalid Basic authentication header" in data["error_description"]
assert response.status_code == 401
async def test_missing_credentials_returns_invalid_request(self, temp_storage):
"""Missing credentials (no body, no Basic header) returns invalid_request."""
request = _make_request(
body={"grant_type": "client_credentials"},
)
response = await _call_oauth_token(request)
import json
data = json.loads(response.body)
assert data["error"] == "invalid_request"
assert "Missing client_id or client_secret" in data["error_description"]
async def test_client_secret_with_colon(self, temp_storage):
"""client_secret containing ':' is handled correctly (split on first ':' only)."""
secret_with_colon = "my:secret:with:colons"
request = _make_request(
headers={"authorization": _encode_basic("my_client_id", secret_with_colon)},
body={"grant_type": "client_credentials"},
)
response = await _call_oauth_token(request)
import json
data = json.loads(response.body)
# Should not get invalid_request about missing credentials
assert data.get("error") != "invalid_request" or "Missing client_id" not in data.get(
"error_description", ""
)
# Should not get invalid_client from Basic Auth parsing
assert data.get("error_description") != "Invalid Basic authentication header"

View File

@@ -0,0 +1,333 @@
"""
Tests for RFC 8707 resource parameter support in OAuth 2.1 flow.
Validates that the resource parameter is:
1. Accepted and stored in authorization codes
2. Passed through to JWT access tokens as aud claim
3. Optional -- existing flows without resource continue to work
"""
import os
import tempfile
from datetime import UTC, datetime, timedelta
import jwt
import pytest
from core.oauth.schemas import AuthorizationCode
from core.oauth.token_manager import TokenManager
@pytest.fixture
def token_manager():
"""Create a fresh TokenManager for tests."""
os.environ["OAUTH_JWT_SECRET_KEY"] = "test_secret_key_resource"
with tempfile.TemporaryDirectory() as tmpdir:
os.environ["OAUTH_STORAGE_PATH"] = tmpdir
yield TokenManager()
os.environ.pop("OAUTH_JWT_SECRET_KEY", None)
os.environ.pop("OAUTH_STORAGE_PATH", None)
@pytest.fixture
def temp_storage():
"""Create temporary storage for OAuth server tests."""
with tempfile.TemporaryDirectory() as tmpdir:
os.environ["OAUTH_STORAGE_PATH"] = tmpdir
os.environ["OAUTH_JWT_SECRET_KEY"] = "test_secret_key_resource"
from core.oauth import client_registry, server, storage, token_manager
client_registry._client_registry = None
token_manager._token_manager = None
storage._storage = None
server._oauth_server = None
yield tmpdir
# Teardown
os.environ.pop("OAUTH_STORAGE_PATH", None)
os.environ.pop("OAUTH_JWT_SECRET_KEY", None)
client_registry._client_registry = None
token_manager._token_manager = None
storage._storage = None
server._oauth_server = None
@pytest.fixture
def oauth_components(temp_storage):
"""Get fresh OAuth components."""
from core.oauth import get_client_registry, get_oauth_server, get_storage, get_token_manager
return {
"server": get_oauth_server(),
"client_registry": get_client_registry(),
"token_manager": get_token_manager(),
"storage": get_storage(),
}
@pytest.fixture
def test_client(oauth_components):
"""Create a test OAuth client."""
client_registry = oauth_components["client_registry"]
client_id, client_secret = client_registry.create_client(
client_name="Resource Test Client",
redirect_uris=["http://localhost:3000/callback"],
grant_types=["authorization_code", "refresh_token"],
allowed_scopes=["read", "write"],
)
return {
"client_id": client_id,
"client_secret": client_secret,
"redirect_uri": "http://localhost:3000/callback",
}
# --- AuthorizationCode schema tests ---
def test_authorization_code_accepts_resource():
"""Test that AuthorizationCode model accepts and stores the resource field."""
auth_code = AuthorizationCode(
code="auth_test123",
client_id="test_client",
redirect_uri="http://localhost:3000/callback",
scope="read write",
code_challenge="abc123",
code_challenge_method="S256",
expires_at=datetime.now(UTC) + timedelta(minutes=5),
resource="https://mcp.example.com",
)
assert auth_code.resource == "https://mcp.example.com"
def test_authorization_code_resource_defaults_to_none():
"""Test that AuthorizationCode resource defaults to None when not provided."""
auth_code = AuthorizationCode(
code="auth_test456",
client_id="test_client",
redirect_uri="http://localhost:3000/callback",
scope="read",
code_challenge="xyz789",
code_challenge_method="S256",
expires_at=datetime.now(UTC) + timedelta(minutes=5),
)
assert auth_code.resource is None
# --- create_authorization_code tests ---
def test_create_authorization_code_with_resource(oauth_components, test_client):
"""Test that create_authorization_code accepts and stores resource parameter."""
from core.oauth import generate_code_challenge, generate_code_verifier
code_verifier = generate_code_verifier()
code_challenge = generate_code_challenge(code_verifier)
oauth_server = oauth_components["server"]
resource_url = "https://mcp.example.com"
code = oauth_server.create_authorization_code(
client_id=test_client["client_id"],
redirect_uri=test_client["redirect_uri"],
scope="read write",
code_challenge=code_challenge,
code_challenge_method="S256",
resource=resource_url,
)
assert code.startswith("auth_")
# Verify resource is stored in the authorization code
stored_code = oauth_components["storage"].get_authorization_code(code)
assert stored_code is not None
assert stored_code.resource == resource_url
def test_create_authorization_code_without_resource(oauth_components, test_client):
"""Test that create_authorization_code works without resource (backward compat)."""
from core.oauth import generate_code_challenge, generate_code_verifier
code_verifier = generate_code_verifier()
code_challenge = generate_code_challenge(code_verifier)
oauth_server = oauth_components["server"]
code = oauth_server.create_authorization_code(
client_id=test_client["client_id"],
redirect_uri=test_client["redirect_uri"],
scope="read write",
code_challenge=code_challenge,
code_challenge_method="S256",
)
assert code.startswith("auth_")
stored_code = oauth_components["storage"].get_authorization_code(code)
assert stored_code is not None
assert stored_code.resource is None
# --- generate_access_token / JWT aud claim tests ---
def test_generate_access_token_with_resource_sets_aud(token_manager):
"""Test that resource parameter flows through to JWT aud claim."""
resource_url = "https://mcp.example.com"
token = token_manager.generate_access_token(
client_id="test_client",
scope="read write",
user_id="user_123",
resource=resource_url,
)
# Decode JWT and verify aud claim
payload = jwt.decode(
token,
"test_secret_key_resource",
algorithms=["HS256"],
audience=resource_url,
)
assert payload["aud"] == resource_url
assert payload["client_id"] == "test_client"
assert payload["sub"] == "user_123"
def test_generate_access_token_without_resource_no_aud(token_manager):
"""Test that JWT has no aud claim when resource is not provided."""
token = token_manager.generate_access_token(
client_id="test_client",
scope="read write",
user_id="user_123",
)
payload = jwt.decode(
token,
"test_secret_key_resource",
algorithms=["HS256"],
)
assert "aud" not in payload
assert payload["client_id"] == "test_client"
def test_generate_access_token_resource_none_no_aud(token_manager):
"""Test that resource=None does not add aud claim."""
token = token_manager.generate_access_token(
client_id="test_client",
scope="read",
resource=None,
)
payload = jwt.decode(
token,
"test_secret_key_resource",
algorithms=["HS256"],
)
assert "aud" not in payload
def test_generate_access_token_resource_empty_string_no_aud(token_manager):
"""Test that empty string resource does not add aud claim."""
token = token_manager.generate_access_token(
client_id="test_client",
scope="read",
resource="",
)
payload = jwt.decode(
token,
"test_secret_key_resource",
algorithms=["HS256"],
)
assert "aud" not in payload
# --- Full flow: resource through code exchange ---
def test_full_flow_resource_to_jwt_aud(oauth_components, test_client):
"""Test that resource flows from auth code creation through to JWT aud claim."""
from core.oauth import generate_code_challenge, generate_code_verifier
code_verifier = generate_code_verifier()
code_challenge = generate_code_challenge(code_verifier)
oauth_server = oauth_components["server"]
resource_url = "https://mcp.example.com"
# Step 1: Create authorization code with resource
code = oauth_server.create_authorization_code(
client_id=test_client["client_id"],
redirect_uri=test_client["redirect_uri"],
scope="read write",
code_challenge=code_challenge,
code_challenge_method="S256",
api_key_id="master",
api_key_project_id="*",
api_key_scope="read write",
resource=resource_url,
)
# Step 2: Exchange code for tokens
token_response = oauth_server.exchange_code_for_tokens(
client_id=test_client["client_id"],
client_secret=test_client["client_secret"],
code=code,
redirect_uri=test_client["redirect_uri"],
code_verifier=code_verifier,
)
# Step 3: Verify JWT contains aud claim
payload = jwt.decode(
token_response.access_token,
os.environ["OAUTH_JWT_SECRET_KEY"],
algorithms=["HS256"],
audience=resource_url,
)
assert payload["aud"] == resource_url
def test_full_flow_without_resource_no_aud(oauth_components, test_client):
"""Test that full flow without resource produces JWT without aud claim."""
from core.oauth import generate_code_challenge, generate_code_verifier
code_verifier = generate_code_verifier()
code_challenge = generate_code_challenge(code_verifier)
oauth_server = oauth_components["server"]
# Step 1: Create authorization code WITHOUT resource
code = oauth_server.create_authorization_code(
client_id=test_client["client_id"],
redirect_uri=test_client["redirect_uri"],
scope="read write",
code_challenge=code_challenge,
code_challenge_method="S256",
api_key_id="master",
api_key_project_id="*",
api_key_scope="read write",
)
# Step 2: Exchange code for tokens
token_response = oauth_server.exchange_code_for_tokens(
client_id=test_client["client_id"],
client_secret=test_client["client_secret"],
code=code,
redirect_uri=test_client["redirect_uri"],
code_verifier=code_verifier,
)
# Step 3: Verify JWT does NOT contain aud claim
payload = jwt.decode(
token_response.access_token,
os.environ["OAUTH_JWT_SECRET_KEY"],
algorithms=["HS256"],
)
assert "aud" not in payload

227
tests/test_oauth_revoke.py Normal file
View File

@@ -0,0 +1,227 @@
"""Tests for OAuth 2.0 Token Revocation endpoint (RFC 7009)."""
import base64
import json
import os
import tempfile
from unittest.mock import AsyncMock, MagicMock
import pytest
def _make_request(headers=None, body=None, content_type="application/x-www-form-urlencoded"):
"""Create a mock Starlette Request for oauth_revoke."""
request = AsyncMock()
request.headers = MagicMock()
_headers = {"content-type": content_type}
if headers:
_headers.update(headers)
request.headers.get = lambda key, default="": _headers.get(key.lower(), default)
form_data = body or {}
form = AsyncMock(return_value=form_data)
request.form = form
if "application/json" in content_type:
request.json = AsyncMock(return_value=form_data)
return request
def _encode_basic(client_id: str, client_secret: str) -> str:
"""Encode client_id:client_secret as Basic Auth header value."""
raw = f"{client_id}:{client_secret}"
encoded = base64.b64encode(raw.encode("utf-8")).decode("utf-8")
return f"Basic {encoded}"
@pytest.fixture
def temp_storage():
"""Create temporary storage for OAuth tests."""
with tempfile.TemporaryDirectory() as tmpdir:
os.environ["OAUTH_STORAGE_PATH"] = tmpdir
os.environ["OAUTH_JWT_SECRET_KEY"] = "test_secret_key_for_revoke_tests"
from core.oauth import client_registry, server, storage, token_manager
client_registry._client_registry = None
token_manager._token_manager = None
storage._storage = None
server._oauth_server = None
yield tmpdir
os.environ.pop("OAUTH_STORAGE_PATH", None)
os.environ.pop("OAUTH_JWT_SECRET_KEY", None)
client_registry._client_registry = None
token_manager._token_manager = None
storage._storage = None
server._oauth_server = None
@pytest.fixture
def registered_client(temp_storage):
"""Create a registered OAuth client and return (client_id, client_secret)."""
from core.oauth import get_client_registry
registry = get_client_registry()
client_id, client_secret = registry.create_client(
client_name="Test Revoke Client",
redirect_uris=["http://localhost:3000/callback"],
)
return client_id, client_secret
async def _call_oauth_revoke(request):
"""Import and call the oauth_revoke endpoint from server.py."""
from server import oauth_revoke
return await oauth_revoke(request)
@pytest.mark.unit
class TestOAuthRevoke:
"""Tests for the /oauth/revoke endpoint (RFC 7009)."""
async def test_valid_token_revocation_returns_200(self, registered_client):
"""Valid token revocation returns 200 with empty body."""
client_id, client_secret = registered_client
request = _make_request(
body={
"client_id": client_id,
"client_secret": client_secret,
"token": "rt_some_refresh_token",
},
)
response = await _call_oauth_revoke(request)
assert response.status_code == 200
data = json.loads(response.body)
assert data == {}
async def test_missing_token_returns_200(self, registered_client):
"""Per RFC 7009, missing token returns 200 (no-op)."""
client_id, client_secret = registered_client
request = _make_request(
body={
"client_id": client_id,
"client_secret": client_secret,
},
)
response = await _call_oauth_revoke(request)
assert response.status_code == 200
data = json.loads(response.body)
assert data == {}
async def test_invalid_client_credentials_returns_401(self, registered_client):
"""Invalid client credentials return 401 with invalid_client error."""
client_id, _ = registered_client
request = _make_request(
body={
"client_id": client_id,
"client_secret": "wrong_secret",
"token": "rt_some_token",
},
)
response = await _call_oauth_revoke(request)
assert response.status_code == 401
data = json.loads(response.body)
assert data["error"] == "invalid_client"
assert "Invalid client credentials" in data["error_description"]
async def test_missing_client_credentials_returns_401(self, temp_storage):
"""Missing client credentials return 401 with invalid_client error."""
request = _make_request(
body={
"token": "rt_some_token",
},
)
response = await _call_oauth_revoke(request)
assert response.status_code == 401
data = json.loads(response.body)
assert data["error"] == "invalid_client"
assert "Client authentication required" in data["error_description"]
async def test_revocation_with_refresh_token_hint(self, registered_client):
"""Revocation with token_type_hint=refresh_token works correctly."""
client_id, client_secret = registered_client
request = _make_request(
body={
"client_id": client_id,
"client_secret": client_secret,
"token": "some_token_value",
"token_type_hint": "refresh_token",
},
)
response = await _call_oauth_revoke(request)
assert response.status_code == 200
data = json.loads(response.body)
assert data == {}
async def test_unknown_token_returns_200(self, registered_client):
"""Per RFC 7009, revoking an unknown/invalid token still returns 200."""
client_id, client_secret = registered_client
request = _make_request(
body={
"client_id": client_id,
"client_secret": client_secret,
"token": "completely_nonexistent_token_xyz",
},
)
response = await _call_oauth_revoke(request)
assert response.status_code == 200
data = json.loads(response.body)
assert data == {}
async def test_basic_auth_works_for_revocation(self, registered_client):
"""Client authentication via Basic Auth header works on the revoke endpoint."""
client_id, client_secret = registered_client
request = _make_request(
headers={"authorization": _encode_basic(client_id, client_secret)},
body={
"token": "rt_some_refresh_token",
},
)
response = await _call_oauth_revoke(request)
assert response.status_code == 200
data = json.loads(response.body)
assert data == {}
async def test_metadata_includes_revocation_endpoint(self, temp_storage):
"""OAuth metadata includes the revocation_endpoint field."""
from server import oauth_metadata
# Create a mock request with a host header
request = AsyncMock()
request.headers = MagicMock()
_headers = {"host": "localhost:8000"}
request.headers.get = lambda key, default="": _headers.get(key.lower(), default)
request.url = MagicMock()
request.url.scheme = "http"
response = await oauth_metadata(request)
data = json.loads(response.body)
assert "revocation_endpoint" in data
assert data["revocation_endpoint"].endswith("/oauth/revoke")
assert "client_secret_basic" in data.get("revocation_endpoint_auth_methods_supported", [])