添加MCP客户端超时配置和连接管理改进

添加了MCP客户端的超时配置功能,包括启动超时和工具调用超时参数。
改进了HTTP客户端和标准IO客户端的超时处理机制,确保请求在指定时间内完成或取消。
增加了对MCP服务器数量的限制,防止配置过多服务器导致系统不稳定。
增强了错误处理机制,当连接异常时能够正确清理资源并移除桥接主机。
添加了对大型MCP消息的大小验证,防止过大的请求导致系统问题。
优化了密钥更改后的处理流程,确保在修改密钥时停用服务器并要求重新测试。
This commit is contained in:
2026-09-03 16:16:58 +08:00
parent 7d5f4023a9
commit 2f7066aa92
11 changed files with 504 additions and 72 deletions
+49 -14
View File
@@ -146,7 +146,7 @@ class McpStdioClient:
timeout_code: str, timeout_code: str,
response_error_code: str = "MCP_TOOL_CALL_FAILED", response_error_code: str = "MCP_TOOL_CALL_FAILED",
) -> dict[str, Any]: ) -> dict[str, Any]:
request_id, pending = self.begin_request(method, params) request_id, pending = self.begin_request(method, params, timeout=timeout)
return self.wait_response( return self.wait_response(
request_id, request_id,
pending, pending,
@@ -156,7 +156,7 @@ class McpStdioClient:
) )
def begin_request( def begin_request(
self, method: str, params: dict[str, Any] self, method: str, params: dict[str, Any], *, timeout: float | None = None
) -> tuple[int, _PendingRequest]: ) -> tuple[int, _PendingRequest]:
self._ensure_running() self._ensure_running()
with self._pending_lock: with self._pending_lock:
@@ -388,6 +388,7 @@ class McpHttpClient:
url: str, url: str,
*, *,
headers: dict[str, str], headers: dict[str, str],
startup_timeout_seconds: float = 15,
on_seen: Callable[[], None], on_seen: Callable[[], None],
on_broken: Callable[[str], None], on_broken: Callable[[str], None],
on_tools_changed: Callable[[], None], on_tools_changed: Callable[[], None],
@@ -407,6 +408,7 @@ class McpHttpClient:
self._stream_started = False self._stream_started = False
self._last_event_id: str | None = None self._last_event_id: str | None = None
self._stop_event = threading.Event() self._stop_event = threading.Event()
self._startup_timeout_seconds = startup_timeout_seconds
def start(self) -> None: def start(self) -> None:
return return
@@ -429,7 +431,7 @@ class McpHttpClient:
timeout_code: str, timeout_code: str,
response_error_code: str = "MCP_TOOL_CALL_FAILED", response_error_code: str = "MCP_TOOL_CALL_FAILED",
) -> dict[str, Any]: ) -> dict[str, Any]:
request_id, pending = self.begin_request(method, params) request_id, pending = self.begin_request(method, params, timeout=timeout)
return self.wait_response( return self.wait_response(
request_id, request_id,
pending, pending,
@@ -439,7 +441,7 @@ class McpHttpClient:
) )
def begin_request( def begin_request(
self, method: str, params: dict[str, Any] self, method: str, params: dict[str, Any], *, timeout: float | None = None
) -> tuple[int, _PendingRequest]: ) -> tuple[int, _PendingRequest]:
with self._pending_lock: with self._pending_lock:
request_id = self._next_id request_id = self._next_id
@@ -454,7 +456,7 @@ class McpHttpClient:
} }
threading.Thread( threading.Thread(
target=self._dispatch_request, target=self._dispatch_request,
args=(request_id, message), args=(request_id, message, timeout),
daemon=True, daemon=True,
).start() ).start()
return request_id, pending return request_id, pending
@@ -526,7 +528,10 @@ class McpHttpClient:
if self._session_id: if self._session_id:
try: try:
request = self._client.build_request( request = self._client.build_request(
"DELETE", self.url, headers=self._request_headers() "DELETE",
self.url,
headers=self._request_headers(),
timeout=min(self._startup_timeout_seconds, 5),
) )
response = self._client.send(request, stream=True) response = self._client.send(request, stream=True)
response.close() response.close()
@@ -539,9 +544,14 @@ class McpHttpClient:
) )
) )
def _dispatch_request(self, request_id: int, message: dict[str, Any]) -> None: def _dispatch_request(
self,
request_id: int,
message: dict[str, Any],
timeout: float | None,
) -> None:
try: try:
response = self._post(message, timeout=None) response = self._post(message, timeout=timeout)
try: try:
self._capture_session(response) self._capture_session(response)
content_type = response.headers.get("content-type", "").lower() content_type = response.headers.get("content-type", "").lower()
@@ -588,7 +598,12 @@ class McpHttpClient:
def _post_notification(self, message: dict[str, Any]) -> None: def _post_notification(self, message: dict[str, Any]) -> None:
try: try:
response = self._post(message, timeout=10) timeout = (
self._startup_timeout_seconds
if message.get("method") == "notifications/initialized"
else 10
)
response = self._post(message, timeout=timeout)
except httpx.HTTPError as exc: except httpx.HTTPError as exc:
raise McpBridgeError( raise McpBridgeError(
"MCP_HTTP_REQUEST_FAILED", "MCP_HTTP_REQUEST_FAILED",
@@ -616,6 +631,7 @@ class McpHttpClient:
self.url, self.url,
content=encoded.encode("utf-8"), content=encoded.encode("utf-8"),
headers=self._request_headers(), headers=self._request_headers(),
timeout=timeout,
) )
return self._client.send(request, stream=True) return self._client.send(request, stream=True)
@@ -714,7 +730,7 @@ class McpLegacySseClient(McpHttpClient):
def start(self) -> None: def start(self) -> None:
threading.Thread(target=self._event_loop, daemon=True).start() threading.Thread(target=self._event_loop, daemon=True).start()
try: try:
endpoint = self._endpoint_ready.get(timeout=15) endpoint = self._endpoint_ready.get(timeout=self._startup_timeout_seconds)
except queue.Empty as exc: except queue.Empty as exc:
raise McpBridgeError( raise McpBridgeError(
"MCP_INITIALIZE_FAILED", "MCP_INITIALIZE_FAILED",
@@ -730,9 +746,14 @@ class McpLegacySseClient(McpHttpClient):
return return
def _dispatch_request(self, request_id: int, message: dict[str, Any]) -> None: def _dispatch_request(
self,
request_id: int,
message: dict[str, Any],
timeout: float | None,
) -> None:
try: try:
response = self._post(message, timeout=10) response = self._post(message, timeout=timeout)
try: try:
if response.status_code not in {200, 202, 204}: if response.status_code not in {200, 202, 204}:
raise McpBridgeError( raise McpBridgeError(
@@ -761,6 +782,8 @@ class McpLegacySseClient(McpHttpClient):
"MCP_INITIALIZE_FAILED", "Legacy MCP endpoint is not ready." "MCP_INITIALIZE_FAILED", "Legacy MCP endpoint is not ready."
) )
encoded = json.dumps(message, ensure_ascii=False, separators=(",", ":")) encoded = json.dumps(message, ensure_ascii=False, separators=(",", ":"))
if len(encoded.encode("utf-8")) > MAX_MCP_MESSAGE_BYTES:
raise McpBridgeError("MCP_TOOL_CALL_FAILED", "MCP request is too large.")
request = self._client.build_request( request = self._client.build_request(
"POST", "POST",
self._endpoint, self._endpoint,
@@ -770,6 +793,7 @@ class McpLegacySseClient(McpHttpClient):
"Accept": "application/json, text/event-stream", "Accept": "application/json, text/event-stream",
"Content-Type": "application/json", "Content-Type": "application/json",
}, },
timeout=timeout,
) )
return self._client.send(request, stream=True) return self._client.send(request, stream=True)
@@ -801,6 +825,8 @@ class McpLegacySseClient(McpHttpClient):
self._endpoint = endpoint self._endpoint = endpoint
continue continue
self._handle_message(_json_rpc_message(data)) self._handle_message(_json_rpc_message(data))
if not self._stopping:
self.on_broken("Legacy MCP SSE stream ended unexpectedly.")
except (McpBridgeError, httpx.HTTPError) as exc: except (McpBridgeError, httpx.HTTPError) as exc:
if self._endpoint is None: if self._endpoint is None:
self._endpoint_ready.put(exc) self._endpoint_ready.put(exc)
@@ -827,7 +853,7 @@ class _McpClient(Protocol):
response_error_code: str = "MCP_TOOL_CALL_FAILED", response_error_code: str = "MCP_TOOL_CALL_FAILED",
) -> dict[str, Any]: ... ) -> dict[str, Any]: ...
def begin_request( def begin_request(
self, method: str, params: dict[str, Any] self, method: str, params: dict[str, Any], *, timeout: float | None = None
) -> tuple[int, _PendingRequest]: ... ) -> tuple[int, _PendingRequest]: ...
def wait_response( def wait_response(
self, self,
@@ -931,6 +957,7 @@ class McpBridge:
client = client_type( client = client_type(
url, url,
headers=headers or {}, headers=headers or {},
startup_timeout_seconds=backend.startup_timeout_seconds,
on_seen=seen, on_seen=seen,
on_broken=broken, on_broken=broken,
on_tools_changed=tools_changed, on_tools_changed=tools_changed,
@@ -989,6 +1016,12 @@ class McpBridge:
discovered = self._discover_tools( discovered = self._discover_tools(
plugin_id, client, backend, declared_permissions, tool_source plugin_id, client, backend, declared_permissions, tool_source
) )
if status.status == PluginHostState.unhealthy:
raise McpBridgeError(
"PLUGIN_HOST_UNAVAILABLE",
status.error or "MCP event stream became unavailable during startup.",
status_code=503,
)
status.status = PluginHostState.ready status.status = PluginHostState.ready
status.tools_count = len(discovered) status.tools_count = len(discovered)
status.last_seen_at = datetime.now(UTC) status.last_seen_at = datetime.now(UTC)
@@ -1019,7 +1052,9 @@ class McpBridge:
) -> Any: ) -> Any:
host = self._host(plugin_id) host = self._host(plugin_id)
rpc_id, pending = host.client.begin_request( rpc_id, pending = host.client.begin_request(
"tools/call", {"name": remote_name, "arguments": arguments} "tools/call",
{"name": remote_name, "arguments": arguments},
timeout=host.backend.tool_timeout_seconds,
) )
call_key = (plugin_id, request_id) call_key = (plugin_id, request_id)
with self._lock: with self._lock:
+109 -28
View File
@@ -13,12 +13,13 @@ from typing import Any
from urllib.parse import urlsplit from urllib.parse import urlsplit
from uuid import uuid4 from uuid import uuid4
from pydantic import BaseModel, ConfigDict, create_model from pydantic import BaseModel, ConfigDict, Field, ValidationError, create_model
from app.agent.permissions import KNOWN_PERMISSIONS from app.agent.permissions import KNOWN_PERMISSIONS
from app.agent.tools import ToolExecutionContext, ToolRegistry from app.agent.tools import ToolExecutionContext, ToolRegistry
from app.contracts import ( from app.contracts import (
McpServer, McpServer,
McpServerConfig,
McpServerCreateRequest, McpServerCreateRequest,
McpServerSecretStatus, McpServerSecretStatus,
McpServerTransport, McpServerTransport,
@@ -40,6 +41,23 @@ _RESERVED_HEADERS = {
"mcp-protocol-version", "mcp-protocol-version",
"mcp-session-id", "mcp-session-id",
} }
_SERVER_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]{0,63}$")
_MAX_MCP_SERVERS = 256
class _McpServerRecord(McpServerConfig):
"""Validated on-disk representation with defaults for older C.1 records."""
version: int = Field(default=1, ge=1)
enabled: bool = False
approved_digest: str | None = Field(
default=None, min_length=64, max_length=64, pattern=r"^[0-9a-f]{64}$"
)
tested_digest: str | None = Field(
default=None, min_length=64, max_length=64, pattern=r"^[0-9a-f]{64}$"
)
last_tested_at: datetime | None = None
last_test_succeeded: bool | None = None
class McpRegistryError(RuntimeError): class McpRegistryError(RuntimeError):
@@ -105,6 +123,13 @@ class McpServerRegistry:
@_serialized_lifecycle @_serialized_lifecycle
def create(self, request: McpServerCreateRequest) -> McpServer: def create(self, request: McpServerCreateRequest) -> McpServer:
self._validate(request) self._validate(request)
with self._lock:
if len(self._records) >= _MAX_MCP_SERVERS:
raise McpRegistryError(
"MCP_SERVER_LIMIT_REACHED",
f"At most {_MAX_MCP_SERVERS} MCP servers can be configured.",
status_code=409,
)
server_id = uuid4().hex[:12] server_id = uuid4().hex[:12]
record = request.model_dump(mode="json") record = request.model_dump(mode="json")
record["name"] = request.name.strip() record["name"] = request.name.strip()
@@ -137,8 +162,8 @@ class McpServerRegistry:
self.disable(server_id) self.disable(server_id)
with self._lock: with self._lock:
previous = self._record(server_id) previous = self._record(server_id)
removed = [ removed_secret_ids = [
(kind, key) self._secret_id(server_id, key, kind)
for kind, old_keys, new_keys in ( for kind, old_keys, new_keys in (
( (
"environment", "environment",
@@ -153,6 +178,13 @@ class McpServerRegistry:
) )
for key in set(old_keys) - set(new_keys) for key in set(old_keys) - set(new_keys)
] ]
try:
self.credentials.delete_many(removed_secret_ids)
except CredentialStoreError as exc:
raise McpRegistryError(
"MCP_SECRET_STORE_ERROR", str(exc), status_code=500
) from exc
with self._lock:
record = request.model_dump(mode="json", exclude={"version"}) record = request.model_dump(mode="json", exclude={"version"})
record["name"] = request.name.strip() record["name"] = request.name.strip()
record["command"] = request.command.strip() if request.command else None record["command"] = request.command.strip() if request.command else None
@@ -170,13 +202,6 @@ class McpServerRegistry:
self._records = updated self._records = updated
self._last_status.pop(server_id, None) self._last_status.pop(server_id, None)
self._summaries.pop(server_id, None) self._summaries.pop(server_id, None)
for kind, key in removed:
try:
self.credentials.delete(self._secret_id(server_id, key, kind))
except CredentialStoreError as exc:
raise McpRegistryError(
"MCP_SECRET_STORE_ERROR", str(exc), status_code=500
) from exc
return self.get(server_id) return self.get(server_id)
@_serialized_lifecycle @_serialized_lifecycle
@@ -192,18 +217,19 @@ class McpServerRegistry:
) )
for key in keys for key in keys
] ]
updated = dict(self._records)
del updated[server_id]
self._write(updated)
self._records = updated
self._last_status.pop(server_id, None)
self._summaries.pop(server_id, None)
try: try:
self.credentials.delete_many(secret_ids) self.credentials.delete_many(secret_ids)
except CredentialStoreError as exc: except CredentialStoreError as exc:
raise McpRegistryError( raise McpRegistryError(
"MCP_SECRET_STORE_ERROR", str(exc), status_code=500 "MCP_SECRET_STORE_ERROR", str(exc), status_code=500
) from exc ) from exc
with self._lock:
updated = dict(self._records)
del updated[server_id]
self._write(updated)
self._records = updated
self._last_status.pop(server_id, None)
self._summaries.pop(server_id, None)
self.bridge.remove(self._host_id(server_id)) self.bridge.remove(self._host_id(server_id))
@_serialized_lifecycle @_serialized_lifecycle
@@ -236,6 +262,9 @@ class McpServerRegistry:
"MCP_SECRET_NOT_DECLARED", "MCP_SECRET_NOT_DECLARED",
"Secret environment key is not declared in this server configuration.", "Secret environment key is not declared in this server configuration.",
) )
if record.get("enabled"):
self.disable(server_id)
self._invalidate_test(server_id)
try: try:
self.credentials.put(self._secret_id(server_id, key, kind), secret) self.credentials.put(self._secret_id(server_id, key, kind), secret)
except CredentialStoreError as exc: except CredentialStoreError as exc:
@@ -254,6 +283,9 @@ class McpServerRegistry:
"MCP_SECRET_NOT_DECLARED", "MCP_SECRET_NOT_DECLARED",
"Secret environment key is not declared in this server configuration.", "Secret environment key is not declared in this server configuration.",
) )
if record.get("enabled"):
self.disable(server_id)
self._invalidate_test(server_id)
try: try:
self.credentials.delete(self._secret_id(server_id, key, kind)) self.credentials.delete(self._secret_id(server_id, key, kind))
except CredentialStoreError as exc: except CredentialStoreError as exc:
@@ -465,17 +497,27 @@ class McpServerRegistry:
self.tools.register(definition, arguments_model, executor) self.tools.register(definition, arguments_model, executor)
def _unavailable(self, server_id: str, message: str) -> None: def _unavailable(self, server_id: str, message: str) -> None:
with self._lock: # A failure may race with enable(). Waiting for the lifecycle mutation makes
for name in self._registered.pop(server_id, []): # sure tools registered immediately before the callback are also removed.
self.tools.unregister(name) with self._lifecycle_lock:
record = self._records.get(server_id) try:
if record is not None: with self._lock:
self._records[server_id] = {**record, "enabled": False} record = self._records.get(server_id)
self._last_status[server_id] = { registered = self._registered.pop(server_id, [])
"status": PluginHostState.unhealthy, for name in registered:
"error": message, self.tools.unregister(name)
} if record is not None and (record.get("enabled") or registered):
self._write() self._records[server_id] = {**record, "enabled": False}
self._last_status[server_id] = {
"status": PluginHostState.unhealthy,
"error": message,
}
self._write()
finally:
# broken() can run on the client's reader/event thread. stop() does
# not join that thread, and setting _stopping before closing the
# transport prevents the close itself from reporting another failure.
self.bridge.remove(self._host_id(server_id))
def _require_launch_allowed( def _require_launch_allowed(
self, record: dict[str, Any], *, require_test: bool self, record: dict[str, Any], *, require_test: bool
@@ -767,6 +809,23 @@ class McpServerRegistry:
status_code=404, status_code=404,
) from exc ) from exc
def _invalidate_test(self, server_id: str) -> None:
"""Make credential changes safe before touching the encrypted store."""
with self._lock:
record = self._record(server_id)
invalidated = {
**record,
"tested_digest": None,
"last_tested_at": None,
"last_test_succeeded": None,
}
updated = {**self._records, server_id: invalidated}
self._write(updated)
self._records = updated
self._last_status.pop(server_id, None)
self._summaries.pop(server_id, None)
@property @property
def _path(self) -> Path: def _path(self) -> Path:
return self.data_dir / "mcp" / "servers.json" return self.data_dir / "mcp" / "servers.json"
@@ -788,7 +847,29 @@ class McpServerRegistry:
"MCP server registry has an invalid format.", "MCP server registry has an invalid format.",
status_code=500, status_code=500,
) )
return value if len(value) > _MAX_MCP_SERVERS:
raise McpRegistryError(
"MCP_REGISTRY_INVALID",
"MCP server registry contains too many records.",
status_code=500,
)
normalized: dict[str, dict[str, Any]] = {}
config_fields = set(McpServerConfig.model_fields)
try:
for server_id, raw in value.items():
if not isinstance(server_id, str) or not _SERVER_ID.fullmatch(server_id):
raise ValueError("invalid server id")
record = _McpServerRecord.model_validate(raw)
config = record.model_dump(mode="json", include=config_fields)
self._validate(McpServerCreateRequest.model_validate(config))
normalized[server_id] = record.model_dump(mode="json")
except (McpRegistryError, ValidationError, ValueError, TypeError) as exc:
raise McpRegistryError(
"MCP_REGISTRY_INVALID",
"MCP server registry contains an invalid record.",
status_code=500,
) from exc
return normalized
def _write(self, records: dict[str, dict[str, Any]] | None = None) -> None: def _write(self, records: dict[str, dict[str, Any]] | None = None) -> None:
temporary = self._path.with_suffix(".tmp") temporary = self._path.with_suffix(".tmp")
+2 -2
View File
@@ -607,7 +607,7 @@ async def put_mcp_server_secret(
request: McpServerSecretWriteRequest, request: McpServerSecretWriteRequest,
kind: str = Query(default="environment", pattern="^(environment|header)$"), kind: str = Query(default="environment", pattern="^(environment|header)$"),
) -> McpServerSecretStatus: ) -> McpServerSecretStatus:
return mcp_call( return await mcp_call_async(
lambda: container.mcp_servers.put_secret( lambda: container.mcp_servers.put_secret(
server_id, key, request.secret.get_secret_value(), kind=kind server_id, key, request.secret.get_secret_value(), kind=kind
) )
@@ -624,7 +624,7 @@ async def delete_mcp_server_secret(
key: str, key: str,
kind: str = Query(default="environment", pattern="^(environment|header)$"), kind: str = Query(default="environment", pattern="^(environment|header)$"),
) -> McpServerSecretStatus: ) -> McpServerSecretStatus:
return mcp_call( return await mcp_call_async(
lambda: container.mcp_servers.delete_secret(server_id, key, kind=kind) lambda: container.mcp_servers.delete_secret(server_id, key, kind=kind)
) )
+69 -20
View File
@@ -1,26 +1,10 @@
import asyncio import asyncio
import threading
from types import SimpleNamespace
from app.main import health, service_status
from app.routes import (
get_index_status,
list_notes,
list_plugins,
list_provider_presets,
list_providers,
list_skills,
)
from app.routes import (
create_provider,
create_task,
delete_provider,
delete_task,
get_provider,
get_task,
list_tasks,
update_provider,
update_task,
)
from app.contracts import ( from app.contracts import (
McpServerSecretStatus,
McpServerSecretWriteRequest,
ProviderCreateRequest, ProviderCreateRequest,
ProviderType, ProviderType,
ProviderUpdateRequest, ProviderUpdateRequest,
@@ -28,6 +12,63 @@ from app.contracts import (
TaskStatus, TaskStatus,
TaskUpdateRequest, TaskUpdateRequest,
) )
from app.main import health, service_status
from app.routes import (
create_provider,
create_task,
delete_provider,
delete_task,
get_index_status,
get_provider,
get_task,
list_notes,
list_plugins,
list_provider_presets,
list_providers,
list_skills,
list_tasks,
update_provider,
update_task,
)
def test_mcp_secret_routes_offload_blocking_lifecycle_work(monkeypatch) -> None:
from app import routes
caller_thread = threading.get_ident()
worker_threads: list[int] = []
class FakeMcpRegistry:
def put_secret(self, server_id, key, secret, *, kind):
worker_threads.append(threading.get_ident())
return McpServerSecretStatus(key=key, configured=True)
def delete_secret(self, server_id, key, *, kind):
worker_threads.append(threading.get_ident())
return McpServerSecretStatus(key=key, configured=False)
monkeypatch.setattr(
routes,
"container",
SimpleNamespace(mcp_servers=FakeMcpRegistry()),
)
written = asyncio.run(
routes.put_mcp_server_secret(
"server-1",
"TOKEN",
McpServerSecretWriteRequest(secret="hidden"),
kind="environment",
)
)
deleted = asyncio.run(
routes.delete_mcp_server_secret(
"server-1", "TOKEN", kind="environment"
)
)
assert written.configured is True
assert deleted.configured is False
assert worker_threads and all(item != caller_thread for item in worker_threads)
def test_health() -> None: def test_health() -> None:
@@ -101,6 +142,14 @@ def test_openapi_contains_documented_frontend_interfaces() -> None:
"/api/plugins/{plugin_id}/settings/{key}/secret", "/api/plugins/{plugin_id}/settings/{key}/secret",
"/api/plugins/{plugin_id}/enable", "/api/plugins/{plugin_id}/enable",
"/api/plugins/{plugin_id}/disable", "/api/plugins/{plugin_id}/disable",
"/api/mcp/servers",
"/api/mcp/servers/{server_id}",
"/api/mcp/servers/{server_id}/tools",
"/api/mcp/servers/{server_id}/trust",
"/api/mcp/servers/{server_id}/test",
"/api/mcp/servers/{server_id}/enable",
"/api/mcp/servers/{server_id}/disable",
"/api/mcp/servers/{server_id}/secrets/{key}",
"/api/providers/test", "/api/providers/test",
"/api/providers/presets", "/api/providers/presets",
"/api/credentials/{credential_id}", "/api/credentials/{credential_id}",
+169 -2
View File
@@ -1,6 +1,7 @@
import asyncio import asyncio
import json import json
import sys import sys
import threading
import time import time
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
@@ -10,8 +11,9 @@ import pytest
from app.agent.tools import ToolExecutionContext, ToolRegistry from app.agent.tools import ToolExecutionContext, ToolRegistry
from app.config import BACKEND_DIR, get_settings from app.config import BACKEND_DIR, get_settings
from app.contracts import McpServerCreateRequest, McpServerUpdateRequest, ToolCall from app.contracts import McpServerCreateRequest, McpServerUpdateRequest, ToolCall
from app.extensions.mcp import McpLegacySseClient
from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry
from app.providers.credentials import EncryptedCredentialStore from app.providers.credentials import CredentialStoreError, EncryptedCredentialStore
SERVER = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo" / "server.py" SERVER = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo" / "server.py"
@@ -59,6 +61,28 @@ def test_registry_requires_current_trust_and_never_returns_secret() -> None:
service.shutdown() service.shutdown()
def test_secret_change_disables_server_and_requires_a_new_connection_test() -> None:
service = registry()
created = service.create(request())
service.put_secret(created.server_id, "TEST_MCP_SECRET", "first")
service.trust(created.server_id, created.command_digest)
service.test(created.server_id)
service.enable(created.server_id)
service.put_secret(created.server_id, "TEST_MCP_SECRET", "second")
current = service.get(created.server_id)
assert current.enabled is False
assert current.last_test_succeeded is None
assert not any(
item.name.startswith(f"mcp.{created.server_id}.")
for item in service.tools.definitions()
)
with pytest.raises(McpRegistryError) as error:
service.enable(created.server_id)
assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED"
service.shutdown()
def test_update_disables_server_and_revokes_command_trust() -> None: def test_update_disables_server_and_revokes_command_trust() -> None:
service = registry() service = registry()
created = service.create(request(secret_environment_keys=[])) created = service.create(request(secret_environment_keys=[]))
@@ -87,6 +111,70 @@ def test_update_disables_server_and_revokes_command_trust() -> None:
service.shutdown() service.shutdown()
def test_update_remains_retryable_when_removed_secret_cleanup_fails(monkeypatch) -> None:
service = registry()
created = service.create(request())
service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry")
def fail_delete_many(_secret_ids: list[str]) -> set[str]:
raise CredentialStoreError("credential store unavailable")
monkeypatch.setattr(service.credentials, "delete_many", fail_delete_many)
with pytest.raises(McpRegistryError) as error:
service.update(
created.server_id,
McpServerUpdateRequest(
**request(secret_environment_keys=[]).model_dump(),
version=created.version,
),
)
current = service.get(created.server_id)
assert error.value.code == "MCP_SECRET_STORE_ERROR"
assert current.version == created.version
assert current.secret_environment == {"TEST_MCP_SECRET": True}
service.shutdown()
def test_delete_keeps_server_retryable_when_secret_cleanup_fails(monkeypatch) -> None:
service = registry()
created = service.create(request())
service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry")
def fail_delete_many(_secret_ids: list[str]) -> set[str]:
raise CredentialStoreError("credential store unavailable")
monkeypatch.setattr(service.credentials, "delete_many", fail_delete_many)
with pytest.raises(McpRegistryError) as error:
service.delete(created.server_id)
current = service.get(created.server_id)
assert error.value.code == "MCP_SECRET_STORE_ERROR"
assert current.server_id == created.server_id
assert current.secret_environment == {"TEST_MCP_SECRET": True}
service.shutdown()
def test_unavailable_server_removes_bridge_host(monkeypatch) -> None:
service = registry()
created = service.create(request(secret_environment_keys=[]))
with service._lock:
service._records[created.server_id] = {
**service._records[created.server_id],
"enabled": True,
}
removed: list[str] = []
monkeypatch.setattr(service.bridge, "remove", removed.append)
service._unavailable(created.server_id, "connection lost")
current = service.get(created.server_id)
assert removed == [f"mcp.{created.server_id}"]
assert current.enabled is False
assert current.status == "unhealthy"
service.shutdown()
def test_production_rejects_process_launch_even_after_approval() -> None: def test_production_rejects_process_launch_even_after_approval() -> None:
service = registry(launch=False) service = registry(launch=False)
created = service.create(request(secret_environment_keys=[])) created = service.create(request(secret_environment_keys=[]))
@@ -141,6 +229,36 @@ def test_registry_rejects_corrupt_persisted_json(tmp_path) -> None:
assert error.value.code == "MCP_REGISTRY_INVALID" assert error.value.code == "MCP_REGISTRY_INVALID"
def test_registry_rejects_structurally_invalid_record(tmp_path) -> None:
path = tmp_path / "mcp"
path.mkdir()
(path / "servers.json").write_text(
json.dumps({"server-1": {"name": "Broken", "transport": "stdio"}}),
encoding="utf-8",
)
with pytest.raises(McpRegistryError) as error:
McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
tmp_path,
allow_process_launch=True,
)
assert error.value.code == "MCP_REGISTRY_INVALID"
def test_registry_rejects_create_before_exceeding_persisted_limit(
monkeypatch,
) -> None:
service = registry()
service.create(request(name="Only server"))
monkeypatch.setattr("app.extensions.mcp_registry._MAX_MCP_SERVERS", 1)
with pytest.raises(McpRegistryError) as error:
service.create(request(name="One too many"))
assert error.value.code == "MCP_SERVER_LIMIT_REACHED"
assert len(service.list()) == 1
service.shutdown()
def test_stdio_command_is_not_parsed_as_a_shell_string() -> None: def test_stdio_command_is_not_parsed_as_a_shell_string() -> None:
service = registry() service = registry()
created = service.create( created = service.create(
@@ -217,6 +335,7 @@ def test_streamable_http_supports_session_headers_secrets_and_tool_summary(
monkeypatch, monkeypatch,
) -> None: ) -> None:
requests: list[httpx.Request] = [] requests: list[httpx.Request] = []
request_timeouts: dict[str, float] = {}
def handler(request_value: httpx.Request) -> httpx.Response: def handler(request_value: httpx.Request) -> httpx.Response:
requests.append(request_value) requests.append(request_value)
@@ -225,6 +344,9 @@ def test_streamable_http_supports_session_headers_secrets_and_tool_summary(
if request_value.method == "DELETE": if request_value.method == "DELETE":
return httpx.Response(405) return httpx.Response(405)
payload = json.loads(request_value.content) payload = json.loads(request_value.content)
timeout = request_value.extensions.get("timeout", {}).get("read")
if isinstance(timeout, (int, float)):
request_timeouts[payload.get("method", "notification")] = float(timeout)
if payload.get("method") == "initialize": if payload.get("method") == "initialize":
response = _http_result( response = _http_result(
payload["id"], payload["id"],
@@ -290,6 +412,9 @@ def test_streamable_http_supports_session_headers_secrets_and_tool_summary(
assert all( assert all(
request.headers.get("authorization") == "Bearer hidden" for request in requests request.headers.get("authorization") == "Bearer hidden" for request in requests
) )
assert request_timeouts["initialize"] == 15
assert request_timeouts["notifications/initialized"] == 15
assert request_timeouts["tools/list"] == 15
enabled = service.enable(created.server_id) enabled = service.enable(created.server_id)
tool_name = service.list_tools(created.server_id)[0].name tool_name = service.list_tools(created.server_id)[0].name
result = asyncio.run( result = asyncio.run(
@@ -301,11 +426,15 @@ def test_streamable_http_supports_session_headers_secrets_and_tool_summary(
assert enabled.enabled is True assert enabled.enabled is True
assert result.success is True assert result.success is True
assert result.output == {"transport": "http"} assert result.output == {"transport": "http"}
assert request_timeouts["tools/call"] == 30
service.disable(created.server_id) service.disable(created.server_id)
service.shutdown() service.shutdown()
class _LegacyEventStream(httpx.SyncByteStream): class _LegacyEventStream(httpx.SyncByteStream):
def __init__(self) -> None:
self.closed = threading.Event()
def __iter__(self): def __iter__(self):
yield b"event: endpoint\ndata: /messages\n\n" yield b"event: endpoint\ndata: /messages\n\n"
time.sleep(0.1) time.sleep(0.1)
@@ -326,17 +455,22 @@ class _LegacyEventStream(httpx.SyncByteStream):
"result": {"tools": []}, "result": {"tools": []},
} }
yield f"data: {json.dumps(tools)}\n\n".encode() yield f"data: {json.dumps(tools)}\n\n".encode()
self.closed.wait()
def close(self) -> None:
self.closed.set()
def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None: def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None:
posted_urls: list[str] = [] posted_urls: list[str] = []
event_stream = _LegacyEventStream()
def handler(request_value: httpx.Request) -> httpx.Response: def handler(request_value: httpx.Request) -> httpx.Response:
if request_value.method == "GET": if request_value.method == "GET":
return httpx.Response( return httpx.Response(
200, 200,
headers={"content-type": "text/event-stream"}, headers={"content-type": "text/event-stream"},
stream=_LegacyEventStream(), stream=event_stream,
) )
posted_urls.append(str(request_value.url)) posted_urls.append(str(request_value.url))
return httpx.Response(202) return httpx.Response(202)
@@ -361,6 +495,39 @@ def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None:
url == "https://legacy.example.test/messages" for url in posted_urls url == "https://legacy.example.test/messages" for url in posted_urls
) )
service.shutdown() service.shutdown()
event_stream.close()
class _EndingLegacyEventStream(httpx.SyncByteStream):
def __iter__(self):
yield b"event: endpoint\ndata: /messages\n\n"
def test_legacy_sse_eof_marks_client_unavailable(monkeypatch) -> None:
def handler(_request_value: httpx.Request) -> httpx.Response:
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
stream=_EndingLegacyEventStream(),
)
real_client = httpx.Client
monkeypatch.setattr(
"app.extensions.mcp.httpx.Client",
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
)
broken = threading.Event()
client = McpLegacySseClient(
"https://legacy.example.test/sse",
headers={},
startup_timeout_seconds=1,
on_seen=lambda: None,
on_broken=lambda _message: broken.set(),
on_tools_changed=lambda: None,
)
client.start()
assert broken.wait(timeout=1)
client.stop()
class _CrossOriginLegacyEventStream(httpx.SyncByteStream): class _CrossOriginLegacyEventStream(httpx.SyncByteStream):
@@ -661,6 +661,7 @@ MCP_SERVER_NOT_FOUND
MCP_SERVER_NAME_INVALID MCP_SERVER_NAME_INVALID
MCP_SERVER_ALREADY_ENABLED MCP_SERVER_ALREADY_ENABLED
MCP_SERVER_VERSION_CONFLICT MCP_SERVER_VERSION_CONFLICT
MCP_SERVER_LIMIT_REACHED
MCP_REGISTRY_WRITE_FAILED MCP_REGISTRY_WRITE_FAILED
MCP_REGISTRY_INVALID MCP_REGISTRY_INVALID
MCP_CONNECTION_TEST_REQUIRED MCP_CONNECTION_TEST_REQUIRED
@@ -707,6 +708,8 @@ stdio 配置使用 `command`、`args`、`environment` 和 `secret_environment_ke
创建或编辑配置后,调用方必须向 `/trust` 回传服务端计算的 `command_digest`。后端只接受与当前 Transport、连接参数、环境/Header 及权限完全一致的摘要;配置变化会撤销旧信任和测试结果。只有当前摘要通过 `/test`,才能调用 `/enable`。测试失败也会持久化时间和失败状态。 创建或编辑配置后,调用方必须向 `/trust` 回传服务端计算的 `command_digest`。后端只接受与当前 Transport、连接参数、环境/Header 及权限完全一致的摘要;配置变化会撤销旧信任和测试结果。只有当前摘要通过 `/test`,才能调用 `/enable`。测试失败也会持久化时间和失败状态。
Secret 明文变化无法进入摘要,因此 Secret 写入和删除采用更严格规则:若 Server 已启用则先停用并注销 Tool,随后清除 `tested_digest` 和最近测试状态。调用方必须用新 Secret 再次执行 `/test`,不能沿用旧凭据的测试结果。
Streamable HTTP 支持 Session ID、`MCP-Protocol-Version`、JSON 或 SSE POST 响应、可选 GET 事件流及 `Last-Event-ID`;旧 SSE 按 endpoint 事件确定 POST 地址,并要求与初始 URL 同源。Secret 接口用 `?kind=environment``?kind=header` 区分类型。HTTP URL 不允许内嵌凭据或 Fragment,配置不得覆盖协议保留 Header。 Streamable HTTP 支持 Session ID、`MCP-Protocol-Version`、JSON 或 SSE POST 响应、可选 GET 事件流及 `Last-Event-ID`;旧 SSE 按 endpoint 事件确定 POST 地址,并要求与初始 URL 同源。Secret 接口用 `?kind=environment``?kind=header` 区分类型。HTTP URL 不允许内嵌凭据或 Fragment,配置不得覆盖协议保留 Header。
stdio 命令始终以 executable 与 args 数组通过 `shell=False` 启动;普通环境变量和加密 Secret 显式注入,不继承 Provider Key、数据库或 Vault 路径。当前 Python Host 仅在 `APP_ENVIRONMENT=development` 时允许启动 stdio;其他环境返回 `403 MCP_SANDBOX_REQUIRED`。远程 HTTP Transport 不创建本机进程,但仍受摘要确认、成功测试、超时、消息限长与 Secret 隔离约束。 stdio 命令始终以 executable 与 args 数组通过 `shell=False` 启动;普通环境变量和加密 Secret 显式注入,不继承 Provider Key、数据库或 Vault 路径。当前 Python Host 仅在 `APP_ENVIRONMENT=development` 时允许启动 stdio;其他环境返回 `403 MCP_SANDBOX_REQUIRED`。远程 HTTP Transport 不创建本机进程,但仍受摘要确认、成功测试、超时、消息限长与 Secret 隔离约束。
@@ -21,6 +21,10 @@ Streamable HTTP 按 MCP 当前规范实现;SSE 仅用于兼容旧 Server,不
Secret 使用带类型和键名哈希的 `mcp.*` 内部 ID 写入 Fernet 凭据存储。API 只返回环境变量或 Header 是否配置,绝不返回明文。删除 Server 或移除 Secret 键会同步清理密文。前端密码框提交后立即清空,不写入 JSON 编辑器、localStorage、普通配置或日志。 Secret 使用带类型和键名哈希的 `mcp.*` 内部 ID 写入 Fernet 凭据存储。API 只返回环境变量或 Header 是否配置,绝不返回明文。删除 Server 或移除 Secret 键会同步清理密文。前端密码框提交后立即清空,不写入 JSON 编辑器、localStorage、普通配置或日志。
跨 Registry 与凭据存储的删除以“失败后可重试”为顺序约束:先原子清理密文,再提交新版本或删除 Registry 记录。凭据存储失败时保留原版本和 Server 记录,避免出现返回 500 但配置已提交、版本无法重试或密文失去清理入口的状态。
写入或删除 Secret 会先停用正在运行的连接、注销动态 Tool,并撤销当前配置的测试通过状态;必须使用新凭据重新测试后才能启用。这样页面展示的凭据状态不会与运行中进程实际持有的旧凭据不一致。
## 3. 启用与运行时规则 ## 3. 启用与运行时规则
一次连接按以下顺序执行: 一次连接按以下顺序执行:
@@ -30,8 +34,12 @@ Secret 使用带类型和键名哈希的 `mcp.*` 内部 ID 写入 Fernet 凭据
3. 只有当前摘要测试成功,启用操作才会启动长期连接并注册 Tool; 3. 只有当前摘要测试成功,启用操作才会启动长期连接并注册 Tool;
4. 停用、删除、超时或异常退出会注销 Tool 并关闭连接。 4. 停用、删除、超时或异常退出会注销 Tool 并关闭连接。
运行期连接异常或 Tool 列表变化时,Registry 会在同一生命周期临界区内标记不可用、注销 Tool,并从 Bridge 移除 Host,不保留后台事件线程或 HTTP Client。Secret 写入和删除可能触发该停用流程,因此对应 API 通过工作线程执行,不阻塞 FastAPI 事件循环。
HTTP Header 中 `Host``Content-Type``MCP-Session-Id` 等协议保留项不可由配置覆盖。URL 不允许内嵌凭据或 Fragment。旧 SSE 返回的 POST endpoint 必须与初始 URL 同源,防止认证 Header 被转发到其他站点。 HTTP Header 中 `Host``Content-Type``MCP-Session-Id` 等协议保留项不可由配置覆盖。URL 不允许内嵌凭据或 Fragment。旧 SSE 返回的 POST endpoint 必须与初始 URL 同源,防止认证 Header 被转发到其他站点。
启动与 Tool 请求超时会同时应用于业务等待和底层 HTTP 请求;旧 SSE 的 endpoint 等待也使用启动超时。非主动结束的旧 SSE 事件流视为 Host 不可用,宿主随后注销 Tool。注册表加载时逐条校验 Server ID、配置字段、Transport 组合和摘要格式,损坏记录统一返回 `MCP_REGISTRY_INVALID`
stdio 命令不经过 Shell,管道、重定向和命令拼接不会被解释。Windows 使用新进程组并通过 `taskkill /T` 回收子树;POSIX 使用独立 session/process group 并向进程组发信号。Python 阶段仍无法提供文件、网络、系统调用或操作系统版本差异下的绝对隔离保证。 stdio 命令不经过 Shell,管道、重定向和命令拼接不会被解释。Windows 使用新进程组并通过 `taskkill /T` 回收子树;POSIX 使用独立 session/process group 并向进程组发信号。Python 阶段仍无法提供文件、网络、系统调用或操作系统版本差异下的绝对隔离保证。
`uvx` 模板使用 `--isolated`、明确的 `--from` 和固定包版本。它只能隔离依赖,不能替代安全沙箱。非开发环境仍拒绝启动 stdio Server,并返回 `403 MCP_SANDBOX_REQUIRED`;远程 HTTP Transport 不创建本机子进程,但仍要求摘要确认和成功测试。第三阶段前的 C.5 将冻结 Tauri/Rust 沙箱设计。 `uvx` 模板使用 `--isolated`、明确的 `--from` 和固定包版本。它只能隔离依赖,不能替代安全沙箱。非开发环境仍拒绝启动 stdio Server,并返回 `403 MCP_SANDBOX_REQUIRED`;远程 HTTP Transport 不创建本机子进程,但仍要求摘要确认和成功测试。第三阶段前的 C.5 将冻结 Tauri/Rust 沙箱设计。
@@ -81,4 +81,15 @@ describe('McpServersView', () => {
expect(confirm).toHaveBeenCalled() expect(confirm).toHaveBeenCalled()
expect(service.deleteMcpServer).toHaveBeenCalledWith('server-1') expect(service.deleteMcpServer).toHaveBeenCalledWith('server-1')
}) })
it('confirms permission changes before updating an existing server', async () => {
const wrapper = await render([server])
vi.mocked(service.updateMcpServer).mockResolvedValue(server)
await wrapper.findAll('button').find(button => button.text().includes('编辑'))!.trigger('click')
await wrapper.get('input[placeholder="network.request, notes.read"]').setValue('notes.read')
await wrapper.get('form').trigger('submit')
await flushPromises()
expect(confirm).toHaveBeenCalledWith(expect.stringContaining('旧测试与授权会失效'))
expect(service.updateMcpServer).toHaveBeenCalled()
})
}) })
+14 -1
View File
@@ -142,7 +142,20 @@ async function save() {
} }
function executionChanged(server: McpServer, input: McpServerInput) { function executionChanged(server: McpServer, input: McpServerInput) {
return JSON.stringify([server.transport, server.command, server.args, server.url, server.headers, Object.keys(server.secret_headers)]) !== JSON.stringify([input.transport, input.command, input.args, input.url, input.headers, input.secret_header_keys]) const sortedEntries = (value: Record<string, string>) => Object.entries(value).sort(([left], [right]) => left.localeCompare(right))
const current = [
server.transport, server.command, server.args, server.url,
sortedEntries(server.headers), sortedEntries(server.environment),
Object.keys(server.secret_headers).sort(), Object.keys(server.secret_environment).sort(),
[...server.permissions].sort(), server.startup_timeout_seconds, server.tool_timeout_seconds,
]
const next = [
input.transport, input.command, input.args, input.url,
sortedEntries(input.headers), sortedEntries(input.environment),
[...input.secret_header_keys].sort(), [...input.secret_environment_keys].sort(),
[...input.permissions].sort(), input.startup_timeout_seconds, input.tool_timeout_seconds,
]
return JSON.stringify(current) !== JSON.stringify(next)
} }
async function approve(server: McpServer): Promise<McpServer | null> { async function approve(server: McpServer): Promise<McpServer | null> {
@@ -74,4 +74,33 @@ describe('FileTreePanel file switching', () => {
expect(editorStore.content).toContain('# 二叉搜索树') expect(editorStore.content).toContain('# 二叉搜索树')
expect(editorStore.currentNoteId).toBe('note-bst') expect(editorStore.currentNoteId).toBe('note-bst')
}) })
it('creates a Markdown note inside the selected folder', async () => {
const router = createRouter({
history: createMemoryHistory(),
routes: [{ path: '/workspace', component: { template: '<div />' } }],
})
await router.push('/workspace')
await router.isReady()
const workspaceStore = useWorkspaceStore()
await workspaceStore.openVault('C:/vault')
const createFile = vi.spyOn(workspaceService, 'createFile').mockResolvedValue({
id: 'note-new', note_id: 'note-new', name: '新笔记.md',
path: '/数据结构/新笔记.md', type: 'file',
})
wrapper = mount(FileTreePanel, { attachTo: document.body, global: { plugins: [router] } })
await wrapper.findAll('.tree-node').find((node) => node.text().includes('数据结构'))!.trigger('click')
await wrapper.get('button[aria-label="新建笔记"]').trigger('click')
await wrapper.get('.new-item input').setValue('新笔记')
await wrapper.get('.new-item').trigger('submit')
await waitForPath('/数据结构/新笔记.md')
await vi.waitFor(() => {
expect(workspaceStore.activeFilePath).toBe('/数据结构/新笔记.md')
})
expect(createFile).toHaveBeenCalledWith('/数据结构', '新笔记.md', '# 新笔记\n\n')
expect(wrapper.findAll('.tree-node').some((node) => node.classes().includes('active') && node.text().includes('新笔记.md'))).toBe(true)
})
}) })
@@ -1,5 +1,5 @@
<script setup lang="ts"> <script setup lang="ts">
import { ref } from 'vue' import { ref, watch } from 'vue'
import { useRouter } from 'vue-router' import { useRouter } from 'vue-router'
import type { FileNode } from '@/contracts' import type { FileNode } from '@/contracts'
import * as workspaceService from '@/services/workspaceService' import * as workspaceService from '@/services/workspaceService'
@@ -15,9 +15,19 @@ const router = useRouter()
const newItemType = ref<'file' | 'folder' | null>(null) const newItemType = ref<'file' | 'folder' | null>(null)
const newItemName = ref('') const newItemName = ref('')
const parentPath = ref('/') const parentPath = ref('/')
const selectedTreePath = ref(workspaceStore.activeFilePath ?? '/')
const selectedFolderPath = ref(
workspaceStore.activeFilePath ? containingFolder(workspaceStore.activeFilePath) : '/',
)
const contextTarget = ref<FileNode | null>(null) const contextTarget = ref<FileNode | null>(null)
const contextMenuPosition = ref({ x: 0, y: 0 }) const contextMenuPosition = ref({ x: 0, y: 0 })
watch(() => workspaceStore.activeFilePath, (path) => {
if (!path) return
selectedTreePath.value = path
selectedFolderPath.value = containingFolder(path)
})
function beginCreate(type: 'file' | 'folder', parent = '/') { function beginCreate(type: 'file' | 'folder', parent = '/') {
newItemType.value = type newItemType.value = type
newItemName.value = '' newItemName.value = ''
@@ -31,19 +41,28 @@ async function createItem() {
const name = rawName.endsWith('.md') ? rawName : `${rawName}.md` const name = rawName.endsWith('.md') ? rawName : `${rawName}.md`
const file = await workspaceService.createFile(parentPath.value, name, `# ${rawName}\n\n`) const file = await workspaceService.createFile(parentPath.value, name, `# ${rawName}\n\n`)
workspaceStore.addFileToTree(parentPath.value, file) workspaceStore.addFileToTree(parentPath.value, file)
selectedTreePath.value = file.path
selectedFolderPath.value = parentPath.value
await editorStore.loadFile(file.path) await editorStore.loadFile(file.path)
workspaceStore.openFile(file.path) workspaceStore.openFile(file.path)
await router.push('/workspace') await router.push('/workspace')
} else { } else {
const folder = await workspaceService.createFolder(parentPath.value, rawName) const folder = await workspaceService.createFolder(parentPath.value, rawName)
workspaceStore.addFileToTree(parentPath.value, folder) workspaceStore.addFileToTree(parentPath.value, folder)
selectedTreePath.value = folder.path
selectedFolderPath.value = folder.path
} }
newItemType.value = null newItemType.value = null
newItemName.value = '' newItemName.value = ''
} }
async function openNode(node: FileNode) { async function openNode(node: FileNode) {
if (node.type === 'folder') return workspaceStore.toggleFolder(node.path) selectedTreePath.value = node.path
if (node.type === 'folder') {
selectedFolderPath.value = node.path
return workspaceStore.toggleFolder(node.path)
}
selectedFolderPath.value = containingFolder(node.path)
// //
const previousPath = workspaceStore.activeFilePath const previousPath = workspaceStore.activeFilePath
const wasOpen = workspaceStore.openFiles.includes(node.path) const wasOpen = workspaceStore.openFiles.includes(node.path)
@@ -61,6 +80,8 @@ async function openNode(node: FileNode) {
function openContextMenu(event: MouseEvent, node: FileNode) { function openContextMenu(event: MouseEvent, node: FileNode) {
event.preventDefault() event.preventDefault()
event.stopPropagation() event.stopPropagation()
selectedTreePath.value = node.path
selectedFolderPath.value = node.type === 'folder' ? node.path : containingFolder(node.path)
contextTarget.value = node contextTarget.value = node
contextMenuPosition.value = { x: event.clientX, y: event.clientY } contextMenuPosition.value = { x: event.clientX, y: event.clientY }
} }
@@ -79,6 +100,12 @@ async function renameTarget() {
await workspaceService.renameFile(oldPath, normalizedName) await workspaceService.renameFile(oldPath, normalizedName)
workspaceStore.renamePath(oldPath, newPath, normalizedName) workspaceStore.renamePath(oldPath, newPath, normalizedName)
editorStore.renameFilePath(oldPath, newPath) editorStore.renameFilePath(oldPath, newPath)
if (selectedTreePath.value === oldPath || selectedTreePath.value.startsWith(`${oldPath}/`)) {
selectedTreePath.value = `${newPath}${selectedTreePath.value.slice(oldPath.length)}`
}
if (selectedFolderPath.value === oldPath || selectedFolderPath.value.startsWith(`${oldPath}/`)) {
selectedFolderPath.value = `${newPath}${selectedFolderPath.value.slice(oldPath.length)}`
}
} }
closeContextMenu() closeContextMenu()
} }
@@ -90,19 +117,28 @@ async function deleteTarget() {
await workspaceService.deleteFile(node.path) await workspaceService.deleteFile(node.path)
const activeWasRemoved = workspaceStore.closePath(node.path) const activeWasRemoved = workspaceStore.closePath(node.path)
workspaceStore.removeFromTree(node.path) workspaceStore.removeFromTree(node.path)
if (selectedTreePath.value === node.path || selectedTreePath.value.startsWith(`${node.path}/`)) {
selectedTreePath.value = containingFolder(node.path)
selectedFolderPath.value = selectedTreePath.value
}
if (activeWasRemoved) { if (activeWasRemoved) {
editorStore.closeFile() editorStore.closeFile()
if (workspaceStore.activeFilePath) await editorStore.loadFile(workspaceStore.activeFilePath) if (workspaceStore.activeFilePath) await editorStore.loadFile(workspaceStore.activeFilePath)
} }
closeContextMenu() closeContextMenu()
} }
function containingFolder(path: string): string {
const separator = path.lastIndexOf('/')
return separator > 0 ? path.slice(0, separator) : '/'
}
</script> </script>
<template> <template>
<section class="file-tree-panel" @click="closeContextMenu"> <section class="file-tree-panel" @click="closeContextMenu">
<div class="toolbar"> <div class="toolbar">
<button type="button" title="新建笔记" aria-label="新建笔记" @click.stop="beginCreate('file')"><AppIcon :icon="DocumentAdd" /></button> <button type="button" title="新建笔记" aria-label="新建笔记" @click.stop="beginCreate('file', selectedFolderPath)"><AppIcon :icon="DocumentAdd" /></button>
<button type="button" title="新建文件夹" aria-label="新建文件夹" @click.stop="beginCreate('folder')"><AppIcon :icon="FolderAdd" /></button> <button type="button" title="新建文件夹" aria-label="新建文件夹" @click.stop="beginCreate('folder', selectedFolderPath)"><AppIcon :icon="FolderAdd" /></button>
</div> </div>
<form v-if="newItemType" class="new-item" @submit.prevent="createItem"> <form v-if="newItemType" class="new-item" @submit.prevent="createItem">
<input v-model="newItemName" :placeholder="newItemType === 'file' ? '笔记名称' : '文件夹名称'" autofocus /> <input v-model="newItemName" :placeholder="newItemType === 'file' ? '笔记名称' : '文件夹名称'" autofocus />
@@ -111,7 +147,7 @@ async function deleteTarget() {
</form> </form>
<div class="tree"> <div class="tree">
<FileTreeNode v-for="node in workspaceStore.fileTree" :key="node.id" :node="node" <FileTreeNode v-for="node in workspaceStore.fileTree" :key="node.id" :node="node"
:active-path="workspaceStore.activeFilePath" @open="openNode" @context-menu="openContextMenu" /> :active-path="selectedTreePath" @open="openNode" @context-menu="openContextMenu" />
</div> </div>
<Teleport to="body"> <Teleport to="body">
<div v-if="contextTarget" class="context-menu" <div v-if="contextTarget" class="context-menu"