Files
sovereign_ai/SovereignAI-Starter/scripts/eval_knowledge_retrieval.py
T

284 lines
12 KiB
Python

"""Evaluate local FTS retrieval with temporary Arabic documents and API cleanup."""
from __future__ import annotations
import argparse
import atexit
import json
import shutil
import sys
import tempfile
from datetime import datetime, timezone
from pathlib import Path
from urllib.error import HTTPError
from urllib.request import Request, urlopen
ROOT = Path(__file__).resolve().parents[1]
DATASET = ROOT / "evals" / "knowledge_retrieval.jsonl"
def post_json(
url: str,
payload: dict,
timeout: float = 15.0,
*,
token: str | None = None,
) -> dict:
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
headers = {"Content-Type": "application/json"}
if token:
headers["Authorization"] = f"Bearer {token}"
request = Request(url, data=body, headers=headers, method="POST")
try:
with urlopen(request, timeout=timeout) as response:
return json.loads(response.read().decode("utf-8"))
except HTTPError as error:
detail = error.read().decode("utf-8", errors="replace")
raise RuntimeError(f"HTTP {error.code} from {url}: {detail[:1000]}") from error
def delete_json(
url: str,
payload: dict,
timeout: float = 15.0,
*,
token: str | None = None,
) -> dict:
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
headers = {"Content-Type": "application/json"}
if token:
headers["Authorization"] = f"Bearer {token}"
request = Request(
url,
data=body,
headers=headers,
method="DELETE",
)
try:
with urlopen(request, timeout=timeout) as response:
return json.loads(response.read().decode("utf-8"))
except HTTPError as error:
detail = error.read().decode("utf-8", errors="replace")
raise RuntimeError(f"HTTP {error.code} from {url}: {detail[:1000]}") from error
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--base-url", default="http://127.0.0.1:8000")
parser.add_argument("--dataset", type=Path, default=DATASET)
parser.add_argument("--output", type=Path, default=None)
parser.add_argument("--request-timeout", type=float, default=180.0)
parser.add_argument("--case", action="append", dest="case_ids")
parser.add_argument(
"--agent-answers",
action="store_true",
help="Also ask the local agent to answer each selected case from indexed knowledge.",
)
args = parser.parse_args()
cases = [
json.loads(line)
for line in args.dataset.read_text(encoding="utf-8").splitlines()
if line.strip()
]
if args.case_ids:
cases = [case for case in cases if case["id"] in args.case_ids]
missing = set(args.case_ids) - {case["id"] for case in cases}
if missing:
parser.error(f"Unknown case ID(s): {', '.join(sorted(missing))}")
if not cases or len({case["id"] for case in cases}) != len(cases):
raise ValueError("Retrieval dataset must be non-empty with unique IDs.")
results: list[dict] = []
index_summary: dict[str, object] = {}
files = sorted({case["document"] for case in cases})
projects_url = args.base_url.rstrip("/") + "/v1/agent/projects"
index_url = args.base_url.rstrip("/") + "/v1/agent/knowledge/index"
search_url = args.base_url.rstrip("/") + "/v1/agent/knowledge/search"
base_url = args.base_url.rstrip("/")
session = post_json(base_url + "/v1/auth/local-session", {})
token = session["access_token"]
def logout() -> None:
try:
post_json(base_url + "/v1/auth/logout", {}, token=token)
except Exception:
pass
atexit.register(logout)
with tempfile.TemporaryDirectory(prefix="sovereignai-rag-eval-") as temporary:
workspace_path = Path(temporary)
for case in cases:
destination = workspace_path / case["document"]
destination.parent.mkdir(parents=True, exist_ok=True)
source_path = case.get("source_path")
if source_path:
source = (ROOT / source_path).resolve(strict=True)
if not source.is_relative_to(ROOT) or not source.is_file():
raise ValueError(f"Evaluation source must be a file inside the project: {source_path}")
shutil.copyfile(source, destination)
elif not destination.exists():
destination.write_text(case["text"], encoding="utf-8")
index_payload = {"workspace_path": str(workspace_path), "files": files}
index_attempted = False
project_registration_attempted = False
project_registered = False
try:
# A temporary directory is outside the default workspace grant. Register
# it explicitly through the same local-user API used by the desktop app.
project_registration_attempted = True
registration = post_json(
projects_url,
{"workspace_path": str(workspace_path)},
timeout=args.request_timeout,
token=token,
)
project_registered = True
if Path(registration.get("path", "")).resolve() != workspace_path.resolve():
raise RuntimeError("The API registered a different fixture workspace.")
index_attempted = True
index_result = post_json(
index_url, index_payload, timeout=args.request_timeout, token=token
)
if len(index_result.get("indexed", [])) != len(files):
raise RuntimeError("The API did not confirm every fixture document.")
indexed_documents = index_result["indexed"]
index_summary = {
"semantic_indexed_documents": sum(
bool(item.get("semantic_indexed")) for item in indexed_documents
),
"documents": len(indexed_documents),
"embedding_models": sorted(
{
item["embedding_model"]
for item in indexed_documents
if item.get("embedding_model")
}
),
}
for case in cases:
response = post_json(
search_url,
{"workspace_path": str(workspace_path), "task": case["query"]},
timeout=args.request_timeout,
token=token,
)
retrieved = response.get("results", [])
ranked_paths = list(dict.fromkeys(item["path"] for item in retrieved))
rank = ranked_paths.index(case["document"]) + 1 if case["document"] in ranked_paths else None
evidence_found = any(
case["expected_evidence"].casefold() in item.get("text", "").casefold()
for item in retrieved
if item.get("path") == case["document"]
)
result = {
"id": case["id"],
"query": case["query"],
"expected_document": case["document"],
"expected_evidence": case["expected_evidence"],
"rank": rank,
"top_paths": ranked_paths,
"expected_document_chunks": [
item.get("chunk")
for item in retrieved
if item.get("path") == case["document"]
],
"evidence_found": evidence_found,
"search_mode": response.get("search_mode", "unknown"),
}
if args.agent_answers:
agent_response = post_json(
base_url + "/v1/agent/run",
{
"workspace_path": str(workspace_path),
"task": (
"ابحث في فهرس المعرفة المحلي عن المقاطع التي تجيب عن السؤال، ثم أجب "
"استنادًا إلى المقاطع فقط واذكر الملف ورقم المقطع. "
f"السؤال: {case['query']}"
),
},
timeout=args.request_timeout,
token=token,
)
result["agent_answer"] = agent_response.get(
"result", agent_response.get("answer", "")
)
result["agent_tool"] = agent_response.get("tool")
result["agent_files"] = agent_response.get("files", [])
result["agent_steps"] = agent_response.get("steps", [])
result["human_rating"] = None
results.append(result)
finally:
original_error = sys.exc_info()[0] is not None
cleanup_errors: list[Exception] = []
if index_attempted:
try:
deletion = delete_json(
index_url,
index_payload,
timeout=args.request_timeout,
token=token,
)
if len(deletion.get("deleted", [])) != len(files):
raise RuntimeError(
"The API did not confirm cleanup for every fixture document."
)
except Exception as cleanup_error:
cleanup_errors.append(cleanup_error)
if project_registration_attempted:
try:
revocation = delete_json(
projects_url,
{"workspace_path": str(workspace_path)},
timeout=args.request_timeout,
token=token,
)
if project_registered and not revocation.get("removed"):
raise RuntimeError(
"The API did not confirm removal of the temporary project registration."
)
except Exception as cleanup_error:
cleanup_errors.append(cleanup_error)
if cleanup_errors:
if original_error:
for cleanup_error in cleanup_errors:
print(f"Cleanup also failed: {cleanup_error}", file=sys.stderr)
else:
raise RuntimeError(
"; ".join(str(error) for error in cleanup_errors)
) from cleanup_errors[0]
logout()
total = len(results)
metrics = {
"cases": total,
"hit_at_1": sum(item["rank"] == 1 for item in results) / total,
"hit_at_3": sum(item["rank"] is not None and item["rank"] <= 3 for item in results) / total,
"mean_reciprocal_rank": sum(1 / item["rank"] if item["rank"] else 0 for item in results) / total,
"evidence_rate": sum(item["evidence_found"] for item in results) / total,
}
timestamp = datetime.now(timezone.utc)
output = args.output or ROOT / "evals" / "results" / f"retrieval_{timestamp.strftime('%Y-%m-%d_%H%M%S')}.json"
report = {
"created_at_utc": timestamp.isoformat(),
"base_url": args.base_url,
"dataset": str(args.dataset.relative_to(ROOT) if args.dataset.is_relative_to(ROOT) else args.dataset),
"indexing": index_summary,
"metrics": metrics,
"results": results,
"note": (
"Small deterministic fixture set. Retrieval metrics cover source ranking and evidence presence, "
"not general RAG quality. Optional agent answers are retained for human review and are not auto-scored."
),
}
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
print(json.dumps({"output": str(output), "metrics": metrics}, ensure_ascii=False, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())