192 lines
9.4 KiB
Python
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)
|