Support turn cost telemetry for custom model providers (#39785)

## What changed

- Route turn-cost queries for non-OpenAI providers through the configured
  provider endpoint and authentication, while retaining the existing OpenAI
  API-key path and excluding Amazon Bedrock.
- Observe turns only when their model provider matches the worker's provider.
- Retry custom-provider authentication failures during periodic availability
  probes and ensure client authentication takes precedence over provider
  headers.

## Testing

- Add coverage for provider matching, custom-provider authentication retries,
  ChatGPT-auth rejection, and header precedence.

GitOrigin-RevId: 04a7b28e8e3e18a510ae6fface5193af40d114c0
This commit is contained in:
xli-oai
2026-08-20 19:14:45 +00:00
committed by copyberry
parent 90c67e6f33
commit 0cc80b8db5
4 changed files with 390 additions and 114 deletions

View File

@@ -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

View File

@@ -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<TurnCostObservation>,
auth_manager: Arc<AuthManager>,
backend: TurnCostBackend,
config: Arc<Config>,
}
enum TurnCostObservationKind {
@@ -72,10 +75,16 @@ struct TurnCostEntry {
struct WorkerRuntime {
config: Arc<Config>,
auth_manager: Arc<AuthManager>,
backend: TurnCostBackend,
turns: HashMap<String, TurnCostEntry>,
}
#[derive(Clone)]
enum TurnCostBackend {
OpenAiApiKey(Arc<AuthManager>),
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<TurnCostObservation>, 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<TurnCostObservation>,
shutdown: CancellationToken,
mut auth_changes: tokio::sync::watch::Receiver<u64>,
mut auth_changes: Option<tokio::sync::watch::Receiver<u64>>,
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<String> = 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<String, ApiKeyTurnCost> = 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<Option<Vec<ApiKeyTurnCost>>, 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);

View File

@@ -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<AuthManager>) -> 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<AuthManager> {
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,

View File

@@ -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<Vec<ApiKeyTurnCost>, 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());
}
}