#!/usr/bin/env python3
"""codex_proxy.py — OpenAI chat/completions → ChatGPT Codex Responses proxy.

agent.py stays on /v1/chat/completions. This process:
  • serves GET  /v1/models
  • serves POST /v1/chat/completions  (streaming SSE, tools)
  • runs official Codex PKCE OAuth (same client as Codex CLI / OpenCode)
  • translates chat messages/tools ↔ Responses API
  • forwards to https://chatgpt.com/backend-api/codex/responses

Usage:
  python3 codex_proxy.py              # listen 127.0.0.1:8787
  python3 codex_proxy.py --login      # browser OAuth, save tokens, exit
  python3 codex_proxy.py --port 9000

Then point agent.py at the proxy:
  /set api http://127.0.0.1:8787/v1
  /set api_key dummy
  /set model gpt-5-codex
"""

from __future__ import annotations

import argparse, base64, hashlib, http.server, json, os, secrets, sys, threading, time
import urllib.error, urllib.parse, urllib.request, webbrowser
from typing import Any

HOST = "127.0.0.1"
PORT = 8787
AUTH_PATH = os.path.join(os.path.expanduser("~"), ".agentfiles", "codex_auth.json")
CODEX_AUTH_FALLBACK = os.path.join(os.path.expanduser("~"), ".codex", "auth.json")
CODEX_URL = "https://chatgpt.com/backend-api/codex/responses"
CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann"
AUTH_URL = "https://auth.openai.com/oauth/authorize"
TOKEN_URL = "https://auth.openai.com/oauth/token"
REDIRECT = "http://localhost:1455/auth/callback"
SCOPE = "openid profile email offline_access"
ORIGINATOR = "codex_proxy"
CODEX_MODELS_URL = "https://chatgpt.com/backend-api/codex/models"
# agent.py reads n_ctx from /props and context_length from /models
CONTEXT_LENGTH = 1_000_000
# Cached live model ids from Codex (filled by fetch_available_models)
_models_cache: list = []
_models_cache_at: float = 0.0
_MODELS_CACHE_TTL = 300.0

# ---------------------------------------------------------------------------
# Auth
# ---------------------------------------------------------------------------

def _jwt_claims(token: str) -> dict:
    try:
        part = (token or "").split(".")[1]
        part += "=" * (-len(part) % 4)
        return json.loads(base64.urlsafe_b64decode(part.encode("ascii")))
    except Exception:
        return {}


def account_id_from(access: str, id_token: str = "") -> str:
    for tok in (access, id_token):
        c = _jwt_claims(tok)
        auth = c.get("https://api.openai.com/auth")
        if isinstance(auth, dict) and auth.get("chatgpt_account_id"):
            return str(auth["chatgpt_account_id"])
        if c.get("chatgpt_account_id"):
            return str(c["chatgpt_account_id"])
    return ""


def load_auth(path: str | None = None) -> dict | None:
    for p in ([path] if path else [AUTH_PATH, CODEX_AUTH_FALLBACK]):
        if not p or not os.path.isfile(p):
            continue
        try:
            data = json.load(open(p, encoding="utf8"))
        except (OSError, json.JSONDecodeError):
            continue
        tokens = data.get("tokens") if isinstance(data.get("tokens"), dict) else data
        access = (tokens.get("access_token") or data.get("access_token")
                  or data.get("access") or "").strip()
        if not access:
            continue
        idt = (tokens.get("id_token") or data.get("id_token") or "").strip()
        return {
            "access_token": access,
            "refresh_token": (tokens.get("refresh_token") or data.get("refresh_token")
                              or data.get("refresh") or "").strip(),
            "id_token": idt,
            "account_id": (tokens.get("account_id") or data.get("account_id")
                           or data.get("accountId") or account_id_from(access, idt)),
            "path": p,
        }
    return None


