From 2b32d11f435c582d201f88fc4197f798e754b4ff Mon Sep 17 00:00:00 2001 From: oliver Date: Wed, 13 May 2026 23:02:52 +0800 Subject: [PATCH] feat(tools): add grep and workspace profile, fix write path - Add grep_tool: ripgrep with Python fallback; avoid --json with -l/-c. - Add workspace_profile_tool: scan, languages, tests, packages, git, CI hints; tests. - write_file: resolve relative paths via workspace root (no data/workspace prefix). - Update path guard and local public tool tests for new write behavior. Co-authored-by: Cursor --- runtime/tools/public/grep_tool.py | 382 ++++++++++++ .../tools/public/workspace_profile_tool.py | 569 ++++++++++++++++++ runtime/tools/public/write_file_tool.py | 20 +- tests/test_local_public_tools.py | 2 +- tests/test_workspace_path_guard.py | 8 +- tests/test_workspace_profile_tool.py | 59 ++ 6 files changed, 1019 insertions(+), 21 deletions(-) create mode 100644 runtime/tools/public/grep_tool.py create mode 100644 runtime/tools/public/workspace_profile_tool.py create mode 100644 tests/test_workspace_profile_tool.py diff --git a/runtime/tools/public/grep_tool.py b/runtime/tools/public/grep_tool.py new file mode 100644 index 00000000..821175d5 --- /dev/null +++ b/runtime/tools/public/grep_tool.py @@ -0,0 +1,382 @@ +from __future__ import annotations + +import json +import re +import shutil +import subprocess +from pathlib import Path +from typing import Any + +from runtime.tools.base import ToolSpec +from runtime.tools.path_guard import resolve_workspace_path + + +def grep_tool() -> ToolSpec: + """Grep / ripgrep — fast file content search with fallback.""" + + # ── engine detection ────────────────────────────────────────── + def _detect_engine() -> str: + return "ripgrep" if shutil.which("rg") else "python" + + # ── ripgrep backend ─────────────────────────────────────────── + def _rg_search( + pattern: str, + root: Path, + *, + case_insensitive: bool = False, + word_regexp: bool = False, + files_with_matches: bool = False, + count_only: bool = False, + context_lines: int = 0, + glob: str | None = None, + max_results: int = 200, + ) -> dict[str, Any]: + # ripgrep forbids --json together with -l / -c (see rg manpage OUTPUT MODES). + use_json = not files_with_matches and not count_only + cmd = ["rg"] + if use_json: + cmd.extend(["--json", "--line-number", "--column"]) + if case_insensitive: + cmd.append("-i") + if word_regexp: + cmd.append("-w") + if files_with_matches: + cmd.append("-l") + if count_only: + cmd.append("-c") + if context_lines > 0: + cmd.extend(["-C", str(context_lines)]) + if glob: + cmd.extend(["-g", glob]) + cmd.extend([pattern, str(root)]) + + try: + proc = subprocess.run( + cmd, capture_output=True, text=True, timeout=30 + ) + except subprocess.TimeoutExpired: + return {"ok": False, "error_code": "timeout", "detail": "rg search timed out after 30s"} + except Exception as exc: + return {"ok": False, "error_code": "rg_failed", "detail": str(exc)} + + if proc.returncode not in (0, 1): + return { + "ok": False, + "error_code": "rg_error", + "detail": (proc.stderr.strip() or f"exit_code={proc.returncode}"), + } + + # ── files-with-matches mode ── + if files_with_matches: + files = [f.strip() for f in proc.stdout.strip().split("\n") if f.strip()] + return { + "ok": True, + "engine": "ripgrep", + "files_with_matches": True, + "files": files, + "total_files": len(files), + } + + # ── count mode ── + if count_only: + counts: dict[str, int] = {} + for line in proc.stdout.strip().split("\n"): + line = line.strip() + if ":" in line: + fp, cnt = line.rsplit(":", 1) + counts[fp.strip()] = int(cnt.strip()) + return { + "ok": True, + "engine": "ripgrep", + "count": True, + "file_counts": counts, + } + + # ── standard match mode (JSON) ── + matches: list[dict[str, Any]] = [] + file_map: dict[str, list[int]] = {} + + for raw_line in proc.stdout.strip().split("\n"): + if not raw_line.strip(): + continue + if len(matches) >= max_results: + break + try: + obj = json.loads(raw_line) + except json.JSONDecodeError: + continue + if obj.get("type") != "match": + continue + data = obj.get("data", {}) + filepath = data.get("path", {}).get("text", "") + line_num = data.get("line_number", 0) + line_text = data.get("lines", {}).get("text", "").rstrip("\n").rstrip("\r") + submatches = data.get("submatches", []) + col = (submatches[0].get("start", 0) + 1) if submatches else 0 + + match_entry: dict[str, Any] = { + "file": filepath, + "line": line_num, + "column": col, + "preview": line_text, + } + matches.append(match_entry) + file_map.setdefault(filepath, []).append(len(matches) - 1) + + # ── attach context lines if requested ── + if context_lines > 0 and matches: + for filepath, indices in file_map.items(): + try: + fp = Path(filepath) + if not fp.is_absolute(): + fp = root / filepath + f_lines = fp.read_text(encoding="utf-8", errors="replace").splitlines() + except Exception: + continue + for idx in indices: + m = matches[idx] + ln = m["line"] - 1 # 0-indexed + start = max(0, ln - context_lines) + end = min(len(f_lines), ln + context_lines + 1) + m["before"] = [f_lines[i] for i in range(start, ln)] + m["after"] = [f_lines[i] for i in range(ln + 1, end)] + + total_files = len(file_map) + return { + "ok": True, + "engine": "ripgrep", + "matches": matches, + "total_matches": len(matches), + "files_with_matches": total_files, + } + + # ── Python fallback backend ─────────────────────────────────── + def _py_search( + pattern: str, + root: Path, + *, + case_insensitive: bool = False, + word_regexp: bool = False, + files_with_matches: bool = False, + count_only: bool = False, + context_lines: int = 0, + glob: str | None = None, + max_results: int = 200, + ) -> dict[str, Any]: + flags = re.IGNORECASE if case_insensitive else 0 + pat_str = rf"\b{pattern}\b" if word_regexp else pattern + try: + matcher = re.compile(pat_str, flags) + except re.error as exc: + return {"ok": False, "error_code": "invalid_regex", "detail": str(exc)} + + try: + if root.is_file(): + file_list = [root] + else: + file_list = sorted(root.rglob(glob or "**/*")) + file_list = [p for p in file_list if p.is_file()] + except Exception as exc: + return {"ok": False, "error_code": "scan_failed", "detail": str(exc)} + + scanned = 0 + matches: list[dict[str, Any]] = [] + file_counts: dict[str, int] = {} + + for fp in file_list: + if len(matches) >= max_results and not (files_with_matches or count_only): + break + scanned += 1 + try: + lines = fp.read_text(encoding="utf-8", errors="replace").splitlines() + except Exception: + continue + + file_hit_count = 0 + for idx, line in enumerate(lines): + m = matcher.search(line) + if not m: + continue + file_hit_count += 1 + try: + rel_path = str(fp.relative_to(root)) + except ValueError: + rel_path = str(fp) + + if files_with_matches: + if file_hit_count == 1: + matches.append({"file": rel_path}) + continue + + if count_only: + file_counts[rel_path] = file_hit_count + continue + + col = m.start() + 1 + match_entry: dict[str, Any] = { + "file": rel_path, + "line": idx + 1, + "column": col, + "preview": line, + } + if context_lines > 0: + start = max(0, idx - context_lines) + end = min(len(lines), idx + context_lines + 1) + match_entry["before"] = [lines[i] for i in range(start, idx)] + match_entry["after"] = [lines[i] for i in range(idx + 1, end)] + matches.append(match_entry) + + if len(matches) >= max_results: + break + + if files_with_matches: + files = sorted(set(m["file"] for m in matches)) + return { + "ok": True, + "engine": "python", + "files_with_matches": True, + "files": files, + "total_files": len(files), + } + + if count_only: + return { + "ok": True, + "engine": "python", + "count": True, + "file_counts": file_counts, + } + + return { + "ok": True, + "engine": "python", + "matches": matches, + "total_matches": len(matches), + "files_with_matches": len(set(m.get("file") for m in matches)), + "scanned_files": scanned, + } + + # ── main handler ────────────────────────────────────────────── + def _handler(args: dict[str, Any]) -> dict[str, Any]: + pattern = str(args.get("pattern") or "").strip() + root_raw = str(args.get("root") or ".").strip() + glob_pat = str(args.get("glob") or "").strip() or None + case_insensitive = bool(args.get("case_insensitive", False)) + word_regexp = bool(args.get("word_regexp", False)) + files_with_matches = bool(args.get("files_with_matches", False)) + count_only = bool(args.get("count", False)) + context_lines = int(args.get("context_lines") or 0) + max_results = int(args.get("max_results") or 200) + + if not pattern: + return {"ok": False, "error_code": "pattern_required", "detail": "pattern is required"} + if max_results <= 0: + max_results = 1 + if context_lines < 0: + context_lines = 0 + + try: + root = resolve_workspace_path(root_raw) + except Exception as exc: + return {"ok": False, "error_code": "invalid_root", "detail": str(exc)} + + engine = _detect_engine() + + if engine == "ripgrep": + result = _rg_search( + pattern, root, + case_insensitive=case_insensitive, + word_regexp=word_regexp, + files_with_matches=files_with_matches, + count_only=count_only, + context_lines=context_lines, + glob=glob_pat, + max_results=max_results, + ) + if result.get("ok"): + return result + rg_error = result.get("detail", "") + else: + rg_error = "" + + py_result = _py_search( + pattern, root, + case_insensitive=case_insensitive, + word_regexp=word_regexp, + files_with_matches=files_with_matches, + count_only=count_only, + context_lines=context_lines, + glob=glob_pat, + max_results=max_results, + ) + if not py_result.get("ok"): + return py_result + + if rg_error and engine == "ripgrep": + py_result["_rg_fallback_reason"] = rg_error + return py_result + + # ── ToolSpec ────────────────────────────────────────────────── + return ToolSpec( + name="grep", + description="Fast file content search (grep). Uses ripgrep (`rg`) if available, falls back to Python regex. Supports regex, case-insensitive, word-regexp, files-with-matches, count, and context lines. Automatically respects .gitignore when using ripgrep.", + parameters={ + "type": "object", + "properties": { + "pattern": { + "type": "string", + "description": "Search pattern (regex).", + }, + "root": { + "type": "string", + "default": ".", + "description": "Directory to search under.", + }, + "glob": { + "type": "string", + "default": "", + "description": "Only search files matching glob, e.g. '*.py' or '*.{py,js}'. When omitted, all files are searched.", + }, + "case_insensitive": { + "type": "boolean", + "default": False, + "description": "Case-insensitive search.", + }, + "word_regexp": { + "type": "boolean", + "default": False, + "description": "Only match whole words (wraps pattern in \\b...\\b).", + }, + "files_with_matches": { + "type": "boolean", + "default": False, + "description": "Only list filenames that contain matches (like `rg -l`).", + }, + "count": { + "type": "boolean", + "default": False, + "description": "Only return match count per file (like `rg -c`).", + }, + "context_lines": { + "type": "integer", + "default": 0, + "description": "Number of context lines before and after each match.", + }, + "max_results": { + "type": "integer", + "default": 200, + "description": "Maximum number of matches to return.", + }, + }, + "required": ["pattern"], + "additionalProperties": False, + }, + handler=_handler, + tags=frozenset({"public", "search", "workspace", "grep"}), + risk_level="low", + read_only=True, + timeout_s=35.0, + ) + + +__all__ = ["grep_tool"] diff --git a/runtime/tools/public/workspace_profile_tool.py b/runtime/tools/public/workspace_profile_tool.py new file mode 100644 index 00000000..d587c715 --- /dev/null +++ b/runtime/tools/public/workspace_profile_tool.py @@ -0,0 +1,569 @@ +from __future__ import annotations + +import os +import subprocess +from pathlib import Path +from typing import Any + +from runtime.tools.base import ToolSpec +from runtime.tools.path_guard import resolve_workspace_path + + +# ── Language extension map ──────────────────────────────────────── +LANGUAGE_MAP: dict[str, str] = { + ".py": "Python", + ".js": "JavaScript", + ".mjs": "JavaScript", + ".cjs": "JavaScript", + ".jsx": "JavaScript React", + ".ts": "TypeScript", + ".tsx": "TypeScript React", + ".go": "Go", + ".java": "Java", + ".rs": "Rust", + ".rb": "Ruby", + ".php": "PHP", + ".swift": "Swift", + ".kt": "Kotlin", + ".kts": "Kotlin", + ".scala": "Scala", + ".cs": "C#", + ".cpp": "C++", + ".cc": "C++", + ".cxx": "C++", + ".c": "C", + ".h": "C/C++ Header", + ".hpp": "C++ Header", + ".sh": "Shell", + ".bash": "Shell", + ".zsh": "Shell", + ".bat": "Batch", + ".cmd": "Batch", + ".ps1": "PowerShell", + ".md": "Markdown", + ".json": "JSON", + ".yaml": "YAML", + ".yml": "YAML", + ".toml": "TOML", + ".xml": "XML", + ".html": "HTML", + ".htm": "HTML", + ".css": "CSS", + ".scss": "SCSS", + ".sass": "Sass", + ".less": "Less", + ".sql": "SQL", + ".env": "Environment Variables", + ".ini": "INI Config", + ".cfg": "Config", + ".txt": "Text", + ".vue": "Vue", + ".svelte": "Svelte", + ".ex": "Elixir", + ".exs": "Elixir", + ".erl": "Erlang", + ".dart": "Dart", + ".lua": "Lua", + ".r": "R", + ".jl": "Julia", + ".hs": "Haskell", + ".zig": "Zig", + ".nim": "Nim", + ".pl": "Perl", + ".pm": "Perl Module", + ".tcl": "Tcl", + ".clj": "Clojure", + ".cljs": "ClojureScript", +} + +# ── Entry point patterns ────────────────────────────────────────── +ENTRY_PATTERNS = frozenset({ + "main.py", "app.py", "index.py", "cli.py", "manage.py", + "wsgi.py", "asgi.py", "run.py", "server.py", + "index.js", "index.ts", "app.js", "app.ts", "server.js", "server.ts", + "main.go", "main.rs", "main.java", "Main.java", + "main.kt", "main.scala", +}) + +# ── Test framework indicators ───────────────────────────────────── +TEST_FILES: dict[str, list[str]] = { + "pytest": ["pytest.ini", "conftest.py", "pyproject.toml"], + "jest": ["jest.config.js", "jest.config.ts", "jest.config.json", "jest.config.mjs"], + "vitest": ["vitest.config.ts", "vitest.config.js"], + "mocha": [".mocharc.yml", ".mocharc.js", ".mocharc.json"], + "playwright": ["playwright.config.ts", "playwright.config.js"], + "cypress": ["cypress.config.ts", "cypress.config.js"], + "rspec": [".rspec"], +} + +# ── Package manager indicators ──────────────────────────────────── +PKG_FILES: dict[str, list[str]] = { + "pip": ["requirements.txt", "setup.py", "setup.cfg"], + "poetry": ["pyproject.toml"], + "uv": ["uv.lock"], + "npm": ["package.json", "package-lock.json"], + "yarn": ["yarn.lock"], + "pnpm": ["pnpm-lock.yaml"], + "bun": ["bun.lockb"], + "cargo": ["Cargo.toml", "Cargo.lock"], + "go_modules": ["go.mod", "go.sum"], + "maven": ["pom.xml"], + "gradle": ["build.gradle", "build.gradle.kts"], + "composer": ["composer.json", "composer.lock"], + "gem": ["Gemfile", "Gemfile.lock"], + "swiftpm": ["Package.swift"], + "mix": ["mix.exs"], +} + +# Directories to skip during scanning +SKIP_DIRS = frozenset({ + ".git", "__pycache__", "node_modules", ".venv", "venv", + ".tox", ".eggs", "eggs", ".mypy_cache", ".pytest_cache", + ".ruff_cache", ".hypothesis", ".coverage", "htmlcov", + "dist", "build", ".next", ".nuxt", ".output", + "target", "vendor", ".bundle", ".gradle", ".idea", ".vscode", + ".svn", ".hg", ".DS_Store", +}) + +# Hidden dirs we still descend into (CI, hooks, local toolchain metadata) +KEEP_DOT_DIRS = frozenset({ + ".github", ".husky", ".cargo", ".gitea", ".gitlab", ".buildkite", ".woodpecker", +}) + +# Hard cap so callers cannot request an unbounded walk +_MAX_FILES_CAP = 50_000 + + +def _walkable_dir(name: str) -> bool: + if name in SKIP_DIRS: + return False + if name.startswith(".") and name not in KEEP_DOT_DIRS: + return False + return True + + +def _scan_workspace(root: Path, max_files: int = 5000) -> tuple[list[dict[str, Any]], int]: + """Scan workspace files, returning (file_infos, skipped_dir_entries).""" + files: list[dict[str, Any]] = [] + skipped_dirs_count = 0 + + try: + for dirpath_str, dirnames, filenames in os.walk(root, followlinks=False): + if len(files) >= max_files: + dirnames[:] = [] + continue + + dirpath = Path(dirpath_str) + before = list(dirnames) + dirnames[:] = [d for d in dirnames if _walkable_dir(d)] + skipped_dirs_count += len(before) - len(dirnames) + + rel = dirpath.relative_to(root) if dirpath != root else Path(".") + + for fname in filenames: + if len(files) >= max_files: + break + if fname.startswith(".") and fname not in ( + ".env", ".gitignore", ".dockerignore", ".editorconfig", ".prettierrc", + ): + continue + + try: + fp = dirpath / fname + st = fp.stat() + except OSError: + continue + + suffix = fp.suffix.lower() + ext = suffix or (fname if "." not in fname else "") + + if st.st_size == 0 and suffix not in (".env", ".gitignore"): + continue + + files.append({ + "path": str(rel / fname) if str(rel) != "." else fname, + "size": st.st_size, + "ext": ext, + "suffix": suffix, + }) + + if len(files) >= max_files: + dirnames[:] = [] + except PermissionError: + pass + + return files, skipped_dirs_count + + +def _build_language_stats( + files: list[dict[str, Any]], +) -> dict[str, Any]: + """Build language breakdown from scanned files.""" + ext_counter: dict[str, int] = {} + ext_bytes: dict[str, int] = {} + lang_files: dict[str, int] = {} + lang_bytes: dict[str, int] = {} + + for f in files: + ext = f["ext"] + ext_counter[ext] = ext_counter.get(ext, 0) + 1 + ext_bytes[ext] = ext_bytes.get(ext, 0) + f["size"] + + for ext, count in ext_counter.items(): + lang = LANGUAGE_MAP.get(ext, ext.lstrip(".").capitalize() if ext else "Other") + lang_files[lang] = lang_files.get(lang, 0) + count + lang_bytes[lang] = lang_bytes.get(lang, 0) + ext_bytes.get(ext, 0) + + total_files = sum(lang_files.values()) + total_bytes = sum(lang_bytes.values()) + + sorted_langs = sorted(lang_files.items(), key=lambda x: -x[1]) + + breakdown = [] + for lang, count in sorted_langs: + b = lang_bytes[lang] + breakdown.append({ + "language": lang, + "files": count, + "bytes": b, + "pct": round(count / total_files * 100, 1) if total_files else 0, + }) + + return { + "total_files": total_files, + "total_bytes": total_bytes, + "languages": breakdown, + "top_language": breakdown[0]["language"] if breakdown else None, + } + + +def _find_largest_files( + files: list[dict[str, Any]], top_n: int = 10, +) -> list[dict[str, Any]]: + """Find the largest files by byte size.""" + sorted_files = sorted(files, key=lambda x: -x["size"]) + result = [] + for f in sorted_files[:top_n]: + result.append({ + "path": f["path"], + "bytes": f["size"], + "ext": f["suffix"] or f["ext"], + }) + return result + + +def _looks_like_test_file(rel: str) -> bool: + pp = Path(rel) + name = pp.name + if name.startswith("test_") or name.endswith("_test.py"): + return True + parts = pp.parts + if "tests" in parts: + return True + return any(part.startswith("test_") for part in parts) + + +def _detect_test_framework(files: list[dict[str, Any]]) -> dict[str, Any]: + """Detect test frameworks used in the project.""" + detected: list[str] = [] + configs: list[str] = [] + test_file_count = 0 + + file_names = {Path(f["path"]).name for f in files} + + for framework, indicators in TEST_FILES.items(): + for ind in indicators: + if ind in file_names: + detected.append(framework) + configs.append(ind) + break + + for f in files: + if _looks_like_test_file(f["path"]): + test_file_count += 1 + + return { + "detected": len(detected) > 0 or test_file_count > 0, + "frameworks": detected if detected else (["unittest"] if test_file_count > 0 else []), + "config_files": configs, + "test_file_count": test_file_count, + } + + +def _detect_package_manager(files: list[dict[str, Any]]) -> dict[str, Any]: + """Detect package managers used in the project.""" + detected: list[str] = [] + configs: list[str] = [] + file_names = {Path(f["path"]).name for f in files} + + for manager, indicators in PKG_FILES.items(): + for ind in indicators: + if ind in file_names: + detected.append(manager) + configs.append(ind) + break + + if any(Path(f["path"]).suffix.lower() == ".csproj" for f in files): + detected.append("nuget") + first = next((f["path"] for f in files if Path(f["path"]).suffix.lower() == ".csproj"), "") + if first: + configs.append(Path(first).name) + + return { + "detected": len(detected) > 0, + "managers": detected, + "config_files": configs, + } + + +def _find_entry_points(files: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Find common project entry points.""" + entries: list[dict[str, Any]] = [] + file_names = {Path(f["path"]).name for f in files} + + for pattern in ENTRY_PATTERNS: + if pattern in file_names: + entries.append({ + "path": pattern, + "type": _classify_entry(pattern), + }) + + for f in files: + parts = Path(f["path"]).parts + if len(parts) >= 2 and parts[0] in ("cmd", "bin"): + name = Path(f["path"]).name + if name not in {e["path"] for e in entries}: + entries.append({ + "path": f["path"], + "type": "cli_entry", + }) + + return entries + + +def _classify_entry(name: str) -> str: + if name.endswith(".py"): + return "python_entry" + elif name.endswith((".js", ".mjs", ".cjs")): + return "js_entry" + elif name.endswith((".ts", ".mts", ".cts")): + return "ts_entry" + elif name.endswith(".go"): + return "go_entry" + elif name.endswith(".rs"): + return "rust_entry" + elif name.endswith((".java", ".kt", ".scala")): + return "jvm_entry" + return "unknown" + + +def _get_git_status(root: Path) -> dict[str, Any]: + """Get git status summary.""" + git_root = root / ".git" + if not git_root.exists(): + return {"detected": False} + + result: dict[str, Any] = {"detected": True} + + try: + # Branch + branch_proc = subprocess.run( + ["git", "-C", str(root), "rev-parse", "--abbrev-ref", "HEAD"], + capture_output=True, text=True, timeout=5, + ) + result["branch"] = branch_proc.stdout.strip() if branch_proc.returncode == 0 else "unknown" + + # Status - porcelain + status_proc = subprocess.run( + ["git", "-C", str(root), "status", "--porcelain"], + capture_output=True, text=True, timeout=5, + ) + lines = [l.strip() for l in status_proc.stdout.strip().split("\n") if l.strip()] + modified = sum(1 for l in lines if l.startswith(" M") or l.startswith("M ") or l.startswith("MM")) + added = sum(1 for l in lines if l.startswith("A ")) + deleted = sum(1 for l in lines if l.startswith(" D") or l.startswith("D ")) + renamed = sum(1 for l in lines if l.startswith("R ")) + untracked = sum(1 for l in lines if l.startswith("??")) + + result["modified"] = modified + result["added"] = added + result["deleted"] = deleted + result["renamed"] = renamed + result["untracked"] = untracked + result["total_uncommitted"] = len(lines) + + # Recent commits + log_proc = subprocess.run( + ["git", "-C", str(root), "log", "--oneline", "-5"], + capture_output=True, text=True, timeout=5, + ) + if log_proc.returncode == 0: + recent = [l.strip() for l in log_proc.stdout.strip().split("\n") if l.strip()] + result["recent_commits"] = recent + + # Last commit + last_proc = subprocess.run( + ["git", "-C", str(root), "log", "-1", "--format=%h %s (%ai)"], + capture_output=True, text=True, timeout=5, + ) + if last_proc.returncode == 0: + result["last_commit"] = last_proc.stdout.strip() + + except Exception as exc: + result["error"] = str(exc) + + return result + + +def _path_suggests_ci(rel: str) -> bool: + p = rel.replace("\\", "/") + pl = p.lower() + if ".github/workflows/" in pl: + return True + name = Path(p).name.lower() + if name in ( + "jenkinsfile", + ".gitlab-ci.yml", + "gitlab-ci.yml", + "azure-pipelines.yml", + ".travis.yml", + "appveyor.yml", + "buildkite.yml", + "bitbucket-pipelines.yml", + ): + return True + parts = tuple(Path(pl).parts) + if ".circleci" in parts or ".woodpecker" in parts: + return True + return any(seg == "jenkinsfile" for seg in parts) + + +def _generate_suggestions( + lang_stats: dict[str, Any], + test: dict[str, Any], + pkg: dict[str, Any], + entries: list[dict[str, Any]], + git: dict[str, Any], + file_count: int, + files: list[dict[str, Any]], +) -> list[str]: + """Generate actionable suggestions based on profile.""" + suggestions: list[str] = [] + + top_lang = lang_stats.get("top_language") + + # Missing test framework + if not test.get("detected"): + suggestions.append("未检测到测试框架。建议引入 pytest 并为关键模块编写测试。") + + # Missing package manager + if not pkg.get("detected"): + if top_lang == "Python": + suggestions.append("未检测到包管理器配置文件。建议添加 pyproject.toml 或 requirements.txt。") + elif top_lang in ("JavaScript", "TypeScript"): + suggestions.append("未检测到包管理器。建议初始化 package.json(npm init)。") + elif top_lang == "Go": + suggestions.append("未检测到 Go module。建议运行 go mod init 。") + + # Many uncommitted changes + uncommitted = git.get("total_uncommitted", 0) + if git.get("detected") and uncommitted > 10: + suggestions.append(f"有 {uncommitted} 个未提交的变更。建议分批次提交以保持历史清晰。") + + # No entry points found + if not entries: + suggestions.append("未检测到常见入口文件。项目结构可能需要明确的主入口。") + + if file_count > 200: + has_ci = any(_path_suggests_ci(f["path"]) for f in files) + if not has_ci: + suggestions.append( + "文件量较多但未发现常见 CI 配置(如 .github/workflows、Jenkinsfile、" + ".gitlab-ci.yml、Azure Pipelines、CircleCI、Bitbucket Pipelines)。可考虑补充自动化流水线。" + ) + + return suggestions + + +def workspace_profile_tool() -> ToolSpec: + """Workspace profile — one-shot codebase overview.""" + + def _handler(args: dict[str, Any]) -> dict[str, Any]: + root_raw = str(args.get("root") or ".").strip() + max_files = int(args.get("max_files") or 5000) + if max_files <= 0: + max_files = 1 + if max_files > _MAX_FILES_CAP: + max_files = _MAX_FILES_CAP + + try: + root = resolve_workspace_path(root_raw) + except Exception as exc: + return {"ok": False, "error_code": "invalid_root", "detail": str(exc)} + + if not root.is_dir(): + return {"ok": False, "error_code": "not_a_directory", "detail": str(root)} + + # ── Scan ────────────────────────────────────────────────── + files, skipped = _scan_workspace(root, max_files=max_files) + + # ── Profile sections ────────────────────────────────────── + lang_stats = _build_language_stats(files) + largest = _find_largest_files(files, top_n=10) + test_info = _detect_test_framework(files) + pkg_info = _detect_package_manager(files) + entries = _find_entry_points(files) + git_info = _get_git_status(root) + suggestions = _generate_suggestions( + lang_stats, test_info, pkg_info, entries, git_info, len(files), files, + ) + + # ── Assemble ────────────────────────────────────────────── + profile = { + "project_name": root.name, + "total_files_scanned": len(files), + "skipped_dirs": skipped, + "truncated": len(files) >= max_files, + "languages": lang_stats, + "largest_files": largest, + "test_framework": test_info, + "package_manager": pkg_info, + "entry_points": entries, + "git": git_info, + "suggestions": suggestions, + } + + return {"ok": True, "profile": profile} + + return ToolSpec( + name="workspace_profile", + description=( + "Quick codebase profiling: language breakdown, largest files, test framework, package manager, " + "entry points, git status, and actionable suggestions. Filesystem walk with common vendor/build " + "directories skipped; not .gitignore-aware (unlike ripgrep). One-shot overview for unfamiliar projects." + ), + parameters={ + "type": "object", + "properties": { + "root": { + "type": "string", + "default": ".", + "description": "Directory to profile (relative to workspace root).", + }, + "max_files": { + "type": "integer", + "default": 5000, + "description": "Maximum files to collect (1–50000; prevents hangs on large repos).", + }, + }, + "additionalProperties": False, + }, + handler=_handler, + tags=frozenset({"public", "profile", "workspace", "analysis", "readonly"}), + risk_level="low", + read_only=True, + timeout_s=60.0, + ) + + +__all__ = ["workspace_profile_tool"] diff --git a/runtime/tools/public/write_file_tool.py b/runtime/tools/public/write_file_tool.py index 70f7e9f9..c17e8f66 100644 --- a/runtime/tools/public/write_file_tool.py +++ b/runtime/tools/public/write_file_tool.py @@ -1,6 +1,5 @@ from __future__ import annotations -from pathlib import Path from typing import Any from runtime.tools.base import ToolSpec @@ -8,27 +7,16 @@ from runtime.tools.path_guard import resolve_workspace_path def write_file_tool() -> ToolSpec: - def _normalize_write_path(path: str) -> str: - raw = str(path or "").strip().strip('"').strip("'") - if not raw: - raise ValueError("path_required") - p = Path(raw) - if p.is_absolute(): - return raw - rel = raw.lstrip("./\\") - if not rel: - raise ValueError("path_required") - return str(Path("data") / "workspace" / rel) - def _handler(args: dict[str, Any]) -> dict[str, Any]: - path = str(args.get("path") or "").strip() + raw = str(args.get("path") or "").strip().strip('"').strip("'") + if not raw: + return {"ok": False, "error": "path_required"} content = str(args.get("content") or "") mode = str(args.get("mode") or "overwrite").strip().lower() try: - normalized = _normalize_write_path(path) + p = resolve_workspace_path(raw) except ValueError as exc: return {"ok": False, "error": str(exc)} - p = resolve_workspace_path(normalized) p.parent.mkdir(parents=True, exist_ok=True) if mode not in ("overwrite", "append"): return {"ok": False, "error": "invalid_mode", "allowed": ["overwrite", "append"]} diff --git a/tests/test_local_public_tools.py b/tests/test_local_public_tools.py index b5533a27..3a1eb5b6 100644 --- a/tests/test_local_public_tools.py +++ b/tests/test_local_public_tools.py @@ -123,7 +123,7 @@ def test_local_tool_integration_roundtrip(monkeypatch) -> None: tmpdir = Path(tempfile.mkdtemp(prefix="local_it_")) monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmpdir)) monkeypatch.setenv("AIA_ENABLE_RUN_COMMAND", "1") - target_rel = "data/workspace/it_sample.txt" + target_rel = "it_sample.txt" out_write = write_spec.handler({"path": "it_sample.txt", "content": "line1\nline2\n", "mode": "overwrite"}) assert out_write.get("ok") is True, out_write diff --git a/tests/test_workspace_path_guard.py b/tests/test_workspace_path_guard.py index 837577e2..0d952be3 100644 --- a/tests/test_workspace_path_guard.py +++ b/tests/test_workspace_path_guard.py @@ -103,7 +103,7 @@ class WorkspacePathGuardTests(unittest.TestCase): p = resolve_workspace_path(str(f)) self.assertEqual(p, f.resolve()) - def test_write_file_relative_path_defaults_to_data_workspace_subdir(self) -> None: + def test_write_file_relative_path_resolves_under_workspace_root(self) -> None: with mock.patch.dict( os.environ, {"OPS_WORKSPACE_ROOT": str(self.root), "OPS_WORKSPACE_EXTRA_ROOTS": "", "OPS_WORKSPACE_ALLOW_ANY_PATH": ""}, @@ -114,11 +114,11 @@ class WorkspacePathGuardTests(unittest.TestCase): with workspace_path_access_scope(None, None): r = spec.handler({"path": "generated.py", "content": "print('ok')\n", "mode": "overwrite"}) self.assertTrue(r.get("ok"), r) - expected = (self.root / "data" / "workspace" / "generated.py").resolve() + expected = (self.root / "generated.py").resolve() self.assertEqual(str(expected), str(r.get("path"))) self.assertTrue(expected.exists()) - def test_write_file_relative_path_uses_workspace_namespace_scope(self) -> None: + def test_write_file_relative_path_with_write_namespace_scope(self) -> None: with mock.patch.dict( os.environ, {"OPS_WORKSPACE_ROOT": str(self.root), "OPS_WORKSPACE_EXTRA_ROOTS": "", "OPS_WORKSPACE_ALLOW_ANY_PATH": ""}, @@ -130,7 +130,7 @@ class WorkspacePathGuardTests(unittest.TestCase): self.assertEqual(current_workspace_write_namespace(), "ops") r = spec.handler({"path": "generated.py", "content": "print('ok')\n", "mode": "overwrite"}) self.assertTrue(r.get("ok"), r) - expected = (self.root / "data" / "workspace" / "generated.py").resolve() + expected = (self.root / "generated.py").resolve() self.assertEqual(str(expected), str(r.get("path"))) self.assertTrue(expected.exists()) diff --git a/tests/test_workspace_profile_tool.py b/tests/test_workspace_profile_tool.py new file mode 100644 index 00000000..2195a357 --- /dev/null +++ b/tests/test_workspace_profile_tool.py @@ -0,0 +1,59 @@ +"""Unit tests for workspace_profile_tool helpers and handler.""" + +from __future__ import annotations + +from pathlib import Path + +import pytest + +from runtime.tools.path_guard import clear_workspace_path_access_for_tests, workspace_path_access_scope +from runtime.tools.public import workspace_profile_tool as mod + + +def test_walkable_dir_github_vs_vendor() -> None: + assert mod._walkable_dir(".github") is True + assert mod._walkable_dir(".gitlab") is True + assert mod._walkable_dir("node_modules") is False + assert mod._walkable_dir(".cache") is False + + +def test_path_suggests_ci() -> None: + assert mod._path_suggests_ci(".github/workflows/ci.yml") + assert mod._path_suggests_ci("pkg/.circleci/config.yml") + assert mod._path_suggests_ci("bitbucket-pipelines.yml") + assert mod._path_suggests_ci("ci/Jenkinsfile") + assert mod._path_suggests_ci(str(Path("x") / ".woodpecker" / "ci.yaml")) + assert not mod._path_suggests_ci("src/main.py") + + +def test_looks_like_test_file() -> None: + assert mod._looks_like_test_file("tests/unit/test_x.py") + assert mod._looks_like_test_file(str(Path("src") / "tests" / "a.py")) + assert mod._looks_like_test_file("test_foo.py") + assert not mod._looks_like_test_file("src/main.py") + + +def test_detect_package_manager_nuget_csproj() -> None: + files = [ + {"path": "src/App.csproj", "size": 12, "ext": ".csproj", "suffix": ".csproj"}, + ] + out = mod._detect_package_manager(files) + assert out["detected"] is True + assert "nuget" in out["managers"] + assert "App.csproj" in out["config_files"] + + +def test_workspace_profile_max_files_clamped(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmp_path)) + (tmp_path / "a.txt").write_text("hello", encoding="utf-8") + clear_workspace_path_access_for_tests() + spec = mod.workspace_profile_tool() + with workspace_path_access_scope(None, None): + r = spec.handler({"root": ".", "max_files": 0}) + assert r.get("ok") is True + assert (r.get("profile") or {}).get("total_files_scanned") == 1 + + with workspace_path_access_scope(None, None): + r2 = spec.handler({"root": ".", "max_files": 999_999}) + assert r2.get("ok") is True + assert (r2.get("profile") or {}).get("truncated") is False