#!/usr/bin/env python3
"""Collect serial Ollama API timings; not a concurrency or capacity benchmark."""
import argparse
import datetime
import json
import sys
import time
import urllib.error
import urllib.request


def request_json(url, body=None, timeout=300):
    payload = None if body is None else json.dumps(body).encode("utf-8")
    request = urllib.request.Request(
        url, data=payload, headers={"Content-Type": "application/json"}
    )
    with urllib.request.urlopen(request, timeout=timeout) as response:
        result = json.load(response)
    if not isinstance(result, dict):
        raise ValueError("Expected a JSON object from Ollama")
    if result.get("error"):
        raise ValueError(str(result["error"]))
    return result


def positive_int(value):
    parsed = int(value)
    if parsed < 1:
        raise argparse.ArgumentTypeError("must be a positive integer")
    return parsed


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--url", default="http://127.0.0.1:11434")
    parser.add_argument("--model", default="qwen2.5:7b-instruct-q4_K_M")
    parser.add_argument("--runs", type=positive_int, default=3)
    parser.add_argument("--context", type=positive_int, default=4096)
    parser.add_argument("--output-tokens", type=positive_int, default=128)
    parser.add_argument("--prompt", default="Explain what a database index does in three short sentences.")
    args = parser.parse_args()
    base = args.url.rstrip("/")
    try:
        version = request_json(base + "/api/version").get("version")
        models = request_json(base + "/api/tags").get("models", [])
        digest = next((m.get("digest") for m in models if m.get("name") == args.model), None)
        for run in range(args.runs + 1):
            start = time.perf_counter()
            result = request_json(base + "/api/chat", {
                "model": args.model,
                "messages": [{"role": "user", "content": args.prompt}],
                "stream": False,
                "keep_alive": "5m",
                "options": {"num_ctx": args.context, "num_predict": args.output_tokens, "temperature": 0},
            })
            seconds = time.perf_counter() - start
            count, duration = result.get("eval_count"), result.get("eval_duration")
            if result.get("done") is not True or not isinstance(count, (int, float)) or not isinstance(duration, (int, float)) or count <= 0 or duration <= 0:
                raise ValueError("Incomplete generation or missing positive token/timing counters")
            print(json.dumps({
                "recorded_at_utc": datetime.datetime.now(datetime.timezone.utc).isoformat(),
                "phase": "warmup" if run == 0 else "measured",
                "run": run,
                "ollama_version": version,
                "model": args.model,
                "model_digest": digest,
                "context": args.context,
                "output_token_limit": args.output_tokens,
                "wall_seconds": round(seconds, 4),
                "load_seconds": result.get("load_duration", 0) / 1e9,
                "prompt_tokens": result.get("prompt_eval_count"),
                "output_tokens": count,
                "decode_tokens_per_second": round(count / duration * 1e9, 2),
                "response": result.get("message", {}).get("content", ""),
            }, ensure_ascii=False), flush=True)
    except (urllib.error.URLError, ValueError, OSError, TimeoutError) as error:
        print("Measurement failed: " + str(error), file=sys.stderr)
        return 1
    return 0


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