factory/mcp_server/tests.py

165 lines
5.3 KiB
Python

import asyncio
from types import SimpleNamespace
from unittest.mock import patch
from django.test import SimpleTestCase
from mcp.server.auth.middleware.auth_context import auth_context_var
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
from mcp.server.auth.provider import AccessToken
from rest_framework.exceptions import AuthenticationFailed
from starlette.testclient import TestClient
from mcp_server.auth import verify_factory_jwt
from mcp_server.server import (
PROTOCOL_REVISION,
create_app,
factory_server_info,
factory_whoami,
mcp,
)
class FactoryJWTVerifierTests(SimpleTestCase):
def test_factory_access_token_is_accepted(self):
user = SimpleNamespace(
pk=42,
name="MCP用户",
is_superuser=False,
get_username=lambda: "mcp-user",
)
authentication = patch("mcp_server.auth.JWTAuthentication").start()
self.addCleanup(patch.stopall)
authentication.return_value.get_validated_token.return_value = {
"exp": 1234567890,
}
authentication.return_value.get_user.return_value = user
access_token = verify_factory_jwt("access-token")
self.assertIsNotNone(access_token)
self.assertEqual(access_token.subject, str(user.pk))
self.assertEqual(
access_token.claims["factory_user"]["username"],
user.get_username(),
)
def test_invalid_token_is_rejected(self):
with patch("mcp_server.auth.JWTAuthentication") as authentication:
authentication.return_value.get_validated_token.side_effect = AuthenticationFailed("invalid token")
self.assertIsNone(verify_factory_jwt("invalid-token"))
class FactoryMCPServerTests(SimpleTestCase):
def test_base_tools_are_registered(self):
tools = asyncio.run(mcp.list_tools())
names = {tool.name for tool in tools}
self.assertEqual(
names,
{
"execute_dataset",
"factory_server_info",
"factory_whoami",
"get_batch_stat",
"get_wpr",
"search_batch_stats",
"search_datasets",
"search_wprs",
},
)
def test_server_info_targets_mcp_v2(self):
result = factory_server_info()
self.assertEqual(result["protocol_revision"], PROTOCOL_REVISION)
self.assertTrue(result["domain_tools_ready"])
def test_whoami_uses_authenticated_request_context(self):
user = {
"id": "42",
"username": "agent-user",
"name": "Agent用户",
"is_superuser": False,
}
authenticated = AuthenticatedUser(
AccessToken(
token="test-token",
client_id="factory-mcp",
scopes=["factory:user"],
subject=user["id"],
claims={"factory_user": user},
)
)
context_token = auth_context_var.set(authenticated)
try:
self.assertEqual(factory_whoami(), user)
finally:
auth_context_var.reset(context_token)
def test_http_endpoint_requires_bearer_token(self):
app = create_app(allowed_hosts=["testserver"])
with TestClient(app) as client:
response = client.post("/mcp", json={})
self.assertEqual(response.status_code, 401)
self.assertEqual(response.json()["error"], "invalid_token")
def test_mcp_v2_request_uses_bearer_identity(self):
user = {
"id": "42",
"username": "agent-user",
"name": "Agent用户",
"is_superuser": False,
}
class TestTokenVerifier:
async def verify_token(self, token):
if token != "valid-token":
return None
return AccessToken(
token=token,
client_id="test-client",
scopes=["factory:user"],
subject=user["id"],
claims={"factory_user": user},
)
app = create_app(
allowed_hosts=["testserver"],
token_verifier=TestTokenVerifier(),
)
request = {
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": {
"name": "factory_whoami",
"arguments": {},
"_meta": {
"io.modelcontextprotocol/protocolVersion": (PROTOCOL_REVISION),
"io.modelcontextprotocol/clientInfo": {
"name": "factory-tests",
"version": "1.0",
},
"io.modelcontextprotocol/clientCapabilities": {},
},
},
}
headers = {
"Authorization": "Bearer valid-token",
"MCP-Protocol-Version": PROTOCOL_REVISION,
"Mcp-Method": "tools/call",
"Mcp-Name": "factory_whoami",
}
with TestClient(app) as client:
response = client.post("/mcp", json=request, headers=headers)
self.assertEqual(response.status_code, 200)
self.assertEqual(
response.json()["result"]["structuredContent"],
user,
)