From 3da9810e87b0c9e3ff8aaa3eb4dd21c5f5009d79 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Fri, 20 Mar 2026 08:14:20 -0700 Subject: [PATCH] feat(llm): Add OpenAI Codex (ChatGPT subscription) as LLM provider (#1461) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(llm): add OpenAI Codex backend config and OAuth session manager Add OpenAiCodex as a new LLM backend variant with config for auth endpoint, API base URL, client ID, and session persistence path. The session manager implements OpenAI's device code auth flow (headless-friendly, no browser required on the server) with automatic token refresh, following the same persistence pattern as the existing NEAR AI session manager. Closes #742 Co-Authored-By: Claude Opus 4.6 * feat(llm): add Responses API client and token-refreshing decorator Native Responses API client for chatgpt.com/backend-api/codex/responses, the endpoint that works with ChatGPT subscription tokens. Handles SSE streaming, text completions, and tool call round-trips. Token-refreshing decorator wraps the provider to pre-emptively refresh OAuth tokens before API calls and retry once on auth failures. Reports zero cost since billing is through subscription. Co-Authored-By: Claude Opus 4.6 * feat(llm): wire OpenAI Codex into provider factory, CLI, and setup wizard Connect the new provider to the LLM factory, add openai_codex to the CLI --backend flag, and add it as an option in the onboarding wizard. Co-Authored-By: Claude Opus 4.6 * fix(llm): address PR #744 review feedback (20 items) Review fixes for the OpenAI Codex provider PR: - Remove dead `generate_pkce()` code (device flow gets PKCE from server) - Fix `refresh_tokens()` to use `.form()` instead of `.json()` per OAuth spec - Inline codex dispatch into `build_provider_chain()` (single async function, no separate `assemble_provider_chain()` helper — matches main's pattern) - Remove Clone from `OpenAiCodexSession`, restrict fields to `pub(crate)` - Propagate HTTP client builder error instead of silent fallback - Redact device code response body from debug log - Change `set_model()` in TokenRefreshingProvider to delegate to inner - Replace hardcoded `/tmp/` test path with `tempfile::tempdir()` - Accept `request_timeout_secs` from config instead of hardcoded 300s - Parse `Retry-After` header on 429 responses (matches nearai_chat.rs pattern) - Reuse `normalize_schema_strict()` for Codex tool definitions - Add warning log for dropped image attachments - Add doc comments on `list_models()` and `include` field - Add `OPENAI_CODEX_API_URL` to `.env.example` - Fix codex error message in `create_llm_provider()` for clarity - Revert unrelated `.worktrees` addition to `.gitignore` - Update `src/llm/CLAUDE.md` with Codex provider docs [skip-regression-check] Co-Authored-By: Claude Opus 4.6 * fix: address review feedback and harden OpenAI Codex provider (takeover #744) Security: - Add SSRF validation (validate_base_url) on OPENAI_CODEX_AUTH_URL and OPENAI_CODEX_API_URL, matching the pattern used by all other base URL configs (regression test for #1103 included) Correctness: - Add missing cache_write_multiplier() and cache_read_discount() trait delegation in TokenRefreshingProvider - Cap device-code polling backoff at 60s to prevent unbounded interval growth on repeated 429 responses - Default expires_in to 3600s when server returns 0, preventing immediately-expired sessions - Fix pre-existing SseEvent::JobResult missing fallback_deliverable field in job_monitor.rs tests Cleanup: - Extract duplicated make_test_jwt() and test_codex_config() into shared codex_test_helpers module Co-Authored-By: Sanjeev-S Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address PR review feedback on OpenAI Codex provider (#1461) - Login command now resolves OPENAI_CODEX_* env overrides even when LLM_BACKEND isn't set to openai_codex (Copilot review) - Setup wizard "Keep current provider?" for codex no longer re-triggers device code login — mirrors Bedrock's keep-and-return pattern (Copilot) - Revert provider init log from info back to debug (Copilot) - Add warning log when token expires_in=0, before defaulting to 3600s (Gemini review) Co-Authored-By: Sanjeev-S Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Sanjeev Suresh Co-authored-by: Claude Opus 4.6 --- .env.example | 9 +- src/app.rs | 1 + src/cli/mod.rs | 11 + .../ironclaw__cli__tests__help_output.snap | 1 + ...li__tests__help_output_without_import.snap | 1 + ...ronclaw__cli__tests__long_help_output.snap | 1 + ...ests__long_help_output_without_import.snap | 1 + src/config/llm.rs | 201 ++- src/config/mod.rs | 4 +- src/llm/CLAUDE.md | 24 +- src/llm/codex_test_helpers.rs | 34 + src/llm/config.rs | 33 + src/llm/mod.rs | 68 +- src/llm/models.rs | 1 + src/llm/openai_codex_provider.rs | 1091 +++++++++++++++++ src/llm/openai_codex_session.rs | 731 +++++++++++ src/llm/rig_adapter.rs | 2 +- src/llm/token_refreshing.rs | 191 +++ src/main.rs | 41 + src/setup/wizard.rs | 45 +- 20 files changed, 2477 insertions(+), 14 deletions(-) create mode 100644 src/llm/codex_test_helpers.rs create mode 100644 src/llm/openai_codex_provider.rs create mode 100644 src/llm/openai_codex_session.rs create mode 100644 src/llm/token_refreshing.rs diff --git a/.env.example b/.env.example index 3fd58ef6..b52412c5 100644 --- a/.env.example +++ b/.env.example @@ -4,7 +4,7 @@ DATABASE_POOL_SIZE=10 # LLM Provider # LLM_BACKEND=nearai # default -# Possible values: nearai, ollama, openai_compatible, openai, anthropic, tinfoil +# Possible values: nearai, ollama, openai_compatible, openai, anthropic, tinfoil, openai_codex # LLM_REQUEST_TIMEOUT_SECS=120 # Increase for local LLMs (Ollama, vLLM, LM Studio) # === Anthropic Direct === @@ -92,6 +92,13 @@ NEARAI_AUTH_URL=https://private.near.ai # long = 1-hour TTL, 2.0× (200%) write surcharge # ANTHROPIC_CACHE_RETENTION=short +# === OpenAI Codex (ChatGPT subscription, OAuth) === +# LLM_BACKEND=openai_codex +# OPENAI_CODEX_MODEL=gpt-5.3-codex # default +# OPENAI_CODEX_CLIENT_ID=app_EMoamEEZ73f0CkXaXp7hrann # override (rare) +# OPENAI_CODEX_AUTH_URL=https://auth.openai.com # override (rare) +# OPENAI_CODEX_API_URL=https://chatgpt.com/backend-api/codex # override (rare) + # For full provider setup guide see docs/LLM_PROVIDERS.md # Channel Configuration diff --git a/src/app.rs b/src/app.rs index 729d2269..df246458 100644 --- a/src/app.rs +++ b/src/app.rs @@ -696,6 +696,7 @@ impl AppBuilder { // fail early with a clear error instead of a confusing runtime failure. if self.config.llm.backend != "nearai" && self.config.llm.backend != "bedrock" + && self.config.llm.backend != "openai_codex" && self.config.llm.provider.is_none() { let backend = &self.config.llm.backend; diff --git a/src/cli/mod.rs b/src/cli/mod.rs index 54779ae1..dffcc2c5 100644 --- a/src/cli/mod.rs +++ b/src/cli/mod.rs @@ -239,6 +239,17 @@ pub enum Command { )] Import(ImportCommand), + /// Authenticate with a provider (re-login) + #[command( + about = "Authenticate with a provider", + long_about = "Re-authenticate with an LLM provider.\nExample: ironclaw login --openai-codex" + )] + Login { + /// Authenticate with OpenAI Codex (ChatGPT subscription) + #[arg(long)] + openai_codex: bool, + }, + /// Run as a sandboxed worker inside a Docker container (internal use). /// This is invoked automatically by the orchestrator, not by users directly. #[command(hide = true)] diff --git a/src/cli/snapshots/ironclaw__cli__tests__help_output.snap b/src/cli/snapshots/ironclaw__cli__tests__help_output.snap index a554acae..81fed592 100644 --- a/src/cli/snapshots/ironclaw__cli__tests__help_output.snap +++ b/src/cli/snapshots/ironclaw__cli__tests__help_output.snap @@ -24,6 +24,7 @@ Commands: status Show system status completion Generate completions import Import from other AI systems + login Authenticate with a provider help Print this message or the help of the given subcommand(s) Options: diff --git a/src/cli/snapshots/ironclaw__cli__tests__help_output_without_import.snap b/src/cli/snapshots/ironclaw__cli__tests__help_output_without_import.snap index 3f3cf4fc..a6237fde 100644 --- a/src/cli/snapshots/ironclaw__cli__tests__help_output_without_import.snap +++ b/src/cli/snapshots/ironclaw__cli__tests__help_output_without_import.snap @@ -23,6 +23,7 @@ Commands: logs View and manage gateway logs status Show system status completion Generate completions + login Authenticate with a provider help Print this message or the help of the given subcommand(s) Options: diff --git a/src/cli/snapshots/ironclaw__cli__tests__long_help_output.snap b/src/cli/snapshots/ironclaw__cli__tests__long_help_output.snap index 99b3ef53..c124bad3 100644 --- a/src/cli/snapshots/ironclaw__cli__tests__long_help_output.snap +++ b/src/cli/snapshots/ironclaw__cli__tests__long_help_output.snap @@ -27,6 +27,7 @@ Commands: status Show system status completion Generate completions import Import from other AI systems + login Authenticate with a provider help Print this message or the help of the given subcommand(s) Options: diff --git a/src/cli/snapshots/ironclaw__cli__tests__long_help_output_without_import.snap b/src/cli/snapshots/ironclaw__cli__tests__long_help_output_without_import.snap index aa7ae8b0..6aa05e75 100644 --- a/src/cli/snapshots/ironclaw__cli__tests__long_help_output_without_import.snap +++ b/src/cli/snapshots/ironclaw__cli__tests__long_help_output_without_import.snap @@ -26,6 +26,7 @@ Commands: logs View and manage gateway logs status Show system status completion Generate completions + login Authenticate with a provider help Print this message or the help of the given subcommand(s) Options: diff --git a/src/config/llm.rs b/src/config/llm.rs index 37fd9c47..cc515611 100644 --- a/src/config/llm.rs +++ b/src/config/llm.rs @@ -37,6 +37,7 @@ impl LlmConfig { }, provider: None, bedrock: None, + openai_codex: None, request_timeout_secs: 120, cheap_model: None, smart_routing_cascade: false, @@ -72,8 +73,12 @@ impl LlmConfig { backend_lower == "nearai" || backend_lower == "near_ai" || backend_lower == "near"; let is_bedrock = backend_lower == "bedrock" || backend_lower == "aws_bedrock" || backend_lower == "aws"; + let is_openai_codex = backend_lower == "openai_codex" + || backend_lower == "openai-codex" + || backend_lower == "codex"; - if !is_nearai && !is_bedrock && registry.find(&backend_lower).is_none() { + if !is_nearai && !is_bedrock && !is_openai_codex && registry.find(&backend_lower).is_none() + { tracing::warn!( "Unknown LLM backend '{}'. Will attempt as openai_compatible fallback.", backend @@ -126,8 +131,8 @@ impl LlmConfig { smart_routing_cascade: parse_optional_env("SMART_ROUTING_CASCADE", true)?, }; - // Resolve registry provider config (for non-NearAI, non-Bedrock backends) - let provider = if is_nearai || is_bedrock { + // Resolve registry provider config (for non-NearAI, non-Bedrock, non-Codex backends) + let provider = if is_nearai || is_bedrock || is_openai_codex { None } else { Some(Self::resolve_registry_provider( @@ -174,6 +179,38 @@ impl LlmConfig { None }; + // Resolve OpenAI Codex config + let openai_codex = if is_openai_codex { + // Model: OPENAI_CODEX_MODEL > OPENAI_MODEL > settings.selected_model > default + let model = optional_env("OPENAI_CODEX_MODEL")? + .or(optional_env("OPENAI_MODEL")?) + .or_else(|| settings.selected_model.clone()) + .unwrap_or_else(|| "gpt-5.3-codex".to_string()); + let auth_endpoint = optional_env("OPENAI_CODEX_AUTH_URL")? + .unwrap_or_else(|| "https://auth.openai.com".to_string()); + validate_base_url(&auth_endpoint, "OPENAI_CODEX_AUTH_URL")?; + let api_base_url = optional_env("OPENAI_CODEX_API_URL")? + .unwrap_or_else(|| "https://chatgpt.com/backend-api/codex".to_string()); + validate_base_url(&api_base_url, "OPENAI_CODEX_API_URL")?; + let client_id = optional_env("OPENAI_CODEX_CLIENT_ID")? + .unwrap_or_else(|| "app_EMoamEEZ73f0CkXaXp7hrann".to_string()); + let session_path = optional_env("OPENAI_CODEX_SESSION_PATH")? + .map(PathBuf::from) + .unwrap_or_else(|| ironclaw_base_dir().join("openai_codex_session.json")); + let token_refresh_margin_secs = + parse_optional_env("OPENAI_CODEX_REFRESH_MARGIN_SECS", 300)?; + Some(OpenAiCodexConfig { + model, + auth_endpoint, + api_base_url, + client_id, + session_path, + token_refresh_margin_secs, + }) + } else { + None + }; + let request_timeout_secs = parse_optional_env("LLM_REQUEST_TIMEOUT_SECS", 120)?; // Generic cheap model (works with any backend). @@ -189,6 +226,8 @@ impl LlmConfig { "nearai".to_string() } else if is_bedrock { "bedrock".to_string() + } else if is_openai_codex { + "openai_codex".to_string() } else if let Some(ref p) = provider { p.provider_id.clone() } else { @@ -198,6 +237,7 @@ impl LlmConfig { nearai, provider, bedrock, + openai_codex, request_timeout_secs, cheap_model, smart_routing_cascade, @@ -1069,4 +1109,159 @@ mod tests { std::env::remove_var("LLM_REQUEST_TIMEOUT_SECS"); } } + + // ── OpenAI Codex tests ────────────────────────────────────────── + + /// Clear all openai-codex-related env vars. + fn clear_openai_codex_env() { + // SAFETY: Only called under ENV_MUTEX in tests. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("OPENAI_CODEX_MODEL"); + std::env::remove_var("OPENAI_MODEL"); + } + } + + #[test] + fn openai_codex_resolves_config() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_openai_codex_env(); + + let settings = Settings { + llm_backend: Some("openai_codex".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + assert_eq!(cfg.backend, "openai_codex"); + let codex = cfg.openai_codex.expect("codex config should be present"); + assert_eq!(codex.model, "gpt-5.3-codex"); // default + assert!( + cfg.provider.is_none(), + "codex should not use registry provider" + ); + } + + #[test] + fn openai_codex_model_env_resolution() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_openai_codex_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::set_var("OPENAI_CODEX_MODEL", "o3-pro"); + } + + let settings = Settings { + llm_backend: Some("openai_codex".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let codex = cfg.openai_codex.expect("codex config should be present"); + assert_eq!(codex.model, "o3-pro"); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("OPENAI_CODEX_MODEL"); + } + } + + #[test] + fn openai_codex_falls_back_to_openai_model() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_openai_codex_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::set_var("OPENAI_MODEL", "gpt-4o"); + } + + let settings = Settings { + llm_backend: Some("openai_codex".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let codex = cfg.openai_codex.expect("codex config should be present"); + assert_eq!(codex.model, "gpt-4o"); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("OPENAI_MODEL"); + } + } + + #[test] + fn openai_codex_falls_back_to_selected_model() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_openai_codex_env(); + + let settings = Settings { + llm_backend: Some("openai_codex".to_string()), + selected_model: Some("gpt-4o-mini".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let codex = cfg.openai_codex.expect("codex config should be present"); + assert_eq!(codex.model, "gpt-4o-mini"); + } + + /// Regression: SSRF validation on OPENAI_CODEX_API_URL (#1103). + #[test] + fn openai_codex_rejects_ssrf_api_url() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_openai_codex_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::set_var( + "OPENAI_CODEX_API_URL", + "http://169.254.169.254/latest/meta-data", + ); + } + + let settings = Settings { + llm_backend: Some("openai_codex".to_string()), + ..Default::default() + }; + + let err = LlmConfig::resolve(&settings).unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("OPENAI_CODEX_API_URL"), + "error should reference the field name: {msg}" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("OPENAI_CODEX_API_URL"); + } + } + + /// Regression: SSRF validation on OPENAI_CODEX_AUTH_URL (#1103). + #[test] + fn openai_codex_rejects_ssrf_auth_url() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_openai_codex_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::set_var("OPENAI_CODEX_AUTH_URL", "http://10.0.0.1"); + } + + let settings = Settings { + llm_backend: Some("openai_codex".to_string()), + ..Default::default() + }; + + let err = LlmConfig::resolve(&settings).unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("OPENAI_CODEX_AUTH_URL"), + "error should reference the field name: {msg}" + ); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("OPENAI_CODEX_AUTH_URL"); + } + } } diff --git a/src/config/mod.rs b/src/config/mod.rs index e704d7dc..e4834a88 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -54,7 +54,7 @@ pub use self::transcription::TranscriptionConfig; pub use self::tunnel::TunnelConfig; pub use self::wasm::WasmConfig; pub use crate::llm::config::{ - BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER, + BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER, OpenAiCodexConfig, RegistryProviderConfig, }; pub use crate::llm::session::SessionConfig; @@ -377,7 +377,7 @@ pub(crate) fn resolve_owner_id(settings: &Settings) -> Result String { + use base64::Engine; + let engine = base64::engine::general_purpose::URL_SAFE_NO_PAD; + + let header = engine.encode(b"{\"alg\":\"RS256\",\"typ\":\"JWT\"}"); + let payload_json = serde_json::json!({ + "sub": "user123", + "https://api.openai.com/auth": { + "chatgpt_account_id": account_id, + }, + }); + let payload = engine.encode(payload_json.to_string().as_bytes()); + let sig = engine.encode(b"fake-signature"); + format!("{header}.{payload}.{sig}") +} + +/// Build a test `OpenAiCodexConfig` with a given session path. +pub(crate) fn test_codex_config(session_path: std::path::PathBuf) -> OpenAiCodexConfig { + OpenAiCodexConfig { + model: "gpt-5.3-codex".to_string(), + auth_endpoint: "https://auth.openai.com".to_string(), + api_base_url: "https://chatgpt.com/backend-api/codex".to_string(), + client_id: "test_client_id".to_string(), + session_path, + token_refresh_margin_secs: 300, + } +} diff --git a/src/llm/config.rs b/src/llm/config.rs index 6ac0060a..aea0478a 100644 --- a/src/llm/config.rs +++ b/src/llm/config.rs @@ -9,6 +9,7 @@ use std::path::PathBuf; use secrecy::SecretString; +use crate::bootstrap::ironclaw_base_dir; use crate::llm::registry::ProviderProtocol; use crate::llm::session::SessionConfig; @@ -102,6 +103,36 @@ pub struct RegistryProviderConfig { pub unsupported_params: Vec, } +/// Configuration for OpenAI Codex (ChatGPT subscription OAuth). +#[derive(Debug, Clone)] +pub struct OpenAiCodexConfig { + /// Model to use (default: "gpt-5.3-codex"). + pub model: String, + /// OAuth authorization server (default: "https://auth.openai.com"). + pub auth_endpoint: String, + /// Responses API base URL (default: "https://chatgpt.com/backend-api/codex"). + pub api_base_url: String, + /// OAuth client ID (default: OpenAI's public Codex client). + pub client_id: String, + /// Path to session file (default: ~/.ironclaw/openai_codex_session.json). + pub session_path: PathBuf, + /// Seconds before expiry to proactively refresh (default: 300). + pub token_refresh_margin_secs: u64, +} + +impl Default for OpenAiCodexConfig { + fn default() -> Self { + Self { + model: "gpt-5.3-codex".to_string(), + auth_endpoint: "https://auth.openai.com".to_string(), + api_base_url: "https://chatgpt.com/backend-api/codex".to_string(), + client_id: "app_EMoamEEZ73f0CkXaXp7hrann".to_string(), + session_path: ironclaw_base_dir().join("openai_codex_session.json"), + token_refresh_margin_secs: 300, + } + } +} + /// Configuration for AWS Bedrock (native Converse API). #[derive(Debug, Clone)] pub struct BedrockConfig { @@ -134,6 +165,8 @@ pub struct LlmConfig { pub provider: Option, /// AWS Bedrock config (populated when backend=bedrock, requires --features bedrock). pub bedrock: Option, + /// OpenAI Codex config (populated when backend=openai_codex). + pub openai_codex: Option, /// HTTP request timeout in seconds for LLM API calls. /// Default: 120. Increase for local LLMs (Ollama, vLLM, LM Studio) that /// need more time for prompt evaluation on consumer hardware. diff --git a/src/llm/mod.rs b/src/llm/mod.rs index 8551cb61..8d75de95 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -20,6 +20,8 @@ pub mod error; pub mod failover; mod nearai_chat; pub mod oauth_helpers; +pub mod openai_codex_provider; +pub mod openai_codex_session; mod provider; mod reasoning; pub mod recording; @@ -29,6 +31,10 @@ pub mod retry; mod rig_adapter; pub mod session; pub mod smart_routing; +mod token_refreshing; + +#[cfg(test)] +mod codex_test_helpers; pub mod image_models; pub mod models; @@ -37,12 +43,14 @@ pub mod vision_models; pub use circuit_breaker::{CircuitBreakerConfig, CircuitBreakerProvider}; pub use config::{ - BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER, + BedrockConfig, CacheRetention, LlmConfig, NearAiConfig, OAUTH_PLACEHOLDER, OpenAiCodexConfig, RegistryProviderConfig, }; pub use error::LlmError; pub use failover::{CooldownConfig, FailoverProvider}; pub use nearai_chat::{DEFAULT_MODEL, ModelInfo, NearAiChatProvider, default_models}; +pub use openai_codex_provider::OpenAiCodexProvider; +pub use openai_codex_session::{OpenAiCodexSession, OpenAiCodexSessionManager}; pub use provider::{ ChatMessage, CompletionRequest, CompletionResponse, ContentPart, FinishReason, ImageUrl, LlmProvider, ModelMetadata, Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse, @@ -59,6 +67,7 @@ pub use retry::{RetryConfig, RetryProvider}; pub use rig_adapter::RigAdapter; pub use session::{SessionConfig, SessionManager, create_session_manager}; pub use smart_routing::{SmartRoutingConfig, SmartRoutingProvider, TaskComplexity}; +pub use token_refreshing::TokenRefreshingProvider; use std::sync::Arc; @@ -97,6 +106,15 @@ pub async fn create_llm_provider( } } + if config.backend == "openai_codex" { + return Err(LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: + "OpenAI Codex uses a dedicated factory path. Use build_provider_chain() instead of create_llm_provider()." + .to_string(), + }); + } + let reg_config = config .provider .as_ref() @@ -374,6 +392,47 @@ fn create_ollama_from_registry( Ok(Arc::new(adapter)) } +/// Create an OpenAI Codex provider with OAuth authentication. +/// +/// This is async because it needs to ensure authentication before +/// creating the provider (which requires a valid Bearer token). +/// +/// Uses the Responses API (`chatgpt.com/backend-api/codex/responses`) +/// instead of the Chat Completions API, matching OpenClaw's approach. +async fn create_openai_codex_provider( + config: &LlmConfig, +) -> Result, LlmError> { + let codex = config + .openai_codex + .as_ref() + .ok_or_else(|| LlmError::AuthFailed { + provider: "openai_codex".to_string(), + })?; + + let session_mgr = Arc::new(OpenAiCodexSessionManager::new(codex.clone())?); + session_mgr.ensure_authenticated().await?; + + let token = session_mgr.get_access_token().await?; + + let provider = Arc::new(OpenAiCodexProvider::new( + &codex.model, + &codex.api_base_url, + token.expose_secret(), + config.request_timeout_secs, + )?); + + tracing::info!( + "Using OpenAI Codex (Responses API, model: {}, base: {})", + codex.model, + codex.api_base_url, + ); + + Ok(Arc::new(TokenRefreshingProvider::new( + provider, + session_mgr, + ))) +} + /// Create a cheap/fast LLM provider for lightweight tasks (heartbeat, routing, evaluation). /// /// Resolution order: @@ -460,7 +519,11 @@ pub async fn build_provider_chain( ), LlmError, > { - let llm = create_llm_provider(config, session.clone()).await?; + let llm: Arc = if config.backend == "openai_codex" { + create_openai_codex_provider(config).await? + } else { + create_llm_provider(config, session.clone()).await? + }; tracing::debug!("LLM provider initialized: {}", llm.model_name()); // 1. Retry @@ -632,6 +695,7 @@ mod tests { request_timeout_secs: 120, cheap_model: None, smart_routing_cascade: true, + openai_codex: None, } } diff --git a/src/llm/models.rs b/src/llm/models.rs index daec9df3..fcf09beb 100644 --- a/src/llm/models.rs +++ b/src/llm/models.rs @@ -347,5 +347,6 @@ pub(crate) fn build_nearai_model_fetch_config() -> crate::config::LlmConfig { request_timeout_secs: 120, cheap_model: None, smart_routing_cascade: false, + openai_codex: None, } } diff --git a/src/llm/openai_codex_provider.rs b/src/llm/openai_codex_provider.rs new file mode 100644 index 00000000..9e3aa955 --- /dev/null +++ b/src/llm/openai_codex_provider.rs @@ -0,0 +1,1091 @@ +//! OpenAI Codex Responses API client. +//! +//! Implements `LlmProvider` using the Responses API at +//! `chatgpt.com/backend-api/codex/responses` -- the endpoint that works +//! with ChatGPT subscription OAuth tokens. +//! +//! This mirrors OpenClaw's Responses API flow translated to Rust. + +use async_trait::async_trait; +use reqwest::Client; +use rust_decimal::Decimal; +use serde::Deserialize; +use tokio::sync::RwLock; + +use crate::error::LlmError; +use crate::llm::provider::{ + ChatMessage, CompletionRequest, CompletionResponse, ContentPart, FinishReason, LlmProvider, + ModelMetadata, Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse, ToolDefinition, +}; + +/// OpenAI Codex Responses API provider. +/// +/// Sends requests to `{api_base_url}/responses` using SSE streaming, +/// with JWT-based auth headers matching OpenClaw's approach. +/// Token + account ID pair, updated atomically. +struct AuthState { + token: String, + account_id: String, +} + +pub struct OpenAiCodexProvider { + client: Client, + model: String, + api_base_url: String, + auth: RwLock, +} + +impl OpenAiCodexProvider { + /// Create a new provider. + /// + /// Extracts the `chatgpt_account_id` from the JWT token. + /// `request_timeout_secs` controls the HTTP client timeout (falls back to 300s). + pub fn new( + model: &str, + api_base_url: &str, + token: &str, + request_timeout_secs: u64, + ) -> Result { + let account_id = extract_account_id(token)?; + Ok(Self { + client: Client::builder() + .timeout(std::time::Duration::from_secs(request_timeout_secs)) + .build() + .map_err(|e| LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to create HTTP client: {e}"), + })?, + model: model.to_string(), + api_base_url: api_base_url.trim_end_matches('/').to_string(), + auth: RwLock::new(AuthState { + token: token.to_string(), + account_id, + }), + }) + } + + /// Update the access token after a refresh. + pub async fn update_token(&self, token: &str) -> Result<(), LlmError> { + let account_id = extract_account_id(token)?; + *self.auth.write().await = AuthState { + token: token.to_string(), + account_id, + }; + tracing::debug!("Updated Codex provider token"); + Ok(()) + } + + /// Build request headers matching OpenClaw's `buildHeaders`. + async fn build_headers(&self) -> Result { + use reqwest::header::{ + ACCEPT, AUTHORIZATION, CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue, USER_AGENT, + }; + + let auth = self.auth.read().await; + + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {}", auth.token)).map_err(|e| { + LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("Invalid token for header: {e}"), + } + })?, + ); + headers.insert( + HeaderName::from_static("chatgpt-account-id"), + HeaderValue::from_str(&auth.account_id).map_err(|e| LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("Invalid account ID for header: {e}"), + })?, + ); + headers.insert( + HeaderName::from_static("openai-beta"), + HeaderValue::from_static("responses=experimental"), + ); + headers.insert( + HeaderName::from_static("originator"), + HeaderValue::from_static("ironclaw"), + ); + headers.insert( + USER_AGENT, + HeaderValue::from_static(concat!("ironclaw/", env!("CARGO_PKG_VERSION"))), + ); + headers.insert(ACCEPT, HeaderValue::from_static("text/event-stream")); + headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + + Ok(headers) + } + + /// Build the request body for the Responses API. + fn build_request_body( + &self, + messages: &[ChatMessage], + tools: Option<&[ToolDefinition]>, + ) -> serde_json::Value { + // Separate system messages into `instructions` + let instructions: String = messages + .iter() + .filter(|m| m.role == Role::System) + .map(|m| m.content.as_str()) + .collect::>() + .join("\n\n"); + + // Convert non-system messages to Responses API format + let input: Vec = messages + .iter() + .filter(|m| m.role != Role::System) + .enumerate() + .flat_map(|(i, m)| convert_message(m, i)) + .collect(); + + let mut body = serde_json::json!({ + "model": self.model, + "store": false, + "stream": true, + "input": input, + "text": { "verbosity": "medium" }, + // Safe for non-reasoning models — API ignores unrecognized include values + "include": ["reasoning.encrypted_content"], + }); + + if !instructions.is_empty() { + body["instructions"] = serde_json::Value::String(instructions); + } + + if let Some(tools) = tools + && !tools.is_empty() + { + let tools_json: Vec = + tools.iter().map(convert_tool_definition).collect(); + body["tools"] = serde_json::Value::Array(tools_json); + body["tool_choice"] = serde_json::Value::String("auto".to_string()); + body["parallel_tool_calls"] = serde_json::Value::Bool(true); + } + + body + } + + /// Send a request and parse the SSE response stream. + async fn send_request(&self, body: serde_json::Value) -> Result { + let url = format!("{}/responses", self.api_base_url); + let headers = self.build_headers().await?; + + tracing::debug!( + url = %url, + model = %self.model, + "Sending Responses API request" + ); + + let response = self + .client + .post(&url) + .headers(headers) + .json(&body) + .send() + .await + .map_err(|e| LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("HTTP request failed: {e}"), + })?; + + let status = response.status(); + if !status.is_success() { + // Extract Retry-After header before consuming the response body. + // Supports both delay-seconds (RFC 7231 §7.1.3) and HTTP-date formats. + let retry_after = response + .headers() + .get("retry-after") + .and_then(|v| v.to_str().ok()) + .and_then(|v| { + if let Ok(secs) = v.trim().parse::() { + return Some(std::time::Duration::from_secs(secs)); + } + if let Ok(dt) = chrono::DateTime::parse_from_rfc2822(v.trim()) { + let now = chrono::Utc::now(); + let delta = dt.signed_duration_since(now); + return Some(std::time::Duration::from_secs( + delta.num_seconds().max(0) as u64 + )); + } + None + }); + + let body_text = response.text().await.unwrap_or_default(); + if status == reqwest::StatusCode::UNAUTHORIZED { + return Err(LlmError::AuthFailed { + provider: "openai_codex".to_string(), + }); + } + if status == reqwest::StatusCode::TOO_MANY_REQUESTS { + return Err(LlmError::RateLimited { + provider: "openai_codex".to_string(), + retry_after, + }); + } + return Err(LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("HTTP {status}: {body_text}"), + }); + } + + // Read the full body and parse SSE events + let body_bytes = response + .bytes() + .await + .map_err(|e| LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to read response body: {e}"), + })?; + + let body_text = String::from_utf8_lossy(&body_bytes); + parse_sse_response(&body_text) + } +} + +#[async_trait] +impl LlmProvider for OpenAiCodexProvider { + fn model_name(&self) -> &str { + &self.model + } + + fn cost_per_token(&self) -> (Decimal, Decimal) { + (Decimal::ZERO, Decimal::ZERO) + } + + fn calculate_cost(&self, _input_tokens: u32, _output_tokens: u32) -> Decimal { + Decimal::ZERO + } + + async fn complete(&self, request: CompletionRequest) -> Result { + let body = self.build_request_body(&request.messages, None); + let parsed = self.send_request(body).await?; + + Ok(CompletionResponse { + content: parsed.text_content, + input_tokens: parsed.input_tokens, + output_tokens: parsed.output_tokens, + finish_reason: parsed.finish_reason, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }) + } + + async fn complete_with_tools( + &self, + request: ToolCompletionRequest, + ) -> Result { + let body = self.build_request_body(&request.messages, Some(&request.tools)); + let parsed = self.send_request(body).await?; + + let finish_reason = if !parsed.tool_calls.is_empty() { + FinishReason::ToolUse + } else { + parsed.finish_reason + }; + + Ok(ToolCompletionResponse { + content: if parsed.text_content.is_empty() { + None + } else { + Some(parsed.text_content) + }, + tool_calls: parsed.tool_calls, + input_tokens: parsed.input_tokens, + output_tokens: parsed.output_tokens, + finish_reason, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }) + } + + /// Returns empty — Codex uses subscription-based access with a fixed model, + /// no model enumeration API is available. + async fn list_models(&self) -> Result, LlmError> { + Ok(vec![]) + } + + async fn model_metadata(&self) -> Result { + Ok(ModelMetadata { + id: self.model.clone(), + context_length: None, + }) + } + + fn set_model(&self, _model: &str) -> Result<(), LlmError> { + Err(LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: "Cannot change model on Codex provider at runtime".to_string(), + }) + } + + fn effective_model_name(&self, _requested_model: Option<&str>) -> String { + self.model.clone() + } +} + +// --------------------------------------------------------------------------- +// JWT account ID extraction +// --------------------------------------------------------------------------- + +/// Extract `chatgpt_account_id` from a JWT token's payload. +/// +/// Matches OpenClaw's `extractAccountId` which reads: +/// `payload["https://api.openai.com/auth"]["chatgpt_account_id"]` +fn extract_account_id(token: &str) -> Result { + let parts: Vec<&str> = token.split('.').collect(); + if parts.len() < 2 { + return Err(LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: "JWT token has fewer than 2 parts".to_string(), + }); + } + + use base64::Engine; + let engine = base64::engine::general_purpose::URL_SAFE_NO_PAD; + + // JWT base64url may need padding + let payload_b64 = parts[1]; + let decoded = engine + .decode(payload_b64) + .map_err(|e| LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to decode JWT payload: {e}"), + })?; + + let payload: serde_json::Value = + serde_json::from_slice(&decoded).map_err(|e| LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to parse JWT payload as JSON: {e}"), + })?; + + let account_id = payload + .get("https://api.openai.com/auth") + .and_then(|auth| auth.get("chatgpt_account_id")) + .and_then(|v| v.as_str()) + .ok_or_else(|| LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: "JWT payload missing chatgpt_account_id claim".to_string(), + })?; + + Ok(account_id.to_string()) +} + +// --------------------------------------------------------------------------- +// Message conversion (matching OpenClaw's convertResponsesMessages) +// --------------------------------------------------------------------------- + +/// Convert a single `ChatMessage` to Responses API `input` items. +/// +/// Returns a Vec because assistant messages with tool_calls produce +/// one `function_call` item per tool call. +fn convert_message(msg: &ChatMessage, index: usize) -> Vec { + match msg.role { + Role::System => { + // System messages are handled separately as `instructions` + vec![] + } + Role::User => { + let image_count = msg + .content_parts + .iter() + .filter(|p| matches!(p, ContentPart::ImageUrl { .. })) + .count(); + if image_count > 0 { + tracing::warn!( + "OpenAI Codex: {} image attachment(s) dropped — Responses API image support not yet implemented", + image_count + ); + } + vec![serde_json::json!({ + "role": "user", + "content": [{ + "type": "input_text", + "text": msg.content, + }], + })] + } + Role::Assistant => { + // Check if this message has tool calls + if let Some(ref tool_calls) = msg.tool_calls { + // Emit one function_call item per tool call + tool_calls + .iter() + .map(|tc| { + let args_str = if tc.arguments.is_string() { + tc.arguments.as_str().unwrap_or("{}").to_string() + } else { + tc.arguments.to_string() + }; + serde_json::json!({ + "type": "function_call", + "call_id": tc.id, + "name": tc.name, + "arguments": args_str, + }) + }) + .collect() + } else { + // Plain text assistant message + vec![serde_json::json!({ + "type": "message", + "role": "assistant", + "id": format!("msg_{index}"), + "status": "completed", + "content": [{ + "type": "output_text", + "text": msg.content, + "annotations": [], + }], + })] + } + } + Role::Tool => { + let call_id = msg.tool_call_id.as_deref().unwrap_or("unknown"); + vec![serde_json::json!({ + "type": "function_call_output", + "call_id": call_id, + "output": msg.content, + })] + } + } +} + +/// Convert a `ToolDefinition` to Responses API tool format. +/// +/// Applies strict-mode schema normalization (same as OpenAI Chat Completions): +/// `additionalProperties: false`, all properties required, optional fields nullable. +fn convert_tool_definition(tool: &ToolDefinition) -> serde_json::Value { + use crate::llm::rig_adapter::normalize_schema_strict; + + serde_json::json!({ + "type": "function", + "name": tool.name, + "description": tool.description, + "parameters": normalize_schema_strict(&tool.parameters), + }) +} + +// --------------------------------------------------------------------------- +// SSE response parsing (matching OpenClaw's processResponsesStream) +// --------------------------------------------------------------------------- + +/// Parsed result from the SSE stream. +#[derive(Debug)] +struct ParsedResponse { + text_content: String, + tool_calls: Vec, + input_tokens: u32, + output_tokens: u32, + finish_reason: FinishReason, +} + +/// SSE event data from the Responses API. +#[derive(Debug, Deserialize)] +struct SseEvent { + #[serde(rename = "type")] + event_type: String, + #[serde(flatten)] + data: serde_json::Value, +} + +/// Tracking state for an in-progress function call. +#[derive(Debug, Default)] +struct FunctionCallState { + call_id: String, + name: String, + arguments: String, +} + +/// Parse the full SSE response body into a `ParsedResponse`. +fn parse_sse_response(body: &str) -> Result { + let mut text_content = String::new(); + let mut tool_calls: Vec = Vec::new(); + let mut input_tokens: u32 = 0; + let mut output_tokens: u32 = 0; + let mut finish_reason = FinishReason::Stop; + let mut active_function_calls: std::collections::HashMap = + std::collections::HashMap::new(); + let mut response_status: Option = None; + + for line in body.lines() { + let line = line.trim(); + + // Skip empty lines and comments + if line.is_empty() || line.starts_with(':') { + continue; + } + + // Parse SSE data lines + let data_str = if let Some(stripped) = line.strip_prefix("data: ") { + stripped.trim() + } else if let Some(stripped) = line.strip_prefix("data:") { + stripped.trim() + } else { + continue; + }; + + // Skip [DONE] marker + if data_str == "[DONE]" { + break; + } + + // Parse JSON + let event: SseEvent = match serde_json::from_str(data_str) { + Ok(e) => e, + Err(e) => { + tracing::trace!(data = data_str, error = %e, "Skipping unparseable SSE event"); + continue; + } + }; + + match event.event_type.as_str() { + // Text output + "response.output_text.delta" => { + if let Some(delta) = event.data.get("delta").and_then(|d| d.as_str()) { + text_content.push_str(delta); + } + } + + // Output item added (could be message or function_call) + "response.output_item.added" => { + if let Some(item) = event.data.get("item") { + let item_type = item.get("type").and_then(|t| t.as_str()).unwrap_or(""); + if item_type == "function_call" { + let item_id = item + .get("id") + .or_else(|| item.get("call_id")) + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let name = item + .get("name") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let call_id = item + .get("call_id") + .and_then(|v| v.as_str()) + .unwrap_or(&item_id) + .to_string(); + active_function_calls.insert( + item_id.clone(), + FunctionCallState { + call_id, + name, + arguments: String::new(), + }, + ); + } + } + } + + // Function call arguments streaming + "response.function_call_arguments.delta" => { + if let Some(delta) = event.data.get("delta").and_then(|d| d.as_str()) { + let item_id = event + .data + .get("item_id") + .and_then(|v| v.as_str()) + .unwrap_or(""); + if let Some(state) = active_function_calls.get_mut(item_id) { + state.arguments.push_str(delta); + } + } + } + + // Function call arguments done + "response.function_call_arguments.done" => { + // Arguments are finalized, item_id used to match + if let Some(args_str) = event.data.get("arguments").and_then(|a| a.as_str()) { + let item_id = event + .data + .get("item_id") + .and_then(|v| v.as_str()) + .unwrap_or(""); + if let Some(state) = active_function_calls.get_mut(item_id) { + state.arguments = args_str.to_string(); + } + } + } + + // Output item done (finalize function call) + "response.output_item.done" => { + if let Some(item) = event.data.get("item") { + let item_type = item.get("type").and_then(|t| t.as_str()).unwrap_or(""); + if item_type == "function_call" { + let item_id = item.get("id").and_then(|v| v.as_str()).unwrap_or(""); + if let Some(state) = active_function_calls.remove(item_id) { + let arguments: serde_json::Value = + serde_json::from_str(&state.arguments).unwrap_or_else(|_| { + serde_json::Value::String(state.arguments.clone()) + }); + tool_calls.push(ToolCall { + id: state.call_id, + name: state.name, + arguments, + }); + } else { + // Fallback: extract directly from the item + let call_id = item + .get("call_id") + .and_then(|v| v.as_str()) + .unwrap_or(item_id) + .to_string(); + let name = item + .get("name") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let args_str = item + .get("arguments") + .and_then(|v| v.as_str()) + .unwrap_or("{}"); + let arguments: serde_json::Value = serde_json::from_str(args_str) + .unwrap_or_else(|_| { + serde_json::Value::String(args_str.to_string()) + }); + tool_calls.push(ToolCall { + id: call_id, + name, + arguments, + }); + } + } + } + } + + // Response completed + "response.completed" => { + if let Some(response) = event.data.get("response") { + // Extract usage + if let Some(usage) = response.get("usage") { + input_tokens = usage + .get("input_tokens") + .and_then(|v| v.as_u64()) + .unwrap_or(0) as u32; + output_tokens = usage + .get("output_tokens") + .and_then(|v| v.as_u64()) + .unwrap_or(0) as u32; + } + // Extract status + if let Some(status) = response.get("status").and_then(|s| s.as_str()) { + response_status = Some(status.to_string()); + } + } + } + + // Response failed + "response.failed" => { + let reason = event + .data + .get("response") + .and_then(|r| r.get("status_details")) + .and_then(|d| d.get("error")) + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + .unwrap_or("Unknown error"); + return Err(LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("Response failed: {reason}"), + }); + } + + // Error event + "error" => { + let code = event + .data + .get("code") + .and_then(|c| c.as_str()) + .unwrap_or("unknown"); + let message = event + .data + .get("message") + .and_then(|m| m.as_str()) + .unwrap_or("Unknown error"); + return Err(LlmError::RequestFailed { + provider: "openai_codex".to_string(), + reason: format!("Error {code}: {message}"), + }); + } + + _ => { + // Ignore unhandled event types (e.g. response.created, + // response.output_item.added for messages, etc.) + } + } + } + + // Finalize any remaining active function calls + for (_, state) in active_function_calls { + if !state.name.is_empty() { + let arguments: serde_json::Value = serde_json::from_str(&state.arguments) + .unwrap_or(serde_json::Value::String(state.arguments)); + tool_calls.push(ToolCall { + id: state.call_id, + name: state.name, + arguments, + }); + } + } + + // Map status to finish reason (matching OpenClaw's mapStopReason) + if !tool_calls.is_empty() { + finish_reason = FinishReason::ToolUse; + } else if let Some(ref status) = response_status { + finish_reason = match status.as_str() { + "completed" => FinishReason::Stop, + "incomplete" => FinishReason::Length, + _ => FinishReason::Stop, + }; + } + + Ok(ParsedResponse { + text_content, + tool_calls, + input_tokens, + output_tokens, + finish_reason, + }) +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + use crate::llm::codex_test_helpers::make_test_jwt; + + #[test] + fn test_extract_account_id_success() { + let jwt = make_test_jwt("acct_abc123"); + let result = extract_account_id(&jwt); + assert!(result.is_ok()); + assert_eq!(result.unwrap(), "acct_abc123"); + } + + #[test] + fn test_extract_account_id_missing_claim() { + use base64::Engine; + let engine = base64::engine::general_purpose::URL_SAFE_NO_PAD; + let header = engine.encode(b"{\"alg\":\"RS256\"}"); + let payload = engine.encode(b"{\"sub\":\"user123\"}"); + let sig = engine.encode(b"sig"); + let jwt = format!("{header}.{payload}.{sig}"); + + let result = extract_account_id(&jwt); + assert!(result.is_err()); + } + + #[test] + fn test_extract_account_id_invalid_jwt() { + let result = extract_account_id("not-a-jwt"); + assert!(result.is_err()); + } + + #[test] + fn test_convert_user_message() { + let msg = ChatMessage::user("Hello world"); + let items = convert_message(&msg, 0); + assert_eq!(items.len(), 1); + assert_eq!(items[0]["role"], "user"); + assert_eq!(items[0]["content"][0]["type"], "input_text"); + assert_eq!(items[0]["content"][0]["text"], "Hello world"); + } + + #[test] + fn test_convert_system_message_excluded() { + let msg = ChatMessage::system("You are helpful"); + let items = convert_message(&msg, 0); + assert!(items.is_empty()); + } + + #[test] + fn test_convert_assistant_text_message() { + let msg = ChatMessage::assistant("Sure, I can help"); + let items = convert_message(&msg, 3); + assert_eq!(items.len(), 1); + assert_eq!(items[0]["type"], "message"); + assert_eq!(items[0]["role"], "assistant"); + assert_eq!(items[0]["id"], "msg_3"); + assert_eq!(items[0]["content"][0]["type"], "output_text"); + } + + #[test] + fn test_convert_assistant_with_tool_calls() { + let tool_calls = vec![ + ToolCall { + id: "call_1".to_string(), + name: "search".to_string(), + arguments: serde_json::json!({"query": "test"}), + }, + ToolCall { + id: "call_2".to_string(), + name: "read".to_string(), + arguments: serde_json::json!({"path": "/tmp"}), + }, + ]; + let msg = + ChatMessage::assistant_with_tool_calls(Some("Let me check".to_string()), tool_calls); + let items = convert_message(&msg, 0); + assert_eq!(items.len(), 2); + assert_eq!(items[0]["type"], "function_call"); + assert_eq!(items[0]["call_id"], "call_1"); + assert_eq!(items[0]["name"], "search"); + assert_eq!(items[1]["type"], "function_call"); + assert_eq!(items[1]["call_id"], "call_2"); + } + + #[test] + fn test_convert_tool_result_message() { + let msg = ChatMessage::tool_result("call_1", "search", "found 3 results"); + let items = convert_message(&msg, 0); + assert_eq!(items.len(), 1); + assert_eq!(items[0]["type"], "function_call_output"); + assert_eq!(items[0]["call_id"], "call_1"); + assert_eq!(items[0]["output"], "found 3 results"); + } + + #[test] + fn test_convert_tool_definition() { + let tool = ToolDefinition { + name: "my_tool".to_string(), + description: "Does things".to_string(), + parameters: serde_json::json!({ + "type": "object", + "properties": { + "x": { "type": "string" } + } + }), + }; + let json = convert_tool_definition(&tool); + assert_eq!(json["type"], "function"); + assert_eq!(json["name"], "my_tool"); + assert_eq!(json["description"], "Does things"); + } + + #[test] + fn test_parse_sse_text_response() { + let sse_body = r#"data: {"type":"response.output_item.added","item":{"type":"message","role":"assistant","id":"msg_1"}} + +data: {"type":"response.output_text.delta","delta":"Hello "} + +data: {"type":"response.output_text.delta","delta":"world!"} + +data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":10,"output_tokens":5}}} + +"#; + let result = parse_sse_response(sse_body); + assert!(result.is_ok()); + let parsed = result.unwrap(); + assert_eq!(parsed.text_content, "Hello world!"); + assert_eq!(parsed.input_tokens, 10); + assert_eq!(parsed.output_tokens, 5); + assert_eq!(parsed.finish_reason, FinishReason::Stop); + assert!(parsed.tool_calls.is_empty()); + } + + #[test] + fn test_parse_sse_tool_call_response() { + let sse_body = r#"data: {"type":"response.output_item.added","item":{"type":"function_call","id":"fc_1","call_id":"call_abc","name":"search"}} + +data: {"type":"response.function_call_arguments.delta","item_id":"fc_1","delta":"{\"query\":"} + +data: {"type":"response.function_call_arguments.delta","item_id":"fc_1","delta":"\"test\"}"} + +data: {"type":"response.output_item.done","item":{"type":"function_call","id":"fc_1","call_id":"call_abc","name":"search","arguments":"{\"query\":\"test\"}"}} + +data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":15,"output_tokens":8}}} + +"#; + let result = parse_sse_response(sse_body); + assert!(result.is_ok()); + let parsed = result.unwrap(); + assert!(parsed.text_content.is_empty()); + assert_eq!(parsed.tool_calls.len(), 1); + assert_eq!(parsed.tool_calls[0].id, "call_abc"); + assert_eq!(parsed.tool_calls[0].name, "search"); + assert_eq!( + parsed.tool_calls[0].arguments, + serde_json::json!({"query": "test"}) + ); + assert_eq!(parsed.finish_reason, FinishReason::ToolUse); + } + + #[test] + fn test_parse_sse_error_response() { + let sse_body = r#"data: {"type":"error","code":"rate_limit_exceeded","message":"Too many requests"} + +"#; + let result = parse_sse_response(sse_body); + assert!(result.is_err()); + let err = result.unwrap_err().to_string(); + assert!(err.contains("rate_limit_exceeded")); + } + + #[test] + fn test_parse_sse_failed_response() { + let sse_body = r#"data: {"type":"response.failed","response":{"status":"failed","status_details":{"error":{"message":"Model overloaded"}}}} + +"#; + let result = parse_sse_response(sse_body); + assert!(result.is_err()); + let err = result.unwrap_err().to_string(); + assert!(err.contains("Model overloaded")); + } + + #[test] + fn test_parse_sse_incomplete_status() { + let sse_body = r#"data: {"type":"response.output_text.delta","delta":"partial"} + +data: {"type":"response.completed","response":{"status":"incomplete","usage":{"input_tokens":5,"output_tokens":2}}} + +"#; + let result = parse_sse_response(sse_body); + assert!(result.is_ok()); + let parsed = result.unwrap(); + assert_eq!(parsed.text_content, "partial"); + assert_eq!(parsed.finish_reason, FinishReason::Length); + } + + #[test] + fn test_parse_sse_done_marker() { + let sse_body = r#"data: {"type":"response.output_text.delta","delta":"hello"} + +data: [DONE] + +data: {"type":"response.output_text.delta","delta":" ignored"} + +"#; + let result = parse_sse_response(sse_body); + assert!(result.is_ok()); + let parsed = result.unwrap(); + assert_eq!(parsed.text_content, "hello"); + } + + #[tokio::test] + async fn test_provider_new() { + let jwt = make_test_jwt("acct_test"); + let provider = OpenAiCodexProvider::new( + "gpt-5.3-codex", + "https://chatgpt.com/backend-api/codex", + &jwt, + 300, + ); + assert!(provider.is_ok()); + let provider = provider.unwrap(); + assert_eq!(provider.model_name(), "gpt-5.3-codex"); + assert_eq!(provider.cost_per_token(), (Decimal::ZERO, Decimal::ZERO)); + assert_eq!(provider.calculate_cost(1000, 500), Decimal::ZERO); + } + + #[tokio::test] + async fn test_update_token() { + let jwt1 = make_test_jwt("acct_old"); + let provider = OpenAiCodexProvider::new( + "gpt-5.3-codex", + "https://chatgpt.com/backend-api/codex", + &jwt1, + 300, + ) + .unwrap(); + + let jwt2 = make_test_jwt("acct_new"); + let result = provider.update_token(&jwt2).await; + assert!(result.is_ok()); + + // Verify account_id was updated + let auth = provider.auth.read().await; + assert_eq!(auth.account_id, "acct_new"); + } + + #[test] + fn test_build_request_body_structure() { + let jwt = make_test_jwt("acct_test"); + let provider = OpenAiCodexProvider::new( + "gpt-5.3-codex", + "https://chatgpt.com/backend-api/codex", + &jwt, + 300, + ) + .unwrap(); + + let messages = vec![ + ChatMessage::system("You are helpful"), + ChatMessage::user("Hello"), + ]; + + let body = provider.build_request_body(&messages, None); + + assert_eq!(body["model"], "gpt-5.3-codex"); + assert_eq!(body["store"], false); + assert_eq!(body["stream"], true); + assert_eq!(body["instructions"], "You are helpful"); + // input should only contain the user message, not system + let input = body["input"].as_array().unwrap(); + assert_eq!(input.len(), 1); + assert_eq!(input[0]["role"], "user"); + // No tools + assert!(body.get("tools").is_none()); + } + + #[test] + fn test_build_request_body_with_tools() { + let jwt = make_test_jwt("acct_test"); + let provider = OpenAiCodexProvider::new( + "gpt-5.3-codex", + "https://chatgpt.com/backend-api/codex", + &jwt, + 300, + ) + .unwrap(); + + let messages = vec![ChatMessage::user("Search for X")]; + let tools = vec![ToolDefinition { + name: "search".to_string(), + description: "Search for things".to_string(), + parameters: serde_json::json!({"type": "object"}), + }]; + + let body = provider.build_request_body(&messages, Some(&tools)); + + assert!(body.get("tools").is_some()); + let tools_arr = body["tools"].as_array().unwrap(); + assert_eq!(tools_arr.len(), 1); + assert_eq!(tools_arr[0]["type"], "function"); + assert_eq!(body["tool_choice"], "auto"); + assert_eq!(body["parallel_tool_calls"], true); + } + + #[test] + fn test_parse_sse_multiple_tool_calls() { + let sse_body = r#"data: {"type":"response.output_item.added","item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"read_file"}} + +data: {"type":"response.function_call_arguments.done","item_id":"fc_1","arguments":"{\"path\":\"/tmp/a\"}"} + +data: {"type":"response.output_item.done","item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"read_file","arguments":"{\"path\":\"/tmp/a\"}"}} + +data: {"type":"response.output_item.added","item":{"type":"function_call","id":"fc_2","call_id":"call_2","name":"read_file"}} + +data: {"type":"response.function_call_arguments.done","item_id":"fc_2","arguments":"{\"path\":\"/tmp/b\"}"} + +data: {"type":"response.output_item.done","item":{"type":"function_call","id":"fc_2","call_id":"call_2","name":"read_file","arguments":"{\"path\":\"/tmp/b\"}"}} + +data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":20,"output_tokens":12}}} + +"#; + let result = parse_sse_response(sse_body); + assert!(result.is_ok()); + let parsed = result.unwrap(); + assert_eq!(parsed.tool_calls.len(), 2); + assert_eq!(parsed.tool_calls[0].id, "call_1"); + assert_eq!(parsed.tool_calls[0].name, "read_file"); + assert_eq!(parsed.tool_calls[1].id, "call_2"); + assert_eq!(parsed.tool_calls[1].name, "read_file"); + assert_eq!(parsed.finish_reason, FinishReason::ToolUse); + } +} diff --git a/src/llm/openai_codex_session.rs b/src/llm/openai_codex_session.rs new file mode 100644 index 00000000..75c5e961 --- /dev/null +++ b/src/llm/openai_codex_session.rs @@ -0,0 +1,731 @@ +//! OAuth 2.0 session manager for OpenAI Codex (ChatGPT subscription). +//! +//! Supports two auth flows: +//! - **Device Code** (primary): Works on headless servers, no browser needed. +//! - **Browser PKCE** (fallback): Standard OAuth for local machines. +//! +//! Tokens are persisted to `~/.ironclaw/openai_codex_session.json` and +//! auto-refreshed before expiry. + +use chrono::{DateTime, Utc}; +use reqwest::Client; +use reqwest::header::{HeaderMap, HeaderValue, USER_AGENT}; +use secrecy::SecretString; +use serde::{Deserialize, Serialize}; +use tokio::sync::{Mutex, RwLock}; + +use crate::config::OpenAiCodexConfig; +use crate::error::LlmError; + +/// Persisted OAuth session data. +/// +/// Note: `Debug` is manually implemented to redact tokens. +#[derive(Serialize, Deserialize)] +pub struct OpenAiCodexSession { + pub(crate) access_token: String, + pub(crate) refresh_token: String, + pub(crate) expires_at: DateTime, + pub(crate) created_at: DateTime, +} + +impl std::fmt::Debug for OpenAiCodexSession { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("OpenAiCodexSession") + .field("access_token", &"[REDACTED]") + .field("refresh_token", &"[REDACTED]") + .field("expires_at", &self.expires_at) + .field("created_at", &self.created_at) + .finish() + } +} + +/// Request body for the device code usercode endpoint. +#[derive(Debug, Serialize)] +struct UserCodeRequest { + client_id: String, +} + +/// Response from the device code usercode endpoint. +#[derive(Debug, Deserialize)] +struct UserCodeResponse { + /// Unique ID for this device auth session. + device_auth_id: String, + /// Code the user enters in their browser. + user_code: String, + /// URL where the user enters the code (may not be present). + #[serde(default = "default_verification_uri")] + verification_uri: String, + /// Polling interval in seconds (OpenAI sends this as a string). + #[serde( + default = "default_interval", + deserialize_with = "deserialize_string_or_u64" + )] + interval: u64, + /// Expiry timestamp (OpenAI sends `expires_at` as ISO-8601). + #[serde(default)] + expires_at: Option, + /// Seconds until the device code expires (standard field, may not be present). + #[serde(default)] + expires_in: Option, +} + +fn default_verification_uri() -> String { + "https://auth.openai.com/codex/device".to_string() +} + +fn default_interval() -> u64 { + 5 +} + +/// Deserialize a value that may be either a string or a number as u64. +fn deserialize_string_or_u64<'de, D>(deserializer: D) -> Result +where + D: serde::Deserializer<'de>, +{ + use serde::de; + + struct StringOrU64; + impl<'de> de::Visitor<'de> for StringOrU64 { + type Value = u64; + fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result { + formatter.write_str("a string or integer") + } + fn visit_u64(self, v: u64) -> Result { + Ok(v) + } + fn visit_str(self, v: &str) -> Result { + v.parse().map_err(de::Error::custom) + } + } + deserializer.deserialize_any(StringOrU64) +} + +impl UserCodeResponse { + /// Get the expiry duration in seconds, from either `expires_in` or `expires_at`. + fn expires_in_secs(&self) -> u64 { + if let Some(secs) = self.expires_in { + return secs; + } + if let Some(ref ts) = self.expires_at + && let Ok(dt) = chrono::DateTime::parse_from_rfc3339(ts) + { + let remaining = dt.signed_duration_since(Utc::now()).num_seconds(); + return remaining.max(0) as u64; + } + 900 // default 15 minutes + } +} + +/// Request body for polling the device auth token endpoint. +#[derive(Debug, Serialize)] +struct DeviceTokenPollRequest { + device_auth_id: String, + user_code: String, +} + +/// Successful response from the device auth token endpoint. +/// Returns an authorization code + PKCE pair for the final token exchange. +#[derive(Debug, Deserialize)] +struct DeviceAuthCodeResponse { + authorization_code: String, + #[allow(dead_code)] + code_challenge: String, + code_verifier: String, +} + +/// Response from the final OAuth token exchange. +#[derive(Debug, Deserialize)] +struct TokenResponse { + access_token: String, + #[serde(default)] + refresh_token: String, + #[serde(default)] + expires_in: u64, + #[serde(default)] + #[allow(dead_code)] + token_type: String, +} + +/// Manages OpenAI Codex OAuth sessions with persistence and auto-refresh. +pub struct OpenAiCodexSessionManager { + config: OpenAiCodexConfig, + client: Client, + session: RwLock>, + renewal_lock: Mutex<()>, +} + +impl OpenAiCodexSessionManager { + /// Create a new session manager. Tries to load existing session from disk. + /// + /// # Errors + /// + /// Returns `LlmError` if the HTTP client cannot be constructed. + pub fn new(config: OpenAiCodexConfig) -> Result { + let mut headers = HeaderMap::new(); + headers.insert( + USER_AGENT, + HeaderValue::from_static(concat!("ironclaw/", env!("CARGO_PKG_VERSION"))), + ); + let client = Client::builder() + .default_headers(headers) + .timeout(std::time::Duration::from_secs(30)) + .build() + .map_err(|e| LlmError::RequestFailed { + provider: "openai_codex".into(), + reason: format!("HTTP client build failed: {e}"), + })?; + + let mgr = Self { + config, + client, + session: RwLock::new(None), + renewal_lock: Mutex::new(()), + }; + + // Try synchronous load from disk during construction + if let Ok(data) = std::fs::read_to_string(&mgr.config.session_path) + && let Ok(session) = serde_json::from_str::(&data) + && let Ok(mut guard) = mgr.session.try_write() + { + *guard = Some(session); + tracing::info!( + "Loaded OpenAI Codex session from {}", + mgr.config.session_path.display() + ); + } + + Ok(mgr) + } + + /// Check if we have a session (may be expired). + pub async fn has_session(&self) -> bool { + self.session.read().await.is_some() + } + + /// Check if the current access token needs refreshing. + pub async fn needs_refresh(&self) -> bool { + let guard = self.session.read().await; + match guard.as_ref() { + None => true, + Some(s) => { + let margin = + chrono::Duration::seconds(self.config.token_refresh_margin_secs as i64); + Utc::now() + margin >= s.expires_at + } + } + } + + /// Get the current access token, refreshing if needed. + /// + /// If the token is within the refresh margin, silently refreshes first. + /// If no session exists, returns an AuthFailed error. + pub async fn get_access_token(&self) -> Result { + if self.needs_refresh().await { + let has_refresh = self + .session + .read() + .await + .as_ref() + .map(|s| !s.refresh_token.is_empty()) + .unwrap_or(false); + if has_refresh { + self.refresh_tokens().await?; + } else { + return Err(LlmError::AuthFailed { + provider: "openai_codex".to_string(), + }); + } + } + + let guard = self.session.read().await; + guard + .as_ref() + .map(|s| SecretString::from(s.access_token.clone())) + .ok_or_else(|| LlmError::AuthFailed { + provider: "openai_codex".to_string(), + }) + } + + /// Ensure we have a valid session. Loads from disk, refreshes, or prompts login. + pub async fn ensure_authenticated(&self) -> Result<(), LlmError> { + // Try loading from disk if we don't have a session + if !self.has_session().await { + let _ = self.load_session().await; + } + + if !self.has_session().await { + // No session at all -- need to authenticate + return self.device_code_login().await; + } + + if self.needs_refresh().await { + // Try refresh; if it fails, re-authenticate + match self.refresh_tokens().await { + Ok(()) => Ok(()), + Err(e) => { + tracing::info!("Token refresh failed ({}), re-authenticating...", e); + self.device_code_login().await + } + } + } else { + Ok(()) + } + } + + /// Run OpenAI's device code auth flow. + /// + /// Uses OpenAI's custom `/api/accounts/deviceauth/*` endpoints (not the standard + /// Auth0 `/oauth/device/code` which is behind Cloudflare managed challenge). + /// + /// Flow: + /// 1. POST `/api/accounts/deviceauth/usercode` → get device_auth_id + user_code + /// 2. Poll POST `/api/accounts/deviceauth/token` → get authorization_code + PKCE + /// 3. Exchange via POST `/oauth/token` → get access_token + refresh_token + pub async fn device_code_login(&self) -> Result<(), LlmError> { + let _guard = self.renewal_lock.lock().await; + + let auth_base = format!("{}/api/accounts", self.config.auth_endpoint); + + // Step 1: Request device code + let usercode_url = format!("{}/deviceauth/usercode", auth_base); + let resp = self + .client + .post(&usercode_url) + .json(&UserCodeRequest { + client_id: self.config.client_id.clone(), + }) + .send() + .await + .map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Device code request failed: {}", e), + })?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Device code request failed: HTTP {} -- {}", status, body), + }); + } + + let body_text = resp + .text() + .await + .map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to read device code response: {}", e), + })?; + tracing::debug!("Device code response received ({} bytes)", body_text.len()); + let device: UserCodeResponse = + serde_json::from_str(&body_text).map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!( + "Failed to parse device code response: {} ({} bytes)", + e, + body_text.len() + ), + })?; + + // Step 2: Display code to user + println!(); + println!("==========================================================="); + println!(" OpenAI Codex Authentication "); + println!("==========================================================="); + println!(); + println!(" 1. Open this URL in any browser:"); + println!(" {}", device.verification_uri); + println!(); + println!(" 2. Enter this code:"); + println!(); + println!(" [ {} ]", device.user_code); + println!(); + let expires_secs = device.expires_in_secs(); + println!( + " Waiting for authorization... (expires in {} min)", + expires_secs / 60 + ); + println!("==========================================================="); + println!(); + + // Step 3: Poll for authorization code + let poll_url = format!("{}/deviceauth/token", auth_base); + let mut interval = std::time::Duration::from_secs(device.interval.max(5)); + let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(expires_secs); + + let auth_code = loop { + tokio::time::sleep(interval).await; + + if tokio::time::Instant::now() >= deadline { + return Err(LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: "Device code authorization timed out".to_string(), + }); + } + + let resp = self + .client + .post(&poll_url) + .json(&DeviceTokenPollRequest { + device_auth_id: device.device_auth_id.clone(), + user_code: device.user_code.clone(), + }) + .send() + .await + .map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Token poll request failed: {}", e), + })?; + + let status = resp.status(); + if status.is_success() { + let code_resp: DeviceAuthCodeResponse = + resp.json() + .await + .map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to parse auth code response: {}", e), + })?; + break code_resp; + } + + // 403 = authorization_pending, keep polling + // 404 = device code not found / not enabled + if status == reqwest::StatusCode::FORBIDDEN { + continue; + } + + if status == reqwest::StatusCode::NOT_FOUND { + return Err(LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: "Device code login is not enabled. Please check your OpenAI account settings.".to_string(), + }); + } + + // Slow down on 429, cap at 60s to avoid unbounded growth + if status == reqwest::StatusCode::TOO_MANY_REQUESTS { + interval = (interval + std::time::Duration::from_secs(5)) + .min(std::time::Duration::from_secs(60)); + continue; + } + + let body = resp.text().await.unwrap_or_default(); + return Err(LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Device auth poll failed: HTTP {} -- {}", status, body), + }); + }; + + // Step 4: Exchange authorization code for tokens (form-encoded, per Auth0 spec) + let token_url = format!("{}/oauth/token", self.config.auth_endpoint); + let resp = self + .client + .post(&token_url) + .form(&[ + ("grant_type", "authorization_code"), + ("code", &auth_code.authorization_code), + ("code_verifier", &auth_code.code_verifier), + ("client_id", &self.config.client_id), + ( + "redirect_uri", + &format!("{}/deviceauth/callback", self.config.auth_endpoint), + ), + ]) + .send() + .await + .map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Token exchange failed: {}", e), + })?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Token exchange failed: HTTP {} -- {}", status, body), + }); + } + + let token_resp: TokenResponse = + resp.json() + .await + .map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to parse token response: {}", e), + })?; + + let session = OpenAiCodexSession { + access_token: token_resp.access_token, + refresh_token: token_resp.refresh_token, + expires_at: Utc::now() + + chrono::Duration::seconds(if token_resp.expires_in > 0 { + token_resp.expires_in + } else { + tracing::warn!("Token response has expires_in=0, defaulting to 3600s"); + 3600 + } as i64), + created_at: Utc::now(), + }; + + self.save_session(&session).await?; + self.set_session(session).await; + + println!(); + println!("Authentication successful!"); + println!(); + Ok(()) + } + + /// Refresh the access token using the refresh token. + pub async fn refresh_tokens(&self) -> Result<(), LlmError> { + let _guard = self.renewal_lock.lock().await; + + // Double-check: another task may have refreshed while we waited on the lock + if !self.needs_refresh().await { + return Ok(()); + } + + let refresh_token = { + let guard = self.session.read().await; + guard + .as_ref() + .map(|s| s.refresh_token.clone()) + .ok_or_else(|| LlmError::AuthFailed { + provider: "openai_codex".to_string(), + })? + }; + + let token_url = format!("{}/oauth/token", self.config.auth_endpoint); + let resp = self + .client + .post(&token_url) + .form(&[ + ("grant_type", "refresh_token"), + ("refresh_token", refresh_token.as_str()), + ("client_id", self.config.client_id.as_str()), + ]) + .send() + .await + .map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Token refresh request failed: {}", e), + })?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Token refresh failed: HTTP {} -- {}", status, body), + }); + } + + let token_resp: TokenResponse = + resp.json() + .await + .map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to parse refresh response: {}", e), + })?; + + let session = OpenAiCodexSession { + access_token: token_resp.access_token, + refresh_token: token_resp.refresh_token, + expires_at: Utc::now() + + chrono::Duration::seconds(if token_resp.expires_in > 0 { + token_resp.expires_in + } else { + tracing::warn!("Token response has expires_in=0, defaulting to 3600s"); + 3600 + } as i64), + created_at: Utc::now(), + }; + + self.save_session(&session).await?; + self.set_session(session).await; + + tracing::debug!("OpenAI Codex token refreshed successfully"); + Ok(()) + } + + /// Save session data to disk with restrictive permissions. + pub async fn save_session(&self, session: &OpenAiCodexSession) -> Result<(), LlmError> { + if let Some(parent) = self.config.session_path.parent() { + tokio::fs::create_dir_all(parent).await.map_err(|e| { + LlmError::Io(std::io::Error::new( + e.kind(), + format!("Failed to create session directory: {}", e), + )) + })?; + } + + let json = + serde_json::to_string_pretty(session).map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to serialize session: {}", e), + })?; + + tokio::fs::write(&self.config.session_path, &json) + .await + .map_err(|e| { + LlmError::Io(std::io::Error::new( + e.kind(), + format!("Failed to write session file: {}", e), + )) + })?; + + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + let perms = std::fs::Permissions::from_mode(0o600); + tokio::fs::set_permissions(&self.config.session_path, perms) + .await + .map_err(|e| { + LlmError::Io(std::io::Error::new( + e.kind(), + format!("Failed to set permissions: {}", e), + )) + })?; + } + + Ok(()) + } + + /// Load session from disk. + pub async fn load_session(&self) -> Result<(), LlmError> { + let data = tokio::fs::read_to_string(&self.config.session_path) + .await + .map_err(|e| { + LlmError::Io(std::io::Error::new( + e.kind(), + format!("Failed to read session file: {}", e), + )) + })?; + + let session: OpenAiCodexSession = + serde_json::from_str(&data).map_err(|e| LlmError::SessionRenewalFailed { + provider: "openai_codex".to_string(), + reason: format!("Failed to parse session file: {}", e), + })?; + + let mut guard = self.session.write().await; + *guard = Some(session); + tracing::info!( + "Loaded OpenAI Codex session from {}", + self.config.session_path.display() + ); + Ok(()) + } + + /// Set session directly (for testing or after auth). + pub async fn set_session(&self, session: OpenAiCodexSession) { + let mut guard = self.session.write().await; + *guard = Some(session); + } + + /// Handle a 401 response by refreshing, or re-authenticating. + pub async fn handle_auth_failure(&self) -> Result<(), LlmError> { + match self.refresh_tokens().await { + Ok(()) => Ok(()), + Err(_) => self.device_code_login().await, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::llm::codex_test_helpers::test_codex_config as test_config; + use tempfile::tempdir; + + #[tokio::test] + async fn test_save_and_load_session() { + let dir = tempdir().unwrap(); + let path = dir.path().join("session.json"); + let config = test_config(path.clone()); + + let mgr = OpenAiCodexSessionManager::new(config).unwrap(); + + // No session initially + assert!(!mgr.has_session().await); + + // Save a session + let session = OpenAiCodexSession { + access_token: "access_abc".to_string(), + refresh_token: "refresh_xyz".to_string(), + expires_at: chrono::Utc::now() + chrono::Duration::hours(1), + created_at: chrono::Utc::now(), + }; + mgr.save_session(&session).await.unwrap(); + mgr.set_session(session).await; + + assert!(mgr.has_session().await); + + // Load from disk in a new manager + let config2 = test_config(path); + let mgr2 = OpenAiCodexSessionManager::new(config2).unwrap(); + mgr2.load_session().await.unwrap(); + assert!(mgr2.has_session().await); + } + + #[tokio::test] + async fn test_needs_refresh_when_near_expiry() { + let dir = tempdir().unwrap(); + let config = test_config(dir.path().join("session.json")); + let mgr = OpenAiCodexSessionManager::new(config).unwrap(); + + // Token expiring in 2 minutes (margin is 300s = 5 min) + let session = OpenAiCodexSession { + access_token: "access_abc".to_string(), + refresh_token: "refresh_xyz".to_string(), + expires_at: chrono::Utc::now() + chrono::Duration::minutes(2), + created_at: chrono::Utc::now(), + }; + mgr.set_session(session).await; + + assert!(mgr.needs_refresh().await); + } + + #[test] + fn device_code_parse_error_redacts_body() { + // Regression: the parse error used to include raw body_text which could + // contain sensitive auth data. Now it only shows byte count. + let body_text = r#"{"secret_token":"sk-12345","error":"unexpected"}"#; + let err: Result = serde_json::from_str(body_text); + assert!(err.is_err()); + let e = err.unwrap_err(); + let error_msg = format!( + "Failed to parse device code response: {} ({} bytes)", + e, + body_text.len() + ); + assert!( + !error_msg.contains("sk-12345"), + "error message must not contain raw body: {error_msg}" + ); + assert!( + error_msg.contains("bytes"), + "error message should show byte count" + ); + } + + #[tokio::test] + async fn test_no_refresh_when_fresh() { + let dir = tempdir().unwrap(); + let config = test_config(dir.path().join("session.json")); + let mgr = OpenAiCodexSessionManager::new(config).unwrap(); + + // Token expiring in 30 minutes (margin is 300s = 5 min) + let session = OpenAiCodexSession { + access_token: "access_abc".to_string(), + refresh_token: "refresh_xyz".to_string(), + expires_at: chrono::Utc::now() + chrono::Duration::minutes(30), + created_at: chrono::Utc::now(), + }; + mgr.set_session(session).await; + + assert!(!mgr.needs_refresh().await); + } +} diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index 26001086..1741e860 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -132,7 +132,7 @@ fn round_f32_to_f64(val: f32) -> f64 { /// /// This is applied as a clone-and-transform at the provider boundary so the /// original tool definitions remain unchanged for other providers. -fn normalize_schema_strict(schema: &JsonValue) -> JsonValue { +pub(crate) fn normalize_schema_strict(schema: &JsonValue) -> JsonValue { let mut schema = schema.clone(); normalize_schema_recursive(&mut schema); schema diff --git a/src/llm/token_refreshing.rs b/src/llm/token_refreshing.rs new file mode 100644 index 00000000..c39ad324 --- /dev/null +++ b/src/llm/token_refreshing.rs @@ -0,0 +1,191 @@ +//! Token-refreshing LlmProvider decorator for OpenAI Codex. +//! +//! Wraps an `OpenAiCodexProvider` and: +//! - Pre-emptively refreshes the OAuth access token before each call if near expiry +//! - Updates the inner provider's token after refresh (no client rebuild needed) +//! - Retries once on `AuthFailed` / `SessionExpired` after refreshing +//! - Overrides `cost_per_token()` to return (0, 0) since billing is through subscription + +use std::sync::Arc; + +use async_trait::async_trait; +use rust_decimal::Decimal; +use secrecy::ExposeSecret; + +use crate::error::LlmError; +use crate::llm::openai_codex_provider::OpenAiCodexProvider; +use crate::llm::openai_codex_session::OpenAiCodexSessionManager; +use crate::llm::provider::{ + CompletionRequest, CompletionResponse, LlmProvider, ModelMetadata, ToolCompletionRequest, + ToolCompletionResponse, +}; + +/// Decorator that refreshes OAuth tokens before API calls and reports zero cost. +/// +/// The inner `OpenAiCodexProvider` manages its own token state, so after a +/// refresh we just call `update_token()` -- no client rebuild is needed. +pub struct TokenRefreshingProvider { + inner: Arc, + session: Arc, +} + +impl TokenRefreshingProvider { + pub fn new(inner: Arc, session: Arc) -> Self { + Self { inner, session } + } + + /// Push a fresh token from the session manager into the inner provider. + async fn update_inner_token(&self) -> Result<(), LlmError> { + let token = self.session.get_access_token().await?; + self.inner.update_token(token.expose_secret()).await?; + tracing::debug!("Updated inner provider token after refresh"); + Ok(()) + } + + /// Best-effort pre-emptive token refresh before an API call. + /// + /// If refresh fails (e.g., no refresh token), we log and continue so the + /// actual request still fires and the retry-on-auth-failure path can kick in. + async fn ensure_fresh_token(&self) { + if self.session.needs_refresh().await { + match self.session.refresh_tokens().await { + Ok(()) => { + if let Err(e) = self.update_inner_token().await { + tracing::warn!( + "Pre-emptive token update failed: {e}, will retry on auth failure" + ); + } + } + Err(e) => { + tracing::warn!( + "Pre-emptive token refresh failed: {e}, will retry on auth failure" + ); + } + } + } + } +} + +#[async_trait] +impl LlmProvider for TokenRefreshingProvider { + fn model_name(&self) -> &str { + self.inner.model_name() + } + + fn cost_per_token(&self) -> (Decimal, Decimal) { + (Decimal::ZERO, Decimal::ZERO) + } + + async fn complete(&self, request: CompletionRequest) -> Result { + self.ensure_fresh_token().await; + + match self.inner.complete(request.clone()).await { + Err(LlmError::AuthFailed { .. } | LlmError::SessionExpired { .. }) => { + tracing::info!("Auth failure during complete(), refreshing and retrying once"); + self.session.handle_auth_failure().await?; + self.update_inner_token().await?; + self.inner.complete(request).await + } + other => other, + } + } + + async fn complete_with_tools( + &self, + request: ToolCompletionRequest, + ) -> Result { + self.ensure_fresh_token().await; + + match self.inner.complete_with_tools(request.clone()).await { + Err(LlmError::AuthFailed { .. } | LlmError::SessionExpired { .. }) => { + tracing::info!( + "Auth failure during complete_with_tools(), refreshing and retrying once" + ); + self.session.handle_auth_failure().await?; + self.update_inner_token().await?; + self.inner.complete_with_tools(request).await + } + other => other, + } + } + + async fn list_models(&self) -> Result, LlmError> { + self.ensure_fresh_token().await; + self.inner.list_models().await + } + + async fn model_metadata(&self) -> Result { + self.ensure_fresh_token().await; + self.inner.model_metadata().await + } + + fn active_model_name(&self) -> String { + self.inner.model_name().to_string() + } + + fn effective_model_name(&self, requested_model: Option<&str>) -> String { + self.inner.effective_model_name(requested_model) + } + + fn set_model(&self, model: &str) -> Result<(), LlmError> { + self.inner.set_model(model) + } + + fn calculate_cost(&self, _input_tokens: u32, _output_tokens: u32) -> Decimal { + Decimal::ZERO + } + + fn cache_write_multiplier(&self) -> Decimal { + self.inner.cache_write_multiplier() + } + + fn cache_read_discount(&self) -> Decimal { + self.inner.cache_read_discount() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::llm::codex_test_helpers::{make_test_jwt, test_codex_config}; + use crate::llm::openai_codex_session::OpenAiCodexSessionManager; + use tempfile::tempdir; + + fn make_provider_and_session() -> (TokenRefreshingProvider, tempfile::TempDir) { + let dir = tempdir().unwrap(); + let config = test_codex_config(dir.path().join("session.json")); + let jwt = make_test_jwt("acct_test"); + let inner = Arc::new( + OpenAiCodexProvider::new(&config.model, &config.api_base_url, &jwt, 300) + .expect("provider creation should succeed"), + ); + let session = Arc::new(OpenAiCodexSessionManager::new(config).unwrap()); + (TokenRefreshingProvider::new(inner, session), dir) + } + + #[test] + fn test_model_name_delegates() { + let (provider, _dir) = make_provider_and_session(); + assert_eq!(provider.model_name(), "gpt-5.3-codex"); + } + + #[test] + fn test_cost_per_token_zero() { + let (provider, _dir) = make_provider_and_session(); + let (input, output) = provider.cost_per_token(); + assert_eq!(input, Decimal::ZERO); + assert_eq!(output, Decimal::ZERO); + } + + #[test] + fn test_calculate_cost_zero() { + let (provider, _dir) = make_provider_and_session(); + assert_eq!(provider.calculate_cost(1000, 500), Decimal::ZERO); + } + + #[test] + fn test_active_model_name_delegates() { + let (provider, _dir) = make_provider_and_session(); + assert_eq!(provider.active_model_name(), "gpt-5.3-codex"); + } +} diff --git a/src/main.rs b/src/main.rs index 9c482e1b..af310fc4 100644 --- a/src/main.rs +++ b/src/main.rs @@ -139,6 +139,47 @@ async fn async_main() -> anyhow::Result<()> { ) .await; } + Some(Command::Login { openai_codex }) => { + init_cli_tracing(); + if *openai_codex { + // Resolve codex config so OPENAI_CODEX_* env overrides are + // honoured even when LLM_BACKEND isn't set to openai_codex. + let codex_config = { + let config = Config::from_env() + .await + .map_err(|e| anyhow::anyhow!("{}", e))?; + config.llm.openai_codex.unwrap_or_else(|| { + use ironclaw::llm::OpenAiCodexConfig; + let mut cfg = OpenAiCodexConfig::default(); + if let Ok(v) = std::env::var("OPENAI_CODEX_AUTH_URL") { + cfg.auth_endpoint = v; + } + if let Ok(v) = std::env::var("OPENAI_CODEX_API_URL") { + cfg.api_base_url = v; + } + if let Ok(v) = std::env::var("OPENAI_CODEX_CLIENT_ID") { + cfg.client_id = v; + } + if let Ok(v) = std::env::var("OPENAI_CODEX_SESSION_PATH") { + cfg.session_path = std::path::PathBuf::from(v); + } + cfg + }) + }; + let mgr = ironclaw::llm::OpenAiCodexSessionManager::new(codex_config) + .map_err(|e| anyhow::anyhow!("{}", e))?; + mgr.device_code_login() + .await + .map_err(|e| anyhow::anyhow!("{}", e))?; + println!( + "OpenAI Codex authentication complete. Set LLM_BACKEND=openai_codex to use it." + ); + } else { + println!("Specify a provider to authenticate with:"); + println!(" ironclaw login --openai-codex (ChatGPT subscription)"); + } + return Ok(()); + } Some(Command::Onboard { skip_auth, channels_only, diff --git a/src/setup/wizard.rs b/src/setup/wizard.rs index 6935a619..aca5b91e 100644 --- a/src/setup/wizard.rs +++ b/src/setup/wizard.rs @@ -3,7 +3,7 @@ //! The wizard guides users through: //! 1. Database connection //! 2. Security (secrets master key) -//! 3. Inference provider (NEAR AI, Anthropic, OpenAI, Ollama, OpenAI-compatible) +//! 3. Inference provider (NEAR AI, Anthropic, OpenAI, OpenAI Codex, Ollama, OpenAI-compatible) //! 4. Model selection //! 5. Embeddings //! 6. Channel configuration @@ -1083,8 +1083,10 @@ impl SetupWizard { print_info(&format!("Current provider: {}", display)); println!(); - let is_known = - current == "nearai" || current == "bedrock" || registry.is_known(¤t); + let is_known = current == "nearai" + || current == "bedrock" + || current == "openai_codex" + || registry.is_known(¤t); if is_known && confirm("Keep current provider?", true).map_err(SetupError::Io)? { if current == "bedrock" { @@ -1093,6 +1095,10 @@ impl SetupWizard { print_info("Keeping existing AWS Bedrock configuration."); return Ok(()); } + if current == "openai_codex" { + print_info("Keeping existing OpenAI Codex configuration."); + return Ok(()); + } return self.run_provider_setup(¤t, ®istry).await; } @@ -1107,7 +1113,7 @@ impl SetupWizard { print_info("Select your inference provider:"); println!(); - // Build menu: NearAI first, then all registry providers with setup hints, then Bedrock + // Build menu: NearAI first, then OpenAI Codex, then registry providers, then Bedrock let selectable = registry.selectable(); let mut options: Vec = Vec::with_capacity(2 + selectable.len()); let mut provider_ids: Vec = Vec::with_capacity(2 + selectable.len()); @@ -1115,6 +1121,9 @@ impl SetupWizard { options.push("NEAR AI - multi-model access via NEAR account".to_string()); provider_ids.push("nearai".to_string()); + options.push("OpenAI Codex - ChatGPT subscription (Plus/Pro/Max)".to_string()); + provider_ids.push("openai_codex".to_string()); + for def in &selectable { let label = format!( "{:<17}- {}", @@ -1158,6 +1167,10 @@ impl SetupWizard { return self.setup_nearai().await; } + if provider_id == "openai_codex" { + return self.setup_openai_codex().await; + } + let def = registry .find(provider_id) .ok_or_else(|| SetupError::Config(format!("Unknown provider: {}", provider_id)))?; @@ -1490,6 +1503,29 @@ impl SetupWizard { Ok(()) } + /// OpenAI Codex (ChatGPT subscription) setup: device code OAuth flow. + async fn setup_openai_codex(&mut self) -> Result<(), SetupError> { + self.settings.llm_backend = Some("openai_codex".to_string()); + if self.settings.selected_model.is_some() { + self.settings.selected_model = None; + } + + use crate::config::OpenAiCodexConfig; + use crate::llm::OpenAiCodexSessionManager; + + let config = OpenAiCodexConfig::default(); + + let mgr = OpenAiCodexSessionManager::new(config).map_err(|e| { + SetupError::Config(format!("OpenAI Codex session manager init failed: {}", e)) + })?; + mgr.device_code_login().await.map_err(|e| { + SetupError::Config(format!("OpenAI Codex authentication failed: {}", e)) + })?; + + print_success("OpenAI Codex configured (ChatGPT subscription)"); + Ok(()) + } + /// Generic Ollama-style setup: just needs a base URL, no API key. fn setup_ollama_generic( &mut self, @@ -2963,6 +2999,7 @@ impl SetupWizard { "ollama" => "Ollama", "openai_compatible" => "OpenAI-compatible", "bedrock" => "AWS Bedrock", + "openai_codex" => "OpenAI Codex", other => other, }; println!(" Provider: {}", display);