feat(avatar): optimize embedded H5 management flow

This commit is contained in:
stefanfeng
2026-08-27 09:28:05 +08:00
parent 6e3fe5a616
commit ef58c5f2d2
10 changed files with 307 additions and 63 deletions
+1 -1
View File
@@ -188,7 +188,7 @@ class KnowledgeDoc(Base):
file_type = Column(String, default="") # pdf | doc | docx | xlsx
file_size = Column(Integer, default=0)
file_url = Column(String, default="")
status = Column(String, default="uploaded") # uploaded | parsing | ready
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
vectorized = Column(Boolean, default=False) # 是否已向量化
embedding_model = Column(String, default="") # 向量模型标识
chunk_count = Column(Integer, default=0) # 切片数量
+44 -24
View File
@@ -1,5 +1,6 @@
import os
import json
import logging
import uuid
from datetime import datetime, timezone
@@ -13,6 +14,7 @@ from responses import ok, fail
import embeddings
router = APIRouter()
logger = logging.getLogger(__name__)
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
@@ -69,6 +71,15 @@ def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = D
.order_by(KnowledgeDoc.created_at.desc())
.all()
)
# Older synchronous uploads could be interrupted after persisting "parsing".
# New uploads are committed only after indexing finishes, so these rows are stale.
stale_docs = [doc for doc in docs if doc.status == "parsing"]
if stale_docs:
for doc in stale_docs:
doc.status = "failed"
doc.vectorized = False
doc.chunk_count = 0
db.commit()
return ok([_doc_payload(d) for d in docs])
@@ -88,6 +99,7 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
with open(path, "wb") as f:
f.write(content)
doc = KnowledgeDoc(
id=uuid.uuid4().hex,
avatar_id=avatar_id,
filename=file.filename,
file_type=ext.lstrip("."),
@@ -95,39 +107,47 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
file_url=f"/api/files/{avatar_id}/{stored}",
status="parsing",
)
db.add(doc)
db.commit()
db.refresh(doc)
# 向量化:抽取文本 -> 分块 -> 调第三方/本地嵌入 -> 存切片
# Complete extraction and embedding before the first database commit so a
# process restart cannot leave a permanent "parsing" row behind.
try:
text = embeddings.extract_text(path, ext)
chunks = embeddings.chunk_text(text)
if chunks:
vectors = embeddings.embed(chunks)
for i, (c, v) in enumerate(zip(chunks, vectors)):
db.add(
KnowledgeChunk(
doc_id=doc.id,
avatar_id=avatar_id,
content=c,
vector=json.dumps(v),
chunk_index=i,
embedding_model=embeddings.MODEL,
)
)
doc.vectorized = True
doc.embedding_model = embeddings.MODEL
doc.chunk_count = len(chunks)
doc.vectorized_at = datetime.now(timezone.utc)
if not chunks:
raise ValueError("文档没有可建立索引的文字内容")
vectors = embeddings.embed(chunks)
if len(vectors) != len(chunks):
raise ValueError("向量服务返回数量与文档分段不一致")
doc.vectorized = True
doc.embedding_model = embeddings.MODEL
doc.chunk_count = len(chunks)
doc.vectorized_at = datetime.now(timezone.utc)
doc.status = "ready"
db.add(doc)
for i, (chunk, vector) in enumerate(zip(chunks, vectors)):
db.add(
KnowledgeChunk(
doc_id=doc.id,
avatar_id=avatar_id,
content=chunk,
vector=json.dumps(vector),
chunk_index=i,
embedding_model=embeddings.MODEL,
)
)
db.commit()
db.refresh(doc)
except Exception as e:
print("vectorize failed:", e)
doc.status = "ready" # 上传成功但向量化失败,仍可展示
except Exception as exc:
db.rollback()
doc.status = "failed"
doc.vectorized = False
doc.embedding_model = ""
doc.chunk_count = 0
doc.vectorized_at = None
db.add(doc)
db.commit()
db.refresh(doc)
logger.exception("knowledge vectorization failed for %s: %s", doc.id, exc)
return ok(_doc_payload(doc))
@@ -2,9 +2,17 @@ from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch
from fastapi.testclient import TestClient
from database import SessionLocal
from main import app
from models import KnowledgeChunk, KnowledgeDoc
from routers.knowledge import _doc_payload
client = TestClient(app)
def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
avatar_id = "avatar-1"
stored_name = "knowledge.md"
@@ -21,3 +29,66 @@ def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
assert _doc_payload(doc)["filePresent"] is False
stored_file.write_text("knowledge", encoding="utf-8")
assert _doc_payload(doc)["filePresent"] is True
def test_upload_marks_vectorization_failure_instead_of_staying_processing(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.embeddings.embed", side_effect=RuntimeError("provider unavailable")),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
headers=context["owner_headers"],
files={"file": ("knowledge.md", b"# Knowledge\n\nTest content", "text/markdown")},
)
payload = response.json()["data"]
assert payload["status"] == "failed"
assert payload["vectorized"] is False
assert payload["chunkCount"] == 0
db = SessionLocal()
try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "failed"
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
db.delete(stored)
db.commit()
finally:
db.close()
def test_markdown_upload_commits_ready_document_and_chunks_together(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.embeddings.embed", return_value=[[1.0, 0.0]]),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
headers=context["owner_headers"],
files={"file": ("knowledge.md", b"# Knowledge\n\nTest content", "text/markdown")},
)
payload = response.json()["data"]
assert payload["status"] == "ready"
assert payload["vectorized"] is True
assert payload["chunkCount"] == 1
db = SessionLocal()
try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "ready"
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)
db.commit()
finally:
db.close()