Merge remote-tracking branch 'origin/main' into codex/provider-policy-hardening

This commit is contained in:
elky
2026-09-05 16:21:10 +08:00
1228 changed files with 241230 additions and 31814 deletions
+2 -1
View File
@@ -11,6 +11,7 @@ aether-ai-formats.workspace = true
aether-contracts.workspace = true
aether-crypto.workspace = true
aether-data-contracts.workspace = true
aether-http.workspace = true
aether-oauth.workspace = true
aether-runtime-state.workspace = true
aether-video-tasks-core.workspace = true
@@ -22,7 +23,6 @@ ed25519-dalek.workspace = true
http.workspace = true
regex.workspace = true
reqwest.workspace = true
rsa = "0.9.10"
serde.workspace = true
serde_json.workspace = true
sha2 = { workspace = true, features = ["oid"] }
@@ -33,4 +33,5 @@ url.workspace = true
uuid.workspace = true
[dev-dependencies]
aws-lc-rs.workspace = true
axum = { version = "0.8", features = ["ws"] }
@@ -91,6 +91,35 @@ struct AgentIdentityAssertionEnvelope {
signature: String,
}
const MAX_AGENT_IDENTITY_PRIVATE_KEY_DER_BYTES: usize = 4 * 1024;
const MAX_AGENT_IDENTITY_ASSERTION_ENVELOPE_BYTES: usize = 64 * 1024;
const MAX_AGENT_IDENTITY_SIGNATURE_BYTES: usize = 128;
const MAX_AGENT_IDENTITY_ENCRYPTED_TASK_ID_BYTES: usize = 64 * 1024;
fn maximum_base64_len_for_decoded_limit(limit: usize) -> usize {
limit
.saturating_add(2)
.checked_div(3)
.unwrap_or(usize::MAX)
.saturating_mul(4)
}
fn decode_standard_base64_with_limit(value: &str, limit: usize) -> Option<Vec<u8>> {
if value.len() > maximum_base64_len_for_decoded_limit(limit) {
return None;
}
let decoded = STANDARD.decode(value).ok()?;
(decoded.len() <= limit).then_some(decoded)
}
fn decode_urlsafe_base64_with_limit(value: &str, limit: usize) -> Option<Vec<u8>> {
if value.len() > maximum_base64_len_for_decoded_limit(limit) {
return None;
}
let decoded = URL_SAFE_NO_PAD.decode(value).ok()?;
(decoded.len() <= limit).then_some(decoded)
}
#[derive(Debug, Deserialize)]
struct AgentTaskRegistrationResponse {
#[serde(default)]
@@ -288,7 +317,9 @@ pub fn codex_agent_identity_authorization_matches_transport(
let Some(encoded) = encoded_agent_identity_assertion(authorization) else {
return false;
};
let Ok(envelope_bytes) = URL_SAFE_NO_PAD.decode(encoded) else {
let Some(envelope_bytes) =
decode_urlsafe_base64_with_limit(encoded, MAX_AGENT_IDENTITY_ASSERTION_ENVELOPE_BYTES)
else {
return false;
};
let Ok(envelope) = serde_json::from_slice::<AgentIdentityAssertionEnvelope>(&envelope_bytes)
@@ -310,7 +341,10 @@ pub fn codex_agent_identity_authorization_matches_transport(
if runtime_id != credentials.runtime_id || task_id != current_task_id || timestamp.is_empty() {
return false;
}
let Ok(signature_bytes) = STANDARD.decode(envelope.signature.trim()) else {
let Some(signature_bytes) = decode_standard_base64_with_limit(
envelope.signature.trim(),
MAX_AGENT_IDENTITY_SIGNATURE_BYTES,
) else {
return false;
};
let Ok(signature) = Signature::from_slice(&signature_bytes) else {
@@ -481,9 +515,11 @@ fn agent_identity_credentials(config: &Value) -> Result<AgentIdentityCredentials
let encoded_private_key =
string_from_maps(root, nested, &["agent_private_key", "agentPrivateKey"])
.ok_or_else(|| "Agent Identity agent_private_key is required".to_string())?;
let private_key_der = STANDARD
.decode(encoded_private_key)
.map_err(|_| "Agent Identity agent_private_key must be base64 PKCS#8".to_string())?;
let private_key_der = decode_standard_base64_with_limit(
&encoded_private_key,
MAX_AGENT_IDENTITY_PRIVATE_KEY_DER_BYTES,
)
.ok_or_else(|| "Agent Identity agent_private_key must be base64 PKCS#8".to_string())?;
let signing_key = SigningKey::from_pkcs8_der(&private_key_der).map_err(|_| {
"Agent Identity agent_private_key must be an Ed25519 PKCS#8 key".to_string()
})?;
@@ -829,9 +865,11 @@ fn decrypt_agent_task_id(
credentials: &AgentIdentityCredentials,
encrypted_task_id: &str,
) -> Result<String, String> {
let ciphertext = STANDARD
.decode(encrypted_task_id.trim())
.map_err(|_| "Agent Identity encrypted_task_id must be base64".to_string())?;
let ciphertext = decode_standard_base64_with_limit(
encrypted_task_id.trim(),
MAX_AGENT_IDENTITY_ENCRYPTED_TASK_ID_BYTES,
)
.ok_or_else(|| "Agent Identity encrypted_task_id must be valid bounded base64".to_string())?;
let seed = credentials.signing_key.to_bytes();
let digest = Sha512::digest(seed);
let mut curve_private_key = [0u8; 32];
@@ -1240,6 +1278,34 @@ mod tests {
));
}
#[test]
fn agent_identity_rejects_oversized_base64_before_decode() {
let encoded_limit = super::maximum_base64_len_for_decoded_limit(
super::MAX_AGENT_IDENTITY_ASSERTION_ENVELOPE_BYTES,
);
let assertion = format!("AgentAssertion {}", "A".repeat(encoded_limit + 1));
assert!(!codex_agent_identity_authorization_matches_transport(
&sample_transport(test_auth_config(Some("task-test"))),
&assertion,
));
let mut config = test_auth_config(Some("task-test"));
let private_key_limit = super::maximum_base64_len_for_decoded_limit(
super::MAX_AGENT_IDENTITY_PRIVATE_KEY_DER_BYTES,
);
config["agent_private_key"] = json!("A".repeat(private_key_limit + 1));
assert!(agent_identity_credentials(&config).is_err());
let credentials =
agent_identity_credentials(&test_auth_config(None)).expect("credentials should parse");
let encrypted_task_id_limit = super::maximum_base64_len_for_decoded_limit(
super::MAX_AGENT_IDENTITY_ENCRYPTED_TASK_ID_BYTES,
);
assert!(
decrypt_agent_task_id(&credentials, &"A".repeat(encrypted_task_id_limit + 1),).is_err()
);
}
#[test]
fn task_rotation_context_rejects_metadata_and_credential_replacement() {
let initial = sample_transport(test_auth_config(Some("task-old")));
@@ -5,8 +5,8 @@ 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)";
pub const ANTIGRAVITY_CLIENT_VERSION: &str = "4.3.0";
pub const ANTIGRAVITY_REQUEST_USER_AGENT: &str = "vscode/1.X.X (Antigravity/4.3.0)";
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";
@@ -109,7 +109,7 @@ pub fn build_antigravity_static_identity_headers(
}
pub fn build_antigravity_static_client_headers(
client_version: Option<&str>,
_client_version: Option<&str>,
session_id: Option<&str>,
) -> BTreeMap<String, String> {
let mut headers = BTreeMap::from([
@@ -125,14 +125,12 @@ pub fn build_antigravity_static_client_headers(
String::from("user-agent"),
String::from(ANTIGRAVITY_REQUEST_USER_AGENT),
),
(
String::from("x-client-version"),
String::from(ANTIGRAVITY_CLIENT_VERSION),
),
]);
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());
}
@@ -273,7 +271,8 @@ mod tests {
use super::{
build_antigravity_static_client_headers, resolve_local_antigravity_request_auth,
AntigravityRequestAuth, AntigravityRequestAuthSupport, ANTIGRAVITY_REQUEST_USER_AGENT,
AntigravityRequestAuth, AntigravityRequestAuthSupport, ANTIGRAVITY_CLIENT_VERSION,
ANTIGRAVITY_REQUEST_USER_AGENT,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
@@ -383,7 +382,7 @@ mod tests {
}
#[test]
fn static_client_headers_use_native_antigravity_cli_user_agent() {
fn static_client_headers_pin_the_known_good_antigravity_identity() {
let headers = build_antigravity_static_client_headers(Some("1.0.16"), Some("session-abc"));
assert_eq!(
@@ -396,7 +395,7 @@ mod tests {
);
assert_eq!(
headers.get("x-client-version").map(String::as_str),
Some("1.0.16")
Some(ANTIGRAVITY_CLIENT_VERSION)
);
assert_eq!(
headers.get("x-vscode-sessionid").map(String::as_str),
@@ -1,6 +1,7 @@
use serde_json::Value;
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::super::transport_proxy_is_locally_supported;
use super::auth::{
resolve_local_antigravity_request_auth, AntigravityRequestAuth, AntigravityRequestAuthSupport,
AntigravityRequestAuthUnsupportedReason, ANTIGRAVITY_PROVIDER_TYPE,
@@ -87,11 +88,11 @@ pub fn classify_local_antigravity_request_support(
AntigravityRequestSideUnsupportedReason::UnsupportedBodyRules,
);
}
if transport.provider.proxy.is_some()
|| transport.endpoint.proxy.is_some()
|| transport.key.proxy.is_some()
|| transport.key.fingerprint.is_some()
{
// A configured proxy is carried by the execution plan itself, so it only
// disqualifies the local request when it cannot be resolved into a usable
// snapshot. Transport profiles stay unsupported because the v1internal
// payload never carries one.
if !transport_proxy_is_locally_supported(transport) || transport.key.fingerprint.is_some() {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig,
);
@@ -114,3 +115,145 @@ pub fn classify_local_antigravity_request_support(
AntigravityRequestSideSupport::Supported(AntigravityRequestSideSpec { auth, request_type })
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::super::request::AntigravityEnvelopeRequestType;
use super::{
classify_local_antigravity_request_support, AntigravityRequestSideSupport,
AntigravityRequestSideUnsupportedReason,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport() -> 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(
r#"{"provider_type":"antigravity","refresh_token":"rt","cloudaicompanionProject":"project-1"}"#
.to_string(),
),
},
}
}
fn classify(transport: &GatewayProviderTransportSnapshot) -> AntigravityRequestSideSupport {
classify_local_antigravity_request_support(
transport,
&json!({"contents": []}),
AntigravityEnvelopeRequestType::Agent,
)
}
fn assert_unsupported_network_config(support: AntigravityRequestSideSupport) {
assert_eq!(
support,
AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig,
)
);
}
#[test]
fn a_resolvable_tunnel_node_proxy_keeps_the_envelope_supported() {
let mut transport = sample_transport();
transport.provider.proxy = Some(json!({
"enabled": true,
"node_id": "702d158b-a432-4694-94cc-3bec13dbbc20",
}));
assert!(matches!(
classify(&transport),
AntigravityRequestSideSupport::Supported(_)
));
}
#[test]
fn a_resolvable_url_proxy_keeps_the_envelope_supported() {
for proxy_owner in ["provider", "endpoint", "key"] {
let mut transport = sample_transport();
let proxy = Some(json!({"enabled": true, "url": "http://127.0.0.1:17000"}));
match proxy_owner {
"provider" => transport.provider.proxy = proxy,
"endpoint" => transport.endpoint.proxy = proxy,
_ => transport.key.proxy = proxy,
}
assert!(
matches!(
classify(&transport),
AntigravityRequestSideSupport::Supported(_)
),
"a {proxy_owner} proxy should not disqualify the antigravity envelope"
);
}
}
#[test]
fn a_proxy_without_a_route_still_disqualifies_the_envelope() {
let mut transport = sample_transport();
transport.provider.proxy = Some(json!({"enabled": true}));
assert_unsupported_network_config(classify(&transport));
}
#[test]
fn a_key_fingerprint_still_disqualifies_the_envelope() {
let mut transport = sample_transport();
transport.key.fingerprint = Some(json!({"transport_profile": "chrome"}));
assert_unsupported_network_config(classify(&transport));
}
}
+107 -11
View File
@@ -1,8 +1,10 @@
use std::collections::BTreeMap;
use super::headers::{
is_aether_internal_header, is_upstream_credential_header, normalize_upstream_accept_encoding,
should_skip_upstream_complete_passthrough_header, should_skip_upstream_passthrough_header,
declared_connection_header_names, is_aether_internal_header, is_upstream_credential_header,
normalize_upstream_accept_encoding, remove_declared_connection_headers,
should_skip_upstream_complete_passthrough_header_with_connection,
should_skip_upstream_passthrough_header_with_connection,
};
use super::snapshot::GatewayProviderTransportSnapshot;
@@ -13,13 +15,17 @@ fn collect_passthrough_headers(
headers: &http::HeaderMap,
extra_headers: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
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) {
if should_skip_upstream_passthrough_header_with_connection(
&key,
&declared_connection_headers,
) {
continue;
}
let Some(value) = normalize_passthrough_header_value(&key, value) else {
@@ -30,7 +36,10 @@ fn collect_passthrough_headers(
for (key, value) in extra_headers {
let normalized_key = key.to_ascii_lowercase();
if should_skip_upstream_passthrough_header(&normalized_key) {
if should_skip_upstream_passthrough_header_with_connection(
&normalized_key,
&declared_connection_headers,
) {
continue;
}
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
@@ -46,13 +55,17 @@ fn collect_complete_passthrough_headers(
headers: &http::HeaderMap,
extra_headers: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
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) {
if should_skip_upstream_complete_passthrough_header_with_connection(
&key,
&declared_connection_headers,
) {
continue;
}
let Some(value) = normalize_passthrough_header_value(&key, value) else {
@@ -63,7 +76,10 @@ fn collect_complete_passthrough_headers(
for (key, value) in extra_headers {
let normalized_key = key.to_ascii_lowercase();
if should_skip_upstream_complete_passthrough_header(&normalized_key) {
if should_skip_upstream_complete_passthrough_header_with_connection(
&normalized_key,
&declared_connection_headers,
) {
continue;
}
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
@@ -101,6 +117,8 @@ pub fn build_passthrough_headers(
.trim()
.to_string()
});
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out.remove("content-length");
out
}
@@ -112,8 +130,10 @@ pub fn build_openai_passthrough_headers(
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
let mut out = build_passthrough_headers(headers, extra_headers, content_type);
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out
}
@@ -122,7 +142,9 @@ pub fn build_complete_passthrough_headers(
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
let mut out = collect_complete_passthrough_headers(headers, extra_headers);
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out.entry("content-type".to_string()).or_insert_with(|| {
content_type
.filter(|value| !value.trim().is_empty())
@@ -130,6 +152,7 @@ pub fn build_complete_passthrough_headers(
.trim()
.to_string()
});
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out.remove("content-length");
out
}
@@ -141,8 +164,10 @@ pub fn build_complete_passthrough_headers_with_auth(
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
let mut out = build_complete_passthrough_headers(headers, extra_headers, content_type);
replace_upstream_auth_headers(&mut out, auth_header, auth_value);
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out
}
@@ -153,6 +178,7 @@ pub fn build_claude_passthrough_headers(
extra_headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> BTreeMap<String, String> {
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
let mut out = build_openai_passthrough_headers(
headers,
auth_header,
@@ -164,7 +190,10 @@ pub fn build_claude_passthrough_headers(
for (name, value) in extra_headers {
let key = name.to_ascii_lowercase();
let value = value.trim();
if value.is_empty() || !should_restore_claude_passthrough_header(&key) {
if value.is_empty()
|| !should_restore_claude_passthrough_header(&key)
|| declared_connection_headers.contains(&key)
{
continue;
}
@@ -185,7 +214,10 @@ pub fn build_claude_passthrough_headers(
};
let key = name.as_str().to_ascii_lowercase();
let value = value.trim();
if value.is_empty() || !should_restore_claude_passthrough_header(&key) {
if value.is_empty()
|| !should_restore_claude_passthrough_header(&key)
|| declared_connection_headers.contains(&key)
{
continue;
}
@@ -202,6 +234,7 @@ pub fn build_claude_passthrough_headers(
out.entry("anthropic-version".to_string())
.or_insert_with(|| DEFAULT_ANTHROPIC_VERSION.to_string());
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out
}
@@ -211,8 +244,10 @@ pub fn build_passthrough_headers_with_auth(
auth_value: &str,
extra_headers: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
let mut out = collect_passthrough_headers(headers, extra_headers);
replace_upstream_auth_headers(&mut out, auth_header, auth_value);
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out.remove("content-length");
out
}
@@ -357,9 +392,9 @@ fn bearer_auth_value(secret: &str) -> String {
#[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,
build_claude_passthrough_headers, build_complete_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,
@@ -495,6 +530,67 @@ mod tests {
);
}
#[test]
fn passthrough_headers_strip_connection_declared_fields() {
let mut headers = http::HeaderMap::new();
headers.append(
http::header::CONNECTION,
http::HeaderValue::from_static("X-Internal-Hop, keep-alive"),
);
headers.append(
http::header::CONNECTION,
http::HeaderValue::from_static("x-extra-hop"),
);
headers.insert(
"x-internal-hop",
http::HeaderValue::from_static("private-value"),
);
headers.insert(
"x-extra-hop",
http::HeaderValue::from_static("private-value-2"),
);
headers.insert("x-public", http::HeaderValue::from_static("ok"));
let extra = BTreeMap::from([
(
"Connection".to_string(),
"X-Extra-From-Connection".to_string(),
),
("X-Extra-From-Connection".to_string(), "secret".to_string()),
]);
let built = build_openai_passthrough_headers(
&headers,
"authorization",
"Bearer upstream",
&extra,
Some("application/json"),
);
assert_eq!(built.get("x-public").map(String::as_str), Some("ok"));
assert!(!built.contains_key("connection"));
assert!(!built.contains_key("x-internal-hop"));
assert!(!built.contains_key("x-extra-hop"));
assert!(built
.keys()
.all(|name| { !name.eq_ignore_ascii_case("x-extra-from-connection") }));
}
#[test]
fn complete_passthrough_headers_strip_connection_declared_fields_from_extra_headers() {
let headers = http::HeaderMap::new();
let extra = BTreeMap::from([
("Connection".to_string(), "x-private-hop".to_string()),
("X-Private-Hop".to_string(), "secret".to_string()),
("x-public".to_string(), "ok".to_string()),
]);
let built = build_complete_passthrough_headers(&headers, &extra, None);
assert_eq!(built.get("x-public").map(String::as_str), Some("ok"));
assert!(!built.contains_key("connection"));
assert!(!built.contains_key("x-private-hop"));
}
#[test]
fn claude_passthrough_headers_preserve_explicit_anthropic_version_override() {
let mut headers = http::HeaderMap::new();
@@ -1,7 +1,7 @@
use aether_ai_formats::formats::matrix::{
request_conversion_kind, request_conversion_requires_enable_flag,
};
use aether_contracts::ProxySnapshot;
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
use serde_json::{json, Map, Value};
use crate::conversion::{
@@ -71,25 +71,24 @@ pub fn build_transport_diagnostics(
) -> 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())
.as_ref()
.map(summarize_transport_profile)
.unwrap_or(Value::Null);
let configured_key_transport_profile = transport
let key_transport_profile_configured = 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
.is_some_and(|value| !value.is_null());
let provider_transport_profile_configured = 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);
.is_some_and(|value| !value.is_null());
let configured_legacy_grok_transport_profile = if transport
.provider
.provider_type
@@ -109,7 +108,8 @@ pub fn build_transport_diagnostics(
&auth_config,
"grok_auth_config",
)
.and_then(|profile| serde_json::to_value(profile).ok())
.as_ref()
.map(summarize_transport_profile)
})
.unwrap_or(Value::Null)
} else {
@@ -132,11 +132,14 @@ pub fn build_transport_diagnostics(
"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,
"endpoint_format_acceptance": summarize_format_acceptance_config(
transport.endpoint.format_acceptance_config.as_ref()
),
"endpoint_has_custom_path": transport.endpoint.custom_path.as_deref()
.is_some_and(|value| !value.trim().is_empty()),
"header_rules_count": json_array_len(transport.endpoint.header_rules.as_ref()),
"header_rules_supported": header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()),
"body_rules": transport.endpoint.body_rules,
"body_rules_count": json_array_len(transport.endpoint.body_rules.as_ref()),
"body_rules_supported": body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()),
"proxy": {
"locally_supported": transport_proxy_is_locally_supported(transport),
@@ -149,9 +152,10 @@ pub fn build_transport_diagnostics(
"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,
"key_fingerprint_configured": transport.key.fingerprint.as_ref()
.is_some_and(|value| !value.is_null()),
"key_transport_profile_configured": key_transport_profile_configured,
"provider_transport_profile_configured": provider_transport_profile_configured,
"configured_legacy_grok_transport_profile": configured_legacy_grok_transport_profile,
"resolved_transport_profile_id": resolved_transport_profile_id,
"resolved_transport_profile": resolved_transport_profile,
@@ -177,6 +181,34 @@ pub fn build_transport_diagnostics(
})
}
fn json_array_len(value: Option<&Value>) -> usize {
value.and_then(Value::as_array).map_or(0, Vec::len)
}
fn summarize_format_acceptance_config(value: Option<&Value>) -> Value {
let Some(object) = value.and_then(Value::as_object) else {
return json!({ "configured": false });
};
json!({
"configured": true,
"enabled": object.get("enabled").and_then(Value::as_bool),
"accept_formats_count": json_array_len(object.get("accept_formats")),
"reject_formats_count": json_array_len(object.get("reject_formats")),
})
}
fn summarize_transport_profile(profile: &ResolvedTransportProfile) -> Value {
json!({
"profile_id": profile.profile_id,
"backend": profile.backend,
"http_mode": profile.http_mode,
"pool_scope": profile.pool_scope,
"has_header_fingerprint": profile.header_fingerprint.as_ref()
.is_some_and(|value| !value.is_null()),
"has_extra": profile.extra.as_ref().is_some_and(|value| !value.is_null()),
})
}
fn summarize_proxy_config(proxy: Option<&Value>) -> Value {
let Some(object) = proxy.and_then(Value::as_object) else {
return Value::Null;
@@ -187,11 +219,16 @@ fn summarize_proxy_config(proxy: Option<&Value>) -> Value {
.and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty());
json!({
"configured": true,
"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_mode": object.get("mode").and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty()),
"has_node_id": object.get("node_id").and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty()),
"has_label": object.get("label").and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty()),
"has_url": has_url,
"has_extra": object.get("extra").is_some_and(|value| !value.is_null()),
})
}
@@ -426,11 +463,13 @@ mod tests {
build_transport_diagnostics(&sample_transport(), "claude:messages", "openai:responses");
assert_eq!(diagnostics["provider_type"], "codex");
assert_eq!(diagnostics["key_fingerprint_configured"], true);
assert_eq!(diagnostics["key_transport_profile_configured"], true);
assert_eq!(diagnostics["resolved_transport_profile_id"], "chrome_136");
assert_eq!(
diagnostics["fingerprint"]["transport_profile"]["profile_id"],
diagnostics["resolved_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)
@@ -489,6 +528,86 @@ mod tests {
);
}
#[test]
fn transport_diagnostics_do_not_serialize_configured_secrets() {
let mut transport = sample_transport();
transport.provider.proxy = Some(json!({
"enabled": true,
"mode": "secret-proxy-mode",
"node_id": "secret-proxy-node",
"label": "secret-proxy-label",
"url": "https://secret-user:[email protected]/secret-path",
"extra": {"token": "secret-proxy-extra"}
}));
transport.provider.config = Some(json!({
"secret": "secret-provider-config",
"fingerprint": {
"transport_profile": {
"profile_id": "safe-profile",
"header_fingerprint": {"authorization": "secret-profile-header"},
"extra": {"token": "secret-profile-extra"}
}
}
}));
transport.endpoint.custom_path = Some("/secret-custom-path".to_string());
transport.endpoint.header_rules = Some(json!([
{"op": "set", "key": "authorization", "value": "secret-header-rule"}
]));
transport.endpoint.body_rules = Some(json!([
{"op": "set", "path": "auth.token", "value": "secret-body-rule"}
]));
transport.endpoint.format_acceptance_config = Some(json!({
"enabled": true,
"accept_formats": ["secret-accepted-format"],
"reject_formats": ["secret-rejected-format"],
"token": "secret-format-config"
}));
transport.key.fingerprint = Some(json!({
"secret": "secret-key-fingerprint",
"transport_profile": {
"profile_id": "safe-key-profile",
"header_fingerprint": {"authorization": "secret-key-profile-header"},
"extra": {"token": "secret-key-profile-extra"}
}
}));
let diagnostics =
build_transport_diagnostics(&transport, "claude:messages", "openai:responses");
let serialized = serde_json::to_string(&diagnostics).unwrap();
for secret in [
"secret-proxy-mode",
"secret-proxy-node",
"secret-proxy-label",
"secret-user",
"secret-pass",
"secret-path",
"secret-proxy-extra",
"secret-provider-config",
"secret-profile-header",
"secret-profile-extra",
"secret-custom-path",
"secret-header-rule",
"secret-body-rule",
"secret-accepted-format",
"secret-rejected-format",
"secret-format-config",
"secret-key-fingerprint",
"secret-key-profile-header",
"secret-key-profile-extra",
] {
assert!(!serialized.contains(secret), "leaked {secret}");
}
assert_eq!(diagnostics["header_rules_count"], 1);
assert_eq!(diagnostics["body_rules_count"], 1);
assert_eq!(diagnostics["proxy"]["provider"]["has_node_id"], true);
assert_eq!(
diagnostics["resolved_transport_profile"]["has_header_fingerprint"],
true
);
assert_eq!(diagnostics["resolved_transport_profile"]["has_extra"], true);
}
#[test]
fn request_trace_proxy_value_sanitizes_url_and_marks_config_source() {
let transport = sample_transport();
@@ -1,15 +1,26 @@
use serde_json::Value;
use std::fmt;
use super::super::snapshot::GatewayProviderTransportSnapshot;
pub const GEMINI_CLI_PROVIDER_TYPE: &str = "gemini_cli";
#[derive(Debug, Clone, Default, PartialEq, Eq)]
#[derive(Clone, Default, PartialEq, Eq)]
pub struct GeminiCliRequestAuth {
pub project_id: Option<String>,
pub session_id: Option<String>,
}
impl fmt::Debug for GeminiCliRequestAuth {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("GeminiCliRequestAuth")
.field("project_id", &self.project_id)
.field("has_session_id", &self.session_id.is_some())
.finish()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum GeminiCliRequestAuthSupport {
Supported(GeminiCliRequestAuth),
@@ -1,13 +1,31 @@
use serde_json::{Map, Value};
use std::fmt;
use super::auth::GeminiCliRequestAuth;
#[derive(Debug, Clone, PartialEq)]
#[derive(Clone, PartialEq)]
pub enum GeminiCliRequestEnvelopeSupport {
Supported(Value),
Unsupported(GeminiCliRequestEnvelopeUnsupportedReason),
}
impl fmt::Debug for GeminiCliRequestEnvelopeSupport {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Supported(body) => formatter
.debug_struct("Supported")
.field(
"body_bytes",
&serde_json::to_vec(body).ok().map(|bytes| bytes.len()),
)
.finish(),
Self::Unsupported(reason) => {
formatter.debug_tuple("Unsupported").field(reason).finish()
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum GeminiCliRequestEnvelopeUnsupportedReason {
NonObjectBody,
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::fmt;
use serde_json::{json, Value};
@@ -17,13 +18,36 @@ pub enum GeminiFilesRequestBodyError {
BodyRulesApplyFailed,
}
#[derive(Debug, Clone, PartialEq)]
#[derive(Clone, PartialEq)]
pub struct GeminiFilesRequestBodyParts {
pub provider_request_body: Option<Value>,
pub provider_request_body_base64: Option<String>,
}
#[derive(Debug, Clone, Copy)]
impl fmt::Debug for GeminiFilesRequestBodyParts {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("GeminiFilesRequestBodyParts")
.field(
"has_provider_request_body",
&self.provider_request_body.is_some(),
)
.field(
"provider_request_body_bytes",
&self
.provider_request_body
.as_ref()
.and_then(|body| serde_json::to_vec(body).ok().map(|bytes| bytes.len())),
)
.field(
"provider_request_body_base64_len",
&self.provider_request_body_base64.as_ref().map(String::len),
)
.finish()
}
}
#[derive(Clone, Copy)]
pub struct GeminiFilesHeadersInput<'a> {
pub headers: &'a http::HeaderMap,
pub auth_header: &'a str,
@@ -35,6 +59,40 @@ pub struct GeminiFilesHeadersInput<'a> {
pub original_body_is_empty: bool,
}
impl fmt::Debug for GeminiFilesHeadersInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("GeminiFilesHeadersInput")
.field(
"request_header_names",
&self
.headers
.keys()
.map(|name| name.as_str())
.collect::<Vec<_>>(),
)
.field("auth_header", &self.auth_header)
.field("has_auth_value", &(!self.auth_value.is_empty()))
.field("has_header_rules", &self.header_rules.is_some())
.field(
"has_provider_request_body",
&self.provider_request_body.is_some(),
)
.field(
"provider_request_body_base64_len",
&self.provider_request_body_base64.map(str::len),
)
.field(
"original_request_body_bytes",
&serde_json::to_vec(self.original_request_body_json)
.ok()
.map(|bytes| bytes.len()),
)
.field("original_body_is_empty", &self.original_body_is_empty)
.finish()
}
}
pub fn gemini_files_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
@@ -135,6 +193,12 @@ pub fn build_gemini_files_headers(
) {
return None;
}
let declared_connection_headers =
crate::headers::declared_connection_header_names(input.headers, &BTreeMap::new());
crate::headers::remove_declared_connection_headers(
&mut provider_request_headers,
&declared_connection_headers,
);
Some(provider_request_headers)
}
@@ -62,9 +62,26 @@ pub fn resolve_local_generic_oauth_transport_authorization(
.map(|token| format!("Bearer {token}"))
}
#[derive(Debug, Clone, Default)]
#[derive(Clone, Default)]
pub struct GenericOAuthRefreshAdapter {
token_url_overrides: BTreeMap<String, String>,
oauth_credentials_overrides: BTreeMap<String, (String, String)>,
}
impl std::fmt::Debug for GenericOAuthRefreshAdapter {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("GenericOAuthRefreshAdapter")
.field(
"token_url_override_provider_types",
&self.token_url_overrides.keys().collect::<Vec<_>>(),
)
.field(
"oauth_credentials_override_provider_types",
&self.oauth_credentials_overrides.keys().collect::<Vec<_>>(),
)
.finish()
}
}
impl GenericOAuthRefreshAdapter {
@@ -78,11 +95,29 @@ impl GenericOAuthRefreshAdapter {
self
}
pub fn with_oauth_credentials_for_tests(
mut self,
provider_type: &str,
client_id: impl Into<String>,
client_secret: impl Into<String>,
) -> Self {
self.oauth_credentials_overrides.insert(
provider_type.trim().to_ascii_lowercase(),
(client_id.into(), client_secret.into()),
);
self
}
fn adapter_for_provider_type(
&self,
provider_type: &'static str,
) -> Option<GenericProviderOAuthAdapter> {
let adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?;
let mut adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?;
if let Some((client_id, client_secret)) =
self.oauth_credentials_overrides.get(provider_type)
{
adapter = adapter.with_oauth_credentials_for_tests(client_id, client_secret);
}
if let Some(token_url) = self.token_url_overrides.get(provider_type) {
return Some(adapter.with_token_url_override(token_url.clone()));
}
@@ -912,7 +947,12 @@ mod tests {
hits: Arc::clone(&hits),
};
let adapter = GenericOAuthRefreshAdapter::default()
.with_token_url_for_tests("antigravity", "https://oauth.example/token");
.with_token_url_for_tests("antigravity", "https://oauth.example/token")
.with_oauth_credentials_for_tests(
"antigravity",
"test-client-id",
"test-client-secret",
);
assert!(adapter.supports(&transport));
assert!(adapter.should_refresh(&transport, None));
+41 -1
View File
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::fmt;
use aether_contracts::{
ResolvedTransportProfile, TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_HTTP_MODE_AUTO,
@@ -33,7 +34,7 @@ pub struct GrokBrowserProfileMetadata {
pub sec_ch_ua_platform: String,
}
#[derive(Debug, Clone)]
#[derive(Clone)]
pub struct GrokHeaderInput<'a> {
pub transport: &'a GatewayProviderTransportSnapshot,
pub transport_profile: Option<&'a ResolvedTransportProfile>,
@@ -45,6 +46,37 @@ pub struct GrokHeaderInput<'a> {
pub original_request_body: &'a Value,
}
impl fmt::Debug for GrokHeaderInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("GrokHeaderInput")
.field("transport", &self.transport)
.field("transport_profile", &self.transport_profile)
.field(
"request_header_names",
&self
.request_headers
.map(|headers| headers.keys().map(|name| name.as_str()).collect::<Vec<_>>()),
)
.field("content_type", &self.content_type)
.field("accept", &self.accept)
.field("has_header_rules", &self.header_rules.is_some())
.field(
"provider_request_body_bytes",
&serde_json::to_vec(self.provider_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field(
"original_request_body_bytes",
&serde_json::to_vec(self.original_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.finish()
}
}
pub fn is_grok_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
@@ -254,6 +286,14 @@ pub fn build_grok_browser_headers(input: GrokHeaderInput<'_>) -> Option<BTreeMap
) {
return None;
}
if let Some(request_headers) = input.request_headers {
let declared_connection_headers =
super::headers::declared_connection_header_names(request_headers, &BTreeMap::new());
super::headers::remove_declared_connection_headers(
&mut headers,
&declared_connection_headers,
);
}
Some(headers)
}
+146 -2
View File
@@ -1,7 +1,35 @@
use std::collections::BTreeMap;
use std::collections::{BTreeMap, BTreeSet};
use aether_contracts::USAGE_SERVER_NOW_UNIX_MS_HEADER;
/// Return the case-insensitive field names nominated by all `Connection`
/// header values in a request. Those fields are hop-by-hop even when their
/// names are application-defined and therefore must not cross to a provider.
pub(crate) fn declared_connection_header_names(
headers: &http::HeaderMap,
extra_headers: &BTreeMap<String, String>,
) -> BTreeSet<String> {
let mut values = headers
.get_all(http::header::CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok())
.collect::<Vec<_>>();
values.extend(
extra_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case(http::header::CONNECTION.as_str()))
.map(|(_, value)| value.as_str()),
);
aether_http::connection_declared_header_names(values)
}
pub(crate) fn is_declared_connection_header(
name: &str,
declared_connection_headers: &BTreeSet<String>,
) -> bool {
declared_connection_headers.contains(&name.trim().to_ascii_lowercase())
}
const UPSTREAM_CREDENTIAL_HEADER_NAMES: &[&str] = &[
"authorization",
"proxy-authorization",
@@ -30,7 +58,10 @@ pub(crate) fn is_aether_internal_header(name: &str) -> bool {
pub fn should_skip_request_header(name: &str) -> bool {
let normalized = name.to_ascii_lowercase();
if is_aether_internal_header(&normalized) {
if is_aether_internal_header(&normalized)
|| is_untrusted_forwarding_metadata_header(&normalized)
|| is_untrusted_routing_override_header(&normalized)
{
return true;
}
matches!(
@@ -49,6 +80,27 @@ pub fn should_skip_request_header(name: &str) -> bool {
)
}
fn is_untrusted_forwarding_metadata_header(normalized_name: &str) -> bool {
matches!(normalized_name, "forwarded" | "via")
|| normalized_name.starts_with("x-forwarded-")
|| normalized_name.starts_with("x_forwarded_")
|| normalized_name.starts_with("x-real-")
|| normalized_name.starts_with("x_real_")
}
fn is_untrusted_routing_override_header(normalized_name: &str) -> bool {
matches!(
normalized_name,
"x-http-method"
| "x-http-method-override"
| "x-method-override"
| "x-override-url"
| "x-rewrite-url"
) || normalized_name.starts_with("x-original-")
|| normalized_name.starts_with("x_original_")
|| normalized_name.starts_with("x-envoy-original-")
}
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
@@ -82,6 +134,14 @@ pub fn should_skip_upstream_passthrough_header(name: &str) -> bool {
|| should_skip_request_header(name)
}
pub(crate) fn should_skip_upstream_passthrough_header_with_connection(
name: &str,
declared_connection_headers: &BTreeSet<String>,
) -> bool {
should_skip_upstream_passthrough_header(name)
|| is_declared_connection_header(name, declared_connection_headers)
}
pub(crate) fn should_skip_upstream_complete_passthrough_header(name: &str) -> bool {
let lower = name.to_ascii_lowercase();
is_upstream_credential_header(&lower)
@@ -103,6 +163,33 @@ pub(crate) fn should_skip_upstream_complete_passthrough_header(name: &str) -> bo
|| should_skip_request_header(name)
}
pub(crate) fn should_skip_upstream_complete_passthrough_header_with_connection(
name: &str,
declared_connection_headers: &BTreeSet<String>,
) -> bool {
should_skip_upstream_complete_passthrough_header(name)
|| is_declared_connection_header(name, declared_connection_headers)
}
pub(crate) fn remove_declared_connection_headers(
headers: &mut BTreeMap<String, String>,
declared_connection_headers: &BTreeSet<String>,
) {
let mut all_declared = declared_connection_headers.clone();
let connection_values = headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("connection"))
.map(|(_, value)| value.as_str())
.collect::<Vec<_>>();
all_declared.extend(aether_http::connection_declared_header_names(
connection_values,
));
headers.retain(|name, _| {
!name.eq_ignore_ascii_case("connection")
&& !is_declared_connection_header(name, &all_declared)
});
}
pub fn normalize_upstream_accept_encoding(value: &str) -> Option<String> {
let mut accepted = Vec::new();
let mut wildcard_allowed = false;
@@ -337,6 +424,44 @@ mod tests {
}
}
#[test]
fn strips_all_client_supplied_forwarding_metadata() {
for header in [
"Forwarded",
"Via",
"X-Forwarded-For",
"X-Forwarded-Host",
"X-Forwarded-Prefix",
"X-Forwarded-Server",
"x_forwarded_for",
"X-Real-IP",
"X-Real-Host",
"x_real_ip",
] {
assert!(should_skip_request_header(header));
assert!(should_skip_upstream_passthrough_header(header));
assert!(should_skip_upstream_complete_passthrough_header(header));
}
}
#[test]
fn strips_client_supplied_upstream_routing_overrides() {
for header in [
"X-HTTP-Method-Override",
"X-Method-Override",
"X-Original-URL",
"X-Original-URI",
"x_original_url",
"X-Rewrite-URL",
"X-Override-URL",
"X-Envoy-Original-Path",
] {
assert!(should_skip_request_header(header));
assert!(should_skip_upstream_passthrough_header(header));
assert!(should_skip_upstream_complete_passthrough_header(header));
}
}
#[test]
fn strips_usage_server_time_header_from_provider_requests() {
for h in [
@@ -365,4 +490,23 @@ mod tests {
);
}
}
#[test]
fn declared_connection_names_are_case_insensitive_and_multi_line() {
let mut headers = http::HeaderMap::new();
headers.append(
http::header::CONNECTION,
http::HeaderValue::from_static("X-Hop, keep-alive"),
);
headers.append(
http::header::CONNECTION,
http::HeaderValue::from_static("x-other-hop"),
);
let names = super::declared_connection_header_names(&headers, &BTreeMap::new());
assert!(names.contains("x-hop"));
assert!(names.contains("keep-alive"));
assert!(names.contains("x-other-hop"));
assert!(super::is_declared_connection_header("X-HOP", &names));
}
}
@@ -1,16 +1,27 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::credentials::{generate_machine_id, KiroAuthConfig};
use std::fmt;
pub const PROVIDER_TYPE: &str = "kiro";
pub const KIRO_AUTH_HEADER: &str = "authorization";
#[derive(Debug, Clone, PartialEq, Eq)]
#[derive(Clone, PartialEq, Eq)]
pub struct KiroBearerAuth {
pub name: &'static str,
pub value: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
impl fmt::Debug for KiroBearerAuth {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("KiroBearerAuth")
.field("name", &self.name)
.field("value", &"[REDACTED]")
.finish()
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct KiroRequestAuth {
pub name: &'static str,
pub value: String,
@@ -18,6 +29,18 @@ pub struct KiroRequestAuth {
pub machine_id: String,
}
impl fmt::Debug for KiroRequestAuth {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("KiroRequestAuth")
.field("name", &self.name)
.field("value", &"[REDACTED]")
.field("auth_config", &self.auth_config)
.field("machine_id", &"[REDACTED]")
.finish()
}
}
pub fn is_kiro_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
@@ -224,6 +247,41 @@ mod tests {
assert!(supports_local_kiro_auth_prerequisites(&sample_transport()));
}
#[test]
fn kiro_request_auth_debug_output_redacts_credentials_and_machine_identity() {
let bearer = resolve_local_kiro_bearer_auth(&sample_transport())
.expect("kiro bearer auth should resolve");
let bearer_debug = format!("{bearer:?}");
assert!(!bearer_debug.contains("upstream-key"));
assert!(bearer_debug.contains("[REDACTED]"));
let mut transport = sample_transport();
transport.key.decrypted_api_key = "__placeholder__".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"access_token":"kiro-request-access-canary",
"expires_at":4102444800,
"refresh_token":"kiro-request-refresh-canary................................................................................................",
"machine_id":"kiro-request-machine-canary",
"profile_arn":"kiro-request-profile-canary",
"api_region":"us-west-2"
}"#
.to_string(),
);
let request_auth =
resolve_local_kiro_request_auth(&transport).expect("kiro request auth should resolve");
let request_debug = format!("{request_auth:?}");
for secret in [
"kiro-request-access-canary",
"kiro-request-refresh-canary",
"kiro-request-machine-canary",
"kiro-request-profile-canary",
] {
assert!(!request_debug.contains(secret), "debug leaked {secret}");
}
assert!(request_debug.contains("[REDACTED]"));
}
#[test]
fn rejects_auth_config_subset() {
let mut transport = sample_transport();
@@ -2,6 +2,7 @@ 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,
};
pub use aether_oauth::provider::providers::{is_valid_kiro_region, normalize_kiro_region};
#[cfg(test)]
mod tests {
@@ -2,7 +2,7 @@ use std::collections::BTreeMap;
use uuid::Uuid;
use super::credentials::KiroAuthConfig;
use super::credentials::{normalize_kiro_region, KiroAuthConfig};
pub const AWS_EVENTSTREAM_CONTENT_TYPE: &str = "application/vnd.amazon.eventstream";
pub const KIRO_PROFILE_ARN_HEADER: &str = "x-amzn-kiro-profile-arn";
@@ -47,7 +47,7 @@ pub fn build_generate_assistant_headers(
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 region = normalize_kiro_region(auth_config.effective_api_region());
let host = format!("q.{region}.amazonaws.com");
BTreeMap::from([
@@ -92,7 +92,7 @@ pub fn build_mcp_headers(
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 region = normalize_kiro_region(auth_config.effective_api_region());
let host = format!("q.{region}.amazonaws.com");
let mut headers = BTreeMap::from([
@@ -134,7 +134,7 @@ pub fn build_list_available_models_headers(
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 region = normalize_kiro_region(auth_config.effective_api_region());
let host = format!("q.{region}.amazonaws.com");
let ide_tag = build_kiro_ide_tag(kiro_version, machine_id);
@@ -331,4 +331,30 @@ mod tests {
);
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
}
#[test]
fn malicious_api_region_cannot_inject_host_header() {
let auth_config = KiroAuthConfig {
auth_method: Some("social".to_string()),
refresh_token: None,
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: Some("attacker.example/".to_string()),
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let headers = build_list_available_models_headers(&auth_config, "machine");
assert_eq!(
headers.get("host").map(String::as_str),
Some("q.us-east-1.amazonaws.com")
);
}
}
@@ -30,7 +30,10 @@ pub use auth::{
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 credentials::{
generate_machine_id, is_valid_kiro_region, normalize_kiro_region, 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,
@@ -1,7 +1,9 @@
use std::collections::BTreeMap;
use std::fmt;
use serde_json::{json, Value};
use super::super::headers::{declared_connection_header_names, remove_declared_connection_headers};
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,
@@ -83,7 +85,7 @@ pub fn build_kiro_provider_request_body(
Some(provider_request_body)
}
#[derive(Clone, Copy, Debug)]
#[derive(Clone, Copy)]
pub struct KiroProviderHeadersInput<'a> {
pub headers: &'a http::HeaderMap,
pub provider_request_body: &'a Value,
@@ -95,6 +97,39 @@ pub struct KiroProviderHeadersInput<'a> {
pub machine_id: &'a str,
}
impl fmt::Debug for KiroProviderHeadersInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("KiroProviderHeadersInput")
.field(
"request_header_names",
&self
.headers
.keys()
.map(|name| name.as_str())
.collect::<Vec<_>>(),
)
.field(
"provider_request_body_bytes",
&serde_json::to_vec(self.provider_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field(
"original_request_body_bytes",
&serde_json::to_vec(self.original_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field("has_header_rules", &self.header_rules.is_some())
.field("auth_header", &self.auth_header)
.field("has_auth_value", &(!self.auth_value.is_empty()))
.field("auth_config", &self.auth_config)
.field("has_machine_id", &(!self.machine_id.is_empty()))
.finish()
}
}
pub fn build_kiro_provider_headers(
input: KiroProviderHeadersInput<'_>,
) -> Option<BTreeMap<String, String>> {
@@ -109,13 +144,16 @@ pub fn build_kiro_provider_headers(
machine_id,
} = input;
let declared_connection_headers = declared_connection_header_names(headers, &BTreeMap::new());
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) {
if should_skip_upstream_passthrough_header(&key)
|| declared_connection_headers.contains(&key)
{
continue;
}
let value = value.trim();
@@ -143,8 +181,10 @@ pub fn build_kiro_provider_headers(
auth_header.trim().to_ascii_lowercase(),
auth_value.trim().to_string(),
);
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out.entry("content-type".to_string())
.or_insert_with(|| "application/json".to_string());
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out.remove("content-length");
Some(out)
}
@@ -1,5 +1,5 @@
use super::super::url::build_passthrough_path_url;
use super::credentials::DEFAULT_REGION;
use super::credentials::{normalize_kiro_region, DEFAULT_REGION};
pub const GENERATE_ASSISTANT_RESPONSE_PATH: &str = "/generateAssistantResponse";
pub const LIST_AVAILABLE_MODELS_PATH: &str = "/ListAvailableModels";
@@ -9,8 +9,7 @@ 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())
.map(normalize_kiro_region)
.unwrap_or(DEFAULT_REGION);
upstream_base_url
.trim()
@@ -105,6 +104,29 @@ mod tests {
);
}
#[test]
fn malicious_region_cannot_escape_configured_origin() {
for region in [
"attacker.example/",
"attacker.example?x=1",
"attacker.example#fragment",
"attacker@example",
] {
assert_eq!(
resolve_kiro_base_url("https://q.{region}.amazonaws.com", Some(region)),
"https://q.us-east-1.amazonaws.com"
);
}
assert_eq!(
build_kiro_list_available_models_url(
"https://q.{region}.amazonaws.com",
Some("attacker.example/"),
)
.as_deref(),
Some("https://q.us-east-1.amazonaws.com/ListAvailableModels?origin=AI_EDITOR")
);
}
#[test]
fn builds_mcp_url_for_latest_kiro_endpoint() {
assert_eq!(
@@ -320,13 +320,28 @@ fn proxy_snapshot_from_value(value: &Value) -> Option<ProxySnapshot> {
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 url = json_string_field(object, "url")
.or_else(|| json_string_field(object, "proxy_url"))
.and_then(|proxy_url| {
proxy_url_with_auth(
&proxy_url,
json_proxy_credential_field(object, "username"),
json_proxy_credential_field(object, "password"),
)
});
let mut extra = Map::new();
for (key, value) in object {
if matches!(
key.as_str(),
"enabled" | "mode" | "node_id" | "label" | "url" | "proxy_url"
"enabled"
| "mode"
| "node_id"
| "label"
| "url"
| "proxy_url"
| "username"
| "password"
) {
continue;
}
@@ -356,6 +371,35 @@ fn json_string_field(object: &Map<String, Value>, key: &str) -> Option<String> {
.map(ToOwned::to_owned)
}
fn json_proxy_credential_field<'a>(object: &'a Map<String, Value>, key: &str) -> Option<&'a str> {
object
.get(key)
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
}
fn proxy_url_with_auth(
proxy_url: &str,
username: Option<&str>,
password: Option<&str>,
) -> Option<String> {
let username = username.filter(|value| !value.is_empty());
let password = password.filter(|value| !value.is_empty());
let mut parsed = url::Url::parse(proxy_url).ok()?;
if !matches!(parsed.scheme(), "http" | "https" | "socks5" | "socks5h")
|| parsed.host_str().is_none()
{
return None;
}
if username.is_none() && password.is_none() {
return Some(parsed.to_string());
}
let username = username.unwrap_or("");
parsed.set_username(username).ok()?;
parsed.set_password(password).ok()?;
Some(parsed.to_string())
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
@@ -520,6 +564,43 @@ mod tests {
assert_eq!(snapshot.extra, Some(json!({"kind":"manual"})));
}
#[test]
fn resolves_authenticated_inline_proxy_without_secret_extra_fields() {
let mut transport = sample_transport();
transport.key.proxy = Some(json!({
"url": "socks5h://proxy.example:1080",
"username": " alice ",
"password": " p:ss ",
"kind": "manual",
}));
let snapshot = resolve_transport_proxy_snapshot(&transport)
.expect("authenticated proxy snapshot should resolve");
assert_eq!(
snapshot.url.as_deref(),
Some("socks5h://%20alice%20:%20p%3Ass%[email protected]:1080")
);
assert_eq!(snapshot.extra, Some(json!({"kind":"manual"})));
}
#[test]
fn resolves_legacy_password_only_inline_proxy() {
let mut transport = sample_transport();
transport.key.proxy = Some(json!({
"url": "http://proxy.example:8080",
"password": "legacy-password",
}));
let snapshot = resolve_transport_proxy_snapshot(&transport)
.expect("password-only proxy snapshot should resolve");
assert_eq!(
snapshot.url.as_deref(),
Some("http://:[email protected]:8080/")
);
assert!(snapshot.extra.is_none());
}
#[tokio::test]
async fn enriches_transport_proxy_snapshot_with_tunnel_owner_hint() {
let state = sample_lookup();
@@ -25,7 +25,9 @@ use super::vertex::{
supports_local_vertex_service_account_auth_resolution, VertexServiceAccountRefreshAdapter,
};
#[derive(Debug, Clone, PartialEq, Eq)]
const LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024;
#[derive(Clone, PartialEq, Eq)]
#[allow(clippy::large_enum_variant)]
pub enum LocalResolvedOAuthRequestAuth {
#[allow(dead_code)]
@@ -36,7 +38,20 @@ pub enum LocalResolvedOAuthRequestAuth {
Kiro(KiroRequestAuth),
}
#[derive(Debug, Clone, PartialEq)]
impl fmt::Debug for LocalResolvedOAuthRequestAuth {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Header { name, .. } => formatter
.debug_struct("Header")
.field("name", name)
.field("value", &"[REDACTED]")
.finish(),
Self::Kiro(auth) => formatter.debug_tuple("Kiro").field(auth).finish(),
}
}
}
#[derive(Clone, PartialEq)]
pub struct LocalOAuthResolution {
pub auth: Option<LocalResolvedOAuthRequestAuth>,
pub refreshed_entry: Option<CachedOAuthEntry>,
@@ -53,6 +68,23 @@ pub struct LocalOAuthResolution {
pub local_refresh_guard: Option<LocalOAuthRefreshCommitGuard>,
}
impl fmt::Debug for LocalOAuthResolution {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("LocalOAuthResolution")
.field("auth", &self.auth)
.field("refreshed_entry", &self.refreshed_entry)
.field("refresh_in_flight", &self.refresh_in_flight)
.field("reused_refresh", &self.reused_refresh)
.field("has_distributed_lease", &self.distributed_lease.is_some())
.field(
"has_local_refresh_guard",
&self.local_refresh_guard.is_some(),
)
.finish()
}
}
#[derive(Clone)]
pub struct LocalOAuthRefreshCommitGuard {
guard: Arc<OwnedMutexGuard<()>>,
@@ -78,7 +110,7 @@ impl PartialEq for LocalOAuthRefreshCommitGuard {
}
}
#[derive(Debug, Clone, PartialEq)]
#[derive(Clone, PartialEq)]
pub struct CachedOAuthEntry {
pub provider_type: String,
pub auth_header_name: String,
@@ -89,7 +121,21 @@ pub struct CachedOAuthEntry {
pub source_fingerprint: Option<String>,
}
#[derive(Debug, Clone, PartialEq)]
impl fmt::Debug for CachedOAuthEntry {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CachedOAuthEntry")
.field("provider_type", &self.provider_type)
.field("auth_header_name", &self.auth_header_name)
.field("auth_header_value", &"[REDACTED]")
.field("expires_at_unix_secs", &self.expires_at_unix_secs)
.field("has_metadata", &self.metadata.is_some())
.field("has_source_fingerprint", &self.source_fingerprint.is_some())
.finish()
}
}
#[derive(Clone, PartialEq)]
pub struct LocalOAuthHttpRequest {
pub request_id: &'static str,
pub method: reqwest::Method,
@@ -99,21 +145,44 @@ pub struct LocalOAuthHttpRequest {
pub body_bytes: Option<Vec<u8>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
impl fmt::Debug for LocalOAuthHttpRequest {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("LocalOAuthHttpRequest")
.field("request_id", &self.request_id)
.field("method", &self.method)
.field("url", &"[REDACTED]")
.field("header_names", &self.headers.keys().collect::<Vec<_>>())
.field("json_body", &self.json_body.as_ref().map(|_| "[REDACTED]"))
.field("body_bytes_len", &self.body_bytes.as_ref().map(Vec::len))
.finish()
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct LocalOAuthHttpResponse {
pub status_code: u16,
pub body_text: String,
}
#[derive(Debug, Error)]
impl fmt::Debug for LocalOAuthHttpResponse {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("LocalOAuthHttpResponse")
.field("status_code", &self.status_code)
.field("body_bytes_len", &self.body_text.len())
.finish()
}
}
#[derive(Error)]
pub enum LocalOAuthRefreshError {
#[error("{provider_type} oauth refresh request failed: {source}")]
#[error("{provider_type} oauth refresh request failed")]
Transport {
provider_type: &'static str,
#[source]
source: reqwest::Error,
error: reqwest::Error,
},
#[error("{provider_type} oauth refresh returned HTTP {status_code}: {body_excerpt}")]
#[error("{provider_type} oauth refresh returned HTTP {status_code}")]
HttpStatus {
provider_type: &'static str,
status_code: u16,
@@ -152,6 +221,64 @@ impl ReqwestLocalOAuthHttpExecutor {
}
}
impl fmt::Debug for LocalOAuthRefreshError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Transport { provider_type, .. } => formatter
.debug_struct("Transport")
.field("provider_type", provider_type)
.field("error", &"[REDACTED]")
.finish(),
Self::HttpStatus {
provider_type,
status_code,
..
} => formatter
.debug_struct("HttpStatus")
.field("provider_type", provider_type)
.field("status_code", status_code)
.field("body_excerpt", &"[REDACTED]")
.finish(),
Self::TransportMessage { provider_type, .. } => formatter
.debug_struct("TransportMessage")
.field("provider_type", provider_type)
.field("message", &"[REDACTED]")
.finish(),
Self::InvalidResponse { provider_type, .. } => formatter
.debug_struct("InvalidResponse")
.field("provider_type", provider_type)
.field("message", &"[REDACTED]")
.finish(),
}
}
}
fn validate_local_oauth_request_url(
provider_type: &'static str,
raw_url: &str,
) -> Result<url::Url, LocalOAuthRefreshError> {
let invalid = |message: &'static str| LocalOAuthRefreshError::TransportMessage {
provider_type,
message: message.to_string(),
};
let url = url::Url::parse(raw_url).map_err(|_| invalid("invalid OAuth endpoint URL"))?;
if url.host().is_none() || !matches!(url.scheme(), "http" | "https") {
return Err(invalid(
"OAuth endpoint URL must use HTTP or HTTPS and include a host",
));
}
if !url.username().is_empty() || url.password().is_some() {
return Err(invalid("OAuth endpoint URL must not include credentials"));
}
if url.fragment().is_some() {
return Err(invalid("OAuth endpoint URL must not include a fragment"));
}
if !aether_http::is_https_or_loopback_http_url(&url) {
return Err(invalid("remote OAuth endpoint URL must use HTTPS"));
}
Ok(url)
}
#[async_trait]
impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor {
async fn execute(
@@ -160,9 +287,8 @@ impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor {
_transport: &GatewayProviderTransportSnapshot,
request: &LocalOAuthHttpRequest,
) -> Result<LocalOAuthHttpResponse, LocalOAuthRefreshError> {
let mut builder = self
.client
.request(request.method.clone(), request.url.as_str());
let request_url = validate_local_oauth_request_url(provider_type, request.url.as_str())?;
let mut builder = self.client.request(request.method.clone(), request_url);
for (name, value) in &request.headers {
builder = builder.header(name, value);
}
@@ -172,23 +298,37 @@ impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor {
builder = builder.body(body_bytes.clone());
}
let response =
let mut response =
builder
.send()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
.map_err(|error| LocalOAuthRefreshError::Transport {
provider_type,
source,
error,
})?;
let status_code = response.status().as_u16();
let body_text =
if response
.content_length()
.is_some_and(|length| length > LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES as u64)
{
return Err(local_oauth_response_too_large(provider_type));
}
let mut body = Vec::new();
while let Some(chunk) =
response
.text()
.chunk()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
.map_err(|error| LocalOAuthRefreshError::Transport {
provider_type,
source,
})?;
error,
})?
{
if chunk.len() > LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES.saturating_sub(body.len()) {
return Err(local_oauth_response_too_large(provider_type));
}
body.extend_from_slice(&chunk);
}
let body_text = String::from_utf8_lossy(&body).to_string();
Ok(LocalOAuthHttpResponse {
status_code,
body_text,
@@ -196,6 +336,13 @@ impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor {
}
}
fn local_oauth_response_too_large(provider_type: &'static str) -> LocalOAuthRefreshError {
LocalOAuthRefreshError::InvalidResponse {
provider_type,
message: format!("response body exceeds {LOCAL_OAUTH_RESPONSE_BODY_LIMIT_BYTES} bytes"),
}
}
pub(crate) struct ProviderOAuthLocalHttpExecutor<'a> {
provider_type: &'static str,
transport: &'a GatewayProviderTransportSnapshot,
@@ -299,8 +446,8 @@ pub(crate) fn oauth_error_to_local_refresh_error(
fn local_refresh_error_to_oauth_error(error: LocalOAuthRefreshError) -> OAuthError {
match error {
LocalOAuthRefreshError::Transport { source, .. } => {
OAuthError::Transport(source.to_string())
LocalOAuthRefreshError::Transport { .. } => {
OAuthError::Transport("oauth refresh request failed".to_string())
}
LocalOAuthRefreshError::TransportMessage { message, .. } => OAuthError::Transport(message),
LocalOAuthRefreshError::HttpStatus {
@@ -965,13 +1112,128 @@ mod tests {
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter,
LocalOAuthRefreshCoordinator, LocalOAuthRefreshError, LocalOAuthResolution,
LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
validate_local_oauth_request_url, CachedOAuthEntry, LocalOAuthHttpExecutor,
LocalOAuthRefreshAdapter, LocalOAuthRefreshCoordinator, LocalOAuthRefreshError,
LocalOAuthResolution, LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
};
use async_trait::async_trait;
use std::sync::Arc;
#[test]
fn local_oauth_transport_requires_https_or_literal_loopback_http() {
for allowed in [
"https://oauth.example.test/token?tenant=one",
"http://localhost:8080/token",
"http://127.42.0.1:8080/token",
"http://[::1]:8080/token",
] {
assert!(
validate_local_oauth_request_url("test", allowed).is_ok(),
"URL should be accepted: {allowed}"
);
}
for rejected in [
"http://oauth.example.test/token",
"http://10.0.0.1/token",
"http://0.0.0.0:8080/token",
"http://[::ffff:127.0.0.1]:8080/token",
"https://[email protected]/token",
"https://oauth.example.test/token#secret",
"file:///tmp/token",
] {
assert!(
validate_local_oauth_request_url("test", rejected).is_err(),
"URL should be rejected: {rejected}"
);
}
}
#[test]
fn local_oauth_url_error_does_not_echo_embedded_credentials() {
let error = validate_local_oauth_request_url(
"test",
"https://sensitive-user:[email protected]/token",
)
.expect_err("userinfo should be rejected")
.to_string();
assert!(!error.contains("sensitive-user"));
assert!(!error.contains("sensitive-password"));
}
#[test]
fn local_oauth_debug_output_redacts_request_response_and_cached_credentials() {
let request = super::LocalOAuthHttpRequest {
request_id: "provider-oauth:test",
method: reqwest::Method::POST,
url: "https://oauth.example.test/token?client_secret=url-secret-canary".to_string(),
headers: std::collections::BTreeMap::from([
(
"authorization".to_string(),
"Bearer request-header-canary".to_string(),
),
("content-type".to_string(), "application/json".to_string()),
]),
json_body: Some(serde_json::json!({
"refresh_token": "request-body-canary"
})),
body_bytes: None,
};
let response = super::LocalOAuthHttpResponse {
status_code: 401,
body_text: "{\"access_token\":\"response-body-canary\"}".to_string(),
};
let entry = super::CachedOAuthEntry {
provider_type: "test".to_string(),
auth_header_name: "authorization".to_string(),
auth_header_value: "Bearer cached-header-canary".to_string(),
expires_at_unix_secs: Some(1),
metadata: Some(serde_json::json!({
"refresh_token": "cached-metadata-canary"
})),
source_fingerprint: Some("non-secret-fingerprint".to_string()),
};
let resolution = super::LocalOAuthResolution {
auth: Some(super::LocalResolvedOAuthRequestAuth::Header {
name: "authorization".to_string(),
value: "Bearer resolution-header-canary".to_string(),
}),
refreshed_entry: Some(entry.clone()),
refresh_in_flight: false,
reused_refresh: false,
distributed_lease: None,
local_refresh_guard: None,
};
let debug = format!("{request:?} {response:?} {entry:?} {resolution:?}");
for secret in [
"url-secret-canary",
"request-header-canary",
"request-body-canary",
"response-body-canary",
"cached-header-canary",
"cached-metadata-canary",
"resolution-header-canary",
] {
assert!(!debug.contains(secret), "debug leaked {secret}");
}
assert!(debug.contains("[REDACTED]"));
}
#[test]
fn local_oauth_refresh_error_debug_and_display_hide_body_and_transport_details() {
let http_error = LocalOAuthRefreshError::HttpStatus {
provider_type: "test",
status_code: 401,
body_excerpt: "{\"access_token\":\"refresh-error-canary\"}".to_string(),
};
let debug = format!("{http_error:?}");
let display = http_error.to_string();
assert!(!debug.contains("refresh-error-canary"));
assert!(!display.contains("refresh-error-canary"));
assert_eq!(display, "test oauth refresh returned HTTP 401");
}
#[derive(Debug)]
struct TestAdapter {
refresh_hits: Arc<AtomicUsize>,
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::fmt;
use serde_json::Value;
@@ -9,7 +10,7 @@ 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)]
#[derive(Clone, Copy)]
pub struct ProviderOpenAiImageHeadersInput<'a> {
pub transport: &'a GatewayProviderTransportSnapshot,
pub headers: &'a http::HeaderMap,
@@ -21,6 +22,39 @@ pub struct ProviderOpenAiImageHeadersInput<'a> {
pub original_request_body: &'a Value,
}
impl fmt::Debug for ProviderOpenAiImageHeadersInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ProviderOpenAiImageHeadersInput")
.field("transport", &self.transport)
.field(
"request_header_names",
&self
.headers
.keys()
.map(|name| name.as_str())
.collect::<Vec<_>>(),
)
.field("auth_header", &self.auth_header)
.field("has_auth_value", &(!self.auth_value.is_empty()))
.field("accept", &self.accept)
.field("has_header_rules", &self.header_rules.is_some())
.field(
"provider_request_body_bytes",
&serde_json::to_vec(self.provider_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field(
"original_request_body_bytes",
&serde_json::to_vec(self.original_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.finish()
}
}
pub fn openai_image_transport_unsupported_reason(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
@@ -98,6 +132,12 @@ pub fn build_openai_image_headers(
) {
return None;
}
let declared_connection_headers =
crate::headers::declared_connection_header_names(input.headers, &BTreeMap::new());
crate::headers::remove_declared_connection_headers(
&mut provider_request_headers,
&declared_connection_headers,
);
Some(provider_request_headers)
}
@@ -5,7 +5,6 @@ pub struct ProviderOAuthTemplate {
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,
@@ -550,7 +549,6 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
authorize_url: aether_oauth::provider::providers::CLAUDE_CODE_AUTHORIZE_URL,
token_url: aether_oauth::provider::providers::CLAUDE_CODE_TOKEN_URL,
client_id: aether_oauth::provider::providers::CLAUDE_CODE_CLIENT_ID,
client_secret: "",
scopes: aether_oauth::provider::providers::CLAUDE_CODE_OAUTH_SCOPES,
redirect_uri: aether_oauth::provider::providers::CLAUDE_CODE_REDIRECT_URI,
use_pkce: true,
@@ -561,7 +559,6 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
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,
@@ -572,7 +569,6 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
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,
@@ -583,7 +579,6 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
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",
@@ -598,7 +593,6 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
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",
@@ -615,7 +609,6 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
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,
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::fmt;
use std::sync::OnceLock;
use aether_ai_formats::ApiOperation;
@@ -19,8 +20,8 @@ use crate::url::{
build_claude_count_tokens_url as build_default_claude_count_tokens_url,
build_claude_messages_url, build_gemini_content_url, build_openai_chat_url,
build_openai_responses_url, build_openai_search_url, build_passthrough_path_url,
normalize_gemini_content_action_path, strip_gateway_credential_query_parameters,
GATEWAY_CREDENTIAL_QUERY_KEYS,
encode_url_path_segment, normalize_gemini_content_action_path,
strip_gateway_credential_query_parameters, GATEWAY_CREDENTIAL_QUERY_KEYS,
};
use crate::vertex::{
build_vertex_api_key_gemini_content_url, build_vertex_api_key_gemini_embedding_url,
@@ -28,7 +29,7 @@ use crate::vertex::{
build_vertex_service_account_gemini_embedding_url, resolve_local_vertex_api_key_query_auth,
resolve_local_vertex_service_account_auth_config,
};
#[derive(Debug, Clone, Copy)]
#[derive(Clone, Copy)]
pub struct TransportRequestUrlParams<'a> {
pub provider_api_format: &'a str,
pub mapped_model: Option<&'a str>,
@@ -38,6 +39,21 @@ pub struct TransportRequestUrlParams<'a> {
pub api_operation: Option<aether_ai_formats::ApiOperation>,
}
impl fmt::Debug for TransportRequestUrlParams<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("TransportRequestUrlParams")
.field("provider_api_format", &self.provider_api_format)
.field("mapped_model", &self.mapped_model)
.field("upstream_is_stream", &self.upstream_is_stream)
.field("has_request_query", &self.request_query.is_some())
.field("request_query_len", &self.request_query.map(str::len))
.field("kiro_api_region", &self.kiro_api_region)
.field("api_operation", &self.api_operation)
.finish()
}
}
pub fn build_transport_request_url(
transport: &GatewayProviderTransportSnapshot,
params: TransportRequestUrlParams<'_>,
@@ -112,9 +128,13 @@ fn build_transport_request_url_inner(
.filter(|value| !value.is_empty());
let custom_path_handles_operation =
custom_path_template.is_some_and(|path| path.contains("{operation}"));
let custom_path = custom_path_template.map(|path| {
expand_custom_path_template(path, build_path_params(params, gemini_embedding_batch))
});
let custom_path = match custom_path_template {
Some(path) => Some(expand_custom_path_template(
path,
build_path_params(params, gemini_embedding_batch),
)?),
None => None,
};
if let Some(path) = custom_path.as_deref() {
let custom_path_is_complete_claude_count_tokens = normalized_provider_api_format
@@ -670,22 +690,24 @@ fn build_gemini_embedding_url(
} else {
"embedContent"
};
let encoded_model = encode_url_path_segment(trimmed_model);
let path = if trimmed_base_url.ends_with("/v1beta") {
format!("/models/{trimmed_model}:{action}")
format!("/models/{encoded_model}:{action}")
} else if trimmed_base_url.contains("/v1beta/models/") {
format!(":{action}")
} else {
format!("/v1beta/models/{trimmed_model}:{action}")
format!("/v1beta/models/{encoded_model}:{action}")
};
build_passthrough_path_url(upstream_base_url, &path, query, &["key"])
}
fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>) -> String {
fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>) -> Option<String> {
if params.is_empty() {
return path.to_string();
return Some(path.to_string());
}
let regex = custom_path_template_regex();
let query_start = path.find('?');
let mut missing_key = false;
let replaced = regex.replace_all(path, |captures: &regex::Captures<'_>| {
let key = captures
@@ -693,7 +715,14 @@ fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>)
.map(|value| value.as_str())
.unwrap_or_default();
match params.get(key).copied() {
Some(value) => value.to_string(),
Some(value) => {
let capture_start = captures.get(0).map_or(0, |value| value.start());
if query_start.is_some_and(|query_start| capture_start > query_start) {
url::form_urlencoded::byte_serialize(value.as_bytes()).collect()
} else {
encode_url_path_segment(value)
}
}
None => {
missing_key = true;
captures
@@ -705,12 +734,31 @@ fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>)
});
if missing_key {
path.to_string()
Some(path.to_string())
} else {
replaced.into_owned()
let replaced = replaced.into_owned();
if dynamic_template_value_created_dot_path_segment(path, &replaced, regex) {
None
} else {
Some(replaced)
}
}
}
fn dynamic_template_value_created_dot_path_segment(
template: &str,
expanded: &str,
regex: &Regex,
) -> bool {
let template_path = template.split_once('?').map_or(template, |(path, _)| path);
let expanded_path = expanded.split_once('?').map_or(expanded, |(path, _)| path);
template_path.split('/').zip(expanded_path.split('/')).any(
|(template_segment, expanded_segment)| {
matches!(expanded_segment, "." | "..") && regex.is_match(template_segment)
},
)
}
fn maybe_add_gemini_stream_alt_sse(
upstream_url: String,
provider_api_format: &str,
@@ -2288,4 +2336,85 @@ mod tests {
"https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001:embedContent?foo=bar"
);
}
#[test]
fn custom_path_template_encodes_dynamic_model_as_one_path_segment() {
let transport = sample_transport(
"custom",
"gemini:generate_content",
"https://generativelanguage.googleapis.com",
Some("/v1beta/models/{model}:{action}"),
);
let url = build_transport_request_url(
&transport,
TransportRequestUrlParams {
provider_api_format: "gemini:generate_content",
mapped_model: Some("model/../../admin?key=attacker#fragment"),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
api_operation: None,
},
)
.expect("custom path should remain on the configured origin");
assert_eq!(
url,
"https://generativelanguage.googleapis.com/v1beta/models/model%2F..%2F..%2Fadmin%3Fkey=attacker%23fragment:generateContent"
);
}
#[test]
fn custom_path_template_rejects_dot_only_dynamic_path_segments() {
let transport = sample_transport(
"custom",
"claude:messages",
"https://api.example.com",
Some("/v1/messages/{model}/invoke"),
);
assert_eq!(
build_transport_request_url(
&transport,
TransportRequestUrlParams {
provider_api_format: "claude:messages",
mapped_model: Some(".."),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
api_operation: None,
},
),
None
);
}
#[test]
fn custom_path_template_query_values_cannot_inject_parameters() {
let transport = sample_transport(
"custom",
"claude:messages",
"https://api.example.com",
Some("/v1/messages?model={model}"),
);
let url = build_transport_request_url(
&transport,
TransportRequestUrlParams {
provider_api_format: "claude:messages",
mapped_model: Some("claude&admin=true#fragment"),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
api_operation: None,
},
)
.expect("query template should build");
assert_eq!(
url,
"https://api.example.com/v1/messages?model=claude%26admin%3Dtrue%23fragment"
);
}
}
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::fmt;
use serde::Serialize;
use serde_json::Value;
@@ -66,7 +67,7 @@ pub struct SameFormatProviderRequestBehavior {
pub report_kind: &'static str,
}
#[derive(Debug, Clone, Copy)]
#[derive(Clone, Copy)]
pub struct SameFormatProviderRequestBodyInput<'a> {
pub body_json: &'a Value,
pub mapped_model: &'a str,
@@ -83,13 +84,64 @@ pub struct SameFormatProviderRequestBodyInput<'a> {
pub enable_model_directives: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
impl fmt::Debug for SameFormatProviderRequestBodyInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("SameFormatProviderRequestBodyInput")
.field(
"body_json_bytes",
&serde_json::to_vec(self.body_json)
.ok()
.map(|bytes| bytes.len()),
)
.field("mapped_model", &self.mapped_model)
.field("client_api_format", &self.client_api_format)
.field("provider_api_format", &self.provider_api_format)
.field("source_model", &self.source_model)
.field("family", &self.family)
.field("has_body_rules", &self.body_rules.is_some())
.field(
"request_header_names",
&self
.request_headers
.map(|headers| headers.keys().map(|name| name.as_str()).collect::<Vec<_>>()),
)
.field("upstream_is_stream", &self.upstream_is_stream)
.field("force_body_stream_field", &self.force_body_stream_field)
.field("has_kiro_auth_config", &self.kiro_auth_config.is_some())
.field("is_claude_code", &self.is_claude_code)
.field("enable_model_directives", &self.enable_model_directives)
.finish()
}
}
#[derive(Clone, PartialEq, Eq, Serialize)]
pub struct SameFormatProviderRequestBodyOutput {
pub body: Value,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub compatibility_edits: Vec<SameFormatProviderCompatibilityEdit>,
}
impl fmt::Debug for SameFormatProviderRequestBodyOutput {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("SameFormatProviderRequestBodyOutput")
.field(
"body_bytes",
&serde_json::to_vec(&self.body).ok().map(|bytes| bytes.len()),
)
.field(
"compatibility_edits",
&self
.compatibility_edits
.iter()
.map(|edit| (edit.field.as_str(), edit.action, edit.detail.len()))
.collect::<Vec<_>>(),
)
.finish()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct SameFormatProviderCompatibilityEdit {
pub field: String,
@@ -106,7 +158,7 @@ pub enum SameFormatProviderCompatibilityEditAction {
OperatorRule,
}
#[derive(Debug, Clone, Copy)]
#[derive(Clone, Copy)]
pub struct SameFormatProviderUpstreamUrlParams<'a> {
pub provider_api_format: &'a str,
pub mapped_model: &'a str,
@@ -117,7 +169,26 @@ pub struct SameFormatProviderUpstreamUrlParams<'a> {
pub provider_request_body: Option<&'a Value>,
}
#[derive(Debug, Clone, Copy)]
impl fmt::Debug for SameFormatProviderUpstreamUrlParams<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("SameFormatProviderUpstreamUrlParams")
.field("provider_api_format", &self.provider_api_format)
.field("mapped_model", &self.mapped_model)
.field("upstream_is_stream", &self.upstream_is_stream)
.field("has_request_query", &self.request_query.is_some())
.field("request_query_len", &self.request_query.map(str::len))
.field("kiro_api_region", &self.kiro_api_region)
.field("api_operation", &self.api_operation)
.field(
"has_provider_request_body",
&self.provider_request_body.is_some(),
)
.finish()
}
}
#[derive(Clone, Copy)]
pub struct SameFormatProviderHeadersInput<'a> {
pub headers: &'a http::HeaderMap,
pub provider_request_body: &'a Value,
@@ -132,6 +203,45 @@ pub struct SameFormatProviderHeadersInput<'a> {
pub kiro_machine_id: Option<&'a str>,
}
impl fmt::Debug for SameFormatProviderHeadersInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("SameFormatProviderHeadersInput")
.field(
"request_header_names",
&self
.headers
.keys()
.map(|name| name.as_str())
.collect::<Vec<_>>(),
)
.field(
"provider_request_body_bytes",
&serde_json::to_vec(self.provider_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field(
"original_request_body_bytes",
&serde_json::to_vec(self.original_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field("has_header_rules", &self.header_rules.is_some())
.field("behavior", &self.behavior)
.field("api_operation", &self.api_operation)
.field("auth_header", &self.auth_header)
.field("has_auth_value", &self.auth_value.is_some())
.field(
"extra_header_names",
&self.extra_headers.keys().collect::<Vec<_>>(),
)
.field("has_kiro_auth_config", &self.kiro_auth_config.is_some())
.field("has_kiro_machine_id", &self.kiro_machine_id.is_some())
.finish()
}
}
pub fn classify_same_format_provider_request_behavior(
transport: &GatewayProviderTransportSnapshot,
params: SameFormatProviderRequestBehaviorParams<'_>,
@@ -721,6 +831,12 @@ pub fn build_same_format_provider_headers(
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
force_identity_accept_encoding(&mut provider_request_headers);
}
let declared_connection_headers =
crate::headers::declared_connection_header_names(input.headers, input.extra_headers);
crate::headers::remove_declared_connection_headers(
&mut provider_request_headers,
&declared_connection_headers,
);
Some(provider_request_headers)
}
+196 -14
View File
@@ -1,3 +1,4 @@
use aether_contracts::redact_url_for_debug;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
@@ -18,7 +19,7 @@ pub struct GatewayProviderTransportSnapshot {
pub key: GatewayProviderTransportKey,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize)]
#[derive(Clone, PartialEq, serde::Serialize)]
pub struct GatewayProviderTransportProvider {
pub id: String,
pub name: String,
@@ -29,13 +30,45 @@ pub struct GatewayProviderTransportProvider {
pub enable_format_conversion: bool,
pub concurrent_limit: Option<i32>,
pub max_retries: Option<i32>,
#[serde(skip_serializing)]
pub proxy: Option<serde_json::Value>,
pub request_timeout_secs: Option<f64>,
pub stream_first_byte_timeout_secs: Option<f64>,
#[serde(skip_serializing)]
pub config: Option<serde_json::Value>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize)]
impl std::fmt::Debug for GatewayProviderTransportProvider {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("GatewayProviderTransportProvider")
.field("id", &self.id)
.field("name", &self.name)
.field("provider_type", &self.provider_type)
.field(
"website",
&self.website.as_deref().map(redact_url_for_debug),
)
.field("is_active", &self.is_active)
.field(
"keep_priority_on_conversion",
&self.keep_priority_on_conversion,
)
.field("enable_format_conversion", &self.enable_format_conversion)
.field("concurrent_limit", &self.concurrent_limit)
.field("max_retries", &self.max_retries)
.field("has_proxy", &self.proxy.is_some())
.field("request_timeout_secs", &self.request_timeout_secs)
.field(
"stream_first_byte_timeout_secs",
&self.stream_first_byte_timeout_secs,
)
.field("has_config", &self.config.is_some())
.finish()
}
}
#[derive(Clone, PartialEq, serde::Serialize)]
pub struct GatewayProviderTransportEndpoint {
pub id: String,
pub provider_id: String,
@@ -44,16 +77,49 @@ pub struct GatewayProviderTransportEndpoint {
pub endpoint_kind: Option<String>,
pub is_active: bool,
pub base_url: String,
#[serde(skip_serializing)]
pub header_rules: Option<serde_json::Value>,
#[serde(skip_serializing)]
pub body_rules: Option<serde_json::Value>,
pub max_retries: Option<i32>,
pub custom_path: Option<String>,
#[serde(skip_serializing)]
pub config: Option<serde_json::Value>,
#[serde(skip_serializing)]
pub format_acceptance_config: Option<serde_json::Value>,
#[serde(skip_serializing)]
pub proxy: Option<serde_json::Value>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize)]
impl std::fmt::Debug for GatewayProviderTransportEndpoint {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("GatewayProviderTransportEndpoint")
.field("id", &self.id)
.field("provider_id", &self.provider_id)
.field("api_format", &self.api_format)
.field("api_family", &self.api_family)
.field("endpoint_kind", &self.endpoint_kind)
.field("is_active", &self.is_active)
.field("base_url", &redact_url_for_debug(&self.base_url))
.field("has_header_rules", &self.header_rules.is_some())
.field("has_body_rules", &self.body_rules.is_some())
.field("max_retries", &self.max_retries)
.field(
"custom_path_len",
&self.custom_path.as_ref().map(String::len),
)
.field("has_config", &self.config.is_some())
.field(
"has_format_acceptance_config",
&self.format_acceptance_config.is_some(),
)
.field("has_proxy", &self.proxy.is_some())
.finish()
}
}
#[derive(Clone, PartialEq, serde::Serialize)]
pub struct GatewayProviderTransportKey {
pub id: String,
pub provider_id: String,
@@ -68,13 +134,42 @@ pub struct GatewayProviderTransportKey {
pub rate_multipliers: Option<serde_json::Value>,
pub global_priority_by_format: Option<serde_json::Value>,
pub expires_at_unix_secs: Option<u64>,
#[serde(skip_serializing)]
pub proxy: Option<serde_json::Value>,
#[serde(skip_serializing)]
pub fingerprint: Option<serde_json::Value>,
#[serde(skip_serializing)]
pub upstream_metadata: Option<serde_json::Value>,
#[serde(skip_serializing)]
pub decrypted_api_key: String,
#[serde(skip_serializing)]
pub decrypted_auth_config: Option<String>,
}
impl std::fmt::Debug for GatewayProviderTransportKey {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("GatewayProviderTransportKey")
.field("id", &self.id)
.field("provider_id", &self.provider_id)
.field("name", &self.name)
.field("auth_type", &self.auth_type)
.field("is_active", &self.is_active)
.field("api_formats", &self.api_formats)
.field("allowed_models", &self.allowed_models)
.field("expires_at_unix_secs", &self.expires_at_unix_secs)
.field("has_proxy", &self.proxy.is_some())
.field("has_fingerprint", &self.fingerprint.is_some())
.field("has_upstream_metadata", &self.upstream_metadata.is_some())
.field("decrypted_api_key", &"[REDACTED]")
.field(
"decrypted_auth_config",
&self.decrypted_auth_config.as_ref().map(|_| "[REDACTED]"),
)
.finish_non_exhaustive()
}
}
#[async_trait]
pub trait ProviderTransportSnapshotSource: Send + Sync {
fn encryption_key(&self) -> Option<&str>;
@@ -335,6 +430,23 @@ mod tests {
)
}
fn seal_bound_provider_credential(
provider_id: &str,
key_id: &str,
field: &str,
plaintext: &str,
) -> String {
let purpose = format!(
"provider-catalog-credential-bound-v2\0provider-id-bytes={}\0{provider_id}\0key-id-bytes={}\0{key_id}\0field={field}",
provider_id.len(),
key_id.len(),
);
let protected = format!("{purpose}\0{plaintext}");
let ciphertext = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &protected)
.expect("bound credential should encrypt");
format!("aether-provider-catalog-credential-v2:aether-runtime-secret-v1:{ciphertext}")
}
#[tokio::test]
async fn reads_decrypted_provider_transport_snapshot() {
let state = read_state();
@@ -411,6 +523,56 @@ mod tests {
);
}
#[tokio::test]
async fn transport_snapshot_debug_and_serialization_exclude_credentials() {
let state = read_state();
let mut snapshot =
read_provider_transport_snapshot(&state, "provider-1", "endpoint-1", "key-1")
.await
.expect("snapshot should read")
.expect("snapshot should exist");
snapshot.provider.proxy = Some(serde_json::json!({"password": "provider-proxy-canary"}));
snapshot.provider.config =
Some(serde_json::json!({"authorization": "provider-config-canary"}));
snapshot.endpoint.header_rules =
Some(serde_json::json!({"authorization": "endpoint-header-canary"}));
snapshot.endpoint.body_rules = Some(serde_json::json!({"token": "endpoint-body-canary"}));
snapshot.endpoint.config = Some(serde_json::json!({"secret": "endpoint-config-canary"}));
snapshot.endpoint.format_acceptance_config =
Some(serde_json::json!({"secret": "endpoint-acceptance-canary"}));
snapshot.endpoint.proxy = Some(serde_json::json!({"password": "endpoint-proxy-canary"}));
snapshot.key.proxy = Some(serde_json::json!({"password": "key-proxy-canary"}));
snapshot.key.fingerprint = Some(serde_json::json!({"cookie": "key-fingerprint-canary"}));
snapshot.key.upstream_metadata =
Some(serde_json::json!({"access_token": "key-metadata-canary"}));
let debug = format!("{snapshot:?}");
let serialized = serde_json::to_string(&snapshot).expect("snapshot should serialize");
for secret in [
"sk-live-openai",
"rt-1",
"provider-proxy-canary",
"provider-config-canary",
"endpoint-header-canary",
"endpoint-body-canary",
"endpoint-config-canary",
"endpoint-acceptance-canary",
"endpoint-proxy-canary",
"key-proxy-canary",
"key-fingerprint-canary",
"key-metadata-canary",
] {
assert!(!debug.contains(secret), "debug leaked {secret}");
assert!(
!serialized.contains(secret),
"serialization leaked {secret}"
);
}
assert!(debug.contains("[REDACTED]"));
assert!(!serialized.contains("decrypted_api_key"));
assert!(!serialized.contains("decrypted_auth_config"));
}
#[tokio::test]
async fn reads_snapshot_when_provider_key_api_key_is_null() {
let mut key = sample_key();
@@ -534,7 +696,7 @@ mod tests {
}
#[tokio::test]
async fn accepts_plaintext_legacy_key_material() {
async fn rejects_plaintext_legacy_key_material() {
let provider = sample_provider();
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-legacy-1".to_string(),
@@ -584,25 +746,45 @@ mod tests {
Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()),
);
let snapshot = read_provider_transport_snapshot(
let error = read_provider_transport_snapshot(
&state,
"provider-1",
"endpoint-legacy-1",
"key-legacy-1",
)
.await
.expect("snapshot read should succeed")
.expect("snapshot should exist");
.expect_err("plaintext credentials must be rejected");
assert_eq!(snapshot.key.decrypted_api_key, "sk-plaintext-openai");
assert_eq!(snapshot.key.decrypted_auth_config, None);
assert!(matches!(error, DataLayerError::UnexpectedValue(message)
if message.contains("provider_api_keys.api_key is not an authenticated ciphertext")));
}
#[test]
fn decrypts_record_bound_v2_credentials_and_rejects_copying() {
let mut key = sample_key();
key.encrypted_api_key = Some(seal_bound_provider_credential(
"provider-1",
"key-1",
"api-key",
"bound-api-key",
));
key.encrypted_auth_config = Some(seal_bound_provider_credential(
"provider-1",
"key-1",
"auth-config",
r#"{"refresh_token":"bound-refresh"}"#,
));
let mapped = map_key(key.clone(), DEVELOPMENT_ENCRYPTION_KEY, &[])
.expect("matching record binding should decrypt");
assert_eq!(mapped.decrypted_api_key, "bound-api-key");
assert_eq!(
snapshot.endpoint.header_rules,
Some(serde_json::json!([
{"action":"set","key":"x-test","value":"1"},
{"action":"set","key":"x-account-id","value":"acc-legacy"}
]))
mapped.decrypted_auth_config.as_deref(),
Some(r#"{"refresh_token":"bound-refresh"}"#)
);
key.id = "key-2".to_string();
assert!(map_key(key, DEVELOPMENT_ENCRYPTION_KEY, &[]).is_err());
}
#[tokio::test]
@@ -8,6 +8,11 @@ use super::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey, GatewayProviderTransportProvider,
};
const PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_FAMILY: &str = "aether-provider-catalog-credential-";
const PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2: &str = "aether-provider-catalog-credential-v2:";
const PROVIDER_CATALOG_CREDENTIAL_PURPOSE_V2: &str = "provider-catalog-credential-bound-v2";
const RUNTIME_SECRET_ENVELOPE_PREFIX: &str = "aether-runtime-secret-v1:";
pub(super) fn map_provider(
provider: StoredProviderCatalogProvider,
) -> GatewayProviderTransportProvider {
@@ -64,6 +69,9 @@ pub(super) fn map_key(
encryption_key,
fallback_encryption_keys,
ciphertext,
&key.provider_id,
&key.id,
"api-key",
"provider_api_keys.api_key",
)
})
@@ -79,6 +87,9 @@ pub(super) fn map_key(
encryption_key,
fallback_encryption_keys,
ciphertext,
&key.provider_id,
&key.id,
"auth-config",
"provider_api_keys.auth_config",
)
})
@@ -125,12 +136,78 @@ fn decrypt_secret(
encryption_key: &str,
fallback_encryption_keys: &[String],
ciphertext: &str,
provider_id: &str,
key_id: &str,
field: &str,
field_name: &str,
) -> Result<String, DataLayerError> {
if should_use_plaintext_secret(ciphertext, field_name) {
return Ok(ciphertext.trim().to_string());
if let Some(runtime_envelope) = ciphertext.strip_prefix(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_V2)
{
let inner_ciphertext = runtime_envelope
.strip_prefix(RUNTIME_SECRET_ENVELOPE_PREFIX)
.ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"{field_name} has an invalid provider catalog credential envelope"
))
})?;
let protected = decrypt_fernet_with_fallbacks(
encryption_key,
fallback_encryption_keys,
inner_ciphertext,
field_name,
)?;
let purpose = provider_catalog_credential_purpose(provider_id, key_id, field);
return protected
.strip_prefix(&purpose)
.and_then(|value| value.strip_prefix('\0'))
.map(ToOwned::to_owned)
.ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"{field_name} provider catalog credential authentication failed"
))
});
}
if ciphertext.starts_with(PROVIDER_CATALOG_CREDENTIAL_ENVELOPE_FAMILY)
|| ciphertext.starts_with("aether-")
{
return Err(DataLayerError::UnexpectedValue(format!(
"{field_name} has an unsupported or incorrectly bound Aether secret envelope"
)));
}
if !looks_like_python_fernet_ciphertext(ciphertext) {
return Err(DataLayerError::UnexpectedValue(format!(
"{field_name} is not an authenticated ciphertext"
)));
}
let plaintext = decrypt_fernet_with_fallbacks(
encryption_key,
fallback_encryption_keys,
ciphertext,
field_name,
)?;
if plaintext.contains('\0') {
return Err(DataLayerError::UnexpectedValue(format!(
"{field_name} legacy ciphertext contains reserved framing"
)));
}
Ok(plaintext)
}
fn provider_catalog_credential_purpose(provider_id: &str, key_id: &str, field: &str) -> String {
format!(
"{PROVIDER_CATALOG_CREDENTIAL_PURPOSE_V2}\0provider-id-bytes={}\0{provider_id}\0key-id-bytes={}\0{key_id}\0field={field}",
provider_id.len(),
key_id.len(),
)
}
fn decrypt_fernet_with_fallbacks(
encryption_key: &str,
fallback_encryption_keys: &[String],
ciphertext: &str,
field_name: &str,
) -> Result<String, DataLayerError> {
match decrypt_python_fernet_ciphertext(encryption_key, ciphertext) {
Ok(value) => Ok(value),
Err(error) => {
@@ -166,29 +243,6 @@ pub(super) fn fallback_encryption_keys(primary_encryption_key: &str) -> Vec<Stri
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,
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::fmt;
use serde_json::Value;
@@ -16,7 +17,7 @@ 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)]
#[derive(Clone, Copy)]
pub struct StandardProviderRequestHeadersInput<'a> {
pub transport: &'a GatewayProviderTransportSnapshot,
pub provider_api_format: &'a str,
@@ -31,13 +32,63 @@ pub struct StandardProviderRequestHeadersInput<'a> {
pub upstream_is_stream: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
impl fmt::Debug for StandardProviderRequestHeadersInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("StandardProviderRequestHeadersInput")
.field("transport", &self.transport)
.field("provider_api_format", &self.provider_api_format)
.field("same_format", &self.same_format)
.field(
"request_header_names",
&self
.headers
.keys()
.map(|name| name.as_str())
.collect::<Vec<_>>(),
)
.field("auth_header", &self.auth_header)
.field("has_auth_value", &(!self.auth_value.is_empty()))
.field(
"extra_header_names",
&self.extra_headers.keys().collect::<Vec<_>>(),
)
.field("has_header_rules", &self.header_rules.is_some())
.field(
"provider_request_body_bytes",
&serde_json::to_vec(self.provider_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field(
"original_request_body_bytes",
&serde_json::to_vec(self.original_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field("upstream_is_stream", &self.upstream_is_stream)
.finish()
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct StandardProviderRequestHeaders {
pub headers: BTreeMap<String, String>,
pub auth_header: String,
pub auth_value: String,
}
impl fmt::Debug for StandardProviderRequestHeaders {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("StandardProviderRequestHeaders")
.field("header_names", &self.headers.keys().collect::<Vec<_>>())
.field("auth_header", &self.auth_header)
.field("has_auth_value", &(!self.auth_value.is_empty()))
.finish()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StandardPlanFallbackAcceptPolicy {
None,
@@ -47,7 +98,7 @@ pub enum StandardPlanFallbackAcceptPolicy {
ProviderEventStreamIfMissing,
}
#[derive(Debug)]
#[derive(Clone)]
pub struct StandardPlanFallbackHeadersInput<'a> {
pub request_headers: &'a http::HeaderMap,
pub existing_provider_request_headers: BTreeMap<String, String>,
@@ -62,6 +113,44 @@ pub struct StandardPlanFallbackHeadersInput<'a> {
pub accept_policy: StandardPlanFallbackAcceptPolicy,
}
impl fmt::Debug for StandardPlanFallbackHeadersInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("StandardPlanFallbackHeadersInput")
.field(
"request_header_names",
&self
.request_headers
.keys()
.map(|name| name.as_str())
.collect::<Vec<_>>(),
)
.field(
"existing_provider_header_names",
&self
.existing_provider_request_headers
.keys()
.collect::<Vec<_>>(),
)
.field("auth_header", &self.auth_header)
.field("has_auth_value", &self.auth_value.is_some())
.field(
"extra_header_names",
&self.extra_headers.keys().collect::<Vec<_>>(),
)
.field("content_type", &self.content_type)
.field("provider_api_format", &self.provider_api_format)
.field("client_api_format", &self.client_api_format)
.field("upstream_is_stream", &self.upstream_is_stream)
.field(
"build_from_request_when_empty",
&self.build_from_request_when_empty,
)
.field("accept_policy", &self.accept_policy)
.finish()
}
}
pub fn build_standard_plan_fallback_openai_chat_url(
upstream_base_url: &str,
request_query: Option<&str>,
@@ -152,6 +241,12 @@ pub fn build_standard_plan_fallback_headers(
force_identity_accept_encoding(&mut headers);
}
let declared_connection_headers = crate::headers::declared_connection_header_names(
input.request_headers,
input.extra_headers,
);
crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers);
headers
}
@@ -301,6 +396,10 @@ pub fn build_standard_provider_request_headers(
force_identity_accept_encoding(&mut headers);
}
let declared_connection_headers =
crate::headers::declared_connection_header_names(input.headers, input.extra_headers);
crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers);
Some(StandardProviderRequestHeaders {
headers,
auth_header,
+101 -7
View File
@@ -5,6 +5,42 @@ use url::Url;
pub(crate) const GATEWAY_CREDENTIAL_QUERY_KEYS: &[&str] = &["key"];
pub(crate) fn encode_url_path_segment(value: &str) -> String {
const HEX: &[u8; 16] = b"0123456789ABCDEF";
let mut encoded = String::with_capacity(value.len());
for byte in value.bytes() {
if byte.is_ascii_alphanumeric()
|| matches!(
byte,
b'-' | b'.'
| b'_'
| b'~'
| b'!'
| b'$'
| b'&'
| b'\''
| b'('
| b')'
| b'*'
| b'+'
| b','
| b';'
| b'='
| b':'
| b'@'
)
{
encoded.push(char::from(byte));
} else {
encoded.push('%');
encoded.push(char::from(HEX[usize::from(byte >> 4)]));
encoded.push(char::from(HEX[usize::from(byte & 0x0f)]));
}
}
encoded
}
pub(crate) fn strip_gateway_credential_query_parameters(query: Option<&str>) -> Option<String> {
let query = query.map(str::trim).filter(|value| !value.is_empty())?;
let mut serializer = form_urlencoded::Serializer::new(String::new());
@@ -161,6 +197,7 @@ pub fn build_gemini_content_url(
if trimmed_base_url.is_empty() || trimmed_model.is_empty() {
return None;
}
let encoded_model = encode_url_path_segment(trimmed_model);
let operation = if stream {
"streamGenerateContent"
@@ -168,12 +205,12 @@ pub fn build_gemini_content_url(
"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}")
format!("{trimmed_base_url}/models/{encoded_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}")
format!("{trimmed_base_url}/v1beta/models/{encoded_model}:{operation}")
};
append_merged_query(&mut url, base_query, None, query, &["key"]);
Some(url)
@@ -221,13 +258,14 @@ pub fn build_gemini_video_predict_long_running_url(
if trimmed_base_url.is_empty() || trimmed_model.is_empty() {
return None;
}
let encoded_model = encode_url_path_segment(trimmed_model);
let mut url = if trimmed_base_url.ends_with("/v1") || trimmed_base_url.ends_with("/v1beta") {
format!("{trimmed_base_url}/models/{trimmed_model}:predictLongRunning")
format!("{trimmed_base_url}/models/{encoded_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")
format!("{trimmed_base_url}/v1beta/models/{encoded_model}:predictLongRunning")
};
append_merged_query(&mut url, base_query, None, query, &["key"]);
Some(url)
@@ -412,9 +450,27 @@ fn bigmodel_coding_models_base_is_supported(base_url: &str) -> bool {
fn looks_like_vertex_ai_host(host: &str) -> bool {
const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com";
let host = host.trim().to_ascii_lowercase();
host == VERTEX_AI_HOST
|| host.ends_with(&format!(".{VERTEX_AI_HOST}"))
|| host.ends_with(&format!("-{VERTEX_AI_HOST}"))
|| host
.strip_suffix(&format!("-{VERTEX_AI_HOST}"))
.is_some_and(is_vertex_region_label)
}
fn is_vertex_region_label(value: &str) -> bool {
!value.is_empty()
&& value.len() <= 63
&& value
.as_bytes()
.first()
.is_some_and(u8::is_ascii_alphanumeric)
&& value
.as_bytes()
.last()
.is_some_and(u8::is_ascii_alphanumeric)
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-')
}
fn split_path_query(path: &str) -> (&str, Option<&str>) {
@@ -504,7 +560,8 @@ mod tests {
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,
build_openai_search_url, build_passthrough_path_url, encode_url_path_segment,
normalize_gemini_content_action_path,
};
#[test]
@@ -868,4 +925,41 @@ mod tests {
)
);
}
#[test]
fn gemini_model_names_cannot_inject_path_query_or_fragment_components() {
let model = "model/../../admin?key=attacker#fragment";
assert_eq!(
build_gemini_content_url(
"https://generativelanguage.googleapis.com/v1beta",
model,
false,
None,
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/v1beta/models/model%2F..%2F..%2Fadmin%3Fkey=attacker%23fragment:generateContent"
)
);
assert_eq!(
build_gemini_video_predict_long_running_url(
"https://generativelanguage.googleapis.com/v1beta",
model,
None,
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/v1beta/models/model%2F..%2F..%2Fadmin%3Fkey=attacker%23fragment:predictLongRunning"
)
);
}
#[test]
fn dynamic_path_segment_encoding_keeps_raw_values_in_one_segment() {
assert_eq!(
encode_url_path_segment("gemini+2.5@preview~/model%2Fraw"),
"gemini+2.5@preview~%2Fmodel%252Fraw"
);
assert_eq!(encode_url_path_segment(".."), "..");
}
}
@@ -1,22 +1,19 @@
use std::collections::BTreeMap;
use aether_crypto::{rsa_pkcs1_sha256_sign, RsaPkcs1Sha256Error};
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::{Digest, Sha256};
use url::form_urlencoded;
use url::{form_urlencoded, Url};
use super::super::oauth_refresh::{
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthHttpRequest, LocalOAuthRefreshAdapter,
LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth,
};
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::context::is_valid_vertex_region;
pub const VERTEX_API_KEY_QUERY_PARAM: &str = "key";
pub const VERTEX_SERVICE_ACCOUNT_AUTH_HEADER: &str = "authorization";
@@ -25,13 +22,23 @@ 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)]
#[derive(Clone, PartialEq, Eq)]
pub struct VertexApiKeyQueryAuth {
pub name: &'static str,
pub value: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
impl std::fmt::Debug for VertexApiKeyQueryAuth {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("VertexApiKeyQueryAuth")
.field("name", &self.name)
.field("value", &"[REDACTED]")
.finish()
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct VertexServiceAccountAuthConfig {
pub client_email: String,
pub private_key: String,
@@ -41,6 +48,20 @@ pub struct VertexServiceAccountAuthConfig {
pub model_regions: BTreeMap<String, String>,
}
impl std::fmt::Debug for VertexServiceAccountAuthConfig {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("VertexServiceAccountAuthConfig")
.field("client_email", &self.client_email)
.field("private_key", &"[REDACTED]")
.field("project_id", &self.project_id)
.field("token_uri", &"[REDACTED]")
.field("region", &self.region)
.field("model_regions", &self.model_regions)
.finish()
}
}
pub fn resolve_local_vertex_api_key_query_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<VertexApiKeyQueryAuth> {
@@ -101,9 +122,8 @@ fn parse_vertex_service_account_auth_config_value(
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 token_uri = resolve_vertex_service_account_token_uri(value.get("token_uri"))?;
let region = json_string(value.get("region")).filter(|value| is_valid_vertex_region(value));
let model_regions = value
.get("model_regions")
.and_then(Value::as_object)
@@ -113,7 +133,7 @@ fn parse_vertex_service_account_auth_config_value(
.filter_map(|(model, region)| {
let model = model.trim();
let region = region.as_str()?.trim();
(!model.is_empty() && !region.is_empty())
(!model.is_empty() && is_valid_vertex_region(region))
.then(|| (model.to_string(), region.to_string()))
})
.collect::<BTreeMap<_, _>>()
@@ -130,6 +150,31 @@ fn parse_vertex_service_account_auth_config_value(
})
}
fn resolve_vertex_service_account_token_uri(value: Option<&Value>) -> Option<String> {
let Some(value) = value else {
return Some(GOOGLE_OAUTH_TOKEN_URL.to_string());
};
let raw = value.as_str()?.trim();
if raw.is_empty() {
return None;
}
let parsed = Url::parse(raw).ok()?;
if parsed.scheme() != "https"
|| !parsed.username().is_empty()
|| parsed.password().is_some()
|| !parsed
.host_str()
.is_some_and(|host| host.eq_ignore_ascii_case("oauth2.googleapis.com"))
|| parsed.port().is_some()
|| parsed.path() != "/token"
|| parsed.query().is_some()
|| parsed.fragment().is_some()
{
return None;
}
Some(GOOGLE_OAUTH_TOKEN_URL.to_string())
}
fn json_string(value: Option<&Value>) -> Option<String> {
value
.and_then(Value::as_str)
@@ -352,29 +397,17 @@ pub fn build_vertex_service_account_assertion(
})?,
);
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}"
),
let signature = rsa_pkcs1_sha256_sign(auth_config.private_key.as_bytes(), message.as_bytes())
.map_err(|error| LocalOAuthRefreshError::InvalidResponse {
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
message: match error {
RsaPkcs1Sha256Error::InvalidPrivateKey => {
"vertex service account private_key parse failed".to_string()
}
}),
}
_ => "vertex service account signing failed".to_string(),
},
})?;
Ok(format!("{message}.{}", URL_SAFE_NO_PAD.encode(signature)))
}
fn service_account_token_expires_soon(expires_at_unix_secs: Option<u64>) -> bool {
@@ -387,7 +420,7 @@ fn service_account_token_expires_soon(expires_at_unix_secs: Option<u64>) -> bool
}
fn body_excerpt(value: &str) -> String {
value.chars().take(500).collect()
aether_oauth::core::redacted_oauth_error_body_excerpt(value)
}
#[cfg(test)]
@@ -399,19 +432,46 @@ mod tests {
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use rsa::pkcs1::{EncodeRsaPrivateKey, LineEnding};
use rsa::rand_core::OsRng;
use rsa::RsaPrivateKey;
use aws_lc_rs::encoding::{AsDer, Pkcs8V1Der};
use aws_lc_rs::rsa::{KeyPair as AwsRsaKeyPair, KeySize};
use aws_lc_rs::signature::{KeyPair as _, UnparsedPublicKey, RSA_PKCS1_2048_8192_SHA256};
use base64::engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD};
use base64::Engine as _;
use serde_json::{json, Value};
use super::{
decode_vertex_service_account_private_key, parse_vertex_service_account_auth_config,
build_vertex_service_account_assertion, parse_vertex_service_account_auth_config,
resolve_local_vertex_api_key_query_auth,
supports_local_vertex_service_account_auth_resolution,
vertex_service_account_credential_fingerprint, VertexServiceAccountRefreshAdapter,
vertex_service_account_credential_fingerprint, VertexApiKeyQueryAuth,
VertexServiceAccountAuthConfig, VertexServiceAccountRefreshAdapter, GOOGLE_OAUTH_TOKEN_URL,
VERTEX_API_KEY_QUERY_PARAM, VERTEX_SERVICE_ACCOUNT_AUTH_HEADER,
VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
};
#[test]
fn vertex_auth_debug_output_redacts_api_keys_and_private_keys() {
let query_auth = VertexApiKeyQueryAuth {
name: VERTEX_API_KEY_QUERY_PARAM,
value: "vertex-api-key-canary".to_string(),
};
let service_account = VertexServiceAccountAuthConfig {
client_email: "[email protected]".to_string(),
private_key: "vertex-private-key-canary".to_string(),
project_id: "project-1".to_string(),
token_uri: GOOGLE_OAUTH_TOKEN_URL.to_string(),
region: None,
model_regions: std::collections::BTreeMap::new(),
};
let query_debug = format!("{query_auth:?}");
assert!(!query_debug.contains("vertex-api-key-canary"));
assert!(query_debug.contains("[REDACTED]"));
let service_account_debug = format!("{service_account:?}");
assert!(!service_account_debug.contains("vertex-private-key-canary"));
assert!(service_account_debug.contains("[REDACTED]"));
}
fn sample_transport() -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
@@ -469,6 +529,32 @@ mod tests {
}
}
fn read_der_tlv<'a>(input: &mut &'a [u8], expected_tag: u8) -> &'a [u8] {
assert_eq!(input.first().copied(), Some(expected_tag));
let length_byte = input[1];
let (header_len, value_len) = if length_byte & 0x80 == 0 {
(2, usize::from(length_byte))
} else {
let length_bytes = usize::from(length_byte & 0x7f);
let value_len = input[2..2 + length_bytes]
.iter()
.fold(0usize, |value, byte| (value << 8) | usize::from(*byte));
(2 + length_bytes, value_len)
};
let end = header_len + value_len;
let value = &input[header_len..end];
*input = &input[end..];
value
}
fn pkcs1_private_key_from_pkcs8(pkcs8: &[u8]) -> Vec<u8> {
let mut input = pkcs8;
let mut sequence = read_der_tlv(&mut input, 0x30);
let _version = read_der_tlv(&mut sequence, 0x02);
let _algorithm = read_der_tlv(&mut sequence, 0x30);
read_der_tlv(&mut sequence, 0x04).to_vec()
}
fn sample_service_account_transport(private_key: &str) -> GatewayProviderTransportSnapshot {
let mut transport = sample_transport();
transport.key.auth_type = "service_account".to_string();
@@ -570,6 +656,7 @@ mod tests {
assert_eq!(config.client_email, "[email protected]");
assert_eq!(config.project_id, "demo-project");
assert_eq!(config.token_uri, GOOGLE_OAUTH_TOKEN_URL);
assert_eq!(config.region.as_deref(), Some("global"));
assert_eq!(
config
@@ -580,6 +667,69 @@ mod tests {
);
}
#[test]
fn service_account_regions_reject_url_syntax() {
let raw = r#"{
"client_email":"[email protected]",
"private_key":"TEST-PRIVATE-KEY",
"project_id":"demo-project",
"region":"attacker.example/",
"model_regions":{
"gemini-2.0-flash":"attacker.example/",
"gemini-2.5-pro":"us-central1"
}
}"#;
let config = parse_vertex_service_account_auth_config(Some(raw))
.expect("service account config should parse");
assert!(config.region.is_none());
assert!(!config.model_regions.contains_key("gemini-2.0-flash"));
assert_eq!(
config
.model_regions
.get("gemini-2.5-pro")
.map(String::as_str),
Some("us-central1")
);
}
#[test]
fn service_account_token_uri_is_limited_to_google_oauth_endpoint() {
let config_with_token_uri = |token_uri: Value| {
serde_json::json!({
"client_email": "[email protected]",
"private_key": "TEST-PRIVATE-KEY",
"project_id": "demo-project",
"token_uri": token_uri,
})
.to_string()
};
let official = parse_vertex_service_account_auth_config(Some(&config_with_token_uri(
Value::String(GOOGLE_OAUTH_TOKEN_URL.to_string()),
)))
.expect("official Google OAuth token URI should be accepted");
assert_eq!(official.token_uri, GOOGLE_OAUTH_TOKEN_URL);
for token_uri in [
Value::String("http://oauth2.googleapis.com/token".to_string()),
Value::String("https://127.0.0.1/token".to_string()),
Value::String("https://oauth2.googleapis.com.evil.example/token".to_string()),
Value::String("https://[email protected]/token".to_string()),
Value::String("https://oauth2.googleapis.com:8443/token".to_string()),
Value::String("https://oauth2.googleapis.com/token/../metadata".to_string()),
Value::String("https://oauth2.googleapis.com/token?target=metadata".to_string()),
Value::String("https://oauth2.googleapis.com/token#fragment".to_string()),
Value::String(String::new()),
Value::Null,
] {
let raw = config_with_token_uri(token_uri.clone());
assert!(
parse_vertex_service_account_auth_config(Some(&raw)).is_none(),
"token URI should be rejected: {token_uri}"
);
}
}
#[test]
fn supports_vertex_service_account_auth_resolution() {
let mut transport = sample_transport();
@@ -600,14 +750,35 @@ mod tests {
}
#[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");
fn signs_with_2048_bit_pkcs1_service_account_private_key() {
let key_pair = AwsRsaKeyPair::generate(KeySize::Rsa2048)
.expect("2048-bit test RSA private key should generate");
let pkcs8 = AsDer::<Pkcs8V1Der<'static>>::as_der(&key_pair)
.expect("test RSA private key should encode as PKCS#8");
let pkcs1 = pkcs1_private_key_from_pkcs8(pkcs8.as_ref());
let private_key = format!(
"-----BEGIN RSA PRIVATE KEY-----\n{}\n-----END RSA PRIVATE KEY-----",
STANDARD.encode(pkcs1)
);
let auth_config = parse_vertex_service_account_auth_config(Some(
&json!({
"client_email": "[email protected]",
"private_key": private_key,
"project_id": "demo-project"
})
.to_string(),
))
.expect("service account config should parse");
let assertion = build_vertex_service_account_assertion(&auth_config, 1_700_000_000)
.expect("PKCS#1 private key should sign");
let parts = assertion.split('.').collect::<Vec<_>>();
assert_eq!(parts.len(), 3);
let message = format!("{}.{}", parts[0], parts[1]);
let signature = URL_SAFE_NO_PAD
.decode(parts[2])
.expect("JWT signature should decode");
UnparsedPublicKey::new(&RSA_PKCS1_2048_8192_SHA256, key_pair.public_key().as_ref())
.verify(message.as_bytes(), &signature)
.expect("AWS-LC signature should verify");
}
}
@@ -5,6 +5,26 @@ use super::super::snapshot::GatewayProviderTransportSnapshot;
const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com";
/// Vertex regions are interpolated into regional service hostnames and path
/// segments. Keep them to one DNS label so imported credential metadata cannot
/// redirect bearer-token requests to another origin.
pub fn is_valid_vertex_region(value: &str) -> bool {
let value = value.trim();
!value.is_empty()
&& value.len() <= 63
&& value
.as_bytes()
.first()
.is_some_and(u8::is_ascii_alphanumeric)
&& value
.as_bytes()
.last()
.is_some_and(u8::is_ascii_alphanumeric)
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-')
}
pub fn looks_like_vertex_ai_host(base_url: &str) -> bool {
let trimmed = base_url.trim();
if trimmed.is_empty() {
@@ -14,6 +34,20 @@ pub fn looks_like_vertex_ai_host(base_url: &str) -> bool {
let Ok(parsed) = Url::parse(trimmed) else {
return false;
};
// A service-account bearer token must never be sent over plaintext HTTP
// or to a URL carrying userinfo/alternate ports. Host matching alone is
// insufficient because an imported endpoint can still select those URL
// forms.
if parsed.scheme() != "https"
|| !parsed.username().is_empty()
|| parsed.password().is_some()
|| parsed.port().is_some()
{
return false;
}
if parsed.query().is_some() || parsed.fragment().is_some() {
return false;
}
let Some(host) = parsed
.host_str()
.map(|value| value.trim().to_ascii_lowercase())
@@ -22,8 +56,9 @@ pub fn looks_like_vertex_ai_host(base_url: &str) -> bool {
};
host == VERTEX_AI_HOST
|| host.ends_with(&format!(".{VERTEX_AI_HOST}"))
|| host.ends_with(&format!("-{VERTEX_AI_HOST}"))
|| host
.strip_suffix(&format!("-{VERTEX_AI_HOST}"))
.is_some_and(is_valid_vertex_region)
}
pub fn is_vertex_api_key_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool {
@@ -101,8 +136,9 @@ fn looks_like_vertex_openai_compat_base(base_url: &str) -> bool {
#[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,
is_valid_vertex_region, 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,
@@ -174,9 +210,41 @@ mod tests {
assert!(looks_like_vertex_ai_host(
"https://us-central1-aiplatform.googleapis.com"
));
assert!(!looks_like_vertex_ai_host(
"https://foo.bar-aiplatform.googleapis.com"
));
assert!(!looks_like_vertex_ai_host(
"https://us-central1-aiplatform.googleapis.com?token=secret"
));
assert!(!looks_like_vertex_ai_host(
"https://us-central1-aiplatform.googleapis.com."
));
assert!(!looks_like_vertex_ai_host(
"http://us-central1-aiplatform.googleapis.com"
));
assert!(!looks_like_vertex_ai_host(
"https://[email protected]"
));
assert!(!looks_like_vertex_ai_host(
"https://us-central1-aiplatform.googleapis.com:8443"
));
assert!(!looks_like_vertex_ai_host("https://example.com"));
}
#[test]
fn rejects_vertex_region_url_syntax() {
for value in [
"attacker.example/",
"us-central1?x=1",
"us.central1",
"-bad",
] {
assert!(!is_valid_vertex_region(value));
}
assert!(is_valid_vertex_region("us-central1"));
assert!(is_valid_vertex_region("global"));
}
#[test]
fn infers_vertex_api_key_context_for_custom_aiplatform_transport() {
assert!(is_vertex_api_key_transport_context(&sample_transport()));
@@ -11,8 +11,9 @@ pub use auth::{
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,
is_valid_vertex_region, 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,
@@ -2,8 +2,9 @@ use std::collections::BTreeMap;
use url::form_urlencoded;
use super::super::url::build_passthrough_path_url;
use super::super::url::{build_passthrough_path_url, encode_url_path_segment};
use super::auth::VertexServiceAccountAuthConfig;
use super::context::is_valid_vertex_region;
pub const VERTEX_API_KEY_BASE_URL: &str = "https://aiplatform.googleapis.com";
@@ -85,7 +86,8 @@ fn build_vertex_api_key_google_model_url(
return None;
}
let path = format!("/v1/publishers/google/models/{trimmed_model}:{trimmed_action}");
let encoded_model = encode_url_path_segment(trimmed_model);
let path = format!("/v1/publishers/google/models/{encoded_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(), &[])
}
@@ -110,8 +112,14 @@ fn build_vertex_service_account_google_model_url(
} else {
format!("https://{region}-aiplatform.googleapis.com")
};
// URL parsers normalize dot-only segments even when the dots are percent-encoded.
if matches!(project_id, "." | "..") {
return None;
}
let encoded_project_id = encode_url_path_segment(project_id);
let encoded_model = encode_url_path_segment(trimmed_model);
let path = format!(
"/v1/projects/{project_id}/locations/{region}/publishers/google/models/{trimmed_model}:{trimmed_action}"
"/v1/projects/{encoded_project_id}/locations/{region}/publishers/google/models/{encoded_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(), &[])
@@ -127,7 +135,7 @@ pub fn resolve_vertex_service_account_region(
.get(trimmed_model)
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|value| is_valid_vertex_region(value))
{
return region.to_string();
}
@@ -138,7 +146,7 @@ pub fn resolve_vertex_service_account_region(
.region
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|value| is_valid_vertex_region(value))
{
return region.to_string();
}
@@ -236,7 +244,7 @@ mod tests {
use super::{
build_vertex_api_key_gemini_content_url, build_vertex_api_key_imagen_content_url,
build_vertex_service_account_gemini_content_url,
build_vertex_service_account_gemini_content_url, resolve_vertex_service_account_region,
};
use crate::vertex::VertexServiceAccountAuthConfig;
@@ -324,4 +332,80 @@ mod tests {
)
);
}
#[test]
fn service_account_region_override_cannot_escape_vertex_origin() {
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("attacker.example/".to_string()),
model_regions: BTreeMap::from([(
"custom-model".to_string(),
"attacker.example/".to_string(),
)]),
};
assert_eq!(
resolve_vertex_service_account_region("custom-model", &auth_config),
"global"
);
let url = build_vertex_service_account_gemini_content_url(
"custom-model",
false,
&auth_config,
None,
)
.expect("service account URL should be built");
assert!(url.starts_with("https://aiplatform.googleapis.com/"));
assert!(!url.contains("attacker.example"));
}
#[test]
fn vertex_resource_components_cannot_rewrite_the_request_path() {
let auth_config = VertexServiceAccountAuthConfig {
client_email: "[email protected]".to_string(),
private_key: "not-used".to_string(),
project_id: "project/../victim?key=attacker".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(
"model/../../admin#fragment",
false,
&auth_config,
None,
)
.as_deref(),
Some(
"https://aiplatform.googleapis.com/v1/projects/project%2F..%2Fvictim%3Fkey=attacker/locations/global/publishers/google/models/model%2F..%2F..%2Fadmin%23fragment:generateContent"
)
);
}
#[test]
fn vertex_rejects_dot_only_project_path_segments() {
let auth_config = VertexServiceAccountAuthConfig {
client_email: "[email protected]".to_string(),
private_key: "not-used".to_string(),
project_id: "..".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-2.5-pro",
false,
&auth_config,
None,
),
None
);
}
}
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::fmt;
use aether_data_contracts::repository::video_tasks::StoredVideoTask;
use aether_video_tasks_core::{
@@ -29,7 +30,7 @@ pub enum ProviderVideoCreateFamily {
Gemini,
}
#[derive(Debug, Clone, Copy)]
#[derive(Clone, Copy)]
pub struct ProviderVideoCreateHeadersInput<'a> {
pub headers: &'a http::HeaderMap,
pub auth_header: &'a str,
@@ -39,6 +40,37 @@ pub struct ProviderVideoCreateHeadersInput<'a> {
pub original_request_body: &'a Value,
}
impl fmt::Debug for ProviderVideoCreateHeadersInput<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ProviderVideoCreateHeadersInput")
.field(
"request_header_names",
&self
.headers
.keys()
.map(|name| name.as_str())
.collect::<Vec<_>>(),
)
.field("auth_header", &self.auth_header)
.field("has_auth_value", &(!self.auth_value.is_empty()))
.field("has_header_rules", &self.header_rules.is_some())
.field(
"provider_request_body_bytes",
&serde_json::to_vec(self.provider_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.field(
"original_request_body_bytes",
&serde_json::to_vec(self.original_request_body)
.ok()
.map(|bytes| bytes.len()),
)
.finish()
}
}
#[async_trait]
pub trait VideoTaskTransportSnapshotLookup: Send + Sync {
async fn read_video_task_provider_transport_snapshot(
@@ -210,6 +242,12 @@ pub fn build_video_create_headers(
) {
return None;
}
let declared_connection_headers =
super::headers::declared_connection_header_names(input.headers, &BTreeMap::new());
super::headers::remove_declared_connection_headers(
&mut provider_request_headers,
&declared_connection_headers,
);
Some(provider_request_headers)
}
@@ -3,6 +3,7 @@ use std::collections::BTreeMap;
use serde_json::{json, Value};
use uuid::Uuid;
use super::headers::{declared_connection_header_names, remove_declared_connection_headers};
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,
@@ -190,13 +191,16 @@ pub fn build_windsurf_cascade_headers(
auth_value: &str,
_upstream_is_stream: bool,
) -> Option<BTreeMap<String, String>> {
let declared_connection_headers = declared_connection_header_names(headers, &BTreeMap::new());
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) {
if should_skip_upstream_passthrough_header(&key)
|| declared_connection_headers.contains(&key)
{
continue;
}
let value = value.trim();
@@ -234,6 +238,7 @@ pub fn build_windsurf_cascade_headers(
if !auth_header.is_empty() {
out.insert(auth_header, auth_value.trim().to_string());
}
remove_declared_connection_headers(&mut out, &declared_connection_headers);
out.remove("content-length");
Some(out)
}
@@ -23,7 +23,7 @@ OUTPUT RULES:
Violating these rules will produce broken output for the end user. Stay in chat-API mode at all times."#;
#[derive(Debug, Clone, PartialEq, Eq)]
#[derive(Clone, PartialEq, Eq)]
pub struct CascadeStep {
pub step_type: u64,
pub status: u64,
@@ -36,12 +36,44 @@ pub struct CascadeStep {
pub usage: Option<CascadeUsage>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
impl fmt::Debug for CascadeStep {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CascadeStep")
.field("step_type", &self.step_type)
.field("status", &self.status)
.field("text_len", &self.text.len())
.field("response_text_len", &self.response_text.len())
.field("modified_text_len", &self.modified_text.len())
.field("thinking_len", &self.thinking.len())
.field("error_text_len", &self.error_text.len())
.field("has_native_tool", &self.native_tool.is_some())
.field("usage", &self.usage)
.finish()
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct CascadeNativeToolStep {
pub kind: String,
pub arguments: Value,
}
impl fmt::Debug for CascadeNativeToolStep {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CascadeNativeToolStep")
.field("kind", &self.kind)
.field(
"arguments_bytes",
&serde_json::to_vec(&self.arguments)
.ok()
.map(|bytes| bytes.len()),
)
.finish()
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct CascadeUsage {
pub input_tokens: u64,
@@ -60,13 +92,23 @@ impl CascadeUsage {
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[derive(Clone, PartialEq, Eq)]
pub struct CascadeImage {
pub base64_data: String,
pub mime_type: String,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
impl fmt::Debug for CascadeImage {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CascadeImage")
.field("base64_data_len", &self.base64_data.len())
.field("mime_type", &self.mime_type)
.finish()
}
}
#[derive(Clone, Default, PartialEq, Eq)]
pub struct SendCascadeMessageOptions {
pub tool_preamble: Option<String>,
pub images: Vec<CascadeImage>,
@@ -75,6 +117,34 @@ pub struct SendCascadeMessageOptions {
pub native_allowlist: Vec<String>,
}
impl fmt::Debug for SendCascadeMessageOptions {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("SendCascadeMessageOptions")
.field(
"tool_preamble_len",
&self.tool_preamble.as_ref().map(String::len),
)
.field("image_count", &self.images.len())
.field(
"image_bytes",
&self
.images
.iter()
.map(|image| image.base64_data.len())
.sum::<usize>(),
)
.field("additional_steps_count", &self.additional_steps.len())
.field(
"additional_steps_bytes",
&self.additional_steps.iter().map(Vec::len).sum::<usize>(),
)
.field("native_mode", &self.native_mode)
.field("native_allowlist_count", &self.native_allowlist.len())
.finish()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CascadeBuildError {
message: String,
@@ -8,19 +8,42 @@ pub enum WireType {
Fixed32 = 5,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[derive(Clone, PartialEq, Eq)]
pub enum FieldValue {
Varint(u64),
Bytes(Vec<u8>),
}
#[derive(Debug, Clone, PartialEq, Eq)]
impl fmt::Debug for FieldValue {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Varint(value) => formatter.debug_tuple("Varint").field(value).finish(),
Self::Bytes(bytes) => formatter
.debug_struct("Bytes")
.field("len", &bytes.len())
.finish(),
}
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct Field {
pub number: u32,
pub wire_type: WireType,
pub value: FieldValue,
}
impl fmt::Debug for Field {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("Field")
.field("number", &self.number)
.field("wire_type", &self.wire_type)
.field("value", &self.value)
.finish()
}
}
impl Field {
pub fn bytes(&self) -> &[u8] {
match &self.value {
@@ -138,7 +161,7 @@ pub fn parse_fields(buf: &[u8]) -> Result<Vec<Field>, ProtoError> {
other => {
return Err(ProtoError::new(format!(
"unknown wire type {other} at offset {pos}"
)))
)));
}
};
let value = match wire_type {