feat(avatar): optimize embedded H5 management flow
This commit is contained in:
@@ -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) # 切片数量
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user