#!/usr/bin/env python3
"""
GoNext MLX embeddings server — an OpenAI-compatible /v1/embeddings endpoint for MLX
embedding models (e.g. Qwen3-Embedding-8B), which `mlx_lm.server` does NOT provide.

The GoNext worker's RAG tools call `POST {url}/v1/embeddings` first, so pointing the
"RAG embedding server URL" (web Settings → Agent → RAG) at this server makes MLX
embeddings work with no other changes.

Run:
    python3 gonext_mlx_embed.py --model ~/mlx-models/Qwen3-Embedding-8B-4bit-DWQ --port 8085

Then set the RAG embedding server URL to http://127.0.0.1:8085

Endpoints:
    POST /v1/embeddings   {"model": <ignored>, "input": "text" | ["t1","t2",...]}
                          → {"object":"list","data":[{"object":"embedding","index":i,
                             "embedding":[...]}], "model": <name>, "usage": {...}}
    GET  /v1/models       → the single loaded model
    GET  /health          → {"status":"ok"}

Pooling: last-token hidden state (Qwen3-Embedding convention), L2-normalized.
Requires: mlx, mlx_lm (already used by the worker). No extra deps.
"""
import argparse
import json
import os
import sys
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

import mlx.core as mx
from mlx_lm import load

# Cap tokens per input so a pathologically long chunk can't OOM / stall the GPU.
MAX_TOKENS = 8192

_model = None
_tok = None
_inner = None            # the transformer body that returns hidden states
_model_name = "mlx-embedding"
_lock = threading.Lock()  # MLX eval isn't guaranteed thread-safe; serialize inference


def _load(model_path: str) -> None:
    global _model, _tok, _inner, _model_name
    print(f"[gonext-mlx-embed] loading {model_path} …", flush=True)
    _model, _tok = load(model_path)
    # mlx_lm wraps the transformer body as `.model` (returns hidden states); `_model(x)`
    # would return logits from the LM head, which we do NOT want.
    _inner = getattr(_model, "model", None) or getattr(_model, "language_model", None)
    if _inner is None:
        raise RuntimeError("Could not find the transformer body on the loaded model.")
    _model_name = os.path.basename(os.path.normpath(model_path))
    # Warm + sanity-check the embedding path so a bad model fails fast at startup.
    v = _embed_one("warmup")
    print(f"[gonext-mlx-embed] ready — {_model_name}, dim={len(v)}", flush=True)


def _embed_one(text: str) -> list:
    ids = _tok.encode(text or " ")
    if not ids:
        ids = _tok.encode(" ")
    if len(ids) > MAX_TOKENS:
        ids = ids[:MAX_TOKENS]
    h = _inner(mx.array([ids]))          # (1, seq, hidden)
    v = h[0, -1, :]                       # last-token pooling
    v = v / mx.maximum(mx.linalg.norm(v), 1e-9)   # L2 normalize
    mx.eval(v)
    return [float(x) for x in v.tolist()]


def _embed_batch(texts: list) -> list:
    with _lock:
        out = [_embed_one(t) for t in texts]
    return out


class Handler(BaseHTTPRequestHandler):
    protocol_version = "HTTP/1.1"

    def log_message(self, *a):  # quieter logs
        pass

    def _send(self, code: int, obj: dict) -> None:
        body = json.dumps(obj).encode("utf-8")
        self.send_response(code)
        self.send_header("Content-Type", "application/json")
        self.send_header("Content-Length", str(len(body)))
        self.end_headers()
        self.wfile.write(body)

    def do_GET(self):
        if self.path == "/health":
            self._send(200, {"status": "ok", "model": _model_name})
        elif self.path in ("/v1/models", "/models"):
            self._send(200, {"object": "list", "data": [
                {"id": _model_name, "object": "model", "owned_by": "gonext-mlx-embed"}
            ]})
        else:
            self._send(404, {"error": {"message": f"unknown path {self.path}"}})

    def do_POST(self):
        if self.path not in ("/v1/embeddings", "/embeddings"):
            self._send(404, {"error": {"message": f"unknown path {self.path}"}})
            return
        try:
            n = int(self.headers.get("Content-Length", 0) or 0)
            payload = json.loads(self.rfile.read(n) or b"{}")
        except Exception as e:  # noqa: BLE001
            self._send(400, {"error": {"message": f"bad JSON: {e}"}})
            return
        inp = payload.get("input")
        if inp is None:
            self._send(400, {"error": {"message": "missing 'input'"}})
            return
        texts = [inp] if isinstance(inp, str) else list(inp)
        if not all(isinstance(t, str) for t in texts):
            self._send(400, {"error": {"message": "'input' must be a string or list of strings"}})
            return
        try:
            vecs = _embed_batch(texts)
        except Exception as e:  # noqa: BLE001
            self._send(500, {"error": {"message": f"embedding failed: {type(e).__name__}: {e}"}})
            return
        total_tokens = sum(len(_tok.encode(t or " ")) for t in texts)
        self._send(200, {
            "object": "list",
            "data": [
                {"object": "embedding", "index": i, "embedding": v}
                for i, v in enumerate(vecs)
            ],
            "model": payload.get("model") or _model_name,
            "usage": {"prompt_tokens": total_tokens, "total_tokens": total_tokens},
        })


def main():
    ap = argparse.ArgumentParser(description="OpenAI-compatible MLX embeddings server.")
    ap.add_argument("--model", required=True, help="Path to the MLX embedding model dir.")
    ap.add_argument("--port", type=int, default=8085)
    ap.add_argument("--host", default="127.0.0.1")
    args = ap.parse_args()
    _load(os.path.expanduser(args.model))
    srv = ThreadingHTTPServer((args.host, args.port), Handler)
    print(f"[gonext-mlx-embed] serving /v1/embeddings on http://{args.host}:{args.port}", flush=True)
    try:
        srv.serve_forever()
    except KeyboardInterrupt:
        print("\n[gonext-mlx-embed] shutting down", flush=True)
        srv.shutdown()


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