"""Reference blah-core-http/1 scoring runtime (Hugging Face causal LMs).

Serves POST /v1/score and GET /healthz for the evals.blah.dev CORE benchmark.
The span logic is nanochat's core_eval.py, verbatim in spirit: answer spans are
found in TOKEN space — the common prefix across a choice set, the common suffix
across a schema set, and the without/with prefix split for greedy pairs.

    python server.py --model gpt2 --port 7990 [--device cpu]

This file doubles as the reference implementation for framework authors: the
whole contract is score_example() plus an HTTP wrapper.
"""

import argparse
import json
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

import torch
import torch.nn.functional as F
from transformers import AutoModelForCausalLM, AutoTokenizer

PROTOCOL = "blah-core-http/1"


def find_common_length(token_sequences, direction):
    """Length of the shared prefix ('left') or suffix ('right')."""
    min_len = min(len(seq) for seq in token_sequences)
    indices = {"left": range(min_len), "right": range(-1, -min_len - 1, -1)}[direction]
    for i, idx in enumerate(indices):
        token = token_sequences[0][idx]
        if not all(seq[idx] == token for seq in token_sequences):
            return i
    return min_len


class Scorer:
    def __init__(self, model_name, device):
        self.tokenizer = AutoTokenizer.from_pretrained(model_name)
        self.model = AutoModelForCausalLM.from_pretrained(model_name).to(device).eval()
        self.device = device
        # GPT-2 has no dedicated BOS; its EOS conventionally serves. Prepending
        # one gives every sequence an autoregressive target from position 0.
        self.bos = self.tokenizer.bos_token_id
        if self.bos is None:
            self.bos = self.tokenizer.eos_token_id
        self.max_len = int(self.model.config.max_position_embeddings)
        self.revision = model_name

    def tokenize(self, prompts):
        return [[self.bos] + self.tokenizer.encode(p) for p in prompts]

    @torch.no_grad()
    def forward_losses(self, sequences):
        """Per-position CE losses and argmax predictions, padded batch."""
        bsz = len(sequences)
        seq_len = max(len(s) for s in sequences)
        input_ids = torch.full((bsz, seq_len), self.bos, dtype=torch.long)
        for i, s in enumerate(sequences):
            input_ids[i, : len(s)] = torch.tensor(s, dtype=torch.long)
        input_ids = input_ids.to(self.device)
        logits = self.model(input_ids).logits
        targets = torch.roll(input_ids, shifts=-1, dims=1)
        losses = F.cross_entropy(
            logits.view(bsz * seq_len, -1), targets.view(bsz * seq_len), reduction="none"
        ).view(bsz, seq_len)
        losses[:, -1] = float("nan")
        predictions = logits.argmax(dim=-1)
        return input_ids, losses, predictions

    def crop(self, tokens, starts, ends):
        """Keep the last max_len tokens, shifting span indices (nanochat crop)."""
        out_t, out_s, out_e = [], [], []
        for t, s, e in zip(tokens, starts, ends):
            if len(t) > self.max_len:
                drop = len(t) - self.max_len
                assert s - drop >= 0 and e - drop >= 0
                out_t.append(t[-self.max_len :])
                out_s.append(s - drop)
                out_e.append(e - drop)
            else:
                out_t.append(t)
                out_s.append(s)
                out_e.append(e)
        return out_t, out_s, out_e

    def score_example(self, example):
        kind = example["kind"]
        prompts = example["prompts"]
        tokens = self.tokenize(prompts)

        if kind in ("choice_set", "suffix_set"):
            if kind == "choice_set":
                start = find_common_length(tokens, "left")
                starts = [start] * len(tokens)
                ends = [len(t) for t in tokens]
            else:
                suffix = find_common_length(tokens, "right")
                ends = [len(t) for t in tokens]
                starts = [e - suffix for e in ends]
            tokens, starts, ends = self.crop(tokens, starts, ends)
            _, losses, _ = self.forward_losses(tokens)
            mean_losses = [
                losses[i, s - 1 : e - 1].mean().item() for i, (s, e) in enumerate(zip(starts, ends))
            ]
            return {"kind": kind, "mean_losses": mean_losses}

        if kind == "greedy_pair":
            without, with_ = tokens
            assert without == with_[: len(without)], "prompt_without must prefix prompt_with"
            starts, ends = [len(without)], [len(with_)]
            seqs, starts, ends = self.crop([with_], starts, ends)
            input_ids, _, predictions = self.forward_losses(seqs)
            s, e = starts[0], ends[0]
            match = bool(torch.all(predictions[0, s - 1 : e - 1] == input_ids[0, s:e]).item())
            return {"kind": kind, "match": match}

        raise ValueError(f"unknown kind: {kind}")


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--model", default="gpt2")
    parser.add_argument("--device", default="cpu")
    parser.add_argument("--host", default="127.0.0.1")
    parser.add_argument("--port", type=int, default=7990)
    args = parser.parse_args()

    scorer = Scorer(args.model, args.device)
    print(f"[core-score] {args.model} ready on {args.host}:{args.port}", flush=True)

    class Handler(BaseHTTPRequestHandler):
        def log_message(self, fmt, *log_args):
            pass

        def _json(self, status, body):
            payload = json.dumps(body).encode()
            self.send_response(status)
            self.send_header("Content-Type", "application/json")
            self.send_header("Content-Length", str(len(payload)))
            self.end_headers()
            self.wfile.write(payload)

        def do_GET(self):
            if self.path == "/healthz":
                self._json(200, {"status": "ok", "protocol": PROTOCOL, "model_revision": scorer.revision})
            else:
                self._json(404, {"error": "not_found"})

        def do_POST(self):
            if self.path != "/v1/score":
                self._json(404, {"error": "not_found"})
                return
            try:
                length = int(self.headers.get("Content-Length", "0"))
                request = json.loads(self.rfile.read(length))
                results = [scorer.score_example(ex) for ex in request["examples"]]
                self._json(200, {"results": results})
            except Exception as err:  # noqa: BLE001 — report, don't die
                self._json(400, {"error": str(err)})

    ThreadingHTTPServer((args.host, args.port), Handler).serve_forever()


if __name__ == "__main__":
    main()