def save_auth(tokens: dict, path: str = AUTH_PATH) -> dict:
    os.makedirs(os.path.dirname(path), exist_ok=True)
    access = tokens.get("access_token") or ""
    idt = tokens.get("id_token") or ""
    out = {
        "access_token": access,
        "refresh_token": tokens.get("refresh_token") or "",
        "id_token": idt,
        "account_id": tokens.get("account_id") or account_id_from(access, idt),
        "expires_at": int(time.time()) + int(tokens.get("expires_in") or 3600),
        "updated_at": time.strftime("%Y-%m-%d %H:%M:%S"),
    }
    with open(path, "w", encoding="utf8") as f:
        json.dump(out, f, indent=2)
    try:
        os.chmod(path, 0o600)
    except OSError:
        pass
    return out


def _post_form(url: str, fields: dict) -> dict:
    data = urllib.parse.urlencode(fields).encode()
    req = urllib.request.Request(
        url, data=data, method="POST",
        headers={"Content-Type": "application/x-www-form-urlencoded", "Accept": "application/json"},
    )
    with urllib.request.urlopen(req, timeout=60) as r:
        return json.loads(r.read().decode())


def refresh_token(refresh: str) -> dict:
    body = _post_form(TOKEN_URL, {
        "grant_type": "refresh_token",
        "client_id": CLIENT_ID,
        "refresh_token": refresh,
    })
    if not body.get("access_token"):
        raise RuntimeError(f"refresh failed: {body}")
    if not body.get("refresh_token"):
        body["refresh_token"] = refresh
    return body


def ensure_auth() -> dict:
    auth = load_auth()
    if not auth:
        raise RuntimeError("No tokens. Run: python3 codex_proxy.py --login")
    claims = _jwt_claims(auth["access_token"])
    exp = claims.get("exp")
    if isinstance(exp, (int, float)) and exp < time.time() + 60:
        if not auth.get("refresh_token"):
            raise RuntimeError("Access token expired and no refresh_token. Re-run --login.")
        body = refresh_token(auth["refresh_token"])
        auth = save_auth(body)
        print("[codex_proxy] refreshed access token", flush=True)
    return auth


def _auth_headers(auth: dict) -> dict:
    h = {
        "Authorization": f"Bearer {auth['access_token']}",
        "Content-Type": "application/json",
        "OpenAI-Beta": "responses=experimental",
        "originator": ORIGINATOR,
        "User-Agent": "codex_proxy/1.0",
        "Accept": "application/json",
    }
    if auth.get("account_id"):
        h["ChatGPT-Account-ID"] = auth["account_id"]
    return h


def _extract_model_ids(payload) -> list[str]:
    """Parse Codex /models JSON into a sorted list of model id/slug strings."""
    items = []
    if isinstance(payload, list):
        items = payload
    elif isinstance(payload, dict):
        for key in ("models", "data", "items"):
            if isinstance(payload.get(key), list):
                items = payload[key]
                break
        else:
            # single nested object
            if "slug" in payload or "id" in payload:
                items = [payload]
    ids: list[str] = []
    seen: set[str] = set()
    for it in items:
        if isinstance(it, str):
            mid = it.strip()
        elif isinstance(it, dict):
            mid = (it.get("slug") or it.get("id") or it.get("name") or "").strip()
            # skip hidden models when the catalog marks them
            vis = it.get("visibility")
            if vis is not None and str(vis).lower() in ("hidden", "none", "false"):
                continue
        else:
            continue
        if mid and mid not in seen:
            seen.add(mid)
            ids.append(mid)
    return sorted(ids)


