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 <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-05-13 23:02:52 +08:00
parent 48190afbc1
commit 2b32d11f43
6 changed files with 1019 additions and 21 deletions

View file

@ -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"]

View file

@ -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 <module-name>。")
# 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"]

View file

@ -1,6 +1,5 @@
from __future__ import annotations from __future__ import annotations
from pathlib import Path
from typing import Any from typing import Any
from runtime.tools.base import ToolSpec 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 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]: 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 "") content = str(args.get("content") or "")
mode = str(args.get("mode") or "overwrite").strip().lower() mode = str(args.get("mode") or "overwrite").strip().lower()
try: try:
normalized = _normalize_write_path(path) p = resolve_workspace_path(raw)
except ValueError as exc: except ValueError as exc:
return {"ok": False, "error": str(exc)} return {"ok": False, "error": str(exc)}
p = resolve_workspace_path(normalized)
p.parent.mkdir(parents=True, exist_ok=True) p.parent.mkdir(parents=True, exist_ok=True)
if mode not in ("overwrite", "append"): if mode not in ("overwrite", "append"):
return {"ok": False, "error": "invalid_mode", "allowed": ["overwrite", "append"]} return {"ok": False, "error": "invalid_mode", "allowed": ["overwrite", "append"]}

View file

@ -123,7 +123,7 @@ def test_local_tool_integration_roundtrip(monkeypatch) -> None:
tmpdir = Path(tempfile.mkdtemp(prefix="local_it_")) tmpdir = Path(tempfile.mkdtemp(prefix="local_it_"))
monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmpdir)) monkeypatch.setenv("OPS_WORKSPACE_ROOT", str(tmpdir))
monkeypatch.setenv("AIA_ENABLE_RUN_COMMAND", "1") 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"}) out_write = write_spec.handler({"path": "it_sample.txt", "content": "line1\nline2\n", "mode": "overwrite"})
assert out_write.get("ok") is True, out_write assert out_write.get("ok") is True, out_write

View file

@ -103,7 +103,7 @@ class WorkspacePathGuardTests(unittest.TestCase):
p = resolve_workspace_path(str(f)) p = resolve_workspace_path(str(f))
self.assertEqual(p, f.resolve()) 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( with mock.patch.dict(
os.environ, os.environ,
{"OPS_WORKSPACE_ROOT": str(self.root), "OPS_WORKSPACE_EXTRA_ROOTS": "", "OPS_WORKSPACE_ALLOW_ANY_PATH": ""}, {"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): with workspace_path_access_scope(None, None):
r = spec.handler({"path": "generated.py", "content": "print('ok')\n", "mode": "overwrite"}) r = spec.handler({"path": "generated.py", "content": "print('ok')\n", "mode": "overwrite"})
self.assertTrue(r.get("ok"), r) 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.assertEqual(str(expected), str(r.get("path")))
self.assertTrue(expected.exists()) 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( with mock.patch.dict(
os.environ, os.environ,
{"OPS_WORKSPACE_ROOT": str(self.root), "OPS_WORKSPACE_EXTRA_ROOTS": "", "OPS_WORKSPACE_ALLOW_ANY_PATH": ""}, {"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") self.assertEqual(current_workspace_write_namespace(), "ops")
r = spec.handler({"path": "generated.py", "content": "print('ok')\n", "mode": "overwrite"}) r = spec.handler({"path": "generated.py", "content": "print('ok')\n", "mode": "overwrite"})
self.assertTrue(r.get("ok"), r) 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.assertEqual(str(expected), str(r.get("path")))
self.assertTrue(expected.exists()) self.assertTrue(expected.exists())

View file

@ -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