fix(paths): stabilize workspace and MCP cwd resolution

This commit is contained in:
JOJO 2026-05-03 15:28:36 +08:00
parent d406244604
commit 13689e80c2
7 changed files with 212 additions and 26 deletions

View File

@ -1,27 +1,41 @@
"""项目路径与目录配置。""" """项目路径与目录配置。"""
import os import os
from pathlib import Path
_REPO_ROOT = Path(__file__).resolve().parents[1]
def _resolve_repo_path(raw_value: str, default: str) -> str:
candidate = str(raw_value or "").strip() or str(default)
path = Path(candidate).expanduser()
if not path.is_absolute():
path = (_REPO_ROOT / path).resolve()
else:
path = path.resolve()
return str(path)
# 默认项目路径,可通过环境变量覆盖以指向宿主机任意目录 # 默认项目路径,可通过环境变量覆盖以指向宿主机任意目录
DEFAULT_PROJECT_PATH = os.environ.get("DEFAULT_PROJECT_PATH", "./project") DEFAULT_PROJECT_PATH = _resolve_repo_path(os.environ.get("DEFAULT_PROJECT_PATH", ""), "./project")
# 宿主机模式工作区配置文件JSON # 宿主机模式工作区配置文件JSON
HOST_WORKSPACES_FILE = os.environ.get("HOST_WORKSPACES_FILE", "./config/host_workspaces.json") HOST_WORKSPACES_FILE = _resolve_repo_path(os.environ.get("HOST_WORKSPACES_FILE", ""), "./config/host_workspaces.json")
# 兼容旧配置:若仍有模块读取 HOST_PROJECT_PATH保留该键实际宿主机路径选择改由 JSON 管理) # 兼容旧配置:若仍有模块读取 HOST_PROJECT_PATH保留该键实际宿主机路径选择改由 JSON 管理)
HOST_PROJECT_PATH = os.environ.get("HOST_PROJECT_PATH", DEFAULT_PROJECT_PATH) HOST_PROJECT_PATH = _resolve_repo_path(os.environ.get("HOST_PROJECT_PATH", ""), DEFAULT_PROJECT_PATH)
PROMPTS_DIR = "./prompts" PROMPTS_DIR = _resolve_repo_path(os.environ.get("PROMPTS_DIR", ""), "./prompts")
DATA_DIR = "./data" DATA_DIR = _resolve_repo_path(os.environ.get("DATA_DIR", ""), "./data")
LOGS_DIR = "./logs" LOGS_DIR = _resolve_repo_path(os.environ.get("LOGS_DIR", ""), "./logs")
AGENT_SKILLS_DIR = "./agentskills" AGENT_SKILLS_DIR = _resolve_repo_path(os.environ.get("AGENT_SKILLS_DIR", ""), "./agentskills")
WORKSPACE_SKILLS_DIRNAME = "skills" WORKSPACE_SKILLS_DIRNAME = "skills"
# 多用户空间 # 多用户空间
USER_SPACE_DIR = "./users" USER_SPACE_DIR = _resolve_repo_path(os.environ.get("USER_SPACE_DIR", ""), "./users")
USERS_DB_FILE = f"{DATA_DIR}/users.json" USERS_DB_FILE = f"{DATA_DIR}/users.json"
INVITE_CODES_FILE = f"{DATA_DIR}/invite_codes.json" INVITE_CODES_FILE = f"{DATA_DIR}/invite_codes.json"
ADMIN_POLICY_FILE = f"{DATA_DIR}/admin_policy.json" ADMIN_POLICY_FILE = f"{DATA_DIR}/admin_policy.json"
# API 专用用户与工作区(与网页用户隔离) # API 专用用户与工作区(与网页用户隔离)
API_USER_SPACE_DIR = "./api/users" API_USER_SPACE_DIR = _resolve_repo_path(os.environ.get("API_USER_SPACE_DIR", ""), "./api/users")
API_USERS_DB_FILE = f"{DATA_DIR}/api_users.json" API_USERS_DB_FILE = f"{DATA_DIR}/api_users.json"
API_TOKENS_FILE = f"{DATA_DIR}/api_tokens.json" API_TOKENS_FILE = f"{DATA_DIR}/api_tokens.json"
API_USAGE_FILE = f"{DATA_DIR}/api_usage.json" API_USAGE_FILE = f"{DATA_DIR}/api_usage.json"

View File

