mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 08:27:46 +08:00
feat(security): harden gateway boundaries and usage policies
Consolidate subscription usage policy enforcement, privacy-safe persistence, and gateway security hardening into one reviewable change. Includes bounded HTTP and execution envelopes, header and protocol guards, DNS and relay validation, authentication and secret projection hardening, secure backup/install paths, and regression coverage.
This commit is contained in:
@@ -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")));
|
||||
|
||||
@@ -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()));
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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: ®ex::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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user