mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +08:00
Merge remote-tracking branch 'origin/main' into codex/provider-policy-hardening
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")));
|
||||
|
||||
@@ -5,8 +5,8 @@ use serde_json::Value;
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
pub const ANTIGRAVITY_PROVIDER_TYPE: &str = "antigravity";
|
||||
pub const ANTIGRAVITY_REQUEST_USER_AGENT: &str =
|
||||
"antigravity/cli/1.0.16 (aidev_client; os_type=linux; arch=arm64; auth_method=consumer)";
|
||||
pub const ANTIGRAVITY_CLIENT_VERSION: &str = "4.3.0";
|
||||
pub const ANTIGRAVITY_REQUEST_USER_AGENT: &str = "vscode/1.X.X (Antigravity/4.3.0)";
|
||||
const ANTIGRAVITY_CLIENT_NAME: &str = "antigravity";
|
||||
const ANTIGRAVITY_GOOG_API_CLIENT: &str = "gl-node/18.18.2 fire/0.8.6 grpc/1.10.x";
|
||||
|
||||
@@ -109,7 +109,7 @@ pub fn build_antigravity_static_identity_headers(
|
||||
}
|
||||
|
||||
pub fn build_antigravity_static_client_headers(
|
||||
client_version: Option<&str>,
|
||||
_client_version: Option<&str>,
|
||||
session_id: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut headers = BTreeMap::from([
|
||||
@@ -125,14 +125,12 @@ pub fn build_antigravity_static_client_headers(
|
||||
String::from("user-agent"),
|
||||
String::from(ANTIGRAVITY_REQUEST_USER_AGENT),
|
||||
),
|
||||
(
|
||||
String::from("x-client-version"),
|
||||
String::from(ANTIGRAVITY_CLIENT_VERSION),
|
||||
),
|
||||
]);
|
||||
|
||||
if let Some(client_version) = client_version
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
headers.insert(String::from("x-client-version"), client_version.to_string());
|
||||
}
|
||||
if let Some(session_id) = session_id.map(str::trim).filter(|value| !value.is_empty()) {
|
||||
headers.insert(String::from("x-vscode-sessionid"), session_id.to_string());
|
||||
}
|
||||
@@ -273,7 +271,8 @@ mod tests {
|
||||
|
||||
use super::{
|
||||
build_antigravity_static_client_headers, resolve_local_antigravity_request_auth,
|
||||
AntigravityRequestAuth, AntigravityRequestAuthSupport, ANTIGRAVITY_REQUEST_USER_AGENT,
|
||||
AntigravityRequestAuth, AntigravityRequestAuthSupport, ANTIGRAVITY_CLIENT_VERSION,
|
||||
ANTIGRAVITY_REQUEST_USER_AGENT,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
@@ -383,7 +382,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn static_client_headers_use_native_antigravity_cli_user_agent() {
|
||||
fn static_client_headers_pin_the_known_good_antigravity_identity() {
|
||||
let headers = build_antigravity_static_client_headers(Some("1.0.16"), Some("session-abc"));
|
||||
|
||||
assert_eq!(
|
||||
@@ -396,7 +395,7 @@ mod tests {
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-client-version").map(String::as_str),
|
||||
Some("1.0.16")
|
||||
Some(ANTIGRAVITY_CLIENT_VERSION)
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-vscode-sessionid").map(String::as_str),
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::super::transport_proxy_is_locally_supported;
|
||||
use super::auth::{
|
||||
resolve_local_antigravity_request_auth, AntigravityRequestAuth, AntigravityRequestAuthSupport,
|
||||
AntigravityRequestAuthUnsupportedReason, ANTIGRAVITY_PROVIDER_TYPE,
|
||||
@@ -87,11 +88,11 @@ pub fn classify_local_antigravity_request_support(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedBodyRules,
|
||||
);
|
||||
}
|
||||
if transport.provider.proxy.is_some()
|
||||
|| transport.endpoint.proxy.is_some()
|
||||
|| transport.key.proxy.is_some()
|
||||
|| transport.key.fingerprint.is_some()
|
||||
{
|
||||
// A configured proxy is carried by the execution plan itself, so it only
|
||||
// disqualifies the local request when it cannot be resolved into a usable
|
||||
// snapshot. Transport profiles stay unsupported because the v1internal
|
||||
// payload never carries one.
|
||||
if !transport_proxy_is_locally_supported(transport) || transport.key.fingerprint.is_some() {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig,
|
||||
);
|
||||
@@ -114,3 +115,145 @@ pub fn classify_local_antigravity_request_support(
|
||||
|
||||
AntigravityRequestSideSupport::Supported(AntigravityRequestSideSpec { auth, request_type })
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::request::AntigravityEnvelopeRequestType;
|
||||
use super::{
|
||||
classify_local_antigravity_request_support, AntigravityRequestSideSupport,
|
||||
AntigravityRequestSideUnsupportedReason,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Antigravity".to_string(),
|
||||
provider_type: "antigravity".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: true,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "gemini:generate_content".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("generate_content".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://daily-cloudcode-pa.googleapis.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "oauth".to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec!["gemini:generate_content".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: Some(
|
||||
r#"{"provider_type":"antigravity","refresh_token":"rt","cloudaicompanionProject":"project-1"}"#
|
||||
.to_string(),
|
||||
),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn classify(transport: &GatewayProviderTransportSnapshot) -> AntigravityRequestSideSupport {
|
||||
classify_local_antigravity_request_support(
|
||||
transport,
|
||||
&json!({"contents": []}),
|
||||
AntigravityEnvelopeRequestType::Agent,
|
||||
)
|
||||
}
|
||||
|
||||
fn assert_unsupported_network_config(support: AntigravityRequestSideSupport) {
|
||||
assert_eq!(
|
||||
support,
|
||||
AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig,
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_resolvable_tunnel_node_proxy_keeps_the_envelope_supported() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.proxy = Some(json!({
|
||||
"enabled": true,
|
||||
"node_id": "702d158b-a432-4694-94cc-3bec13dbbc20",
|
||||
}));
|
||||
|
||||
assert!(matches!(
|
||||
classify(&transport),
|
||||
AntigravityRequestSideSupport::Supported(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_resolvable_url_proxy_keeps_the_envelope_supported() {
|
||||
for proxy_owner in ["provider", "endpoint", "key"] {
|
||||
let mut transport = sample_transport();
|
||||
let proxy = Some(json!({"enabled": true, "url": "http://127.0.0.1:17000"}));
|
||||
match proxy_owner {
|
||||
"provider" => transport.provider.proxy = proxy,
|
||||
"endpoint" => transport.endpoint.proxy = proxy,
|
||||
_ => transport.key.proxy = proxy,
|
||||
}
|
||||
|
||||
assert!(
|
||||
matches!(
|
||||
classify(&transport),
|
||||
AntigravityRequestSideSupport::Supported(_)
|
||||
),
|
||||
"a {proxy_owner} proxy should not disqualify the antigravity envelope"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_proxy_without_a_route_still_disqualifies_the_envelope() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.proxy = Some(json!({"enabled": true}));
|
||||
|
||||
assert_unsupported_network_config(classify(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_key_fingerprint_still_disqualifies_the_envelope() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.fingerprint = Some(json!({"transport_profile": "chrome"}));
|
||||
|
||||
assert_unsupported_network_config(classify(&transport));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::headers::{
|
||||
is_aether_internal_header, is_upstream_credential_header, normalize_upstream_accept_encoding,
|
||||
should_skip_upstream_complete_passthrough_header, should_skip_upstream_passthrough_header,
|
||||
declared_connection_header_names, is_aether_internal_header, is_upstream_credential_header,
|
||||
normalize_upstream_accept_encoding, remove_declared_connection_headers,
|
||||
should_skip_upstream_complete_passthrough_header_with_connection,
|
||||
should_skip_upstream_passthrough_header_with_connection,
|
||||
};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
@@ -13,13 +15,17 @@ fn collect_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
|
||||
let mut out = BTreeMap::new();
|
||||
for (name, value) in headers.iter() {
|
||||
let Ok(value) = value.to_str() else {
|
||||
continue;
|
||||
};
|
||||
let key = name.as_str().to_ascii_lowercase();
|
||||
if should_skip_upstream_passthrough_header(&key) {
|
||||
if should_skip_upstream_passthrough_header_with_connection(
|
||||
&key,
|
||||
&declared_connection_headers,
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = normalize_passthrough_header_value(&key, value) else {
|
||||
@@ -30,7 +36,10 @@ fn collect_passthrough_headers(
|
||||
|
||||
for (key, value) in extra_headers {
|
||||
let normalized_key = key.to_ascii_lowercase();
|
||||
if should_skip_upstream_passthrough_header(&normalized_key) {
|
||||
if should_skip_upstream_passthrough_header_with_connection(
|
||||
&normalized_key,
|
||||
&declared_connection_headers,
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
|
||||
@@ -46,13 +55,17 @@ fn collect_complete_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
|
||||
let mut out = BTreeMap::new();
|
||||
for (name, value) in headers.iter() {
|
||||
let Ok(value) = value.to_str() else {
|
||||
continue;
|
||||
};
|
||||
let key = name.as_str().to_ascii_lowercase();
|
||||
if should_skip_upstream_complete_passthrough_header(&key) {
|
||||
if should_skip_upstream_complete_passthrough_header_with_connection(
|
||||
&key,
|
||||
&declared_connection_headers,
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = normalize_passthrough_header_value(&key, value) else {
|
||||
@@ -63,7 +76,10 @@ fn collect_complete_passthrough_headers(
|
||||
|
||||
for (key, value) in extra_headers {
|
||||
let normalized_key = key.to_ascii_lowercase();
|
||||
if should_skip_upstream_complete_passthrough_header(&normalized_key) {
|
||||
if should_skip_upstream_complete_passthrough_header_with_connection(
|
||||
&normalized_key,
|
||||
&declared_connection_headers,
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
|
||||
@@ -101,6 +117,8 @@ pub fn build_passthrough_headers(
|
||||
.trim()
|
||||
.to_string()
|
||||
});
|
||||
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
|
||||
remove_declared_connection_headers(&mut out, &declared_connection_headers);
|
||||
out.remove("content-length");
|
||||
out
|
||||
}
|
||||
@@ -112,8 +130,10 @@ pub fn build_openai_passthrough_headers(
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
|
||||
let mut out = build_passthrough_headers(headers, extra_headers, content_type);
|
||||
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
|
||||
remove_declared_connection_headers(&mut out, &declared_connection_headers);
|
||||
out
|
||||
}
|
||||
|
||||
@@ -122,7 +142,9 @@ pub fn build_complete_passthrough_headers(
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
|
||||
let mut out = collect_complete_passthrough_headers(headers, extra_headers);
|
||||
remove_declared_connection_headers(&mut out, &declared_connection_headers);
|
||||
out.entry("content-type".to_string()).or_insert_with(|| {
|
||||
content_type
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
@@ -130,6 +152,7 @@ pub fn build_complete_passthrough_headers(
|
||||
.trim()
|
||||
.to_string()
|
||||
});
|
||||
remove_declared_connection_headers(&mut out, &declared_connection_headers);
|
||||
out.remove("content-length");
|
||||
out
|
||||
}
|
||||
@@ -141,8 +164,10 @@ pub fn build_complete_passthrough_headers_with_auth(
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
|
||||
let mut out = build_complete_passthrough_headers(headers, extra_headers, content_type);
|
||||
replace_upstream_auth_headers(&mut out, auth_header, auth_value);
|
||||
remove_declared_connection_headers(&mut out, &declared_connection_headers);
|
||||
out
|
||||
}
|
||||
|
||||
@@ -153,6 +178,7 @@ pub fn build_claude_passthrough_headers(
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
|
||||
let mut out = build_openai_passthrough_headers(
|
||||
headers,
|
||||
auth_header,
|
||||
@@ -164,7 +190,10 @@ pub fn build_claude_passthrough_headers(
|
||||
for (name, value) in extra_headers {
|
||||
let key = name.to_ascii_lowercase();
|
||||
let value = value.trim();
|
||||
if value.is_empty() || !should_restore_claude_passthrough_header(&key) {
|
||||
if value.is_empty()
|
||||
|| !should_restore_claude_passthrough_header(&key)
|
||||
|| declared_connection_headers.contains(&key)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -185,7 +214,10 @@ pub fn build_claude_passthrough_headers(
|
||||
};
|
||||
let key = name.as_str().to_ascii_lowercase();
|
||||
let value = value.trim();
|
||||
if value.is_empty() || !should_restore_claude_passthrough_header(&key) {
|
||||
if value.is_empty()
|
||||
|| !should_restore_claude_passthrough_header(&key)
|
||||
|| declared_connection_headers.contains(&key)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -202,6 +234,7 @@ pub fn build_claude_passthrough_headers(
|
||||
|
||||
out.entry("anthropic-version".to_string())
|
||||
.or_insert_with(|| DEFAULT_ANTHROPIC_VERSION.to_string());
|
||||
remove_declared_connection_headers(&mut out, &declared_connection_headers);
|
||||
out
|
||||
}
|
||||
|
||||
@@ -211,8 +244,10 @@ pub fn build_passthrough_headers_with_auth(
|
||||
auth_value: &str,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let declared_connection_headers = declared_connection_header_names(headers, extra_headers);
|
||||
let mut out = collect_passthrough_headers(headers, extra_headers);
|
||||
replace_upstream_auth_headers(&mut out, auth_header, auth_value);
|
||||
remove_declared_connection_headers(&mut out, &declared_connection_headers);
|
||||
out.remove("content-length");
|
||||
out
|
||||
}
|
||||
@@ -357,9 +392,9 @@ fn bearer_auth_value(secret: &str) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_claude_passthrough_headers, build_complete_passthrough_headers_with_auth,
|
||||
build_openai_passthrough_headers, resolve_local_openai_bearer_auth,
|
||||
resolve_local_standard_auth,
|
||||
build_claude_passthrough_headers, build_complete_passthrough_headers,
|
||||
build_complete_passthrough_headers_with_auth, build_openai_passthrough_headers,
|
||||
resolve_local_openai_bearer_auth, resolve_local_standard_auth,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
@@ -495,6 +530,67 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn passthrough_headers_strip_connection_declared_fields() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.append(
|
||||
http::header::CONNECTION,
|
||||
http::HeaderValue::from_static("X-Internal-Hop, keep-alive"),
|
||||
);
|
||||
headers.append(
|
||||
http::header::CONNECTION,
|
||||
http::HeaderValue::from_static("x-extra-hop"),
|
||||
);
|
||||
headers.insert(
|
||||
"x-internal-hop",
|
||||
http::HeaderValue::from_static("private-value"),
|
||||
);
|
||||
headers.insert(
|
||||
"x-extra-hop",
|
||||
http::HeaderValue::from_static("private-value-2"),
|
||||
);
|
||||
headers.insert("x-public", http::HeaderValue::from_static("ok"));
|
||||
|
||||
let extra = BTreeMap::from([
|
||||
(
|
||||
"Connection".to_string(),
|
||||
"X-Extra-From-Connection".to_string(),
|
||||
),
|
||||
("X-Extra-From-Connection".to_string(), "secret".to_string()),
|
||||
]);
|
||||
let built = build_openai_passthrough_headers(
|
||||
&headers,
|
||||
"authorization",
|
||||
"Bearer upstream",
|
||||
&extra,
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
assert_eq!(built.get("x-public").map(String::as_str), Some("ok"));
|
||||
assert!(!built.contains_key("connection"));
|
||||
assert!(!built.contains_key("x-internal-hop"));
|
||||
assert!(!built.contains_key("x-extra-hop"));
|
||||
assert!(built
|
||||
.keys()
|
||||
.all(|name| { !name.eq_ignore_ascii_case("x-extra-from-connection") }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn complete_passthrough_headers_strip_connection_declared_fields_from_extra_headers() {
|
||||
let headers = http::HeaderMap::new();
|
||||
let extra = BTreeMap::from([
|
||||
("Connection".to_string(), "x-private-hop".to_string()),
|
||||
("X-Private-Hop".to_string(), "secret".to_string()),
|
||||
("x-public".to_string(), "ok".to_string()),
|
||||
]);
|
||||
|
||||
let built = build_complete_passthrough_headers(&headers, &extra, None);
|
||||
|
||||
assert_eq!(built.get("x-public").map(String::as_str), Some("ok"));
|
||||
assert!(!built.contains_key("connection"));
|
||||
assert!(!built.contains_key("x-private-hop"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_passthrough_headers_preserve_explicit_anthropic_version_override() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use aether_ai_formats::formats::matrix::{
|
||||
request_conversion_kind, request_conversion_requires_enable_flag,
|
||||
};
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
|
||||
use serde_json::{json, Map, Value};
|
||||
|
||||
use crate::conversion::{
|
||||
@@ -71,25 +71,24 @@ pub fn build_transport_diagnostics(
|
||||
) -> Value {
|
||||
let resolved_transport_profile_id = resolve_transport_profile_id(transport);
|
||||
let resolved_transport_profile = resolve_transport_profile(transport)
|
||||
.and_then(|profile| serde_json::to_value(profile).ok())
|
||||
.as_ref()
|
||||
.map(summarize_transport_profile)
|
||||
.unwrap_or(Value::Null);
|
||||
let configured_key_transport_profile = transport
|
||||
let key_transport_profile_configured = transport
|
||||
.key
|
||||
.fingerprint
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get("transport_profile"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null);
|
||||
let configured_provider_transport_profile = transport
|
||||
.is_some_and(|value| !value.is_null());
|
||||
let provider_transport_profile_configured = transport
|
||||
.provider
|
||||
.config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("fingerprint"))
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get("transport_profile"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null);
|
||||
.is_some_and(|value| !value.is_null());
|
||||
let configured_legacy_grok_transport_profile = if transport
|
||||
.provider
|
||||
.provider_type
|
||||
@@ -109,7 +108,8 @@ pub fn build_transport_diagnostics(
|
||||
&auth_config,
|
||||
"grok_auth_config",
|
||||
)
|
||||
.and_then(|profile| serde_json::to_value(profile).ok())
|
||||
.as_ref()
|
||||
.map(summarize_transport_profile)
|
||||
})
|
||||
.unwrap_or(Value::Null)
|
||||
} else {
|
||||
@@ -132,11 +132,14 @@ pub fn build_transport_diagnostics(
|
||||
"key_is_active": transport.key.is_active,
|
||||
"provider_enable_format_conversion": transport.provider.enable_format_conversion,
|
||||
"provider_keep_priority_on_conversion": transport.provider.keep_priority_on_conversion,
|
||||
"endpoint_format_acceptance_config": transport.endpoint.format_acceptance_config,
|
||||
"endpoint_custom_path": transport.endpoint.custom_path,
|
||||
"header_rules": transport.endpoint.header_rules,
|
||||
"endpoint_format_acceptance": summarize_format_acceptance_config(
|
||||
transport.endpoint.format_acceptance_config.as_ref()
|
||||
),
|
||||
"endpoint_has_custom_path": transport.endpoint.custom_path.as_deref()
|
||||
.is_some_and(|value| !value.trim().is_empty()),
|
||||
"header_rules_count": json_array_len(transport.endpoint.header_rules.as_ref()),
|
||||
"header_rules_supported": header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()),
|
||||
"body_rules": transport.endpoint.body_rules,
|
||||
"body_rules_count": json_array_len(transport.endpoint.body_rules.as_ref()),
|
||||
"body_rules_supported": body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()),
|
||||
"proxy": {
|
||||
"locally_supported": transport_proxy_is_locally_supported(transport),
|
||||
@@ -149,9 +152,10 @@ pub fn build_transport_diagnostics(
|
||||
"has_oauth_config": has_oauth_config,
|
||||
"oauth_request_auth_resolution_supported": oauth_resolution_supported,
|
||||
},
|
||||
"fingerprint": transport.key.fingerprint,
|
||||
"configured_key_transport_profile": configured_key_transport_profile,
|
||||
"configured_provider_transport_profile": configured_provider_transport_profile,
|
||||
"key_fingerprint_configured": transport.key.fingerprint.as_ref()
|
||||
.is_some_and(|value| !value.is_null()),
|
||||
"key_transport_profile_configured": key_transport_profile_configured,
|
||||
"provider_transport_profile_configured": provider_transport_profile_configured,
|
||||
"configured_legacy_grok_transport_profile": configured_legacy_grok_transport_profile,
|
||||
"resolved_transport_profile_id": resolved_transport_profile_id,
|
||||
"resolved_transport_profile": resolved_transport_profile,
|
||||
@@ -177,6 +181,34 @@ pub fn build_transport_diagnostics(
|
||||
})
|
||||
}
|
||||
|
||||
fn json_array_len(value: Option<&Value>) -> usize {
|
||||
value.and_then(Value::as_array).map_or(0, Vec::len)
|
||||
}
|
||||
|
||||
fn summarize_format_acceptance_config(value: Option<&Value>) -> Value {
|
||||
let Some(object) = value.and_then(Value::as_object) else {
|
||||
return json!({ "configured": false });
|
||||
};
|
||||
json!({
|
||||
"configured": true,
|
||||
"enabled": object.get("enabled").and_then(Value::as_bool),
|
||||
"accept_formats_count": json_array_len(object.get("accept_formats")),
|
||||
"reject_formats_count": json_array_len(object.get("reject_formats")),
|
||||
})
|
||||
}
|
||||
|
||||
fn summarize_transport_profile(profile: &ResolvedTransportProfile) -> Value {
|
||||
json!({
|
||||
"profile_id": profile.profile_id,
|
||||
"backend": profile.backend,
|
||||
"http_mode": profile.http_mode,
|
||||
"pool_scope": profile.pool_scope,
|
||||
"has_header_fingerprint": profile.header_fingerprint.as_ref()
|
||||
.is_some_and(|value| !value.is_null()),
|
||||
"has_extra": profile.extra.as_ref().is_some_and(|value| !value.is_null()),
|
||||
})
|
||||
}
|
||||
|
||||
fn summarize_proxy_config(proxy: Option<&Value>) -> Value {
|
||||
let Some(object) = proxy.and_then(Value::as_object) else {
|
||||
return Value::Null;
|
||||
@@ -187,11 +219,16 @@ fn summarize_proxy_config(proxy: Option<&Value>) -> Value {
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty());
|
||||
json!({
|
||||
"configured": true,
|
||||
"enabled": object.get("enabled").cloned().unwrap_or(Value::Null),
|
||||
"mode": object.get("mode").cloned().unwrap_or(Value::Null),
|
||||
"node_id": object.get("node_id").cloned().unwrap_or(Value::Null),
|
||||
"label": object.get("label").cloned().unwrap_or(Value::Null),
|
||||
"has_mode": object.get("mode").and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty()),
|
||||
"has_node_id": object.get("node_id").and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty()),
|
||||
"has_label": object.get("label").and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty()),
|
||||
"has_url": has_url,
|
||||
"has_extra": object.get("extra").is_some_and(|value| !value.is_null()),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -426,11 +463,13 @@ mod tests {
|
||||
build_transport_diagnostics(&sample_transport(), "claude:messages", "openai:responses");
|
||||
|
||||
assert_eq!(diagnostics["provider_type"], "codex");
|
||||
assert_eq!(diagnostics["key_fingerprint_configured"], true);
|
||||
assert_eq!(diagnostics["key_transport_profile_configured"], true);
|
||||
assert_eq!(diagnostics["resolved_transport_profile_id"], "chrome_136");
|
||||
assert_eq!(
|
||||
diagnostics["fingerprint"]["transport_profile"]["profile_id"],
|
||||
diagnostics["resolved_transport_profile"]["profile_id"],
|
||||
"chrome_136"
|
||||
);
|
||||
assert_eq!(diagnostics["resolved_transport_profile_id"], "chrome_136");
|
||||
assert_eq!(
|
||||
diagnostics["request_pair"]["conversion_enabled"],
|
||||
Value::Bool(true)
|
||||
@@ -489,6 +528,86 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_diagnostics_do_not_serialize_configured_secrets() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.proxy = Some(json!({
|
||||
"enabled": true,
|
||||
"mode": "secret-proxy-mode",
|
||||
"node_id": "secret-proxy-node",
|
||||
"label": "secret-proxy-label",
|
||||
"url": "https://secret-user:[email protected]/secret-path",
|
||||
"extra": {"token": "secret-proxy-extra"}
|
||||
}));
|
||||
transport.provider.config = Some(json!({
|
||||
"secret": "secret-provider-config",
|
||||
"fingerprint": {
|
||||
"transport_profile": {
|
||||
"profile_id": "safe-profile",
|
||||
"header_fingerprint": {"authorization": "secret-profile-header"},
|
||||
"extra": {"token": "secret-profile-extra"}
|
||||
}
|
||||
}
|
||||
}));
|
||||
transport.endpoint.custom_path = Some("/secret-custom-path".to_string());
|
||||
transport.endpoint.header_rules = Some(json!([
|
||||
{"op": "set", "key": "authorization", "value": "secret-header-rule"}
|
||||
]));
|
||||
transport.endpoint.body_rules = Some(json!([
|
||||
{"op": "set", "path": "auth.token", "value": "secret-body-rule"}
|
||||
]));
|
||||
transport.endpoint.format_acceptance_config = Some(json!({
|
||||
"enabled": true,
|
||||
"accept_formats": ["secret-accepted-format"],
|
||||
"reject_formats": ["secret-rejected-format"],
|
||||
"token": "secret-format-config"
|
||||
}));
|
||||
transport.key.fingerprint = Some(json!({
|
||||
"secret": "secret-key-fingerprint",
|
||||
"transport_profile": {
|
||||
"profile_id": "safe-key-profile",
|
||||
"header_fingerprint": {"authorization": "secret-key-profile-header"},
|
||||
"extra": {"token": "secret-key-profile-extra"}
|
||||
}
|
||||
}));
|
||||
|
||||
let diagnostics =
|
||||
build_transport_diagnostics(&transport, "claude:messages", "openai:responses");
|
||||
let serialized = serde_json::to_string(&diagnostics).unwrap();
|
||||
|
||||
for secret in [
|
||||
"secret-proxy-mode",
|
||||
"secret-proxy-node",
|
||||
"secret-proxy-label",
|
||||
"secret-user",
|
||||
"secret-pass",
|
||||
"secret-path",
|
||||
"secret-proxy-extra",
|
||||
"secret-provider-config",
|
||||
"secret-profile-header",
|
||||
"secret-profile-extra",
|
||||
"secret-custom-path",
|
||||
"secret-header-rule",
|
||||
"secret-body-rule",
|
||||
"secret-accepted-format",
|
||||
"secret-rejected-format",
|
||||
"secret-format-config",
|
||||
"secret-key-fingerprint",
|
||||
"secret-key-profile-header",
|
||||
"secret-key-profile-extra",
|
||||
] {
|
||||
assert!(!serialized.contains(secret), "leaked {secret}");
|
||||
}
|
||||
assert_eq!(diagnostics["header_rules_count"], 1);
|
||||
assert_eq!(diagnostics["body_rules_count"], 1);
|
||||
assert_eq!(diagnostics["proxy"]["provider"]["has_node_id"], true);
|
||||
assert_eq!(
|
||||
diagnostics["resolved_transport_profile"]["has_header_fingerprint"],
|
||||
true
|
||||
);
|
||||
assert_eq!(diagnostics["resolved_transport_profile"]["has_extra"], true);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_trace_proxy_value_sanitizes_url_and_marks_config_source() {
|
||||
let transport = sample_transport();
|
||||
|
||||
@@ -1,15 +1,26 @@
|
||||
use serde_json::Value;
|
||||
use std::fmt;
|
||||
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
pub const GEMINI_CLI_PROVIDER_TYPE: &str = "gemini_cli";
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
#[derive(Clone, Default, PartialEq, Eq)]
|
||||
pub struct GeminiCliRequestAuth {
|
||||
pub project_id: Option<String>,
|
||||
pub session_id: Option<String>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for GeminiCliRequestAuth {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("GeminiCliRequestAuth")
|
||||
.field("project_id", &self.project_id)
|
||||
.field("has_session_id", &self.session_id.is_some())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GeminiCliRequestAuthSupport {
|
||||
Supported(GeminiCliRequestAuth),
|
||||
|
||||
@@ -1,13 +1,31 @@
|
||||
use serde_json::{Map, Value};
|
||||
use std::fmt;
|
||||
|
||||
use super::auth::GeminiCliRequestAuth;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub enum GeminiCliRequestEnvelopeSupport {
|
||||
Supported(Value),
|
||||
Unsupported(GeminiCliRequestEnvelopeUnsupportedReason),
|
||||
}
|
||||
|
||||
impl fmt::Debug for GeminiCliRequestEnvelopeSupport {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Supported(body) => formatter
|
||||
.debug_struct("Supported")
|
||||
.field(
|
||||
"body_bytes",
|
||||
&serde_json::to_vec(body).ok().map(|bytes| bytes.len()),
|
||||
)
|
||||
.finish(),
|
||||
Self::Unsupported(reason) => {
|
||||
formatter.debug_tuple("Unsupported").field(reason).finish()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GeminiCliRequestEnvelopeUnsupportedReason {
|
||||
NonObjectBody,
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::fmt;
|
||||
|
||||
use serde_json::{json, Value};
|
||||
|
||||
@@ -17,13 +18,36 @@ pub enum GeminiFilesRequestBodyError {
|
||||
BodyRulesApplyFailed,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct GeminiFilesRequestBodyParts {
|
||||
pub provider_request_body: Option<Value>,
|
||||
pub provider_request_body_base64: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
impl fmt::Debug for GeminiFilesRequestBodyParts {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("GeminiFilesRequestBodyParts")
|
||||
.field(
|
||||
"has_provider_request_body",
|
||||
&self.provider_request_body.is_some(),
|
||||
)
|
||||
.field(
|
||||
"provider_request_body_bytes",
|
||||
&self
|
||||
.provider_request_body
|
||||
.as_ref()
|
||||
.and_then(|body| serde_json::to_vec(body).ok().map(|bytes| bytes.len())),
|
||||
)
|
||||
.field(
|
||||
"provider_request_body_base64_len",
|
||||
&self.provider_request_body_base64.as_ref().map(String::len),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
pub struct GeminiFilesHeadersInput<'a> {
|
||||
pub headers: &'a http::HeaderMap,
|
||||
pub auth_header: &'a str,
|
||||
@@ -35,6 +59,40 @@ pub struct GeminiFilesHeadersInput<'a> {
|
||||
pub original_body_is_empty: bool,
|
||||
}
|
||||
|
||||
impl fmt::Debug for GeminiFilesHeadersInput<'_> {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("GeminiFilesHeadersInput")
|
||||
.field(
|
||||
"request_header_names",
|
||||
&self
|
||||
.headers
|
||||
.keys()
|
||||
.map(|name| name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
)
|
||||
.field("auth_header", &self.auth_header)
|
||||
.field("has_auth_value", &(!self.auth_value.is_empty()))
|
||||
.field("has_header_rules", &self.header_rules.is_some())
|
||||
.field(
|
||||
"has_provider_request_body",
|
||||
&self.provider_request_body.is_some(),
|
||||
)
|
||||
.field(
|
||||
"provider_request_body_base64_len",
|
||||
&self.provider_request_body_base64.map(str::len),
|
||||
)
|
||||
.field(
|
||||
"original_request_body_bytes",
|
||||
&serde_json::to_vec(self.original_request_body_json)
|
||||
.ok()
|
||||
.map(|bytes| bytes.len()),
|
||||
)
|
||||
.field("original_body_is_empty", &self.original_body_is_empty)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn gemini_files_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
@@ -135,6 +193,12 @@ pub fn build_gemini_files_headers(
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
let declared_connection_headers =
|
||||
crate::headers::declared_connection_header_names(input.headers, &BTreeMap::new());
|
||||
crate::headers::remove_declared_connection_headers(
|
||||
&mut provider_request_headers,
|
||||
&declared_connection_headers,
|
||||
);
|
||||
Some(provider_request_headers)
|
||||
}
|
||||
|
||||
|
||||
@@ -62,9 +62,26 @@ pub fn resolve_local_generic_oauth_transport_authorization(
|
||||
.map(|token| format!("Bearer {token}"))
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
#[derive(Clone, Default)]
|
||||
pub struct GenericOAuthRefreshAdapter {
|
||||
token_url_overrides: BTreeMap<String, String>,
|
||||
oauth_credentials_overrides: BTreeMap<String, (String, String)>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for GenericOAuthRefreshAdapter {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("GenericOAuthRefreshAdapter")
|
||||
.field(
|
||||
"token_url_override_provider_types",
|
||||
&self.token_url_overrides.keys().collect::<Vec<_>>(),
|
||||
)
|
||||
.field(
|
||||
"oauth_credentials_override_provider_types",
|
||||
&self.oauth_credentials_overrides.keys().collect::<Vec<_>>(),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl GenericOAuthRefreshAdapter {
|
||||
@@ -78,11 +95,29 @@ impl GenericOAuthRefreshAdapter {
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_oauth_credentials_for_tests(
|
||||
mut self,
|
||||
provider_type: &str,
|
||||
client_id: impl Into<String>,
|
||||
client_secret: impl Into<String>,
|
||||
) -> Self {
|
||||
self.oauth_credentials_overrides.insert(
|
||||
provider_type.trim().to_ascii_lowercase(),
|
||||
(client_id.into(), client_secret.into()),
|
||||
);
|
||||
self
|
||||
}
|
||||
|
||||
fn adapter_for_provider_type(
|
||||
&self,
|
||||
provider_type: &'static str,
|
||||
) -> Option<GenericProviderOAuthAdapter> {
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?;
|
||||
let mut adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?;
|
||||
if let Some((client_id, client_secret)) =
|
||||
self.oauth_credentials_overrides.get(provider_type)
|
||||
{
|
||||
adapter = adapter.with_oauth_credentials_for_tests(client_id, client_secret);
|
||||
}
|
||||
if let Some(token_url) = self.token_url_overrides.get(provider_type) {
|
||||
return Some(adapter.with_token_url_override(token_url.clone()));
|
||||
}
|
||||
@@ -912,7 +947,12 @@ mod tests {
|
||||
hits: Arc::clone(&hits),
|
||||
};
|
||||
let adapter = GenericOAuthRefreshAdapter::default()
|
||||
.with_token_url_for_tests("antigravity", "https://oauth.example/token");
|
||||
.with_token_url_for_tests("antigravity", "https://oauth.example/token")
|
||||
.with_oauth_credentials_for_tests(
|
||||
"antigravity",
|
||||
"test-client-id",
|
||||
"test-client-secret",
|
||||
);
|
||||
|
||||
assert!(adapter.supports(&transport));
|
||||
assert!(adapter.should_refresh(&transport, None));
|
||||
|
||||
@@ -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