181 lines
6.0 KiB
Python
181 lines
6.0 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import http.client
|
|
import socket
|
|
import tempfile
|
|
import threading
|
|
import unittest
|
|
import urllib.error
|
|
import urllib.request
|
|
from pathlib import Path
|
|
|
|
from project_bus.auth import StaticTokenAuthProvider, TokenIdentity
|
|
from project_bus.db import Database
|
|
from project_bus.mcp import MCPApplication
|
|
from project_bus.server import ProjectBusHTTPServer
|
|
from project_bus.service import ProjectBusService
|
|
|
|
|
|
class MCPHTTPTestCase(unittest.TestCase):
|
|
def setUp(self) -> None:
|
|
self.temporary = tempfile.TemporaryDirectory()
|
|
service = ProjectBusService(Database(Path(self.temporary.name) / "bus.db"))
|
|
service.initialize()
|
|
auth = StaticTokenAuthProvider(
|
|
{"bootstrap-secret-at-least-24": TokenIdentity("system", "system", "System")}
|
|
)
|
|
self.server = ProjectBusHTTPServer(
|
|
("127.0.0.1", 0),
|
|
MCPApplication(service),
|
|
auth,
|
|
{"https://chat.example"},
|
|
request_read_timeout_seconds=0.2,
|
|
)
|
|
self.thread = threading.Thread(target=self.server.serve_forever, daemon=True)
|
|
self.thread.start()
|
|
self.url = f"http://127.0.0.1:{self.server.server_port}/mcp"
|
|
|
|
def tearDown(self) -> None:
|
|
self.server.shutdown()
|
|
self.server.server_close()
|
|
self.thread.join()
|
|
self.temporary.cleanup()
|
|
|
|
def request(
|
|
self,
|
|
payload: dict,
|
|
*,
|
|
token: str | None = None,
|
|
origin: str | None = None,
|
|
) -> tuple[int, dict]:
|
|
headers = {"Content-Type": "application/json"}
|
|
if token:
|
|
headers["Authorization"] = f"Bearer {token}"
|
|
if origin:
|
|
headers["Origin"] = origin
|
|
request = urllib.request.Request(
|
|
self.url,
|
|
data=json.dumps(payload).encode(),
|
|
headers=headers,
|
|
method="POST",
|
|
)
|
|
try:
|
|
response = urllib.request.urlopen(request, timeout=2)
|
|
return response.status, json.loads(response.read())
|
|
except urllib.error.HTTPError as error:
|
|
return error.code, json.loads(error.read())
|
|
|
|
def test_initialize_and_list_tools(self) -> None:
|
|
status, initialized = self.request(
|
|
{"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}}
|
|
)
|
|
self.assertEqual(status, 200)
|
|
self.assertEqual(initialized["result"]["protocolVersion"], "2025-03-26")
|
|
status, listed = self.request(
|
|
{"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}}
|
|
)
|
|
self.assertEqual(status, 200)
|
|
names = {tool["name"] for tool in listed["result"]["tools"]}
|
|
self.assertIn("sync_since", names)
|
|
self.assertIn("close_gate", names)
|
|
|
|
def test_authenticated_tool_call(self) -> None:
|
|
status, response = self.request(
|
|
{
|
|
"jsonrpc": "2.0",
|
|
"id": 1,
|
|
"method": "tools/call",
|
|
"params": {
|
|
"name": "create_project",
|
|
"arguments": {
|
|
"project_id": "demo",
|
|
"name": "Demo",
|
|
"idempotency_key": "create-demo",
|
|
},
|
|
},
|
|
},
|
|
token="bootstrap-secret-at-least-24",
|
|
)
|
|
self.assertEqual(status, 200)
|
|
self.assertFalse(response["result"]["isError"])
|
|
self.assertEqual(response["result"]["structuredContent"]["project_id"], "demo")
|
|
|
|
def test_tool_call_requires_authentication(self) -> None:
|
|
status, response = self.request(
|
|
{
|
|
"jsonrpc": "2.0",
|
|
"id": 1,
|
|
"method": "tools/call",
|
|
"params": {"name": "sync_since", "arguments": {"project_id": "demo"}},
|
|
}
|
|
)
|
|
self.assertEqual(status, 401)
|
|
self.assertEqual(response["error"], "unauthorized")
|
|
|
|
def test_origin_allowlist(self) -> None:
|
|
status, response = self.request(
|
|
{"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}},
|
|
origin="https://evil.example",
|
|
)
|
|
self.assertEqual(status, 403)
|
|
self.assertEqual(response["error"], "origin_not_allowed")
|
|
|
|
def test_deeply_nested_json_returns_parse_error(self) -> None:
|
|
body = b"[" * 200_000 + b"]" * 200_000
|
|
connection = http.client.HTTPConnection(
|
|
"127.0.0.1", self.server.server_port, timeout=2
|
|
)
|
|
connection.request(
|
|
"POST",
|
|
"/mcp",
|
|
body=body,
|
|
headers={"Content-Type": "application/json"},
|
|
)
|
|
response = connection.getresponse()
|
|
payload = json.loads(response.read())
|
|
connection.close()
|
|
|
|
self.assertEqual(response.status, 200)
|
|
self.assertEqual(payload["error"]["code"], -32700)
|
|
|
|
def test_brackets_inside_json_string_do_not_count_as_nesting(self) -> None:
|
|
status, response = self.request(
|
|
{
|
|
"jsonrpc": "2.0",
|
|
"id": 1,
|
|
"method": "initialize",
|
|
"params": {"text": "[{" * 500 + "}]" * 500},
|
|
}
|
|
)
|
|
self.assertEqual(status, 200)
|
|
self.assertIn("result", response)
|
|
|
|
def test_incomplete_body_times_out(self) -> None:
|
|
connection = socket.create_connection(
|
|
("127.0.0.1", self.server.server_port), timeout=2
|
|
)
|
|
connection.settimeout(2)
|
|
connection.sendall(
|
|
b"POST /mcp HTTP/1.1\r\n"
|
|
b"Host: 127.0.0.1\r\n"
|
|
b"Content-Type: application/json\r\n"
|
|
b"Content-Length: 100\r\n"
|
|
b"Connection: close\r\n\r\n"
|
|
b'{"jsonrpc":'
|
|
)
|
|
response = b""
|
|
while True:
|
|
chunk = connection.recv(4_096)
|
|
if not chunk:
|
|
break
|
|
response += chunk
|
|
connection.close()
|
|
|
|
self.assertIn(b"HTTP/1.1 408 Request Timeout", response)
|
|
self.assertIn(b'"error":"request_body_timeout"', response)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|