Files
odysseus/scripts/eval_followup_domain_continuity.py
T
pewdiepie-archdaemon 297ad19248 Harden maintainer-preview harness and review fixes
Unify conversational domain routing, preserve artifact completion evidence, replace provisional tool-round prose with terminal synthesis, and resolve verified maintainer review findings across search, frontend module identity, path policy, configuration, and built-in skill startup.
2026-09-21 06:54:03 +00:00

210 lines
8.0 KiB
Python

#!/usr/bin/env python3
"""Evaluate follow-up domain continuity against an OpenAI-compatible API."""
from __future__ import annotations
import argparse
import asyncio
import json
import random
import sqlite3
import sys
from collections import Counter
from pathlib import Path
import httpx
ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from src.tool_policy import ToolPolicy
from src.tool_routing_experiment import MODEL_CHOICE_MODE, select_experiment_inventory
from src.tool_schemas import FUNCTION_TOOL_SCHEMAS
from src.turn_contract import (
_families_for_tool,
canonical_tool,
requested_capabilities,
resolve_full_inventory_contract,
resolve_turn_contract,
)
FAMILY_TOOL = {
"search_browser": "web_search",
"notes": "manage_notes",
"documents": "manage_documents",
"email": "search_emails",
"calendar": "manage_calendar",
"tasks": "manage_tasks",
"skills": "manage_skills",
"memory": "manage_memory",
}
SUBJECTS = [
"Rocket League", "Project Juniper", "Aurora Seven", "Kyoto Railway Museum",
"Framework Laptop", "Blue Harbor", "Atlas Report", "Orion Browser",
"Maple Invoice", "Cobalt Launch", "Sakura Booking", "Nimbus Checklist",
]
FOLLOWUPS = [
"I searched {subject} and cannot find it",
"{subject} is not showing for me",
"where is {subject} listed",
"I looked for {subject} but got nothing",
"why does {subject} not appear in the results",
"still no sign of {subject}",
"the result for {subject} seems to be missing",
"I tried again and {subject} is absent",
]
def history(subject: str, family: str) -> list[dict]:
tool = FAMILY_TOOL[family]
return [
{"role": "user", "content": f"Find {subject}"},
{"role": "assistant", "content": f"I found information about {subject}.", "metadata": {
"tool_events": [{"tool": tool, "exit_code": 0, "error": False}],
}},
]
def build_cases(count: int, seed: int) -> list[dict]:
rng = random.Random(seed)
families = tuple(FAMILY_TOOL)
cases = []
for index in range(count):
source = families[index % len(families)]
subject = f"{rng.choice(SUBJECTS)} {index + 1}"
if index % 4:
prompt = rng.choice(FOLLOWUPS).format(subject=subject)
expected = source
kind = "continuation"
else:
target = families[(families.index(source) + 1 + index) % len(families)]
noun = {
"search_browser": "the web", "notes": "my notes",
"documents": "my documents", "email": "my email",
"calendar": "my calendar", "tasks": "my tasks",
"skills": "my skills", "memory": "my memories",
}[target]
prompt = f"Search {noun} for {subject} instead"
expected = target
kind = "explicit_switch"
cases.append({
"id": index + 1, "kind": kind, "source": source,
"expected": expected, "subject": subject, "prompt": prompt,
})
rng.shuffle(cases)
return cases
def contract_for(case: dict):
prior = history(case["subject"], case["source"])
policy = ToolPolicy()
capabilities = requested_capabilities(case["prompt"], prior)
inventory = resolve_full_inventory_contract(
schemas=FUNCTION_TOOL_SCHEMAS, policy=policy,
)
routed = resolve_turn_contract(
capabilities=capabilities, schemas=FUNCTION_TOOL_SCHEMAS, policy=policy,
)
contract = select_experiment_inventory(
inventory, routed, prior, MODEL_CHOICE_MODE, user_text=case["prompt"],
)
return prior, capabilities, contract
def endpoint_from_database(path: Path, session_id: str) -> tuple[str, str, dict]:
with sqlite3.connect(f"file:{path}?mode=ro", uri=True) as db:
row = db.execute(
"select endpoint_url, model, headers from sessions where id=?", (session_id,),
).fetchone()
if not row:
raise SystemExit(f"Session {session_id} was not found in {path}")
return row[0], row[1], json.loads(row[2] or "{}")
async def run_case(client, semaphore, endpoint, model, headers, case):
prior, capabilities, contract = contract_for(case)
schemas = contract.schemas()
exposed = {
family for schema in schemas
for family in _families_for_tool(canonical_tool(schema["function"]["name"]))
}
allowed_families = {case["expected"]}
if case["expected"] == "email":
allowed_families.add("contacts")
result = {**case, "capabilities": sorted(capabilities), "exposed": sorted(exposed)}
result["contract_ok"] = bool(exposed) and exposed <= allowed_families
payload = {
"model": model,
"messages": [
{"role": "system", "content": "Continue the conversation. Use an offered tool when evidence is needed."},
*[{"role": row["role"], "content": row["content"]} for row in prior],
{"role": "user", "content": case["prompt"]},
],
"tools": schemas,
"tool_choice": "auto",
"temperature": 0.7,
"max_tokens": 96,
}
async with semaphore:
for attempt in range(4):
try:
response = await client.post(endpoint, headers=headers, json=payload)
if response.status_code == 429 and attempt < 3:
await asyncio.sleep(1.5 * (attempt + 1))
continue
response.raise_for_status()
message = response.json()["choices"][0]["message"]
calls = message.get("tool_calls") or []
tools = [canonical_tool(call["function"]["name"]) for call in calls]
called_families = sorted({
family for tool in tools for family in _families_for_tool(tool)
})
result.update({"tools": tools, "called_families": called_families})
result["cross_domain"] = bool(set(called_families) - allowed_families)
result["expected_call"] = case["expected"] in called_families
return result
except Exception as exc:
if attempt == 3:
result["error"] = f"{type(exc).__name__}: {exc}"
return result
await asyncio.sleep(0.5 * (attempt + 1))
async def main(args):
endpoint, model, headers = endpoint_from_database(args.database, args.session_id)
cases = build_cases(args.count, args.seed)
timeout = httpx.Timeout(45, connect=10)
semaphore = asyncio.Semaphore(args.concurrency)
async with httpx.AsyncClient(timeout=timeout) as client:
results = await asyncio.gather(*(
run_case(client, semaphore, endpoint, model, headers, case)
for case in cases
))
summary = {
"count": len(results),
"contract_pass": sum(bool(row.get("contract_ok")) for row in results),
"cross_domain_calls": sum(bool(row.get("cross_domain")) for row in results),
"expected_tool_calls": sum(bool(row.get("expected_call")) for row in results),
"no_tool": sum(not row.get("tools") and not row.get("error") for row in results),
"errors": sum("error" in row for row in results),
"kinds": Counter(row["kind"] for row in results),
}
output = {"summary": summary, "results": results}
args.output.write_text(json.dumps(output, indent=2, default=dict) + "\n")
print(json.dumps(summary, indent=2, default=dict))
raise SystemExit(1 if summary["contract_pass"] != len(results) or summary["cross_domain_calls"] else 0)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--count", type=int, default=500)
parser.add_argument("--concurrency", type=int, default=12)
parser.add_argument("--seed", type=int, default=20260921)
parser.add_argument("--database", type=Path, required=True)
parser.add_argument("--session-id", required=True)
parser.add_argument("--output", type=Path, required=True)
asyncio.run(main(parser.parse_args()))