144 lines
5.3 KiB
Python
144 lines
5.3 KiB
Python
"""Durable, serial knowledge-document indexing for the avatar knowledge base."""
|
|
|
|
import json
|
|
import logging
|
|
import os
|
|
import queue
|
|
import threading
|
|
from datetime import datetime, timezone
|
|
|
|
from database import SessionLocal
|
|
from models import KnowledgeChunk, KnowledgeDoc
|
|
import embeddings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
UPLOAD_DIR = os.path.abspath(
|
|
os.getenv("UPLOAD_DIR", os.path.join(BACKEND_DIR, "routers", "uploads"))
|
|
)
|
|
|
|
|
|
class KnowledgeVectorizer:
|
|
"""Indexes one document at a time so slow providers cannot block uploads."""
|
|
|
|
def __init__(self):
|
|
self._queue: queue.Queue[str] = queue.Queue()
|
|
self._queued: set[str] = set()
|
|
self._lock = threading.Lock()
|
|
self._thread: threading.Thread | None = None
|
|
|
|
def start(self):
|
|
if self._thread and self._thread.is_alive():
|
|
return
|
|
self._thread = threading.Thread(
|
|
target=self._run, name="knowledge-vectorizer", daemon=True
|
|
)
|
|
self._thread.start()
|
|
db = SessionLocal()
|
|
try:
|
|
# A process restart must not abandon documents already accepted by upload.
|
|
for (doc_id,) in db.query(KnowledgeDoc.id).filter(KnowledgeDoc.status == "parsing"):
|
|
self.enqueue(doc_id)
|
|
finally:
|
|
db.close()
|
|
|
|
def enqueue(self, doc_id: str):
|
|
with self._lock:
|
|
if doc_id in self._queued:
|
|
return
|
|
self._queued.add(doc_id)
|
|
self._queue.put(doc_id)
|
|
|
|
def _run(self):
|
|
while True:
|
|
doc_id = self._queue.get()
|
|
try:
|
|
self.vectorize_document(doc_id)
|
|
except Exception:
|
|
logger.exception("Unexpected knowledge vectorizer failure for %s", doc_id)
|
|
finally:
|
|
with self._lock:
|
|
self._queued.discard(doc_id)
|
|
self._queue.task_done()
|
|
|
|
def vectorize_document(self, doc_id: str):
|
|
db = SessionLocal()
|
|
try:
|
|
doc = db.get(KnowledgeDoc, doc_id)
|
|
if not doc or doc.status != "parsing":
|
|
return
|
|
|
|
stored_name = os.path.basename(doc.file_url or "")
|
|
path = os.path.join(UPLOAD_DIR, doc.avatar_id, stored_name)
|
|
if not stored_name or not os.path.isfile(path):
|
|
raise FileNotFoundError("原文件不可用,请重新上传")
|
|
|
|
self._set_progress(db, doc, "extracting", 8)
|
|
text = embeddings.extract_text(path, f".{doc.file_type}")
|
|
self._set_progress(db, doc, "chunking", 22)
|
|
chunks = embeddings.chunk_text(text)
|
|
if not chunks:
|
|
raise ValueError("文档没有可建立索引的文字内容")
|
|
self._set_progress(db, doc, "embedding", 30)
|
|
|
|
def embedding_progress(done: int, total: int):
|
|
percent = 30 + int((done / max(1, total)) * 65)
|
|
self._set_progress(db, doc, "embedding", min(percent, 95))
|
|
|
|
vectors = embeddings.embed(chunks, on_progress=embedding_progress)
|
|
if len(vectors) != len(chunks):
|
|
raise ValueError("向量服务返回数量与文档分段不一致")
|
|
|
|
# Commit the document and every chunk together. Chat only sees complete indexes.
|
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
|
|
db.add_all(
|
|
[
|
|
KnowledgeChunk(
|
|
doc_id=doc.id,
|
|
avatar_id=doc.avatar_id,
|
|
content=chunk,
|
|
vector=json.dumps(vector),
|
|
chunk_index=index,
|
|
embedding_model=embeddings.MODEL,
|
|
)
|
|
for index, (chunk, vector) in enumerate(zip(chunks, vectors))
|
|
]
|
|
)
|
|
doc.vectorized = True
|
|
doc.embedding_model = embeddings.MODEL
|
|
doc.chunk_count = len(chunks)
|
|
doc.vectorized_at = datetime.now(timezone.utc)
|
|
doc.status = "ready"
|
|
doc.error_message = ""
|
|
doc.index_stage = "ready"
|
|
doc.index_progress = 100
|
|
db.commit()
|
|
logger.info("Knowledge document %s indexed with %s chunks", doc.id, len(chunks))
|
|
except Exception as exc:
|
|
db.rollback()
|
|
failed_doc = db.get(KnowledgeDoc, doc_id)
|
|
if failed_doc:
|
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == failed_doc.id).delete()
|
|
failed_doc.status = "failed"
|
|
failed_doc.vectorized = False
|
|
failed_doc.embedding_model = ""
|
|
failed_doc.chunk_count = 0
|
|
failed_doc.vectorized_at = None
|
|
failed_doc.error_message = str(exc)[:300] or "建立知识索引失败"
|
|
failed_doc.index_stage = "failed"
|
|
failed_doc.index_progress = 0
|
|
db.commit()
|
|
logger.exception("Knowledge vectorization failed for %s: %s", doc_id, exc)
|
|
finally:
|
|
db.close()
|
|
|
|
@staticmethod
|
|
def _set_progress(db, doc, stage: str, progress: int):
|
|
doc.index_stage = stage
|
|
doc.index_progress = progress
|
|
db.commit()
|
|
|
|
|
|
knowledge_vectorizer = KnowledgeVectorizer()
|