Complete local hybrid search and improve agent reliability
This commit is contained in:
@@ -0,0 +1,202 @@
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
from fastapi import HTTPException
|
||||
|
||||
_TEST_DATA_DIR = None
|
||||
if "SOVEREIGNAI_DATA_DIR" not in os.environ:
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-skills-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from app.main import AgentRequest, _execute_agent, app, safe_arithmetic
|
||||
|
||||
|
||||
class AgentSkillTests(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls) -> None:
|
||||
cls.client = TestClient(app)
|
||||
cls.workspace = str(Path(__file__).resolve().parents[1])
|
||||
|
||||
def test_skill_catalog_discloses_scope_and_permissions(self) -> None:
|
||||
response = self.client.get("/v1/agent/skills")
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
skills = {item["id"]: item for item in response.json()["skills"]}
|
||||
self.assertEqual(set(skills), {"code_explain", "code_review", "test_plan"})
|
||||
self.assertNotIn("propose_file_change", skills["code_explain"]["allowed_tools"])
|
||||
self.assertIn("propose_file_change", skills["code_review"]["allowed_tools"])
|
||||
self.assertEqual(response.json()["default"], None)
|
||||
|
||||
def test_safe_arithmetic_accepts_common_unicode_operator_symbols(self) -> None:
|
||||
self.assertEqual(safe_arithmetic("137 × 29"), 3973.0)
|
||||
self.assertEqual(safe_arithmetic("12 ÷ 3 − 1"), 3.0)
|
||||
|
||||
def test_explicit_calculator_request_uses_bounded_local_calculator(self) -> None:
|
||||
with patch("app.main.get_completion", new=AsyncMock()) as model:
|
||||
result = asyncio.run(
|
||||
_execute_agent(
|
||||
AgentRequest(
|
||||
task="استخدم الحاسبة المتاحة لحساب 137 × 29، ثم أجب بالناتج فقط.",
|
||||
model="qwen2.5:1.5b-instruct-q4_K_M",
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
model.assert_not_awaited()
|
||||
self.assertEqual(result["tool"], "calculator")
|
||||
self.assertEqual(result["result"], 3973.0)
|
||||
|
||||
def test_explain_skill_is_sent_to_model_and_hides_file_write_tool(self) -> None:
|
||||
completion = {"choices": [{"message": {"content": "شرح مختصر."}}]}
|
||||
with patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model:
|
||||
result = asyncio.run(
|
||||
_execute_agent(
|
||||
AgentRequest(
|
||||
task="اشرح بنية المشروع باختصار",
|
||||
workspace_path=self.workspace,
|
||||
skill_id="code_explain",
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
payload = model.await_args.args[0]
|
||||
tool_names = [tool["function"]["name"] for tool in payload["tools"]]
|
||||
system = payload["messages"][0]["content"]
|
||||
self.assertIn("شرح الكود", system)
|
||||
self.assertIn("ميّز بين ما قرأته وما استنتجته", system)
|
||||
self.assertEqual(tool_names, ["calculator", "search_workspace", "search_knowledge"])
|
||||
self.assertEqual(result["skill"], "code_explain")
|
||||
|
||||
def test_review_skill_exposes_preview_tool_but_never_applies_it(self) -> None:
|
||||
completion = {"choices": [{"message": {"content": "سأعرض النتائج."}}]}
|
||||
with patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model:
|
||||
result = asyncio.run(
|
||||
_execute_agent(
|
||||
AgentRequest(
|
||||
task="راجع الملفات دون تعديل",
|
||||
workspace_path=self.workspace,
|
||||
skill_id="code_review",
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
payload = model.await_args.args[0]
|
||||
tool_names = [tool["function"]["name"] for tool in payload["tools"]]
|
||||
self.assertIn("propose_file_change", tool_names)
|
||||
self.assertIn("لا تطبق الكتابة", payload["messages"][0]["content"])
|
||||
self.assertEqual(result["skill"], "code_review")
|
||||
self.assertNotIn("proposal", result)
|
||||
|
||||
def test_preselected_file_is_read_once_without_redundant_search_call(self) -> None:
|
||||
completion = {"choices": [{"message": {"content": "المهارات مسجلة في قاموس محلي."}}]}
|
||||
with patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model:
|
||||
result = asyncio.run(
|
||||
_execute_agent(
|
||||
AgentRequest(
|
||||
task="اشرح الملف المحدد",
|
||||
workspace_path=self.workspace,
|
||||
workspace_files=["app/skills.py"],
|
||||
skill_id="code_explain",
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(model.await_count, 1)
|
||||
self.assertEqual(result["files"], ["app/skills.py"])
|
||||
self.assertIn("Curated, local agent skills", model.await_args.args[0]["messages"][1]["content"])
|
||||
|
||||
def test_explicit_knowledge_search_is_prefetched_before_model_answer(self) -> None:
|
||||
completion = {"choices": [{"message": {"content": "المرحلة 5 تضيف الفهرسة المحلية."}}]}
|
||||
match = {"path": "ROADMAP.md", "chunk": 2, "text": "SQLite FTS5 local index", "excerpt": "SQLite FTS5"}
|
||||
with (
|
||||
patch("app.main.knowledge.search", return_value=[match]) as retrieve,
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
result = asyncio.run(
|
||||
_execute_agent(
|
||||
AgentRequest(
|
||||
task="ابحث في فهرس المعرفة عن المرحلة 5",
|
||||
workspace_path=self.workspace,
|
||||
skill_id="test_plan",
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
retrieve.assert_called_once()
|
||||
model_payload = model.await_args.args[0]
|
||||
self.assertEqual(model_payload["max_tokens"], 384)
|
||||
self.assertIn("SQLite FTS5 local index", model_payload["messages"][1]["content"])
|
||||
self.assertNotIn(
|
||||
"search_knowledge",
|
||||
[tool["function"]["name"] for tool in model_payload["tools"]],
|
||||
)
|
||||
self.assertEqual(result["tool"], "search_knowledge")
|
||||
self.assertEqual(result["files"], ["ROADMAP.md"])
|
||||
|
||||
def test_test_plan_skill_only_advertises_workspace_search(self) -> None:
|
||||
completion = {"choices": [{"message": {"content": "ثلاث حالات اختبار مقترحة."}}]}
|
||||
with patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model:
|
||||
asyncio.run(
|
||||
_execute_agent(
|
||||
AgentRequest(
|
||||
task="أنشئ خطة اختبار",
|
||||
workspace_path=self.workspace,
|
||||
skill_id="test_plan",
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
tools = model.await_args.args[0]["tools"]
|
||||
self.assertEqual(
|
||||
[tool["function"]["name"] for tool in tools],
|
||||
["search_workspace", "search_knowledge"],
|
||||
)
|
||||
|
||||
def test_server_rejects_tool_call_outside_active_skill_permissions(self) -> None:
|
||||
completion = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call-forbidden",
|
||||
"function": {
|
||||
"name": "propose_file_change",
|
||||
"arguments": '{"path":"new.py","operation":"create","content":"print(1)"}',
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
with patch("app.main.get_completion", new=AsyncMock(return_value=completion)):
|
||||
with self.assertRaises(HTTPException) as error:
|
||||
asyncio.run(
|
||||
_execute_agent(
|
||||
AgentRequest(
|
||||
task="اشرح المشروع",
|
||||
workspace_path=self.workspace,
|
||||
skill_id="code_explain",
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(error.exception.status_code, 422)
|
||||
|
||||
def test_unknown_skill_is_rejected_by_request_contract(self) -> None:
|
||||
response = self.client.post(
|
||||
"/v1/agent/run",
|
||||
json={"task": "سؤال", "skill_id": "execute_shell"},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 422)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,95 @@
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from uuid import UUID
|
||||
from unittest.mock import patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
# Importing the API initializes its SQLite schema. Keep this test process isolated
|
||||
# from the real local conversation database.
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-api-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from app.main import app
|
||||
|
||||
|
||||
class ApiErrorContractTests(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls) -> None:
|
||||
cls.client = TestClient(app)
|
||||
|
||||
def test_health_response_has_correlation_id(self) -> None:
|
||||
response = self.client.get("/health")
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
UUID(response.headers["x-request-id"])
|
||||
self.assertEqual(response.json()["status"], "ok")
|
||||
|
||||
def test_not_found_uses_normalized_error_contract(self) -> None:
|
||||
response = self.client.get("/no-such-route")
|
||||
|
||||
self.assertEqual(response.status_code, 404)
|
||||
self.assertEqual(response.json()["error"]["code"], "not_found")
|
||||
self.assertEqual(response.json()["request_id"], response.headers["x-request-id"])
|
||||
|
||||
def test_validation_error_is_normalized_and_does_not_echo_input(self) -> None:
|
||||
sentinel = "sensitive-validation-input"
|
||||
response = self.client.post(
|
||||
"/v1/chat/completions", json={"messages": sentinel}
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 422)
|
||||
payload = response.json()
|
||||
self.assertEqual(payload["error"]["code"], "invalid_request")
|
||||
self.assertEqual(payload["request_id"], response.headers["x-request-id"])
|
||||
self.assertNotIn(sentinel, response.text)
|
||||
self.assertNotIn("input", payload["detail"][0])
|
||||
|
||||
def test_upstream_timeout_uses_gateway_timeout_contract(self) -> None:
|
||||
with patch(
|
||||
"app.main.get_completion",
|
||||
side_effect=HTTPException(
|
||||
status_code=504, detail="انتهت مهلة انتظار خادم النموذج المحلي."
|
||||
),
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/chat/completions",
|
||||
json={"messages": [{"role": "user", "content": "مرحبا"}]},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 504)
|
||||
self.assertEqual(response.json()["error"]["code"], "upstream_timeout")
|
||||
self.assertEqual(response.json()["request_id"], response.headers["x-request-id"])
|
||||
|
||||
def test_model_list_exposes_verified_capabilities(self) -> None:
|
||||
class MetadataProvider:
|
||||
name = "test"
|
||||
default_model = "vision-test"
|
||||
base_url = "http://local"
|
||||
|
||||
async def list_models(self):
|
||||
return ["vision-test", "unknown-test"]
|
||||
|
||||
async def describe_models(self, model_names):
|
||||
return {
|
||||
"vision-test": {
|
||||
"verified": True,
|
||||
"capabilities": ["completion", "vision"],
|
||||
},
|
||||
"unknown-test": {"verified": False, "capabilities": []},
|
||||
}
|
||||
|
||||
with patch("app.main.get_model_provider", return_value=MetadataProvider()):
|
||||
response = self.client.get("/v1/models")
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
models = {item["id"]: item for item in response.json()["data"]}
|
||||
self.assertEqual(models["vision-test"]["capabilities"], ["completion", "vision"])
|
||||
self.assertTrue(models["vision-test"]["capabilities_verified"])
|
||||
self.assertFalse(models["unknown-test"]["capabilities_verified"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,19 @@
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.embeddings import EmbeddingUnavailable, embed_texts
|
||||
|
||||
|
||||
class EmbeddingTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_empty_model_setting_reports_disabled_embedding(self) -> None:
|
||||
with patch.dict("os.environ", {"KNOWLEDGE_EMBEDDING_MODEL": ""}):
|
||||
with self.assertRaises(EmbeddingUnavailable):
|
||||
await embed_texts(["sample"])
|
||||
|
||||
async def test_empty_input_needs_no_embedding_model(self) -> None:
|
||||
with patch.dict("os.environ", {"KNOWLEDGE_EMBEDDING_MODEL": ""}):
|
||||
self.assertEqual(await embed_texts([]), [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,264 @@
|
||||
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 PIL import Image, ImageDraw
|
||||
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-knowledge-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from app import knowledge
|
||||
from fastapi.testclient import TestClient
|
||||
from app.main import app
|
||||
|
||||
|
||||
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 = TestClient(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_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 = TestClient(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 = TestClient(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()
|
||||
@@ -0,0 +1,44 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from io import BytesIO
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from app.local_ocr import LocalOCRError, _sort_reading_order, recognize_image_text
|
||||
|
||||
|
||||
def box(left: int, top: int, right: int, bottom: int):
|
||||
return [[left, top], [right, top], [right, bottom], [left, bottom]]
|
||||
|
||||
|
||||
class LocalOcrTests(unittest.TestCase):
|
||||
def test_text_detections_are_ordered_by_rows_and_language_direction(self) -> None:
|
||||
detections = [
|
||||
(box(10, 100, 100, 120), "Place: Amman", 0.9),
|
||||
(box(400, 100, 520, 120), "المكان: عمان", 0.8),
|
||||
(box(10, 10, 100, 30), "COMMUNITY", 0.9),
|
||||
(box(400, 10, 520, 30), "لقاء القراءة", 0.8),
|
||||
]
|
||||
|
||||
ordered = _sort_reading_order(detections)
|
||||
|
||||
self.assertEqual(
|
||||
[text for _box, text, _confidence in ordered],
|
||||
["لقاء القراءة", "COMMUNITY", "المكان: عمان", "Place: Amman"],
|
||||
)
|
||||
|
||||
def test_image_pixel_limit_is_checked_before_loading_ocr_model(self) -> None:
|
||||
output = BytesIO()
|
||||
Image.new("RGB", (1500, 1400), "white").save(output, format="JPEG")
|
||||
|
||||
with patch("app.local_ocr._get_reader") as get_reader:
|
||||
with self.assertRaises(LocalOCRError):
|
||||
recognize_image_text(output.getvalue(), label="large test")
|
||||
|
||||
get_reader.assert_not_called()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,94 @@
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
_TEST_DATA_DIR = None
|
||||
if "SOVEREIGNAI_DATA_DIR" not in os.environ:
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-mixed-pdf-index-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app import database, knowledge
|
||||
from app.main import app
|
||||
from tests.test_pdf_analysis import make_mixed_pdf
|
||||
|
||||
|
||||
class MixedPdfKnowledgeIntegrationTests(unittest.TestCase):
|
||||
def test_indexes_and_searches_digital_and_scanned_pages_together(self) -> None:
|
||||
data_dir = Path(os.environ["SOVEREIGNAI_DATA_DIR"])
|
||||
data_dir.mkdir(parents=True, exist_ok=True)
|
||||
workspace = data_dir / "mixed-pdf-project"
|
||||
workspace.mkdir(exist_ok=True)
|
||||
pdf_path = workspace / "mixed.pdf"
|
||||
pdf_path.write_bytes(make_mixed_pdf())
|
||||
database_path = data_dir / "mixed-pdf-knowledge.sqlite3"
|
||||
client = TestClient(app)
|
||||
ocr = {
|
||||
"engine": "easyocr-local-ar-en",
|
||||
"text": "ScannedIndexMarker ArabicPageText",
|
||||
"lines": [],
|
||||
"average_confidence": 0.85,
|
||||
"elapsed_seconds": 0.1,
|
||||
}
|
||||
|
||||
async def fake_embeddings(texts, *, model=None):
|
||||
return [[1.0, 0.0] for _ in texts]
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(database, "DATABASE_PATH", database_path),
|
||||
patch.dict(os.environ, {"KNOWLEDGE_EMBEDDING_MODEL": "granite-embedding:278m"}),
|
||||
patch("app.main.embeddings.embed_texts", new=AsyncMock(side_effect=fake_embeddings)),
|
||||
):
|
||||
database.initialize_database()
|
||||
knowledge.initialize()
|
||||
with patch("app.main.recognize_image_text", return_value=ocr):
|
||||
indexed = client.post(
|
||||
"/v1/agent/knowledge/index",
|
||||
json={"workspace_path": str(workspace), "files": ["mixed.pdf"]},
|
||||
)
|
||||
digital = client.post(
|
||||
"/v1/agent/knowledge/search",
|
||||
json={"workspace_path": str(workspace), "task": "IndexMarker"},
|
||||
)
|
||||
scanned = client.post(
|
||||
"/v1/agent/knowledge/search",
|
||||
json={"workspace_path": str(workspace), "task": "ScannedIndexMarker"},
|
||||
)
|
||||
semantic = client.post(
|
||||
"/v1/agent/knowledge/search",
|
||||
json={"workspace_path": str(workspace), "task": "صياغة مختلفة بلا كلمات مشتركة"},
|
||||
)
|
||||
removed = client.request(
|
||||
"DELETE",
|
||||
"/v1/agent/knowledge/index",
|
||||
json={"workspace_path": str(workspace), "files": ["mixed.pdf"]},
|
||||
)
|
||||
|
||||
self.assertEqual(indexed.status_code, 200, indexed.text)
|
||||
self.assertEqual(indexed.json()["indexed"][0]["ocr_pages"], 1)
|
||||
self.assertTrue(indexed.json()["indexed"][0]["semantic_indexed"])
|
||||
self.assertEqual(digital.status_code, 200, digital.text)
|
||||
self.assertIn("IndexMarker", digital.json()["results"][0]["text"])
|
||||
self.assertEqual(scanned.status_code, 200, scanned.text)
|
||||
self.assertIn("ScannedIndexMarker", scanned.json()["results"][0]["text"])
|
||||
self.assertEqual(semantic.json()["search_mode"], "hybrid")
|
||||
self.assertTrue(semantic.json()["results"])
|
||||
self.assertEqual(removed.status_code, 200, removed.text)
|
||||
self.assertTrue(removed.json()["deleted"][0]["deleted"])
|
||||
self.assertFalse(
|
||||
knowledge.has_embeddings(
|
||||
user_id=database.LOCAL_USER_ID,
|
||||
workspace_path=workspace,
|
||||
model="granite-embedding:278m",
|
||||
)
|
||||
)
|
||||
finally:
|
||||
pdf_path.unlink(missing_ok=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,266 @@
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
import base64
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from PIL import Image, ImageDraw
|
||||
from pypdf import PdfReader, PdfWriter
|
||||
from fastapi.testclient import TestClient
|
||||
from app.local_ocr import LocalOCRError
|
||||
from app.pdf_documents import extract_pdf_pages_text
|
||||
|
||||
_TEST_DATA_DIR = None
|
||||
if "SOVEREIGNAI_DATA_DIR" not in os.environ:
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-pdf-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from app.main import app
|
||||
|
||||
|
||||
def make_pdf(text: str | None) -> bytes:
|
||||
stream = (
|
||||
b"BT /F1 12 Tf 72 720 Td (" + text.encode("ascii") + b") Tj ET"
|
||||
if text is not None
|
||||
else b""
|
||||
)
|
||||
objects = [
|
||||
b"<< /Type /Catalog /Pages 2 0 R >>",
|
||||
b"<< /Type /Pages /Kids [3 0 R] /Count 1 >>",
|
||||
b"<< /Type /Page /Parent 2 0 R /MediaBox [0 0 612 792] /Resources << /Font << /F1 4 0 R >> >> /Contents 5 0 R >>",
|
||||
b"<< /Type /Font /Subtype /Type1 /BaseFont /Helvetica >>",
|
||||
b"<< /Length " + str(len(stream)).encode("ascii") + b" >>\nstream\n" + stream + b"\nendstream",
|
||||
]
|
||||
data = bytearray(b"%PDF-1.4\n")
|
||||
offsets = [0]
|
||||
for number, body in enumerate(objects, 1):
|
||||
offsets.append(len(data))
|
||||
data.extend(f"{number} 0 obj\n".encode("ascii") + body + b"\nendobj\n")
|
||||
xref_offset = len(data)
|
||||
data.extend(f"xref\n0 {len(objects) + 1}\n".encode("ascii"))
|
||||
data.extend(b"0000000000 65535 f \n")
|
||||
for offset in offsets[1:]:
|
||||
data.extend(f"{offset:010d} 00000 n \n".encode("ascii"))
|
||||
data.extend(
|
||||
f"trailer\n<< /Size {len(objects) + 1} /Root 1 0 R >>\nstartxref\n{xref_offset}\n%%EOF\n".encode("ascii")
|
||||
)
|
||||
return bytes(data)
|
||||
|
||||
|
||||
def make_scanned_pdf(page_count: int = 1) -> bytes:
|
||||
pages = []
|
||||
for index in range(page_count):
|
||||
image = Image.new("RGB", (900, 1165), "white")
|
||||
draw = ImageDraw.Draw(image)
|
||||
draw.text((60, 80), f"Community Reading Meetup\nPlace: Amman\nPage: {index + 1}", fill="black")
|
||||
pages.append(image)
|
||||
output = BytesIO()
|
||||
pages[0].save(
|
||||
output,
|
||||
format="PDF",
|
||||
resolution=144,
|
||||
save_all=True,
|
||||
append_images=pages[1:],
|
||||
)
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
def make_mixed_pdf() -> bytes:
|
||||
writer = PdfWriter()
|
||||
for page_data in (make_pdf("Digital page: IndexMarker"), make_scanned_pdf()):
|
||||
writer.append_pages_from_reader(PdfReader(BytesIO(page_data)))
|
||||
output = BytesIO()
|
||||
writer.write(output)
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
class PdfAnalysisTests(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls) -> None:
|
||||
cls.client = TestClient(app)
|
||||
|
||||
def test_extracts_pdf_text_and_sends_page_number_to_model(self) -> None:
|
||||
completion = {"choices": [{"message": {"content": "يتحدث الملف عن لقاء مجتمعي في عمّان."}}]}
|
||||
with patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model:
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
data={"question": "لخص محتوى الملف"},
|
||||
files={"files": ("meetup.pdf", make_pdf("Community Meetup Amman"), "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
self.assertEqual(response.json()["files"], ["meetup.pdf"])
|
||||
sent_prompt = model.await_args.args[0]["messages"][1]["content"]
|
||||
self.assertIn("--- صفحة 1 ---", sent_prompt)
|
||||
self.assertIn("Community Meetup Amman", sent_prompt)
|
||||
self.assertIn("meetup.pdf", sent_prompt)
|
||||
|
||||
def test_extracts_text_per_page_and_marks_scanned_pages(self) -> None:
|
||||
parsed = extract_pdf_pages_text(make_mixed_pdf(), "mixed.pdf")
|
||||
|
||||
self.assertEqual(parsed["total_pages"], 2)
|
||||
self.assertTrue(parsed["pages"][0]["has_text"])
|
||||
self.assertIn("IndexMarker", parsed["pages"][0]["text"])
|
||||
self.assertFalse(parsed["pages"][1]["has_text"])
|
||||
|
||||
def test_mixed_pdf_analysis_combines_digital_text_ocr_and_page_image(self) -> None:
|
||||
provider = MagicMock()
|
||||
provider.default_model = "gemma4:e2b"
|
||||
provider.list_models = AsyncMock(return_value=["ministral-3:3b"])
|
||||
completion = {"choices": [{"message": {"content": "يلخص المستند الرقمي والممسوح."}}]}
|
||||
ocr = {
|
||||
"engine": "easyocr-local-ar-en",
|
||||
"text": "ScannedPageMarker Community Reading Meetup",
|
||||
"lines": [],
|
||||
"average_confidence": 0.8,
|
||||
"elapsed_seconds": 1.0,
|
||||
}
|
||||
with (
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.recognize_image_text", return_value=ocr),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
data={"question": "ما محتوى الملف؟"},
|
||||
files={"files": ("mixed.pdf", make_mixed_pdf(), "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
body = response.json()
|
||||
self.assertEqual(body["pages_rendered"], 1)
|
||||
self.assertIn("ScannedPageMarker", body["result"])
|
||||
content = model.await_args.args[0]["messages"][1]["content"]
|
||||
self.assertTrue(any("IndexMarker" in part.get("text", "") for part in content))
|
||||
self.assertTrue(any("ScannedPageMarker" in part.get("text", "") for part in content))
|
||||
self.assertEqual(len([part for part in content if part["type"] == "image_url"]), 1)
|
||||
|
||||
def test_scanned_pdf_is_rendered_and_routed_to_local_vision_model(self) -> None:
|
||||
provider = MagicMock()
|
||||
provider.default_model = "gemma4:e2b"
|
||||
provider.list_models = AsyncMock(return_value=["gemma4:e2b", "ministral-3:3b"])
|
||||
completion = {"choices": [{"message": {"content": "Community Reading Meetup — Amman."}}]}
|
||||
with (
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.recognize_image_text", side_effect=LocalOCRError("OCR weights absent")),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
data={"question": "اقرأ اسم اللقاء والمكان."},
|
||||
files={"files": ("meetup-scan.pdf", make_scanned_pdf(4), "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
body = response.json()
|
||||
self.assertEqual(body["tool"], "local-pdf-vision-analysis")
|
||||
self.assertEqual(body["model"], "ministral-3:3b")
|
||||
self.assertTrue(body["auto_routed"])
|
||||
self.assertEqual(body["pages_rendered"], 3)
|
||||
self.assertIn("أول 3 صفحات فقط", body["result"])
|
||||
content = model.await_args.args[0]["messages"][1]["content"]
|
||||
system = model.await_args.args[0]["messages"][0]["content"]
|
||||
images = [part for part in content if part["type"] == "image_url"]
|
||||
self.assertEqual(len(images), 3)
|
||||
self.assertIn("دون ترجمتها", system)
|
||||
image_data = images[0]["image_url"].split(",", 1)[1]
|
||||
self.assertTrue(base64.b64decode(image_data).startswith(b"\xff\xd8\xff"))
|
||||
|
||||
def test_scanned_pdf_ocr_text_is_given_to_vision_model_and_returned(self) -> None:
|
||||
provider = MagicMock()
|
||||
provider.default_model = "gemma4:e2b"
|
||||
provider.list_models = AsyncMock(return_value=["ministral-3:3b"])
|
||||
completion = {"choices": [{"message": {"content": "الإعلان عن لقاء قراءة."}}]}
|
||||
ocr = {
|
||||
"engine": "easyocr-local-ar-en",
|
||||
"text": "لقاء القراءة المجتمعي\nCommunity Reading Meetup\nPlace: Amman",
|
||||
"lines": [],
|
||||
"average_confidence": 0.75,
|
||||
"elapsed_seconds": 1.2,
|
||||
}
|
||||
with (
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.recognize_image_text", return_value=ocr),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
data={"question": "انسخ عنوان الإعلان ومكانه."},
|
||||
files={"files": ("meetup-scan.pdf", make_scanned_pdf(), "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
body = response.json()
|
||||
self.assertEqual(body["ocr_engine"], "easyocr-local-ar-en")
|
||||
self.assertEqual(body["ocr_pages"], 1)
|
||||
self.assertIn("Community Reading Meetup", body["result"])
|
||||
content = model.await_args.args[0]["messages"][1]["content"]
|
||||
self.assertTrue(any("Amman" in part.get("text", "") for part in content))
|
||||
|
||||
def test_image_ocr_text_is_fused_with_the_visual_model_request(self) -> None:
|
||||
provider = MagicMock()
|
||||
provider.default_model = "gemma4:e2b"
|
||||
provider.list_models = AsyncMock(return_value=["ministral-3:3b"])
|
||||
completion = {"choices": [{"message": {"content": "المكان عمّان."}}]}
|
||||
ocr = {
|
||||
"engine": "easyocr-local-ar-en",
|
||||
"text": "Place: Amman\nالمكان: عمان",
|
||||
"lines": [],
|
||||
"average_confidence": 0.82,
|
||||
"elapsed_seconds": 2.1,
|
||||
}
|
||||
image_output = BytesIO()
|
||||
Image.new("RGB", (100, 100), "white").save(image_output, format="PNG")
|
||||
with (
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.recognize_image_text", return_value=ocr),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/agent/images/analyze",
|
||||
data={"question": "ما المكان المكتوب؟"},
|
||||
files={"files": ("notice.png", image_output.getvalue(), "image/png")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
self.assertEqual(response.json()["ocr_engine"], "easyocr-local-ar-en")
|
||||
self.assertIn("Place: Amman", response.json()["result"])
|
||||
user_parts = model.await_args.args[0]["messages"][1]["content"]
|
||||
self.assertIn("Place: Amman", user_parts[0]["text"])
|
||||
|
||||
def test_rejects_invalid_pdf_signature(self) -> None:
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
files={"files": ("broken.pdf", b"not a pdf", "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 415)
|
||||
self.assertIn("ترويسة ملف PDF", response.json()["detail"])
|
||||
|
||||
def test_workspace_pdf_is_listed_indexed_and_searchable(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "meetup.pdf").write_bytes(make_pdf("Community Meetup Amman"))
|
||||
listed = self.client.post(
|
||||
"/v1/agent/workspace/files", json={"workspace_path": str(root)}
|
||||
)
|
||||
indexed = self.client.post(
|
||||
"/v1/agent/knowledge/index",
|
||||
json={"workspace_path": str(root), "files": ["meetup.pdf"]},
|
||||
)
|
||||
searched = self.client.post(
|
||||
"/v1/agent/knowledge/search",
|
||||
json={"workspace_path": str(root), "task": "Community Meetup"},
|
||||
)
|
||||
|
||||
self.assertEqual(listed.status_code, 200, listed.text)
|
||||
self.assertIn("meetup.pdf", listed.json()["files"])
|
||||
self.assertEqual(indexed.status_code, 200, indexed.text)
|
||||
self.assertEqual(searched.status_code, 200, searched.text)
|
||||
self.assertIn("Community Meetup Amman", searched.json()["results"][0]["text"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,96 @@
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
|
||||
# Importing the API initializes SQLite; keep this test process away from user data.
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-timeout-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from app.main import AgentRequest, run_agent_stream
|
||||
from app.model_provider import OllamaProvider
|
||||
|
||||
|
||||
class TimeoutAndCancellationTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_provider_timeout_is_bounded_and_reported_as_504(self) -> None:
|
||||
observed_timeouts: list[float] = []
|
||||
|
||||
class TimeoutClient:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_args):
|
||||
return None
|
||||
|
||||
async def post(self, *_args, **_kwargs):
|
||||
request = httpx.Request("POST", "http://127.0.0.1/chat/completions")
|
||||
raise httpx.ReadTimeout("test timeout", request=request)
|
||||
|
||||
def client_factory(*, timeout):
|
||||
observed_timeouts.append(timeout)
|
||||
return TimeoutClient()
|
||||
|
||||
provider = OllamaProvider("http://127.0.0.1:11434/v1", "test-model")
|
||||
with patch("app.model_provider.httpx.AsyncClient", side_effect=client_factory):
|
||||
with self.assertRaises(HTTPException) as caught:
|
||||
await provider.complete({"messages": []}, timeout_seconds=0.025)
|
||||
|
||||
self.assertEqual(observed_timeouts, [0.025])
|
||||
self.assertEqual(caught.exception.status_code, 504)
|
||||
|
||||
async def test_provider_cancellation_propagates_and_closes_client(self) -> None:
|
||||
started = asyncio.Event()
|
||||
closed = asyncio.Event()
|
||||
|
||||
class WaitingClient:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_args):
|
||||
closed.set()
|
||||
return None
|
||||
|
||||
async def post(self, *_args, **_kwargs):
|
||||
started.set()
|
||||
await asyncio.Future()
|
||||
|
||||
provider = OllamaProvider("http://127.0.0.1:11434/v1", "test-model")
|
||||
with patch("app.model_provider.httpx.AsyncClient", return_value=WaitingClient()):
|
||||
task = asyncio.create_task(provider.complete({"messages": []}))
|
||||
await asyncio.wait_for(started.wait(), timeout=1)
|
||||
task.cancel()
|
||||
with self.assertRaises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
self.assertTrue(closed.is_set())
|
||||
|
||||
async def test_disconnecting_agent_stream_cancels_agent_task(self) -> None:
|
||||
started = asyncio.Event()
|
||||
cancelled = asyncio.Event()
|
||||
|
||||
async def waiting_agent(_request, *, report_progress):
|
||||
await report_progress("بدأ الاختبار")
|
||||
started.set()
|
||||
try:
|
||||
await asyncio.Future()
|
||||
except asyncio.CancelledError:
|
||||
cancelled.set()
|
||||
raise
|
||||
|
||||
request = AgentRequest(task="اختبار إلغاء البث")
|
||||
with patch("app.main._execute_agent", side_effect=waiting_agent):
|
||||
response = await run_agent_stream(request)
|
||||
stream = response.body_iterator
|
||||
first_event = await asyncio.wait_for(anext(stream), timeout=1)
|
||||
self.assertIn("event: progress", first_event)
|
||||
self.assertTrue(started.is_set())
|
||||
await stream.aclose()
|
||||
await asyncio.wait_for(cancelled.wait(), timeout=1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,74 @@
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
_TEST_DATA_DIR = None
|
||||
if "SOVEREIGNAI_DATA_DIR" not in os.environ:
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-web-search-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.main import app
|
||||
|
||||
|
||||
class _FakeSearchClient:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, traceback):
|
||||
return None
|
||||
|
||||
async def get(self, *args, **kwargs):
|
||||
return SimpleNamespace(
|
||||
content=b"<html>mock search</html>",
|
||||
text="<html>mock search</html>",
|
||||
raise_for_status=lambda: None,
|
||||
)
|
||||
|
||||
|
||||
class WebSearchApiTests(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls) -> None:
|
||||
cls.client = TestClient(app)
|
||||
|
||||
def test_search_returns_per_source_and_total_fetch_durations(self) -> None:
|
||||
candidates = [
|
||||
{"title": "Source A", "url": "https://a.example/article", "snippet": "A excerpt"},
|
||||
{"title": "Source B", "url": "https://b.example/article", "snippet": "B excerpt"},
|
||||
]
|
||||
|
||||
async def read_page(url: str) -> tuple[str, str, str]:
|
||||
await asyncio.sleep(0.01)
|
||||
title = "A article" if "a.example" in url else "B article"
|
||||
return url, title, f"Content for {title}"
|
||||
|
||||
provider = MagicMock(default_model="test-model")
|
||||
completion = {"choices": [{"message": {"content": "ملخص موثق."}}]}
|
||||
with (
|
||||
patch("app.main.httpx.AsyncClient", return_value=_FakeSearchClient()),
|
||||
patch("app.main.parse_duckduckgo_results", return_value=candidates),
|
||||
patch("app.main._validate_public_http_url", side_effect=lambda url: url),
|
||||
patch("app.main._read_public_page", side_effect=read_page),
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)),
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/web/search",
|
||||
json={"query": "اختبار البحث", "max_results": 2},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
body = response.json()
|
||||
self.assertEqual(body["result"], "ملخص موثق.")
|
||||
self.assertGreaterEqual(body["source_fetch_ms"], 0)
|
||||
self.assertEqual(len(body["sources"]), 2)
|
||||
self.assertTrue(all(source["fetch_ms"] >= 0 for source in body["sources"]))
|
||||
self.assertTrue(all(source["status"] == "read" for source in body["sources"]))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,57 @@
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
_TEST_DATA_DIR = None
|
||||
if "SOVEREIGNAI_DATA_DIR" not in os.environ:
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-web-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from app.main import _PageText, _validate_public_http_url
|
||||
|
||||
|
||||
class WebSecurityTests(unittest.TestCase):
|
||||
def test_rejects_local_and_private_targets(self) -> None:
|
||||
for url in (
|
||||
"http://127.0.0.1/",
|
||||
"http://10.0.0.5/",
|
||||
"http://192.168.1.1/",
|
||||
"http://169.254.169.254/latest/meta-data/",
|
||||
"http://[::1]/",
|
||||
"http://router.local/",
|
||||
"file:///etc/passwd",
|
||||
):
|
||||
with self.subTest(url=url), self.assertRaises(HTTPException):
|
||||
_validate_public_http_url(url)
|
||||
|
||||
def test_rejects_credentials_and_unapproved_ports(self) -> None:
|
||||
for url in (
|
||||
"https://user:pass@example.com/",
|
||||
"http://example.com:8080/",
|
||||
):
|
||||
with self.subTest(url=url), self.assertRaises(HTTPException):
|
||||
_validate_public_http_url(url)
|
||||
|
||||
def test_allows_public_ip_and_preserves_url(self) -> None:
|
||||
url = "https://8.8.8.8/dns-query?q=hello"
|
||||
|
||||
self.assertEqual(_validate_public_http_url(url), url)
|
||||
|
||||
def test_html_extractor_omits_script_and_style_contents(self) -> None:
|
||||
parser = _PageText()
|
||||
parser.feed(
|
||||
"<html><head><title>Research</title><style>hidden css</style></head>"
|
||||
"<body><p>Visible text</p><script>secret instruction</script></body></html>"
|
||||
)
|
||||
text = " ".join(" ".join(parser.parts).split())
|
||||
|
||||
self.assertEqual(" ".join(parser.title.split()), "Research")
|
||||
self.assertIn("Visible text", text)
|
||||
self.assertNotIn("hidden css", text)
|
||||
self.assertNotIn("secret instruction", text)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,84 @@
|
||||
"""Tests for preview-only workspace edits and explicit one-time application."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from app import workspace
|
||||
|
||||
|
||||
class WorkspaceChangeTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.temporary = tempfile.TemporaryDirectory()
|
||||
self.root = Path(self.temporary.name).resolve()
|
||||
(self.root / "src").mkdir()
|
||||
|
||||
def tearDown(self) -> None:
|
||||
self.temporary.cleanup()
|
||||
|
||||
def test_create_preview_does_not_write_until_apply_and_is_one_time(self) -> None:
|
||||
preview = workspace.create_change_preview(
|
||||
self.root,
|
||||
"src/new.py",
|
||||
"create",
|
||||
"print('hello')\n",
|
||||
)
|
||||
target = self.root / "src" / "new.py"
|
||||
|
||||
self.assertFalse(target.exists())
|
||||
self.assertIn("+print('hello')", str(preview["diff"]))
|
||||
|
||||
result = workspace.apply_change_preview(str(preview["token"]))
|
||||
self.assertEqual(result["status"], "applied")
|
||||
self.assertEqual(target.read_text(encoding="utf-8"), "print('hello')\n")
|
||||
with self.assertRaisesRegex(ValueError, "انتهت صلاحية"):
|
||||
workspace.apply_change_preview(str(preview["token"]))
|
||||
|
||||
def test_update_preview_detects_external_changes_before_apply(self) -> None:
|
||||
target = self.root / "src" / "existing.py"
|
||||
target.write_text("old = 1\n", encoding="utf-8")
|
||||
preview = workspace.create_change_preview(
|
||||
self.root,
|
||||
"src/existing.py",
|
||||
"update",
|
||||
"new = 2\n",
|
||||
)
|
||||
self.assertIn("-old = 1", str(preview["diff"]))
|
||||
self.assertIn("+new = 2", str(preview["diff"]))
|
||||
|
||||
target.write_text("external = 3\n", encoding="utf-8")
|
||||
with self.assertRaisesRegex(ValueError, "تغير الملف"):
|
||||
workspace.apply_change_preview(str(preview["token"]))
|
||||
self.assertEqual(target.read_text(encoding="utf-8"), "external = 3\n")
|
||||
|
||||
def test_rejects_escape_hidden_unsupported_and_missing_parent_paths(self) -> None:
|
||||
invalid = [
|
||||
("../outside.py", "create", "escape"),
|
||||
(".env", "create", "hidden"),
|
||||
("src/image.png", "create", "unsupported"),
|
||||
("missing/new.py", "create", "parent does not exist"),
|
||||
]
|
||||
for relative_path, operation, _reason in invalid:
|
||||
with self.subTest(path=relative_path):
|
||||
with self.assertRaises(ValueError):
|
||||
workspace.create_change_preview(
|
||||
self.root, relative_path, operation, "content\n"
|
||||
)
|
||||
|
||||
def test_rejects_create_overwrite_and_update_of_missing_file(self) -> None:
|
||||
target = self.root / "src" / "existing.py"
|
||||
target.write_text("value = 1\n", encoding="utf-8")
|
||||
with self.assertRaisesRegex(ValueError, "موجود بالفعل"):
|
||||
workspace.create_change_preview(
|
||||
self.root, "src/existing.py", "create", "value = 2\n"
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "غير موجود"):
|
||||
workspace.create_change_preview(
|
||||
self.root, "src/missing.py", "update", "value = 2\n"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user