mirror of
https://github.com/openai/codex.git
synced 2026-08-23 13:09:46 +00:00
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:
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user