714 lines
26 KiB
Python
714 lines
26 KiB
Python
"""技能地址(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)
|