def fetch_available_models(force: bool = False) -> list[str]:
    """Query https://chatgpt.com/backend-api/codex/models for this account's catalog."""
    global _models_cache, _models_cache_at
    now = time.time()
    if not force and _models_cache and (now - _models_cache_at) < _MODELS_CACHE_TTL:
        return list(_models_cache)
    auth = ensure_auth()
    urls = [
        CODEX_MODELS_URL,
        CODEX_MODELS_URL + "?client_version=0.0.0",
    ]
    last_err = None
    for url in urls:
        try:
            req = urllib.request.Request(url, headers=_auth_headers(auth), method="GET")
            with urllib.request.urlopen(req, timeout=30) as r:
                raw = r.read().decode("utf8", "replace")
            payload = json.loads(raw)
            ids = _extract_model_ids(payload)
            if ids:
                _models_cache = ids
                _models_cache_at = now
                print(f"[codex_proxy] fetched {len(ids)} models from {url}", flush=True)
                return list(ids)
            last_err = f"empty catalog from {url}: {raw[:200]}"
        except Exception as e:
            last_err = str(e)
            print(f"[codex_proxy] models fetch failed ({url}): {e}", flush=True)
    if _models_cache:
        print(f"[codex_proxy] using stale model cache ({len(_models_cache)}); last error: {last_err}",
              flush=True)
        return list(_models_cache)
    raise RuntimeError(f"Could not fetch Codex models: {last_err}")


def _parse_callback_query(query: str) -> tuple[str | None, str | None, str | None]:
    """Return (code, state, error) from a query string or full redirect URL."""
    if "://" in query or query.startswith("/"):
        query = urllib.parse.urlparse(query).query
    q = urllib.parse.parse_qs(query)
    code = (q.get("code") or [None])[0]
    st = (q.get("state") or [None])[0]
    err = None
    if q.get("error"):
        err = (q.get("error_description") or q["error"])[0]
    return code, st, err


def _exchange_code(code: str, verifier: str) -> dict:
    body = _post_form(TOKEN_URL, {
        "grant_type": "authorization_code",
        "client_id": CLIENT_ID,
        "code": code,
        "redirect_uri": REDIRECT,
        "code_verifier": verifier,
    })
    if not body.get("access_token"):
        raise RuntimeError(f"token exchange failed: {body}")
    return save_auth(body)


