424 lines
15 KiB
Python
424 lines
15 KiB
Python
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"] == "文件不能超过 50MB"
|
|
enqueue.assert_not_called()
|
|
assert not list((tmp_path / context["avatar"].id).glob("*"))
|
|
|
|
|
|
def test_multipart_upload_reassembles_file_before_queuing_indexing(
|
|
tmp_path: Path,
|
|
authorization_context,
|
|
):
|
|
context = authorization_context
|
|
avatar_id = context["avatar"].id
|
|
content = b"0123456789"
|
|
with (
|
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
|
patch("routers.knowledge.MULTIPART_CHUNK_BYTES", 4),
|
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
|
):
|
|
created = client.post(
|
|
f"/api/avatar/{avatar_id}/knowledge/uploads",
|
|
headers=context["owner_headers"],
|
|
json={"filename": "large.pdf", "fileSize": len(content), "totalChunks": 3},
|
|
).json()["data"]
|
|
|
|
for index, chunk in enumerate((content[:4], content[4:8], content[8:])):
|
|
response = client.post(
|
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/chunks/{index}",
|
|
headers=context["owner_headers"],
|
|
files={"file": (f"chunk-{index}", chunk, "application/octet-stream")},
|
|
)
|
|
assert response.json()["code"] == 200
|
|
|
|
completed = client.post(
|
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/complete",
|
|
headers=context["owner_headers"],
|
|
).json()["data"]
|
|
|
|
assert completed["status"] == "parsing"
|
|
assert completed["fileSize"] == len(content)
|
|
enqueue.assert_called_once_with(completed["id"])
|
|
stored_path = tmp_path / avatar_id / Path(completed["fileUrl"]).name
|
|
assert stored_path.read_bytes() == content
|
|
assert not (tmp_path / ".multipart" / avatar_id / created["uploadId"]).exists()
|
|
|
|
db = SessionLocal()
|
|
try:
|
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == completed["id"]).one()
|
|
db.delete(stored)
|
|
db.commit()
|
|
finally:
|
|
db.close()
|
|
|
|
|
|
def test_multipart_upload_rejects_incomplete_parts(
|
|
tmp_path: Path,
|
|
authorization_context,
|
|
):
|
|
context = authorization_context
|
|
avatar_id = context["avatar"].id
|
|
with (
|
|
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
|
patch("routers.knowledge.MULTIPART_CHUNK_BYTES", 4),
|
|
patch("routers.knowledge.knowledge_vectorizer.enqueue") as enqueue,
|
|
):
|
|
created = client.post(
|
|
f"/api/avatar/{avatar_id}/knowledge/uploads",
|
|
headers=context["owner_headers"],
|
|
json={"filename": "large.pdf", "fileSize": 6, "totalChunks": 2},
|
|
).json()["data"]
|
|
client.post(
|
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/chunks/0",
|
|
headers=context["owner_headers"],
|
|
files={"file": ("chunk-0", b"0123", "application/octet-stream")},
|
|
)
|
|
response = client.post(
|
|
f"/api/avatar/{avatar_id}/knowledge/uploads/{created['uploadId']}/complete",
|
|
headers=context["owner_headers"],
|
|
)
|
|
|
|
assert response.json()["code"] == 400
|
|
assert response.json()["message"] == "文件分片尚未上传完整"
|
|
enqueue.assert_not_called()
|
|
|
|
|
|
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 stored.index_stage == "ready"
|
|
assert stored.index_progress == 100
|
|
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_background_vectorizer_uses_ocr_for_image_only_pdf(
|
|
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": ("scanned.pdf", b"image-only-pdf", "application/pdf")},
|
|
)
|
|
|
|
payload = response.json()["data"]
|
|
progress = []
|
|
with (
|
|
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
|
|
patch("services.knowledge_vectorizer.embeddings.extract_text", return_value=""),
|
|
patch(
|
|
"services.knowledge_vectorizer.extract_scanned_pdf_text",
|
|
side_effect=lambda _db, _avatar, _path, on_progress: (
|
|
on_progress(1, 2), on_progress(2, 2), "扫描页文字"
|
|
)[-1],
|
|
) as ocr,
|
|
patch("services.knowledge_vectorizer.embeddings.embed", return_value=[[1.0, 0.0]]),
|
|
patch.object(knowledge_vectorizer, "_set_progress", wraps=knowledge_vectorizer._set_progress) as set_progress,
|
|
):
|
|
knowledge_vectorizer.vectorize_document(payload["id"])
|
|
progress = [(call.args[2], call.args[3]) for call in set_progress.call_args_list]
|
|
|
|
db = SessionLocal()
|
|
try:
|
|
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
|
assert stored.status == "ready"
|
|
assert stored.chunk_count == 1
|
|
assert ("ocr", 18) in progress
|
|
assert ("ocr", 28) in progress
|
|
ocr.assert_called_once()
|
|
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
|
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()
|