refactor(workspace): enforce layered crate boundaries

This commit is contained in:
elky
2026-07-15 23:47:19 +08:00
parent a728c090a9
commit 8616fe6ee2
969 changed files with 40187 additions and 27240 deletions
@@ -0,0 +1,406 @@
use std::collections::BTreeMap;
use serde_json::Value;
use super::super::snapshot::GatewayProviderTransportSnapshot;
pub const ANTIGRAVITY_PROVIDER_TYPE: &str = "antigravity";
pub const ANTIGRAVITY_REQUEST_USER_AGENT: &str =
"antigravity/cli/1.0.16 (aidev_client; os_type=linux; arch=arm64; auth_method=consumer)";
const ANTIGRAVITY_CLIENT_NAME: &str = "antigravity";
const ANTIGRAVITY_GOOG_API_CLIENT: &str = "gl-node/18.18.2 fire/0.8.6 grpc/1.10.x";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AntigravityRequestAuth {
pub project_id: String,
pub client_version: Option<String>,
pub session_id: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AntigravityRequestAuthSupport {
Supported(AntigravityRequestAuth),
Unsupported(AntigravityRequestAuthUnsupportedReason),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AntigravityRequestAuthUnsupportedReason {
WrongProviderType,
MissingAuthConfig,
InvalidAuthConfigJson,
ComplexDynamicAuthConfig,
MissingProjectId,
}
pub fn resolve_local_antigravity_request_auth(
transport: &GatewayProviderTransportSnapshot,
) -> AntigravityRequestAuthSupport {
if !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(ANTIGRAVITY_PROVIDER_TYPE)
{
return AntigravityRequestAuthSupport::Unsupported(
AntigravityRequestAuthUnsupportedReason::WrongProviderType,
);
}
let Some(raw_auth_config) = transport
.key
.decrypted_auth_config
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return AntigravityRequestAuthSupport::Unsupported(
AntigravityRequestAuthUnsupportedReason::MissingAuthConfig,
);
};
let Ok(auth_config) = serde_json::from_str::<Value>(raw_auth_config) else {
return AntigravityRequestAuthSupport::Unsupported(
AntigravityRequestAuthUnsupportedReason::InvalidAuthConfigJson,
);
};
if contains_blocked_auth_fields(&auth_config) {
return AntigravityRequestAuthSupport::Unsupported(
AntigravityRequestAuthUnsupportedReason::ComplexDynamicAuthConfig,
);
}
let upstream_metadata = transport.key.upstream_metadata.as_ref();
let Some(project_id) = find_antigravity_string(
upstream_metadata,
&auth_config,
ANTIGRAVITY_PROJECT_ID_PATHS,
) else {
return AntigravityRequestAuthSupport::Unsupported(
AntigravityRequestAuthUnsupportedReason::MissingProjectId,
);
};
let client_version = find_antigravity_string(
upstream_metadata,
&auth_config,
ANTIGRAVITY_CLIENT_VERSION_PATHS,
);
let session_id = find_antigravity_string(
upstream_metadata,
&auth_config,
ANTIGRAVITY_SESSION_ID_PATHS,
);
AntigravityRequestAuthSupport::Supported(AntigravityRequestAuth {
project_id,
client_version,
session_id,
})
}
pub fn build_antigravity_static_identity_headers(
auth: &AntigravityRequestAuth,
) -> BTreeMap<String, String> {
build_antigravity_static_client_headers(
auth.client_version.as_deref(),
auth.session_id.as_deref(),
)
}
pub fn build_antigravity_static_client_headers(
client_version: Option<&str>,
session_id: Option<&str>,
) -> BTreeMap<String, String> {
let mut headers = BTreeMap::from([
(
String::from("x-client-name"),
String::from(ANTIGRAVITY_CLIENT_NAME),
),
(
String::from("x-goog-api-client"),
String::from(ANTIGRAVITY_GOOG_API_CLIENT),
),
(
String::from("user-agent"),
String::from(ANTIGRAVITY_REQUEST_USER_AGENT),
),
]);
if let Some(client_version) = client_version
.map(str::trim)
.filter(|value| !value.is_empty())
{
headers.insert(String::from("x-client-version"), client_version.to_string());
}
if let Some(session_id) = session_id.map(str::trim).filter(|value| !value.is_empty()) {
headers.insert(String::from("x-vscode-sessionid"), session_id.to_string());
}
headers
}
const ANTIGRAVITY_PROJECT_ID_PATHS: &[&[&str]] = &[
&["project_id"],
&["projectId"],
&["project", "id"],
&["project", "project_id"],
&["project", "projectId"],
&["cloudaicompanionProject"],
&["cloudaicompanionProject", "id"],
&["cloudAiCompanionProject"],
&["cloudAiCompanionProject", "id"],
&["antigravity", "project_id"],
&["antigravity", "projectId"],
&["antigravity", "project", "id"],
&["antigravity", "cloudaicompanionProject"],
&["antigravity", "cloudaicompanionProject", "id"],
&["antigravity", "cloudAiCompanionProject"],
&["antigravity", "cloudAiCompanionProject", "id"],
&["metadata", "project_id"],
&["metadata", "projectId"],
&["metadata", "cloudaicompanionProject"],
&["metadata", "cloudaicompanionProject", "id"],
&["metadata", "cloudAiCompanionProject"],
&["metadata", "cloudAiCompanionProject", "id"],
];
const ANTIGRAVITY_CLIENT_VERSION_PATHS: &[&[&str]] = &[
&["client_version"],
&["clientVersion"],
&["antigravity", "client_version"],
&["antigravity", "clientVersion"],
&["metadata", "client_version"],
&["metadata", "clientVersion"],
];
const ANTIGRAVITY_SESSION_ID_PATHS: &[&[&str]] = &[
&["session_id"],
&["sessionId"],
&["antigravity", "session_id"],
&["antigravity", "sessionId"],
&["metadata", "session_id"],
&["metadata", "sessionId"],
];
fn find_antigravity_string(
upstream_metadata: Option<&Value>,
auth_config: &Value,
paths: &[&[&str]],
) -> Option<String> {
upstream_metadata
.and_then(|metadata| find_string_by_paths(metadata, paths))
.or_else(|| find_string_by_paths(auth_config, paths))
}
fn find_string_by_paths(value: &Value, paths: &[&[&str]]) -> Option<String> {
for path in paths {
let mut current = value;
let mut matched = true;
for segment in *path {
let Some(next) = current.get(*segment) else {
matched = false;
break;
};
current = next;
}
if !matched {
continue;
}
if let Some(string) = current
.as_str()
.map(str::trim)
.filter(|item| !item.is_empty())
{
return Some(string.to_string());
}
if let Some(string) = current
.as_object()
.and_then(|object| {
object
.get("id")
.or_else(|| object.get("project_id"))
.or_else(|| object.get("projectId"))
})
.and_then(Value::as_str)
.map(str::trim)
.filter(|item| !item.is_empty())
{
return Some(string.to_string());
}
}
None
}
fn contains_blocked_auth_fields(value: &Value) -> bool {
match value {
Value::Object(map) => map.iter().any(|(key, inner)| {
is_blocked_auth_key(key.as_str()) || contains_blocked_auth_fields(inner)
}),
Value::Array(items) => items.iter().any(contains_blocked_auth_fields),
_ => false,
}
}
fn is_blocked_auth_key(key: &str) -> bool {
matches!(
key.trim().to_ascii_lowercase().as_str(),
"private_key"
| "privateKey"
| "private_key_id"
| "privateKeyId"
| "service_account"
| "serviceAccount"
| "service_account_json"
| "serviceAccountJson"
| "service_account_key"
| "serviceAccountKey"
| "credential_source"
| "credentialSource"
| "token_url"
| "tokenUrl"
| "auth_uri"
| "authUri"
| "subject"
| "audience"
)
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
build_antigravity_static_client_headers, resolve_local_antigravity_request_auth,
AntigravityRequestAuth, AntigravityRequestAuthSupport, ANTIGRAVITY_REQUEST_USER_AGENT,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport(auth_config: &str) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Antigravity".to_string(),
provider_type: "antigravity".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "gemini:generate_content".to_string(),
api_family: Some("gemini".to_string()),
endpoint_kind: Some("generate_content".to_string()),
is_active: true,
base_url: "https://daily-cloudcode-pa.googleapis.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "oauth".to_string(),
is_active: true,
api_formats: Some(vec!["gemini:generate_content".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: Some(auth_config.to_string()),
},
}
}
#[test]
fn resolves_cloudaicompanion_project_object_from_auth_config() {
let transport = sample_transport(
r#"{
"provider_type":"antigravity",
"refresh_token":"rt",
"cloudaicompanionProject":{"id":"project-from-auth-config"}
}"#,
);
assert_eq!(
resolve_local_antigravity_request_auth(&transport),
AntigravityRequestAuthSupport::Supported(AntigravityRequestAuth {
project_id: "project-from-auth-config".to_string(),
client_version: None,
session_id: None,
})
);
}
#[test]
fn resolves_identity_from_antigravity_upstream_metadata() {
let mut transport = sample_transport(
r#"{
"provider_type":"antigravity",
"refresh_token":"rt"
}"#,
);
transport.key.upstream_metadata = Some(json!({
"antigravity": {
"project_id": "project-from-metadata",
"client_version": "1.99.0",
"session_id": "session-from-metadata"
}
}));
assert_eq!(
resolve_local_antigravity_request_auth(&transport),
AntigravityRequestAuthSupport::Supported(AntigravityRequestAuth {
project_id: "project-from-metadata".to_string(),
client_version: Some("1.99.0".to_string()),
session_id: Some("session-from-metadata".to_string()),
})
);
}
#[test]
fn static_client_headers_use_native_antigravity_cli_user_agent() {
let headers = build_antigravity_static_client_headers(Some("1.0.16"), Some("session-abc"));
assert_eq!(
headers.get("user-agent").map(String::as_str),
Some(ANTIGRAVITY_REQUEST_USER_AGENT)
);
assert_eq!(
headers.get("x-client-name").map(String::as_str),
Some("antigravity")
);
assert_eq!(
headers.get("x-client-version").map(String::as_str),
Some("1.0.16")
);
assert_eq!(
headers.get("x-vscode-sessionid").map(String::as_str),
Some("session-abc")
);
}
}
@@ -0,0 +1,25 @@
mod auth;
mod policy;
mod request;
mod url;
pub use auth::{
build_antigravity_static_client_headers, build_antigravity_static_identity_headers,
resolve_local_antigravity_request_auth, AntigravityRequestAuth, AntigravityRequestAuthSupport,
AntigravityRequestAuthUnsupportedReason, ANTIGRAVITY_PROVIDER_TYPE,
ANTIGRAVITY_REQUEST_USER_AGENT,
};
pub use policy::{
classify_local_antigravity_request_support, is_antigravity_provider_transport,
AntigravityRequestSideSpec, AntigravityRequestSideSupport,
AntigravityRequestSideUnsupportedReason,
};
pub use request::{
build_antigravity_safe_v1internal_request, classify_antigravity_safe_request_body,
AntigravityEnvelopeRequestType, AntigravityRequestEnvelopeSupport,
AntigravityRequestEnvelopeUnsupportedReason,
};
pub use url::{
build_antigravity_v1internal_url, AntigravityRequestUrlAction,
ANTIGRAVITY_V1INTERNAL_PATH_TEMPLATE,
};
@@ -0,0 +1,116 @@
use serde_json::Value;
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::auth::{
resolve_local_antigravity_request_auth, AntigravityRequestAuth, AntigravityRequestAuthSupport,
AntigravityRequestAuthUnsupportedReason, ANTIGRAVITY_PROVIDER_TYPE,
};
use super::request::{
classify_antigravity_safe_request_body, AntigravityEnvelopeRequestType,
AntigravityRequestEnvelopeUnsupportedReason,
};
use crate::rules::{body_rules_have_enabled_rules, header_rules_have_enabled_rules};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AntigravityRequestSideSpec {
pub auth: AntigravityRequestAuth,
pub request_type: AntigravityEnvelopeRequestType,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AntigravityRequestSideSupport {
Supported(AntigravityRequestSideSpec),
Unsupported(AntigravityRequestSideUnsupportedReason),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AntigravityRequestSideUnsupportedReason {
InactiveTransport,
WrongProviderType,
UnsupportedApiFormat,
UnsupportedCustomPath,
UnsupportedHeaderRules,
UnsupportedBodyRules,
UnsupportedNetworkConfig,
UnsupportedAuth(AntigravityRequestAuthUnsupportedReason),
UnsupportedEnvelope(AntigravityRequestEnvelopeUnsupportedReason),
}
pub fn is_antigravity_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(ANTIGRAVITY_PROVIDER_TYPE)
}
pub fn classify_local_antigravity_request_support(
transport: &GatewayProviderTransportSnapshot,
request_body: &Value,
request_type: AntigravityEnvelopeRequestType,
) -> AntigravityRequestSideSupport {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::InactiveTransport,
);
}
if !is_antigravity_provider_transport(transport) {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::WrongProviderType,
);
}
let endpoint_format =
aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
if endpoint_format != "gemini:generate_content" {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedApiFormat,
);
}
if transport
.endpoint
.custom_path
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
{
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedCustomPath,
);
}
if header_rules_have_enabled_rules(transport.endpoint.header_rules.as_ref()) {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedHeaderRules,
);
}
if body_rules_have_enabled_rules(transport.endpoint.body_rules.as_ref()) {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedBodyRules,
);
}
if transport.provider.proxy.is_some()
|| transport.endpoint.proxy.is_some()
|| transport.key.proxy.is_some()
|| transport.key.fingerprint.is_some()
{
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig,
);
}
let auth = match resolve_local_antigravity_request_auth(transport) {
AntigravityRequestAuthSupport::Supported(auth) => auth,
AntigravityRequestAuthSupport::Unsupported(reason) => {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedAuth(reason),
);
}
};
if let Err(reason) = classify_antigravity_safe_request_body(request_body) {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedEnvelope(reason),
);
}
AntigravityRequestSideSupport::Supported(AntigravityRequestSideSpec { auth, request_type })
}
@@ -0,0 +1,367 @@
use serde_json::{Map, Value};
use super::auth::{AntigravityRequestAuth, ANTIGRAVITY_REQUEST_USER_AGENT};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AntigravityEnvelopeRequestType {
Agent,
Checkpoint,
EndpointTest,
}
impl AntigravityEnvelopeRequestType {
fn as_str(self) -> &'static str {
match self {
Self::Agent => "agent",
Self::Checkpoint => "checkpoint",
Self::EndpointTest => "endpoint_test",
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum AntigravityRequestEnvelopeSupport {
Supported(Value),
Unsupported(AntigravityRequestEnvelopeUnsupportedReason),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AntigravityRequestEnvelopeUnsupportedReason {
NonObjectBody,
MissingContents,
MissingRequestId,
MissingModel,
}
pub fn classify_antigravity_safe_request_body(
request_body: &Value,
) -> Result<(), AntigravityRequestEnvelopeUnsupportedReason> {
let Value::Object(map) = request_body else {
return Err(AntigravityRequestEnvelopeUnsupportedReason::NonObjectBody);
};
if !map.contains_key("contents") && existing_v1internal_request_object(map).is_none() {
return Err(AntigravityRequestEnvelopeUnsupportedReason::MissingContents);
}
Ok(())
}
pub fn build_antigravity_safe_v1internal_request(
auth: &AntigravityRequestAuth,
request_id: &str,
model: &str,
request_body: &Value,
request_type: AntigravityEnvelopeRequestType,
) -> AntigravityRequestEnvelopeSupport {
if request_id.trim().is_empty() {
return AntigravityRequestEnvelopeSupport::Unsupported(
AntigravityRequestEnvelopeUnsupportedReason::MissingRequestId,
);
}
if model.trim().is_empty() {
return AntigravityRequestEnvelopeSupport::Unsupported(
AntigravityRequestEnvelopeUnsupportedReason::MissingModel,
);
}
if let Err(reason) = classify_antigravity_safe_request_body(request_body) {
return AntigravityRequestEnvelopeSupport::Unsupported(reason);
}
let Value::Object(source) = request_body else {
return AntigravityRequestEnvelopeSupport::Unsupported(
AntigravityRequestEnvelopeUnsupportedReason::NonObjectBody,
);
};
if let Some(existing_request) = existing_v1internal_request_object(source) {
let mut inner_request: Map<String, Value> = existing_request.clone();
inner_request.remove("model");
inner_request.remove("safetySettings");
inner_request.remove("safety_settings");
let request_id = non_empty_string_field(source, "requestId").unwrap_or(request_id);
let user_agent =
non_empty_string_field(source, "userAgent").unwrap_or(ANTIGRAVITY_REQUEST_USER_AGENT);
let request_type =
existing_v1internal_request_type(source).unwrap_or_else(|| request_type.as_str());
return AntigravityRequestEnvelopeSupport::Supported(serde_json::json!({
"project": auth.project_id,
"requestId": request_id,
"request": Value::Object(inner_request),
"model": model,
"userAgent": user_agent,
"requestType": request_type,
}));
}
let mut inner_request: Map<String, Value> = source.clone();
inner_request.remove("model");
inner_request.remove("safetySettings");
inner_request.remove("safety_settings");
AntigravityRequestEnvelopeSupport::Supported(serde_json::json!({
"project": auth.project_id,
"requestId": request_id,
"request": Value::Object(inner_request),
"model": model,
"userAgent": ANTIGRAVITY_REQUEST_USER_AGENT,
"requestType": request_type.as_str(),
}))
}
fn existing_v1internal_request_object(source: &Map<String, Value>) -> Option<&Map<String, Value>> {
source
.get("request")
.and_then(Value::as_object)
.filter(|request| request.contains_key("contents"))
}
fn non_empty_string_field<'a>(source: &'a Map<String, Value>, key: &str) -> Option<&'a str> {
source
.get(key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn existing_v1internal_request_type(source: &Map<String, Value>) -> Option<&str> {
match non_empty_string_field(source, "requestType")? {
"agent" => Some("agent"),
"checkpoint" => Some("checkpoint"),
"endpoint_test" => Some("endpoint_test"),
_ => None,
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
build_antigravity_safe_v1internal_request, classify_antigravity_safe_request_body,
AntigravityEnvelopeRequestType, AntigravityRequestAuth, AntigravityRequestEnvelopeSupport,
};
use crate::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT;
fn sample_auth() -> AntigravityRequestAuth {
AntigravityRequestAuth {
project_id: "project-ant-123".to_string(),
client_version: None,
session_id: None,
}
}
#[test]
fn real_agent_request_preserves_antigravity_agent_fields() {
let request_body = json!({
"model": "client-side-model-should-not-be-nested",
"contents": [
{
"role": "user",
"parts": [
{ "text": "Reply with OK only." }
]
}
],
"systemInstruction": {
"role": "user",
"parts": [
{ "text": "Antigravity agent system prompt" }
]
},
"generationConfig": {
"maxOutputTokens": 8192,
"thinkingConfig": {
"includeThoughts": true,
"thinkingBudget": 4000
}
},
"toolConfig": {
"functionCallingConfig": {
"mode": "VALIDATED"
}
},
"tools": [
{
"functionDeclarations": [
{
"name": "run_command",
"description": "Run a command",
"parameters": {
"type": "object",
"properties": {
"cmd": { "type": "string" }
},
"required": ["cmd"]
}
}
]
}
],
"labels": {
"trajectory_id": "trajectory-123",
"used_claude": "false"
},
"sessionId": "session-ant-123",
"safetySettings": [
{ "category": "HARM_CATEGORY_UNSPECIFIED" }
]
});
assert_eq!(
classify_antigravity_safe_request_body(&request_body),
Ok(())
);
let envelope = match build_antigravity_safe_v1internal_request(
&sample_auth(),
"request-ant-agent-123",
"gemini-3.5-flash-low",
&request_body,
AntigravityEnvelopeRequestType::Agent,
) {
AntigravityRequestEnvelopeSupport::Supported(envelope) => envelope,
AntigravityRequestEnvelopeSupport::Unsupported(reason) => {
panic!("real agent envelope should be supported: {reason:?}")
}
};
assert_eq!(envelope["project"], "project-ant-123");
assert_eq!(envelope["requestId"], "request-ant-agent-123");
assert_eq!(envelope["model"], "gemini-3.5-flash-low");
assert_eq!(envelope["userAgent"], ANTIGRAVITY_REQUEST_USER_AGENT);
assert_eq!(envelope["requestType"], "agent");
assert!(envelope["request"].get("model").is_none());
assert!(envelope["request"].get("safetySettings").is_none());
assert_eq!(
envelope["request"]["systemInstruction"]["parts"][0]["text"],
"Antigravity agent system prompt"
);
assert_eq!(
envelope["request"]["generationConfig"]["thinkingConfig"]["thinkingBudget"],
4000
);
assert_eq!(
envelope["request"]["toolConfig"]["functionCallingConfig"]["mode"],
"VALIDATED"
);
assert_eq!(
envelope["request"]["tools"][0]["functionDeclarations"][0]["name"],
"run_command"
);
assert_eq!(
envelope["request"]["labels"]["trajectory_id"],
"trajectory-123"
);
assert_eq!(envelope["request"]["sessionId"], "session-ant-123");
}
#[test]
fn checkpoint_request_type_builds_checkpoint_envelope() {
let request_body = json!({
"contents": [
{
"role": "user",
"parts": [
{ "text": "checkpoint context" }
]
}
],
"generationConfig": {
"maxOutputTokens": 8192,
"thinkingConfig": {
"includeThoughts": true,
"thinkingBudget": 4000
}
},
"toolConfig": {
"functionCallingConfig": {
"mode": "NONE"
}
}
});
let envelope = match build_antigravity_safe_v1internal_request(
&sample_auth(),
"request-ant-checkpoint-123",
"gemini-3.5-flash-low",
&request_body,
AntigravityEnvelopeRequestType::Checkpoint,
) {
AntigravityRequestEnvelopeSupport::Supported(envelope) => envelope,
AntigravityRequestEnvelopeSupport::Unsupported(reason) => {
panic!("checkpoint envelope should be supported: {reason:?}")
}
};
assert_eq!(envelope["requestType"], "checkpoint");
assert_eq!(
envelope["request"]["toolConfig"]["functionCallingConfig"]["mode"],
"NONE"
);
}
#[test]
fn existing_v1internal_envelope_is_not_double_wrapped() {
let request_body = json!({
"project": "client-side-project",
"requestId": "client-request-id-123",
"model": "gemini-3.5-flash-low",
"userAgent": "antigravity",
"requestType": "checkpoint",
"request": {
"contents": [
{
"role": "user",
"parts": [
{ "text": "checkpoint context" }
]
}
],
"generationConfig": {
"thinkingConfig": {
"includeThoughts": true
}
},
"toolConfig": {
"functionCallingConfig": {
"mode": "NONE"
}
}
}
});
assert_eq!(
classify_antigravity_safe_request_body(&request_body),
Ok(())
);
let envelope = match build_antigravity_safe_v1internal_request(
&sample_auth(),
"trace-request-id-should-not-overwrite-client-id",
"mapped-antigravity-model",
&request_body,
AntigravityEnvelopeRequestType::Agent,
) {
AntigravityRequestEnvelopeSupport::Supported(envelope) => envelope,
AntigravityRequestEnvelopeSupport::Unsupported(reason) => {
panic!("existing v1internal envelope should be supported: {reason:?}")
}
};
assert_eq!(envelope["project"], "project-ant-123");
assert_eq!(envelope["requestId"], "client-request-id-123");
assert_eq!(envelope["model"], "mapped-antigravity-model");
assert_eq!(envelope["userAgent"], "antigravity");
assert_eq!(envelope["requestType"], "checkpoint");
assert!(envelope["request"].get("request").is_none());
assert_eq!(
envelope["request"]["contents"][0]["parts"][0]["text"],
"checkpoint context"
);
assert_eq!(
envelope["request"]["toolConfig"]["functionCallingConfig"]["mode"],
"NONE"
);
}
}
@@ -0,0 +1,73 @@
use std::collections::BTreeMap;
use url::form_urlencoded;
pub const ANTIGRAVITY_V1INTERNAL_PATH_TEMPLATE: &str = "/v1internal:{action}";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AntigravityRequestUrlAction {
GenerateContent,
StreamGenerateContent,
}
impl AntigravityRequestUrlAction {
fn as_str(self) -> &'static str {
match self {
Self::GenerateContent => "generateContent",
Self::StreamGenerateContent => "streamGenerateContent",
}
}
fn is_stream(self) -> bool {
matches!(self, Self::StreamGenerateContent)
}
}
pub fn build_antigravity_v1internal_url(
base_url: &str,
action: AntigravityRequestUrlAction,
query: Option<&BTreeMap<String, String>>,
) -> Option<String> {
let trimmed_base = base_url.trim();
if trimmed_base.is_empty() {
return None;
}
let path = ANTIGRAVITY_V1INTERNAL_PATH_TEMPLATE.replace("{action}", action.as_str());
let mut url = format!("{}{}", trimmed_base.trim_end_matches('/'), path);
let mut params = BTreeMap::new();
if let Some(query) = query {
for (key, value) in query {
let key = key.trim();
let value = value.trim();
if key.is_empty()
|| value.is_empty()
|| key.eq_ignore_ascii_case("beta")
|| key.eq_ignore_ascii_case("key")
{
continue;
}
params.insert(key.to_string(), value.to_string());
}
}
if action.is_stream() {
params
.entry(String::from("alt"))
.or_insert_with(|| String::from("sse"));
}
if !params.is_empty() {
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in params {
serializer.append_pair(key.as_str(), value.as_str());
}
let query_string = serializer.finish();
if !query_string.is_empty() {
url.push('?');
url.push_str(&query_string);
}
}
Some(url)
}
@@ -0,0 +1,599 @@
use std::collections::BTreeMap;
use super::headers::{
normalize_upstream_accept_encoding, should_skip_upstream_complete_passthrough_header,
should_skip_upstream_passthrough_header,
};
use super::snapshot::GatewayProviderTransportSnapshot;
const DEFAULT_ANTHROPIC_VERSION: &str = "2023-06-01";
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
fn collect_passthrough_headers(
headers: &http::HeaderMap,
extra_headers: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
let mut out = BTreeMap::new();
for (name, value) in headers.iter() {
let Ok(value) = value.to_str() else {
continue;
};
let key = name.as_str().to_ascii_lowercase();
if should_skip_upstream_passthrough_header(&key) {
continue;
}
let Some(value) = normalize_passthrough_header_value(&key, value) else {
continue;
};
out.insert(key, value);
}
for (key, value) in extra_headers {
let normalized_key = key.to_ascii_lowercase();
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
continue;
};
out.insert(normalized_key, value);
}
out
}
fn collect_complete_passthrough_headers(
headers: &http::HeaderMap,
extra_headers: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
let mut out = BTreeMap::new();
for (name, value) in headers.iter() {
let Ok(value) = value.to_str() else {
continue;
};
let key = name.as_str().to_ascii_lowercase();
if should_skip_upstream_complete_passthrough_header(&key) {
continue;
}
let Some(value) = normalize_passthrough_header_value(&key, value) else {
continue;
};
out.insert(key, value);
}
for (key, value) in extra_headers {
let normalized_key = key.to_ascii_lowercase();
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
continue;
};
out.insert(normalized_key, value);
}
out
}
fn normalize_passthrough_header_value(key: &str, value: &str) -> Option<String> {
let value = value.trim();
if value.is_empty() {
return None;
}
if key.eq_ignore_ascii_case("accept-encoding") {
return normalize_upstream_accept_encoding(value);
}
Some(value.to_string())
}
pub fn build_passthrough_headers(
headers: &http::HeaderMap,
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let mut out = collect_passthrough_headers(headers, extra_headers);
out.entry("content-type".to_string()).or_insert_with(|| {
content_type
.filter(|value| !value.trim().is_empty())
.unwrap_or("application/json")
.trim()
.to_string()
});
out.remove("content-length");
out
}
pub fn build_openai_passthrough_headers(
headers: &http::HeaderMap,
auth_header: &str,
auth_value: &str,
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let mut out = build_passthrough_headers(headers, extra_headers, content_type);
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
out
}
pub fn build_complete_passthrough_headers(
headers: &http::HeaderMap,
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let mut out = collect_complete_passthrough_headers(headers, extra_headers);
out.entry("content-type".to_string()).or_insert_with(|| {
content_type
.filter(|value| !value.trim().is_empty())
.unwrap_or("application/json")
.trim()
.to_string()
});
out.remove("content-length");
out
}
pub fn build_complete_passthrough_headers_with_auth(
headers: &http::HeaderMap,
auth_header: &str,
auth_value: &str,
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let mut out = build_complete_passthrough_headers(headers, extra_headers, content_type);
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
out
}
pub fn build_claude_passthrough_headers(
headers: &http::HeaderMap,
auth_header: &str,
auth_value: &str,
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let mut out = build_openai_passthrough_headers(
headers,
auth_header,
auth_value,
extra_headers,
content_type,
);
for (name, value) in headers.iter() {
let Ok(value) = value.to_str() else {
continue;
};
let key = name.as_str().to_ascii_lowercase();
let value = value.trim();
if value.is_empty() || !should_restore_claude_passthrough_header(&key) {
continue;
}
if key == "anthropic-beta" {
let merged = merge_comma_header_values(out.get(&key).map(String::as_str), Some(value));
if let Some(merged) = merged {
out.insert(key, merged);
}
continue;
}
out.entry(key).or_insert_with(|| value.to_string());
}
out.entry("anthropic-version".to_string())
.or_insert_with(|| DEFAULT_ANTHROPIC_VERSION.to_string());
out
}
pub fn build_passthrough_headers_with_auth(
headers: &http::HeaderMap,
auth_header: &str,
auth_value: &str,
extra_headers: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
let mut out = collect_passthrough_headers(headers, extra_headers);
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
out.remove("content-length");
out
}
pub fn ensure_upstream_auth_header(
headers: &mut BTreeMap<String, String>,
auth_header: &str,
auth_value: &str,
) {
let header_name = auth_header.trim().to_ascii_lowercase();
let header_value = auth_value.trim();
if header_name.is_empty() || header_value.is_empty() {
return;
}
if headers
.get(&header_name)
.map(|value| value.trim().is_empty())
.unwrap_or(true)
{
headers.insert(header_name, header_value.to_string());
}
}
fn should_restore_claude_passthrough_header(name: &str) -> bool {
name.starts_with("anthropic-") || name.starts_with("x-stainless-") || name == "x-app"
}
fn merge_comma_header_values(left: Option<&str>, right: Option<&str>) -> Option<String> {
let mut merged = Vec::new();
for raw in [left, right].into_iter().flatten() {
for token in raw.split(',') {
let token = token.trim();
if token.is_empty() || merged.iter().any(|existing: &String| existing == token) {
continue;
}
merged.push(token.to_string());
}
}
if merged.is_empty() {
None
} else {
Some(merged.join(","))
}
}
pub fn resolve_local_openai_bearer_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
let auth_type = resolve_local_auth_type_for_transport_format(transport);
if !matches!(auth_type.as_str(), "api_key" | "bearer") {
return None;
}
let secret = resolved_local_secret(transport)?;
Some(("authorization".to_string(), bearer_auth_value(secret)))
}
pub fn resolve_local_standard_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
let auth_type = resolve_local_auth_type_for_transport_format(transport);
let secret = resolved_local_secret(transport)?;
match auth_type.as_str() {
"api_key" => Some(("x-api-key".to_string(), secret.to_string())),
"bearer" => Some(("authorization".to_string(), bearer_auth_value(secret))),
_ => None,
}
}
pub fn resolve_local_gemini_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
let auth_type = resolve_local_auth_type_for_transport_format(transport);
let secret = resolved_local_secret(transport)?;
match auth_type.as_str() {
"api_key" => Some(("x-goog-api-key".to_string(), secret.to_string())),
"bearer" => Some(("authorization".to_string(), bearer_auth_value(secret))),
_ => None,
}
}
pub fn resolve_local_auth_type_for_transport_format(
transport: &GatewayProviderTransportSnapshot,
) -> String {
let default_auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
let api_format = aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
let Some(overrides) = transport
.key
.auth_type_by_format
.as_ref()
.and_then(serde_json::Value::as_object)
else {
return default_auth_type;
};
overrides
.get(&api_format)
.or_else(|| overrides.get(transport.endpoint.api_format.trim()))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.map(str::to_ascii_lowercase)
.filter(|value| matches!(value.as_str(), "api_key" | "bearer"))
.unwrap_or(default_auth_type)
}
fn resolved_local_secret(transport: &GatewayProviderTransportSnapshot) -> Option<&str> {
let secret = transport.key.decrypted_api_key.trim();
if !secret.is_empty() && secret != PLACEHOLDER_API_KEY {
Some(secret)
} else if transport.key.decrypted_auth_config.is_some() {
None
} else {
Some("")
}
}
fn bearer_auth_value(secret: &str) -> String {
if secret.is_empty() {
String::new()
} else {
format!("Bearer {secret}")
}
}
#[cfg(test)]
mod tests {
use super::{
build_claude_passthrough_headers, build_complete_passthrough_headers_with_auth,
build_openai_passthrough_headers, resolve_local_openai_bearer_auth,
resolve_local_standard_auth,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use std::collections::BTreeMap;
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "provider".to_string(),
provider_type: "custom".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "claude:messages".to_string(),
api_family: Some("claude".to_string()),
endpoint_kind: Some("chat".to_string()),
is_active: true,
base_url: "https://example.test".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn claude_passthrough_headers_restore_stripped_anthropic_headers() {
let mut headers = http::HeaderMap::new();
headers.insert(
"anthropic-beta",
http::HeaderValue::from_static("prompt-caching-2024-07-31,context-1m-2025-08-07"),
);
headers.insert(
"x-stainless-runtime-version",
http::HeaderValue::from_static("v22.14.0"),
);
headers.insert("x-app", http::HeaderValue::from_static("cli"));
let built = build_claude_passthrough_headers(
&headers,
"x-api-key",
"sk-upstream-claude",
&BTreeMap::from([("anthropic-beta".to_string(), "custom-beta".to_string())]),
Some("application/json"),
);
assert_eq!(
built.get("anthropic-version").map(String::as_str),
Some("2023-06-01")
);
assert_eq!(
built.get("anthropic-beta").map(String::as_str),
Some("custom-beta,prompt-caching-2024-07-31,context-1m-2025-08-07")
);
assert_eq!(
built.get("x-stainless-runtime-version").map(String::as_str),
Some("v22.14.0")
);
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
assert_eq!(
built.get("x-api-key").map(String::as_str),
Some("sk-upstream-claude")
);
}
#[test]
fn passthrough_headers_preserve_supported_response_compression() {
let mut headers = http::HeaderMap::new();
headers.insert(
http::header::ACCEPT_ENCODING,
http::HeaderValue::from_static("gzip, br"),
);
let built = build_openai_passthrough_headers(
&headers,
"authorization",
"Bearer upstream",
&BTreeMap::new(),
Some("application/json"),
);
assert_eq!(
built.get("accept-encoding").map(String::as_str),
Some("gzip")
);
}
#[test]
fn claude_passthrough_headers_preserve_explicit_anthropic_version_override() {
let mut headers = http::HeaderMap::new();
headers.insert(
"anthropic-version",
http::HeaderValue::from_static("2024-01-01"),
);
let built = build_claude_passthrough_headers(
&headers,
"authorization",
"Bearer upstream-token",
&BTreeMap::new(),
Some("application/json"),
);
assert_eq!(
built.get("anthropic-version").map(String::as_str),
Some("2024-01-01")
);
}
#[test]
fn complete_passthrough_headers_preserve_business_headers() {
let mut headers = http::HeaderMap::new();
headers.insert(
"anthropic-beta",
http::HeaderValue::from_static("prompt-caching-2024-07-31"),
);
headers.insert(
"x-stainless-runtime-version",
http::HeaderValue::from_static("v24.0.0"),
);
headers.insert("x-app", http::HeaderValue::from_static("cli"));
headers.insert(
"authorization",
http::HeaderValue::from_static("Bearer client-token"),
);
let built = build_complete_passthrough_headers_with_auth(
&headers,
"x-api-key",
"sk-upstream",
&BTreeMap::new(),
Some("application/json"),
);
assert_eq!(
built.get("anthropic-beta").map(String::as_str),
Some("prompt-caching-2024-07-31")
);
assert_eq!(
built.get("x-stainless-runtime-version").map(String::as_str),
Some("v24.0.0")
);
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
assert_eq!(built.get("authorization"), None);
assert_eq!(
built.get("x-api-key").map(String::as_str),
Some("sk-upstream")
);
}
#[test]
fn local_standard_auth_keeps_header_shape_for_placeholder_secret() {
assert_eq!(
resolve_local_standard_auth(&sample_transport()),
Some(("authorization".to_string(), String::new()))
);
}
#[test]
fn local_standard_auth_keeps_header_shape_for_empty_secret() {
let mut transport = sample_transport();
transport.key.auth_type = "api_key".to_string();
transport.key.decrypted_api_key = String::new();
assert_eq!(
resolve_local_standard_auth(&transport),
Some(("x-api-key".to_string(), String::new()))
);
}
#[test]
fn local_standard_auth_defers_to_auth_config_when_raw_secret_is_empty() {
let mut transport = sample_transport();
transport.key.decrypted_auth_config =
Some(r#"{"access_token":"cached-token"}"#.to_string());
assert!(resolve_local_standard_auth(&transport).is_none());
}
#[test]
fn local_openai_bearer_auth_maps_api_key_to_bearer_authorization() {
let mut transport = sample_transport();
transport.key.auth_type = "api_key".to_string();
transport.key.decrypted_api_key = "sk-openai".to_string();
assert_eq!(
resolve_local_openai_bearer_auth(&transport),
Some(("authorization".to_string(), "Bearer sk-openai".to_string(),))
);
}
#[test]
fn local_openai_bearer_auth_preserves_bearer_header_shape() {
let mut transport = sample_transport();
transport.key.auth_type = "bearer".to_string();
transport.key.decrypted_api_key = "sk-openai".to_string();
assert_eq!(
resolve_local_openai_bearer_auth(&transport),
Some(("authorization".to_string(), "Bearer sk-openai".to_string(),))
);
}
#[test]
fn local_standard_auth_uses_format_auth_type_override() {
let mut transport = sample_transport();
transport.key.auth_type = "api_key".to_string();
transport.key.auth_type_by_format = Some(serde_json::json!({
"claude:messages": "bearer"
}));
transport.key.decrypted_api_key = "sk-claude".to_string();
assert_eq!(
resolve_local_standard_auth(&transport),
Some(("authorization".to_string(), "Bearer sk-claude".to_string(),))
);
}
#[test]
fn local_gemini_auth_falls_back_to_default_when_other_format_is_overridden() {
let mut transport = sample_transport();
transport.endpoint.api_format = "gemini:generate_content".to_string();
transport.key.auth_type = "api_key".to_string();
transport.key.auth_type_by_format = Some(serde_json::json!({
"claude:messages": "bearer"
}));
transport.key.decrypted_api_key = "sk-gemini".to_string();
assert_eq!(
super::resolve_local_gemini_auth(&transport),
Some(("x-goog-api-key".to_string(), "sk-gemini".to_string(),))
);
}
}
@@ -0,0 +1,784 @@
use std::collections::BTreeMap;
use serde_json::Value;
use url::form_urlencoded;
const UNSAFE_AUTH_CONFIG_HEADER_NAMES: &[&str] = &["content-length", "host", "proxy-authorization"];
const RUNTIME_ONLY_AUTH_CONFIG_HEADER_NAMES: &[&str] = &[
"api-key",
"authorization",
"content-type",
"cookie",
"x-api-key",
"x-goog-api-key",
];
const UNSAFE_AUTH_CONFIG_QUERY_NAMES: &[&str] = &[
"access_token",
"api_key",
"apikey",
"authorization",
"key",
"token",
];
const SENSITIVE_AUTH_CONFIG_KEYS: &[&str] = &[
"access_token",
"api_key",
"apikey",
"authorization",
"client_email",
"client_id",
"client_secret",
"expires_at",
"id_token",
"key",
"private_key",
"refresh_token",
"service_account",
"token",
"token_uri",
];
const IGNORABLE_AUTH_CONFIG_METADATA_KEYS: &[&str] = &[
"account_id",
"account_name",
"account_user_id",
"auth_method",
"access_token_import_temporary",
"email",
"expires_at",
"is_fedramp",
"model_regions",
"organizations",
"plan_type",
"project_id",
"provider_type",
"refresh_token_import_error",
"region",
"scope",
"tier",
"token_type",
"updated_at",
"user_id",
"workspace_id",
"workspace_name",
];
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LocalAuthConfigSafeSubset {
pub headers: BTreeMap<String, String>,
pub query: BTreeMap<String, String>,
pub path: Option<String>,
}
impl LocalAuthConfigSafeSubset {
fn is_empty(&self) -> bool {
self.headers.is_empty() && self.query.is_empty() && self.path.is_none()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum LocalAuthConfigAbsorption {
Missing,
Unsupported,
Absorbed {
base_url: String,
header_rules: Option<Value>,
custom_path: Option<String>,
},
}
pub fn apply_local_auth_config_header_overrides(
headers: &mut BTreeMap<String, String>,
raw_auth_config: Option<&str>,
) {
let Some(raw_auth_config) = raw_auth_config
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return;
};
let Ok(parsed) = serde_json::from_str::<Value>(raw_auth_config) else {
return;
};
let Some(object) = parsed.as_object() else {
return;
};
let mut overrides = BTreeMap::new();
collect_auth_config_header_overrides(object, &mut overrides);
for (key, value) in overrides {
headers.insert(key, value);
}
}
pub fn absorb_local_auth_config_safe_subset(
base_url: &str,
header_rules: Option<Value>,
custom_path: Option<String>,
raw_auth_config: Option<&str>,
) -> LocalAuthConfigAbsorption {
let Some(raw_auth_config) = raw_auth_config
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return LocalAuthConfigAbsorption::Missing;
};
let subset = match parse_local_auth_config_safe_subset(raw_auth_config) {
Ok(subset) => subset,
Err(()) => return LocalAuthConfigAbsorption::Unsupported,
};
if subset.is_empty() {
return LocalAuthConfigAbsorption::Unsupported;
}
let header_rules = match merge_auth_config_header_rules(header_rules, &subset.headers) {
Some(rules) => rules,
None => return LocalAuthConfigAbsorption::Unsupported,
};
let base_url = match merge_auth_config_base_url(base_url, &subset.query) {
Some(value) => value,
None => return LocalAuthConfigAbsorption::Unsupported,
};
let custom_path = match merge_auth_config_custom_path(custom_path, subset.path) {
Some(path) => path,
None => return LocalAuthConfigAbsorption::Unsupported,
};
LocalAuthConfigAbsorption::Absorbed {
base_url,
header_rules,
custom_path,
}
}
fn parse_local_auth_config_safe_subset(raw: &str) -> Result<LocalAuthConfigSafeSubset, ()> {
let parsed: Value = serde_json::from_str(raw).map_err(|_| ())?;
let object = parsed.as_object().ok_or(())?;
let mut headers = BTreeMap::new();
let mut query = BTreeMap::new();
let mut path = None;
parse_local_auth_config_object(object, &mut headers, &mut query, &mut path, true)?;
if headers
.keys()
.any(|key| RUNTIME_ONLY_AUTH_CONFIG_HEADER_NAMES.contains(&key.as_str()))
{
return Err(());
}
Ok(LocalAuthConfigSafeSubset {
headers,
query,
path,
})
}
fn collect_auth_config_header_overrides(
object: &serde_json::Map<String, Value>,
out: &mut BTreeMap<String, String>,
) {
for (key, value) in object {
let normalized = key.trim().to_ascii_lowercase();
match normalized.as_str() {
"headers" | "extra_headers" | "extraheaders" => {
merge_header_string_map_lenient(out, value);
}
"transport" | "request" => {
if let Some(nested) = value.as_object() {
collect_auth_config_header_overrides(nested, out);
}
}
_ => {}
}
}
}
fn parse_local_auth_config_object(
object: &serde_json::Map<String, Value>,
headers: &mut BTreeMap<String, String>,
query: &mut BTreeMap<String, String>,
path: &mut Option<String>,
allow_metadata: bool,
) -> Result<(), ()> {
for (key, value) in object {
let normalized = key.trim().to_ascii_lowercase();
match normalized.as_str() {
"headers" | "extra_headers" | "extraheaders" => {
merge_header_string_map(headers, value)?
}
"query" | "query_params" | "queryparams" => {
merge_string_map(query, value, normalize_auth_config_query_key)?
}
"path" | "custom_path" => {
let value = value.as_str().ok_or(())?;
let normalized = normalize_auth_config_path(value).ok_or(())?;
*path = Some(normalized);
}
"custompath" => {
let value = value.as_str().ok_or(())?;
let normalized = normalize_auth_config_path(value).ok_or(())?;
*path = Some(normalized);
}
"transport" | "request" => {
let nested = value.as_object().ok_or(())?;
parse_local_auth_config_object(nested, headers, query, path, false)?;
}
_ if allow_metadata && is_ignorable_auth_config_metadata_key(&normalized) => {}
_ if allow_metadata && is_sensitive_auth_config_key(&normalized) => return Err(()),
_ => return Err(()),
}
}
Ok(())
}
fn merge_string_map(
out: &mut BTreeMap<String, String>,
value: &Value,
normalize_key: fn(&str) -> Option<String>,
) -> Result<(), ()> {
let object = value.as_object().ok_or(())?;
for (raw_key, raw_value) in object {
let key = normalize_key(raw_key).ok_or(())?;
let value = parse_static_auth_config_value(raw_value).ok_or(())?;
out.insert(key, value);
}
Ok(())
}
fn parse_static_auth_config_value(value: &Value) -> Option<String> {
match value {
Value::String(raw) => {
let normalized = raw.trim();
if normalized.is_empty() {
None
} else {
Some(normalized.to_string())
}
}
Value::Number(raw) => Some(raw.to_string()),
Value::Bool(raw) => Some(raw.to_string()),
_ => None,
}
}
fn merge_header_string_map(out: &mut BTreeMap<String, String>, value: &Value) -> Result<(), ()> {
let object = value.as_object().ok_or(())?;
for (raw_key, raw_value) in object {
let key = normalize_auth_config_header_name(raw_key).ok_or(())?;
let value = parse_static_auth_config_header_value(raw_value).ok_or(())?;
out.insert(key, value);
}
Ok(())
}
fn merge_header_string_map_lenient(out: &mut BTreeMap<String, String>, value: &Value) {
let Some(object) = value.as_object() else {
return;
};
for (raw_key, raw_value) in object {
let Some(key) = normalize_auth_config_header_name(raw_key) else {
continue;
};
let Some(value) = parse_static_auth_config_header_value(raw_value) else {
continue;
};
out.insert(key, value);
}
}
fn parse_static_auth_config_header_value(value: &Value) -> Option<String> {
let value = parse_static_auth_config_value(value)?;
http::header::HeaderValue::from_str(&value)
.is_ok()
.then_some(value)
}
fn normalize_auth_config_header_name(raw: &str) -> Option<String> {
let value = raw.trim().to_ascii_lowercase();
if value.is_empty()
|| value.chars().any(|char| char.is_ascii_control())
|| UNSAFE_AUTH_CONFIG_HEADER_NAMES.contains(&value.as_str())
{
return None;
}
http::header::HeaderName::from_bytes(value.as_bytes())
.ok()
.map(|name| name.as_str().to_string())
}
fn normalize_auth_config_query_key(raw: &str) -> Option<String> {
let value = raw.trim();
if value.is_empty()
|| value.chars().any(|char| matches!(char, '&' | '=' | '#'))
|| value.chars().any(|char| char.is_ascii_control())
|| UNSAFE_AUTH_CONFIG_QUERY_NAMES
.iter()
.any(|blocked| value.eq_ignore_ascii_case(blocked))
{
return None;
}
Some(value.to_string())
}
fn is_sensitive_auth_config_key(key: &str) -> bool {
SENSITIVE_AUTH_CONFIG_KEYS
.iter()
.any(|blocked| key.eq_ignore_ascii_case(blocked))
}
fn is_ignorable_auth_config_metadata_key(key: &str) -> bool {
IGNORABLE_AUTH_CONFIG_METADATA_KEYS
.iter()
.any(|allowed| key.eq_ignore_ascii_case(allowed))
}
fn normalize_auth_config_path(raw: &str) -> Option<String> {
let value = raw.trim();
if value.is_empty()
|| !value.starts_with('/')
|| value.contains("://")
|| value
.chars()
.any(|char| matches!(char, '{' | '}' | '$' | '#'))
|| value.chars().any(|char| char.is_ascii_control())
{
return None;
}
Some(value.to_string())
}
fn merge_auth_config_header_rules(
existing_rules: Option<Value>,
headers: &BTreeMap<String, String>,
) -> Option<Option<Value>> {
if headers.is_empty() {
return Some(existing_rules);
}
let mut merged = match existing_rules {
Some(Value::Array(items)) => items,
Some(_) => return None,
None => Vec::new(),
};
for (key, value) in headers {
merged.push(serde_json::json!({
"action": "set",
"key": key,
"value": value,
}));
}
Some(Some(Value::Array(merged)))
}
fn merge_auth_config_custom_path(
existing_custom_path: Option<String>,
path_override: Option<String>,
) -> Option<Option<String>> {
let base_path = path_override.or(existing_custom_path);
let Some(base_path) = base_path else {
return Some(None);
};
let (path_only, query) = split_path_and_query(&base_path)?;
if query.is_empty() {
return Some(Some(path_only));
}
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in query {
serializer.append_pair(&key, &value);
}
Some(Some(format!("{path_only}?{}", serializer.finish())))
}
fn merge_auth_config_base_url(base_url: &str, query: &BTreeMap<String, String>) -> Option<String> {
if query.is_empty() {
return Some(base_url.to_string());
}
let raw_base_url = base_url.trim();
let had_implicit_root = raw_base_url
.split_once("://")
.map(|(_, rest)| {
let authority = rest.split_once('?').map(|(head, _)| head).unwrap_or(rest);
!authority.contains('/')
})
.unwrap_or(false);
let mut url = url::Url::parse(raw_base_url).ok()?;
let mut merged = BTreeMap::new();
for (key, value) in url.query_pairs() {
let value = value.trim();
if value.is_empty() {
return None;
}
merged.insert(key.into_owned(), value.to_string());
}
for (key, value) in query {
merged.insert(key.clone(), value.clone());
}
if merged.is_empty() {
url.set_query(None);
return Some(url.to_string());
}
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in merged {
serializer.append_pair(&key, &value);
}
url.set_query(Some(&serializer.finish()));
let mut normalized = url.to_string();
if had_implicit_root {
normalized = normalized.replacen("/?", "?", 1);
}
Some(normalized)
}
fn split_path_and_query(path: &str) -> Option<(String, BTreeMap<String, String>)> {
let normalized = normalize_auth_config_path(path)?;
let (path_only, query_part) = if let Some((path, query)) = normalized.split_once('?') {
(path.to_string(), Some(query))
} else {
(normalized, None)
};
if path_only.is_empty() {
return None;
}
let mut query = BTreeMap::new();
if let Some(query_part) = query_part.filter(|value| !value.trim().is_empty()) {
for (key, value) in form_urlencoded::parse(query_part.as_bytes()) {
let key = normalize_auth_config_query_key(key.as_ref())?;
let value = value.trim();
if value.is_empty() {
return None;
}
query.insert(key, value.to_string());
}
}
Some((path_only, query))
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
absorb_local_auth_config_safe_subset, apply_local_auth_config_header_overrides,
LocalAuthConfigAbsorption,
};
#[test]
fn absorbs_static_headers_and_query_into_existing_transport_fields() {
let result = absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
Some(json!([{"action":"set","key":"x-base","value":"1"}])),
None,
Some(
r#"{
"headers": {"x-account-id": "acc-1"},
"query": {"tenant": "demo"}
}"#,
),
);
let LocalAuthConfigAbsorption::Absorbed {
base_url,
header_rules,
custom_path,
} = result
else {
panic!("auth_config should be absorbed");
};
assert_eq!(base_url, "https://api.openai.example/v1?tenant=demo");
assert_eq!(
header_rules,
Some(json!([
{"action":"set","key":"x-base","value":"1"},
{"action":"set","key":"x-account-id","value":"acc-1"}
]))
);
assert_eq!(custom_path, None);
}
#[test]
fn absorbs_path_override_and_query_aliases() {
let result = absorb_local_auth_config_safe_subset(
"https://generativelanguage.googleapis.com/v1beta",
None,
Some("/v1beta/models/original:generateContent".to_string()),
Some(
r#"{
"extra_headers": {"x-tenant": "demo"},
"query_params": {"alt": "sse"},
"custom_path": "/v1beta/models/gemini-2.5-pro:streamGenerateContent"
}"#,
),
);
let LocalAuthConfigAbsorption::Absorbed {
base_url,
header_rules,
custom_path,
} = result
else {
panic!("auth_config should be absorbed");
};
assert_eq!(
base_url,
"https://generativelanguage.googleapis.com/v1beta?alt=sse"
);
assert_eq!(
header_rules,
Some(json!([{"action":"set","key":"x-tenant","value":"demo"}]))
);
assert_eq!(
custom_path.as_deref(),
Some("/v1beta/models/gemini-2.5-pro:streamGenerateContent")
);
}
#[test]
fn rejects_unknown_keys_and_reserved_headers() {
assert_eq!(
absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
None,
None,
Some(r#"{"provider_type":"custom"}"#),
),
LocalAuthConfigAbsorption::Unsupported
);
assert_eq!(
absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
None,
None,
Some(r#"{"headers":{"host":"api.example.test"}}"#),
),
LocalAuthConfigAbsorption::Unsupported
);
assert_eq!(
absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
None,
None,
Some(r#"{"query":{"key":"secret"}}"#),
),
LocalAuthConfigAbsorption::Unsupported
);
}
#[test]
fn applies_header_overrides_even_when_auth_config_has_refresh_token() {
let mut headers = std::collections::BTreeMap::from([
(
"authorization".to_string(),
"Bearer direct-token".to_string(),
),
("content-type".to_string(), "application/json".to_string()),
]);
apply_local_auth_config_header_overrides(
&mut headers,
Some(
r#"{
"refresh_token": "rt-1",
"headers": {
"authorization": "Bearer imported-session",
"content-type": "text/plain",
"host": "blocked.example"
}
}"#,
),
);
assert_eq!(
headers.get("authorization"),
Some(&"Bearer imported-session".to_string())
);
assert_eq!(headers.get("content-type"), Some(&"text/plain".to_string()));
assert!(!headers.contains_key("host"));
}
#[test]
fn ignores_invalid_auth_config_header_values_when_applying_overrides() {
let mut headers = std::collections::BTreeMap::new();
apply_local_auth_config_header_overrides(
&mut headers,
Some(r#"{"headers":{"authorization":"Bearer ok","x-bad":"line\nbreak"}}"#),
);
assert_eq!(headers.get("authorization"), Some(&"Bearer ok".to_string()));
assert!(!headers.contains_key("x-bad"));
}
#[test]
fn keeps_imported_authorization_headers_for_runtime_override() {
assert_eq!(
absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
None,
None,
Some(
r#"{
"provider_type": "codex",
"access_token_import_temporary": true,
"headers": {
"authorization": "Bearer imported-session",
"chatgpt-account-id": "acct-1"
}
}"#,
),
),
LocalAuthConfigAbsorption::Unsupported
);
let mut headers = std::collections::BTreeMap::from([(
"authorization".to_string(),
"Bearer direct-token".to_string(),
)]);
apply_local_auth_config_header_overrides(
&mut headers,
Some(
r#"{
"provider_type": "codex",
"access_token_import_temporary": true,
"headers": {
"authorization": "Bearer imported-session",
"chatgpt-account-id": "acct-1"
}
}"#,
),
);
assert_eq!(
headers.get("authorization"),
Some(&"Bearer imported-session".to_string())
);
assert_eq!(
headers.get("chatgpt-account-id"),
Some(&"acct-1".to_string())
);
}
#[test]
fn absorbs_query_only_configs_into_base_url_for_dynamic_path_formats() {
let result = absorb_local_auth_config_safe_subset(
"https://generativelanguage.googleapis.com/v1beta",
None,
None,
Some(r#"{"query":{"alt":"sse"}}"#),
);
let LocalAuthConfigAbsorption::Absorbed {
base_url,
header_rules,
custom_path,
} = result
else {
panic!("query-only auth_config should be absorbed");
};
assert_eq!(
base_url,
"https://generativelanguage.googleapis.com/v1beta?alt=sse"
);
assert_eq!(header_rules, None);
assert_eq!(custom_path, None);
}
#[test]
fn absorbs_camel_case_transport_keys_with_ignorable_metadata() {
let result = absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
None,
None,
Some(
r#"{
"email": "[email protected]",
"plan_type": "plus",
"request": {
"extraHeaders": {"x-org-id": "org-1"},
"queryParams": {"tenant": "demo", "retry": 2, "stream": true},
"customPath": "/v1/responses"
}
}"#,
),
);
let LocalAuthConfigAbsorption::Absorbed {
base_url,
header_rules,
custom_path,
} = result
else {
panic!("camelCase auth_config should be absorbed");
};
assert_eq!(
base_url,
"https://api.openai.example/v1?retry=2&stream=true&tenant=demo"
);
assert_eq!(
header_rules,
Some(json!([{"action":"set","key":"x-org-id","value":"org-1"}]))
);
assert_eq!(custom_path.as_deref(), Some("/v1/responses"));
}
#[test]
fn rejects_sensitive_oauth_fields_even_with_transport_subset() {
assert_eq!(
absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
None,
None,
Some(
r#"{
"headers": {"x-org-id": "org-1"},
"refresh_token": "rt-1"
}"#,
),
),
LocalAuthConfigAbsorption::Unsupported
);
assert_eq!(
absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
None,
None,
Some(
r#"{
"query": {"tenant": "demo"},
"access_token": "at-1"
}"#,
),
),
LocalAuthConfigAbsorption::Unsupported
);
}
#[test]
fn rejects_metadata_only_auth_config_without_transport_subset() {
assert_eq!(
absorb_local_auth_config_safe_subset(
"https://api.openai.example/v1",
None,
None,
Some(
r#"{
"email": "[email protected]",
"plan_type": "plus",
"workspace_name": "demo"
}"#,
),
),
LocalAuthConfigAbsorption::Unsupported
);
}
}
@@ -0,0 +1,133 @@
use super::snapshot::GatewayProviderTransportSnapshot;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ProviderTransportSnapshotCacheKey {
provider_id: String,
endpoint_id: String,
key_id: String,
}
impl ProviderTransportSnapshotCacheKey {
pub fn new(provider_id: &str, endpoint_id: &str, key_id: &str) -> Option<Self> {
let provider_id = provider_id.trim();
let endpoint_id = endpoint_id.trim();
let key_id = key_id.trim();
if provider_id.is_empty() || endpoint_id.is_empty() || key_id.is_empty() {
return None;
}
Some(Self {
provider_id: provider_id.to_string(),
endpoint_id: endpoint_id.to_string(),
key_id: key_id.to_string(),
})
}
}
pub fn provider_transport_snapshot_looks_refreshed(
current: &GatewayProviderTransportSnapshot,
refreshed: &GatewayProviderTransportSnapshot,
) -> bool {
current.key.decrypted_api_key != refreshed.key.decrypted_api_key
|| current.key.decrypted_auth_config != refreshed.key.decrypted_auth_config
|| current.key.expires_at_unix_secs != refreshed.key.expires_at_unix_secs
}
#[cfg(test)]
mod tests {
use super::{provider_transport_snapshot_looks_refreshed, ProviderTransportSnapshotCacheKey};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_snapshot() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Provider".to_string(),
provider_type: "openai".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "openai".to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://example.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "Key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: Some(1),
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "sk-test".to_string(),
decrypted_auth_config: Some("{\"token\":\"x\"}".to_string()),
},
}
}
#[test]
fn cache_key_requires_non_empty_segments() {
assert!(ProviderTransportSnapshotCacheKey::new("provider", "endpoint", "key").is_some());
assert!(ProviderTransportSnapshotCacheKey::new("", "endpoint", "key").is_none());
assert!(ProviderTransportSnapshotCacheKey::new("provider", " ", "key").is_none());
assert!(ProviderTransportSnapshotCacheKey::new("provider", "endpoint", "").is_none());
}
#[test]
fn refresh_detection_tracks_key_material_and_expiry() {
let current = sample_snapshot();
let mut refreshed = current.clone();
assert!(!provider_transport_snapshot_looks_refreshed(
&current, &refreshed
));
refreshed.key.decrypted_api_key = "sk-updated".to_string();
assert!(provider_transport_snapshot_looks_refreshed(
&current, &refreshed
));
let mut refreshed = current.clone();
refreshed.key.decrypted_auth_config = Some("{\"token\":\"y\"}".to_string());
assert!(provider_transport_snapshot_looks_refreshed(
&current, &refreshed
));
let mut refreshed = current.clone();
refreshed.key.expires_at_unix_secs = Some(2);
assert!(provider_transport_snapshot_looks_refreshed(
&current, &refreshed
));
}
}
@@ -0,0 +1,8 @@
use super::super::auth::resolve_local_standard_auth;
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::super::supports_local_oauth_request_auth_resolution;
pub fn supports_local_claude_code_auth(transport: &GatewayProviderTransportSnapshot) -> bool {
resolve_local_standard_auth(transport).is_some_and(|(_, value)| !value.trim().is_empty())
|| supports_local_oauth_request_auth_resolution(transport)
}
@@ -0,0 +1,386 @@
use aether_contracts::{
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
};
use serde_json::{Map, Value};
use sha2::{Digest, Sha256};
use uuid::Uuid;
// Chrome impersonate profiles
const CHROME_IMPERSONATE_PROFILES: &[&str] = &[
"chrome110",
"chrome116",
"chrome119",
"chrome120",
"chrome123",
"chrome124",
"chrome131",
"chrome133",
];
const CHROME_VERSIONS: &[(&str, &str)] = &[
("chrome110", "110.0.5481.177"),
("chrome116", "116.0.5845.188"),
("chrome119", "119.0.6045.214"),
("chrome120", "120.0.6099.216"),
("chrome123", "123.0.6312.122"),
("chrome124", "124.0.6367.243"),
("chrome131", "131.0.6778.265"),
("chrome133", "133.0.6943.142"),
];
// (os, arch, platform_token, platform_info)
const PLATFORM_VARIANTS: &[(&str, &str, &str, &str)] = &[
("Linux", "x64", "X11; Linux x86_64", "Linux x86_64"),
("Linux", "arm64", "X11; Linux arm64", "Linux arm64"),
(
"Windows",
"x64",
"Windows NT 10.0; Win64; x64",
"Windows x64",
),
(
"MacOS",
"x64",
"Macintosh; Intel Mac OS X 10_15_7",
"Darwin x64",
),
(
"MacOS",
"arm64",
"Macintosh; ARM Mac OS X 14_0_0",
"Darwin arm64",
),
];
const STAINLESS_PACKAGE_VERSIONS: &[&str] = &["0.68.0", "0.69.0", "0.70.0", "0.71.0"];
const NODE_VERSIONS: &[&str] = &["v20.18.1", "v22.12.0", "v22.14.0", "v24.13.0"];
const ELECTRON_VERSIONS: &[&str] = &["35.5.1", "36.7.1", "37.3.0", "38.7.0", "39.2.3"];
const STAINLESS_TIMEOUTS: &[&str] = &["600", "900"];
const CLAUDE_CODE_TRANSPORT_PROFILE_ID: &str = "claude_code_nodejs";
/// Deterministic hash-based index picker, compatible with Python implementation.
/// Each `slot` produces a different selection from the same seed.
struct SeededPicker {
seed_bytes: [u8; 32],
}
impl SeededPicker {
fn new(seed: &str) -> Self {
let mut hasher = Sha256::new();
hasher.update(seed.as_bytes());
Self {
seed_bytes: hasher.finalize().into(),
}
}
/// Pick an index from `[0, len)` using a specific slot.
/// Different slots produce independent-looking selections from the same seed.
fn pick(&self, slot: u8, len: usize) -> usize {
if len == 0 {
return 0;
}
// Hash seed_bytes + slot to get a new digest, take first 8 bytes as u64
let mut hasher = Sha256::new();
hasher.update(self.seed_bytes);
hasher.update([slot]);
let hash = hasher.finalize();
let value = u64::from_be_bytes(hash[..8].try_into().unwrap());
(value % len as u64) as usize
}
}
fn chrome_version_for_profile(profile: &str) -> &'static str {
for (p, v) in CHROME_VERSIONS {
if p.eq_ignore_ascii_case(profile) {
return v;
}
}
"120.0.6099.216"
}
fn build_user_agent(platform_token: &str, chrome_version: &str, electron_version: &str) -> String {
format!(
"Mozilla/5.0 ({platform_token}) AppleWebKit/537.36 (KHTML, like Gecko) \
Chrome/{chrome_version} Electron/{electron_version} Safari/537.36"
)
}
fn resolve_platform_token(os: &str, arch: &str) -> &'static str {
let os_lower = os.to_ascii_lowercase();
let arch_lower = arch.to_ascii_lowercase();
if os_lower.starts_with("win") {
return "Windows NT 10.0; Win64; x64";
}
if matches!(os_lower.as_str(), "darwin" | "mac" | "macos") {
return if matches!(arch_lower.as_str(), "arm64" | "aarch64") {
"Macintosh; ARM Mac OS X 14_0_0"
} else {
"Macintosh; Intel Mac OS X 10_15_7"
};
}
if matches!(arch_lower.as_str(), "arm64" | "aarch64") {
"X11; Linux arm64"
} else {
"X11; Linux x86_64"
}
}
/// Generate a complete Claude Code transport fingerprint from a seed.
pub fn generate_fingerprint(seed: &str) -> Value {
wrap_header_fingerprint(generate_header_fingerprint(seed))
}
fn generate_header_fingerprint(seed: &str) -> Value {
let picker = SeededPicker::new(seed);
let impersonate =
CHROME_IMPERSONATE_PROFILES[picker.pick(0, CHROME_IMPERSONATE_PROFILES.len())];
let chrome_version = chrome_version_for_profile(impersonate);
let node_version = NODE_VERSIONS[picker.pick(1, NODE_VERSIONS.len())];
let electron_version = ELECTRON_VERSIONS[picker.pick(2, ELECTRON_VERSIONS.len())];
let platform = PLATFORM_VARIANTS[picker.pick(3, PLATFORM_VARIANTS.len())];
let (stainless_os, stainless_arch, platform_token, platform_info) = platform;
let stainless_package_version =
STAINLESS_PACKAGE_VERSIONS[picker.pick(4, STAINLESS_PACKAGE_VERSIONS.len())];
let stainless_timeout = STAINLESS_TIMEOUTS[picker.pick(5, STAINLESS_TIMEOUTS.len())];
let vscode_session_id = Uuid::new_v5(
&Uuid::NAMESPACE_URL,
format!("aether:fingerprint:{seed}").as_bytes(),
)
.simple()
.to_string();
let user_agent = build_user_agent(platform_token, chrome_version, electron_version);
serde_json::json!({
"impersonate": impersonate,
"stainless_package_version": stainless_package_version,
"stainless_os": stainless_os,
"stainless_arch": stainless_arch,
"stainless_runtime_version": node_version,
"stainless_timeout": stainless_timeout,
"node_version": node_version,
"chrome_version": chrome_version,
"electron_version": electron_version,
"vscode_session_id": vscode_session_id,
"platform_info": platform_info,
"user_agent": user_agent,
})
}
fn wrap_header_fingerprint(header_fingerprint: Value) -> Value {
serde_json::json!({
"transport_profile": {
"profile_id": CLAUDE_CODE_TRANSPORT_PROFILE_ID,
"backend": TRANSPORT_BACKEND_REQWEST_RUSTLS,
"http_mode": TRANSPORT_HTTP_MODE_AUTO,
"pool_scope": TRANSPORT_POOL_SCOPE_KEY,
"header_fingerprint": header_fingerprint,
}
})
}
pub fn header_fingerprint_from_fingerprint(fingerprint: &Value) -> Option<&Map<String, Value>> {
fingerprint
.get("transport_profile")
.and_then(Value::as_object)
.and_then(|profile| profile.get("header_fingerprint"))
.and_then(Value::as_object)
}
/// Generate a random (non-deterministic) fingerprint.
pub fn generate_random_fingerprint() -> Value {
let random_seed = Uuid::new_v4().to_string();
generate_fingerprint(&random_seed)
}
/// Sanitize an existing fingerprint JSON, filling missing fields with
/// deterministic fallbacks derived from `key_id`.
pub fn sanitize_fingerprint(raw: &Value, key_id: &str) -> Value {
let generated = generate_header_fingerprint(key_id);
let gen_map = generated.as_object().unwrap();
let raw_map = header_fingerprint_from_fingerprint(raw);
let mut out = Map::new();
// Start with generated values, then overlay non-empty raw values
for (key, gen_value) in gen_map {
let value = raw_map
.and_then(|raw_map| raw_map.get(key))
.and_then(Value::as_str)
.map(str::trim)
.filter(|v| !v.is_empty())
.map(|v| Value::String(v.to_string()))
.unwrap_or_else(|| gen_value.clone());
out.insert(key.clone(), value);
}
// Normalize impersonate to known profile
let impersonate = out
.get("impersonate")
.and_then(Value::as_str)
.unwrap_or_default()
.to_ascii_lowercase();
let is_known = CHROME_IMPERSONATE_PROFILES
.iter()
.any(|p| p.eq_ignore_ascii_case(&impersonate));
if !is_known {
out.insert("impersonate".to_string(), gen_map["impersonate"].clone());
}
// Ensure chrome_version matches impersonate profile
let profile = out
.get("impersonate")
.and_then(Value::as_str)
.unwrap_or_default();
let chrome_version = chrome_version_for_profile(profile);
out.insert(
"chrome_version".to_string(),
Value::String(chrome_version.to_string()),
);
// Rebuild user_agent if missing
let has_ua = out
.get("user_agent")
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|v| !v.is_empty());
if !has_ua {
let os = out
.get("stainless_os")
.and_then(Value::as_str)
.unwrap_or("Linux");
let arch = out
.get("stainless_arch")
.and_then(Value::as_str)
.unwrap_or("x64");
let electron = out
.get("electron_version")
.and_then(Value::as_str)
.unwrap_or("38.7.0");
let platform_token = resolve_platform_token(os, arch);
out.insert(
"user_agent".to_string(),
Value::String(build_user_agent(platform_token, chrome_version, electron)),
);
}
wrap_header_fingerprint(Value::Object(out))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deterministic_generation_from_seed() {
let fp1 = generate_fingerprint("key-abc-123");
let fp2 = generate_fingerprint("key-abc-123");
assert_eq!(fp1, fp2, "same seed should produce identical fingerprint");
}
#[test]
fn different_seeds_produce_different_fingerprints() {
let fp1 = generate_fingerprint("key-1");
let fp2 = generate_fingerprint("key-2");
// At least one field should differ (statistically near-certain)
assert_ne!(fp1, fp2);
}
#[test]
fn generated_fingerprint_has_all_fields() {
let fp = generate_fingerprint("test-key");
let expected_keys = [
"impersonate",
"stainless_package_version",
"stainless_os",
"stainless_arch",
"stainless_runtime_version",
"stainless_timeout",
"node_version",
"chrome_version",
"electron_version",
"vscode_session_id",
"platform_info",
"user_agent",
];
let map = header_fingerprint_from_fingerprint(&fp).unwrap();
for key in expected_keys {
assert!(map.contains_key(key), "missing field: {key}");
let value = map[key].as_str().unwrap();
assert!(!value.is_empty(), "empty field: {key}");
}
}
#[test]
fn sanitize_preserves_user_overrides() {
let raw = serde_json::json!({
"transport_profile": {
"profile_id": "claude_code_nodejs",
"header_fingerprint": {
"stainless_os": "MacOS",
"stainless_arch": "arm64",
"stainless_timeout": "900",
"user_agent": "Custom-Agent/1.0"
}
}
});
let sanitized = sanitize_fingerprint(&raw, "test-key");
let map = header_fingerprint_from_fingerprint(&sanitized).unwrap();
assert_eq!(map["stainless_os"].as_str(), Some("MacOS"));
assert_eq!(map["stainless_arch"].as_str(), Some("arm64"));
assert_eq!(map["stainless_timeout"].as_str(), Some("900"));
assert_eq!(map["user_agent"].as_str(), Some("Custom-Agent/1.0"));
// Other fields should be filled from generation
assert!(map.contains_key("impersonate"));
assert!(map.contains_key("stainless_package_version"));
}
#[test]
fn sanitize_fills_missing_fields_from_seed() {
let raw = serde_json::json!({});
let sanitized = sanitize_fingerprint(&raw, "test-key");
let generated = generate_header_fingerprint("test-key");
// All fields should match generated since raw is empty
let s = header_fingerprint_from_fingerprint(&sanitized).unwrap();
let g = generated.as_object().unwrap();
for key in g.keys() {
assert!(s.contains_key(key), "sanitized missing key: {key}");
assert!(
!s[key].as_str().unwrap().is_empty(),
"sanitized empty key: {key}"
);
}
}
#[test]
fn sanitize_normalizes_unknown_impersonate_profile() {
let raw = serde_json::json!({
"transport_profile": {
"profile_id": "claude_code_nodejs",
"header_fingerprint": {
"impersonate": "firefox99"
}
}
});
let sanitized = sanitize_fingerprint(&raw, "test-key");
let profile = header_fingerprint_from_fingerprint(&sanitized).unwrap()["impersonate"]
.as_str()
.unwrap();
assert!(
CHROME_IMPERSONATE_PROFILES
.iter()
.any(|p| p.eq_ignore_ascii_case(profile)),
"should normalize to known profile, got: {profile}"
);
}
#[test]
fn random_fingerprint_differs_each_call() {
let fp1 = generate_random_fingerprint();
let fp2 = generate_random_fingerprint();
assert_ne!(fp1, fp2);
}
}
@@ -0,0 +1,17 @@
mod auth;
mod fingerprint;
mod policy;
mod request;
mod url;
pub use auth::supports_local_claude_code_auth;
pub use fingerprint::{
generate_fingerprint, generate_random_fingerprint, header_fingerprint_from_fingerprint,
sanitize_fingerprint,
};
pub use policy::{
local_claude_code_transport_unsupported_reason_with_network,
supports_local_claude_code_transport_with_network,
};
pub use request::{build_claude_code_passthrough_headers, sanitize_claude_code_request_body};
pub use url::build_claude_code_messages_url;
@@ -0,0 +1,162 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::super::{
body_rules_are_locally_supported, header_rules_are_locally_supported,
resolve_transport_profile, supports_local_oauth_request_auth_resolution,
transport_profile_is_configured, transport_proxy_is_locally_supported,
};
use super::auth::supports_local_claude_code_auth;
pub fn local_claude_code_transport_unsupported_reason_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<&'static str> {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return if !transport.provider.is_active {
Some("provider_inactive")
} else if !transport.endpoint.is_active {
Some("endpoint_inactive")
} else {
Some("key_inactive")
};
}
if !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("claude_code")
{
return Some("transport_provider_type_unsupported");
}
if !transport
.endpoint
.api_format
.trim()
.eq_ignore_ascii_case(api_format.trim())
{
return Some("transport_api_format_mismatch");
}
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
return Some("transport_header_rules_unsupported");
}
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
return Some("transport_body_rules_unsupported");
}
if transport.key.decrypted_auth_config.is_some()
&& !supports_local_oauth_request_auth_resolution(transport)
{
return Some("transport_oauth_resolution_unsupported");
}
if !supports_local_claude_code_auth(transport) {
return Some("transport_auth_unavailable");
}
if !transport_proxy_is_locally_supported(transport) {
return Some("transport_proxy_unsupported");
}
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
{
return Some("transport_profile_unsupported");
}
None
}
pub fn supports_local_claude_code_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
local_claude_code_transport_unsupported_reason_with_network(transport, api_format).is_none()
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
local_claude_code_transport_unsupported_reason_with_network,
supports_local_claude_code_transport_with_network,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Claude Code".to_string(),
provider_type: "claude_code".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "claude:messages".to_string(),
api_family: Some("claude".to_string()),
endpoint_kind: Some("cli".to_string()),
is_active: true,
base_url: "https://api.anthropic.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: Some(vec!["claude:messages".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: Some(json!({"transport_profile":"chrome_136"})),
upstream_metadata: None,
decrypted_api_key: "sk-ant-123".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn supports_claude_code_transport_when_auth_and_profile_are_valid() {
assert!(supports_local_claude_code_transport_with_network(
&sample_transport(),
"claude:messages"
));
}
#[test]
fn reports_auth_unavailable_for_claude_code_without_local_auth() {
let mut transport = sample_transport();
transport.key.auth_type = "api_key".to_string();
transport.key.decrypted_api_key = "__placeholder__".to_string();
assert_eq!(
local_claude_code_transport_unsupported_reason_with_network(
&transport,
"claude:messages"
),
Some("transport_auth_unavailable")
);
}
}
@@ -0,0 +1,327 @@
use std::collections::{BTreeMap, BTreeSet};
use serde_json::{Map, Value};
use super::super::auth::build_openai_passthrough_headers;
use super::fingerprint::header_fingerprint_from_fingerprint;
const DEFAULT_ANTHROPIC_VERSION: &str = "2023-06-01";
const DEFAULT_ACCEPT: &str = "application/json";
const STREAM_HELPER_METHOD: &str = "stream";
const DUMMY_THINKING_SIGNATURE: &str = "skip_thought_signature_validator";
const REQUIRED_BETA_TOKENS: &[&str] = &[
"claude-code-20250219",
"oauth-2025-04-20",
"interleaved-thinking-2025-05-14",
];
const EXCLUDED_BETA_TOKENS: &[&str] = &["context-1m-2025-08-07"];
/// Fingerprint field -> HTTP header mapping.
/// Every stainless / identity dimension that can vary per-key is listed here.
const FINGERPRINT_HEADER_MAP: &[(&str, &str)] = &[
("stainless_package_version", "x-stainless-package-version"),
("stainless_os", "x-stainless-os"),
("stainless_arch", "x-stainless-arch"),
("stainless_runtime_version", "x-stainless-runtime-version"),
("stainless_timeout", "x-stainless-timeout"),
("user_agent", "user-agent"),
];
pub fn build_claude_code_passthrough_headers(
headers: &http::HeaderMap,
auth_header: &str,
auth_value: &str,
extra_headers: &BTreeMap<String, String>,
stream: bool,
fingerprint: Option<&Value>,
) -> BTreeMap<String, String> {
let mut out = build_openai_passthrough_headers(
headers,
auth_header,
auth_value,
extra_headers,
Some("application/json"),
);
// -- Anthropic protocol headers --
out.insert("accept".to_string(), DEFAULT_ACCEPT.to_string());
out.insert(
"anthropic-version".to_string(),
DEFAULT_ANTHROPIC_VERSION.to_string(),
);
// Read incoming anthropic-beta directly from the original HeaderMap because the
// upstream passthrough filter now strips `anthropic-*` headers to avoid leaking
// them to non-Anthropic upstreams.
let incoming_anthropic_beta = headers
.get("anthropic-beta")
.and_then(|value| value.to_str().ok());
out.insert(
"anthropic-beta".to_string(),
merge_anthropic_beta_tokens(incoming_anthropic_beta),
);
out.insert(
"anthropic-dangerous-direct-browser-access".to_string(),
"true".to_string(),
);
out.insert("x-app".to_string(), "cli".to_string());
// -- Stainless SDK identity headers --
// Fixed values: these don't vary per fingerprint.
out.insert("x-stainless-lang".to_string(), "js".to_string());
out.insert("x-stainless-runtime".to_string(), "node".to_string());
out.insert("x-stainless-retry-count".to_string(), "0".to_string());
// Defaults for fingerprint-overridable fields (used when no fingerprint is present).
out.insert(
"x-stainless-package-version".to_string(),
"0.70.0".to_string(),
);
out.insert("x-stainless-os".to_string(), "Linux".to_string());
out.insert("x-stainless-arch".to_string(), "arm64".to_string());
out.insert(
"x-stainless-runtime-version".to_string(),
"v24.13.0".to_string(),
);
out.insert("x-stainless-timeout".to_string(), "600".to_string());
if stream {
out.insert(
"x-stainless-helper-method".to_string(),
STREAM_HELPER_METHOD.to_string(),
);
} else {
out.remove("x-stainless-helper-method");
}
// Override from the formal transport profile header fingerprint.
if let Some(fp) = fingerprint.and_then(header_fingerprint_from_fingerprint) {
for &(fp_key, header_key) in FINGERPRINT_HEADER_MAP {
if let Some(value) = fp
.get(fp_key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|v| !v.is_empty())
{
out.insert(header_key.to_string(), value.to_string());
}
}
}
out
}
pub fn sanitize_claude_code_request_body(body: &mut Value) {
let Some(body_object) = body.as_object_mut() else {
return;
};
let thinking_enabled = body_object
.get("thinking")
.and_then(Value::as_object)
.and_then(|thinking| thinking.get("type"))
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| matches!(value.to_ascii_lowercase().as_str(), "enabled" | "adaptive"));
let Some(messages) = body_object
.get_mut("messages")
.and_then(Value::as_array_mut)
else {
return;
};
for message in messages {
let Some(message_object) = message.as_object_mut() else {
continue;
};
let role = message_object
.get("role")
.and_then(Value::as_str)
.map(str::trim)
.unwrap_or_default()
.to_string();
let Some(content) = message_object
.get_mut("content")
.and_then(Value::as_array_mut)
else {
continue;
};
let mut filtered = Vec::with_capacity(content.len());
for block in std::mem::take(content) {
let Value::Object(block_object) = block else {
filtered.push(block);
continue;
};
if keep_claude_code_block(&block_object, &role, thinking_enabled) {
filtered.push(Value::Object(block_object));
}
}
*content = filtered;
}
}
fn keep_claude_code_block(
block_object: &Map<String, Value>,
role: &str,
thinking_enabled: bool,
) -> bool {
let block_type = block_object
.get("type")
.and_then(Value::as_str)
.map(str::trim)
.unwrap_or_default();
if matches!(block_type, "thinking" | "redacted_thinking") {
let signature = block_object
.get("signature")
.and_then(Value::as_str)
.map(str::trim)
.unwrap_or_default();
return thinking_enabled
&& role.eq_ignore_ascii_case("assistant")
&& !signature.is_empty()
&& signature != DUMMY_THINKING_SIGNATURE;
}
if block_type.is_empty() && block_object.contains_key("thinking") {
return false;
}
true
}
fn merge_anthropic_beta_tokens(incoming: Option<&str>) -> String {
let mut seen = BTreeSet::new();
let mut merged = Vec::new();
for token in REQUIRED_BETA_TOKENS {
append_beta_token(&mut seen, &mut merged, token);
}
for token in incoming.unwrap_or_default().split(',') {
let token = token.trim();
if EXCLUDED_BETA_TOKENS
.iter()
.any(|excluded| token.eq_ignore_ascii_case(excluded))
{
continue;
}
append_beta_token(&mut seen, &mut merged, token);
}
merged.join(",")
}
fn append_beta_token(seen: &mut BTreeSet<String>, merged: &mut Vec<String>, token: &str) {
let normalized = token.trim();
if normalized.is_empty() {
return;
}
let key = normalized.to_ascii_lowercase();
if seen.insert(key) {
merged.push(normalized.to_string());
}
}
#[cfg(test)]
mod tests {
use super::{build_claude_code_passthrough_headers, sanitize_claude_code_request_body};
use serde_json::json;
use std::collections::BTreeMap;
#[test]
fn claude_code_headers_use_transport_profile_header_fingerprint_and_merge_required_betas() {
let mut headers = http::HeaderMap::new();
headers.insert(
"anthropic-beta",
http::HeaderValue::from_static("context-1m-2025-08-07,custom-beta"),
);
headers.insert(
"user-agent",
http::HeaderValue::from_static("Claude-Code/Test"),
);
let built = build_claude_code_passthrough_headers(
&headers,
"authorization",
"Bearer upstream-token",
&BTreeMap::new(),
true,
Some(&json!({
"transport_profile": {
"profile_id": "claude_code_nodejs",
"header_fingerprint": {
"user_agent":"Claude-Code/9.9",
"stainless_package_version":"1.0.5",
"stainless_runtime_version":"v22.12.0",
"stainless_timeout":"900"
}
}
})),
);
assert_eq!(
built.get("anthropic-beta").map(String::as_str),
Some(
"claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14,custom-beta"
)
);
assert_eq!(
built.get("anthropic-version").map(String::as_str),
Some("2023-06-01")
);
assert_eq!(
built.get("accept").map(String::as_str),
Some("application/json")
);
assert_eq!(
built.get("x-stainless-helper-method").map(String::as_str),
Some("stream")
);
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
assert_eq!(
built.get("x-stainless-package-version").map(String::as_str),
Some("1.0.5")
);
assert_eq!(
built.get("x-stainless-runtime-version").map(String::as_str),
Some("v22.12.0")
);
assert_eq!(
built.get("x-stainless-timeout").map(String::as_str),
Some("900")
);
assert_eq!(
built.get("user-agent").map(String::as_str),
Some("Claude-Code/9.9")
);
assert_eq!(
built.get("authorization").map(String::as_str),
Some("Bearer upstream-token")
);
}
#[test]
fn claude_code_body_sanitizer_drops_invalid_thinking_blocks() {
let mut body = json!({
"thinking": {"type":"enabled"},
"messages": [{
"role":"assistant",
"content":[
{"type":"thinking","thinking":"keep","signature":"sig_valid"},
{"type":"thinking","thinking":"drop-empty","signature":""},
{"type":"redacted_thinking","data":"keep-redacted","signature":"sig_redacted"},
{"type":"redacted_thinking","data":"drop-no-signature"},
{"thinking":"drop-no-type"},
{"type":"text","text":"ok"}
]
}]
});
sanitize_claude_code_request_body(&mut body);
assert_eq!(
body["messages"][0]["content"],
json!([
{"type":"thinking","thinking":"keep","signature":"sig_valid"},
{"type":"redacted_thinking","data":"keep-redacted","signature":"sig_redacted"},
{"type":"text","text":"ok"}
])
);
}
}
@@ -0,0 +1,81 @@
use std::collections::BTreeMap;
use url::form_urlencoded;
pub fn build_claude_code_messages_url(upstream_base_url: &str, query: Option<&str>) -> String {
let (trimmed_base_url, base_query) = split_query(upstream_base_url.trim());
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
let mut url =
if trimmed_base_url.ends_with("/v1/messages") || trimmed_base_url.ends_with("/messages") {
trimmed_base_url.to_string()
} else if trimmed_base_url.ends_with("/v1") {
format!("{trimmed_base_url}/messages")
} else {
format!("{trimmed_base_url}/v1/messages")
};
append_merged_query(&mut url, base_query, query);
url
}
fn split_query(value: &str) -> (&str, Option<&str>) {
value
.split_once('?')
.map(|(base, query)| (base, Some(query)))
.unwrap_or((value, None))
}
fn append_merged_query(url: &mut String, base_query: Option<&str>, request_query: Option<&str>) {
let Some(query) = merge_query_layers(base_query, request_query) else {
return;
};
if url.contains('?') {
url.push('&');
} else {
url.push('?');
}
url.push_str(&query);
}
fn merge_query_layers(base_query: Option<&str>, request_query: Option<&str>) -> Option<String> {
let mut merged = BTreeMap::new();
for source in [base_query, request_query] {
let Some(source) = source.map(str::trim).filter(|value| !value.is_empty()) else {
continue;
};
for (key, value) in form_urlencoded::parse(source.as_bytes()) {
merged.insert(key.into_owned(), value.into_owned());
}
}
if merged.is_empty() {
return None;
}
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in merged {
serializer.append_pair(&key, &value);
}
Some(serializer.finish())
}
#[cfg(test)]
mod tests {
use super::build_claude_code_messages_url;
#[test]
fn keeps_existing_messages_suffix_without_duplication() {
assert_eq!(
build_claude_code_messages_url("https://api.anthropic.com/v1/messages", None),
"https://api.anthropic.com/v1/messages"
);
}
#[test]
fn appends_messages_and_merges_query() {
assert_eq!(
build_claude_code_messages_url(
"https://api.anthropic.com/v1?beta=true",
Some("foo=bar"),
),
"https://api.anthropic.com/v1/messages?beta=true&foo=bar"
);
}
}
@@ -0,0 +1,864 @@
use aether_ai_formats::formats::matrix::{
api_data_format_id, request_conversion_kind, request_conversion_requires_enable_flag,
RequestConversionKind,
};
use aether_ai_formats::normalize_api_format_alias;
use crate::antigravity::is_antigravity_provider_transport;
use crate::auth::{
resolve_local_gemini_auth, resolve_local_openai_bearer_auth, resolve_local_standard_auth,
};
use crate::kiro::{
is_kiro_claude_messages_transport, local_kiro_request_transport_unsupported_reason_with_network,
};
use crate::policy::{
local_gemini_transport_unsupported_reason_with_network,
local_openai_chat_transport_unsupported_reason,
local_standard_transport_unsupported_reason_with_network,
};
use crate::vertex::{
is_vertex_api_key_transport_context, is_vertex_transport_context,
local_vertex_gemini_transport_unsupported_reason_with_network,
resolve_local_vertex_api_key_query_auth, VERTEX_API_KEY_QUERY_PARAM,
};
use crate::windsurf::{
is_windsurf_provider_transport,
local_windsurf_request_transport_unsupported_reason_with_network,
resolve_windsurf_cascade_auth,
};
use crate::GatewayProviderTransportSnapshot;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CandidateTransportPolicyFacts<'a> {
pub endpoint_api_format: &'a str,
pub global_model_name: &'a str,
pub selected_provider_model_name: &'a str,
pub mapping_matched_model: Option<&'a str>,
}
pub fn request_conversion_enabled_for_transport(
transport: &GatewayProviderTransportSnapshot,
client_api_format: &str,
provider_api_format: &str,
) -> bool {
let client_api_format = normalize_api_format_alias(client_api_format);
let provider_api_format = normalize_api_format_alias(provider_api_format);
if client_api_format == provider_api_format {
return true;
}
let conversion_kind =
request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str());
if conversion_kind.is_none()
&& !same_data_format_transport_pair(
client_api_format.as_str(),
provider_api_format.as_str(),
)
{
return false;
}
if !request_conversion_requires_enable_flag(
client_api_format.as_str(),
provider_api_format.as_str(),
) {
return true;
}
transport.provider.enable_format_conversion
|| endpoint_accepts_client_api_format(transport, client_api_format.as_str())
}
pub fn request_pair_allowed_for_transport(
transport: &GatewayProviderTransportSnapshot,
client_api_format: &str,
provider_api_format: &str,
) -> bool {
let client_api_format = normalize_api_format_alias(client_api_format);
let provider_api_format = normalize_api_format_alias(provider_api_format);
if client_api_format == provider_api_format {
return true;
}
let conversion_kind =
request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str());
if conversion_kind.is_none()
&& !same_data_format_transport_pair(
client_api_format.as_str(),
provider_api_format.as_str(),
)
{
return false;
}
if is_kiro_claude_messages_transport(transport, &provider_api_format) {
return request_conversion_enabled_for_transport(
transport,
client_api_format.as_str(),
provider_api_format.as_str(),
) && local_kiro_request_transport_unsupported_reason_with_network(transport)
.is_none();
}
request_conversion_enabled_for_transport(
transport,
client_api_format.as_str(),
provider_api_format.as_str(),
)
}
fn same_data_format_transport_pair(client_api_format: &str, provider_api_format: &str) -> bool {
if aether_ai_formats::api_format_alias_matches(client_api_format, provider_api_format) {
return false;
}
matches!(
(
api_data_format_id(client_api_format),
api_data_format_id(provider_api_format)
),
(Some("embedding"), Some("embedding")) | (Some("rerank"), Some("rerank"))
)
}
pub fn request_conversion_transport_supported(
transport: &GatewayProviderTransportSnapshot,
kind: RequestConversionKind,
) -> bool {
request_conversion_transport_unsupported_reason(transport, kind).is_none()
}
pub fn request_conversion_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
_kind: RequestConversionKind,
) -> Option<&'static str> {
if is_kiro_claude_messages_transport(transport, &transport.endpoint.api_format) {
return local_kiro_request_transport_unsupported_reason_with_network(transport);
}
if is_windsurf_provider_transport(transport)
&& normalize_api_format_alias(&transport.endpoint.api_format) == "openai:chat"
{
return local_windsurf_request_transport_unsupported_reason_with_network(transport);
}
if is_antigravity_provider_transport(transport)
&& normalize_api_format_alias(&transport.endpoint.api_format) == "gemini:generate_content"
{
return None;
}
match normalize_api_format_alias(&transport.endpoint.api_format).as_str() {
"openai:chat" => local_openai_chat_transport_unsupported_reason(transport),
"openai:responses" | "openai:responses:compact" => {
local_standard_transport_unsupported_reason_with_network(
transport,
transport.endpoint.api_format.trim(),
)
}
"claude:messages" => {
local_standard_transport_unsupported_reason_with_network(transport, "claude:messages")
}
"gemini:generate_content" if is_vertex_transport_context(transport) => {
local_vertex_gemini_transport_unsupported_reason_with_network(transport)
}
"gemini:generate_content" => local_gemini_transport_unsupported_reason_with_network(
transport,
"gemini:generate_content",
),
_ => Some("transport_api_format_unsupported"),
}
}
pub fn request_pair_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
client_api_format: &str,
provider_api_format: &str,
) -> Option<&'static str> {
let client_api_format = normalize_api_format_alias(client_api_format);
let provider_api_format = normalize_api_format_alias(provider_api_format);
if let Some(kind) =
request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str())
{
return request_conversion_transport_unsupported_reason(transport, kind);
}
if !same_data_format_transport_pair(client_api_format.as_str(), provider_api_format.as_str()) {
return Some("transport_api_format_unsupported");
}
match provider_api_format.as_str() {
"gemini:embedding" => {
if is_vertex_transport_context(transport) {
local_vertex_gemini_transport_unsupported_reason_with_network(transport)
} else {
local_gemini_transport_unsupported_reason_with_network(
transport,
"gemini:embedding",
)
}
}
"openai:embedding"
| "jina:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
| "openai:rerank"
| "jina:rerank" => local_standard_transport_unsupported_reason_with_network(
transport,
provider_api_format.as_str(),
),
_ => Some("transport_api_format_unsupported"),
}
}
pub fn request_conversion_direct_auth(
transport: &GatewayProviderTransportSnapshot,
_kind: RequestConversionKind,
) -> Option<(String, String)> {
request_direct_auth_for_provider_format(transport, transport.endpoint.api_format.as_str())
}
pub fn request_pair_direct_auth(
transport: &GatewayProviderTransportSnapshot,
provider_api_format: &str,
) -> Option<(String, String)> {
request_direct_auth_for_provider_format(transport, provider_api_format)
}
fn request_direct_auth_for_provider_format(
transport: &GatewayProviderTransportSnapshot,
provider_api_format: &str,
) -> Option<(String, String)> {
match normalize_api_format_alias(provider_api_format).as_str() {
"openai:chat" if is_windsurf_provider_transport(transport) => {
resolve_windsurf_cascade_auth(transport)
}
"openai:chat"
| "openai:responses"
| "openai:responses:compact"
| "openai:embedding"
| "jina:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
| "openai:rerank"
| "jina:rerank" => resolve_local_openai_bearer_auth(transport),
"gemini:generate_content" | "gemini:embedding" => {
if is_vertex_api_key_transport_context(transport) {
resolve_local_vertex_api_key_query_auth(transport)
.map(|auth| (VERTEX_API_KEY_QUERY_PARAM.to_string(), auth.value))
} else {
resolve_local_gemini_auth(transport)
}
}
"claude:messages" => resolve_local_standard_auth(transport),
_ => None,
}
}
pub fn candidate_common_transport_skip_reason(
transport: &GatewayProviderTransportSnapshot,
candidate: CandidateTransportPolicyFacts<'_>,
requested_model: Option<&str>,
) -> Option<&'static str> {
let requested_model = requested_model.unwrap_or_default();
if !transport.provider.is_active {
return Some("provider_inactive");
}
if !transport.endpoint.is_active {
return Some("endpoint_inactive");
}
if !transport.key.is_active {
return Some("key_inactive");
}
let endpoint_api_format = transport.endpoint.api_format.trim();
if !aether_ai_formats::api_format_alias_matches(
candidate.endpoint_api_format,
endpoint_api_format,
) {
return Some("endpoint_api_format_changed");
}
if !transport_key_supports_api_format(transport, endpoint_api_format) {
return Some("key_api_format_disabled");
}
if !transport_key_allows_candidate_model(transport, requested_model, candidate) {
return Some("key_model_disabled");
}
None
}
pub fn candidate_transport_pair_skip_reason(
transport: &GatewayProviderTransportSnapshot,
normalized_client_api_format: &str,
) -> Option<&'static str> {
let endpoint_api_format = transport.endpoint.api_format.trim();
if aether_ai_formats::api_format_alias_matches(
endpoint_api_format,
normalized_client_api_format,
) {
return None;
}
if let Some(skip_reason) =
disabled_format_conversion_skip_reason(transport, normalized_client_api_format)
{
return Some(skip_reason);
}
if !request_pair_allowed_for_transport(
transport,
normalized_client_api_format,
endpoint_api_format,
) {
return Some("transport_unsupported");
}
None
}
fn disabled_format_conversion_skip_reason(
transport: &GatewayProviderTransportSnapshot,
normalized_client_api_format: &str,
) -> Option<&'static str> {
let endpoint_api_format = transport.endpoint.api_format.trim();
if aether_ai_formats::api_format_alias_matches(
endpoint_api_format,
normalized_client_api_format,
) {
return None;
}
request_conversion_kind(normalized_client_api_format, endpoint_api_format)?;
if request_conversion_requires_enable_flag(normalized_client_api_format, endpoint_api_format)
&& !request_conversion_enabled_for_transport(
transport,
normalized_client_api_format,
endpoint_api_format,
)
{
return Some("format_conversion_disabled");
}
None
}
fn transport_key_supports_api_format(
transport: &GatewayProviderTransportSnapshot,
endpoint_api_format: &str,
) -> bool {
let inherits_provider_api_formats =
crate::provider_types::fixed_provider_key_inherits_api_formats(
transport.provider.provider_type.as_str(),
transport.key.auth_type.as_str(),
transport.key.decrypted_auth_config.as_deref(),
);
if inherits_provider_api_formats {
return true;
}
match transport.key.api_formats.as_deref() {
None => true,
Some(formats) => formats.iter().any(|value| {
aether_ai_formats::api_format_permission_covers(value, endpoint_api_format)
}),
}
}
fn transport_key_allows_candidate_model(
transport: &GatewayProviderTransportSnapshot,
requested_model: &str,
candidate: CandidateTransportPolicyFacts<'_>,
) -> bool {
let Some(allowed_models) = transport.key.allowed_models.as_deref() else {
return true;
};
let requested_model = requested_model.trim();
let global_model_name = candidate.global_model_name.trim();
let selected_provider_model_name = candidate.selected_provider_model_name.trim();
let mapping_matched_model = candidate
.mapping_matched_model
.map(str::trim)
.filter(|value| !value.is_empty());
let requested_base_model = aether_ai_formats::model_directive_base_model(requested_model);
for allowed_model in allowed_models.iter().map(String::as_str).map(str::trim) {
if allowed_model.is_empty() {
continue;
}
if allowed_model == requested_model
|| requested_base_model
.as_deref()
.is_some_and(|base_model| allowed_model == base_model)
|| allowed_model == global_model_name
|| allowed_model == selected_provider_model_name
|| mapping_matched_model.is_some_and(|value| value == allowed_model)
{
return true;
}
}
false
}
fn endpoint_accepts_client_api_format(
transport: &GatewayProviderTransportSnapshot,
client_api_format: &str,
) -> bool {
let Some(config) = transport
.endpoint
.format_acceptance_config
.as_ref()
.and_then(serde_json::Value::as_object)
else {
return false;
};
if !config
.get("enabled")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
return false;
}
if config
.get("reject_formats")
.is_some_and(|value| json_format_list_contains(value, client_api_format))
{
return false;
}
match config.get("accept_formats") {
Some(value) => json_format_list_contains(value, client_api_format),
None => true,
}
}
fn json_format_list_contains(value: &serde_json::Value, api_format: &str) -> bool {
let Some(items) = value.as_array() else {
return false;
};
items.iter().any(|item| {
item.as_str().is_some_and(|candidate| {
aether_ai_formats::api_format_alias_matches(candidate, api_format)
})
})
}
#[cfg(test)]
mod tests {
use super::{
candidate_common_transport_skip_reason, candidate_transport_pair_skip_reason,
request_conversion_direct_auth, request_conversion_enabled_for_transport,
request_conversion_transport_supported, request_pair_allowed_for_transport,
request_pair_direct_auth, CandidateTransportPolicyFacts,
};
use aether_ai_formats::formats::matrix::RequestConversionKind;
use serde_json::json;
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn transport_snapshot(
provider_type: &str,
api_format: &str,
auth_type: &str,
enable_format_conversion: bool,
format_acceptance_config: Option<serde_json::Value>,
) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: format!("provider-{provider_type}"),
name: provider_type.to_string(),
provider_type: provider_type.to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: format!("provider-{provider_type}"),
api_format: api_format.to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: if provider_type == "kiro" {
"https://q.{region}.amazonaws.com".to_string()
} else if provider_type == "vertex_ai" {
"https://aiplatform.googleapis.com".to_string()
} else {
"https://api.example.com".to_string()
},
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: format!("provider-{provider_type}"),
name: "key".to_string(),
auth_type: auth_type.to_string(),
is_active: true,
api_formats: Some(vec![api_format.to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
}
}
fn candidate_facts(api_format: &str) -> CandidateTransportPolicyFacts<'_> {
CandidateTransportPolicyFacts {
endpoint_api_format: api_format,
global_model_name: "gpt-4.1",
selected_provider_model_name: "provider-gpt-4.1",
mapping_matched_model: Some("gpt-4.1-mini"),
}
}
#[test]
fn conversion_helpers_follow_transport_api_format() {
let transport = transport_snapshot("openai", "openai:chat", "bearer", true, None);
assert!(request_conversion_transport_supported(
&transport,
RequestConversionKind::ToOpenAIChat
));
assert_eq!(
request_conversion_direct_auth(&transport, RequestConversionKind::ToOpenAIChat),
Some(("authorization".to_string(), "Bearer secret".to_string()))
);
}
#[test]
fn endpoint_level_format_acceptance_enables_cross_format_pair_without_provider_flag() {
let transport = transport_snapshot(
"custom",
"openai:responses",
"bearer",
false,
Some(json!({
"enabled": true,
"accept_formats": ["claude:messages"],
})),
);
assert!(request_conversion_enabled_for_transport(
&transport,
"claude:messages",
"openai:responses"
));
assert!(request_pair_allowed_for_transport(
&transport,
"claude:messages",
"openai:responses"
));
assert!(!request_pair_allowed_for_transport(
&transport,
"gemini:generate_content",
"openai:responses"
));
}
#[test]
fn endpoint_reject_formats_override_endpoint_cross_format_enablement() {
let transport = transport_snapshot(
"custom",
"openai:responses",
"bearer",
false,
Some(json!({
"enabled": true,
"reject_formats": ["claude:messages"],
})),
);
assert!(!request_conversion_enabled_for_transport(
&transport,
"claude:messages",
"openai:responses"
));
}
#[test]
fn cross_format_pair_requires_provider_or_endpoint_enablement() {
let disabled = transport_snapshot("custom", "openai:responses", "bearer", false, None);
assert!(!request_conversion_enabled_for_transport(
&disabled,
"claude:messages",
"openai:responses",
));
assert_eq!(
candidate_transport_pair_skip_reason(&disabled, "claude:messages"),
Some("format_conversion_disabled")
);
let provider_enabled =
transport_snapshot("custom", "openai:responses", "bearer", true, None);
assert!(request_conversion_enabled_for_transport(
&provider_enabled,
"claude:messages",
"openai:responses",
));
assert_eq!(
candidate_transport_pair_skip_reason(&provider_enabled, "claude:messages"),
None
);
let endpoint_enabled = transport_snapshot(
"custom",
"openai:responses",
"bearer",
false,
Some(json!({
"enabled": true,
"accept_formats": ["claude:messages"],
})),
);
assert!(request_conversion_enabled_for_transport(
&endpoint_enabled,
"claude:messages",
"openai:responses",
));
assert_eq!(
candidate_transport_pair_skip_reason(&endpoint_enabled, "claude:messages"),
None
);
}
#[test]
fn vertex_gemini_transport_supports_cross_format_conversion_with_query_auth() {
let transport = transport_snapshot(
"vertex_ai",
"gemini:generate_content",
"api_key",
true,
None,
);
assert!(request_conversion_transport_supported(
&transport,
RequestConversionKind::ToGeminiStandard
));
assert_eq!(
request_conversion_direct_auth(&transport, RequestConversionKind::ToGeminiStandard),
Some(("key".to_string(), "secret".to_string()))
);
}
#[test]
fn antigravity_gemini_transport_supports_standard_cross_format_conversion_via_envelope() {
let transport = transport_snapshot(
"antigravity",
"gemini:generate_content",
"oauth",
true,
None,
);
assert!(request_pair_allowed_for_transport(
&transport,
"openai:chat",
"gemini:generate_content"
));
assert!(request_pair_allowed_for_transport(
&transport,
"openai:responses",
"gemini:generate_content"
));
assert!(request_conversion_transport_supported(
&transport,
RequestConversionKind::ToGeminiStandard
));
}
#[test]
fn vertex_gemini_embedding_transport_supports_openai_embedding_conversion() {
let transport = transport_snapshot("vertex_ai", "gemini:embedding", "api_key", true, None);
assert!(request_pair_allowed_for_transport(
&transport,
"openai:embedding",
"gemini:embedding"
));
assert_eq!(
request_pair_direct_auth(&transport, "gemini:embedding"),
Some(("key".to_string(), "secret".to_string()))
);
}
#[test]
fn kiro_claude_messages_transport_supports_cross_format_conversion_via_envelope() {
let transport = transport_snapshot("kiro", "claude:messages", "bearer", true, None);
assert!(request_pair_allowed_for_transport(
&transport,
"openai:chat",
"claude:messages"
));
assert!(request_conversion_transport_supported(
&transport,
RequestConversionKind::ToClaudeStandard
));
}
#[test]
fn windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() {
let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None);
transport.key.decrypted_api_key = "devin-session-token$abc".to_string();
transport.key.decrypted_auth_config = Some(r#"{"provider_type":"windsurf"}"#.to_string());
assert!(request_pair_allowed_for_transport(
&transport,
"claude:messages",
"openai:chat"
));
assert!(request_pair_allowed_for_transport(
&transport,
"openai:responses",
"openai:chat"
));
assert!(request_conversion_transport_supported(
&transport,
RequestConversionKind::ToOpenAIChat
));
assert_eq!(
request_conversion_direct_auth(&transport, RequestConversionKind::ToOpenAIChat),
Some((
"authorization".to_string(),
"Bearer devin-session-token$abc".to_string()
))
);
}
#[test]
fn candidate_common_transport_policy_checks_active_state_format_and_allowed_models() {
let mut transport = transport_snapshot("custom", "openai:chat", "bearer", true, None);
transport.key.allowed_models = Some(vec!["gpt-4.1-mini".to_string()]);
assert_eq!(
candidate_common_transport_skip_reason(
&transport,
candidate_facts("openai:chat"),
Some("gpt-4.1")
),
None
);
assert_eq!(
candidate_common_transport_skip_reason(
&transport,
candidate_facts("claude:messages"),
Some("gpt-4.1")
),
Some("endpoint_api_format_changed")
);
transport.key.allowed_models = Some(vec!["other-model".to_string()]);
assert_eq!(
candidate_common_transport_skip_reason(
&transport,
candidate_facts("openai:chat"),
Some("gpt-4.1")
),
Some("key_model_disabled")
);
}
#[test]
fn candidate_common_transport_policy_allows_model_directive_base_model() {
let mut transport = transport_snapshot("custom", "openai:responses", "bearer", true, None);
transport.key.allowed_models = Some(vec!["gpt-5.5".to_string()]);
assert_eq!(
candidate_common_transport_skip_reason(
&transport,
CandidateTransportPolicyFacts {
endpoint_api_format: "openai:responses",
global_model_name: "gpt-5",
selected_provider_model_name: "provider-gpt-5",
mapping_matched_model: None,
},
Some("gpt-5.5-xhigh"),
),
None
);
}
#[test]
fn fixed_provider_oauth_keys_inherit_endpoint_api_formats_for_candidate_policy() {
let mut transport = transport_snapshot("codex", "openai:responses", "oauth", true, None);
transport.key.api_formats = Some(vec!["openai:image".to_string()]);
assert_eq!(
candidate_common_transport_skip_reason(
&transport,
candidate_facts("openai:responses"),
None,
),
None
);
}
#[test]
fn responses_key_permission_covers_search_without_changing_endpoint_identity() {
let mut transport = transport_snapshot("custom", "openai:search", "bearer", true, None);
transport.key.api_formats = Some(vec!["openai:responses".to_string()]);
assert_eq!(
candidate_common_transport_skip_reason(
&transport,
candidate_facts("openai:search"),
None,
),
None
);
assert_eq!(
candidate_common_transport_skip_reason(
&transport,
candidate_facts("openai:responses"),
None,
),
Some("endpoint_api_format_changed")
);
}
#[test]
fn candidate_transport_pair_policy_reports_disabled_conversion_and_unsupported_pairs() {
let transport = transport_snapshot("custom", "openai:responses", "bearer", false, None);
assert_eq!(
candidate_transport_pair_skip_reason(&transport, "openai:chat"),
Some("format_conversion_disabled")
);
assert_eq!(
candidate_transport_pair_skip_reason(&transport, "gemini:video"),
Some("transport_unsupported")
);
let enabled_transport =
transport_snapshot("custom", "openai:responses", "bearer", true, None);
assert_eq!(
candidate_transport_pair_skip_reason(&enabled_transport, "openai:chat"),
None
);
}
}
@@ -0,0 +1,511 @@
use aether_ai_formats::formats::matrix::{
request_conversion_kind, request_conversion_requires_enable_flag,
};
use aether_contracts::ProxySnapshot;
use serde_json::{json, Map, Value};
use crate::conversion::{
request_conversion_enabled_for_transport, request_conversion_transport_unsupported_reason,
request_pair_allowed_for_transport,
};
use crate::grok::grok_browser_resolved_transport_profile_from_auth_config;
use crate::network::{
resolve_transport_profile, resolve_transport_profile_id, transport_proxy_is_locally_supported,
};
use crate::policy::{
local_gemini_transport_unsupported_reason_with_network,
local_openai_chat_transport_unsupported_reason,
local_standard_transport_unsupported_reason_with_network,
};
use crate::rules::{body_rules_are_locally_supported, header_rules_are_locally_supported};
use crate::same_format_provider::same_format_provider_transport_unsupported_reason_for_trace;
use crate::snapshot::GatewayProviderTransportSnapshot;
pub fn build_request_trace_proxy_value(
transport: Option<&GatewayProviderTransportSnapshot>,
resolved_proxy: Option<&ProxySnapshot>,
) -> Option<Value> {
let resolved_proxy = resolved_proxy?;
let mut object = Map::new();
if let Some(node_id) = trimmed_non_empty(resolved_proxy.node_id.as_deref()) {
object.insert("node_id".to_string(), Value::String(node_id));
}
if let Some(node_name) = trimmed_non_empty(resolved_proxy.label.as_deref()) {
object.insert("node_name".to_string(), Value::String(node_name));
}
if let Some(url) = sanitize_trace_proxy_url(resolved_proxy.url.as_deref()) {
object.insert("url".to_string(), Value::String(url));
}
if let Some(source) = resolve_request_trace_proxy_source(transport, true) {
object.insert("source".to_string(), Value::String(source.to_string()));
}
(!object.is_empty()).then_some(Value::Object(object))
}
pub fn append_transport_diagnostics_to_value(
value: Value,
transport: Option<&GatewayProviderTransportSnapshot>,
client_api_format: &str,
provider_api_format: &str,
) -> Value {
let Value::Object(mut object) = value else {
return value;
};
object.insert(
"transport_diagnostics".to_string(),
transport
.map(|transport| {
build_transport_diagnostics(transport, client_api_format, provider_api_format)
})
.unwrap_or_else(|| json!({ "transport_snapshot_available": false })),
);
Value::Object(object)
}
pub fn build_transport_diagnostics(
transport: &GatewayProviderTransportSnapshot,
client_api_format: &str,
provider_api_format: &str,
) -> Value {
let resolved_transport_profile_id = resolve_transport_profile_id(transport);
let resolved_transport_profile = resolve_transport_profile(transport)
.and_then(|profile| serde_json::to_value(profile).ok())
.unwrap_or(Value::Null);
let configured_key_transport_profile = transport
.key
.fingerprint
.as_ref()
.and_then(Value::as_object)
.and_then(|value| value.get("transport_profile"))
.cloned()
.unwrap_or(Value::Null);
let configured_provider_transport_profile = transport
.provider
.config
.as_ref()
.and_then(|value| value.get("fingerprint"))
.and_then(Value::as_object)
.and_then(|value| value.get("transport_profile"))
.cloned()
.unwrap_or(Value::Null);
let configured_legacy_grok_transport_profile = if transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok")
{
transport
.key
.decrypted_auth_config
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| serde_json::from_str::<Value>(value).ok())
.and_then(|value| value.as_object().cloned())
.and_then(|auth_config| {
grok_browser_resolved_transport_profile_from_auth_config(
&auth_config,
"grok_auth_config",
)
.and_then(|profile| serde_json::to_value(profile).ok())
})
.unwrap_or(Value::Null)
} else {
Value::Null
};
let has_oauth_config = transport.key.decrypted_auth_config.is_some();
let oauth_resolution_supported =
!has_oauth_config || crate::supports_local_oauth_request_auth_resolution(transport);
let request_transport_unsupported_reason = resolve_request_transport_unsupported_reason(
transport,
client_api_format,
provider_api_format,
);
json!({
"transport_snapshot_available": true,
"provider_type": transport.provider.provider_type,
"provider_is_active": transport.provider.is_active,
"endpoint_is_active": transport.endpoint.is_active,
"key_is_active": transport.key.is_active,
"provider_enable_format_conversion": transport.provider.enable_format_conversion,
"provider_keep_priority_on_conversion": transport.provider.keep_priority_on_conversion,
"endpoint_format_acceptance_config": transport.endpoint.format_acceptance_config,
"endpoint_custom_path": transport.endpoint.custom_path,
"header_rules": transport.endpoint.header_rules,
"header_rules_supported": header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()),
"body_rules": transport.endpoint.body_rules,
"body_rules_supported": body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()),
"proxy": {
"locally_supported": transport_proxy_is_locally_supported(transport),
"provider": summarize_proxy_config(transport.provider.proxy.as_ref()),
"endpoint": summarize_proxy_config(transport.endpoint.proxy.as_ref()),
"key": summarize_proxy_config(transport.key.proxy.as_ref()),
},
"auth": {
"key_auth_type": transport.key.auth_type,
"has_oauth_config": has_oauth_config,
"oauth_request_auth_resolution_supported": oauth_resolution_supported,
},
"fingerprint": transport.key.fingerprint,
"configured_key_transport_profile": configured_key_transport_profile,
"configured_provider_transport_profile": configured_provider_transport_profile,
"configured_legacy_grok_transport_profile": configured_legacy_grok_transport_profile,
"resolved_transport_profile_id": resolved_transport_profile_id,
"resolved_transport_profile": resolved_transport_profile,
"request_pair": {
"client_api_format": client_api_format,
"provider_api_format": provider_api_format,
"requires_conversion_enable_flag": request_conversion_requires_enable_flag(
client_api_format,
provider_api_format,
),
"conversion_enabled": request_conversion_enabled_for_transport(
transport,
client_api_format,
provider_api_format,
),
"pair_allowed": request_pair_allowed_for_transport(
transport,
client_api_format,
provider_api_format,
),
"transport_unsupported_reason": request_transport_unsupported_reason,
},
})
}
fn summarize_proxy_config(proxy: Option<&Value>) -> Value {
let Some(object) = proxy.and_then(Value::as_object) else {
return Value::Null;
};
let has_url = object
.get("url")
.or_else(|| object.get("proxy_url"))
.and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty());
json!({
"enabled": object.get("enabled").cloned().unwrap_or(Value::Null),
"mode": object.get("mode").cloned().unwrap_or(Value::Null),
"node_id": object.get("node_id").cloned().unwrap_or(Value::Null),
"label": object.get("label").cloned().unwrap_or(Value::Null),
"has_url": has_url,
})
}
fn resolve_request_trace_proxy_source(
transport: Option<&GatewayProviderTransportSnapshot>,
has_resolved_proxy: bool,
) -> Option<&'static str> {
let transport = transport?;
if transport_has_explicit_proxy(transport.key.proxy.as_ref()) {
return Some("key");
}
if transport_has_explicit_proxy(transport.endpoint.proxy.as_ref()) {
return Some("endpoint");
}
if transport_has_explicit_proxy(transport.provider.proxy.as_ref()) {
return Some("provider");
}
has_resolved_proxy.then_some("system")
}
fn transport_has_explicit_proxy(proxy: Option<&Value>) -> bool {
let Some(object) = proxy.and_then(Value::as_object) else {
return false;
};
let enabled = object
.get("enabled")
.and_then(Value::as_bool)
.unwrap_or(true);
if !enabled {
return false;
}
object
.get("node_id")
.or_else(|| object.get("url"))
.or_else(|| object.get("proxy_url"))
.and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty())
}
fn sanitize_trace_proxy_url(url: Option<&str>) -> Option<String> {
let raw = url.map(str::trim).filter(|value| !value.is_empty())?;
let parsed = url::Url::parse(raw).ok()?;
let scheme = parsed.scheme().trim();
let host = parsed.host_str()?.trim();
if scheme.is_empty() || host.is_empty() {
return None;
}
let mut safe = format!("{scheme}://{host}");
if let Some(port) = parsed.port() {
safe.push(':');
safe.push_str(port.to_string().as_str());
}
Some(safe)
}
fn trimmed_non_empty(value: Option<&str>) -> Option<String> {
value
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn resolve_request_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
client_api_format: &str,
provider_api_format: &str,
) -> Option<&'static str> {
let client_api_format = client_api_format.trim().to_ascii_lowercase();
let provider_api_format = provider_api_format.trim().to_ascii_lowercase();
if client_api_format == provider_api_format {
if let Some(skip_reason) = same_format_provider_transport_unsupported_reason_for_trace(
transport,
provider_api_format.as_str(),
) {
return Some(skip_reason);
}
return match provider_api_format.as_str() {
"openai:chat" => local_openai_chat_transport_unsupported_reason(transport),
"gemini:generate_content" => local_gemini_transport_unsupported_reason_with_network(
transport,
provider_api_format.as_str(),
),
_ => local_standard_transport_unsupported_reason_with_network(
transport,
provider_api_format.as_str(),
),
};
}
match request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str()) {
Some(kind) => request_conversion_transport_unsupported_reason(transport, kind),
None => Some("transport_api_format_unsupported"),
}
}
#[cfg(test)]
mod tests {
use super::{build_request_trace_proxy_value, build_transport_diagnostics};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use aether_contracts::ProxySnapshot;
use serde_json::{json, Value};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "RightCode".to_string(),
provider_type: "codex".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: Some(json!({"enabled": true, "mode": "node", "node_id": "proxy-node-1"})),
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "openai:responses".to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://example.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: Some("/v1/responses".to_string()),
config: None,
format_acceptance_config: Some(json!({
"enabled": true,
"accept_formats": ["claude:messages"]
})),
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "codex".to_string(),
auth_type: "oauth".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: Some(json!({
"transport_profile": {
"profile_id": "chrome_136",
"header_fingerprint": {
"user_agent": "Mozilla/5.0"
}
}
})),
upstream_metadata: None,
decrypted_api_key: "sk-test".to_string(),
decrypted_auth_config: None,
},
}
}
fn sample_claude_code_transport_without_auth() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-cc-1".to_string(),
name: "NekoCode".to_string(),
provider_type: "claude_code".to_string(),
website: Some("https://nekocode.ai".to_string()),
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-cc-1".to_string(),
provider_id: "provider-cc-1".to_string(),
api_format: "claude:messages".to_string(),
api_family: Some("claude".to_string()),
endpoint_kind: Some("cli".to_string()),
is_active: true,
base_url: "https://api.anthropic.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-cc-1".to_string(),
provider_id: "provider-cc-1".to_string(),
name: "CC".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: Some(vec!["claude:messages".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn transport_diagnostics_include_format_and_network_policy() {
let diagnostics =
build_transport_diagnostics(&sample_transport(), "claude:messages", "openai:responses");
assert_eq!(diagnostics["provider_type"], "codex");
assert_eq!(
diagnostics["fingerprint"]["transport_profile"]["profile_id"],
"chrome_136"
);
assert_eq!(diagnostics["resolved_transport_profile_id"], "chrome_136");
assert_eq!(
diagnostics["request_pair"]["conversion_enabled"],
Value::Bool(true)
);
assert!(diagnostics["request_pair"]["transport_unsupported_reason"].is_null());
}
#[test]
fn transport_diagnostics_use_provider_private_same_format_reason() {
let diagnostics = build_transport_diagnostics(
&sample_claude_code_transport_without_auth(),
"claude:messages",
"claude:messages",
);
assert_eq!(
diagnostics["request_pair"]["transport_unsupported_reason"],
Value::String("transport_auth_unavailable".to_string())
);
}
fn sample_grok_transport_with_legacy_user_agent() -> GatewayProviderTransportSnapshot {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.fingerprint = None;
transport.provider.config = None;
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token",
"user_agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/137.0.0.0 Safari/537.36"
})
.to_string(),
);
transport
}
#[test]
fn transport_diagnostics_include_legacy_grok_transport_profile() {
let diagnostics = build_transport_diagnostics(
&sample_grok_transport_with_legacy_user_agent(),
"openai:chat",
"openai:chat",
);
assert_eq!(
diagnostics["configured_legacy_grok_transport_profile"]["profile_id"],
"chrome137"
);
assert_eq!(
diagnostics["resolved_transport_profile"]["profile_id"],
"chrome137"
);
assert_eq!(
diagnostics["resolved_transport_profile"]["backend"],
"browser_wreq"
);
}
#[test]
fn request_trace_proxy_value_sanitizes_url_and_marks_config_source() {
let transport = sample_transport();
let proxy = ProxySnapshot {
enabled: Some(true),
mode: Some("node".to_string()),
node_id: Some("proxy-node-1".to_string()),
label: Some("Primary proxy".to_string()),
url: Some("https://user:[email protected]:9443/path".to_string()),
extra: None,
};
let value = build_request_trace_proxy_value(Some(&transport), Some(&proxy)).unwrap();
assert_eq!(value["node_id"], "proxy-node-1");
assert_eq!(value["node_name"], "Primary proxy");
assert_eq!(value["url"], "https://example.test:9443");
assert_eq!(value["source"], "provider");
}
}
@@ -0,0 +1,272 @@
use serde_json::Value;
use super::super::snapshot::GatewayProviderTransportSnapshot;
pub const GEMINI_CLI_PROVIDER_TYPE: &str = "gemini_cli";
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct GeminiCliRequestAuth {
pub project_id: Option<String>,
pub session_id: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum GeminiCliRequestAuthSupport {
Supported(GeminiCliRequestAuth),
Unsupported(GeminiCliRequestAuthUnsupportedReason),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum GeminiCliRequestAuthUnsupportedReason {
WrongProviderType,
InvalidAuthConfigJson,
}
pub fn is_gemini_cli_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(GEMINI_CLI_PROVIDER_TYPE)
}
pub fn resolve_local_gemini_cli_request_auth(
transport: &GatewayProviderTransportSnapshot,
) -> GeminiCliRequestAuthSupport {
if !is_gemini_cli_provider_transport(transport) {
return GeminiCliRequestAuthSupport::Unsupported(
GeminiCliRequestAuthUnsupportedReason::WrongProviderType,
);
}
let metadata = transport.key.upstream_metadata.as_ref();
let metadata_project_id = metadata.and_then(resolve_project_id_from_value);
let metadata_session_id = metadata.and_then(resolve_session_id_from_value);
let Some(raw_auth_config) = transport
.key
.decrypted_auth_config
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return GeminiCliRequestAuthSupport::Supported(GeminiCliRequestAuth {
project_id: metadata_project_id,
session_id: metadata_session_id,
});
};
let Ok(auth_config) = serde_json::from_str::<Value>(raw_auth_config) else {
return GeminiCliRequestAuthSupport::Unsupported(
GeminiCliRequestAuthUnsupportedReason::InvalidAuthConfigJson,
);
};
GeminiCliRequestAuthSupport::Supported(GeminiCliRequestAuth {
project_id: metadata_project_id.or_else(|| resolve_project_id_from_value(&auth_config)),
session_id: metadata_session_id.or_else(|| resolve_session_id_from_value(&auth_config)),
})
}
pub fn resolve_gemini_cli_project_id(
transport: &GatewayProviderTransportSnapshot,
) -> Option<String> {
match resolve_local_gemini_cli_request_auth(transport) {
GeminiCliRequestAuthSupport::Supported(auth) => auth.project_id,
GeminiCliRequestAuthSupport::Unsupported(_) => None,
}
}
fn resolve_project_id_from_value(value: &Value) -> Option<String> {
find_string_by_paths(
value,
&[
&["project"],
&["project_id"],
&["projectId"],
&["project", "id"],
&["project", "project_id"],
&["project", "projectId"],
&["cloudaicompanionProject"],
&["cloudaicompanionProject", "id"],
&["cloudaicompanion_project"],
&["cloudaicompanion_project", "id"],
&["gemini_cli", "project"],
&["gemini_cli", "project_id"],
&["gemini_cli", "projectId"],
&["gemini_cli", "cloudaicompanionProject"],
&["gemini_cli", "cloudaicompanionProject", "id"],
&["geminiCli", "project"],
&["geminiCli", "project_id"],
&["geminiCli", "projectId"],
&["metadata", "project"],
&["metadata", "project_id"],
&["metadata", "projectId"],
],
)
}
fn resolve_session_id_from_value(value: &Value) -> Option<String> {
find_string_by_paths(
value,
&[
&["session_id"],
&["sessionId"],
&["gemini_cli", "session_id"],
&["gemini_cli", "sessionId"],
&["geminiCli", "session_id"],
&["geminiCli", "sessionId"],
&["metadata", "session_id"],
&["metadata", "sessionId"],
],
)
}
fn find_string_by_paths(value: &Value, paths: &[&[&str]]) -> Option<String> {
for path in paths {
let mut current = value;
let mut matched = true;
for segment in *path {
let Some(next) = current.get(*segment) else {
matched = false;
break;
};
current = next;
}
if !matched {
continue;
}
if let Some(string) = current
.as_str()
.map(str::trim)
.filter(|item| !item.is_empty())
{
return Some(string.to_string());
}
}
None
}
#[cfg(test)]
mod tests {
use super::{
resolve_gemini_cli_project_id, resolve_local_gemini_cli_request_auth, GeminiCliRequestAuth,
GeminiCliRequestAuthSupport,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Gemini CLI".to_string(),
provider_type: "gemini_cli".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "gemini:generate_content".to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://cloudcode-pa.googleapis.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "oauth".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "__oauth__".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn extracts_project_and_session_metadata_when_available() {
let mut transport = sample_transport();
transport.key.decrypted_auth_config = Some(
r#"{
"cloudaicompanionProject": {"id": "project-123"},
"metadata": {"sessionId": "session-123"}
}"#
.to_string(),
);
assert_eq!(
resolve_local_gemini_cli_request_auth(&transport),
GeminiCliRequestAuthSupport::Supported(GeminiCliRequestAuth {
project_id: Some("project-123".to_string()),
session_id: Some("session-123".to_string()),
})
);
}
#[test]
fn project_id_prefers_upstream_metadata_then_auth_config() {
let mut transport = sample_transport();
transport.key.upstream_metadata = Some(serde_json::json!({
"gemini_cli": {
"project_id": "metadata-project"
}
}));
transport.key.decrypted_auth_config = Some(r#"{"project_id":"auth-project"}"#.to_string());
assert_eq!(
resolve_gemini_cli_project_id(&transport).as_deref(),
Some("metadata-project")
);
transport.key.upstream_metadata = None;
assert_eq!(
resolve_gemini_cli_project_id(&transport).as_deref(),
Some("auth-project")
);
}
#[test]
fn missing_metadata_still_supports_request_envelope() {
let transport = sample_transport();
assert_eq!(
resolve_local_gemini_cli_request_auth(&transport),
GeminiCliRequestAuthSupport::Supported(GeminiCliRequestAuth {
project_id: None,
session_id: None,
})
);
}
}
@@ -0,0 +1,22 @@
mod auth;
mod policy;
mod request;
mod url;
pub use auth::{
is_gemini_cli_provider_transport, resolve_gemini_cli_project_id,
resolve_local_gemini_cli_request_auth, GeminiCliRequestAuth, GeminiCliRequestAuthSupport,
GeminiCliRequestAuthUnsupportedReason, GEMINI_CLI_PROVIDER_TYPE,
};
pub use policy::gemini_cli_v1internal_requires_upstream_streaming;
pub use request::{
build_gemini_cli_v1internal_request, classify_gemini_cli_v1internal_request_body,
GeminiCliRequestEnvelopeSupport, GeminiCliRequestEnvelopeUnsupportedReason,
};
pub use url::{
build_gemini_cli_v1internal_url, GeminiCliRequestUrlAction,
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT,
GEMINI_CLI_V1INTERNAL_PATH_TEMPLATE,
};
pub const GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME: &str = "gemini_cli:v1internal";
@@ -0,0 +1,29 @@
pub fn gemini_cli_v1internal_requires_upstream_streaming(
provider_api_format: &str,
client_requires_streaming: bool,
) -> bool {
client_requires_streaming
&& aether_ai_formats::normalize_api_format_alias(provider_api_format)
== "gemini:generate_content"
}
#[cfg(test)]
mod tests {
use super::gemini_cli_v1internal_requires_upstream_streaming;
#[test]
fn v1internal_generate_content_requires_upstream_streaming_for_stream_clients() {
assert!(gemini_cli_v1internal_requires_upstream_streaming(
"gemini:generate_content",
true,
));
assert!(!gemini_cli_v1internal_requires_upstream_streaming(
"gemini:generate_content",
false,
));
assert!(!gemini_cli_v1internal_requires_upstream_streaming(
"openai:chat",
true,
));
}
}
@@ -0,0 +1,243 @@
use serde_json::{Map, Value};
use super::auth::GeminiCliRequestAuth;
#[derive(Debug, Clone, PartialEq)]
pub enum GeminiCliRequestEnvelopeSupport {
Supported(Value),
Unsupported(GeminiCliRequestEnvelopeUnsupportedReason),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum GeminiCliRequestEnvelopeUnsupportedReason {
NonObjectBody,
MissingContents,
MissingUserPromptId,
MissingModel,
}
pub fn classify_gemini_cli_v1internal_request_body(
request_body: &Value,
) -> Result<(), GeminiCliRequestEnvelopeUnsupportedReason> {
let Value::Object(map) = request_body else {
return Err(GeminiCliRequestEnvelopeUnsupportedReason::NonObjectBody);
};
if !map.contains_key("contents") && existing_v1internal_request_object(map).is_none() {
return Err(GeminiCliRequestEnvelopeUnsupportedReason::MissingContents);
}
Ok(())
}
pub fn build_gemini_cli_v1internal_request(
auth: &GeminiCliRequestAuth,
user_prompt_id: &str,
model: &str,
request_body: &Value,
) -> GeminiCliRequestEnvelopeSupport {
if user_prompt_id.trim().is_empty() {
return GeminiCliRequestEnvelopeSupport::Unsupported(
GeminiCliRequestEnvelopeUnsupportedReason::MissingUserPromptId,
);
}
if model.trim().is_empty() {
return GeminiCliRequestEnvelopeSupport::Unsupported(
GeminiCliRequestEnvelopeUnsupportedReason::MissingModel,
);
}
if let Err(reason) = classify_gemini_cli_v1internal_request_body(request_body) {
return GeminiCliRequestEnvelopeSupport::Unsupported(reason);
}
let Value::Object(source) = request_body else {
return GeminiCliRequestEnvelopeSupport::Unsupported(
GeminiCliRequestEnvelopeUnsupportedReason::NonObjectBody,
);
};
let existing_request = existing_v1internal_request_object(source);
let mut inner_request: Map<String, Value> =
existing_request.cloned().unwrap_or_else(|| source.clone());
sanitize_inner_request(&mut inner_request);
maybe_insert_session_id(&mut inner_request, auth.session_id.as_deref());
let project = non_empty_string_field(source, "project")
.map(ToOwned::to_owned)
.or_else(|| {
auth.project_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
});
let user_prompt_id = non_empty_string_field(source, "user_prompt_id")
.or_else(|| non_empty_string_field(source, "userPromptId"))
.unwrap_or(user_prompt_id)
.trim()
.to_string();
let mut envelope = Map::new();
envelope.insert("model".to_string(), Value::String(model.trim().to_string()));
if let Some(project) = project {
envelope.insert("project".to_string(), Value::String(project));
}
envelope.insert("user_prompt_id".to_string(), Value::String(user_prompt_id));
envelope.insert("request".to_string(), Value::Object(inner_request));
GeminiCliRequestEnvelopeSupport::Supported(Value::Object(envelope))
}
fn sanitize_inner_request(inner_request: &mut Map<String, Value>) {
inner_request.remove("model");
inner_request.remove("stream");
}
fn maybe_insert_session_id(inner_request: &mut Map<String, Value>, session_id: Option<&str>) {
let Some(session_id) = session_id.map(str::trim).filter(|value| !value.is_empty()) else {
return;
};
if inner_request.contains_key("session_id") || inner_request.contains_key("sessionId") {
return;
}
inner_request.insert(
"session_id".to_string(),
Value::String(session_id.to_string()),
);
}
fn existing_v1internal_request_object(source: &Map<String, Value>) -> Option<&Map<String, Value>> {
source
.get("request")
.and_then(Value::as_object)
.filter(|request| request.contains_key("contents"))
}
fn non_empty_string_field<'a>(source: &'a Map<String, Value>, key: &str) -> Option<&'a str> {
source
.get(key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
build_gemini_cli_v1internal_request, classify_gemini_cli_v1internal_request_body,
GeminiCliRequestEnvelopeSupport,
};
use crate::gemini_cli::GeminiCliRequestAuth;
fn sample_auth() -> GeminiCliRequestAuth {
GeminiCliRequestAuth {
project_id: Some("project-123".to_string()),
session_id: Some("session-123".to_string()),
}
}
#[test]
fn wraps_generate_content_body_in_gemini_cli_v1internal_envelope() {
let request_body = json!({
"model": "client-model",
"contents": [
{"role": "user", "parts": [{"text": "hello"}]}
],
"stream": true,
"generationConfig": {"temperature": 0.2}
});
assert_eq!(
classify_gemini_cli_v1internal_request_body(&request_body),
Ok(())
);
assert_eq!(
build_gemini_cli_v1internal_request(
&sample_auth(),
"trace-123",
"gemini-2.5-pro",
&request_body,
),
GeminiCliRequestEnvelopeSupport::Supported(json!({
"model": "gemini-2.5-pro",
"project": "project-123",
"user_prompt_id": "trace-123",
"request": {
"contents": [
{"role": "user", "parts": [{"text": "hello"}]}
],
"generationConfig": {"temperature": 0.2},
"session_id": "session-123"
}
}))
);
}
#[test]
fn preserves_existing_v1internal_request_shape_without_antigravity_fields() {
let request_body = json!({
"model": "old-model",
"project": "project-from-body",
"user_prompt_id": "prompt-from-body",
"request": {
"model": "nested-client-model",
"contents": [
{"role": "user", "parts": [{"text": "hello"}]}
],
"stream": false,
"labels": {"source": "test"}
},
"userAgent": "antigravity",
"requestType": "agent"
});
assert_eq!(
build_gemini_cli_v1internal_request(
&sample_auth(),
"trace-123",
"gemini-2.5-pro",
&request_body,
),
GeminiCliRequestEnvelopeSupport::Supported(json!({
"model": "gemini-2.5-pro",
"project": "project-from-body",
"user_prompt_id": "prompt-from-body",
"request": {
"contents": [
{"role": "user", "parts": [{"text": "hello"}]}
],
"labels": {"source": "test"},
"session_id": "session-123"
}
}))
);
}
#[test]
fn omits_optional_project_and_session_when_metadata_is_absent() {
let request_body = json!({
"contents": [
{"role": "user", "parts": [{"text": "hello"}]}
]
});
assert_eq!(
build_gemini_cli_v1internal_request(
&GeminiCliRequestAuth::default(),
"trace-123",
"gemini-2.5-pro",
&request_body,
),
GeminiCliRequestEnvelopeSupport::Supported(json!({
"model": "gemini-2.5-pro",
"user_prompt_id": "trace-123",
"request": {
"contents": [
{"role": "user", "parts": [{"text": "hello"}]}
]
}
}))
);
}
}
@@ -0,0 +1,120 @@
use std::collections::BTreeMap;
use url::form_urlencoded;
pub const GEMINI_CLI_USER_AGENT: &str = "GeminiCLI/0.1.5 (Windows; AMD64)";
pub const GEMINI_CLI_V1INTERNAL_PATH_TEMPLATE: &str = "/v1internal:{action}";
pub const GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH: &str = "/v1internal:retrieveUserQuota";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum GeminiCliRequestUrlAction {
GenerateContent,
StreamGenerateContent,
}
impl GeminiCliRequestUrlAction {
fn as_str(self) -> &'static str {
match self {
Self::GenerateContent => "generateContent",
Self::StreamGenerateContent => "streamGenerateContent",
}
}
fn is_stream(self) -> bool {
matches!(self, Self::StreamGenerateContent)
}
}
pub fn build_gemini_cli_v1internal_url(
base_url: &str,
action: GeminiCliRequestUrlAction,
query: Option<&BTreeMap<String, String>>,
) -> Option<String> {
let trimmed_base = base_url.trim();
if trimmed_base.is_empty() {
return None;
}
let path = GEMINI_CLI_V1INTERNAL_PATH_TEMPLATE.replace("{action}", action.as_str());
let mut url = format!("{}{}", trimmed_base.trim_end_matches('/'), path);
let mut params = BTreeMap::new();
if let Some(query) = query {
for (key, value) in query {
let key = key.trim();
let value = value.trim();
if key.is_empty()
|| value.is_empty()
|| key.eq_ignore_ascii_case("beta")
|| key.eq_ignore_ascii_case("key")
{
continue;
}
params.insert(key.to_string(), value.to_string());
}
}
if action.is_stream() {
params
.entry(String::from("alt"))
.or_insert_with(|| String::from("sse"));
}
if !params.is_empty() {
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in params {
serializer.append_pair(key.as_str(), value.as_str());
}
let query_string = serializer.finish();
if !query_string.is_empty() {
url.push('?');
url.push_str(&query_string);
}
}
Some(url)
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use super::{
build_gemini_cli_v1internal_url, GeminiCliRequestUrlAction,
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH,
};
#[test]
fn builds_gemini_cli_stream_url_with_alt_sse() {
let query = BTreeMap::from([
("foo".to_string(), "bar".to_string()),
("key".to_string(), "blocked".to_string()),
]);
assert_eq!(
build_gemini_cli_v1internal_url(
"https://cloudcode-pa.googleapis.com/",
GeminiCliRequestUrlAction::StreamGenerateContent,
Some(&query),
)
.as_deref(),
Some("https://cloudcode-pa.googleapis.com/v1internal:streamGenerateContent?alt=sse&foo=bar")
);
}
#[test]
fn builds_gemini_cli_sync_url_without_stream_query() {
assert_eq!(
build_gemini_cli_v1internal_url(
"https://cloudcode-pa.googleapis.com",
GeminiCliRequestUrlAction::GenerateContent,
None,
)
.as_deref(),
Some("https://cloudcode-pa.googleapis.com/v1internal:generateContent")
);
assert_eq!(
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH,
"/v1internal:retrieveUserQuota"
);
}
}
@@ -0,0 +1,283 @@
use std::collections::BTreeMap;
use serde_json::{json, Value};
use crate::auth::{build_passthrough_headers_with_auth, resolve_local_gemini_auth};
use crate::policy::local_gemini_transport_unsupported_reason_with_network;
use crate::rules::{
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
body_rules_have_enabled_rules,
};
use crate::snapshot::GatewayProviderTransportSnapshot;
use crate::url::build_gemini_files_passthrough_url;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum GeminiFilesRequestBodyError {
BodyRulesUnsupportedForBinaryUpload,
BodyRulesApplyFailed,
}
#[derive(Debug, Clone, PartialEq)]
pub struct GeminiFilesRequestBodyParts {
pub provider_request_body: Option<Value>,
pub provider_request_body_base64: Option<String>,
}
#[derive(Debug, Clone, Copy)]
pub struct GeminiFilesHeadersInput<'a> {
pub headers: &'a http::HeaderMap,
pub auth_header: &'a str,
pub auth_value: &'a str,
pub header_rules: Option<&'a Value>,
pub provider_request_body: Option<&'a Value>,
pub provider_request_body_base64: Option<&'a str>,
pub original_request_body_json: &'a Value,
pub original_body_is_empty: bool,
}
pub fn gemini_files_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<&'static str> {
local_gemini_transport_unsupported_reason_with_network(transport, api_format)
}
pub fn resolve_gemini_files_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
resolve_local_gemini_auth(transport)
}
pub fn build_gemini_files_upstream_url(
transport: &GatewayProviderTransportSnapshot,
request_path: &str,
request_query: Option<&str>,
) -> Option<String> {
let custom_path = transport
.endpoint
.custom_path
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty());
let passthrough_path = custom_path.unwrap_or(request_path);
build_gemini_files_passthrough_url(
&transport.endpoint.base_url,
passthrough_path,
request_query,
)
}
pub fn build_gemini_files_request_body(
body_json: &Value,
body_base64: Option<&str>,
body_is_empty: bool,
is_upload: bool,
body_rules: Option<&Value>,
request_headers: Option<&http::HeaderMap>,
) -> Result<GeminiFilesRequestBodyParts, GeminiFilesRequestBodyError> {
let mut provider_request_body = if is_upload && !body_is_empty && body_base64.is_none() {
Some(body_json.clone())
} else {
None
};
let provider_request_body_base64 = if is_upload {
body_base64
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
} else {
None
};
if provider_request_body_base64.is_some() && body_rules_have_enabled_rules(body_rules) {
return Err(GeminiFilesRequestBodyError::BodyRulesUnsupportedForBinaryUpload);
}
if let Some(body) = provider_request_body.as_mut() {
if !apply_local_body_rules_with_request_headers(
body,
body_rules,
Some(body_json),
request_headers,
) {
return Err(GeminiFilesRequestBodyError::BodyRulesApplyFailed);
}
}
Ok(GeminiFilesRequestBodyParts {
provider_request_body,
provider_request_body_base64,
})
}
pub fn build_gemini_files_headers(
input: GeminiFilesHeadersInput<'_>,
) -> Option<BTreeMap<String, String>> {
let mut provider_request_headers = build_passthrough_headers_with_auth(
input.headers,
input.auth_header,
input.auth_value,
&BTreeMap::new(),
);
let null_original_request_body = Value::Null;
let base64_original_request_body = input
.provider_request_body_base64
.map(|body_bytes_b64| json!({ "body_bytes_b64": body_bytes_b64 }));
let original_request_body = base64_original_request_body
.as_ref()
.or_else(|| (!input.original_body_is_empty).then_some(input.original_request_body_json))
.unwrap_or(&null_original_request_body);
if !apply_local_header_rules_with_request_headers(
&mut provider_request_headers,
input.header_rules,
&[input.auth_header, "content-type"],
input.provider_request_body.unwrap_or(original_request_body),
Some(original_request_body),
Some(input.headers),
) {
return None;
}
Some(provider_request_headers)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Provider".to_string(),
provider_type: "gemini".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "gemini:files".to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://generativelanguage.googleapis.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn builds_gemini_files_url_from_request_path_and_strips_client_key() {
let url = build_gemini_files_upstream_url(
&sample_transport(),
"/v1beta/files/demo",
Some("key=client-key&alt=json"),
)
.expect("url should build");
assert_eq!(
url,
"https://generativelanguage.googleapis.com/v1beta/files/demo?alt=json"
);
}
#[test]
fn builds_json_upload_body_and_applies_rules() {
let body = build_gemini_files_request_body(
&json!({"display_name": "demo"}),
None,
false,
true,
Some(&json!([
{"action":"set","path":"metadata.source","value":"local"}
])),
None,
)
.expect("body should build");
assert_eq!(
body.provider_request_body
.as_ref()
.and_then(|value| value.pointer("/metadata/source")),
Some(&json!("local"))
);
assert!(body.provider_request_body_base64.is_none());
}
#[test]
fn rejects_body_rules_for_binary_upload() {
assert_eq!(
build_gemini_files_request_body(
&json!({}),
Some("YWJj"),
false,
true,
Some(&json!([{"action":"set","path":"x","value":1}])),
None,
),
Err(GeminiFilesRequestBodyError::BodyRulesUnsupportedForBinaryUpload)
);
}
#[test]
fn builds_headers_using_binary_body_as_original_context() {
let headers = build_gemini_files_headers(GeminiFilesHeadersInput {
headers: &http::HeaderMap::new(),
auth_header: "x-goog-api-key",
auth_value: "secret",
header_rules: Some(&json!([
{"action":"set","key":"x-upload-mode","value":"binary"}
])),
provider_request_body: None,
provider_request_body_base64: Some("YWJj"),
original_request_body_json: &json!({}),
original_body_is_empty: false,
})
.expect("headers should build");
assert_eq!(
headers.get("x-goog-api-key").map(String::as_str),
Some("secret")
);
assert_eq!(
headers.get("x-upload-mode").map(String::as_str),
Some("binary")
);
}
}
@@ -0,0 +1,497 @@
use std::collections::BTreeMap;
use aether_oauth::provider::providers::{
GenericProviderOAuthAdapter, GENERIC_PROVIDER_OAUTH_TEMPLATES,
};
use aether_oauth::provider::{ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthTokenSet};
use async_trait::async_trait;
use serde_json::Value;
use super::oauth_refresh::{
oauth_error_to_local_refresh_error, provider_oauth_transport_context_from_snapshot,
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter, LocalOAuthRefreshError,
LocalResolvedOAuthRequestAuth, ProviderOAuthLocalHttpExecutor,
};
use super::snapshot::GatewayProviderTransportSnapshot;
const AUTH_HEADER_NAME: &str = "authorization";
const OAUTH_REFRESH_SKEW_SECS: u64 = 120;
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
pub fn supports_local_generic_oauth_request_auth_resolution(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
&& generic_provider_type(transport.provider.provider_type.as_str()).is_some()
}
#[derive(Debug, Clone, Default)]
pub struct GenericOAuthRefreshAdapter {
token_url_overrides: BTreeMap<String, String>,
}
impl GenericOAuthRefreshAdapter {
pub fn with_token_url_for_tests(
mut self,
provider_type: &str,
token_url: impl Into<String>,
) -> Self {
self.token_url_overrides
.insert(provider_type.trim().to_ascii_lowercase(), token_url.into());
self
}
fn adapter_for_provider_type(
&self,
provider_type: &'static str,
) -> Option<GenericProviderOAuthAdapter> {
let adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?;
if let Some(token_url) = self.token_url_overrides.get(provider_type) {
return Some(adapter.with_token_url_override(token_url.clone()));
}
Some(adapter)
}
fn auth_config_from_transport(transport: &GatewayProviderTransportSnapshot) -> Option<Value> {
transport
.key
.decrypted_auth_config
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| serde_json::from_str::<Value>(value).ok())
}
fn auth_config_from_entry(
transport: &GatewayProviderTransportSnapshot,
entry: &CachedOAuthEntry,
) -> Option<Value> {
entry
.metadata
.as_ref()
.filter(|_| {
entry
.provider_type
.eq_ignore_ascii_case(transport.provider.provider_type.as_str())
})
.cloned()
}
fn auth_config_updated_at(auth_config: &Value) -> Option<u64> {
auth_config
.as_object()
.and_then(|object| object.get("updated_at"))
.and_then(|value| parse_u64_value(Some(value)))
}
fn base_auth_config(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Option<Value> {
let cached = entry.and_then(|cached| Self::auth_config_from_entry(transport, cached));
let transport_auth = Self::auth_config_from_transport(transport);
match (cached, transport_auth) {
(Some(cached), Some(transport_auth)) => {
let cached_updated_at = Self::auth_config_updated_at(&cached);
let transport_updated_at = Self::auth_config_updated_at(&transport_auth);
if transport_updated_at > cached_updated_at {
Some(transport_auth)
} else {
Some(cached)
}
}
(Some(cached), None) => Some(cached),
(None, Some(transport_auth)) => Some(transport_auth),
(None, None) => None,
}
}
fn resolve_direct_header(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<LocalResolvedOAuthRequestAuth> {
if !supports_local_generic_oauth_request_auth_resolution(transport) {
return None;
}
if let Some(value) =
auth_config_authorization_header(transport.key.decrypted_auth_config.as_deref())
{
return Some(LocalResolvedOAuthRequestAuth::Header {
name: AUTH_HEADER_NAME.to_string(),
value,
});
}
let secret = transport.key.decrypted_api_key.trim();
if secret.is_empty() || secret == PLACEHOLDER_API_KEY {
return None;
}
let auth_config = Self::auth_config_from_transport(transport);
let refreshable = auth_config
.as_ref()
.and_then(refresh_token_from_auth_config)
.is_some();
if refreshable && auth_config_expires_soon(auth_config.as_ref()) {
return None;
}
Some(LocalResolvedOAuthRequestAuth::Header {
name: AUTH_HEADER_NAME.to_string(),
value: format!("Bearer {secret}"),
})
}
fn build_cached_entry(
provider_type: &'static str,
refreshed: ProviderOAuthTokenSet,
) -> CachedOAuthEntry {
CachedOAuthEntry {
provider_type: provider_type.to_string(),
auth_header_name: AUTH_HEADER_NAME.to_string(),
auth_header_value: refreshed.token_set.bearer_header_value(),
expires_at_unix_secs: refreshed.token_set.expires_at_unix_secs,
metadata: Some(refreshed.auth_config),
}
}
}
#[async_trait]
impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
fn provider_type(&self) -> &'static str {
"generic_oauth"
}
fn supports(&self, transport: &GatewayProviderTransportSnapshot) -> bool {
supports_local_generic_oauth_request_auth_resolution(transport)
}
fn resolve_cached(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: &CachedOAuthEntry,
) -> Option<LocalResolvedOAuthRequestAuth> {
if !entry
.provider_type
.eq_ignore_ascii_case(transport.provider.provider_type.as_str())
{
return None;
}
if let Some(value) =
auth_config_authorization_header(transport.key.decrypted_auth_config.as_deref())
{
return Some(LocalResolvedOAuthRequestAuth::Header {
name: AUTH_HEADER_NAME.to_string(),
value,
});
}
if expires_at_requires_refresh(entry.expires_at_unix_secs) {
return None;
}
let name = entry.auth_header_name.trim();
let value = entry.auth_header_value.trim();
if name.is_empty() || value.is_empty() {
return None;
}
Some(LocalResolvedOAuthRequestAuth::Header {
name: name.to_ascii_lowercase(),
value: value.to_string(),
})
}
fn resolve_without_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<LocalResolvedOAuthRequestAuth> {
self.resolve_direct_header(transport)
}
fn should_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> bool {
if !supports_local_generic_oauth_request_auth_resolution(transport) {
return false;
}
if entry
.and_then(|cached| self.resolve_cached(transport, cached))
.is_some()
|| self.resolve_direct_header(transport).is_some()
{
return false;
}
self.base_auth_config(transport, entry)
.as_ref()
.and_then(refresh_token_from_auth_config)
.is_some()
}
async fn refresh(
&self,
executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
let Some(provider_type) = generic_provider_type(transport.provider.provider_type.as_str())
else {
return Ok(None);
};
let Some(auth_config) = self.base_auth_config(transport, entry) else {
return Ok(None);
};
let Some(refresh_token) = refresh_token_from_auth_config(&auth_config) else {
tracing::warn!(
key_id = %transport.key.id,
provider_id = %transport.provider.id,
provider_type,
"gateway generic oauth refresh skipped because auth_config has no refresh_token"
);
return Ok(None);
};
let Some(adapter) = self.adapter_for_provider_type(provider_type) else {
return Ok(None);
};
tracing::info!(
key_id = %transport.key.id,
provider_id = %transport.provider.id,
endpoint_id = %transport.endpoint.id,
provider_type,
request_refresh_token_len = refresh_token.len(),
"gateway generic oauth refresh delegated to provider oauth adapter"
);
let oauth_executor =
ProviderOAuthLocalHttpExecutor::new(provider_type, transport, executor);
let ctx = provider_oauth_transport_context_from_snapshot(transport);
let account = ProviderOAuthAccount {
provider_type: provider_type.to_string(),
access_token: current_access_token(transport, entry).unwrap_or_default(),
expires_at_unix_secs: auth_config_expires_at(&auth_config),
auth_config,
identity: BTreeMap::new(),
};
let refreshed = adapter
.refresh(&oauth_executor, &ctx, &account)
.await
.map_err(|error| oauth_error_to_local_refresh_error(provider_type, error))?;
tracing::info!(
key_id = %transport.key.id,
provider_id = %transport.provider.id,
endpoint_id = %transport.endpoint.id,
provider_type,
expires_at_unix_secs = ?refreshed.token_set.expires_at_unix_secs,
response_has_refresh_token = refreshed.token_set.refresh_token.is_some(),
"gateway generic oauth refresh succeeded"
);
Ok(Some(Self::build_cached_entry(provider_type, refreshed)))
}
}
fn generic_provider_type(provider_type: &str) -> Option<&'static str> {
let normalized = provider_type.trim();
GENERIC_PROVIDER_OAUTH_TEMPLATES
.iter()
.find(|template| normalized.eq_ignore_ascii_case(template.provider_type))
.map(|template| template.provider_type)
}
fn refresh_token_from_auth_config(auth_config: &Value) -> Option<String> {
auth_config
.as_object()
.and_then(|object| object.get("refresh_token"))
.and_then(non_empty_string)
}
fn auth_config_expires_at(auth_config: &Value) -> Option<u64> {
auth_config
.as_object()
.and_then(|object| object.get("expires_at"))
.and_then(|value| parse_u64_value(Some(value)))
}
fn auth_config_expires_soon(auth_config: Option<&Value>) -> bool {
expires_at_requires_refresh(auth_config.and_then(auth_config_expires_at))
}
fn expires_at_requires_refresh(expires_at_unix_secs: Option<u64>) -> bool {
expires_at_unix_secs
.map(|expires_at_unix_secs| {
aether_oauth::core::current_unix_secs()
>= expires_at_unix_secs.saturating_sub(OAUTH_REFRESH_SKEW_SECS)
})
.unwrap_or(false)
}
fn parse_u64_value(value: Option<&Value>) -> Option<u64> {
match value? {
Value::Number(number) => number.as_u64(),
Value::String(string) => string.trim().parse::<u64>().ok(),
_ => None,
}
}
fn non_empty_string(value: &Value) -> Option<String> {
value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn auth_config_authorization_header(raw_auth_config: Option<&str>) -> Option<String> {
let mut headers = BTreeMap::new();
crate::auth_config::apply_local_auth_config_header_overrides(&mut headers, raw_auth_config);
headers
.remove(AUTH_HEADER_NAME)
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
fn current_access_token(
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Option<String> {
entry
.and_then(|entry| {
entry
.auth_header_value
.trim()
.strip_prefix("Bearer ")
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
.or_else(|| {
let secret = transport.key.decrypted_api_key.trim();
(!secret.is_empty() && secret != PLACEHOLDER_API_KEY).then(|| secret.to_string())
})
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::super::oauth_refresh::{
CachedOAuthEntry, LocalOAuthRefreshAdapter, LocalResolvedOAuthRequestAuth,
};
use super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::GenericOAuthRefreshAdapter;
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Codex".to_string(),
provider_type: "codex".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "openai:responses".to_string(),
api_family: Some("openai".to_string()),
endpoint_kind: Some("responses".to_string()),
is_active: true,
base_url: "https://chatgpt.com/backend-api/codex".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "OAuth headers".to_string(),
auth_type: "oauth".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: Some(
json!({
"provider_type": "codex",
"access_token_import_temporary": true,
"headers": {
"Authorization": "Bearer imported-session",
"Host": "blocked.example"
}
})
.to_string(),
),
},
}
}
#[test]
fn resolves_imported_authorization_header_without_api_key_secret() {
let adapter = GenericOAuthRefreshAdapter::default();
let auth = adapter
.resolve_without_refresh(&sample_transport())
.expect("auth_config authorization header should resolve");
assert_eq!(
auth,
LocalResolvedOAuthRequestAuth::Header {
name: "authorization".to_string(),
value: "Bearer imported-session".to_string(),
}
);
}
#[test]
fn auth_config_authorization_header_overrides_cached_oauth_entry() {
let adapter = GenericOAuthRefreshAdapter::default();
let entry = CachedOAuthEntry {
provider_type: "codex".to_string(),
auth_header_name: "authorization".to_string(),
auth_header_value: "Bearer refreshed-access-token".to_string(),
expires_at_unix_secs: Some(u64::MAX),
metadata: None,
};
let auth = adapter
.resolve_cached(&sample_transport(), &entry)
.expect("auth_config authorization header should override cache");
assert_eq!(
auth,
LocalResolvedOAuthRequestAuth::Header {
name: "authorization".to_string(),
value: "Bearer imported-session".to_string(),
}
);
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,304 @@
use std::collections::BTreeMap;
use aether_contracts::USAGE_SERVER_NOW_UNIX_MS_HEADER;
pub fn should_skip_request_header(name: &str) -> bool {
let normalized = name.to_ascii_lowercase();
matches!(
normalized.as_str(),
"connection"
| "keep-alive"
| "proxy-authenticate"
| "proxy-authorization"
| "proxy-connection"
| "te"
| "trailer"
| "transfer-encoding"
| "upgrade"
| "x-aether-execution-path"
| "x-aether-dependency-reason"
| "x-aether-execution-loop-guard"
| "x-aether-control-execute-fallback"
| "x-aether-rate-limit-preflight"
| USAGE_SERVER_NOW_UNIX_MS_HEADER
)
}
pub fn should_skip_upstream_passthrough_header(name: &str) -> bool {
let lower = name.to_ascii_lowercase();
// Anthropic SDK (stainless) client metadata and Anthropic-specific headers
// (anthropic-version / anthropic-beta / anthropic-dangerous-direct-browser-access / ...).
// These are only meaningful when the upstream is Anthropic. The Claude Code
// adapter re-reads what it needs directly from the original HeaderMap and
// injects its own values *after* this filter runs, so stripping both prefixes
// here prevents leakage to any other upstream (OpenAI/Gemini/Codex/...).
if lower.starts_with("x-stainless-") || lower.starts_with("anthropic-") {
return true;
}
matches!(
lower.as_str(),
"authorization"
| "x-api-key"
| "x-goog-api-key"
| "host"
| "content-length"
| "transfer-encoding"
| "connection"
| "content-encoding"
| "x-real-ip"
| "x-real-proto"
| "x-forwarded-for"
| "x-forwarded-proto"
| "x-forwarded-scheme"
| "x-forwarded-host"
| "x-forwarded-port"
// Claude CLI client identifier; re-injected by the Claude Code adapter
// when the upstream is Anthropic, filtered for everybody else.
| "x-app"
) || should_skip_request_header(name)
}
pub(crate) fn should_skip_upstream_complete_passthrough_header(name: &str) -> bool {
let lower = name.to_ascii_lowercase();
matches!(
lower.as_str(),
"authorization"
| "x-api-key"
| "x-goog-api-key"
| "host"
| "content-length"
| "transfer-encoding"
| "connection"
| "content-encoding"
| "x-real-ip"
| "x-real-proto"
| "x-forwarded-for"
| "x-forwarded-proto"
| "x-forwarded-scheme"
| "x-forwarded-host"
| "x-forwarded-port"
) || should_skip_request_header(name)
}
pub fn normalize_upstream_accept_encoding(value: &str) -> Option<String> {
let mut accepted = Vec::new();
let mut wildcard_allowed = false;
let mut gzip_disabled = false;
let mut deflate_disabled = false;
let mut identity_disabled = false;
for item in value.split(',') {
let Some((token, normalized_item, enabled)) = parse_accept_encoding_item(item) else {
continue;
};
match token.as_str() {
"gzip" if enabled => accepted.push(normalized_item),
"gzip" => gzip_disabled = true,
"deflate" if enabled => accepted.push(normalized_item),
"deflate" => deflate_disabled = true,
"identity" if enabled => accepted.push(normalized_item),
"identity" => identity_disabled = true,
"*" if enabled => wildcard_allowed = true,
_ => {}
}
}
if !accepted.is_empty() {
return Some(accepted.join(", "));
}
if wildcard_allowed && !gzip_disabled {
Some("gzip".to_string())
} else if wildcard_allowed && !deflate_disabled {
Some("deflate".to_string())
} else if wildcard_allowed && !identity_disabled {
Some("identity".to_string())
} else {
None
}
}
fn parse_accept_encoding_item(raw_item: &str) -> Option<(String, String, bool)> {
let mut parts = raw_item.trim().split(';');
let token = parts.next()?.trim().to_ascii_lowercase();
if token.is_empty() {
return None;
}
let mut enabled = true;
let mut normalized = token.clone();
for raw_param in parts {
let param = raw_param.trim();
if param.is_empty() {
continue;
}
let Some((name, value)) = param.split_once('=') else {
continue;
};
if name.trim().eq_ignore_ascii_case("q") {
let value = value.trim();
if q_value_is_zero(value) {
enabled = false;
continue;
}
normalized.push_str(";q=");
normalized.push_str(value);
}
}
Some((token, normalized, enabled))
}
fn q_value_is_zero(value: &str) -> bool {
value
.trim_matches('"')
.parse::<f32>()
.is_ok_and(|q| q <= 0.0)
}
pub fn force_identity_accept_encoding(headers: &mut BTreeMap<String, String>) {
if let Some(existing_key) = headers
.keys()
.find(|key| key.eq_ignore_ascii_case("accept-encoding"))
.cloned()
{
headers.remove(&existing_key);
}
headers.insert("accept-encoding".to_string(), "identity".to_string());
}
#[cfg(test)]
mod tests {
use super::{
force_identity_accept_encoding, normalize_upstream_accept_encoding,
should_skip_request_header, should_skip_upstream_complete_passthrough_header,
should_skip_upstream_passthrough_header,
};
use aether_contracts::USAGE_SERVER_NOW_UNIX_MS_HEADER;
use std::collections::BTreeMap;
#[test]
fn strips_all_stainless_headers() {
let stainless = [
"x-stainless-arch",
"x-stainless-lang",
"x-stainless-os",
"x-stainless-package-version",
"x-stainless-retry-count",
"x-stainless-runtime",
"x-stainless-runtime-version",
"x-stainless-timeout",
"x-stainless-helper-method",
"X-Stainless-Arch",
"X-STAINLESS-FUTURE-HEADER",
];
for h in stainless {
assert!(
should_skip_upstream_passthrough_header(h),
"should skip {h}"
);
}
}
#[test]
fn accept_encoding_is_not_classified_as_hop_by_hop_passthrough_skip() {
assert!(!should_skip_upstream_passthrough_header("accept-encoding"));
assert!(!should_skip_upstream_complete_passthrough_header(
"accept-encoding"
));
}
#[test]
fn normalizes_accept_encoding_to_supported_upstream_codecs() {
assert_eq!(
normalize_upstream_accept_encoding("gzip, br").as_deref(),
Some("gzip")
);
assert_eq!(
normalize_upstream_accept_encoding("br, deflate").as_deref(),
Some("deflate")
);
assert_eq!(
normalize_upstream_accept_encoding("identity").as_deref(),
Some("identity")
);
assert_eq!(
normalize_upstream_accept_encoding("gzip;q=0.5, br").as_deref(),
Some("gzip;q=0.5")
);
assert_eq!(
normalize_upstream_accept_encoding("gzip;q=0, br, deflate").as_deref(),
Some("deflate")
);
assert_eq!(
normalize_upstream_accept_encoding("gzip;q=0, br").as_deref(),
None
);
assert_eq!(
normalize_upstream_accept_encoding("*").as_deref(),
Some("gzip")
);
assert_eq!(normalize_upstream_accept_encoding("br"), None);
}
#[test]
fn force_identity_accept_encoding_replaces_existing_casing() {
let mut headers = BTreeMap::from([("Accept-Encoding".to_string(), "gzip".to_string())]);
force_identity_accept_encoding(&mut headers);
assert_eq!(
headers.get("accept-encoding").map(String::as_str),
Some("identity")
);
assert!(!headers.contains_key("Accept-Encoding"));
}
#[test]
fn strips_anthropic_and_claude_cli_identity_headers() {
let anthropic = [
"anthropic-version",
"anthropic-beta",
"anthropic-dangerous-direct-browser-access",
"Anthropic-Version",
"ANTHROPIC-FUTURE-HEADER",
"x-app",
"X-App",
];
for h in anthropic {
assert!(
should_skip_upstream_passthrough_header(h),
"should skip {h}"
);
}
}
#[test]
fn strips_usage_server_time_header_from_provider_requests() {
for h in [
USAGE_SERVER_NOW_UNIX_MS_HEADER,
"X-Aether-Server-Now-Unix-Ms",
] {
assert!(should_skip_request_header(h), "should skip {h}");
assert!(
should_skip_upstream_passthrough_header(h),
"should skip passthrough {h}"
);
assert!(
should_skip_upstream_complete_passthrough_header(h),
"should skip complete passthrough {h}"
);
}
}
#[test]
fn allows_normal_headers_through() {
let allowed = ["user-agent", "accept", "content-type", "x-custom-header"];
for h in allowed {
assert!(
!should_skip_upstream_passthrough_header(h),
"should allow {h}"
);
}
}
}
@@ -0,0 +1,353 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::credentials::{generate_machine_id, KiroAuthConfig};
pub const PROVIDER_TYPE: &str = "kiro";
pub const KIRO_AUTH_HEADER: &str = "authorization";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct KiroBearerAuth {
pub name: &'static str,
pub value: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct KiroRequestAuth {
pub name: &'static str,
pub value: String,
pub auth_config: KiroAuthConfig,
pub machine_id: String,
}
pub fn is_kiro_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(PROVIDER_TYPE)
}
pub fn is_kiro_claude_messages_transport(
transport: &GatewayProviderTransportSnapshot,
provider_api_format: &str,
) -> bool {
is_kiro_provider_transport(transport)
&& aether_ai_formats::api_format_alias_matches(provider_api_format, "claude:messages")
}
pub fn build_kiro_request_auth_from_config(
auth_config: KiroAuthConfig,
fallback_secret: Option<&str>,
) -> Option<KiroRequestAuth> {
let cached_token_needs_refresh = auth_config.cached_access_token_requires_refresh(120);
let fallback_secret = fallback_secret
.map(str::trim)
.filter(|value| !value.is_empty() && *value != "__placeholder__");
if cached_token_needs_refresh && auth_config.can_refresh_access_token() {
return None;
}
let token = auth_config
.cached_access_token()
.filter(|_| !cached_token_needs_refresh)
.or(fallback_secret)?;
let machine_id = generate_machine_id(&auth_config, Some(token))?;
Some(KiroRequestAuth {
name: KIRO_AUTH_HEADER,
value: format!("Bearer {token}"),
auth_config,
machine_id,
})
}
pub fn resolve_local_kiro_bearer_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<KiroBearerAuth> {
if !is_kiro_provider_transport(transport) {
return None;
}
if transport.key.decrypted_auth_config.is_some() {
return None;
}
if !kiro_auth_type_supported(transport.key.auth_type.as_str()) {
return None;
}
let secret = transport.key.decrypted_api_key.trim();
if secret.is_empty() {
return None;
}
Some(KiroBearerAuth {
name: KIRO_AUTH_HEADER,
value: format!("Bearer {secret}"),
})
}
pub fn supports_local_kiro_auth_prerequisites(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
resolve_local_kiro_bearer_auth(transport).is_some()
}
pub fn resolve_local_kiro_request_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<KiroRequestAuth> {
if !is_kiro_provider_transport(transport) {
return None;
}
if !kiro_auth_type_supported(transport.key.auth_type.as_str()) {
return None;
}
let auth_config = KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
.unwrap_or(KiroAuthConfig {
auth_method: None,
refresh_token: None,
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: None,
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
});
let fallback_secret = transport
.key
.decrypted_api_key
.trim()
.strip_prefix("__placeholder__")
.map(|_| "")
.unwrap_or(transport.key.decrypted_api_key.trim());
build_kiro_request_auth_from_config(auth_config, Some(fallback_secret))
}
pub fn supports_local_kiro_request_auth_resolution(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
resolve_local_kiro_request_auth(transport).is_some()
|| KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
.is_some_and(|auth_config| {
is_kiro_provider_transport(transport)
&& kiro_auth_type_supported(transport.key.auth_type.as_str())
&& auth_config.can_refresh_access_token()
})
}
fn kiro_auth_type_supported(auth_type: &str) -> bool {
matches!(
auth_type.trim().to_ascii_lowercase().as_str(),
"bearer" | "oauth"
)
}
#[cfg(test)]
mod tests {
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
resolve_local_kiro_bearer_auth, resolve_local_kiro_request_auth,
supports_local_kiro_auth_prerequisites, supports_local_kiro_request_auth_resolution,
KIRO_AUTH_HEADER,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Kiro".to_string(),
provider_type: "kiro".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "claude:messages".to_string(),
api_family: Some("claude".to_string()),
endpoint_kind: Some("cli".to_string()),
is_active: true,
base_url: "https://kiro.example".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: Some(vec!["claude:messages".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "upstream-key".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn resolves_bearer_auth_for_known_kiro_subset() {
let auth = resolve_local_kiro_bearer_auth(&sample_transport())
.expect("kiro bearer auth should resolve");
assert_eq!(auth.name, KIRO_AUTH_HEADER);
assert_eq!(auth.value, "Bearer upstream-key");
assert!(supports_local_kiro_auth_prerequisites(&sample_transport()));
}
#[test]
fn rejects_auth_config_subset() {
let mut transport = sample_transport();
transport.key.decrypted_auth_config = Some("{\"mode\":\"custom\"}".to_string());
assert!(resolve_local_kiro_bearer_auth(&transport).is_none());
assert!(!supports_local_kiro_auth_prerequisites(&transport));
}
#[test]
fn rejects_non_bearer_subset() {
let mut transport = sample_transport();
transport.key.auth_type = "api_key".to_string();
assert!(resolve_local_kiro_bearer_auth(&transport).is_none());
}
#[test]
fn resolves_request_auth_when_legacy_oauth_auth_type_is_used() {
let mut transport = sample_transport();
transport.key.auth_type = "oauth".to_string();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"access_token":"cached-token",
"expires_at":4102444800,
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
"machine_id":"123e4567-e89b-12d3-a456-426614174000",
"api_region":"us-west-2"
}"#
.to_string(),
);
let auth = resolve_local_kiro_request_auth(&transport)
.expect("request auth should resolve from legacy oauth auth_type");
assert_eq!(auth.value, "Bearer cached-token");
assert!(supports_local_kiro_request_auth_resolution(&transport));
}
#[test]
fn resolves_request_auth_from_cached_access_token() {
let mut transport = sample_transport();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"access_token":"cached-token",
"expires_at":4102444800,
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
"machine_id":"123e4567-e89b-12d3-a456-426614174000",
"api_region":"us-west-2"
}"#
.to_string(),
);
let auth = resolve_local_kiro_request_auth(&transport)
.expect("request auth should resolve from cached token");
assert_eq!(auth.name, KIRO_AUTH_HEADER);
assert_eq!(auth.value, "Bearer cached-token");
assert_eq!(auth.auth_config.effective_api_region(), "us-west-2");
assert_eq!(
auth.machine_id,
"123e4567e89b12d3a456426614174000123e4567e89b12d3a456426614174000"
);
}
#[test]
fn skips_expired_cached_access_token_without_fallback_secret() {
let mut transport = sample_transport();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"access_token":"expired-token",
"expires_at": 1,
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
}"#
.to_string(),
);
assert!(resolve_local_kiro_request_auth(&transport).is_none());
}
#[test]
fn skips_refreshable_cached_access_token_without_expiry() {
let mut transport = sample_transport();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"access_token":"cached-token-without-expiry",
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
}"#
.to_string(),
);
assert!(resolve_local_kiro_request_auth(&transport).is_none());
assert!(supports_local_kiro_request_auth_resolution(&transport));
}
#[test]
fn refreshable_expired_cached_access_token_does_not_fallback_to_decrypted_api_key() {
let mut transport = sample_transport();
transport.key.decrypted_api_key = "stale-upstream-token".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"access_token":"expired-token",
"expires_at": 1,
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
}"#
.to_string(),
);
assert!(resolve_local_kiro_request_auth(&transport).is_none());
assert!(supports_local_kiro_request_auth_resolution(&transport));
}
#[test]
fn supports_refresh_only_resolution_without_cached_access_token() {
let mut transport = sample_transport();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
}"#
.to_string(),
);
assert!(resolve_local_kiro_request_auth(&transport).is_none());
assert!(supports_local_kiro_request_auth_resolution(&transport));
}
}
@@ -0,0 +1,714 @@
use std::collections::BTreeSet;
use serde_json::{json, Map, Value};
use tracing::warn;
use uuid::Uuid;
const SYSTEM_CHUNKED_POLICY: &str = "When the Write or Edit tool has content size limits, always comply silently. Never suggest bypassing these limits via alternative tools. Never ask the user whether to switch approaches. Complete all chunked operations without commentary.";
const WRITE_TOOL_DESCRIPTION_SUFFIX: &str = "- IMPORTANT: If the content to write exceeds 150 lines, you MUST only write the first 50 lines using this tool, then use `Edit` tool to append the remaining content in chunks of no more than 50 lines each. If needed, leave a unique placeholder to help append content. Do NOT attempt to write all content at once.";
const EDIT_TOOL_DESCRIPTION_SUFFIX: &str = "- IMPORTANT: If the `new_string` content exceeds 50 lines, you MUST split it into multiple Edit calls, each replacing no more than 50 lines at a time. If used to append content, leave a unique placeholder to help append content. On the final chunk, do NOT include the placeholder.";
pub fn convert_claude_messages_to_conversation_state(
request_body: &Value,
model: &str,
) -> Option<Value> {
let model_id = model.trim();
if model_id.is_empty() {
return None;
}
let messages = request_body.get("messages")?.as_array()?;
if messages.is_empty() {
return None;
}
let conversation_id = request_body
.get("metadata")
.and_then(Value::as_object)
.and_then(|metadata| {
metadata
.get("user_id")
.or_else(|| metadata.get("userId"))
.and_then(Value::as_str)
})
.and_then(extract_session_id)
.unwrap_or_else(|| Uuid::new_v4().to_string());
let agent_continuation_id = Uuid::new_v4().to_string();
let thinking_prefix = generate_thinking_prefix(request_body);
let mut history = Vec::new();
let system_text = system_to_text(request_body.get("system"));
if !system_text.is_empty() {
history.push(json!({
"userInputMessage": {
"content": format!("{system_text}\n{SYSTEM_CHUNKED_POLICY}"),
"modelId": model_id,
"origin": "AI_EDITOR"
}
}));
history.push(json!({
"assistantResponseMessage": {
"content": "I will follow these instructions."
}
}));
}
let last_is_assistant = messages
.last()
.and_then(Value::as_object)
.and_then(|message| message.get("role"))
.and_then(Value::as_str)
.is_some_and(|role| role == "assistant");
let history_end_index = if last_is_assistant {
messages.len()
} else {
messages.len().saturating_sub(1)
};
let mut user_buffer = Vec::new();
for message in &messages[..history_end_index] {
let Some(message) = message.as_object() else {
continue;
};
match message.get("role").and_then(Value::as_str) {
Some("user") => user_buffer.push(message),
Some("assistant") => {
if let Some(user_item) = flush_user_buffer(&mut user_buffer, model_id) {
history.push(user_item);
} else if history.is_empty()
|| history
.last()
.and_then(Value::as_object)
.is_some_and(|item| item.contains_key("assistantResponseMessage"))
{
history.push(json!({
"userInputMessage": {
"content": "Continue.",
"modelId": model_id,
"origin": "AI_EDITOR"
}
}));
}
if let Some(assistant_item) = convert_assistant_message(message) {
history.push(json!({"assistantResponseMessage": assistant_item}));
}
}
_ => {}
}
}
if let Some(tail_user) = flush_user_buffer(&mut user_buffer, model_id) {
history.push(tail_user);
history.push(json!({"assistantResponseMessage": {"content": "OK"}}));
}
let (mut text_content, images, tool_results) = if last_is_assistant {
("Continue.".to_string(), Vec::new(), Vec::new())
} else {
let last = messages.last()?.as_object()?;
if last.get("role").and_then(Value::as_str) != Some("user") {
return None;
}
process_message_content(last.get("content"))
};
let mut tools = convert_tools(request_body.get("tools"));
let mut history_tool_names = BTreeSet::new();
let mut history_tool_result_ids = BTreeSet::new();
let mut history_tool_use_ids = BTreeSet::new();
for item in &history {
let Some(item) = item.as_object() else {
continue;
};
if let Some(user_input) = item.get("userInputMessage").and_then(Value::as_object) {
if let Some(results) = user_input
.get("userInputMessageContext")
.and_then(Value::as_object)
.and_then(|ctx| ctx.get("toolResults"))
.and_then(Value::as_array)
{
for result in results {
if let Some(tool_use_id) = result
.as_object()
.and_then(|result| result.get("toolUseId"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
history_tool_result_ids.insert(tool_use_id.to_string());
}
}
}
}
if let Some(assistant) = item
.get("assistantResponseMessage")
.and_then(Value::as_object)
{
if let Some(tool_uses) = assistant.get("toolUses").and_then(Value::as_array) {
for tool_use in tool_uses {
let Some(tool_use) = tool_use.as_object() else {
continue;
};
if let Some(name) = tool_use
.get("name")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
history_tool_names.insert(name.to_string());
}
if let Some(tool_use_id) = tool_use
.get("toolUseId")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
history_tool_use_ids.insert(tool_use_id.to_string());
}
}
}
}
}
let existing_tool_names = tools
.iter()
.filter_map(|tool| {
tool.get("toolSpecification")
.and_then(Value::as_object)
.and_then(|spec| spec.get("name"))
.and_then(Value::as_str)
.map(|name| name.to_ascii_lowercase())
})
.collect::<BTreeSet<_>>();
for tool_name in history_tool_names {
if !existing_tool_names.contains(&tool_name.to_ascii_lowercase()) {
tools.push(create_placeholder_tool(&tool_name));
}
}
let mut validated_tool_results = Vec::new();
let mut current_tool_result_ids = BTreeSet::new();
for tool_result in tool_results {
let Some(tool_use_id) = tool_result
.get("toolUseId")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
continue;
};
if !history_tool_use_ids.contains(tool_use_id)
|| history_tool_result_ids.contains(tool_use_id)
{
continue;
}
current_tool_result_ids.insert(tool_use_id.to_string());
validated_tool_results.push(tool_result);
}
let orphaned_tool_use_ids = history_tool_use_ids
.difference(&history_tool_result_ids)
.filter(|tool_use_id| !current_tool_result_ids.contains(*tool_use_id))
.cloned()
.collect::<BTreeSet<_>>();
if !orphaned_tool_use_ids.is_empty() {
warn!(
"kiro: removing {} orphaned tool_use(s) from history",
orphaned_tool_use_ids.len()
);
for item in &mut history {
let Some(item) = item.as_object_mut() else {
continue;
};
let Some(assistant) = item
.get_mut("assistantResponseMessage")
.and_then(Value::as_object_mut)
else {
continue;
};
let Some(tool_uses) = assistant.get_mut("toolUses").and_then(Value::as_array_mut)
else {
continue;
};
tool_uses.retain(|tool_use| {
!tool_use
.get("toolUseId")
.and_then(Value::as_str)
.is_some_and(|tool_use_id| orphaned_tool_use_ids.contains(tool_use_id))
});
if tool_uses.is_empty() {
assistant.remove("toolUses");
}
}
}
let mut user_context = Map::new();
if !tools.is_empty() {
user_context.insert("tools".to_string(), Value::Array(tools));
}
if !validated_tool_results.is_empty() {
user_context.insert(
"toolResults".to_string(),
Value::Array(validated_tool_results),
);
}
if let Some(thinking_prefix) = thinking_prefix.as_deref() {
if !has_thinking_tags(&text_content) {
text_content = format!("{thinking_prefix}\n{text_content}");
}
}
let mut user_input = Map::new();
user_input.insert(
"userInputMessageContext".to_string(),
Value::Object(user_context),
);
user_input.insert("content".to_string(), Value::String(text_content));
user_input.insert("modelId".to_string(), Value::String(model_id.to_string()));
user_input.insert("origin".to_string(), Value::String("AI_EDITOR".to_string()));
if !images.is_empty() {
user_input.insert("images".to_string(), Value::Array(images));
}
Some(json!({
"agentContinuationId": agent_continuation_id,
"agentTaskType": "vibe",
"chatTriggerType": "MANUAL",
"currentMessage": {
"userInputMessage": Value::Object(user_input)
},
"conversationId": conversation_id,
"history": history,
}))
}
fn extract_session_id(user_id: &str) -> Option<String> {
let position = user_id.find("session_")?;
let candidate = user_id.get(position + "session_".len()..position + "session_".len() + 36)?;
(candidate.matches('-').count() == 4).then(|| candidate.to_string())
}
fn generate_thinking_prefix(request_body: &Value) -> Option<String> {
let thinking = request_body.get("thinking")?.as_object()?;
match thinking.get("type").and_then(Value::as_str).map(str::trim) {
Some("enabled") => {
let budget_tokens = thinking
.get("budget_tokens")
.and_then(Value::as_i64)
.unwrap_or_default();
Some(format!(
"<thinking_mode>enabled</thinking_mode><max_thinking_length>{budget_tokens}</max_thinking_length>"
))
}
Some("adaptive") => {
let effort = request_body
.get("output_config")
.and_then(Value::as_object)
.and_then(|cfg| cfg.get("effort"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("high");
Some(format!(
"<thinking_mode>adaptive</thinking_mode><thinking_effort>{effort}</thinking_effort>"
))
}
_ => None,
}
}
fn has_thinking_tags(content: &str) -> bool {
content.contains("<thinking_mode>") || content.contains("<max_thinking_length>")
}
fn system_to_text(system: Option<&Value>) -> String {
match system {
Some(Value::String(text)) => text.clone(),
Some(Value::Array(items)) => items
.iter()
.filter_map(|item| {
item.as_object()
.and_then(|item| item.get("text"))
.and_then(Value::as_str)
.map(ToOwned::to_owned)
})
.collect::<Vec<_>>()
.join("\n"),
_ => String::new(),
}
}
fn flush_user_buffer(user_buffer: &mut Vec<&Map<String, Value>>, model_id: &str) -> Option<Value> {
if user_buffer.is_empty() {
return None;
}
let mut parts = Vec::new();
let mut images = Vec::new();
let mut tool_results = Vec::new();
for message in user_buffer.drain(..) {
let (text, mut message_images, mut message_tool_results) =
process_message_content(message.get("content"));
if !text.is_empty() {
parts.push(text);
}
images.append(&mut message_images);
tool_results.append(&mut message_tool_results);
}
let mut payload = Map::new();
payload.insert("content".to_string(), Value::String(parts.join("\n")));
payload.insert("modelId".to_string(), Value::String(model_id.to_string()));
payload.insert("origin".to_string(), Value::String("AI_EDITOR".to_string()));
if !images.is_empty() {
payload.insert("images".to_string(), Value::Array(images));
}
if !tool_results.is_empty() {
payload.insert(
"userInputMessageContext".to_string(),
json!({"toolResults": tool_results}),
);
}
Some(json!({"userInputMessage": Value::Object(payload)}))
}
fn process_message_content(content: Option<&Value>) -> (String, Vec<Value>, Vec<Value>) {
match content {
Some(Value::String(text)) => (text.clone(), Vec::new(), Vec::new()),
Some(Value::Array(blocks)) => {
let mut text_parts = Vec::new();
let mut images = Vec::new();
let mut tool_results = Vec::new();
for block in blocks {
let Some(block) = block.as_object() else {
continue;
};
match block
.get("type")
.and_then(Value::as_str)
.unwrap_or_default()
{
"text" => {
if let Some(text) = block.get("text").and_then(Value::as_str) {
text_parts.push(text.to_string());
}
}
"image" => {
let Some(source) = block.get("source").and_then(Value::as_object) else {
continue;
};
let Some(format) = source
.get("media_type")
.or_else(|| source.get("mediaType"))
.and_then(Value::as_str)
.and_then(image_format)
else {
continue;
};
let Some(bytes) = source.get("data").and_then(Value::as_str) else {
continue;
};
images.push(json!({
"format": format,
"source": {"bytes": bytes}
}));
}
"tool_result" => {
let Some(tool_use_id) = block
.get("tool_use_id")
.or_else(|| block.get("toolUseId"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
continue;
};
let text = match block.get("content") {
Some(Value::String(text)) => text.clone(),
Some(Value::Array(items)) => items
.iter()
.filter_map(|item| {
item.as_object()
.filter(|item| {
item.get("type").and_then(Value::as_str) == Some("text")
})
.and_then(|item| item.get("text"))
.and_then(Value::as_str)
.map(ToOwned::to_owned)
})
.collect::<Vec<_>>()
.join("\n"),
Some(other) => {
serde_json::to_string(other).unwrap_or_else(|_| other.to_string())
}
None => String::new(),
};
let is_error = block
.get("is_error")
.or_else(|| block.get("isError"))
.and_then(Value::as_bool)
.unwrap_or(false);
tool_results.push(json!({
"toolUseId": tool_use_id,
"content": [{"text": text}],
"status": if is_error { "error" } else { "success" },
"isError": is_error,
}));
}
_ => {}
}
}
(text_parts.join(""), images, tool_results)
}
_ => (String::new(), Vec::new(), Vec::new()),
}
}
fn image_format(media_type: &str) -> Option<&'static str> {
let (prefix, suffix) = media_type.split_once('/')?;
if prefix != "image" {
return None;
}
match suffix.trim().to_ascii_lowercase().as_str() {
"jpeg" => Some("jpeg"),
"png" => Some("png"),
"gif" => Some("gif"),
"webp" => Some("webp"),
"jpg" => Some("jpeg"),
_ => None,
}
}
fn clean_tool_schema(value: &Value) -> Value {
match value {
Value::Object(object) => {
let mut out = Map::new();
for (key, inner) in object {
if key == "additionalProperties" {
continue;
}
if key == "required" && inner.as_array().is_some_and(|items| items.is_empty()) {
continue;
}
out.insert(key.clone(), clean_tool_schema(inner));
}
Value::Object(out)
}
Value::Array(items) => Value::Array(items.iter().map(clean_tool_schema).collect()),
_ => value.clone(),
}
}
fn convert_tools(tools: Option<&Value>) -> Vec<Value> {
let Some(tools) = tools.and_then(Value::as_array) else {
return Vec::new();
};
tools
.iter()
.filter_map(|tool| {
let tool = tool.as_object()?;
let name = tool.get("name")?.as_str()?.trim();
if name.is_empty() {
return None;
}
let mut description = tool
.get("description")
.and_then(Value::as_str)
.map(str::trim)
.unwrap_or_default()
.to_string();
let suffix = match name {
"Write" => Some(WRITE_TOOL_DESCRIPTION_SUFFIX),
"Edit" => Some(EDIT_TOOL_DESCRIPTION_SUFFIX),
_ => None,
};
if let Some(suffix) = suffix {
description = if description.is_empty() {
suffix.to_string()
} else {
format!("{description}\n{suffix}")
};
}
if description.len() > 10_000 {
description.truncate(10_000);
}
let input_schema = tool
.get("input_schema")
.or_else(|| tool.get("inputSchema"))
.filter(|value| value.is_object())
.map(clean_tool_schema)
.unwrap_or_else(|| json!({}));
Some(json!({
"toolSpecification": {
"name": name,
"description": description,
"inputSchema": {
"json": input_schema
}
}
}))
})
.collect()
}
fn create_placeholder_tool(name: &str) -> Value {
json!({
"toolSpecification": {
"name": name,
"description": "Tool used in conversation history",
"inputSchema": {
"json": {
"type": "object",
"properties": {}
}
}
}
})
}
fn convert_assistant_message(message: &Map<String, Value>) -> Option<Value> {
let content = message.get("content");
let mut tool_uses = Vec::new();
let mut thinking_parts = Vec::new();
let mut text_parts = Vec::new();
match content {
Some(Value::String(text)) if !text.is_empty() => {
text_parts.push(text.clone());
}
Some(Value::Array(blocks)) => {
for block in blocks {
let Some(block) = block.as_object() else {
continue;
};
match block
.get("type")
.and_then(Value::as_str)
.unwrap_or_default()
{
"thinking" => {
if let Some(thinking) = block.get("thinking").and_then(Value::as_str) {
if !thinking.is_empty() {
thinking_parts.push(thinking.to_string());
}
}
}
"text" => {
if let Some(text) = block.get("text").and_then(Value::as_str) {
if !text.is_empty() {
text_parts.push(text.to_string());
}
}
}
"tool_use" => {
let Some(tool_use_id) = block
.get("id")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
continue;
};
let Some(name) = block
.get("name")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
continue;
};
let input = block
.get("input")
.filter(|value| value.is_object())
.cloned()
.unwrap_or_else(|| json!({}));
tool_uses.push(json!({
"toolUseId": tool_use_id,
"name": name,
"input": input
}));
}
_ => {}
}
}
}
_ => {}
}
let thinking_str = thinking_parts.join("");
let text_str = text_parts.join("");
let mut content_str = if thinking_str.is_empty() {
text_str
} else if text_str.is_empty() {
format!("<thinking>{thinking_str}</thinking>")
} else {
format!("<thinking>{thinking_str}</thinking>\n\n{text_str}")
};
if content_str.is_empty() && !tool_uses.is_empty() {
content_str = " ".to_string();
}
if content_str.is_empty() && tool_uses.is_empty() {
return None;
}
let mut out = Map::new();
out.insert("content".to_string(), Value::String(content_str));
if !tool_uses.is_empty() {
out.insert("toolUses".to_string(), Value::Array(tool_uses));
}
Some(Value::Object(out))
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::convert_claude_messages_to_conversation_state;
#[test]
fn converts_simple_claude_request_into_conversation_state() {
let conversation_state = convert_claude_messages_to_conversation_state(
&json!({
"messages": [
{"role":"user","content":"hello"}
],
"thinking": {"type": "enabled", "budget_tokens": 128},
"tools": [
{"name":"Write","description":"write file","input_schema":{"type":"object","properties":{},"required":[]}}
]
}),
"claude-sonnet-4-upstream",
)
.expect("conversation state should build");
assert_eq!(
conversation_state
.get("currentMessage")
.and_then(|value| value.get("userInputMessage"))
.and_then(|value| value.get("content"))
.and_then(|value| value.as_str()),
Some(
"<thinking_mode>enabled</thinking_mode><max_thinking_length>128</max_thinking_length>\nhello"
)
);
assert_eq!(
conversation_state
.get("currentMessage")
.and_then(|value| value.get("userInputMessage"))
.and_then(|value| value.get("userInputMessageContext"))
.and_then(|value| value.get("tools"))
.and_then(|value| value.as_array())
.map(Vec::len),
Some(1)
);
}
}
@@ -0,0 +1,154 @@
pub use aether_oauth::provider::providers::{
generate_kiro_machine_id as generate_machine_id,
normalize_kiro_machine_id as normalize_machine_id, KiroAuthConfig, DEFAULT_REGION,
};
#[cfg(test)]
mod tests {
use super::{generate_machine_id, normalize_machine_id, KiroAuthConfig, DEFAULT_REGION};
#[test]
fn normalizes_uuid_machine_id() {
assert_eq!(
normalize_machine_id("123e4567-e89b-12d3-a456-426614174000").as_deref(),
Some("123e4567e89b12d3a456426614174000123e4567e89b12d3a456426614174000")
);
}
#[test]
fn hashes_refresh_token_into_machine_id() {
let auth_config = KiroAuthConfig {
auth_method: None,
refresh_token: Some("r".repeat(128)),
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: None,
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let machine_id = generate_machine_id(&auth_config, None).expect("machine id should exist");
assert_eq!(machine_id.len(), 64);
}
#[test]
fn parses_auth_config_aliases() {
let auth_config = KiroAuthConfig::from_raw_json(Some(
r#"{
"authMethod":"identity_center",
"refreshToken":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
"expires_at": 4102444800,
"profileArn":"arn:aws:bedrock:demo",
"apiRegion":"us-west-2",
"clientId":"cid",
"clientSecret":"secret",
"machineId":"123e4567-e89b-12d3-a456-426614174000",
"kiroVersion":"1.2.3",
"systemVersion":"darwin#24.6.0",
"nodeVersion":"22.21.1",
"accessToken":"cached-token"
}"#,
))
.expect("auth config should parse");
assert_eq!(auth_config.auth_method.as_deref(), Some("idc"));
assert_eq!(
auth_config.refresh_token.as_deref(),
Some(
"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
)
);
assert_eq!(auth_config.expires_at, Some(4_102_444_800));
assert_eq!(
auth_config.profile_arn.as_deref(),
Some("arn:aws:bedrock:demo")
);
assert_eq!(auth_config.client_id.as_deref(), Some("cid"));
assert_eq!(auth_config.client_secret.as_deref(), Some("secret"));
assert_eq!(auth_config.effective_api_region(), "us-west-2");
assert_eq!(auth_config.effective_kiro_version(), "1.2.3");
assert_eq!(auth_config.effective_system_version(), "darwin#24.6.0");
assert_eq!(auth_config.effective_node_version(), "22.21.1");
assert_eq!(auth_config.access_token.as_deref(), Some("cached-token"));
assert!(auth_config.is_idc_auth());
assert!(auth_config.profile_arn_for_payload().is_none());
assert_eq!(auth_config.effective_auth_region(), "us-east-1");
assert!(auth_config.can_refresh_access_token());
assert_eq!(DEFAULT_REGION, "us-east-1");
}
#[test]
fn preserves_external_idp_auth_method_for_header_selection() {
let auth_config = KiroAuthConfig::from_raw_json(Some(
r#"{
"authMethod":"external_idp",
"refreshToken":"rt-1",
"clientId":"cid",
"clientSecret":"secret",
"profileArn":"arn:aws:bedrock:demo"
}"#,
))
.expect("auth config should parse");
assert_eq!(auth_config.auth_method.as_deref(), Some("external_idp"));
assert!(auth_config.is_idc_auth());
assert!(auth_config.uses_external_idp_token_type());
assert!(auth_config.profile_arn_for_payload().is_none());
assert_eq!(
auth_config.profile_arn_for_mcp(),
Some("arn:aws:bedrock:demo")
);
}
#[test]
fn infers_idc_when_client_credentials_exist() {
let auth_config = KiroAuthConfig::from_raw_json(Some(
r#"{
"refreshToken":"rt-1",
"clientId":"cid",
"clientSecret":"secret",
"profileArn":"arn:aws:bedrock:demo"
}"#,
))
.expect("auth config should parse");
assert!(auth_config.is_idc_auth());
assert!(!auth_config.uses_external_idp_token_type());
assert!(auth_config.profile_arn_for_payload().is_none());
assert_eq!(
auth_config.profile_arn_for_mcp(),
Some("arn:aws:bedrock:demo")
);
}
#[test]
fn round_trips_json_value() {
let auth_config = KiroAuthConfig::from_raw_json(Some(
r#"{
"auth_method":"social",
"refreshToken":"rt-1....................................................................................................",
"expires_at": 4102444800,
"profileArn":"arn:aws:bedrock:demo",
"region":"eu-north-1",
"apiRegion":"us-west-2",
"machineId":"123e4567-e89b-12d3-a456-426614174000",
"kiroVersion":"1.2.3",
"systemVersion":"darwin#24.6.0",
"nodeVersion":"22.21.1",
"accessToken":"cached-token"
}"#,
))
.expect("auth config should parse");
let value = auth_config.to_json_value();
let reparsed = KiroAuthConfig::from_json_value(&value).expect("auth config should reparse");
assert_eq!(reparsed, auth_config);
}
}
@@ -0,0 +1,334 @@
use std::collections::BTreeMap;
use uuid::Uuid;
use super::credentials::KiroAuthConfig;
pub const AWS_EVENTSTREAM_CONTENT_TYPE: &str = "application/vnd.amazon.eventstream";
pub const KIRO_PROFILE_ARN_HEADER: &str = "x-amzn-kiro-profile-arn";
pub const KIRO_TOKEN_TYPE_HEADER: &str = "TokenType";
pub const KIRO_EXTERNAL_IDP_TOKEN_TYPE: &str = "EXTERNAL_IDP";
const AWS_SDK_JS_MAIN_VERSION: &str = "1.0.27";
const AWS_SDK_JS_LIST_MODELS_VERSION: &str = "1.0.0";
const CODEWHISPERER_OPTOUT: &str = "true";
const KIRO_AGENT_MODE: &str = "vibe";
fn build_kiro_ide_tag(kiro_version: &str, machine_id: &str) -> String {
if machine_id.trim().is_empty() {
format!("KiroIDE-{kiro_version}")
} else {
format!("KiroIDE-{kiro_version}-{machine_id}")
}
}
fn build_x_amz_user_agent_main(kiro_version: &str, machine_id: &str) -> String {
format!(
"aws-sdk-js/{AWS_SDK_JS_MAIN_VERSION} {}",
build_kiro_ide_tag(kiro_version, machine_id)
)
}
fn build_user_agent_main(
system_version: &str,
node_version: &str,
kiro_version: &str,
machine_id: &str,
) -> String {
format!(
"aws-sdk-js/{AWS_SDK_JS_MAIN_VERSION} ua/2.1 os/{system_version} lang/js md/nodejs#{node_version} api/codewhispererstreaming#{AWS_SDK_JS_MAIN_VERSION} m/E {}",
build_kiro_ide_tag(kiro_version, machine_id)
)
}
pub fn build_generate_assistant_headers(
auth_config: &KiroAuthConfig,
machine_id: &str,
) -> BTreeMap<String, String> {
let kiro_version = auth_config.effective_kiro_version();
let system_version = auth_config.effective_system_version();
let node_version = auth_config.effective_node_version();
let region = auth_config.effective_api_region();
let host = format!("q.{region}.amazonaws.com");
BTreeMap::from([
(
"accept".to_string(),
AWS_EVENTSTREAM_CONTENT_TYPE.to_string(),
),
(
"amz-sdk-invocation-id".to_string(),
Uuid::new_v4().to_string(),
),
(
"amz-sdk-request".to_string(),
"attempt=1; max=3".to_string(),
),
("connection".to_string(), "close".to_string()),
("content-type".to_string(), "application/json".to_string()),
("host".to_string(), host),
(
"user-agent".to_string(),
build_user_agent_main(system_version, node_version, kiro_version, machine_id),
),
(
"x-amz-user-agent".to_string(),
build_x_amz_user_agent_main(kiro_version, machine_id),
),
(
"x-amzn-codewhisperer-optout".to_string(),
CODEWHISPERER_OPTOUT.to_string(),
),
(
"x-amzn-kiro-agent-mode".to_string(),
KIRO_AGENT_MODE.to_string(),
),
])
}
pub fn build_mcp_headers(
auth_config: &KiroAuthConfig,
machine_id: &str,
) -> BTreeMap<String, String> {
let kiro_version = auth_config.effective_kiro_version();
let system_version = auth_config.effective_system_version();
let node_version = auth_config.effective_node_version();
let region = auth_config.effective_api_region();
let host = format!("q.{region}.amazonaws.com");
let mut headers = BTreeMap::from([
("accept".to_string(), "application/json".to_string()),
(
"amz-sdk-invocation-id".to_string(),
Uuid::new_v4().to_string(),
),
(
"amz-sdk-request".to_string(),
"attempt=1; max=3".to_string(),
),
("connection".to_string(), "close".to_string()),
("content-type".to_string(), "application/json".to_string()),
("host".to_string(), host),
(
"user-agent".to_string(),
build_user_agent_main(system_version, node_version, kiro_version, machine_id),
),
(
"x-amz-user-agent".to_string(),
build_x_amz_user_agent_main(kiro_version, machine_id),
),
(
"x-amzn-codewhisperer-optout".to_string(),
CODEWHISPERER_OPTOUT.to_string(),
),
]);
if let Some(profile_arn) = auth_config.profile_arn_for_mcp() {
headers.insert(KIRO_PROFILE_ARN_HEADER.to_string(), profile_arn.to_string());
}
headers
}
pub fn build_list_available_models_headers(
auth_config: &KiroAuthConfig,
machine_id: &str,
) -> BTreeMap<String, String> {
let kiro_version = auth_config.effective_kiro_version();
let system_version = auth_config.effective_system_version();
let node_version = auth_config.effective_node_version();
let region = auth_config.effective_api_region();
let host = format!("q.{region}.amazonaws.com");
let ide_tag = build_kiro_ide_tag(kiro_version, machine_id);
BTreeMap::from([
("accept".to_string(), "application/json".to_string()),
(
"amz-sdk-invocation-id".to_string(),
Uuid::new_v4().to_string(),
),
(
"amz-sdk-request".to_string(),
"attempt=1; max=1".to_string(),
),
("connection".to_string(), "close".to_string()),
("host".to_string(), host),
(
"user-agent".to_string(),
format!(
"aws-sdk-js/{AWS_SDK_JS_LIST_MODELS_VERSION} ua/2.1 os/{system_version} lang/js md/nodejs#{node_version} api/codewhispererruntime#1.0.0 m/N,E {ide_tag}"
),
),
(
"x-amz-user-agent".to_string(),
format!("aws-sdk-js/{AWS_SDK_JS_LIST_MODELS_VERSION} {ide_tag}"),
),
])
}
#[cfg(test)]
mod tests {
use super::super::credentials::KiroAuthConfig;
use super::{
build_generate_assistant_headers, build_list_available_models_headers, build_mcp_headers,
AWS_EVENTSTREAM_CONTENT_TYPE, KIRO_PROFILE_ARN_HEADER, KIRO_TOKEN_TYPE_HEADER,
};
#[test]
fn builds_generate_assistant_headers_for_region() {
let auth_config = KiroAuthConfig {
auth_method: None,
refresh_token: None,
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: Some("us-west-2".to_string()),
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: Some("1.2.3".to_string()),
system_version: Some("darwin#24.6.0".to_string()),
node_version: Some("22.21.1".to_string()),
access_token: None,
};
let headers = build_generate_assistant_headers(&auth_config, "machine-123");
assert_eq!(
headers.get("accept").map(String::as_str),
Some(AWS_EVENTSTREAM_CONTENT_TYPE)
);
assert_eq!(
headers.get("host").map(String::as_str),
Some("q.us-west-2.amazonaws.com")
);
assert_eq!(
headers.get("x-amzn-kiro-agent-mode").map(String::as_str),
Some("vibe")
);
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
}
#[test]
fn builds_list_available_models_headers() {
let auth_config = KiroAuthConfig {
auth_method: Some("social".to_string()),
refresh_token: None,
expires_at: None,
profile_arn: Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo".into()),
region: None,
auth_region: None,
api_region: Some("us-east-1".to_string()),
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: Some("0.12.155".to_string()),
system_version: Some("darwin#24.6.0".to_string()),
node_version: Some("22.21.1".to_string()),
access_token: None,
};
let headers = build_list_available_models_headers(&auth_config, "machine-123");
assert_eq!(
headers.get("host").map(String::as_str),
Some("q.us-east-1.amazonaws.com")
);
assert_eq!(
headers.get("amz-sdk-request").map(String::as_str),
Some("attempt=1; max=1")
);
assert_eq!(
headers.get("x-amz-user-agent").map(String::as_str),
Some("aws-sdk-js/1.0.0 KiroIDE-0.12.155-machine-123")
);
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
}
#[test]
fn keeps_generate_assistant_headers_without_external_idp_token_type() {
let auth_config = KiroAuthConfig {
auth_method: Some("idc".to_string()),
refresh_token: None,
expires_at: None,
profile_arn: Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo".into()),
region: None,
auth_region: None,
api_region: Some("us-east-1".to_string()),
client_id: Some("client-id".to_string()),
client_secret: Some("client-secret".to_string()),
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let headers = build_generate_assistant_headers(&auth_config, "machine-123");
assert_eq!(
headers.get("accept").map(String::as_str),
Some(AWS_EVENTSTREAM_CONTENT_TYPE)
);
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
}
#[test]
fn builds_mcp_headers_with_profile_arn_for_social_auth() {
let auth_config = KiroAuthConfig {
auth_method: Some("social".to_string()),
refresh_token: None,
expires_at: None,
profile_arn: Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo".into()),
region: None,
auth_region: None,
api_region: Some("us-east-1".to_string()),
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: Some("0.3.210".to_string()),
system_version: Some("darwin#24.6.0".to_string()),
node_version: Some("22.21.1".to_string()),
access_token: None,
};
let headers = build_mcp_headers(&auth_config, "machine-123");
assert_eq!(
headers.get("accept").map(String::as_str),
Some("application/json")
);
assert_eq!(
headers.get(KIRO_PROFILE_ARN_HEADER).map(String::as_str),
Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo")
);
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
}
#[test]
fn builds_mcp_headers_with_profile_arn_for_idc_auth() {
let auth_config = KiroAuthConfig {
auth_method: Some("idc".to_string()),
refresh_token: None,
expires_at: None,
profile_arn: Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo".into()),
region: None,
auth_region: None,
api_region: Some("us-west-2".to_string()),
client_id: Some("client-id".to_string()),
client_secret: Some("client-secret".to_string()),
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let headers = build_mcp_headers(&auth_config, "machine-123");
assert_eq!(
headers.get("host").map(String::as_str),
Some("q.us-west-2.amazonaws.com")
);
assert_eq!(
headers.get(KIRO_PROFILE_ARN_HEADER).map(String::as_str),
Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo")
);
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
}
}
@@ -0,0 +1,55 @@
mod auth;
mod converter;
mod credentials;
mod headers;
mod policy;
mod refresh;
mod request;
mod url;
use crate::provider_types::{
ProviderApiFormatInheritance, ProviderLocalEmbeddingSupport, ProviderRuntimePolicy,
};
pub const RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrConfiguredBearer,
enable_format_conversion_by_default: true,
allow_auth_channel_mismatch_by_default: true,
oauth_is_bearer_like: true,
supports_model_fetch: true,
supports_local_openai_chat_transport: false,
supports_local_same_format_transport: false,
local_embedding_support: ProviderLocalEmbeddingSupport::None,
};
pub use auth::{
build_kiro_request_auth_from_config, is_kiro_claude_messages_transport,
is_kiro_provider_transport, resolve_local_kiro_bearer_auth, resolve_local_kiro_request_auth,
supports_local_kiro_auth_prerequisites, supports_local_kiro_request_auth_resolution,
KiroBearerAuth, KiroRequestAuth, KIRO_AUTH_HEADER, PROVIDER_TYPE,
};
pub use converter::convert_claude_messages_to_conversation_state;
pub use credentials::{generate_machine_id, normalize_machine_id, KiroAuthConfig};
pub use headers::{
build_generate_assistant_headers, build_list_available_models_headers, build_mcp_headers,
AWS_EVENTSTREAM_CONTENT_TYPE, KIRO_EXTERNAL_IDP_TOKEN_TYPE, KIRO_PROFILE_ARN_HEADER,
KIRO_TOKEN_TYPE_HEADER,
};
pub use policy::{
local_kiro_request_transport_unsupported_reason_with_network,
supports_local_kiro_request_transport, supports_local_kiro_request_transport_with_network,
};
pub use refresh::KiroOAuthRefreshAdapter;
pub use request::{
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
body_rules_are_locally_supported, build_kiro_provider_headers,
build_kiro_provider_request_body, header_rules_are_locally_supported,
supports_local_kiro_request_shape, KiroProviderHeadersInput,
};
pub use url::{
build_kiro_generate_assistant_response_url, build_kiro_list_available_models_url,
build_kiro_mcp_url, build_kiro_mcp_url_from_resolved_url, resolve_kiro_base_url,
GENERATE_ASSISTANT_RESPONSE_PATH, KIRO_ENVELOPE_NAME, LIST_AVAILABLE_MODELS_PATH, MCP_PATH,
MCP_STREAM_PATH,
};
@@ -0,0 +1,203 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::super::{
resolve_transport_profile, transport_profile_is_configured,
transport_proxy_is_locally_supported,
};
use super::{
body_rules_are_locally_supported, header_rules_are_locally_supported,
supports_local_kiro_request_auth_resolution, supports_local_kiro_request_shape, PROVIDER_TYPE,
};
pub fn local_kiro_request_transport_unsupported_reason_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> Option<&'static str> {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return if !transport.provider.is_active {
Some("provider_inactive")
} else if !transport.endpoint.is_active {
Some("endpoint_inactive")
} else {
Some("key_inactive")
};
}
if !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(PROVIDER_TYPE)
{
return Some("transport_provider_type_unsupported");
}
if !aether_ai_formats::api_format_alias_matches(
&transport.endpoint.api_format,
"claude:messages",
) {
return Some("transport_api_format_mismatch");
}
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
return Some("transport_header_rules_unsupported");
}
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
return Some("transport_body_rules_unsupported");
}
if transport.key.decrypted_auth_config.is_some()
&& !supports_local_kiro_request_auth_resolution(transport)
{
return Some("transport_oauth_resolution_unsupported");
}
if !supports_local_kiro_request_auth_resolution(transport) {
return Some("transport_auth_unavailable");
}
if !transport_proxy_is_locally_supported(transport) {
return Some("transport_proxy_unsupported");
}
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
{
return Some("transport_profile_unsupported");
}
None
}
pub fn supports_local_kiro_request_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return false;
}
if !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(PROVIDER_TYPE)
{
return false;
}
if !aether_ai_formats::api_format_alias_matches(
&transport.endpoint.api_format,
"claude:messages",
) {
return false;
}
supports_local_kiro_request_shape(
transport.endpoint.header_rules.as_ref(),
transport.endpoint.body_rules.as_ref(),
) && supports_local_kiro_request_auth_resolution(transport)
}
pub fn supports_local_kiro_request_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
local_kiro_request_transport_unsupported_reason_with_network(transport).is_none()
}
#[cfg(test)]
mod tests {
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
local_kiro_request_transport_unsupported_reason_with_network,
supports_local_kiro_request_transport, supports_local_kiro_request_transport_with_network,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Kiro".to_string(),
provider_type: "kiro".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "claude:messages".to_string(),
api_family: Some("claude".to_string()),
endpoint_kind: Some("messages".to_string()),
is_active: true,
base_url: "https://kiro.example".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: Some(vec!["claude:messages".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: Some(
r#"{
"access_token":"cached-token",
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
"machine_id":"123e4567-e89b-12d3-a456-426614174000"
}"#
.to_string(),
),
},
}
}
#[test]
fn supports_kiro_request_transport_when_cached_access_token_exists() {
assert!(supports_local_kiro_request_transport(&sample_transport()));
assert!(supports_local_kiro_request_transport_with_network(
&sample_transport()
));
}
#[test]
fn supports_kiro_request_transport_when_refresh_only_auth_exists() {
let mut transport = sample_transport();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
}"#
.to_string(),
);
assert!(supports_local_kiro_request_transport(&transport));
assert!(supports_local_kiro_request_transport_with_network(
&transport
));
}
#[test]
fn reports_auth_unavailable_when_kiro_key_cannot_resolve_request_auth() {
let mut transport = sample_transport();
transport.key.auth_type = "api_key".to_string();
transport.key.decrypted_auth_config = None;
assert_eq!(
local_kiro_request_transport_unsupported_reason_with_network(&transport),
Some("transport_auth_unavailable")
);
}
}
@@ -0,0 +1,446 @@
use aether_oauth::provider::providers::KiroProviderOAuthAdapter as CoreKiroProviderOAuthAdapter;
use async_trait::async_trait;
use super::super::oauth_refresh::{
oauth_error_to_local_refresh_error, provider_oauth_transport_context_from_snapshot,
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter, LocalOAuthRefreshError,
LocalResolvedOAuthRequestAuth, ProviderOAuthLocalHttpExecutor,
};
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::auth::{
build_kiro_request_auth_from_config, resolve_local_kiro_request_auth, PROVIDER_TYPE,
};
use super::credentials::KiroAuthConfig;
#[cfg(test)]
const IDC_AMZ_USER_AGENT: &str = "aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
#[derive(Debug, Clone, Default)]
pub struct KiroOAuthRefreshAdapter {
social_refresh_base_url: Option<String>,
idc_refresh_base_url: Option<String>,
}
impl KiroOAuthRefreshAdapter {
pub fn with_refresh_base_urls(
mut self,
social_refresh_base_url: Option<String>,
idc_refresh_base_url: Option<String>,
) -> Self {
self.social_refresh_base_url = social_refresh_base_url;
self.idc_refresh_base_url = idc_refresh_base_url;
self
}
pub async fn refresh_auth_config(
&self,
executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot,
auth_config: &KiroAuthConfig,
) -> Result<KiroAuthConfig, LocalOAuthRefreshError> {
let adapter = CoreKiroProviderOAuthAdapter::default().with_refresh_base_urls(
self.social_refresh_base_url.clone(),
self.idc_refresh_base_url.clone(),
);
let oauth_executor =
ProviderOAuthLocalHttpExecutor::new(PROVIDER_TYPE, transport, executor);
let ctx = provider_oauth_transport_context_from_snapshot(transport);
adapter
.refresh_auth_config(&oauth_executor, &ctx, auth_config)
.await
.map_err(|error| oauth_error_to_local_refresh_error(PROVIDER_TYPE, error))
}
fn auth_config_from_entry(entry: &CachedOAuthEntry) -> Option<KiroAuthConfig> {
entry
.metadata
.as_ref()
.filter(|_| entry.provider_type.eq_ignore_ascii_case(PROVIDER_TYPE))
.and_then(KiroAuthConfig::from_json_value)
}
fn base_auth_config(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Option<KiroAuthConfig> {
entry.and_then(Self::auth_config_from_entry).or_else(|| {
KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
})
}
fn build_cached_entry(auth_config: &KiroAuthConfig) -> Option<CachedOAuthEntry> {
let request_auth = build_kiro_request_auth_from_config(auth_config.clone(), None)?;
Some(CachedOAuthEntry {
provider_type: PROVIDER_TYPE.to_string(),
auth_header_name: request_auth.name.to_string(),
auth_header_value: request_auth.value,
expires_at_unix_secs: auth_config.expires_at,
metadata: Some(auth_config.to_json_value()),
})
}
fn refreshable_auth_config(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Option<KiroAuthConfig> {
let auth_config = self.base_auth_config(transport, entry)?;
auth_config
.can_refresh_access_token()
.then_some(auth_config)
}
}
#[async_trait]
impl LocalOAuthRefreshAdapter for KiroOAuthRefreshAdapter {
fn provider_type(&self) -> &'static str {
PROVIDER_TYPE
}
fn resolve_cached(
&self,
_transport: &GatewayProviderTransportSnapshot,
entry: &CachedOAuthEntry,
) -> Option<LocalResolvedOAuthRequestAuth> {
let auth_config = Self::auth_config_from_entry(entry)?;
let request_auth = build_kiro_request_auth_from_config(auth_config, None)?;
Some(LocalResolvedOAuthRequestAuth::Kiro(request_auth))
}
fn resolve_without_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<LocalResolvedOAuthRequestAuth> {
resolve_local_kiro_request_auth(transport).map(LocalResolvedOAuthRequestAuth::Kiro)
}
fn should_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> bool {
entry
.and_then(|cached| self.resolve_cached(transport, cached))
.is_none()
&& self.resolve_without_refresh(transport).is_none()
&& self.refreshable_auth_config(transport, entry).is_some()
}
async fn refresh(
&self,
executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
let Some(auth_config) = self.refreshable_auth_config(transport, entry) else {
return Ok(None);
};
let refreshed = self
.refresh_auth_config(executor, transport, &auth_config)
.await?;
Ok(Self::build_cached_entry(&refreshed))
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use super::super::super::oauth_refresh::{
LocalOAuthRefreshAdapter, LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
};
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{KiroOAuthRefreshAdapter, IDC_AMZ_USER_AGENT};
use axum::body::to_bytes;
use axum::extract::Request;
use axum::response::IntoResponse;
use axum::routing::any;
use axum::{Json, Router};
use http::StatusCode;
use serde_json::{json, Value};
use tokio::task::JoinHandle;
#[derive(Debug, Clone)]
struct SeenRefreshRequest {
body: Value,
authorization: String,
host: String,
user_agent: String,
x_amz_user_agent: String,
}
fn sample_transport(raw_auth_config: &str) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Kiro".to_string(),
provider_type: "kiro".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "claude:messages".to_string(),
api_family: Some("claude".to_string()),
endpoint_kind: Some("cli".to_string()),
is_active: true,
base_url: "https://kiro.example".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: Some(vec!["claude:messages".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: Some(raw_auth_config.to_string()),
},
}
}
async fn start_server(app: Router) -> (String, JoinHandle<()>) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let addr = listener
.local_addr()
.expect("listener should expose local addr");
let handle = tokio::spawn(async move {
axum::serve(listener, app).await.expect("server should run");
});
(format!("http://{addr}"), handle)
}
#[tokio::test]
async fn refreshes_social_token_via_adapter() {
let seen_request = Arc::new(Mutex::new(None::<SeenRefreshRequest>));
let seen_request_clone = Arc::clone(&seen_request);
let server = Router::new().route(
"/refreshToken",
any(move |request: Request| {
let seen_request_inner = Arc::clone(&seen_request_clone);
async move {
let (parts, body) = request.into_parts();
let raw_body = to_bytes(body, usize::MAX).await.expect("body should read");
let body: Value =
serde_json::from_slice(&raw_body).expect("body should parse as json");
*seen_request_inner.lock().expect("mutex should lock") =
Some(SeenRefreshRequest {
body,
authorization: parts
.headers
.get("authorization")
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
host: parts
.headers
.get("host")
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
user_agent: parts
.headers
.get("user-agent")
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
x_amz_user_agent: parts
.headers
.get("x-amz-user-agent")
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
});
(
StatusCode::OK,
Json(json!({
"accessToken": "cached-kiro-access-token",
"refreshToken": "s".repeat(120),
"expiresIn": 3600,
"profileArn": "arn:aws:bedrock:demo"
})),
)
.into_response()
}
}),
);
let (server_url, server_handle) = start_server(server).await;
let adapter =
KiroOAuthRefreshAdapter::default().with_refresh_base_urls(Some(server_url), None);
let transport = sample_transport(
r#"{
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
"machine_id":"123e4567-e89b-12d3-a456-426614174000",
"kiro_version":"1.2.3"
}"#,
);
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
let entry = adapter
.refresh(&executor, &transport, None)
.await
.expect("refresh should succeed")
.expect("cached entry should exist");
let resolved = adapter
.resolve_cached(&transport, &entry)
.expect("cached entry should resolve");
let seen_request = seen_request
.lock()
.expect("mutex should lock")
.clone()
.expect("refresh request should be captured");
assert_eq!(seen_request.body["refreshToken"], json!("r".repeat(120)));
assert_eq!(seen_request.authorization, "");
assert!(!seen_request.user_agent.is_empty());
assert_eq!(seen_request.x_amz_user_agent, "");
assert!(!seen_request.host.trim().is_empty());
match resolved {
LocalResolvedOAuthRequestAuth::Kiro(auth) => {
assert_eq!(auth.value, "Bearer cached-kiro-access-token");
assert_eq!(
auth.auth_config.profile_arn.as_deref(),
Some("arn:aws:bedrock:demo")
);
assert!(auth.auth_config.expires_at.is_some());
}
other => panic!("unexpected resolved auth: {other:?}"),
}
server_handle.abort();
}
#[tokio::test]
async fn refreshes_idc_token_via_adapter() {
let seen_request = Arc::new(Mutex::new(None::<SeenRefreshRequest>));
let seen_request_clone = Arc::clone(&seen_request);
let server = Router::new().route(
"/token",
any(move |request: Request| {
let seen_request_inner = Arc::clone(&seen_request_clone);
async move {
let (parts, body) = request.into_parts();
let raw_body = to_bytes(body, usize::MAX).await.expect("body should read");
let body: Value =
serde_json::from_slice(&raw_body).expect("body should parse as json");
*seen_request_inner.lock().expect("mutex should lock") =
Some(SeenRefreshRequest {
body,
authorization: parts
.headers
.get("authorization")
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
host: parts
.headers
.get("host")
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
user_agent: parts
.headers
.get("user-agent")
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
x_amz_user_agent: parts
.headers
.get("x-amz-user-agent")
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
});
(
StatusCode::OK,
Json(json!({
"accessToken": "cached-idc-access-token",
"refreshToken": "i".repeat(120),
"expiresIn": 1800
})),
)
.into_response()
}
}),
);
let (server_url, server_handle) = start_server(server).await;
let adapter =
KiroOAuthRefreshAdapter::default().with_refresh_base_urls(None, Some(server_url));
let transport = sample_transport(
r#"{
"auth_method":"identity_center",
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
"client_id":"cid",
"client_secret":"secret",
"profile_arn":"arn:aws:bedrock:demo"
}"#,
);
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
let entry = adapter
.refresh(&executor, &transport, None)
.await
.expect("refresh should succeed")
.expect("cached entry should exist");
let resolved = adapter
.resolve_cached(&transport, &entry)
.expect("cached entry should resolve");
let seen_request = seen_request
.lock()
.expect("mutex should lock")
.clone()
.expect("refresh request should be captured");
assert_eq!(
seen_request.body["grantType"].as_str(),
Some("refresh_token")
);
assert_eq!(seen_request.body["clientId"].as_str(), Some("cid"));
assert_eq!(seen_request.user_agent, "node");
assert_eq!(seen_request.x_amz_user_agent, IDC_AMZ_USER_AGENT);
assert!(!seen_request.host.trim().is_empty());
match resolved {
LocalResolvedOAuthRequestAuth::Kiro(auth) => {
assert_eq!(auth.value, "Bearer cached-idc-access-token");
assert!(auth.auth_config.profile_arn_for_payload().is_none());
assert!(auth.auth_config.expires_at.is_some());
}
other => panic!("unexpected resolved auth: {other:?}"),
}
server_handle.abort();
}
}
@@ -0,0 +1,309 @@
use std::collections::BTreeMap;
use serde_json::{json, Value};
pub use super::super::rules::{
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
body_rules_are_locally_supported, header_rules_are_locally_supported,
};
use super::super::should_skip_upstream_passthrough_header;
use super::converter::convert_claude_messages_to_conversation_state;
use super::credentials::KiroAuthConfig;
use super::headers::build_generate_assistant_headers;
pub fn supports_local_kiro_request_shape(
header_rules: Option<&Value>,
body_rules: Option<&Value>,
) -> bool {
header_rules_are_locally_supported(header_rules) && body_rules_are_locally_supported(body_rules)
}
pub fn build_kiro_provider_request_body(
body_json: &Value,
mapped_model: &str,
auth_config: &KiroAuthConfig,
body_rules: Option<&Value>,
request_headers: Option<&http::HeaderMap>,
) -> Option<Value> {
let conversation_state =
convert_claude_messages_to_conversation_state(body_json, mapped_model)?;
let mut provider_request_body = json!({
"conversationState": conversation_state
});
let mut inference_config = serde_json::Map::new();
if let Some(max_tokens) = body_json
.get("max_tokens")
.and_then(|value| {
value
.as_i64()
.or_else(|| value.as_u64().map(|value| value as i64))
})
.filter(|value| *value > 0)
{
inference_config.insert("maxTokens".to_string(), Value::from(max_tokens));
}
if let Some(temperature) = body_json
.get("temperature")
.and_then(Value::as_f64)
.filter(|value| *value >= 0.0)
{
inference_config.insert("temperature".to_string(), Value::from(temperature));
}
if let Some(top_p) = body_json
.get("top_p")
.and_then(Value::as_f64)
.filter(|value| *value > 0.0)
{
inference_config.insert("topP".to_string(), Value::from(top_p));
}
if !inference_config.is_empty() {
provider_request_body.as_object_mut()?.insert(
"inferenceConfig".to_string(),
Value::Object(inference_config),
);
}
if let Some(profile_arn) = auth_config.profile_arn_for_payload() {
provider_request_body.as_object_mut()?.insert(
"profileArn".to_string(),
Value::String(profile_arn.to_string()),
);
}
if !apply_local_body_rules_with_request_headers(
&mut provider_request_body,
body_rules,
Some(body_json),
request_headers,
) {
return None;
}
Some(provider_request_body)
}
#[derive(Clone, Copy, Debug)]
pub struct KiroProviderHeadersInput<'a> {
pub headers: &'a http::HeaderMap,
pub provider_request_body: &'a Value,
pub original_request_body: &'a Value,
pub header_rules: Option<&'a Value>,
pub auth_header: &'a str,
pub auth_value: &'a str,
pub auth_config: &'a KiroAuthConfig,
pub machine_id: &'a str,
}
pub fn build_kiro_provider_headers(
input: KiroProviderHeadersInput<'_>,
) -> Option<BTreeMap<String, String>> {
let KiroProviderHeadersInput {
headers,
provider_request_body,
original_request_body,
header_rules,
auth_header,
auth_value,
auth_config,
machine_id,
} = input;
let mut out = BTreeMap::new();
for (name, value) in headers {
let Ok(value) = value.to_str() else {
continue;
};
let key = name.as_str().to_ascii_lowercase();
if should_skip_upstream_passthrough_header(&key) {
continue;
}
let value = value.trim();
if value.is_empty() {
continue;
}
out.insert(key, value.to_string());
}
if !apply_local_header_rules_with_request_headers(
&mut out,
header_rules,
&[auth_header, "content-type"],
provider_request_body,
Some(original_request_body),
Some(headers),
) {
return None;
}
for (key, value) in build_generate_assistant_headers(auth_config, machine_id) {
out.insert(key, value);
}
out.insert(
auth_header.trim().to_ascii_lowercase(),
auth_value.trim().to_string(),
);
out.entry("content-type".to_string())
.or_insert_with(|| "application/json".to_string());
out.remove("content-length");
Some(out)
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::super::credentials::KiroAuthConfig;
use super::{
build_kiro_provider_headers, build_kiro_provider_request_body,
supports_local_kiro_request_shape, KiroProviderHeadersInput,
};
#[test]
fn supports_empty_local_request_shape() {
assert!(supports_local_kiro_request_shape(None, None));
}
#[test]
fn rejects_unsupported_rule_shape() {
assert!(!supports_local_kiro_request_shape(
Some(&json!({"action":"set"})),
None
));
}
#[test]
fn supports_simple_header_and_body_rules() {
assert!(supports_local_kiro_request_shape(
Some(&json!([{"action":"set","key":"x-provider-extra","value":"1"}])),
Some(&json!([{"action":"set","path":"debugTag","value":true}]))
));
}
#[test]
fn wraps_claude_request_into_kiro_payload_before_body_rules() {
let auth_config = KiroAuthConfig {
auth_method: None,
refresh_token: Some("r".repeat(128)),
expires_at: None,
profile_arn: Some("arn:aws:bedrock:demo".to_string()),
region: None,
auth_region: None,
api_region: Some("us-east-1".to_string()),
client_id: None,
client_secret: None,
machine_id: Some("123e4567-e89b-12d3-a456-426614174000".to_string()),
kiro_version: None,
system_version: None,
node_version: None,
access_token: Some("cached-token".to_string()),
};
let payload = build_kiro_provider_request_body(
&json!({
"messages": [{"role":"user","content":"hello"}],
"max_tokens": 64
}),
"claude-sonnet-4-upstream",
&auth_config,
Some(&json!([
{"action":"set","path":"debugTag","value":"kiro-local"}
])),
None,
)
.expect("payload should build");
assert!(payload.get("conversationState").is_some());
assert_eq!(
payload
.get("inferenceConfig")
.and_then(|value| value.get("maxTokens")),
Some(&json!(64))
);
assert_eq!(
payload.get("profileArn"),
Some(&json!("arn:aws:bedrock:demo"))
);
assert_eq!(payload.get("debugTag"), Some(&json!("kiro-local")));
}
#[test]
fn applies_header_rules_before_kiro_extra_headers() {
let auth_config = KiroAuthConfig {
auth_method: None,
refresh_token: Some("r".repeat(128)),
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: Some("us-east-1".to_string()),
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: Some("cached-token".to_string()),
};
let headers = build_kiro_provider_headers(KiroProviderHeadersInput {
headers: &http::HeaderMap::new(),
provider_request_body: &json!({"conversationState": {}}),
original_request_body: &json!({"messages": []}),
header_rules: Some(&json!([
{"action":"set","key":"accept","value":"text/plain"},
{"action":"set","key":"x-endpoint-tag","value":"kiro-local"}
])),
auth_header: "authorization",
auth_value: "Bearer cached-token",
auth_config: &auth_config,
machine_id: "machine-123",
})
.expect("headers should build");
assert_eq!(
headers.get("accept").map(String::as_str),
Some("application/vnd.amazon.eventstream")
);
assert_eq!(
headers.get("authorization").map(String::as_str),
Some("Bearer cached-token")
);
assert_eq!(
headers.get("x-endpoint-tag").map(String::as_str),
Some("kiro-local")
);
}
#[test]
fn omits_profile_arn_for_idc_auth() {
let auth_config = KiroAuthConfig {
auth_method: Some("identity_center".to_string()),
refresh_token: Some("r".repeat(128)),
expires_at: None,
profile_arn: Some("arn:aws:bedrock:demo".to_string()),
region: None,
auth_region: None,
api_region: Some("us-east-1".to_string()),
client_id: Some("cid".to_string()),
client_secret: Some("secret".to_string()),
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: Some("cached-token".to_string()),
};
let payload = build_kiro_provider_request_body(
&json!({
"messages": [{"role":"user","content":"hello"}]
}),
"claude-sonnet-4-upstream",
&auth_config,
None,
None,
)
.expect("payload should build");
assert!(payload.get("profileArn").is_none());
}
}
@@ -0,0 +1,138 @@
use super::super::url::build_passthrough_path_url;
use super::credentials::DEFAULT_REGION;
pub const GENERATE_ASSISTANT_RESPONSE_PATH: &str = "/generateAssistantResponse";
pub const LIST_AVAILABLE_MODELS_PATH: &str = "/ListAvailableModels";
pub const MCP_PATH: &str = "/mcp";
pub const MCP_STREAM_PATH: &str = "/mcp/stream";
pub const KIRO_ENVELOPE_NAME: &str = "kiro:generateAssistantResponse";
pub fn resolve_kiro_base_url(upstream_base_url: &str, api_region: Option<&str>) -> String {
let region = api_region
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_REGION);
upstream_base_url
.trim()
.replace("{region}", region)
.trim_end_matches('/')
.to_string()
}
pub fn build_kiro_generate_assistant_response_url(
upstream_base_url: &str,
query: Option<&str>,
api_region: Option<&str>,
) -> Option<String> {
let upstream_base_url = resolve_kiro_base_url(upstream_base_url, api_region);
build_passthrough_path_url(
upstream_base_url.as_str(),
GENERATE_ASSISTANT_RESPONSE_PATH,
query,
&[],
)
}
pub fn build_kiro_mcp_url(upstream_base_url: &str, api_region: Option<&str>) -> Option<String> {
let upstream_base_url = resolve_kiro_base_url(upstream_base_url, api_region);
build_passthrough_path_url(upstream_base_url.as_str(), MCP_PATH, None, &[])
}
pub fn build_kiro_list_available_models_url(
upstream_base_url: &str,
api_region: Option<&str>,
) -> Option<String> {
let mut serializer = url::form_urlencoded::Serializer::new(String::new());
serializer.append_pair("origin", "AI_EDITOR");
let query = serializer.finish();
let upstream_base_url = resolve_kiro_base_url(upstream_base_url, api_region);
build_passthrough_path_url(
upstream_base_url.as_str(),
LIST_AVAILABLE_MODELS_PATH,
Some(&query),
&[],
)
}
pub fn build_kiro_mcp_url_from_resolved_url(resolved_url: &str) -> Option<String> {
let mut parsed = url::Url::parse(resolved_url).ok()?;
parsed.set_path(MCP_PATH);
parsed.set_query(None);
Some(parsed.to_string())
}
#[cfg(test)]
mod tests {
use super::{
build_kiro_generate_assistant_response_url, build_kiro_list_available_models_url,
build_kiro_mcp_url, build_kiro_mcp_url_from_resolved_url, resolve_kiro_base_url,
GENERATE_ASSISTANT_RESPONSE_PATH, KIRO_ENVELOPE_NAME, LIST_AVAILABLE_MODELS_PATH, MCP_PATH,
MCP_STREAM_PATH,
};
#[test]
fn exposes_kiro_request_constants() {
assert_eq!(
GENERATE_ASSISTANT_RESPONSE_PATH,
"/generateAssistantResponse"
);
assert_eq!(LIST_AVAILABLE_MODELS_PATH, "/ListAvailableModels");
assert_eq!(MCP_PATH, "/mcp");
assert_eq!(MCP_STREAM_PATH, "/mcp/stream");
assert_eq!(KIRO_ENVELOPE_NAME, "kiro:generateAssistantResponse");
}
#[test]
fn builds_generate_assistant_response_url() {
assert_eq!(
build_kiro_generate_assistant_response_url(
"https://kiro.{region}.example?tenant=demo",
Some("stream=true"),
Some("us-west-2")
)
.as_deref(),
Some(
"https://kiro.us-west-2.example/generateAssistantResponse?stream=true&tenant=demo"
)
);
}
#[test]
fn resolves_region_placeholder_in_base_url() {
assert_eq!(
resolve_kiro_base_url("https://kiro.{region}.example/", Some("us-west-2")),
"https://kiro.us-west-2.example"
);
}
#[test]
fn builds_mcp_url_for_latest_kiro_endpoint() {
assert_eq!(
build_kiro_mcp_url("https://q.{region}.amazonaws.com", Some("eu-west-1")).as_deref(),
Some("https://q.eu-west-1.amazonaws.com/mcp")
);
}
#[test]
fn builds_list_available_models_url_for_profile() {
assert_eq!(
build_kiro_list_available_models_url(
"https://q.{region}.amazonaws.com",
Some("us-west-2")
)
.as_deref(),
Some("https://q.us-west-2.amazonaws.com/ListAvailableModels?origin=AI_EDITOR")
);
}
#[test]
fn rewrites_generate_assistant_url_to_mcp_url() {
assert_eq!(
build_kiro_mcp_url_from_resolved_url(
"https://q.us-east-1.amazonaws.com/generateAssistantResponse?beta=true"
)
.as_deref(),
Some("https://q.us-east-1.amazonaws.com/mcp")
);
}
}
+154
View File
@@ -0,0 +1,154 @@
pub mod antigravity;
pub mod auth;
mod auth_config;
mod cache;
pub mod claude_code;
pub mod conversion;
mod diagnostics;
pub mod gemini_cli;
mod gemini_files;
mod generic_oauth;
pub mod grok;
mod headers;
pub mod kiro;
mod network;
pub mod oauth_refresh;
mod openai_image;
pub mod policy;
pub mod provider_types;
mod request_body;
mod request_url;
pub mod rules;
pub mod same_format_provider;
pub mod snapshot;
mod standard;
pub mod url;
pub mod vertex;
mod video;
pub mod windsurf;
pub use aether_oauth as oauth;
pub use auth::{build_passthrough_headers, ensure_upstream_auth_header};
pub use auth_config::apply_local_auth_config_header_overrides;
pub use cache::{provider_transport_snapshot_looks_refreshed, ProviderTransportSnapshotCacheKey};
pub use conversion::{
candidate_common_transport_skip_reason, candidate_transport_pair_skip_reason,
request_conversion_direct_auth, request_conversion_enabled_for_transport,
request_conversion_transport_supported, request_conversion_transport_unsupported_reason,
request_pair_allowed_for_transport, request_pair_direct_auth,
request_pair_transport_unsupported_reason, CandidateTransportPolicyFacts,
};
pub use diagnostics::{
append_transport_diagnostics_to_value, build_request_trace_proxy_value,
build_transport_diagnostics,
};
pub use gemini_cli::{
build_gemini_cli_v1internal_request, build_gemini_cli_v1internal_url,
classify_gemini_cli_v1internal_request_body, gemini_cli_v1internal_requires_upstream_streaming,
is_gemini_cli_provider_transport, resolve_gemini_cli_project_id,
resolve_local_gemini_cli_request_auth, GeminiCliRequestAuth, GeminiCliRequestAuthSupport,
GeminiCliRequestAuthUnsupportedReason, GeminiCliRequestEnvelopeSupport,
GeminiCliRequestEnvelopeUnsupportedReason, GeminiCliRequestUrlAction, GEMINI_CLI_PROVIDER_TYPE,
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT,
GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME, GEMINI_CLI_V1INTERNAL_PATH_TEMPLATE,
};
pub use gemini_files::{
build_gemini_files_headers, build_gemini_files_request_body, build_gemini_files_upstream_url,
gemini_files_transport_unsupported_reason, resolve_gemini_files_auth, GeminiFilesHeadersInput,
GeminiFilesRequestBodyError, GeminiFilesRequestBodyParts,
};
pub use generic_oauth::{
supports_local_generic_oauth_request_auth_resolution, GenericOAuthRefreshAdapter,
};
pub use grok::{
build_grok_app_chat_body, build_grok_browser_headers, build_grok_upstream_url, grok_base_url,
grok_browser_profile_id_from_user_agent,
grok_browser_profile_metadata_from_resolved_transport_profile,
grok_browser_resolved_transport_profile,
grok_browser_resolved_transport_profile_from_auth_config,
grok_browser_transport_fingerprint_from_auth_config, is_grok_provider_transport,
resolve_grok_session_auth, GrokBrowserProfileMetadata, GrokHeaderInput, GROK_CHAT_PATH,
GROK_DEFAULT_BASE_URL, GROK_DEFAULT_BROWSER_PROFILE, GROK_DEFAULT_USER_AGENT,
GROK_INTERNAL_HEADER, GROK_RATE_LIMITS_PATH,
};
pub use headers::{should_skip_request_header, should_skip_upstream_passthrough_header};
pub use network::{
resolve_transport_execution_timeouts, resolve_transport_profile, resolve_transport_profile_id,
resolve_transport_proxy_snapshot, resolve_transport_proxy_snapshot_with_tunnel_affinity,
transport_profile_is_configured, transport_proxy_is_locally_supported,
TransportTunnelAffinityLookup, TransportTunnelAttachmentOwner,
};
pub use oauth_refresh::{
supports_local_oauth_request_auth_resolution, CachedOAuthEntry, LocalOAuthHttpExecutor,
LocalOAuthHttpRequest, LocalOAuthHttpResponse, LocalOAuthRefreshCoordinator,
LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
};
pub use openai_image::{
build_openai_image_headers, build_openai_image_upstream_url,
openai_image_transport_unsupported_reason, resolve_openai_image_auth,
ProviderOpenAiImageHeadersInput,
};
pub use policy::{
local_gemini_transport_unsupported_reason,
local_gemini_transport_unsupported_reason_with_network,
local_openai_chat_transport_unsupported_reason, local_standard_transport_unsupported_reason,
local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport,
supports_local_gemini_transport_with_network, supports_local_standard_transport,
};
pub use request_body::{
apply_transport_request_body_semantics, TransportRequestBodySemanticsError,
};
pub use request_url::{
build_cross_format_openai_chat_upstream_url, build_cross_format_openai_responses_upstream_url,
build_kiro_cross_format_upstream_url, build_local_openai_chat_upstream_url,
build_local_openai_responses_upstream_url, build_transport_request_url,
build_transport_request_url_for_request_body, gemini_embedding_request_body_uses_batch,
TransportRequestUrlParams,
};
pub use rules::{
apply_local_body_rules, apply_local_body_rules_with_request_headers, apply_local_header_rules,
apply_local_header_rules_with_request_headers, body_rules_are_locally_supported,
body_rules_handle_path, body_rules_have_enabled_rules, header_rules_are_locally_supported,
header_rules_have_enabled_rules,
};
pub use same_format_provider::{
build_same_format_provider_headers, build_same_format_provider_request_body,
build_same_format_provider_request_body_with_compatibility_report,
build_same_format_provider_upstream_url, classify_same_format_provider_request_behavior,
resolve_same_format_provider_direct_auth, same_format_provider_transport_supported,
same_format_provider_transport_unsupported_reason,
same_format_provider_transport_unsupported_reason_for_trace,
should_try_same_format_provider_oauth_auth, SameFormatProviderCompatibilityEdit,
SameFormatProviderCompatibilityEditAction, SameFormatProviderFamily,
SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior,
SameFormatProviderRequestBehaviorParams, SameFormatProviderRequestBodyInput,
SameFormatProviderRequestBodyOutput, SameFormatProviderUpstreamUrlParams,
};
pub use snapshot::{
read_provider_transport_snapshot, GatewayProviderTransportSnapshot,
ProviderTransportSnapshotSource,
};
pub use standard::{
apply_standard_provider_request_body_rules,
apply_standard_provider_request_body_rules_with_request_headers,
build_standard_plan_fallback_headers, build_standard_plan_fallback_openai_chat_url,
build_standard_plan_fallback_openai_responses_url, build_standard_provider_request_headers,
StandardPlanFallbackAcceptPolicy, StandardPlanFallbackHeadersInput,
StandardProviderRequestHeaders, StandardProviderRequestHeadersInput,
};
pub use vertex::{
is_vertex_api_key_transport_context, is_vertex_service_account_transport_context,
is_vertex_transport_context, uses_vertex_api_key_query_auth,
};
pub use video::{
build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url,
reconstruct_local_video_task_snapshot, resolve_local_video_task_transport,
resolve_video_create_auth, video_create_transport_unsupported_reason,
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, VideoTaskTransportSnapshotLookup,
};
pub use windsurf::{
build_windsurf_cascade_headers, build_windsurf_cascade_request_body,
build_windsurf_cascade_upstream_url, is_windsurf_provider_transport,
local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH,
WINDSURF_ENVELOPE_NAME,
};
@@ -0,0 +1,848 @@
use aether_contracts::{
ExecutionTimeouts, ProxySnapshot, ResolvedTransportProfile, TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
};
use async_trait::async_trait;
use serde_json::{json, Map, Value};
use tracing::warn;
use crate::grok::grok_browser_resolved_transport_profile_from_auth_config;
use super::snapshot::GatewayProviderTransportSnapshot;
const TUNNEL_BASE_URL_EXTRA_KEY: &str = "tunnel_base_url";
const TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY: &str = "tunnel_owner_instance_id";
const TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY: &str = "tunnel_owner_observed_at_unix_secs";
const DEFAULT_PROVIDER_STREAM_FIRST_BYTE_TIMEOUT_SECS: f64 = 30.0;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TransportTunnelAttachmentOwner {
pub gateway_instance_id: String,
pub relay_base_url: String,
pub observed_at_unix_secs: u64,
}
#[async_trait]
pub trait TransportTunnelAffinityLookup: Send + Sync {
async fn lookup_tunnel_attachment_owner(
&self,
node_id: &str,
) -> Result<Option<TransportTunnelAttachmentOwner>, String>;
}
pub fn resolve_transport_execution_timeouts(
transport: &GatewayProviderTransportSnapshot,
) -> Option<ExecutionTimeouts> {
Some(ExecutionTimeouts {
total_ms: transport
.provider
.request_timeout_secs
.filter(|value| value.is_finite() && *value > 0.0)
.map(timeout_secs_to_ms),
first_byte_ms: Some(timeout_secs_to_ms(
transport
.provider
.stream_first_byte_timeout_secs
.filter(|value| value.is_finite() && *value > 0.0)
.unwrap_or(DEFAULT_PROVIDER_STREAM_FIRST_BYTE_TIMEOUT_SECS),
)),
..ExecutionTimeouts::default()
})
}
fn timeout_secs_to_ms(secs: f64) -> u64 {
((secs * 1000.0).round() as u64).max(1)
}
pub fn resolve_transport_proxy_snapshot(
transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
let raw = effective_proxy_config(transport)?;
proxy_snapshot_from_value(raw)
}
pub async fn resolve_transport_proxy_snapshot_with_tunnel_affinity(
lookup: &dyn TransportTunnelAffinityLookup,
transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
let mut snapshot = resolve_transport_proxy_snapshot(transport)?;
let Some(node_id) = snapshot
.node_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Some(snapshot);
};
let owner = match lookup.lookup_tunnel_attachment_owner(node_id).await {
Ok(owner) => owner,
Err(error) => {
warn!(error = %error, node_id = node_id, "failed to load tunnel attachment owner");
None
}
};
let Some(owner) = owner else {
return Some(snapshot);
};
let mut extra = snapshot
.extra
.take()
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
let configured_tunnel_base_url = extra
.get(TUNNEL_BASE_URL_EXTRA_KEY)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
if configured_tunnel_base_url.is_none() {
extra.insert(
TUNNEL_BASE_URL_EXTRA_KEY.to_string(),
Value::String(owner.relay_base_url.clone()),
);
}
extra.insert(
TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY.to_string(),
Value::String(owner.gateway_instance_id),
);
extra.insert(
TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY.to_string(),
json!(owner.observed_at_unix_secs),
);
snapshot.extra = Some(Value::Object(extra));
Some(snapshot)
}
pub fn transport_proxy_is_locally_supported(transport: &GatewayProviderTransportSnapshot) -> bool {
let has_configured_proxy = transport.provider.proxy.is_some()
|| transport.endpoint.proxy.is_some()
|| transport.key.proxy.is_some();
if !has_configured_proxy {
return true;
}
let Some(snapshot) = resolve_transport_proxy_snapshot(transport) else {
return false;
};
if snapshot.enabled == Some(false) {
return true;
}
snapshot
.url
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty())
|| snapshot
.node_id
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
pub fn resolve_transport_profile_id(
transport: &GatewayProviderTransportSnapshot,
) -> Option<String> {
resolve_transport_profile(transport).map(|profile| profile.profile_id)
}
pub fn resolve_transport_profile(
transport: &GatewayProviderTransportSnapshot,
) -> Option<ResolvedTransportProfile> {
resolve_transport_profile_from_fingerprint(transport.key.fingerprint.as_ref()).or_else(|| {
resolve_transport_profile_from_provider_config(transport.provider.config.as_ref())
.or_else(|| resolve_grok_browser_transport_profile(transport))
})
}
fn resolve_grok_browser_transport_profile(
transport: &GatewayProviderTransportSnapshot,
) -> Option<ResolvedTransportProfile> {
if !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok")
{
return None;
}
let auth_config = transport
.key
.decrypted_auth_config
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| serde_json::from_str::<Value>(value).ok())?;
let object = auth_config.as_object()?;
let has_session = json_string_field(object, "sso_token")
.or_else(|| json_string_field(object, "access_token"))
.or_else(|| json_string_field(object, "token"))
.is_some();
if !has_session {
return None;
}
grok_browser_resolved_transport_profile_from_auth_config(object, "grok_auth_config")
}
fn resolve_transport_profile_from_provider_config(
config: Option<&Value>,
) -> Option<ResolvedTransportProfile> {
let fingerprint = config?.get("fingerprint");
resolve_transport_profile_from_fingerprint(fingerprint)
}
fn resolve_transport_profile_from_fingerprint(
fingerprint: Option<&Value>,
) -> Option<ResolvedTransportProfile> {
let fingerprint = fingerprint?;
fingerprint
.get("transport_profile")
.and_then(parse_transport_profile_value)
}
pub fn transport_profile_is_configured(transport: &GatewayProviderTransportSnapshot) -> bool {
transport_profile_configured_in_fingerprint(transport.key.fingerprint.as_ref())
|| transport_profile_configured_in_provider_config(transport.provider.config.as_ref())
}
fn transport_profile_configured_in_provider_config(config: Option<&Value>) -> bool {
let fingerprint = config.and_then(|value| value.get("fingerprint"));
transport_profile_configured_in_fingerprint(fingerprint)
}
fn transport_profile_configured_in_fingerprint(fingerprint: Option<&Value>) -> bool {
fingerprint
.and_then(|value| value.get("transport_profile"))
.is_some_and(|value| !value.is_null())
}
fn parse_transport_profile_value(value: &Value) -> Option<ResolvedTransportProfile> {
if let Some(profile_id) = value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
{
return Some(ResolvedTransportProfile {
profile_id: profile_id.to_string(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string(),
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
header_fingerprint: None,
extra: None,
});
}
let object = value.as_object()?;
let profile_id = json_string_field(object, "profile_id")
.or_else(|| json_string_field(object, "id"))
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())?;
let backend = json_string_field(object, "backend")
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.unwrap_or_else(|| TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string());
let http_mode = json_string_field(object, "http_mode")
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.unwrap_or_else(|| TRANSPORT_HTTP_MODE_AUTO.to_string());
let pool_scope = json_string_field(object, "pool_scope")
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.unwrap_or_else(|| TRANSPORT_POOL_SCOPE_KEY.to_string());
let header_fingerprint = object.get("header_fingerprint").cloned();
let extra = object.get("extra").cloned();
Some(ResolvedTransportProfile {
profile_id,
backend,
http_mode,
pool_scope,
header_fingerprint,
extra,
})
}
fn effective_proxy_config(transport: &GatewayProviderTransportSnapshot) -> Option<&Value> {
[
transport.key.proxy.as_ref(),
transport.endpoint.proxy.as_ref(),
transport.provider.proxy.as_ref(),
]
.into_iter()
.flatten()
.find(|candidate| proxy_enabled(candidate))
}
fn proxy_enabled(value: &Value) -> bool {
value
.as_object()
.and_then(|object| object.get("enabled"))
.and_then(Value::as_bool)
.unwrap_or(true)
}
fn proxy_snapshot_from_value(value: &Value) -> Option<ProxySnapshot> {
let object = value.as_object()?;
let enabled = object.get("enabled").and_then(Value::as_bool);
let mode = json_string_field(object, "mode");
let node_id = json_string_field(object, "node_id");
let label = json_string_field(object, "label");
let url = json_string_field(object, "url").or_else(|| json_string_field(object, "proxy_url"));
let mut extra = Map::new();
for (key, value) in object {
if matches!(
key.as_str(),
"enabled" | "mode" | "node_id" | "label" | "url" | "proxy_url"
) {
continue;
}
extra.insert(key.clone(), value.clone());
}
Some(ProxySnapshot {
enabled,
mode,
node_id,
label,
url,
extra: if extra.is_empty() {
None
} else {
Some(Value::Object(extra))
},
})
}
fn json_string_field(object: &Map<String, Value>, key: &str) -> Option<String> {
object
.get(key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use async_trait::async_trait;
use serde_json::{json, Value};
use super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
resolve_transport_execution_timeouts, resolve_transport_profile,
resolve_transport_profile_id, resolve_transport_proxy_snapshot,
resolve_transport_proxy_snapshot_with_tunnel_affinity, transport_profile_is_configured,
transport_proxy_is_locally_supported, TransportTunnelAffinityLookup,
TransportTunnelAttachmentOwner,
};
use aether_contracts::TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE;
#[derive(Default)]
struct TestTunnelAffinityLookup {
owners: BTreeMap<String, TransportTunnelAttachmentOwner>,
}
#[async_trait]
impl TransportTunnelAffinityLookup for TestTunnelAffinityLookup {
async fn lookup_tunnel_attachment_owner(
&self,
node_id: &str,
) -> Result<Option<TransportTunnelAttachmentOwner>, String> {
Ok(self.owners.get(node_id).cloned())
}
}
fn sample_lookup() -> TestTunnelAffinityLookup {
let mut owners = BTreeMap::new();
owners.insert(
"proxy-node-1".to_string(),
TransportTunnelAttachmentOwner {
gateway_instance_id: "gateway-b".to_string(),
relay_base_url: "http://gateway-b.internal".to_string(),
observed_at_unix_secs: 4_102_444_800u64,
},
);
TestTunnelAffinityLookup { owners }
}
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "provider".to_string(),
provider_type: "custom".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: Some(json!({"url":"http://provider-proxy:8080"})),
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "openai:chat".to_string(),
api_family: Some("openai".to_string()),
endpoint_kind: Some("chat".to_string()),
is_active: true,
base_url: "https://api.openai.example".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: Some(json!({"enabled":false,"url":"http://endpoint-proxy:8080"})),
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: Some(json!({"node_id":"proxy-node-1","kind":"manual"})),
fingerprint: Some(json!({"transport_profile":"chrome_136"})),
upstream_metadata: None,
decrypted_api_key: "sk-test".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn transport_execution_timeouts_use_provider_defaults_when_unset() {
let transport = sample_transport();
let timeouts = resolve_transport_execution_timeouts(&transport)
.expect("default provider timeouts should resolve");
assert_eq!(timeouts.total_ms, None);
assert_eq!(timeouts.first_byte_ms, Some(30_000));
}
#[test]
fn transport_execution_timeouts_preserve_configured_values_independently() {
let mut transport = sample_transport();
transport.provider.request_timeout_secs = Some(12.0);
let timeouts = resolve_transport_execution_timeouts(&transport)
.expect("provider timeouts should resolve");
assert_eq!(timeouts.total_ms, Some(12_000));
assert_eq!(timeouts.first_byte_ms, Some(30_000));
}
#[test]
fn transport_execution_timeouts_preserve_the_configurable_maximum() {
let mut transport = sample_transport();
transport.provider.request_timeout_secs =
Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS as f64);
let timeouts = resolve_transport_execution_timeouts(&transport)
.expect("provider timeouts should resolve");
assert_eq!(
timeouts.total_ms,
Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_MS)
);
}
#[test]
fn transport_execution_timeouts_preserve_configured_first_byte_value() {
let mut transport = sample_transport();
transport.provider.stream_first_byte_timeout_secs = Some(7.5);
let timeouts = resolve_transport_execution_timeouts(&transport)
.expect("provider timeouts should resolve");
assert_eq!(timeouts.total_ms, None);
assert_eq!(timeouts.first_byte_ms, Some(7_500));
}
#[test]
fn resolves_transport_proxy_with_key_precedence() {
let snapshot = resolve_transport_proxy_snapshot(&sample_transport())
.expect("proxy snapshot should resolve");
assert_eq!(snapshot.node_id.as_deref(), Some("proxy-node-1"));
assert_eq!(snapshot.url, None);
assert_eq!(snapshot.extra, Some(json!({"kind":"manual"})));
}
#[tokio::test]
async fn enriches_transport_proxy_snapshot_with_tunnel_owner_hint() {
let state = sample_lookup();
let snapshot =
resolve_transport_proxy_snapshot_with_tunnel_affinity(&state, &sample_transport())
.await
.expect("proxy snapshot should resolve");
assert_eq!(snapshot.node_id.as_deref(), Some("proxy-node-1"));
assert_eq!(
snapshot
.extra
.as_ref()
.and_then(|value| value.get("tunnel_base_url"))
.and_then(Value::as_str),
Some("http://gateway-b.internal")
);
assert_eq!(
snapshot
.extra
.as_ref()
.and_then(|value| value.get("tunnel_owner_instance_id"))
.and_then(Value::as_str),
Some("gateway-b")
);
}
#[tokio::test]
async fn preserves_explicit_tunnel_base_url_when_owner_hint_exists() {
let mut transport = sample_transport();
transport.key.proxy = Some(json!({
"node_id": "proxy-node-1",
"kind": "manual",
"tunnel_base_url": "http://configured-gateway.internal",
}));
let state = sample_lookup();
let snapshot = resolve_transport_proxy_snapshot_with_tunnel_affinity(&state, &transport)
.await
.expect("proxy snapshot should resolve");
assert_eq!(
snapshot
.extra
.as_ref()
.and_then(|value| value.get("tunnel_base_url"))
.and_then(Value::as_str),
Some("http://configured-gateway.internal")
);
assert_eq!(
snapshot
.extra
.as_ref()
.and_then(|value| value.get("tunnel_owner_instance_id"))
.and_then(Value::as_str),
Some("gateway-b")
);
}
#[test]
fn resolves_transport_profile_id_from_key_fingerprint() {
assert_eq!(
resolve_transport_profile_id(&sample_transport()).as_deref(),
Some("chrome_136")
);
assert!(transport_proxy_is_locally_supported(&sample_transport()));
}
#[test]
fn resolves_transport_profile_from_key_fingerprint_before_provider_default() {
let mut transport = sample_transport();
transport.provider.config = Some(json!({
"fingerprint": {"transport_profile": "provider_profile"}
}));
transport.key.fingerprint = Some(json!({
"transport_profile": {
"profile_id": "key_profile",
"backend": "reqwest_rustls",
"http_mode": "http1_only"
}
}));
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "key_profile");
assert_eq!(profile.backend, "reqwest_rustls");
assert_eq!(profile.http_mode, "http1_only");
assert_eq!(profile.pool_scope, "key");
}
#[test]
fn resolves_transport_profile_from_provider_default() {
let mut transport = sample_transport();
transport.key.fingerprint = None;
transport.provider.config = Some(json!({
"fingerprint": {"transport_profile": "provider_profile"}
}));
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "provider_profile");
assert_eq!(profile.backend, "reqwest_rustls");
}
#[test]
fn maps_string_transport_profile_to_resolved_profile() {
let profile = resolve_transport_profile(&sample_transport()).expect("profile");
assert_eq!(profile.profile_id, "chrome_136");
assert_eq!(profile.backend, "reqwest_rustls");
assert_eq!(profile.http_mode, "auto");
assert_eq!(profile.pool_scope, "key");
}
#[test]
fn resolves_h2c_prior_knowledge_transport_profile() {
let mut transport = sample_transport();
transport.key.fingerprint = Some(json!({
"transport_profile": {
"profile_id": "mock-h2c",
"backend": "reqwest_rustls",
"http_mode": "h2c_prior_knowledge"
}
}));
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "mock-h2c");
assert_eq!(profile.http_mode, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE);
}
#[test]
fn resolves_no_transport_profile_without_fingerprint_configuration() {
let mut transport = sample_transport();
transport.key.fingerprint = None;
transport.provider.config = None;
assert!(resolve_transport_profile(&transport).is_none());
assert!(!transport_profile_is_configured(&transport));
}
#[test]
fn resolves_grok_browser_transport_profile_from_session_auth_config() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.fingerprint = None;
transport.provider.config = None;
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token",
"browser_profile": "chrome136",
"cf_clearance": "clearance"
})
.to_string(),
);
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "chrome136");
assert_eq!(profile.backend, "browser_wreq");
assert_eq!(profile.http_mode, "auto");
assert_eq!(profile.pool_scope, "key");
assert_eq!(
profile
.extra
.as_ref()
.and_then(|value| value.get("browser_profile"))
.and_then(Value::as_str),
Some("chrome136")
);
}
#[test]
fn resolves_grok_browser_transport_profile_default_from_session_auth_config() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.fingerprint = None;
transport.provider.config = None;
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token"
})
.to_string(),
);
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "chrome136");
assert_eq!(profile.backend, "browser_wreq");
assert_eq!(
profile
.extra
.as_ref()
.and_then(|value| value.get("source"))
.and_then(Value::as_str),
Some("grok_auth_config")
);
}
#[test]
fn resolves_grok_browser_transport_profile_normalizes_auth_config_alias() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.fingerprint = None;
transport.provider.config = None;
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token",
"browser_profile": "Chrome-137"
})
.to_string(),
);
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "chrome137");
assert_eq!(
profile
.extra
.as_ref()
.and_then(|value| value.get("browser_profile"))
.and_then(Value::as_str),
Some("chrome137")
);
}
#[test]
fn resolves_grok_browser_transport_profile_from_legacy_user_agent() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.fingerprint = None;
transport.provider.config = None;
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token",
"user_agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/137.0.0.0 Safari/537.36"
})
.to_string(),
);
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "chrome137");
assert_eq!(
profile
.extra
.as_ref()
.and_then(|value| value.get("browser_profile"))
.and_then(Value::as_str),
Some("chrome137")
);
}
#[test]
fn key_fingerprint_wins_over_grok_auth_config_fallback() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.provider.config = None;
transport.key.fingerprint = Some(json!({
"transport_profile": {
"profile_id": "chrome136",
"backend": "browser_wreq",
"extra": {"browser_profile": "chrome136", "source": "key"}
}
}));
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token",
"browser_profile": "chrome137"
})
.to_string(),
);
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "chrome136");
assert_eq!(
profile
.extra
.as_ref()
.and_then(|value| value.get("source"))
.and_then(Value::as_str),
Some("key")
);
}
#[test]
fn provider_fingerprint_wins_over_grok_auth_config_fallback() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.fingerprint = None;
transport.provider.config = Some(json!({
"fingerprint": {
"transport_profile": {
"profile_id": "chrome136",
"backend": "browser_wreq",
"extra": {"browser_profile": "chrome136", "source": "provider"}
}
}
}));
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token",
"browser_profile": "chrome137"
})
.to_string(),
);
let profile = resolve_transport_profile(&transport).expect("profile");
assert_eq!(profile.profile_id, "chrome136");
assert_eq!(
profile
.extra
.as_ref()
.and_then(|value| value.get("source"))
.and_then(Value::as_str),
Some("provider")
);
}
#[test]
fn rejects_unsupported_grok_auth_config_browser_profile() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.fingerprint = None;
transport.provider.config = None;
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token",
"browser_profile": "safari999"
})
.to_string(),
);
assert!(resolve_transport_profile(&transport).is_none());
}
#[test]
fn rejects_unsupported_grok_auth_config_user_agent_profile() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.fingerprint = None;
transport.provider.config = None;
transport.key.decrypted_auth_config = Some(
json!({
"sso_token": "sso-token",
"user_agent": "Mozilla/5.0 Version/18.0 Safari/605.1.15"
})
.to_string(),
);
assert!(resolve_transport_profile(&transport).is_none());
}
}
@@ -0,0 +1,791 @@
use std::collections::BTreeMap;
use std::fmt;
use std::sync::Arc;
use aether_oauth::core::OAuthError;
use aether_oauth::network::{
OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse, OAuthNetworkContext,
};
use aether_oauth::provider::ProviderOAuthTransportContext;
use aether_runtime_state::RuntimeState;
use async_trait::async_trait;
use serde_json::Value;
use thiserror::Error;
use tokio::sync::Mutex;
use super::generic_oauth::supports_local_generic_oauth_request_auth_resolution;
pub use super::generic_oauth::GenericOAuthRefreshAdapter;
use super::kiro::{
supports_local_kiro_request_auth_resolution, KiroOAuthRefreshAdapter, KiroRequestAuth,
};
use super::snapshot::GatewayProviderTransportSnapshot;
use super::vertex::{
supports_local_vertex_service_account_auth_resolution, VertexServiceAccountRefreshAdapter,
};
#[derive(Debug, Clone, PartialEq, Eq)]
#[allow(clippy::large_enum_variant)]
pub enum LocalResolvedOAuthRequestAuth {
#[allow(dead_code)]
Header {
name: String,
value: String,
},
Kiro(KiroRequestAuth),
}
#[derive(Debug, Clone, PartialEq)]
pub struct LocalOAuthResolution {
pub auth: Option<LocalResolvedOAuthRequestAuth>,
pub refreshed_entry: Option<CachedOAuthEntry>,
pub refresh_in_flight: bool,
}
#[derive(Debug, Clone, PartialEq)]
pub struct CachedOAuthEntry {
pub provider_type: String,
pub auth_header_name: String,
pub auth_header_value: String,
pub expires_at_unix_secs: Option<u64>,
pub metadata: Option<Value>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct LocalOAuthHttpRequest {
pub request_id: &'static str,
pub method: reqwest::Method,
pub url: String,
pub headers: BTreeMap<String, String>,
pub json_body: Option<Value>,
pub body_bytes: Option<Vec<u8>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LocalOAuthHttpResponse {
pub status_code: u16,
pub body_text: String,
}
#[derive(Debug, Error)]
pub enum LocalOAuthRefreshError {
#[error("{provider_type} oauth refresh request failed: {source}")]
Transport {
provider_type: &'static str,
#[source]
source: reqwest::Error,
},
#[error("{provider_type} oauth refresh returned HTTP {status_code}: {body_excerpt}")]
HttpStatus {
provider_type: &'static str,
status_code: u16,
body_excerpt: String,
},
#[error("{provider_type} oauth refresh transport failed: {message}")]
TransportMessage {
provider_type: &'static str,
message: String,
},
#[error("{provider_type} oauth refresh returned invalid response: {message}")]
InvalidResponse {
provider_type: &'static str,
message: String,
},
}
#[async_trait]
pub trait LocalOAuthHttpExecutor: Send + Sync {
async fn execute(
&self,
provider_type: &'static str,
transport: &GatewayProviderTransportSnapshot,
request: &LocalOAuthHttpRequest,
) -> Result<LocalOAuthHttpResponse, LocalOAuthRefreshError>;
}
#[derive(Debug, Clone)]
pub struct ReqwestLocalOAuthHttpExecutor {
client: reqwest::Client,
}
impl ReqwestLocalOAuthHttpExecutor {
pub fn new(client: reqwest::Client) -> Self {
Self { client }
}
}
#[async_trait]
impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor {
async fn execute(
&self,
provider_type: &'static str,
_transport: &GatewayProviderTransportSnapshot,
request: &LocalOAuthHttpRequest,
) -> Result<LocalOAuthHttpResponse, LocalOAuthRefreshError> {
let mut builder = self
.client
.request(request.method.clone(), request.url.as_str());
for (name, value) in &request.headers {
builder = builder.header(name, value);
}
if let Some(json_body) = request.json_body.as_ref() {
builder = builder.json(json_body);
} else if let Some(body_bytes) = request.body_bytes.as_ref() {
builder = builder.body(body_bytes.clone());
}
let response =
builder
.send()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
provider_type,
source,
})?;
let status_code = response.status().as_u16();
let body_text =
response
.text()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
provider_type,
source,
})?;
Ok(LocalOAuthHttpResponse {
status_code,
body_text,
})
}
}
pub(crate) struct ProviderOAuthLocalHttpExecutor<'a> {
provider_type: &'static str,
transport: &'a GatewayProviderTransportSnapshot,
inner: &'a dyn LocalOAuthHttpExecutor,
}
impl<'a> ProviderOAuthLocalHttpExecutor<'a> {
pub(crate) fn new(
provider_type: &'static str,
transport: &'a GatewayProviderTransportSnapshot,
inner: &'a dyn LocalOAuthHttpExecutor,
) -> Self {
Self {
provider_type,
transport,
inner,
}
}
}
#[async_trait]
impl OAuthHttpExecutor for ProviderOAuthLocalHttpExecutor<'_> {
async fn execute(&self, request: OAuthHttpRequest) -> Result<OAuthHttpResponse, OAuthError> {
let response = self
.inner
.execute(
self.provider_type,
self.transport,
&LocalOAuthHttpRequest {
request_id: "provider-oauth:local-refresh-token",
method: request.method,
url: request.url,
headers: request.headers,
json_body: request.json_body,
body_bytes: request.body_bytes,
},
)
.await
.map_err(local_refresh_error_to_oauth_error)?;
let json_body = serde_json::from_str::<Value>(&response.body_text).ok();
Ok(OAuthHttpResponse {
status_code: response.status_code,
body_text: response.body_text,
json_body,
})
}
}
pub(crate) fn provider_oauth_transport_context_from_snapshot(
transport: &GatewayProviderTransportSnapshot,
) -> ProviderOAuthTransportContext {
ProviderOAuthTransportContext {
provider_id: transport.provider.id.clone(),
provider_type: transport.provider.provider_type.clone(),
endpoint_id: Some(transport.endpoint.id.clone()),
key_id: Some(transport.key.id.clone()),
auth_type: Some(transport.key.auth_type.clone()),
decrypted_api_key: Some(transport.key.decrypted_api_key.clone()),
decrypted_auth_config: transport.key.decrypted_auth_config.clone(),
provider_config: transport.provider.config.clone(),
endpoint_config: transport.endpoint.config.clone(),
key_config: None,
network: OAuthNetworkContext::provider_operation(None),
}
}
pub(crate) fn oauth_error_to_local_refresh_error(
provider_type: &'static str,
error: OAuthError,
) -> LocalOAuthRefreshError {
match error {
OAuthError::HttpStatus {
status_code,
body_excerpt,
} => LocalOAuthRefreshError::HttpStatus {
provider_type,
status_code,
body_excerpt,
},
OAuthError::Transport(message) => LocalOAuthRefreshError::TransportMessage {
provider_type,
message,
},
OAuthError::InvalidRequest(message)
| OAuthError::InvalidResponse(message)
| OAuthError::Storage(message)
| OAuthError::UnsupportedProvider(message) => LocalOAuthRefreshError::InvalidResponse {
provider_type,
message,
},
OAuthError::InvalidState => LocalOAuthRefreshError::InvalidResponse {
provider_type,
message: "oauth state is invalid or expired".to_string(),
},
OAuthError::EncryptionUnavailable => LocalOAuthRefreshError::InvalidResponse {
provider_type,
message: "oauth encryption unavailable".to_string(),
},
}
}
fn local_refresh_error_to_oauth_error(error: LocalOAuthRefreshError) -> OAuthError {
match error {
LocalOAuthRefreshError::Transport { source, .. } => {
OAuthError::Transport(source.to_string())
}
LocalOAuthRefreshError::TransportMessage { message, .. } => OAuthError::Transport(message),
LocalOAuthRefreshError::HttpStatus {
status_code,
body_excerpt,
..
} => OAuthError::HttpStatus {
status_code,
body_excerpt,
},
LocalOAuthRefreshError::InvalidResponse { message, .. } => {
OAuthError::InvalidResponse(message)
}
}
}
#[async_trait]
pub trait LocalOAuthRefreshAdapter: Send + Sync {
fn provider_type(&self) -> &'static str;
fn supports(&self, transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(self.provider_type())
}
fn resolve_cached(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: &CachedOAuthEntry,
) -> Option<LocalResolvedOAuthRequestAuth>;
fn resolve_without_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<LocalResolvedOAuthRequestAuth>;
fn should_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> bool;
async fn refresh(
&self,
executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError>;
}
pub struct LocalOAuthRefreshCoordinator {
adapters: Vec<Arc<dyn LocalOAuthRefreshAdapter>>,
cache: Mutex<BTreeMap<String, CachedOAuthEntry>>,
key_locks: Mutex<BTreeMap<String, Arc<Mutex<()>>>>,
}
impl fmt::Debug for LocalOAuthRefreshCoordinator {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("LocalOAuthRefreshCoordinator")
.field("adapter_count", &self.adapters.len())
.finish()
}
}
impl Default for LocalOAuthRefreshCoordinator {
fn default() -> Self {
Self::new()
}
}
impl LocalOAuthRefreshCoordinator {
const DISTRIBUTED_REFRESH_LOCK_TTL_MS: u64 = 30_000;
pub fn new() -> Self {
Self {
adapters: vec![
Arc::new(KiroOAuthRefreshAdapter::default()),
Arc::new(VertexServiceAccountRefreshAdapter),
Arc::new(GenericOAuthRefreshAdapter::default()),
],
cache: Mutex::new(BTreeMap::new()),
key_locks: Mutex::new(BTreeMap::new()),
}
}
async fn lock_for_key(&self, key_id: &str) -> Arc<Mutex<()>> {
let mut key_locks = self.key_locks.lock().await;
key_locks
.entry(key_id.to_string())
.or_insert_with(|| Arc::new(Mutex::new(())))
.clone()
}
async fn cached_entry(&self, key_id: &str) -> Option<CachedOAuthEntry> {
self.cache.lock().await.get(key_id).cloned()
}
async fn insert_cached_entry(&self, key_id: &str, entry: CachedOAuthEntry) {
self.cache.lock().await.insert(key_id.to_string(), entry);
}
pub async fn store_cached_entry(&self, key_id: &str, entry: CachedOAuthEntry) {
self.insert_cached_entry(key_id, entry).await;
}
pub async fn invalidate_cached_entry(&self, key_id: &str) -> bool {
self.cache.lock().await.remove(key_id).is_some()
}
pub async fn resolve_with_result(
&self,
executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot,
distributed_lock: Option<&RuntimeState>,
distributed_owner: Option<&str>,
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
self.resolve_with_result_mode(
executor,
transport,
distributed_lock,
distributed_owner,
false,
)
.await
}
pub async fn force_refresh_with_result(
&self,
executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot,
distributed_lock: Option<&RuntimeState>,
distributed_owner: Option<&str>,
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
self.resolve_with_result_mode(
executor,
transport,
distributed_lock,
distributed_owner,
true,
)
.await
}
async fn resolve_with_result_mode(
&self,
executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot,
distributed_lock: Option<&RuntimeState>,
distributed_owner: Option<&str>,
force_refresh: bool,
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
let Some(adapter) = self
.adapters
.iter()
.find(|adapter| adapter.supports(transport))
else {
return Ok(None);
};
let key_id = transport.key.id.trim();
let cached_entry = if key_id.is_empty() {
None
} else {
self.cached_entry(key_id).await
};
if !force_refresh {
if let Some(auth) = cached_entry
.as_ref()
.and_then(|entry| adapter.resolve_cached(transport, entry))
{
return Ok(Some(LocalOAuthResolution::resolved(auth, None)));
}
if let Some(auth) = adapter.resolve_without_refresh(transport) {
return Ok(Some(LocalOAuthResolution::resolved(auth, None)));
}
if !adapter.should_refresh(transport, cached_entry.as_ref()) {
return Ok(None);
}
}
if key_id.is_empty() {
return Ok(None);
}
let key_lock = self.lock_for_key(key_id).await;
let _key_guard = key_lock.lock().await;
let cached_entry = self.cached_entry(key_id).await;
if !force_refresh {
if let Some(auth) = cached_entry
.as_ref()
.and_then(|entry| adapter.resolve_cached(transport, entry))
{
return Ok(Some(LocalOAuthResolution::resolved(auth, None)));
}
if let Some(auth) = adapter.resolve_without_refresh(transport) {
return Ok(Some(LocalOAuthResolution::resolved(auth, None)));
}
if !adapter.should_refresh(transport, cached_entry.as_ref()) {
return Ok(None);
}
}
let distributed_lease = match (distributed_lock, distributed_owner) {
(Some(lock), Some(owner)) if !owner.trim().is_empty() => {
match lock
.lock_try_acquire(
&format!("provider_oauth_refresh_lock:{key_id}"),
owner,
std::time::Duration::from_millis(Self::DISTRIBUTED_REFRESH_LOCK_TTL_MS),
)
.await
{
Ok(Some(lease)) => Some(lease),
Ok(None) => return Ok(Some(LocalOAuthResolution::refresh_in_flight())),
Err(err) => {
tracing::warn!(
key_id = %key_id,
provider_type = adapter.provider_type(),
error = ?err,
"gateway local oauth refresh distributed lock unavailable"
);
None
}
}
}
_ => None,
};
// Forced refresh still needs the latest rotated refresh_token as input.
// Otherwise a second overlapping refresh can acquire the lock after the
// first one completes, then immediately retry with the stale token that
// came from the original transport snapshot.
let refresh_entry = cached_entry.as_ref();
let refresh_result = adapter.refresh(executor, transport, refresh_entry).await;
if let (Some(lock), Some(lease)) = (distributed_lock, distributed_lease.as_ref()) {
if let Err(err) = lock.lock_release(lease).await {
tracing::warn!(
key_id = %key_id,
provider_type = adapter.provider_type(),
error = ?err,
"gateway local oauth refresh distributed lock release failed"
);
}
}
let Some(refreshed_entry) = refresh_result? else {
return Ok(None);
};
Ok(adapter
.resolve_cached(transport, &refreshed_entry)
.map(|auth| LocalOAuthResolution::resolved(auth, Some(refreshed_entry))))
}
pub fn with_adapters_for_tests(adapters: Vec<Arc<dyn LocalOAuthRefreshAdapter>>) -> Self {
Self {
adapters,
cache: Mutex::new(BTreeMap::new()),
key_locks: Mutex::new(BTreeMap::new()),
}
}
}
impl LocalOAuthResolution {
fn resolved(
auth: LocalResolvedOAuthRequestAuth,
refreshed_entry: Option<CachedOAuthEntry>,
) -> Self {
Self {
auth: Some(auth),
refreshed_entry,
refresh_in_flight: false,
}
}
fn refresh_in_flight() -> Self {
Self {
auth: None,
refreshed_entry: None,
refresh_in_flight: true,
}
}
}
pub fn supports_local_oauth_request_auth_resolution(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
supports_local_kiro_request_auth_resolution(transport)
|| supports_local_vertex_service_account_auth_resolution(transport)
|| supports_local_generic_oauth_request_auth_resolution(transport)
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter,
LocalOAuthRefreshCoordinator, LocalOAuthRefreshError, LocalOAuthResolution,
LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
};
use async_trait::async_trait;
use std::sync::Arc;
#[derive(Debug)]
struct TestAdapter {
refresh_hits: Arc<AtomicUsize>,
refresh_with_entry_hits: Arc<AtomicUsize>,
}
#[async_trait]
impl LocalOAuthRefreshAdapter for TestAdapter {
fn provider_type(&self) -> &'static str {
"test-oauth"
}
fn resolve_cached(
&self,
_transport: &GatewayProviderTransportSnapshot,
entry: &CachedOAuthEntry,
) -> Option<LocalResolvedOAuthRequestAuth> {
(entry.provider_type == "test-oauth").then(|| LocalResolvedOAuthRequestAuth::Header {
name: entry.auth_header_name.clone(),
value: entry.auth_header_value.clone(),
})
}
fn resolve_without_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<LocalResolvedOAuthRequestAuth> {
let secret = transport.key.decrypted_api_key.trim();
(!secret.is_empty() && secret != "__placeholder__").then(|| {
LocalResolvedOAuthRequestAuth::Header {
name: "authorization".to_string(),
value: format!("Bearer {secret}"),
}
})
}
fn should_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> bool {
entry.is_none() && transport.key.decrypted_api_key.trim() == "__placeholder__"
}
async fn refresh(
&self,
_executor: &dyn LocalOAuthHttpExecutor,
_transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
self.refresh_hits.fetch_add(1, Ordering::SeqCst);
if entry.is_some() {
self.refresh_with_entry_hits.fetch_add(1, Ordering::SeqCst);
}
Ok(Some(CachedOAuthEntry {
provider_type: "test-oauth".to_string(),
auth_header_name: "authorization".to_string(),
auth_header_value: "Bearer refreshed-token".to_string(),
expires_at_unix_secs: Some(4_102_444_800),
metadata: None,
}))
}
}
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "test".to_string(),
provider_type: "test-oauth".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "claude:messages".to_string(),
api_family: Some("claude".to_string()),
endpoint_kind: Some("cli".to_string()),
is_active: true,
base_url: "https://example.test".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: Some("{\"refresh_token\":\"rt-1\"}".to_string()),
},
}
}
#[tokio::test]
async fn coordinator_reuses_runtime_cached_refresh_result() {
let refresh_hits = Arc::new(AtomicUsize::new(0));
let refresh_with_entry_hits = Arc::new(AtomicUsize::new(0));
let coordinator =
LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(TestAdapter {
refresh_hits: Arc::clone(&refresh_hits),
refresh_with_entry_hits: Arc::clone(&refresh_with_entry_hits),
})]);
let transport = sample_transport();
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
let first = coordinator
.resolve_with_result(&executor, &transport, None, None)
.await
.expect("first resolve should succeed");
coordinator
.insert_cached_entry(
transport.key.id.as_str(),
first
.as_ref()
.and_then(|result| result.refreshed_entry.clone())
.expect("first resolve should provide cached entry"),
)
.await;
let second = coordinator
.resolve_with_result(&executor, &transport, None, None)
.await
.expect("second resolve should succeed");
assert_eq!(refresh_hits.load(Ordering::SeqCst), 1);
assert_eq!(
first,
Some(LocalOAuthResolution {
auth: Some(LocalResolvedOAuthRequestAuth::Header {
name: "authorization".to_string(),
value: "Bearer refreshed-token".to_string(),
}),
refreshed_entry: Some(CachedOAuthEntry {
provider_type: "test-oauth".to_string(),
auth_header_name: "authorization".to_string(),
auth_header_value: "Bearer refreshed-token".to_string(),
expires_at_unix_secs: Some(4_102_444_800),
metadata: None,
}),
refresh_in_flight: false,
})
);
assert_eq!(
second,
Some(LocalOAuthResolution {
auth: Some(LocalResolvedOAuthRequestAuth::Header {
name: "authorization".to_string(),
value: "Bearer refreshed-token".to_string(),
}),
refreshed_entry: None,
refresh_in_flight: false,
})
);
}
#[tokio::test]
async fn coordinator_force_refresh_bypasses_runtime_cache() {
let refresh_hits = Arc::new(AtomicUsize::new(0));
let refresh_with_entry_hits = Arc::new(AtomicUsize::new(0));
let coordinator =
LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(TestAdapter {
refresh_hits: Arc::clone(&refresh_hits),
refresh_with_entry_hits: Arc::clone(&refresh_with_entry_hits),
})]);
let transport = sample_transport();
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
let first = coordinator
.resolve_with_result(&executor, &transport, None, None)
.await
.expect("initial resolve should succeed");
coordinator
.insert_cached_entry(
transport.key.id.as_str(),
first
.as_ref()
.and_then(|result| result.refreshed_entry.clone())
.expect("first resolve should provide cached entry"),
)
.await;
let forced = coordinator
.force_refresh_with_result(&executor, &transport, None, None)
.await
.expect("forced refresh should succeed");
assert!(first.and_then(|result| result.refreshed_entry).is_some());
assert!(forced.and_then(|result| result.refreshed_entry).is_some());
assert_eq!(refresh_hits.load(Ordering::SeqCst), 2);
assert_eq!(refresh_with_entry_hits.load(Ordering::SeqCst), 1);
}
}
@@ -0,0 +1,419 @@
use std::collections::BTreeMap;
use serde_json::Value;
use crate::auth::{build_passthrough_headers_with_auth, resolve_local_openai_bearer_auth};
use crate::grok::{is_grok_provider_transport, resolve_grok_session_auth};
use crate::policy::local_standard_transport_unsupported_reason_with_network;
use crate::rules::apply_local_header_rules_with_request_headers;
use crate::snapshot::GatewayProviderTransportSnapshot;
use crate::url::build_openai_image_url;
#[derive(Debug, Clone, Copy)]
pub struct ProviderOpenAiImageHeadersInput<'a> {
pub transport: &'a GatewayProviderTransportSnapshot,
pub headers: &'a http::HeaderMap,
pub auth_header: &'a str,
pub auth_value: &'a str,
pub accept: Option<&'a str>,
pub header_rules: Option<&'a Value>,
pub provider_request_body: &'a Value,
pub original_request_body: &'a Value,
}
pub fn openai_image_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<&'static str> {
let reason = local_standard_transport_unsupported_reason_with_network(transport, api_format);
if is_dedicated_openai_image_provider(transport)
&& matches!(
reason,
Some("transport_provider_type_unsupported")
| Some("transport_oauth_resolution_unsupported")
)
{
return None;
}
reason
}
fn is_dedicated_openai_image_provider(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("chatgpt_web")
|| transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("codex")
|| is_grok_provider_transport(transport)
}
pub fn resolve_openai_image_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
if is_grok_provider_transport(transport) {
return resolve_grok_session_auth(transport);
}
resolve_local_openai_bearer_auth(transport)
}
pub fn build_openai_image_upstream_url(
transport: &GatewayProviderTransportSnapshot,
request_path: Option<&str>,
request_query: Option<&str>,
) -> String {
build_openai_image_url(&transport.endpoint.base_url, request_path, request_query)
}
pub fn build_openai_image_headers(
input: ProviderOpenAiImageHeadersInput<'_>,
) -> Option<BTreeMap<String, String>> {
let mut provider_request_headers = build_passthrough_headers_with_auth(
input.headers,
input.auth_header,
input.auth_value,
&BTreeMap::new(),
);
provider_request_headers.insert("content-type".to_string(), "application/json".to_string());
if let Some(accept) = input.accept {
provider_request_headers.insert("accept".to_string(), accept.to_string());
} else {
provider_request_headers.remove("accept");
}
crate::apply_local_auth_config_header_overrides(
&mut provider_request_headers,
input.transport.key.decrypted_auth_config.as_deref(),
);
if !apply_local_header_rules_with_request_headers(
&mut provider_request_headers,
input.header_rules,
&[input.auth_header, "content-type", "accept"],
input.provider_request_body,
Some(input.original_request_body),
Some(input.headers),
) {
return None;
}
Some(provider_request_headers)
}
#[cfg(test)]
mod tests {
use http::HeaderMap;
use serde_json::json;
use super::{
build_openai_image_headers, build_openai_image_upstream_url,
openai_image_transport_unsupported_reason, ProviderOpenAiImageHeadersInput,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Provider".to_string(),
provider_type: "codex".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "openai:image".to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://api.openai.com/v1".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn codex_openai_image_url_uses_images_surface() {
let url = build_openai_image_upstream_url(
&sample_transport(),
Some("/v1/images/generations"),
Some("trace=1"),
);
assert_eq!(url, "https://api.openai.com/v1/images/generations?trace=1");
}
#[test]
fn codex_openai_image_edit_url_uses_images_edit_surface() {
let url = build_openai_image_upstream_url(
&sample_transport(),
Some("/v1/images/edits"),
Some("trace=1"),
);
assert_eq!(url, "https://api.openai.com/v1/images/edits?trace=1");
}
#[test]
fn standard_openai_image_url_uses_images_surface() {
let mut transport = sample_transport();
transport.provider.provider_type = "openai".to_string();
let url = build_openai_image_upstream_url(
&transport,
Some("/v1/images/generations"),
Some("trace=1"),
);
assert_eq!(url, "https://api.openai.com/v1/images/generations?trace=1");
}
#[test]
fn standard_openai_image_url_preserves_edit_surface() {
let mut transport = sample_transport();
transport.provider.provider_type = "openai".to_string();
let url =
build_openai_image_upstream_url(&transport, Some("/v1/images/edits"), Some("trace=1"));
assert_eq!(url, "https://api.openai.com/v1/images/edits?trace=1");
}
#[test]
fn chatgpt_web_is_supported_by_dedicated_openai_image_transport_policy() {
let mut transport = sample_transport();
transport.provider.provider_type = "chatgpt_web".to_string();
assert_eq!(
openai_image_transport_unsupported_reason(&transport, "openai:image"),
None
);
}
#[test]
fn grok_is_supported_by_dedicated_openai_image_transport_policy() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
assert_eq!(
openai_image_transport_unsupported_reason(&transport, "openai:image"),
None
);
}
#[test]
fn codex_is_supported_by_dedicated_openai_image_transport_policy() {
let mut transport = sample_transport();
transport.provider.provider_type = "codex".to_string();
transport.key.auth_type = "oauth".to_string();
transport.key.decrypted_auth_config = Some(json!({"access_token":"token"}).to_string());
assert_eq!(
openai_image_transport_unsupported_reason(&transport, "openai:image"),
None
);
}
#[test]
fn grok_oauth_session_is_supported_by_dedicated_openai_image_transport_policy() {
let mut transport = sample_transport();
transport.provider.provider_type = "grok".to_string();
transport.key.auth_type = "oauth".to_string();
transport.key.decrypted_api_key = String::new();
transport.key.decrypted_auth_config = Some(json!({"sso_token":"abc"}).to_string());
assert_eq!(
openai_image_transport_unsupported_reason(&transport, "openai:image"),
None
);
}
#[test]
fn chatgpt_web_oauth_is_supported_by_dedicated_openai_image_transport_policy() {
let mut transport = sample_transport();
transport.provider.provider_type = "chatgpt_web".to_string();
transport.key.auth_type = "oauth".to_string();
transport.key.decrypted_api_key = String::new();
transport.key.decrypted_auth_config = Some(json!({"access_token":"token"}).to_string());
assert_eq!(
openai_image_transport_unsupported_reason(&transport, "openai:image"),
None
);
}
#[test]
fn builds_json_eventstream_headers_and_applies_rules() {
let transport = sample_transport();
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
transport: &transport,
headers: &HeaderMap::new(),
auth_header: "authorization",
auth_value: "Bearer secret",
accept: Some("text/event-stream"),
header_rules: Some(&json!([
{"action":"set","key":"x-image-route","value":"codex"}
])),
provider_request_body: &json!({"model":"gpt-5.4-mini"}),
original_request_body: &json!({"prompt":"draw"}),
})
.expect("headers should build");
assert_eq!(
headers.get("authorization"),
Some(&"Bearer secret".to_string())
);
assert_eq!(
headers.get("content-type"),
Some(&"application/json".to_string())
);
assert_eq!(
headers.get("accept"),
Some(&"text/event-stream".to_string())
);
assert_eq!(headers.get("x-image-route"), Some(&"codex".to_string()));
}
#[test]
fn auth_config_headers_override_default_authorization() {
let mut transport = sample_transport();
transport.key.decrypted_auth_config = Some(
json!({
"refresh_token": "rt-1",
"headers": {
"authorization": "Bearer imported-session",
"chatgpt-account-id": "acct-1"
}
})
.to_string(),
);
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
transport: &transport,
headers: &HeaderMap::new(),
auth_header: "authorization",
auth_value: "Bearer refreshed-access-token",
accept: Some("text/event-stream"),
header_rules: None,
provider_request_body: &json!({"model":"gpt-5.4-mini"}),
original_request_body: &json!({"prompt":"draw"}),
})
.expect("headers should build");
assert_eq!(
headers.get("authorization"),
Some(&"Bearer imported-session".to_string())
);
assert_eq!(
headers.get("chatgpt-account-id"),
Some(&"acct-1".to_string())
);
}
#[test]
fn standard_openai_compatible_image_headers_can_request_json() {
let transport = sample_transport();
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
transport: &transport,
headers: &HeaderMap::new(),
auth_header: "authorization",
auth_value: "Bearer secret",
accept: Some("application/json"),
header_rules: None,
provider_request_body: &json!({
"model": "upstream-image-model",
"prompt": "draw a city",
}),
original_request_body: &json!({"prompt":"draw a city"}),
})
.expect("headers should build");
assert_eq!(headers.get("accept"), Some(&"application/json".to_string()));
assert_eq!(
headers.get("content-type"),
Some(&"application/json".to_string())
);
}
#[test]
fn codex_images_omits_explicit_accept_header() {
let transport = sample_transport();
let mut request_headers = HeaderMap::new();
request_headers.insert("accept", "text/event-stream".parse().expect("valid header"));
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
transport: &transport,
headers: &request_headers,
auth_header: "authorization",
auth_value: "Bearer secret",
accept: None,
header_rules: Some(&json!([
{"action":"set","key":"accept","value":"application/json"}
])),
provider_request_body: &json!({
"model": "gpt-image-2",
"prompt": "draw a city",
}),
original_request_body: &json!({"prompt":"draw a city"}),
})
.expect("headers should build");
assert!(!headers.contains_key("accept"));
}
#[test]
fn standard_openai_compatible_image_url_supports_aether_api_root() {
let mut transport = sample_transport();
transport.provider.provider_type = "custom".to_string();
transport.endpoint.base_url = "https://upstream-aether.example/v1".to_string();
let url = build_openai_image_upstream_url(
&transport,
Some("/v1/images/generations"),
Some("trace=1"),
);
assert_eq!(
url,
"https://upstream-aether.example/v1/images/generations?trace=1"
);
}
}
@@ -0,0 +1,395 @@
use super::provider_types::{
provider_type_supports_local_embedding_transport,
provider_type_supports_local_openai_chat_transport,
provider_type_supports_local_same_format_transport,
};
use super::snapshot::GatewayProviderTransportSnapshot;
use super::{
body_rules_are_locally_supported, header_rules_are_locally_supported,
resolve_transport_profile, supports_local_oauth_request_auth_resolution,
transport_profile_is_configured, transport_proxy_is_locally_supported,
};
pub fn supports_local_openai_chat_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
local_openai_chat_transport_unsupported_reason(transport).is_none()
}
pub fn supports_local_standard_transport(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
local_standard_transport_unsupported_reason(transport, api_format).is_none()
}
pub fn supports_local_gemini_transport(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
local_gemini_transport_unsupported_reason(transport, api_format).is_none()
}
pub fn supports_local_standard_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
local_standard_transport_unsupported_reason_with_network(transport, api_format).is_none()
}
pub fn supports_local_gemini_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
local_gemini_transport_unsupported_reason_with_network(transport, api_format).is_none()
}
pub fn local_openai_chat_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
) -> Option<&'static str> {
if !transport.provider.is_active {
return Some("provider_inactive");
}
if !transport.endpoint.is_active {
return Some("endpoint_inactive");
}
if !transport.key.is_active {
return Some("key_inactive");
}
if !transport
.endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:chat")
{
return Some("transport_api_format_mismatch");
}
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
return Some("transport_header_rules_unsupported");
}
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
return Some("transport_body_rules_unsupported");
}
if transport.key.decrypted_auth_config.is_some()
&& !supports_local_oauth_request_auth_resolution(transport)
{
return Some("transport_oauth_resolution_unsupported");
}
if !transport_proxy_is_locally_supported(transport) {
return Some("transport_proxy_unsupported");
}
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
{
return Some("transport_profile_unsupported");
}
if !provider_type_supports_local_openai_chat_transport(&transport.provider.provider_type) {
return Some("transport_provider_type_unsupported");
}
None
}
pub fn local_standard_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<&'static str> {
local_same_format_transport_unsupported_reason(
transport,
api_format,
false,
provider_type_supports_local_same_format_transport,
)
}
pub fn local_gemini_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<&'static str> {
local_same_format_transport_unsupported_reason(
transport,
api_format,
false,
provider_type_supports_local_same_format_transport,
)
}
pub fn local_standard_transport_unsupported_reason_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<&'static str> {
local_same_format_transport_unsupported_reason(
transport,
api_format,
true,
provider_type_supports_local_same_format_transport,
)
}
pub fn local_gemini_transport_unsupported_reason_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<&'static str> {
local_same_format_transport_unsupported_reason(
transport,
api_format,
true,
provider_type_supports_local_same_format_transport,
)
}
fn local_same_format_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
allow_network_passthrough: bool,
provider_type_supported: fn(&str) -> bool,
) -> Option<&'static str> {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return if !transport.provider.is_active {
Some("provider_inactive")
} else if !transport.endpoint.is_active {
Some("endpoint_inactive")
} else {
Some("key_inactive")
};
}
if !same_api_format(&transport.endpoint.api_format, api_format) {
return Some("transport_api_format_mismatch");
}
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
return Some("transport_header_rules_unsupported");
}
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
return Some("transport_body_rules_unsupported");
}
if transport.key.decrypted_auth_config.is_some()
&& !supports_local_oauth_request_auth_resolution(transport)
{
return Some("transport_oauth_resolution_unsupported");
}
let has_custom_path = transport
.endpoint
.custom_path
.as_deref()
.is_some_and(|value| !value.trim().is_empty());
if has_custom_path && !allow_network_passthrough {
return Some("transport_custom_path_unsupported");
}
if allow_network_passthrough {
if !transport_proxy_is_locally_supported(transport) {
return Some("transport_proxy_unsupported");
}
if transport_profile_is_configured(transport)
&& resolve_transport_profile(transport).is_none()
{
return Some("transport_profile_unsupported");
}
} else if transport.provider.proxy.is_some()
|| transport.endpoint.proxy.is_some()
|| transport.key.proxy.is_some()
|| transport_profile_is_configured(transport)
{
return Some("transport_proxy_or_profile_unsupported");
}
if !provider_type_supported(&transport.provider.provider_type) {
return Some("transport_provider_type_unsupported");
}
if aether_ai_formats::is_embedding_api_format(api_format) {
if !endpoint_kind_allows_embedding(transport.endpoint.endpoint_kind.as_deref()) {
return Some("transport_endpoint_kind_unsupported");
}
if !provider_type_supports_local_embedding_transport(
&transport.provider.provider_type,
api_format,
) {
return Some("transport_provider_type_unsupported");
}
}
if aether_ai_formats::is_rerank_api_format(api_format) {
if !endpoint_kind_allows_rerank(transport.endpoint.endpoint_kind.as_deref()) {
return Some("transport_endpoint_kind_unsupported");
}
if !provider_type_supports_local_embedding_transport(
&transport.provider.provider_type,
api_format,
) {
return Some("transport_provider_type_unsupported");
}
}
None
}
fn endpoint_kind_allows_embedding(endpoint_kind: Option<&str>) -> bool {
endpoint_kind
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| {
matches!(
value.to_ascii_lowercase().as_str(),
"embedding" | "embeddings" | "multimodal_embedding" | "multimodal_embeddings"
)
})
.unwrap_or(true)
}
fn endpoint_kind_allows_rerank(endpoint_kind: Option<&str>) -> bool {
endpoint_kind
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| matches!(value.to_ascii_lowercase().as_str(), "rerank" | "reranking"))
.unwrap_or(true)
}
fn same_api_format(left: &str, right: &str) -> bool {
aether_ai_formats::api_format_alias_matches(left, right)
}
#[cfg(test)]
mod tests {
use super::local_standard_transport_unsupported_reason_with_network;
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport(
provider_type: &str,
api_format: &str,
endpoint_kind: Option<&str>,
) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "provider".to_string(),
provider_type: provider_type.to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: api_format.to_string(),
api_family: None,
endpoint_kind: endpoint_kind.map(ToOwned::to_owned),
is_active: true,
base_url: "https://provider.example".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "sk-test".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn rejects_unsupported_embedding_provider_format_pairs() {
let openai_on_gemini = sample_transport("openai", "gemini:embedding", Some("embedding"));
let gemini_on_openai = sample_transport("gemini", "openai:embedding", Some("embedding"));
let chat_marked_embedding = sample_transport("openai", "openai:embedding", Some("chat"));
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&openai_on_gemini,
"gemini:embedding"
),
Some("transport_provider_type_unsupported")
);
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&gemini_on_openai,
"openai:embedding"
),
Some("transport_provider_type_unsupported")
);
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&chat_marked_embedding,
"openai:embedding"
),
Some("transport_endpoint_kind_unsupported")
);
}
#[test]
fn accepts_supported_embedding_provider_format_pairs() {
for (provider_type, api_format) in [
("openai", "openai:embedding"),
("gemini", "gemini:embedding"),
("google", "gemini:embedding"),
("jina", "jina:embedding"),
("doubao", "doubao:embedding"),
("volcengine", "doubao:embedding"),
("aliyun", "aliyun:multimodal_embedding"),
("dashscope", "aliyun:multimodal_embedding"),
("custom", "openai:embedding"),
("custom", "gemini:embedding"),
("custom", "jina:embedding"),
("custom", "doubao:embedding"),
("custom", "aliyun:multimodal_embedding"),
] {
let transport = sample_transport(provider_type, api_format, Some("embedding"));
assert_eq!(
local_standard_transport_unsupported_reason_with_network(&transport, api_format),
None,
"{provider_type} should support {api_format}"
);
}
}
#[test]
fn embedding_policy_accepts_embedding_endpoint_kind_aliases_only() {
for endpoint_kind in [None, Some(""), Some(" embedding "), Some("EMBEDDINGS")] {
let transport = sample_transport("openai", "openai:embedding", endpoint_kind);
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&transport,
"openai:embedding"
),
None,
"endpoint kind {endpoint_kind:?} should be accepted"
);
}
for endpoint_kind in [Some("chat"), Some("responses"), Some("image")] {
let transport = sample_transport("openai", "openai:embedding", endpoint_kind);
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&transport,
"openai:embedding"
),
Some("transport_endpoint_kind_unsupported"),
"endpoint kind {endpoint_kind:?} should fail closed"
);
}
}
}
@@ -0,0 +1,961 @@
#[derive(Debug, Clone, Copy)]
pub struct ProviderOAuthTemplate {
pub provider_type: &'static str,
pub display_name: &'static str,
pub authorize_url: &'static str,
pub token_url: &'static str,
pub client_id: &'static str,
pub client_secret: &'static str,
pub scopes: &'static [&'static str],
pub redirect_uri: &'static str,
pub use_pkce: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FixedProviderEndpointConfigValue {
String(&'static str),
Bool(bool),
I64(i64),
}
impl FixedProviderEndpointConfigValue {
pub fn to_json_value(self) -> serde_json::Value {
match self {
Self::String(value) => serde_json::Value::String(value.to_string()),
Self::Bool(value) => serde_json::Value::Bool(value),
Self::I64(value) => serde_json::Value::Number(value.into()),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FixedProviderEndpointConfigDefault {
pub key: &'static str,
pub value: FixedProviderEndpointConfigValue,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FixedProviderEndpointTemplate {
pub item_key: &'static str,
pub api_format: &'static str,
pub custom_path: Option<&'static str>,
pub config_defaults: &'static [FixedProviderEndpointConfigDefault],
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderApiFormatInheritance {
None,
OAuth,
OAuthOrBearer,
OAuthOrServiceAccount,
OAuthOrConfiguredBearer,
}
impl ProviderApiFormatInheritance {
pub fn key_inherits_api_formats(
self,
auth_type: &str,
decrypted_auth_config: Option<&str>,
) -> bool {
let auth_type = auth_type.trim().to_ascii_lowercase();
match self {
Self::None => false,
Self::OAuth => auth_type == "oauth",
Self::OAuthOrBearer => auth_type == "oauth" || auth_type == "bearer",
Self::OAuthOrServiceAccount => {
auth_type == "oauth" || auth_type == "service_account" || auth_type == "vertex_ai"
}
Self::OAuthOrConfiguredBearer => {
auth_type == "oauth"
|| auth_type == "bearer"
&& decrypted_auth_config
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderLocalEmbeddingSupport {
None,
AnyKnown,
OpenAi,
Gemini,
Jina,
Doubao,
Aliyun,
}
impl ProviderLocalEmbeddingSupport {
pub fn supports_api_format(self, api_format: &str) -> bool {
let api_format = aether_ai_formats::normalize_api_format_alias(api_format);
match self {
Self::None => false,
Self::AnyKnown => matches!(
api_format.as_str(),
"openai:embedding"
| "openai:rerank"
| "gemini:embedding"
| "jina:embedding"
| "jina:rerank"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
),
Self::OpenAi => matches!(api_format.as_str(), "openai:embedding" | "openai:rerank"),
Self::Gemini => api_format == "gemini:embedding",
Self::Jina => matches!(api_format.as_str(), "jina:embedding" | "jina:rerank"),
Self::Doubao => api_format == "doubao:embedding",
Self::Aliyun => api_format == "aliyun:multimodal_embedding",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ProviderRuntimePolicy {
pub fixed_provider: bool,
pub api_format_inheritance: ProviderApiFormatInheritance,
pub enable_format_conversion_by_default: bool,
pub allow_auth_channel_mismatch_by_default: bool,
pub oauth_is_bearer_like: bool,
pub supports_model_fetch: bool,
pub supports_local_openai_chat_transport: bool,
pub supports_local_same_format_transport: bool,
pub local_embedding_support: ProviderLocalEmbeddingSupport,
}
impl ProviderRuntimePolicy {
pub const fn standard() -> Self {
Self {
fixed_provider: false,
api_format_inheritance: ProviderApiFormatInheritance::None,
enable_format_conversion_by_default: false,
allow_auth_channel_mismatch_by_default: false,
oauth_is_bearer_like: false,
supports_model_fetch: true,
supports_local_openai_chat_transport: true,
supports_local_same_format_transport: true,
local_embedding_support: ProviderLocalEmbeddingSupport::None,
}
}
pub fn key_inherits_api_formats(
self,
auth_type: &str,
decrypted_auth_config: Option<&str>,
) -> bool {
self.api_format_inheritance
.key_inherits_api_formats(auth_type, decrypted_auth_config)
}
pub fn supports_local_embedding_transport(self, api_format: &str) -> bool {
self.local_embedding_support.supports_api_format(api_format)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FixedProviderTemplate {
pub provider_type: &'static str,
pub version: u32,
pub base_url: &'static str,
pub endpoints: &'static [FixedProviderEndpointTemplate],
pub runtime_policy: ProviderRuntimePolicy,
}
const EMPTY_ENDPOINT_CONFIG_DEFAULTS: &[FixedProviderEndpointConfigDefault] = &[];
const FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS: &[FixedProviderEndpointConfigDefault] =
&[FixedProviderEndpointConfigDefault {
key: "upstream_stream_policy",
value: FixedProviderEndpointConfigValue::String("force_stream"),
}];
const AUTO_STREAM_ENDPOINT_CONFIG_DEFAULTS: &[FixedProviderEndpointConfigDefault] =
&[FixedProviderEndpointConfigDefault {
key: "upstream_stream_policy",
value: FixedProviderEndpointConfigValue::String("auto"),
}];
const STANDARD_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy::standard();
const CUSTOM_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
local_embedding_support: ProviderLocalEmbeddingSupport::AnyKnown,
..STANDARD_RUNTIME_POLICY
};
const OPENAI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
local_embedding_support: ProviderLocalEmbeddingSupport::OpenAi,
..STANDARD_RUNTIME_POLICY
};
const GEMINI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
local_embedding_support: ProviderLocalEmbeddingSupport::Gemini,
..STANDARD_RUNTIME_POLICY
};
const JINA_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
local_embedding_support: ProviderLocalEmbeddingSupport::Jina,
..STANDARD_RUNTIME_POLICY
};
const DOUBAO_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
local_embedding_support: ProviderLocalEmbeddingSupport::Doubao,
..STANDARD_RUNTIME_POLICY
};
const ALIYUN_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
local_embedding_support: ProviderLocalEmbeddingSupport::Aliyun,
..STANDARD_RUNTIME_POLICY
};
const CLAUDE_CODE_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
enable_format_conversion_by_default: true,
oauth_is_bearer_like: true,
supports_model_fetch: false,
supports_local_openai_chat_transport: false,
supports_local_same_format_transport: false,
..STANDARD_RUNTIME_POLICY
};
const CODEX_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
enable_format_conversion_by_default: true,
supports_model_fetch: false,
supports_local_openai_chat_transport: false,
..STANDARD_RUNTIME_POLICY
};
const CHATGPT_WEB_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrBearer,
enable_format_conversion_by_default: true,
oauth_is_bearer_like: true,
supports_model_fetch: false,
supports_local_openai_chat_transport: false,
supports_local_same_format_transport: false,
..STANDARD_RUNTIME_POLICY
};
const GEMINI_CLI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
oauth_is_bearer_like: true,
supports_local_openai_chat_transport: false,
..STANDARD_RUNTIME_POLICY
};
const VERTEX_AI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrServiceAccount,
enable_format_conversion_by_default: true,
supports_model_fetch: false,
supports_local_openai_chat_transport: false,
supports_local_same_format_transport: false,
local_embedding_support: ProviderLocalEmbeddingSupport::Gemini,
..STANDARD_RUNTIME_POLICY
};
const ANTIGRAVITY_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
enable_format_conversion_by_default: true,
oauth_is_bearer_like: true,
supports_model_fetch: false,
supports_local_openai_chat_transport: false,
supports_local_same_format_transport: false,
..STANDARD_RUNTIME_POLICY
};
const GROK_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
enable_format_conversion_by_default: true,
supports_model_fetch: false,
supports_local_openai_chat_transport: false,
supports_local_same_format_transport: false,
..STANDARD_RUNTIME_POLICY
};
const WINDSURF_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrBearer,
enable_format_conversion_by_default: true,
oauth_is_bearer_like: true,
supports_model_fetch: false,
supports_local_openai_chat_transport: false,
supports_local_same_format_transport: false,
..STANDARD_RUNTIME_POLICY
};
const CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "claude_code",
version: 1,
base_url: "https://api.anthropic.com",
endpoints: &[FixedProviderEndpointTemplate {
item_key: "claude:messages",
api_format: "claude:messages",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
}],
runtime_policy: CLAUDE_CODE_RUNTIME_POLICY,
};
const CODEX_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "codex",
version: 1,
base_url: "https://chatgpt.com/backend-api/codex",
endpoints: &[
FixedProviderEndpointTemplate {
item_key: "openai:responses",
api_format: "openai:responses",
custom_path: None,
config_defaults: FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "openai:responses:compact",
api_format: "openai:responses:compact",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "openai:search",
api_format: "openai:search",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "openai:image",
api_format: "openai:image",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
],
runtime_policy: CODEX_RUNTIME_POLICY,
};
const CHATGPT_WEB_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "chatgpt_web",
version: 1,
base_url: "https://chatgpt.com",
endpoints: &[FixedProviderEndpointTemplate {
item_key: "openai:image",
api_format: "openai:image",
custom_path: None,
config_defaults: FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS,
}],
runtime_policy: CHATGPT_WEB_RUNTIME_POLICY,
};
const KIRO_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "kiro",
version: 1,
base_url: "https://q.{region}.amazonaws.com",
endpoints: &[FixedProviderEndpointTemplate {
item_key: "claude:messages",
api_format: "claude:messages",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
}],
runtime_policy: crate::kiro::RUNTIME_POLICY,
};
const GEMINI_CLI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "gemini_cli",
version: 3,
base_url: "https://cloudcode-pa.googleapis.com",
endpoints: &[FixedProviderEndpointTemplate {
item_key: "gemini:generate_content",
api_format: "gemini:generate_content",
custom_path: Some("/v1internal:{action}"),
config_defaults: AUTO_STREAM_ENDPOINT_CONFIG_DEFAULTS,
}],
runtime_policy: GEMINI_CLI_RUNTIME_POLICY,
};
const VERTEX_AI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "vertex_ai",
version: 1,
base_url: "https://aiplatform.googleapis.com",
endpoints: &[
FixedProviderEndpointTemplate {
item_key: "gemini:generate_content",
api_format: "gemini:generate_content",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "gemini:embedding",
api_format: "gemini:embedding",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "claude:messages",
api_format: "claude:messages",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
],
runtime_policy: VERTEX_AI_RUNTIME_POLICY,
};
const ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "antigravity",
version: 2,
base_url: "https://daily-cloudcode-pa.googleapis.com",
endpoints: &[FixedProviderEndpointTemplate {
item_key: "gemini:generate_content",
api_format: "gemini:generate_content",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
}],
runtime_policy: ANTIGRAVITY_RUNTIME_POLICY,
};
const GROK_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "grok",
version: 1,
base_url: "https://grok.com",
endpoints: &[
FixedProviderEndpointTemplate {
item_key: "openai:chat",
api_format: "openai:chat",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "openai:responses",
api_format: "openai:responses",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "claude:messages",
api_format: "claude:messages",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "openai:image",
api_format: "openai:image",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
],
runtime_policy: GROK_RUNTIME_POLICY,
};
const WINDSURF_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "windsurf",
version: 1,
base_url: "https://server.codeium.com",
endpoints: &[FixedProviderEndpointTemplate {
item_key: "openai:chat",
api_format: "openai:chat",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
}],
runtime_policy: WINDSURF_RUNTIME_POLICY,
};
pub fn provider_type_is_fixed(provider_type: &str) -> bool {
provider_runtime_policy(provider_type).fixed_provider
}
pub fn fixed_provider_key_inherits_api_formats(
provider_type: &str,
auth_type: &str,
decrypted_auth_config: Option<&str>,
) -> bool {
provider_runtime_policy(provider_type)
.key_inherits_api_formats(auth_type, decrypted_auth_config)
}
pub fn provider_type_enables_format_conversion_by_default(provider_type: &str) -> bool {
provider_runtime_policy(provider_type).enable_format_conversion_by_default
}
pub fn provider_type_allows_auth_channel_mismatch_by_default(provider_type: &str) -> bool {
provider_runtime_policy(provider_type).allow_auth_channel_mismatch_by_default
}
pub fn provider_type_oauth_is_bearer_like(provider_type: &str) -> bool {
provider_runtime_policy(provider_type).oauth_is_bearer_like
}
pub fn provider_runtime_policy(provider_type: &str) -> ProviderRuntimePolicy {
if let Some(template) = fixed_provider_template(provider_type) {
return template.runtime_policy;
}
match provider_type.trim().to_ascii_lowercase().as_str() {
"custom" => CUSTOM_RUNTIME_POLICY,
"openai" => OPENAI_RUNTIME_POLICY,
"gemini" | "google" => GEMINI_RUNTIME_POLICY,
"jina" => JINA_RUNTIME_POLICY,
"doubao" | "volcengine" => DOUBAO_RUNTIME_POLICY,
"aliyun" | "dashscope" => ALIYUN_RUNTIME_POLICY,
_ => STANDARD_RUNTIME_POLICY,
}
}
pub fn fixed_provider_template(provider_type: &str) -> Option<&'static FixedProviderTemplate> {
match provider_type.trim().to_ascii_lowercase().as_str() {
"claude_code" => Some(&CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE),
"codex" => Some(&CODEX_FIXED_PROVIDER_TEMPLATE),
"chatgpt_web" => Some(&CHATGPT_WEB_FIXED_PROVIDER_TEMPLATE),
"kiro" => Some(&KIRO_FIXED_PROVIDER_TEMPLATE),
"grok" => Some(&GROK_FIXED_PROVIDER_TEMPLATE),
"gemini_cli" => Some(&GEMINI_CLI_FIXED_PROVIDER_TEMPLATE),
"vertex_ai" => Some(&VERTEX_AI_FIXED_PROVIDER_TEMPLATE),
"antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE),
"windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE),
_ => None,
}
}
pub fn fixed_provider_endpoint_template_by_api_format(
provider_type: &str,
api_format: &str,
) -> Option<&'static FixedProviderEndpointTemplate> {
let normalized = aether_ai_formats::normalize_api_format_alias(api_format);
fixed_provider_template(provider_type)?
.endpoints
.iter()
.find(|item| item.api_format.eq_ignore_ascii_case(&normalized))
}
pub fn provider_type_supports_model_fetch(provider_type: &str) -> bool {
provider_runtime_policy(provider_type).supports_model_fetch
}
pub fn provider_type_supports_local_openai_chat_transport(provider_type: &str) -> bool {
provider_runtime_policy(provider_type).supports_local_openai_chat_transport
}
pub fn provider_type_supports_local_same_format_transport(provider_type: &str) -> bool {
provider_runtime_policy(provider_type).supports_local_same_format_transport
}
pub fn provider_type_supports_local_embedding_transport(
provider_type: &str,
api_format: &str,
) -> bool {
provider_runtime_policy(provider_type).supports_local_embedding_transport(api_format)
}
pub fn is_codex_cli_backend_url(url: &str) -> bool {
let url = url.trim().to_ascii_lowercase();
url.contains("/codex") && (url.contains("/backend-api/") || url.contains("/backendapi/"))
}
pub fn provider_type_is_fixed_for_admin_oauth(provider_type: &str) -> bool {
provider_type_is_fixed(provider_type)
}
pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<ProviderOAuthTemplate> {
match provider_type.trim().to_ascii_lowercase().as_str() {
"claude_code" => Some(ProviderOAuthTemplate {
provider_type: "claude_code",
display_name: "ClaudeCode",
authorize_url: "https://claude.ai/oauth/authorize",
token_url: "https://console.anthropic.com/v1/oauth/token",
client_id: "9d1c250a-e61b-44d9-88ed-5944d1962f5e",
client_secret: "",
scopes: &["org:create_api_key", "user:profile", "user:inference"],
redirect_uri: "http://localhost:54545/callback",
use_pkce: true,
}),
"codex" => Some(ProviderOAuthTemplate {
provider_type: "codex",
display_name: "Codex",
authorize_url: "https://auth.openai.com/oauth/authorize",
token_url: "https://auth.openai.com/oauth/token",
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
client_secret: "",
scopes: &["openid", "email", "profile", "offline_access"],
redirect_uri: "http://localhost:1455/auth/callback",
use_pkce: true,
}),
"chatgpt_web" => Some(ProviderOAuthTemplate {
provider_type: "chatgpt_web",
display_name: "ChatGPT Web",
authorize_url: "https://auth.openai.com/oauth/authorize",
token_url: "https://auth.openai.com/oauth/token",
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
client_secret: "",
scopes: &["openid", "email", "profile", "offline_access"],
redirect_uri: "http://localhost:1455/auth/callback",
use_pkce: true,
}),
"gemini_cli" => Some(ProviderOAuthTemplate {
provider_type: "gemini_cli",
display_name: "GeminiCli",
authorize_url: "https://accounts.google.com/o/oauth2/v2/auth",
token_url: "https://oauth2.googleapis.com/token",
client_id: "681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com",
client_secret: "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",
scopes: &[
"https://www.googleapis.com/auth/cloud-platform",
"https://www.googleapis.com/auth/userinfo.email",
"https://www.googleapis.com/auth/userinfo.profile",
],
redirect_uri: "http://localhost:8085/oauth2callback",
use_pkce: false,
}),
"antigravity" => Some(ProviderOAuthTemplate {
provider_type: "antigravity",
display_name: "Antigravity",
authorize_url: "https://accounts.google.com/o/oauth2/v2/auth",
token_url: "https://oauth2.googleapis.com/token",
client_id: "1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com",
client_secret: "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",
scopes: &[
"https://www.googleapis.com/auth/cloud-platform",
"https://www.googleapis.com/auth/userinfo.email",
"https://www.googleapis.com/auth/userinfo.profile",
"https://www.googleapis.com/auth/cclog",
"https://www.googleapis.com/auth/experimentsandconfigs",
],
redirect_uri: "http://localhost:51121/oauth2callback",
use_pkce: true,
}),
"windsurf" => Some(ProviderOAuthTemplate {
provider_type: "windsurf",
display_name: "Windsurf",
authorize_url: "https://windsurf.com/windsurf/signin",
token_url: "https://register.windsurf.com/exa.seat_management_pb.SeatManagementService/RegisterUser",
client_id: "3GUryQ7ldAeKEuD2obYnppsnmj58eP5u",
client_secret: "",
scopes: &[],
redirect_uri: "show-auth-token",
use_pkce: false,
}),
_ => None,
}
}
pub const ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES: &[&str] = &[
"claude_code",
"codex",
"chatgpt_web",
"gemini_cli",
"antigravity",
"windsurf",
];
#[cfg(test)]
mod tests {
use super::{
fixed_provider_endpoint_template_by_api_format, fixed_provider_key_inherits_api_formats,
fixed_provider_template, provider_runtime_policy, provider_type_admin_oauth_template,
provider_type_allows_auth_channel_mismatch_by_default, provider_type_oauth_is_bearer_like,
provider_type_supports_local_embedding_transport,
provider_type_supports_local_same_format_transport, FixedProviderEndpointConfigValue,
ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES,
};
#[test]
fn codex_fixed_provider_template_includes_codex_companion_endpoints() {
let template = fixed_provider_template("codex").expect("codex template should exist");
assert_eq!(template.base_url, "https://chatgpt.com/backend-api/codex");
assert_eq!(template.version, 1);
assert_eq!(
template
.endpoints
.iter()
.map(|item| item.api_format)
.collect::<Vec<_>>(),
vec![
"openai:responses",
"openai:responses:compact",
"openai:search",
"openai:image"
]
);
let image_template =
fixed_provider_endpoint_template_by_api_format("codex", "openai:image")
.expect("codex image endpoint should exist");
assert!(image_template.config_defaults.is_empty());
let search_template =
fixed_provider_endpoint_template_by_api_format("codex", "openai:search")
.expect("codex search endpoint should exist");
assert!(search_template.config_defaults.is_empty());
}
#[test]
fn chatgpt_web_fixed_provider_template_only_exposes_openai_image() {
let template =
fixed_provider_template("chatgpt_web").expect("chatgpt_web template should exist");
assert_eq!(template.base_url, "https://chatgpt.com");
assert_eq!(template.version, 1);
assert_eq!(
template
.endpoints
.iter()
.map(|item| item.api_format)
.collect::<Vec<_>>(),
vec!["openai:image"]
);
let image_template =
fixed_provider_endpoint_template_by_api_format("chatgpt_web", "openai:image")
.expect("chatgpt_web image endpoint should exist");
assert_eq!(
image_template
.config_defaults
.iter()
.map(|item| (item.key, item.value))
.collect::<Vec<_>>(),
vec![(
"upstream_stream_policy",
FixedProviderEndpointConfigValue::String("force_stream")
)]
);
}
#[test]
fn grok_fixed_provider_template_exposes_chat_responses_messages_and_image() {
let template = fixed_provider_template("grok").expect("grok template should exist");
assert_eq!(template.base_url, "https://grok.com");
assert_eq!(template.version, 1);
assert_eq!(
template
.endpoints
.iter()
.map(|item| item.api_format)
.collect::<Vec<_>>(),
vec![
"openai:chat",
"openai:responses",
"claude:messages",
"openai:image"
]
);
assert!(!template.runtime_policy.supports_model_fetch);
assert!(!template.runtime_policy.supports_local_openai_chat_transport);
assert!(!template.runtime_policy.supports_local_same_format_transport);
}
#[test]
fn gemini_cli_fixed_provider_template_uses_v1internal_endpoint_path() {
let template =
fixed_provider_template("gemini_cli").expect("gemini_cli template should exist");
assert_eq!(template.base_url, "https://cloudcode-pa.googleapis.com");
assert_eq!(template.version, 3);
let endpoint =
fixed_provider_endpoint_template_by_api_format("gemini_cli", "gemini:generate_content")
.expect("gemini_cli generateContent endpoint should exist");
assert_eq!(endpoint.custom_path, Some("/v1internal:{action}"));
assert_eq!(
endpoint
.config_defaults
.iter()
.map(|item| (item.key, item.value))
.collect::<Vec<_>>(),
vec![(
"upstream_stream_policy",
FixedProviderEndpointConfigValue::String("auto")
)]
);
}
#[test]
fn antigravity_fixed_provider_template_uses_daily_cloudcode_endpoint() {
let template =
fixed_provider_template("antigravity").expect("antigravity template should exist");
assert_eq!(
template.base_url,
"https://daily-cloudcode-pa.googleapis.com"
);
assert_eq!(template.version, 2);
let endpoint = fixed_provider_endpoint_template_by_api_format(
"antigravity",
"gemini:generate_content",
)
.expect("antigravity generateContent endpoint should exist");
assert_eq!(endpoint.custom_path, None);
}
#[test]
fn windsurf_fixed_provider_template_exposes_openai_chat() {
let template = fixed_provider_template("windsurf").expect("windsurf template should exist");
assert_eq!(template.provider_type, "windsurf");
assert_eq!(template.base_url, "https://server.codeium.com");
assert_eq!(template.version, 1);
assert_eq!(
template
.endpoints
.iter()
.map(|item| item.api_format)
.collect::<Vec<_>>(),
vec!["openai:chat"]
);
assert!(
fixed_provider_endpoint_template_by_api_format("windsurf", "openai:chat").is_some()
);
let policy = provider_runtime_policy("windsurf");
assert!(policy.fixed_provider);
assert!(policy.enable_format_conversion_by_default);
assert!(policy.oauth_is_bearer_like);
assert!(!policy.supports_model_fetch);
assert!(!policy.supports_local_same_format_transport);
}
#[test]
fn windsurf_admin_oauth_template_is_advertised() {
let template =
provider_type_admin_oauth_template("windsurf").expect("windsurf oauth template");
assert_eq!(template.provider_type, "windsurf");
assert_eq!(template.display_name, "Windsurf");
assert_eq!(
template.authorize_url,
"https://windsurf.com/windsurf/signin"
);
assert_eq!(template.redirect_uri, "show-auth-token");
assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf"));
}
#[test]
fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() {
assert!(fixed_provider_key_inherits_api_formats(
"codex", "oauth", None
));
assert!(fixed_provider_key_inherits_api_formats(
"chatgpt_web",
"oauth",
None
));
assert!(fixed_provider_key_inherits_api_formats(
"kiro",
"bearer",
Some("{}")
));
assert!(fixed_provider_key_inherits_api_formats(
"vertex_ai",
"service_account",
None
));
assert!(!fixed_provider_key_inherits_api_formats(
"kiro", "bearer", None
));
assert!(!fixed_provider_key_inherits_api_formats(
"custom", "oauth", None
));
}
#[test]
fn kiro_allows_auth_channel_mismatch_by_default() {
let policy = provider_runtime_policy("kiro");
assert!(policy.fixed_provider);
assert!(policy.enable_format_conversion_by_default);
assert!(policy.oauth_is_bearer_like);
assert!(policy.supports_model_fetch);
assert!(!policy.supports_local_openai_chat_transport);
assert!(!policy.supports_local_same_format_transport);
assert!(policy.key_inherits_api_formats("oauth", None));
assert!(policy.key_inherits_api_formats("bearer", Some("{}")));
assert!(!policy.key_inherits_api_formats("bearer", None));
assert!(provider_type_allows_auth_channel_mismatch_by_default(
"kiro"
));
assert!(provider_type_allows_auth_channel_mismatch_by_default(
" KIRO "
));
assert!(!provider_type_allows_auth_channel_mismatch_by_default(
"claude_code"
));
assert!(!provider_type_allows_auth_channel_mismatch_by_default(
"custom"
));
}
#[test]
fn runtime_policy_preserves_other_fixed_provider_behavior() {
let codex = provider_runtime_policy("codex");
assert!(codex.fixed_provider);
assert!(codex.enable_format_conversion_by_default);
assert!(!codex.oauth_is_bearer_like);
assert!(!codex.supports_model_fetch);
assert!(!codex.supports_local_openai_chat_transport);
assert!(codex.supports_local_same_format_transport);
let gemini_cli = provider_runtime_policy("gemini_cli");
assert!(gemini_cli.fixed_provider);
assert!(!gemini_cli.enable_format_conversion_by_default);
assert!(provider_type_oauth_is_bearer_like("gemini_cli"));
assert!(gemini_cli.supports_model_fetch);
assert!(!gemini_cli.supports_local_openai_chat_transport);
assert!(gemini_cli.supports_local_same_format_transport);
}
#[test]
fn chatgpt_web_does_not_use_generic_same_format_transport() {
assert!(!provider_type_supports_local_same_format_transport(
"chatgpt_web"
));
}
#[test]
fn provider_type_supports_only_matching_embedding_formats() {
for (provider_type, api_format) in [
("openai", "openai:embedding"),
("custom", "openai:embedding"),
("gemini", "gemini:embedding"),
("google", "gemini:embedding"),
("vertex_ai", "gemini:embedding"),
("jina", "jina:embedding"),
("doubao", "doubao:embedding"),
("volcengine", "doubao:embedding"),
("aliyun", "aliyun:multimodal_embedding"),
("dashscope", "aliyun:multimodal_embedding"),
] {
assert!(
provider_type_supports_local_embedding_transport(provider_type, api_format),
"{provider_type} should support {api_format}"
);
}
for (provider_type, api_format) in [
("openai", "gemini:embedding"),
("gemini", "openai:embedding"),
("vertex_ai", "openai:embedding"),
("jina", "doubao:embedding"),
("doubao", "jina:embedding"),
("aliyun", "openai:embedding"),
("openai", "aliyun:multimodal_embedding"),
("claude_code", "openai:embedding"),
("openai", "openai:chat"),
] {
assert!(
!provider_type_supports_local_embedding_transport(provider_type, api_format),
"{provider_type} should not support {api_format}"
);
}
assert!(provider_type_supports_local_embedding_transport(
" Google ",
"GEMINI:EMBEDDING"
));
}
#[test]
fn vertex_fixed_provider_template_includes_gemini_embedding_endpoint() {
let template =
fixed_provider_template("vertex_ai").expect("vertex_ai template should exist");
assert_eq!(
template
.endpoints
.iter()
.map(|item| item.api_format)
.collect::<Vec<_>>(),
vec![
"gemini:generate_content",
"gemini:embedding",
"claude:messages",
]
);
assert!(
fixed_provider_endpoint_template_by_api_format("vertex_ai", "gemini:embedding")
.is_some()
);
}
}
@@ -0,0 +1,416 @@
use serde_json::{Map, Value};
use crate::snapshot::GatewayProviderTransportSnapshot;
use crate::vertex::is_vertex_transport_context;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TransportRequestBodySemanticsError {
message: &'static str,
}
impl TransportRequestBodySemanticsError {
const fn new(message: &'static str) -> Self {
Self { message }
}
pub const fn message(&self) -> &'static str {
self.message
}
}
impl std::fmt::Display for TransportRequestBodySemanticsError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.message)
}
}
impl std::error::Error for TransportRequestBodySemanticsError {}
pub fn apply_transport_request_body_semantics(
provider_request_body: &mut Value,
transport: &GatewayProviderTransportSnapshot,
provider_api_format: &str,
) -> Result<(), TransportRequestBodySemanticsError> {
let provider_api_format = aether_ai_formats::normalize_api_format_alias(provider_api_format);
if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) {
apply_vertex_gemini_embedding_body_semantics(provider_request_body)?;
}
Ok(())
}
fn apply_vertex_gemini_embedding_body_semantics(
provider_request_body: &mut Value,
) -> Result<(), TransportRequestBodySemanticsError> {
let object = provider_request_body.as_object_mut().ok_or_else(|| {
TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding request body must be a JSON object",
)
})?;
if object.contains_key("instances") {
validate_existing_vertex_predict_body(object)?;
object.remove("model");
return Ok(());
}
let next = build_vertex_predict_body_from_gemini_embedding_object(object)?;
*object = next;
Ok(())
}
fn build_vertex_predict_body_from_gemini_embedding_object(
object: &Map<String, Value>,
) -> Result<Map<String, Value>, TransportRequestBodySemanticsError> {
if let Some(requests) = object.get("requests") {
if object.keys().any(|key| key != "requests") {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding batch body cannot mix requests with other top-level fields",
));
}
let request_items = requests.as_array().ok_or_else(|| {
TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding requests must be an array",
)
})?;
let request_objects = request_items
.iter()
.map(Value::as_object)
.collect::<Option<Vec<_>>>()
.ok_or_else(|| {
TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding requests must be an array of objects",
)
})?;
if request_objects.is_empty() {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding requests must contain at least one item",
));
}
return build_vertex_predict_body_from_gemini_embedding_items(&request_objects);
}
build_vertex_predict_body_from_gemini_embedding_items(&[object])
}
fn build_vertex_predict_body_from_gemini_embedding_items(
items: &[&Map<String, Value>],
) -> Result<Map<String, Value>, TransportRequestBodySemanticsError> {
if items.iter().any(|item| {
item.keys().any(|key| {
!matches!(
key.as_str(),
"model"
| "content"
| "taskType"
| "title"
| "outputDimensionality"
| "autoTruncate"
)
})
}) {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding body contains fields that cannot be mapped to predict instances",
));
}
let instances = items
.iter()
.map(|item| build_vertex_predict_instance(item))
.collect::<Option<Vec<_>>>()
.ok_or_else(|| {
TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding body must contain text content parts",
)
})?;
if instances.is_empty() {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding body must contain at least one instance",
));
}
let mut output = Map::new();
output.insert("instances".to_string(), Value::Array(instances));
let mut parameters = Map::new();
insert_shared_parameter(items, &mut parameters, "outputDimensionality")?;
insert_shared_parameter(items, &mut parameters, "autoTruncate")?;
if !parameters.is_empty() {
output.insert("parameters".to_string(), Value::Object(parameters));
}
Ok(output)
}
fn validate_existing_vertex_predict_body(
object: &Map<String, Value>,
) -> Result<(), TransportRequestBodySemanticsError> {
if object
.keys()
.any(|key| !matches!(key.as_str(), "model" | "instances" | "parameters"))
{
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding predict body contains unsupported top-level fields",
));
}
let Some(instances) = object.get("instances").and_then(Value::as_array) else {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding predict body must contain an instances array",
));
};
if instances.is_empty() {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding predict body must contain at least one instance",
));
}
if object
.get("parameters")
.is_some_and(|parameters| !parameters.is_object())
{
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding predict parameters must be an object",
));
}
Ok(())
}
fn build_vertex_predict_instance(item: &Map<String, Value>) -> Option<Value> {
let content = gemini_embedding_content_text(item.get("content")?)?;
let mut instance = Map::new();
instance.insert("content".to_string(), Value::String(content));
if let Some(task_type) = item.get("taskType") {
instance.insert(
"task_type".to_string(),
Value::String(task_type.as_str()?.to_string()),
);
}
if let Some(title) = item.get("title") {
instance.insert(
"title".to_string(),
Value::String(title.as_str()?.to_string()),
);
}
Some(Value::Object(instance))
}
fn gemini_embedding_content_text(content: &Value) -> Option<String> {
let parts = content
.as_object()?
.get("parts")?
.as_array()?
.iter()
.filter_map(|part| part.as_object()?.get("text")?.as_str())
.filter(|text| !text.trim().is_empty())
.collect::<Vec<_>>();
if parts.is_empty() {
return None;
}
Some(parts.join(""))
}
fn insert_shared_parameter(
items: &[&Map<String, Value>],
parameters: &mut Map<String, Value>,
key: &str,
) -> Result<(), TransportRequestBodySemanticsError> {
let mut value: Option<Value> = None;
for item in items {
let Some(next) = item.get(key) else {
continue;
};
match &value {
Some(current) if current != next => {
return Err(TransportRequestBodySemanticsError::new(
"Vertex Gemini embedding batch items must use the same shared parameters",
));
}
None => value = Some(next.clone()),
_ => {}
}
}
if let Some(value) = value {
parameters.insert(key.to_string(), value);
}
Ok(())
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::apply_transport_request_body_semantics;
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport(provider_type: &str, base_url: &str) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "provider".to_string(),
provider_type: provider_type.to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "gemini:embedding".to_string(),
api_family: Some("gemini".to_string()),
endpoint_kind: Some("embedding".to_string()),
is_active: true,
base_url: base_url.to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: Some(vec!["gemini:embedding".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn vertex_gemini_embedding_single_body_uses_predict_contract() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"model": "gemini-embedding-2",
"content": {"parts": [{"text": "hello"}]},
"taskType": "RETRIEVAL_QUERY",
"outputDimensionality": 768
});
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect("body semantics should apply");
assert!(body.get("model").is_none());
assert!(body.get("content").is_none());
assert_eq!(body["instances"][0]["content"], "hello");
assert_eq!(body["instances"][0]["task_type"], "RETRIEVAL_QUERY");
assert_eq!(body["parameters"]["outputDimensionality"], 768);
}
#[test]
fn gemini_api_embedding_single_body_keeps_model_for_developer_api() {
let transport =
sample_transport("gemini", "https://generativelanguage.googleapis.com/v1beta");
let mut body = json!({
"model": "gemini-embedding-2",
"content": {"parts": [{"text": "hello"}]}
});
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect("developer API body should pass through");
assert_eq!(body["model"], "gemini-embedding-2");
}
#[test]
fn vertex_gemini_embedding_batch_body_uses_predict_instances() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"requests": [
{
"model": "models/gemini-embedding-2",
"content": {"parts": [{"text": "hello"}]}
},
{
"model": "models/gemini-embedding-2",
"content": {"parts": [{"text": "world"}]}
}
]
});
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect("batch body semantics should apply");
assert!(body.get("requests").is_none());
assert_eq!(body["instances"][0]["content"], "hello");
assert_eq!(body["instances"][1]["content"], "world");
}
#[test]
fn vertex_gemini_embedding_existing_predict_body_removes_duplicate_model() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"model": "gemini-embedding-2",
"instances": [
{"content": "hello", "task_type": "RETRIEVAL_QUERY"}
],
"parameters": {
"outputDimensionality": 768
}
});
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect("existing predict body should be accepted");
assert!(body.get("model").is_none());
assert_eq!(body["instances"][0]["content"], "hello");
assert_eq!(body["parameters"]["outputDimensionality"], 768);
}
#[test]
fn vertex_gemini_embedding_existing_predict_body_rejects_unconsumed_fields() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"model": "gemini-embedding-2",
"instances": [
{"content": "hello"}
],
"input": "this field would not be consumed by Vertex predict"
});
let error =
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect_err("predict body must not carry unconsumed OpenAI fields");
assert!(error.message().contains("unsupported top-level fields"));
assert!(body.get("model").is_some());
}
#[test]
fn vertex_gemini_embedding_rejects_unconverted_openai_body() {
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
let mut body = json!({
"model": "gemini-embedding-2",
"input": "hello"
});
let error =
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
.expect_err("OpenAI embedding body must not be sent to Vertex native predict");
assert!(error.message().contains("cannot be mapped"));
assert!(body.get("input").is_some());
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1 @@
pub use aether_ai_formats::provider_compat::proxy::rules::*;
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,249 @@
use aether_crypto::{decrypt_python_fernet_ciphertext, looks_like_python_fernet_ciphertext};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data_contracts::DataLayerError;
use super::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey, GatewayProviderTransportProvider,
};
pub(super) fn map_provider(
provider: StoredProviderCatalogProvider,
) -> GatewayProviderTransportProvider {
GatewayProviderTransportProvider {
id: provider.id,
name: provider.name,
provider_type: provider.provider_type,
website: provider.website,
is_active: provider.is_active,
keep_priority_on_conversion: provider.keep_priority_on_conversion,
enable_format_conversion: provider.enable_format_conversion,
concurrent_limit: provider.concurrent_limit,
max_retries: provider.max_retries,
proxy: normalize_optional_json(provider.proxy),
request_timeout_secs: provider.request_timeout_secs,
stream_first_byte_timeout_secs: provider.stream_first_byte_timeout_secs,
config: normalize_optional_json(provider.config),
}
}
pub(super) fn map_endpoint(
endpoint: StoredProviderCatalogEndpoint,
) -> GatewayProviderTransportEndpoint {
GatewayProviderTransportEndpoint {
id: endpoint.id,
provider_id: endpoint.provider_id,
api_format: endpoint.api_format,
api_family: endpoint.api_family,
endpoint_kind: endpoint.endpoint_kind,
is_active: endpoint.is_active,
base_url: endpoint.base_url,
header_rules: normalize_optional_json(endpoint.header_rules),
body_rules: normalize_optional_json(endpoint.body_rules),
max_retries: endpoint.max_retries,
custom_path: endpoint.custom_path,
config: normalize_optional_json(endpoint.config),
format_acceptance_config: normalize_optional_json(endpoint.format_acceptance_config),
proxy: normalize_optional_json(endpoint.proxy),
}
}
pub(super) fn map_key(
key: StoredProviderCatalogKey,
encryption_key: &str,
fallback_encryption_keys: &[String],
) -> Result<GatewayProviderTransportKey, DataLayerError> {
let decrypted_api_key = key
.encrypted_api_key
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|ciphertext| {
decrypt_secret(
encryption_key,
fallback_encryption_keys,
ciphertext,
"provider_api_keys.api_key",
)
})
.transpose()?
.unwrap_or_default();
let decrypted_auth_config = key
.encrypted_auth_config
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|ciphertext| {
decrypt_secret(
encryption_key,
fallback_encryption_keys,
ciphertext,
"provider_api_keys.auth_config",
)
})
.transpose()?;
Ok(GatewayProviderTransportKey {
id: key.id,
provider_id: key.provider_id,
name: key.name,
auth_type: key.auth_type,
is_active: key.is_active,
api_formats: normalize_string_list(
normalize_optional_json(key.api_formats),
"provider_api_keys.api_formats",
)?,
auth_type_by_format: normalize_optional_json(key.auth_type_by_format),
allow_auth_channel_mismatch_formats: normalize_optional_json(
key.allow_auth_channel_mismatch_formats,
),
allowed_models: normalize_string_list(
normalize_optional_json(key.allowed_models),
"provider_api_keys.allowed_models",
)?,
capabilities: normalize_optional_json(key.capabilities),
rate_multipliers: normalize_optional_json(key.rate_multipliers),
global_priority_by_format: normalize_optional_json(key.global_priority_by_format),
expires_at_unix_secs: key.expires_at_unix_secs,
proxy: normalize_optional_json(key.proxy),
fingerprint: normalize_optional_json(key.fingerprint),
upstream_metadata: normalize_optional_json(key.upstream_metadata),
decrypted_api_key,
decrypted_auth_config,
})
}
fn normalize_optional_json(value: Option<serde_json::Value>) -> Option<serde_json::Value> {
match value {
Some(serde_json::Value::Null) | None => None,
Some(value) => Some(value),
}
}
fn decrypt_secret(
encryption_key: &str,
fallback_encryption_keys: &[String],
ciphertext: &str,
field_name: &str,
) -> Result<String, DataLayerError> {
if should_use_plaintext_secret(ciphertext, field_name) {
return Ok(ciphertext.trim().to_string());
}
match decrypt_python_fernet_ciphertext(encryption_key, ciphertext) {
Ok(value) => Ok(value),
Err(error) => {
for fallback_encryption_key in fallback_encryption_keys {
if let Ok(value) =
decrypt_python_fernet_ciphertext(fallback_encryption_key, ciphertext)
{
return Ok(value);
}
}
Err(DataLayerError::UnexpectedValue(format!(
"failed to decrypt {field_name}: {error}"
)))
}
}
}
pub(super) fn fallback_encryption_keys(primary_encryption_key: &str) -> Vec<String> {
let mut keys = Vec::new();
for env_key in ["AETHER_GATEWAY_DATA_ENCRYPTION_KEY", "ENCRYPTION_KEY"] {
let Ok(value) = std::env::var(env_key) else {
continue;
};
let value = value.trim();
if value.is_empty()
|| value == primary_encryption_key
|| keys.iter().any(|existing| existing == value)
{
continue;
}
keys.push(value.to_string());
}
keys
}
fn should_use_plaintext_secret(ciphertext: &str, field_name: &str) -> bool {
let ciphertext = ciphertext.trim();
if ciphertext.is_empty() {
return false;
}
match field_name {
"provider_api_keys.api_key" => {
if ciphertext.starts_with('{') || ciphertext.starts_with('[') {
return false;
}
!looks_like_python_fernet_ciphertext(ciphertext)
}
"provider_api_keys.auth_config" => {
if ciphertext.starts_with('{') || ciphertext.starts_with('[') {
return true;
}
false
}
_ => false,
}
}
fn normalize_string_list(
raw: Option<serde_json::Value>,
field_name: &str,
) -> Result<Option<Vec<String>>, DataLayerError> {
let Some(raw) = raw else {
return Ok(None);
};
normalize_string_list_value(&raw, field_name)
}
fn normalize_string_list_value(
raw: &serde_json::Value,
field_name: &str,
) -> Result<Option<Vec<String>>, DataLayerError> {
match raw {
serde_json::Value::Null => Ok(None),
serde_json::Value::Array(items) => normalize_string_list_array(items, field_name).map(Some),
serde_json::Value::String(raw) => normalize_embedded_string_list(raw, field_name),
_ => Err(DataLayerError::UnexpectedValue(format!(
"{field_name} is not a JSON array"
))),
}
}
fn normalize_embedded_string_list(
raw: &str,
field_name: &str,
) -> Result<Option<Vec<String>>, DataLayerError> {
let raw = raw.trim();
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
return Ok(None);
}
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
return normalize_string_list_value(&decoded, field_name);
}
Ok(Some(vec![raw.to_string()]))
}
fn normalize_string_list_array(
items: &[serde_json::Value],
field_name: &str,
) -> Result<Vec<String>, DataLayerError> {
let mut values = Vec::with_capacity(items.len());
for item in items {
let Some(value) = item.as_str() else {
return Err(DataLayerError::UnexpectedValue(format!(
"{field_name} contains a non-string item"
)));
};
let value = value.trim();
if !value.is_empty() {
values.push(value.to_string());
}
}
Ok(values)
}
@@ -0,0 +1,703 @@
use std::collections::BTreeMap;
use serde_json::Value;
use crate::auth::{
build_claude_passthrough_headers, build_complete_passthrough_headers_with_auth,
build_openai_passthrough_headers, build_passthrough_headers, ensure_upstream_auth_header,
};
use crate::headers::force_identity_accept_encoding;
use crate::rules::{
apply_local_body_rules, apply_local_body_rules_with_request_headers,
apply_local_header_rules_with_request_headers,
};
use crate::snapshot::GatewayProviderTransportSnapshot;
use crate::url::{build_openai_chat_url, build_openai_responses_url};
use crate::vertex::uses_vertex_api_key_query_auth;
#[derive(Debug, Clone, Copy)]
pub struct StandardProviderRequestHeadersInput<'a> {
pub transport: &'a GatewayProviderTransportSnapshot,
pub provider_api_format: &'a str,
pub same_format: bool,
pub headers: &'a http::HeaderMap,
pub auth_header: &'a str,
pub auth_value: &'a str,
pub extra_headers: &'a BTreeMap<String, String>,
pub header_rules: Option<&'a Value>,
pub provider_request_body: &'a Value,
pub original_request_body: &'a Value,
pub upstream_is_stream: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StandardProviderRequestHeaders {
pub headers: BTreeMap<String, String>,
pub auth_header: String,
pub auth_value: String,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StandardPlanFallbackAcceptPolicy {
None,
TextEventStreamIfStreaming,
TextEventStreamIfStreamingOrWildcard,
TextEventStreamRequired,
ProviderEventStreamIfMissing,
}
#[derive(Debug)]
pub struct StandardPlanFallbackHeadersInput<'a> {
pub request_headers: &'a http::HeaderMap,
pub existing_provider_request_headers: BTreeMap<String, String>,
pub auth_header: Option<&'a str>,
pub auth_value: Option<&'a str>,
pub extra_headers: &'a BTreeMap<String, String>,
pub content_type: Option<&'a str>,
pub provider_api_format: &'a str,
pub client_api_format: &'a str,
pub upstream_is_stream: bool,
pub build_from_request_when_empty: bool,
pub accept_policy: StandardPlanFallbackAcceptPolicy,
}
pub fn build_standard_plan_fallback_openai_chat_url(
upstream_base_url: &str,
request_query: Option<&str>,
) -> String {
build_openai_chat_url(upstream_base_url, request_query)
}
pub fn build_standard_plan_fallback_openai_responses_url(
upstream_base_url: &str,
request_query: Option<&str>,
compact: bool,
) -> String {
build_openai_responses_url(upstream_base_url, request_query, compact)
}
pub fn build_standard_plan_fallback_headers(
input: StandardPlanFallbackHeadersInput<'_>,
) -> BTreeMap<String, String> {
let auth_pair = input.auth_header.zip(input.auth_value);
let mut headers = if !input.existing_provider_request_headers.is_empty() {
input.existing_provider_request_headers
} else if input.build_from_request_when_empty {
match auth_pair {
Some((auth_header, auth_value))
if input.provider_api_format == input.client_api_format =>
{
build_complete_passthrough_headers_with_auth(
input.request_headers,
auth_header,
auth_value,
input.extra_headers,
input.content_type,
)
}
Some((auth_header, auth_value)) if input.provider_api_format.starts_with("claude:") => {
build_claude_passthrough_headers(
input.request_headers,
auth_header,
auth_value,
input.extra_headers,
input.content_type,
)
}
Some((auth_header, auth_value)) => build_openai_passthrough_headers(
input.request_headers,
auth_header,
auth_value,
input.extra_headers,
input.content_type,
),
None => build_passthrough_headers(
input.request_headers,
input.extra_headers,
input.content_type,
),
}
} else {
input.existing_provider_request_headers
};
if let Some((auth_header, auth_value)) = auth_pair {
ensure_upstream_auth_header(&mut headers, auth_header, auth_value);
}
match input.accept_policy {
StandardPlanFallbackAcceptPolicy::None => {}
StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreaming => {
if input.upstream_is_stream {
headers
.entry("accept".to_string())
.or_insert_with(|| "text/event-stream".to_string());
}
}
StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreamingOrWildcard => {
if input.upstream_is_stream {
set_accept_if_missing_or_wildcard(&mut headers, "text/event-stream");
}
}
StandardPlanFallbackAcceptPolicy::TextEventStreamRequired => {
headers.insert("accept".to_string(), "text/event-stream".to_string());
}
StandardPlanFallbackAcceptPolicy::ProviderEventStreamIfMissing => {
set_accept_if_missing_or_wildcard(&mut headers, "application/vnd.amazon.eventstream");
}
}
if input.upstream_is_stream {
force_identity_accept_encoding(&mut headers);
}
headers
}
fn set_accept_if_missing_or_wildcard(headers: &mut BTreeMap<String, String>, value: &str) {
let Some(existing_key) = headers
.keys()
.find(|key| key.eq_ignore_ascii_case("accept"))
.cloned()
else {
headers.insert("accept".to_string(), value.to_string());
return;
};
if headers
.get(&existing_key)
.is_some_and(|existing_value| accept_is_wildcard_only(existing_value))
{
headers.insert(existing_key, value.to_string());
}
}
fn accept_is_wildcard_only(value: &str) -> bool {
let mut saw_value = false;
for raw_part in value.split(',') {
let media_type = raw_part.trim().split(';').next().unwrap_or_default().trim();
if media_type.is_empty() {
continue;
}
saw_value = true;
if media_type != "*/*" {
return false;
}
}
saw_value
}
pub fn apply_standard_provider_request_body_rules(
mut provider_request_body: Value,
body_rules: Option<&Value>,
original_request_body: &Value,
) -> Option<Value> {
if !apply_local_body_rules(
&mut provider_request_body,
body_rules,
Some(original_request_body),
) {
return None;
}
Some(provider_request_body)
}
pub fn apply_standard_provider_request_body_rules_with_request_headers(
mut provider_request_body: Value,
body_rules: Option<&Value>,
original_request_body: &Value,
request_headers: &http::HeaderMap,
) -> Option<Value> {
if !apply_local_body_rules_with_request_headers(
&mut provider_request_body,
body_rules,
Some(original_request_body),
Some(request_headers),
) {
return None;
}
Some(provider_request_body)
}
pub fn build_standard_provider_request_headers(
input: StandardProviderRequestHeadersInput<'_>,
) -> Option<StandardProviderRequestHeaders> {
let uses_vertex_query_auth =
uses_vertex_api_key_query_auth(input.transport, input.provider_api_format);
let mut headers = if input.same_format {
build_complete_passthrough_headers_with_auth(
input.headers,
input.auth_header,
input.auth_value,
input.extra_headers,
Some("application/json"),
)
} else if input.provider_api_format.starts_with("claude:") {
build_claude_passthrough_headers(
input.headers,
input.auth_header,
input.auth_value,
input.extra_headers,
Some("application/json"),
)
} else {
build_openai_passthrough_headers(
input.headers,
input.auth_header,
input.auth_value,
input.extra_headers,
Some("application/json"),
)
};
crate::apply_local_auth_config_header_overrides(
&mut headers,
input.transport.key.decrypted_auth_config.as_deref(),
);
let protected_headers = if uses_vertex_query_auth {
&["content-type"][..]
} else {
&[input.auth_header, "content-type"][..]
};
if !apply_local_header_rules_with_request_headers(
&mut headers,
input.header_rules,
protected_headers,
input.provider_request_body,
Some(input.original_request_body),
Some(input.headers),
) {
return None;
}
let (auth_header, auth_value) = if uses_vertex_query_auth {
headers.remove("x-goog-api-key");
(String::new(), String::new())
} else {
ensure_upstream_auth_header(&mut headers, input.auth_header, input.auth_value);
(input.auth_header.to_string(), input.auth_value.to_string())
};
if input.upstream_is_stream {
headers
.entry("accept".to_string())
.or_insert_with(|| "text/event-stream".to_string());
force_identity_accept_encoding(&mut headers);
}
Some(StandardProviderRequestHeaders {
headers,
auth_header,
auth_value,
})
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use http::HeaderMap;
use serde_json::json;
use super::{
apply_standard_provider_request_body_rules, build_standard_plan_fallback_headers,
build_standard_plan_fallback_openai_chat_url,
build_standard_plan_fallback_openai_responses_url, build_standard_provider_request_headers,
StandardPlanFallbackAcceptPolicy, StandardPlanFallbackHeadersInput,
StandardProviderRequestHeadersInput,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport(api_format: &str) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Provider".to_string(),
provider_type: "openai".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: api_format.to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://api.example.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "bearer".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn builds_same_format_headers_with_complete_passthrough_and_stream_accept() {
let mut request_headers = HeaderMap::new();
request_headers.insert("x-client", "demo".parse().expect("header"));
request_headers.insert(
http::header::ACCEPT_ENCODING,
"gzip, br".parse().expect("header"),
);
let transport = sample_transport("openai:chat");
let resolved =
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
transport: &transport,
provider_api_format: "openai:chat",
same_format: true,
headers: &request_headers,
auth_header: "authorization",
auth_value: "Bearer secret",
extra_headers: &BTreeMap::new(),
header_rules: None,
provider_request_body: &json!({"model":"gpt-5"}),
original_request_body: &json!({"model":"gpt-5"}),
upstream_is_stream: true,
})
.expect("headers should build");
assert_eq!(resolved.auth_header, "authorization");
assert_eq!(resolved.auth_value, "Bearer secret");
assert_eq!(
resolved.headers.get("authorization"),
Some(&"Bearer secret".to_string())
);
assert_eq!(
resolved.headers.get("accept"),
Some(&"text/event-stream".to_string())
);
assert_eq!(
resolved.headers.get("accept-encoding"),
Some(&"identity".to_string())
);
assert_eq!(resolved.headers.get("x-client"), Some(&"demo".to_string()));
}
#[test]
fn builds_sync_headers_preserving_supported_accept_encoding() {
let mut request_headers = HeaderMap::new();
request_headers.insert(
http::header::ACCEPT_ENCODING,
"gzip, br".parse().expect("header"),
);
let transport = sample_transport("openai:chat");
let resolved =
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
transport: &transport,
provider_api_format: "openai:chat",
same_format: true,
headers: &request_headers,
auth_header: "authorization",
auth_value: "Bearer secret",
extra_headers: &BTreeMap::new(),
header_rules: None,
provider_request_body: &json!({"model":"gpt-5"}),
original_request_body: &json!({"model":"gpt-5"}),
upstream_is_stream: false,
})
.expect("headers should build");
assert_eq!(
resolved.headers.get("accept-encoding"),
Some(&"gzip".to_string())
);
}
#[test]
fn applies_header_rules_after_base_headers() {
let transport = sample_transport("claude:messages");
let resolved =
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
transport: &transport,
provider_api_format: "claude:messages",
same_format: false,
headers: &HeaderMap::new(),
auth_header: "x-api-key",
auth_value: "secret",
extra_headers: &BTreeMap::new(),
header_rules: Some(&json!([
{"action":"set","key":"x-route","value":"standard"}
])),
provider_request_body: &json!({"model":"claude"}),
original_request_body: &json!({"model":"claude"}),
upstream_is_stream: false,
})
.expect("headers should build");
assert_eq!(
resolved.headers.get("x-api-key"),
Some(&"secret".to_string())
);
assert_eq!(
resolved.headers.get("x-route"),
Some(&"standard".to_string())
);
assert_eq!(
resolved.headers.get("anthropic-version"),
Some(&"2023-06-01".to_string())
);
}
#[test]
fn endpoint_header_rules_do_not_override_protected_authorization() {
let mut transport = sample_transport("openai:responses");
transport.endpoint.header_rules = Some(json!([
{"action":"set","key":"authorization","value":"Bearer imported-session"}
]));
let resolved =
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
transport: &transport,
provider_api_format: "openai:responses",
same_format: true,
headers: &HeaderMap::new(),
auth_header: "authorization",
auth_value: "Bearer jwt-access-token",
extra_headers: &BTreeMap::new(),
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &json!({"model":"gpt-test"}),
original_request_body: &json!({"model":"gpt-test"}),
upstream_is_stream: false,
})
.expect("headers should build");
assert_eq!(
resolved.headers.get("authorization"),
Some(&"Bearer jwt-access-token".to_string())
);
}
#[test]
fn auth_config_headers_can_override_authorization_when_refresh_token_is_present() {
let mut transport = sample_transport("openai:responses");
transport.key.decrypted_auth_config = Some(
json!({
"refresh_token": "rt-1",
"headers": {
"authorization": "Bearer imported-session"
}
})
.to_string(),
);
let resolved =
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
transport: &transport,
provider_api_format: "openai:responses",
same_format: true,
headers: &HeaderMap::new(),
auth_header: "authorization",
auth_value: "Bearer refreshed-access-token",
extra_headers: &BTreeMap::new(),
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &json!({"model":"gpt-test"}),
original_request_body: &json!({"model":"gpt-test"}),
upstream_is_stream: false,
})
.expect("headers should build");
assert_eq!(
resolved.headers.get("authorization"),
Some(&"Bearer imported-session".to_string())
);
}
#[test]
fn applies_standard_body_rules_to_surface_built_body() {
let body = apply_standard_provider_request_body_rules(
json!({"model":"gpt-5"}),
Some(&json!([
{"action":"set","path":"metadata.source","value":"standard"}
])),
&json!({"model":"client"}),
)
.expect("body rules should apply");
assert_eq!(body["metadata"]["source"], json!("standard"));
}
#[test]
fn builds_plan_fallback_headers_from_request_when_enabled() {
let mut request_headers = HeaderMap::new();
request_headers.insert("x-client", "demo".parse().expect("header"));
let headers = build_standard_plan_fallback_headers(StandardPlanFallbackHeadersInput {
request_headers: &request_headers,
existing_provider_request_headers: BTreeMap::new(),
auth_header: Some("authorization"),
auth_value: Some("Bearer secret"),
extra_headers: &BTreeMap::new(),
content_type: Some("application/json"),
provider_api_format: "openai:chat",
client_api_format: "openai:chat",
upstream_is_stream: true,
build_from_request_when_empty: true,
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreaming,
});
assert_eq!(
headers.get("authorization"),
Some(&"Bearer secret".to_string())
);
assert_eq!(headers.get("x-client"), Some(&"demo".to_string()));
assert_eq!(
headers.get("accept"),
Some(&"text/event-stream".to_string())
);
}
#[test]
fn stream_fallback_headers_treat_wildcard_accept_as_absent() {
let mut request_headers = HeaderMap::new();
request_headers.insert(http::header::ACCEPT, "*/*".parse().expect("header"));
request_headers.insert(
http::header::ACCEPT_ENCODING,
"gzip, br".parse().expect("header"),
);
let headers = build_standard_plan_fallback_headers(StandardPlanFallbackHeadersInput {
request_headers: &request_headers,
existing_provider_request_headers: BTreeMap::new(),
auth_header: Some("authorization"),
auth_value: Some("Bearer secret"),
extra_headers: &BTreeMap::new(),
content_type: Some("application/json"),
provider_api_format: "openai:chat",
client_api_format: "openai:chat",
upstream_is_stream: true,
build_from_request_when_empty: true,
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreamingOrWildcard,
});
assert_eq!(
headers.get("accept"),
Some(&"text/event-stream".to_string())
);
assert_eq!(
headers.get("accept-encoding"),
Some(&"identity".to_string())
);
}
#[test]
fn stream_fallback_headers_preserve_wildcard_in_missing_only_mode() {
let mut request_headers = HeaderMap::new();
request_headers.insert(http::header::ACCEPT, "*/*".parse().expect("header"));
let headers = build_standard_plan_fallback_headers(StandardPlanFallbackHeadersInput {
request_headers: &request_headers,
existing_provider_request_headers: BTreeMap::new(),
auth_header: Some("authorization"),
auth_value: Some("Bearer secret"),
extra_headers: &BTreeMap::new(),
content_type: Some("application/json"),
provider_api_format: "gemini:generate_content",
client_api_format: "openai:responses",
upstream_is_stream: true,
build_from_request_when_empty: true,
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreaming,
});
assert_eq!(headers.get("accept"), Some(&"*/*".to_string()));
}
#[test]
fn stream_fallback_headers_preserve_explicit_accept() {
let mut existing_headers = BTreeMap::new();
existing_headers.insert("accept".to_string(), "application/json".to_string());
let headers = build_standard_plan_fallback_headers(StandardPlanFallbackHeadersInput {
request_headers: &HeaderMap::new(),
existing_provider_request_headers: existing_headers,
auth_header: Some("authorization"),
auth_value: Some("Bearer secret"),
extra_headers: &BTreeMap::new(),
content_type: Some("application/json"),
provider_api_format: "openai:chat",
client_api_format: "openai:chat",
upstream_is_stream: true,
build_from_request_when_empty: false,
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreamingOrWildcard,
});
assert_eq!(headers.get("accept"), Some(&"application/json".to_string()));
}
#[test]
fn plan_fallback_headers_preserve_empty_existing_mode() {
let headers = build_standard_plan_fallback_headers(StandardPlanFallbackHeadersInput {
request_headers: &HeaderMap::new(),
existing_provider_request_headers: BTreeMap::new(),
auth_header: Some("authorization"),
auth_value: Some("Bearer secret"),
extra_headers: &BTreeMap::new(),
content_type: Some("application/json"),
provider_api_format: "openai:responses",
client_api_format: "openai:responses",
upstream_is_stream: true,
build_from_request_when_empty: false,
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreaming,
});
assert_eq!(
headers.get("authorization"),
Some(&"Bearer secret".to_string())
);
assert_eq!(
headers.get("accept"),
Some(&"text/event-stream".to_string())
);
assert!(!headers.contains_key("content-type"));
}
#[test]
fn plan_fallback_url_helpers_route_openai_surfaces() {
assert_eq!(
build_standard_plan_fallback_openai_chat_url("https://api.example.com/v1", Some("x=1")),
"https://api.example.com/v1/chat/completions?x=1"
);
assert_eq!(
build_standard_plan_fallback_openai_responses_url(
"https://api.example.com/v1",
Some("x=1"),
true,
),
"https://api.example.com/v1/responses/compact?x=1"
);
}
}
+752
View File
@@ -0,0 +1,752 @@
use std::collections::BTreeMap;
use url::form_urlencoded;
use url::Url;
pub fn build_openai_chat_url(upstream_base_url: &str, query: Option<&str>) -> String {
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
let trimmed = trimmed.trim_end_matches('/');
let mut url = format!("{trimmed}/chat/completions");
append_merged_query(&mut url, base_query, None, query, &[]);
url
}
pub fn build_openai_responses_url(
upstream_base_url: &str,
query: Option<&str>,
compact: bool,
) -> String {
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
let trimmed = trimmed.trim_end_matches('/');
let suffix = if compact {
"/responses/compact"
} else {
"/responses"
};
let mut url = format!("{trimmed}{suffix}");
append_merged_query(&mut url, base_query, None, query, &[]);
url
}
pub fn build_openai_search_url(upstream_base_url: &str, query: Option<&str>) -> String {
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
let trimmed = trimmed.trim_end_matches('/');
let mut url = format!("{trimmed}/alpha/search");
append_merged_query(&mut url, base_query, None, query, &[]);
url
}
pub fn build_openai_image_url(
upstream_base_url: &str,
request_path: Option<&str>,
query: Option<&str>,
) -> String {
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
let trimmed = trimmed.trim_end_matches('/');
let suffix = openai_image_path_suffix(request_path);
let mut url = if openai_image_base_includes_operation_path(trimmed) {
trimmed.to_string()
} else {
format!("{trimmed}{suffix}")
};
append_merged_query(&mut url, base_query, None, query, &[]);
url
}
fn openai_image_path_suffix(request_path: Option<&str>) -> &'static str {
match request_path
.map(str::trim)
.map(|value| value.trim_end_matches('/'))
{
Some("/v1/images/edits") | Some("/images/edits") => "/images/edits",
_ => "/images/generations",
}
}
fn openai_image_base_includes_operation_path(base_url: &str) -> bool {
let path = Url::parse(base_url)
.ok()
.map(|url| url.path().trim_end_matches('/').to_string())
.unwrap_or_else(|| base_url.trim_end_matches('/').to_string());
path.ends_with("/images/generations") || path.ends_with("/images/edits")
}
pub fn build_claude_messages_url(upstream_base_url: &str, query: Option<&str>) -> String {
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
let trimmed = trimmed.trim_end_matches('/');
let mut url = format!("{trimmed}/messages");
append_merged_query(&mut url, base_query, None, query, &[]);
url
}
pub fn build_gemini_content_url(
upstream_base_url: &str,
model: &str,
stream: bool,
query: Option<&str>,
) -> Option<String> {
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
let trimmed_model = model.trim();
if trimmed_base_url.is_empty() || trimmed_model.is_empty() {
return None;
}
let operation = if stream {
"streamGenerateContent"
} else {
"generateContent"
};
let mut url = if trimmed_base_url.ends_with("/v1") || trimmed_base_url.ends_with("/v1beta") {
format!("{trimmed_base_url}/models/{trimmed_model}:{operation}")
} else if gemini_content_base_url_contains_model_path(trimmed_base_url) {
let trimmed_base_url = strip_gemini_content_action(trimmed_base_url);
format!("{trimmed_base_url}:{operation}")
} else {
format!("{trimmed_base_url}/v1beta/models/{trimmed_model}:{operation}")
};
append_merged_query(&mut url, base_query, None, query, &["key"]);
Some(url)
}
pub fn normalize_gemini_content_action_path(path: &str, stream: bool) -> String {
let trimmed = path.trim();
let (path, query) = split_path_query(trimmed);
let action = if stream {
"streamGenerateContent"
} else {
"generateContent"
};
let normalized = strip_gemini_content_action(path);
let normalized = if normalized.len() == path.len() {
path.to_string()
} else {
format!("{normalized}:{action}")
};
match query {
Some(query) => format!("{normalized}?{query}"),
None => normalized,
}
}
fn strip_gemini_content_action(value: &str) -> &str {
value
.strip_suffix(":streamGenerateContent")
.or_else(|| value.strip_suffix(":generateContent"))
.unwrap_or(value)
}
fn gemini_content_base_url_contains_model_path(value: &str) -> bool {
value.contains("/v1/models/") || value.contains("/v1beta/models/")
}
pub fn build_gemini_video_predict_long_running_url(
upstream_base_url: &str,
model: &str,
query: Option<&str>,
) -> Option<String> {
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
let trimmed_model = model.trim();
if trimmed_base_url.is_empty() || trimmed_model.is_empty() {
return None;
}
let mut url = if trimmed_base_url.ends_with("/v1") || trimmed_base_url.ends_with("/v1beta") {
format!("{trimmed_base_url}/models/{trimmed_model}:predictLongRunning")
} else if gemini_content_base_url_contains_model_path(trimmed_base_url) {
format!("{trimmed_base_url}:predictLongRunning")
} else {
format!("{trimmed_base_url}/v1beta/models/{trimmed_model}:predictLongRunning")
};
append_merged_query(&mut url, base_query, None, query, &["key"]);
Some(url)
}
pub fn build_passthrough_path_url(
upstream_base_url: &str,
path: &str,
query: Option<&str>,
blocked_keys: &[&str],
) -> Option<String> {
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
let trimmed_path = path.trim();
if trimmed_base_url.is_empty() || trimmed_path.is_empty() {
return None;
}
let (trimmed_path, path_query) = split_path_query(trimmed_path);
let normalized_base_url =
if trimmed_base_url.ends_with("/v1beta") && trimmed_path.starts_with("/v1beta") {
trimmed_base_url.trim_end_matches("/v1beta")
} else if trimmed_base_url.ends_with("/v1") && trimmed_path.starts_with("/v1/") {
trimmed_base_url.trim_end_matches("/v1")
} else {
trimmed_base_url
};
let mut url = format!("{normalized_base_url}{trimmed_path}");
append_merged_query(&mut url, base_query, path_query, query, blocked_keys);
Some(url)
}
pub fn build_bigmodel_coding_models_url(upstream_base_url: &str) -> Option<String> {
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
if trimmed_base_url.is_empty() || !bigmodel_coding_models_base_is_supported(trimmed_base_url) {
return None;
}
let path = Url::parse(trimmed_base_url)
.ok()
.map(|url| url.path().trim_end_matches('/').to_string())
.unwrap_or_else(|| trimmed_base_url.trim_end_matches('/').to_string());
let mut url = if path.ends_with("/models") {
trimmed_base_url.to_string()
} else {
format!("{trimmed_base_url}/models")
};
append_merged_query(&mut url, base_query, None, None, &[]);
Some(url)
}
pub fn build_openai_compatible_models_url(upstream_base_url: &str) -> Option<String> {
if let Some(url) = build_bigmodel_coding_models_url(upstream_base_url) {
return Some(url);
}
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
if trimmed_base_url.is_empty() {
return None;
}
let mut url = if trimmed_base_url.ends_with("/models") {
trimmed_base_url.to_string()
} else {
format!("{trimmed_base_url}/models")
};
append_merged_query(&mut url, base_query, None, None, &[]);
Some(url)
}
pub fn build_gemini_files_passthrough_url(
upstream_base_url: &str,
path: &str,
query: Option<&str>,
) -> Option<String> {
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
let trimmed_path = path.trim();
if trimmed_base_url.is_empty() || trimmed_path.is_empty() {
return None;
}
let (trimmed_path, path_query) = split_path_query(trimmed_path);
let normalized_base_url = if trimmed_base_url.ends_with("/v1beta")
&& (trimmed_path.starts_with("/v1beta/") || trimmed_path.starts_with("/upload/v1beta/"))
{
trimmed_base_url.trim_end_matches("/v1beta")
} else {
trimmed_base_url
};
let mut url = format!("{normalized_base_url}{trimmed_path}");
append_merged_query(&mut url, base_query, path_query, query, &["key"]);
Some(url)
}
fn split_base_url_query(base_url: &str) -> (&str, Option<&str>) {
let trimmed = base_url.trim();
trimmed
.split_once('?')
.map(|(base, query)| (base, Some(query)))
.unwrap_or((trimmed, None))
}
pub(crate) fn google_openai_compat_base_includes_api_root(base_url: &str) -> bool {
let Ok(parsed) = Url::parse(base_url.trim()) else {
return false;
};
let Some(host) = parsed.host_str().map(|value| value.to_ascii_lowercase()) else {
return false;
};
let path = parsed.path().trim_end_matches('/');
if host == "generativelanguage.googleapis.com" {
return path == "/v1beta/openai" || path == "/v1/openai";
}
if looks_like_vertex_ai_host(&host) {
return path.ends_with("/endpoints/openapi");
}
false
}
pub fn openai_compatible_base_includes_api_root(base_url: &str) -> bool {
let trimmed = base_url.trim().trim_end_matches('/');
trimmed.ends_with("/v1")
|| google_openai_compat_base_includes_api_root(trimmed)
|| bigmodel_coding_base_includes_api_root(trimmed)
|| openai_compatible_base_includes_unversioned_api_root(trimmed)
}
pub fn v1_compatible_base_includes_api_root(base_url: &str) -> bool {
let trimmed = base_url.trim().trim_end_matches('/');
trimmed.ends_with("/v1") || openai_compatible_base_includes_unversioned_api_root(trimmed)
}
pub fn openai_compatible_base_includes_unversioned_api_root(base_url: &str) -> bool {
let trimmed = base_url.trim().trim_end_matches('/');
let path = Url::parse(trimmed)
.ok()
.map(|url| url.path().trim_end_matches('/').to_ascii_lowercase())
.unwrap_or_else(|| {
trimmed
.split_once('/')
.map(|(_, path)| format!("/{path}"))
.unwrap_or_default()
.trim_end_matches('/')
.to_ascii_lowercase()
});
!path.is_empty()
}
fn bigmodel_coding_base_includes_api_root(base_url: &str) -> bool {
let Ok(parsed) = Url::parse(base_url.trim()) else {
return false;
};
let Some(host) = parsed.host_str().map(|value| value.to_ascii_lowercase()) else {
return false;
};
host == "open.bigmodel.cn" && parsed.path().trim_end_matches('/') == "/api/coding/paas/v4"
}
fn bigmodel_coding_models_base_is_supported(base_url: &str) -> bool {
let Ok(parsed) = Url::parse(base_url.trim()) else {
return false;
};
let Some(host) = parsed.host_str().map(|value| value.to_ascii_lowercase()) else {
return false;
};
if host != "open.bigmodel.cn" {
return false;
}
matches!(
parsed.path().trim_end_matches('/'),
"/api/coding/paas/v4" | "/api/coding/paas/v4/models"
)
}
fn looks_like_vertex_ai_host(host: &str) -> bool {
const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com";
host == VERTEX_AI_HOST
|| host.ends_with(&format!(".{VERTEX_AI_HOST}"))
|| host.ends_with(&format!("-{VERTEX_AI_HOST}"))
}
fn split_path_query(path: &str) -> (&str, Option<&str>) {
path.split_once('?')
.map(|(path, query)| (path, Some(query)))
.unwrap_or((path, None))
}
fn append_merged_query(
url: &mut String,
base_query: Option<&str>,
path_query: Option<&str>,
request_query: Option<&str>,
blocked_keys: &[&str],
) {
let Some(query) = merge_query_layers(base_query, path_query, request_query, blocked_keys)
else {
return;
};
if url.contains('?') {
url.push('&');
} else {
url.push('?');
}
url.push_str(&query);
}
fn merge_query_layers(
base_query: Option<&str>,
path_query: Option<&str>,
request_query: Option<&str>,
blocked_keys: &[&str],
) -> Option<String> {
if blocked_keys.is_empty()
&& path_query.is_none()
&& base_query.is_none()
&& request_query
.map(str::trim)
.is_some_and(|value| !value.is_empty())
{
return request_query
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
}
let mut merged = BTreeMap::new();
for source in [base_query, path_query, request_query] {
merge_query_string(&mut merged, source, blocked_keys);
}
if merged.is_empty() {
return None;
}
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in merged {
serializer.append_pair(&key, &value);
}
Some(serializer.finish())
}
fn merge_query_string(
out: &mut BTreeMap<String, String>,
query: Option<&str>,
blocked_keys: &[&str],
) {
let Some(query) = query.map(str::trim).filter(|value| !value.is_empty()) else {
return;
};
for (key, value) in form_urlencoded::parse(query.as_bytes()) {
if blocked_keys
.iter()
.any(|blocked| key.as_ref().eq_ignore_ascii_case(blocked))
{
continue;
}
out.insert(key.into_owned(), value.into_owned());
}
}
#[cfg(test)]
mod tests {
use super::{
build_bigmodel_coding_models_url, build_claude_messages_url, build_gemini_content_url,
build_gemini_files_passthrough_url, build_gemini_video_predict_long_running_url,
build_openai_chat_url, build_openai_compatible_models_url, build_openai_image_url,
build_openai_responses_url, build_openai_search_url, build_passthrough_path_url,
normalize_gemini_content_action_path,
};
#[test]
fn merges_base_url_query_for_same_format_urls() {
assert_eq!(
build_openai_chat_url(
"https://api.openai.example/v1?tenant=demo",
Some("mode=fast&tenant=override")
),
"https://api.openai.example/v1/chat/completions?mode=fast&tenant=override"
);
}
#[test]
fn openai_chat_url_preserves_google_openai_compat_roots() {
assert_eq!(
build_openai_chat_url(
"https://generativelanguage.googleapis.com/v1beta/openai",
Some("trace=1")
),
"https://generativelanguage.googleapis.com/v1beta/openai/chat/completions?trace=1"
);
assert_eq!(
build_openai_chat_url(
"https://aiplatform.googleapis.com/v1/projects/project-1/locations/global/endpoints/openapi",
None,
),
"https://aiplatform.googleapis.com/v1/projects/project-1/locations/global/endpoints/openapi/chat/completions"
);
}
#[test]
fn openai_urls_preserve_bigmodel_coding_api_root() {
assert_eq!(
build_openai_chat_url(
"https://open.bigmodel.cn/api/coding/paas/v4",
Some("trace=1")
),
"https://open.bigmodel.cn/api/coding/paas/v4/chat/completions?trace=1"
);
assert_eq!(
build_openai_responses_url("https://open.bigmodel.cn/api/coding/paas/v4", None, false),
"https://open.bigmodel.cn/api/coding/paas/v4/responses"
);
}
#[test]
fn openai_urls_preserve_unversioned_api_root() {
assert_eq!(
build_openai_chat_url("https://proxy.example.com/api", Some("trace=1")),
"https://proxy.example.com/api/chat/completions?trace=1"
);
assert_eq!(
build_openai_chat_url("https://proxy.example.com/openai", None),
"https://proxy.example.com/openai/chat/completions"
);
assert_eq!(
build_openai_chat_url("https://proxy.example.com", None),
"https://proxy.example.com/chat/completions"
);
assert_eq!(
build_openai_chat_url("https://api.deepseek.com", None),
"https://api.deepseek.com/chat/completions"
);
assert_eq!(
build_openai_responses_url("https://proxy.example.com/api", None, false),
"https://proxy.example.com/api/responses"
);
assert_eq!(
build_openai_responses_url("https://api.deepseek.com", None, false),
"https://api.deepseek.com/responses"
);
assert_eq!(
build_openai_image_url(
"https://proxy.example.com/api",
Some("/v1/images/generations"),
None
),
"https://proxy.example.com/api/images/generations"
);
}
#[test]
fn claude_messages_url_preserves_v1_and_unversioned_api_roots() {
assert_eq!(
build_claude_messages_url("https://api.anthropic.example/v1", Some("trace=1")),
"https://api.anthropic.example/v1/messages?trace=1"
);
assert_eq!(
build_claude_messages_url("https://proxy.example.com/api", None),
"https://proxy.example.com/api/messages"
);
assert_eq!(
build_claude_messages_url("https://proxy.example.com/anthropic", None),
"https://proxy.example.com/anthropic/messages"
);
assert_eq!(
build_claude_messages_url("https://api.anthropic.example", None),
"https://api.anthropic.example/messages"
);
}
#[test]
fn bigmodel_coding_models_url_uses_models_resource() {
assert_eq!(
build_bigmodel_coding_models_url(
"https://open.bigmodel.cn/api/coding/paas/v4?tenant=demo"
)
.as_deref(),
Some("https://open.bigmodel.cn/api/coding/paas/v4/models?tenant=demo")
);
assert_eq!(
build_bigmodel_coding_models_url("https://open.bigmodel.cn/api/coding/paas/v4/models")
.as_deref(),
Some("https://open.bigmodel.cn/api/coding/paas/v4/models")
);
}
#[test]
fn openai_compatible_models_url_preserves_unversioned_api_root() {
assert_eq!(
build_openai_compatible_models_url("https://proxy.example.com/api?tenant=demo")
.as_deref(),
Some("https://proxy.example.com/api/models?tenant=demo")
);
assert_eq!(
build_openai_compatible_models_url("https://proxy.example.com/openai").as_deref(),
Some("https://proxy.example.com/openai/models")
);
assert_eq!(
build_openai_compatible_models_url("https://proxy.example.com").as_deref(),
Some("https://proxy.example.com/models")
);
}
#[test]
fn openai_responses_url_preserves_codex_path_prefix() {
assert_eq!(
build_openai_responses_url("https://tiger.bookapi.cc/codex", None, false),
"https://tiger.bookapi.cc/codex/responses"
);
assert_eq!(
build_openai_responses_url("https://tiger.bookapi.cc/codex?tenant=demo", None, true),
"https://tiger.bookapi.cc/codex/responses/compact?tenant=demo"
);
}
#[test]
fn openai_search_url_preserves_api_and_codex_roots() {
assert_eq!(
build_openai_search_url(
"https://api.openai.com/v1?tenant=base",
Some("trace=1&tenant=request")
),
"https://api.openai.com/v1/alpha/search?tenant=request&trace=1"
);
assert_eq!(
build_openai_search_url("https://chatgpt.com/backend-api/codex/", None),
"https://chatgpt.com/backend-api/codex/alpha/search"
);
}
#[test]
fn openai_image_url_uses_images_surface() {
assert_eq!(
build_openai_image_url(
"https://api.openai.example/v1?tenant=demo",
Some("/v1/images/generations"),
Some("trace=1")
),
"https://api.openai.example/v1/images/generations?tenant=demo&trace=1"
);
assert_eq!(
build_openai_image_url(
"https://api.openai.example/v1",
Some("/v1/images/edits"),
None
),
"https://api.openai.example/v1/images/edits"
);
}
#[test]
fn merges_base_url_query_for_dynamic_gemini_content_urls() {
assert_eq!(
build_gemini_content_url(
"https://generativelanguage.googleapis.com/v1beta?alt=sse",
"gemini-2.5-pro",
true,
Some("foo=bar&key=secret")
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:streamGenerateContent?alt=sse&foo=bar"
)
);
}
#[test]
fn gemini_content_urls_rewrite_existing_base_action_for_stream_mode() {
assert_eq!(
build_gemini_content_url(
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent",
"ignored-model",
true,
Some("foo=bar")
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:streamGenerateContent?foo=bar"
)
);
assert_eq!(
build_gemini_content_url(
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:streamGenerateContent",
"ignored-model",
false,
Some("foo=bar")
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent?foo=bar"
)
);
assert_eq!(
build_gemini_content_url(
"https://generativelanguage.googleapis.com/v1/models/gemini-2.5-pro:generateContent",
"ignored-model",
true,
Some("foo=bar")
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/v1/models/gemini-2.5-pro:streamGenerateContent?foo=bar"
)
);
}
#[test]
fn normalizes_gemini_content_action_in_custom_paths() {
assert_eq!(
normalize_gemini_content_action_path(
"/v1beta/models/gemini-2.5-pro:generateContent?alt=sse",
true
),
"/v1beta/models/gemini-2.5-pro:streamGenerateContent?alt=sse"
);
assert_eq!(
normalize_gemini_content_action_path(
"/v1beta/models/gemini-2.5-pro:streamGenerateContent",
false
),
"/v1beta/models/gemini-2.5-pro:generateContent"
);
}
#[test]
fn merges_base_path_and_request_query_for_passthrough_paths() {
assert_eq!(
build_passthrough_path_url(
"https://api.openai.example/v1?tenant=demo",
"/videos/generations?variant=video",
Some("size=1024"),
&[]
)
.as_deref(),
Some(
"https://api.openai.example/v1/videos/generations?size=1024&tenant=demo&variant=video"
)
);
}
#[test]
fn passthrough_path_does_not_duplicate_openai_v1_root() {
assert_eq!(
build_passthrough_path_url(
"https://api.openai.example/v1?tenant=demo",
"/v1/chat/completions?variant=chat",
Some("trace=1"),
&[]
)
.as_deref(),
Some("https://api.openai.example/v1/chat/completions?tenant=demo&trace=1&variant=chat")
);
}
#[test]
fn merges_base_url_query_for_gemini_files_passthrough_urls() {
assert_eq!(
build_gemini_files_passthrough_url(
"https://generativelanguage.googleapis.com/v1beta?alt=media",
"/upload/v1beta/files?uploadType=resumable",
Some("key=secret&pageSize=10")
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/upload/v1beta/files?alt=media&pageSize=10&uploadType=resumable"
)
);
}
#[test]
fn merges_base_url_query_for_gemini_video_urls() {
assert_eq!(
build_gemini_video_predict_long_running_url(
"https://generativelanguage.googleapis.com/v1beta?alt=sse",
"veo-3.0-generate-preview",
Some("foo=bar&key=secret")
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/v1beta/models/veo-3.0-generate-preview:predictLongRunning?alt=sse&foo=bar"
)
);
}
}
@@ -0,0 +1,485 @@
use std::collections::BTreeMap;
use async_trait::async_trait;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine as _;
use rsa::pkcs1::DecodeRsaPrivateKey;
use rsa::pkcs1v15::SigningKey;
use rsa::pkcs8::DecodePrivateKey;
use rsa::signature::{SignatureEncoding, Signer};
use rsa::RsaPrivateKey;
use serde_json::{json, Value};
use sha2::Sha256;
use url::form_urlencoded;
use super::super::oauth_refresh::{
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthHttpRequest, LocalOAuthRefreshAdapter,
LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth,
};
use super::super::snapshot::GatewayProviderTransportSnapshot;
pub const VERTEX_API_KEY_QUERY_PARAM: &str = "key";
pub const VERTEX_SERVICE_ACCOUNT_AUTH_HEADER: &str = "authorization";
pub const VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE: &str = "vertex_ai";
pub const GOOGLE_OAUTH_TOKEN_URL: &str = "https://oauth2.googleapis.com/token";
const GOOGLE_CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
const SERVICE_ACCOUNT_REFRESH_SKEW_SECS: u64 = 120;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VertexApiKeyQueryAuth {
pub name: &'static str,
pub value: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VertexServiceAccountAuthConfig {
pub client_email: String,
pub private_key: String,
pub project_id: String,
pub token_uri: String,
pub region: Option<String>,
pub model_regions: BTreeMap<String, String>,
}
pub fn resolve_local_vertex_api_key_query_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<VertexApiKeyQueryAuth> {
if !super::is_vertex_api_key_transport_context(transport) {
return None;
}
if transport.key.decrypted_auth_config.is_some() {
return None;
}
if !transport
.key
.auth_type
.trim()
.eq_ignore_ascii_case("api_key")
{
return None;
}
let secret = transport.key.decrypted_api_key.trim();
if secret.is_empty() {
return None;
}
Some(VertexApiKeyQueryAuth {
name: VERTEX_API_KEY_QUERY_PARAM,
value: secret.to_string(),
})
}
pub fn resolve_local_vertex_service_account_auth_config(
transport: &GatewayProviderTransportSnapshot,
) -> Option<VertexServiceAccountAuthConfig> {
if !super::is_vertex_service_account_transport_context(transport) {
return None;
}
parse_vertex_service_account_auth_config(transport.key.decrypted_auth_config.as_deref())
}
pub fn supports_local_vertex_service_account_auth_resolution(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
resolve_local_vertex_service_account_auth_config(transport).is_some()
}
pub fn parse_vertex_service_account_auth_config(
raw: Option<&str>,
) -> Option<VertexServiceAccountAuthConfig> {
let raw = raw.map(str::trim).filter(|value| !value.is_empty())?;
let value: Value = serde_json::from_str(raw).ok()?;
parse_vertex_service_account_auth_config_value(&value)
}
fn parse_vertex_service_account_auth_config_value(
value: &Value,
) -> Option<VertexServiceAccountAuthConfig> {
let client_email = json_string(value.get("client_email"))?;
let private_key = json_string(value.get("private_key"))?;
let project_id = json_string(value.get("project_id"))?;
let token_uri =
json_string(value.get("token_uri")).unwrap_or_else(|| GOOGLE_OAUTH_TOKEN_URL.to_string());
let region = json_string(value.get("region"));
let model_regions = value
.get("model_regions")
.and_then(Value::as_object)
.map(|items| {
items
.iter()
.filter_map(|(model, region)| {
let model = model.trim();
let region = region.as_str()?.trim();
(!model.is_empty() && !region.is_empty())
.then(|| (model.to_string(), region.to_string()))
})
.collect::<BTreeMap<_, _>>()
})
.unwrap_or_default();
Some(VertexServiceAccountAuthConfig {
client_email,
private_key,
project_id,
token_uri,
region,
model_regions,
})
}
fn json_string(value: Option<&Value>) -> Option<String> {
value
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
#[derive(Debug, Clone, Default)]
pub struct VertexServiceAccountRefreshAdapter;
#[async_trait]
impl LocalOAuthRefreshAdapter for VertexServiceAccountRefreshAdapter {
fn provider_type(&self) -> &'static str {
VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE
}
fn supports(&self, transport: &GatewayProviderTransportSnapshot) -> bool {
supports_local_vertex_service_account_auth_resolution(transport)
}
fn resolve_cached(
&self,
_transport: &GatewayProviderTransportSnapshot,
entry: &CachedOAuthEntry,
) -> Option<LocalResolvedOAuthRequestAuth> {
if !entry
.provider_type
.eq_ignore_ascii_case(VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE)
{
return None;
}
if service_account_token_expires_soon(entry.expires_at_unix_secs) {
return None;
}
let name = entry.auth_header_name.trim();
let value = entry.auth_header_value.trim();
if name.is_empty() || value.is_empty() {
return None;
}
Some(LocalResolvedOAuthRequestAuth::Header {
name: name.to_ascii_lowercase(),
value: value.to_string(),
})
}
fn resolve_without_refresh(
&self,
_transport: &GatewayProviderTransportSnapshot,
) -> Option<LocalResolvedOAuthRequestAuth> {
None
}
fn should_refresh(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> bool {
supports_local_vertex_service_account_auth_resolution(transport)
&& entry
.and_then(|cached| self.resolve_cached(transport, cached))
.is_none()
}
async fn refresh(
&self,
executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot,
_entry: Option<&CachedOAuthEntry>,
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
let Some(auth_config) = resolve_local_vertex_service_account_auth_config(transport) else {
return Ok(None);
};
let now = aether_oauth::core::current_unix_secs();
let assertion = build_vertex_service_account_assertion(&auth_config, now)?;
let body = form_urlencoded::Serializer::new(String::new())
.append_pair("grant_type", "urn:ietf:params:oauth:grant-type:jwt-bearer")
.append_pair("assertion", &assertion)
.finish();
let response = executor
.execute(
VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
transport,
&LocalOAuthHttpRequest {
request_id: "vertex_ai:service-account-token",
method: reqwest::Method::POST,
url: auth_config.token_uri.clone(),
headers: BTreeMap::from([(
"content-type".to_string(),
"application/x-www-form-urlencoded".to_string(),
)]),
json_body: None,
body_bytes: Some(body.into_bytes()),
},
)
.await?;
if response.status_code != 200 {
return Err(LocalOAuthRefreshError::HttpStatus {
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
status_code: response.status_code,
body_excerpt: body_excerpt(&response.body_text),
});
}
let body_json: Value = serde_json::from_str(&response.body_text).map_err(|err| {
LocalOAuthRefreshError::InvalidResponse {
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
message: format!("vertex service account token response is not JSON: {err}"),
}
})?;
let access_token = json_string(body_json.get("access_token")).ok_or_else(|| {
LocalOAuthRefreshError::InvalidResponse {
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
message: "vertex service account token response missing access_token".to_string(),
}
})?;
let expires_in = body_json
.get("expires_in")
.and_then(Value::as_u64)
.unwrap_or(3600);
Ok(Some(CachedOAuthEntry {
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE.to_string(),
auth_header_name: VERTEX_SERVICE_ACCOUNT_AUTH_HEADER.to_string(),
auth_header_value: format!("Bearer {access_token}"),
expires_at_unix_secs: Some(now.saturating_add(expires_in)),
metadata: Some(json!({
"project_id": auth_config.project_id,
"client_email": auth_config.client_email,
})),
}))
}
}
pub fn build_vertex_service_account_assertion(
auth_config: &VertexServiceAccountAuthConfig,
now_unix_secs: u64,
) -> Result<String, LocalOAuthRefreshError> {
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#);
let payload = URL_SAFE_NO_PAD.encode(
serde_json::to_string(&json!({
"iss": auth_config.client_email,
"sub": auth_config.client_email,
"scope": GOOGLE_CLOUD_PLATFORM_SCOPE,
"aud": auth_config.token_uri,
"iat": now_unix_secs,
"exp": now_unix_secs.saturating_add(3600),
}))
.map_err(|err| LocalOAuthRefreshError::InvalidResponse {
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
message: format!("vertex service account jwt payload encode failed: {err}"),
})?,
);
let message = format!("{header}.{payload}");
let private_key = decode_vertex_service_account_private_key(auth_config.private_key.as_str())?;
let signing_key = SigningKey::<Sha256>::new(private_key);
let signature = signing_key.sign(message.as_bytes());
Ok(format!(
"{message}.{}",
URL_SAFE_NO_PAD.encode(signature.to_bytes())
))
}
fn decode_vertex_service_account_private_key(
private_key_pem: &str,
) -> Result<RsaPrivateKey, LocalOAuthRefreshError> {
match RsaPrivateKey::from_pkcs8_pem(private_key_pem) {
Ok(private_key) => Ok(private_key),
Err(pkcs8_err) => RsaPrivateKey::from_pkcs1_pem(private_key_pem).map_err(|pkcs1_err| {
LocalOAuthRefreshError::InvalidResponse {
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
message: format!(
"vertex service account private_key parse failed: pkcs8: {pkcs8_err}; pkcs1: {pkcs1_err}"
),
}
}),
}
}
fn service_account_token_expires_soon(expires_at_unix_secs: Option<u64>) -> bool {
expires_at_unix_secs
.map(|expires_at_unix_secs| {
aether_oauth::core::current_unix_secs()
>= expires_at_unix_secs.saturating_sub(SERVICE_ACCOUNT_REFRESH_SKEW_SECS)
})
.unwrap_or(true)
}
fn body_excerpt(value: &str) -> String {
value.chars().take(500).collect()
}
#[cfg(test)]
mod tests {
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use rsa::pkcs1::{EncodeRsaPrivateKey, LineEnding};
use rsa::rand_core::OsRng;
use rsa::RsaPrivateKey;
use super::{
decode_vertex_service_account_private_key, parse_vertex_service_account_auth_config,
resolve_local_vertex_api_key_query_auth,
supports_local_vertex_service_account_auth_resolution, VERTEX_API_KEY_QUERY_PARAM,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Vertex".to_string(),
provider_type: "vertex_ai".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "gemini:generate_content".to_string(),
api_family: Some("gemini".to_string()),
endpoint_kind: Some("chat".to_string()),
is_active: true,
base_url: "https://aiplatform.googleapis.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: Some(vec!["gemini:generate_content".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "vertex-secret".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn resolves_query_auth_for_vertex_api_key_subset() {
let auth = resolve_local_vertex_api_key_query_auth(&sample_transport())
.expect("vertex api key query auth should resolve");
assert_eq!(auth.name, VERTEX_API_KEY_QUERY_PARAM);
assert_eq!(auth.value, "vertex-secret");
}
#[test]
fn rejects_non_api_key_transport() {
let mut transport = sample_transport();
transport.key.auth_type = "service_account".to_string();
assert!(resolve_local_vertex_api_key_query_auth(&transport).is_none());
}
#[test]
fn rejects_vertex_auth_config_transport() {
let mut transport = sample_transport();
transport.key.decrypted_auth_config = Some("{\"project_id\":\"demo-project\"}".to_string());
assert!(resolve_local_vertex_api_key_query_auth(&transport).is_none());
}
#[test]
fn resolves_query_auth_for_custom_aiplatform_transport() {
let mut transport = sample_transport();
transport.provider.provider_type = "custom".to_string();
transport.endpoint.api_format = "gemini:generate_content".to_string();
let auth = resolve_local_vertex_api_key_query_auth(&transport)
.expect("custom aiplatform transport should resolve");
assert_eq!(auth.value, "vertex-secret");
}
#[test]
fn parses_service_account_auth_config() {
let config = parse_vertex_service_account_auth_config(Some(
r#"{
"client_email":"[email protected]",
"private_key":"TEST-PRIVATE-KEY",
"project_id":"demo-project",
"region":"global",
"model_regions":{"gemini-2.0-flash":"us-central1"}
}"#,
))
.expect("service account config should parse");
assert_eq!(config.client_email, "[email protected]");
assert_eq!(config.project_id, "demo-project");
assert_eq!(config.region.as_deref(), Some("global"));
assert_eq!(
config
.model_regions
.get("gemini-2.0-flash")
.map(String::as_str),
Some("us-central1")
);
}
#[test]
fn supports_vertex_service_account_auth_resolution() {
let mut transport = sample_transport();
transport.key.auth_type = "service_account".to_string();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"client_email":"[email protected]",
"private_key":"TEST-PRIVATE-KEY",
"project_id":"demo-project"
}"#
.to_string(),
);
assert!(supports_local_vertex_service_account_auth_resolution(
&transport
));
}
#[test]
fn decodes_pkcs1_service_account_private_key() {
let mut rng = OsRng;
let private_key = RsaPrivateKey::new(&mut rng, 1024)
.expect("test RSA private key should generate")
.to_pkcs1_pem(LineEnding::LF)
.expect("test RSA private key should encode as PKCS#1 PEM");
decode_vertex_service_account_private_key(private_key.as_str())
.expect("PKCS#1 private key should decode");
}
}
@@ -0,0 +1,239 @@
use url::Url;
use super::super::auth::resolve_local_auth_type_for_transport_format;
use super::super::snapshot::GatewayProviderTransportSnapshot;
const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com";
pub fn looks_like_vertex_ai_host(base_url: &str) -> bool {
let trimmed = base_url.trim();
if trimmed.is_empty() {
return false;
}
let Ok(parsed) = Url::parse(trimmed) else {
return false;
};
let Some(host) = parsed
.host_str()
.map(|value| value.trim().to_ascii_lowercase())
else {
return false;
};
host == VERTEX_AI_HOST
|| host.ends_with(&format!(".{VERTEX_AI_HOST}"))
|| host.ends_with(&format!("-{VERTEX_AI_HOST}"))
}
pub fn is_vertex_api_key_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool {
if is_vertex_provider_type(transport) {
return resolve_local_auth_type_for_transport_format(transport)
.eq_ignore_ascii_case("api_key");
}
if !is_vertex_host_format_context(transport) {
return false;
}
resolve_local_auth_type_for_transport_format(transport).eq_ignore_ascii_case("api_key")
}
pub fn is_vertex_service_account_transport_context(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
if !is_vertex_provider_type(transport) && !is_vertex_host_format_context(transport) {
return false;
}
matches!(
resolve_local_auth_type_for_transport_format(transport).as_str(),
"service_account" | "vertex_ai"
)
}
pub fn is_vertex_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool {
is_vertex_api_key_transport_context(transport)
|| is_vertex_service_account_transport_context(transport)
}
pub fn uses_vertex_api_key_query_auth(
transport: &GatewayProviderTransportSnapshot,
provider_api_format: &str,
) -> bool {
is_vertex_api_key_transport_context(transport)
&& provider_api_format
.trim()
.to_ascii_lowercase()
.starts_with("gemini:")
}
fn is_vertex_provider_type(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
}
fn is_vertex_host_format_context(transport: &GatewayProviderTransportSnapshot) -> bool {
if !looks_like_vertex_ai_host(&transport.endpoint.base_url) {
return false;
}
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
endpoint_api_format.starts_with("gemini:")
|| endpoint_api_format.starts_with("claude:")
|| (endpoint_api_format.starts_with("openai:")
&& looks_like_vertex_openai_compat_base(&transport.endpoint.base_url))
}
fn looks_like_vertex_openai_compat_base(base_url: &str) -> bool {
let Ok(parsed) = Url::parse(base_url.trim()) else {
return false;
};
parsed
.path()
.trim_end_matches('/')
.ends_with("/endpoints/openapi")
}
#[cfg(test)]
mod tests {
use super::{
is_vertex_api_key_transport_context, is_vertex_service_account_transport_context,
is_vertex_transport_context, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Vertex".to_string(),
provider_type: "custom".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "gemini:generate_content".to_string(),
api_family: Some("gemini".to_string()),
endpoint_kind: Some("cli".to_string()),
is_active: true,
base_url: "https://aiplatform.googleapis.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: Some("/v1/publishers/google/models/{model}:{action}".to_string()),
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: Some(vec!["gemini:generate_content".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "vertex-secret".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn detects_vertex_host() {
assert!(looks_like_vertex_ai_host(
"https://aiplatform.googleapis.com"
));
assert!(looks_like_vertex_ai_host(
"https://us-central1-aiplatform.googleapis.com"
));
assert!(!looks_like_vertex_ai_host("https://example.com"));
}
#[test]
fn infers_vertex_api_key_context_for_custom_aiplatform_transport() {
assert!(is_vertex_api_key_transport_context(&sample_transport()));
}
#[test]
fn rejects_non_api_key_custom_aiplatform_transport() {
let mut transport = sample_transport();
transport.key.auth_type = "bearer".to_string();
assert!(!is_vertex_api_key_transport_context(&transport));
}
#[test]
fn infers_vertex_service_account_context_for_fixed_provider() {
let mut transport = sample_transport();
transport.provider.provider_type = "vertex_ai".to_string();
transport.key.auth_type = "service_account".to_string();
assert!(is_vertex_service_account_transport_context(&transport));
assert!(is_vertex_transport_context(&transport));
assert!(!is_vertex_api_key_transport_context(&transport));
}
#[test]
fn detects_vertex_query_auth_usage_for_gemini_formats() {
let transport = sample_transport();
assert!(uses_vertex_api_key_query_auth(
&transport,
"gemini:generate_content"
));
assert!(!uses_vertex_api_key_query_auth(
&transport,
"claude:messages"
));
}
#[test]
fn infers_vertex_service_account_context_for_openai_compat_endpoint_root() {
let mut transport = sample_transport();
transport.endpoint.api_format = "openai:chat".to_string();
transport.endpoint.base_url =
"https://aiplatform.googleapis.com/v1/projects/project-1/locations/global/endpoints/openapi"
.to_string();
transport.key.auth_type = "service_account".to_string();
assert!(is_vertex_service_account_transport_context(&transport));
assert!(is_vertex_transport_context(&transport));
}
#[test]
fn does_not_infer_vertex_context_for_generic_openai_format_on_aiplatform_root() {
let mut transport = sample_transport();
transport.endpoint.api_format = "openai:chat".to_string();
transport.endpoint.base_url = "https://aiplatform.googleapis.com".to_string();
transport.key.auth_type = "service_account".to_string();
assert!(!is_vertex_service_account_transport_context(&transport));
assert!(!is_vertex_transport_context(&transport));
}
}
@@ -0,0 +1,33 @@
mod auth;
mod context;
mod policy;
mod url;
pub use auth::{
parse_vertex_service_account_auth_config, resolve_local_vertex_api_key_query_auth,
resolve_local_vertex_service_account_auth_config,
supports_local_vertex_service_account_auth_resolution, VertexApiKeyQueryAuth,
VertexServiceAccountAuthConfig, VertexServiceAccountRefreshAdapter, VERTEX_API_KEY_QUERY_PARAM,
VERTEX_SERVICE_ACCOUNT_AUTH_HEADER, VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
};
pub use context::{
is_vertex_api_key_transport_context, is_vertex_service_account_transport_context,
is_vertex_transport_context, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth,
};
pub use policy::{
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
local_vertex_gemini_transport_unsupported_reason_with_network,
supports_local_vertex_api_key_gemini_transport,
supports_local_vertex_api_key_gemini_transport_with_network,
supports_local_vertex_api_key_imagen_transport,
supports_local_vertex_api_key_imagen_transport_with_network,
supports_local_vertex_gemini_transport_with_network,
};
pub use url::{
build_vertex_api_key_gemini_content_url, build_vertex_api_key_gemini_embedding_url,
build_vertex_api_key_imagen_content_url, build_vertex_service_account_gemini_content_url,
build_vertex_service_account_gemini_embedding_url, resolve_vertex_service_account_region,
VERTEX_API_KEY_BASE_URL,
};
pub const PROVIDER_TYPE: &str = "vertex_ai";
@@ -0,0 +1,356 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::super::{
body_rules_are_locally_supported, header_rules_are_locally_supported,
resolve_transport_profile, supports_local_oauth_request_auth_resolution,
transport_profile_is_configured, transport_proxy_is_locally_supported,
};
use super::auth::{
resolve_local_vertex_api_key_query_auth, supports_local_vertex_service_account_auth_resolution,
};
fn is_vertex_transport_family(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
|| super::looks_like_vertex_ai_host(&transport.endpoint.base_url)
}
pub fn local_vertex_api_key_gemini_transport_unsupported_reason_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> Option<&'static str> {
local_vertex_gemini_transport_unsupported_reason_with_network_impl(transport, true)
}
pub fn local_vertex_gemini_transport_unsupported_reason_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> Option<&'static str> {
local_vertex_gemini_transport_unsupported_reason_with_network_impl(transport, false)
}
fn local_vertex_gemini_transport_unsupported_reason_with_network_impl(
transport: &GatewayProviderTransportSnapshot,
require_api_key: bool,
) -> Option<&'static str> {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return if !transport.provider.is_active {
Some("provider_inactive")
} else if !transport.endpoint.is_active {
Some("endpoint_inactive")
} else {
Some("key_inactive")
};
}
let endpoint_api_format =
aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
if !matches!(
endpoint_api_format.as_str(),
"gemini:generate_content" | "gemini:embedding"
) {
return Some("transport_api_format_mismatch");
}
if !is_vertex_transport_family(transport) {
return Some("transport_provider_type_unsupported");
}
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
return Some("transport_header_rules_unsupported");
}
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
return Some("transport_body_rules_unsupported");
}
let has_api_key_auth = resolve_local_vertex_api_key_query_auth(transport).is_some();
let has_service_account_auth = supports_local_vertex_service_account_auth_resolution(transport)
&& supports_local_oauth_request_auth_resolution(transport);
if require_api_key {
if !has_api_key_auth {
return Some("transport_auth_unavailable");
}
} else if !has_api_key_auth && !has_service_account_auth {
return Some("transport_auth_unavailable");
}
if !transport_proxy_is_locally_supported(transport) {
return Some("transport_proxy_unsupported");
}
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
{
return Some("transport_profile_unsupported");
}
None
}
pub fn supports_local_vertex_api_key_gemini_transport(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
supports_local_vertex_api_key_same_format_transport(
transport,
&["gemini:generate_content"],
false,
)
}
pub fn supports_local_vertex_api_key_gemini_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(transport).is_none()
}
pub fn supports_local_vertex_gemini_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
local_vertex_gemini_transport_unsupported_reason_with_network(transport).is_none()
}
pub fn supports_local_vertex_api_key_imagen_transport(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
supports_local_vertex_api_key_same_format_transport(
transport,
&["gemini:generate_content"],
false,
)
}
pub fn supports_local_vertex_api_key_imagen_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
supports_local_vertex_api_key_same_format_transport(
transport,
&["gemini:generate_content"],
true,
)
}
fn supports_local_vertex_api_key_same_format_transport(
transport: &GatewayProviderTransportSnapshot,
api_formats: &[&str],
allow_network_passthrough: bool,
) -> bool {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return false;
}
let endpoint_api_format =
aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
if !api_formats
.iter()
.any(|api_format| endpoint_api_format.eq_ignore_ascii_case(api_format))
{
return false;
}
if !super::is_vertex_api_key_transport_context(transport) {
return false;
}
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref())
|| !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref())
{
return false;
}
if resolve_local_vertex_api_key_query_auth(transport).is_none() {
return false;
}
let has_custom_path = transport
.endpoint
.custom_path
.as_deref()
.is_some_and(|value: &str| !value.trim().is_empty());
if has_custom_path && !allow_network_passthrough {
return false;
}
if allow_network_passthrough {
if !transport_proxy_is_locally_supported(transport) {
return false;
}
if transport_profile_is_configured(transport)
&& resolve_transport_profile(transport).is_none()
{
return false;
}
} else if transport.provider.proxy.is_some()
|| transport.endpoint.proxy.is_some()
|| transport.key.proxy.is_some()
|| transport_profile_is_configured(transport)
{
return false;
}
true
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
local_vertex_gemini_transport_unsupported_reason_with_network,
supports_local_vertex_api_key_gemini_transport,
supports_local_vertex_api_key_gemini_transport_with_network,
supports_local_vertex_gemini_transport_with_network,
};
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Vertex".to_string(),
provider_type: "vertex_ai".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "gemini:generate_content".to_string(),
api_family: Some("gemini".to_string()),
endpoint_kind: Some("generate_content".to_string()),
is_active: true,
base_url: "https://aiplatform.googleapis.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: Some(vec!["gemini:generate_content".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "vertex-secret".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn supports_vertex_api_key_same_format_subset() {
assert!(supports_local_vertex_api_key_gemini_transport(
&sample_transport()
));
}
#[test]
fn supports_vertex_api_key_gemini_generate_content_subset() {
let mut transport = sample_transport();
transport.endpoint.api_format = "gemini:generate_content".to_string();
assert!(supports_local_vertex_api_key_gemini_transport(&transport));
}
#[test]
fn supports_custom_aiplatform_gemini_generate_content_subset() {
let mut transport = sample_transport();
transport.provider.provider_type = "custom".to_string();
transport.endpoint.api_format = "gemini:generate_content".to_string();
assert!(supports_local_vertex_api_key_gemini_transport(&transport));
}
#[test]
fn rejects_vertex_service_account_from_api_key_subset() {
let mut transport = sample_transport();
transport.key.auth_type = "service_account".to_string();
transport.key.decrypted_auth_config = Some("{\"project_id\":\"demo-project\"}".to_string());
assert!(!supports_local_vertex_api_key_gemini_transport(&transport));
}
#[test]
fn supports_vertex_service_account_gemini_transport_with_network() {
let mut transport = sample_transport();
transport.key.auth_type = "service_account".to_string();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"client_email":"[email protected]",
"private_key":"TEST-PRIVATE-KEY",
"project_id":"demo-project"
}"#
.to_string(),
);
assert!(!supports_local_vertex_api_key_gemini_transport_with_network(&transport));
assert!(supports_local_vertex_gemini_transport_with_network(
&transport
));
}
#[test]
fn supports_vertex_service_account_gemini_embedding_transport_with_network() {
let mut transport = sample_transport();
transport.endpoint.api_format = "gemini:embedding".to_string();
transport.endpoint.endpoint_kind = Some("embedding".to_string());
transport.key.api_formats = Some(vec!["gemini:embedding".to_string()]);
transport.key.auth_type = "service_account".to_string();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"client_email":"[email protected]",
"private_key":"TEST-PRIVATE-KEY",
"project_id":"demo-project"
}"#
.to_string(),
);
assert!(supports_local_vertex_gemini_transport_with_network(
&transport
));
}
#[test]
fn allows_network_passthrough_for_custom_path_with_local_proxy_support() {
let mut transport = sample_transport();
transport.endpoint.custom_path =
Some("/v1/publishers/google/models/gemini-2.5-pro:generateContent".to_string());
transport.key.proxy = Some(json!({"url":"http://proxy.example:8080"}));
transport.key.fingerprint = Some(json!({"transport_profile":"chrome_136"}));
assert!(!supports_local_vertex_api_key_gemini_transport(&transport));
assert!(supports_local_vertex_api_key_gemini_transport_with_network(
&transport
));
}
#[test]
fn reports_auth_unavailable_when_vertex_api_key_query_auth_is_missing() {
let mut transport = sample_transport();
transport.key.auth_type = "service_account".to_string();
transport.key.decrypted_api_key = "__placeholder__".to_string();
assert_eq!(
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(&transport),
Some("transport_auth_unavailable")
);
assert_eq!(
local_vertex_gemini_transport_unsupported_reason_with_network(&transport),
Some("transport_auth_unavailable")
);
}
}
@@ -0,0 +1,327 @@
use std::collections::BTreeMap;
use url::form_urlencoded;
use super::super::url::build_passthrough_path_url;
use super::auth::VertexServiceAccountAuthConfig;
pub const VERTEX_API_KEY_BASE_URL: &str = "https://aiplatform.googleapis.com";
pub fn build_vertex_api_key_gemini_content_url(
model: &str,
stream: bool,
api_key: &str,
request_query: Option<&str>,
) -> Option<String> {
let action = if stream {
"streamGenerateContent"
} else {
"generateContent"
};
build_vertex_api_key_google_model_url(model, action, stream, api_key, request_query)
}
pub fn build_vertex_api_key_imagen_content_url(
model: &str,
stream: bool,
api_key: &str,
request_query: Option<&str>,
) -> Option<String> {
let action = if stream {
"streamGenerateContent"
} else {
"generateContent"
};
build_vertex_api_key_google_model_url(model, action, stream, api_key, request_query)
}
pub fn build_vertex_api_key_gemini_embedding_url(
model: &str,
api_key: &str,
request_query: Option<&str>,
) -> Option<String> {
build_vertex_api_key_google_model_url(model, "predict", false, api_key, request_query)
}
pub fn build_vertex_service_account_gemini_content_url(
model: &str,
stream: bool,
auth_config: &VertexServiceAccountAuthConfig,
request_query: Option<&str>,
) -> Option<String> {
let action = if stream {
"streamGenerateContent"
} else {
"generateContent"
};
build_vertex_service_account_google_model_url(model, action, stream, auth_config, request_query)
}
pub fn build_vertex_service_account_gemini_embedding_url(
model: &str,
auth_config: &VertexServiceAccountAuthConfig,
request_query: Option<&str>,
) -> Option<String> {
build_vertex_service_account_google_model_url(
model,
"predict",
false,
auth_config,
request_query,
)
}
fn build_vertex_api_key_google_model_url(
model: &str,
action: &str,
stream: bool,
api_key: &str,
request_query: Option<&str>,
) -> Option<String> {
let trimmed_model = model.trim();
let trimmed_action = action.trim();
let trimmed_api_key = api_key.trim();
if trimmed_model.is_empty() || trimmed_action.is_empty() || trimmed_api_key.is_empty() {
return None;
}
let path = format!("/v1/publishers/google/models/{trimmed_model}:{trimmed_action}");
let merged_query = build_vertex_api_key_query(trimmed_api_key, request_query, stream);
build_passthrough_path_url(VERTEX_API_KEY_BASE_URL, &path, merged_query.as_deref(), &[])
}
fn build_vertex_service_account_google_model_url(
model: &str,
action: &str,
stream: bool,
auth_config: &VertexServiceAccountAuthConfig,
request_query: Option<&str>,
) -> Option<String> {
let trimmed_model = model.trim();
let trimmed_action = action.trim();
let project_id = auth_config.project_id.trim();
if trimmed_model.is_empty() || trimmed_action.is_empty() || project_id.is_empty() {
return None;
}
let region = resolve_vertex_service_account_region(trimmed_model, auth_config);
let base_url = if region == "global" {
VERTEX_API_KEY_BASE_URL.to_string()
} else {
format!("https://{region}-aiplatform.googleapis.com")
};
let path = format!(
"/v1/projects/{project_id}/locations/{region}/publishers/google/models/{trimmed_model}:{trimmed_action}"
);
let merged_query = build_vertex_service_account_query(request_query, stream);
build_passthrough_path_url(&base_url, &path, merged_query.as_deref(), &[])
}
pub fn resolve_vertex_service_account_region(
model: &str,
auth_config: &VertexServiceAccountAuthConfig,
) -> String {
let trimmed_model = model.trim();
if let Some(region) = auth_config
.model_regions
.get(trimmed_model)
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
return region.to_string();
}
if let Some(region) = default_vertex_model_region(trimmed_model) {
return region.to_string();
}
if let Some(region) = auth_config
.region
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
return region.to_string();
}
"global".to_string()
}
fn default_vertex_model_region(model: &str) -> Option<&'static str> {
if model.starts_with("gemini-3.") || model == "gemini-3-pro-image-preview" {
return Some("global");
}
if matches!(
model,
"gemini-2.0-flash"
| "gemini-2.0-flash-exp"
| "gemini-2.0-flash-001"
| "gemini-2.0-pro-exp"
| "gemini-2.0-flash-exp-image-generation"
| "gemini-1.5-pro"
| "gemini-1.5-pro-001"
| "gemini-1.5-pro-002"
| "gemini-1.5-flash"
| "gemini-1.5-flash-001"
| "gemini-1.5-flash-002"
| "imagen-3.0-generate-001"
| "imagen-3.0-fast-generate-001"
) {
return Some("us-central1");
}
None
}
fn build_vertex_api_key_query(
api_key: &str,
request_query: Option<&str>,
stream: bool,
) -> Option<String> {
let mut merged = BTreeMap::new();
merge_query_string(&mut merged, request_query);
merged.remove("beta");
merged.insert("key".to_string(), api_key.to_string());
if stream {
merged
.entry("alt".to_string())
.or_insert_with(|| "sse".to_string());
}
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in merged {
serializer.append_pair(&key, &value);
}
let query = serializer.finish();
if query.is_empty() {
None
} else {
Some(query)
}
}
fn build_vertex_service_account_query(request_query: Option<&str>, stream: bool) -> Option<String> {
let mut merged = BTreeMap::new();
merge_query_string(&mut merged, request_query);
merged.remove("beta");
merged.remove("key");
if stream {
merged
.entry("alt".to_string())
.or_insert_with(|| "sse".to_string());
}
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in merged {
serializer.append_pair(&key, &value);
}
let query = serializer.finish();
if query.is_empty() {
None
} else {
Some(query)
}
}
fn merge_query_string(out: &mut BTreeMap<String, String>, query: Option<&str>) {
let Some(query) = query.map(str::trim).filter(|value| !value.is_empty()) else {
return;
};
for (key, value) in form_urlencoded::parse(query.as_bytes()) {
out.insert(key.into_owned(), value.into_owned());
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use super::{
build_vertex_api_key_gemini_content_url, build_vertex_api_key_imagen_content_url,
build_vertex_service_account_gemini_content_url,
};
use crate::vertex::VertexServiceAccountAuthConfig;
#[test]
fn builds_vertex_gemini_api_key_stream_url() {
assert_eq!(
build_vertex_api_key_gemini_content_url(
"gemini-2.5-pro",
true,
"vertex-secret",
Some("foo=bar&beta=v1")
)
.as_deref(),
Some(
"https://aiplatform.googleapis.com/v1/publishers/google/models/gemini-2.5-pro:streamGenerateContent?alt=sse&foo=bar&key=vertex-secret"
)
);
}
#[test]
fn builds_vertex_imagen_api_key_sync_url() {
assert_eq!(
build_vertex_api_key_imagen_content_url(
"imagen-3.0-generate-001",
false,
"vertex-secret",
Some("view=full")
)
.as_deref(),
Some(
"https://aiplatform.googleapis.com/v1/publishers/google/models/imagen-3.0-generate-001:generateContent?key=vertex-secret&view=full"
)
);
}
#[test]
fn builds_vertex_service_account_gemini_sync_url() {
let auth_config = VertexServiceAccountAuthConfig {
client_email: "[email protected]".to_string(),
private_key: "not-used".to_string(),
project_id: "demo-project".to_string(),
token_uri: "https://oauth2.googleapis.com/token".to_string(),
region: None,
model_regions: BTreeMap::new(),
};
assert_eq!(
build_vertex_service_account_gemini_content_url(
"gemini-3.1-pro-preview",
false,
&auth_config,
Some("foo=bar&beta=1&key=client-key")
)
.as_deref(),
Some(
"https://aiplatform.googleapis.com/v1/projects/demo-project/locations/global/publishers/google/models/gemini-3.1-pro-preview:generateContent?foo=bar"
)
);
}
#[test]
fn builds_vertex_service_account_gemini_stream_url_with_model_region_override() {
let auth_config = VertexServiceAccountAuthConfig {
client_email: "[email protected]".to_string(),
private_key: "not-used".to_string(),
project_id: "demo-project".to_string(),
token_uri: "https://oauth2.googleapis.com/token".to_string(),
region: Some("global".to_string()),
model_regions: BTreeMap::from([(
"gemini-2.0-flash".to_string(),
"us-central1".to_string(),
)]),
};
assert_eq!(
build_vertex_service_account_gemini_content_url(
"gemini-2.0-flash",
true,
&auth_config,
Some("foo=bar")
)
.as_deref(),
Some(
"https://us-central1-aiplatform.googleapis.com/v1/projects/demo-project/locations/us-central1/publishers/google/models/gemini-2.0-flash:streamGenerateContent?alt=sse&foo=bar"
)
);
}
}
@@ -0,0 +1,520 @@
use std::collections::BTreeMap;
use aether_data_contracts::repository::video_tasks::StoredVideoTask;
use aether_video_tasks_core::{
LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput,
};
use async_trait::async_trait;
use serde_json::Value;
use super::auth::{
build_passthrough_headers_with_auth, resolve_local_gemini_auth,
resolve_local_openai_bearer_auth,
};
use super::network::{resolve_transport_execution_timeouts, resolve_transport_profile};
use super::policy::{
local_gemini_transport_unsupported_reason_with_network,
local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport,
supports_local_standard_transport,
};
use super::rules::{
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
};
use super::snapshot::GatewayProviderTransportSnapshot;
use super::url::{build_gemini_video_predict_long_running_url, build_passthrough_path_url};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderVideoCreateFamily {
OpenAi,
Gemini,
}
#[derive(Debug, Clone, Copy)]
pub struct ProviderVideoCreateHeadersInput<'a> {
pub headers: &'a http::HeaderMap,
pub auth_header: &'a str,
pub auth_value: &'a str,
pub header_rules: Option<&'a Value>,
pub provider_request_body: &'a Value,
pub original_request_body: &'a Value,
}
#[async_trait]
pub trait VideoTaskTransportSnapshotLookup: Send + Sync {
async fn read_video_task_provider_transport_snapshot(
&self,
provider_id: &str,
endpoint_id: &str,
key_id: &str,
) -> Result<Option<GatewayProviderTransportSnapshot>, String>;
}
pub fn resolve_local_video_task_transport(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
model_name: Option<String>,
) -> Option<LocalVideoTaskTransport> {
let api_format = api_format.trim();
let (auth_header, auth_value) = match api_format {
"openai:video" => {
if !supports_local_standard_transport(transport, api_format) {
return None;
}
resolve_local_openai_bearer_auth(transport)?
}
"gemini:video" => {
if !supports_local_gemini_transport(transport, api_format) {
return None;
}
resolve_local_gemini_auth(transport)?
}
_ => return None,
};
Some(LocalVideoTaskTransport::from_bridge_input(
LocalVideoTaskTransportBridgeInput {
upstream_base_url: transport.endpoint.base_url.clone(),
provider_name: Some(transport.provider.name.clone()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
auth_header,
auth_value,
content_type: Some("application/json".to_string()),
model_name,
proxy: None,
transport_profile: resolve_transport_profile(transport),
timeouts: resolve_transport_execution_timeouts(transport),
},
))
}
pub fn video_create_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
family: ProviderVideoCreateFamily,
api_format: &str,
) -> Option<&'static str> {
match family {
ProviderVideoCreateFamily::OpenAi => {
local_standard_transport_unsupported_reason_with_network(transport, api_format)
}
ProviderVideoCreateFamily::Gemini => {
local_gemini_transport_unsupported_reason_with_network(transport, api_format)
}
}
}
pub fn resolve_video_create_auth(
transport: &GatewayProviderTransportSnapshot,
family: ProviderVideoCreateFamily,
) -> Option<(String, String)> {
match family {
ProviderVideoCreateFamily::OpenAi => resolve_local_openai_bearer_auth(transport),
ProviderVideoCreateFamily::Gemini => resolve_local_gemini_auth(transport),
}
}
pub fn build_video_create_request_body(
body_json: &Value,
family: ProviderVideoCreateFamily,
mapped_model: &str,
body_rules: Option<&Value>,
request_headers: Option<&http::HeaderMap>,
) -> Option<Value> {
let mut provider_request_body = match family {
ProviderVideoCreateFamily::OpenAi => {
let mut provider_request_body = body_json.as_object().cloned().unwrap_or_default();
provider_request_body
.insert("model".to_string(), Value::String(mapped_model.to_string()));
Value::Object(provider_request_body)
}
ProviderVideoCreateFamily::Gemini => body_json.clone(),
};
if !apply_local_body_rules_with_request_headers(
&mut provider_request_body,
body_rules,
Some(body_json),
request_headers,
) {
return None;
}
Some(provider_request_body)
}
pub fn build_video_create_upstream_url(
transport: &GatewayProviderTransportSnapshot,
request_path: &str,
request_query: Option<&str>,
mapped_model: &str,
family: ProviderVideoCreateFamily,
) -> Option<String> {
let custom_path = transport
.endpoint
.custom_path
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty());
if let Some(path) = custom_path {
let blocked_keys = match family {
ProviderVideoCreateFamily::OpenAi => &[][..],
ProviderVideoCreateFamily::Gemini => &["key"][..],
};
return build_passthrough_path_url(
&transport.endpoint.base_url,
path,
request_query,
blocked_keys,
);
}
match family {
ProviderVideoCreateFamily::OpenAi => build_passthrough_path_url(
&transport.endpoint.base_url,
openai_video_api_root_request_path(request_path),
request_query,
&[],
),
ProviderVideoCreateFamily::Gemini => build_gemini_video_predict_long_running_url(
&transport.endpoint.base_url,
mapped_model,
request_query,
),
}
}
fn openai_video_api_root_request_path(request_path: &str) -> &str {
if request_path.starts_with("/v1/") {
&request_path[3..]
} else {
request_path
}
}
pub fn build_video_create_headers(
input: ProviderVideoCreateHeadersInput<'_>,
) -> Option<BTreeMap<String, String>> {
let mut provider_request_headers = build_passthrough_headers_with_auth(
input.headers,
input.auth_header,
input.auth_value,
&BTreeMap::new(),
);
if !apply_local_header_rules_with_request_headers(
&mut provider_request_headers,
input.header_rules,
&[input.auth_header, "content-type"],
input.provider_request_body,
Some(input.original_request_body),
Some(input.headers),
) {
return None;
}
Some(provider_request_headers)
}
pub async fn reconstruct_local_video_task_snapshot(
lookup: &dyn VideoTaskTransportSnapshotLookup,
task: &StoredVideoTask,
) -> Result<Option<LocalVideoTaskSnapshot>, String> {
let provider_api_format = task
.provider_api_format
.as_deref()
.unwrap_or_default()
.trim();
if !matches!(provider_api_format, "openai:video" | "gemini:video") {
return Ok(None);
}
let Some(provider_id) = task.provider_id.as_deref() else {
return Ok(None);
};
let Some(endpoint_id) = task.endpoint_id.as_deref() else {
return Ok(None);
};
let Some(key_id) = task.key_id.as_deref() else {
return Ok(None);
};
let Some(transport) = lookup
.read_video_task_provider_transport_snapshot(provider_id, endpoint_id, key_id)
.await?
else {
return Ok(None);
};
let Some(local_transport) =
resolve_local_video_task_transport(&transport, provider_api_format, task.model.clone())
else {
return Ok(None);
};
Ok(LocalVideoTaskSnapshot::from_stored_task_with_transport(
task,
local_transport,
))
}
#[cfg(test)]
mod tests {
use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskStatus};
use aether_video_tasks_core::LocalVideoTaskSnapshot;
use async_trait::async_trait;
use serde_json::json;
use super::{
build_video_create_headers, build_video_create_request_body,
build_video_create_upstream_url, reconstruct_local_video_task_snapshot,
resolve_local_video_task_transport, ProviderVideoCreateFamily,
ProviderVideoCreateHeadersInput, VideoTaskTransportSnapshotLookup,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport(api_format: &str, auth_type: &str) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Provider One".to_string(),
provider_type: "openai".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: Some(30.0),
stream_first_byte_timeout_secs: Some(5.0),
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: api_format.to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://example.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: auth_type.to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
upstream_metadata: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
}
}
fn sample_stored_video_task() -> StoredVideoTask {
StoredVideoTask {
id: "task-1".to_string(),
short_id: Some("short-1".to_string()),
request_id: "request-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
username: Some("user".to_string()),
api_key_name: Some("key".to_string()),
external_task_id: Some("upstream-task-1".to_string()),
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("key-1".to_string()),
client_api_format: Some("openai:video".to_string()),
provider_api_format: Some("openai:video".to_string()),
format_converted: false,
model: Some("sora".to_string()),
prompt: Some("generate".to_string()),
original_request_body: Some(json!({"prompt": "generate"})),
duration_seconds: None,
resolution: None,
aspect_ratio: None,
size: Some("1024x1024".to_string()),
status: VideoTaskStatus::Submitted,
progress_percent: 0,
progress_message: None,
retry_count: 0,
poll_interval_seconds: 10,
next_poll_at_unix_secs: None,
poll_count: 0,
max_poll_count: 360,
created_at_unix_ms: 1,
submitted_at_unix_secs: Some(1),
completed_at_unix_secs: None,
updated_at_unix_secs: 1,
error_code: None,
error_message: None,
video_url: None,
request_metadata: None,
}
}
struct TestLookup(Option<GatewayProviderTransportSnapshot>);
#[async_trait]
impl VideoTaskTransportSnapshotLookup for TestLookup {
async fn read_video_task_provider_transport_snapshot(
&self,
_provider_id: &str,
_endpoint_id: &str,
_key_id: &str,
) -> Result<Option<GatewayProviderTransportSnapshot>, String> {
Ok(self.0.clone())
}
}
#[test]
fn resolves_openai_video_transport() {
let transport = resolve_local_video_task_transport(
&sample_transport("openai:video", "api_key"),
"openai:video",
Some("sora".to_string()),
)
.expect("transport");
assert_eq!(
transport.headers.get("authorization").map(String::as_str),
Some("Bearer secret")
);
assert_eq!(transport.model_name.as_deref(), Some("sora"));
assert_eq!(transport.provider_id, "provider-1");
}
#[test]
fn resolves_gemini_video_transport() {
let transport = resolve_local_video_task_transport(
&sample_transport("gemini:video", "api_key"),
"gemini:video",
Some("veo".to_string()),
)
.expect("transport");
assert_eq!(
transport.headers.get("x-goog-api-key").map(String::as_str),
Some("secret")
);
assert_eq!(transport.model_name.as_deref(), Some("veo"));
assert_eq!(transport.endpoint_id, "endpoint-1");
}
#[test]
fn builds_openai_video_create_request_body_with_mapped_model() {
let body = build_video_create_request_body(
&json!({"prompt": "make a clip", "model": "client-model"}),
ProviderVideoCreateFamily::OpenAi,
"upstream-video-model",
None,
None,
)
.expect("body should build");
assert_eq!(body.get("prompt"), Some(&json!("make a clip")));
assert_eq!(body.get("model"), Some(&json!("upstream-video-model")));
}
#[test]
fn builds_openai_video_create_url_from_api_root_base() {
let mut transport = sample_transport("openai:video", "bearer");
transport.endpoint.base_url = "https://api.openai.example/v1".to_string();
let url = build_video_create_upstream_url(
&transport,
"/v1/videos",
Some("trace=1"),
"sora-upstream",
ProviderVideoCreateFamily::OpenAi,
)
.expect("url should build");
assert_eq!(url, "https://api.openai.example/v1/videos?trace=1");
}
#[test]
fn builds_gemini_video_create_url_and_removes_client_key_query() {
let transport = sample_transport("gemini:video", "api_key");
let url = build_video_create_upstream_url(
&transport,
"/v1beta/models/client-model:predictLongRunning",
Some("key=client-key&trace=1"),
"veo-upstream",
ProviderVideoCreateFamily::Gemini,
)
.expect("url should build");
assert_eq!(
url,
"https://example.com/v1beta/models/veo-upstream:predictLongRunning?trace=1"
);
}
#[test]
fn builds_video_create_headers_with_auth_and_rules() {
let provider_request_body = json!({"prompt": "make a clip"});
let original_request_body = provider_request_body.clone();
let headers = build_video_create_headers(ProviderVideoCreateHeadersInput {
headers: &http::HeaderMap::new(),
auth_header: "authorization",
auth_value: "Bearer secret",
header_rules: Some(&json!([
{"action":"set","key":"x-provider-tag","value":"video"}
])),
provider_request_body: &provider_request_body,
original_request_body: &original_request_body,
})
.expect("headers should build");
assert_eq!(
headers.get("authorization").map(String::as_str),
Some("Bearer secret")
);
assert_eq!(
headers.get("x-provider-tag").map(String::as_str),
Some("video")
);
}
#[test]
fn rejects_mismatched_video_transport_format() {
let transport = sample_transport("openai:chat", "bearer");
assert!(resolve_local_video_task_transport(&transport, "openai:video", None).is_none());
}
#[tokio::test]
async fn reconstructs_openai_video_snapshot_via_lookup_trait() {
let lookup = TestLookup(Some(sample_transport("openai:video", "bearer")));
let snapshot = reconstruct_local_video_task_snapshot(&lookup, &sample_stored_video_task())
.await
.expect("lookup should succeed")
.expect("snapshot");
match snapshot {
LocalVideoTaskSnapshot::OpenAi(seed) => {
assert_eq!(seed.transport.provider_id, "provider-1");
assert_eq!(seed.transport.model_name.as_deref(), Some("sora"));
}
LocalVideoTaskSnapshot::Gemini(_) => panic!("expected openai snapshot"),
}
}
}
@@ -0,0 +1,582 @@
use std::collections::BTreeMap;
use serde_json::{json, Value};
use uuid::Uuid;
use crate::rules::{
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
body_rules_are_locally_supported, header_rules_are_locally_supported,
};
use crate::snapshot::GatewayProviderTransportSnapshot;
use crate::url::build_passthrough_path_url;
use crate::{
resolve_transport_profile, should_skip_upstream_passthrough_header,
supports_local_oauth_request_auth_resolution, transport_profile_is_configured,
transport_proxy_is_locally_supported,
};
pub mod cascade;
pub mod models;
pub mod proto;
pub const PROVIDER_TYPE: &str = "windsurf";
pub const WINDSURF_ENVELOPE_NAME: &str = "windsurf:GetChatMessage";
pub const GET_CHAT_MESSAGE_PATH: &str = "/exa.api_server_pb.ApiServerService/GetChatMessage";
const DEFAULT_IDE_VERSION: &str = "1.9600.41";
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
pub fn is_windsurf_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(PROVIDER_TYPE)
}
pub fn local_windsurf_request_transport_unsupported_reason_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> Option<&'static str> {
if !transport.provider.is_active {
return Some("provider_inactive");
}
if !transport.endpoint.is_active {
return Some("endpoint_inactive");
}
if !transport.key.is_active {
return Some("key_inactive");
}
if !is_windsurf_provider_transport(transport) {
return Some("transport_provider_type_unsupported");
}
if !transport
.endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:chat")
{
return Some("transport_api_format_mismatch");
}
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
return Some("transport_header_rules_unsupported");
}
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
return Some("transport_body_rules_unsupported");
}
if transport.key.decrypted_auth_config.is_some()
&& !supports_local_oauth_request_auth_resolution(transport)
&& !supports_local_windsurf_request_auth_resolution(transport)
{
return Some("transport_oauth_resolution_unsupported");
}
if !transport_proxy_is_locally_supported(transport) {
return Some("transport_proxy_unsupported");
}
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
{
return Some("transport_profile_unsupported");
}
None
}
pub fn supports_local_windsurf_request_auth_resolution(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
resolve_windsurf_cascade_auth(transport).is_some()
}
pub fn resolve_windsurf_cascade_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
if !is_windsurf_provider_transport(transport) {
return None;
}
let auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
if !matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer") {
return None;
}
let secret = transport.key.decrypted_api_key.trim();
if secret.is_empty() || secret == PLACEHOLDER_API_KEY {
return None;
}
Some(("authorization".to_string(), format!("Bearer {secret}")))
}
pub fn build_windsurf_cascade_upstream_url(
upstream_base_url: &str,
query: Option<&str>,
) -> Option<String> {
build_passthrough_path_url(upstream_base_url, GET_CHAT_MESSAGE_PATH, query, &[])
}
pub fn build_windsurf_cascade_request_body(
body_json: &Value,
mapped_model: &str,
auth_value: &str,
body_rules: Option<&Value>,
request_headers: Option<&http::HeaderMap>,
upstream_is_stream: bool,
) -> Option<Value> {
let mapped_model = mapped_model.trim();
if mapped_model.is_empty() {
return None;
}
let messages = body_json.get("messages")?.as_array()?.clone();
if messages.is_empty() {
return None;
}
let conversation_id =
extract_conversation_id(body_json).unwrap_or_else(|| Uuid::new_v4().to_string());
let message_text =
latest_message_snapshot_text(&messages).unwrap_or_else(|| "Continue.".to_string());
let mut provider_request_body = json!({
"metadata": windsurf_metadata_from_auth(auth_value),
"model": mapped_model,
"modelName": mapped_model,
"stream": upstream_is_stream,
"conversationId": conversation_id,
"message": message_text,
"messages": messages,
});
if let Some(max_tokens) = body_json
.get("max_tokens")
.or_else(|| body_json.get("maxTokens"))
{
provider_request_body
.as_object_mut()?
.insert("maxTokens".to_string(), max_tokens.clone());
}
if let Some(temperature) = body_json.get("temperature") {
provider_request_body
.as_object_mut()?
.insert("temperature".to_string(), temperature.clone());
}
if let Some(top_p) = body_json.get("top_p").or_else(|| body_json.get("topP")) {
provider_request_body
.as_object_mut()?
.insert("topP".to_string(), top_p.clone());
}
for field in [
"tools",
"tool_choice",
"toolChoice",
"parallel_tool_calls",
"response_format",
] {
if let Some(value) = body_json.get(field) {
provider_request_body
.as_object_mut()?
.insert(field.to_string(), value.clone());
}
}
if !apply_local_body_rules_with_request_headers(
&mut provider_request_body,
body_rules,
Some(body_json),
request_headers,
) {
return None;
}
Some(provider_request_body)
}
pub fn build_windsurf_cascade_headers(
headers: &http::HeaderMap,
provider_request_body: &Value,
original_request_body: &Value,
header_rules: Option<&Value>,
auth_header: &str,
auth_value: &str,
_upstream_is_stream: bool,
) -> Option<BTreeMap<String, String>> {
let mut out = BTreeMap::new();
for (name, value) in headers {
let Ok(value) = value.to_str() else {
continue;
};
let key = name.as_str().to_ascii_lowercase();
if should_skip_upstream_passthrough_header(&key) {
continue;
}
let value = value.trim();
if !value.is_empty() {
out.insert(key, value.to_string());
}
}
let auth_header = auth_header.trim().to_ascii_lowercase();
if !apply_local_header_rules_with_request_headers(
&mut out,
header_rules,
&[
auth_header.as_str(),
"content-type",
"connect-protocol-version",
],
provider_request_body,
Some(original_request_body),
Some(headers),
) {
return None;
}
out.insert(
"content-type".to_string(),
"application/connect+json".to_string(),
);
out.insert("connect-protocol-version".to_string(), "1".to_string());
out.insert(
"user-agent".to_string(),
format!("windsurf/{DEFAULT_IDE_VERSION}"),
);
out.insert("accept".to_string(), "application/connect+json".to_string());
if !auth_header.is_empty() {
out.insert(auth_header, auth_value.trim().to_string());
}
out.remove("content-length");
Some(out)
}
fn windsurf_metadata_from_auth(auth_value: &str) -> Value {
json!({
"apiKey": auth_secret_from_header_value(auth_value),
"ideName": "windsurf",
"ideVersion": DEFAULT_IDE_VERSION,
"extensionName": "windsurf",
"extensionVersion": DEFAULT_IDE_VERSION,
"locale": "en",
})
}
fn auth_secret_from_header_value(auth_value: &str) -> String {
let value = auth_value.trim();
value
.strip_prefix("Bearer ")
.or_else(|| value.strip_prefix("bearer "))
.unwrap_or(value)
.trim()
.to_string()
}
fn extract_conversation_id(body_json: &Value) -> Option<String> {
let object = body_json.as_object()?;
string_value(object.get("conversation_id"))
.or_else(|| string_value(object.get("conversationId")))
.or_else(|| string_value(object.get("session_id")))
.or_else(|| string_value(object.get("sessionId")))
.or_else(|| {
object
.get("metadata")
.and_then(Value::as_object)
.and_then(|metadata| {
string_value(metadata.get("conversation_id"))
.or_else(|| string_value(metadata.get("conversationId")))
.or_else(|| string_value(metadata.get("session_id")))
.or_else(|| string_value(metadata.get("sessionId")))
})
})
}
fn string_value(value: Option<&Value>) -> Option<String> {
value
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn latest_message_snapshot_text(messages: &[Value]) -> Option<String> {
messages
.iter()
.rev()
.filter_map(Value::as_object)
.find_map(|message| {
let role = message.get("role").and_then(Value::as_str)?;
match role {
"user" => openai_content_to_text(message.get("content"))
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()),
"tool" => {
let content = openai_content_to_text(message.get("content"))
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())?;
let tool_call_id = message
.get("tool_call_id")
.and_then(Value::as_str)
.unwrap_or("unknown");
Some(format!(
"<tool_result tool_call_id=\"{}\">\n{content}\n</tool_result>",
escape_xml_attr(tool_call_id)
))
}
"assistant" => openai_content_to_text(message.get("content"))
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()),
_ => None,
}
})
}
fn openai_content_to_text(value: Option<&Value>) -> Option<String> {
match value? {
Value::String(text) => Some(text.clone()),
Value::Array(items) => {
let parts = items
.iter()
.filter_map(|item| {
item.as_object()
.and_then(|object| object.get("text"))
.and_then(Value::as_str)
})
.map(str::trim)
.filter(|value| !value.is_empty())
.collect::<Vec<_>>();
(!parts.is_empty()).then(|| parts.join("\n"))
}
_ => None,
}
}
fn escape_xml_attr(value: &str) -> String {
value
.replace('&', "&amp;")
.replace('"', "&quot;")
.replace('<', "&lt;")
.replace('>', "&gt;")
}
#[cfg(test)]
mod tests {
use http::HeaderMap;
use serde_json::json;
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
build_windsurf_cascade_headers, build_windsurf_cascade_request_body,
build_windsurf_cascade_upstream_url,
local_windsurf_request_transport_unsupported_reason_with_network,
resolve_windsurf_cascade_auth, GET_CHAT_MESSAGE_PATH,
};
fn sample_windsurf_transport(auth_type: &str) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-windsurf".to_string(),
name: "Windsurf".to_string(),
provider_type: "windsurf".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-windsurf-chat".to_string(),
provider_id: "provider-windsurf".to_string(),
api_format: "openai:chat".to_string(),
api_family: Some("openai".to_string()),
endpoint_kind: Some("chat".to_string()),
is_active: true,
base_url: "https://server.codeium.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-windsurf".to_string(),
provider_id: "provider-windsurf".to_string(),
name: "[email protected]".to_string(),
auth_type: auth_type.to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "devin-session-token$abc".to_string(),
decrypted_auth_config: Some(r#"{"provider_type":"windsurf"}"#.to_string()),
upstream_metadata: None,
},
}
}
#[test]
fn builds_windsurf_cascade_url() {
assert_eq!(
build_windsurf_cascade_upstream_url("https://server.codeium.com", Some("debug=1"))
.as_deref(),
Some(
"https://server.codeium.com/exa.api_server_pb.ApiServerService/GetChatMessage?debug=1"
)
);
assert!(GET_CHAT_MESSAGE_PATH.ends_with("/GetChatMessage"));
}
#[test]
fn builds_cascade_request_body_with_metadata_and_messages() {
let body = build_windsurf_cascade_request_body(
&json!({
"model": "gpt-5",
"conversation_id": "conv-1",
"messages": [
{"role": "system", "content": "brief"},
{"role": "user", "content": [{"type": "text", "text": "hello"}]}
],
"max_tokens": 128
}),
"windsurf-model",
"Bearer devin-session-token$abc",
None,
None,
true,
)
.expect("body should build");
assert_eq!(body["metadata"]["apiKey"], json!("devin-session-token$abc"));
assert_eq!(body["modelName"], json!("windsurf-model"));
assert_eq!(body["stream"], json!(true));
assert_eq!(body["conversationId"], json!("conv-1"));
assert_eq!(body["message"], json!("hello"));
assert_eq!(body["maxTokens"], json!(128));
}
#[test]
fn preserves_openai_tool_fields_for_native_windsurf_runtime() {
let body = build_windsurf_cascade_request_body(
&json!({
"model": "gpt-5-5-low",
"messages": [
{"role": "user", "content": "read Cargo.toml"}
],
"tools": [{
"type": "function",
"function": {
"name": "Read",
"parameters": {
"type": "object",
"properties": {"file_path": {"type": "string"}}
}
}
}],
"tool_choice": "required",
"parallel_tool_calls": false,
"response_format": {"type": "json_object"}
}),
"gpt-5-5-low",
"Bearer devin-session-token$abc",
None,
None,
false,
)
.expect("body should build");
assert_eq!(body["tools"][0]["function"]["name"], json!("Read"));
assert_eq!(body["tool_choice"], json!("required"));
assert_eq!(body["parallel_tool_calls"], json!(false));
assert_eq!(body["response_format"]["type"], json!("json_object"));
}
#[test]
fn builds_cascade_request_body_message_snapshot_from_latest_tool_result() {
let body = build_windsurf_cascade_request_body(
&json!({
"model": "gpt-5-5-low",
"messages": [
{"role": "user", "content": "read Cargo.toml"},
{
"role": "assistant",
"content": null,
"tool_calls": [{
"id": "call_1",
"type": "function",
"function": {"name": "Read", "arguments": "{\"file_path\":\"Cargo.toml\"}"}
}]
},
{"role": "tool", "tool_call_id": "call_1", "content": "workspace Cargo.toml content"}
]
}),
"gpt-5-5-low",
"Bearer devin-session-token$abc",
None,
None,
false,
)
.expect("body should build");
assert!(body["message"]
.as_str()
.expect("message should be a string")
.contains(r#"<tool_result tool_call_id="call_1">"#));
assert!(body["message"]
.as_str()
.expect("message should be a string")
.contains("workspace Cargo.toml content"));
}
#[test]
fn builds_cascade_headers_with_connect_protocol_and_auth() {
let headers = build_windsurf_cascade_headers(
&HeaderMap::new(),
&json!({"metadata": {"apiKey": "secret"}}),
&json!({"messages": []}),
None,
"authorization",
"Bearer secret",
false,
)
.expect("headers should build");
assert_eq!(
headers.get("connect-protocol-version").map(String::as_str),
Some("1")
);
assert_eq!(
headers.get("authorization").map(String::as_str),
Some("Bearer secret")
);
assert_eq!(
headers.get("content-type").map(String::as_str),
Some("application/connect+json")
);
assert_eq!(
headers.get("accept").map(String::as_str),
Some("application/connect+json")
);
}
#[test]
fn oauth_windsurf_transport_resolves_direct_bearer_auth() {
let transport = sample_windsurf_transport("oauth");
assert_eq!(
local_windsurf_request_transport_unsupported_reason_with_network(&transport),
None
);
assert_eq!(
resolve_windsurf_cascade_auth(&transport),
Some((
"authorization".to_string(),
"Bearer devin-session-token$abc".to_string()
))
);
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,384 @@
#[derive(Debug, Clone, PartialEq)]
pub struct WindsurfModel {
pub canonical_name: &'static str,
pub enum_value: u32,
pub model_uid: Option<&'static str>,
pub credit_multiplier: f32,
pub provider: &'static str,
pub deprecated: bool,
}
#[rustfmt::skip]
const MODELS: &[WindsurfModel] = &[
WindsurfModel { canonical_name: "claude-3.5-sonnet", enum_value: 166, model_uid: None, credit_multiplier: 2.0, provider: "anthropic", deprecated: true },
WindsurfModel { canonical_name: "claude-3.7-sonnet", enum_value: 226, model_uid: None, credit_multiplier: 2.0, provider: "anthropic", deprecated: true },
WindsurfModel { canonical_name: "claude-3.7-sonnet-thinking", enum_value: 227, model_uid: None, credit_multiplier: 3.0, provider: "anthropic", deprecated: true },
WindsurfModel { canonical_name: "claude-4-sonnet", enum_value: 281, model_uid: Some("MODEL_CLAUDE_4_SONNET"), credit_multiplier: 2.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4-sonnet-thinking", enum_value: 282, model_uid: Some("MODEL_CLAUDE_4_SONNET_THINKING"), credit_multiplier: 3.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4-opus", enum_value: 290, model_uid: Some("MODEL_CLAUDE_4_OPUS"), credit_multiplier: 4.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4-opus-thinking", enum_value: 291, model_uid: Some("MODEL_CLAUDE_4_OPUS_THINKING"), credit_multiplier: 5.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4.1-opus", enum_value: 328, model_uid: Some("MODEL_CLAUDE_4_1_OPUS"), credit_multiplier: 4.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4.1-opus-thinking", enum_value: 329, model_uid: Some("MODEL_CLAUDE_4_1_OPUS_THINKING"), credit_multiplier: 5.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4.5-haiku", enum_value: 0, model_uid: Some("MODEL_PRIVATE_11"), credit_multiplier: 1.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4.5-sonnet", enum_value: 353, model_uid: Some("MODEL_PRIVATE_2"), credit_multiplier: 2.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4.5-sonnet-thinking", enum_value: 354, model_uid: Some("MODEL_PRIVATE_3"), credit_multiplier: 3.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4.5-opus", enum_value: 391, model_uid: Some("MODEL_CLAUDE_4_5_OPUS"), credit_multiplier: 4.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-4.5-opus-thinking", enum_value: 392, model_uid: Some("MODEL_CLAUDE_4_5_OPUS_THINKING"), credit_multiplier: 5.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-sonnet-4.6", enum_value: 0, model_uid: Some("claude-sonnet-4-6"), credit_multiplier: 4.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-sonnet-4.6-thinking", enum_value: 0, model_uid: Some("claude-sonnet-4-6-thinking"), credit_multiplier: 6.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-sonnet-4.6-1m", enum_value: 0, model_uid: Some("claude-sonnet-4-6-1m"), credit_multiplier: 12.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-sonnet-4.6-thinking-1m", enum_value: 0, model_uid: Some("claude-sonnet-4-6-thinking-1m"), credit_multiplier: 16.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4.6", enum_value: 0, model_uid: Some("claude-opus-4-6"), credit_multiplier: 6.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4.6-thinking", enum_value: 0, model_uid: Some("claude-opus-4-6-thinking"), credit_multiplier: 8.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4-7-medium", enum_value: 0, model_uid: Some("claude-opus-4-7-medium"), credit_multiplier: 8.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4-7-low", enum_value: 0, model_uid: Some("claude-opus-4-7-low"), credit_multiplier: 6.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4-7-high", enum_value: 0, model_uid: Some("claude-opus-4-7-high"), credit_multiplier: 10.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4-7-xhigh", enum_value: 0, model_uid: Some("claude-opus-4-7-xhigh"), credit_multiplier: 12.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4-7-medium-thinking", enum_value: 0, model_uid: Some("claude-opus-4-7-medium-thinking"), credit_multiplier: 10.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4-7-high-thinking", enum_value: 0, model_uid: Some("claude-opus-4-7-high-thinking"), credit_multiplier: 12.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4-7-xhigh-thinking", enum_value: 0, model_uid: Some("claude-opus-4-7-xhigh-thinking"), credit_multiplier: 16.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "claude-opus-4-7-max", enum_value: 0, model_uid: Some("claude-opus-4-7-max"), credit_multiplier: 16.0, provider: "anthropic", deprecated: false },
WindsurfModel { canonical_name: "gpt-4o", enum_value: 109, model_uid: Some("MODEL_CHAT_GPT_4O_2024_08_06"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-4o-mini", enum_value: 113, model_uid: None, credit_multiplier: 0.5, provider: "openai", deprecated: true },
WindsurfModel { canonical_name: "gpt-4.1", enum_value: 259, model_uid: Some("MODEL_CHAT_GPT_4_1_2025_04_14"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-4.1-mini", enum_value: 260, model_uid: None, credit_multiplier: 0.5, provider: "openai", deprecated: true },
WindsurfModel { canonical_name: "gpt-4.1-nano", enum_value: 261, model_uid: None, credit_multiplier: 0.25, provider: "openai", deprecated: true },
WindsurfModel { canonical_name: "gpt-5", enum_value: 340, model_uid: Some("MODEL_PRIVATE_6"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5-medium", enum_value: 0, model_uid: Some("MODEL_PRIVATE_7"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5-high", enum_value: 0, model_uid: Some("MODEL_PRIVATE_8"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5-mini", enum_value: 337, model_uid: None, credit_multiplier: 0.25, provider: "openai", deprecated: true },
WindsurfModel { canonical_name: "gpt-5-codex", enum_value: 346, model_uid: Some("MODEL_CHAT_GPT_5_CODEX"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1", enum_value: 0, model_uid: Some("MODEL_PRIVATE_12"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-low", enum_value: 0, model_uid: Some("MODEL_PRIVATE_13"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-medium", enum_value: 0, model_uid: Some("MODEL_PRIVATE_14"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-high", enum_value: 0, model_uid: Some("MODEL_PRIVATE_15"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-fast", enum_value: 0, model_uid: Some("MODEL_PRIVATE_20"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-low-fast", enum_value: 0, model_uid: Some("MODEL_PRIVATE_21"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-medium-fast", enum_value: 0, model_uid: Some("MODEL_PRIVATE_22"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-high-fast", enum_value: 0, model_uid: Some("MODEL_PRIVATE_23"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-codex-low", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_LOW"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-codex-medium", enum_value: 0, model_uid: Some("MODEL_PRIVATE_9"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-codex-mini-low", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_MINI_LOW"), credit_multiplier: 0.25, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-codex-mini", enum_value: 0, model_uid: Some("MODEL_PRIVATE_19"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-codex-max-low", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_MAX_LOW"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-codex-max-medium", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_MAX_MEDIUM"), credit_multiplier: 1.25, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.1-codex-max-high", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_MAX_HIGH"), credit_multiplier: 1.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2", enum_value: 401, model_uid: Some("MODEL_GPT_5_2_MEDIUM"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-none", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_NONE"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-low", enum_value: 400, model_uid: Some("MODEL_GPT_5_2_LOW"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-high", enum_value: 402, model_uid: Some("MODEL_GPT_5_2_HIGH"), credit_multiplier: 3.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-xhigh", enum_value: 403, model_uid: Some("MODEL_GPT_5_2_XHIGH"), credit_multiplier: 8.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-none-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_NONE_PRIORITY"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-low-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_LOW_PRIORITY"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-medium-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_MEDIUM_PRIORITY"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-high-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_HIGH_PRIORITY"), credit_multiplier: 6.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-xhigh-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_XHIGH_PRIORITY"), credit_multiplier: 16.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-codex-low", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_LOW"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-codex-medium", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_MEDIUM"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-codex-high", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_HIGH"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-codex-xhigh", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_XHIGH"), credit_multiplier: 3.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-codex-low-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_LOW_PRIORITY"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-codex-medium-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_MEDIUM_PRIORITY"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-codex-high-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_HIGH_PRIORITY"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.2-codex-xhigh-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_XHIGH_PRIORITY"), credit_multiplier: 6.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.3-codex", enum_value: 0, model_uid: Some("gpt-5-3-codex-medium"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-none", enum_value: 0, model_uid: Some("gpt-5-4-none"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-low", enum_value: 0, model_uid: Some("gpt-5-4-low"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-medium", enum_value: 0, model_uid: Some("gpt-5-4-medium"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-high", enum_value: 0, model_uid: Some("gpt-5-4-high"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-xhigh", enum_value: 0, model_uid: Some("gpt-5-4-xhigh"), credit_multiplier: 8.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-mini-low", enum_value: 0, model_uid: Some("gpt-5-4-mini-low"), credit_multiplier: 1.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-mini-medium", enum_value: 0, model_uid: Some("gpt-5-4-mini-medium"), credit_multiplier: 1.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-mini-high", enum_value: 0, model_uid: Some("gpt-5-4-mini-high"), credit_multiplier: 4.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.4-mini-xhigh", enum_value: 0, model_uid: Some("gpt-5-4-mini-xhigh"), credit_multiplier: 12.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5", enum_value: 0, model_uid: Some("gpt-5-5-medium"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-none", enum_value: 0, model_uid: Some("gpt-5-5-none"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-low", enum_value: 0, model_uid: Some("gpt-5-5-low"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-medium", enum_value: 0, model_uid: Some("gpt-5-5-medium"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-high", enum_value: 0, model_uid: Some("gpt-5-5-high"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-xhigh", enum_value: 0, model_uid: Some("gpt-5-5-xhigh"), credit_multiplier: 8.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-none-fast", enum_value: 0, model_uid: Some("gpt-5-5-none-priority"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-low-fast", enum_value: 0, model_uid: Some("gpt-5-5-low-priority"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-medium-fast", enum_value: 0, model_uid: Some("gpt-5-5-medium-priority"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-high-fast", enum_value: 0, model_uid: Some("gpt-5-5-high-priority"), credit_multiplier: 8.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.5-xhigh-fast", enum_value: 0, model_uid: Some("gpt-5-5-xhigh-priority"), credit_multiplier: 16.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.3-codex-low", enum_value: 0, model_uid: Some("gpt-5-3-codex-low"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.3-codex-high", enum_value: 0, model_uid: Some("gpt-5-3-codex-high"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.3-codex-xhigh", enum_value: 0, model_uid: Some("gpt-5-3-codex-xhigh"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.3-codex-low-fast", enum_value: 0, model_uid: Some("gpt-5-3-codex-low-priority"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.3-codex-medium-fast", enum_value: 0, model_uid: Some("gpt-5-3-codex-medium-priority"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.3-codex-high-fast", enum_value: 0, model_uid: Some("gpt-5-3-codex-high-priority"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-5.3-codex-xhigh-fast", enum_value: 0, model_uid: Some("gpt-5-3-codex-xhigh-priority"), credit_multiplier: 6.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gpt-oss-120b", enum_value: 0, model_uid: Some("MODEL_GPT_OSS_120B"), credit_multiplier: 0.25, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "o3-mini", enum_value: 207, model_uid: None, credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "o3", enum_value: 218, model_uid: Some("MODEL_CHAT_O3"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "o3-high", enum_value: 0, model_uid: Some("MODEL_CHAT_O3_HIGH"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "o3-pro", enum_value: 294, model_uid: None, credit_multiplier: 4.0, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "o4-mini", enum_value: 264, model_uid: None, credit_multiplier: 0.5, provider: "openai", deprecated: false },
WindsurfModel { canonical_name: "gemini-2.5-pro", enum_value: 246, model_uid: Some("MODEL_GOOGLE_GEMINI_2_5_PRO"), credit_multiplier: 1.0, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "gemini-2.5-flash", enum_value: 312, model_uid: Some("MODEL_GOOGLE_GEMINI_2_5_FLASH"), credit_multiplier: 0.5, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "gemini-3.0-pro", enum_value: 412, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_PRO_LOW"), credit_multiplier: 1.0, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "gemini-3.0-flash-minimal", enum_value: 0, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_FLASH_MINIMAL"), credit_multiplier: 0.75, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "gemini-3.0-flash-low", enum_value: 0, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_FLASH_LOW"), credit_multiplier: 1.0, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "gemini-3.0-flash", enum_value: 415, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_FLASH_MEDIUM"), credit_multiplier: 1.0, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "gemini-3.0-flash-high", enum_value: 0, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_FLASH_HIGH"), credit_multiplier: 1.75, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "gemini-3.1-pro-low", enum_value: 0, model_uid: Some("gemini-3-1-pro-low"), credit_multiplier: 1.0, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "gemini-3.1-pro-high", enum_value: 0, model_uid: Some("gemini-3-1-pro-high"), credit_multiplier: 2.0, provider: "google", deprecated: false },
WindsurfModel { canonical_name: "deepseek-v3", enum_value: 205, model_uid: None, credit_multiplier: 0.5, provider: "deepseek", deprecated: true },
WindsurfModel { canonical_name: "deepseek-v3-2", enum_value: 409, model_uid: None, credit_multiplier: 0.5, provider: "deepseek", deprecated: true },
WindsurfModel { canonical_name: "deepseek-r1", enum_value: 206, model_uid: None, credit_multiplier: 1.0, provider: "deepseek", deprecated: true },
WindsurfModel { canonical_name: "grok-3", enum_value: 217, model_uid: Some("MODEL_XAI_GROK_3"), credit_multiplier: 1.0, provider: "xai", deprecated: false },
WindsurfModel { canonical_name: "grok-3-mini", enum_value: 234, model_uid: None, credit_multiplier: 0.5, provider: "xai", deprecated: true },
WindsurfModel { canonical_name: "grok-3-mini-thinking", enum_value: 0, model_uid: Some("MODEL_XAI_GROK_3_MINI_REASONING"), credit_multiplier: 0.125, provider: "xai", deprecated: false },
WindsurfModel { canonical_name: "grok-code-fast-1", enum_value: 0, model_uid: Some("MODEL_PRIVATE_4"), credit_multiplier: 0.5, provider: "xai", deprecated: false },
WindsurfModel { canonical_name: "qwen-3", enum_value: 324, model_uid: None, credit_multiplier: 0.5, provider: "alibaba", deprecated: true },
WindsurfModel { canonical_name: "kimi-k2", enum_value: 323, model_uid: Some("MODEL_KIMI_K2"), credit_multiplier: 0.5, provider: "moonshot", deprecated: false },
WindsurfModel { canonical_name: "kimi-k2-thinking", enum_value: 394, model_uid: Some("MODEL_KIMI_K2_THINKING"), credit_multiplier: 1.0, provider: "moonshot", deprecated: false },
WindsurfModel { canonical_name: "kimi-k2.5", enum_value: 0, model_uid: Some("kimi-k2-5"), credit_multiplier: 1.0, provider: "moonshot", deprecated: false },
WindsurfModel { canonical_name: "kimi-k2-6", enum_value: 0, model_uid: Some("kimi-k2-6"), credit_multiplier: 1.0, provider: "moonshot", deprecated: false },
WindsurfModel { canonical_name: "glm-4.7", enum_value: 417, model_uid: Some("MODEL_GLM_4_7"), credit_multiplier: 0.25, provider: "zhipu", deprecated: false },
WindsurfModel { canonical_name: "glm-4.7-fast", enum_value: 418, model_uid: Some("MODEL_GLM_4_7_FAST"), credit_multiplier: 0.5, provider: "zhipu", deprecated: false },
WindsurfModel { canonical_name: "glm-5", enum_value: 0, model_uid: Some("glm-5"), credit_multiplier: 1.5, provider: "zhipu", deprecated: false },
WindsurfModel { canonical_name: "glm-5.1", enum_value: 0, model_uid: Some("glm-5-1"), credit_multiplier: 1.5, provider: "zhipu", deprecated: false },
WindsurfModel { canonical_name: "minimax-m2.5", enum_value: 419, model_uid: Some("MODEL_MINIMAX_M2_1"), credit_multiplier: 1.0, provider: "minimax", deprecated: false },
WindsurfModel { canonical_name: "swe-1.5", enum_value: 377, model_uid: Some("MODEL_SWE_1_5_SLOW"), credit_multiplier: 0.5, provider: "windsurf", deprecated: false },
WindsurfModel { canonical_name: "swe-1.5-fast", enum_value: 359, model_uid: Some("MODEL_SWE_1_5"), credit_multiplier: 0.5, provider: "windsurf", deprecated: false },
WindsurfModel { canonical_name: "swe-1.5-thinking", enum_value: 369, model_uid: Some("MODEL_SWE_1_5_THINKING"), credit_multiplier: 0.75, provider: "windsurf", deprecated: false },
WindsurfModel { canonical_name: "swe-1.6", enum_value: 420, model_uid: Some("MODEL_SWE_1_6"), credit_multiplier: 0.5, provider: "windsurf", deprecated: false },
WindsurfModel { canonical_name: "swe-1.6-fast", enum_value: 421, model_uid: Some("MODEL_SWE_1_6_FAST"), credit_multiplier: 0.5, provider: "windsurf", deprecated: false },
WindsurfModel { canonical_name: "adaptive", enum_value: 0, model_uid: Some("adaptive"), credit_multiplier: 1.0, provider: "windsurf", deprecated: true },
WindsurfModel { canonical_name: "arena-fast", enum_value: 0, model_uid: Some("arena-fast"), credit_multiplier: 0.5, provider: "windsurf", deprecated: true },
WindsurfModel { canonical_name: "arena-smart", enum_value: 0, model_uid: Some("arena-smart"), credit_multiplier: 1.0, provider: "windsurf", deprecated: true },
];
#[rustfmt::skip]
const ALIASES: &[(&str, &str)] = &[
("claude-3-5-haiku-20241022", "claude-4.5-haiku"),
("claude-3-5-haiku-latest", "claude-4.5-haiku"),
("claude-3-5-sonnet-20240620", "claude-3.5-sonnet"),
("claude-3-5-sonnet-20241022", "claude-3.5-sonnet"),
("claude-3-5-sonnet-latest", "claude-3.5-sonnet"),
("claude-3-7-sonnet-20250219", "claude-3.7-sonnet"),
("claude-3-7-sonnet-latest", "claude-3.7-sonnet"),
("claude-4.6", "claude-sonnet-4.6"),
("claude-4.6-1m", "claude-sonnet-4.6-1m"),
("claude-4.6-thinking", "claude-sonnet-4.6-thinking"),
("claude-4.6-thinking-1m", "claude-sonnet-4.6-thinking-1m"),
("claude-haiku-3-5", "claude-4.5-haiku"),
("claude-haiku-3-5-latest", "claude-4.5-haiku"),
("claude-haiku-4-5", "claude-4.5-haiku"),
("claude-haiku-4-5-20251001", "claude-4.5-haiku"),
("claude-haiku-4-5-latest", "claude-4.5-haiku"),
("claude-haiku-4.5", "claude-4.5-haiku"),
("claude-haiku-4.5-latest", "claude-4.5-haiku"),
("claude-opus-4-0", "claude-4-opus"),
("claude-opus-4-1", "claude-4.1-opus"),
("claude-opus-4-1-20250805", "claude-4.1-opus"),
("claude-opus-4-20250514", "claude-4-opus"),
("claude-opus-4-5", "claude-4.5-opus"),
("claude-opus-4-5-20251101", "claude-4.5-opus"),
("claude-opus-4-5-latest", "claude-4.5-opus"),
("claude-opus-4-6", "claude-opus-4.6"),
("claude-opus-4-6-thinking", "claude-opus-4.6-thinking"),
("claude-opus-4-7", "claude-opus-4-7-medium"),
("claude-opus-4-7-latest", "claude-opus-4-7-medium"),
("claude-opus-4-7-thinking", "claude-opus-4-7-medium-thinking"),
("claude-opus-4.5", "claude-4.5-opus"),
("claude-opus-4.5-thinking", "claude-4.5-opus-thinking"),
("claude-opus-4.7", "claude-opus-4-7-medium"),
("claude-opus-4.7-high", "claude-opus-4-7-high"),
("claude-opus-4.7-high-thinking", "claude-opus-4-7-high-thinking"),
("claude-opus-4.7-low", "claude-opus-4-7-low"),
("claude-opus-4.7-max", "claude-opus-4-7-max"),
("claude-opus-4.7-medium", "claude-opus-4-7-medium"),
("claude-opus-4.7-medium-thinking", "claude-opus-4-7-medium-thinking"),
("claude-opus-4.7-thinking", "claude-opus-4-7-medium-thinking"),
("claude-opus-4.7-xhigh", "claude-opus-4-7-xhigh"),
("claude-opus-4.7-xhigh-thinking", "claude-opus-4-7-xhigh-thinking"),
("claude-sonnet-4-0", "claude-4-sonnet"),
("claude-sonnet-4-20250514", "claude-4-sonnet"),
("claude-sonnet-4-5", "claude-4.5-sonnet"),
("claude-sonnet-4-5-20250929", "claude-4.5-sonnet"),
("claude-sonnet-4-5-latest", "claude-4.5-sonnet"),
("claude-sonnet-4-6", "claude-sonnet-4.6"),
("claude-sonnet-4-6-1m", "claude-sonnet-4.6-1m"),
("claude-sonnet-4-6-thinking", "claude-sonnet-4.6-thinking"),
("claude-sonnet-4-6-thinking-1m", "claude-sonnet-4.6-thinking-1m"),
("claude-sonnet-4.5", "claude-4.5-sonnet"),
("claude-sonnet-4.5-thinking", "claude-4.5-sonnet-thinking"),
("gpt-4.1-2025-04-14", "gpt-4.1"),
("gpt-4.1-mini-2025-04-14", "gpt-4.1-mini"),
("gpt-4.1-nano-2025-04-14", "gpt-4.1-nano"),
("gpt-4o-2024-05-13", "gpt-4o"),
("gpt-4o-2024-08-06", "gpt-4o"),
("gpt-4o-2024-11-20", "gpt-4o"),
("gpt-4o-mini-2024-07-18", "gpt-4o-mini"),
("gpt-5-2-codex-medium", "gpt-5.2-codex-medium"),
("gpt-5-2-medium", "gpt-5.2"),
("gpt-5-2025-08-07", "gpt-5"),
("gpt-5-3-codex-high", "gpt-5.3-codex-high"),
("gpt-5-3-codex-high-priority", "gpt-5.3-codex-high-fast"),
("gpt-5-3-codex-low", "gpt-5.3-codex-low"),
("gpt-5-3-codex-low-priority", "gpt-5.3-codex-low-fast"),
("gpt-5-3-codex-medium", "gpt-5.3-codex"),
("gpt-5-3-codex-medium-priority", "gpt-5.3-codex-medium-fast"),
("gpt-5-3-codex-xhigh", "gpt-5.3-codex-xhigh"),
("gpt-5-3-codex-xhigh-priority", "gpt-5.3-codex-xhigh-fast"),
("gpt-5-4-high", "gpt-5.4-high"),
("gpt-5-4-low", "gpt-5.4-low"),
("gpt-5-4-medium", "gpt-5.4-medium"),
("gpt-5-4-mini-high", "gpt-5.4-mini-high"),
("gpt-5-4-mini-low", "gpt-5.4-mini-low"),
("gpt-5-4-mini-medium", "gpt-5.4-mini-medium"),
("gpt-5-4-mini-xhigh", "gpt-5.4-mini-xhigh"),
("gpt-5-4-none", "gpt-5.4-none"),
("gpt-5-4-xhigh", "gpt-5.4-xhigh"),
("gpt-5-5", "gpt-5.5-medium"),
("gpt-5-5-high", "gpt-5.5-high"),
("gpt-5-5-high-priority", "gpt-5.5-high-fast"),
("gpt-5-5-low", "gpt-5.5-low"),
("gpt-5-5-low-priority", "gpt-5.5-low-fast"),
("gpt-5-5-medium", "gpt-5.5-medium"),
("gpt-5-5-medium-priority", "gpt-5.5-medium-fast"),
("gpt-5-5-none", "gpt-5.5-none"),
("gpt-5-5-none-priority", "gpt-5.5-none-fast"),
("gpt-5-5-xhigh", "gpt-5.5-xhigh"),
("gpt-5-5-xhigh-priority", "gpt-5.5-xhigh-fast"),
("gpt-5-pro-2025-10-06", "gpt-5-high"),
("gpt-5.2-codex", "gpt-5.2-codex-medium"),
("gpt-5.2-medium", "gpt-5.2"),
("gpt-5.3-codex-medium", "gpt-5.3-codex"),
("gpt-5.4", "gpt-5.4-medium"),
("gpt-5.5", "gpt-5.5-medium"),
("haiku-4.5", "claude-4.5-haiku"),
("kimi-k2-5", "kimi-k2.5"),
("minimax-m2-5", "minimax-m2.5"),
("model_claude_4_5_sonnet", "claude-4.5-sonnet"),
("model_claude_4_5_sonnet_thinking", "claude-4.5-sonnet-thinking"),
("o4.7", "claude-opus-4-7-medium"),
("opus-4", "claude-4-opus"),
("opus-4-7", "claude-opus-4-7-medium"),
("opus-4.1", "claude-4.1-opus"),
("opus-4.6", "claude-opus-4.6"),
("opus-4.6-thinking", "claude-opus-4.6-thinking"),
("opus-4.7", "claude-opus-4-7-medium"),
("opus-4.7-thinking", "claude-opus-4-7-medium-thinking"),
("sonnet-3.5", "claude-3.5-sonnet"),
("sonnet-3.7", "claude-3.7-sonnet"),
("sonnet-4", "claude-4-sonnet"),
("sonnet-4.5", "claude-4.5-sonnet"),
("sonnet-4.5-thinking", "claude-4.5-sonnet-thinking"),
("sonnet-4.6", "claude-sonnet-4.6"),
("sonnet-4.6-1m", "claude-sonnet-4.6-1m"),
("sonnet-4.6-thinking", "claude-sonnet-4.6-thinking"),
("swe-1-6", "swe-1.6"),
("swe-1-6-fast", "swe-1.6-fast"),
("ws-haiku", "claude-4.5-haiku"),
("ws-opus", "claude-opus-4.6"),
("ws-opus-thinking", "claude-opus-4.6-thinking"),
("ws-sonnet", "claude-sonnet-4.6"),
("ws-sonnet-thinking", "claude-sonnet-4.6-thinking"),
];
pub fn windsurf_models() -> &'static [WindsurfModel] {
MODELS
}
pub fn resolve_windsurf_model(name: &str) -> Option<WindsurfModel> {
let normalized = name.trim().to_ascii_lowercase();
if normalized.is_empty() {
return None;
}
let canonical = ALIASES
.iter()
.find_map(|(alias, canonical)| (*alias == normalized).then_some(*canonical))
.unwrap_or(normalized.as_str());
MODELS
.iter()
.find(|model| model_matches(model, canonical))
.cloned()
}
fn model_matches(model: &WindsurfModel, value: &str) -> bool {
model.canonical_name.eq_ignore_ascii_case(value)
|| model
.model_uid
.is_some_and(|uid| uid.eq_ignore_ascii_case(value))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn resolves_gpt55_cloud_alias_to_windsurf_model_uid() {
let model = resolve_windsurf_model("gpt-5-5-low").expect("model should resolve");
assert_eq!(model.canonical_name, "gpt-5.5-low");
assert_eq!(model.model_uid, Some("gpt-5-5-low"));
assert_eq!(model.enum_value, 0);
assert_eq!(model.credit_multiplier, 1.0);
}
#[test]
fn resolves_claude_opus_47_bare_alias_to_medium() {
let model = resolve_windsurf_model("claude-opus-4.7").expect("model should resolve");
assert_eq!(model.canonical_name, "claude-opus-4-7-medium");
assert_eq!(model.model_uid, Some("claude-opus-4-7-medium"));
assert_eq!(model.enum_value, 0);
assert_eq!(model.credit_multiplier, 8.0);
}
#[test]
fn resolves_priority_alias_to_fast_variant() {
let model = resolve_windsurf_model("gpt-5-5-low-priority").expect("model should resolve");
assert_eq!(model.canonical_name, "gpt-5.5-low-fast");
assert_eq!(model.model_uid, Some("gpt-5-5-low-priority"));
assert_eq!(model.credit_multiplier, 2.0);
}
#[test]
fn resolves_full_gpt55_effort_ladder_and_priority_aliases() {
let none = resolve_windsurf_model("gpt-5-5-none").expect("none should resolve");
assert_eq!(none.canonical_name, "gpt-5.5-none");
assert_eq!(none.model_uid, Some("gpt-5-5-none"));
assert_eq!(none.credit_multiplier, 1.0);
let high = resolve_windsurf_model("gpt-5.5-high").expect("high should resolve");
assert_eq!(high.canonical_name, "gpt-5.5-high");
assert_eq!(high.model_uid, Some("gpt-5-5-high"));
assert_eq!(high.credit_multiplier, 4.0);
let xhigh_fast = resolve_windsurf_model("gpt-5-5-xhigh-priority")
.expect("xhigh priority should resolve");
assert_eq!(xhigh_fast.canonical_name, "gpt-5.5-xhigh-fast");
assert_eq!(xhigh_fast.model_uid, Some("gpt-5-5-xhigh-priority"));
assert_eq!(xhigh_fast.credit_multiplier, 16.0);
}
#[test]
fn resolves_windsurfapi_catalog_aliases_beyond_gpt55() {
let gpt52_medium = resolve_windsurf_model("gpt-5.2-medium").expect("gpt-5.2 medium alias");
assert_eq!(gpt52_medium.canonical_name, "gpt-5.2");
assert_eq!(gpt52_medium.model_uid, Some("MODEL_GPT_5_2_MEDIUM"));
let haiku = resolve_windsurf_model("claude-haiku-4-5-20251001").expect("dated haiku alias");
assert_eq!(haiku.canonical_name, "claude-4.5-haiku");
assert_eq!(haiku.model_uid, Some("MODEL_PRIVATE_11"));
let uid = resolve_windsurf_model("MODEL_GPT_5_2_LOW").expect("model uid alias");
assert_eq!(uid.canonical_name, "gpt-5.2-low");
assert_eq!(uid.enum_value, 400);
let cursor = resolve_windsurf_model("ws-opus").expect("cursor alias");
assert_eq!(cursor.canonical_name, "claude-opus-4.6");
}
#[test]
fn windsurf_catalog_covers_current_windsurfapi_model_set() {
assert_eq!(MODELS.len(), 139);
assert!(ALIASES.len() >= 100);
}
}
@@ -0,0 +1,248 @@
use std::fmt;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WireType {
Varint = 0,
Fixed64 = 1,
Len = 2,
Fixed32 = 5,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum FieldValue {
Varint(u64),
Bytes(Vec<u8>),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Field {
pub number: u32,
pub wire_type: WireType,
pub value: FieldValue,
}
impl Field {
pub fn bytes(&self) -> &[u8] {
match &self.value {
FieldValue::Bytes(bytes) => bytes.as_slice(),
FieldValue::Varint(_) => &[],
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ProtoError {
message: String,
}
impl ProtoError {
fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
}
impl fmt::Display for ProtoError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.message)
}
}
impl std::error::Error for ProtoError {}
pub fn encode_varint(value: u64) -> Vec<u8> {
let mut value = value;
let mut out = Vec::new();
loop {
let mut byte = (value & 0x7f) as u8;
value >>= 7;
if value != 0 {
byte |= 0x80;
}
out.push(byte);
if value == 0 {
break;
}
}
out
}
fn decode_varint(buf: &[u8], offset: usize) -> Result<(u64, usize), ProtoError> {
let mut value = 0u64;
let mut shift = 0u32;
let mut pos = offset;
while pos < buf.len() {
let byte = buf[pos];
pos += 1;
value |= u64::from(byte & 0x7f) << shift;
if byte & 0x80 == 0 {
return Ok((value, pos - offset));
}
shift += 7;
if shift >= 64 {
return Err(ProtoError::new("varint overflow"));
}
}
Err(ProtoError::new("truncated varint"))
}
fn tag(field: u32, wire_type: WireType) -> Vec<u8> {
encode_varint((u64::from(field) << 3) | wire_type as u64)
}
pub fn write_varint_field(field: u32, value: u64) -> Vec<u8> {
let mut out = tag(field, WireType::Varint);
out.extend(encode_varint(value));
out
}
pub fn write_string_field(field: u32, value: &str) -> Vec<u8> {
let bytes = value.as_bytes();
let mut out = tag(field, WireType::Len);
out.extend(encode_varint(bytes.len() as u64));
out.extend(bytes);
out
}
pub fn write_message_field(field: u32, value: &[u8]) -> Vec<u8> {
if value.is_empty() {
return Vec::new();
}
let mut out = tag(field, WireType::Len);
out.extend(encode_varint(value.len() as u64));
out.extend(value);
out
}
pub fn write_bool_field(field: u32, value: bool) -> Vec<u8> {
if value {
write_varint_field(field, 1)
} else {
Vec::new()
}
}
pub fn parse_fields(buf: &[u8]) -> Result<Vec<Field>, ProtoError> {
let mut fields = Vec::new();
let mut pos = 0usize;
while pos < buf.len() {
let (tag, tag_len) = decode_varint(buf, pos)?;
pos += tag_len;
let number = (tag >> 3) as u32;
let wire_type = match tag & 0x07 {
0 => WireType::Varint,
1 => WireType::Fixed64,
2 => WireType::Len,
5 => WireType::Fixed32,
other => {
return Err(ProtoError::new(format!(
"unknown wire type {other} at offset {pos}"
)))
}
};
let value = match wire_type {
WireType::Varint => {
let (value, value_len) = decode_varint(buf, pos)?;
pos += value_len;
FieldValue::Varint(value)
}
WireType::Len => {
let (len, len_len) = decode_varint(buf, pos)?;
pos += len_len;
let len = usize::try_from(len)
.map_err(|_| ProtoError::new("length-delimited field too large"))?;
if pos + len > buf.len() {
return Err(ProtoError::new(format!(
"truncated len-delimited field {number} at offset {pos}"
)));
}
let bytes = buf[pos..pos + len].to_vec();
pos += len;
FieldValue::Bytes(bytes)
}
WireType::Fixed64 => {
if pos + 8 > buf.len() {
return Err(ProtoError::new(format!("truncated fixed64 field {number}")));
}
let bytes = buf[pos..pos + 8].to_vec();
pos += 8;
FieldValue::Bytes(bytes)
}
WireType::Fixed32 => {
if pos + 4 > buf.len() {
return Err(ProtoError::new(format!("truncated fixed32 field {number}")));
}
let bytes = buf[pos..pos + 4].to_vec();
pos += 4;
FieldValue::Bytes(bytes)
}
};
fields.push(Field {
number,
wire_type,
value,
});
}
Ok(fields)
}
pub fn get_field(fields: &[Field], number: u32, wire_type: Option<WireType>) -> Option<&Field> {
fields
.iter()
.find(|field| field.number == number && wire_type.is_none_or(|ty| field.wire_type == ty))
}
pub fn get_all_fields(fields: &[Field], number: u32) -> Vec<&Field> {
fields
.iter()
.filter(|field| field.number == number)
.collect()
}
pub fn get_varint(fields: &[Field], number: u32) -> Option<u64> {
match get_field(fields, number, Some(WireType::Varint))?.value {
FieldValue::Varint(value) => Some(value),
FieldValue::Bytes(_) => None,
}
}
pub fn get_string(fields: &[Field], number: u32) -> Option<String> {
let field = get_field(fields, number, Some(WireType::Len))?;
String::from_utf8(field.bytes().to_vec()).ok()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn encodes_varint_and_string_fields_like_windsurfapi() {
assert_eq!(encode_varint(300), vec![0xac, 0x02]);
assert_eq!(
write_string_field(3, "abc"),
vec![0x1a, 0x03, b'a', b'b', b'c']
);
assert_eq!(write_bool_field(2, false), Vec::<u8>::new());
assert_eq!(write_bool_field(2, true), vec![0x10, 0x01]);
}
#[test]
fn parses_repeated_len_delimited_fields() {
let mut bytes = Vec::new();
bytes.extend(write_string_field(1, "alpha"));
bytes.extend(write_string_field(1, "beta"));
bytes.extend(write_varint_field(2, 42));
let fields = parse_fields(&bytes).expect("fields should parse");
assert_eq!(get_all_fields(&fields, 1).len(), 2);
assert_eq!(get_string(&fields, 1).as_deref(), Some("alpha"));
assert_eq!(get_varint(&fields, 2), Some(42));
}
#[test]
fn rejects_truncated_len_delimited_field() {
let err = parse_fields(&[0x0a, 0x05, b'a']).expect_err("must reject truncated field");
assert!(err.to_string().contains("truncated"));
}
}