def oauth_login(timeout: int = 300) -> dict:
    """PKCE login with a real callback server + manual paste fallback.

    OpenAI redirects to http://localhost:1455/auth/callback?code=…&state=…
    After workspace selection the browser must hit that URL on *this* machine.
    If the page spins (WSL/SSH/remote), copy the address bar URL and paste it here.
    """
    verifier = secrets.token_urlsafe(64)
    challenge = base64.urlsafe_b64encode(
        hashlib.sha256(verifier.encode("ascii")).digest()
    ).decode("ascii").rstrip("=")
    state = secrets.token_urlsafe(16)
    params = {
        "response_type": "code",
        "client_id": CLIENT_ID,
        "redirect_uri": REDIRECT,
        "scope": SCOPE,
        "code_challenge": challenge,
        "code_challenge_method": "S256",
        "state": state,
        "id_token_add_organizations": "true",
        "codex_cli_simplified_flow": "true",
        "originator": ORIGINATOR,
    }
    url = AUTH_URL + "?" + urllib.parse.urlencode(params)
    result: dict[str, Any] = {"code": None, "state": None, "error": None}
    done = threading.Event()

    class Handler(http.server.BaseHTTPRequestHandler):
        def log_message(self, fmt, *args):
            sys.stderr.write("[oauth-callback] " + (fmt % args) + "\n")

        def _reply(self, html: bytes, code: int = 200):
            self.send_response(code)
            self.send_header("Content-Type", "text/html; charset=utf-8")
            self.send_header("Content-Length", str(len(html)))
            self.send_header("Connection", "close")
            self.end_headers()
            self.wfile.write(html)

        def do_GET(self):
            parsed = urllib.parse.urlparse(self.path)
            path = parsed.path.rstrip("/") or "/"
            # Favicon / stray probes must not finish the flow
            if path not in ("/auth/callback", "/callback", "/"):
                if "code=" not in (parsed.query or ""):
                    self._reply(b"ok", 200)
                    return
            code, st, err = _parse_callback_query(parsed.query)
            print(f"[oauth-callback] path={path!r} code={'yes' if code else 'no'} "
                  f"state_match={st == state} error={err!r}", flush=True)
            if err:
                result["error"] = err
                self._reply(b"<html><body><h2>Login failed</h2><pre>"
                            + err.encode("utf8", "replace") + b"</pre></body></html>")
                done.set()
                return
            if not code:
                # Landing without code (e.g. GET /) — keep waiting
                self._reply(b"<html><body><p>Waiting for OAuth redirect...</p></body></html>")
                return
            result["code"] = code
            result["state"] = st
            self._reply(
                b"<html><body style='font-family:sans-serif'>"
                b"<h2>codex_proxy login complete</h2>"
                b"<p>You can close this tab and return to the terminal.</p>"
                b"</body></html>"
            )
            done.set()

        def do_HEAD(self):
            self.send_response(200)
            self.end_headers()

    # Prefer dual-stack when available so localhost → ::1 still hits us
    httpd = None
    last_err = None
    for bind_host in ("127.0.0.1", "localhost", "0.0.0.0"):
        try:
            httpd = http.server.ThreadingHTTPServer((bind_host, 1455), Handler)
            httpd.allow_reuse_address = True
            print(f"[oauth] callback listening on http://{bind_host}:1455/auth/callback", flush=True)
            break
        except OSError as e:
            last_err = e
            httpd = None
    if httpd is None:
        raise RuntimeError(
            f"Port 1455 busy or unbindable ({last_err}). "
            "Stop other Codex/OpenCode login servers, or paste the redirect URL manually when prompted."
        )

    def serve():
        while not done.is_set():
            httpd.timeout = 0.5
            try:
                httpd.handle_request()
            except Exception as e:
                if not done.is_set():
                    print(f"[oauth] serve error: {e}", flush=True)

    threading.Thread(target=serve, daemon=True).start()

    print("\n=== Codex OAuth ===", flush=True)
    print("1. Browser should open. Sign in and select your workspace.", flush=True)
    print("2. You should land on http://localhost:1455/auth/callback?code=…", flush=True)
    print("3. If the browser spins or shows connection refused, copy the FULL", flush=True)
    print("   URL from the address bar and paste it below, then press Enter.", flush=True)
    print(f"\nAuth URL:\n{url}\n", flush=True)
    try:
        webbrowser.open(url)
    except Exception:
        pass

    def wait_for_paste():
        # stdin paste races the HTTP callback; whichever finishes first wins
        try:
            line = input("Paste redirect URL (or press Enter to wait for browser):\n> ").strip()
        except EOFError:
            return
        if not line or done.is_set():
            return
        code, st, err = _parse_callback_query(line)
        if err:
            result["error"] = err
            done.set()
            return
        if code:
            result["code"] = code
            result["state"] = st or state
            print("[oauth] accepted pasted redirect URL", flush=True)
            done.set()
        else:
            print("[oauth] no code= in that input; still waiting for browser…", flush=True)

    paste_thread = threading.Thread(target=wait_for_paste, daemon=True)
    paste_thread.start()

    if not done.wait(timeout):
        try:
            httpd.server_close()
        except Exception:
            pass
        raise TimeoutError(
            "OAuth timed out — callback never received.\n"
            "Common causes: WSL/SSH (browser is on another host), port 1455 blocked,\n"
            "or another app holding 1455. Re-run and paste the redirect URL, or run login\n"
            "on the same machine as the browser."
        )
    try:
        httpd.server_close()
    except Exception:
        pass

    if result["error"]:
        raise RuntimeError(result["error"])
    if not result["code"]:
        raise RuntimeError("OAuth finished without an authorization code")
    if result["state"] and result["state"] != state:
        raise RuntimeError(
            f"OAuth state mismatch (got {result['state']!r}, expected {state!r}). "
            "Another login may be using port 1455."
        )
    print("[oauth] exchanging code for tokens…", flush=True)
    return _exchange_code(result["code"], verifier)


# ---------------------------------------------------------------------------
# Chat ↔ Responses translation
# ---------------------------------------------------------------------------

def _text(content) -> str:
    if content is None:
        return ""
    if isinstance(content, str):
        return content
    if isinstance(content, list):
        parts = []
        for p in content:
            if isinstance(p, dict):
                parts.append(p.get("text") or "")
            else:
                parts.append(str(p))
        return "".join(parts)
    return str(content)


