deerflow-code/offline-backend-20260512/backend/app/gateway/routers/skill_addresses.py
2026-09-07 18:24:55 +08:00

714 lines
26 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""技能地址(URL/IP)批量替换 API (技能地址管理).
Scans the **running** skills directory (``skills/{public,custom}``) — the files
the agent actually reads — for ``http(s)://`` URLs and IPv4 (``ip[:port]``)
addresses, aggregates each unique address with every place it appears (across
``.md`` and ``.py``, across files), and lets an admin rewrite them in bulk.
Why this shape: the same address routinely lives in both a skill's ``SKILL.md``
(documented endpoint) and its ``scripts/*.py`` (the actual call), so the natural
unit is *"one address → replace all of its occurrences at once"*. The pure
scan/replace logic lives in ``deerflow.skills.address_scan`` (harness layer,
substring-safe span replacement); this router only walks the tree, enforces
admin auth + path safety, backs up before writing, and refreshes the skills
system-prompt cache so edits take effect immediately.
**Scale / availability (150+ skills):** every filesystem walk / read / write is
offloaded to a worker thread via ``asyncio.to_thread`` so the event loop is
never blocked while scanning or rewriting — the Gateway stays responsive. The
streaming endpoints additionally emit NDJSON progress (``已处理 X / 共 N``) so the
UI can render a progress bar for both the scan and the apply phases.
Routes (prefix ``/api/skill-addresses``, **admin only**):
GET /scan?kinds=url,ip&skill=<name> read-only aggregated scan (one shot)
GET /scan/stream?kinds=...&skill=... scan with per-skill NDJSON progress
POST /preview dry-run diff of a replacement set
POST /apply write changes (one shot, with backup)
POST /apply/stream write changes with per-file progress
GET /backups list restorable backups (newest first)
POST /rollback/{backup_id} restore a backup
"""
from __future__ import annotations
import asyncio
import json
import logging
import shutil
from collections.abc import AsyncIterator
from datetime import UTC, datetime
from pathlib import Path
from fastapi import APIRouter, HTTPException, Query, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, ConfigDict, Field
from app.gateway.deps import get_optional_user_from_request
from deerflow.agents.lead_agent.prompt import refresh_skills_system_prompt_cache_async
from deerflow.skills.address_scan import (
VALID_KINDS,
apply_replacements,
infer_kind,
is_valid_replacement,
normalize_kinds,
scan_occurrences,
)
from deerflow.skills.storage import get_or_new_skill_storage
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/skill-addresses", tags=["skill-addresses"])
_BACKUP_DIRNAME = ".address-edits"
_SCAN_SUBDIRS = ("public", "custom")
_SCAN_EXTS = (".md", ".py")
# ── schemas ────────────────────────────────────────────────────────────────
class OccurrenceModel(BaseModel):
skill: str
category: str # "public" | "custom"
file: str # POSIX path relative to the skills root
line: int
col: int
snippet: str
class AddressGroup(BaseModel):
address: str
kind: str # "url" | "ip"
count: int
occurrences: list[OccurrenceModel]
class ScanResponse(BaseModel):
groups: list[AddressGroup]
scanned_files: int
scanned_skills: int
kinds: list[str]
scanned_at: str
class Replacement(BaseModel):
model_config = ConfigDict(populate_by_name=True)
from_: str = Field(alias="from", min_length=1)
to: str = Field(min_length=1)
kind: str | None = None
class PreviewRequest(BaseModel):
replacements: list[Replacement] = Field(default_factory=list)
class PreviewFileHit(BaseModel):
file: str
line: int
before: str
after: str
class PreviewItem(BaseModel):
model_config = ConfigDict(populate_by_name=True)
from_: str = Field(serialization_alias="from")
to: str
kind: str
count: int
files: list[PreviewFileHit]
class PreviewResponse(BaseModel):
items: list[PreviewItem]
total: int
warnings: list[str]
expected_total: int
class ApplyRequest(BaseModel):
replacements: list[Replacement] = Field(default_factory=list)
expected_total: int | None = None
backup: bool = True
class ChangedFile(BaseModel):
file: str
count: int
class ApplyResponse(BaseModel):
changed_files: list[ChangedFile]
replaced_total: int
backup_id: str | None
warnings: list[str]
class ReplacementSummary(BaseModel):
model_config = ConfigDict(populate_by_name=True)
from_: str = Field(serialization_alias="from")
to: str
kind: str
class BackupEntry(BaseModel):
backup_id: str
files: int
created_at: str
replaced_total: int = 0
replacements: list[ReplacementSummary] = Field(default_factory=list)
class BackupListResponse(BaseModel):
backups: list[BackupEntry]
class RollbackResponse(BaseModel):
backup_id: str
restored_files: list[str]
# ── auth / path helpers ──────────────────────────────────────────────────────
async def _require_admin(request: Request) -> None:
"""Admin-only. Lets the call through when auth is disabled (no user)."""
user = await get_optional_user_from_request(request)
if user is None:
return
if getattr(user, "system_role", None) != "admin":
raise HTTPException(status_code=403, detail="技能地址管理仅限管理员")
def _skills_root() -> Path:
return Path(get_or_new_skill_storage().get_skills_root_path()).resolve()
def _is_within(root: Path, p: Path) -> bool:
try:
p.resolve().relative_to(root)
return True
except (ValueError, OSError):
return False
def _skill_of(root: Path, path: Path) -> tuple[str, str, str]:
"""Return ``(category, skill_name, posix_rel_path)`` for a scanned file."""
rel = path.relative_to(root)
parts = rel.parts
category = parts[0] if parts else ""
skill_name = parts[1] if len(parts) > 1 else ""
return category, skill_name, rel.as_posix()
def _read_text(path: Path) -> str | None:
"""Read UTF-8 text preserving original line endings; skip binary/bad files."""
try:
with open(path, "r", encoding="utf-8", newline="") as fh:
return fh.read()
except (UnicodeDecodeError, OSError):
return None
def _write_text(path: Path, text: str) -> None:
with open(path, "w", encoding="utf-8", newline="") as fh:
fh.write(text)
# ── synchronous workers (run via asyncio.to_thread; never block the loop) ─────
def _iter_skill_files(root: Path, skill: str | None = None) -> list[Path]:
"""All ``*.md`` / ``*.py`` under ``{root}/{public,custom}``.
Hidden directories (``.address-edits`` backups, ``.history``) are skipped.
When *skill* is given, only that skill's directory is walked.
"""
files: list[Path] = []
for sub in _SCAN_SUBDIRS:
base = root / sub
if skill:
base = base / skill
if not base.is_dir():
continue
for p in base.rglob("*"):
if not p.is_file() or p.suffix.lower() not in _SCAN_EXTS:
continue
try:
rel_parts = p.relative_to(root).parts
except ValueError:
continue
if any(part.startswith(".") for part in rel_parts):
continue
files.append(p)
return files
def _group_files_by_skill(root: Path, files: list[Path]) -> dict[tuple[str, str], list[Path]]:
"""Bucket scanned files under their owning ``(category, skill_name)``."""
buckets: dict[tuple[str, str], list[Path]] = {}
for p in files:
category, skill_name, _ = _skill_of(root, p)
buckets.setdefault((category, skill_name), []).append(p)
return buckets
def _occurrence_kind(value: str) -> str:
return infer_kind(value) or "url"
def _merge_occurrences(groups: dict[str, AddressGroup], pairs: list[tuple[str, OccurrenceModel]]) -> None:
"""Merge ``(address, occurrence)`` pairs into the aggregate group map."""
for address, occ in pairs:
g = groups.get(address)
if g is None:
g = AddressGroup(address=address, kind=_occurrence_kind(address), count=0, occurrences=[])
groups[address] = g
g.count += 1
g.occurrences.append(occ)
def _plan_apply(root: Path, mapping: dict[str, str]) -> tuple[list[tuple[Path, str, int, str]], int]:
"""Compute the rewrite plan without touching disk → (plan, current_total)."""
plan: list[tuple[Path, str, int, str]] = []
current_total = 0
for path in _iter_skill_files(root):
text = _read_text(path)
if text is None:
continue
new_text, cnt = apply_replacements(text, mapping, VALID_KINDS)
if cnt <= 0:
continue
if not _is_within(root, path):
raise HTTPException(status_code=400, detail=f"拒绝写入技能目录之外的文件:{path}")
_, _, rel = _skill_of(root, path)
plan.append((path, new_text, cnt, rel))
current_total += cnt
return plan, current_total
def _backup_file(root: Path, backup_id: str, path: Path, rel: str) -> None:
dest = root / _BACKUP_DIRNAME / backup_id / rel
dest.parent.mkdir(parents=True, exist_ok=True)
shutil.copy2(path, dest)
_MANIFEST_SUFFIX = ".manifest.json"
def _manifest_path(root: Path, backup_id: str) -> Path:
"""Manifest lives as a *sibling* of the backup dir (``{id}.manifest.json``),
not inside it — so it is neither restored back into the skills tree on
rollback nor counted as a backed-up file."""
return root / _BACKUP_DIRNAME / f"{backup_id}{_MANIFEST_SUFFIX}"
def _write_manifest(
root: Path,
backup_id: str,
norm: list[tuple[str, str, str]],
changed: list["ChangedFile"],
replaced_total: int,
created_at: str,
) -> None:
"""Record what this edit changed so the UI can show a meaningful, per-edit
history (原地址 → 新地址, 文件数, 处数) and roll a single record back."""
data = {
"created_at": created_at,
"replaced_total": replaced_total,
"replacements": [{"from": frm, "to": to, "kind": kind} for frm, to, kind in norm],
"changed_files": [{"file": c.file, "count": c.count} for c in changed],
}
try: # best-effort: a metadata write failure must not fail the apply itself
path = _manifest_path(root, backup_id)
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w", encoding="utf-8") as fh:
json.dump(data, fh, ensure_ascii=False)
except OSError:
logger.warning("Failed to write skill-address backup manifest for %s", backup_id, exc_info=True)
def _read_manifest(root: Path, backup_id: str) -> dict | None:
try:
with open(_manifest_path(root, backup_id), "r", encoding="utf-8") as fh:
return json.load(fh)
except (OSError, json.JSONDecodeError):
return None
# ── scan extraction shared by one-shot + stream ───────────────────────────────
def _scan_files_with_addresses(root: Path, files: list[Path], wanted: tuple[str, ...]) -> tuple[list[tuple[str, OccurrenceModel]], int]:
"""Scan files → ([(address, occurrence)], files_read)."""
pairs: list[tuple[str, OccurrenceModel]] = []
read = 0
for path in files:
text = _read_text(path)
if text is None:
continue
read += 1
category, skill_name, rel = _skill_of(root, path)
for o in scan_occurrences(text, wanted):
pairs.append(
(
o.value,
OccurrenceModel(skill=skill_name, category=category, file=rel, line=o.line, col=o.col, snippet=o.snippet),
)
)
return pairs, read
def _ordered_groups(groups: dict[str, AddressGroup]) -> list[AddressGroup]:
return sorted(groups.values(), key=lambda g: (-g.count, g.address))
def _build_mapping(replacements: list[Replacement]) -> tuple[dict[str, str], list[tuple[str, str, str]], list[str]]:
"""Validate replacements → ``(mapping, [(from,to,kind)], warnings)``.
Drops (with a warning) any entry that is empty, duplicated, not a recognized
URL/IP, a no-op (``from == to``), or whose ``to`` is malformed for its kind.
"""
mapping: dict[str, str] = {}
norm: list[tuple[str, str, str]] = []
warnings: list[str] = []
seen: set[str] = set()
for r in replacements:
frm = r.from_.strip()
to = r.to.strip()
if not frm:
continue
if frm in seen:
warnings.append(f"重复的原始地址已忽略:{frm}")
continue
kind = (r.kind or "").strip().lower() or infer_kind(frm)
if kind not in VALID_KINDS:
warnings.append(f"无法识别为 URL/IP,已忽略:{frm}")
continue
if to == frm:
warnings.append(f"替换值与原值相同,已忽略:{frm}")
continue
if not is_valid_replacement(kind, to):
warnings.append(f"替换值不是合法的 {kind},已忽略:{to}")
continue
seen.add(frm)
mapping[frm] = to
norm.append((frm, to, kind))
return mapping, norm, warnings
def _ndjson(obj: dict) -> bytes:
return (json.dumps(obj, ensure_ascii=False) + "\n").encode("utf-8")
# ── routes: scan ─────────────────────────────────────────────────────────────
@router.get("/scan", response_model=ScanResponse)
async def scan_addresses(
request: Request,
kinds: str | None = Query(default=None, description="逗号分隔:url,ip(默认全部)"),
skill: str | None = Query(default=None, description="仅扫描指定技能目录"),
) -> ScanResponse:
await _require_admin(request)
root = _skills_root()
wanted = normalize_kinds(kinds)
files = await asyncio.to_thread(_iter_skill_files, root, skill)
buckets = _group_files_by_skill(root, files)
groups: dict[str, AddressGroup] = {}
scanned = 0
for paths in buckets.values():
pairs, read = await asyncio.to_thread(_scan_files_with_addresses, root, paths, wanted)
_merge_occurrences(groups, pairs)
scanned += read
return ScanResponse(
groups=_ordered_groups(groups),
scanned_files=scanned,
scanned_skills=len(buckets),
kinds=list(wanted),
scanned_at=datetime.now(UTC).isoformat(),
)
@router.get("/scan/stream")
async def scan_addresses_stream(
request: Request,
kinds: str | None = Query(default=None),
skill: str | None = Query(default=None),
) -> StreamingResponse:
"""Scan with per-skill NDJSON progress.
Emits ``{"type":"start","total":N}`` then one
``{"type":"progress","done":i,"total":N,"skill":name}`` per skill, finally
``{"type":"result", ...ScanResponse...}``.
"""
await _require_admin(request)
root = _skills_root()
wanted = normalize_kinds(kinds)
async def gen() -> AsyncIterator[bytes]:
files = await asyncio.to_thread(_iter_skill_files, root, skill)
buckets = _group_files_by_skill(root, files)
total = len(buckets)
yield _ndjson({"type": "start", "total": total})
groups: dict[str, AddressGroup] = {}
scanned = 0
done = 0
for (category, skill_name), paths in buckets.items():
pairs, read = await asyncio.to_thread(_scan_files_with_addresses, root, paths, wanted)
_merge_occurrences(groups, pairs)
scanned += read
done += 1
yield _ndjson({"type": "progress", "done": done, "total": total, "skill": skill_name or category})
result = ScanResponse(
groups=_ordered_groups(groups),
scanned_files=scanned,
scanned_skills=total,
kinds=list(wanted),
scanned_at=datetime.now(UTC).isoformat(),
)
yield _ndjson({"type": "result", **result.model_dump()})
return StreamingResponse(gen(), media_type="application/x-ndjson")
# ── routes: preview ──────────────────────────────────────────────────────────
def _compute_preview(root: Path, mapping: dict[str, str], norm: list[tuple[str, str, str]]) -> tuple[list[PreviewItem], int]:
counts: dict[str, int] = {frm: 0 for frm, _, _ in norm}
files: dict[str, list[PreviewFileHit]] = {frm: [] for frm, _, _ in norm}
total = 0
if mapping:
for path in _iter_skill_files(root):
text = _read_text(path)
if text is None:
continue
_, _, rel = _skill_of(root, path)
for o in scan_occurrences(text, VALID_KINDS):
if o.value not in mapping:
continue
to = mapping[o.value]
counts[o.value] += 1
total += 1
files[o.value].append(PreviewFileHit(file=rel, line=o.line, before=o.snippet, after=o.snippet.replace(o.value, to)))
items = [PreviewItem(from_=frm, to=to, kind=kind, count=counts[frm], files=files[frm]) for frm, to, kind in norm]
return items, total
@router.post("/preview", response_model=PreviewResponse)
async def preview_replacements(request: Request, body: PreviewRequest) -> PreviewResponse:
await _require_admin(request)
root = _skills_root()
mapping, norm, warnings = _build_mapping(body.replacements)
items, total = await asyncio.to_thread(_compute_preview, root, mapping, norm)
return PreviewResponse(items=items, total=total, warnings=warnings, expected_total=total)
# ── routes: apply ────────────────────────────────────────────────────────────
def _new_backup_id() -> str:
return datetime.now(UTC).strftime("%Y%m%d-%H%M%S-%f")
@router.post("/apply", response_model=ApplyResponse)
async def apply_replacements_route(request: Request, body: ApplyRequest) -> ApplyResponse:
await _require_admin(request)
root = _skills_root()
mapping, norm, warnings = _build_mapping(body.replacements)
if not mapping:
return ApplyResponse(changed_files=[], replaced_total=0, backup_id=None, warnings=warnings)
plan, current_total = await asyncio.to_thread(_plan_apply, root, mapping)
if body.expected_total is not None and body.expected_total != current_total:
raise HTTPException(
status_code=409,
detail=f"文件自上次扫描后已变化(预期 {body.expected_total} 处,现为 {current_total} 处),请重新扫描后再应用。",
)
if not plan:
return ApplyResponse(changed_files=[], replaced_total=0, backup_id=None, warnings=warnings)
backup_id: str | None = None
if body.backup:
backup_id = _new_backup_id()
for path, _new, _cnt, rel in plan:
await asyncio.to_thread(_backup_file, root, backup_id, path, rel)
changed: list[ChangedFile] = []
replaced_total = 0
for path, new_text, cnt, rel in plan:
await asyncio.to_thread(_write_text, path, new_text)
changed.append(ChangedFile(file=rel, count=cnt))
replaced_total += cnt
if backup_id:
await asyncio.to_thread(_write_manifest, root, backup_id, norm, changed, replaced_total, datetime.now(UTC).isoformat())
await _refresh_cache_safe()
logger.info("Skill address apply: %d files, %d replacements, backup=%s", len(changed), replaced_total, backup_id)
return ApplyResponse(changed_files=changed, replaced_total=replaced_total, backup_id=backup_id, warnings=warnings)
@router.post("/apply/stream")
async def apply_replacements_stream(request: Request, body: ApplyRequest) -> StreamingResponse:
"""Apply with per-file NDJSON progress.
Validation + the optimistic-lock check happen *before* streaming (so a 409
is a normal HTTP error). Then: ``{"type":"start","total_files":M,...}`` →
one ``{"type":"progress","done":k,...}`` per written file →
``{"type":"result", ...ApplyResponse...}``.
"""
await _require_admin(request)
root = _skills_root()
mapping, norm, warnings = _build_mapping(body.replacements)
plan: list[tuple[Path, str, int, str]] = []
current_total = 0
if mapping:
plan, current_total = await asyncio.to_thread(_plan_apply, root, mapping)
if body.expected_total is not None and body.expected_total != current_total:
raise HTTPException(
status_code=409,
detail=f"文件自上次扫描后已变化(预期 {body.expected_total} 处,现为 {current_total} 处),请重新扫描后再应用。",
)
async def gen() -> AsyncIterator[bytes]:
if not plan:
result = ApplyResponse(changed_files=[], replaced_total=0, backup_id=None, warnings=warnings)
yield _ndjson({"type": "result", **result.model_dump()})
return
backup_id: str | None = None
if body.backup:
backup_id = _new_backup_id()
yield _ndjson({"type": "start", "total_files": len(plan), "total_replacements": current_total})
changed: list[ChangedFile] = []
replaced_total = 0
done = 0
for path, new_text, cnt, rel in plan:
if backup_id:
await asyncio.to_thread(_backup_file, root, backup_id, path, rel)
await asyncio.to_thread(_write_text, path, new_text)
changed.append(ChangedFile(file=rel, count=cnt))
replaced_total += cnt
done += 1
yield _ndjson({"type": "progress", "done": done, "total_files": len(plan), "file": rel, "count": cnt})
if backup_id:
await asyncio.to_thread(_write_manifest, root, backup_id, norm, changed, replaced_total, datetime.now(UTC).isoformat())
await _refresh_cache_safe()
logger.info("Skill address apply(stream): %d files, %d replacements, backup=%s", len(changed), replaced_total, backup_id)
result = ApplyResponse(changed_files=changed, replaced_total=replaced_total, backup_id=backup_id, warnings=warnings)
yield _ndjson({"type": "result", **result.model_dump()})
return StreamingResponse(gen(), media_type="application/x-ndjson")
# ── routes: backups / rollback ───────────────────────────────────────────────
def _list_backups(root: Path) -> list[BackupEntry]:
backup_root = root / _BACKUP_DIRNAME
entries: list[BackupEntry] = []
if backup_root.is_dir():
for d in backup_root.iterdir():
if not d.is_dir():
continue
file_count = sum(1 for p in d.rglob("*") if p.is_file())
manifest = _read_manifest(root, d.name)
created = ""
replaced_total = 0
replacements: list[ReplacementSummary] = []
if manifest:
created = str(manifest.get("created_at") or "")
replaced_total = int(manifest.get("replaced_total") or 0)
replacements = [
ReplacementSummary(from_=str(r.get("from", "")), to=str(r.get("to", "")), kind=str(r.get("kind", "")))
for r in manifest.get("replacements", [])
if isinstance(r, dict)
]
if not created:
try:
created = datetime.fromtimestamp(d.stat().st_mtime, UTC).isoformat()
except OSError:
created = ""
entries.append(
BackupEntry(
backup_id=d.name,
files=file_count,
created_at=created,
replaced_total=replaced_total,
replacements=replacements,
)
)
entries.sort(key=lambda e: e.backup_id, reverse=True)
return entries
@router.get("/backups", response_model=BackupListResponse)
async def list_backups(request: Request) -> BackupListResponse:
await _require_admin(request)
root = _skills_root()
entries = await asyncio.to_thread(_list_backups, root)
return BackupListResponse(backups=entries)
def _restore_backup(root: Path, backup_root: Path) -> list[str]:
restored: list[str] = []
for src in backup_root.rglob("*"):
if not src.is_file():
continue
rel = src.relative_to(backup_root)
dest = root / rel
if not _is_within(root, dest):
continue
dest.parent.mkdir(parents=True, exist_ok=True)
shutil.copy2(src, dest)
restored.append(rel.as_posix())
return restored
@router.post("/rollback/{backup_id}", response_model=RollbackResponse)
async def rollback(request: Request, backup_id: str) -> RollbackResponse:
await _require_admin(request)
root = _skills_root()
if "/" in backup_id or "\\" in backup_id or backup_id in (".", ".."):
raise HTTPException(status_code=400, detail="非法的备份编号")
backup_root = (root / _BACKUP_DIRNAME / backup_id).resolve()
if not _is_within(root, backup_root) or not backup_root.is_dir():
raise HTTPException(status_code=404, detail="备份不存在")
restored = await asyncio.to_thread(_restore_backup, root, backup_root)
await _refresh_cache_safe()
logger.info("Skill address rollback: backup=%s, restored %d files", backup_id, len(restored))
return RollbackResponse(backup_id=backup_id, restored_files=restored)
# ── misc ─────────────────────────────────────────────────────────────────────
async def _refresh_cache_safe() -> None:
try:
await refresh_skills_system_prompt_cache_async()
except Exception: # noqa: BLE001 — cache refresh is best-effort
logger.warning("Failed to refresh skills prompt cache", exc_info=True)