#![allow(clippy::unused_async)]
use axum::body::Body;
use axum::http::{HeaderMap, HeaderValue, StatusCode};
use axum::response::Response;
use futures_util::StreamExt;
use std::collections::BTreeMap;
use crate::metrics::Surface;
use crate::proxy::{
AppState, error_response, maybe_mpp_challenge, relay_response_headers, request_routing_context,
retry_after_duration,
};
use crate::subscription::{SubscriptionProvider, SubscriptionToken};
pub async fn forward_subscription_openai(
state: &AppState,
headers: &HeaderMap,
body: serde_json::Value,
routing_body: &serde_json::Value,
path: &str,
surface: Surface,
) -> Response {
forward_subscription_openai_inner(
state,
headers,
body,
routing_body,
path,
surface,
SubscriptionResponseShape::Passthrough,
)
.await
}
pub async fn forward_codex_chat_completions(
state: &AppState,
headers: &HeaderMap,
body: serde_json::Value,
routing_body: &serde_json::Value,
surface: Surface,
) -> Response {
forward_subscription_openai_inner(
state,
headers,
body,
routing_body,
"/v1/responses",
surface,
SubscriptionResponseShape::ChatCompletion,
)
.await
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum SubscriptionResponseShape {
Passthrough,
ChatCompletion,
}
async fn forward_subscription_openai_inner(
state: &AppState,
headers: &HeaderMap,
mut body: serde_json::Value,
routing_body: &serde_json::Value,
path: &str,
surface: Surface,
response_shape: SubscriptionResponseShape,
) -> Response {
if let Some(resp) = maybe_mpp_challenge(state, headers, path) {
return resp;
}
let claims = match crate::proxy::authenticate_client(state, headers) {
Ok(claims) => claims,
Err(response) => return *response,
};
if let Err(e) = state.token_manager.enforce_request_budget(&claims.sub) {
return crate::token_http::budget_error_response(&e);
}
crate::audit::record_authorised_request(state, &claims, surface, path, Some(routing_body));
let Some(provider) = state.upstream_provider.subscription_provider() else {
return error_response(
StatusCode::INTERNAL_SERVER_ERROR,
"api_error",
"active upstream is not a subscription provider",
);
};
let pinned_account = match state.token_manager.account_for(&claims.sub) {
Ok(account) => account,
Err(error) => {
return error_response(
StatusCode::INTERNAL_SERVER_ERROR,
"api_error",
&format!("failed to resolve token account binding: {error}"),
);
}
};
let routing_context = request_routing_context(headers, routing_body, pinned_account);
let selected = if let Some(router) = state.account_router.as_ref() {
match router.select_subscription(&routing_context) {
Ok(selected) => selected,
Err(error) => {
return error_response(
StatusCode::SERVICE_UNAVAILABLE,
"account_unavailable",
&error.to_string(),
);
}
}
} else {
let Some(reader) = state.subscription_reader.as_ref() else {
return error_response(
StatusCode::INTERNAL_SERVER_ERROR,
"api_error",
"subscription credentials reader is not configured",
);
};
let disk_token = match reader.read_token() {
Ok(token) => token,
Err(e) => {
return error_response(
StatusCode::BAD_GATEWAY,
"authentication_error",
&format!("failed to read {provider} subscription credentials: {e}"),
);
}
};
crate::accounts::SelectedSubscriptionAccount {
name: "primary".to_string(),
token: disk_token,
}
};
let now_ms = chrono::Utc::now().timestamp_millis();
let sub_token = state
.subscription_cache
.get_fresh_for(
&state.client,
provider,
&selected.name,
selected.token,
now_ms,
)
.await;
let selected_account = Some(selected.name);
let emulated_output_limit = (crate::capabilities::subscription(provider, None)
.output_token_limit
== crate::capabilities::Capability::Emulated)
.then(|| {
body.get("max_output_tokens")
.and_then(serde_json::Value::as_u64)
})
.flatten();
let stream_requested = body
.get("stream")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false);
normalize_subscription_request(provider, &mut body);
let serialized = match serde_json::to_vec(&body) {
Ok(v) => v,
Err(e) => {
return error_response(
StatusCode::INTERNAL_SERVER_ERROR,
"api_error",
&format!("failed to serialize subscription request body: {e}"),
);
}
};
let bytes_sent = serialized.len() as u64;
let base_url = state
.subscription_base_url
.clone()
.unwrap_or_else(|| sub_token.base_url(provider));
let upstream_url = join_subscription_url(provider, &base_url, path);
let mut upstream_req = state
.client
.post(upstream_url)
.header("content-type", "application/json")
.header(
"authorization",
format!("Bearer {}", sub_token.access_token),
)
.body(serialized);
for (name, value) in subscription_headers(provider, &sub_token) {
upstream_req = upstream_req.header(name, value);
}
let correlation_id = crate::request_log::correlation_id(headers);
let upstream_resp = match state
.request_log
.send_upstream(&correlation_id, &state.client, upstream_req)
.await
{
Ok(resp) => resp,
Err(e) => {
state
.metrics
.record_request(surface, 502, selected_account.as_deref());
return error_response(
StatusCode::BAD_GATEWAY,
"api_error",
&format!("{provider} subscription upstream request failed: {e}"),
);
}
};
let status = StatusCode::from_u16(upstream_resp.status().as_u16())
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
state
.metrics
.record_request(surface, status.as_u16(), selected_account.as_deref());
state
.subscription_cache
.record_status(provider, status.as_u16());
let retry_after = retry_after_duration(upstream_resp.headers());
if status == StatusCode::TOO_MANY_REQUESTS
&& let (Some(router), Some(account)) =
(state.account_router.as_ref(), selected_account.as_deref())
{
router.report_failure_with_retry_after(
account,
"subscription upstream returned 429",
retry_after,
);
}
let content_type = upstream_resp
.headers()
.get("content-type")
.cloned()
.unwrap_or_else(|| HeaderValue::from_static("application/json"));
let response_headers = relay_response_headers(upstream_resp.headers());
let codex = provider == SubscriptionProvider::Codex;
if stream_requested || (!codex && is_event_stream(&content_type)) {
let stream_content_type = if codex {
HeaderValue::from_static("text/event-stream")
} else {
content_type
};
let requested_model = routing_body
.get("model")
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
let include_usage = routing_body
.pointer("/stream_options/include_usage")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false);
let stop_sequences = crate::stop_sequences::from_value(routing_body.get("stop"));
let mut translator = crate::responses::ResponsesChatStreamTranslator::new(requested_model)
.with_include_usage(include_usage)
.with_stop_sequences(stop_sequences)
.with_output_token_limit(emulated_output_limit);
let mut rewriter = crate::output_limit::ResponsesStreamRewriter::new(
requested_model,
emulated_output_limit,
);
let rewrite_passthrough =
codex && response_shape == SubscriptionResponseShape::Passthrough && rewriter.active();
let response_log = std::sync::Arc::clone(&state.request_log);
let mut usage = status.is_success().then(|| {
crate::usage::UsageTracker::new(state.token_manager.clone(), claims.sub.clone())
});
let stream = upstream_resp.bytes_stream().map(move |chunk| {
chunk.map_or_else(
|error| Err(std::io::Error::other(error)),
|bytes| {
response_log.record_upstream_body(&correlation_id, &bytes);
if let Some(tracker) = &mut usage {
tracker.feed(&bytes);
}
if codex && response_shape == SubscriptionResponseShape::ChatCompletion {
Ok(bytes::Bytes::from(translator.push(&bytes).join("")))
} else if rewrite_passthrough {
Ok(bytes::Bytes::from(rewriter.push(&bytes)))
} else {
Ok(bytes)
}
},
)
});
let mut response = Response::new(Body::from_stream(stream));
*response.status_mut() = status;
*response.headers_mut() = response_headers;
response
.headers_mut()
.insert("content-type", stream_content_type);
return response;
}
let upstream_body = match upstream_resp.bytes().await {
Ok(bytes) => bytes,
Err(e) => {
state
.metrics
.record_request(surface, 502, selected_account.as_deref());
return error_response(
StatusCode::BAD_GATEWAY,
"api_error",
&format!("{provider} subscription upstream body read failed: {e}"),
);
}
};
state
.request_log
.record_upstream_body(&correlation_id, &upstream_body);
state
.metrics
.record_bytes(bytes_sent, upstream_body.len() as u64);
if status.is_success() {
let mut usage =
crate::usage::UsageTracker::new(state.token_manager.clone(), claims.sub.clone());
usage.feed(&upstream_body);
}
let mut response_body = upstream_body;
let mut upstream_model: Option<String> = None;
if codex && status.is_success() {
if let Some(json) = codex_sse_to_response_json(&response_body) {
response_body = bytes::Bytes::from(json);
}
if response_shape == SubscriptionResponseShape::ChatCompletion {
let requested_model = routing_body
.get("model")
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
let parsed = match serde_json::from_slice::<serde_json::Value>(&response_body) {
Ok(value) => value,
Err(error) => {
return error_response(
StatusCode::BAD_GATEWAY,
"api_error",
&format!(
"Codex subscription upstream returned an invalid response: {error}"
),
);
}
};
let mut translated =
crate::responses::response_to_chat_completion(&parsed, requested_model);
crate::responses::enforce_chat_stop(
&mut translated,
&crate::stop_sequences::from_value(routing_body.get("stop")),
);
if let Some(limit) = emulated_output_limit {
crate::output_limit::enforce_chat_limit(&mut translated, limit);
}
upstream_model = translated
.get(crate::output_limit::UPSTREAM_MODEL_FIELD)
.and_then(serde_json::Value::as_str)
.map(str::to_string);
response_body = bytes::Bytes::from(
serde_json::to_vec(&translated).expect("JSON values always serialize"),
);
} else if let Ok(mut parsed) = serde_json::from_slice::<serde_json::Value>(&response_body) {
let requested_model = routing_body
.get("model")
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
upstream_model =
crate::output_limit::preserve_model_identity(&mut parsed, requested_model);
if let Some(limit) = emulated_output_limit {
crate::output_limit::enforce_response_limit(&mut parsed, limit);
}
response_body = bytes::Bytes::from(
serde_json::to_vec(&parsed).expect("JSON values always serialize"),
);
}
let mut response = Response::new(Body::from(response_body));
*response.status_mut() = status;
*response.headers_mut() = response_headers;
response
.headers_mut()
.insert("content-type", HeaderValue::from_static("application/json"));
if let Some(served) = upstream_model.as_deref()
&& let Ok(value) = HeaderValue::from_str(served)
{
response
.headers_mut()
.insert(crate::output_limit::UPSTREAM_MODEL_HEADER, value);
}
return response;
}
let mut response = Response::new(Body::from(response_body));
*response.status_mut() = status;
*response.headers_mut() = response_headers;
response.headers_mut().insert("content-type", content_type);
response
}
fn codex_sse_to_response_json(body: &[u8]) -> Option<Vec<u8>> {
let text = std::str::from_utf8(body).ok()?;
let mut completed: Option<serde_json::Value> = None;
let mut output = BTreeMap::<u64, serde_json::Value>::new();
for line in text.lines() {
let Some(payload) = line.strip_prefix("data:") else {
continue;
};
let payload = payload.trim();
if payload.is_empty() || payload == "[DONE]" {
continue;
}
let Ok(event) = serde_json::from_str::<serde_json::Value>(payload) else {
continue;
};
match event.get("type").and_then(serde_json::Value::as_str) {
Some("response.output_item.added" | "response.output_item.done") => {
if let Some(item) = event.get("item") {
let index = event
.get("output_index")
.and_then(serde_json::Value::as_u64)
.unwrap_or(output.len() as u64);
output.insert(index, item.clone());
}
}
Some("response.output_text.delta") => {
update_codex_output_text(&mut output, &event, false);
}
Some("response.output_text.done") => {
update_codex_output_text(&mut output, &event, true);
}
Some("response.completed") => {
if let Some(response) = event.get("response") {
completed = Some(response.clone());
}
}
_ => {}
}
}
if let Some(response) = completed.as_mut() {
let missing_output = response
.get("output")
.and_then(serde_json::Value::as_array)
.is_none_or(Vec::is_empty);
if missing_output && !output.is_empty() {
response["output"] = serde_json::Value::Array(output.into_values().collect());
}
}
completed.and_then(|value| serde_json::to_vec(&value).ok())
}
fn update_codex_output_text(
output: &mut BTreeMap<u64, serde_json::Value>,
event: &serde_json::Value,
done: bool,
) {
let output_index = event
.get("output_index")
.and_then(serde_json::Value::as_u64)
.unwrap_or(0);
let content_index = event
.get("content_index")
.and_then(serde_json::Value::as_u64)
.and_then(|index| usize::try_from(index).ok())
.unwrap_or(0);
let item_id = event
.get("item_id")
.and_then(serde_json::Value::as_str)
.unwrap_or("");
let item = output.entry(output_index).or_insert_with(|| {
serde_json::json!({
"id": item_id,
"type": "message",
"status": "in_progress",
"role": "assistant",
"content": []
})
});
let Some(content) = item
.get_mut("content")
.and_then(serde_json::Value::as_array_mut)
else {
return;
};
content.resize(content_index + 1, serde_json::Value::Null);
if content[content_index].is_null() {
content[content_index] =
serde_json::json!({"type": "output_text", "text": "", "annotations": []});
}
let text = if done {
event.get("text")
} else {
event.get("delta")
}
.and_then(serde_json::Value::as_str)
.unwrap_or("");
if done {
content[content_index]["text"] = serde_json::Value::String(text.to_string());
item["status"] = serde_json::Value::String("completed".to_string());
} else if let Some(current) = content[content_index]["text"].as_str() {
content[content_index]["text"] = serde_json::Value::String(format!("{current}{text}"));
}
}
fn subscription_headers(
provider: SubscriptionProvider,
token: &SubscriptionToken,
) -> Vec<(&'static str, String)> {
let mut out = Vec::new();
if provider == SubscriptionProvider::Codex {
if let Some(account_id) = token.account_id.as_deref() {
out.push(("chatgpt-account-id", account_id.to_string()));
}
out.push(("openai-beta", "responses=experimental".to_string()));
out.push(("originator", "codex_cli_rs".to_string()));
out.push((
"version",
std::env::var("CODEX_CLIENT_VERSION").unwrap_or_else(|_| "0.144.1".to_string()),
));
}
out
}
fn join_subscription_url(provider: SubscriptionProvider, base_url: &str, path: &str) -> String {
let base = base_url.trim_end_matches('/');
match provider {
SubscriptionProvider::Codex => {
let suffix = path.strip_prefix("/v1").unwrap_or(path);
format!("{base}{suffix}")
}
_ => {
if base.ends_with("/v1") {
let suffix = path.strip_prefix("/v1").unwrap_or(path);
format!("{base}{suffix}")
} else {
format!("{base}{path}")
}
}
}
}
pub async fn subscription_models(state: &AppState) -> serde_json::Value {
match state.upstream_provider.subscription_provider() {
Some(provider) => crate::model_routing::pinned_model_catalog(state, provider).await,
None => serde_json::json!({"object": "list", "data": []}),
}
}
fn is_event_stream(content_type: &HeaderValue) -> bool {
content_type
.to_str()
.is_ok_and(|value| value.to_ascii_lowercase().contains("text/event-stream"))
}
fn input_item_text(item: &serde_json::Value) -> Option<String> {
match item.get("content") {
Some(serde_json::Value::String(s)) => Some(s.clone()),
Some(serde_json::Value::Array(parts)) => {
let text: String = parts
.iter()
.filter_map(|p| p.get("text").and_then(serde_json::Value::as_str))
.collect::<Vec<_>>()
.join("");
(!text.is_empty()).then_some(text)
}
_ => None,
}
}
fn normalize_subscription_request(provider: SubscriptionProvider, body: &mut serde_json::Value) {
crate::openai::reconcile_subscription_parameters(provider, body);
if provider == SubscriptionProvider::Codex {
normalize_codex_responses_body(body);
}
}
fn normalize_codex_responses_body(body: &mut serde_json::Value) {
let Some(obj) = body.as_object_mut() else {
return;
};
obj.entry("reasoning").or_insert_with(
|| serde_json::json!({"effort": crate::clients::DEFAULT_OPENAI_REASONING_EFFORT}),
);
obj.insert("stream".to_string(), serde_json::Value::Bool(true));
obj.insert("store".to_string(), serde_json::Value::Bool(false));
obj.remove("max_output_tokens");
if let Some(input) = obj.get("input") {
let normalized = crate::responses::normalize_input_items(input);
obj.insert("input".to_string(), normalized);
}
let mut hoisted: Vec<String> = Vec::new();
if let Some(serde_json::Value::Array(items)) = obj.get_mut("input") {
items.retain(
|item| match item.get("role").and_then(serde_json::Value::as_str) {
Some("system" | "developer") => {
if let Some(text) = input_item_text(item) {
hoisted.push(text);
}
false
}
_ => true,
},
);
}
let mut parts: Vec<String> = Vec::new();
if let Some(existing) = obj.get("instructions").and_then(serde_json::Value::as_str)
&& !existing.trim().is_empty()
{
parts.push(existing.to_string());
}
parts.extend(hoisted);
let instructions = if parts.is_empty() {
"You are a helpful assistant.".to_string()
} else {
parts.join("\n\n")
};
obj.insert(
"instructions".to_string(),
serde_json::Value::String(instructions),
);
}
#[cfg(test)]
#[path = "subscription_proxy_tests.rs"]
mod tests;