Files
sovereign_ai/SovereignAI-Starter/tests/test_knowledge.py
T

318 lines
13 KiB
Python

from __future__ import annotations
import base64
import hashlib
import os
import tempfile
import unittest
from io import BytesIO
from pathlib import Path
from unittest.mock import patch
from uuid import uuid4
from PIL import Image, ImageDraw
from fastapi.testclient import TestClient
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-knowledge-tests-")
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
from app import auth, database, knowledge
from app.main import app
from tests.api_client import authenticated_client
class KnowledgeIndexTests(unittest.TestCase):
def setUp(self) -> None:
self.temp = tempfile.TemporaryDirectory()
self.base = Path(self.temp.name)
self.database = self.base / "knowledge.sqlite3"
self.workspace = self.base / "project"
self.workspace.mkdir()
self.other_workspace = self.base / "other-project"
self.other_workspace.mkdir()
with patch.object(knowledge.database, "DATABASE_PATH", self.database):
knowledge.initialize()
def tearDown(self) -> None:
self.temp.cleanup()
def _index(self, text: str, root: Path | None = None) -> None:
selected_root = root or self.workspace
source = selected_root / "docs" / "guide.md"
source.parent.mkdir(parents=True, exist_ok=True)
source.write_text(text, encoding="utf-8")
raw = text.encode("utf-8")
with patch.object(knowledge.database, "DATABASE_PATH", self.database):
knowledge.index_document(
user_id="local-user",
workspace_path=selected_root,
relative_path="docs/guide.md",
text=text,
content_hash=hashlib.sha256(raw).hexdigest(),
)
def _search(self, query: str, root: Path | None = None) -> list[dict[str, object]]:
with patch.object(knowledge.database, "DATABASE_PATH", self.database):
return knowledge.search(
query,
user_id="local-user",
workspace_path=root or self.workspace,
)
def test_search_returns_source_and_matching_chunk(self) -> None:
self._index("SQLite keeps the local conversation history in a durable database.")
result = self._search("conversation database")
self.assertTrue(result)
self.assertEqual(result[0]["path"], "docs/guide.md")
self.assertIn("conversation", str(result[0]["text"]))
def test_search_limits_duplicate_chunks_per_document_to_keep_source_diversity(self) -> None:
documents = {
"docs/a.md": "sharedmarker " * 700,
"docs/b.md": "sharedmarker evidence from a second document.",
}
with patch.object(knowledge.database, "DATABASE_PATH", self.database):
for relative_path, text in documents.items():
source = self.workspace / relative_path
source.parent.mkdir(parents=True, exist_ok=True)
source.write_text(text, encoding="utf-8")
knowledge.index_document(
user_id="local-user",
workspace_path=self.workspace,
relative_path=relative_path,
text=text,
content_hash=hashlib.sha256(text.encode("utf-8")).hexdigest(),
)
with patch.object(knowledge.database, "DATABASE_PATH", self.database):
results = knowledge.search(
"sharedmarker",
user_id="local-user",
workspace_path=self.workspace,
limit=5,
)
paths = [str(item["path"]) for item in results]
self.assertIn("docs/b.md", paths)
self.assertLessEqual(paths.count("docs/a.md"), knowledge.MAX_CHUNKS_PER_DOCUMENT)
def test_reindex_replaces_old_chunks_and_search_is_workspace_scoped(self) -> None:
self._index("oldkeyword exists in the first document.")
self._index("newkeyword replaces the first document.")
self.assertEqual(self._search("oldkeyword"), [])
self.assertTrue(self._search("newkeyword"))
self.assertEqual(self._search("newkeyword", self.other_workspace), [])
def test_search_discards_a_source_that_changed_after_indexing(self) -> None:
self._index("stalephrase was present when the index was created.")
(self.workspace / "docs" / "guide.md").write_text(
"newphrase replaced the indexed document.", encoding="utf-8"
)
self.assertEqual(self._search("stalephrase"), [])
self.assertEqual(self._search("stalephrase"), [])
self._index("newphrase is now indexed.")
self.assertTrue(self._search("newphrase"))
def test_arabic_audio_synonyms_retrieve_transcription_source(self) -> None:
self._index("يرسل مسار الصوت التسجيل إلى Groq باستخدام Whisper للتفريغ.")
results = self._search("أي خدمة تفرغ التسجيلات الصوتية؟")
self.assertTrue(results)
self.assertEqual(results[0]["path"], "docs/guide.md")
self.assertIn("Whisper", str(results[0]["text"]))
def test_chunking_preserves_overlap_for_boundary_context(self) -> None:
text = "a" * (knowledge.CHUNK_SIZE - 2) + " boundaryphrase " + "b" * 80
chunks = knowledge._chunks(text)
self.assertGreaterEqual(len(chunks), 2)
self.assertIn("boundaryphrase", "".join(chunks))
def test_api_indexes_selected_file_and_returns_matching_source(self) -> None:
file = self.workspace / "guide.md"
file.write_text("The project stores conversations in SQLite.", encoding="utf-8")
client = authenticated_client(app)
indexed = client.post(
"/v1/agent/knowledge/index",
json={"workspace_path": str(self.workspace), "files": ["guide.md"]},
)
result = client.post(
"/v1/agent/knowledge/search",
json={
"workspace_path": str(self.workspace),
"task": "conversations SQLite",
},
)
self.assertEqual(indexed.status_code, 200, indexed.text)
self.assertEqual(indexed.json()["storage"], "local_sqlite_fts5")
self.assertEqual(result.status_code, 200, result.text)
self.assertEqual(result.json()["results"][0]["path"], "guide.md")
deleted = client.request(
"DELETE",
"/v1/agent/knowledge/index",
json={"workspace_path": str(self.workspace), "files": ["guide.md"]},
)
self.assertEqual(deleted.status_code, 200, deleted.text)
self.assertTrue(deleted.json()["deleted"][0]["deleted"])
def test_api_knowledge_index_isolated_between_account_sessions(self) -> None:
marker = "OwnerScopedKnowledgeMarker"
(self.workspace / "private.md").write_text(marker, encoding="utf-8")
owner_id = auth.create_account(
f"{uuid4().hex}@example.test", "account one secure passphrase"
)
other_id = auth.create_account(
f"{uuid4().hex}@example.test", "account two secure passphrase"
)
self.addCleanup(self._delete_test_users, owner_id, other_id)
owner_token, _ = auth.issue_session(owner_id)
other_token, _ = auth.issue_session(other_id)
owner_client = TestClient(
app, headers={"Authorization": f"Bearer {owner_token}"}
)
other_client = TestClient(
app, headers={"Authorization": f"Bearer {other_token}"}
)
payload = {
"workspace_path": str(self.workspace),
"files": ["private.md"],
}
indexed = owner_client.post("/v1/agent/knowledge/index", json=payload)
owner_results = owner_client.post(
"/v1/agent/knowledge/search",
json={"workspace_path": str(self.workspace), "task": marker},
)
other_results = other_client.post(
"/v1/agent/knowledge/search",
json={"workspace_path": str(self.workspace), "task": marker},
)
self.assertEqual(indexed.status_code, 200, indexed.text)
self.assertEqual(owner_results.status_code, 200, owner_results.text)
self.assertEqual(other_results.status_code, 200, other_results.text)
self.assertEqual(owner_results.json()["results"][0]["path"], "private.md")
self.assertEqual(other_results.json()["results"], [])
knowledge.delete_document(
user_id=owner_id,
workspace_path=self.workspace,
relative_path="private.md",
)
@staticmethod
def _delete_test_users(*user_ids: str) -> None:
with database._connect() as connection:
connection.executemany(
"DELETE FROM users WHERE id=?", [(user_id,) for user_id in user_ids]
)
def test_api_ocr_indexes_a_scanned_pdf_and_retrieves_its_text(self) -> None:
pdf_path = self.workspace / "notice.pdf"
image = Image.new("RGB", (900, 1165), "white")
ImageDraw.Draw(image).text((60, 80), "scanned notice", fill="black")
output = BytesIO()
image.save(output, format="PDF", resolution=144)
pdf_path.write_bytes(output.getvalue())
client = authenticated_client(app)
ocr = {
"engine": "easyocr-local-ar-en",
"text": "Community Reading Meetup Place Amman",
"lines": [],
"average_confidence": 0.8,
"elapsed_seconds": 1.0,
}
with patch("app.main.recognize_image_text", return_value=ocr):
indexed = client.post(
"/v1/agent/knowledge/index",
json={"workspace_path": str(self.workspace), "files": ["notice.pdf"]},
)
searched = client.post(
"/v1/agent/knowledge/search",
json={"workspace_path": str(self.workspace), "task": "Community Meetup Amman"},
)
self.assertEqual(indexed.status_code, 200, indexed.text)
self.assertEqual(indexed.json()["indexed"][0]["ocr_pages"], 1)
self.assertEqual(searched.status_code, 200, searched.text)
self.assertEqual(searched.json()["results"][0]["path"], "notice.pdf")
self.assertIn("Community Reading Meetup", searched.json()["results"][0]["text"])
def test_api_indexes_digital_and_scanned_pages_from_mixed_pdf(self) -> None:
pdf_path = self.workspace / "mixed.pdf"
pdf_path.write_bytes(b"%PDF-mixed-test")
client = authenticated_client(app)
rendered = {
"pages": [
{
"page": 2,
"mime_type": "image/jpeg",
"data": base64.b64encode(b"rendered jpeg").decode("ascii"),
}
],
"total_pages": 2,
"truncated": False,
}
ocr = {
"engine": "easyocr-local-ar-en",
"text": "ScannedPageMarker",
"lines": [],
"average_confidence": 0.85,
"elapsed_seconds": 0.9,
}
with (
patch(
"app.main.extract_pdf_pages_text",
return_value={
"pages": [
{"page": 1, "text": "DigitalPageMarker", "has_text": True},
{"page": 2, "text": "", "has_text": False},
],
"total_pages": 2,
"truncated": False,
},
),
patch("app.main.render_scanned_pdf_pages", return_value=rendered),
patch("app.main.recognize_image_text", return_value=ocr),
):
indexed = client.post(
"/v1/agent/knowledge/index",
json={"workspace_path": str(self.workspace), "files": ["mixed.pdf"]},
)
digital_result = client.post(
"/v1/agent/knowledge/search",
json={"workspace_path": str(self.workspace), "task": "DigitalPageMarker"},
)
scanned_result = client.post(
"/v1/agent/knowledge/search",
json={"workspace_path": str(self.workspace), "task": "ScannedPageMarker"},
)
deleted = client.request(
"DELETE",
"/v1/agent/knowledge/index",
json={"workspace_path": str(self.workspace), "files": ["mixed.pdf"]},
)
self.assertEqual(indexed.status_code, 200, indexed.text)
self.assertEqual(indexed.json()["indexed"][0]["ocr_pages"], 1)
self.assertEqual(digital_result.status_code, 200, digital_result.text)
self.assertEqual(scanned_result.status_code, 200, scanned_result.text)
self.assertIn("DigitalPageMarker", digital_result.json()["results"][0]["text"])
self.assertIn("ScannedPageMarker", scanned_result.json()["results"][0]["text"])
self.assertEqual(deleted.status_code, 200, deleted.text)
if __name__ == "__main__":
unittest.main()