194 lines
11 KiB
Python
194 lines
11 KiB
Python
"""Asynchronous, file-backed assistant export and ordinary Wiki restore."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from pathlib import Path
|
|
from uuid import uuid4
|
|
|
|
from fastapi import APIRouter, File, HTTPException, Request, UploadFile
|
|
from fastapi.responses import FileResponse
|
|
|
|
from app.gateway.llmwiki_index_scheduler import spawn_local_wiki_index_job
|
|
from app.gateway.routers.assistant_knowledge import _actor, _list_export_manifests, _package_root, _read_export_manifest, _store, _write_export_manifest
|
|
from deerflow.assistant_knowledge.archive import PackageArchive
|
|
from deerflow.assistant_knowledge.package_io import package_path, write_jsonl_record
|
|
from deerflow.integrations.weknora.local_index.normalize import normalize_wiki_page
|
|
from deerflow.integrations.weknora.runtime import build_weknora_client, get_resolved_llmwiki_runtime
|
|
from deerflow.persistence.assistant_knowledge.store import _vector_blob_from_payload
|
|
|
|
router = APIRouter(tags=["wiki-packages"])
|
|
|
|
|
|
def job_view(row):
|
|
if row["mapping_id"].startswith("assistant:"):
|
|
base_id = row["mapping_id"].split(":", 1)[1]
|
|
row["package_download_url"] = f"/api/assistant-knowledge/bases/{base_id}/export-package/{row['id']}/download"
|
|
elif row["mapping_id"].startswith("restore:"):
|
|
row["package_download_url"] = None
|
|
return row
|
|
|
|
|
|
async def export_assistant(app, base_id, job_id):
|
|
scope = f"assistant:{base_id}"
|
|
path = package_path(_package_root(), job_id)
|
|
temp = path.with_suffix(".tmp")
|
|
counts = {}
|
|
try:
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="running", phase="exporting")
|
|
with temp.open("w", encoding="utf-8", newline="\n") as handle:
|
|
async for kind, payload in app.state.assistant_knowledge_store.iter_export_records(base_id):
|
|
await asyncio.to_thread(write_jsonl_record, handle, kind, payload)
|
|
counts[kind] = counts.get(kind, 0) + 1
|
|
if counts.get("wiki_page") and not counts.get("vector"):
|
|
raise ValueError("该助手知识库尚无可导出的 Wiki 向量,请先从已向量化的普通知识库同步")
|
|
await asyncio.to_thread(temp.replace, path)
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="completed", phase="package_ready", counts=counts)
|
|
except Exception as exc:
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="failed", phase="failed", counts=counts, error=str(exc))
|
|
|
|
|
|
@router.post("/api/assistant-knowledge/bases/{base_id}/export-package", status_code=202)
|
|
async def create_export(request: Request, base_id: str):
|
|
await _actor(request, admin=True)
|
|
if not await _store(request).get_base(base_id):
|
|
raise HTTPException(404, "助手知识库不存在")
|
|
job_id = str(uuid4())
|
|
row = _write_export_manifest(job_id=job_id, mapping_id=f"assistant:{base_id}", status="queued", phase="queued")
|
|
spawn_local_wiki_index_job(request.app, export_assistant(request.app, base_id, job_id))
|
|
return job_view(row)
|
|
|
|
|
|
@router.get("/api/assistant-knowledge/bases/{base_id}/export-package")
|
|
async def list_exports(request: Request, base_id: str):
|
|
await _actor(request, admin=True)
|
|
return {"jobs": [job_view(row) for row in await asyncio.to_thread(_list_export_manifests, f"assistant:{base_id}", limit=100)]}
|
|
|
|
|
|
@router.get("/api/assistant-knowledge/bases/{base_id}/export-package/{job_id}/download")
|
|
async def download_export(request: Request, base_id: str, job_id: str):
|
|
await _actor(request, admin=True)
|
|
row = _read_export_manifest(job_id)
|
|
if not row or row["mapping_id"] != f"assistant:{base_id}" or row["status"] != "completed" or not row["file_exists"]:
|
|
raise HTTPException(404, "导出包尚未就绪")
|
|
return FileResponse(package_path(_package_root(), job_id), media_type="application/x-ndjson", filename=f"wiki-vector-{job_id}.jsonl")
|
|
|
|
|
|
async def restore_package(app, mapping, job_id, path: Path):
|
|
scope = f"restore:{mapping['id']}"
|
|
archive = None
|
|
counts = {"wiki_page": 0, "vector": 0}
|
|
store = app.state.llmwiki_index_store
|
|
owner = f"package:{job_id}"
|
|
leased = False
|
|
try:
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="running", phase="validating")
|
|
archive = await asyncio.to_thread(PackageArchive, path)
|
|
embedding = app.state.llmwiki_embedding
|
|
if embedding is None:
|
|
raise ValueError("请配置与数据包相同的 Wiki 编码模型,以便复用向量进行检索")
|
|
if embedding.dimensions is None:
|
|
await embedding.embed_query("Wiki embedding dimension probe")
|
|
# Validate the complete package before the first remote mutation.
|
|
slugs = {p["slug"] for p in archive["wiki_pages"]}
|
|
vector_slugs = set()
|
|
for raw in archive["vectors"]:
|
|
slug = str(raw.get("wiki_slug") or raw.get("page_slug") or raw.get("slug") or "")
|
|
if slug not in slugs or _vector_blob_from_payload(raw) is None:
|
|
raise ValueError(f"向量数据或 Wiki 引用无效:{slug}")
|
|
if raw.get("embedding_fingerprint") != embedding.fingerprint or raw.get("embedding_dimensions") != embedding.dimensions:
|
|
raise ValueError("数据包编码模型指纹与当前配置不一致;请使用同一模型配置,禁止静默重新编码")
|
|
vector_slugs.add(slug)
|
|
if not vector_slugs:
|
|
raise ValueError("数据包中没有 Wiki 向量")
|
|
for page in archive["wiki_pages"]:
|
|
if page["slug"] not in vector_slugs and page.get("page_type") not in {"index", "directory"} and (page.get("content") or "").strip():
|
|
raise ValueError(f"Wiki 缺少向量:{page['slug']}")
|
|
leased = await store.try_acquire_sync_lease(mapping["id"], owner=owner, sync_id=job_id, lease_seconds=900, fingerprint=embedding.fingerprint)
|
|
if not leased:
|
|
raise ValueError("目标知识库正在同步,请稍后重试")
|
|
client = build_weknora_client(get_resolved_llmwiki_runtime(app.state.config))
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="running", phase="restoring")
|
|
folders = {(): ""}
|
|
for raw in archive["wiki_pages"]:
|
|
if not await store.renew_sync_lease(mapping["id"], owner=owner, lease_seconds=900):
|
|
raise RuntimeError("导入租约已失效,请重新发起")
|
|
# Restore the provider's Wiki schema, never raw documents; raw
|
|
# uploads would trigger generation and embedding a second time.
|
|
categories = raw.get("category_path") or []
|
|
if isinstance(categories, str):
|
|
categories = [part for part in categories.split("/") if part]
|
|
parent = ()
|
|
for name in categories:
|
|
path_key = (*parent, name)
|
|
if path_key not in folders:
|
|
folder = await client.create_wiki_folder(mapping["weknora_id"], name=name, parent_id=folders[parent])
|
|
folders[path_key] = folder["id"]
|
|
parent = path_key
|
|
payload = {key: value for key, value in raw.items() if key not in {"id", "tenant_id", "knowledge_base_id", "folder_id", "origins"}}
|
|
payload["folder_id"] = folders[parent]
|
|
if raw.get("tags"):
|
|
payload["page_metadata"] = {**(raw.get("page_metadata") or {}), "tags": raw["tags"]}
|
|
remote = await client.create_wiki_page(mapping["weknora_id"], payload)
|
|
page = normalize_wiki_page(mapping["id"], {**raw, **remote})
|
|
vectors = []
|
|
for value in archive.vectors_for(raw["slug"]):
|
|
blob, dimensions = _vector_blob_from_payload(value)
|
|
vectors.append({**value, "vector_blob": blob, "embedding_dimensions": dimensions})
|
|
await store.replace_page_vectors(page=page, vectors=vectors, fingerprint=embedding.fingerprint, sync_id=job_id)
|
|
counts["wiki_page"] += 1
|
|
counts["vector"] += len(vectors)
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="running", phase="restoring", counts=counts)
|
|
await client.rebuild_wiki_links(mapping["weknora_id"])
|
|
await store.finish_sync(mapping["id"], owner=owner, sync_id=job_id, success=True, remote_page_count=counts["wiki_page"])
|
|
leased = False
|
|
if app.state.llmwiki_vector_cache:
|
|
app.state.llmwiki_vector_cache.invalidate(mapping["id"])
|
|
from app.gateway.knowledge_transfer import sync_to_global
|
|
|
|
await sync_to_global(app, mapping)
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="completed", phase="completed", counts=counts)
|
|
except Exception as exc:
|
|
if leased:
|
|
await store.finish_sync(mapping["id"], owner=owner, sync_id=job_id, success=False, remote_page_count=counts["wiki_page"], error=str(exc))
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="failed", phase="failed", counts=counts, error=str(exc))
|
|
finally:
|
|
if archive is not None:
|
|
await asyncio.to_thread(archive.close)
|
|
|
|
|
|
@router.post("/api/llmwiki/knowledge-bases/{mapping_id}/import-package", status_code=202)
|
|
async def upload_restore(request: Request, mapping_id: str, file: UploadFile = File(...)):
|
|
actor = await _actor(request, admin=True)
|
|
mapping = await request.app.state.llmwiki_store.get_authorized(mapping_id, actor or "system", write=True, is_admin=True)
|
|
if not mapping or mapping.get("is_conversation_deposit") or mapping.get("owner_user_id") == "system":
|
|
raise HTTPException(404, "目标普通知识库不存在或不可导入")
|
|
client = build_weknora_client(get_resolved_llmwiki_runtime(request.app.state.config))
|
|
remote = await client.get_knowledge_base(mapping["weknora_id"])
|
|
if not (remote.get("indexing_strategy") or {}).get("wiki_enabled"):
|
|
raise HTTPException(422, "请先开启目标知识库的 Wiki")
|
|
existing = await client.list_wiki_pages(mapping["weknora_id"], page_size=1)
|
|
if existing.get("total"):
|
|
raise HTTPException(409, "请导入到新建的空 Wiki 知识库,避免覆盖已有页面")
|
|
job_id = str(uuid4())
|
|
scope = f"restore:{mapping_id}"
|
|
row = _write_export_manifest(job_id=job_id, mapping_id=scope, status="queued", phase="uploading")
|
|
path = package_path(_package_root(), job_id)
|
|
try:
|
|
with path.open("wb") as handle:
|
|
while data := await file.read(1024 * 1024):
|
|
await asyncio.to_thread(handle.write, data)
|
|
except Exception as exc:
|
|
_write_export_manifest(job_id=job_id, mapping_id=scope, status="failed", phase="upload_failed", error=str(exc))
|
|
raise
|
|
finally:
|
|
await file.close()
|
|
spawn_local_wiki_index_job(request.app, restore_package(request.app, mapping, job_id, path))
|
|
return job_view(row)
|
|
|
|
|
|
@router.get("/api/llmwiki/knowledge-bases/{mapping_id}/import-package")
|
|
async def list_restores(request: Request, mapping_id: str):
|
|
await _actor(request, admin=True)
|
|
return {"jobs": [job_view(row) for row in await asyncio.to_thread(_list_export_manifests, f"restore:{mapping_id}", limit=100)]}
|