#!/usr/bin/env python3
"""
OpenRouter -> Ollama Fallback Proxy
Listens on 127.0.0.1:11435, exposes an OpenAI-compatible /v1/chat/completions endpoint.

Strategy:
  1. Try OpenRouter models in order (primary -> secondary)
  2. On 429 (rate limit) or other transient error -> try next model
  3. If ALL OpenRouter models fail -> fall back to local Ollama
  4. Cooldown tracking to avoid hammering rate-limited models
"""
import json
import os
import sys
import time
import asyncio
import logging
from pathlib import Path
from http.server import HTTPServer, BaseHTTPRequestHandler
import urllib.request
import urllib.error

# Load config
CFG_FILE = Path.home() / ".hermes" / "openrouter_proxy.json"
ENV_FILE = Path.home() / ".hermes" / "openrouter.env"
STATE_FILE = Path.home() / ".hermes" / "openrouter_proxy_state.json"

if not CFG_FILE.exists():
    sys.exit(f"Config not found: {CFG_FILE}")

cfg = json.loads(CFG_FILE.read_text())

# Load API key
api_key = None
for line in ENV_FILE.read_text().splitlines():
    if line.startswith("OPENROUTER_API_KEY="):
        api_key = line.split("=", 1)[1].strip()
        break
if not api_key:
    sys.exit("OPENROUTER_API_KEY not found in env file")

# State: cooldowns per model
state = {"cooldowns": {}}
if STATE_FILE.exists():
    try:
        state = json.loads(STATE_FILE.read_text())
    except Exception:
        state = {"cooldowns": {}}

cooldown_seconds = cfg.get("rate_limit_cooldown_seconds", 60)
all_openrouter_models = cfg["openrouter_models_primary"] + cfg["openrouter_models_secondary"]

logging.basicConfig(level=logging.INFO,
                    format="%(asctime)s %(levelname)s: %(message)s")
log = logging.getLogger("openrouter_proxy")


def sanitize(model_id: str) -> str:
    """Hermes rejects model IDs with ':' or '/' — sanitize for Hermes-compatible clients."""
    return model_id.replace(":", "_").replace("/", "-")


# Build the sanitized -> real map at module load
ALL_MODEL_MAP = {}
for m in all_openrouter_models:
    ALL_MODEL_MAP[sanitize(m)] = m
ALL_MODEL_MAP[sanitize(cfg["ollama_model"])] = cfg["ollama_model"]


def is_cooling_down(model: str) -> bool:
    until = state["cooldowns"].get(model, 0)
    return time.time() < until


def put_on_cooldown(model: str):
    state["cooldowns"][model] = time.time() + cooldown_seconds
    STATE_FILE.write_text(json.dumps(state, indent=2))


def cleanup_cooldowns():
    now = time.time()
    state["cooldowns"] = {m: t for m, t in state["cooldowns"].items() if t > now}
    STATE_FILE.write_text(json.dumps(state, indent=2))


async def call_openrouter(model: str, body: dict) -> tuple[int, dict]:
    """Returns (status_code, json_response). Caller inspects for 429."""
    url = "https://openrouter.ai/api/v1/chat/completions"
    payload = json.dumps({**body, "model": model}).encode("utf-8")
    req = urllib.request.Request(
        url, data=payload, method="POST",
        headers={
            "Authorization": f"Bearer {api_key}",
            "Content-Type": "application/json",
            "HTTP-Referer": "https://hermes-agent.local",
            "X-Title": "Hermes Agent OpenRouter Proxy",
        }
    )
    try:
        loop = asyncio.get_event_loop()
        resp = await loop.run_in_executor(
            None,
            lambda: urllib.request.urlopen(req, timeout=120),
        )
        return resp.status, json.loads(resp.read().decode("utf-8"))
    except urllib.error.HTTPError as e:
        try:
            err_body = json.loads(e.read().decode("utf-8"))
        except Exception:
            err_body = {"error": str(e)}
        return e.code, err_body
    except Exception as e:
        return 599, {"error": str(e)}


async def call_ollama(body: dict) -> tuple[int, dict]:
    """Forward to local Ollama (OpenAI-compatible endpoint)."""
    url = cfg["ollama_url"] + "/chat/completions"
    payload = json.dumps({**body, "model": cfg["ollama_model"]}).encode("utf-8")
    req = urllib.request.Request(url, data=payload, method="POST",
                                  headers={"Content-Type": "application/json"})
    try:
        loop = asyncio.get_event_loop()
        resp = await loop.run_in_executor(
            None,
            lambda: urllib.request.urlopen(req, timeout=600),  # local is slow
        )
        return resp.status, json.loads(resp.read().decode("utf-8"))
    except urllib.error.HTTPError as e:
        try:
            err_body = json.loads(e.read().decode("utf-8"))
        except Exception:
            err_body = {"error": str(e)}
        return e.code, err_body
    except Exception as e:
        return 599, {"error": str(e)}