def chat_to_responses(body: dict) -> dict:
    instructions = None
    items = []
    for m in body.get("messages") or []:
        role = m.get("role")
        if role == "system":
            t = _text(m.get("content"))
            if t:
                instructions = f"{instructions}\n\n{t}" if instructions else t
        elif role == "user":
            items.append({"role": "user", "content": _text(m.get("content"))})
        elif role == "assistant":
            t = _text(m.get("content"))
            if t:
                items.append({"role": "assistant", "content": t})
            for tc in m.get("tool_calls") or []:
                fn = tc.get("function") or {}
                items.append({
                    "type": "function_call",
                    "call_id": tc.get("id") or "",
                    "name": fn.get("name") or "",
                    "arguments": fn.get("arguments") or "{}",
                })
        elif role == "tool":
            items.append({
                "type": "function_call_output",
                "call_id": m.get("tool_call_id") or "",
                "output": _text(m.get("content")),
            })
    out: dict[str, Any] = {
        "model": body.get("model") or "gpt-5-codex",
        "input": items,
        "stream": True,
        "store": False,
    }
    if instructions:
        out["instructions"] = instructions
    tools = body.get("tools") or []
    if tools:
        rtools = []
        for t in tools:
            if t.get("type") == "function" and "function" in t:
                fn = t["function"]
                rtools.append({
                    "type": "function",
                    "name": fn.get("name", ""),
                    "description": fn.get("description") or "",
                    "parameters": fn.get("parameters") or {"type": "object", "properties": {}},
                })
            elif t.get("type") == "function" and "name" in t:
                rtools.append(t)
        if rtools:
            out["tools"] = rtools
            out["tool_choice"] = body.get("tool_choice") or "auto"
            out["parallel_tool_calls"] = True
    return out


def sse_line(obj: dict) -> bytes:
    return f"data: {json.dumps(obj, ensure_ascii=False)}\n\n".encode("utf-8")


def _emit_text(wfile, model: str, text: str):
    if not text:
        return
    wfile.write(sse_line({
        "id": "chatcmpl-proxy",
        "object": "chat.completion.chunk",
        "model": model,
        "choices": [{"index": 0, "delta": {"content": text}, "finish_reason": None}],
    }))
    wfile.flush()


def _emit_done(wfile, model: str, finish: str, usage: dict | None = None):
    chunk = {
        "id": "chatcmpl-proxy",
        "object": "chat.completion.chunk",
        "model": model,
        "choices": [{"index": 0, "delta": {}, "finish_reason": finish}],
    }
    if usage:
        chunk["usage"] = usage
    wfile.write(sse_line(chunk))
    wfile.write(b"data: [DONE]\n\n")
    wfile.flush()


