feat(avatar): show multi-file knowledge upload progress
This commit is contained in:
@@ -54,6 +54,8 @@ def init_db():
|
||||
("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"),
|
||||
("knowledge_docs", "vectorized_at", "TIMESTAMP"),
|
||||
("knowledge_docs", "error_message", "VARCHAR DEFAULT ''"),
|
||||
("knowledge_docs", "index_stage", "VARCHAR DEFAULT ''"),
|
||||
("knowledge_docs", "index_progress", "INTEGER DEFAULT 0"),
|
||||
("avatars", "owner_id", "VARCHAR DEFAULT ''"),
|
||||
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
|
||||
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
|
||||
|
||||
@@ -51,7 +51,7 @@ def _hash_embedding(texts, dim=EMBED_DIM):
|
||||
return vecs
|
||||
|
||||
|
||||
def embed(texts):
|
||||
def embed(texts, on_progress=None):
|
||||
"""返回 list[list[float]],与输入顺序一致。"""
|
||||
if not texts:
|
||||
return []
|
||||
@@ -64,6 +64,7 @@ def embed(texts):
|
||||
except ValueError:
|
||||
batch_size = 10
|
||||
embeddings = []
|
||||
total = len(texts)
|
||||
for start in range(0, len(texts), batch_size):
|
||||
batch = texts[start:start + batch_size]
|
||||
payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
|
||||
@@ -84,8 +85,13 @@ def embed(texts):
|
||||
if len(items) != len(batch):
|
||||
raise ValueError("embedding response count does not match request")
|
||||
embeddings.extend(item["embedding"] for item in items)
|
||||
if on_progress:
|
||||
on_progress(len(embeddings), total)
|
||||
return embeddings
|
||||
return _hash_embedding(texts)
|
||||
vectors = _hash_embedding(texts)
|
||||
if on_progress:
|
||||
on_progress(len(vectors), len(texts))
|
||||
return vectors
|
||||
|
||||
|
||||
def cosine(a, b):
|
||||
|
||||
@@ -191,6 +191,8 @@ class KnowledgeDoc(Base):
|
||||
file_url = Column(String, default="")
|
||||
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
|
||||
error_message = Column(String, default="") # 建立索引失败原因
|
||||
index_stage = Column(String, default="") # queued | extracting | chunking | embedding | ready | failed
|
||||
index_progress = Column(Integer, default=0) # 0-100
|
||||
vectorized = Column(Boolean, default=False) # 是否已向量化
|
||||
embedding_model = Column(String, default="") # 向量模型标识
|
||||
chunk_count = Column(Integer, default=0) # 切片数量
|
||||
@@ -207,6 +209,8 @@ class KnowledgeDoc(Base):
|
||||
"fileUrl": self.file_url,
|
||||
"status": self.status,
|
||||
"errorMessage": self.error_message or "",
|
||||
"indexStage": self.index_stage or "",
|
||||
"indexProgress": int(self.index_progress or 0),
|
||||
"vectorized": bool(self.vectorized),
|
||||
"embeddingModel": self.embedding_model,
|
||||
"chunkCount": self.chunk_count,
|
||||
|
||||
@@ -102,6 +102,8 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
|
||||
file_size=file_size,
|
||||
file_url=f"/api/files/{avatar_id}/{stored}",
|
||||
status="parsing",
|
||||
index_stage="queued",
|
||||
index_progress=0,
|
||||
)
|
||||
|
||||
# Persist and acknowledge the upload first. Extraction and embeddings may take
|
||||
@@ -134,6 +136,8 @@ def retry_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db
|
||||
doc.chunk_count = 0
|
||||
doc.vectorized_at = None
|
||||
doc.error_message = ""
|
||||
doc.index_stage = "queued"
|
||||
doc.index_progress = 0
|
||||
db.commit()
|
||||
db.refresh(doc)
|
||||
knowledge_vectorizer.enqueue(doc.id)
|
||||
|
||||
@@ -74,11 +74,19 @@ class KnowledgeVectorizer:
|
||||
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("文档没有可建立索引的文字内容")
|
||||
vectors = embeddings.embed(chunks)
|
||||
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("向量服务返回数量与文档分段不一致")
|
||||
|
||||
@@ -103,6 +111,8 @@ class KnowledgeVectorizer:
|
||||
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:
|
||||
@@ -116,10 +126,18 @@ class KnowledgeVectorizer:
|
||||
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()
|
||||
|
||||
@@ -49,6 +49,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
||||
texts = [f"chunk-{index}" for index in range(14)]
|
||||
batch_sizes = []
|
||||
requested_urls = []
|
||||
progress_updates = []
|
||||
|
||||
def fake_urlopen(request, timeout):
|
||||
self.assertEqual(timeout, 30)
|
||||
@@ -68,7 +69,10 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
||||
"EMBEDDING_MODEL": "text-embedding-v4",
|
||||
"EMBEDDING_BATCH_SIZE": "10",
|
||||
}), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen):
|
||||
result = embeddings.embed(texts)
|
||||
result = embeddings.embed(
|
||||
texts,
|
||||
on_progress=lambda completed, total: progress_updates.append((completed, total)),
|
||||
)
|
||||
|
||||
self.assertEqual(batch_sizes, [10, 4])
|
||||
self.assertEqual(requested_urls, [
|
||||
@@ -76,6 +80,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
|
||||
"https://embedding.example/v1/embeddings",
|
||||
])
|
||||
self.assertEqual(result, [[float(index)] for index in range(14)])
|
||||
self.assertEqual(progress_updates, [(10, 14), (14, 14)])
|
||||
|
||||
def test_full_embedding_endpoint_is_not_modified(self):
|
||||
self.assertEqual(
|
||||
|
||||
@@ -116,6 +116,8 @@ def test_background_vectorizer_commits_ready_document_and_chunks_together(
|
||||
assert stored.status == "ready"
|
||||
assert stored.vectorized is True
|
||||
assert stored.chunk_count == 1
|
||||
assert stored.index_stage == "ready"
|
||||
assert stored.index_progress == 100
|
||||
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 1
|
||||
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
||||
db.delete(stored)
|
||||
|
||||
Reference in New Issue
Block a user