Compare commits

...
12 changed files with 480 additions and 102 deletions
+3
View File
@@ -53,6 +53,9 @@ def init_db():
("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"), ("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"),
("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"), ("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"),
("knowledge_docs", "vectorized_at", "TIMESTAMP"), ("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 ''"), ("avatars", "owner_id", "VARCHAR DEFAULT ''"),
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"), ("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"), ("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
+8 -2
View File
@@ -51,7 +51,7 @@ def _hash_embedding(texts, dim=EMBED_DIM):
return vecs return vecs
def embed(texts): def embed(texts, on_progress=None):
"""返回 list[list[float]],与输入顺序一致。""" """返回 list[list[float]],与输入顺序一致。"""
if not texts: if not texts:
return [] return []
@@ -64,6 +64,7 @@ def embed(texts):
except ValueError: except ValueError:
batch_size = 10 batch_size = 10
embeddings = [] embeddings = []
total = len(texts)
for start in range(0, len(texts), batch_size): for start in range(0, len(texts), batch_size):
batch = texts[start:start + batch_size] batch = texts[start:start + batch_size]
payload = json.dumps({"input": batch, "model": model}).encode("utf-8") payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
@@ -84,8 +85,13 @@ def embed(texts):
if len(items) != len(batch): if len(items) != len(batch):
raise ValueError("embedding response count does not match request") raise ValueError("embedding response count does not match request")
embeddings.extend(item["embedding"] for item in items) embeddings.extend(item["embedding"] for item in items)
if on_progress:
on_progress(len(embeddings), total)
return embeddings 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): def cosine(a, b):
+2
View File
@@ -20,6 +20,7 @@ import routers.chat
import routers.takeover import routers.takeover
from responses import ok from responses import ok
from services.chat_attachment_service import purge_expired_chat_attachments from services.chat_attachment_service import purge_expired_chat_attachments
from services.knowledge_vectorizer import knowledge_vectorizer
from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -131,6 +132,7 @@ def on_startup():
init_db() init_db()
seed() seed()
knowledge_vectorizer.start()
# Release stale resources when startup is invoked again by a reload/test. # Release stale resources when startup is invoked again by a reload/test.
stop_takeover_scheduler() stop_takeover_scheduler()
+6
View File
@@ -190,6 +190,9 @@ class KnowledgeDoc(Base):
file_size = Column(Integer, default=0) file_size = Column(Integer, default=0)
file_url = Column(String, default="") file_url = Column(String, default="")
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed 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) # 是否已向量化 vectorized = Column(Boolean, default=False) # 是否已向量化
embedding_model = Column(String, default="") # 向量模型标识 embedding_model = Column(String, default="") # 向量模型标识
chunk_count = Column(Integer, default=0) # 切片数量 chunk_count = Column(Integer, default=0) # 切片数量
@@ -205,6 +208,9 @@ class KnowledgeDoc(Base):
"fileSize": self.file_size, "fileSize": self.file_size,
"fileUrl": self.file_url, "fileUrl": self.file_url,
"status": self.status, "status": self.status,
"errorMessage": self.error_message or "",
"indexStage": self.index_stage or "",
"indexProgress": int(self.index_progress or 0),
"vectorized": bool(self.vectorized), "vectorized": bool(self.vectorized),
"embeddingModel": self.embedding_model, "embeddingModel": self.embedding_model,
"chunkCount": self.chunk_count, "chunkCount": self.chunk_count,
+53 -61
View File
@@ -1,8 +1,5 @@
import os import os
import json
import logging
import uuid import uuid
from datetime import datetime, timezone
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException
from pydantic import BaseModel from pydantic import BaseModel
@@ -12,16 +9,16 @@ from database import get_db
from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User
from responses import ok, fail from responses import ok, fail
import embeddings import embeddings
from services.knowledge_vectorizer import knowledge_vectorizer
router = APIRouter() router = APIRouter()
logger = logging.getLogger(__name__)
BASE_DIR = os.path.dirname(os.path.abspath(__file__)) BASE_DIR = os.path.dirname(os.path.abspath(__file__))
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads"))) UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
os.makedirs(UPLOAD_DIR, exist_ok=True) os.makedirs(UPLOAD_DIR, exist_ok=True)
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"} ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
MAX_UPLOAD_BYTES = 10 * 1024 * 1024 MAX_UPLOAD_BYTES = 50 * 1024 * 1024
UPLOAD_CHUNK_BYTES = 1024 * 1024
class QAIn(BaseModel): class QAIn(BaseModel):
@@ -71,15 +68,6 @@ def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = D
.order_by(KnowledgeDoc.created_at.desc()) .order_by(KnowledgeDoc.created_at.desc())
.all() .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]) return ok([_doc_payload(d) for d in docs])
@@ -93,65 +81,69 @@ async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization
os.makedirs(avatar_dir, exist_ok=True) os.makedirs(avatar_dir, exist_ok=True)
stored = f"{uuid.uuid4().hex}{ext}" stored = f"{uuid.uuid4().hex}{ext}"
path = os.path.join(avatar_dir, stored) path = os.path.join(avatar_dir, stored)
content = await file.read() file_size = 0
if len(content) > MAX_UPLOAD_BYTES: try:
return fail("文件不能超过 10MB", code=400) # Stream large files to disk so a 100MB upload does not occupy 100MB RAM.
with open(path, "wb") as f: with open(path, "wb") as f:
f.write(content) while chunk := await file.read(UPLOAD_CHUNK_BYTES):
file_size += len(chunk)
if file_size > MAX_UPLOAD_BYTES:
raise ValueError("文件不能超过 50MB")
f.write(chunk)
except ValueError as exc:
if os.path.exists(path):
os.remove(path)
return fail(str(exc), code=400)
doc = KnowledgeDoc( doc = KnowledgeDoc(
id=uuid.uuid4().hex, id=uuid.uuid4().hex,
avatar_id=avatar_id, avatar_id=avatar_id,
filename=file.filename, filename=file.filename,
file_type=ext.lstrip("."), file_type=ext.lstrip("."),
file_size=len(content), file_size=file_size,
file_url=f"/api/files/{avatar_id}/{stored}", file_url=f"/api/files/{avatar_id}/{stored}",
status="parsing", status="parsing",
index_stage="queued",
index_progress=0,
) )
# Complete extraction and embedding before the first database commit so a # Persist and acknowledge the upload first. Extraction and embeddings may take
# process restart cannot leave a permanent "parsing" row behind. # minutes for a PDF and must never consume the browser request timeout.
try: db.add(doc)
text = embeddings.extract_text(path, ext) db.commit()
chunks = embeddings.chunk_text(text) db.refresh(doc)
if not chunks: knowledge_vectorizer.enqueue(doc.id)
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 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)) return ok(_doc_payload(doc))
@router.post("/avatar/{avatar_id}/knowledge/docs/{doc_id}/retry")
def retry_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
doc = db.query(KnowledgeDoc).filter(
KnowledgeDoc.id == doc_id, KnowledgeDoc.avatar_id == avatar_id
).first()
if not doc:
return fail("文档不存在", code=404)
if doc.vectorized and doc.status == "ready":
return ok(_doc_payload(doc))
stored_name = os.path.basename(doc.file_url or "")
if not stored_name or not os.path.isfile(os.path.join(UPLOAD_DIR, avatar_id, stored_name)):
return fail("原文件不可用,请重新上传", code=400)
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
doc.status = "parsing"
doc.vectorized = False
doc.embedding_model = ""
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)
return ok(_doc_payload(doc))
@router.delete("/avatar/{avatar_id}/knowledge/docs/{doc_id}") @router.delete("/avatar/{avatar_id}/knowledge/docs/{doc_id}")
def delete_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)): def delete_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization) _require_owned_avatar(db, avatar_id, authorization)
@@ -0,0 +1,143 @@
"""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()
@@ -49,6 +49,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
texts = [f"chunk-{index}" for index in range(14)] texts = [f"chunk-{index}" for index in range(14)]
batch_sizes = [] batch_sizes = []
requested_urls = [] requested_urls = []
progress_updates = []
def fake_urlopen(request, timeout): def fake_urlopen(request, timeout):
self.assertEqual(timeout, 30) self.assertEqual(timeout, 30)
@@ -68,7 +69,10 @@ class RemoteEmbeddingTests(unittest.TestCase):
"EMBEDDING_MODEL": "text-embedding-v4", "EMBEDDING_MODEL": "text-embedding-v4",
"EMBEDDING_BATCH_SIZE": "10", "EMBEDDING_BATCH_SIZE": "10",
}), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen): }), 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(batch_sizes, [10, 4])
self.assertEqual(requested_urls, [ self.assertEqual(requested_urls, [
@@ -76,6 +80,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
"https://embedding.example/v1/embeddings", "https://embedding.example/v1/embeddings",
]) ])
self.assertEqual(result, [[float(index)] for index in range(14)]) 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): def test_full_embedding_endpoint_is_not_modified(self):
self.assertEqual( self.assertEqual(
@@ -8,6 +8,7 @@ from database import SessionLocal
from main import app from main import app
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair
from routers.knowledge import _doc_payload from routers.knowledge import _doc_payload
from services.knowledge_vectorizer import knowledge_vectorizer
client = TestClient(app) client = TestClient(app)
@@ -31,14 +32,14 @@ def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path):
assert _doc_payload(doc)["filePresent"] is True assert _doc_payload(doc)["filePresent"] is True
def test_upload_marks_vectorization_failure_instead_of_staying_processing( def test_upload_returns_before_background_vectorization(
tmp_path: Path, tmp_path: Path,
authorization_context, authorization_context,
): ):
context = authorization_context context = authorization_context
with ( with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)), patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.embeddings.embed", side_effect=RuntimeError("provider unavailable")), patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
): ):
response = client.post( response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs", f"/api/avatar/{context['avatar'].id}/knowledge/docs",
@@ -47,14 +48,15 @@ def test_upload_marks_vectorization_failure_instead_of_staying_processing(
) )
payload = response.json()["data"] payload = response.json()["data"]
assert payload["status"] == "failed" assert payload["status"] == "parsing"
assert payload["vectorized"] is False assert payload["vectorized"] is False
assert payload["chunkCount"] == 0 assert payload["chunkCount"] == 0
enqueue.assert_called_once_with(payload["id"])
db = SessionLocal() db = SessionLocal()
try: try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one() stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "failed" assert stored.status == "parsing"
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0 assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
db.delete(stored) db.delete(stored)
db.commit() db.commit()
@@ -62,14 +64,37 @@ def test_upload_marks_vectorization_failure_instead_of_staying_processing(
db.close() db.close()
def test_markdown_upload_commits_ready_document_and_chunks_together( def test_upload_rejects_oversize_file_before_queuing_indexing(
tmp_path: Path, tmp_path: Path,
authorization_context, authorization_context,
): ):
context = authorization_context context = authorization_context
with ( with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)), patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.embeddings.embed", return_value=[[1.0, 0.0]]), patch("routers.knowledge.MAX_UPLOAD_BYTES", 4),
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
headers=context["owner_headers"],
files={"file": ("oversize.md", b"12345", "text/markdown")},
)
payload = response.json()
assert payload["code"] == 400
assert payload["message"] == "文件不能超过 50MB"
enqueue.assert_not_called()
assert not list((tmp_path / context["avatar"].id).glob("*"))
def test_background_vectorizer_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.knowledge_vectorizer.enqueue"),
): ):
response = client.post( response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs", f"/api/avatar/{context['avatar'].id}/knowledge/docs",
@@ -78,14 +103,21 @@ def test_markdown_upload_commits_ready_document_and_chunks_together(
) )
payload = response.json()["data"] payload = response.json()["data"]
assert payload["status"] == "ready" assert payload["status"] == "parsing"
assert payload["vectorized"] is True with (
assert payload["chunkCount"] == 1 patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
patch("services.knowledge_vectorizer.embeddings.embed", return_value=[[1.0, 0.0]]),
):
knowledge_vectorizer.vectorize_document(payload["id"])
db = SessionLocal() db = SessionLocal()
try: try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one() stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "ready" 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 assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 1
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete() db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
db.delete(stored) db.delete(stored)
@@ -94,6 +126,87 @@ def test_markdown_upload_commits_ready_document_and_chunks_together(
db.close() db.close()
def test_background_vectorizer_keeps_failure_reason_for_retry(
tmp_path: Path,
authorization_context,
):
context = authorization_context
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.knowledge_vectorizer.enqueue"),
):
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"]
with (
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
patch("services.knowledge_vectorizer.embeddings.embed", side_effect=RuntimeError("provider unavailable")),
):
knowledge_vectorizer.vectorize_document(payload["id"])
db = SessionLocal()
try:
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
assert stored.status == "failed"
assert stored.error_message == "provider unavailable"
assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0
db.delete(stored)
db.commit()
finally:
db.close()
def test_retry_queues_a_failed_document_again(
tmp_path: Path,
authorization_context,
):
context = authorization_context
document_id = f"retry-doc-{context['suffix']}"
avatar_dir = tmp_path / context["avatar"].id
avatar_dir.mkdir()
(avatar_dir / "retry.md").write_text("retry content", encoding="utf-8")
db = SessionLocal()
try:
db.add(
KnowledgeDoc(
id=document_id,
avatar_id=context["avatar"].id,
filename="retry.md",
file_type="md",
file_url=f"/api/files/{context['avatar'].id}/retry.md",
status="failed",
error_message="provider unavailable",
)
)
db.commit()
finally:
db.close()
with (
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/knowledge/docs/{document_id}/retry",
headers=context["owner_headers"],
)
payload = response.json()["data"]
assert payload["status"] == "parsing"
assert payload["errorMessage"] == ""
enqueue.assert_called_once_with(document_id)
db = SessionLocal()
try:
db.query(KnowledgeDoc).filter(KnowledgeDoc.id == document_id).delete()
db.commit()
finally:
db.close()
def test_each_avatar_has_an_independent_document_and_qa_scope(authorization_context): def test_each_avatar_has_an_independent_document_and_qa_scope(authorization_context):
context = authorization_context context = authorization_context
first_avatar_id = context["avatar"].id first_avatar_id = context["avatar"].id
@@ -117,7 +117,7 @@ location /api/ {
proxy_set_header X-Forwarded-Proto $scheme; proxy_set_header X-Forwarded-Proto $scheme;
proxy_buffering off; proxy_buffering off;
proxy_read_timeout 300s; proxy_read_timeout 300s;
client_max_body_size 20m; client_max_body_size 100m;
} }
``` ```
+4
View File
@@ -23,6 +23,10 @@ http {
root /usr/share/nginx/html; root /usr/share/nginx/html;
index index.html; index index.html;
# Keep the application gateway aligned with the production edge gateway.
# Without this Nginx rejects ordinary PDF uploads with HTTP 413 before
# FastAPI can return its user-facing file-size validation message.
client_max_body_size 100m;
# SPA 兜底(hash 路由下深链接也可正常加载) # SPA 兜底(hash 路由下深链接也可正常加载)
location / { location / {
+15 -2
View File
@@ -305,6 +305,9 @@ export interface KnowledgeDoc {
vectorized?: boolean vectorized?: boolean
embeddingModel?: string embeddingModel?: string
chunkCount?: number chunkCount?: number
errorMessage?: string
indexStage?: string
indexProgress?: number
createdAt: string createdAt: string
} }
@@ -331,11 +334,18 @@ export const getKnowledgeDocs = (avatarId: string) =>
request.get<KnowledgeDoc[]>(`/avatar/${avatarId}/knowledge/docs`) request.get<KnowledgeDoc[]>(`/avatar/${avatarId}/knowledge/docs`)
// 上传文档(支持 md/txt/pdf/doc/docx/xlsx) // 上传文档(支持 md/txt/pdf/doc/docx/xlsx)
export const uploadKnowledgeDoc = (avatarId: string, file: File) => { export const uploadKnowledgeDoc = (
avatarId: string,
file: File,
onUploadProgress?: (loaded: number, total: number) => void
) => {
const form = new FormData() const form = new FormData()
form.append('file', file) form.append('file', file)
return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, { return request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs`, form, {
headers: { 'Content-Type': 'multipart/form-data' } headers: { 'Content-Type': 'multipart/form-data' },
// A slow mobile uplink must not be mistaken for a failed upload.
timeout: 10 * 60 * 1000,
onUploadProgress: (event) => onUploadProgress?.(event.loaded, event.total || file.size)
}) })
} }
@@ -343,6 +353,9 @@ export const uploadKnowledgeDoc = (avatarId: string, file: File) => {
export const deleteKnowledgeDoc = (avatarId: string, docId: string) => export const deleteKnowledgeDoc = (avatarId: string, docId: string) =>
request.delete(`/avatar/${avatarId}/knowledge/docs/${docId}`) request.delete(`/avatar/${avatarId}/knowledge/docs/${docId}`)
export const retryKnowledgeDoc = (avatarId: string, docId: string) =>
request.post<KnowledgeDoc>(`/avatar/${avatarId}/knowledge/docs/${docId}/retry`)
// 标准问答对列表 // 标准问答对列表
export const getQAPairs = (avatarId: string) => export const getQAPairs = (avatarId: string) =>
request.get<QAPair[]>(`/avatar/${avatarId}/knowledge/qa`) request.get<QAPair[]>(`/avatar/${avatarId}/knowledge/qa`)
+117 -26
View File
@@ -15,7 +15,7 @@
<template v-else> <template v-else>
<div class="tab-switcher" role="tablist" aria-label="知识库类型"> <div class="tab-switcher" role="tablist" aria-label="知识库类型">
<button class="tab-btn" :class="{ active: activeTab === 'docs' }" role="tab" :aria-selected="activeTab === 'docs'" @click="activeTab = 'docs'">文档知识库 <b>{{ docs.length }}</b></button> <button class="tab-btn" :class="{ active: activeTab === 'docs' }" role="tab" :aria-selected="activeTab === 'docs'" @click="activeTab = 'docs'">文档知识库 <b>{{ displayDocs.length }}</b></button>
<button class="tab-btn" :class="{ active: activeTab === 'qa' }" role="tab" :aria-selected="activeTab === 'qa'" @click="activeTab = 'qa'">标准问答对 <b>{{ qaPairs.length }}</b></button> <button class="tab-btn" :class="{ active: activeTab === 'qa' }" role="tab" :aria-selected="activeTab === 'qa'" @click="activeTab = 'qa'">标准问答对 <b>{{ qaPairs.length }}</b></button>
</div> </div>
@@ -25,14 +25,14 @@
<div class="upload-icon">📥</div> <div class="upload-icon">📥</div>
<p class="upload-title"><span class="upload-link">点击上传</span></p> <p class="upload-title"><span class="upload-link">点击上传</span></p>
<p class="upload-hint">支持 MD / TXT / PDF / DOC / DOCX / XLSX,上传后自动向量化</p> <p class="upload-hint">支持 MD / TXT / PDF / DOC / DOCX / XLSX,上传后自动向量化</p>
<input ref="fileInput" type="file" accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" /> <input ref="fileInput" type="file" multiple accept=".md,.txt,.pdf,.doc,.docx,.xlsx" class="hidden-input" @change="onFileChange" />
</div> </div>
<p v-if="uploading" class="uploading-text">上传并向量化中…</p> <p v-if="uploading" class="uploading-text">{{ pendingUploads.length }} 个文件正在上传</p>
<p v-if="uploadError" class="error-text">{{ uploadError }}</p> <p v-if="uploadError" class="error-text">{{ uploadError }}</p>
</div> </div>
<div v-if="docs.length" class="mobile-card-list"> <div v-if="displayDocs.length" class="mobile-card-list">
<article v-for="doc in docs" :key="doc.id" class="knowledge-card"> <article v-for="doc in displayDocs" :key="doc.id" class="knowledge-card">
<div class="card-icon">{{ fileEmoji(doc.fileType) }}</div> <div class="card-icon">{{ fileEmoji(doc.fileType) }}</div>
<div class="card-content"> <div class="card-content">
<div class="card-title-row"> <div class="card-title-row">
@@ -41,8 +41,14 @@
</div> </div>
<p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p> <p class="card-meta">{{ doc.fileType.toUpperCase() }} · {{ formatSize(doc.fileSize) }} · {{ formatDate(doc.createdAt) }}</p>
<p class="card-detail">{{ documentState(doc).detail }}</p> <p class="card-detail">{{ documentState(doc).detail }}</p>
<div v-if="documentState(doc).progress !== undefined" class="progress-track" :aria-label="`${documentState(doc).label} ${documentState(doc).progress}%`">
<span class="progress-fill" :style="{ width: `${documentState(doc).progress}%` }"></span>
</div>
</div>
<div class="card-actions">
<button v-if="documentState(doc).tone === 'failed'" class="card-retry" @click="retryDoc(doc.id)">重新索引</button>
<button v-if="!doc.localUploading" class="card-delete" @click="removeDoc(doc.id)">{{ doc.localOnly ? '移除' : '删除' }}</button>
</div> </div>
<button class="card-delete" @click="removeDoc(doc.id)">删除</button>
</article> </article>
</div> </div>
<div v-else class="card-empty">📂 暂无文档,先上传一个知识文件</div> <div v-else class="card-empty">📂 暂无文档,先上传一个知识文件</div>
@@ -78,7 +84,7 @@
</template> </template>
<script setup lang="ts"> <script setup lang="ts">
import { ref, onMounted, computed } from 'vue' import { ref, onMounted, onUnmounted, computed } from 'vue'
import { useRoute, useRouter } from 'vue-router' import { useRoute, useRouter } from 'vue-router'
import { useAvatarStore } from '@/store/avatar' import { useAvatarStore } from '@/store/avatar'
import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js' import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js'
@@ -87,6 +93,7 @@ import {
getKnowledgeDocs, getKnowledgeDocs,
uploadKnowledgeDoc, uploadKnowledgeDoc,
deleteKnowledgeDoc, deleteKnowledgeDoc,
retryKnowledgeDoc,
getQAPairs, getQAPairs,
deleteQAPair, deleteQAPair,
searchKnowledge, searchKnowledge,
@@ -102,18 +109,28 @@ const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, store.
const activeTab = ref<'docs' | 'qa'>('docs') const activeTab = ref<'docs' | 'qa'>('docs')
const docs = ref<any[]>([]) const docs = ref<any[]>([])
const pendingUploads = ref<any[]>([])
const qaPairs = ref<any[]>([]) const qaPairs = ref<any[]>([])
const uploading = ref(false) const uploading = computed(() => pendingUploads.value.some((doc) => doc.localUploading))
const uploadError = ref('') const uploadError = ref('')
const dragOver = ref(false) const dragOver = ref(false)
const fileInput = ref<HTMLInputElement | null>(null) const fileInput = ref<HTMLInputElement | null>(null)
let documentPollingTimer: ReturnType<typeof setInterval> | undefined
const query = ref('') const query = ref('')
const searching = ref(false) const searching = ref(false)
const searched = ref(false) const searched = ref(false)
const searchResults = ref<any[]>([]) const searchResults = ref<any[]>([])
const displayDocs = computed(() => [...pendingUploads.value, ...docs.value])
const documentState = (doc: any) => { const documentState = (doc: any) => {
if (doc.localUploading) {
return { tone: 'pending', label: '上传中', detail: `正在上传 ${doc.uploadProgress || 0}%`, progress: doc.uploadProgress || 0 }
}
if (doc.localOnly) {
return { tone: 'failed', label: '上传失败', detail: doc.errorMessage || '文件未上传成功,请移除后重试' }
}
if (doc.filePresent === false) { if (doc.filePresent === false) {
return { tone: 'missing', label: '文件缺失', detail: '原文件不可用,请删除后重新上传' } return { tone: 'missing', label: '文件缺失', detail: '原文件不可用,请删除后重新上传' }
} }
@@ -121,9 +138,33 @@ const documentState = (doc: any) => {
return { tone: 'ready', label: '已入库', detail: `已切分 ${doc.chunkCount || 0} 段,可用于对话` } return { tone: 'ready', label: '已入库', detail: `已切分 ${doc.chunkCount || 0} 段,可用于对话` }
} }
if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) { if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
return { tone: 'pending', label: '处理中', detail: '正在解析并建立知识索引' } const stage = String(doc.indexStage || 'queued').toLowerCase()
const labels: Record<string, string> = {
queued: '等待处理', extracting: '解析文档', chunking: '切分文本', embedding: '向量化中'
}
const progress = Math.max(0, Math.min(99, Number(doc.indexProgress || 0)))
return { tone: 'pending', label: labels[stage] || '处理中', detail: `${labels[stage] || '正在建立知识索引'} ${progress}%`, progress }
} }
return { tone: 'failed', label: '处理失败', detail: '未能建立知识索引,请删除后重新上传' } return { tone: 'failed', label: '处理失败', detail: doc.errorMessage || '未能建立知识索引,请重新索引或重新上传' }
}
const hasPendingDocuments = () => docs.value.some((doc) =>
['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())
)
const stopDocumentPolling = () => {
if (documentPollingTimer) {
clearInterval(documentPollingTimer)
documentPollingTimer = undefined
}
}
const startDocumentPolling = () => {
if (documentPollingTimer || !hasPendingDocuments()) return
documentPollingTimer = setInterval(async () => {
await loadDocs()
if (!hasPendingDocuments()) stopDocumentPolling()
}, 2000)
} }
const loadDocs = async () => { const loadDocs = async () => {
@@ -131,6 +172,7 @@ const loadDocs = async () => {
try { try {
const res: any = await getKnowledgeDocs(avatarId.value) const res: any = await getKnowledgeDocs(avatarId.value)
docs.value = unwrapListData(res) docs.value = unwrapListData(res)
startDocumentPolling()
} catch (e) { } catch (e) {
console.error(e) console.error(e)
} }
@@ -149,40 +191,82 @@ const loadQA = async () => {
const triggerFile = () => fileInput.value?.click() const triggerFile = () => fileInput.value?.click()
const onFileChange = (e: Event) => { const onFileChange = (e: Event) => {
const f = (e.target as HTMLInputElement).files?.[0] const files = Array.from((e.target as HTMLInputElement).files || [])
if (f) doUpload(f) if (files.length) uploadFiles(files)
;(e.target as HTMLInputElement).value = '' ;(e.target as HTMLInputElement).value = ''
} }
const onDrop = (e: DragEvent) => { const onDrop = (e: DragEvent) => {
dragOver.value = false dragOver.value = false
const f = e.dataTransfer?.files?.[0] const files = Array.from(e.dataTransfer?.files || [])
if (f) doUpload(f) if (files.length) uploadFiles(files)
} }
const doUpload = async (file: File) => { const uploadFiles = (files: File[]) => {
uploadError.value = '' uploadError.value = ''
const ext = '.' + (file.name.split('.').pop() || '').toLowerCase()
if (!['.md', '.txt', '.pdf', '.doc', '.docx', '.xlsx'].includes(ext)) {
uploadError.value = `不支持的类型:${ext},仅支持 md/txt/pdf/doc/docx/xlsx`
return
}
if (!avatarId.value) { if (!avatarId.value) {
uploadError.value = '请先创建数字分身' uploadError.value = '请先创建数字分身'
return return
} }
uploading.value = true for (const file of files) {
const ext = '.' + (file.name.split('.').pop() || '').toLowerCase()
if (!['.md', '.txt', '.pdf', '.doc', '.docx', '.xlsx'].includes(ext)) {
uploadError.value = `不支持的类型:${ext},仅支持 md/txt/pdf/doc/docx/xlsx`
continue
}
void uploadOne(file, ext)
}
}
const uploadOne = async (file: File, ext: string) => {
if (!avatarId.value) return
const localId = `upload-${Date.now()}-${Math.random().toString(16).slice(2)}`
const card = {
id: localId,
filename: file.name,
fileType: ext.slice(1),
fileSize: file.size,
createdAt: new Date().toISOString(),
localUploading: true,
localOnly: true,
uploadProgress: 0,
errorMessage: ''
}
pendingUploads.value.unshift(card)
try { try {
await uploadKnowledgeDoc(avatarId.value, file) const created: any = await uploadKnowledgeDoc(avatarId.value, file, (loaded, total) => {
const current = pendingUploads.value.find((doc) => doc.id === localId)
if (current) current.uploadProgress = Math.min(99, Math.round((loaded / Math.max(1, total)) * 100))
})
pendingUploads.value = pendingUploads.value.filter((doc) => doc.id !== localId)
docs.value = [created, ...docs.value.filter((doc) => doc.id !== created.id)]
startDocumentPolling()
} catch (e: any) {
const current = pendingUploads.value.find((doc) => doc.id === localId)
if (current) {
current.localUploading = false
current.errorMessage = e?.message || '上传失败'
}
}
}
const retryDoc = async (id: string) => {
if (!avatarId.value) return
uploadError.value = ''
try {
await retryKnowledgeDoc(avatarId.value, id)
await loadDocs() await loadDocs()
} catch (e: any) { } catch (e: any) {
uploadError.value = e?.message || '上传失败' uploadError.value = e?.message || '重新索引失败'
} finally {
uploading.value = false
} }
} }
const removeDoc = async (id: string) => { const removeDoc = async (id: string) => {
const local = pendingUploads.value.find((doc) => doc.id === id)
if (local?.localOnly) {
pendingUploads.value = pendingUploads.value.filter((doc) => doc.id !== id)
return
}
if (!avatarId.value) return if (!avatarId.value) return
await deleteKnowledgeDoc(avatarId.value, id) await deleteKnowledgeDoc(avatarId.value, id)
await loadDocs() await loadDocs()
@@ -257,6 +341,8 @@ onMounted(async () => {
if (avatarId.value) store.currentAvatarId = avatarId.value if (avatarId.value) store.currentAvatarId = avatarId.value
await Promise.all([loadDocs(), loadQA()]) await Promise.all([loadDocs(), loadQA()])
}) })
onUnmounted(stopDocumentPolling)
</script> </script>
<style scoped> <style scoped>
@@ -300,7 +386,12 @@ onMounted(async () => {
.status-pill.missing { color: #B91C1C; background: #FEF2F2; } .status-pill.missing { color: #B91C1C; background: #FEF2F2; }
.status-pill.failed { color: #B91C1C; background: #FEF2F2; } .status-pill.failed { color: #B91C1C; background: #FEF2F2; }
.card-meta, .card-detail { margin: 5px 0 0; color: #9398AE; font-size: 11px; line-height: 1.4; }.card-detail { color: #8B6B58; } .card-meta, .card-detail { margin: 5px 0 0; color: #9398AE; font-size: 11px; line-height: 1.4; }.card-detail { color: #8B6B58; }
.card-delete { flex: 0 0 auto; align-self: center; border: 0; color: #EF4444; background: #FEF2F2; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; } .progress-track { width: 100%; height: 4px; margin-top: 8px; overflow: hidden; border-radius: 999px; background: #FDE7D1; }
.progress-fill { display: block; height: 100%; border-radius: inherit; background: linear-gradient(90deg, #FB923C, #F97316); transition: width .25s ease; }
.card-actions { flex: 0 0 auto; display: flex; flex-direction: column; align-items: stretch; gap: 6px; }
.card-delete, .card-retry { align-self: center; border: 0; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; white-space: nowrap; }
.card-delete { color: #EF4444; background: #FEF2F2; }
.card-retry { color: #C15F18; background: #FFF3E6; }
.card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; } .card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; }
.qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; } .qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; }
.qa-card .card-content, .qa-card .card-content,