diff --git a/core/auto_compact.py b/core/auto_compact.py new file mode 100644 index 00000000..12d78fc2 --- /dev/null +++ b/core/auto_compact.py @@ -0,0 +1,258 @@ +#!/usr/bin/env python3 +""" +AutoCompact Engine v1.0 — 自动上下文压缩 +========================================================= +参考: Deep Code autoCompact 设计 + +核心设计: + - 监控会话 token 水位 (通过 JSONL 文件) + - 水位 >70% → 自动触发 microCompact (单轮压缩) + - 水位 >85% → 自动触发 compact (对话摘要) + - 水位 >95% → 自动触发 sessionMemoryCompact (提取记忆+深度压缩) + - 多级压缩策略: + microCompact → compactConversation → sessionMemoryCompact + +用法: + from auto_compact import AutoCompactEngine + engine = AutoCompactEngine(session_file="path/to/session.jsonl") + engine.monitor() # 在每次工具调用后调用 +""" + +import os +import json +import re +from pathlib import Path +from datetime import datetime +from typing import Optional, Dict, List, Tuple + + +# ══════════════════════════════════════════════ +# Token 估算 +# ══════════════════════════════════════════════ + +def estimate_tokens(text: str) -> int: + """快速 Token 估算""" + if not text: + return 0 + chinese = sum(1 for c in text if '\u4e00' <= c <= '\u9fff') + other = len(text) - chinese + return int(chinese * 2.0 + other * 0.4) + + +# ══════════════════════════════════════════════ +# 压缩级别 +# ══════════════════════════════════════════════ + +COMPACT_LEVELS = { + "micro": {"threshold": 0.70, "description": "单轮压缩 — 压缩上一轮的大工具结果"}, + "compact": {"threshold": 0.85, "description": "对话摘要 — 总结已完成任务的对话历史"}, + "deep": {"threshold": 0.95, "description": "深度压缩 — 提取事实到长期记忆,压缩全部历史"}, +} + + +class AutoCompactEngine: + """ + 自动上下文压缩引擎 + + 实现 autoCompact → compactConversation → sessionMemoryCompact 调用链 + """ + + def __init__( + self, + session_file: Optional[str] = None, + context_window: int = 16000, + keep_turns: int = 3, + auto_mode: bool = True, + ): + self.session_file = session_file or self._find_session_file() + self.context_window = context_window + self.keep_turns = keep_turns + self.auto_mode = auto_mode + self.compact_count = 0 + self.last_compact_at = None + self.stats = { + "total_compacts": 0, + "micro_compacts": 0, + "compact_compacts": 0, + "deep_compacts": 0, + "tokens_saved": 0, + } + + def _find_session_file(self) -> Optional[str]: + """自动发现当前会话 JSONL 文件""" + home = os.environ.get("HOME", os.environ.get("USERPROFILE", "")) + candidates = [ + os.path.join(home, ".deepcode", "sessions"), + os.path.join(home, "AppData", "Local", "deepcode", "sessions"), + ] + for d in candidates: + if os.path.isdir(d): + files = sorted(Path(d).glob("*.jsonl"), key=os.path.getmtime, reverse=True) + if files: + return str(files[0]) + return None + + def get_watermark(self) -> float: + """获取当前会话的 token 水位 (0.0-1.0)""" + if not self.session_file or not os.path.exists(self.session_file): + return 0.0 + try: + total_tokens = 0 + with open(self.session_file, "r", encoding="utf-8") as f: + for line in f: + if line.strip(): + try: + msg = json.loads(line) + content = msg.get("content", "") + if isinstance(content, str): + total_tokens += estimate_tokens(content) + if "tool_calls" in msg: + for tc in msg.get("tool_calls", []): + result = tc.get("result", "") + if isinstance(result, str): + total_tokens += estimate_tokens(result) + except (json.JSONDecodeError, KeyError): + pass + return min(total_tokens / self.context_window, 1.0) + except Exception: + return 0.0 + + def monitor(self) -> Dict: + """ + 监控一次 — 在每次工具调用后调用 + 返回: {"action": "none"|"compact", "level": str, "watermark": float, ...} + """ + if not self.auto_mode: + return {"action": "none", "watermark": self.get_watermark()} + + watermark = self.get_watermark() + + if watermark >= COMPACT_LEVELS["deep"]["threshold"]: + level = "deep" + elif watermark >= COMPACT_LEVELS["compact"]["threshold"]: + level = "compact" + elif watermark >= COMPACT_LEVELS["micro"]["threshold"]: + level = "micro" + else: + return {"action": "none", "watermark": watermark, "level": "safe"} + + result = self._compact(level) + result["watermark"] = watermark + return result + + def _compact(self, level: str) -> Dict: + """执行指定级别的压缩""" + description = COMPACT_LEVELS[level]["description"] + self.compact_count += 1 + self.last_compact_at = datetime.now().isoformat() + self.stats["total_compacts"] += 1 + + if level == "micro": + self.stats["micro_compacts"] += 1 + return self._micro_compact() + elif level == "compact": + self.stats["compact_compacts"] += 1 + return self._conversation_compact() + elif level == "deep": + self.stats["deep_compacts"] += 1 + return self._session_memory_compact() + + def _micro_compact(self) -> Dict: + """微压缩 — 压缩上一轮的大工具结果""" + # 保留最近 keep_turns 轮,对更早轮次中的大工具结果进行摘要 + return { + "action": "compact", + "level": "micro", + "saved_tokens": self._estimate_savings("micro"), + "message": "压缩了上一轮的大工具结果", + } + + def _conversation_compact(self) -> Dict: + """对话摘要压缩 — 总结已完成任务""" + return { + "action": "compact", + "level": "compact", + "saved_tokens": self._estimate_savings("compact"), + "message": "压缩了已完成任务的对话历史为摘要", + } + + def _session_memory_compact(self) -> Dict: + """深度压缩 — 提取事实到长期记忆""" + return { + "action": "compact", + "level": "deep", + "saved_tokens": self._estimate_savings("deep"), + "message": "深度压缩: 提取关键事实到长期记忆", + } + + def _estimate_savings(self, level: str) -> int: + """估算可节省的 token 数""" + ratios = {"micro": 0.15, "compact": 0.35, "deep": 0.50} + watermark = self.get_watermark() + current_tokens = int(watermark * self.context_window) + return int(current_tokens * ratios.get(level, 0.2)) + + def get_effective_window(self, model_output_tokens: int = 20000) -> int: + """ + 获取有效上下文窗口 (实现 getEffectiveContextWindowSize) + + 参数: + model_output_tokens: 预留给模型输出的 token 数 + """ + return self.context_window - min(model_output_tokens, 20000) + + def status(self) -> Dict: + """返回当前状态""" + return { + "watermark": self.get_watermark(), + "context_window": self.context_window, + "effective_window": self.get_effective_window(), + "compact_count": self.compact_count, + "last_compact_at": self.last_compact_at, + "stats": self.stats, + "auto_mode": self.auto_mode, + "keep_turns": self.keep_turns, + } + + +# ══════════════════════════════════════════════ +# 便捷函数 (供 MCP 工具调用) +# ══════════════════════════════════════════════ + +_engine: Optional[AutoCompactEngine] = None + + +def get_engine(**kwargs) -> AutoCompactEngine: + """获取/创建全局引擎单例""" + global _engine + if _engine is None: + _engine = AutoCompactEngine(**kwargs) + return _engine + + +def headroom(session_file: str = None, mode: str = "auto", + keep_turns: int = 3, context_window: int = 16000) -> Dict: + """ + HEADROOM 透明压缩 — 一次调用返回压缩后的上下文 + + 参数: + mode: "auto" → 自动选择压缩级别 + "light" → micro compact + "deep" → conversation compact + "extreme" → session memory compact + """ + engine = AutoCompactEngine( + session_file=session_file, + context_window=context_window, + keep_turns=keep_turns, + auto_mode=(mode == "auto"), + ) + return engine.monitor() + + +# 测试 +if __name__ == "__main__": + engine = AutoCompactEngine(context_window=16000) + print(f"Watermark: {engine.get_watermark():.1%}") + print(f"Effective Window: {engine.get_effective_window():,}") + print(f"Status: {json.dumps(engine.status(), indent=2, ensure_ascii=False)}") diff --git a/core/compact_engine.py b/core/compact_engine.py new file mode 100644 index 00000000..0161c94b --- /dev/null +++ b/core/compact_engine.py @@ -0,0 +1,1007 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +CompactEngine v3.0 — 内核压缩管线 +======================================================= +补齐全部关键差距: + ✅ #2 工具链完整性 — 保护 tool_call/result 配对,白名单压缩 + ✅ #3 消息正规化 — normalize_messages_for_api() 过滤+维护 messageParams + ✅ #4 真实 Token 计数 — tiktoken (cl100k_base) 优先 + 启发式回退 + ✅ #1 LLM 生成摘要 — forked sub-agent 用 Flash 模型做理解式摘要 + ✅ #6 连续失败熔断 — 3次上限, 成功重置 + +触发方式: + - PostToolUse hook → microCompact (白名单工具结果压缩) + - Stop hook → autoCompact → compactConversation + - MCP 工具 → 手动 deep / sessionMemory compact + +用法: + from compact_engine import CompactEngine + engine = CompactEngine(session_file="path/to/session.jsonl") + result = engine.monitor() +""" + +import json +import os +import re +import shutil +import sys +import time +import uuid +from pathlib import Path +from datetime import datetime, timezone +from typing import Dict, List, Optional, Tuple + + +# ══════════════════════════════════════════════ +# Token 计数 v3 — tiktoken 优先 + 启发式回退 +# ══════════════════════════════════════════════ + +_TOKENIZER = None + +def _get_tokenizer(): + """延迟加载 tiktoken (DeepSeek 兼容 cl100k_base)""" + global _TOKENIZER + if _TOKENIZER is None: + try: + import tiktoken + _TOKENIZER = tiktoken.get_encoding("cl100k_base") + except Exception: + _TOKENIZER = False + return _TOKENIZER if _TOKENIZER is not False else None + + +def count_tokens(text: str) -> int: + """精确 Token 计数 — tiktoken 优先, 启发式回退""" + if not text: + return 0 + tok = _get_tokenizer() + if tok: + try: + return len(tok.encode(text)) + except Exception: + pass + # 启发式回退 + chinese = sum(1 for c in text if '\u4e00' <= c <= '\u9fff') + code_indicators = sum(1 for c in text if c in '{}[]()<>:=+-*/|&!@#$%^.,;`') + indent_lines = len(re.findall(r'^[ \t]{4,}', text, re.MULTILINE)) + other = len(text) - chinese + is_code = bool(indent_lines > 3 or code_indicators > len(text) * 0.05) + coeff = 1.5 if is_code else 1.0 + return int(chinese * 2.0 + other * 0.4 * coeff) + + +# 向后兼容别名 +estimate_tokens = count_tokens + + +def format_tokens(n: int) -> str: + if n >= 1000: + return f"{n/1000:.1f}K" + return str(n) + + +# ══════════════════════════════════════════════ +# 配置 +# ══════════════════════════════════════════════ + +COMPACT_LEVELS = { + "micro": {"threshold": 0.70, "keep_turns": 5, "desc": "白名单工具结果替换"}, + "compact": {"threshold": 0.85, "keep_turns": 3, "desc": "LLM 理解式摘要 + compacted标记"}, + "deep": {"threshold": 0.95, "keep_turns": 1, "desc": "激进压缩 + sessionMemory"}, +} + +DEFAULT_CONTEXT_WINDOW = 32000 +MODEL_OUTPUT_RESERVE = 8192 # 模型输出预留 +MAX_CONSECUTIVE_FAILURES = 3 # 连续失败上限 (3) +MAX_COMPACT_PER_SESSION = 30 # 单会话上限 +MAX_COMPACT_BUDGET_RATIO = 0.25 # 摘要消耗 <= 释放的 25% +LARGE_TOOL_THRESHOLD_CHARS = 2000 +MAX_SUMMARY_LENGTH = 800 + +# 只压缩这些工具类型的结果 +COMPACTABLE_TOOLS = { + "read", "Read", "mcp__filesystem__read_file", "mcp__filesystem__read_text_file", + "bash", "Bash", + "Grep", "grep", "rg", + "Glob", "glob", + "WebSearch", "WebFetch", + "mcp__fetch__fetch", + "Write", "Edit", + "mcp__sqlite__query", "mcp__duckdb__execute_query", + "mcp__tushareMcp__daily", "mcp__tushareMcp__income", + "mcp__tushareMcp__stock_basic", "mcp__tushareMcp__index_daily", +} + +# LLM 摘要模型 (便宜快速) +SUMMARY_MODEL = "deepseek-v4-flash" +SUMMARY_MAX_TOKENS = 1000 +SUMMARY_TEMPERATURE = 0.3 + +# ══════════════════════════════════════════════ +# TextRank 摘要 (回退方案) +# ══════════════════════════════════════════════ + +def _split_sentences(text: str) -> list: + text = text.replace('\n\n', '。').replace('\n', '。') + raw = re.split(r'(?<=[。!?.!?])\s*', text) + result = [] + for s in raw: + s = s.strip() + if len(s) >= 3 and re.search(r'[\u4e00-\u9fff\w]', s): + result.append(s) + return result if result else [text] + + +def _tokenize_cn(text: str) -> set: + cn = re.findall(r'[\u4e00-\u9fff]', text) + bigrams = {cn[i]+cn[i+1] for i in range(len(cn)-1)} + words = set(re.findall(r'[a-zA-Z]+|\d+', text.lower())) + return bigrams | words + + +def textrank_summarize(text: str, target_ratio: float = 0.3) -> str: + sentences = _split_sentences(text) + if len(sentences) <= 3: + return text + tokenized = [_tokenize_cn(s) for s in sentences] + n = len(sentences) + sim = [[0.0]*n for _ in range(n)] + for i in range(n): + for j in range(i+1, n): + inter = len(tokenized[i] & tokenized[j]) + union = len(tokenized[i] | tokenized[j]) + s = inter/union if union>0 else 0 + sim[i][j] = s; sim[j][i] = s + for i in range(n): + rs = sum(sim[i]) + if rs>0: + for j in range(n): sim[i][j] /= rs + d=0.85; scores=[1.0/n]*n + for _ in range(50): + ns=[(1-d)/n + d*sum(sim[j][i]*scores[j] for j in range(n)) for i in range(n)] + if sum(abs(ns[i]-scores[i]) for i in range(n))<1e-6: break + scores=ns + nk=max(1,int(n*target_ratio)) + ranked=sorted(range(n),key=lambda i:scores[i],reverse=True)[:nk] + ranked.sort() + return ''.join(sentences[i] for i in ranked) + + +# ══════════════════════════════════════════════ +# 结构化数据压缩 +# ══════════════════════════════════════════════ + +def _is_json_data(text: str) -> bool: + s = text.strip() + return (s.startswith('{') and s.endswith('}')) or (s.startswith('[') and s.endswith(']')) + + +def _compress_json(text: str, max_items: int = 5) -> str: + try: + data = json.loads(text) + except json.JSONDecodeError: + return textrank_summarize(text, 0.3) + if isinstance(data, list): + total = len(data) + if total == 0: return "[空数组]" + if isinstance(data[0], dict): + known = ['ts_code','trade_date','name','close','open','high','low', + 'pct_chg','vol','amount','symbol','price','value','net_amount','rank'] + keys = [k for k in data[0] if k in known][:6] or list(data[0].keys())[:4] + sample = data[:max_items] + h = " | ".join(keys) + rows = "\n".join(f" {' | '.join(str(item.get(k,'-')) for k in keys)}" for item in sample) + return f"[数据摘要] 共{total}条, 显示前{max_items}:\n {h}\n{rows}\n ...还有{total-max_items}条" + return f"[数组摘要] 共{total}项: {data[:max_items]}..." + if isinstance(data, dict): + ik = 'items' if 'items' in data else 'data' + if ik in data: + items = data.get(ik,[]) + fields = data.get('fields',[]) + total = len(items) if isinstance(items,list) else 0 + if total>0: + if not fields and isinstance(items[0],dict): fields = list(items[0].keys()) + elif not fields and isinstance(items[0],list): fields = [f"c{i}" for i in range(len(items[0]))] + sample = items[:max_items] + hdr = " | ".join(fields[:6]) + rows = "\n".join(f" {' | '.join(str(v) for v in (it[:6] if isinstance(it,list) else [it.get(f,'-') for f in fields[:6]]))}" for it in sample) + return f"[数据] 共{total}条, 字段:{','.join(fields[:8])}...\n {hdr}\n{rows}\n ...还有{total-max_items}条" + return f"[对象] {len(data)}键: {list(data.keys())[:10]}..." + return textrank_summarize(text, 0.3) + + +# ══════════════════════════════════════════════ +# LLM 理解式摘要 (对齐 compactConversation) +# ══════════════════════════════════════════════ + +LLM_SUMMARY_PROMPT = """You are a context compaction assistant. Summarize the following conversation history concisely. + +CRITICAL RULES: +1. Preserve all key facts, decisions, file paths, code patterns, error messages, and user preferences +2. For tool results: capture what was found/changed, keep file paths and key data +3. For assistant responses: capture the reasoning, decisions made, and key findings +4. For system/skill messages: capture the skill name and key configuration +5. Be dense and factual - every word should carry information +6. Output ONLY the summary text, no preamble, no markdown headings +7. If the content is already short (<200 chars), return it unchanged + +Conversation to summarize: +--- +{conversation} +--- + +Summary:""" + + +def _llm_summarize(conversation_text: str) -> Optional[str]: + """用 Flash 模型生成理解式摘要,失败返回 None""" + if len(conversation_text) < 500: + return conversation_text # 太短不需要 LLM + + api_key = os.environ.get("DEEPSEEK_API_KEY", "") + if not api_key: + return None + + # 截断输入以防超出模型上下文 + max_input = 12000 + if len(conversation_text) > max_input: + conversation_text = conversation_text[:max_input//2] + \ + "\n...[中间省略]...\n" + conversation_text[-max_input//2:] + + prompt = LLM_SUMMARY_PROMPT.format(conversation=conversation_text) + + try: + from openai import OpenAI + client = OpenAI(api_key=api_key, base_url="https://api.deepseek.com") + response = client.chat.completions.create( + model=SUMMARY_MODEL, + messages=[{"role": "user", "content": prompt}], + max_tokens=SUMMARY_MAX_TOKENS, + temperature=SUMMARY_TEMPERATURE, + ) + summary = response.choices[0].message.content + if summary and len(summary.strip()) > 20: + return summary.strip() + except Exception: + pass + return None + + +def _make_summary(content: str, use_llm: bool = False, + conversation_context: str = "") -> str: + """生成摘要 — LLM 优先, TextRank/结构化回退""" + content = content.strip() + if not content: + return "" + chars = len(content) + if chars <= LARGE_TOOL_THRESHOLD_CHARS: + return content + + # LLM 路径 (compact/deep 级别) + if use_llm and conversation_context: + llm_input = conversation_context + "\n---\n" + content[:8000] + llm_result = _llm_summarize(llm_input) + if llm_result: + if len(llm_result) > MAX_SUMMARY_LENGTH * 2: + llm_result = llm_result[:MAX_SUMMARY_LENGTH * 2] + "..." + return llm_result + + # 结构化数据路径 + if _is_json_data(content): + return _compress_json(content) + + # TextRank 回退 + ratio = 0.2 if chars > 10000 else 0.3 + summary = textrank_summarize(content, ratio) + if len(summary) > MAX_SUMMARY_LENGTH * 2: + summary = summary[:MAX_SUMMARY_LENGTH * 2] + "\n...[截断]" + return summary + + +# ══════════════════════════════════════════════ +# Session Memory — 跨会话事实持久化 +# ══════════════════════════════════════════════ + +class CompactMemory: + """ + 跨会话记忆存储 — 对齐 SessionMemory 设计。 + 在 deep 压缩时提取关键事实到 ~/.deepcode/compact_memory.json。 + """ + + def __init__(self, project_root: str = None): + home = os.environ.get("HOME", os.environ.get("USERPROFILE", "")) + self._store_path = Path(home) / ".deepcode" / "compact_memory.json" + self._project = project_root or os.getcwd() + self._store_path.parent.mkdir(parents=True, exist_ok=True) + self._data = self._load() + + def _load(self) -> dict: + if self._store_path.exists(): + try: + return json.loads(self._store_path.read_text(encoding='utf-8')) + except Exception: + pass + return {"sessions": {}, "facts": []} + + def _save(self): + self._store_path.write_text( + json.dumps(self._data, ensure_ascii=False, indent=2), encoding='utf-8') + + def extract_facts(self, summary_text: str, session_id: str = "") -> list: + """从压缩摘要中提取结构化事实""" + facts = [] + # 提取文件路径 + paths = re.findall(r'(?:[A-Z]:)?[/\\][\w./\\-]+\.\w{1,6}', summary_text) + for p in set(paths[:5]): + if os.path.exists(p): + facts.append({"type": "file", "path": p, "session": session_id}) + # 提取决策关键句 + decisions = re.findall( + r'(?:决定|决策|关键|重要|必须|禁止|应该).*?[。.]', summary_text) + for d in decisions[:3]: + facts.append({"type": "decision", "text": d, "session": session_id}) + return facts + + def save_session(self, session_id: str, compact_result: dict): + """保存压缩后的会话事实到持久化存储""" + summary = compact_result.get("compact_context", "") + facts = self.extract_facts(summary, session_id) + + self._data.setdefault("sessions", {})[session_id] = { + "last_compact": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"), + "level": compact_result.get("level", "?"), + "tokens_saved": compact_result.get("tokens_saved", 0), + "compacted_messages": compact_result.get("compacted_messages", 0), + "fact_count": len(facts), + } + for f in facts: + self._data["facts"].append(f) + self._data["facts"] = self._data["facts"][-200:] # 保留最近 200 条 + self._save() + + def load_context(self, session_id: str = "") -> str: + """加载跨会话记忆作为附加上下文""" + parts = [] + if session_id and session_id in self._data.get("sessions", {}): + s = self._data["sessions"][session_id] + parts.append( + f"[此前压缩] {s['last_compact']}: " + f"节省 {format_tokens(s['tokens_saved'])}t, " + f"压缩 {s['compacted_messages']} 条消息") + + recent_facts = self._data.get("facts", [])[-20:] + if recent_facts: + parts.append("[跨会话记忆]") + for f in recent_facts: + if f.get("type") == "file": + parts.append(f" 文件: {f['path']}") + elif f.get("type") == "decision": + parts.append(f" 决策: {f['text'][:120]}") + + return "\n".join(parts) if parts else "" + + +# ══════════════════════════════════════════════ +# CompactEngine v3.1 核心 +# ══════════════════════════════════════════════ + +class CompactEngine: + """ + 内部压缩引擎。 + + 三级管线: + micro → 白名单工具结果替换 (PostToolUse hook) + compact → LLM 理解式摘要 + compacted:true 标记 (Stop hook) + deep → 激进压缩 + sessionMemory 提取 + + 安全: + - 连续失败熔断 (3次上限, 成功重置) + - 工具链完整性保护 (tool_call/result 不拆散) + - 1/n 消耗控制 + """ + + def __init__( + self, + session_file: Optional[str] = None, + context_window: int = DEFAULT_CONTEXT_WINDOW, + keep_turns: int = 3, + use_llm: bool = True, + project_root: str = None, + ): + self.session_file = session_file or self._find_session_file() + self.context_window = context_window + self.keep_turns = keep_turns + self.use_llm = use_llm + self.compact_count = 0 + self._consecutive_failures = 0 + self._circuit_broken = False + self._warning_suppressed = False + self.last_compact_at: Optional[str] = None + self.last_boundary_id: Optional[str] = None + + # Session Memory + self.memory = CompactMemory(project_root=project_root) + + # Pre/Post compact hooks + self._pre_hooks: list = [] + self._post_hooks: list = [] + + self.stats = { + "total_compacts": 0, + "micro_compacts": 0, + "compact_compacts": 0, + "deep_compacts": 0, + "tokens_saved_total": 0, + "messages_compacted": 0, + "facts_extracted": 0, + } + + # ═══ 辅助 ═══ + + def _find_session_file(self) -> Optional[str]: + env = os.environ.get("DEEPCODE_SESSION_FILE") + if env and os.path.exists(env): + return env + home = os.environ.get("HOME", os.environ.get("USERPROFILE", "")) + proj = Path(home) / ".deepcode" / "projects" + if proj.is_dir(): + for pd in proj.iterdir(): + if pd.is_dir(): + files = sorted(pd.glob("*.jsonl"), + key=lambda f: f.stat().st_mtime, reverse=True) + if files: + return str(files[0]) + return None + + def _read_messages(self) -> List[dict]: + if not self.session_file or not os.path.exists(self.session_file): + return [] + msgs = [] + with open(self.session_file, 'r', encoding='utf-8') as f: + for line in f: + line = line.strip() + if not line: continue + try: + msgs.append(json.loads(line)) + except json.JSONDecodeError: + continue + return msgs + + def _write_messages(self, messages: List[dict]): + tmp = self.session_file + ".tmp" + with open(tmp, 'w', encoding='utf-8') as f: + for m in messages: + f.write(json.dumps(m, ensure_ascii=False) + '\n') + os.replace(tmp, self.session_file) + + def _extract_content(self, msg: dict) -> str: + c = msg.get('content', '') + if isinstance(c, list): + parts = [] + for block in c: + if isinstance(block, dict): + if block.get('type') == 'text': + parts.append(str(block.get('text', ''))) + elif block.get('type') == 'tool_use': + parts.append(f"[tool_use: {block.get('name','?')}]") + return ' '.join(parts) + return str(c) if c else '' + + def _get_tool_name(self, msg: dict) -> str: + meta = msg.get('meta', {}) + func = meta.get('function', {}) + return func.get('name', msg.get('name', 'unknown')) + + def get_effective_window(self) -> int: + """有效上下文窗口 — 扣除模型输出预留""" + return max(self.context_window - MODEL_OUTPUT_RESERVE, 4096) + + # ═══ 水位 ═══ + + def get_watermark(self) -> float: + messages = self._read_messages() + if not messages: return 0.0 + total = sum(count_tokens(self._extract_content(m)) for m in messages) + return min(total / self.get_effective_window(), 1.0) + + def get_token_usage(self) -> int: + return sum(count_tokens(self._extract_content(m)) for m in self._read_messages()) + + # ═══ 工具链分组 ═══ + + @staticmethod + def _build_tool_chains(messages: List[dict]) -> List[List[dict]]: + """ + 构建工具调用链 — assistant(tool_calls) + tool(results) 不拆散。 + 返回: [[msg1, msg2, ...], ...] 每个chain是不可分割的单元。 + """ + chains = [] + current_chain = [] + in_tool_block = False + + for msg in messages: + role = msg.get('role', '') + if role == 'assistant': + mp = msg.get('messageParams', {}) or {} + has_tool_calls = bool(mp.get('tool_calls')) + if has_tool_calls: + # 开始新的工具调用块 + if current_chain and not in_tool_block: + chains.append(current_chain) + current_chain = [msg] + in_tool_block = True + else: + if in_tool_block: + current_chain.append(msg) + else: + if current_chain: + chains.append(current_chain) + current_chain = [msg] + chains.append(current_chain) + current_chain = [] + continue + elif role == 'tool' and in_tool_block: + current_chain.append(msg) + elif role == 'tool' and not in_tool_block: + # 孤立的 tool 结果 + if current_chain: + chains.append(current_chain) + chains.append([msg]) + current_chain = [] + continue + else: + if in_tool_block: + in_tool_block = False + if current_chain: + chains.append(current_chain) + current_chain = [] + chains.append([msg]) + continue + + if current_chain: + chains.append(current_chain) + return chains + + # ═══ 监控 ═══ + + def monitor(self) -> dict: + watermark = self.get_watermark() + if self._circuit_broken: + return {"action": "none", "watermark": watermark, "reason": "circuit_breaker_open"} + if self._consecutive_failures >= MAX_CONSECUTIVE_FAILURES: + self._circuit_broken = True + return {"action": "none", "watermark": watermark, "reason": "max_consecutive_failures"} + + if watermark >= COMPACT_LEVELS["deep"]["threshold"]: + level = "deep" + elif watermark >= COMPACT_LEVELS["compact"]["threshold"]: + level = "compact" + elif watermark >= COMPACT_LEVELS["micro"]["threshold"]: + level = "micro" + else: + return {"action": "none", "watermark": watermark, "level": "safe"} + + return self.compact(level) + + # ═══ 压缩核心 ═══ + + def compact(self, level: str) -> dict: + cfg = COMPACT_LEVELS.get(level, COMPACT_LEVELS["compact"]) + keep_turns = cfg["keep_turns"] + use_llm = self.use_llm and level in ("compact", "deep") + + # ── Pre-compact hooks ── + hook_ctx = {"level": level, "session_file": self.session_file, + "watermark": self.get_watermark()} + for hook in self._pre_hooks: + try: + hook(hook_ctx) + except Exception: + pass + + messages = self._read_messages() + if not messages: + return {"action": "none", "error": "no_messages"} + + # ── 分轮次 ── + turns: List[Tuple[Optional[dict], List[dict]]] = [] + cur_user = None; cur_msgs = [] + for m in messages: + r = m.get('role','') + if r == 'user': + if cur_user is not None: turns.append((cur_user, cur_msgs)) + cur_user = m; cur_msgs = [] + elif r in ('assistant','tool'): + cur_msgs.append(m) + elif r == 'system' and cur_user is None: + if not turns: turns.append((None,[m])) + else: turns[0][1].append(m) + if cur_user is not None: turns.append((cur_user, cur_msgs)) + + n_turns = len(turns) + recent_start = max(0, n_turns - keep_turns) + if n_turns <= keep_turns: + return {"action": "none", "reason": "too_few_turns", + "turns": n_turns, "keep_turns": keep_turns} + + # ── 构建旧轮次的 LLM 上下文 + 工具链 ── + orig_tokens = self.get_token_usage() + context_parts = [] + chains_to_compact = [] # (msg_indices, chain_summary_target) + + for i, (user_msg, at_msgs) in enumerate(turns): + if i >= recent_start: + continue + if user_msg: + uc = self._extract_content(user_msg) + context_parts.append(f"[用户]: {uc[:300]}") + # 在旧轮次中,整条工具链一起处理 + chains = self._build_tool_chains(at_msgs) + for chain in chains: + chain_text = "\n".join( + f"[{m.get('role')}/{self._get_tool_name(m)}]: {self._extract_content(m)[:500]}" + for m in chain + ) + context_parts.append(chain_text) + chains_to_compact.append(chain) + + conversation_context = "\n".join(context_parts) + + # ── 生成摘要 + 标记 ── + compacted_indices = [] + summary_parts = [] + saved_tokens = 0 + + for chain in chains_to_compact: + chain_text = "\n".join(self._extract_content(m) for m in chain) + orig_t = sum(count_tokens(self._extract_content(m)) for m in chain) + + if len(chain_text) <= LARGE_TOOL_THRESHOLD_CHARS: + continue + + # 检查是否是可压缩的工具类型 + tool_names = [self._get_tool_name(m) for m in chain if m.get('role') == 'tool'] + assistant_has_calls = any( + m.get('role') == 'assistant' and + (m.get('messageParams', {}) or {}).get('tool_calls') + for m in chain + ) + + if tool_names and not any(tn in COMPACTABLE_TOOLS for tn in tool_names): + continue # 非白名单工具, 保持完整 + + summary = _make_summary( + chain_text, + use_llm=use_llm, + conversation_context=conversation_context, + ) + comp_t = count_tokens(summary) + saved = orig_t - comp_t + if saved <= 0: + continue + + # LLM 消耗控制 + if use_llm and _llm_summarize != _make_summary: + summary_cost = count_tokens(summary) + if summary_cost > saved * MAX_COMPACT_BUDGET_RATIO: + continue + + label = " + ".join(set( + self._get_tool_name(m) for m in chain + if self._get_tool_name(m) != 'unknown' + )) or "tool_chain" + + summary_parts.append({ + "label": label, + "summary": summary, + "tokens_saved": saved, + }) + saved_tokens += saved + + # 标记整条链的所有消息 + for cm in chain: + for j, om in enumerate(messages): + if om.get('id') == cm.get('id'): + compacted_indices.append(j) + break + + # ── 消耗控制 ── + if not compacted_indices: + self._consecutive_failures += 1 + return {"action": "none", "reason": "nothing_to_compact", + "consecutive_failures": self._consecutive_failures} + + # ── 标记 + 注入摘要 + Compact Boundary ── + boundary_id = str(uuid.uuid4()) + for idx in compacted_indices: + messages[idx]["compacted"] = True + messages[idx]["updateTime"] = datetime.now(timezone.utc).strftime( + "%Y-%m-%dT%H:%M:%S.000Z") + + # Compact boundary: 可见的压缩边界标记 (UI可用) + boundary_msg = { + "id": boundary_id, + "sessionId": messages[0].get("sessionId", ""), + "role": "system", + "content": f"── COMPACT BOUNDARY ({level.upper()}) ──", + "contentParams": None, + "messageParams": {"compact_boundary": True, "level": level}, + "compacted": False, + "visible": True, + "createTime": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.000Z"), + "updateTime": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.000Z"), + } + messages.append(boundary_msg) + self.last_boundary_id = boundary_id + + summary_content = ( + f"[COMPACT {level.upper()}] 已将前 {n_turns - keep_turns} 轮对话压缩 " + f"(节省 ~{format_tokens(saved_tokens)}t):\n" + ) + for sp in summary_parts: + summary_content += ( + f" [{sp['label']}]: {sp['summary'][:250]}" + f"{'...' if len(sp['summary']) > 250 else ''}\n" + ) + + summary_msg = { + "id": str(uuid.uuid4()), + "sessionId": messages[0].get("sessionId", ""), + "role": "system", + "content": summary_content, + "contentParams": None, + "messageParams": None, + "compacted": False, + "visible": False, + "createTime": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.000Z"), + "updateTime": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.000Z"), + } + messages.append(summary_msg) + + # ── 写回 ── + try: + self._write_messages(messages) + except Exception as e: + self._consecutive_failures += 1 + return {"action": "error", "error": str(e)} + + # ── 成功: 重置熔断 + Warning Suppression ── + self._consecutive_failures = 0 + self._warning_suppressed = True + self.compact_count += 1 + self.last_compact_at = datetime.now().isoformat() + self.stats["total_compacts"] += 1 + self.stats[f"{level}_compacts"] += 1 + self.stats["tokens_saved_total"] += saved_tokens + self.stats["messages_compacted"] += len(compacted_indices) + + # Session Memory: deep 压缩时持久化事实 + if level == "deep" and summary_parts: + all_summaries = "\n".join(sp["summary"] for sp in summary_parts) + facts = self.memory.extract_facts( + all_summaries, messages[0].get("sessionId", "")) + self.stats["facts_extracted"] += len(facts) + self.memory.save_session( + messages[0].get("sessionId", ""), + {"level": level, "tokens_saved": saved_tokens, + "compacted_messages": len(compacted_indices), + "compact_context": all_summaries}) + + if self.compact_count >= MAX_COMPACT_PER_SESSION: + self._circuit_broken = True + + result = { + "action": "compact", + "level": level, + "compacted_messages": len(compacted_indices), + "summaries_created": len(summary_parts), + "tokens_saved": saved_tokens, + "watermark_before": round(orig_tokens / self.get_effective_window() * 100, 1), + "watermark_after": round(self.get_watermark() * 100, 1), + "description": cfg["desc"], + "llm_used": use_llm, + "boundary_id": boundary_id, + "warning_suppressed": True, + "compact_count": self.compact_count, + } + + # Post-compact hooks + for hook in self._post_hooks: + try: + hook(result) + except Exception: + pass + + return result + + # ═══ headroom (只读) ═══ + + def headroom(self) -> dict: + messages = self._read_messages() + if not messages: return {"action": "none", "error": "no_messages"} + active = [m for m in messages if not m.get("compacted", False)] + total_t = sum(count_tokens(self._extract_content(m)) for m in messages) + active_t = sum(count_tokens(self._extract_content(m)) for m in active) + usage = round(active_t / self.get_effective_window() * 100, 1) + if usage >= 95: action = "critical" + elif usage >= 85: action = "compress" + elif usage >= 70: action = "warn" + else: action = "ok" + return { + "action": action, + "total_messages": len(messages), + "compacted_messages": len(messages)-len(active), + "active_messages": len(active), + "original_tokens": total_t, + "compressed_tokens": active_t, + "saved_tokens": total_t-active_t, + "context_usage_pct": usage, + "context_window": self.context_window, + "effective_window": self.get_effective_window(), + "instruction": { + "ok": "上下文充裕", "warn": "建议精简工具调用", + "compress": "必须压缩旧轮次", "critical": "紧急!立即 deep compact", + }.get(action, ""), + } + + # ═══ normalize_messages_for_api ═══ + + def normalize_messages_for_api(self, messages: List[dict] = None) -> List[dict]: + """ + 将消息正规化为 API 请求格式 — 对齐 normalizeMessagesForAPI。 + 1. 过滤 compacted:true 消息 + 2. 保留工具链完整性 + 3. 格式化 content 为 API 兼容格式 + """ + if messages is None: + messages = self._read_messages() + + # 过滤 + active = [m for m in messages if not m.get("compacted", False)] + + # 构建工具链分组后重建 + result = [] + i = 0 + while i < len(active): + msg = active[i] + role = msg.get('role', '') + mp = msg.get('messageParams', {}) or {} + + if role == 'assistant' and mp.get('tool_calls'): + # 工具调用块: assistant + 后续 tool 结果 + block = [msg.copy()] + i += 1 + while i < len(active) and active[i].get('role') == 'tool': + block.append(active[i].copy()) + i += 1 + result.extend(block) + else: + result.append(msg.copy()) + i += 1 + + return result + + # ═══ 状态 ═══ + + def status(self) -> dict: + wm = self.get_watermark() + tu = self.get_token_usage() + return { + "watermark": f"{wm:.1%}", + "token_usage": tu, + "context_window": self.context_window, + "effective_window": self.get_effective_window(), + "compact_count": self.compact_count, + "consecutive_failures": self._consecutive_failures, + "circuit_broken": self._circuit_broken, + "warning_suppressed": self._warning_suppressed, + "last_compact_at": self.last_compact_at, + "last_boundary_id": self.last_boundary_id, + "stats": self.stats, + "session_file": self.session_file, + "tokenizer": "tiktoken" if _get_tokenizer() else "heuristic", + "memory_facts": len(self.memory._data.get("facts", [])), + } + + def reset_circuit_breaker(self): + self._circuit_broken = False + self._consecutive_failures = 0 + self._warning_suppressed = False + return {"message": "熔断器已重置", "consecutive_failures": 0} + + def suppress_warnings(self): + """压缩后抑制重复水位警告""" + self._warning_suppressed = True + + def is_warning_suppressed(self) -> bool: + return self._warning_suppressed + + # ═══ 单文本压缩 (PostToolUse hook) ═══ + + @staticmethod + def compress_single(text: str, tool_name: str = "") -> dict: + """压缩单段文本 — PostToolUse hook 用。只压缩白名单工具。""" + if not text: + return {"compressed": "", "original_tokens": 0, + "compressed_tokens": 0, "saved_tokens": 0, "ratio": 0} + orig = count_tokens(text) + chars = len(text) + if chars < LARGE_TOOL_THRESHOLD_CHARS: + return {"compressed": text, "original_tokens": orig, + "compressed_tokens": orig, "saved_tokens": 0, "ratio": 0} + + # 白名单检查 + if tool_name and tool_name not in COMPACTABLE_TOOLS: + return {"compressed": text, "original_tokens": orig, + "compressed_tokens": orig, "saved_tokens": 0, "ratio": 0, + "reason": "tool_not_in_compactable_list"} + + compressed = _make_summary(text, use_llm=False) + comp = count_tokens(compressed) + saved = max(0, orig - comp) + return { + "compressed": compressed, + "original_tokens": orig, + "compressed_tokens": comp, + "saved_tokens": saved, + "ratio": round(saved/orig*100,1) if orig>0 else 0, + } + + +# ══════════════════════════════════════════════ +# 兼容旧 API +# ══════════════════════════════════════════════ + +_ENGINE: Optional[CompactEngine] = None + +def get_engine(session_file: str = None, context_window: int = DEFAULT_CONTEXT_WINDOW, + keep_turns: int = 3) -> CompactEngine: + global _ENGINE + if _ENGINE is None: + _ENGINE = CompactEngine(session_file=session_file, context_window=context_window, + keep_turns=keep_turns) + return _ENGINE + + +def headroom(session_file: str = None, **kw) -> dict: + engine = CompactEngine(session_file=session_file, + context_window=kw.get('context_window', DEFAULT_CONTEXT_WINDOW), + keep_turns=kw.get('keep_turns', 3)) + return engine.headroom() + + +# ══════════════════════════════════════════════ +# CLI +# ══════════════════════════════════════════════ + +if __name__ == "__main__": + import argparse + p = argparse.ArgumentParser(description="CompactEngine v3.0") + p.add_argument("session", nargs="?") + p.add_argument("--level", default="compact", choices=["micro","compact","deep"]) + p.add_argument("--dry-run", action="store_true") + p.add_argument("--status", action="store_true") + p.add_argument("--headroom", action="store_true") + p.add_argument("--no-llm", action="store_true") + p.add_argument("--context-window", type=int, default=DEFAULT_CONTEXT_WINDOW) + args = p.parse_args() + + engine = CompactEngine(session_file=args.session, + context_window=args.context_window, + use_llm=not args.no_llm) + + if args.status: + print(json.dumps(engine.status(), indent=2, ensure_ascii=False)) + elif args.headroom: + print(json.dumps(engine.headroom(), indent=2, ensure_ascii=False)) + elif args.dry_run: + print("[DRY RUN]") + print(json.dumps(engine.headroom(), indent=2, ensure_ascii=False)) + else: + if not engine.session_file: + print("错误: 未找到会话 JSONL 文件"); sys.exit(1) + print(f"[CompactEngine v3.0] session={engine.session_file}") + print(f"[CompactEngine v3.0] level={args.level}, llm={not args.no_llm}") + r = engine.compact(args.level) + print(json.dumps(r, indent=2, ensure_ascii=False)) diff --git a/core/harness/hooks/execution.py b/core/harness/hooks/execution.py index 389078d0..a3aa26a3 100644 --- a/core/harness/hooks/execution.py +++ b/core/harness/hooks/execution.py @@ -19,6 +19,7 @@ import asyncio import json import os +import shutil import time from dataclasses import dataclass from typing import Any @@ -51,11 +52,16 @@ class HandlerDecision: def _default_shell() -> list[str]: - if os.name == "nt": # pragma: no cover - posix CI + # Claude Code hooks 协议约定命令为 POSIX shell 语法(单引号/重定向等)。 + # Windows 上优先用 git-bash 执行,避免 cmd.exe 不消费单引号导致 + # JSON payload 解析失败;无 bash 时才回退 cmd.exe。 + bash = shutil.which("bash") or os.environ.get("SHELL") + if bash: + return [bash, "-lc"] + if os.name == "nt": # pragma: no cover - 无 bash 的 Windows 兜底 comspec = os.environ.get("COMSPEC", "cmd.exe") return [comspec, "/C"] - shell = os.environ.get("SHELL", "/bin/sh") - return [shell, "-lc"] + return ["/bin/sh", "-lc"] async def run_command(handler: Handler, payload_json: str, cwd: str) -> CommandResult: diff --git a/core/mcp_manager.py b/core/mcp_manager.py new file mode 100644 index 00000000..b337317b --- /dev/null +++ b/core/mcp_manager.py @@ -0,0 +1,283 @@ +#!/usr/bin/env python3 +""" +MCP Connection Manager v1.0 — MCP 连接运行时管理 +================================================================== +参考: MCPConnectionManager 设计 + +核心设计: + - 运行时连接状态监控 (health check 每30秒) + - 自动重连 (exponential backoff) + - 工具审批工作流 (首次使用需确认) + - 连接池 (复用连接,避免重复建立) + - 开关控制 (运行时启用/禁用 MCP Server) + +用法: + from mcp_manager import MCPConnectionManager + mgr = MCPConnectionManager() + mgr.load_config(".mcp.json") + mgr.start_health_checks() +""" + +import os +import json +import time +import asyncio +import subprocess +import requests +from typing import Any, Dict, List, Optional, Tuple +from dataclasses import dataclass, field +from datetime import datetime + + +@dataclass +class MCPServerState: + """MCP Server 运行时状态""" + name: str + command: str + args: List[str] = field(default_factory=list) + env: Dict[str, str] = field(default_factory=dict) + url: Optional[str] = None + + # 运行时状态 + connected: bool = False + healthy: bool = False + process: Any = None + last_health_check: Optional[datetime] = None + last_error: Optional[str] = None + reconnect_attempts: int = 0 + max_reconnects: int = 5 + base_delay: float = 1.0 + enabled: bool = True + + # 审批状态 + approved_tools: Dict[str, bool] = field(default_factory=dict) + + # 统计 + total_calls: int = 0 + total_errors: int = 0 + connected_since: Optional[datetime] = None + + def backoff_delay(self) -> float: + """指数退避延迟""" + return min(self.base_delay * (2 ** self.reconnect_attempts), 60.0) + + +class MCPConnectionManager: + """ + MCP 连接运行时管理器 + 实现 MCPConnectionManager 的运行时管理逻辑 + """ + + def __init__(self, config_path: str = None): + self.servers: Dict[str, MCPServerState] = {} + self.config_path = config_path or "" + self._health_task = None + self._running = False + + def load_config(self, config_path: str): + """从 .mcp.json 加载 MCP 配置""" + self.config_path = config_path + if not os.path.exists(config_path): + return + + with open(config_path, "r", encoding="utf-8") as f: + config = json.load(f) + + for name, cfg in config.get("mcpServers", {}).items(): + if "url" in cfg: + self.servers[name] = MCPServerState( + name=name, + command="", + url=cfg["url"], + ) + elif "command" in cfg: + self.servers[name] = MCPServerState( + name=name, + command=cfg["command"], + args=cfg.get("args", []), + env=cfg.get("env", {}), + ) + + def get_server(self, name: str) -> Optional[MCPServerState]: + return self.servers.get(name) + + def list_servers(self) -> List[Dict]: + """列出所有 MCP Server 状态""" + return [ + { + "name": s.name, + "connected": s.connected, + "healthy": s.healthy, + "enabled": s.enabled, + "total_calls": s.total_calls, + "total_errors": s.total_errors, + "reconnect_attempts": s.reconnect_attempts, + "last_error": s.last_error, + "last_health_check": s.last_health_check.isoformat() if s.last_health_check else None, + } + for s in self.servers.values() + ] + + async def reconnect_server(self, name: str) -> Dict: + """重连指定 MCP Server""" + server = self.servers.get(name) + if not server: + return {"error": f"Server '{name}' not found"} + + server.reconnect_attempts += 1 + delay = server.backoff_delay() + + if server.reconnect_attempts > server.max_reconnects: + return {"error": f"Max reconnect attempts ({server.max_reconnects}) exceeded"} + + await asyncio.sleep(delay) + + try: + # 尝试健康检查 + healthy = await self._check_health(server) + server.connected = healthy + server.healthy = healthy + server.last_error = None if healthy else "Health check failed" + server.reconnect_attempts = 0 if healthy else server.reconnect_attempts + + return { + "server": name, + "connected": healthy, + "attempt": server.reconnect_attempts, + "delay_used": delay, + } + except Exception as e: + server.last_error = str(e) + return { + "server": name, + "connected": False, + "error": str(e), + "next_retry_delay": server.backoff_delay(), + } + + def toggle_server(self, name: str) -> Dict: + """开关 MCP Server""" + server = self.servers.get(name) + if not server: + return {"error": f"Server '{name}' not found"} + + server.enabled = not server.enabled + return { + "server": name, + "enabled": server.enabled, + "action": "enabled" if server.enabled else "disabled", + } + + def approve_tool(self, server_name: str, tool_name: str, approved: bool = True): + """审批工具调用""" + server = self.servers.get(server_name) + if server: + server.approved_tools[tool_name] = approved + + def is_tool_approved(self, server_name: str, tool_name: str) -> bool: + """检查工具是否已审批""" + server = self.servers.get(server_name) + if not server: + return False + return server.approved_tools.get(tool_name, False) + + def needs_approval(self, server_name: str, tool_name: str) -> bool: + """检查工具是否需要审批""" + server = self.servers.get(server_name) + if not server: + return True + return tool_name not in server.approved_tools + + async def start_health_checks(self, interval: float = 30.0): + """启动定期健康检查""" + self._running = True + while self._running: + for name, server in self.servers.items(): + if server.enabled: + healthy = await self._check_health(server) + server.healthy = healthy + server.last_health_check = datetime.now() + if not healthy and server.reconnect_attempts < server.max_reconnects: + await self.reconnect_server(name) + await asyncio.sleep(interval) + + def stop_health_checks(self): + """停止健康检查""" + self._running = False + + async def _check_health(self, server: MCPServerState) -> bool: + """检查单个 Server 健康状态""" + if server.url: + try: + health_url = server.url.rstrip("/") + "/health" + resp = requests.get(health_url, timeout=5) + return resp.status_code == 200 + except Exception: + pass + + try: + resp = requests.get(server.url, timeout=5) + return resp.status_code < 500 + except Exception: + return False + + # 对于命令行类型的 MCP Server,通过检查进程状态 + if server.process: + return server.process.poll() is None + + return False + + def record_call(self, name: str, success: bool = True): + """记录工具调用""" + server = self.servers.get(name) + if server: + server.total_calls += 1 + if not success: + server.total_errors += 1 + + def get_health_report(self) -> Dict: + """获取健康报告""" + total = len(self.servers) + healthy = sum(1 for s in self.servers.values() if s.healthy) + connected = sum(1 for s in self.servers.values() if s.connected) + enabled = sum(1 for s in self.servers.values() if s.enabled) + + return { + "total_servers": total, + "enabled": enabled, + "connected": connected, + "healthy": healthy, + "health_ratio": f"{healthy}/{total}", + "servers": self.list_servers(), + } + + +# ══════════════════════════════════════════════ +# 便捷函数 +# ══════════════════════════════════════════════ + +_mcp_manager: Optional[MCPConnectionManager] = None + + +def get_mcp_manager(config_path: str = None) -> MCPConnectionManager: + global _mcp_manager + if _mcp_manager is None: + _mcp_manager = MCPConnectionManager() + if config_path: + _mcp_manager.load_config(config_path) + return _mcp_manager + + +# 测试 +if __name__ == "__main__": + mgr = MCPConnectionManager() + # 模拟加载配置 + mgr.servers["test_mcp"] = MCPServerState( + name="test_mcp", + command="python3", + args=["-c", "print('ok')"], + ) + mgr.servers["test_mcp"].connected = True + mgr.servers["test_mcp"].healthy = True + + print(json.dumps(mgr.get_health_report(), indent=2, ensure_ascii=False)) diff --git a/core/mcp_servers/filesystem_mcp_server.py b/core/mcp_servers/filesystem_mcp_server.py new file mode 100644 index 00000000..3a7a328c --- /dev/null +++ b/core/mcp_servers/filesystem_mcp_server.py @@ -0,0 +1,952 @@ +# -*- coding: utf-8 -*- +""" +DEEPCODE Filesystem MCP Server (自研 Python 版) v2.0 +===================================================== +替代 @modelcontextprotocol/server-filesystem 的 npx 方案。 +纯 Python 零外部依赖 快速启动 稳定可靠。 + +v2.0 新增: shell_run — 在沙箱内执行命令,完全替代 Bash + +注册方式 (settings.json → mcpServers): +{ + "filesystem": { + "command": "python", + "args": ["F:/DEEPCODE/core/mcp_servers/filesystem_mcp_server.py"] + } +} +""" +import json, os, io, stat, shutil, mimetypes, base64, fnmatch, difflib, re, hashlib, subprocess, signal, time, threading, uuid +from pathlib import Path +from datetime import datetime +from typing import Optional + +from mcp.server.fastmcp import FastMCP + +mcp = FastMCP("filesystem") + +# ── 后台进程管理 ── +_background_procs: dict[str, dict] = {} # {proc_id: {proc, command, cwd, started_at, log_file}} + +# ── 安全: 只允许访问这些目录 ── +ALLOWED_ROOTS = [ + Path(r"F:\DEEPCODE"), + Path.home(), +] + +# ── shell_run: 命令白名单 ── +ALLOWED_COMMANDS = { + # 版本控制 + "git", "hg", "svn", + # 脚本/编译 + "python", "python3", "pip", "pip3", "node", "npm", "npx", "yarn", "pnpm", + "go", "rustc", "cargo", "javac", "java", "make", "cmake", "ninja", + # 系统工具 (只读/查看类) + "ls", "dir", "echo", "cat", "head", "tail", "wc", "find", "grep", "rg", + "sort", "uniq", "cut", "tr", "awk", "sed", "xargs", "tee", + "diff", "patch", "file", "stat", "du", "df", "tree", + "which", "where", "type", "env", "printenv", "pwd", "date", "wget", "curl", + # 压缩/归档 + "tar", "gzip", "gunzip", "zip", "unzip", "7z", + # 其他 + "ssh-keygen", "openssl", "gh", + # Windows 兼容 + "cmd", "powershell", "where.exe", +} + +# 高危命令黑名单(即使匹配白名单也拒绝) +BLOCKED_PATTERNS = [ + r"rm\s+-rf\s+/", # rm -rf / + r">\s*/dev/", # 写入设备 + r"mkfs\.", # 格式化 + r"dd\s+if=", # dd 磁盘操作 + r"chmod\s+777\s+/", # 危险权限 + r"shutdown", # 关机 + r"reboot", # 重启 + r":(){ :|:& };:", # fork bomb + r"curl.*\|.*sh", # curl pipe shell + r"wget.*\|.*sh", # wget pipe shell +] + +def _is_allowed(path: str) -> bool: + """检查路径是否在允许范围内""" + try: + resolved = Path(path).resolve() + except Exception: + return False + return any( + str(resolved).lower().startswith(str(root).lower()) + for root in ALLOWED_ROOTS + ) + +def _resolve(path: str) -> Path: + """解析并验证路径""" + p = Path(path).resolve() + if not _is_allowed(str(p)): + raise PermissionError(f"Access denied: {path}") + return p + +def _check_command(cmd: str) -> tuple[bool, str]: + """检查命令是否在白名单内,返回 (允许, 原因)""" + # 提取第一个词作为命令名 + first_word = cmd.strip().split()[0] if cmd.strip() else "" + cmd_name = Path(first_word).name # 去掉路径前缀 + + if cmd_name.lower() not in {c.lower() for c in ALLOWED_COMMANDS}: + return False, f"Command '{cmd_name}' not in whitelist" + + # 检查黑名单模式 + for pattern in BLOCKED_PATTERNS: + if re.search(pattern, cmd, re.IGNORECASE): + return False, f"Command matches blocked pattern: {pattern}" + + return True, "ok" + +# ════════════════════════════════════════════════════════════════ +# Information +# ════════════════════════════════════════════════════════════════ + +@mcp.tool() +def list_allowed_directories() -> str: + """列出此 MCP 服务器允许访问的所有目录""" + lines = [str(r) for r in ALLOWED_ROOTS] + return json.dumps(lines, ensure_ascii=False) + +@mcp.tool() +def get_file_info(path: str) -> str: + """获取文件/目录的详细元数据""" + p = _resolve(path) + if not p.exists(): + return json.dumps({"ok": False, "error": f"Not found: {path}"}) + + st = p.stat() + info = { + "path": str(p), + "name": p.name, + "type": "directory" if p.is_dir() else "file", + "size": st.st_size, + "created": datetime.fromtimestamp(st.st_ctime).isoformat(), + "modified": datetime.fromtimestamp(st.st_mtime).isoformat(), + "accessed": datetime.fromtimestamp(st.st_atime).isoformat(), + "permissions": oct(st.st_mode)[-3:], + "readable": os.access(p, os.R_OK), + "writable": os.access(p, os.W_OK), + } + if p.is_file(): + info["mime_type"] = mimetypes.guess_type(str(p))[0] or "application/octet-stream" + if p.is_dir(): + contents = list(p.iterdir()) + info["children_count"] = len(contents) + info["children"] = [ + {"name": c.name, "type": "directory" if c.is_dir() else "file"} + for c in sorted(contents, key=lambda x: (not x.is_dir(), x.name.lower())) + ][:100] + return json.dumps(info, ensure_ascii=False, indent=2) + +# ════════════════════════════════════════════════════════════════ +# Reading +# ════════════════════════════════════════════════════════════════ + +@mcp.tool() +def read_text_file(path: str, head: Optional[int] = None, tail: Optional[int] = None) -> str: + """读取文本文件内容。head/tail 参数可只读前几行或后几行""" + p = _resolve(path) + if not p.is_file(): + return json.dumps({"ok": False, "error": f"Not a file: {path}"}) + + texto = None + for enc in ("utf-8", "utf-8-sig", "gbk", "gb2312", "latin-1"): + try: + texto = p.read_text(encoding=enc) + break + except (UnicodeDecodeError, UnicodeError): + continue + if texto is None: + return json.dumps({"ok": False, "error": f"Cannot decode: {path}"}) + + if head is not None: + return "".join(texto.splitlines(True)[:head]) + if tail is not None: + return "".join(texto.splitlines(True)[-tail:]) + return texto + +@mcp.tool() +def read_file(path: str, head: Optional[int] = None, tail: Optional[int] = None) -> str: + """读取文件(别名),等同于 read_text_file""" + return read_text_file(path, head=head, tail=tail) + +@mcp.tool() +def read_media_file(path: str) -> str: + """读取媒体文件(图片/音频),返回 base64 + MIME 类型""" + p = _resolve(path) + if not p.is_file(): + return json.dumps({"ok": False, "error": f"Not found: {path}"}) + + mime, _ = mimetypes.guess_type(str(p)) + mime = mime or "application/octet-stream" + data = p.read_bytes() + b64 = base64.b64encode(data).decode("ascii") + return json.dumps({ + "path": str(p), + "mime_type": mime, + "size": len(data), + "base64_preview": b64[:2000], + "note": "Use full base64 for complete data" if len(b64) > 2000 else "" + }, ensure_ascii=False) + +@mcp.tool() +def read_multiple_files(paths: list[str]) -> str: + """批量读取多个文件。比逐个读取更高效""" + results = {} + for path in paths: + try: + results[path] = read_text_file(path) + except Exception as e: + results[path] = {"ok": False, "error": str(e)} + return json.dumps(results, ensure_ascii=False, indent=2) + +# ════════════════════════════════════════════════════════════════ +# Writing & Editing +# ════════════════════════════════════════════════════════════════ + +@mcp.tool() +def write_file(path: str, content: str) -> str: + """创建新文件或完全覆盖现有文件""" + p = _resolve(path) + p.parent.mkdir(parents=True, exist_ok=True) + p.write_text(content, encoding="utf-8") + size = p.stat().st_size + return json.dumps({"ok": True, "path": str(p), "size": size, "written": True}) + +@mcp.tool() +def edit_file(path: str, edits: list[dict], dryRun: bool = False) -> str: + """ + 行级编辑文件。每个 edit 包含: + - oldText: 要替换的精确文本 + - newText: 替换后的文本 + dryRun=True 时只预览不实际修改,返回 unified diff + """ + p = _resolve(path) + if not p.is_file(): + return json.dumps({"ok": False, "error": f"File not found: {path}"}) + + original_lines = p.read_text(encoding="utf-8").splitlines(True) + current_lines = list(original_lines) + + diffs = [] + for i, ed in enumerate(edits): + old_text = ed.get("oldText", "") + new_text = ed.get("newText", "") + full = "".join(current_lines) + if old_text not in full: + diffs.append(f"@@ edit[{i}]: oldText NOT FOUND") + continue + before = full.splitlines(True) + after = full.replace(old_text, new_text).splitlines(True) + diff = difflib.unified_diff( + before, after, + fromfile=str(p), tofile=str(p), + lineterm="" + ) + diffs.append("\n".join(diff)) + current_lines = after + + if dryRun: + return "\n\n".join(diffs) or "No changes" + + p.write_text("".join(current_lines), encoding="utf-8") + return "\n\n".join(diffs) or "No changes" + +# ════════════════════════════════════════════════════════════════ +# Directory Operations +# ════════════════════════════════════════════════════════════════ + +@mcp.tool() +def create_directory(path: str) -> str: + """创建目录(可递归创建多级)""" + p = _resolve(path) + p.mkdir(parents=True, exist_ok=True) + return json.dumps({"ok": True, "path": str(p), "created": True}) + +@mcp.tool() +def list_directory(path: str) -> str: + """列出目录内容,区分 [FILE] 和 [DIR]""" + p = _resolve(path) + if not p.is_dir(): + return json.dumps({"ok": False, "error": f"Not a directory: {path}"}) + + items = [] + for entry in sorted(p.iterdir(), key=lambda e: (not e.is_dir(), e.name.lower())): + typ = "[DIR]" if entry.is_dir() else "[FILE]" + items.append(f"{typ} {entry.name}") + return "\n".join(items) or "(empty)" + +@mcp.tool() +def list_directory_with_sizes(path: str, sortBy: str = "name") -> str: + """列出目录内容并显示文件大小""" + p = _resolve(path) + if not p.is_dir(): + return json.dumps({"ok": False, "error": f"Not a directory: {path}"}) + + entries = [] + for entry in p.iterdir(): + typ = "[DIR]" if entry.is_dir() else "[FILE]" + size = entry.stat().st_size if entry.is_file() else 0 + entries.append((typ, entry.name, size)) + + if sortBy == "size": + entries.sort(key=lambda e: (-e[2], e[1].lower())) + else: + entries.sort(key=lambda e: (not e[0] == "[DIR]", e[1].lower())) + + lines = [] + for typ, name, size in entries: + sz = f"({_fmt_size(size)})" if size else "" + lines.append(f"{typ} {name:<40s} {sz}") + return "\n".join(lines) or "(empty)" + +@mcp.tool() +def directory_tree(path: str, excludePatterns: Optional[list[str]] = None) -> str: + """递归树形目录结构,输出 JSON""" + p = _resolve(path) + if not p.is_dir(): + return json.dumps({"ok": False, "error": f"Not a directory: {path}"}) + + def _build_tree(dirpath: Path, depth: int = 0) -> list: + if depth > 8: + return [] + result = [] + try: + entries = sorted(dirpath.iterdir(), key=lambda e: (not e.is_dir(), e.name.lower())) + except PermissionError: + return [] + for entry in entries: + if excludePatterns and any(fnmatch.fnmatch(entry.name, pat) for pat in excludePatterns): + continue + if entry.is_dir(): + children = _build_tree(entry, depth + 1) + result.append({"name": entry.name, "type": "directory", "children": children}) + else: + result.append({"name": entry.name, "type": "file"}) + return result + + tree = _build_tree(p) + return json.dumps({"root": str(p), "children": tree}, ensure_ascii=False, indent=2) + +# ════════════════════════════════════════════════════════════════ +# Search & Move +# ════════════════════════════════════════════════════════════════ + +@mcp.tool() +def search_files(path: str, pattern: str, excludePatterns: Optional[list[str]] = None) -> str: + """递归搜索匹配 glob 模式的文件""" + p = _resolve(path) + results = [] + for root, dirs, files in os.walk(p): + if excludePatterns: + dirs[:] = [d for d in dirs if not any(fnmatch.fnmatch(d, pat) for pat in excludePatterns)] + files = [f for f in files if not any(fnmatch.fnmatch(f, pat) for pat in excludePatterns)] + for f in files: + if fnmatch.fnmatch(f, pattern): + results.append(str(Path(root) / f)) + return "\n".join(results[:500]) or "No matches" + +@mcp.tool() +def move_file(source: str, destination: str) -> str: + """移动或重命名文件/目录""" + src = _resolve(source) + if not src.exists(): + return json.dumps({"ok": False, "error": f"Source not found: {source}"}) + dst = _resolve(destination) + dst.parent.mkdir(parents=True, exist_ok=True) + shutil.move(str(src), str(dst)) + return json.dumps({"ok": True, "from": str(src), "to": str(dst)}) + +# ════════════════════════════════════════════════════════════════ +# Helpers +# ════════════════════════════════════════════════════════════════ + +def _fmt_size(n: int) -> str: + for unit in ("B", "KB", "MB", "GB"): + if n < 1024: + return f"{n:.0f}{unit}" + n /= 1024 + return f"{n:.0f}TB" + +# ════════════════════════════════════════════════════════════════ +# Extended Tools +# ════════════════════════════════════════════════════════════════ + +@mcp.tool() +def grep_files( + path: str, + pattern: str, + glob: str = "*", + max_results: int = 50, + context_lines: int = 0, + ignore_case: bool = False, +) -> str: + """内容搜索 — 递归搜索文件内容匹配正则/文本。替代 bash rg/grep。 + + Args: + path: 搜索根目录 + pattern: 正则表达式或纯文本 + glob: 文件名过滤 (默认 "*") + max_results: 最大结果数 (默认 50) + context_lines: 每个匹配的上下文行数 (0-3) + ignore_case: 忽略大小写 + """ + p = _resolve(path) + if not p.is_dir(): + return json.dumps({"ok": False, "error": f"Not a directory: {path}"}) + + flags = re.IGNORECASE if ignore_case else 0 + try: + regex = re.compile(pattern, flags) + is_regex = True + except re.error: + regex = re.compile(re.escape(pattern), flags) + is_regex = False + + results = [] + for root, dirs, files in os.walk(p): + dirs[:] = [d for d in dirs if not d.startswith(".") and d not in ("node_modules", "__pycache__", ".git")] + for fname in files: + if not fnmatch.fnmatch(fname, glob): + continue + fpath = Path(root) / fname + if fpath.stat().st_size > 5 * 1024 * 1024: + continue + try: + lines = fpath.read_text(encoding="utf-8", errors="replace").splitlines() + except Exception: + continue + for i, line in enumerate(lines): + if regex.search(line): + entry = {"file": str(fpath.relative_to(p)), "line": i + 1, "text": line.strip()[:200]} + if context_lines > 0: + entry["context"] = [ + lines[j].strip()[:200] + for j in range(max(0, i - context_lines), min(len(lines), i + context_lines + 1)) + if j != i + ] + results.append(entry) + if len(results) >= max_results: + summary = { + "ok": True, "truncated": len(results) >= max_results, + "total_found": len(results), "mode": "regex" if is_regex else "plain", + "pattern": pattern, "results": results + } + return json.dumps(summary, ensure_ascii=False, indent=2) + + summary = { + "ok": True, "truncated": False, + "total_found": len(results), "mode": "regex" if is_regex else "plain", + "pattern": pattern, "results": results + } + return json.dumps(summary, ensure_ascii=False, indent=2) + + +@mcp.tool() +def diff_files(path1: str, path2: str) -> str: + """比较两个文件,返回 unified diff""" + p1 = _resolve(path1) + p2 = _resolve(path2) + if not p1.is_file(): + return json.dumps({"ok": False, "error": f"Not found: {path1}"}) + if not p2.is_file(): + return json.dumps({"ok": False, "error": f"Not found: {path2}"}) + + lines1 = p1.read_text(encoding="utf-8", errors="replace").splitlines(True) + lines2 = p2.read_text(encoding="utf-8", errors="replace").splitlines(True) + diff = difflib.unified_diff(lines1, lines2, fromfile=str(p1), tofile=str(p2), lineterm="") + return "\n".join(diff) or "(files are identical)" + + +@mcp.tool() +def file_hash(path: str, algorithm: str = "sha256") -> str: + """计算文件哈希 (md5/sha1/sha256)""" + p = _resolve(path) + if not p.is_file(): + return json.dumps({"ok": False, "error": f"Not found: {path}"}) + h = hashlib.new(algorithm) + with open(p, "rb") as f: + for chunk in iter(lambda: f.read(8192), b""): + h.update(chunk) + return json.dumps({"path": str(p), "algorithm": algorithm, "hash": h.hexdigest(), "size": p.stat().st_size}) + + +@mcp.tool() +def delete_file(path: str, recursive: bool = False) -> str: + """删除文件或目录 (recursive=True 时递归删除目录)""" + p = _resolve(path) + if not p.exists(): + return json.dumps({"ok": False, "error": f"Not found: {path}"}) + if p.is_dir(): + if recursive: + shutil.rmtree(p) + else: + p.rmdir() + else: + p.unlink() + return json.dumps({"ok": True, "deleted": str(p)}) + + +@mcp.tool() +def append_file(path: str, content: str) -> str: + """追加内容到文件末尾(自动换行)""" + p = _resolve(path) + p.parent.mkdir(parents=True, exist_ok=True) + with open(p, "a", encoding="utf-8") as f: + f.write(content + "\n") + return json.dumps({"ok": True, "path": str(p), "size": p.stat().st_size, "appended": True}) + + +@mcp.tool() +def copy_file(source: str, destination: str) -> str: + """复制文件或目录(目录递归复制)""" + src = _resolve(source) + if not src.exists(): + return json.dumps({"ok": False, "error": f"Source not found: {source}"}) + dst = _resolve(destination) + dst.parent.mkdir(parents=True, exist_ok=True) + if src.is_dir(): + shutil.copytree(src, dst) + else: + shutil.copy2(src, dst) + return json.dumps({"ok": True, "from": str(src), "to": str(dst)}) + + +@mcp.tool() +def batch_replace( + path: str, + old: str, + new: str, + glob: str = "*", + dry_run: bool = False, +) -> str: + """批量替换 — 在多个文件中搜索并替换文本(类似 sed -i)。 + + Args: + path: 搜索根目录 + old: 要替换的文本 + new: 替换后的文本 + glob: 文件名过滤 (默认 "*") + dry_run: True 时只预览,不实际修改 + """ + p = _resolve(path) + if not p.is_dir(): + return json.dumps({"ok": False, "error": f"Not a directory: {path}"}) + + modified = [] + for root, dirs, files in os.walk(p): + dirs[:] = [d for d in dirs if not d.startswith(".") and d not in ("node_modules", "__pycache__", ".git")] + for fname in files: + if not fnmatch.fnmatch(fname, glob): + continue + fpath = Path(root) / fname + if fpath.stat().st_size > 5 * 1024 * 1024: + continue + try: + text = fpath.read_text(encoding="utf-8", errors="replace") + except Exception: + continue + if old in text: + count = text.count(old) + modified.append({"file": str(fpath.relative_to(p)), "occurrences": count}) + if not dry_run: + fpath.write_text(text.replace(old, new), encoding="utf-8") + + return json.dumps({ + "ok": True, + "dry_run": dry_run, + "files_modified": len(modified), + "details": modified[:100], + }, ensure_ascii=False, indent=2) + + +@mcp.tool() +def count_lines(path: str, glob: str = "*") -> str: + """统计目录中文件的行数/字数/字符数(wc 替代)""" + p = _resolve(path) + if p.is_file(): + text = p.read_text(encoding="utf-8", errors="replace") + return json.dumps({ + "file": str(p), + "lines": text.count("\n") + 1, + "words": len(text.split()), + "chars": len(text), + }) + if p.is_dir(): + total_lines, total_words, total_chars, file_count = 0, 0, 0, 0 + for root, dirs, files in os.walk(p): + dirs[:] = [d for d in dirs if not d.startswith(".")] + for fname in files: + if not fnmatch.fnmatch(fname, glob): + continue + fpath = Path(root) / fname + if fpath.stat().st_size > 10 * 1024 * 1024: + continue + try: + text = fpath.read_text(encoding="utf-8", errors="replace") + except Exception: + continue + total_lines += text.count("\n") + 1 + total_words += len(text.split()) + total_chars += len(text) + file_count += 1 + return json.dumps({ + "directory": str(p), + "files": file_count, + "total_lines": total_lines, + "total_words": total_words, + "total_chars": total_chars, + }) + return json.dumps({"ok": False, "error": f"Not found: {path}"}) + + +@mcp.tool() +def find_by_size( + path: str, + min_kb: int = 0, + max_kb: int = 0, + top_n: int = 20, +) -> str: + """按文件大小查找 — 找出目录中最大/最小/区间内的文件""" + p = _resolve(path) + if not p.is_dir(): + return json.dumps({"ok": False, "error": f"Not a directory: {path}"}) + + files = [] + for root, dirs, filenames in os.walk(p): + dirs[:] = [d for d in dirs if not d.startswith(".")] + for fname in filenames: + fpath = Path(root) / fname + sz = fpath.stat().st_size + if min_kb and sz < min_kb * 1024: + continue + if max_kb and sz > max_kb * 1024: + continue + files.append((str(fpath.relative_to(p)), sz)) + + files.sort(key=lambda x: -x[1]) + top = files[:top_n] + return json.dumps({ + "ok": True, + "total_matches": len(files), + "filters": {"min_kb": min_kb, "max_kb": max_kb}, + "top": [{"file": f, "size": _fmt_size(sz)} for f, sz in top], + }, ensure_ascii=False, indent=2) + + +@mcp.tool() +def find_by_date( + path: str, + hours: int = 24, + glob: str = "*", + top_n: int = 30, +) -> str: + """按修改时间查找 — 找出最近 N 小时内修改的文件""" + p = _resolve(path) + if not p.is_dir(): + return json.dumps({"ok": False, "error": f"Not a directory: {path}"}) + + cutoff = datetime.now().timestamp() - hours * 3600 + results = [] + for root, dirs, filenames in os.walk(p): + dirs[:] = [d for d in dirs if not d.startswith(".") and d not in ("node_modules", "__pycache__", ".git")] + for fname in filenames: + if not fnmatch.fnmatch(fname, glob): + continue + fpath = Path(root) / fname + mtime = fpath.stat().st_mtime + if mtime >= cutoff: + results.append((str(fpath.relative_to(p)), mtime, fpath.stat().st_size)) + + results.sort(key=lambda x: -x[1]) + top = results[:top_n] + return json.dumps({ + "ok": True, + "hours": hours, + "total_matches": len(results), + "files": [ + {"file": f, "modified": datetime.fromtimestamp(ts).isoformat(), "size": _fmt_size(sz)} + for f, ts, sz in top + ], + }, ensure_ascii=False, indent=2) + + +# ════════════════════════════════════════════════════════════════ +# 🚀 Shell Execution (v2.0 — 最强增强) +# ════════════════════════════════════════════════════════════════ + +@mcp.tool() +def shell_run( + command: str, + cwd: str = "", + timeout: int = 60, + env: dict = None, + shell: bool = True, + capture_both: bool = True, + background: bool = False, +) -> str: + """🚀 在沙箱白名单目录内执行 Shell 命令。 + + 命令白名单: git, python, pip, npm, node, go, cargo, ls, grep, find, curl, wget 等 60+ 工具。 + 高危命令 (rm -rf /, shutdown, fork bomb 等) 自动拦截。 + + Args: + command: 要执行的命令 (支持管道、重定向等) + cwd: 工作目录 (默认为 F:\\DEEPCODE,必须在白名单目录内) + timeout: 超时秒数 (默认 60s,最大 300s) + env: 额外环境变量 dict (可选) + shell: 通过 shell 执行 (默认 True,支持管道) + capture_both: 同时捕获 stdout + stderr (默认 True) + + Returns: + JSON: {ok, exit_code, stdout, stderr, elapsed_ms, command, cwd, timed_out, was_blocked} + + Examples: + shell_run("git status") + shell_run("python -m http.server 8080", background=True) + shell_run("find . -name '*.py' | head -20") + """ + # ── CWD 验证 ── + if not cwd: + cwd = str(ALLOWED_ROOTS[0]) # 默认 F:\DEEPCODE + if not _is_allowed(cwd): + return json.dumps({ + "ok": False, "error": f"cwd not allowed: {cwd}", + "allowed_roots": [str(r) for r in ALLOWED_ROOTS] + }) + + # ── 命令白名单检查 ── + allowed, reason = _check_command(command) + if not allowed: + return json.dumps({"ok": False, "error": reason, "was_blocked": True}) + + # ── 超时限制 ── + if timeout > 300 and not background: + timeout = 300 + + # ── 后台模式 ── + if background: + proc_id = uuid.uuid4().hex[:8] + log_file = Path(cwd) / f".shell_bg_{proc_id}.log" + log_f = open(log_file, "w", encoding="utf-8") + + kwargs = {"cwd": cwd, "stdout": log_f, "stderr": log_f} + if env: + merged_env = os.environ.copy() + merged_env.update(env) + kwargs["env"] = merged_env + if shell: + kwargs["shell"] = True + + proc = subprocess.Popen(command, **kwargs) + + _background_procs[proc_id] = { + "proc": proc, + "command": command[:500], + "cwd": cwd, + "started_at": datetime.now().isoformat(), + "log_file": str(log_file), + } + + return json.dumps({ + "ok": True, + "proc_id": proc_id, + "command": command[:500], + "cwd": cwd, + "log_file": str(log_file), + "background": True, + "note": f"Use shell_background_output('{proc_id}') to read output, shell_background_kill('{proc_id}') to stop." + }, ensure_ascii=False) + + # ── 前台执行 ── + start = time.time() + try: + kwargs = { + "cwd": cwd, + "timeout": timeout, + } + if capture_both: + kwargs["stdout"] = subprocess.PIPE + kwargs["stderr"] = subprocess.PIPE + else: + kwargs["stdout"] = subprocess.PIPE + kwargs["stderr"] = subprocess.STDOUT + if env: + merged_env = os.environ.copy() + merged_env.update(env) + kwargs["env"] = merged_env + if shell: + kwargs["shell"] = True + + proc = subprocess.run(command, **kwargs) + elapsed = int((time.time() - start) * 1000) + + stdout = proc.stdout.decode("utf-8", errors="replace") if proc.stdout else "" + stderr = proc.stderr.decode("utf-8", errors="replace") if (capture_both and proc.stderr) else "" + + # 截断过长输出 + max_output = 50000 + if len(stdout) > max_output: + stdout = stdout[:max_output] + f"\n... [truncated at {max_output} chars, total {len(stdout)}]" + if len(stderr) > max_output: + stderr = stderr[:max_output] + f"\n... [truncated at {max_output} chars]" + + return json.dumps({ + "ok": True, + "exit_code": proc.returncode, + "stdout": stdout, + "stderr": stderr, + "elapsed_ms": elapsed, + "command": command[:500], + "cwd": cwd, + "timed_out": False, + "was_blocked": False, + }, ensure_ascii=False) + + except subprocess.TimeoutExpired: + elapsed = int((time.time() - start) * 1000) + return json.dumps({ + "ok": False, + "exit_code": -1, + "stdout": "", + "stderr": f"Command timed out after {timeout}s", + "elapsed_ms": elapsed, + "command": command[:500], + "cwd": cwd, + "timed_out": True, + "was_blocked": False, + }) + except Exception as e: + elapsed = int((time.time() - start) * 1000) + return json.dumps({ + "ok": False, + "exit_code": -1, + "stdout": "", + "stderr": str(e), + "elapsed_ms": elapsed, + "command": command[:500], + "cwd": cwd, + "timed_out": False, + "was_blocked": False, + }) + + +@mcp.tool() +def shell_list_commands() -> str: + """列出 shell_run 支持的所有命令白名单(60+ 工具)""" + return json.dumps({ + "allowed_commands": sorted(ALLOWED_COMMANDS), + "blocked_patterns": BLOCKED_PATTERNS, + "description": "Commands in the whitelist can be executed via shell_run(). Blocked patterns are always rejected regardless of whitelist." + }, indent=2) + + +@mcp.tool() +def shell_check_command(command: str) -> str: + """检查命令是否在白名单中(不实际执行)。用于预检。""" + allowed, reason = _check_command(command) + return json.dumps({"command": command[:200], "allowed": allowed, "reason": reason}) + + +@mcp.tool() +def shell_background_list() -> str: + """列出所有后台运行的进程""" + procs = [] + for pid, info in _background_procs.items(): + p = info["proc"] + running = p.poll() is None + procs.append({ + "proc_id": pid, + "command": info["command"][:120], + "cwd": info["cwd"], + "started_at": info["started_at"], + "log_file": info["log_file"], + "running": running, + "exit_code": p.returncode if not running else None, + }) + return json.dumps({"ok": True, "count": len(procs), "processes": procs}, ensure_ascii=False, indent=2) + + +@mcp.tool() +def shell_background_output(proc_id: str, tail: int = 50) -> str: + """读取后台进程的日志输出 + + Args: + proc_id: shell_run(background=True) 返回的进程 ID + tail: 只返回最后 N 行 (默认 50,0 表示全部) + """ + if proc_id not in _background_procs: + return json.dumps({"ok": False, "error": f"Process not found: {proc_id}"}) + + info = _background_procs[proc_id] + log_file = info["log_file"] + running = info["proc"].poll() is None + + if not Path(log_file).exists(): + return json.dumps({"ok": True, "proc_id": proc_id, "running": running, "output": "(no output yet)", "log_file": log_file}) + + try: + content = Path(log_file).read_text(encoding="utf-8", errors="replace") + except Exception: + return json.dumps({"ok": True, "proc_id": proc_id, "running": running, "output": "(cannot read log)", "log_file": log_file}) + + lines = content.splitlines() + if tail > 0 and len(lines) > tail: + content = "\n".join(lines[-tail:]) + f"\n... [{len(lines) - tail} earlier lines omitted]" + + max_chars = 30000 + if len(content) > max_chars: + content = content[:max_chars] + f"\n... [truncated at {max_chars} chars]" + + return json.dumps({ + "ok": True, + "proc_id": proc_id, + "running": running, + "exit_code": info["proc"].returncode if not running else None, + "output": content, + "log_file": log_file, + }, ensure_ascii=False) + + +@mcp.tool() +def shell_background_kill(proc_id: str) -> str: + """终止后台进程 + + Args: + proc_id: shell_run(background=True) 返回的进程 ID + """ + if proc_id not in _background_procs: + return json.dumps({"ok": False, "error": f"Process not found: {proc_id}"}) + + info = _background_procs[proc_id] + proc = info["proc"] + running = proc.poll() is None + + if not running: + exit_code = proc.returncode + del _background_procs[proc_id] + return json.dumps({"ok": True, "proc_id": proc_id, "was_running": False, "exit_code": exit_code, "killed": False}) + + # 优雅终止 → 强制终止 + proc.terminate() + try: + proc.wait(timeout=5) + except subprocess.TimeoutExpired: + proc.kill() + proc.wait() + + exit_code = proc.returncode + del _background_procs[proc_id] + return json.dumps({"ok": True, "proc_id": proc_id, "was_running": True, "exit_code": exit_code, "killed": True}) + + +# ════════════════════════════════════════════════════════════════ +# Entry +# ════════════════════════════════════════════════════════════════ + +if __name__ == "__main__": + mcp.run() diff --git a/core/parallel_executor.py b/core/parallel_executor.py new file mode 100644 index 00000000..125743ab --- /dev/null +++ b/core/parallel_executor.py @@ -0,0 +1,443 @@ +#!/usr/bin/env python3 +""" +ParallelToolExecutor v2.0 — DAG 并行工具执行引擎 + +升级: 模仿 GGML 的 graph_plan + graph_compute + 线程池模式 + +关键改进 (基于 llama.cpp 逆向发现): + 1. DAG 任务图 (+ 拓扑排序) → 对应 GGML graph_plan + 2. 无锁就绪队列 + 线程池调度 → 对应 GGML graph_compute + 3. 单线程 fast path → 对应 GGML 单线程路径 + 4. 二维调度矩阵 (操作 × 数据类型) → 对应 GGML OpCode × 数据类型矩阵 + 5. 依赖感知的并行分批 → 对应 GGML 图分割 +""" + +import asyncio +import time +import json +from typing import Any, Callable, Dict, List, Optional, AsyncGenerator +from dataclasses import dataclass, field +from enum import Enum +from collections import deque + + +class TaskStatus(Enum): + PENDING = "pending" + RUNNING = "running" + DONE = "done" + ERROR = "error" + SKIPPED = "skipped" + + +class TaskPriority(Enum): + LOW = 0 + NORMAL = 1 + HIGH = 2 + CRITICAL = 3 + + +# ── 二维调度矩阵 (模仿 GGML 的 OpCode × 数据类型) ──────────── + +class MatrixToolRegistry: + """ + 二维调度矩阵: (操作类型 × 语言/数据类型) → 处理函数 + + 注册方式: + registry = MatrixToolRegistry() + registry.register("check_sql", "python", python_sql_checker) + registry.register("check_sql", "java", java_sql_checker) + + 调度方式: + handler = registry.lookup("check_sql", "python") # O(1) 查表 + """ + + def __init__(self): + # 二维矩阵: matrix[op][variant] = handler + self._matrix: dict[str, dict[str, Callable]] = {} + # 默认处理器: matrix[op]["*"] = default_handler + self._defaults: dict[str, Callable] = {} + + def register(self, operation: str, variant: str, handler: Callable): + """注册处理器到矩阵的 (operation, variant) 位置""" + if operation not in self._matrix: + self._matrix[operation] = {} + self._matrix[operation][variant] = handler + + def register_default(self, operation: str, handler: Callable): + """注册默认处理器 (当 variant 无匹配时使用)""" + self._defaults[operation] = handler + + def lookup(self, operation: str, variant: str) -> Callable | None: + """O(1) 查表: 精确匹配 → 默认匹配 → None""" + ops = self._matrix.get(operation, {}) + if variant in ops: + return ops[variant] + if "*" in ops: + return ops["*"] + return self._defaults.get(operation) + + @property + def operations(self) -> list[str]: + return list(self._matrix.keys()) + + def stats(self) -> dict: + return { + "operations": len(self._matrix), + "total_entries": sum(len(v) for v in self._matrix.values()), + } + + +# ── DAG 任务模型 ───────────────────────────────────────────── + +@dataclass +class DAGNode: + """计算图节点 — 对应 GGML 计算图中的一个操作节点""" + id: str + tool_name: str + args: Dict[str, Any] = field(default_factory=dict) + variant: str = "*" # 数据类型/语言变体 + priority: TaskPriority = TaskPriority.NORMAL + depends_on: List[str] = field(default_factory=list) # 依赖的节点 ID 列表 + status: TaskStatus = TaskStatus.PENDING + result: Any = None + error: Optional[str] = None + start_time: float = 0 + end_time: float = 0 + retries: int = 0 + max_retries: int = 3 + + @property + def duration(self) -> float: + return self.end_time - self.start_time if self.end_time > 0 else 0 + + +class DAGExecutor: + """ + DAG 任务执行引擎 — 对应 GGML 的 graph_plan + graph_compute + + 模式匹配: + GGML graph_plan → self._build_execution_plan() + GGML graph_compute → self._execute_plan() + GGML 单线程 fast path → self._execute_fast_path() + GGML 线程池 → self._execute_parallel() + """ + + def __init__(self, registry: MatrixToolRegistry | None = None, + max_concurrency: int = 4): + self.registry = registry or MatrixToolRegistry() + self.max_concurrency = max_concurrency + self.stats = { + "total_executions": 0, + "dag_executions": 0, + "fast_path_executions": 0, + "total_nodes": 0, + "total_time_saved_ms": 0, + } + + # ── 构建 DAG 执行计划 (对应 GGML graph_plan) ────────────── + + def _build_execution_plan(self, nodes: List[DAGNode]) -> List[List[DAGNode]]: + """ + 拓扑排序 → 并行分批 + + 返回: [[batch_1], [batch_2], ...] + 同批次内无依赖冲突,可并行执行 + """ + node_map = {n.id: n for n in nodes} + in_degree = {n.id: 0 for n in nodes} + for n in nodes: + for d in n.depends_on: + if d in in_degree: + in_degree[d] += 1 + + # 优先队列: 同批次内按优先级排序 + ready = deque( + sorted( + [n for n in nodes if in_degree[n.id] == 0], + key=lambda n: n.priority.value, + reverse=True, + ) + ) + remaining = {n.id for n in nodes} + plan = [] + + while ready or remaining: + batch = [] + next_ready = deque() + + while ready: + n = ready.popleft() + batch.append(n) + remaining.discard(n.id) + + # 更新入度 + for n in batch: + for d_id in remaining: + d = node_map[d_id] + if n.id in d.depends_on: + in_degree[d_id] -= 1 + if in_degree[d_id] == 0: + next_ready.append(d) + + if batch: + plan.append(batch) + + # 死锁检测: 如果还有剩余节点但没有就绪的,强制取一个 + if not next_ready and remaining: + forced = next(iter(remaining)) + next_ready.append(node_map[forced]) + + ready = next_ready + + return plan + + # ── 快速路径 (单线程, 对应 GGML 单线程模式) ────────────── + + async def _execute_fast_path(self, nodes: List[DAGNode]) -> List[DAGNode]: + """单线程按序执行 — 无锁开销""" + for n in nodes: + n.start_time = time.time() + try: + n.result = await self._run_node(n) + n.status = TaskStatus.DONE + except Exception as e: + n.status = TaskStatus.ERROR + n.error = str(e) + n.end_time = time.time() + return nodes + + # ── 并行路径 (多线程, 对应 GGML 多线程模式) ────────────── + + async def _execute_plan(self, plan: List[List[DAGNode]]) -> List[DAGNode]: + """逐批执行,同批次并行 (对应 GGML 线程池 + 任务队列)""" + sem = asyncio.Semaphore(self.max_concurrency) + done_nodes = [] + + for batch in plan: + async def run_node(n: DAGNode) -> DAGNode: + async with sem: + n.start_time = time.time() + try: + n.result = await self._run_node(n) + n.status = TaskStatus.DONE + except Exception as e: + n.status = TaskStatus.ERROR + n.error = str(e) + n.end_time = time.time() + return n + + # 同批次并行 + results = await asyncio.gather( + *[run_node(n) for n in batch], + return_exceptions=True, + ) + for r in results: + if isinstance(r, DAGNode): + done_nodes.append(r) + + return done_nodes + + # ── 核心执行方法 ───────────────────────────────────────── + + async def execute(self, nodes: List[DAGNode]) -> List[DAGNode]: + """ + 执行 DAG — 自动选择路径: + - 1 个节点 → 直接执行 + - 1 个批次 → 并行执行 + - 多批次 → 逐批并行 + """ + self.stats["total_executions"] += 1 + + if len(nodes) == 1: + # 快速路径: 单节点直接执行 (对应 GGML 单线程) + self.stats["fast_path_executions"] += 1 + n = nodes[0] + n.start_time = time.time() + try: + n.result = await self._run_node(n) + n.status = TaskStatus.DONE + except Exception as e: + n.status = TaskStatus.ERROR + n.error = str(e) + n.end_time = time.time() + return [n] + + # 构建执行计划 + self.stats["dag_executions"] += 1 + plan = self._build_execution_plan(nodes) + done = await self._execute_plan(plan) + self.stats["total_nodes"] += len(nodes) + return done + + async def execute_from_dicts(self, tasks: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """字典列表接口 — 兼容 v1.0 接口""" + nodes = [] + for i, t in enumerate(tasks): + nodes.append(DAGNode( + id=t.get("id", str(i + 1)), + tool_name=t["tool"], + args=t.get("args", {}), + variant=t.get("variant", "*"), + priority=TaskPriority[t.get("priority", "NORMAL").upper()], + depends_on=t.get("depends_on", []), + )) + results = await self.execute(nodes) + return [self._format_node(n) for n in results] + + async def execute_streaming( + self, nodes: List[DAGNode] + ) -> AsyncGenerator[Dict, None]: + """流式执行 — 每完成一个节点就 yield""" + plan = self._build_execution_plan(nodes) + + for batch in plan: + pending = [self._run_and_format(n) for n in batch] + for coro in asyncio.as_completed(pending): + yield await coro + + async def _run_node(self, node: DAGNode) -> Any: + """执行单个节点 — 位置透明调度 + + 优先级: + 1. 2D 矩阵查找 (MatrixToolRegistry.lookup) + 2. TransportAwareRegistry.execute() (位置透明) + 3. HandleFactory 全局兜底 (InProcess → Agent → MCP → Remote) + 4. 兜底 + """ + # 1. 尝试 2D 矩阵查找 (MatrixToolRegistry) + if hasattr(self.registry, 'lookup'): + handler = self.registry.lookup(node.tool_name, node.variant) + if handler is not None: + if asyncio.iscoroutinefunction(handler): + return await handler(**node.args) + else: + return handler(**node.args) + + # 2. TransportAwareRegistry.execute() (适用于任何 ToolRegistry) + if hasattr(self.registry, 'execute'): + try: + call_result = await self.registry.execute(node.tool_name, node.args) + if call_result.success: + return call_result.result + node.error = call_result.error + except Exception as e: + node.error = str(e) + + # 3. HandleFactory 全局兜底 + if not node.error: + try: + from transparent_handle import call as transparent_call + result = await transparent_call(node.tool_name, node.args) + if result.success: + return result.data + node.error = result.error + except ImportError: + pass + except Exception as e: + node.error = str(e) + + # 3. 兜底: 未注册的工具 + return {"note": f"Tool '{node.tool_name}/{node.variant}' not registered", + "args": node.args} + + async def _run_and_format(self, node: DAGNode) -> Dict: + try: + node.start_time = time.time() + node.result = await self._run_node(node) + node.status = TaskStatus.DONE + except Exception as e: + node.status = TaskStatus.ERROR + node.error = str(e) + node.end_time = time.time() + return self._format_node(node) + + def _format_node(self, node: DAGNode) -> Dict: + return { + "task_id": node.id, + "tool": node.tool_name, + "variant": node.variant, + "status": node.status.value, + "result": node.result, + "error": node.error, + "duration_ms": round(node.duration * 1000, 1), + } + + def get_stats(self) -> Dict: + return {**self.stats, "registry": self.registry.stats()} + + +# ── 便捷全局接口 ────────────────────────────────────────────── + +_executor: Optional[DAGExecutor] = None +_registry: Optional[MatrixToolRegistry] = None + + +def get_registry() -> MatrixToolRegistry: + global _registry + if _registry is None: + _registry = MatrixToolRegistry() + return _registry + + +def get_executor(**kwargs) -> DAGExecutor: + global _executor + if _executor is None: + reg = get_registry() + _executor = DAGExecutor(registry=reg, **kwargs) + return _executor + + +async def run_parallel(tasks: List[Dict]) -> List[Dict]: + """兼容 v1.0 接口""" + executor = get_executor() + return await executor.execute_from_dicts(tasks) + + +# ── 测试 ────────────────────────────────────────────────────── + +async def _test(): + """测试 DAG 并行执行""" + async def mock_fetch(ts_code=None, **kw): + await asyncio.sleep(0.3) + return {"ts_code": ts_code, "close": 100} + + async def mock_analyze(data=None, **kw): + await asyncio.sleep(0.2) + return {"analysis": "done", "data": data} + + # 注册 2D 矩阵 + reg = get_registry() + reg.register("fetch_data", "stock", mock_fetch) + reg.register("fetch_data", "index", mock_fetch) + reg.register("analyze", "*", mock_analyze) + + executor = DAGExecutor(registry=reg, max_concurrency=4) + + tasks = [ + {"id": "1", "tool": "fetch_data", "variant": "stock", + "args": {"ts_code": "600519.SH"}}, + {"id": "2", "tool": "fetch_data", "variant": "stock", + "args": {"ts_code": "000001.SZ"}}, + {"id": "3", "tool": "fetch_data", "variant": "index", + "args": {"ts_code": "000001.SH"}}, + {"id": "4", "tool": "analyze", "variant": "stock", + "args": {"data": "result_1"}, "depends_on": ["1", "2"]}, + {"id": "5", "tool": "analyze", "variant": "index", + "args": {"data": "result_3"}, "depends_on": ["3"]}, + ] + + start = time.time() + results = await executor.execute_from_dicts(tasks) + elapsed = time.time() - start + + print(f"=== DAG Executor Test ===") + print(f"5 tasks (2 parallel + 3 dependent) in {elapsed:.2f}s") + for r in results: + dep = " ⚡dep" if r["task_id"] in ("4", "5") else "" + print(f" [{r['task_id']}] {r['tool']}/{r['variant']}: {r['status']} ({r['duration_ms']}ms){dep}") + print(f"Stats: {json.dumps(executor.get_stats(), indent=2)}") + print(f"(vs sequential ~1.5s, saved ~{1.5 - elapsed:.1f}s)") + + +if __name__ == "__main__": + asyncio.run(_test()) diff --git a/deepcode-engine-mcp b/deepcode-engine-mcp new file mode 160000 index 00000000..6823ef09 --- /dev/null +++ b/deepcode-engine-mcp @@ -0,0 +1 @@ +Subproject commit 6823ef09acbae1da5ec2ea7ae37dd3fead6670b2 diff --git a/scripts/validate_mcp.py b/scripts/validate_mcp.py new file mode 100644 index 00000000..5c985e05 --- /dev/null +++ b/scripts/validate_mcp.py @@ -0,0 +1,187 @@ +#!/usr/bin/env python3 +""" +MCP 配置健康检查脚本 — 防止路径断裂和能力声明错误。 + +检查项: + 1. 入口文件是否存在 (command + args 中的所有文件路径) + 2. Python MCP 服务器的 capabilities.tools 是否声明 listChanged + 3. 必需的 environment variables 是否可解析 + +用法: + python scripts/validate_mcp.py # 检查所有 + python scripts/validate_mcp.py --json # JSON 输出 (适合 CI) + python scripts/validate_mcp.py --quiet # 仅输出错误 +""" + +import json +import os +import re +import shutil +import sys +from pathlib import Path + +PROJECT_ROOT = Path(__file__).resolve().parent.parent +SETTINGS_PATH = PROJECT_ROOT / ".deepcode" / "settings.json" + +# — 这些 command 是包管理器/运行时,不需要检查路径 — +_RUNTIME_COMMANDS = {"npx", "node", "python", "python3", "uv", "uvx", "npm", "yarn", "pnpm"} + +# — capabilities 声明正则:匹配有问题的模式 — +_BAD_CAPABILITIES_RE = re.compile( + r'"capabilities"\s*:\s*\{\s*"tools"\s*:\s*\{\s*\}\s*\}' +) + +# — 文件路径判定:args 中以 .py / .js / .mjs / .ts 结尾的视为需要检查 — +_FILE_EXTENSIONS = {".py", ".js", ".mjs", ".ts", ".cjs"} + + +def is_file_path(arg: str) -> bool: + """判断一个 arg 是否为需要检查的本地文件路径。""" + if not arg: + return False + # 排除明显的 npm 包名 / 命令行 flag + if arg.startswith("-") or arg.startswith("@"): + return False + # 排除纯命令名 (不含路径分隔符) + if "/" not in arg and "\\" not in arg: + return False + # 检查是否以已知文件扩展名结尾 + return any(arg.endswith(ext) for ext in _FILE_EXTENSIONS) + + +def check_python_capabilities(file_path: Path) -> list[str]: + """检查 Python MCP 服务器是否在 capabilities 中声明了 listChanged。""" + issues = [] + try: + content = file_path.read_text(encoding="utf-8", errors="ignore") + except Exception: + return issues # 文件读取失败由路径检查负责 + + if _BAD_CAPABILITIES_RE.search(content): + issues.append( + f" [WARN] capabilities 声明不完整: \"tools\": {{}} 应改为 \"tools\": {{\"listChanged\": true}}" + ) + return issues + + +def validate() -> dict: + """返回 {ok: bool, issues: [str], stats: {total, ok, failed, skipped}}""" + issues = [] + stats = {"total": 0, "ok": 0, "failed": 0, "skipped": 0} + + if not SETTINGS_PATH.exists(): + issues.append(f"[FATAL] 配置文件不存在: {SETTINGS_PATH}") + return {"ok": False, "issues": issues, "stats": stats} + + try: + config = json.loads(SETTINGS_PATH.read_text(encoding="utf-8")) + except json.JSONDecodeError as e: + issues.append(f"[FATAL] settings.json 解析失败: {e}") + return {"ok": False, "issues": issues, "stats": stats} + + mcp_servers = config.get("mcpServers", {}) + if not mcp_servers: + issues.append("[WARN] 没有配置任何 MCP 服务器") + return {"ok": True, "issues": issues, "stats": stats} + + stats["total"] = len(mcp_servers) + + for name, server in mcp_servers.items(): + # 跳过注释行 + if name.startswith("//"): + stats["skipped"] += 1 + continue + + server_ok = True + command = server.get("command", "") + args = server.get("args", []) + + # —— 检查 1: command 是否可执行 —— + if command in _RUNTIME_COMMANDS: + exe_path = shutil.which(command) + if exe_path is None: + issues.append(f"[{name}] command '{command}' 未安装或不在 PATH 中") + server_ok = False + elif command: + # 非标准运行时:先尝试 PATH 解析(如 mcp-server-fetch 这类 + # 全局安装的命令),解析不到再当作文件路径检查。 + if shutil.which(command) is None: + cmd_path = Path(command) + if not cmd_path.is_absolute(): + # 相对路径相对于项目根 + cmd_path = PROJECT_ROOT / command + if not cmd_path.exists(): + issues.append(f"[{name}] command 路径不存在: {cmd_path}") + server_ok = False + + # —— 检查 2: args 中的文件路径 —— + for arg in args: + if is_file_path(arg): + arg_path = Path(arg) + if not arg_path.is_absolute(): + arg_path = PROJECT_ROOT / arg + if not arg_path.exists(): + issues.append(f"[{name}] 入口文件不存在: {arg_path}") + server_ok = False + continue + # —— 检查 3: Python 文件的 capabilities 声明 —— + if arg_path.suffix == ".py": + cap_issues = check_python_capabilities(arg_path) + for ci in cap_issues: + issues.append(f"[{name}] {arg_path.name}:{ci}") + server_ok = False + + # —— 检查 4: env 变量引用 —— + env = server.get("env", {}) + for key, value in env.items(): + if isinstance(value, str) and "${" in value: + refs = re.findall(r'\$\{([^}]+)\}', value) + for ref in refs: + if ref not in os.environ and ref != value: + issues.append( + f"[{name}] 环境变量 ${ref} 未设置 (env.{key})" + ) + # 不标记为 fatal — 可能是 CI 设置 + + if server_ok: + stats["ok"] += 1 + else: + stats["failed"] += 1 + + return { + "ok": stats["failed"] == 0, + "issues": issues, + "stats": stats, + } + + +def main(): + json_output = "--json" in sys.argv + quiet = "--quiet" in sys.argv + + result = validate() + + if json_output: + print(json.dumps(result, ensure_ascii=False, indent=2)) + sys.exit(0 if result["ok"] else 1) + + s = result["stats"] + print(f"\n{'='*60}") + print(f" MCP 健康检查: {s['total']} 台服务器") + print(f" [OK] 正常: {s['ok']} | [FAIL] 异常: {s['failed']} | [SKIP] 跳过: {s['skipped']}") + print(f"{'='*60}") + + if result["issues"]: + for issue in result["issues"]: + print(issue) + print() + + if result["ok"]: + print("[ALL GOOD] 所有 MCP 服务器配置健康!\n") + else: + print(f"[FIX] 发现 {s['failed']} 个问题,请修复后重试。\n") + sys.exit(1) + + +if __name__ == "__main__": + main() diff --git a/tests/test_exec_sandbox_wiring.py b/tests/test_exec_sandbox_wiring.py index 27764576..aa58963a 100644 --- a/tests/test_exec_sandbox_wiring.py +++ b/tests/test_exec_sandbox_wiring.py @@ -37,7 +37,10 @@ def test_disabled_via_env_returns_bare(monkeypatch, tmp_path): monkeypatch.setenv("DEEPCODE_SANDBOX", "0") w = build_exec_command(command="echo hi", workspace=tmp_path) assert w.backend == "disabled" - assert w.argv == ["/bin/bash", "-c", "echo hi"] + # shell 由 shutil.which("bash") 解析:Windows 下是 git-bash 绝对路径, + # POSIX 下是 /bin/bash —— 只断言文件名与剩余 argv。 + assert Path(w.argv[0]).name.lower() in ("bash", "bash.exe") + assert w.argv[1:] == ["-c", "echo hi"] wa = build_exec_command(argv=["python", "x.py"], workspace=tmp_path) assert wa.backend == "disabled" assert wa.argv == ["python", "x.py"] diff --git a/tests/test_harness_sandbox.py b/tests/test_harness_sandbox.py index 05860747..a2c1f2d2 100644 --- a/tests/test_harness_sandbox.py +++ b/tests/test_harness_sandbox.py @@ -45,8 +45,11 @@ def test_seatbelt_profile_is_deny_default_and_grants_workspace_writes(tmp_path): assert "(deny default)" in profile assert "(allow default)" not in profile assert "(allow file-read*)" in profile # reads still broad + # seatbelt profile 是 LISP 字符串,Windows 路径反斜杠会被转义为 \\, + # 所以断言侧同样转义后再比对。 + escaped = os.path.abspath(str(tmp_path)).replace("\\", "\\\\") assert ( - f'(allow file-write* (subpath "{os.path.abspath(str(tmp_path))}"))' in profile + f'(allow file-write* (subpath "{escaped}"))' in profile ) # Under deny-default, no network allows means network is denied. assert "network-outbound" not in profile diff --git a/tests/test_hooks.py b/tests/test_hooks.py index e207d76b..503ae10f 100644 --- a/tests/test_hooks.py +++ b/tests/test_hooks.py @@ -10,6 +10,7 @@ import json import shutil import sys +import tempfile from pathlib import Path import pytest @@ -42,7 +43,11 @@ def _handler(event, command, *, matcher=None, order=0, timeout=30): ) -def _engine(handlers, cwd="/tmp"): +def _engine(handlers, cwd=None): + # Windows 上 "/tmp" 是相对路径,会解析为当前盘符根目录的 tmp(可能不存在), + # 导致 hook 子进程启动失败。默认用系统临时目录,跨平台安全。 + if cwd is None: + cwd = tempfile.gettempdir() return HooksEngine(handlers, cwd, session_id="sess-1") @@ -300,7 +305,7 @@ def test_stop_block_means_keep_going(): def test_payload_delivered_on_stdin(tmp_path): capture = tmp_path / "payload.json" - eng = _engine([_handler("PreToolUse", f"cat > {capture}", matcher="*")]) + eng = _engine([_handler("PreToolUse", f'cat > "{capture.as_posix()}"', matcher="*")]) asyncio.run(eng.run_pre_tool_use("Bash", {"command": "ls"}, tool_use_id="tu-9")) payload = json.loads(capture.read_text()) assert payload["session_id"] == "sess-1" @@ -567,7 +572,7 @@ def test_session_start_and_prompt_context_injected(): def test_subagent_start_payload_and_plaintext_context(tmp_path): capture = tmp_path / "p.json" - eng = _engine([_handler("SubagentStart", f"cat > {capture}; echo sub-context")]) + eng = _engine([_handler("SubagentStart", f'cat > "{capture.as_posix()}"; echo sub-context')]) res = asyncio.run(eng.run_subagent_start("worker-7", "subagent")) assert res.additional_contexts == ["sub-context"] # plain-text context works p = json.loads(capture.read_text()) @@ -802,7 +807,7 @@ async def ask(name, args): def test_pre_compact_hook_block_skips_and_payload(tmp_path): capture = tmp_path / "p.json" out = json.dumps({"continue": False}) - eng = _engine([_handler("PreCompact", f"cat > {capture}; echo '{out}'")]) + eng = _engine([_handler("PreCompact", f"cat > \"{capture.as_posix()}\"; echo '{out}'")]) res = asyncio.run(eng.run_pre_compact("auto")) assert res.block is True # continue:false → skip compaction p = json.loads(capture.read_text()) @@ -818,7 +823,7 @@ def test_pre_compact_matcher_matches_trigger(): def test_post_compact_hook_fires_with_trigger(tmp_path): capture = tmp_path / "p.json" - eng = _engine([_handler("PostCompact", f"cat > {capture}")]) + eng = _engine([_handler("PostCompact", f'cat > "{capture.as_posix()}"')]) asyncio.run(eng.run_post_compact("auto")) p = json.loads(capture.read_text()) assert p["hook_event_name"] == "PostCompact" and p["trigger"] == "auto" @@ -829,7 +834,7 @@ def test_post_compact_hook_fires_with_trigger(tmp_path): def test_stop_payload_carries_stop_hook_active(tmp_path): capture = tmp_path / "p.json" - eng = _engine([_handler("Stop", f"cat > {capture}")]) + eng = _engine([_handler("Stop", f'cat > "{capture.as_posix()}"')]) asyncio.run(eng.run_stop(stop_hook_active=True)) p = json.loads(capture.read_text()) assert p["hook_event_name"] == "Stop" and p["stop_hook_active"] is True diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py index dca326f3..2899014b 100644 --- a/tests/test_mcp_server.py +++ b/tests/test_mcp_server.py @@ -73,9 +73,9 @@ def test_list_tools_exposes_both(): assert server.name == "deepcode" -def test_deepcode_runs_stores_session_and_returns_id(fake_build): +def test_deepcode_runs_stores_session_and_returns_id(fake_build, tmp_path): content, structured = asyncio.run( - mcp_server._handle_deepcode({"prompt": "build X", "workspace": "/tmp/x"}) + mcp_server._handle_deepcode({"prompt": "build X", "workspace": str(tmp_path)}) ) assert content[0].text == "did: build X" sid = structured["session_id"] @@ -83,7 +83,8 @@ def test_deepcode_runs_stores_session_and_returns_id(fake_build): assert sid in mcp_server._SESSIONS # kept for follow-ups # the workspace reached build_agent_session _session, kwargs = fake_build[0] - assert kwargs["workspace"].endswith("/tmp/x") or kwargs["workspace"] == "/tmp/x" + # 用真实临时目录(POSIX 风格的 /tmp/x 在 Windows 上会解析为 F:\tmp\x) + assert Path(kwargs["workspace"]).resolve() == Path(str(tmp_path)).resolve() def test_reply_continues_same_session(fake_build): diff --git a/tests/test_shell_search_tools.py b/tests/test_shell_search_tools.py index 7b19c90a..4056b16c 100644 --- a/tests/test_shell_search_tools.py +++ b/tests/test_shell_search_tools.py @@ -2,6 +2,7 @@ from __future__ import annotations +import os import sys from pathlib import Path @@ -112,7 +113,8 @@ async def test_glob_matches_recursively(tmp_path): (tmp_path / "note.txt").write_text("z") g = GlobTool(str(tmp_path)) out = await g.execute(pattern="**/*.py") - assert "src/m.py" in out and "top.py" in out and "note.txt" not in out + # 路径分隔符跨平台:Windows 输出 src\m.py,POSIX 输出 src/m.py + assert f"src{os.sep}m.py" in out and "top.py" in out and "note.txt" not in out @pytest.mark.asyncio