async def handle_chat(body: dict) -> tuple[int, dict]:
    """Main handler: try OpenRouter models, fall back to Ollama."""
    cleanup_cooldowns()

    # Map sanitized name back to real OpenRouter/Ollama model ID
    requested_sanitized = body.get("model", "")
    real_model = ALL_MODEL_MAP.get(requested_sanitized, requested_sanitized)
    log.info(f"Chat request: {requested_sanitized!r} -> real: {real_model!r}")

    # Build ordered list: prefer requested model first, then fallbacks
    models_to_try = []
    if real_model in all_openrouter_models:
        models_to_try.append(real_model)
    for m in all_openrouter_models:
        if m not in models_to_try:
            models_to_try.append(m)

    for model in models_to_try:
        if is_cooling_down(model):
            log.info(f"Skip {model} (cooldown)")
            continue
        log.info(f"Trying OpenRouter: {model}")
        status, resp = await call_openrouter(model, body)
        if status == 200:
            log.info(f"OK from {model}")
            return status, resp
        if status == 429:
            log.warning(f"429 from {model}, putting on cooldown")
            put_on_cooldown(model)
            continue
        log.warning(f"Error {status} from {model}: {resp.get('error', '')[:100]}")
        continue

    log.warning("All OpenRouter models failed/rate-limited. Falling back to local Ollama.")
    status, resp = await call_ollama(body)
    if status == 200:
        # Annotate that we used the fallback
        if "choices" in resp and resp["choices"]:
            resp["choices"][0].setdefault("message", {})["content"] = \
                "[ollama-fallback] " + (resp["choices"][0].get("message", {}).get("content", "") or "")
    return status, resp


# HTTP handler
class ProxyHandler(BaseHTTPRequestHandler):
    def log_message(self, format, *args):
        log.info(f"{self.address_string()} - {format % args}")

    def do_POST(self):
        log.info(f"do_POST called: path={self.path}, headers={dict(self.headers)}")
        if self.path not in ("/v1/chat/completions", "/chat/completions"):
            self.send_error(404, "Only /v1/chat/completions supported")
            return
        length = int(self.headers.get("Content-Length", 0))
        try:
            raw = self.rfile.read(length)
            body = json.loads(raw)
            log.info(f"INCOMING body: {body}")
        except Exception as e:
            log.error(f"Bad JSON: {e}")
            self.send_error(400, f"Bad JSON: {e}")
            return

        # Run async handler
        loop = asyncio.new_event_loop()
        asyncio.set_event_loop(loop)
        try:
            status, resp = loop.run_until_complete(handle_chat(body))
        finally:
            loop.close()

        self.send_response(status)
        self.send_header("Content-Type", "application/json")
        payload = json.dumps(resp).encode("utf-8")
        self.send_header("Content-Length", str(len(payload)))
        self.end_headers()
        self.wfile.write(payload)

    def do_GET(self):
        # Use the global module-level sanitize + map
        if self.path in ("/v1/models", "/api/v1/models"):
            self.send_response(200)
            self.send_header("Content-Type", "application/json")
            # Include all the standard OpenAI/Ollama model fields a strict client might check
            payload = json.dumps({
                "object": "list",
                "data": [
                    {"id": sanitized, "object": "model", "created": 1700000000,
                     "owned_by": "ollama" if sanitized == "qwen3_4b" else "openrouter",
                     "name": sanitized, "model": sanitized,
                     "modified_at": "2026-08-01T00:00:00Z",
                     "size": 4000000000,
                     "details": {"format": "gguf", "family": "llama" if sanitized == "qwen3_4b" else "gemma",
                                 "parameter_size": "4B" if sanitized == "qwen3_4b" else "31B",
                                 "quantization_level": "Q4_K_M"}}
                    for sanitized in ALL_MODEL_MAP.keys()
                ]
            }).encode("utf-8")
            self.send_header("Content-Length", str(len(payload)))
            self.end_headers()
            self.wfile.write(payload)
            return
        # Ollama-style tags list (for clients that check it)
        if self.path == "/api/tags":
            self.send_response(200)
            self.send_header("Content-Type", "application/json")
            payload = json.dumps({
                "models": [
                    {"name": cfg["ollama_model"], "model": cfg["ollama_model"], "modified_at": "2026-01-01T00:00:00Z", "size": 2500000000}
                ]
            }).encode("utf-8")
            self.send_header("Content-Length", str(len(payload)))
            self.end_headers()
            self.wfile.write(payload)
            return
        # Version (Ollama-style)
        if self.path in ("/version", "/props", "/api/version"):
            self.send_response(200)
            self.send_header("Content-Type", "application/json")
            payload = json.dumps({"version": "0.0.1", "name": "openrouter-proxy"}).encode("utf-8")
            self.send_header("Content-Length", str(len(payload)))
            self.end_headers()
            self.wfile.write(payload)
            return
        # Health / root
        if self.path in ("/", "/health"):
            self.send_response(200)
            self.send_header("Content-Type", "application/json")
            payload = json.dumps({
                "status": "ok",
                "models_available": len(all_openrouter_models),
                "cooldowns": state["cooldowns"],
            }).encode("utf-8")
            self.send_header("Content-Length", str(len(payload)))
            self.end_headers()
            self.wfile.write(payload)
            return
        self.send_error(404)


def main():
    host, port = cfg["proxy_listen"].split(":")
    port = int(port)
    server = HTTPServer((host, port), ProxyHandler)
    log.info(f"OpenRouter Proxy listening on http://{host}:{port}")
    log.info(f"Models: {all_openrouter_models}")
    log.info(f"Ollama fallback: {cfg['ollama_url']} ({cfg['ollama_model']})")
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        log.info("Stopping...")
        server.shutdown()


if __name__ == "__main__":
    main()
