//! OpenAI Codex OAuth Provider //! //! Implements OAuth authentication flow for OpenAI Codex API. //! Supports PKCE (Proof Key for Code Exchange) for secure authentication. use super::error::{ create_auth_error, create_config_error, create_token_refresh_error, ProviderError, }; use reqwest::Client; use serde::{Deserialize, Serialize}; use std::error::Error; use std::path::PathBuf; // OAuth Constants const OPENAI_AUTH_URL: &str = "https://auth.openai.com/oauth/authorize"; const OPENAI_TOKEN_URL: &str = "https://auth.openai.com/oauth/token"; const OPENAI_CLIENT_ID: &str = "app_EMoamEEZ73f0CkXaXp7hrann"; const DEFAULT_CALLBACK_PORT: u16 = 1455; const CODEX_API_BASE_URL: &str = "https://chatgpt.com/backend-api/codex"; const DEFAULT_API_BASE_URL: &str = "https://api.openai.com"; /// Codex OAuth credentials storage /// /// Stores OAuth tokens and user information for Codex authentication. /// Compatible with CLIProxyAPI's CodexTokenStorage format and Codex CLI official format. /// /// Supports multiple field name formats: /// - snake_case: `refresh_token`, `access_token`, `id_token`, `account_id`, `last_refresh` /// - camelCase: `refreshToken`, `accessToken`, `idToken`, `accountId`, `lastRefresh` /// /// 同时兼容 Codex CLI 的 API Key 登录格式: /// - `api_key` / `apiKey` /// - `api_base_url` / `apiBaseUrl` #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CodexCredentials { /// JWT ID token containing user claims #[serde(default, skip_serializing_if = "Option::is_none", alias = "idToken")] pub id_token: Option, /// OAuth2 access token for API access #[serde( default, skip_serializing_if = "Option::is_none", alias = "accessToken" )] pub access_token: Option, /// Refresh token for obtaining new access tokens #[serde( default, skip_serializing_if = "Option::is_none", alias = "refreshToken" )] pub refresh_token: Option, /// API Key(Codex CLI 支持通过 API Key 登录) #[serde(default, skip_serializing_if = "Option::is_none", alias = "apiKey")] pub api_key: Option, /// API Base URL(可选) #[serde(default, skip_serializing_if = "Option::is_none", alias = "apiBaseUrl")] pub api_base_url: Option, /// OpenAI account identifier #[serde(default, skip_serializing_if = "Option::is_none", alias = "accountId")] pub account_id: Option, /// Timestamp of last token refresh #[serde( default, skip_serializing_if = "Option::is_none", alias = "lastRefresh" )] pub last_refresh: Option, /// User email address #[serde(default, skip_serializing_if = "Option::is_none")] pub email: Option, /// Authentication provider type (always "codex") #[serde(default = "default_type")] pub r#type: String, /// Token expiration timestamp (RFC3339 format) /// Supports: `expired`, `expires_at`, `expiresAt` #[serde( default, skip_serializing_if = "Option::is_none", rename = "expired", alias = "expires_at", alias = "expiresAt" )] pub expires_at: Option, } fn default_type() -> String { "codex".to_string() } impl Default for CodexCredentials { fn default() -> Self { Self { id_token: None, access_token: None, refresh_token: None, api_key: None, api_base_url: None, account_id: None, last_refresh: None, email: None, r#type: default_type(), expires_at: None, } } } /// PKCE codes for OAuth2 authorization #[derive(Debug, Clone)] pub struct PKCECodes { /// Cryptographically random string for code verification pub code_verifier: String, /// SHA256 hash of code_verifier, base64url-encoded pub code_challenge: String, } impl PKCECodes { /// Generate new PKCE codes pub fn generate() -> Result> { use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; use rand::RngCore; use sha2::{Digest, Sha256}; // Generate 96 random bytes for code verifier let mut bytes = [0u8; 96]; rand::thread_rng().fill_bytes(&mut bytes); let code_verifier = URL_SAFE_NO_PAD.encode(bytes); // Generate code challenge using S256 method let mut hasher = Sha256::new(); hasher.update(code_verifier.as_bytes()); let hash = hasher.finalize(); let code_challenge = URL_SAFE_NO_PAD.encode(hash); Ok(Self { code_verifier, code_challenge, }) } } /// OAuth callback result #[derive(Debug, Clone)] pub struct OAuthCallbackResult { /// Authorization code from OAuth callback pub code: String, /// State parameter for CSRF protection pub state: String, /// Error message if authentication failed pub error: Option, } /// OAuth server for handling OAuth callbacks pub struct OAuthServer { port: u16, shutdown_tx: Option>, } impl OAuthServer { /// Create a new OAuth server on the specified port pub fn new(port: u16) -> Self { Self { port, shutdown_tx: None, } } /// Start the OAuth server and wait for a callback /// /// Returns the authorization code and state from the OAuth callback. /// The server will automatically shut down after receiving a callback or timeout. pub async fn wait_for_callback( &mut self, timeout: std::time::Duration, ) -> Result> { use axum::{extract::Query, response::Html, routing::get, Router}; use std::collections::HashMap; use tokio::sync::oneshot; let (result_tx, result_rx) = oneshot::channel::(); let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); self.shutdown_tx = Some(shutdown_tx); // Wrap result_tx in Arc for sharing across requests let result_tx = std::sync::Arc::new(tokio::sync::Mutex::new(Some(result_tx))); let result_tx_clone = result_tx.clone(); let callback_handler = move |Query(params): Query>| { let result_tx = result_tx_clone.clone(); async move { let code = params.get("code").cloned().unwrap_or_default(); let state = params.get("state").cloned().unwrap_or_default(); let error = params.get("error").cloned(); let result = OAuthCallbackResult { code, state, error: error.clone(), }; // Send result (ignore if already sent) if let Some(tx) = result_tx.lock().await.take() { let _ = tx.send(result); } // Return success HTML if error.is_some() { Html(OAUTH_ERROR_HTML.to_string()) } else { Html(OAUTH_SUCCESS_HTML.to_string()) } } }; let app = Router::new().route("/auth/callback", get(callback_handler)); let addr = std::net::SocketAddr::from(([127, 0, 0, 1], self.port)); let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| { if e.kind() == std::io::ErrorKind::AddrInUse { format!( "Port {} is already in use. Please close any application using this port.", self.port ) } else { format!("Failed to bind to port {}: {}", self.port, e) } })?; tracing::info!( "[CODEX] OAuth server listening on http://127.0.0.1:{}", self.port ); // Spawn server with graceful shutdown let server = axum::serve(listener, app).with_graceful_shutdown(async move { let _ = shutdown_rx.await; }); tokio::spawn(async move { if let Err(e) = server.await { tracing::error!("[CODEX] OAuth server error: {}", e); } }); // Wait for callback with timeout let result = tokio::time::timeout(timeout, result_rx).await; // Trigger shutdown if let Some(tx) = self.shutdown_tx.take() { let _ = tx.send(()); } match result { Ok(Ok(callback_result)) => { if let Some(ref error) = callback_result.error { Err(format!("OAuth error: {}", error).into()) } else { Ok(callback_result) } } Ok(Err(_)) => Err("OAuth callback channel closed unexpectedly".into()), Err(_) => { Err("OAuth callback timeout - no response received within the time limit".into()) } } } } // HTML templates for OAuth callback responses const OAUTH_SUCCESS_HTML: &str = r#" Authentication Successful

