diff --git a/digital-avatar-app/backend/models.py b/digital-avatar-app/backend/models.py index 5376835..345c454 100644 --- a/digital-avatar-app/backend/models.py +++ b/digital-avatar-app/backend/models.py @@ -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) # 切片数量 diff --git a/digital-avatar-app/backend/routers/knowledge.py b/digital-avatar-app/backend/routers/knowledge.py index 58a9eec..d336a2a 100644 --- a/digital-avatar-app/backend/routers/knowledge.py +++ b/digital-avatar-app/backend/routers/knowledge.py @@ -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)) diff --git a/digital-avatar-app/backend/tests/test_knowledge_storage.py b/digital-avatar-app/backend/tests/test_knowledge_storage.py index 89c6f3d..8e89fde 100644 --- a/digital-avatar-app/backend/tests/test_knowledge_storage.py +++ b/digital-avatar-app/backend/tests/test_knowledge_storage.py @@ -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() diff --git a/digital-avatar-app/src/store/user.ts b/digital-avatar-app/src/store/user.ts index 6bf56d5..3aa4ac0 100644 --- a/digital-avatar-app/src/store/user.ts +++ b/digital-avatar-app/src/store/user.ts @@ -10,6 +10,7 @@ import { type SmsLoginResult, type UserProfile } from '@/api' +import { clearHuihuiEmbeddedMode, markHuihuiEmbeddedMode } from '@/utils/embed-mode' const TOKEN_KEY = 'hh_app_token' const USER_KEY = 'hh_app_user' @@ -70,16 +71,23 @@ export const useUserStore = defineStore('smsuser', () => { // 短信登录 const login = async (phone: string, code: string) => { - return acceptLogin(await loginBySms(phone, code)) + const result = await loginBySms(phone, code) + clearHuihuiEmbeddedMode() + return acceptLogin(result) } // 账号密码登录 const loginByPwd = async (account: string, password: string) => { - return acceptLogin(await loginByPassword(account, password)) + const result = await loginByPassword(account, password) + clearHuihuiEmbeddedMode() + return acceptLogin(result) } - const loginByToken = async (huihuiToken: string) => - acceptLogin(await loginByHuihuiToken(huihuiToken)) + const loginByToken = async (huihuiToken: string) => { + const result = await loginByHuihuiToken(huihuiToken) + markHuihuiEmbeddedMode() + return acceptLogin(result) + } // 退出 const logout = async () => { @@ -88,6 +96,7 @@ export const useUserStore = defineStore('smsuser', () => { } catch { /* 忽略网络错误,本地清除即可 */ } + clearHuihuiEmbeddedMode() clearSession() } diff --git a/digital-avatar-app/src/utils/embed-mode.ts b/digital-avatar-app/src/utils/embed-mode.ts new file mode 100644 index 0000000..3975308 --- /dev/null +++ b/digital-avatar-app/src/utils/embed-mode.ts @@ -0,0 +1,13 @@ +const HUIHUI_EMBED_MODE_KEY = 'hh_huihui_embed_mode' + +export function markHuihuiEmbeddedMode(): void { + sessionStorage.setItem(HUIHUI_EMBED_MODE_KEY, '1') +} + +export function clearHuihuiEmbeddedMode(): void { + sessionStorage.removeItem(HUIHUI_EMBED_MODE_KEY) +} + +export function isHuihuiEmbeddedMode(): boolean { + return sessionStorage.getItem(HUIHUI_EMBED_MODE_KEY) === '1' +} diff --git a/digital-avatar-app/src/views/AuthorizationManage.vue b/digital-avatar-app/src/views/AuthorizationManage.vue index a6510a7..0046f18 100644 --- a/digital-avatar-app/src/views/AuthorizationManage.vue +++ b/digital-avatar-app/src/views/AuthorizationManage.vue @@ -1,6 +1,6 @@