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

192 lines
9.4 KiB
Python

"""Local document → Wiki → vectors, independent of the WeKnora service."""
from __future__ import annotations
import asyncio
import hashlib
import json
from pathlib import Path
from typing import Any
from deerflow.config.runtime_paths import runtime_home
from deerflow.knowledge.llm_extractor import extract_document_llm_multi
from deerflow.utils.file_conversion import CONVERTIBLE_EXTENSIONS, convert_file_to_markdown
ALLOWED_EXTENSIONS = CONVERTIBLE_EXTENSIONS | {".md", ".markdown", ".txt", ".csv"}
MAX_FILE_BYTES = 100 * 1024 * 1024
MAX_TEXT_CHARS = 4_000_000
_tasks: dict[str, asyncio.Task] = {}
_slots = asyncio.Semaphore(2)
_queue_lock = asyncio.Lock()
def file_job_dir(job_id: str) -> Path:
from uuid import UUID
return runtime_home() / "assistant-knowledge" / "files" / str(UUID(job_id))
def save_json(path: Path, value: Any) -> None:
temporary = path.with_suffix(".tmp")
temporary.write_text(json.dumps(value, ensure_ascii=False), encoding="utf-8")
temporary.replace(path)
def read_json(path: Path) -> Any:
return json.loads(path.read_text(encoding="utf-8"))
def split_text(text: str, size: int):
"""Bound every extraction/embedding input without truncating the document."""
for offset in range(0, len(text), size):
yield text[offset : offset + size]
def notes_package(notes: list[dict], source: dict) -> dict:
from deerflow.persistence.assistant_knowledge.store import _entity_alignment_key
merged: dict[str, dict] = {}
for raw in notes:
key = _entity_alignment_key(raw["title"])
if key not in merged:
merged[key] = dict(raw)
continue
note = merged[key]
note["content_md"] = "\n\n".join(dict.fromkeys([note["content_md"], raw["content_md"]]))
for field in ("tags", "entities", "relations"):
note[field] = [*note.get(field, []), *raw.get(field, [])]
note["tags"] = list(dict.fromkeys(note["tags"]))
pages, entities, relations = [], [], []
for index, note in enumerate(merged.values()):
slug = f"files/{source['digest']}/{index}"
pages.append(
{
"slug": slug,
"title": note["title"],
"content": note["content_md"],
"page_type": "article",
"tags": note.get("tags", []),
"category_path": [note.get("category") or "知识文档"],
"source_refs": [{"type": "file", "title": source["filename"]}],
}
)
for entity in note.get("entities", []):
entities.append({"id": entity["name"], "name": entity["name"], "entity_type": entity.get("type", "concept"), "aliases": entity.get("aliases", []), "source_wiki_slug": slug})
for relation in note.get("relations", []):
relations.append({"source": relation["from"], "target": relation["to"], "predicate": relation.get("type", "related_to"), "source_wiki_slug": slug})
# Persist real sibling links, understood by the shared Wiki reader.
for index, page in enumerate(pages):
others = [p for p in pages[max(0, index - 4) : index + 5] if p["slug"] != page["slug"]]
page["out_links"] = [p["slug"] for p in others]
page["in_links"] = list(page["out_links"])
if others:
page["content"] += "\n\n## 相关知识\n" + "\n".join(f"- [[{p['slug']}|{p['title']}]]" for p in others)
return {"source": source, "wiki_pages": pages, "entities": entities, "relations": relations, "vectors": []}
async def run_file_job(app, job_id: str) -> None:
store = app.state.assistant_knowledge_store
root = file_job_dir(job_id)
try:
async with _slots:
manifest = await asyncio.to_thread(read_json, root / "manifest.json")
await store.update_import_job(job_id, status="running", phase="parsing")
package_file = root / "wiki.json"
if package_file.exists():
package = await asyncio.to_thread(read_json, package_file)
else:
text_file = root / "parsed.md"
if not text_file.exists():
original = root / manifest["stored_name"]
if original.suffix in CONVERTIBLE_EXTENSIONS:
converted = await asyncio.to_thread(lambda: asyncio.run(convert_file_to_markdown(original)))
if converted is None:
raise ValueError("无法解析文档,请检查文件是否损坏或是否需要 OCR")
text = await asyncio.to_thread(converted.read_text, encoding="utf-8")
else:
raw = await asyncio.to_thread(original.read_bytes)
try:
text = raw.decode("utf-8-sig")
except UnicodeDecodeError:
text = raw.decode("gb18030")
if not text.strip() or len(text) > MAX_TEXT_CHARS:
raise ValueError("文件没有可提取文本,或解析正文超过 400 万字符,请拆分文件")
await asyncio.to_thread(text_file.write_text, text, encoding="utf-8")
else:
text = await asyncio.to_thread(text_file.read_text, encoding="utf-8")
notes = []
parts = list(split_text(text, 10000))
for index, part in enumerate(parts):
await store.update_import_job(job_id, phase="generating_wiki", counts={"parts_total": len(parts), "parts_done": index})
cache = root / f"notes-{index}.json"
if cache.exists():
drafts = await asyncio.to_thread(read_json, cache)
else:
from dataclasses import asdict
extracted = await asyncio.wait_for(
extract_document_llm_multi(
part,
doc_id=f"{job_id}:{index}",
doc_title=manifest["filename"],
max_input_chars=12000,
max_notes=6,
),
timeout=300,
)
if not extracted:
raise ValueError("Wiki 生成失败,请在设置中检查可用的语言模型,然后重试")
drafts = [asdict(note) for note in extracted]
await asyncio.to_thread(save_json, cache, drafts)
notes.extend(drafts)
package = notes_package(notes, {"type": "file", "filename": manifest["filename"], "digest": manifest["digest"]})
await asyncio.to_thread(save_json, package_file, package)
embedding = getattr(app.state, "llmwiki_embedding", None)
if embedding is None:
raise ValueError("Wiki 已生成;请管理员在设置中配置本地编码器或编码服务,然后重试向量化")
vectors = []
sections = [(p, i, section) for p in package["wiki_pages"] for i, section in enumerate(split_text(p["content"], 1600))]
for index, (page, section_index, content) in enumerate(sections):
await store.update_import_job(job_id, phase="vectorizing", counts={"wiki_pages": len(package["wiki_pages"]), "vectors_total": len(sections), "vectors": index})
digest = hashlib.sha256(content.encode()).hexdigest()
cache = root / f"vector-{digest}.json"
record = await asyncio.to_thread(read_json, cache) if cache.exists() else None
if record is None or record["embedding_fingerprint"] != embedding.fingerprint:
values = await embedding.embed_texts([content])
record = {"vector": values[0], "embedding_model": embedding.config.model, "embedding_fingerprint": embedding.fingerprint, "content_hash": digest}
await asyncio.to_thread(save_json, cache, record)
vectors.append({**record, "wiki_slug": page["slug"], "section_index": section_index, "section_content": content})
package["vectors"] = vectors
await store.update_import_job(job_id, phase="importing")
await store.import_package_for_job(job_id, package=package, created_by=manifest["actor"])
except asyncio.CancelledError:
await store.update_import_job(job_id, status="failed", phase="interrupted", error="服务停止,已保留处理进度,可重试", completed=True)
raise
except Exception as exc:
await store.update_import_job(job_id, status="failed", phase="failed", error=str(exc) or type(exc).__name__, completed=True)
def start_file_job(app, job_id: str) -> bool:
if job_id in _tasks:
return False
task = asyncio.create_task(run_file_job(app, job_id))
_tasks[job_id] = task
task.add_done_callback(lambda _: _tasks.pop(job_id, None))
return True
async def stop_file_jobs() -> None:
tasks = list(_tasks.values())
for task in tasks:
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
async def retry_file_job(app, job_id: str) -> bool:
async with _queue_lock:
if job_id in _tasks:
return False
await app.state.assistant_knowledge_store.update_import_job(job_id, status="queued", phase="queued")
return start_file_job(app, job_id)