| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183 |
- """
- tests/test_mcp_registry.py — MCP server 注册中心 P1(docs/MCP_SKILL_DESIGN.md)。
- 覆盖:
- - 注册中心 CRUD(文件锁 json 权威源)、校验、脱敏。
- - stdio argv allowlist(评审#16)。
- - 风险分级 medium 默认 + 破坏性动词升 high(评审#15)。
- - inject_into_config 只注入启用的 http server(评审#14 禁用即不可用)。
- - ToolGateway MCP 风险修复。
- - API 端点 loopback 门 + CRUD + probe。
- """
- from __future__ import annotations
- import json
- import os
- import pytest
- os.environ.setdefault("AGENTPAAS_DATABASE_URL", "sqlite:///:memory:")
- os.environ.setdefault("AGENTPAAS_TESTING", "1")
- @pytest.fixture()
- def reg(tmp_path, monkeypatch):
- """隔离 MCP_FILE 到 tmp,避免动到真实 ~/.agentpaas/mcp_servers.json。"""
- from agentpaas.engine import mcp_registry as r
- monkeypatch.setattr(r, "MCP_FILE", str(tmp_path / "mcp_servers.json"))
- return r
- # ── CRUD + 校验 ──
- def test_upsert_list_get_remove(reg):
- reg.upsert_server({
- "id": "arxiv", "name": "arXiv", "transport": "http",
- "url": "http://localhost:9000/mcp",
- "auth": {"kind": "env_bearer", "env_key": "ARXIV_KEY"},
- })
- lst = reg.list_servers()
- assert len(lst) == 1 and lst[0]["id"] == "arxiv"
- got = reg.get_server("arxiv")
- assert got["url"] == "http://localhost:9000/mcp"
- assert got["auth"]["env_key"] == "ARXIV_KEY" # 变量名保留
- assert "value" not in got["auth"] # 评审#13: 永不出现凭证值
- assert reg.set_enabled("arxiv", False)
- assert reg.get_server("arxiv")["enabled"] is False
- assert reg.remove_server("arxiv")
- assert reg.get_server("arxiv") is None
- def test_validate_rejects_bad(reg):
- assert reg.validate_server({"id": "x", "transport": "http"}) # 缺 url
- assert reg.validate_server({"id": "x y", "transport": "http", "url": "u"}) # 非法 id
- assert reg.validate_server({"id": "x", "transport": "ftp"}) # 非法 transport
- def test_stdio_argv_allowlist(reg):
- # rm 不在 allowlist → 拒绝
- errs = reg.validate_server({"id": "evil", "transport": "stdio",
- "command_argv": ["rm", "-rf", "/"]})
- assert any("allowlist" in e for e in errs)
- # npx 在 allowlist → 通过
- assert reg.validate_server({"id": "ok", "transport": "stdio",
- "command_argv": ["npx", "some-mcp"]}) == []
- # 显式加白
- assert reg.validate_server({"id": "ok2", "transport": "stdio",
- "command_argv": ["mybin"]},
- argv_allowlist=reg.DEFAULT_ARGV_ALLOWLIST | {"mybin"}) == []
- def test_risk_classification(reg):
- ov = {"special_tool": "low"}
- assert reg.classify_mcp_tool("search_papers", {}) == "medium" # 默认 medium
- assert reg.classify_mcp_tool("delete_paper", {}) == "high" # 动词升 high
- assert reg.classify_mcp_tool("send_email", {}) == "high"
- assert reg.classify_mcp_tool("special_tool", ov) == "low" # override 优先
- # ── 运行时注入桥(评审#14)──
- def test_inject_only_enabled_http(reg):
- reg.upsert_server({"id": "on", "transport": "http", "url": "http://h/mcp",
- "auth": {"kind": "none"}})
- reg.upsert_server({"id": "off", "transport": "http", "url": "http://h2/mcp",
- "auth": {"kind": "none"}, "enabled": False})
- reg.upsert_server({"id": "stdio1", "transport": "stdio",
- "command_argv": ["npx", "x"]})
- cfg = reg.inject_into_config({"name": "a", "type": "simple"})
- nodes = cfg["app"]["mcp"]["custom"]["nodes"]
- assert "on" in nodes # 启用的 http 注入
- assert "off" not in nodes # 禁用的不注入(评审#14)
- assert "stdio1" not in nodes # P1 运行时桥不接 stdio
- def test_inject_credential_from_env(reg, monkeypatch):
- monkeypatch.setenv("MY_MCP_KEY", "secret123")
- reg.upsert_server({"id": "s", "transport": "http", "url": "http://h/mcp",
- "auth": {"kind": "env_bearer", "env_key": "MY_MCP_KEY"}})
- cfg = reg.inject_into_config({"name": "a"})
- hdr = cfg["app"]["mcp"]["custom"]["nodes"]["s"]["headers"]
- assert hdr["Authorization"] == "Bearer secret123"
- def test_inject_no_servers_passthrough(reg):
- cfg = {"name": "a", "type": "simple"}
- assert reg.inject_into_config(cfg) == cfg
- # ── ToolGateway MCP 风险修复(评审#15)──
- def test_gateway_mcp_risk():
- from lambdagent.tool_gateway import classify_tool_call, RiskLevel
- assert classify_tool_call("mcp_search", "")[0] == RiskLevel.MEDIUM
- assert classify_tool_call("arxiv.search", "")[0] == RiskLevel.MEDIUM
- assert classify_tool_call("mcp_delete_file", "")[0] == RiskLevel.HIGH
- assert classify_tool_call("github.create_issue", "")[0] == RiskLevel.HIGH
- # ── API 端点 ──
- @pytest.fixture()
- def api(tmp_path, monkeypatch):
- import secrets
- from fastapi.testclient import TestClient
- from agentpaas.api.app import app
- from agentpaas.api.middleware.auth import hash_key
- from agentpaas.db.models import Database, gen_id, now_utc
- from agentpaas.engine import mcp_registry as r
- import agentpaas.db.session as _session_mod
- monkeypatch.setattr(r, "MCP_FILE", str(tmp_path / "mcp.json"))
- prev = _session_mod._db
- _session_mod._db = Database("sqlite:///:memory:")
- db = _session_mod._db
- tid, uid = gen_id("tn_"), gen_id("usr_")
- raw = f"ap_{secrets.token_hex(16)}"
- now = now_utc()
- db.execute("INSERT INTO tenants (id,name,plan,status,created_at) VALUES (?,'t','free','active',?)", (tid, now))
- db.execute("INSERT INTO users (id,tenant_id,email,role,created_at) VALUES (?,?,'','admin',?)", (uid, tid, now))
- db.execute("INSERT INTO api_keys (id,tenant_id,user_id,key_hash,key_prefix,name,scopes,rate_limit,status,created_at) "
- "VALUES (?,?,?,?,?,'t',?,600,'active',?)",
- (gen_id("key_"), tid, uid, hash_key(raw), raw[:8], json.dumps(["agents:*"]), now))
- db.commit()
- with TestClient(app) as c:
- yield c, raw
- _session_mod._db = prev
- def _auth(k):
- return {"Authorization": f"Bearer {k}"}
- def test_api_crud_and_list(api):
- c, key = api
- r = c.post("/api/v1/mcp", headers=_auth(key),
- json={"id": "arxiv", "name": "arXiv", "transport": "http",
- "url": "http://localhost:9000/mcp"})
- assert r.status_code == 201, r.text[:200]
- lst = c.get("/api/v1/mcp", headers=_auth(key)).json()
- assert any(s["id"] == "arxiv" for s in lst["servers"])
- # 启停
- assert c.put("/api/v1/mcp/arxiv", headers=_auth(key), json={"enabled": False}).status_code == 200
- assert c.get("/api/v1/mcp", headers=_auth(key)).json()["servers"][0]["enabled"] is False
- assert c.delete("/api/v1/mcp/arxiv", headers=_auth(key)).status_code == 200
- def test_api_add_rejects_bad_argv(api):
- c, key = api
- r = c.post("/api/v1/mcp", headers=_auth(key),
- json={"id": "evil", "transport": "stdio", "command_argv": ["rm", "-rf", "/"]})
- assert r.status_code == 400 and "allowlist" in r.json()["detail"]
- def test_api_probe_unreachable(api):
- c, key = api
- c.post("/api/v1/mcp", headers=_auth(key),
- json={"id": "dead", "transport": "http", "url": "http://127.0.0.1:1/mcp"})
- r = c.post("/api/v1/mcp/dead/probe", headers=_auth(key))
- assert r.status_code == 200 and r.json()["status"] == "unreachable"
- def test_api_probe_unknown_404(api):
- c, key = api
- assert c.post("/api/v1/mcp/nope/probe", headers=_auth(key)).status_code == 404
|