165 lines
5.3 KiB
Python
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,
|
|
)
|