#!/usr/bin/env python3
"""Engine comparison harness for rapidmlx.com/compare (2026-09 run).

Talks to any OpenAI-compatible server (Rapid-MLX, oMLX, Ollama, LM Studio)
over /v1/chat/completions on loopback. Standard library only, so the same file
runs against every engine from any Python >= 3.9.

Scenarios (pick with --scenario):
  a  single-stream decode tok/s (3 runs, same prompt)
  b  cold TTFT on ~2k and ~16k token prompts (unique nonce at the very start of
     every prompt, so no prefix/prompt cache can hit; 3 runs each)
  c  agent session: >=20k-token system prompt + 30 tool schemas, then 10 user
     turns replayed from a fixed transcript; per-turn TTFT. Run it once on a
     fresh server (--phase first), restart the server, warm it with an
     unrelated prompt, and run it again (--phase after_restart).
  d  4 concurrent streams, aggregate completion tok/s (2 runs)
  e  tool-call success on 30 fixed prompts (non-streaming)

Definitions
  TTFT        request sent -> first SSE chunk whose delta carries non-empty
              content, reasoning (reasoning_content / reasoning) or tool_calls.
              Role-only chunks do not count. If a stream never carries such a
              chunk, TTFT falls back to end-of-stream and is flagged.
  decode tps  (completion_tokens - 1) / (t_last_chunk - t_first_token), with
              completion_tokens taken from the engine's own usage report
              (stream_options.include_usage). Never client-estimated; a run
              without usage is recorded with decode_tps = null.
  Thinking    Qwen3.5 thinks by default; we leave every engine at its default,
              so reasoning tokens count as generated tokens for all of them.

Every request and its timings are written to the --out JSON verbatim.
"""

from __future__ import annotations

import argparse
import http.client
import json
import platform
import random
import statistics
import sys
import threading
import time
import urllib.parse
import uuid

SEED = 20260926

# ---------------------------------------------------------------- text corpus
_SENTENCES = [
    "The build system caches every compiled module by the hash of its inputs.",
    "A failing test should describe the behaviour it protects, not the code it calls.",
    "Most production incidents begin with a configuration change that looked harmless.",
    "The scheduler admits new requests only when the memory budget allows it.",
    "Logs are written as structured records so that they can be queried later.",
    "Every public function validates its arguments before touching shared state.",
    "The migration runs in small batches to keep lock times below one second.",
    "Retries use exponential backoff with jitter to avoid synchronized storms.",
    "A reviewer should be able to reproduce a benchmark from the committed script.",
    "The parser accepts both the legacy format and the new versioned envelope.",
    "Timeouts are set per call site rather than globally for the whole client.",
    "The cache key includes the tokenizer revision, because templates change.",
    "Feature flags default to off and are removed within two releases.",
    "The worker pool shrinks when the queue has been empty for thirty seconds.",
    "Error messages tell the user what happened and what they can do next.",
    "The index is rebuilt nightly and verified against a checksum manifest.",
    "Large files are streamed in chunks instead of being loaded into memory.",
    "The API returns a stable error code alongside the human-readable message.",
    "Integration tests run against a real database started in a container.",
    "The release checklist is generated from the changelog and the open issues.",
    "Metrics are sampled every ten seconds and aggregated per minute.",
    "The client library pins its dependencies to tested minor versions.",
    "A deadlock was traced to two locks acquired in opposite orders.",
    "Every endpoint documents its rate limit in the response headers.",
    "The formatter runs before the linter so that style noise never blocks review.",
    "Secrets are read from the keychain at start-up and never logged.",
    "The dashboard shows the median and the ninety-fifth percentile side by side.",
    "The fallback path is exercised by a test that forces the primary to fail.",
    "Old snapshots are pruned when they are older than the retention window.",
    "The command prints a short summary and writes the full report to disk.",
    "Configuration is loaded from a file, then environment, then flags.",
    "The benchmark discards the first run because it includes model loading.",
    "Pagination uses opaque cursors so that the ordering can change safely.",
    "The upgrade path was tested from each of the last three releases.",
    "A health check reports ready only after the model has finished warming up.",
    "The queue is drained gracefully on shutdown before the process exits.",
    "Unicode input is normalized before it is compared or stored.",
    "The profiler showed that most time was spent serializing the response.",
    "A small fixture file replaced a network call in the unit tests.",
    "The documentation example is executed in CI so that it cannot rot.",
]


