Record retrieval and Gemma evaluation results
This commit is contained in:
@@ -3,6 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import atexit
|
||||
import json
|
||||
import shutil
|
||||
import tempfile
|
||||
@@ -15,19 +16,37 @@ ROOT = Path(__file__).resolve().parents[1]
|
||||
DATASET = ROOT / "evals" / "knowledge_retrieval.jsonl"
|
||||
|
||||
|
||||
def post_json(url: str, payload: dict, timeout: float = 15.0) -> dict:
|
||||
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")
|
||||
request = Request(url, data=body, headers={"Content-Type": "application/json"}, method="POST")
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if token:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
request = Request(url, data=body, headers=headers, method="POST")
|
||||
with urlopen(request, timeout=timeout) as response:
|
||||
return json.loads(response.read().decode("utf-8"))
|
||||
|
||||
|
||||
def delete_json(url: str, payload: dict, timeout: float = 15.0) -> dict:
|
||||
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={"Content-Type": "application/json"},
|
||||
headers=headers,
|
||||
method="DELETE",
|
||||
)
|
||||
with urlopen(request, timeout=timeout) as response:
|
||||
@@ -39,6 +58,7 @@ def main() -> int:
|
||||
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)
|
||||
args = parser.parse_args()
|
||||
|
||||
cases = [
|
||||
@@ -50,9 +70,21 @@ def main() -> int:
|
||||
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})
|
||||
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:
|
||||
@@ -68,13 +100,31 @@ def main() -> int:
|
||||
destination.write_text(case["text"], encoding="utf-8")
|
||||
index_payload = {"workspace_path": str(workspace_path), "files": files}
|
||||
try:
|
||||
index_result = post_json(index_url, index_payload)
|
||||
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))
|
||||
@@ -93,13 +143,21 @@ def main() -> int:
|
||||
"rank": rank,
|
||||
"top_paths": ranked_paths,
|
||||
"evidence_found": evidence_found,
|
||||
"search_mode": response.get("search_mode", "unknown"),
|
||||
}
|
||||
)
|
||||
finally:
|
||||
deletion = delete_json(index_url, index_payload)
|
||||
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.")
|
||||
|
||||
logout()
|
||||
|
||||
total = len(results)
|
||||
metrics = {
|
||||
"cases": total,
|
||||
@@ -114,9 +172,10 @@ def main() -> int:
|
||||
"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; measures lexical retrieval only, not answer quality or general RAG quality.",
|
||||
"note": "Small deterministic fixture set; measures configured retrieval (keyword or hybrid) and evidence presence, not answer quality or general RAG quality.",
|
||||
}
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
output.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||
|
||||
Reference in New Issue
Block a user