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 Avatar, KnowledgeChunk, KnowledgeDoc, QAPair from routers.knowledge import _doc_payload from services.knowledge_vectorizer import knowledge_vectorizer client = TestClient(app) def test_doc_payload_reports_whether_the_persisted_file_exists(tmp_path: Path): avatar_id = "avatar-1" stored_name = "knowledge.md" doc = SimpleNamespace( avatar_id=avatar_id, file_url=f"/api/files/{avatar_id}/{stored_name}", to_dict=lambda: {"id": "doc-1", "fileUrl": f"/api/files/{avatar_id}/{stored_name}"}, ) stored_dir = tmp_path / avatar_id stored_dir.mkdir() stored_file = stored_dir / stored_name with patch("routers.knowledge.UPLOAD_DIR", str(tmp_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_returns_before_background_vectorization( tmp_path: Path, authorization_context, ): context = authorization_context 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", headers=context["owner_headers"], files={"file": ("knowledge.md", b"# Knowledge\n\nTest content", "text/markdown")}, ) payload = response.json()["data"] assert payload["status"] == "parsing" assert payload["vectorized"] is False assert payload["chunkCount"] == 0 enqueue.assert_called_once_with(payload["id"]) db = SessionLocal() try: stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one() assert stored.status == "parsing" assert db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).count() == 0 db.delete(stored) db.commit() finally: db.close() def test_upload_rejects_oversize_file_before_queuing_indexing( tmp_path: Path, authorization_context, ): context = authorization_context with ( patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)), 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"] == "文件不能超过 100MB" 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( 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"] == "parsing" with ( 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() try: stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one() assert stored.status == "ready" assert stored.vectorized is True assert stored.chunk_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.delete(stored) db.commit() finally: 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): context = authorization_context first_avatar_id = context["avatar"].id second_avatar_id = f"knowledge-second-{context['suffix']}" first_doc_id = f"knowledge-first-doc-{context['suffix']}" second_doc_id = f"knowledge-second-doc-{context['suffix']}" first_qa_id = f"knowledge-first-qa-{context['suffix']}" second_qa_id = f"knowledge-second-qa-{context['suffix']}" db = SessionLocal() try: db.add_all( [ Avatar( id=second_avatar_id, owner_id=context["owner"].huihui_user_id, name="独立知识库分身", status="active", config={}, ), KnowledgeDoc( id=first_doc_id, avatar_id=first_avatar_id, filename="first.md", status="ready", vectorized=True, ), KnowledgeDoc( id=second_doc_id, avatar_id=second_avatar_id, filename="second.md", status="ready", vectorized=True, ), QAPair( id=first_qa_id, avatar_id=first_avatar_id, question="第一个分身问题", answer="第一个分身答案", ), QAPair( id=second_qa_id, avatar_id=second_avatar_id, question="第二个分身问题", answer="第二个分身答案", ), ] ) db.commit() finally: db.close() try: first_docs = client.get( f"/api/avatar/{first_avatar_id}/knowledge/docs", headers=context["owner_headers"], ).json()["data"] second_docs = client.get( f"/api/avatar/{second_avatar_id}/knowledge/docs", headers=context["owner_headers"], ).json()["data"] first_qa = client.get( f"/api/avatar/{first_avatar_id}/knowledge/qa", headers=context["owner_headers"], ).json()["data"] second_qa = client.get( f"/api/avatar/{second_avatar_id}/knowledge/qa", headers=context["owner_headers"], ).json()["data"] assert [item["id"] for item in first_docs if item["id"] == first_doc_id] == [first_doc_id] assert second_doc_id not in {item["id"] for item in first_docs} assert [item["id"] for item in second_docs] == [second_doc_id] assert first_qa_id in {item["id"] for item in first_qa} assert second_qa_id not in {item["id"] for item in first_qa} assert [item["id"] for item in second_qa] == [second_qa_id] finally: db = SessionLocal() try: db.query(QAPair).filter(QAPair.id.in_([first_qa_id, second_qa_id])).delete( synchronize_session=False ) db.query(KnowledgeDoc).filter( KnowledgeDoc.id.in_([first_doc_id, second_doc_id]) ).delete(synchronize_session=False) db.query(Avatar).filter(Avatar.id == second_avatar_id).delete() db.commit() finally: db.close()