632 lines
25 KiB
Python
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())
|