Files
MemRelay/backend/tests/test_mcp.py
T

632 lines
25 KiB
Python

import asyncio
import json
import socket
import threading
import time
from datetime import UTC, datetime
from typing import Any
from unittest.mock import patch
import httpx
import uvicorn
from fastapi.testclient import TestClient
from fastmcp import Client
from fastmcp.client.transports import StreamableHttpTransport
from sqlalchemy import select
from memrelay.models import CurationChange, UsageProfile
from tests.conftest import csrf_headers
from tests.fakes import FakeBasicMemory
from tests.test_vault import FakeBwRunner, configure_vault, install_fake_vault
def create_mcp_token(client: TestClient, name: str, access_mode: str) -> dict[str, Any]:
response = client.post(
"/api/v1/tokens",
json={"name": name, "access_mode": access_mode},
headers=csrf_headers(client),
)
assert response.status_code == 201
return response.json()
def asgi_transport(app, token: str) -> StreamableHttpTransport: # type: ignore[no-untyped-def]
def factory(**kwargs) -> httpx.AsyncClient: # type: ignore[no-untyped-def]
headers = dict(kwargs.get("headers") or {})
headers["Authorization"] = f"Bearer {token}"
return httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
headers=headers,
timeout=kwargs.get("timeout"),
follow_redirects=True,
)
return StreamableHttpTransport(
"http://testserver/mcp", auth=token, httpx_client_factory=factory
)
def result_data(result) -> Any: # type: ignore[no-untyped-def]
data = result.data
return data.model_dump(mode="json") if hasattr(data, "model_dump") else data
def test_mcp_get_sse_stream_returns_405_not_404(initialized_client: TestClient) -> None:
# 客户端 initialize 后会 GET /mcp 尝试打开服务端主动 SSE 流;无状态模式不支持,
# 规范要求 405(客户端静默容忍)。若落入 SPA 兜底返回 404,客户端会把它当作
# 会话失效并在连续重试后禁用整个连接。
response = initialized_client.get("/mcp", headers={"Accept": "text/event-stream"})
assert response.status_code == 405
assert "POST" in response.headers.get("allow", "")
def test_mcp_requires_token(initialized_client: TestClient) -> None:
response = initialized_client.post(
"/mcp",
headers={"Accept": "application/json, text/event-stream"},
json={
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-06-18",
"capabilities": {},
"clientInfo": {"name": "test", "version": "1"},
},
},
)
assert response.status_code == 401
def test_streamable_http_mcp_operates_over_a_real_tcp_server(
initialized_client: TestClient,
) -> None:
client = initialized_client
token = create_mcp_token(client, "TCP MCP", "read_write")
listener = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
listener.bind(("127.0.0.1", 0))
listener.listen(128)
port = int(listener.getsockname()[1])
server = uvicorn.Server(
uvicorn.Config(
client.app,
host="127.0.0.1",
port=port,
log_level="error",
lifespan="off",
)
)
thread = threading.Thread(target=lambda: server.run(sockets=[listener]), daemon=True)
thread.start()
deadline = time.monotonic() + 10
while not server.started and thread.is_alive() and time.monotonic() < deadline:
time.sleep(0.01)
assert server.started
async def exercise() -> None:
transport = StreamableHttpTransport(
f"http://127.0.0.1:{port}/mcp",
auth=token["token"],
)
async with Client(transport) as mcp:
names = {tool.name for tool in await mcp.list_tools()}
assert {
"capabilities_get",
"curation_status",
"curation_history",
"curated_document_get",
"curation_run",
} <= names
capabilities = result_data(await mcp.call_tool("capabilities_get", {}))
assert capabilities["curation_enabled"] is False
status = result_data(await mcp.call_tool("curation_status", {}))
assert status["settings"]["enabled"] is False
history = result_data(await mcp.call_tool("curation_history", {"limit": 5}))
assert history == []
resources = await mcp.list_resources()
assert {str(item.uri) for item in resources} >= {
"memrelay://guide",
"memrelay://capabilities",
"memrelay://projects",
}
try:
asyncio.run(exercise())
finally:
server.should_exit = True
thread.join(timeout=10)
listener.close()
assert not thread.is_alive()
def test_mcp_tools_resources_and_scope_enforcement(initialized_client: TestClient) -> None:
client = initialized_client
client.app.state.basic_memory = FakeBasicMemory()
project = client.post(
"/api/v1/projects",
json={"name": "MCP Project"},
headers=csrf_headers(client),
).json()
stored_file = client.app.state.settings.files_dir / "mcp-metadata.txt"
stored_file.write_text("MCP metadata", encoding="utf-8")
client.post("/api/v1/files/scan", headers=csrf_headers(client))
file_record = client.get(
"/api/v1/files", params={"query": "mcp-metadata.txt"}
).json()[0]
client.patch(
f"/api/v1/files/{file_record['id']}",
json={
"description": "Original MCP description",
"version": "1.0",
"purpose": "reference",
"tags": ["mcp", "stable"],
"project_id": project["id"],
},
headers=csrf_headers(client),
)
read_token = create_mcp_token(client, "Reader", "read_only")
write_token = create_mcp_token(client, "Writer", "read_write")
verified = asyncio.run(client.app.state.mcp_server.auth.verify_token(read_token["token"]))
assert verified is not None
assert verified.scopes == ["read"]
authenticated_probe = client.post(
"/mcp",
headers={
"Accept": "application/json, text/event-stream",
"Authorization": f"Bearer {read_token['token']}",
},
json={
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-06-18",
"capabilities": {},
"clientInfo": {"name": "test", "version": "1"},
},
},
)
assert authenticated_probe.status_code != 401, authenticated_probe.text
async def exercise() -> None:
async with Client(asgi_transport(client.app, read_token["token"])) as reader:
names = {tool.name for tool in await reader.list_tools()}
expected_read = {
"capabilities_get",
"project_resolve",
"project_list",
"project_export_prepare",
"context_get",
"memory_list",
"memory_search",
"memory_get",
"checkpoint_get",
"secret_capabilities",
"secret_list",
"secret_item_get",
"secret_search",
"secret_get",
"secret_totp",
"file_search",
"file_list",
"file_get",
"file_upload_status",
"curation_status",
"curation_history",
"curated_document_list",
"curated_document_get",
"curated_document_history",
"curated_document_diff",
"curation_settings_get",
"curation_schedule_list",
"curation_schedule_preview",
"curation_model_connection_status",
"curation_model_list",
"usage_profile_list",
"git_connection_get",
"git_mode_migration_preview",
"memory_repository_list",
"memory_repository_status",
"memory_repository_history",
"memory_repository_diff",
"memory_version_get",
"memory_repository_remote_inspect",
"dashboard_get",
"system_settings_get",
"token_list",
}
write_tools = {
"project_create",
"project_resolve_or_create",
"project_update",
"project_archive",
"project_delete",
"memory_save",
"memory_archive",
"memory_delete",
"checkpoint_save",
"checkpoint_delete",
"secret_sync",
"secret_resolve",
"secret_save",
"secret_delete",
"file_directory_create",
"file_upload_prepare",
"file_metadata_update",
"file_move",
"file_delete",
"file_scan",
"curation_run",
"curation_source_submit",
"curated_document_update",
"curated_document_revert",
"curation_cancel",
"curation_retry",
"curation_settings_update",
"curation_schedule_save",
"curation_schedule_delete",
"curation_schedule_reset_defaults",
"curation_model_connection_save",
"curation_model_connection_test",
"usage_profile_save",
"usage_profile_token_bind",
"usage_profile_delete",
"git_connection_save",
"git_mode_migration_execute",
"git_connection_test",
"git_connection_delete",
"memory_repository_create",
"memory_repository_sync",
"memory_version_restore",
"memory_snapshot_restore",
"memory_remote_import",
"memory_remote_bootstrap",
"memory_repository_unbind",
"memory_repository_archive_purge",
"system_settings_update",
"token_create",
"token_reveal",
"token_revoke",
"token_delete",
}
assert expected_read <= names
assert names.isdisjoint(write_tools)
capabilities = result_data(await reader.call_tool("capabilities_get", {}))
assert capabilities["memory"] is True
assert capabilities["vault"] is False
assert capabilities["semantic_search"]["status"] == "available"
dashboard = result_data(await reader.call_tool("dashboard_get", {}))
assert "curation" in dashboard
assert "memory_git" in dashboard
assert "database" in dashboard
resolved = result_data(
await reader.call_tool("project_resolve", {"name": "MCP Project"})
)
assert resolved["id"] == project["id"]
read_only_directory = "D:/ReadOnly/MCP Project"
resolved_with_directory = result_data(
await reader.call_tool(
"project_resolve",
{"name": "MCP Project", "directory": read_only_directory},
)
)
assert resolved_with_directory["id"] == project["id"]
assert read_only_directory.casefold() not in resolved_with_directory["aliases"]
try:
await reader.call_tool(
"memory_save",
{
"request_id": "denied",
"scope": "global",
"memory_type": "rule",
"title": "Denied",
"content": "Must not write.",
},
)
except Exception as exc:
message = str(exc).casefold()
assert "authorization" in message or "scope" in message or "unknown tool" in message
else:
raise AssertionError("Read-only token unexpectedly called a write tool")
guide = await reader.read_resource("memrelay://guide")
assert "不把猜测" in guide[0].text
assert "project_resolve_or_create" in guide[0].text
assert "AI 自动整理" not in guide[0].text
assert "记忆版本历史" not in guide[0].text
status_with_timestamp = client.app.state.curation.model_status()
status_with_timestamp["last_check_at"] = datetime.now(UTC)
with patch.object(
client.app.state.curation,
"model_status",
return_value=status_with_timestamp,
):
capabilities_resource = await reader.read_resource("memrelay://capabilities")
resource_data = json.loads(capabilities_resource[0].text)
assert resource_data["semantic_search"]["status"] == "available"
assert resource_data["tools"] == capabilities["tools"]
assert resource_data["curation"]["model"]["last_check_at"].endswith("+00:00")
async with Client(asgi_transport(client.app, write_token["token"])) as writer:
writer_names = {tool.name for tool in await writer.list_tools()}
assert write_tools <= writer_names
saved = result_data(
await writer.call_tool(
"memory_save",
{
"request_id": "mcp-save-1",
"scope": "project",
"project_id": project["id"],
"memory_type": "decision",
"title": "MCP contract",
"content": "Use Streamable HTTP.",
},
)
)
assert saved["revision"] == 1
context = result_data(
await writer.call_tool(
"context_get", {"project_id": project["id"], "query": "HTTP"}
)
)
assert saved["id"] in context["included_memory_ids"]
upload = result_data(
await writer.call_tool("file_upload_prepare", {"path": "mcp/file.txt", "size": 4})
)
assert upload["upload_url"].startswith("http://testserver/")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=client.app),
base_url="http://testserver",
) as uploader:
uploaded = await uploader.put(upload["upload_url"], content=b"mcp!")
assert uploaded.status_code == 200, uploaded.text
partial_metadata = result_data(
await writer.call_tool(
"file_metadata_update",
{
"file_id": file_record["id"],
"description": "Updated through MCP",
},
)
)
assert partial_metadata["description"] == "Updated through MCP"
assert partial_metadata["version"] == "1.0"
assert partial_metadata["purpose"] == "reference"
assert partial_metadata["tags"] == ["mcp", "stable"]
assert partial_metadata["project_id"] == project["id"]
cleared_metadata = result_data(
await writer.call_tool(
"file_metadata_update",
{
"file_id": file_record["id"],
"clear_version": True,
"clear_project": True,
"tags": [],
},
)
)
assert cleared_metadata["version"] is None
assert cleared_metadata["project_id"] is None
assert cleared_metadata["purpose"] == "reference"
assert cleared_metadata["tags"] == []
created_project = result_data(
await writer.call_tool(
"project_create",
{
"name": "Created by MCP",
"git_remote": "https://example.com/team/repo.git",
"aliases": ["C:/code/repo"],
},
)
)
updated_project = result_data(
await writer.call_tool(
"project_update",
{"project_id": created_project["id"], "name": "Updated by MCP"},
)
)
assert updated_project["name"] == "Updated by MCP"
resolved_or_created = result_data(
await writer.call_tool(
"project_resolve_or_create",
{
"git_remote": "git@example.test:team/created-from-workflow.git",
"directory": "E:/Projects/CreatedFromWorkflow",
},
)
)
assert resolved_or_created["created"] is True
assert resolved_or_created["project"]["name"] == "created-from-workflow"
assert "e:/projects/createdfromworkflow" in resolved_or_created["project"]["aliases"]
resolved_again = result_data(
await writer.call_tool(
"project_resolve_or_create",
{
"name": "A different local folder name",
"git_remote": "https://example.test/team/created-from-workflow.git",
"directory": "F:/Projects/CreatedFromWorkflow",
},
)
)
assert resolved_again["created"] is False
assert resolved_again["project"]["id"] == resolved_or_created["project"]["id"]
assert "f:/projects/createdfromworkflow" in resolved_again["project"]["aliases"]
checkpoint = result_data(
await writer.call_tool(
"checkpoint_save",
{
"project_id": project["id"],
"request_id": "mcp-checkpoint-1",
"content": "## Completed\n\nVerified the complete MCP workflow.",
},
)
)
export = result_data(
await writer.call_tool("project_export_prepare", {"project_id": project["id"]})
)
assert export["url"].startswith("http://testserver/")
assert "aidocs/" in export["powershell_command"]
listed_memories = result_data(
await writer.call_tool("memory_list", {"project_id": project["id"]})
)
assert saved["id"] in {item["id"] for item in listed_memories}
# Checkpoints leave the pending bucket and get their own lifecycle filter.
assert checkpoint["id"] not in {item["id"] for item in listed_memories}
listed_checkpoints = result_data(
await writer.call_tool(
"memory_list",
{"project_id": project["id"], "lifecycle": "checkpoint"},
)
)
assert checkpoint["id"] in {item["id"] for item in listed_checkpoints}
assert result_data(await writer.call_tool("dashboard_get", {}))["checkpoints"] == 1
await writer.call_tool("checkpoint_delete", {"checkpoint_id": checkpoint["id"]})
generated_token = result_data(
await writer.call_tool(
"token_create", {"name": "Generated token", "access_mode": "read_only"}
)
)
assert generated_token["token"].startswith("mcp_")
profile = result_data(
await writer.call_tool(
"usage_profile_save",
{"name": "MCP profile", "description": "Bound through MCP"},
)
)
bound = result_data(
await writer.call_tool(
"usage_profile_token_bind",
{"profile_id": profile["id"], "token_ids": [generated_token["id"]]},
)
)
assert bound["token_count"] == 1
schedules = result_data(await writer.call_tool("curation_schedule_reset_defaults", {}))
assert {item["trigger_type"] for item in schedules} == {
"workspace_quiet",
"workspace_fallback",
"global_curation",
}
await writer.call_tool(
"usage_profile_token_bind",
{"profile_id": profile["id"], "token_ids": []},
)
await writer.call_tool("usage_profile_delete", {"profile_id": profile["id"]})
await writer.call_tool("token_delete", {"token_id": generated_token["id"]})
await writer.call_tool(
"project_archive", {"project_id": created_project["id"], "archived": True}
)
await writer.call_tool("project_delete", {"project_id": created_project["id"]})
asyncio.run(exercise())
with client.app.state.session_factory() as db:
shared = db.scalar(select(UsageProfile).where(UsageProfile.is_shared.is_(True)))
agent_changes = db.scalars(
select(CurationChange).where(
CurationChange.project_id == project["id"],
CurationChange.origin == "agent",
)
).all()
assert shared is not None and agent_changes
assert {item.usage_profile_id for item in agent_changes} == {shared.id}
def test_mcp_password_vault_management(initialized_client: TestClient) -> None:
client = initialized_client
runner = FakeBwRunner()
install_fake_vault(client, runner)
configure_vault(client)
token = create_mcp_token(client, "Vault manager", "read_write")
async def exercise() -> None:
async with Client(asgi_transport(client.app, token["token"])) as mcp:
listed = result_data(await mcp.call_tool("secret_list", {}))
assert len(listed["results"]) == 2
assert "git-password" not in str(listed)
detail = result_data(await mcp.call_tool("secret_item_get", {"item_id": "item-1"}))
assert detail["password"] == "git-password"
saved = result_data(
await mcp.call_tool(
"secret_save",
{
"name": "AI Created Credential",
"username": "automation",
"password": "generated-password",
"uris": ["https://automation.example.com"],
"fields": [
{"name": "api_token", "value": "generated-token", "hidden": True}
],
"lookup_key": "automation-service",
},
)
)
assert saved["created"] is True
assert saved["item"]["password"] == "generated-password"
deleted = result_data(
await mcp.call_tool("secret_delete", {"item_id": saved["item"]["id"]})
)
assert deleted == {"deleted": True}
asyncio.run(exercise())
def test_all_profiles_token_includes_profiles_created_later(
initialized_client: TestClient,
) -> None:
client = initialized_client
client.app.state.basic_memory = FakeBasicMemory()
token_response = client.post(
"/api/v1/tokens",
json={"name": "All profiles", "access_mode": "read_write", "all_profiles": True},
headers=csrf_headers(client),
)
assert token_response.status_code == 201
token = token_response.json()
future_profile = client.post(
"/api/v1/curation/profiles",
json={"name": "Created after token"},
headers=csrf_headers(client),
).json()
future_memory = client.post(
"/api/v1/memories",
json={
"request_id": "future-profile-memory",
"scope": "global",
"usage_profile_id": future_profile["id"],
"memory_type": "preference",
"title": "Future profile preference",
"content": "This profile was created after the MCP token.",
"tags": [],
},
headers=csrf_headers(client),
).json()
async def exercise() -> None:
async with Client(asgi_transport(client.app, token["token"])) as mcp:
listed = result_data(await mcp.call_tool("memory_list", {"scope": "global"}))
assert future_memory["id"] in {item["id"] for item in listed}
saved = result_data(
await mcp.call_tool(
"memory_save",
{
"request_id": "all-profile-targeted-write",
"scope": "global",
"usage_profile_id": future_profile["id"],
"memory_type": "task_list",
"title": "Future profile tasks",
"content": "- [ ] Verify all-profile access",
},
)
)
assert saved["usage_profile_id"] == future_profile["id"]
assert saved["memory_type"] == "task_list"
asyncio.run(exercise())