def filler(n_chars: int, seed: int) -> str:
    """Deterministic English-like filler of roughly n_chars characters."""
    rng = random.Random(seed)
    out: list[str] = []
    size = 0
    para: list[str] = []
    while size < n_chars:
        s = rng.choice(_SENTENCES)
        para.append(s)
        size += len(s) + 1
        if len(para) >= 6:
            out.append(" ".join(para))
            para = []
    if para:
        out.append(" ".join(para))
    return "\n\n".join(out)


# ~5.15 chars/token for this corpus with the Qwen3.5 tokenizer (calibrated; the
# engine-reported prompt_tokens are recorded for every request regardless).
CHARS_PER_TOKEN = 5.15


def long_prompt(target_tokens: int, nonce: str, seed: int) -> str:
    body = filler(int(target_tokens * CHARS_PER_TOKEN), seed)
    return (
        f"Session {nonce}.\n\nRead the following engineering notes.\n\n{body}\n\n"
        "In one sentence, what is the most common theme in these notes?"
    )


# ------------------------------------------------------------ agent session
def agent_tools() -> list[dict]:
    rng = random.Random(SEED + 7)
    names = [
        "read_file", "write_file", "edit_file", "list_directory", "glob_search",
        "grep_search", "run_shell", "run_tests", "git_status", "git_diff",
        "git_commit", "git_log", "open_pull_request", "fetch_url", "web_search",
        "create_todo", "update_todo", "list_todos", "ask_user", "spawn_subagent",
        "read_notebook", "edit_notebook", "format_code", "lint_code",
        "type_check", "install_package", "start_server", "stop_server",
        "read_logs", "summarize_changes",
    ]
    tools = []
    for i, name in enumerate(names):
        desc = (
            f"{name.replace('_', ' ').capitalize()}. "
            + filler(2200, SEED + 100 + i).replace("\n\n", " ")
        )
        props = {
            "path": {"type": "string", "description": "Absolute path. " + rng.choice(_SENTENCES)},
            "query": {"type": "string", "description": "Search text or command. " + rng.choice(_SENTENCES)},
            "limit": {"type": "integer", "description": "Maximum number of results. " + rng.choice(_SENTENCES)},
            "dry_run": {"type": "boolean", "description": "Preview without side effects. " + rng.choice(_SENTENCES)},
        }
        tools.append({
            "type": "function",
            "function": {
                "name": name,
                "description": desc,
                "parameters": {"type": "object", "properties": props, "required": ["path"]},
            },
        })
    return tools


def agent_system() -> str:
    return (
        "You are a coding agent working in a user's repository on their Mac. "
        "Follow the project conventions below. Use the tools to inspect files "
        "before editing them. Keep answers short.\n\n# Project conventions\n\n"
        + filler(26000, SEED + 3)
    )


AGENT_TURNS = [
    ("List the files in the src directory.", "I'll look at src/ first. It holds server.py, cache.py, cli.py and a tests/ folder."),
    ("What does cache.py do?", "cache.py keeps a prefix cache keyed by token hashes and evicts least-recently-used entries."),
    ("Is there a test for eviction?", "Yes, tests/test_cache.py::test_lru_eviction fills the cache past its limit and checks the oldest entry is gone."),
    ("Run that test.", "The test passed in 0.4 seconds."),
    ("Add a test for the case where the cache is empty.", "Added test_evict_on_empty_cache, which asserts eviction on an empty cache is a no-op."),
    ("Run the whole test file.", "All 12 tests in tests/test_cache.py passed."),
    ("Show me the git diff.", "The diff adds one test function of 9 lines to tests/test_cache.py."),
    ("Write a commit message for it.", "test(cache): cover eviction on an empty cache"),
    ("Now check server.py for any TODO comments.", "server.py has two TODOs: one about request timeouts and one about logging the model name."),
    ("Summarize what we did in this session.", "We reviewed the cache module, added an empty-cache eviction test, ran the suite, and found two TODOs in server.py."),
]