def stream_codex_as_chat(codex_body: dict, auth: dict, wfile, model: str):
    headers = {
        "Content-Type": "application/json",
        "Authorization": f"Bearer {auth['access_token']}",
        "OpenAI-Beta": "responses=experimental",
        "originator": ORIGINATOR,
        "User-Agent": "codex_proxy/1.0",
    }
    if auth.get("account_id"):
        headers["ChatGPT-Account-ID"] = auth["account_id"]

    # Role message for agent.py (expects assistant role on first chunk sometimes)
    wfile.write(sse_line({
        "id": "chatcmpl-proxy",
        "object": "chat.completion.chunk",
        "model": model,
        "choices": [{"index": 0, "delta": {"role": "assistant"}, "finish_reason": None}],
    }))
    wfile.flush()

    print(f"[codex] → model={codex_body.get('model')} account={auth.get('account_id')!r} "
          f"tools={len(codex_body.get('tools') or [])} input_items={len(codex_body.get('input') or [])}",
          flush=True)

    req = urllib.request.Request(CODEX_URL, json.dumps(codex_body).encode(), headers)
    content = ""
    calls: dict[int, dict] = {}
    call_index: dict[str, int] = {}
    usage: dict = {}
    seen_types: list[str] = []
    raw_preview: list[str] = []

    try:
        with urllib.request.urlopen(req, timeout=600) as r:
            print(f"[codex] ← HTTP {getattr(r, 'status', 200)}", flush=True)
            for raw in r:
                line = raw.decode("utf8", "replace").rstrip()
                if len(raw_preview) < 8 and line:
                    raw_preview.append(line[:300])
                if not line:
                    continue
                # Some gateways send "event: foo" then "data: {...}"
                if line.startswith("event:"):
                    continue
                if line.startswith("data:"):
                    payload = line[5:].lstrip()
                else:
                    # bare JSON line
                    payload = line
                if payload == "[DONE]":
                    break
                try:
                    ev = json.loads(payload)
                except json.JSONDecodeError:
                    continue
                if not isinstance(ev, dict):
                    continue
                et = ev.get("type") or ""
                if et and et not in seen_types:
                    seen_types.append(et)
                    print(f"[codex] event {et}", flush=True)

                if et in ("response.output_text.delta", "response.content_part.delta"):
                    delta = ev.get("delta") or ""
                    if isinstance(delta, dict):
                        delta = delta.get("text") or delta.get("content") or ""
                    content += delta
                    _emit_text(wfile, model, delta)
                elif et == "response.output_text.done":
                    # full text if deltas were skipped
                    full = ev.get("text") or ""
                    if full and not content:
                        content = full
                        _emit_text(wfile, model, full)
                elif et == "response.output_item.added":
                    item = ev.get("item") or {}
                    if item.get("type") == "function_call":
                        idx = len(calls)
                        cid = item.get("call_id") or item.get("id") or f"call_{idx}"
                        call_index[item.get("id") or cid] = idx
                        call_index[cid] = idx
                        name = item.get("name") or ""
                        calls[idx] = {
                            "id": cid, "type": "function",
                            "function": {"name": name, "arguments": item.get("arguments") or ""},
                        }
                        wfile.write(sse_line({
                            "id": "chatcmpl-proxy",
                            "object": "chat.completion.chunk",
                            "model": model,
                            "choices": [{"index": 0, "delta": {
                                "tool_calls": [{"index": idx, "id": cid, "type": "function",
                                                "function": {"name": name, "arguments": ""}}],
                            }, "finish_reason": None}],
                        }))
                        wfile.flush()
                    elif item.get("type") == "message":
                        for part in item.get("content") or []:
                            if part.get("type") in ("output_text", "text") and part.get("text"):
                                if not content:
                                    content += part["text"]
                                    _emit_text(wfile, model, part["text"])
                elif et == "response.function_call_arguments.delta":
                    item_id = ev.get("item_id") or ""
                    idx = call_index.get(item_id)
                    if idx is None and calls:
                        idx = max(calls)
                    if idx is not None and idx in calls:
                        d = ev.get("delta") or ""
                        calls[idx]["function"]["arguments"] += d
                        wfile.write(sse_line({
                            "id": "chatcmpl-proxy",
                            "object": "chat.completion.chunk",
                            "model": model,
                            "choices": [{"index": 0, "delta": {
                                "tool_calls": [{"index": idx, "function": {"arguments": d}}],
                            }, "finish_reason": None}],
                        }))
                        wfile.flush()
                elif et == "response.function_call_arguments.done":
                    item_id = ev.get("item_id") or ""
                    idx = call_index.get(item_id)
                    if idx is not None and idx in calls:
                        if ev.get("arguments") is not None:
                            calls[idx]["function"]["arguments"] = ev["arguments"]
                        if ev.get("name"):
                            calls[idx]["function"]["name"] = ev["name"]
                elif et == "response.output_item.done":
                    item = ev.get("item") or {}
                    if item.get("type") == "function_call":
                        cid = item.get("call_id") or item.get("id") or ""
                        idx = call_index.get(item.get("id") or cid)
                        if idx is None:
                            idx = len(calls)
                            call_index[item.get("id") or cid] = idx
                            call_index[cid] = idx
                            calls[idx] = {"id": cid, "type": "function",
                                          "function": {"name": "", "arguments": ""}}
                        if item.get("name"):
                            calls[idx]["function"]["name"] = item["name"]
                        if item.get("arguments") is not None:
                            calls[idx]["function"]["arguments"] = item["arguments"]
                        if item.get("call_id"):
                            calls[idx]["id"] = item["call_id"]
                    elif item.get("type") == "message":
                        for part in item.get("content") or []:
                            if part.get("type") in ("output_text", "text") and part.get("text"):
                                if part["text"] not in content:
                                    content += part["text"]
                                    _emit_text(wfile, model, part["text"])
                elif et == "response.completed":
                    resp = ev.get("response") or {}
                    if "usage" in resp:
                        u = resp["usage"]
                        usage = {
                            "prompt_tokens": u.get("input_tokens", 0),
                            "completion_tokens": u.get("output_tokens", 0),
                            "total_tokens": u.get("total_tokens", 0),
                        }
                    if not content and not calls:
                        for item in resp.get("output") or []:
                            if item.get("type") == "message":
                                for part in item.get("content") or []:
                                    if part.get("type") in ("output_text", "text"):
                                        t = part.get("text") or ""
                                        content += t
                                        _emit_text(wfile, model, t)
                            elif item.get("type") == "function_call":
                                idx = len(calls)
                                calls[idx] = {
                                    "id": item.get("call_id") or f"call_{idx}",
                                    "type": "function",
                                    "function": {
                                        "name": item.get("name") or "",
                                        "arguments": item.get("arguments") or "{}",
                                    },
                                }
                                wfile.write(sse_line({
                                    "id": "chatcmpl-proxy",
                                    "object": "chat.completion.chunk",
                                    "model": model,
                                    "choices": [{"index": 0, "delta": {
                                        "tool_calls": [{
                                            "index": idx,
                                            "id": calls[idx]["id"],
                                            "type": "function",
                                            "function": dict(calls[idx]["function"]),
                                        }],
                                    }, "finish_reason": None}],
                                }))
                                wfile.flush()
                elif et in ("response.failed", "error"):
                    err = ev.get("error") or ev
                    msg = f"[codex_proxy upstream error] {err}"
                    print(msg, flush=True)
                    content += msg
                    _emit_text(wfile, model, msg)
    except urllib.error.HTTPError as e:
        err = e.read().decode("utf8", "replace")
        msg = f"[codex_proxy HTTP {e.code}] {err[:2000]}"
        print(msg, flush=True)
        _emit_text(wfile, model, msg)
        _emit_done(wfile, model, "stop")
        return
    except Exception as e:
        msg = f"[codex_proxy error] {e}"
        print(msg, flush=True)
        _emit_text(wfile, model, msg)
        _emit_done(wfile, model, "stop")
        return

    if not content and not calls:
        diag = (
            f"[codex_proxy] empty upstream response. events={seen_types or '(none)'} "
            f"preview={raw_preview[:3]}"
        )
        print(diag, flush=True)
        _emit_text(wfile, model, diag)

    print(f"[codex] done content_len={len(content)} tool_calls={len(calls)} usage={usage}",
          flush=True)
    _emit_done(wfile, model, "tool_calls" if calls else "stop", usage or None)


