diff --git a/codex-rs/app-server/src/request_processors/thread_lifecycle.rs b/codex-rs/app-server/src/request_processors/thread_lifecycle.rs index c777f9403b..422adbb578 100644 --- a/codex-rs/app-server/src/request_processors/thread_lifecycle.rs +++ b/codex-rs/app-server/src/request_processors/thread_lifecycle.rs @@ -314,9 +314,12 @@ pub(super) async fn ensure_listener_task_running( }; if let Some(worker) = &turn_cost_worker { - worker.observe_event(conversation_id, &event, || { - conversation.session_telemetry() - }); + worker.observe_event( + conversation_id, + config.as_ref(), + &event, + || conversation.session_telemetry(), + ); } // Track the event before emitting any typed translations diff --git a/codex-rs/app-server/src/turn_cost_worker.rs b/codex-rs/app-server/src/turn_cost_worker.rs index db46941161..08c9b82bf0 100644 --- a/codex-rs/app-server/src/turn_cost_worker.rs +++ b/codex-rs/app-server/src/turn_cost_worker.rs @@ -1,10 +1,12 @@ use codex_backend_client::ApiKeyTurnCost; use codex_backend_client::ApiKeyTurnCostStatus; use codex_backend_client::Client as BackendClient; +use codex_backend_client::RequestError; use codex_config::types::OtelExporterKind; use codex_core::config::Config; use codex_login::AuthManager; -use codex_login::CodexAuth; +use codex_model_provider::SharedModelProvider; +use codex_model_provider::create_model_provider; use codex_otel::SessionTelemetry; use codex_protocol::ThreadId; use codex_protocol::auth::AuthMode; @@ -35,7 +37,8 @@ pub(crate) struct TurnCostWorker { #[derive(Clone)] pub(crate) struct TurnCostWorkerHandle { sender: mpsc::Sender, - auth_manager: Arc, + backend: TurnCostBackend, + config: Arc, } enum TurnCostObservationKind { @@ -72,10 +75,16 @@ struct TurnCostEntry { struct WorkerRuntime { config: Arc, - auth_manager: Arc, + backend: TurnCostBackend, turns: HashMap, } +#[derive(Clone)] +enum TurnCostBackend { + OpenAiApiKey(Arc), + ModelProvider(SharedModelProvider), +} + #[derive(Clone, Copy, Debug, PartialEq, Eq)] enum BackendAvailability { AwaitingAuthChange, @@ -89,15 +98,24 @@ impl TurnCostWorker { if !matches!( config.otel.exporter, OtelExporterKind::OtlpHttp { .. } | OtelExporterKind::OtlpGrpc { .. } - ) || !config.model_provider.is_openai() + ) || config.model_provider.is_amazon_bedrock() { return None; } + let is_openai = config.model_provider.is_openai(); + let backend = if is_openai { + TurnCostBackend::OpenAiApiKey(Arc::clone(&auth_manager)) + } else { + TurnCostBackend::ModelProvider(create_model_provider( + config.model_provider.clone(), + Some(Arc::clone(&auth_manager)), + )) + }; let (sender, receiver) = mpsc::channel(OBSERVATION_CHANNEL_CAPACITY); let shutdown = CancellationToken::new(); let runtime = WorkerRuntime { config: Arc::clone(&config), - auth_manager: Arc::clone(&auth_manager), + backend: backend.clone(), turns: HashMap::new(), }; let worker_shutdown = shutdown.clone(); @@ -107,7 +125,8 @@ impl TurnCostWorker { Some(Self { handle: TurnCostWorkerHandle { sender, - auth_manager, + backend, + config, }, shutdown, _task: task, @@ -133,15 +152,21 @@ impl TurnCostWorkerHandle { pub(crate) fn observe_event( &self, thread_id: ThreadId, + thread_config: &Config, event: &Event, session_telemetry: impl FnOnce() -> SessionTelemetry, ) { - let Some(auth) = self.auth_manager.auth_cached() else { - return; - }; - if !auth.is_api_key_auth() { + if thread_config.model_provider != self.config.model_provider { return; } + if let TurnCostBackend::OpenAiApiKey(auth_manager) = &self.backend { + let Some(auth) = auth_manager.auth_cached() else { + return; + }; + if !auth.is_api_key_auth() { + return; + } + } let kind = match &event.msg { EventMsg::TurnStarted(_) => TurnCostObservationKind::Started { session_telemetry: Box::new(session_telemetry()), @@ -161,7 +186,12 @@ impl TurnCostWorkerHandle { impl WorkerRuntime { async fn run(self, receiver: mpsc::Receiver, shutdown: CancellationToken) { - let auth_changes = self.auth_manager.auth_change_receiver(); + let auth_changes = match &self.backend { + TurnCostBackend::OpenAiApiKey(auth_manager) => { + Some(auth_manager.auth_change_receiver()) + } + TurnCostBackend::ModelProvider(_) => None, + }; let backend_availability = self.probe_backend().await; self.run_with_backend_availability(receiver, shutdown, auth_changes, backend_availability) .await; @@ -171,7 +201,7 @@ impl WorkerRuntime { mut self, mut receiver: mpsc::Receiver, shutdown: CancellationToken, - mut auth_changes: tokio::sync::watch::Receiver, + mut auth_changes: Option>, mut backend_availability: BackendAvailability, ) { let mut ticker = tokio::time::interval(POLL_INTERVAL); @@ -180,7 +210,12 @@ impl WorkerRuntime { tokio::select! { biased; _ = shutdown.cancelled() => break, - changed = auth_changes.changed() => { + changed = async { + match auth_changes.as_mut() { + Some(auth_changes) => auth_changes.changed().await, + None => std::future::pending().await, + } + } => { if changed.is_err() { break; } @@ -215,44 +250,21 @@ impl WorkerRuntime { } async fn probe_backend(&self) -> BackendAvailability { - let Some(auth) = self.auth_manager.auth().await else { - return BackendAvailability::AwaitingAuthChange; - }; - if !auth.is_api_key_auth() { - return BackendAvailability::AwaitingAuthChange; - } - let provider = match self - .config - .model_provider - .to_api_provider(Some(AuthMode::ApiKey)) - { - Ok(provider) => provider, - Err(error) => { - warn!("failed to resolve OpenAI API-key provider headers: {error}"); - return BackendAvailability::RetryProbe; - } - }; - let client = BackendClient::from_auth( - self.config.chatgpt_base_url.clone(), - &auth, - self.config.http_client_factory(), - ); let probe_turn_ids = [uuid::Uuid::new_v4().to_string()]; - match tokio::time::timeout( - REQUEST_TIMEOUT, - client.query_api_key_turn_costs(&probe_turn_ids, &provider.headers), - ) - .await - { - Ok(Ok(_)) => BackendAvailability::Ready, + match tokio::time::timeout(REQUEST_TIMEOUT, self.query_turn_costs(&probe_turn_ids)).await { + Ok(Ok(Some(_))) => BackendAvailability::Ready, + Ok(Ok(None)) => match self.backend { + TurnCostBackend::OpenAiApiKey(_) => BackendAvailability::AwaitingAuthChange, + TurnCostBackend::ModelProvider(_) => BackendAvailability::Disabled, + }, Ok(Err(error)) => match error.status().map(|status| status.as_u16()) { - Some(401 | 403) => { + Some(401 | 403) if matches!(self.backend, TurnCostBackend::OpenAiApiKey(_)) => { tracing::debug!( "turn cost worker waiting for auth change after backend availability check: {error}" ); BackendAvailability::AwaitingAuthChange } - Some(429) => BackendAvailability::RetryProbe, + Some(401 | 403 | 429) => BackendAvailability::RetryProbe, Some(400..=499) => { tracing::debug!( "turn cost worker disabled by backend availability check: {error}" @@ -316,12 +328,6 @@ impl WorkerRuntime { } async fn poll_due(&mut self) { - let Some(auth) = self.auth_manager.auth().await else { - return; - }; - if !auth.is_api_key_auth() { - return; - } let now = Instant::now(); let due_turn_ids: Vec = self .turns @@ -333,46 +339,26 @@ impl WorkerRuntime { .map(|(turn_id, _)| turn_id.clone()) .collect(); if !due_turn_ids.is_empty() { - self.poll_api_key_entries(&due_turn_ids, &auth).await; + self.poll_api_key_entries(&due_turn_ids).await; } } - async fn poll_api_key_entries(&mut self, turn_ids: &[String], auth: &CodexAuth) { - let provider = match self - .config - .model_provider - .to_api_provider(Some(AuthMode::ApiKey)) - { - Ok(provider) => provider, - Err(error) => { - warn!("failed to resolve OpenAI API-key provider headers: {error}"); - self.retry_entries(turn_ids); - return; - } - }; - let client = BackendClient::from_auth( - self.config.chatgpt_base_url.clone(), - auth, - self.config.http_client_factory(), - ); - let costs = match tokio::time::timeout( - REQUEST_TIMEOUT, - client.query_api_key_turn_costs(turn_ids, &provider.headers), - ) - .await - { - Ok(Ok(costs)) => costs, - Ok(Err(error)) => { - warn!("failed to query OpenAI API-key turn costs: {error}"); - self.retry_entries(turn_ids); - return; - } - Err(_) => { - warn!("timed out querying OpenAI API-key turn costs"); - self.retry_entries(turn_ids); - return; - } - }; + async fn poll_api_key_entries(&mut self, turn_ids: &[String]) { + let costs = + match tokio::time::timeout(REQUEST_TIMEOUT, self.query_turn_costs(turn_ids)).await { + Ok(Ok(Some(costs))) => costs, + Ok(Ok(None)) => return, + Ok(Err(error)) => { + warn!("failed to query API-key turn costs: {error}"); + self.retry_entries(turn_ids); + return; + } + Err(_) => { + warn!("timed out querying API-key turn costs"); + self.retry_entries(turn_ids); + return; + } + }; let costs_by_turn: HashMap = costs .into_iter() .map(|cost| (cost.turn_id.clone(), cost)) @@ -386,6 +372,64 @@ impl WorkerRuntime { } } + async fn query_turn_costs( + &self, + turn_ids: &[String], + ) -> Result>, RequestError> { + match &self.backend { + TurnCostBackend::OpenAiApiKey(auth_manager) => { + let Some(auth) = auth_manager.auth().await else { + return Ok(None); + }; + if !auth.is_api_key_auth() { + return Ok(None); + } + let provider = self + .config + .model_provider + .to_api_provider(Some(AuthMode::ApiKey)) + .map_err(|error| RequestError::Other(error.into()))?; + let client = BackendClient::from_auth( + self.config.chatgpt_base_url.clone(), + &auth, + self.config.http_client_factory(), + ); + client + .query_api_key_turn_costs(turn_ids, &provider.headers) + .await + .map(Some) + } + TurnCostBackend::ModelProvider(model_provider) => { + if model_provider.info().requires_openai_auth { + let Some(auth) = model_provider.auth().await else { + return Ok(None); + }; + if !auth.is_api_key_auth() { + return Ok(None); + } + } + let provider = model_provider + .api_provider() + .await + .map_err(|error| RequestError::Other(error.into()))?; + let auth = model_provider + .api_auth() + .await + .map_err(|error| RequestError::Other(error.into()))?; + let endpoint = provider.url_for_path("analytics/codex/turn-costs"); + let client = BackendClient::new( + provider.base_url.clone(), + self.config.http_client_factory(), + ) + .with_auth_provider(auth); + client + .query_api_key_turn_costs_at(&endpoint, turn_ids, &provider.headers) + .await + .map(Some) + } + } + } + fn process_api_key_cost(&mut self, turn_id: &str, cost: &ApiKeyTurnCost) { if cost.status != ApiKeyTurnCostStatus::Priced { self.retry_entry(turn_id); diff --git a/codex-rs/app-server/src/turn_cost_worker_tests.rs b/codex-rs/app-server/src/turn_cost_worker_tests.rs index d9f7009e5a..2fa57fbd60 100644 --- a/codex-rs/app-server/src/turn_cost_worker_tests.rs +++ b/codex-rs/app-server/src/turn_cost_worker_tests.rs @@ -3,20 +3,82 @@ use codex_backend_client::ApiKeyResponseCost; use codex_core::config::ConfigBuilder; use codex_login::AuthCredentialsStoreMode; use codex_login::AuthKeyringBackendKind; +use codex_login::CodexAuth; use codex_login::login_with_api_key; +use codex_model_provider_info::ModelProviderInfo; use codex_otel::TelemetryAuthMode; use codex_protocol::protocol::SessionSource; +use codex_protocol::protocol::TurnStartedEvent; use pretty_assertions::assert_eq; use tempfile::TempDir; use tokio::time::timeout; use wiremock::Mock; use wiremock::MockServer; use wiremock::ResponseTemplate; +use wiremock::matchers::header; use wiremock::matchers::method; use wiremock::matchers::path; const TURN_COST_PATH: &str = "/v1/analytics/codex/turn-costs"; +#[tokio::test] +async fn handle_observes_only_matching_model_provider() { + let codex_home = TempDir::new().expect("temporary Codex home"); + let mut config = ConfigBuilder::default() + .codex_home(codex_home.path().to_path_buf()) + .build() + .await + .expect("test config"); + let model_provider = ModelProviderInfo { + name: "provider-a".to_string(), + base_url: Some("https://provider-a.example/v1".to_string()), + ..Default::default() + }; + config.model_provider = model_provider.clone(); + let config = Arc::new(config); + let (sender, mut receiver) = mpsc::channel(OBSERVATION_CHANNEL_CAPACITY); + let handle = TurnCostWorkerHandle { + sender, + backend: TurnCostBackend::ModelProvider(create_model_provider( + model_provider.clone(), + /*auth_manager*/ None, + )), + config: Arc::clone(&config), + }; + let thread_id = ThreadId::new(); + let event = Event { + id: "turn-1".to_string(), + msg: EventMsg::TurnStarted(TurnStartedEvent { + turn_id: "turn-1".to_string(), + trace_id: None, + started_at: None, + model_context_window: None, + collaboration_mode_kind: Default::default(), + }), + }; + let mut mismatched_config = config.as_ref().clone(); + mismatched_config.model_provider = ModelProviderInfo { + base_url: Some("https://provider-b.example/v1".to_string()), + ..model_provider + }; + + handle.observe_event(thread_id, &mismatched_config, &event, || { + panic!("telemetry should not be captured for a mismatched provider") + }); + assert!(receiver.try_recv().is_err()); + + handle.observe_event(thread_id, config.as_ref(), &event, || { + test_session_telemetry(thread_id) + }); + let observation = receiver.recv().await.expect("matching observation"); + assert_eq!(observation.thread_id, thread_id); + assert_eq!(observation.turn_id, "turn-1"); + assert!(matches!( + observation.kind, + TurnCostObservationKind::Started { .. } + )); +} + #[tokio::test] async fn worker_waits_for_late_api_key_login() { let server = MockServer::start().await; @@ -29,18 +91,7 @@ async fn worker_waits_for_late_api_key_login() { .mount(&server) .await; let auth_home = TempDir::new().expect("temporary auth home"); - let auth_manager = Arc::new( - AuthManager::new( - auth_home.path().to_path_buf(), - /*enable_codex_api_key_env*/ false, - AuthCredentialsStoreMode::File, - /*forced_chatgpt_workspace_id*/ None, - /*chatgpt_base_url*/ None, - AuthKeyringBackendKind::default(), - codex_login::test_support::transport_default_auth_route_config(), - ) - .await, - ); + let auth_manager = auth_manager_at(auth_home.path()).await; let runtime = test_runtime(&server, Arc::clone(&auth_manager)).await; let (_sender, receiver) = mpsc::channel(OBSERVATION_CHANNEL_CAPACITY); let shutdown = CancellationToken::new(); @@ -67,6 +118,110 @@ async fn worker_waits_for_late_api_key_login() { server.verify().await; } +#[tokio::test] +async fn custom_provider_auth_failure_retries_without_auth_changes() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/analytics/codex/turn-costs")) + .and(header("authorization", "Bearer sk-old")) + .respond_with(ResponseTemplate::new(401)) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/analytics/codex/turn-costs")) + .and(header("authorization", "Bearer sk-new-token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "turns": [] + }))) + .expect(1) + .mount(&server) + .await; + + let provider_auth_home = TempDir::new().expect("temporary provider auth home"); + login_with_api_key( + provider_auth_home.path(), + "sk-old", + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("write initial provider auth"); + let provider_auth_manager = auth_manager_at(provider_auth_home.path()).await; + let codex_home = TempDir::new().expect("temporary Codex home"); + let mut config = ConfigBuilder::default() + .codex_home(codex_home.path().to_path_buf()) + .build() + .await + .expect("test config"); + config.model_provider = ModelProviderInfo { + name: "custom-provider".to_string(), + base_url: Some(server.uri()), + requires_openai_auth: true, + ..Default::default() + }; + let backend = TurnCostBackend::ModelProvider(create_model_provider( + config.model_provider.clone(), + Some(Arc::clone(&provider_auth_manager)), + )); + let runtime = WorkerRuntime { + config: Arc::new(config), + backend, + turns: HashMap::new(), + }; + assert_eq!( + runtime.probe_backend().await, + BackendAvailability::RetryProbe + ); + login_with_api_key( + provider_auth_home.path(), + "sk-new-token", + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("update provider auth"); + provider_auth_manager.reload().await; + assert_eq!(runtime.probe_backend().await, BackendAvailability::Ready); + + server.verify().await; +} + +#[tokio::test] +async fn custom_provider_does_not_send_chatgpt_auth_for_turn_costs() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/analytics/codex/turn-costs")) + .respond_with(ResponseTemplate::new(500)) + .mount(&server) + .await; + let auth_manager = + AuthManager::from_auth_for_testing(CodexAuth::create_dummy_chatgpt_auth_for_testing()); + let codex_home = TempDir::new().expect("temporary Codex home"); + let mut config = ConfigBuilder::default() + .codex_home(codex_home.path().to_path_buf()) + .build() + .await + .expect("test config"); + config.model_provider = ModelProviderInfo { + name: "custom-provider".to_string(), + base_url: Some(server.uri()), + requires_openai_auth: true, + ..Default::default() + }; + let backend = TurnCostBackend::ModelProvider(create_model_provider( + config.model_provider.clone(), + Some(Arc::clone(&auth_manager)), + )); + let runtime = WorkerRuntime { + config: Arc::new(config), + backend, + turns: HashMap::new(), + }; + + assert_eq!(runtime.probe_backend().await, BackendAvailability::Disabled); + let requests = server.received_requests().await.expect("received requests"); + assert!(requests.is_empty()); +} + #[tokio::test] async fn transient_probe_failure_keeps_worker_alive() { let server = MockServer::start().await; @@ -82,7 +237,7 @@ async fn transient_probe_failure_keeps_worker_alive() { assert_eq!(backend_availability, BackendAvailability::RetryProbe); let (_sender, receiver) = mpsc::channel(OBSERVATION_CHANNEL_CAPACITY); let shutdown = CancellationToken::new(); - let auth_changes = auth_manager.auth_change_receiver(); + let auth_changes = Some(auth_manager.auth_change_receiver()); let mut task = tokio::spawn(runtime.run_with_backend_availability( receiver, shutdown.clone(), @@ -204,13 +359,29 @@ async fn test_runtime(server: &MockServer, auth_manager: Arc) -> Wo .await .expect("test config"); config.chatgpt_base_url = server.uri(); + let backend = TurnCostBackend::OpenAiApiKey(Arc::clone(&auth_manager)); WorkerRuntime { config: Arc::new(config), - auth_manager, + backend, turns: HashMap::new(), } } +async fn auth_manager_at(codex_home: &std::path::Path) -> Arc { + Arc::new( + AuthManager::new( + codex_home.to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await, + ) +} + fn test_session_telemetry(thread_id: ThreadId) -> SessionTelemetry { SessionTelemetry::new( thread_id, diff --git a/codex-rs/backend-client/src/client/turn_usage.rs b/codex-rs/backend-client/src/client/turn_usage.rs index 99387f43ad..b3b5b27c12 100644 --- a/codex-rs/backend-client/src/client/turn_usage.rs +++ b/codex-rs/backend-client/src/client/turn_usage.rs @@ -62,21 +62,37 @@ impl Client { url.set_path("/v1/analytics/codex/turn-costs"); url.set_query(None); url.set_fragment(None); - let url = url.to_string(); - let mut headers = self.headers(); - for header_name in ["openai-organization", "openai-project"] { - if let Some(value) = provider_headers.get(header_name) { - headers.insert(header_name, value.clone()); - } - } + let provider_scope_headers = provider_headers + .iter() + .filter(|(name, _)| { + ["openai-organization", "openai-project"] + .iter() + .any(|allowed| name.as_str().eq_ignore_ascii_case(allowed)) + }) + .map(|(name, value)| (name.clone(), value.clone())) + .collect(); + self.query_api_key_turn_costs_at(url.as_ref(), turn_ids, &provider_scope_headers) + .await + } + + /// Queries an API-key turn-cost endpoint chosen by the caller, attaching + /// the supplied provider headers in addition to this client's auth. + pub async fn query_api_key_turn_costs_at( + &self, + url: &str, + turn_ids: &[String], + provider_headers: &HeaderMap, + ) -> Result, RequestError> { + let mut headers = provider_headers.clone(); + headers.extend(self.headers()); let request = self - .request(Method::POST, &url) + .request(Method::POST, url) .headers(headers) .header(CONTENT_TYPE, HeaderValue::from_static("application/json")) .json(&ApiKeyTurnCostsRequest { turn_ids }); - let (body, content_type) = self.exec_request_detailed(request, "POST", &url).await?; + let (body, content_type) = self.exec_request_detailed(request, "POST", url).await?; let response: ApiKeyTurnCostsResponse = self - .decode_json(&url, &content_type, &body) + .decode_json(url, &content_type, &body) .map_err(RequestError::Other)?; Ok(response.turns) } @@ -185,4 +201,46 @@ mod tests { ] ); } + + #[tokio::test] + async fn custom_turn_cost_queries_apply_client_auth_after_provider_headers() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/analytics/codex/turn-costs")) + .and(header("authorization", "Bearer sk-test")) + .and(header("chatgpt-account-id", "account-test")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "turns": [] + }))) + .expect(1) + .mount(&server) + .await; + + let auth = CodexAuth::from_api_key("sk-test"); + let client = Client::from_auth( + server.uri(), + &auth, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .with_chatgpt_account_id("account-test"); + let mut provider_headers = HeaderMap::new(); + provider_headers.insert( + "authorization", + HeaderValue::from_static("Bearer provider-override"), + ); + provider_headers.insert( + "chatgpt-account-id", + HeaderValue::from_static("provider-account-override"), + ); + let costs = client + .query_api_key_turn_costs_at( + &format!("{}/analytics/codex/turn-costs", server.uri()), + &["turn-one".to_string()], + &provider_headers, + ) + .await + .expect("query custom-provider turn costs"); + + assert_eq!(costs, Vec::new()); + } }