@ -1,6 +1,20 @@
"""子智能体相关配置。""" """子智能体相关配置。"""
import os import os
from pathlib import Path
_REPO_ROOT = Path(__file__).resolve().parents[1]
def _resolve_repo_path(raw_value: str, default: str) -> str:
candidate = str(raw_value or "").strip() or str(default)
path = Path(candidate).expanduser()
if not path.is_absolute():
path = (_REPO_ROOT / path).resolve()
else:
path = path.resolve()
return str(path)
# 子智能体服务 # 子智能体服务
SUB_AGENT_SERVICE_BASE_URL = os.environ.get("SUB_AGENT_SERVICE_URL", "http://127.0.0.1:8092") SUB_AGENT_SERVICE_BASE_URL = os.environ.get("SUB_AGENT_SERVICE_URL", "http://127.0.0.1:8092")
@ -8,9 +22,9 @@ SUB_AGENT_DEFAULT_TIMEOUT = int(os.environ.get("SUB_AGENT_DEFAULT_TIMEOUT", "180
SUB_AGENT_STATUS_POLL_INTERVAL = float(os.environ.get("SUB_AGENT_STATUS_POLL_INTERVAL", "2.0")) SUB_AGENT_STATUS_POLL_INTERVAL = float(os.environ.get("SUB_AGENT_STATUS_POLL_INTERVAL", "2.0"))
# 存储与并发限制 # 存储与并发限制
SUB_AGENT_TASKS_BASE_DIR = os.environ.get("SUB_AGENT_TASKS_BASE_DIR", "./sub_agent/tasks") SUB_AGENT_TASKS_BASE_DIR = _resolve_repo_path(os.environ.get("SUB_AGENT_TASKS_BASE_DIR", ""), "./sub_agent/tasks")
SUB_AGENT_PROJECT_RESULTS_DIR = os.environ.get("SUB_AGENT_PROJECT_RESULTS_DIR", "./project/sub_agent_results") SUB_AGENT_PROJECT_RESULTS_DIR = _resolve_repo_path(os.environ.get("SUB_AGENT_PROJECT_RESULTS_DIR", ""), "./project/sub_agent_results")
SUB_AGENT_STATE_FILE = os.environ.get("SUB_AGENT_STATE_FILE", "./data/sub_agents.json") SUB_AGENT_STATE_FILE = _resolve_repo_path(os.environ.get("SUB_AGENT_STATE_FILE", ""), "./data/sub_agents.json")
SUB_AGENT_MAX_ACTIVE = int(os.environ.get("SUB_AGENT_MAX_ACTIVE", "5")) SUB_AGENT_MAX_ACTIVE = int(os.environ.get("SUB_AGENT_MAX_ACTIVE", "5"))
__all__ = [ __all__ = [

View File

@ -396,6 +396,13 @@ class MainTerminal(MainTerminalCommandMixin, MainTerminalContextMixin, MainTermi
except Exception: except Exception:
pass pass
# MCP stdio 客户端是长连接子进程,切换工作区后需重建,确保后续 cwd 与路径映射一致。
if getattr(self, "mcp_client_manager", None):
try:
self.mcp_client_manager.close_all_clients()
except Exception:
pass
# 强制下次请求重新同步 skills含 AGENTS.md 场景) # 强制下次请求重新同步 skills含 AGENTS.md 场景)
self._skills_synced_project_path = None self._skills_synced_project_path = None

View File

@ -211,7 +211,12 @@ class _StdioMCPClient:
if not use_container: if not use_container:
env = dict(os.environ) env = dict(os.environ)
env.update(env_raw) env.update(env_raw)
return [command_raw, *args_raw], cwd_raw, env resolved_cwd = cwd_raw
if not resolved_cwd and session is not None:
workspace_hint = self._normalize_workspace_path(getattr(session, "workspace_path", ""))
if workspace_hint:
resolved_cwd = workspace_hint
return [command_raw, *args_raw], resolved_cwd, env
container_name = str(getattr(session, "container_name", "") or "").strip() container_name = str(getattr(session, "container_name", "") or "").strip()
if not container_name: if not container_name:
@ -731,6 +736,7 @@ class MCPClientManager:
self.registry = registry self.registry = registry
self.protocol_version = str(protocol_version or MCP_PROTOCOL_VERSION) self.protocol_version = str(protocol_version or MCP_PROTOCOL_VERSION)
self.container_session = container_session self.container_session = container_session
self._container_session_signature = self._compute_session_signature(container_session)
self._latest_alias_map: Dict[str, MCPToolBinding] = {} self._latest_alias_map: Dict[str, MCPToolBinding] = {}
self._client_pool: Dict[str, MCPClientPoolEntry] = {} self._client_pool: Dict[str, MCPClientPoolEntry] = {}
self._pool_lock = threading.RLock() self._pool_lock = threading.RLock()
@ -741,11 +747,27 @@ class MCPClientManager:
except Exception: except Exception:
pass pass
@staticmethod
def _compute_session_signature(session: Optional["ContainerHandle"]) -> Optional[Tuple[str, str, str, str, str]]:
if not session:
return None
return (
str(getattr(session, "mode", "") or ""),
str(getattr(session, "workspace_path", "") or ""),
str(getattr(session, "mount_path", "") or ""),
str(getattr(session, "container_name", "") or ""),
str(getattr(session, "sandbox_bin", "") or ""),
)
def set_container_session(self, session: Optional["ContainerHandle"]) -> None: def set_container_session(self, session: Optional["ContainerHandle"]) -> None:
if session is self.container_session: next_signature = self._compute_session_signature(session)
if next_signature == self._container_session_signature:
# 引用对象可能相同(甚至被原地修改),这里仍刷新引用,便于后续读取最新字段。
self.container_session = session
return return
self.close_all_clients() self.close_all_clients()
self.container_session = session self.container_session = session
self._container_session_signature = next_signature
@staticmethod @staticmethod
def _server_signature(server: Dict[str, Any]) -> str: def _server_signature(server: Dict[str, Any]) -> str:

View File

@ -4,16 +4,16 @@ import httpx
import json import json
from typing import Dict, Optional, Any, List from typing import Dict, Optional, Any, List
from datetime import datetime from datetime import datetime
from pathlib import Path
import re import re
try: try:
from config import TAVILY_API_KEY, SEARCH_MAX_RESULTS, OUTPUT_FORMATS from config import TAVILY_API_KEY, SEARCH_MAX_RESULTS, OUTPUT_FORMATS, DATA_DIR
except ImportError: except ImportError:
import sys import sys
from pathlib import Path
project_root = Path(__file__).resolve().parents[1] project_root = Path(__file__).resolve().parents[1]
if str(project_root) not in sys.path: if str(project_root) not in sys.path:
sys.path.insert(0, str(project_root)) sys.path.insert(0, str(project_root))
from config import TAVILY_API_KEY, SEARCH_MAX_RESULTS, OUTPUT_FORMATS from config import TAVILY_API_KEY, SEARCH_MAX_RESULTS, OUTPUT_FORMATS, DATA_DIR
class SearchEngine: class SearchEngine:
def __init__(self): def __init__(self):
@ -286,19 +286,16 @@ class SearchEngine:
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
filename = f"search_{timestamp}.json" filename = f"search_{timestamp}.json"
file_path = f"./data/searches/{filename}" file_path = Path(DATA_DIR).expanduser().resolve() / "searches" / filename
file_path.parent.mkdir(parents=True, exist_ok=True)
# 确保目录存在
import os
os.makedirs(os.path.dirname(file_path), exist_ok=True)
# 保存结果 # 保存结果
with open(file_path, 'w', encoding='utf-8') as f: with file_path.open('w', encoding='utf-8') as f:
json.dump(results, f, ensure_ascii=False, indent=2) json.dump(results, f, ensure_ascii=False, indent=2)
print(f"{OUTPUT_FORMATS['file']} 搜索结果已保存到: {file_path}") print(f"{OUTPUT_FORMATS['file']} 搜索结果已保存到: {file_path}")
return file_path return str(file_path)
def load_results(self, filename: str) -> Optional[Dict]: def load_results(self, filename: str) -> Optional[Dict]:
""" """
@ -310,10 +307,10 @@ class SearchEngine:
Returns: Returns:
搜索结果字典或None 搜索结果字典或None
""" """
file_path = f"./data/searches/{filename}" file_path = Path(DATA_DIR).expanduser().resolve() / "searches" / filename
try: try:
with open(file_path, 'r', encoding='utf-8') as f: with file_path.open('r', encoding='utf-8') as f:
return json.load(f) return json.load(f)
except FileNotFoundError: except FileNotFoundError:
print(f"{OUTPUT_FORMATS['error']} 文件不存在: {file_path}") print(f"{OUTPUT_FORMATS['error']} 文件不存在: {file_path}")

View File

@ -0,0 +1,87 @@
from __future__ import annotations
import json
import os
import subprocess
import sys
import tempfile
import unittest
from pathlib import Path
class ConfigPathsResolutionTest(unittest.TestCase):
def setUp(self):
self.repo_root = Path(__file__).resolve().parents[1]
def _load_paths(self, *, cwd: Path, extra_env: dict | None = None) -> dict:
code = """
import json
import config
print(json.dumps({
"DEFAULT_PROJECT_PATH": config.DEFAULT_PROJECT_PATH,
"HOST_WORKSPACES_FILE": config.HOST_WORKSPACES_FILE,
"DATA_DIR": config.DATA_DIR,
"LOGS_DIR": config.LOGS_DIR,
"USER_SPACE_DIR": config.USER_SPACE_DIR,
"API_USER_SPACE_DIR": config.API_USER_SPACE_DIR,
"SUB_AGENT_STATE_FILE": config.SUB_AGENT_STATE_FILE,
}, ensure_ascii=False))
"""
env = dict(os.environ)
env["PYTHONPATH"] = str(self.repo_root)
if extra_env:
env.update(extra_env)
completed = subprocess.run(
[sys.executable, "-c", code],
cwd=str(cwd),
env=env,
capture_output=True,
text=True,
check=True,
)
lines = [line.strip() for line in (completed.stdout or "").splitlines() if line.strip()]
self.assertTrue(lines, f"stdout is empty, stderr={completed.stderr}")
return json.loads(lines[-1])
def test_default_paths_are_resolved_from_repo_root(self):
with tempfile.TemporaryDirectory() as td:
data = self._load_paths(cwd=Path(td))
for key, raw in data.items():
value = Path(str(raw))
self.assertTrue(value.is_absolute(), f"{key} is not absolute: {raw}")
self.assertTrue(
str(value).startswith(str(self.repo_root)),
f"{key} should be anchored to repo root: {raw}",
)
def test_relative_env_overrides_also_anchor_to_repo_root(self):
with tempfile.TemporaryDirectory() as td:
data = self._load_paths(
cwd=Path(td),
extra_env={
"DEFAULT_PROJECT_PATH": "./workspace_rel",
"DATA_DIR": "./data_rel",
"LOGS_DIR": "./logs_rel",
"USER_SPACE_DIR": "./users_rel",
"HOST_WORKSPACES_FILE": "./config/host_workspaces_rel.json",
"SUB_AGENT_STATE_FILE": "./data/sub_agents_rel.json",
},
)
self.assertEqual(data["DEFAULT_PROJECT_PATH"], str((self.repo_root / "workspace_rel").resolve()))
self.assertEqual(data["DATA_DIR"], str((self.repo_root / "data_rel").resolve()))
self.assertEqual(data["LOGS_DIR"], str((self.repo_root / "logs_rel").resolve()))
self.assertEqual(data["USER_SPACE_DIR"], str((self.repo_root / "users_rel").resolve()))
self.assertEqual(
data["HOST_WORKSPACES_FILE"],
str((self.repo_root / "config" / "host_workspaces_rel.json").resolve()),
)
self.assertEqual(
data["SUB_AGENT_STATE_FILE"],
str((self.repo_root / "data" / "sub_agents_rel.json").resolve()),
)
if __name__ == "__main__":
unittest.main()

View File

@ -478,6 +478,51 @@ class MCPIntegrationTest(unittest.TestCase):
self.assertNotIn("/opt/homebrew/bin/npx", launch_cmd) self.assertNotIn("/opt/homebrew/bin/npx", launch_cmd)
self.assertIn("/Users/jojo/Desktop", launch_cmd) self.assertIn("/Users/jojo/Desktop", launch_cmd)
def test_stdio_host_mode_uses_workspace_as_default_cwd(self):
workspace = (self.project_dir / "workspace_a").resolve()
workspace.mkdir(parents=True, exist_ok=True)
fake_session = SimpleNamespace(
mode="host",
workspace_path=str(workspace),
mount_path=str(workspace),
container_name=None,
sandbox_bin="docker",
)
server = {
"command": sys.executable,
"args": ["-V"],
"cwd": "",
"env": {},
}
client = _StdioMCPClient(server, timeout_seconds=10, protocol_version="2025-06-18", container_session=fake_session)
_, cwd, _ = client._prepare_launch()
self.assertEqual(cwd, str(workspace))
def test_set_container_session_detects_inplace_workspace_change(self):
manager = MCPClientManager(self.registry)
close_calls = []
original_close = manager.close_all_clients
manager.close_all_clients = lambda: close_calls.append("closed") # type: ignore[assignment]
try:
session = SimpleNamespace(
mode="host",
workspace_path=str((self.project_dir / "ws_a").resolve()),
mount_path="/workspace",
container_name="",
sandbox_bin="docker",
)
manager.set_container_session(session)
self.assertEqual(len(close_calls), 1)
session.workspace_path = str((self.project_dir / "ws_b").resolve())
manager.set_container_session(session)
self.assertEqual(len(close_calls), 2)
manager.set_container_session(session)
self.assertEqual(len(close_calls), 2)
finally:
manager.close_all_clients = original_close # type: ignore[assignment]
def test_stateful_stdio_server_reuses_persistent_session(self): def test_stateful_stdio_server_reuses_persistent_session(self):
stateful_script = self._write_temp_server( stateful_script = self._write_temp_server(
"stateful_mcp_server.py", "stateful_mcp_server.py",