# ---------------------------------------------------------------------------
# HTTP server
# ---------------------------------------------------------------------------

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

    def log_message(self, fmt, *args):
        sys.stderr.write("[codex_proxy] " + (fmt % args) + "\n")

    def _read_json(self) -> dict:
        n = int(self.headers.get("Content-Length") or 0)
        raw = self.rfile.read(n) if n else b"{}"
        return json.loads(raw.decode("utf8") or "{}")

    def _send(self, code: int, body: bytes, content_type: str = "application/json"):
        self.send_response(code)
        self.send_header("Content-Type", content_type)
        self.send_header("Content-Length", str(len(body)))
        self.send_header("Access-Control-Allow-Origin", "*")
        self.end_headers()
        self.wfile.write(body)

    def do_OPTIONS(self):
        self.send_response(204)
        self.send_header("Access-Control-Allow-Origin", "*")
        self.send_header("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
        self.send_header("Access-Control-Allow-Headers", "*")
        self.end_headers()

    def do_GET(self):
        path = urllib.parse.urlparse(self.path).path.rstrip("/") or "/"
        if path in ("/health", "/v1/health"):
            self._send(200, b'{"ok":true}')
        elif path in ("/props", "/v1/props"):
            # llama-server–style props; agent.py uses n_ctx for context_size
            props = {
                "n_ctx": CONTEXT_LENGTH,
                "default_generation_settings": {"n_ctx": CONTEXT_LENGTH},
            }
            self._send(200, json.dumps(props).encode())
        elif path in ("/v1/models", "/models"):
            try:
                ids = fetch_available_models()
            except Exception as e:
                err = json.dumps({"error": {"message": str(e), "type": "models_error"}}).encode()
                self._send(502, err)
                return
            models = {
                "object": "list",
                "data": [
                    {
                        "id": m,
                        "object": "model",
                        "owned_by": "openai",
                        "context_length": CONTEXT_LENGTH,
                        "top_provider": {"context_length": CONTEXT_LENGTH},
                    }
                    for m in ids
                ],
            }
            self._send(200, json.dumps(models).encode())
        else:
            self._send(404, b'{"error":"not found"}')

    def do_POST(self):
        path = urllib.parse.urlparse(self.path).path.rstrip("/") or "/"
        if path not in ("/v1/chat/completions", "/chat/completions"):
            self._send(404, b'{"error":"not found"}')
            return
        try:
            body = self._read_json()
            auth = ensure_auth()
            codex_body = chat_to_responses(body)
            model = body.get("model") or "gpt-5-codex"
            self.close_connection = True
            self.send_response(200)
            self.send_header("Content-Type", "text/event-stream; charset=utf-8")
            self.send_header("Cache-Control", "no-cache")
            self.send_header("Access-Control-Allow-Origin", "*")
            self.send_header("Connection", "close")
            self.end_headers()
            stream_codex_as_chat(codex_body, auth, self.wfile, model)
        except Exception as e:
            try:
                msg = json.dumps({"error": {"message": str(e), "type": "proxy_error"}}).encode()
                self._send(500, msg)
            except Exception:
                try:
                    self.wfile.write(sse_line({"error": {"message": str(e)}}))
                    self.wfile.write(b"data: [DONE]\n\n")
                except Exception:
                    pass


def main():
    global AUTH_PATH
    ap = argparse.ArgumentParser(description="Chat Completions → Codex Responses proxy")
    ap.add_argument("--host", default=HOST)
    ap.add_argument("--port", type=int, default=PORT)
    ap.add_argument("--login", action="store_true", help="OAuth login and exit")
    ap.add_argument("--auth-file", default=AUTH_PATH)
    args = ap.parse_args()
    AUTH_PATH = args.auth_file

    if args.login:
        auth = oauth_login()
        print(f"Saved tokens to {AUTH_PATH}")
        print(f"account={auth.get('account_id')}")
        return

    if not load_auth():
        print("No tokens found. Run: python3 codex_proxy.py --login", file=sys.stderr)
        sys.exit(1)

    class QuietServer(http.server.ThreadingHTTPServer):
        allow_reuse_address = True

    server = QuietServer((args.host, args.port), Handler)
    print(f"codex_proxy listening on http://{args.host}:{args.port}/v1")
    print(f"auth file: {AUTH_PATH}")
    print("agent.py:")
    print("  /set api http://127.0.0.1:%d/v1" % args.port)
    print("  /set api_key dummy")
    try:
        ids = fetch_available_models(force=True)
        print(f"  # {len(ids)} models from Codex catalog:")
        for m in ids:
            print(f"  /set model {m}")
    except Exception as e:
        print(f"  # could not list models: {e}")
        print("  # retry: curl http://127.0.0.1:%d/v1/models" % args.port)
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        print("\nbye")


if __name__ == "__main__":
    main()