Authentication Successful!

You can close this window and return to ProxyCast.

"#; const OAUTH_ERROR_HTML: &str = r#" Authentication Failed

Authentication Failed

Please close this window and try again.

"#; /// Codex OAuth Provider /// /// Handles OAuth authentication and API calls for OpenAI Codex. pub struct CodexProvider { /// OAuth credentials pub credentials: CodexCredentials, /// HTTP client for API requests pub client: Client, /// Path to credentials file pub creds_path: Option, /// OAuth callback port pub callback_port: u16, } impl Default for CodexProvider { fn default() -> Self { Self { credentials: CodexCredentials::default(), client: Client::new(), creds_path: None, callback_port: DEFAULT_CALLBACK_PORT, } } } impl CodexProvider { /// Create a new CodexProvider instance pub fn new() -> Self { Self::default() } /// Create a new CodexProvider with a custom HTTP client pub fn with_client(client: Client) -> Self { Self { client, ..Self::default() } } /// Get the default credentials file path pub fn default_creds_path() -> PathBuf { dirs::home_dir() .unwrap_or_else(|| PathBuf::from(".")) .join(".codex") .join("auth.json") } /// Get the OAuth authorization URL pub fn get_auth_url(&self) -> &'static str { OPENAI_AUTH_URL } /// Get the OAuth token URL pub fn get_token_url(&self) -> &'static str { OPENAI_TOKEN_URL } /// Get the OAuth client ID pub fn get_client_id(&self) -> &'static str { OPENAI_CLIENT_ID } /// Get the redirect URI for OAuth callback pub fn get_redirect_uri(&self) -> String { format!("http://localhost:{}/auth/callback", self.callback_port) } /// Get the API base URL pub fn get_api_base_url(&self) -> &'static str { CODEX_API_BASE_URL } /// 获取已配置的 API Key(trim 后的非空值) fn get_api_key(&self) -> Option<&str> { self.credentials .api_key .as_deref() .map(|s| s.trim()) .filter(|s| !s.is_empty()) } fn build_responses_url(base_url: &str) -> String { let base = base_url.trim_end_matches('/'); if base.ends_with("/v1") { format!("{}/responses", base) } else { format!("{}/v1/responses", base) } } /// Load credentials from the default path pub async fn load_credentials(&mut self) -> Result<(), Box> { let path = Self::default_creds_path(); self.load_credentials_from_path_internal(&path).await } /// Load credentials from a specific path pub async fn load_credentials_from_path( &mut self, path: &str, ) -> Result<(), Box> { let path = PathBuf::from(path); self.load_credentials_from_path_internal(&path).await } async fn load_credentials_from_path_internal( &mut self, path: &PathBuf, ) -> Result<(), Box> { if tokio::fs::try_exists(&path).await.unwrap_or(false) { let content = tokio::fs::read_to_string(&path).await?; // 尝试解析凭证文件 let creds: CodexCredentials = serde_json::from_str(&content).map_err(|e| { tracing::error!("[CODEX] 凭证文件解析失败: {}. 文件路径: {:?}", e, path); format!("凭证文件格式错误: {}", e) })?; // 检查关键字段 let has_api_key = creds .api_key .as_deref() .map(|s| !s.trim().is_empty()) .unwrap_or(false); if creds.refresh_token.is_none() && !has_api_key { tracing::warn!( "[CODEX] 凭证文件缺少 refresh_token/api_key 字段。支持的字段名: refresh_token, refreshToken, api_key, apiKey" ); // 打印文件中的顶级字段名,帮助调试 if let Ok(json_value) = serde_json::from_str::(&content) { if let Some(obj) = json_value.as_object() { let keys: Vec<&String> = obj.keys().collect(); tracing::info!("[CODEX] 凭证文件包含的字段: {:?}", keys); } } } tracing::info!( "[CODEX] 凭证加载成功: has_access={}, has_refresh={}, has_api_key={}, email={:?}, path={:?}", creds.access_token.is_some(), creds.refresh_token.is_some(), has_api_key, creds.email, path ); self.credentials = creds; self.creds_path = Some(path.clone()); } else { tracing::warn!("[CODEX] 凭证文件不存在: {:?}", path); return Err(format!("凭证文件不存在: {:?}", path).into()); } Ok(()) } /// Save credentials to file pub async fn save_credentials(&self) -> Result<(), Box> { let path = self .creds_path .clone() .unwrap_or_else(Self::default_creds_path); if let Some(parent) = path.parent() { tokio::fs::create_dir_all(parent).await?; } let content = serde_json::to_string_pretty(&self.credentials)?; tokio::fs::write(&path, content).await?; tracing::info!("[CODEX] Credentials saved to {:?}", path); Ok(()) } /// Check if the access token is expired pub fn is_token_expired(&self) -> bool { // API Key 模式:不涉及过期概念 if self.get_api_key().is_some() { return false; } if let Some(expires_str) = &self.credentials.expires_at { if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expires_str) { let now = chrono::Utc::now(); // Consider expired if less than 5 minutes remaining return expires < now + chrono::Duration::minutes(5); } } // If no expiry info, assume expired to be safe true } /// Check if credentials are valid (has access token and not expired) pub fn is_valid(&self) -> bool { if self.get_api_key().is_some() { return true; } self.credentials.access_token.is_some() && !self.is_token_expired() } /// Generate the OAuth authorization URL with PKCE pub fn generate_auth_url( &self, state: &str, pkce_codes: &PKCECodes, ) -> Result> { let params = [ ("client_id", OPENAI_CLIENT_ID), ("response_type", "code"), ("redirect_uri", &self.get_redirect_uri()), ("scope", "openid email profile offline_access"), ("state", state), ("code_challenge", &pkce_codes.code_challenge), ("code_challenge_method", "S256"), ("prompt", "login"), ("id_token_add_organizations", "true"), ("codex_cli_simplified_flow", "true"), ]; let query = serde_urlencoded::to_string(¶ms)?; Ok(format!("{}?{}", OPENAI_AUTH_URL, query)) } /// Generate a random state string for CSRF protection pub fn generate_state() -> Result> { use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; use rand::RngCore; let mut bytes = [0u8; 32]; rand::thread_rng().fill_bytes(&mut bytes); Ok(URL_SAFE_NO_PAD.encode(bytes)) } /// Exchange authorization code for tokens pub async fn exchange_code_for_tokens( &mut self, code: &str, pkce_codes: &PKCECodes, ) -> Result<(), Box> { let params = [ ("grant_type", "authorization_code"), ("client_id", OPENAI_CLIENT_ID), ("code", code), ("redirect_uri", &self.get_redirect_uri()), ("code_verifier", &pkce_codes.code_verifier), ]; let resp = self .client .post(OPENAI_TOKEN_URL) .header("Content-Type", "application/x-www-form-urlencoded") .header("Accept", "application/json") .form(¶ms) .send() .await?; if !resp.status().is_success() { let status = resp.status(); let body = resp.text().await.unwrap_or_default(); return Err(format!("Token exchange failed: {} - {}", status, body).into()); } let data: serde_json::Value = resp.json().await?; // Parse token response let access_token = data["access_token"] .as_str() .ok_or("No access_token in response")? .to_string(); let refresh_token = data["refresh_token"].as_str().map(|s| s.to_string()); let id_token = data["id_token"].as_str().map(|s| s.to_string()); let expires_in = data["expires_in"].as_i64().unwrap_or(3600); // Parse ID token to extract user info let (account_id, email) = if let Some(ref id_token) = id_token { parse_jwt_claims(id_token) } else { (None, None) }; // Calculate expiration time let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in); self.credentials = CodexCredentials { id_token, access_token: Some(access_token), refresh_token, api_key: None, api_base_url: None, account_id, last_refresh: Some(chrono::Utc::now().to_rfc3339()), email, r#type: "codex".to_string(), expires_at: Some(expires_at.to_rfc3339()), }; // Save credentials self.save_credentials().await?; tracing::info!( "[CODEX] Token exchange successful, email={:?}", self.credentials.email ); Ok(()) } /// Refresh the access token using the refresh token /// /// Supports three authentication modes (in priority order): /// 1. **API Key Mode**: Returns the API key directly (no refresh needed) /// 2. **OAuth Mode**: Refreshes the access token using the refresh token /// 3. **Access Token Mode**: Returns the existing access token (may be expired) /// /// # Returns /// * `Ok(String)` - The access token or API key /// * `Err` - If no credentials are available /// /// # Examples /// ```no_run /// // API Key mode /// provider.credentials.api_key = Some("sk-test".to_string()); /// let token = provider.refresh_token().await?; // Returns "sk-test" /// /// // OAuth mode /// provider.credentials.refresh_token = Some("refresh_token".to_string()); /// let token = provider.refresh_token().await?; // Refreshes and returns new access_token /// /// // Access Token mode (fallback) /// provider.credentials.access_token = Some("access_token".to_string()); /// let token = provider.refresh_token().await?; // Returns "access_token" (with warning) /// ``` pub async fn refresh_token(&mut self) -> Result> { // 1. API Key 模式无需刷新(优先级最高) if let Some(api_key) = self.get_api_key() { return Ok(api_key.to_string()); } // 2. 无 refresh_token 时的降级处理 if self.credentials.refresh_token.is_none() { // 2a. 有 access_token:返回(可能过期,由上层处理) if let Some(ref access_token) = self.credentials.access_token { tracing::warn!("[CODEX] 没有 refresh_token,返回现有 access_token(可能已过期)"); return Ok(access_token.clone()); } // 2b. 无任何凭证:清晰的错误指导 return Err(create_config_error( "没有可用的认证凭证。请配置以下任一方式:\n\ 1. API Key 模式:在凭证文件中添加 api_key/apiKey 字段\n\ 2. OAuth 模式:使用 OAuth 登录获取 refresh_token\n\ 3. Access Token 模式:在凭证文件中添加 access_token/accessToken 字段", )); } // 3. OAuth 刷新流程(标准流程) let refresh_token = self.credentials.refresh_token.as_ref().unwrap(); tracing::info!("[CODEX] 正在刷新 access token"); let params = [ ("client_id", OPENAI_CLIENT_ID), ("grant_type", "refresh_token"), ("refresh_token", refresh_token.as_str()), ("scope", "openid profile email"), ]; let resp = self .client .post(OPENAI_TOKEN_URL) .header("Content-Type", "application/x-www-form-urlencoded") .header("Accept", "application/json") .form(¶ms) .send() .await .map_err(|e| Box::new(ProviderError::from(e)) as Box)?; if !resp.status().is_success() { let status = resp.status().as_u16(); let body = resp.text().await.unwrap_or_default(); tracing::error!("[CODEX] Token refresh failed: {} - {}", status, body); // Mark credentials as invalid on refresh failure self.mark_invalid(); return Err(create_token_refresh_error(status, &body, "CODEX")); } let data: serde_json::Value = resp .json() .await .map_err(|e| Box::new(ProviderError::from(e)) as Box)?; // Update credentials let new_access_token = data["access_token"] .as_str() .ok_or_else(|| create_auth_error("响应中没有 access_token"))? .to_string(); self.credentials.access_token = Some(new_access_token.clone()); if let Some(rt) = data["refresh_token"].as_str() { self.credentials.refresh_token = Some(rt.to_string()); } if let Some(id_token) = data["id_token"].as_str() { self.credentials.id_token = Some(id_token.to_string()); let (account_id, email) = parse_jwt_claims(id_token); if account_id.is_some() { self.credentials.account_id = account_id; } if email.is_some() { self.credentials.email = email; } } let expires_in = data["expires_in"].as_i64().unwrap_or(3600); let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in); self.credentials.expires_at = Some(expires_at.to_rfc3339()); self.credentials.last_refresh = Some(chrono::Utc::now().to_rfc3339()); // Save updated credentials self.save_credentials().await?; tracing::info!("[CODEX] Token refresh successful"); Ok(new_access_token) } /// Refresh token with retry mechanism /// /// Attempts to refresh the token up to `max_retries` times with linear backoff (1s, 2s, 3s). /// Marks credentials as invalid if all retries fail. /// /// # Arguments /// * `max_retries` - Maximum number of retry attempts (typically 3) /// /// # Returns /// * `Ok(String)` - The new access token on success /// * `Err` - Error if all retries fail pub async fn refresh_token_with_retry( &mut self, max_retries: u32, ) -> Result> { let mut last_error = None; for attempt in 0..max_retries { if attempt > 0 { // Linear backoff: 1s, 2s, 3s, ... (as per Requirements 8.2) let delay = std::time::Duration::from_secs((attempt) as u64); tracing::info!( "[CODEX] Retry attempt {}/{} after {:?}", attempt + 1, max_retries, delay ); tokio::time::sleep(delay).await; } match self.refresh_token().await { Ok(token) => { if attempt > 0 { tracing::info!( "[CODEX] Token refresh succeeded on attempt {}", attempt + 1 ); } return Ok(token); } Err(e) => { tracing::warn!( "[CODEX] Token refresh attempt {}/{} failed: {}", attempt + 1, max_retries, e ); last_error = Some(e); } } } // All retries failed - mark as invalid self.mark_invalid(); tracing::error!( "[CODEX] Token refresh failed after {} attempts", max_retries ); Err(last_error.unwrap_or_else(|| create_auth_error("Token 刷新失败,请重新登录"))) } /// Check if token needs refresh (expiring within the specified duration) pub fn needs_refresh(&self, lead_time: chrono::Duration) -> bool { // API Key 模式无需刷新 if self.get_api_key().is_some() { return false; } if self.credentials.access_token.is_none() { return true; } if let Some(expires_str) = &self.credentials.expires_at { if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expires_str) { let now = chrono::Utc::now(); return expires < now + lead_time; } } // If no expiry info, assume needs refresh true } /// Ensure token is valid, refreshing if necessary /// /// This is the recommended method to call before making API requests. /// It will automatically refresh the token if it's expired or about to expire. pub async fn ensure_valid_token(&mut self) -> Result> { // 兼容 Codex CLI 的 API Key 登录:auth.json 只有 api_key,没有 refresh_token if let Some(api_key) = self.get_api_key() { return Ok(api_key.to_string()); } // Refresh if token expires within 5 minutes let lead_time = chrono::Duration::minutes(5); if self.needs_refresh(lead_time) { tracing::info!("[CODEX] Token needs refresh, attempting refresh with retry"); self.refresh_token_with_retry(3).await } else { self.credentials .access_token .clone() .ok_or_else(|| create_config_error("没有可用的 access_token")) } } /// Mark credentials as invalid (e.g., after refresh failure) pub fn mark_invalid(&mut self) { tracing::warn!("[CODEX] Marking credentials as invalid"); self.credentials.access_token = None; self.credentials.expires_at = None; } /// Get the access token, refreshing if necessary pub async fn get_access_token(&mut self) -> Result> { // API Key 模式直接返回 if let Some(api_key) = self.get_api_key() { return Ok(api_key.to_string()); } if self.is_token_expired() { self.refresh_token().await?; } self.credentials .access_token .clone() .ok_or_else(|| create_config_error("没有可用的 access_token")) } /// Perform OAuth login flow /// /// Opens a browser for OAuth authentication and waits for the callback. /// Returns the email of the authenticated user on success. pub async fn oauth_login(&mut self) -> Result> { tracing::info!("[CODEX] Starting OAuth login flow"); // Generate PKCE codes and state let pkce_codes = PKCECodes::generate()?; let state = Self::generate_state()?; // Generate authorization URL let auth_url = self.generate_auth_url(&state, &pkce_codes)?; // Start OAuth server let mut oauth_server = OAuthServer::new(self.callback_port); // Open browser tracing::info!("[CODEX] Opening browser for authentication"); if let Err(e) = open::that(&auth_url) { tracing::warn!( "[CODEX] Failed to open browser: {}. Please open the URL manually.", e ); println!( "Please open the following URL in your browser:\n{}", auth_url ); } // Wait for callback (5 minute timeout) let timeout = std::time::Duration::from_secs(300); let callback_result = oauth_server.wait_for_callback(timeout).await?; // Verify state if callback_result.state != state { return Err("OAuth state mismatch - possible CSRF attack".into()); } // Exchange code for tokens self.exchange_code_for_tokens(&callback_result.code, &pkce_codes) .await?; let email = self .credentials .email .clone() .unwrap_or_else(|| "unknown".to_string()); tracing::info!("[CODEX] OAuth login successful for {}", email); Ok(email) } /// Perform OAuth login without opening browser (for headless/SSH environments) /// /// Returns the authorization URL that the user should open manually. pub fn start_oauth_login( &self, ) -> Result<(String, PKCECodes, String), Box> { let pkce_codes = PKCECodes::generate()?; let state = Self::generate_state()?; let auth_url = self.generate_auth_url(&state, &pkce_codes)?; Ok((auth_url, pkce_codes, state)) } /// Complete OAuth login after receiving callback pub async fn complete_oauth_login( &mut self, code: &str, pkce_codes: &PKCECodes, expected_state: &str, received_state: &str, ) -> Result> { // Verify state if received_state != expected_state { return Err("OAuth state mismatch - possible CSRF attack".into()); } // Exchange code for tokens self.exchange_code_for_tokens(code, pkce_codes).await?; let email = self .credentials .email .clone() .unwrap_or_else(|| "unknown".to_string()); tracing::info!("[CODEX] OAuth login completed for {}", email); Ok(email) } /// Call the Codex API for chat completions /// /// Routes GPT model requests through the Codex OAuth endpoint. /// The request should be in OpenAI chat completion format. pub async fn call_api( &self, request: &serde_json::Value, ) -> Result> { enum AuthMode { ApiKey, OAuth, } let (token, mode) = match self.get_api_key() { Some(api_key) => (api_key, AuthMode::ApiKey), None => ( self.credentials .access_token .as_deref() .ok_or("No access token or api_key available")?, AuthMode::OAuth, ), }; // Build the Codex API URL let url = match mode { AuthMode::ApiKey => { let base_url = self .credentials .api_base_url .as_deref() .map(|s| s.trim()) .filter(|s| !s.is_empty()) .unwrap_or(DEFAULT_API_BASE_URL); Self::build_responses_url(base_url) } AuthMode::OAuth => format!("{}/responses", CODEX_API_BASE_URL), }; // Transform OpenAI chat completion request to Codex format let codex_request = transform_to_codex_format(request)?; tracing::debug!("[CODEX] Calling API: {}", url); let mut req = self .client .post(&url) .header("Authorization", format!("Bearer {}", token)) .header("Content-Type", "application/json") .header("Accept", "text/event-stream") .header("Openai-Beta", "responses=experimental") .json(&codex_request); if matches!(mode, AuthMode::OAuth) { req = req .header("Version", "0.21.0") .header( "User-Agent", "codex_cli_rs/0.50.0 (Mac OS 26.0.1; arm64) Apple_Terminal/464", ) .header("Originator", "codex_cli_rs") .header("Session_id", uuid::Uuid::new_v4().to_string()) // Add account ID header if available .header( "Chatgpt-Account-Id", self.credentials.account_id.as_deref().unwrap_or(""), ); } let resp = req.send().await?; Ok(resp) } /// Call the Codex API with streaming response pub async fn call_api_stream( &self, request: &serde_json::Value, ) -> Result> { // Same as call_api - Codex always returns SSE stream self.call_api(request).await } /// Check if this provider supports the given model pub fn supports_model(model: &str) -> bool { let model_lower = model.to_lowercase(); model_lower.starts_with("gpt-") || model_lower.starts_with("o1") || model_lower.starts_with("o3") || model_lower.starts_with("o4") || model_lower.contains("codex") } } /// Parse JWT token to extract account_id and email /// /// Extracts user information from the JWT ID token returned by OpenAI OAuth. /// The account_id is extracted from the `chatgpt_account_id` field in the /// `https://api.openai.com/auth` claim, which is required for Codex API calls. fn parse_jwt_claims(token: &str) -> (Option, Option) { use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; let parts: Vec<&str> = token.split('.').collect(); if parts.len() != 3 { tracing::warn!( "[CODEX] Invalid JWT token format: expected 3 parts, got {}", parts.len() ); return (None, None); } // Decode payload (second part) - JWT uses URL-safe base64 without padding let payload = match URL_SAFE_NO_PAD.decode(parts[1]) { Ok(bytes) => bytes, Err(_) => { // Try with padding added let padded = format!("{}{}", parts[1], "=".repeat((4 - parts[1].len() % 4) % 4)); match base64::engine::general_purpose::URL_SAFE.decode(&padded) { Ok(bytes) => bytes, Err(e) => { tracing::warn!("[CODEX] Failed to decode JWT payload: {}", e); return (None, None); } } } }; let claims: serde_json::Value = match serde_json::from_slice(&payload) { Ok(v) => v, Err(e) => { tracing::warn!("[CODEX] Failed to parse JWT claims: {}", e); return (None, None); } }; // Extract email from standard claim let email = claims["email"].as_str().map(|s| s.to_string()); // Extract account_id from OpenAI-specific claims // Priority: chatgpt_account_id > user_id > sub // The chatgpt_account_id is the correct field for Codex API calls let auth_info = &claims["https://api.openai.com/auth"]; let account_id = auth_info["chatgpt_account_id"] .as_str() .or_else(|| auth_info["user_id"].as_str()) .or_else(|| claims["sub"].as_str()) .map(|s| s.to_string()); tracing::debug!( "[CODEX] JWT parsed: email={:?}, account_id={:?}", email, account_id ); (account_id, email) } /// Transform OpenAI chat completion request to Codex format fn transform_to_codex_format( request: &serde_json::Value, ) -> Result> { let model = request["model"].as_str().unwrap_or("gpt-4o"); let messages = request["messages"].as_array(); let stream = request["stream"].as_bool().unwrap_or(true); // Build input array from messages let mut input = Vec::new(); let mut instructions = None; if let Some(msgs) = messages { for msg in msgs { let role = msg["role"].as_str().unwrap_or("user"); let content = &msg["content"]; match role { "system" => { // System messages become instructions if let Some(text) = content.as_str() { instructions = Some(text.to_string()); } } "user" | "assistant" => { let content_parts = if let Some(text) = content.as_str() { vec![serde_json::json!({"type": "input_text", "text": text})] } else if let Some(arr) = content.as_array() { arr.iter() .filter_map(|part| { if let Some(text) = part["text"].as_str() { Some(serde_json::json!({"type": "input_text", "text": text})) } else { None } }) .collect() } else { vec![] }; input.push(serde_json::json!({ "type": "message", "role": role, "content": content_parts })); } "tool" => { // Tool results let tool_call_id = msg["tool_call_id"].as_str().unwrap_or(""); let output = content.as_str().unwrap_or(""); input.push(serde_json::json!({ "type": "function_call_output", "call_id": tool_call_id, "output": output })); } _ => {} } } } // Build tools array if present let tools = request["tools"].as_array().map(|tools| { tools .iter() .filter_map(|tool| { let func = &tool["function"]; Some(serde_json::json!({ "type": "function", "name": func["name"], "description": func["description"], "parameters": func["parameters"] })) }) .collect::>() }); // Build the Codex request let mut codex_request = serde_json::json!({ "model": model, "input": input, "stream": stream }); if let Some(inst) = instructions { codex_request["instructions"] = serde_json::json!(inst); } if let Some(t) = tools { codex_request["tools"] = serde_json::json!(t); } // Copy over other parameters if let Some(temp) = request["temperature"].as_f64() { codex_request["temperature"] = serde_json::json!(temp); } if let Some(max_tokens) = request["max_tokens"].as_i64() { codex_request["max_output_tokens"] = serde_json::json!(max_tokens); } if let Some(top_p) = request["top_p"].as_f64() { codex_request["top_p"] = serde_json::json!(top_p); } // Handle reasoning effort for o1/o3/o4 models if let Some(reasoning) = request.get("reasoning") { codex_request["reasoning"] = reasoning.clone(); } Ok(codex_request) } #[cfg(test)] mod tests { use super::*; #[test] fn test_codex_credentials_default() { let creds = CodexCredentials::default(); assert!(creds.access_token.is_none()); assert!(creds.refresh_token.is_none()); assert!(creds.api_key.is_none()); assert_eq!(creds.r#type, "codex"); } #[test] fn test_codex_credentials_serialization() { let creds = CodexCredentials { access_token: Some("test_token".to_string()), refresh_token: Some("test_refresh".to_string()), email: Some("test@example.com".to_string()), ..Default::default() }; let json = serde_json::to_string(&creds).unwrap(); assert!(json.contains("test_token")); assert!(json.contains("test@example.com")); let parsed: CodexCredentials = serde_json::from_str(&json).unwrap(); assert_eq!(parsed.access_token, creds.access_token); assert_eq!(parsed.email, creds.email); } #[test] fn test_codex_credentials_camel_case_alias() { // 测试 camelCase 字段名的支持(Codex CLI 官方格式) let json = r#"{ "idToken": "test_id_token", "accessToken": "test_access_token", "refreshToken": "test_refresh_token", "accountId": "test_account_id", "lastRefresh": "2024-01-01T00:00:00Z", "email": "test@example.com", "type": "codex", "expiresAt": "2024-12-31T23:59:59Z" }"#; let creds: CodexCredentials = serde_json::from_str(json).unwrap(); assert_eq!(creds.id_token, Some("test_id_token".to_string())); assert_eq!(creds.access_token, Some("test_access_token".to_string())); assert_eq!(creds.refresh_token, Some("test_refresh_token".to_string())); assert_eq!(creds.account_id, Some("test_account_id".to_string())); assert_eq!(creds.last_refresh, Some("2024-01-01T00:00:00Z".to_string())); assert_eq!(creds.email, Some("test@example.com".to_string())); assert_eq!(creds.expires_at, Some("2024-12-31T23:59:59Z".to_string())); } #[test] fn test_codex_credentials_snake_case() { // 测试 snake_case 字段名的支持(CLIProxyAPI 格式) let json = r#"{ "id_token": "test_id_token", "access_token": "test_access_token", "refresh_token": "test_refresh_token", "account_id": "test_account_id", "last_refresh": "2024-01-01T00:00:00Z", "email": "test@example.com", "type": "codex", "expired": "2024-12-31T23:59:59Z" }"#; let creds: CodexCredentials = serde_json::from_str(json).unwrap(); assert_eq!(creds.id_token, Some("test_id_token".to_string())); assert_eq!(creds.access_token, Some("test_access_token".to_string())); assert_eq!(creds.refresh_token, Some("test_refresh_token".to_string())); assert_eq!(creds.account_id, Some("test_account_id".to_string())); assert_eq!(creds.last_refresh, Some("2024-01-01T00:00:00Z".to_string())); assert_eq!(creds.email, Some("test@example.com".to_string())); assert_eq!(creds.expires_at, Some("2024-12-31T23:59:59Z".to_string())); } #[test] fn test_codex_credentials_api_key_fields() { let json = r#"{ "api_key": "sk-test", "api_base_url": "https://api.openai.com/v1" }"#; let creds: CodexCredentials = serde_json::from_str(json).unwrap(); assert_eq!(creds.api_key, Some("sk-test".to_string())); assert_eq!( creds.api_base_url, Some("https://api.openai.com/v1".to_string()) ); let json2 = r#"{ "apiKey": "sk-test-2", "apiBaseUrl": "https://example.com/v1" }"#; let creds2: CodexCredentials = serde_json::from_str(json2).unwrap(); assert_eq!(creds2.api_key, Some("sk-test-2".to_string())); assert_eq!( creds2.api_base_url, Some("https://example.com/v1".to_string()) ); } #[test] fn test_codex_credentials_expires_at_alias() { // 测试 expires_at 字段的多种别名 let json1 = r#"{"expired": "2024-12-31T23:59:59Z"}"#; let json2 = r#"{"expires_at": "2024-12-31T23:59:59Z"}"#; let json3 = r#"{"expiresAt": "2024-12-31T23:59:59Z"}"#; let creds1: CodexCredentials = serde_json::from_str(json1).unwrap(); let creds2: CodexCredentials = serde_json::from_str(json2).unwrap(); let creds3: CodexCredentials = serde_json::from_str(json3).unwrap(); assert_eq!(creds1.expires_at, Some("2024-12-31T23:59:59Z".to_string())); assert_eq!(creds2.expires_at, Some("2024-12-31T23:59:59Z".to_string())); assert_eq!(creds3.expires_at, Some("2024-12-31T23:59:59Z".to_string())); } #[test] fn test_pkce_generation() { let pkce = PKCECodes::generate().unwrap(); assert!(!pkce.code_verifier.is_empty()); assert!(!pkce.code_challenge.is_empty()); // Verifier should be 128 chars (96 bytes base64 encoded) assert_eq!(pkce.code_verifier.len(), 128); } #[test] fn test_codex_provider_default() { let provider = CodexProvider::new(); assert_eq!(provider.callback_port, DEFAULT_CALLBACK_PORT); assert!(provider.credentials.access_token.is_none()); } #[test] fn test_build_responses_url() { assert_eq!( CodexProvider::build_responses_url("https://api.openai.com"), "https://api.openai.com/v1/responses" ); assert_eq!( CodexProvider::build_responses_url("https://api.openai.com/v1"), "https://api.openai.com/v1/responses" ); assert_eq!( CodexProvider::build_responses_url("https://example.com/v1/"), "https://example.com/v1/responses" ); } #[tokio::test] async fn test_ensure_valid_token_prefers_api_key() { let mut provider = CodexProvider::new(); provider.credentials.api_key = Some("sk-test".to_string()); let token = provider.ensure_valid_token().await.unwrap(); assert_eq!(token, "sk-test"); } #[test] fn test_generate_auth_url() { let provider = CodexProvider::new(); let pkce = PKCECodes::generate().unwrap(); let state = "test_state"; let url = provider.generate_auth_url(state, &pkce).unwrap(); assert!(url.starts_with(OPENAI_AUTH_URL)); assert!(url.contains("client_id=")); assert!(url.contains("code_challenge=")); assert!(url.contains("state=test_state")); } #[test] fn test_parse_jwt_claims_with_sub() { // Mock JWT with only sub claim (fallback case) // Header: {"alg":"RS256","typ":"JWT"} // Payload: {"email":"test@example.com","sub":"user123"} let mock_jwt = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJlbWFpbCI6InRlc3RAZXhhbXBsZS5jb20iLCJzdWIiOiJ1c2VyMTIzIn0.signature"; let (account_id, email) = parse_jwt_claims(mock_jwt); assert_eq!(email, Some("test@example.com".to_string())); assert_eq!(account_id, Some("user123".to_string())); } #[test] fn test_parse_jwt_claims_with_chatgpt_account_id() { // Mock JWT with chatgpt_account_id in https://api.openai.com/auth claim // This is the preferred field for Codex API calls // Payload: {"email":"test@example.com","sub":"user123","https://api.openai.com/auth":{"chatgpt_account_id":"chatgpt_acc_123","user_id":"uid_456"}} use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#); let payload = URL_SAFE_NO_PAD.encode(r#"{"email":"test@example.com","sub":"user123","https://api.openai.com/auth":{"chatgpt_account_id":"chatgpt_acc_123","user_id":"uid_456"}}"#); let mock_jwt = format!("{}.{}.signature", header, payload); let (account_id, email) = parse_jwt_claims(&mock_jwt); assert_eq!(email, Some("test@example.com".to_string())); // Should prefer chatgpt_account_id over user_id and sub assert_eq!(account_id, Some("chatgpt_acc_123".to_string())); } #[test] fn test_parse_jwt_claims_with_user_id() { // Mock JWT with user_id but no chatgpt_account_id // Payload: {"email":"test@example.com","sub":"user123","https://api.openai.com/auth":{"user_id":"uid_456"}} use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#); let payload = URL_SAFE_NO_PAD.encode(r#"{"email":"test@example.com","sub":"user123","https://api.openai.com/auth":{"user_id":"uid_456"}}"#); let mock_jwt = format!("{}.{}.signature", header, payload); let (account_id, email) = parse_jwt_claims(&mock_jwt); assert_eq!(email, Some("test@example.com".to_string())); // Should use user_id when chatgpt_account_id is not present assert_eq!(account_id, Some("uid_456".to_string())); } #[test] fn test_parse_jwt_claims_invalid_token() { // Invalid JWT format let (account_id, email) = parse_jwt_claims("invalid.token"); assert_eq!(account_id, None); assert_eq!(email, None); // Empty token let (account_id, email) = parse_jwt_claims(""); assert_eq!(account_id, None); assert_eq!(email, None); } #[test] fn test_is_token_expired() { let mut provider = CodexProvider::new(); // No expiry - should be considered expired assert!(provider.is_token_expired()); // API Key 模式 - 不应视为过期 provider.credentials.api_key = Some("sk-test".to_string()); assert!(!provider.is_token_expired()); // Expired token provider.credentials.api_key = None; provider.credentials.expires_at = Some("2020-01-01T00:00:00Z".to_string()); assert!(provider.is_token_expired()); // Valid token (far future) provider.credentials.expires_at = Some("2099-01-01T00:00:00Z".to_string()); assert!(!provider.is_token_expired()); } #[test] fn test_supports_model() { // GPT models should be supported assert!(CodexProvider::supports_model("gpt-4")); assert!(CodexProvider::supports_model("gpt-4o")); assert!(CodexProvider::supports_model("gpt-4-turbo")); assert!(CodexProvider::supports_model("GPT-4")); // Case insensitive // O-series models should be supported assert!(CodexProvider::supports_model("o1")); assert!(CodexProvider::supports_model("o1-preview")); assert!(CodexProvider::supports_model("o3")); assert!(CodexProvider::supports_model("o4-mini")); // Codex models should be supported (contains "codex") assert!(CodexProvider::supports_model("codex-mini")); assert!(CodexProvider::supports_model("gpt-4-codex")); // Non-GPT models should not be supported assert!(!CodexProvider::supports_model("claude-3")); assert!(!CodexProvider::supports_model("gemini-pro")); assert!(!CodexProvider::supports_model("llama-2")); } #[test] fn test_transform_to_codex_format_basic() { let request = serde_json::json!({ "model": "gpt-4o", "messages": [ {"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "Hello!"} ], "stream": true }); let result = transform_to_codex_format(&request).unwrap(); assert_eq!(result["model"], "gpt-4o"); assert_eq!(result["stream"], true); assert_eq!(result["instructions"], "You are a helpful assistant."); let input = result["input"].as_array().unwrap(); assert_eq!(input.len(), 1); // Only user message, system becomes instructions assert_eq!(input[0]["role"], "user"); } #[test] fn test_transform_to_codex_format_with_tools() { let request = serde_json::json!({ "model": "gpt-4o", "messages": [ {"role": "user", "content": "What's the weather?"} ], "tools": [ { "type": "function", "function": { "name": "get_weather", "description": "Get weather info", "parameters": {"type": "object"} } } ] }); let result = transform_to_codex_format(&request).unwrap(); let tools = result["tools"].as_array().unwrap(); assert_eq!(tools.len(), 1); assert_eq!(tools[0]["name"], "get_weather"); assert_eq!(tools[0]["description"], "Get weather info"); } #[test] fn test_transform_to_codex_format_with_parameters() { let request = serde_json::json!({ "model": "gpt-4o", "messages": [{"role": "user", "content": "Hi"}], "temperature": 0.7, "max_tokens": 1000, "top_p": 0.9 }); let result = transform_to_codex_format(&request).unwrap(); assert_eq!(result["temperature"], 0.7); assert_eq!(result["max_output_tokens"], 1000); assert_eq!(result["top_p"], 0.9); } #[tokio::test] async fn test_refresh_token_with_only_access_token() { // 场景:只有 access_token(无 refresh_token 和 api_key) let mut provider = CodexProvider::new(); provider.credentials.access_token = Some("test_access_token".to_string()); provider.credentials.refresh_token = None; provider.credentials.api_key = None; let result = provider.refresh_token().await; assert!(result.is_ok()); assert_eq!(result.unwrap(), "test_access_token"); } #[tokio::test] async fn test_refresh_token_with_no_credentials() { // 场景:无任何凭证(api_key、refresh_token、access_token 均为 None) let mut provider = CodexProvider::new(); provider.credentials.api_key = None; provider.credentials.refresh_token = None; provider.credentials.access_token = None; let result = provider.refresh_token().await; assert!(result.is_err()); let error_msg = result.unwrap_err().to_string(); assert!(error_msg.contains("没有可用的认证凭证")); assert!(error_msg.contains("API Key 模式")); assert!(error_msg.contains("OAuth 模式")); assert!(error_msg.contains("Access Token 模式")); } #[tokio::test] async fn test_api_key_priority_over_refresh_token() { // 场景:同时有 api_key 和 refresh_token let mut provider = CodexProvider::new(); provider.credentials.api_key = Some("sk-test-api-key".to_string()); provider.credentials.refresh_token = Some("test_refresh_token".to_string()); provider.credentials.access_token = Some("test_access_token".to_string()); let result = provider.refresh_token().await; assert!(result.is_ok()); // 应该返回 API Key(优先级最高) assert_eq!(result.unwrap(), "sk-test-api-key"); } #[tokio::test] async fn test_refresh_token_with_expired_access_token() { // 场景:只有 access_token(已过期) let mut provider = CodexProvider::new(); provider.credentials.access_token = Some("expired_access_token".to_string()); provider.credentials.expires_at = Some("2020-01-01T00:00:00Z".to_string()); provider.credentials.refresh_token = None; provider.credentials.api_key = None; let result = provider.refresh_token().await; assert!(result.is_ok()); // 应该返回 access_token(即使已过期,由上层处理) assert_eq!(result.unwrap(), "expired_access_token"); } }