# ---------------------------------------------------------- tool-call suite
_TC_TOOLS = [
    {"type": "function", "function": {"name": "get_weather", "description": "Get the current weather for a city.", "parameters": {"type": "object", "properties": {"city": {"type": "string"}, "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}}, "required": ["city"]}}},
    {"type": "function", "function": {"name": "calculate", "description": "Evaluate an arithmetic expression and return the number.", "parameters": {"type": "object", "properties": {"expression": {"type": "string"}}, "required": ["expression"]}}},
    {"type": "function", "function": {"name": "search_web", "description": "Search the web and return the top results.", "parameters": {"type": "object", "properties": {"query": {"type": "string"}, "max_results": {"type": "integer"}}, "required": ["query"]}}},
    {"type": "function", "function": {"name": "send_email", "description": "Send an email.", "parameters": {"type": "object", "properties": {"to": {"type": "string"}, "subject": {"type": "string"}, "body": {"type": "string"}}, "required": ["to", "subject", "body"]}}},
    {"type": "function", "function": {"name": "create_calendar_event", "description": "Create a calendar event.", "parameters": {"type": "object", "properties": {"title": {"type": "string"}, "date": {"type": "string", "description": "YYYY-MM-DD"}, "time": {"type": "string", "description": "HH:MM, 24h"}}, "required": ["title", "date"]}}},
    {"type": "function", "function": {"name": "read_file", "description": "Read a file from disk.", "parameters": {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]}}},
    {"type": "function", "function": {"name": "run_shell", "description": "Run a shell command and return its output.", "parameters": {"type": "object", "properties": {"command": {"type": "string"}}, "required": ["command"]}}},
    {"type": "function", "function": {"name": "convert_currency", "description": "Convert an amount between currencies.", "parameters": {"type": "object", "properties": {"amount": {"type": "number"}, "from_currency": {"type": "string"}, "to_currency": {"type": "string"}}, "required": ["amount", "from_currency", "to_currency"]}}},
]

# (prompt, expected function, {arg: substring that must appear, case-insensitive})
TOOLCALL_CASES = [
    ("What's the weather in Paris right now?", "get_weather", {"city": "paris"}),
    ("Is it raining in Tokyo? Use fahrenheit.", "get_weather", {"city": "tokyo", "unit": "fahrenheit"}),
    ("How warm is it in Buenos Aires today, in celsius?", "get_weather", {"city": "buenos aires"}),
    ("Tell me the current weather for Nairobi.", "get_weather", {"city": "nairobi"}),
    ("Weather check: Reykjavik.", "get_weather", {"city": "reykjavik"}),
    ("What is 1234 * 5678? Use the calculator.", "calculate", {"expression": "1234"}),
    ("Compute (17 + 25) / 6 with the calculator tool.", "calculate", {"expression": "17"}),
    ("Use the calculator to find 2 to the power of 20.", "calculate", {"expression": "2"}),
    ("Calculate 15% of 240 using the tool.", "calculate", {"expression": "240"}),
    ("Search the web for the latest MLX release notes.", "search_web", {"query": "mlx"}),
    ("Find me three articles about Apple Silicon memory bandwidth.", "search_web", {"query": "apple silicon"}),
    ("Look up who won the 2022 World Cup.", "search_web", {"query": "world cup"}),
    ("Search for 'continuous batching LLM inference' and return 5 results.", "search_web", {"query": "continuous batching"}),
    ("Email alice@example.com with subject 'Lunch' and say I'll be 10 minutes late.", "send_email", {"to": "alice@example.com", "subject": "lunch"}),
    ("Send bob@example.org an email titled 'Report ready' telling him the Q3 report is attached.", "send_email", {"to": "bob@example.org", "subject": "report"}),
    ("Write to team@example.com, subject 'Standup moved', body: standup is at 10:30 tomorrow.", "send_email", {"to": "team@example.com", "subject": "standup"}),
    ("Put 'Dentist' on my calendar for 2026-10-14 at 09:30.", "create_calendar_event", {"title": "dentist", "date": "2026-10-14"}),
    ("Create an event called 'Project kickoff' on 2026-11-02 at 14:00.", "create_calendar_event", {"title": "kickoff", "date": "2026-11-02"}),
    ("Schedule 'Call with Mom' for 2026-10-05.", "create_calendar_event", {"title": "mom", "date": "2026-10-05"}),
    ("Read the file /etc/hosts.", "read_file", {"path": "/etc/hosts"}),
    ("Show me what's in ~/project/README.md.", "read_file", {"path": "readme.md"}),
    ("Open /var/log/system.log and read it.", "read_file", {"path": "/var/log/system.log"}),
    ("Run `ls -la` in the current directory.", "run_shell", {"command": "ls"}),
    ("Check disk usage with df -h.", "run_shell", {"command": "df"}),
    ("Run the command `git status` for me.", "run_shell", {"command": "git status"}),
    ("How many lines are in main.py? Run wc -l main.py.", "run_shell", {"command": "wc"}),
    ("Convert 100 US dollars to euros.", "convert_currency", {"from_currency": "usd", "to_currency": "eur"}),
    ("How much is 2500 JPY in GBP?", "convert_currency", {"from_currency": "jpy", "to_currency": "gbp"}),
    ("Convert 42.5 CHF into CAD.", "convert_currency", {"from_currency": "chf", "to_currency": "cad"}),
    ("I need 75 euros converted to Australian dollars.", "convert_currency", {"from_currency": "eur", "to_currency": "aud"}),
]


# ------------------------------------------------------------------- client
class Client:
    def __init__(self, base_url: str, model: str, api_key: str | None, timeout: float, extra_body: dict | None = None):
        u = urllib.parse.urlparse(base_url)
        assert u.hostname in ("127.0.0.1", "localhost", "::1"), "loopback only"
        self.host, self.port = u.hostname, u.port or 80
        self.prefix = u.path.rstrip("/")
        self.model = model
        self.api_key = api_key
        self.timeout = timeout
        self.extra_body = extra_body or {}

    def _headers(self) -> dict:
        h = {"Content-Type": "application/json"}
        if self.api_key:
            h["Authorization"] = f"Bearer {self.api_key}"
        return h

    def chat_stream(self, messages, max_tokens, tools=None, extra=None) -> dict:
        body = {
            "model": self.model, "messages": messages, "max_tokens": max_tokens,
            "temperature": 0, "stream": True,
            "stream_options": {"include_usage": True},
        }
        if tools:
            body["tools"] = tools
        body.update(self.extra_body)
        if extra:
            body.update(extra)
        payload = json.dumps(body).encode()
        conn = http.client.HTTPConnection(self.host, self.port, timeout=self.timeout)
        t0 = time.perf_counter()
        conn.request("POST", self.prefix + "/chat/completions", payload, self._headers())
        resp = conn.getresponse()
        rec = {"http_status": resp.status, "request_bytes": len(payload)}
        if resp.status != 200:
            rec["error"] = resp.read()[:2000].decode("utf-8", "replace")
            conn.close()
            return rec
        t_first = None
        t_last = None
        usage = None
        n_chunks = 0
        text_chars = 0
        content_text: list[str] = []
        reasoning_text: list[str] = []
        finish = None
        buf = b""
        while True:
            line = resp.readline()
            if not line:
                break
            line = line.strip()
            if not line.startswith(b"data:"):
                continue
            data = line[5:].strip()
            if data == b"[DONE]":
                break
            try:
                obj = json.loads(data)
            except json.JSONDecodeError:
                buf += data
                continue
            now = time.perf_counter()
            if obj.get("usage"):
                usage = obj["usage"]
            for ch in obj.get("choices") or []:
                d = ch.get("delta") or {}
                piece = (d.get("content") or "") + (d.get("reasoning_content") or "") + (d.get("reasoning") or "")
                content_text.append(d.get("content") or "")
                reasoning_text.append((d.get("reasoning_content") or "") + (d.get("reasoning") or ""))
                if piece or d.get("tool_calls"):
                    n_chunks += 1
                    text_chars += len(piece)
                    if t_first is None:
                        t_first = now
                    t_last = now
                if ch.get("finish_reason"):
                    finish = ch["finish_reason"]
        t_end = time.perf_counter()
        conn.close()
        rec.update({
            "ttft_s": (t_first - t0) if t_first is not None else (t_end - t0),
            "ttft_fallback_end_of_stream": t_first is None,
            "total_s": t_end - t0,
            "gen_window_s": (t_last - t_first) if (t_first and t_last) else None,
            "content_chunks": n_chunks,
            "text_chars": text_chars,
            "content_text": "".join(content_text),
            "reasoning_text": "".join(reasoning_text),
            "finish_reason": finish,
            "usage": usage,
            "t_start_perf": t0,
            "t_end_perf": t_end,
        })
        ct = (usage or {}).get("completion_tokens")
        if ct and rec["gen_window_s"] and ct > 1:
            rec["decode_tps"] = (ct - 1) / rec["gen_window_s"]
        else:
            rec["decode_tps"] = None
        return rec

    def chat(self, messages, max_tokens, tools=None) -> dict:
        body = {"model": self.model, "messages": messages, "max_tokens": max_tokens, "temperature": 0, "stream": False}
        if tools:
            body["tools"] = tools
        body.update(self.extra_body)
        conn = http.client.HTTPConnection(self.host, self.port, timeout=self.timeout)
        t0 = time.perf_counter()
        conn.request("POST", self.prefix + "/chat/completions", json.dumps(body).encode(), self._headers())
        resp = conn.getresponse()
        raw = resp.read()
        dt = time.perf_counter() - t0
        conn.close()
        try:
            obj = json.loads(raw)
        except json.JSONDecodeError:
            obj = {"unparsed": raw[:2000].decode("utf-8", "replace")}
        return {"http_status": resp.status, "latency_s": dt, "response": obj}


def med(xs):
    xs = [x for x in xs if x is not None]
    return statistics.median(xs) if xs else None


# ---------------------------------------------------------------- scenarios
def warmup(c: Client) -> dict:
    return c.chat_stream([{"role": "user", "content": f"Say hello. ({uuid.uuid4().hex[:8]})"}], 16)


def scen_a(c: Client) -> dict:
    msgs = [{"role": "user", "content": "Write a detailed, 800-word essay on the history of the bicycle, from the 1817 draisine to modern carbon frames."}]
    runs = [c.chat_stream(msgs, 512) for _ in range(3)]
    return {"runs": runs, "median_decode_tps": med([r.get("decode_tps") for r in runs]), "max_tokens": 512}


def scen_b(c: Client) -> dict:
    out = {}
    for label, target in (("2k", 2000), ("16k", 16000)):
        runs = []
        for i in range(3):
            nonce = uuid.uuid4().hex
            p = long_prompt(target, nonce, SEED + target + i)
            r = c.chat_stream([{"role": "user", "content": p}], 16)
            r["nonce"] = nonce
            runs.append(r)
        out[label] = {"runs": runs, "median_ttft_s": med([r.get("ttft_s") for r in runs]),
                      "prompt_tokens": [((r.get("usage") or {}).get("prompt_tokens")) for r in runs]}
    return out


def scen_c(c: Client) -> dict:
    system = agent_system()
    tools = agent_tools()
    msgs = [{"role": "system", "content": system}]
    turns = []
    for i, (user, canned) in enumerate(AGENT_TURNS):
        msgs.append({"role": "user", "content": user})
        try:
            r = c.chat_stream(list(msgs), 16, tools=tools)
        except Exception as exc:  # e.g. the server died mid-session: record it, stop the session
            turns.append({"turn": i + 1, "error": f"{type(exc).__name__}: {exc}", "ttft_s": None})
            break
        r["turn"] = i + 1
        turns.append(r)
        msgs.append({"role": "assistant", "content": canned})
    return {"turns": turns, "ttft_s": [t.get("ttft_s") for t in turns],
            "prompt_tokens": [((t.get("usage") or {}).get("prompt_tokens")) for t in turns]}


def scen_d(c: Client) -> dict:
    topics = ["the history of the printing press", "how tides work", "the life cycle of a star", "the invention of the telephone"]
    reps = []
    for _rep in range(2):
        results = [None] * 4

        def worker(i):
            results[i] = c.chat_stream([{"role": "user", "content": f"Write a detailed 600-word explainer on {topics[i]}."}], 256)

        ths = [threading.Thread(target=worker, args=(i,)) for i in range(4)]
        t0 = time.perf_counter()
        for t in ths:
            t.start()
        for t in ths:
            t.join()
        t1 = time.perf_counter()
        toks = [((r.get("usage") or {}).get("completion_tokens")) for r in results]
        agg = (sum(toks) / (t1 - t0)) if all(toks) else None
        reps.append({"streams": results, "wall_s": t1 - t0, "completion_tokens": toks, "aggregate_tps": agg})
    return {"reps": reps, "median_aggregate_tps": med([r["aggregate_tps"] for r in reps]), "max_tokens": 256}


def grade_toolcall(resp: dict, fn: str, want: dict) -> tuple[bool, str]:
    try:
        msg = resp["choices"][0]["message"]
    except (KeyError, IndexError, TypeError):
        return False, "no choices"
    calls = msg.get("tool_calls") or []
    if not calls:
        return False, "no tool_calls"
    f = calls[0].get("function") or {}
    if f.get("name") != fn:
        return False, f"wrong function {f.get('name')!r}"
    args = f.get("arguments")
    if isinstance(args, str):
        try:
            args = json.loads(args)
        except json.JSONDecodeError:
            return False, "arguments not valid JSON"
    if not isinstance(args, dict):
        return False, "arguments not an object"
    for k, sub in want.items():
        v = args.get(k)
        if v is None or sub.lower() not in str(v).lower():
            return False, f"arg {k}={v!r} missing {sub!r}"
    return True, "ok"


def scen_e(c: Client) -> dict:
    cases = []
    for prompt, fn, want in TOOLCALL_CASES:
        r = c.chat([{"role": "user", "content": prompt}], 4096, tools=_TC_TOOLS)
        ok, why = grade_toolcall(r.get("response") or {}, fn, want)
        cases.append({"prompt": prompt, "expected": fn, "want_args": want, "pass": ok, "reason": why, **r})
    n = sum(1 for x in cases if x["pass"])
    return {"cases": cases, "passed": n, "total": len(cases)}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--engine", required=True)
    ap.add_argument("--base-url", required=True, help="e.g. http://127.0.0.1:8000/v1")
    ap.add_argument("--model", required=True)
    ap.add_argument("--api-key")
    ap.add_argument("--scenario", required=True, choices=list("abcde"))
    ap.add_argument("--phase", default="first", help="label for scenario c: first | after_restart")
    ap.add_argument("--no-warmup", action="store_true")
    ap.add_argument("--timeout", type=float, default=900)
    ap.add_argument("--thinking", choices=["default", "on"], default="default",
                    help="'on' sends chat_template_kwargs.enable_thinking=true on every request, "
                         "so engines that turn thinking off by default (Rapid-MLX does for casual requests) "
                         "generate the same kind of output as the rest")
    ap.add_argument("--out", required=True)
    a = ap.parse_args()
    extra = {"chat_template_kwargs": {"enable_thinking": True}} if a.thinking == "on" else None
    c = Client(a.base_url, a.model, a.api_key, a.timeout, extra)
    rec = {"engine": a.engine, "model": a.model, "scenario": a.scenario, "phase": a.phase, "thinking": a.thinking,
           "started_unix": time.time(), "python": platform.python_version()}
    if not a.no_warmup:
        rec["warmup"] = warmup(c)
    fn = {"a": scen_a, "b": scen_b, "c": scen_c, "d": scen_d, "e": scen_e}[a.scenario]
    rec["result"] = fn(c)
    rec["finished_unix"] = time.time()
    with open(a.out, "w") as f:
        json.dump(rec, f, indent=1)
    summary = {k: v for k, v in rec["result"].items() if not isinstance(v, (list, dict))}
    if a.scenario == "b":
        summary = {k: v["median_ttft_s"] for k, v in rec["result"].items()}
    if a.scenario == "c":
        summary = {"ttft_s": [None if x is None else round(x, 3) for x in rec["result"]["ttft_s"]], "prompt_tokens": rec["result"]["prompt_tokens"]}
    print(json.dumps({"engine": a.engine, "scenario": a.scenario, "phase": a.phase, **summary}))


if __name__ == "__main__":
    sys.exit(main())
