""" LiteLLM custom provider — OpenAI Codex via chatgpt.com/backend-api. Auth tokens are read from CODEX_AUTH_FILE (default /root/.codex/auth.json) and auto-refreshed using the OpenAI OAuth endpoint before expiry. """ import base64 import json import os import time from typing import AsyncIterator, Callable, Iterator, Optional, Union import httpx import litellm from litellm.llms.custom_llm import CustomLLM, CustomLLMError from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.utils import GenericStreamingChunk, Usage from litellm.utils import ModelResponse AUTH_FILE = os.environ.get("CODEX_AUTH_FILE", "/root/.codex/auth.json") # Self-register so the Router can validate this provider before custom_llm_setup() runs if "codex" not in litellm.provider_list: litellm.provider_list.append("codex") if "codex" not in litellm._custom_providers: litellm._custom_providers.append("codex") TOKEN_URL = "https://auth.openai.com/oauth/token" CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann" API_URL = "https://chatgpt.com/backend-api/codex/responses" # --------------------------------------------------------------------------- # Auth helpers # --------------------------------------------------------------------------- def _load_auth() -> dict: with open(AUTH_FILE) as f: return json.load(f) def _save_auth(data: dict) -> None: with open(AUTH_FILE, "w") as f: json.dump(data, f, indent=2) def _jwt_exp(token: str) -> int: try: payload = token.split(".")[1] payload += "=" * (4 - len(payload) % 4) return json.loads(base64.urlsafe_b64decode(payload)).get("exp", 0) except Exception: return 0 def _refresh(auth: dict) -> dict: resp = httpx.post( TOKEN_URL, json={ "client_id": CLIENT_ID, "grant_type": "refresh_token", "refresh_token": auth["tokens"]["refresh_token"], }, timeout=30, ) resp.raise_for_status() data = resp.json() auth["tokens"]["access_token"] = data["access_token"] if "refresh_token" in data: auth["tokens"]["refresh_token"] = data["refresh_token"] auth["last_refresh"] = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()) _save_auth(auth) return auth def _get_auth() -> dict: auth = _load_auth() if time.time() > _jwt_exp(auth["tokens"]["access_token"]) - 60: auth = _refresh(auth) return auth # --------------------------------------------------------------------------- # Request / response helpers # --------------------------------------------------------------------------- def _split_messages(messages: list) -> tuple[str, list]: """Extract system/developer messages as instructions; return (instructions, input_messages).""" instructions_parts = [] input_msgs = [] for msg in messages: role = msg.get("role", "user") content = msg.get("content", "") if role in ("system", "developer"): instructions_parts.append(content) else: input_msgs.append({"role": role, "content": content}) return "\n\n".join(instructions_parts), input_msgs def _resolve_model(model: str) -> str: """Translate internal cx-* names back to the real Codex API model names (gpt-*).""" if model.startswith("cx-"): return "gpt-" + model[3:] return model def _build_body(model: str, messages: list, optional_params: dict, stream: bool) -> dict: instructions, input_msgs = _split_messages(messages) body: dict = { "model": _resolve_model(model), "instructions": instructions or "You are a helpful assistant.", "input": input_msgs, "stream": True, # API only supports streaming; non-stream calls aggregate internally "store": False, "reasoning": {"effort": optional_params.get("reasoning_effort", "medium")}, } return body def _auth_headers(auth: dict) -> dict: return { "Authorization": f"Bearer {auth['tokens']['access_token']}", "ChatGPT-Account-Id": auth["tokens"]["account_id"], "Content-Type": "application/json", "User-Agent": "opencode/1.0", "originator": "opencode", } def _iter_sse(response: httpx.Response) -> Iterator[GenericStreamingChunk]: for line in response.iter_lines(): if not line or line.startswith(":") or not line.startswith("data: "): continue raw = line[6:] if raw == "[DONE]": break try: event = json.loads(raw) except json.JSONDecodeError: continue etype = event.get("type", "") if etype == "response.output_text.delta": yield GenericStreamingChunk( text=event.get("delta", ""), is_finished=False, finish_reason="", usage=None, ) elif etype in ("response.completed", "response.done"): usage_data = event.get("response", {}).get("usage", {}) yield GenericStreamingChunk( text="", is_finished=True, finish_reason="stop", usage={"prompt_tokens": usage_data.get("input_tokens", 0), "completion_tokens": usage_data.get("output_tokens", 0), "total_tokens": usage_data.get("input_tokens", 0) + usage_data.get("output_tokens", 0)}, ) break async def _aiter_sse(response: httpx.Response) -> AsyncIterator[GenericStreamingChunk]: async for line in response.aiter_lines(): if not line or line.startswith(":") or not line.startswith("data: "): continue raw = line[6:] if raw == "[DONE]": break try: event = json.loads(raw) except json.JSONDecodeError: continue etype = event.get("type", "") if etype == "response.output_text.delta": yield GenericStreamingChunk( text=event.get("delta", ""), is_finished=False, finish_reason="", usage=None, ) elif etype in ("response.completed", "response.done"): usage_data = event.get("response", {}).get("usage", {}) yield GenericStreamingChunk( text="", is_finished=True, finish_reason="stop", usage={"prompt_tokens": usage_data.get("input_tokens", 0), "completion_tokens": usage_data.get("output_tokens", 0), "total_tokens": usage_data.get("input_tokens", 0) + usage_data.get("output_tokens", 0)}, ) break # --------------------------------------------------------------------------- # LiteLLM CustomLLM subclass # --------------------------------------------------------------------------- class CodexProvider(CustomLLM): def completion( self, model: str, messages: list, api_base: str, custom_prompt_dict: dict, model_response: ModelResponse, print_verbose: Callable, encoding, api_key, logging_obj, optional_params: dict, acompletion=None, litellm_params=None, logger_fn=None, headers={}, timeout: Optional[Union[float, httpx.Timeout]] = None, client: Optional[HTTPHandler] = None, ) -> ModelResponse: model_name = model.split("/", 1)[-1] if "/" in model else model try: auth = _get_auth() except Exception as e: raise CustomLLMError(status_code=401, message=f"Codex auth error: {e}") body = _build_body(model_name, messages, optional_params, stream=True) text_parts = [] prompt_tokens = completion_tokens = 0 try: with httpx.Client(timeout=120) as c: with c.stream("POST", API_URL, json=body, headers=_auth_headers(auth)) as resp: resp.raise_for_status() for chunk in _iter_sse(resp): text_parts.append(chunk["text"]) if chunk.get("usage"): prompt_tokens = chunk["usage"].get("prompt_tokens", 0) completion_tokens = chunk["usage"].get("completion_tokens", 0) except httpx.HTTPStatusError as e: raise CustomLLMError(status_code=e.response.status_code, message=str(e)) model_response.choices[0].message.content = "".join(text_parts) # type: ignore model_response.choices[0].finish_reason = "stop" # type: ignore model_response.model = _resolve_model(model_name) model_response.usage = Usage( prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=prompt_tokens + completion_tokens, ) return model_response def streaming( self, model: str, messages: list, api_base: str, custom_prompt_dict: dict, model_response: ModelResponse, print_verbose: Callable, encoding, api_key, logging_obj, optional_params: dict, acompletion=None, litellm_params=None, logger_fn=None, headers={}, timeout: Optional[Union[float, httpx.Timeout]] = None, client: Optional[HTTPHandler] = None, ) -> Iterator[GenericStreamingChunk]: model_name = model.split("/", 1)[-1] if "/" in model else model try: auth = _get_auth() except Exception as e: raise CustomLLMError(status_code=401, message=f"Codex auth error: {e}") body = _build_body(model_name, messages, optional_params, stream=True) try: with httpx.Client(timeout=120) as c: with c.stream("POST", API_URL, json=body, headers=_auth_headers(auth)) as resp: resp.raise_for_status() yield from _iter_sse(resp) except httpx.HTTPStatusError as e: raise CustomLLMError(status_code=e.response.status_code, message=str(e)) async def acompletion( self, model: str, messages: list, api_base: str, custom_prompt_dict: dict, model_response: ModelResponse, print_verbose: Callable, encoding, api_key, logging_obj, optional_params: dict, acompletion=None, litellm_params=None, logger_fn=None, headers={}, timeout: Optional[Union[float, httpx.Timeout]] = None, client: Optional[AsyncHTTPHandler] = None, ) -> ModelResponse: model_name = model.split("/", 1)[-1] if "/" in model else model try: auth = _get_auth() except Exception as e: raise CustomLLMError(status_code=401, message=f"Codex auth error: {e}") body = _build_body(model_name, messages, optional_params, stream=True) text_parts = [] prompt_tokens = completion_tokens = 0 try: async with httpx.AsyncClient(timeout=120) as c: async with c.stream("POST", API_URL, json=body, headers=_auth_headers(auth)) as resp: resp.raise_for_status() async for chunk in _aiter_sse(resp): text_parts.append(chunk["text"]) if chunk.get("usage"): prompt_tokens = chunk["usage"].get("prompt_tokens", 0) completion_tokens = chunk["usage"].get("completion_tokens", 0) except httpx.HTTPStatusError as e: raise CustomLLMError(status_code=e.response.status_code, message=str(e)) model_response.choices[0].message.content = "".join(text_parts) # type: ignore model_response.choices[0].finish_reason = "stop" # type: ignore model_response.model = _resolve_model(model_name) model_response.usage = Usage( prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=prompt_tokens + completion_tokens, ) return model_response async def astreaming( self, model: str, messages: list, api_base: str, custom_prompt_dict: dict, model_response: ModelResponse, print_verbose: Callable, encoding, api_key, logging_obj, optional_params: dict, acompletion=None, litellm_params=None, logger_fn=None, headers={}, timeout: Optional[Union[float, httpx.Timeout]] = None, client: Optional[AsyncHTTPHandler] = None, ) -> AsyncIterator[GenericStreamingChunk]: model_name = model.split("/", 1)[-1] if "/" in model else model try: auth = _get_auth() except Exception as e: raise CustomLLMError(status_code=401, message=f"Codex auth error: {e}") body = _build_body(model_name, messages, optional_params, stream=True) try: async with httpx.AsyncClient(timeout=120) as c: async with c.stream("POST", API_URL, json=body, headers=_auth_headers(auth)) as resp: resp.raise_for_status() async for chunk in _aiter_sse(resp): yield chunk except httpx.HTTPStatusError as e: raise CustomLLMError(status_code=e.response.status_code, message=str(e)) codex_provider = CodexProvider()