mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
feat(gateway): harden provider request execution
Preserve exact request payloads and model client surface and API operation explicitly. Add Anthropic compatibility profiles, bounded stream commitment, and scoped OAuth retry behavior across provider transports.
This commit is contained in:
@@ -0,0 +1,376 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
/// Provider-side compatibility applied to otherwise same-format Anthropic requests.
|
||||
///
|
||||
/// Native Anthropic endpoints should remain transparent. The legacy Claude Code
|
||||
/// profile is opt-in, except for the existing `claude_code` provider type where it
|
||||
/// remains the backwards-compatible default. Endpoint config takes precedence over
|
||||
/// provider config. The canonical field is `anthropic.compatibility_profile`;
|
||||
/// explicitly Anthropic-namespaced legacy spellings remain accepted.
|
||||
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum AnthropicCompatibilityProfile {
|
||||
#[default]
|
||||
#[serde(
|
||||
alias = "native",
|
||||
alias = "anthropic",
|
||||
alias = "none",
|
||||
alias = "transparent"
|
||||
)]
|
||||
NativeTransparent,
|
||||
#[serde(
|
||||
alias = "claude_code",
|
||||
alias = "claude-code",
|
||||
alias = "legacy",
|
||||
alias = "same_format_compat"
|
||||
)]
|
||||
ClaudeCodeLegacy,
|
||||
}
|
||||
|
||||
impl AnthropicCompatibilityProfile {
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::NativeTransparent => "native_transparent",
|
||||
Self::ClaudeCodeLegacy => "claude_code_legacy",
|
||||
}
|
||||
}
|
||||
|
||||
pub const fn uses_claude_code_compatibility(self) -> bool {
|
||||
matches!(self, Self::ClaudeCodeLegacy)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
|
||||
#[error("unknown Anthropic compatibility profile")]
|
||||
pub struct AnthropicCompatibilityProfileConfigError;
|
||||
|
||||
/// Validate the optional Anthropic compatibility fields in a provider or
|
||||
/// endpoint config without resolving provider defaults or mutating config.
|
||||
pub fn validate_anthropic_compatibility_profile_config(
|
||||
config: Option<&Value>,
|
||||
) -> Result<(), AnthropicCompatibilityProfileConfigError> {
|
||||
match profile_from_config(config) {
|
||||
ConfiguredProfile::Absent | ConfiguredProfile::Valid(_) => Ok(()),
|
||||
ConfiguredProfile::Invalid => Err(AnthropicCompatibilityProfileConfigError),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_anthropic_compatibility_profile(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> AnthropicCompatibilityProfile {
|
||||
if !aether_ai_formats::api_format_alias_matches(provider_api_format, "claude:messages") {
|
||||
return AnthropicCompatibilityProfile::NativeTransparent;
|
||||
}
|
||||
|
||||
for (scope, resolution) in [
|
||||
(
|
||||
"endpoint",
|
||||
profile_from_config(transport.endpoint.config.as_ref()),
|
||||
),
|
||||
(
|
||||
"provider",
|
||||
profile_from_config(transport.provider.config.as_ref()),
|
||||
),
|
||||
] {
|
||||
match resolution {
|
||||
ConfiguredProfile::Valid(profile) => return profile,
|
||||
ConfiguredProfile::Invalid => {
|
||||
tracing::warn!(
|
||||
event_name = "anthropic_compatibility_profile_invalid",
|
||||
log_type = "ops",
|
||||
provider_id = %transport.provider.id,
|
||||
endpoint_id = %transport.endpoint.id,
|
||||
config_scope = scope,
|
||||
"invalid Anthropic compatibility profile; using native transparent behavior"
|
||||
);
|
||||
return AnthropicCompatibilityProfile::NativeTransparent;
|
||||
}
|
||||
ConfiguredProfile::Absent => {}
|
||||
}
|
||||
}
|
||||
|
||||
if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("claude_code")
|
||||
{
|
||||
AnthropicCompatibilityProfile::ClaudeCodeLegacy
|
||||
} else {
|
||||
AnthropicCompatibilityProfile::NativeTransparent
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum ConfiguredProfile {
|
||||
Absent,
|
||||
Valid(AnthropicCompatibilityProfile),
|
||||
Invalid,
|
||||
}
|
||||
|
||||
fn profile_from_config(config: Option<&Value>) -> ConfiguredProfile {
|
||||
let Some(config) = config.and_then(Value::as_object) else {
|
||||
return ConfiguredProfile::Absent;
|
||||
};
|
||||
|
||||
for (container, fields) in [
|
||||
(
|
||||
config.get("anthropic").and_then(Value::as_object),
|
||||
&["compatibility_profile", "compatibilityProfile", "profile"][..],
|
||||
),
|
||||
(
|
||||
config
|
||||
.get("anthropic_compatibility")
|
||||
.and_then(Value::as_object),
|
||||
&["profile", "compatibility_profile", "compatibilityProfile"][..],
|
||||
),
|
||||
(
|
||||
config
|
||||
.get("anthropicCompatibility")
|
||||
.and_then(Value::as_object),
|
||||
&["profile", "compatibilityProfile", "compatibility_profile"][..],
|
||||
),
|
||||
] {
|
||||
let Some(container) = container else {
|
||||
continue;
|
||||
};
|
||||
for &field in fields {
|
||||
if let Some(profile) = container.get(field).and_then(parse_profile_value) {
|
||||
return ConfiguredProfile::Valid(profile);
|
||||
}
|
||||
if container.contains_key(field) {
|
||||
return ConfiguredProfile::Invalid;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for field in [
|
||||
"anthropic_compatibility_profile",
|
||||
"anthropicCompatibilityProfile",
|
||||
] {
|
||||
if let Some(value) = config.get(field) {
|
||||
return parse_profile_value(value)
|
||||
.map(ConfiguredProfile::Valid)
|
||||
.unwrap_or(ConfiguredProfile::Invalid);
|
||||
}
|
||||
}
|
||||
|
||||
let Some(claude_code_advanced) = config
|
||||
.get("claude_code_advanced")
|
||||
.and_then(Value::as_object)
|
||||
else {
|
||||
return ConfiguredProfile::Absent;
|
||||
};
|
||||
for field in ["compatibility_profile", "compatibilityProfile"] {
|
||||
if let Some(value) = claude_code_advanced.get(field) {
|
||||
return parse_profile_value(value)
|
||||
.map(ConfiguredProfile::Valid)
|
||||
.unwrap_or(ConfiguredProfile::Invalid);
|
||||
}
|
||||
}
|
||||
ConfiguredProfile::Absent
|
||||
}
|
||||
|
||||
fn parse_profile_value(value: &Value) -> Option<AnthropicCompatibilityProfile> {
|
||||
if let Some(object) = value.as_object() {
|
||||
return ["kind", "name", "profile"]
|
||||
.into_iter()
|
||||
.find_map(|field| object.get(field).and_then(parse_profile_value));
|
||||
}
|
||||
serde_json::from_value(value.clone()).ok()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
resolve_anthropic_compatibility_profile, validate_anthropic_compatibility_profile_config,
|
||||
AnthropicCompatibilityProfile,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
fn sample_transport(provider_type: &str) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "provider".to_string(),
|
||||
provider_type: provider_type.to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "claude:messages".to_string(),
|
||||
api_family: Some("claude".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://api.anthropic.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validates_configured_profile_without_resolving_provider_defaults() {
|
||||
assert!(validate_anthropic_compatibility_profile_config(None).is_ok());
|
||||
assert!(
|
||||
validate_anthropic_compatibility_profile_config(Some(&json!({
|
||||
"anthropic_compatibility": {"profile": "native_transparent"}
|
||||
})))
|
||||
.is_ok()
|
||||
);
|
||||
assert!(
|
||||
validate_anthropic_compatibility_profile_config(Some(&json!({
|
||||
"anthropic": {"compatibility_profile": "claude_code_legacy"}
|
||||
})))
|
||||
.is_ok()
|
||||
);
|
||||
|
||||
let error = validate_anthropic_compatibility_profile_config(Some(&json!({
|
||||
"anthropic_compatibility": {"profile": "claude_cod_typo"}
|
||||
})))
|
||||
.expect_err("unknown profile should be rejected");
|
||||
assert_eq!(error.to_string(), "unknown Anthropic compatibility profile");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ignores_generic_compatibility_fields_owned_by_other_transports() {
|
||||
let config = json!({
|
||||
"compatibility_profile": "strict",
|
||||
"compatibility": {"profile": "v2"},
|
||||
"adaptation": {"profile": "legacy"}
|
||||
});
|
||||
|
||||
assert!(validate_anthropic_compatibility_profile_config(Some(&config)).is_ok());
|
||||
let mut transport = sample_transport("custom");
|
||||
transport.provider.config = Some(config);
|
||||
assert_eq!(
|
||||
resolve_anthropic_compatibility_profile(&transport, "claude:messages"),
|
||||
AnthropicCompatibilityProfile::NativeTransparent
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn native_anthropic_is_transparent_by_default() {
|
||||
let transport = sample_transport("custom");
|
||||
|
||||
assert_eq!(
|
||||
resolve_anthropic_compatibility_profile(&transport, "claude:messages"),
|
||||
AnthropicCompatibilityProfile::NativeTransparent
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_code_keeps_legacy_compatibility_by_default() {
|
||||
let transport = sample_transport("claude_code");
|
||||
|
||||
assert_eq!(
|
||||
resolve_anthropic_compatibility_profile(&transport, "claude:messages"),
|
||||
AnthropicCompatibilityProfile::ClaudeCodeLegacy
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn endpoint_profile_overrides_provider_and_legacy_defaults() {
|
||||
let mut transport = sample_transport("claude_code");
|
||||
transport.provider.config = Some(json!({
|
||||
"anthropic": {"compatibility_profile": "claude_code_legacy"}
|
||||
}));
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic_compatibility": {
|
||||
"profile": "native_transparent"
|
||||
}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
resolve_anthropic_compatibility_profile(&transport, "claude:messages"),
|
||||
AnthropicCompatibilityProfile::NativeTransparent
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_profile_can_explicitly_enable_legacy_compatibility() {
|
||||
let mut transport = sample_transport("custom");
|
||||
transport.provider.config = Some(json!({
|
||||
"anthropic": {
|
||||
"compatibility_profile": "claude_code"
|
||||
}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
resolve_anthropic_compatibility_profile(&transport, "claude:messages"),
|
||||
AnthropicCompatibilityProfile::ClaudeCodeLegacy
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_profile_is_ignored_for_non_anthropic_formats() {
|
||||
let mut transport = sample_transport("claude_code");
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic": {"compatibility_profile": "claude_code_legacy"}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
resolve_anthropic_compatibility_profile(&transport, "openai:chat"),
|
||||
AnthropicCompatibilityProfile::NativeTransparent
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_explicit_profile_fails_closed_instead_of_using_legacy_default() {
|
||||
let mut transport = sample_transport("claude_code");
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic_compatibility": {"profile": "native_transparnt"}
|
||||
}));
|
||||
transport.provider.config = Some(json!({
|
||||
"anthropic_compatibility": {"profile": "claude_code_legacy"}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
resolve_anthropic_compatibility_profile(&transport, "claude:messages"),
|
||||
AnthropicCompatibilityProfile::NativeTransparent
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -1,8 +1,8 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::headers::{
|
||||
normalize_upstream_accept_encoding, should_skip_upstream_complete_passthrough_header,
|
||||
should_skip_upstream_passthrough_header,
|
||||
is_aether_internal_header, is_upstream_credential_header, normalize_upstream_accept_encoding,
|
||||
should_skip_upstream_complete_passthrough_header, should_skip_upstream_passthrough_header,
|
||||
};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
@@ -30,6 +30,9 @@ 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) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
|
||||
continue;
|
||||
};
|
||||
@@ -60,6 +63,9 @@ 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) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
|
||||
continue;
|
||||
};
|
||||
@@ -136,7 +142,7 @@ pub fn build_complete_passthrough_headers_with_auth(
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = build_complete_passthrough_headers(headers, extra_headers, content_type);
|
||||
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
|
||||
replace_upstream_auth_headers(&mut out, auth_header, auth_value);
|
||||
out
|
||||
}
|
||||
|
||||
@@ -155,6 +161,24 @@ pub fn build_claude_passthrough_headers(
|
||||
content_type,
|
||||
);
|
||||
|
||||
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) {
|
||||
continue;
|
||||
}
|
||||
|
||||
if key == "anthropic-beta" {
|
||||
let merged = merge_comma_header_values(out.get(&key).map(String::as_str), Some(value));
|
||||
if let Some(merged) = merged {
|
||||
out.insert(key, merged);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
out.insert(key, value.to_string());
|
||||
}
|
||||
|
||||
for (name, value) in headers.iter() {
|
||||
let Ok(value) = value.to_str() else {
|
||||
continue;
|
||||
@@ -188,7 +212,7 @@ pub fn build_passthrough_headers_with_auth(
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = collect_passthrough_headers(headers, extra_headers);
|
||||
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
|
||||
replace_upstream_auth_headers(&mut out, auth_header, auth_value);
|
||||
out.remove("content-length");
|
||||
out
|
||||
}
|
||||
@@ -200,7 +224,9 @@ pub fn ensure_upstream_auth_header(
|
||||
) {
|
||||
let header_name = auth_header.trim().to_ascii_lowercase();
|
||||
let header_value = auth_value.trim();
|
||||
if header_name.is_empty() || header_value.is_empty() {
|
||||
headers.retain(|name, _| !is_aether_internal_header(name));
|
||||
if header_name.is_empty() || header_value.is_empty() || is_aether_internal_header(&header_name)
|
||||
{
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -213,6 +239,16 @@ pub fn ensure_upstream_auth_header(
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn replace_upstream_auth_headers(
|
||||
headers: &mut BTreeMap<String, String>,
|
||||
auth_header: &str,
|
||||
auth_value: &str,
|
||||
) {
|
||||
headers
|
||||
.retain(|name, _| !is_upstream_credential_header(name) && !is_aether_internal_header(name));
|
||||
ensure_upstream_auth_header(headers, auth_header, auth_value);
|
||||
}
|
||||
|
||||
fn should_restore_claude_passthrough_header(name: &str) -> bool {
|
||||
name.starts_with("anthropic-") || name.starts_with("x-stainless-") || name == "x-app"
|
||||
}
|
||||
@@ -405,7 +441,14 @@ mod tests {
|
||||
&headers,
|
||||
"x-api-key",
|
||||
"sk-upstream-claude",
|
||||
&BTreeMap::from([("anthropic-beta".to_string(), "custom-beta".to_string())]),
|
||||
&BTreeMap::from([
|
||||
("anthropic-beta".to_string(), "custom-beta".to_string()),
|
||||
("authorization".to_string(), "Bearer extra".to_string()),
|
||||
(
|
||||
"x-aether-auth-user-id".to_string(),
|
||||
"user-private".to_string(),
|
||||
),
|
||||
]),
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
@@ -422,6 +465,8 @@ mod tests {
|
||||
Some("v22.14.0")
|
||||
);
|
||||
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
|
||||
assert_eq!(built.get("authorization"), None);
|
||||
assert!(built.keys().all(|name| !name.starts_with("x-aether-")));
|
||||
assert_eq!(
|
||||
built.get("x-api-key").map(String::as_str),
|
||||
Some("sk-upstream-claude")
|
||||
@@ -488,12 +533,34 @@ mod tests {
|
||||
"authorization",
|
||||
http::HeaderValue::from_static("Bearer client-token"),
|
||||
);
|
||||
headers.insert("api-key", http::HeaderValue::from_static("client-api-key"));
|
||||
headers.insert(
|
||||
"x-api-key",
|
||||
http::HeaderValue::from_static("client-x-api-key"),
|
||||
);
|
||||
headers.insert("cookie", http::HeaderValue::from_static("session=client"));
|
||||
headers.insert(
|
||||
"proxy-authorization",
|
||||
http::HeaderValue::from_static("Basic client-proxy"),
|
||||
);
|
||||
headers.insert(
|
||||
"x-aether-auth-user-id",
|
||||
http::HeaderValue::from_static("user-private"),
|
||||
);
|
||||
|
||||
let built = build_complete_passthrough_headers_with_auth(
|
||||
&headers,
|
||||
"x-api-key",
|
||||
"sk-upstream",
|
||||
&BTreeMap::new(),
|
||||
&BTreeMap::from([
|
||||
("authorization".to_string(), "Bearer extra".to_string()),
|
||||
("cookie".to_string(), "session=extra".to_string()),
|
||||
("x-api-key".to_string(), "extra-x-api-key".to_string()),
|
||||
(
|
||||
"x-aether-auth-api-key-id".to_string(),
|
||||
"key-private".to_string(),
|
||||
),
|
||||
]),
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
@@ -507,6 +574,10 @@ mod tests {
|
||||
);
|
||||
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
|
||||
assert_eq!(built.get("authorization"), None);
|
||||
assert_eq!(built.get("api-key"), None);
|
||||
assert_eq!(built.get("cookie"), None);
|
||||
assert_eq!(built.get("proxy-authorization"), None);
|
||||
assert!(built.keys().all(|name| !name.starts_with("x-aether-")));
|
||||
assert_eq!(
|
||||
built.get("x-api-key").map(String::as_str),
|
||||
Some("sk-upstream")
|
||||
|
||||
@@ -2,149 +2,19 @@ use aether_contracts::{
|
||||
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
use sha2::{Digest, Sha256};
|
||||
use uuid::Uuid;
|
||||
|
||||
// Chrome impersonate profiles
|
||||
const CHROME_IMPERSONATE_PROFILES: &[&str] = &[
|
||||
"chrome110",
|
||||
"chrome116",
|
||||
"chrome119",
|
||||
"chrome120",
|
||||
"chrome123",
|
||||
"chrome124",
|
||||
"chrome131",
|
||||
"chrome133",
|
||||
];
|
||||
use super::profile::current_claude_code_transport_identity_profile;
|
||||
|
||||
const CHROME_VERSIONS: &[(&str, &str)] = &[
|
||||
("chrome110", "110.0.5481.177"),
|
||||
("chrome116", "116.0.5845.188"),
|
||||
("chrome119", "119.0.6045.214"),
|
||||
("chrome120", "120.0.6099.216"),
|
||||
("chrome123", "123.0.6312.122"),
|
||||
("chrome124", "124.0.6367.243"),
|
||||
("chrome131", "131.0.6778.265"),
|
||||
("chrome133", "133.0.6943.142"),
|
||||
];
|
||||
|
||||
// (os, arch, platform_token, platform_info)
|
||||
const PLATFORM_VARIANTS: &[(&str, &str, &str, &str)] = &[
|
||||
("Linux", "x64", "X11; Linux x86_64", "Linux x86_64"),
|
||||
("Linux", "arm64", "X11; Linux arm64", "Linux arm64"),
|
||||
(
|
||||
"Windows",
|
||||
"x64",
|
||||
"Windows NT 10.0; Win64; x64",
|
||||
"Windows x64",
|
||||
),
|
||||
(
|
||||
"MacOS",
|
||||
"x64",
|
||||
"Macintosh; Intel Mac OS X 10_15_7",
|
||||
"Darwin x64",
|
||||
),
|
||||
(
|
||||
"MacOS",
|
||||
"arm64",
|
||||
"Macintosh; ARM Mac OS X 14_0_0",
|
||||
"Darwin arm64",
|
||||
),
|
||||
];
|
||||
|
||||
const STAINLESS_PACKAGE_VERSIONS: &[&str] = &["0.68.0", "0.69.0", "0.70.0", "0.71.0"];
|
||||
const NODE_VERSIONS: &[&str] = &["v20.18.1", "v22.12.0", "v22.14.0", "v24.13.0"];
|
||||
const ELECTRON_VERSIONS: &[&str] = &["35.5.1", "36.7.1", "37.3.0", "38.7.0", "39.2.3"];
|
||||
const STAINLESS_TIMEOUTS: &[&str] = &["600", "900"];
|
||||
const CLAUDE_CODE_TRANSPORT_PROFILE_ID: &str = "claude_code_nodejs";
|
||||
|
||||
/// Deterministic hash-based index picker, compatible with Python implementation.
|
||||
/// Each `slot` produces a different selection from the same seed.
|
||||
struct SeededPicker {
|
||||
seed_bytes: [u8; 32],
|
||||
}
|
||||
|
||||
impl SeededPicker {
|
||||
fn new(seed: &str) -> Self {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(seed.as_bytes());
|
||||
Self {
|
||||
seed_bytes: hasher.finalize().into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Pick an index from `[0, len)` using a specific slot.
|
||||
/// Different slots produce independent-looking selections from the same seed.
|
||||
fn pick(&self, slot: u8, len: usize) -> usize {
|
||||
if len == 0 {
|
||||
return 0;
|
||||
}
|
||||
// Hash seed_bytes + slot to get a new digest, take first 8 bytes as u64
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(self.seed_bytes);
|
||||
hasher.update([slot]);
|
||||
let hash = hasher.finalize();
|
||||
let value = u64::from_be_bytes(hash[..8].try_into().unwrap());
|
||||
(value % len as u64) as usize
|
||||
}
|
||||
}
|
||||
|
||||
fn chrome_version_for_profile(profile: &str) -> &'static str {
|
||||
for (p, v) in CHROME_VERSIONS {
|
||||
if p.eq_ignore_ascii_case(profile) {
|
||||
return v;
|
||||
}
|
||||
}
|
||||
"120.0.6099.216"
|
||||
}
|
||||
|
||||
fn build_user_agent(platform_token: &str, chrome_version: &str, electron_version: &str) -> String {
|
||||
format!(
|
||||
"Mozilla/5.0 ({platform_token}) AppleWebKit/537.36 (KHTML, like Gecko) \
|
||||
Chrome/{chrome_version} Electron/{electron_version} Safari/537.36"
|
||||
)
|
||||
}
|
||||
|
||||
fn resolve_platform_token(os: &str, arch: &str) -> &'static str {
|
||||
let os_lower = os.to_ascii_lowercase();
|
||||
let arch_lower = arch.to_ascii_lowercase();
|
||||
|
||||
if os_lower.starts_with("win") {
|
||||
return "Windows NT 10.0; Win64; x64";
|
||||
}
|
||||
if matches!(os_lower.as_str(), "darwin" | "mac" | "macos") {
|
||||
return if matches!(arch_lower.as_str(), "arm64" | "aarch64") {
|
||||
"Macintosh; ARM Mac OS X 14_0_0"
|
||||
} else {
|
||||
"Macintosh; Intel Mac OS X 10_15_7"
|
||||
};
|
||||
}
|
||||
if matches!(arch_lower.as_str(), "arm64" | "aarch64") {
|
||||
"X11; Linux arm64"
|
||||
} else {
|
||||
"X11; Linux x86_64"
|
||||
}
|
||||
}
|
||||
|
||||
/// Generate a complete Claude Code transport fingerprint from a seed.
|
||||
/// Generate the transport-profile metadata and per-key pool partition from a
|
||||
/// stable key seed. HTTP identity headers are owned by the versioned profile
|
||||
/// and are not overridden from this stored value.
|
||||
pub fn generate_fingerprint(seed: &str) -> Value {
|
||||
wrap_header_fingerprint(generate_header_fingerprint(seed))
|
||||
}
|
||||
|
||||
fn generate_header_fingerprint(seed: &str) -> Value {
|
||||
let picker = SeededPicker::new(seed);
|
||||
|
||||
let impersonate =
|
||||
CHROME_IMPERSONATE_PROFILES[picker.pick(0, CHROME_IMPERSONATE_PROFILES.len())];
|
||||
let chrome_version = chrome_version_for_profile(impersonate);
|
||||
let node_version = NODE_VERSIONS[picker.pick(1, NODE_VERSIONS.len())];
|
||||
let electron_version = ELECTRON_VERSIONS[picker.pick(2, ELECTRON_VERSIONS.len())];
|
||||
let platform = PLATFORM_VARIANTS[picker.pick(3, PLATFORM_VARIANTS.len())];
|
||||
let (stainless_os, stainless_arch, platform_token, platform_info) = platform;
|
||||
let stainless_package_version =
|
||||
STAINLESS_PACKAGE_VERSIONS[picker.pick(4, STAINLESS_PACKAGE_VERSIONS.len())];
|
||||
let stainless_timeout = STAINLESS_TIMEOUTS[picker.pick(5, STAINLESS_TIMEOUTS.len())];
|
||||
|
||||
let profile = *current_claude_code_transport_identity_profile();
|
||||
let vscode_session_id = Uuid::new_v5(
|
||||
&Uuid::NAMESPACE_URL,
|
||||
format!("aether:fingerprint:{seed}").as_bytes(),
|
||||
@@ -152,32 +22,36 @@ fn generate_header_fingerprint(seed: &str) -> Value {
|
||||
.simple()
|
||||
.to_string();
|
||||
|
||||
let user_agent = build_user_agent(platform_token, chrome_version, electron_version);
|
||||
|
||||
serde_json::json!({
|
||||
"impersonate": impersonate,
|
||||
"stainless_package_version": stainless_package_version,
|
||||
"stainless_os": stainless_os,
|
||||
"stainless_arch": stainless_arch,
|
||||
"stainless_runtime_version": node_version,
|
||||
"stainless_timeout": stainless_timeout,
|
||||
"node_version": node_version,
|
||||
"chrome_version": chrome_version,
|
||||
"electron_version": electron_version,
|
||||
"identity_profile_version": profile.version().as_str(),
|
||||
"cli_version": profile.cli_version(),
|
||||
"billing_cli_version": profile.billing_cli_version(),
|
||||
"stainless_lang": profile.stainless_lang(),
|
||||
"stainless_package_version": profile.stainless_package_version(),
|
||||
"stainless_os": profile.stainless_os(),
|
||||
"stainless_arch": profile.stainless_arch(),
|
||||
"stainless_runtime": profile.stainless_runtime(),
|
||||
"stainless_runtime_version": profile.stainless_runtime_version(),
|
||||
"stainless_retry_count": profile.stainless_retry_count(),
|
||||
"stainless_timeout": profile.stainless_timeout(),
|
||||
"vscode_session_id": vscode_session_id,
|
||||
"platform_info": platform_info,
|
||||
"user_agent": user_agent,
|
||||
"user_agent": profile.user_agent(),
|
||||
})
|
||||
}
|
||||
|
||||
fn wrap_header_fingerprint(header_fingerprint: Value) -> Value {
|
||||
let profile = *current_claude_code_transport_identity_profile();
|
||||
serde_json::json!({
|
||||
"transport_profile": {
|
||||
"profile_id": CLAUDE_CODE_TRANSPORT_PROFILE_ID,
|
||||
"profile_id": profile.transport_profile_id(),
|
||||
"backend": TRANSPORT_BACKEND_REQWEST_RUSTLS,
|
||||
"http_mode": TRANSPORT_HTTP_MODE_AUTO,
|
||||
"pool_scope": TRANSPORT_POOL_SCOPE_KEY,
|
||||
"header_fingerprint": header_fingerprint,
|
||||
"extra": {
|
||||
"claude_code_identity_profile_version": profile.version().as_str(),
|
||||
"claude_code_cli_version": profile.cli_version(),
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -192,82 +66,31 @@ pub fn header_fingerprint_from_fingerprint(fingerprint: &Value) -> Option<&Map<S
|
||||
|
||||
/// Generate a random (non-deterministic) fingerprint.
|
||||
pub fn generate_random_fingerprint() -> Value {
|
||||
let random_seed = Uuid::new_v4().to_string();
|
||||
generate_fingerprint(&random_seed)
|
||||
generate_fingerprint(&Uuid::new_v4().to_string())
|
||||
}
|
||||
|
||||
/// Sanitize an existing fingerprint JSON, filling missing fields with
|
||||
/// deterministic fallbacks derived from `key_id`.
|
||||
/// Upgrade stored transport metadata to the current typed identity profile.
|
||||
/// Only the per-key pool partition (kept under its legacy session-id field) is
|
||||
/// retained; fixed CLI and Stainless values remain one coherent version set.
|
||||
pub fn sanitize_fingerprint(raw: &Value, key_id: &str) -> Value {
|
||||
let generated = generate_header_fingerprint(key_id);
|
||||
let gen_map = generated.as_object().unwrap();
|
||||
let raw_map = header_fingerprint_from_fingerprint(raw);
|
||||
let mut generated = generate_header_fingerprint(key_id);
|
||||
let Some(generated) = generated.as_object_mut() else {
|
||||
return generate_fingerprint(key_id);
|
||||
};
|
||||
|
||||
let mut out = Map::new();
|
||||
|
||||
// Start with generated values, then overlay non-empty raw values
|
||||
for (key, gen_value) in gen_map {
|
||||
let value = raw_map
|
||||
.and_then(|raw_map| raw_map.get(key))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|v| !v.is_empty())
|
||||
.map(|v| Value::String(v.to_string()))
|
||||
.unwrap_or_else(|| gen_value.clone());
|
||||
out.insert(key.clone(), value);
|
||||
}
|
||||
|
||||
// Normalize impersonate to known profile
|
||||
let impersonate = out
|
||||
.get("impersonate")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default()
|
||||
.to_ascii_lowercase();
|
||||
let is_known = CHROME_IMPERSONATE_PROFILES
|
||||
.iter()
|
||||
.any(|p| p.eq_ignore_ascii_case(&impersonate));
|
||||
if !is_known {
|
||||
out.insert("impersonate".to_string(), gen_map["impersonate"].clone());
|
||||
}
|
||||
|
||||
// Ensure chrome_version matches impersonate profile
|
||||
let profile = out
|
||||
.get("impersonate")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let chrome_version = chrome_version_for_profile(profile);
|
||||
out.insert(
|
||||
"chrome_version".to_string(),
|
||||
Value::String(chrome_version.to_string()),
|
||||
);
|
||||
|
||||
// Rebuild user_agent if missing
|
||||
let has_ua = out
|
||||
.get("user_agent")
|
||||
if let Some(session_id) = header_fingerprint_from_fingerprint(raw)
|
||||
.and_then(|raw| raw.get("vscode_session_id"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|v| !v.is_empty());
|
||||
if !has_ua {
|
||||
let os = out
|
||||
.get("stainless_os")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("Linux");
|
||||
let arch = out
|
||||
.get("stainless_arch")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("x64");
|
||||
let electron = out
|
||||
.get("electron_version")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("38.7.0");
|
||||
let platform_token = resolve_platform_token(os, arch);
|
||||
out.insert(
|
||||
"user_agent".to_string(),
|
||||
Value::String(build_user_agent(platform_token, chrome_version, electron)),
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
generated.insert(
|
||||
"vscode_session_id".to_string(),
|
||||
Value::String(session_id.to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
wrap_header_fingerprint(Value::Object(out))
|
||||
wrap_header_fingerprint(Value::Object(generated.clone()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -282,101 +105,51 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn different_seeds_produce_different_fingerprints() {
|
||||
fn different_seeds_produce_different_session_fingerprints() {
|
||||
let fp1 = generate_fingerprint("key-1");
|
||||
let fp2 = generate_fingerprint("key-2");
|
||||
// At least one field should differ (statistically near-certain)
|
||||
assert_ne!(fp1, fp2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_fingerprint_has_all_fields() {
|
||||
fn generated_fingerprint_matches_current_identity_profile() {
|
||||
let fp = generate_fingerprint("test-key");
|
||||
let expected_keys = [
|
||||
"impersonate",
|
||||
"stainless_package_version",
|
||||
"stainless_os",
|
||||
"stainless_arch",
|
||||
"stainless_runtime_version",
|
||||
"stainless_timeout",
|
||||
"node_version",
|
||||
"chrome_version",
|
||||
"electron_version",
|
||||
"vscode_session_id",
|
||||
"platform_info",
|
||||
"user_agent",
|
||||
];
|
||||
let map = header_fingerprint_from_fingerprint(&fp).unwrap();
|
||||
for key in expected_keys {
|
||||
assert!(map.contains_key(key), "missing field: {key}");
|
||||
let value = map[key].as_str().unwrap();
|
||||
assert!(!value.is_empty(), "empty field: {key}");
|
||||
}
|
||||
}
|
||||
let map = header_fingerprint_from_fingerprint(&fp).expect("header fingerprint");
|
||||
|
||||
#[test]
|
||||
fn sanitize_preserves_user_overrides() {
|
||||
let raw = serde_json::json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "claude_code_nodejs",
|
||||
"header_fingerprint": {
|
||||
"stainless_os": "MacOS",
|
||||
"stainless_arch": "arm64",
|
||||
"stainless_timeout": "900",
|
||||
"user_agent": "Custom-Agent/1.0"
|
||||
}
|
||||
}
|
||||
});
|
||||
let sanitized = sanitize_fingerprint(&raw, "test-key");
|
||||
let map = header_fingerprint_from_fingerprint(&sanitized).unwrap();
|
||||
assert_eq!(map["stainless_os"].as_str(), Some("MacOS"));
|
||||
assert_eq!(map["stainless_arch"].as_str(), Some("arm64"));
|
||||
assert_eq!(map["stainless_timeout"].as_str(), Some("900"));
|
||||
assert_eq!(map["user_agent"].as_str(), Some("Custom-Agent/1.0"));
|
||||
// Other fields should be filled from generation
|
||||
assert!(map.contains_key("impersonate"));
|
||||
assert!(map.contains_key("stainless_package_version"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitize_fills_missing_fields_from_seed() {
|
||||
let raw = serde_json::json!({});
|
||||
let sanitized = sanitize_fingerprint(&raw, "test-key");
|
||||
let generated = generate_header_fingerprint("test-key");
|
||||
// All fields should match generated since raw is empty
|
||||
let s = header_fingerprint_from_fingerprint(&sanitized).unwrap();
|
||||
let g = generated.as_object().unwrap();
|
||||
for key in g.keys() {
|
||||
assert!(s.contains_key(key), "sanitized missing key: {key}");
|
||||
assert!(
|
||||
!s[key].as_str().unwrap().is_empty(),
|
||||
"sanitized empty key: {key}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitize_normalizes_unknown_impersonate_profile() {
|
||||
let raw = serde_json::json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "claude_code_nodejs",
|
||||
"header_fingerprint": {
|
||||
"impersonate": "firefox99"
|
||||
}
|
||||
}
|
||||
});
|
||||
let sanitized = sanitize_fingerprint(&raw, "test-key");
|
||||
let profile = header_fingerprint_from_fingerprint(&sanitized).unwrap()["impersonate"]
|
||||
.as_str()
|
||||
.unwrap();
|
||||
assert!(
|
||||
CHROME_IMPERSONATE_PROFILES
|
||||
.iter()
|
||||
.any(|p| p.eq_ignore_ascii_case(profile)),
|
||||
"should normalize to known profile, got: {profile}"
|
||||
assert_eq!(map["identity_profile_version"], "2026-04");
|
||||
assert_eq!(map["cli_version"], "2.1.161");
|
||||
assert_eq!(map["billing_cli_version"], "2.1.161");
|
||||
assert_eq!(map["stainless_package_version"], "0.94.0");
|
||||
assert_eq!(map["stainless_runtime_version"], "v24.3.0");
|
||||
assert_eq!(map["user_agent"], "claude-cli/2.1.161 (external, cli)");
|
||||
assert_eq!(
|
||||
fp["transport_profile"]["extra"]["claude_code_identity_profile_version"],
|
||||
"2026-04"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitize_upgrades_stale_identity_as_one_version_set() {
|
||||
let raw = serde_json::json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "claude_code_nodejs",
|
||||
"header_fingerprint": {
|
||||
"stainless_package_version": "0.68.0",
|
||||
"stainless_runtime_version": "v20.18.1",
|
||||
"user_agent": "Mozilla/5.0 stale",
|
||||
"vscode_session_id": "existing-session"
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
let sanitized = sanitize_fingerprint(&raw, "test-key");
|
||||
let map = header_fingerprint_from_fingerprint(&sanitized).expect("header fingerprint");
|
||||
assert_eq!(map["stainless_package_version"], "0.94.0");
|
||||
assert_eq!(map["stainless_runtime_version"], "v24.3.0");
|
||||
assert_eq!(map["user_agent"], "claude-cli/2.1.161 (external, cli)");
|
||||
assert_eq!(map["vscode_session_id"], "existing-session");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn random_fingerprint_differs_each_call() {
|
||||
let fp1 = generate_random_fingerprint();
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
mod auth;
|
||||
mod fingerprint;
|
||||
mod policy;
|
||||
mod profile;
|
||||
mod request;
|
||||
mod url;
|
||||
|
||||
@@ -13,5 +14,13 @@ pub use policy::{
|
||||
local_claude_code_transport_unsupported_reason_with_network,
|
||||
supports_local_claude_code_transport_with_network,
|
||||
};
|
||||
pub use request::{build_claude_code_passthrough_headers, sanitize_claude_code_request_body};
|
||||
pub use profile::{
|
||||
current_claude_code_transport_identity_profile, ClaudeCodeBodyCapabilityGate,
|
||||
ClaudeCodeTransportIdentityProfile, ClaudeCodeTransportIdentityProfileVersion,
|
||||
CLAUDE_CODE_CONTEXT_MANAGEMENT_BETA, CLAUDE_CODE_TRANSPORT_IDENTITY_2026_04,
|
||||
};
|
||||
pub use request::{
|
||||
build_claude_code_passthrough_headers, sanitize_claude_code_request_body,
|
||||
sanitize_claude_code_request_body_for_beta_header,
|
||||
};
|
||||
pub use url::build_claude_code_messages_url;
|
||||
|
||||
@@ -0,0 +1,322 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use aether_ai_formats::ApiOperation;
|
||||
|
||||
pub const CLAUDE_CODE_CONTEXT_MANAGEMENT_BETA: &str = "context-management-2025-06-27";
|
||||
|
||||
const MESSAGE_BETAS_2026_04: &[&str] = &[
|
||||
"claude-code-20250219",
|
||||
"oauth-2025-04-20",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
"prompt-caching-scope-2026-01-05",
|
||||
"effort-2025-11-24",
|
||||
CLAUDE_CODE_CONTEXT_MANAGEMENT_BETA,
|
||||
"extended-cache-ttl-2025-04-11",
|
||||
];
|
||||
const COUNT_TOKENS_BETAS_2026_04: &[&str] = &[
|
||||
"claude-code-20250219",
|
||||
"oauth-2025-04-20",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
"prompt-caching-scope-2026-01-05",
|
||||
"effort-2025-11-24",
|
||||
CLAUDE_CODE_CONTEXT_MANAGEMENT_BETA,
|
||||
"extended-cache-ttl-2025-04-11",
|
||||
"token-counting-2024-11-01",
|
||||
];
|
||||
const DROPPED_BETAS_2026_04: &[&str] = &[];
|
||||
const BODY_CAPABILITY_GATES_2026_04: &[ClaudeCodeBodyCapabilityGate] =
|
||||
&[ClaudeCodeBodyCapabilityGate {
|
||||
body_field: "context_management",
|
||||
beta_token: CLAUDE_CODE_CONTEXT_MANAGEMENT_BETA,
|
||||
inject_when_thinking_enabled: true,
|
||||
default_edit_type: Some("clear_thinking_20251015"),
|
||||
}];
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ClaudeCodeTransportIdentityProfileVersion {
|
||||
V2026_04,
|
||||
}
|
||||
|
||||
impl ClaudeCodeTransportIdentityProfileVersion {
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::V2026_04 => "2026-04",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct ClaudeCodeBodyCapabilityGate {
|
||||
pub body_field: &'static str,
|
||||
pub beta_token: &'static str,
|
||||
pub inject_when_thinking_enabled: bool,
|
||||
pub default_edit_type: Option<&'static str>,
|
||||
}
|
||||
|
||||
/// Versioned upstream identity used when Aether is intentionally acting as a
|
||||
/// Claude Code transport. Native Anthropic transports never resolve this
|
||||
/// profile and therefore keep their original headers and body untouched.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct ClaudeCodeTransportIdentityProfile {
|
||||
version: ClaudeCodeTransportIdentityProfileVersion,
|
||||
transport_profile_id: &'static str,
|
||||
anthropic_version: &'static str,
|
||||
cli_version: &'static str,
|
||||
stainless_lang: &'static str,
|
||||
stainless_package_version: &'static str,
|
||||
stainless_os: &'static str,
|
||||
stainless_arch: &'static str,
|
||||
stainless_runtime: &'static str,
|
||||
stainless_runtime_version: &'static str,
|
||||
stainless_retry_count: &'static str,
|
||||
stainless_timeout: &'static str,
|
||||
message_required_betas: &'static [&'static str],
|
||||
count_tokens_required_betas: &'static [&'static str],
|
||||
preserve_incoming_betas: bool,
|
||||
dropped_betas: &'static [&'static str],
|
||||
body_capability_gates: &'static [ClaudeCodeBodyCapabilityGate],
|
||||
}
|
||||
|
||||
pub const CLAUDE_CODE_TRANSPORT_IDENTITY_2026_04: ClaudeCodeTransportIdentityProfile =
|
||||
ClaudeCodeTransportIdentityProfile {
|
||||
version: ClaudeCodeTransportIdentityProfileVersion::V2026_04,
|
||||
transport_profile_id: "claude_code_nodejs",
|
||||
anthropic_version: "2023-06-01",
|
||||
cli_version: "2.1.161",
|
||||
stainless_lang: "js",
|
||||
stainless_package_version: "0.94.0",
|
||||
stainless_os: "Linux",
|
||||
stainless_arch: "arm64",
|
||||
stainless_runtime: "node",
|
||||
stainless_runtime_version: "v24.3.0",
|
||||
stainless_retry_count: "0",
|
||||
stainless_timeout: "600",
|
||||
message_required_betas: MESSAGE_BETAS_2026_04,
|
||||
count_tokens_required_betas: COUNT_TOKENS_BETAS_2026_04,
|
||||
preserve_incoming_betas: true,
|
||||
dropped_betas: DROPPED_BETAS_2026_04,
|
||||
body_capability_gates: BODY_CAPABILITY_GATES_2026_04,
|
||||
};
|
||||
|
||||
pub const fn current_claude_code_transport_identity_profile(
|
||||
) -> &'static ClaudeCodeTransportIdentityProfile {
|
||||
&CLAUDE_CODE_TRANSPORT_IDENTITY_2026_04
|
||||
}
|
||||
|
||||
impl ClaudeCodeTransportIdentityProfile {
|
||||
pub const fn version(self) -> ClaudeCodeTransportIdentityProfileVersion {
|
||||
self.version
|
||||
}
|
||||
|
||||
pub const fn transport_profile_id(self) -> &'static str {
|
||||
self.transport_profile_id
|
||||
}
|
||||
|
||||
pub const fn cli_version(self) -> &'static str {
|
||||
self.cli_version
|
||||
}
|
||||
|
||||
pub const fn billing_cli_version(self) -> &'static str {
|
||||
self.cli_version
|
||||
}
|
||||
|
||||
pub fn user_agent(self) -> String {
|
||||
format!("claude-cli/{} (external, cli)", self.cli_version)
|
||||
}
|
||||
|
||||
pub const fn stainless_package_version(self) -> &'static str {
|
||||
self.stainless_package_version
|
||||
}
|
||||
|
||||
pub const fn stainless_lang(self) -> &'static str {
|
||||
self.stainless_lang
|
||||
}
|
||||
|
||||
pub const fn stainless_os(self) -> &'static str {
|
||||
self.stainless_os
|
||||
}
|
||||
|
||||
pub const fn stainless_arch(self) -> &'static str {
|
||||
self.stainless_arch
|
||||
}
|
||||
|
||||
pub const fn stainless_runtime(self) -> &'static str {
|
||||
self.stainless_runtime
|
||||
}
|
||||
|
||||
pub const fn stainless_runtime_version(self) -> &'static str {
|
||||
self.stainless_runtime_version
|
||||
}
|
||||
|
||||
pub const fn stainless_retry_count(self) -> &'static str {
|
||||
self.stainless_retry_count
|
||||
}
|
||||
|
||||
pub const fn stainless_timeout(self) -> &'static str {
|
||||
self.stainless_timeout
|
||||
}
|
||||
|
||||
pub fn required_beta_tokens(self, operation: Option<ApiOperation>) -> &'static [&'static str] {
|
||||
if operation == Some(ApiOperation::ClaudeCountTokens) {
|
||||
self.count_tokens_required_betas
|
||||
} else {
|
||||
self.message_required_betas
|
||||
}
|
||||
}
|
||||
|
||||
pub const fn preserves_incoming_betas(self) -> bool {
|
||||
self.preserve_incoming_betas
|
||||
}
|
||||
|
||||
pub const fn dropped_beta_tokens(self) -> &'static [&'static str] {
|
||||
self.dropped_betas
|
||||
}
|
||||
|
||||
pub const fn body_capability_gates(self) -> &'static [ClaudeCodeBodyCapabilityGate] {
|
||||
self.body_capability_gates
|
||||
}
|
||||
|
||||
pub fn body_capability_gate(self, field: &str) -> Option<ClaudeCodeBodyCapabilityGate> {
|
||||
self.body_capability_gates
|
||||
.iter()
|
||||
.copied()
|
||||
.find(|gate| gate.body_field == field)
|
||||
}
|
||||
|
||||
pub fn apply_fixed_headers(self, headers: &mut BTreeMap<String, String>, stream: bool) {
|
||||
for (name, value) in [
|
||||
("accept", "application/json"),
|
||||
("anthropic-version", self.anthropic_version),
|
||||
("anthropic-dangerous-direct-browser-access", "true"),
|
||||
("x-app", "cli"),
|
||||
("x-stainless-lang", self.stainless_lang),
|
||||
(
|
||||
"x-stainless-package-version",
|
||||
self.stainless_package_version,
|
||||
),
|
||||
("x-stainless-os", self.stainless_os),
|
||||
("x-stainless-arch", self.stainless_arch),
|
||||
("x-stainless-runtime", self.stainless_runtime),
|
||||
(
|
||||
"x-stainless-runtime-version",
|
||||
self.stainless_runtime_version,
|
||||
),
|
||||
("x-stainless-retry-count", self.stainless_retry_count),
|
||||
("x-stainless-timeout", self.stainless_timeout),
|
||||
] {
|
||||
headers.insert(name.to_string(), value.to_string());
|
||||
}
|
||||
headers.insert("user-agent".to_string(), self.user_agent());
|
||||
if stream {
|
||||
headers.insert(
|
||||
"x-stainless-helper-method".to_string(),
|
||||
"stream".to_string(),
|
||||
);
|
||||
} else {
|
||||
headers.remove("x-stainless-helper-method");
|
||||
}
|
||||
}
|
||||
|
||||
pub fn apply_beta_policy(
|
||||
self,
|
||||
headers: &mut BTreeMap<String, String>,
|
||||
operation: Option<ApiOperation>,
|
||||
) {
|
||||
let incoming = headers.get("anthropic-beta").map(String::as_str);
|
||||
let merged = self.merge_beta_tokens(incoming, operation);
|
||||
if merged.is_empty() {
|
||||
headers.remove("anthropic-beta");
|
||||
} else {
|
||||
headers.insert("anthropic-beta".to_string(), merged);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn merge_beta_tokens(
|
||||
self,
|
||||
incoming: Option<&str>,
|
||||
operation: Option<ApiOperation>,
|
||||
) -> String {
|
||||
let mut seen = BTreeSet::new();
|
||||
let mut merged = Vec::new();
|
||||
|
||||
for token in self.required_beta_tokens(operation) {
|
||||
self.append_beta_token(&mut seen, &mut merged, token);
|
||||
}
|
||||
if self.preserve_incoming_betas {
|
||||
for token in incoming.unwrap_or_default().split(',') {
|
||||
self.append_beta_token(&mut seen, &mut merged, token);
|
||||
}
|
||||
}
|
||||
merged.join(",")
|
||||
}
|
||||
|
||||
pub fn beta_header_enables_body_field(self, beta_header: &str, field: &str) -> bool {
|
||||
let Some(gate) = self.body_capability_gate(field) else {
|
||||
return true;
|
||||
};
|
||||
beta_header
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.any(|token| token.eq_ignore_ascii_case(gate.beta_token))
|
||||
}
|
||||
|
||||
fn append_beta_token(self, seen: &mut BTreeSet<String>, merged: &mut Vec<String>, token: &str) {
|
||||
let token = token.trim();
|
||||
if token.is_empty()
|
||||
|| self
|
||||
.dropped_betas
|
||||
.iter()
|
||||
.any(|dropped| token.eq_ignore_ascii_case(dropped))
|
||||
{
|
||||
return;
|
||||
}
|
||||
if seen.insert(token.to_ascii_lowercase()) {
|
||||
merged.push(token.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::current_claude_code_transport_identity_profile;
|
||||
use aether_ai_formats::ApiOperation;
|
||||
|
||||
#[test]
|
||||
fn profile_versions_cli_user_agent_stainless_and_billing_together() {
|
||||
let profile = *current_claude_code_transport_identity_profile();
|
||||
|
||||
assert_eq!(profile.version().as_str(), "2026-04");
|
||||
assert_eq!(profile.cli_version(), "2.1.161");
|
||||
assert_eq!(profile.billing_cli_version(), profile.cli_version());
|
||||
assert_eq!(
|
||||
profile.user_agent(),
|
||||
format!("claude-cli/{} (external, cli)", profile.cli_version())
|
||||
);
|
||||
assert_eq!(profile.stainless_package_version(), "0.94.0");
|
||||
assert_eq!(profile.stainless_runtime_version(), "v24.3.0");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn profile_preserves_context_1m_and_adds_operation_specific_betas() {
|
||||
let profile = *current_claude_code_transport_identity_profile();
|
||||
let messages = profile.merge_beta_tokens(Some("context-1m-2025-08-07,custom"), None);
|
||||
|
||||
assert!(messages
|
||||
.split(',')
|
||||
.any(|token| token == "context-1m-2025-08-07"));
|
||||
assert!(messages.split(',').any(|token| token == "custom"));
|
||||
assert!(!messages
|
||||
.split(',')
|
||||
.any(|token| token == "token-counting-2024-11-01"));
|
||||
|
||||
let count_tokens = profile.merge_beta_tokens(
|
||||
Some("context-1m-2025-08-07"),
|
||||
Some(ApiOperation::ClaudeCountTokens),
|
||||
);
|
||||
assert!(count_tokens
|
||||
.split(',')
|
||||
.any(|token| token == "token-counting-2024-11-01"));
|
||||
assert!(profile.dropped_beta_tokens().is_empty());
|
||||
assert!(profile.preserves_incoming_betas());
|
||||
}
|
||||
}
|
||||
@@ -1,31 +1,15 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::OnceLock;
|
||||
|
||||
use regex::Regex;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::super::auth::build_openai_passthrough_headers;
|
||||
use super::fingerprint::header_fingerprint_from_fingerprint;
|
||||
use super::profile::{
|
||||
current_claude_code_transport_identity_profile, ClaudeCodeTransportIdentityProfile,
|
||||
};
|
||||
|
||||
const DEFAULT_ANTHROPIC_VERSION: &str = "2023-06-01";
|
||||
const DEFAULT_ACCEPT: &str = "application/json";
|
||||
const STREAM_HELPER_METHOD: &str = "stream";
|
||||
const DUMMY_THINKING_SIGNATURE: &str = "skip_thought_signature_validator";
|
||||
const REQUIRED_BETA_TOKENS: &[&str] = &[
|
||||
"claude-code-20250219",
|
||||
"oauth-2025-04-20",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
];
|
||||
const EXCLUDED_BETA_TOKENS: &[&str] = &["context-1m-2025-08-07"];
|
||||
|
||||
/// Fingerprint field -> HTTP header mapping.
|
||||
/// Every stainless / identity dimension that can vary per-key is listed here.
|
||||
const FINGERPRINT_HEADER_MAP: &[(&str, &str)] = &[
|
||||
("stainless_package_version", "x-stainless-package-version"),
|
||||
("stainless_os", "x-stainless-os"),
|
||||
("stainless_arch", "x-stainless-arch"),
|
||||
("stainless_runtime_version", "x-stainless-runtime-version"),
|
||||
("stainless_timeout", "x-stainless-timeout"),
|
||||
("user_agent", "user-agent"),
|
||||
];
|
||||
|
||||
pub fn build_claude_code_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
@@ -33,7 +17,6 @@ pub fn build_claude_code_passthrough_headers(
|
||||
auth_value: &str,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
stream: bool,
|
||||
fingerprint: Option<&Value>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = build_openai_passthrough_headers(
|
||||
headers,
|
||||
@@ -43,77 +26,58 @@ pub fn build_claude_code_passthrough_headers(
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
// -- Anthropic protocol headers --
|
||||
out.insert("accept".to_string(), DEFAULT_ACCEPT.to_string());
|
||||
out.insert(
|
||||
"anthropic-version".to_string(),
|
||||
DEFAULT_ANTHROPIC_VERSION.to_string(),
|
||||
);
|
||||
// Read incoming anthropic-beta directly from the original HeaderMap because the
|
||||
// upstream passthrough filter now strips `anthropic-*` headers to avoid leaking
|
||||
// them to non-Anthropic upstreams.
|
||||
let incoming_anthropic_beta = headers
|
||||
// The common passthrough filter intentionally strips Anthropic identity
|
||||
// headers. Restore only the client beta input; the versioned profile owns
|
||||
// all fixed identity values and the final beta policy.
|
||||
let mut incoming_beta_values = headers
|
||||
.get("anthropic-beta")
|
||||
.and_then(|value| value.to_str().ok());
|
||||
out.insert(
|
||||
"anthropic-beta".to_string(),
|
||||
merge_anthropic_beta_tokens(incoming_anthropic_beta),
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.into_iter()
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<Vec<_>>();
|
||||
incoming_beta_values.extend(
|
||||
extra_headers
|
||||
.iter()
|
||||
.filter(|(name, _)| name.eq_ignore_ascii_case("anthropic-beta"))
|
||||
.map(|(_, value)| value.trim())
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
);
|
||||
out.insert(
|
||||
"anthropic-dangerous-direct-browser-access".to_string(),
|
||||
"true".to_string(),
|
||||
);
|
||||
out.insert("x-app".to_string(), "cli".to_string());
|
||||
|
||||
// -- Stainless SDK identity headers --
|
||||
// Fixed values: these don't vary per fingerprint.
|
||||
out.insert("x-stainless-lang".to_string(), "js".to_string());
|
||||
out.insert("x-stainless-runtime".to_string(), "node".to_string());
|
||||
out.insert("x-stainless-retry-count".to_string(), "0".to_string());
|
||||
|
||||
// Defaults for fingerprint-overridable fields (used when no fingerprint is present).
|
||||
out.insert(
|
||||
"x-stainless-package-version".to_string(),
|
||||
"0.70.0".to_string(),
|
||||
);
|
||||
out.insert("x-stainless-os".to_string(), "Linux".to_string());
|
||||
out.insert("x-stainless-arch".to_string(), "arm64".to_string());
|
||||
out.insert(
|
||||
"x-stainless-runtime-version".to_string(),
|
||||
"v24.13.0".to_string(),
|
||||
);
|
||||
out.insert("x-stainless-timeout".to_string(), "600".to_string());
|
||||
|
||||
if stream {
|
||||
out.insert(
|
||||
"x-stainless-helper-method".to_string(),
|
||||
STREAM_HELPER_METHOD.to_string(),
|
||||
);
|
||||
} else {
|
||||
out.remove("x-stainless-helper-method");
|
||||
if !incoming_beta_values.is_empty() {
|
||||
out.insert("anthropic-beta".to_string(), incoming_beta_values.join(","));
|
||||
}
|
||||
|
||||
// Override from the formal transport profile header fingerprint.
|
||||
if let Some(fp) = fingerprint.and_then(header_fingerprint_from_fingerprint) {
|
||||
for &(fp_key, header_key) in FINGERPRINT_HEADER_MAP {
|
||||
if let Some(value) = fp
|
||||
.get(fp_key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|v| !v.is_empty())
|
||||
{
|
||||
out.insert(header_key.to_string(), value.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
let profile = *current_claude_code_transport_identity_profile();
|
||||
profile.apply_fixed_headers(&mut out, stream);
|
||||
profile.apply_beta_policy(&mut out, None);
|
||||
|
||||
out
|
||||
}
|
||||
|
||||
pub fn sanitize_claude_code_request_body(body: &mut Value) {
|
||||
let profile = *current_claude_code_transport_identity_profile();
|
||||
let beta_header = profile.merge_beta_tokens(None, None);
|
||||
sanitize_claude_code_request_body_for_beta_header(body, &beta_header, profile);
|
||||
}
|
||||
|
||||
pub fn sanitize_claude_code_request_body_for_beta_header(
|
||||
body: &mut Value,
|
||||
beta_header: &str,
|
||||
profile: ClaudeCodeTransportIdentityProfile,
|
||||
) {
|
||||
let Some(body_object) = body.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
|
||||
synchronize_billing_header_version(body_object, profile.billing_cli_version());
|
||||
for gate in profile.body_capability_gates() {
|
||||
if !profile.beta_header_enables_body_field(beta_header, gate.body_field) {
|
||||
body_object.remove(gate.body_field);
|
||||
}
|
||||
}
|
||||
|
||||
let thinking_enabled = body_object
|
||||
.get("thinking")
|
||||
.and_then(Value::as_object)
|
||||
@@ -122,6 +86,26 @@ pub fn sanitize_claude_code_request_body(body: &mut Value) {
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| matches!(value.to_ascii_lowercase().as_str(), "enabled" | "adaptive"));
|
||||
|
||||
if let Some(gate) = profile.body_capability_gate("context_management") {
|
||||
if thinking_enabled
|
||||
&& gate.inject_when_thinking_enabled
|
||||
&& profile.beta_header_enables_body_field(beta_header, gate.body_field)
|
||||
&& !body_object.contains_key(gate.body_field)
|
||||
{
|
||||
if let Some(default_edit_type) = gate.default_edit_type {
|
||||
body_object.insert(
|
||||
gate.body_field.to_string(),
|
||||
serde_json::json!({
|
||||
"edits": [{
|
||||
"type": default_edit_type,
|
||||
"keep": "all"
|
||||
}]
|
||||
}),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let Some(messages) = body_object
|
||||
.get_mut("messages")
|
||||
.and_then(Value::as_array_mut)
|
||||
@@ -160,6 +144,37 @@ pub fn sanitize_claude_code_request_body(body: &mut Value) {
|
||||
}
|
||||
}
|
||||
|
||||
fn synchronize_billing_header_version(body: &mut Map<String, Value>, cli_version: &str) {
|
||||
static CC_VERSION: OnceLock<Regex> = OnceLock::new();
|
||||
let Some(system) = body.get_mut("system").and_then(Value::as_array_mut) else {
|
||||
return;
|
||||
};
|
||||
let regex = CC_VERSION.get_or_init(|| {
|
||||
Regex::new(r"cc_version=\d+\.\d+\.\d+").expect("billing version regex must compile")
|
||||
});
|
||||
let replacement = format!("cc_version={cli_version}");
|
||||
|
||||
for block in system {
|
||||
let Some(block) = block.as_object_mut() else {
|
||||
continue;
|
||||
};
|
||||
let Some(text) = block
|
||||
.get("text")
|
||||
.and_then(Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if !text.starts_with("x-anthropic-billing-header") {
|
||||
continue;
|
||||
}
|
||||
let updated = regex.replace_all(&text, replacement.as_str());
|
||||
if updated != text {
|
||||
block.insert("text".to_string(), Value::String(updated.into_owned()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn keep_claude_code_block(
|
||||
block_object: &Map<String, Value>,
|
||||
role: &str,
|
||||
@@ -187,46 +202,18 @@ fn keep_claude_code_block(
|
||||
true
|
||||
}
|
||||
|
||||
fn merge_anthropic_beta_tokens(incoming: Option<&str>) -> String {
|
||||
let mut seen = BTreeSet::new();
|
||||
let mut merged = Vec::new();
|
||||
|
||||
for token in REQUIRED_BETA_TOKENS {
|
||||
append_beta_token(&mut seen, &mut merged, token);
|
||||
}
|
||||
for token in incoming.unwrap_or_default().split(',') {
|
||||
let token = token.trim();
|
||||
if EXCLUDED_BETA_TOKENS
|
||||
.iter()
|
||||
.any(|excluded| token.eq_ignore_ascii_case(excluded))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
append_beta_token(&mut seen, &mut merged, token);
|
||||
}
|
||||
|
||||
merged.join(",")
|
||||
}
|
||||
|
||||
fn append_beta_token(seen: &mut BTreeSet<String>, merged: &mut Vec<String>, token: &str) {
|
||||
let normalized = token.trim();
|
||||
if normalized.is_empty() {
|
||||
return;
|
||||
}
|
||||
let key = normalized.to_ascii_lowercase();
|
||||
if seen.insert(key) {
|
||||
merged.push(normalized.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{build_claude_code_passthrough_headers, sanitize_claude_code_request_body};
|
||||
use super::{
|
||||
build_claude_code_passthrough_headers, sanitize_claude_code_request_body,
|
||||
sanitize_claude_code_request_body_for_beta_header,
|
||||
};
|
||||
use crate::claude_code::current_claude_code_transport_identity_profile;
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn claude_code_headers_use_transport_profile_header_fingerprint_and_merge_required_betas() {
|
||||
fn claude_code_headers_use_versioned_identity_and_merge_preserved_betas() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert(
|
||||
"anthropic-beta",
|
||||
@@ -242,23 +229,12 @@ mod tests {
|
||||
"Bearer upstream-token",
|
||||
&BTreeMap::new(),
|
||||
true,
|
||||
Some(&json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "claude_code_nodejs",
|
||||
"header_fingerprint": {
|
||||
"user_agent":"Claude-Code/9.9",
|
||||
"stainless_package_version":"1.0.5",
|
||||
"stainless_runtime_version":"v22.12.0",
|
||||
"stainless_timeout":"900"
|
||||
}
|
||||
}
|
||||
})),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
built.get("anthropic-beta").map(String::as_str),
|
||||
Some(
|
||||
"claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14,custom-beta"
|
||||
"claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14,prompt-caching-scope-2026-01-05,effort-2025-11-24,context-management-2025-06-27,extended-cache-ttl-2025-04-11,context-1m-2025-08-07,custom-beta"
|
||||
)
|
||||
);
|
||||
assert_eq!(
|
||||
@@ -276,19 +252,19 @@ mod tests {
|
||||
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
|
||||
assert_eq!(
|
||||
built.get("x-stainless-package-version").map(String::as_str),
|
||||
Some("1.0.5")
|
||||
Some("0.94.0")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("x-stainless-runtime-version").map(String::as_str),
|
||||
Some("v22.12.0")
|
||||
Some("v24.3.0")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("x-stainless-timeout").map(String::as_str),
|
||||
Some("900")
|
||||
Some("600")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("user-agent").map(String::as_str),
|
||||
Some("Claude-Code/9.9")
|
||||
Some("claude-cli/2.1.161 (external, cli)")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("authorization").map(String::as_str),
|
||||
@@ -324,4 +300,74 @@ mod tests {
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn context_management_body_is_gated_by_the_matching_beta_token() {
|
||||
let profile = *current_claude_code_transport_identity_profile();
|
||||
let original = json!({
|
||||
"context_management": {
|
||||
"edits": [{"type":"clear_thinking_20251015", "keep":"all"}]
|
||||
},
|
||||
"messages": []
|
||||
});
|
||||
|
||||
let mut without_beta = original.clone();
|
||||
sanitize_claude_code_request_body_for_beta_header(
|
||||
&mut without_beta,
|
||||
"oauth-2025-04-20",
|
||||
profile,
|
||||
);
|
||||
assert!(without_beta.get("context_management").is_none());
|
||||
|
||||
let mut with_beta = original.clone();
|
||||
sanitize_claude_code_request_body_for_beta_header(
|
||||
&mut with_beta,
|
||||
"oauth-2025-04-20, context-management-2025-06-27",
|
||||
profile,
|
||||
);
|
||||
assert_eq!(with_beta, original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_profile_keeps_header_and_injected_context_management_in_sync() {
|
||||
let headers = build_claude_code_passthrough_headers(
|
||||
&http::HeaderMap::new(),
|
||||
"authorization",
|
||||
"Bearer upstream-token",
|
||||
&BTreeMap::new(),
|
||||
false,
|
||||
);
|
||||
let mut body = json!({
|
||||
"thinking": {"type":"adaptive"},
|
||||
"messages": []
|
||||
});
|
||||
|
||||
sanitize_claude_code_request_body(&mut body);
|
||||
|
||||
assert!(headers["anthropic-beta"]
|
||||
.split(',')
|
||||
.any(|token| token == "context-management-2025-06-27"));
|
||||
assert_eq!(
|
||||
body["context_management"],
|
||||
json!({"edits":[{"type":"clear_thinking_20251015", "keep":"all"}]})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn billing_attribution_uses_the_profile_cli_version() {
|
||||
let mut body = json!({
|
||||
"system": [{
|
||||
"type":"text",
|
||||
"text":"x-anthropic-billing-header: cc_version=2.0.0.abc; cc_entrypoint=cli;"
|
||||
}],
|
||||
"messages": []
|
||||
});
|
||||
|
||||
sanitize_claude_code_request_body(&mut body);
|
||||
|
||||
assert_eq!(
|
||||
body["system"][0]["text"],
|
||||
"x-anthropic-billing-header: cc_version=2.1.161.abc; cc_entrypoint=cli;"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,14 +5,19 @@ use url::form_urlencoded;
|
||||
pub fn build_claude_code_messages_url(upstream_base_url: &str, query: Option<&str>) -> String {
|
||||
let (trimmed_base_url, base_query) = split_query(upstream_base_url.trim());
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
let mut url =
|
||||
if trimmed_base_url.ends_with("/v1/messages") || trimmed_base_url.ends_with("/messages") {
|
||||
trimmed_base_url.to_string()
|
||||
} else if trimmed_base_url.ends_with("/v1") {
|
||||
format!("{trimmed_base_url}/messages")
|
||||
} else {
|
||||
format!("{trimmed_base_url}/v1/messages")
|
||||
};
|
||||
let mut url = if trimmed_base_url.ends_with("/messages/count_tokens") {
|
||||
trimmed_base_url
|
||||
.strip_suffix("/count_tokens")
|
||||
.unwrap_or(trimmed_base_url)
|
||||
.to_string()
|
||||
} else if trimmed_base_url.ends_with("/v1/messages") || trimmed_base_url.ends_with("/messages")
|
||||
{
|
||||
trimmed_base_url.to_string()
|
||||
} else if trimmed_base_url.ends_with("/v1") {
|
||||
format!("{trimmed_base_url}/messages")
|
||||
} else {
|
||||
format!("{trimmed_base_url}/v1/messages")
|
||||
};
|
||||
append_merged_query(&mut url, base_query, query);
|
||||
url
|
||||
}
|
||||
@@ -66,6 +71,13 @@ mod tests {
|
||||
build_claude_code_messages_url("https://api.anthropic.com/v1/messages", None),
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_code_messages_url(
|
||||
"https://api.anthropic.com/v1/messages/count_tokens?tenant=base",
|
||||
Some("trace=1"),
|
||||
),
|
||||
"https://api.anthropic.com/v1/messages?tenant=base&trace=1"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -26,6 +26,42 @@ pub fn supports_local_generic_oauth_request_auth_resolution(
|
||||
&& generic_provider_type(transport.provider.provider_type.as_str()).is_some()
|
||||
}
|
||||
|
||||
pub fn resolve_local_generic_oauth_transport_authorization(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<String> {
|
||||
if !supports_local_generic_oauth_request_auth_resolution(transport) {
|
||||
return None;
|
||||
}
|
||||
if let Some(value) =
|
||||
auth_config_authorization_header(transport.key.decrypted_auth_config.as_deref())
|
||||
{
|
||||
return if bearer_access_token(&value).is_some() {
|
||||
Some(value)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
}
|
||||
|
||||
let auth_config = GenericOAuthRefreshAdapter::auth_config_from_transport(transport);
|
||||
let refreshable = auth_config
|
||||
.as_ref()
|
||||
.and_then(refresh_token_from_auth_config)
|
||||
.is_some();
|
||||
if refreshable && auth_config_expires_soon(auth_config.as_ref()) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if !secret.is_empty() && secret != PLACEHOLDER_API_KEY {
|
||||
return Some(format!("Bearer {secret}"));
|
||||
}
|
||||
|
||||
auth_config
|
||||
.as_ref()
|
||||
.and_then(access_token_from_auth_config)
|
||||
.map(|token| format!("Bearer {token}"))
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct GenericOAuthRefreshAdapter {
|
||||
token_url_overrides: BTreeMap<String, String>,
|
||||
@@ -114,48 +150,23 @@ impl GenericOAuthRefreshAdapter {
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
if !supports_local_generic_oauth_request_auth_resolution(transport) {
|
||||
return None;
|
||||
}
|
||||
|
||||
if let Some(value) =
|
||||
auth_config_authorization_header(transport.key.decrypted_auth_config.as_deref())
|
||||
{
|
||||
return Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: AUTH_HEADER_NAME.to_string(),
|
||||
value,
|
||||
});
|
||||
}
|
||||
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if secret.is_empty() || secret == PLACEHOLDER_API_KEY {
|
||||
return None;
|
||||
}
|
||||
|
||||
let auth_config = Self::auth_config_from_transport(transport);
|
||||
let refreshable = auth_config
|
||||
.as_ref()
|
||||
.and_then(refresh_token_from_auth_config)
|
||||
.is_some();
|
||||
if refreshable && auth_config_expires_soon(auth_config.as_ref()) {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: AUTH_HEADER_NAME.to_string(),
|
||||
value: format!("Bearer {secret}"),
|
||||
value: resolve_local_generic_oauth_transport_authorization(transport)?,
|
||||
})
|
||||
}
|
||||
|
||||
fn build_cached_entry(
|
||||
provider_type: &'static str,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
refreshed: ProviderOAuthTokenSet,
|
||||
mut refreshed: ProviderOAuthTokenSet,
|
||||
) -> CachedOAuthEntry {
|
||||
let auth_header_value = refreshed.token_set.bearer_header_value();
|
||||
synchronize_authorization_overrides(&mut refreshed.auth_config, &auth_header_value);
|
||||
CachedOAuthEntry {
|
||||
provider_type: provider_type.to_string(),
|
||||
auth_header_name: AUTH_HEADER_NAME.to_string(),
|
||||
auth_header_value: refreshed.token_set.bearer_header_value(),
|
||||
auth_header_value,
|
||||
expires_at_unix_secs: refreshed.token_set.expires_at_unix_secs,
|
||||
metadata: Some(refreshed.auth_config),
|
||||
source_fingerprint: Some(generic_oauth_transport_source_fingerprint(transport)),
|
||||
@@ -184,12 +195,14 @@ impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
|
||||
{
|
||||
return None;
|
||||
}
|
||||
if let Some(value) =
|
||||
auth_config_authorization_header(transport.key.decrypted_auth_config.as_deref())
|
||||
if auth_config_authorization_header(transport.key.decrypted_auth_config.as_deref())
|
||||
.is_some()
|
||||
{
|
||||
return Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: AUTH_HEADER_NAME.to_string(),
|
||||
value,
|
||||
return resolve_local_generic_oauth_transport_authorization(transport).map(|value| {
|
||||
LocalResolvedOAuthRequestAuth::Header {
|
||||
name: AUTH_HEADER_NAME.to_string(),
|
||||
value,
|
||||
}
|
||||
});
|
||||
}
|
||||
if !generic_oauth_cached_entry_matches_transport(transport, entry) {
|
||||
@@ -211,6 +224,26 @@ impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_fenced_cached(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
generic_oauth_successor_entry_matches_transport(transport, entry)
|
||||
.then(|| resolved_entry_header(entry))
|
||||
.flatten()
|
||||
}
|
||||
|
||||
fn resolve_refreshed(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
generic_oauth_entry_belongs_to_transport(transport, entry)
|
||||
.then(|| resolved_entry_header(entry))
|
||||
.flatten()
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
@@ -240,6 +273,39 @@ impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
|
||||
.is_some()
|
||||
}
|
||||
|
||||
fn refresh_fingerprint(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
generic_oauth_refresh_fingerprint(transport, entry)
|
||||
}
|
||||
|
||||
fn cached_entry_from_transport(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<CachedOAuthEntry> {
|
||||
let provider_type = generic_provider_type(transport.provider.provider_type.as_str())?;
|
||||
let LocalResolvedOAuthRequestAuth::Header { name, value } =
|
||||
self.resolve_direct_header(transport)?
|
||||
else {
|
||||
return None;
|
||||
};
|
||||
let metadata = Self::auth_config_from_transport(transport);
|
||||
let expires_at_unix_secs = transport
|
||||
.key
|
||||
.expires_at_unix_secs
|
||||
.or_else(|| metadata.as_ref().and_then(auth_config_expires_at));
|
||||
Some(CachedOAuthEntry {
|
||||
provider_type: provider_type.to_string(),
|
||||
auth_header_name: name,
|
||||
auth_header_value: value,
|
||||
expires_at_unix_secs,
|
||||
metadata,
|
||||
source_fingerprint: Some(generic_oauth_transport_source_fingerprint(transport)),
|
||||
})
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
@@ -311,20 +377,33 @@ impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
|
||||
fn generic_oauth_transport_source_fingerprint(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> String {
|
||||
let provider_type = transport.provider.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
|
||||
let auth_config = transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.unwrap_or_default();
|
||||
let api_key = transport.key.decrypted_api_key.as_str();
|
||||
generic_oauth_credential_fingerprint(
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.key.auth_type.as_str(),
|
||||
auth_config,
|
||||
transport.key.decrypted_api_key.as_str(),
|
||||
)
|
||||
}
|
||||
|
||||
fn generic_oauth_credential_fingerprint(
|
||||
provider_type: &str,
|
||||
auth_type: &str,
|
||||
auth_config: &str,
|
||||
access_token: &str,
|
||||
) -> String {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = auth_type.trim().to_ascii_lowercase();
|
||||
let mut digest = Sha256::new();
|
||||
for field in [
|
||||
provider_type.as_bytes(),
|
||||
auth_type.as_bytes(),
|
||||
auth_config.as_bytes(),
|
||||
api_key.as_bytes(),
|
||||
access_token.as_bytes(),
|
||||
] {
|
||||
digest.update((field.len() as u64).to_be_bytes());
|
||||
digest.update(field);
|
||||
@@ -340,6 +419,91 @@ fn generic_oauth_cached_entry_matches_transport(
|
||||
entry.source_fingerprint.as_deref() == Some(transport_fingerprint.as_str())
|
||||
}
|
||||
|
||||
fn generic_oauth_entry_belongs_to_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> bool {
|
||||
entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(transport.provider.provider_type.as_str())
|
||||
&& generic_oauth_cached_entry_matches_transport(transport, entry)
|
||||
}
|
||||
|
||||
fn generic_oauth_successor_entry_matches_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> bool {
|
||||
generic_oauth_entry_belongs_to_transport(transport, entry)
|
||||
&& !expires_at_requires_refresh(entry.expires_at_unix_secs)
|
||||
&& resolved_entry_header(entry).is_some()
|
||||
&& entry.metadata.is_some()
|
||||
}
|
||||
|
||||
fn generic_oauth_refresh_fingerprint(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
if !supports_local_generic_oauth_request_auth_resolution(transport) {
|
||||
return None;
|
||||
}
|
||||
entry
|
||||
.filter(|entry| generic_oauth_successor_entry_matches_transport(transport, entry))
|
||||
.and_then(|entry| {
|
||||
let metadata = serde_json::to_string(entry.metadata.as_ref()?).ok()?;
|
||||
let access_token = bearer_access_token(entry.auth_header_value.as_str())?;
|
||||
Some(generic_oauth_credential_fingerprint(
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.key.auth_type.as_str(),
|
||||
metadata.as_str(),
|
||||
access_token,
|
||||
))
|
||||
})
|
||||
.or_else(|| Some(generic_oauth_transport_source_fingerprint(transport)))
|
||||
}
|
||||
|
||||
fn resolved_entry_header(entry: &CachedOAuthEntry) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
let name = entry.auth_header_name.trim();
|
||||
let value = entry.auth_header_value.trim();
|
||||
if name.is_empty() || value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: name.to_ascii_lowercase(),
|
||||
value: value.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
fn bearer_access_token(authorization: &str) -> Option<&str> {
|
||||
let mut parts = authorization.split_ascii_whitespace();
|
||||
let scheme = parts.next()?;
|
||||
let token = parts.next()?;
|
||||
(scheme.eq_ignore_ascii_case("bearer") && parts.next().is_none()).then_some(token)
|
||||
}
|
||||
|
||||
fn synchronize_authorization_overrides(auth_config: &mut Value, authorization: &str) {
|
||||
let Some(object) = auth_config.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
for (key, value) in object.iter_mut() {
|
||||
match key.trim().to_ascii_lowercase().as_str() {
|
||||
"headers" | "extra_headers" | "extraheaders" => {
|
||||
let Some(headers) = value.as_object_mut() else {
|
||||
continue;
|
||||
};
|
||||
for (header_name, header_value) in headers.iter_mut() {
|
||||
if header_name.trim().eq_ignore_ascii_case(AUTH_HEADER_NAME) {
|
||||
*header_value = Value::String(authorization.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
"transport" | "request" => {
|
||||
synchronize_authorization_overrides(value, authorization);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn generic_provider_type(provider_type: &str) -> Option<&'static str> {
|
||||
let normalized = provider_type.trim();
|
||||
GENERIC_PROVIDER_OAUTH_TEMPLATES
|
||||
@@ -355,6 +519,13 @@ fn refresh_token_from_auth_config(auth_config: &Value) -> Option<String> {
|
||||
.and_then(non_empty_string)
|
||||
}
|
||||
|
||||
fn access_token_from_auth_config(auth_config: &Value) -> Option<String> {
|
||||
let object = auth_config.as_object()?;
|
||||
["access_token", "accessToken"]
|
||||
.iter()
|
||||
.find_map(|field| object.get(*field).and_then(non_empty_string))
|
||||
}
|
||||
|
||||
fn auth_config_expires_at(auth_config: &Value) -> Option<u64> {
|
||||
auth_config
|
||||
.as_object()
|
||||
@@ -423,10 +594,16 @@ fn current_access_token(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::oauth_refresh::{
|
||||
CachedOAuthEntry, LocalOAuthRefreshAdapter, LocalResolvedOAuthRequestAuth,
|
||||
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthHttpRequest, LocalOAuthHttpResponse,
|
||||
LocalOAuthRefreshAdapter, LocalOAuthRefreshCoordinator, LocalOAuthRefreshError,
|
||||
LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
@@ -434,9 +611,35 @@ mod tests {
|
||||
};
|
||||
use super::{
|
||||
current_access_token, generic_oauth_transport_source_fingerprint,
|
||||
GenericOAuthRefreshAdapter,
|
||||
resolve_local_generic_oauth_transport_authorization, GenericOAuthRefreshAdapter,
|
||||
};
|
||||
|
||||
#[derive(Debug)]
|
||||
struct StaticTokenExecutor {
|
||||
hits: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthHttpExecutor for StaticTokenExecutor {
|
||||
async fn execute(
|
||||
&self,
|
||||
_provider_type: &'static str,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
_request: &LocalOAuthHttpRequest,
|
||||
) -> Result<LocalOAuthHttpResponse, LocalOAuthRefreshError> {
|
||||
self.hits.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(LocalOAuthHttpResponse {
|
||||
status_code: 200,
|
||||
body_text: json!({
|
||||
"access_token": "fresh-access-token",
|
||||
"expires_in": 3600,
|
||||
"token_type": "Bearer"
|
||||
})
|
||||
.to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
@@ -519,6 +722,58 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_authorization_uses_one_effective_bearer_generation() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "api-access-token".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"accessToken": "legacy-access-token",
|
||||
"request": {
|
||||
"extraHeaders": {
|
||||
"Authorization": "Bearer nested-override-token"
|
||||
}
|
||||
}
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_local_generic_oauth_transport_authorization(&transport).as_deref(),
|
||||
Some("Bearer nested-override-token")
|
||||
);
|
||||
|
||||
transport.key.decrypted_auth_config =
|
||||
Some(json!({"accessToken": "legacy-access-token"}).to_string());
|
||||
assert_eq!(
|
||||
resolve_local_generic_oauth_transport_authorization(&transport).as_deref(),
|
||||
Some("Bearer api-access-token")
|
||||
);
|
||||
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
assert_eq!(
|
||||
resolve_local_generic_oauth_transport_authorization(&transport).as_deref(),
|
||||
Some("Bearer legacy-access-token")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_bearer_authorization_override_does_not_fall_back_to_stale_token() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "api-access-token".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"access_token": "legacy-access-token",
|
||||
"headers": {"Authorization": "Basic imported-session"}
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(resolve_local_generic_oauth_transport_authorization(&transport).is_none());
|
||||
assert!(GenericOAuthRefreshAdapter::default()
|
||||
.resolve_without_refresh(&transport)
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_config_authorization_header_overrides_cached_oauth_entry() {
|
||||
let adapter = GenericOAuthRefreshAdapter::default();
|
||||
@@ -636,4 +891,95 @@ mod tests {
|
||||
Some("refreshed-access-a")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fenced_force_reuses_successor_when_refresh_token_does_not_rotate() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "stale-access-token".to_string();
|
||||
transport.key.expires_at_unix_secs = Some(u64::MAX);
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "stable-refresh-token",
|
||||
"expires_at": u64::MAX,
|
||||
"headers": {"Authorization": "Bearer stale-top-level"},
|
||||
"request": {
|
||||
"extraHeaders": {"authorization": "Bearer stale-request"}
|
||||
},
|
||||
"transport": {
|
||||
"extra_headers": {"AUTHORIZATION": "Bearer stale-transport"}
|
||||
}
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
let hits = Arc::new(AtomicUsize::new(0));
|
||||
let coordinator = LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(
|
||||
GenericOAuthRefreshAdapter::default()
|
||||
.with_token_url_for_tests("codex", "https://oauth.example/token"),
|
||||
)]);
|
||||
let executor = StaticTokenExecutor {
|
||||
hits: Arc::clone(&hits),
|
||||
};
|
||||
let expected = coordinator
|
||||
.refresh_fingerprint_for_transport(&transport)
|
||||
.expect("refreshable transport should have a generation fence");
|
||||
|
||||
let first = coordinator
|
||||
.force_refresh_with_result_fenced(
|
||||
&executor,
|
||||
&transport,
|
||||
None,
|
||||
None,
|
||||
Some(expected.as_str()),
|
||||
)
|
||||
.await
|
||||
.expect("first refresh should succeed")
|
||||
.expect("first refresh should resolve");
|
||||
let first_entry = first
|
||||
.refreshed_entry
|
||||
.as_ref()
|
||||
.expect("first refresh should return a cache entry");
|
||||
assert_eq!(first_entry.auth_header_value, "Bearer fresh-access-token");
|
||||
let metadata = first_entry
|
||||
.metadata
|
||||
.as_ref()
|
||||
.expect("generic refresh should preserve auth metadata");
|
||||
assert_eq!(metadata["refresh_token"], "stable-refresh-token");
|
||||
assert_eq!(
|
||||
metadata["headers"]["Authorization"],
|
||||
"Bearer fresh-access-token"
|
||||
);
|
||||
assert_eq!(
|
||||
metadata["request"]["extraHeaders"]["authorization"],
|
||||
"Bearer fresh-access-token"
|
||||
);
|
||||
assert_eq!(
|
||||
metadata["transport"]["extra_headers"]["AUTHORIZATION"],
|
||||
"Bearer fresh-access-token"
|
||||
);
|
||||
coordinator
|
||||
.store_cached_entry(&transport.key.id, first_entry.clone())
|
||||
.await;
|
||||
|
||||
let follower = coordinator
|
||||
.force_refresh_with_result_fenced(
|
||||
&executor,
|
||||
&transport,
|
||||
None,
|
||||
None,
|
||||
Some(expected.as_str()),
|
||||
)
|
||||
.await
|
||||
.expect("follower should reuse the completed refresh")
|
||||
.expect("follower should resolve");
|
||||
assert!(follower.reused_refresh);
|
||||
assert_eq!(
|
||||
follower
|
||||
.refreshed_entry
|
||||
.as_ref()
|
||||
.map(|entry| entry.auth_header_value.as_str()),
|
||||
Some("Bearer fresh-access-token")
|
||||
);
|
||||
assert_eq!(hits.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,8 +2,37 @@ use std::collections::BTreeMap;
|
||||
|
||||
use aether_contracts::USAGE_SERVER_NOW_UNIX_MS_HEADER;
|
||||
|
||||
const UPSTREAM_CREDENTIAL_HEADER_NAMES: &[&str] = &[
|
||||
"authorization",
|
||||
"proxy-authorization",
|
||||
"api-key",
|
||||
"x-api-key",
|
||||
"x-goog-api-key",
|
||||
"cookie",
|
||||
"cookie2",
|
||||
"set-cookie",
|
||||
];
|
||||
|
||||
pub(crate) fn upstream_credential_header_names() -> &'static [&'static str] {
|
||||
UPSTREAM_CREDENTIAL_HEADER_NAMES
|
||||
}
|
||||
|
||||
pub(crate) fn is_upstream_credential_header(name: &str) -> bool {
|
||||
let name = name.trim();
|
||||
UPSTREAM_CREDENTIAL_HEADER_NAMES
|
||||
.iter()
|
||||
.any(|candidate| name.eq_ignore_ascii_case(candidate))
|
||||
}
|
||||
|
||||
pub(crate) fn is_aether_internal_header(name: &str) -> bool {
|
||||
name.trim().to_ascii_lowercase().starts_with("x-aether-")
|
||||
}
|
||||
|
||||
pub fn should_skip_request_header(name: &str) -> bool {
|
||||
let normalized = name.to_ascii_lowercase();
|
||||
if is_aether_internal_header(&normalized) {
|
||||
return true;
|
||||
}
|
||||
matches!(
|
||||
normalized.as_str(),
|
||||
"connection"
|
||||
@@ -15,11 +44,7 @@ pub fn should_skip_request_header(name: &str) -> bool {
|
||||
| "trailer"
|
||||
| "transfer-encoding"
|
||||
| "upgrade"
|
||||
| "x-aether-execution-path"
|
||||
| "x-aether-dependency-reason"
|
||||
| "x-aether-execution-loop-guard"
|
||||
| "x-aether-control-execute-fallback"
|
||||
| "x-aether-rate-limit-preflight"
|
||||
| "set-cookie"
|
||||
| USAGE_SERVER_NOW_UNIX_MS_HEADER
|
||||
)
|
||||
}
|
||||
@@ -35,12 +60,10 @@ pub fn should_skip_upstream_passthrough_header(name: &str) -> bool {
|
||||
if lower.starts_with("x-stainless-") || lower.starts_with("anthropic-") {
|
||||
return true;
|
||||
}
|
||||
matches!(
|
||||
lower.as_str(),
|
||||
"authorization"
|
||||
| "x-api-key"
|
||||
| "x-goog-api-key"
|
||||
| "host"
|
||||
is_upstream_credential_header(&lower)
|
||||
|| matches!(
|
||||
lower.as_str(),
|
||||
"host"
|
||||
| "content-length"
|
||||
| "transfer-encoding"
|
||||
| "connection"
|
||||
@@ -55,29 +78,29 @@ pub fn should_skip_upstream_passthrough_header(name: &str) -> bool {
|
||||
// Claude CLI client identifier; re-injected by the Claude Code adapter
|
||||
// when the upstream is Anthropic, filtered for everybody else.
|
||||
| "x-app"
|
||||
) || should_skip_request_header(name)
|
||||
)
|
||||
|| should_skip_request_header(name)
|
||||
}
|
||||
|
||||
pub(crate) fn should_skip_upstream_complete_passthrough_header(name: &str) -> bool {
|
||||
let lower = name.to_ascii_lowercase();
|
||||
matches!(
|
||||
lower.as_str(),
|
||||
"authorization"
|
||||
| "x-api-key"
|
||||
| "x-goog-api-key"
|
||||
| "host"
|
||||
| "content-length"
|
||||
| "transfer-encoding"
|
||||
| "connection"
|
||||
| "content-encoding"
|
||||
| "x-real-ip"
|
||||
| "x-real-proto"
|
||||
| "x-forwarded-for"
|
||||
| "x-forwarded-proto"
|
||||
| "x-forwarded-scheme"
|
||||
| "x-forwarded-host"
|
||||
| "x-forwarded-port"
|
||||
) || should_skip_request_header(name)
|
||||
is_upstream_credential_header(&lower)
|
||||
|| matches!(
|
||||
lower.as_str(),
|
||||
"host"
|
||||
| "content-length"
|
||||
| "transfer-encoding"
|
||||
| "connection"
|
||||
| "content-encoding"
|
||||
| "x-real-ip"
|
||||
| "x-real-proto"
|
||||
| "x-forwarded-for"
|
||||
| "x-forwarded-proto"
|
||||
| "x-forwarded-scheme"
|
||||
| "x-forwarded-host"
|
||||
| "x-forwarded-port"
|
||||
)
|
||||
|| should_skip_request_header(name)
|
||||
}
|
||||
|
||||
pub fn normalize_upstream_accept_encoding(value: &str) -> Option<String> {
|
||||
@@ -170,9 +193,9 @@ pub fn force_identity_accept_encoding(headers: &mut BTreeMap<String, String>) {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
force_identity_accept_encoding, normalize_upstream_accept_encoding,
|
||||
should_skip_request_header, should_skip_upstream_complete_passthrough_header,
|
||||
should_skip_upstream_passthrough_header,
|
||||
force_identity_accept_encoding, is_upstream_credential_header,
|
||||
normalize_upstream_accept_encoding, should_skip_request_header,
|
||||
should_skip_upstream_complete_passthrough_header, should_skip_upstream_passthrough_header,
|
||||
};
|
||||
use aether_contracts::USAGE_SERVER_NOW_UNIX_MS_HEADER;
|
||||
use std::collections::BTreeMap;
|
||||
@@ -273,6 +296,47 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strips_all_client_credential_carriers_from_passthrough() {
|
||||
for header in [
|
||||
"authorization",
|
||||
"proxy-authorization",
|
||||
"api-key",
|
||||
"x-api-key",
|
||||
"x-goog-api-key",
|
||||
"cookie",
|
||||
"cookie2",
|
||||
"set-cookie",
|
||||
"Authorization",
|
||||
"COOKIE",
|
||||
] {
|
||||
assert!(is_upstream_credential_header(header), "credential {header}");
|
||||
assert!(
|
||||
should_skip_upstream_passthrough_header(header),
|
||||
"normal passthrough should strip {header}"
|
||||
);
|
||||
assert!(
|
||||
should_skip_upstream_complete_passthrough_header(header),
|
||||
"complete passthrough should strip {header}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strips_all_aether_owned_headers_from_provider_requests() {
|
||||
for header in [
|
||||
"x-aether-gateway",
|
||||
"x-aether-auth-user-id",
|
||||
"x-aether-auth-api-key-id",
|
||||
"x-aether-auth-balance-remaining",
|
||||
"X-Aether-Tunnel-Forwarded-By",
|
||||
] {
|
||||
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 [
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use aether_oauth::provider::providers::KiroProviderOAuthAdapter as CoreKiroProviderOAuthAdapter;
|
||||
use async_trait::async_trait;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::super::oauth_refresh::{
|
||||
oauth_error_to_local_refresh_error, provider_oauth_transport_context_from_snapshot,
|
||||
@@ -51,11 +52,15 @@ impl KiroOAuthRefreshAdapter {
|
||||
.map_err(|error| oauth_error_to_local_refresh_error(PROVIDER_TYPE, error))
|
||||
}
|
||||
|
||||
fn auth_config_from_entry(entry: &CachedOAuthEntry) -> Option<KiroAuthConfig> {
|
||||
fn auth_config_from_entry(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<KiroAuthConfig> {
|
||||
entry
|
||||
.metadata
|
||||
.as_ref()
|
||||
.filter(|_| entry.provider_type.eq_ignore_ascii_case(PROVIDER_TYPE))
|
||||
.filter(|_| kiro_cached_entry_matches_transport(transport, entry))
|
||||
.and_then(KiroAuthConfig::from_json_value)
|
||||
}
|
||||
|
||||
@@ -64,12 +69,17 @@ impl KiroOAuthRefreshAdapter {
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<KiroAuthConfig> {
|
||||
entry.and_then(Self::auth_config_from_entry).or_else(|| {
|
||||
KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
|
||||
})
|
||||
entry
|
||||
.and_then(|entry| Self::auth_config_from_entry(transport, entry))
|
||||
.or_else(|| {
|
||||
KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
|
||||
})
|
||||
}
|
||||
|
||||
fn build_cached_entry(auth_config: &KiroAuthConfig) -> Option<CachedOAuthEntry> {
|
||||
fn build_cached_entry(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
auth_config: &KiroAuthConfig,
|
||||
) -> Option<CachedOAuthEntry> {
|
||||
let request_auth = build_kiro_request_auth_from_config(auth_config.clone(), None)?;
|
||||
Some(CachedOAuthEntry {
|
||||
provider_type: PROVIDER_TYPE.to_string(),
|
||||
@@ -77,7 +87,21 @@ impl KiroOAuthRefreshAdapter {
|
||||
auth_header_value: request_auth.value,
|
||||
expires_at_unix_secs: auth_config.expires_at,
|
||||
metadata: Some(auth_config.to_json_value()),
|
||||
source_fingerprint: None,
|
||||
source_fingerprint: Some(kiro_transport_credential_fingerprint(transport)),
|
||||
})
|
||||
}
|
||||
|
||||
fn build_cached_entry_from_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<CachedOAuthEntry> {
|
||||
let request_auth = resolve_local_kiro_request_auth(transport)?;
|
||||
Some(CachedOAuthEntry {
|
||||
provider_type: PROVIDER_TYPE.to_string(),
|
||||
auth_header_name: request_auth.name.to_string(),
|
||||
auth_header_value: request_auth.value,
|
||||
expires_at_unix_secs: request_auth.auth_config.expires_at,
|
||||
metadata: Some(request_auth.auth_config.to_json_value()),
|
||||
source_fingerprint: Some(kiro_transport_credential_fingerprint(transport)),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -101,10 +125,10 @@ impl LocalOAuthRefreshAdapter for KiroOAuthRefreshAdapter {
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
let auth_config = Self::auth_config_from_entry(entry)?;
|
||||
let auth_config = Self::auth_config_from_entry(transport, entry)?;
|
||||
let request_auth = build_kiro_request_auth_from_config(auth_config, None)?;
|
||||
Some(LocalResolvedOAuthRequestAuth::Kiro(request_auth))
|
||||
}
|
||||
@@ -128,6 +152,24 @@ impl LocalOAuthRefreshAdapter for KiroOAuthRefreshAdapter {
|
||||
&& self.refreshable_auth_config(transport, entry).is_some()
|
||||
}
|
||||
|
||||
fn refresh_fingerprint(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
self.supports(transport).then(|| {
|
||||
kiro_successor_refresh_fingerprint(transport, entry)
|
||||
.unwrap_or_else(|| kiro_transport_credential_fingerprint(transport))
|
||||
})
|
||||
}
|
||||
|
||||
fn cached_entry_from_transport(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<CachedOAuthEntry> {
|
||||
Self::build_cached_entry_from_transport(transport)
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
@@ -140,10 +182,75 @@ impl LocalOAuthRefreshAdapter for KiroOAuthRefreshAdapter {
|
||||
let refreshed = self
|
||||
.refresh_auth_config(executor, transport, &auth_config)
|
||||
.await?;
|
||||
Ok(Self::build_cached_entry(&refreshed))
|
||||
Ok(Self::build_cached_entry(transport, &refreshed))
|
||||
}
|
||||
}
|
||||
|
||||
fn kiro_transport_credential_fingerprint(transport: &GatewayProviderTransportSnapshot) -> String {
|
||||
kiro_credential_fingerprint(
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.key.auth_type.as_str(),
|
||||
transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.unwrap_or_default(),
|
||||
transport.key.decrypted_api_key.as_str(),
|
||||
)
|
||||
}
|
||||
|
||||
fn kiro_successor_refresh_fingerprint(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
let entry = entry.filter(|entry| kiro_cached_entry_matches_transport(transport, entry))?;
|
||||
let metadata = serde_json::to_string(entry.metadata.as_ref()?).ok()?;
|
||||
let access_token = bearer_access_token(entry.auth_header_value.as_str())?;
|
||||
Some(kiro_credential_fingerprint(
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.key.auth_type.as_str(),
|
||||
metadata.as_str(),
|
||||
access_token,
|
||||
))
|
||||
}
|
||||
|
||||
fn kiro_cached_entry_matches_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> bool {
|
||||
entry.provider_type.eq_ignore_ascii_case(PROVIDER_TYPE)
|
||||
&& entry.source_fingerprint.as_deref()
|
||||
== Some(kiro_transport_credential_fingerprint(transport).as_str())
|
||||
}
|
||||
|
||||
fn kiro_credential_fingerprint(
|
||||
provider_type: &str,
|
||||
auth_type: &str,
|
||||
auth_config: &str,
|
||||
access_token: &str,
|
||||
) -> String {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = auth_type.trim().to_ascii_lowercase();
|
||||
let mut digest = Sha256::new();
|
||||
for field in [
|
||||
provider_type.as_bytes(),
|
||||
auth_type.as_bytes(),
|
||||
auth_config.as_bytes(),
|
||||
access_token.as_bytes(),
|
||||
] {
|
||||
digest.update((field.len() as u64).to_be_bytes());
|
||||
digest.update(field);
|
||||
}
|
||||
format!("{:x}", digest.finalize())
|
||||
}
|
||||
|
||||
fn bearer_access_token(authorization: &str) -> Option<&str> {
|
||||
let mut parts = authorization.split_ascii_whitespace();
|
||||
let scheme = parts.next()?;
|
||||
let token = parts.next()?;
|
||||
(scheme.eq_ignore_ascii_case("bearer") && parts.next().is_none()).then_some(token)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::{Arc, Mutex};
|
||||
@@ -155,7 +262,7 @@ mod tests {
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::{KiroOAuthRefreshAdapter, IDC_AMZ_USER_AGENT};
|
||||
use super::{KiroAuthConfig, KiroOAuthRefreshAdapter, IDC_AMZ_USER_AGENT};
|
||||
use axum::body::to_bytes;
|
||||
use axum::extract::Request;
|
||||
use axum::response::IntoResponse;
|
||||
@@ -244,6 +351,72 @@ mod tests {
|
||||
(format!("http://{addr}"), handle)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cached_entry_is_bound_to_kiro_credential_generation() {
|
||||
let stable_refresh_token = "r".repeat(120);
|
||||
let source_transport = sample_transport(
|
||||
&json!({
|
||||
"refresh_token": stable_refresh_token,
|
||||
"machine_id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"kiro_version": "1.2.3"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
let refreshed_config = KiroAuthConfig::from_raw_json(Some(
|
||||
&json!({
|
||||
"refresh_token": "s".repeat(120),
|
||||
"access_token": "fresh-kiro-access-token",
|
||||
"expires_at": u64::MAX,
|
||||
"machine_id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"kiro_version": "1.2.3"
|
||||
})
|
||||
.to_string(),
|
||||
))
|
||||
.expect("refreshed Kiro config should parse");
|
||||
let entry =
|
||||
KiroOAuthRefreshAdapter::build_cached_entry(&source_transport, &refreshed_config)
|
||||
.expect("refreshed Kiro entry should build");
|
||||
let adapter = KiroOAuthRefreshAdapter::default();
|
||||
|
||||
assert!(adapter.resolve_cached(&source_transport, &entry).is_some());
|
||||
assert_ne!(
|
||||
adapter.refresh_fingerprint(&source_transport, None),
|
||||
adapter.refresh_fingerprint(&source_transport, Some(&entry))
|
||||
);
|
||||
|
||||
let replacement_transport = sample_transport(
|
||||
&json!({
|
||||
"refresh_token": "admin-refresh-token",
|
||||
"access_token": "admin-access-token",
|
||||
"expires_at": u64::MAX,
|
||||
"machine_id": "123e4567-e89b-12d3-a456-426614174001",
|
||||
"kiro_version": "1.2.3"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
assert!(adapter
|
||||
.resolve_cached(&replacement_transport, &entry)
|
||||
.is_none());
|
||||
let selected_refresh_config = adapter
|
||||
.base_auth_config(&replacement_transport, Some(&entry))
|
||||
.expect("replacement transport should provide refresh config");
|
||||
assert_eq!(
|
||||
selected_refresh_config.refresh_token.as_deref(),
|
||||
Some("admin-refresh-token")
|
||||
);
|
||||
assert_eq!(
|
||||
selected_refresh_config.access_token.as_deref(),
|
||||
Some("admin-access-token")
|
||||
);
|
||||
|
||||
let replacement_entry =
|
||||
LocalOAuthRefreshAdapter::cached_entry_from_transport(&adapter, &replacement_transport)
|
||||
.expect("persisted replacement should reconstruct a generation-bound entry");
|
||||
assert!(adapter
|
||||
.resolve_cached(&replacement_transport, &replacement_entry)
|
||||
.is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn refreshes_social_token_via_adapter() {
|
||||
let seen_request = Arc::new(Mutex::new(None::<SeenRefreshRequest>));
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
mod agent_identity;
|
||||
mod anthropic_compat;
|
||||
pub mod antigravity;
|
||||
pub mod auth;
|
||||
mod auth_config;
|
||||
@@ -46,6 +47,10 @@ pub use agent_identity::{
|
||||
CODEX_AGENT_IDENTITY_AUTH_MODE, CODEX_AGENT_IDENTITY_CACHED_ENTRY_PROVIDER_TYPE,
|
||||
CODEX_AGENT_IDENTITY_PROVIDER_TYPE, CODEX_AGENT_IDENTITY_TASK_REGISTRATION_REQUEST_ID,
|
||||
};
|
||||
pub use anthropic_compat::{
|
||||
resolve_anthropic_compatibility_profile, validate_anthropic_compatibility_profile_config,
|
||||
AnthropicCompatibilityProfile, AnthropicCompatibilityProfileConfigError,
|
||||
};
|
||||
pub use auth::{build_passthrough_headers, ensure_upstream_auth_header};
|
||||
pub use auth_config::apply_local_auth_config_header_overrides;
|
||||
pub use cache::{provider_transport_snapshot_looks_refreshed, ProviderTransportSnapshotCacheKey};
|
||||
@@ -76,6 +81,7 @@ pub use gemini_files::{
|
||||
GeminiFilesRequestBodyError, GeminiFilesRequestBodyParts,
|
||||
};
|
||||
pub use generic_oauth::{
|
||||
resolve_local_generic_oauth_transport_authorization,
|
||||
supports_local_generic_oauth_request_auth_resolution, GenericOAuthRefreshAdapter,
|
||||
};
|
||||
pub use grok::{
|
||||
@@ -122,7 +128,7 @@ pub use request_url::{
|
||||
build_kiro_cross_format_upstream_url, build_local_openai_chat_upstream_url,
|
||||
build_local_openai_responses_upstream_url, build_transport_request_url,
|
||||
build_transport_request_url_for_request_body, gemini_embedding_request_body_uses_batch,
|
||||
TransportRequestUrlParams,
|
||||
transport_supports_api_operation, TransportRequestUrlParams,
|
||||
};
|
||||
pub use rules::{
|
||||
apply_local_body_rules, apply_local_body_rules_with_request_headers, apply_local_header_rules,
|
||||
@@ -134,6 +140,8 @@ pub use same_format_provider::{
|
||||
build_same_format_provider_headers, build_same_format_provider_request_body,
|
||||
build_same_format_provider_request_body_with_compatibility_report,
|
||||
build_same_format_provider_upstream_url, classify_same_format_provider_request_behavior,
|
||||
classify_same_format_provider_request_behavior_for_operation,
|
||||
enforce_same_format_provider_api_operation_body_policy,
|
||||
resolve_same_format_provider_direct_auth, same_format_provider_transport_supported,
|
||||
same_format_provider_transport_unsupported_reason,
|
||||
same_format_provider_transport_unsupported_reason_for_trace,
|
||||
|
||||
@@ -6,6 +6,7 @@ use async_trait::async_trait;
|
||||
use serde_json::{json, Map, Value};
|
||||
use tracing::warn;
|
||||
|
||||
use crate::claude_code::current_claude_code_transport_identity_profile;
|
||||
use crate::grok::grok_browser_resolved_transport_profile_from_auth_config;
|
||||
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
@@ -152,9 +153,38 @@ pub fn resolve_transport_profile_id(
|
||||
pub fn resolve_transport_profile(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
resolve_transport_profile_from_fingerprint(transport.key.fingerprint.as_ref()).or_else(|| {
|
||||
resolve_transport_profile_from_provider_config(transport.provider.config.as_ref())
|
||||
.or_else(|| resolve_grok_browser_transport_profile(transport))
|
||||
let configured = resolve_transport_profile_from_fingerprint(transport.key.fingerprint.as_ref())
|
||||
.or_else(|| {
|
||||
resolve_transport_profile_from_provider_config(transport.provider.config.as_ref())
|
||||
});
|
||||
if configured.is_some() || transport_profile_is_configured(transport) {
|
||||
return configured;
|
||||
}
|
||||
|
||||
resolve_claude_code_transport_profile(transport)
|
||||
.or_else(|| resolve_grok_browser_transport_profile(transport))
|
||||
}
|
||||
|
||||
fn resolve_claude_code_transport_profile(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("claude_code")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let identity_profile = *current_claude_code_transport_identity_profile();
|
||||
Some(ResolvedTransportProfile {
|
||||
profile_id: identity_profile.transport_profile_id().to_string(),
|
||||
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -595,6 +625,64 @@ mod tests {
|
||||
assert_eq!(profile.backend, "reqwest_rustls");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_typed_claude_code_transport_profile_when_unconfigured() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "claude_code".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("typed Claude Code profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "claude_code_nodejs");
|
||||
assert_eq!(profile.backend, "reqwest_rustls");
|
||||
assert_eq!(profile.http_mode, "auto");
|
||||
assert_eq!(profile.pool_scope, "key");
|
||||
assert!(profile.header_fingerprint.is_none());
|
||||
assert!(profile.extra.is_none());
|
||||
assert!(!transport_profile_is_configured(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_claude_code_transport_profiles_precede_typed_default() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "claude_code".to_string();
|
||||
transport.provider.config = Some(json!({
|
||||
"fingerprint": {"transport_profile": "provider_claude_profile"}
|
||||
}));
|
||||
transport.key.fingerprint = Some(json!({
|
||||
"transport_profile": "key_claude_profile"
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
resolve_transport_profile(&transport)
|
||||
.expect("key profile")
|
||||
.profile_id,
|
||||
"key_claude_profile"
|
||||
);
|
||||
|
||||
transport.key.fingerprint = None;
|
||||
assert_eq!(
|
||||
resolve_transport_profile(&transport)
|
||||
.expect("provider profile")
|
||||
.profile_id,
|
||||
"provider_claude_profile"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_explicit_claude_code_transport_profile_blocks_typed_default() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "claude_code".to_string();
|
||||
transport.provider.config = None;
|
||||
transport.key.fingerprint = Some(json!({
|
||||
"transport_profile": {"backend": "reqwest_rustls"}
|
||||
}));
|
||||
|
||||
assert!(transport_profile_is_configured(&transport));
|
||||
assert!(resolve_transport_profile(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn maps_string_transport_profile_to_resolved_profile() {
|
||||
let profile = resolve_transport_profile(&sample_transport()).expect("profile");
|
||||
|
||||
@@ -12,7 +12,7 @@ use aether_runtime_state::{RuntimeLockLease, RuntimeState};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::sync::{Mutex, OwnedMutexGuard};
|
||||
|
||||
use super::agent_identity::{is_codex_agent_identity_transport, CodexAgentIdentityRefreshAdapter};
|
||||
use super::generic_oauth::supports_local_generic_oauth_request_auth_resolution;
|
||||
@@ -47,6 +47,35 @@ pub struct LocalOAuthResolution {
|
||||
/// Held until the caller persists `refreshed_entry`. The lease TTL remains
|
||||
/// the cancellation fallback if the caller is dropped.
|
||||
pub distributed_lease: Option<RuntimeLockLease>,
|
||||
/// Keeps memory-only refreshes singleflight until the caller validates the
|
||||
/// credential fence and publishes or discards `refreshed_entry`.
|
||||
#[doc(hidden)]
|
||||
pub local_refresh_guard: Option<LocalOAuthRefreshCommitGuard>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct LocalOAuthRefreshCommitGuard {
|
||||
guard: Arc<OwnedMutexGuard<()>>,
|
||||
}
|
||||
|
||||
impl LocalOAuthRefreshCommitGuard {
|
||||
fn new(guard: OwnedMutexGuard<()>) -> Self {
|
||||
Self {
|
||||
guard: Arc::new(guard),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for LocalOAuthRefreshCommitGuard {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.write_str("LocalOAuthRefreshCommitGuard")
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for LocalOAuthRefreshCommitGuard {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
Arc::ptr_eq(&self.guard, &other.guard)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
@@ -369,6 +398,12 @@ pub trait LocalOAuthRefreshAdapter: Send + Sync {
|
||||
false
|
||||
}
|
||||
|
||||
/// Whether another gateway instance can observe this refresh after the
|
||||
/// caller persists its result into the provider transport record.
|
||||
fn shares_refresh_through_transport_persistence(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
@@ -388,6 +423,7 @@ pub struct LocalOAuthRefreshCoordinator {
|
||||
struct RefreshBackoffState {
|
||||
failures: u32,
|
||||
retry_after: Instant,
|
||||
refresh_fingerprint: Option<String>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LocalOAuthRefreshCoordinator {
|
||||
@@ -424,6 +460,19 @@ impl LocalOAuthRefreshCoordinator {
|
||||
}
|
||||
}
|
||||
|
||||
/// Captures the refresh generation represented by this transport snapshot.
|
||||
/// The coordinator cache is intentionally excluded: callers use this value
|
||||
/// as the fence for the credential generation that produced their request.
|
||||
pub fn refresh_fingerprint_for_transport(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<String> {
|
||||
self.adapters
|
||||
.iter()
|
||||
.find(|adapter| adapter.supports(transport))
|
||||
.and_then(|adapter| adapter.refresh_fingerprint(transport, None))
|
||||
}
|
||||
|
||||
async fn lock_for_key(&self, key_id: &str) -> Arc<Mutex<()>> {
|
||||
let mut key_locks = self.key_locks.lock().await;
|
||||
key_locks
|
||||
@@ -530,6 +579,11 @@ impl LocalOAuthRefreshCoordinator {
|
||||
} else {
|
||||
self.cached_entry(key_id).await
|
||||
};
|
||||
let shares_refresh_through_transport_persistence =
|
||||
adapter.shares_refresh_through_transport_persistence();
|
||||
let pre_lock_local_refresh_fingerprint = force_refresh
|
||||
.then(|| adapter.refresh_fingerprint(transport, cached_entry.as_ref()))
|
||||
.flatten();
|
||||
if !force_refresh {
|
||||
if let Some(auth) = cached_entry
|
||||
.as_ref()
|
||||
@@ -548,7 +602,7 @@ impl LocalOAuthRefreshCoordinator {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if force_refresh {
|
||||
if force_refresh && shares_refresh_through_transport_persistence {
|
||||
if let Some(resolution) = Self::resolve_if_refresh_fence_advanced(
|
||||
adapter.as_ref(),
|
||||
transport,
|
||||
@@ -558,25 +612,46 @@ impl LocalOAuthRefreshCoordinator {
|
||||
return Ok(Some(resolution));
|
||||
}
|
||||
}
|
||||
if let Some(error) = self.backoff_error(key_id, adapter.provider_type()).await {
|
||||
let refresh_fingerprint = adapter.refresh_fingerprint(transport, cached_entry.as_ref());
|
||||
if let Some(error) = self
|
||||
.backoff_error(
|
||||
key_id,
|
||||
adapter.provider_type(),
|
||||
refresh_fingerprint.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
return Err(error);
|
||||
}
|
||||
|
||||
let key_lock = self.lock_for_key(key_id).await;
|
||||
let _key_guard = key_lock.lock().await;
|
||||
let key_guard = key_lock.lock_owned().await;
|
||||
|
||||
let cached_entry = self.cached_entry(key_id).await;
|
||||
if force_refresh {
|
||||
let winner_fingerprint = if shares_refresh_through_transport_persistence {
|
||||
expected_refresh_fingerprint
|
||||
} else {
|
||||
pre_lock_local_refresh_fingerprint.as_deref()
|
||||
};
|
||||
if let Some(resolution) = Self::resolve_if_refresh_fence_advanced(
|
||||
adapter.as_ref(),
|
||||
transport,
|
||||
cached_entry.as_ref(),
|
||||
expected_refresh_fingerprint,
|
||||
winner_fingerprint,
|
||||
) {
|
||||
return Ok(Some(resolution));
|
||||
}
|
||||
}
|
||||
if let Some(error) = self.backoff_error(key_id, adapter.provider_type()).await {
|
||||
let refresh_fingerprint = adapter.refresh_fingerprint(transport, cached_entry.as_ref());
|
||||
if let Some(error) = self
|
||||
.backoff_error(
|
||||
key_id,
|
||||
adapter.provider_type(),
|
||||
refresh_fingerprint.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
return Err(error);
|
||||
}
|
||||
if !force_refresh {
|
||||
@@ -594,40 +669,48 @@ impl LocalOAuthRefreshCoordinator {
|
||||
}
|
||||
}
|
||||
|
||||
let distributed_lease = match (distributed_lock, distributed_owner) {
|
||||
(Some(lock), Some(owner)) if !owner.trim().is_empty() => {
|
||||
match lock
|
||||
.lock_try_acquire(
|
||||
&format!("provider_oauth_refresh_lock:{key_id}"),
|
||||
owner,
|
||||
std::time::Duration::from_millis(Self::DISTRIBUTED_REFRESH_LOCK_TTL_MS),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(lease)) => Some(lease),
|
||||
Ok(None) => return Ok(Some(LocalOAuthResolution::refresh_in_flight())),
|
||||
Err(err) => {
|
||||
tracing::warn!(
|
||||
key_id = %key_id,
|
||||
provider_type = adapter.provider_type(),
|
||||
error = ?err,
|
||||
"gateway local oauth refresh distributed lock unavailable"
|
||||
);
|
||||
if adapter.requires_distributed_refresh_lock() {
|
||||
let error = LocalOAuthRefreshError::TransportMessage {
|
||||
provider_type: adapter.provider_type(),
|
||||
message: "distributed refresh lock is unavailable".to_string(),
|
||||
};
|
||||
if adapter.should_backoff_after_error(&error) {
|
||||
self.record_refresh_failure(key_id).await;
|
||||
let distributed_lease = if !shares_refresh_through_transport_persistence {
|
||||
None
|
||||
} else {
|
||||
match (distributed_lock, distributed_owner) {
|
||||
(Some(lock), Some(owner)) if !owner.trim().is_empty() => {
|
||||
match lock
|
||||
.lock_try_acquire(
|
||||
&format!("provider_oauth_refresh_lock:{key_id}"),
|
||||
owner,
|
||||
std::time::Duration::from_millis(Self::DISTRIBUTED_REFRESH_LOCK_TTL_MS),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(lease)) => Some(lease),
|
||||
Ok(None) => return Ok(Some(LocalOAuthResolution::refresh_in_flight())),
|
||||
Err(err) => {
|
||||
tracing::warn!(
|
||||
key_id = %key_id,
|
||||
provider_type = adapter.provider_type(),
|
||||
error = ?err,
|
||||
"gateway local oauth refresh distributed lock unavailable"
|
||||
);
|
||||
if adapter.requires_distributed_refresh_lock() {
|
||||
let error = LocalOAuthRefreshError::TransportMessage {
|
||||
provider_type: adapter.provider_type(),
|
||||
message: "distributed refresh lock is unavailable".to_string(),
|
||||
};
|
||||
if adapter.should_backoff_after_error(&error) {
|
||||
self.record_refresh_failure(
|
||||
key_id,
|
||||
refresh_fingerprint.as_deref(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
return Err(error);
|
||||
None
|
||||
}
|
||||
None
|
||||
}
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
|
||||
// Forced refresh still needs the latest rotated refresh_token as input.
|
||||
@@ -653,7 +736,8 @@ impl LocalOAuthRefreshCoordinator {
|
||||
}
|
||||
Err(error) => {
|
||||
if adapter.should_backoff_after_error(&error) {
|
||||
self.record_refresh_failure(key_id).await;
|
||||
self.record_refresh_failure(key_id, refresh_fingerprint.as_deref())
|
||||
.await;
|
||||
}
|
||||
Self::release_distributed_lease(
|
||||
distributed_lock,
|
||||
@@ -665,14 +749,9 @@ impl LocalOAuthRefreshCoordinator {
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
// In production the distributed lease is held through the gateway's
|
||||
// DB CAS. Do not publish a provisional task before that CAS succeeds;
|
||||
// otherwise a waiter could consume an assertion that loses the CAS.
|
||||
// Lock-free/test callers retain the historical in-memory behavior.
|
||||
if distributed_lease.is_none() {
|
||||
self.insert_cached_entry(key_id, refreshed_entry.clone())
|
||||
.await;
|
||||
}
|
||||
// Cache publication belongs to the caller after durable persistence.
|
||||
// The result still carries the entry so lock-free callers can inspect
|
||||
// it or explicitly commit it with `store_cached_entry`.
|
||||
let Some(auth) = adapter.resolve_refreshed(transport, &refreshed_entry) else {
|
||||
Self::release_distributed_lease(
|
||||
distributed_lock,
|
||||
@@ -683,10 +762,16 @@ impl LocalOAuthRefreshCoordinator {
|
||||
.await;
|
||||
return Ok(None);
|
||||
};
|
||||
let local_refresh_guard = if shares_refresh_through_transport_persistence {
|
||||
None
|
||||
} else {
|
||||
Some(LocalOAuthRefreshCommitGuard::new(key_guard))
|
||||
};
|
||||
Ok(Some(LocalOAuthResolution::refreshed(
|
||||
auth,
|
||||
refreshed_entry,
|
||||
distributed_lease,
|
||||
local_refresh_guard,
|
||||
)))
|
||||
}
|
||||
|
||||
@@ -728,8 +813,16 @@ impl LocalOAuthRefreshCoordinator {
|
||||
&self,
|
||||
key_id: &str,
|
||||
provider_type: &'static str,
|
||||
refresh_fingerprint: Option<&str>,
|
||||
) -> Option<LocalOAuthRefreshError> {
|
||||
let backoff = self.refresh_backoff.lock().await;
|
||||
let mut backoff = self.refresh_backoff.lock().await;
|
||||
if backoff
|
||||
.get(key_id)
|
||||
.is_some_and(|state| state.refresh_fingerprint.as_deref() != refresh_fingerprint)
|
||||
{
|
||||
backoff.remove(key_id);
|
||||
return None;
|
||||
}
|
||||
let state = backoff.get(key_id)?;
|
||||
let remaining = state.retry_after.checked_duration_since(Instant::now())?;
|
||||
Some(LocalOAuthRefreshError::InvalidResponse {
|
||||
@@ -742,14 +835,19 @@ impl LocalOAuthRefreshCoordinator {
|
||||
})
|
||||
}
|
||||
|
||||
async fn record_refresh_failure(&self, key_id: &str) {
|
||||
async fn record_refresh_failure(&self, key_id: &str, refresh_fingerprint: Option<&str>) {
|
||||
let mut backoff = self.refresh_backoff.lock().await;
|
||||
let state = backoff
|
||||
.entry(key_id.to_string())
|
||||
.or_insert(RefreshBackoffState {
|
||||
failures: 0,
|
||||
retry_after: Instant::now(),
|
||||
refresh_fingerprint: refresh_fingerprint.map(ToOwned::to_owned),
|
||||
});
|
||||
if state.refresh_fingerprint.as_deref() != refresh_fingerprint {
|
||||
state.failures = 0;
|
||||
state.refresh_fingerprint = refresh_fingerprint.map(ToOwned::to_owned);
|
||||
}
|
||||
state.failures = state.failures.saturating_add(1);
|
||||
let exponent = state.failures.saturating_sub(1).min(4);
|
||||
let delay = Duration::from_millis(500u64.saturating_mul(1u64 << exponent));
|
||||
@@ -800,6 +898,7 @@ impl LocalOAuthResolution {
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
local_refresh_guard: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -807,6 +906,7 @@ impl LocalOAuthResolution {
|
||||
auth: LocalResolvedOAuthRequestAuth,
|
||||
refreshed_entry: CachedOAuthEntry,
|
||||
distributed_lease: Option<RuntimeLockLease>,
|
||||
local_refresh_guard: Option<LocalOAuthRefreshCommitGuard>,
|
||||
) -> Self {
|
||||
Self {
|
||||
auth: Some(auth),
|
||||
@@ -814,6 +914,7 @@ impl LocalOAuthResolution {
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease,
|
||||
local_refresh_guard,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -829,6 +930,7 @@ impl LocalOAuthResolution {
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: true,
|
||||
distributed_lease: None,
|
||||
local_refresh_guard: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -839,6 +941,7 @@ impl LocalOAuthResolution {
|
||||
refresh_in_flight: true,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
local_refresh_guard: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -855,6 +958,7 @@ pub fn supports_local_oauth_request_auth_resolution(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
|
||||
use std::time::Duration;
|
||||
|
||||
use super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
@@ -878,6 +982,13 @@ mod tests {
|
||||
struct FencedTestAdapter {
|
||||
refresh_hits: Arc<AtomicUsize>,
|
||||
fail_refresh: Arc<AtomicBool>,
|
||||
generation: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct MemoryOnlyFencedTestAdapter {
|
||||
refresh_hits: Arc<AtomicUsize>,
|
||||
fingerprint_hits: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -978,7 +1089,12 @@ mod tests {
|
||||
) -> Option<String> {
|
||||
entry
|
||||
.and_then(|entry| entry.source_fingerprint.clone())
|
||||
.or_else(|| Some("generation-1".to_string()))
|
||||
.or_else(|| {
|
||||
Some(format!(
|
||||
"generation-{}",
|
||||
self.generation.load(Ordering::SeqCst)
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn should_backoff_after_error(&self, _error: &LocalOAuthRefreshError) -> bool {
|
||||
@@ -1004,7 +1120,75 @@ mod tests {
|
||||
auth_header_value: "stale-winner-cache-value".to_string(),
|
||||
expires_at_unix_secs: None,
|
||||
metadata: None,
|
||||
source_fingerprint: Some("generation-2".to_string()),
|
||||
source_fingerprint: Some(format!(
|
||||
"generation-{}",
|
||||
self.generation.load(Ordering::SeqCst).saturating_add(1)
|
||||
)),
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for MemoryOnlyFencedTestAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
"test-oauth"
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: entry.auth_header_name.clone(),
|
||||
value: entry.auth_header_value.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
None
|
||||
}
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
_entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn refresh_fingerprint(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
self.fingerprint_hits.fetch_add(1, Ordering::SeqCst);
|
||||
entry
|
||||
.and_then(|entry| entry.source_fingerprint.clone())
|
||||
.or_else(|| Some("transport-generation".to_string()))
|
||||
}
|
||||
|
||||
fn shares_refresh_through_transport_persistence(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
_executor: &dyn LocalOAuthHttpExecutor,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
_entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
|
||||
let hit = self.refresh_hits.fetch_add(1, Ordering::SeqCst) + 1;
|
||||
Ok(Some(CachedOAuthEntry {
|
||||
provider_type: "test-oauth".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: format!("Bearer refreshed-token-{hit}"),
|
||||
expires_at_unix_secs: Some(4_102_444_800),
|
||||
metadata: None,
|
||||
source_fingerprint: Some(format!("local-generation-{hit}")),
|
||||
}))
|
||||
}
|
||||
}
|
||||
@@ -1082,8 +1266,12 @@ mod tests {
|
||||
.resolve_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.expect("first resolve should succeed");
|
||||
assert!(coordinator
|
||||
.cached_entry(transport.key.id.as_str())
|
||||
.await
|
||||
.is_none());
|
||||
coordinator
|
||||
.insert_cached_entry(
|
||||
.store_cached_entry(
|
||||
transport.key.id.as_str(),
|
||||
first
|
||||
.as_ref()
|
||||
@@ -1115,6 +1303,7 @@ mod tests {
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
local_refresh_guard: None,
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
@@ -1128,6 +1317,7 @@ mod tests {
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
local_refresh_guard: None,
|
||||
})
|
||||
);
|
||||
}
|
||||
@@ -1149,7 +1339,7 @@ mod tests {
|
||||
.await
|
||||
.expect("initial resolve should succeed");
|
||||
coordinator
|
||||
.insert_cached_entry(
|
||||
.store_cached_entry(
|
||||
transport.key.id.as_str(),
|
||||
first
|
||||
.as_ref()
|
||||
@@ -1168,6 +1358,98 @@ mod tests {
|
||||
assert_eq!(refresh_with_entry_hits.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn memory_only_force_refresh_does_not_reuse_preexisting_cache_as_winner() {
|
||||
let refresh_hits = Arc::new(AtomicUsize::new(0));
|
||||
let fingerprint_hits = Arc::new(AtomicUsize::new(0));
|
||||
let coordinator = Arc::new(LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![
|
||||
Arc::new(MemoryOnlyFencedTestAdapter {
|
||||
refresh_hits: Arc::clone(&refresh_hits),
|
||||
fingerprint_hits: Arc::clone(&fingerprint_hits),
|
||||
}),
|
||||
]));
|
||||
let transport = sample_transport();
|
||||
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
|
||||
coordinator
|
||||
.store_cached_entry(
|
||||
transport.key.id.as_str(),
|
||||
CachedOAuthEntry {
|
||||
provider_type: "test-oauth".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "Bearer rejected-token".to_string(),
|
||||
expires_at_unix_secs: Some(4_102_444_800),
|
||||
metadata: None,
|
||||
source_fingerprint: Some("preexisting-local-generation".to_string()),
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
let mut forced = coordinator
|
||||
.force_refresh_with_result_fenced(
|
||||
&executor,
|
||||
&transport,
|
||||
None,
|
||||
None,
|
||||
Some("transport-generation"),
|
||||
)
|
||||
.await
|
||||
.expect("memory-only force refresh should succeed")
|
||||
.expect("memory-only force refresh should resolve");
|
||||
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 1);
|
||||
assert!(!forced.reused_refresh);
|
||||
assert!(forced.local_refresh_guard.is_some());
|
||||
assert_eq!(
|
||||
forced.auth,
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: "authorization".to_string(),
|
||||
value: "Bearer refreshed-token-1".to_string(),
|
||||
})
|
||||
);
|
||||
|
||||
fingerprint_hits.store(0, Ordering::SeqCst);
|
||||
let follower_coordinator = Arc::clone(&coordinator);
|
||||
let follower_transport = transport.clone();
|
||||
let follower = tokio::spawn(async move {
|
||||
follower_coordinator
|
||||
.force_refresh_with_result_fenced(
|
||||
&ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new()),
|
||||
&follower_transport,
|
||||
None,
|
||||
None,
|
||||
Some("transport-generation"),
|
||||
)
|
||||
.await
|
||||
.expect("memory-only follower should succeed")
|
||||
.expect("memory-only follower should resolve")
|
||||
});
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while fingerprint_hits.load(Ordering::SeqCst) == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("follower should capture its pre-lock fingerprint");
|
||||
|
||||
coordinator
|
||||
.store_cached_entry(
|
||||
transport.key.id.as_str(),
|
||||
forced
|
||||
.refreshed_entry
|
||||
.clone()
|
||||
.expect("leader should provide the memory-only entry"),
|
||||
)
|
||||
.await;
|
||||
forced.local_refresh_guard.take();
|
||||
let follower = tokio::time::timeout(Duration::from_secs(1), follower)
|
||||
.await
|
||||
.expect("follower should unblock after cache publication")
|
||||
.expect("follower task should join");
|
||||
|
||||
assert!(follower.reused_refresh);
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fenced_force_refresh_reuses_the_winner() {
|
||||
let refresh_hits = Arc::new(AtomicUsize::new(0));
|
||||
@@ -1175,6 +1457,7 @@ mod tests {
|
||||
FencedTestAdapter {
|
||||
refresh_hits: Arc::clone(&refresh_hits),
|
||||
fail_refresh: Arc::new(AtomicBool::new(false)),
|
||||
generation: Arc::new(AtomicUsize::new(1)),
|
||||
},
|
||||
)]);
|
||||
let transport = sample_transport();
|
||||
@@ -1191,6 +1474,15 @@ mod tests {
|
||||
.await
|
||||
.expect("first refresh should succeed")
|
||||
.expect("first refresh should resolve");
|
||||
coordinator
|
||||
.store_cached_entry(
|
||||
transport.key.id.as_str(),
|
||||
first
|
||||
.refreshed_entry
|
||||
.clone()
|
||||
.expect("first refresh should return an entry to persist"),
|
||||
)
|
||||
.await;
|
||||
let waiter = coordinator
|
||||
.force_refresh_with_result_fenced(
|
||||
&executor,
|
||||
@@ -1221,10 +1513,12 @@ mod tests {
|
||||
async fn refresh_failure_enters_bounded_negative_backoff() {
|
||||
let refresh_hits = Arc::new(AtomicUsize::new(0));
|
||||
let fail_refresh = Arc::new(AtomicBool::new(true));
|
||||
let generation = Arc::new(AtomicUsize::new(1));
|
||||
let coordinator = LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(
|
||||
FencedTestAdapter {
|
||||
refresh_hits: Arc::clone(&refresh_hits),
|
||||
fail_refresh: Arc::clone(&fail_refresh),
|
||||
generation: Arc::clone(&generation),
|
||||
},
|
||||
)]);
|
||||
let transport = sample_transport();
|
||||
@@ -1241,11 +1535,11 @@ mod tests {
|
||||
assert!(second.to_string().contains("temporarily backed off"));
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 1);
|
||||
fail_refresh.store(false, Ordering::SeqCst);
|
||||
coordinator.invalidate_cached_entry("key-1").await;
|
||||
generation.store(2, Ordering::SeqCst);
|
||||
assert!(coordinator
|
||||
.force_refresh_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.expect("replacement should refresh immediately")
|
||||
.expect("new credential generation should bypass old backoff")
|
||||
.is_some());
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
@@ -363,7 +363,7 @@ const GEMINI_CLI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderT
|
||||
|
||||
const VERTEX_AI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "vertex_ai",
|
||||
version: 1,
|
||||
version: 2,
|
||||
base_url: "https://aiplatform.googleapis.com",
|
||||
endpoints: &[
|
||||
FixedProviderEndpointTemplate {
|
||||
@@ -378,12 +378,6 @@ const VERTEX_AI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTe
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "claude:messages",
|
||||
api_format: "claude:messages",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
],
|
||||
runtime_policy: VERTEX_AI_RUNTIME_POLICY,
|
||||
};
|
||||
@@ -936,26 +930,27 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_fixed_provider_template_includes_gemini_embedding_endpoint() {
|
||||
fn vertex_fixed_provider_template_exposes_only_implemented_gemini_endpoints() {
|
||||
let template =
|
||||
fixed_provider_template("vertex_ai").expect("vertex_ai template should exist");
|
||||
|
||||
assert_eq!(template.version, 2);
|
||||
assert_eq!(
|
||||
template
|
||||
.endpoints
|
||||
.iter()
|
||||
.map(|item| item.api_format)
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"gemini:generate_content",
|
||||
"gemini:embedding",
|
||||
"claude:messages",
|
||||
]
|
||||
vec!["gemini:generate_content", "gemini:embedding"]
|
||||
);
|
||||
|
||||
assert!(
|
||||
fixed_provider_endpoint_template_by_api_format("vertex_ai", "gemini:embedding")
|
||||
.is_some()
|
||||
);
|
||||
assert!(
|
||||
fixed_provider_endpoint_template_by_api_format("vertex_ai", "claude:messages")
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::OnceLock;
|
||||
|
||||
use aether_ai_formats::ApiOperation;
|
||||
use regex::Regex;
|
||||
use serde_json::Value;
|
||||
use url::form_urlencoded;
|
||||
@@ -15,9 +16,11 @@ use crate::gemini_cli::{
|
||||
};
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
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,
|
||||
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,
|
||||
@@ -32,6 +35,7 @@ pub struct TransportRequestUrlParams<'a> {
|
||||
pub upstream_is_stream: bool,
|
||||
pub request_query: Option<&'a str>,
|
||||
pub kiro_api_region: Option<&'a str>,
|
||||
pub api_operation: Option<aether_ai_formats::ApiOperation>,
|
||||
}
|
||||
|
||||
pub fn build_transport_request_url(
|
||||
@@ -70,23 +74,62 @@ fn build_transport_request_url_inner(
|
||||
let provider_api_format = params.provider_api_format.trim().to_ascii_lowercase();
|
||||
let normalized_provider_api_format =
|
||||
aether_ai_formats::normalize_api_format_alias(&provider_api_format);
|
||||
let sanitized_claude_request_query = (normalized_provider_api_format == "claude:messages")
|
||||
.then(|| strip_gateway_credential_query_parameters(params.request_query))
|
||||
.flatten();
|
||||
let params = if normalized_provider_api_format == "claude:messages" {
|
||||
TransportRequestUrlParams {
|
||||
request_query: sanitized_claude_request_query.as_deref(),
|
||||
..params
|
||||
}
|
||||
} else {
|
||||
params
|
||||
};
|
||||
let is_claude_count_tokens = normalized_provider_api_format == "claude:messages"
|
||||
&& params.api_operation == Some(ApiOperation::ClaudeCountTokens);
|
||||
if !transport_supports_api_operation(
|
||||
transport,
|
||||
normalized_provider_api_format.as_str(),
|
||||
params.api_operation,
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
if is_claude_count_tokens {
|
||||
if let Some(url) = build_configured_claude_count_tokens_url(transport, params.request_query)
|
||||
{
|
||||
return Some(url);
|
||||
}
|
||||
}
|
||||
if let Some(url) = build_transport_hook_url(transport, params) {
|
||||
return Some(url);
|
||||
}
|
||||
|
||||
let custom_path = transport
|
||||
let custom_path_template = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|path| {
|
||||
expand_custom_path_template(path, build_path_params(params, gemini_embedding_batch))
|
||||
});
|
||||
.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))
|
||||
});
|
||||
|
||||
if let Some(path) = custom_path.as_deref() {
|
||||
let blocked_keys = if normalized_provider_api_format.starts_with("gemini:") {
|
||||
&["key"][..]
|
||||
let custom_path_is_complete_claude_count_tokens = normalized_provider_api_format
|
||||
== "claude:messages"
|
||||
&& !custom_path_handles_operation
|
||||
&& path
|
||||
.split_once('?')
|
||||
.map(|(path, _)| path)
|
||||
.unwrap_or(path)
|
||||
.trim_end_matches('/')
|
||||
.ends_with("/messages/count_tokens");
|
||||
let blocked_keys = if normalized_provider_api_format.starts_with("gemini:")
|
||||
|| normalized_provider_api_format == "claude:messages"
|
||||
{
|
||||
GATEWAY_CREDENTIAL_QUERY_KEYS
|
||||
} else {
|
||||
&[][..]
|
||||
};
|
||||
@@ -97,12 +140,19 @@ fn build_transport_request_url_inner(
|
||||
} else {
|
||||
path.to_string()
|
||||
};
|
||||
let url = build_passthrough_path_url(
|
||||
let mut url = build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
normalized_path.as_str(),
|
||||
params.request_query,
|
||||
blocked_keys,
|
||||
)?;
|
||||
if is_claude_count_tokens && !custom_path_handles_operation {
|
||||
url = build_default_claude_count_tokens_url(&url, None);
|
||||
} else if params.api_operation == Some(ApiOperation::ClaudeMessagesCreate)
|
||||
&& custom_path_is_complete_claude_count_tokens
|
||||
{
|
||||
url = build_claude_messages_url(&url, None);
|
||||
}
|
||||
return Some(maybe_add_gemini_stream_alt_sse(
|
||||
url,
|
||||
&provider_api_format,
|
||||
@@ -139,10 +189,14 @@ fn build_transport_request_url_inner(
|
||||
"openai:rerank" | "jina:rerank" => {
|
||||
build_provider_rerank_v1_url(&transport.endpoint.base_url, params.request_query)
|
||||
}
|
||||
"claude:messages" => Some(build_claude_messages_url(
|
||||
&transport.endpoint.base_url,
|
||||
params.request_query,
|
||||
)),
|
||||
"claude:messages" => Some(if is_claude_count_tokens {
|
||||
build_default_claude_count_tokens_url(
|
||||
&transport.endpoint.base_url,
|
||||
params.request_query,
|
||||
)
|
||||
} else {
|
||||
build_claude_messages_url(&transport.endpoint.base_url, params.request_query)
|
||||
}),
|
||||
"gemini:generate_content" => build_gemini_content_url(
|
||||
&transport.endpoint.base_url,
|
||||
params.mapped_model?,
|
||||
@@ -186,6 +240,7 @@ pub fn build_local_openai_chat_upstream_url(
|
||||
upstream_is_stream: false,
|
||||
request_query,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -206,6 +261,7 @@ pub fn build_cross_format_openai_chat_upstream_url(
|
||||
upstream_is_stream,
|
||||
request_query,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -228,6 +284,7 @@ pub fn build_local_openai_responses_upstream_url(
|
||||
upstream_is_stream: false,
|
||||
request_query,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -249,6 +306,7 @@ pub fn build_cross_format_openai_responses_upstream_url(
|
||||
upstream_is_stream,
|
||||
request_query,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -269,6 +327,7 @@ pub fn build_kiro_cross_format_upstream_url(
|
||||
upstream_is_stream,
|
||||
request_query,
|
||||
kiro_api_region: Some(api_region),
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -291,10 +350,15 @@ fn build_transport_hook_url(
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("claude_code")
|
||||
{
|
||||
return Some(build_claude_code_messages_url(
|
||||
&transport.endpoint.base_url,
|
||||
params.request_query,
|
||||
));
|
||||
let messages_url =
|
||||
build_claude_code_messages_url(&transport.endpoint.base_url, params.request_query);
|
||||
return Some(
|
||||
if params.api_operation == Some(aether_ai_formats::ApiOperation::ClaudeCountTokens) {
|
||||
build_default_claude_count_tokens_url(&messages_url, None)
|
||||
} else {
|
||||
messages_url
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
let normalized_provider_api_format =
|
||||
@@ -408,6 +472,9 @@ fn build_path_params(
|
||||
}
|
||||
let provider_api_format =
|
||||
aether_ai_formats::normalize_api_format_alias(params.provider_api_format);
|
||||
if let Some(operation) = params.api_operation {
|
||||
path_params.insert("operation", operation.as_str());
|
||||
}
|
||||
if provider_api_format == "gemini:generate_content" || provider_api_format == "gemini:embedding"
|
||||
{
|
||||
path_params.insert(
|
||||
@@ -428,6 +495,79 @@ fn build_path_params(
|
||||
path_params
|
||||
}
|
||||
|
||||
pub fn transport_supports_api_operation(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
operation: Option<ApiOperation>,
|
||||
) -> bool {
|
||||
if operation != Some(ApiOperation::ClaudeCountTokens) {
|
||||
return true;
|
||||
}
|
||||
if aether_ai_formats::normalize_api_format_alias(provider_api_format) != "claude:messages" {
|
||||
return false;
|
||||
}
|
||||
|
||||
anthropic_count_tokens_supported(transport)
|
||||
}
|
||||
|
||||
fn anthropic_count_tokens_supported(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
// Private message adapters do not implement Anthropic's token-counting
|
||||
// operation. A config flag cannot make their request envelopes compatible.
|
||||
if crate::kiro::is_kiro_provider_transport(transport)
|
||||
|| crate::grok::is_grok_provider_transport(transport)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
let Some(operations) = anthropic_transport_config_field(transport, "supported_operations")
|
||||
else {
|
||||
return true;
|
||||
};
|
||||
operations.as_array().is_some_and(|operations| {
|
||||
operations.iter().any(|operation| {
|
||||
operation
|
||||
.as_str()
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case("count_tokens"))
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn build_configured_claude_count_tokens_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let path = anthropic_transport_config_field(transport, "count_tokens_path")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
let path = path?;
|
||||
let path = path
|
||||
.starts_with('/')
|
||||
.then(|| path.to_string())
|
||||
.unwrap_or_else(|| format!("/{path}"));
|
||||
build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
path.as_str(),
|
||||
request_query,
|
||||
GATEWAY_CREDENTIAL_QUERY_KEYS,
|
||||
)
|
||||
}
|
||||
|
||||
fn anthropic_transport_config_field<'a>(
|
||||
transport: &'a GatewayProviderTransportSnapshot,
|
||||
field: &str,
|
||||
) -> Option<&'a Value> {
|
||||
let from_config = |config: Option<&'a Value>| {
|
||||
config
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|config| config.get("anthropic"))
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|anthropic| anthropic.get(field))
|
||||
};
|
||||
from_config(transport.endpoint.config.as_ref())
|
||||
.or_else(|| from_config(transport.provider.config.as_ref()))
|
||||
}
|
||||
|
||||
fn normalize_gemini_embedding_action_path(path: &str, batch: bool) -> String {
|
||||
if batch {
|
||||
path.replace(":embedContent", ":batchEmbedContents")
|
||||
@@ -571,6 +711,8 @@ fn custom_path_template_regex() -> &'static Regex {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use aether_ai_formats::ApiOperation;
|
||||
|
||||
use super::{
|
||||
build_kiro_cross_format_upstream_url, build_transport_request_url,
|
||||
build_transport_request_url_for_request_body, TransportRequestUrlParams,
|
||||
@@ -660,6 +802,7 @@ mod tests {
|
||||
upstream_is_stream: true,
|
||||
request_query: Some("foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("vertex hook url");
|
||||
@@ -697,6 +840,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("foo=bar&beta=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("vertex service account hook url");
|
||||
@@ -738,6 +882,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("foo=bar&beta=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
Some(&provider_request_body),
|
||||
)
|
||||
@@ -767,6 +912,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=blocked&beta=true&foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -792,6 +938,7 @@ mod tests {
|
||||
upstream_is_stream: true,
|
||||
request_query: Some("foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -839,6 +986,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
Some(&batch_body),
|
||||
)
|
||||
@@ -866,6 +1014,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("tenant=demo"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("openai responses url");
|
||||
@@ -890,6 +1039,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("tenant=demo"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("openai search url");
|
||||
@@ -917,6 +1067,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-key&foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("expanded custom path url");
|
||||
@@ -944,6 +1095,7 @@ mod tests {
|
||||
upstream_is_stream: true,
|
||||
request_query: Some("key=client-key&foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("stream custom path url");
|
||||
@@ -968,6 +1120,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("sync custom path url");
|
||||
@@ -992,6 +1145,7 @@ mod tests {
|
||||
upstream_is_stream: true,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("v1 stream custom path url");
|
||||
@@ -1019,6 +1173,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("fallback custom path url");
|
||||
@@ -1026,6 +1181,374 @@ mod tests {
|
||||
assert_eq!(url, "https://api.example.com/v1/messages/{model}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strips_gateway_query_key_from_claude_messages_url() {
|
||||
let transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://api.anthropic.example/v1?region=us",
|
||||
None,
|
||||
);
|
||||
|
||||
let url = build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=gateway-secret&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeMessagesCreate),
|
||||
},
|
||||
)
|
||||
.expect("messages url");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://api.anthropic.example/v1/messages?region=us&trace=1"
|
||||
);
|
||||
assert!(!url.contains("gateway-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routes_claude_count_tokens_as_an_operation_on_messages_format() {
|
||||
let mut transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://api.anthropic.example/v1",
|
||||
None,
|
||||
);
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic": {
|
||||
"supported_operations": ["messages", "count_tokens"]
|
||||
}
|
||||
}));
|
||||
|
||||
let url = build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=gateway-secret&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.expect("count_tokens url");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://api.anthropic.example/v1/messages/count_tokens?trace=1"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn count_tokens_default_preserves_custom_anthropic_prefix() {
|
||||
let transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://proxy.example/anthropic?key=base-secret&tenant=base",
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("KEY=client-secret&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://proxy.example/anthropic/messages/count_tokens?key=base-secret&tenant=base&trace=1"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn count_tokens_uses_operation_aware_custom_path() {
|
||||
let mut transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://proxy.example/anthropic",
|
||||
Some("/operations/{operation}"),
|
||||
);
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic": {"supported_operations": ["messages", "count_tokens"]}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-secret&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://proxy.example/anthropic/operations/count_tokens?trace=1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn count_tokens_derives_from_messages_only_custom_path() {
|
||||
let transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://api.anthropic.example",
|
||||
Some("/custom/v1/messages"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.anthropic.example/custom/v1/messages/count_tokens")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn count_tokens_keeps_complete_custom_count_tokens_path() {
|
||||
let transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://api.anthropic.example",
|
||||
Some("/custom/v1/messages/count_tokens?key=path-secret"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-secret&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://api.anthropic.example/custom/v1/messages/count_tokens?key=path-secret&trace=1"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn messages_create_normalizes_complete_custom_count_tokens_path() {
|
||||
let transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://api.anthropic.example",
|
||||
Some("/custom/v1/messages/count_tokens?key=path-secret"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-secret&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeMessagesCreate),
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.anthropic.example/custom/v1/messages?key=path-secret&trace=1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn private_anthropic_adapters_reject_count_tokens_even_when_config_claims_support() {
|
||||
for provider_type in ["kiro", "grok"] {
|
||||
let mut transport = sample_transport(
|
||||
provider_type,
|
||||
"claude:messages",
|
||||
"https://private.example",
|
||||
None,
|
||||
);
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic": {
|
||||
"supported_operations": ["messages", "count_tokens"],
|
||||
"count_tokens_path": "/v1/messages/count_tokens"
|
||||
}
|
||||
}));
|
||||
|
||||
assert!(build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: Some("us-east-1"),
|
||||
api_operation: Some(ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_code_custom_root_keeps_messages_and_count_tokens_on_v1_surface() {
|
||||
let mut transport = sample_transport(
|
||||
"claude_code",
|
||||
"claude:messages",
|
||||
"https://proxy.example?key=base-secret",
|
||||
None,
|
||||
);
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic": {"supported_operations": ["messages", "count_tokens"]}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-secret&trace=messages"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeMessagesCreate),
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://proxy.example/v1/messages?key=base-secret&trace=messages")
|
||||
);
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-secret&trace=count"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://proxy.example/v1/messages/count_tokens?key=base-secret&trace=count")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn count_tokens_config_fields_fall_back_from_endpoint_to_provider() {
|
||||
let mut transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://api.anthropic.example/v1?key=base-secret",
|
||||
None,
|
||||
);
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic": {"profile": "native_transparent"}
|
||||
}));
|
||||
transport.provider.config = Some(json!({
|
||||
"anthropic": {
|
||||
"supported_operations": ["messages", "count_tokens"],
|
||||
"count_tokens_path": "/v1/messages/provider_count_tokens?key=path-secret"
|
||||
}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-secret&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://api.anthropic.example/v1/messages/provider_count_tokens?key=path-secret&trace=1"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_count_tokens_when_provider_capability_excludes_it() {
|
||||
let mut transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://api.anthropic.example/v1",
|
||||
None,
|
||||
);
|
||||
transport.endpoint.config = Some(json!({
|
||||
"anthropic": {"profile": "native_transparent"}
|
||||
}));
|
||||
transport.provider.config = Some(json!({
|
||||
"anthropic": {"supported_operations": ["messages"]}
|
||||
}));
|
||||
|
||||
assert!(build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
},
|
||||
)
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_configured_query_credentials_and_strips_client_credentials() {
|
||||
let transport = sample_transport(
|
||||
"custom",
|
||||
"claude:messages",
|
||||
"https://proxy.example/anthropic?key=base-secret&tenant=base",
|
||||
Some("/messages?key=path-secret&variant=custom"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "claude:messages",
|
||||
mapped_model: Some("claude-sonnet-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("KEY=client-secret&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://proxy.example/anthropic/messages?key=path-secret&tenant=base&trace=1&variant=custom"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_cross_format_helper_uses_region_specific_generate_assistant_url() {
|
||||
let transport = sample_transport(
|
||||
@@ -1040,7 +1563,7 @@ mod tests {
|
||||
"claude-sonnet-4",
|
||||
"claude:messages",
|
||||
true,
|
||||
Some("conversationId=abc"),
|
||||
Some("key=gateway-secret&conversationId=abc"),
|
||||
"us-west-2",
|
||||
)
|
||||
.expect("kiro url");
|
||||
@@ -1049,6 +1572,7 @@ mod tests {
|
||||
"https://codewhisperer.us-west-2.amazonaws.com/generateAssistantResponse"
|
||||
));
|
||||
assert!(url.contains("conversationId=abc"));
|
||||
assert!(!url.contains("gateway-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1093,6 +1617,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("tenant=demo"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1107,6 +1632,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1121,6 +1647,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-key&foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1137,6 +1664,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1151,6 +1679,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1182,6 +1711,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-key&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1196,6 +1726,7 @@ mod tests {
|
||||
upstream_is_stream: true,
|
||||
request_query: Some("key=client-key&trace=2"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1221,6 +1752,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-key"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1246,6 +1778,7 @@ mod tests {
|
||||
upstream_is_stream: true,
|
||||
request_query: Some("key=client-aether-key&trace=1&beta=true"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1279,6 +1812,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1293,6 +1827,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1332,6 +1867,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-key&foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
Some(&batch_body),
|
||||
)
|
||||
@@ -1368,6 +1904,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
Some(&batch_body),
|
||||
)
|
||||
@@ -1397,6 +1934,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("tenant=demo"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1411,6 +1949,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1454,6 +1993,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("tenant=request&trace=1"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1468,6 +2008,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("trace=2"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1482,6 +2023,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("key=client-key&trace=3"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1496,6 +2038,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: Some("trace=4"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.as_deref(),
|
||||
@@ -1520,6 +2063,7 @@ mod tests {
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.is_none());
|
||||
@@ -1542,6 +2086,7 @@ mod tests {
|
||||
upstream_is_stream: true,
|
||||
request_query: Some("key=client-key&foo=bar"),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("expanded custom embedding path url");
|
||||
|
||||
@@ -3,15 +3,22 @@ use std::collections::BTreeMap;
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::anthropic_compat::{
|
||||
resolve_anthropic_compatibility_profile, AnthropicCompatibilityProfile,
|
||||
};
|
||||
use crate::antigravity::is_antigravity_provider_transport;
|
||||
use crate::auth::{
|
||||
build_complete_passthrough_headers, build_complete_passthrough_headers_with_auth,
|
||||
resolve_local_gemini_auth, resolve_local_openai_bearer_auth, resolve_local_standard_auth,
|
||||
replace_upstream_auth_headers, resolve_local_gemini_auth, resolve_local_openai_bearer_auth,
|
||||
resolve_local_standard_auth,
|
||||
};
|
||||
use crate::claude_code::{
|
||||
build_claude_code_passthrough_headers, current_claude_code_transport_identity_profile,
|
||||
local_claude_code_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::claude_code::build_claude_code_passthrough_headers;
|
||||
use crate::claude_code::local_claude_code_transport_unsupported_reason_with_network;
|
||||
use crate::gemini_cli::is_gemini_cli_provider_transport;
|
||||
use crate::grok::{is_grok_provider_transport, resolve_grok_session_auth};
|
||||
use crate::headers::{force_identity_accept_encoding, upstream_credential_header_names};
|
||||
use crate::kiro::{
|
||||
build_kiro_provider_headers, build_kiro_provider_request_body, is_kiro_provider_transport,
|
||||
local_kiro_request_transport_unsupported_reason_with_network, KiroAuthConfig,
|
||||
@@ -29,10 +36,8 @@ use crate::vertex::{
|
||||
is_vertex_service_account_transport_context, is_vertex_transport_context,
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::{
|
||||
build_transport_request_url_for_request_body, ensure_upstream_auth_header,
|
||||
TransportRequestUrlParams,
|
||||
};
|
||||
|
||||
use crate::{build_transport_request_url_for_request_body, TransportRequestUrlParams};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SameFormatProviderFamily {
|
||||
@@ -52,6 +57,8 @@ pub struct SameFormatProviderRequestBehavior {
|
||||
pub is_antigravity: bool,
|
||||
pub is_gemini_cli: bool,
|
||||
pub is_claude_code: bool,
|
||||
pub is_claude_code_transport: bool,
|
||||
pub anthropic_compatibility_profile: AnthropicCompatibilityProfile,
|
||||
pub is_vertex: bool,
|
||||
pub is_kiro: bool,
|
||||
pub upstream_is_stream: bool,
|
||||
@@ -106,6 +113,7 @@ pub struct SameFormatProviderUpstreamUrlParams<'a> {
|
||||
pub upstream_is_stream: bool,
|
||||
pub request_query: Option<&'a str>,
|
||||
pub kiro_api_region: Option<&'a str>,
|
||||
pub api_operation: Option<aether_ai_formats::ApiOperation>,
|
||||
pub provider_request_body: Option<&'a Value>,
|
||||
}
|
||||
|
||||
@@ -116,10 +124,10 @@ pub struct SameFormatProviderHeadersInput<'a> {
|
||||
pub original_request_body: &'a Value,
|
||||
pub header_rules: Option<&'a Value>,
|
||||
pub behavior: SameFormatProviderRequestBehavior,
|
||||
pub api_operation: Option<aether_ai_formats::ApiOperation>,
|
||||
pub auth_header: Option<&'a str>,
|
||||
pub auth_value: Option<&'a str>,
|
||||
pub extra_headers: &'a BTreeMap<String, String>,
|
||||
pub key_fingerprint: Option<&'a Value>,
|
||||
pub kiro_auth_config: Option<&'a KiroAuthConfig>,
|
||||
pub kiro_machine_id: Option<&'a str>,
|
||||
}
|
||||
@@ -127,14 +135,25 @@ pub struct SameFormatProviderHeadersInput<'a> {
|
||||
pub fn classify_same_format_provider_request_behavior(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
params: SameFormatProviderRequestBehaviorParams<'_>,
|
||||
) -> SameFormatProviderRequestBehavior {
|
||||
classify_same_format_provider_request_behavior_for_operation(transport, params, None)
|
||||
}
|
||||
|
||||
pub fn classify_same_format_provider_request_behavior_for_operation(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
params: SameFormatProviderRequestBehaviorParams<'_>,
|
||||
api_operation: Option<aether_ai_formats::ApiOperation>,
|
||||
) -> SameFormatProviderRequestBehavior {
|
||||
let is_antigravity = is_antigravity_provider_transport(transport);
|
||||
let is_gemini_cli = is_gemini_cli_provider_transport(transport);
|
||||
let is_claude_code = transport
|
||||
let is_claude_code_transport = transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("claude_code");
|
||||
let anthropic_compatibility_profile =
|
||||
resolve_anthropic_compatibility_profile(transport, params.provider_api_format);
|
||||
let is_claude_code = anthropic_compatibility_profile.uses_claude_code_compatibility();
|
||||
let is_vertex = is_vertex_transport_context(transport);
|
||||
let is_kiro = is_kiro_provider_transport(transport);
|
||||
let gemini_cli_requires_upstream_streaming = is_gemini_cli
|
||||
@@ -142,18 +161,23 @@ pub fn classify_same_format_provider_request_behavior(
|
||||
params.provider_api_format,
|
||||
params.require_streaming,
|
||||
);
|
||||
let upstream_is_stream = aether_ai_formats::resolve_upstream_is_stream_for_provider(
|
||||
transport.endpoint.config.as_ref(),
|
||||
transport.provider.provider_type.as_str(),
|
||||
params.provider_api_format,
|
||||
params.require_streaming,
|
||||
is_kiro || is_antigravity || gemini_cli_requires_upstream_streaming,
|
||||
let operation_requires_sync = matches!(
|
||||
api_operation,
|
||||
Some(aether_ai_formats::ApiOperation::ClaudeCountTokens)
|
||||
);
|
||||
let force_body_stream_field =
|
||||
aether_ai_formats::api_format_uses_body_stream_field(params.provider_api_format)
|
||||
&& aether_ai_formats::endpoint_config_forces_upstream_stream_policy(
|
||||
transport.endpoint.config.as_ref(),
|
||||
);
|
||||
let upstream_is_stream = !operation_requires_sync
|
||||
&& aether_ai_formats::resolve_upstream_is_stream_for_provider(
|
||||
transport.endpoint.config.as_ref(),
|
||||
transport.provider.provider_type.as_str(),
|
||||
params.provider_api_format,
|
||||
params.require_streaming,
|
||||
is_kiro || is_antigravity || gemini_cli_requires_upstream_streaming,
|
||||
);
|
||||
let force_body_stream_field = !operation_requires_sync
|
||||
&& aether_ai_formats::api_format_uses_body_stream_field(params.provider_api_format)
|
||||
&& aether_ai_formats::endpoint_config_forces_upstream_stream_policy(
|
||||
transport.endpoint.config.as_ref(),
|
||||
);
|
||||
let report_kind = if is_kiro && !params.require_streaming {
|
||||
"claude_cli_sync_finalize"
|
||||
} else if (is_gemini_cli || is_antigravity) && !params.require_streaming {
|
||||
@@ -170,6 +194,8 @@ pub fn classify_same_format_provider_request_behavior(
|
||||
is_antigravity,
|
||||
is_gemini_cli,
|
||||
is_claude_code,
|
||||
is_claude_code_transport,
|
||||
anthropic_compatibility_profile,
|
||||
is_vertex,
|
||||
is_kiro,
|
||||
upstream_is_stream,
|
||||
@@ -196,6 +222,20 @@ pub fn build_same_format_provider_request_body_with_compatibility_report(
|
||||
})
|
||||
}
|
||||
|
||||
pub fn enforce_same_format_provider_api_operation_body_policy(
|
||||
body: &mut Value,
|
||||
api_operation: Option<aether_ai_formats::ApiOperation>,
|
||||
) -> bool {
|
||||
if !matches!(
|
||||
api_operation,
|
||||
Some(aether_ai_formats::ApiOperation::ClaudeCountTokens)
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
body.as_object_mut()
|
||||
.is_some_and(|object| object.remove("stream").is_some())
|
||||
}
|
||||
|
||||
fn build_same_format_provider_request_body_inner(
|
||||
input: SameFormatProviderRequestBodyInput<'_>,
|
||||
mut compatibility_edits: Option<&mut Vec<SameFormatProviderCompatibilityEdit>>,
|
||||
@@ -503,6 +543,7 @@ pub fn build_same_format_provider_upstream_url(
|
||||
upstream_is_stream: params.upstream_is_stream,
|
||||
request_query: params.request_query,
|
||||
kiro_api_region: params.kiro_api_region,
|
||||
api_operation: params.api_operation,
|
||||
},
|
||||
params.provider_request_body,
|
||||
)
|
||||
@@ -526,14 +567,13 @@ pub fn build_same_format_provider_headers(
|
||||
|
||||
let auth_header = input.auth_header.unwrap_or_default();
|
||||
let auth_value = input.auth_value.unwrap_or_default();
|
||||
let mut provider_request_headers = if input.behavior.is_claude_code {
|
||||
let mut provider_request_headers = if input.behavior.is_claude_code_transport {
|
||||
build_claude_code_passthrough_headers(
|
||||
input.headers,
|
||||
auth_header,
|
||||
auth_value,
|
||||
input.extra_headers,
|
||||
input.behavior.upstream_is_stream,
|
||||
input.key_fingerprint,
|
||||
)
|
||||
} else if input.behavior.is_vertex {
|
||||
build_complete_passthrough_headers(
|
||||
@@ -551,11 +591,8 @@ pub fn build_same_format_provider_headers(
|
||||
)
|
||||
};
|
||||
|
||||
let protected_headers = input
|
||||
.auth_header
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.map(|value| vec![value, "content-type"])
|
||||
.unwrap_or_else(|| vec!["content-type"]);
|
||||
let mut protected_headers = upstream_credential_header_names().to_vec();
|
||||
protected_headers.push("content-type");
|
||||
if !apply_local_header_rules_with_request_headers(
|
||||
&mut provider_request_headers,
|
||||
input.header_rules,
|
||||
@@ -567,10 +604,28 @@ pub fn build_same_format_provider_headers(
|
||||
return None;
|
||||
}
|
||||
if let (Some(auth_header), Some(auth_value)) = (input.auth_header, input.auth_value) {
|
||||
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
|
||||
replace_upstream_auth_headers(&mut provider_request_headers, auth_header, auth_value);
|
||||
} else {
|
||||
replace_upstream_auth_headers(&mut provider_request_headers, "", "");
|
||||
}
|
||||
if input.behavior.upstream_is_stream {
|
||||
let claude_code_profile = *current_claude_code_transport_identity_profile();
|
||||
if input.behavior.is_claude_code_transport {
|
||||
claude_code_profile.apply_fixed_headers(
|
||||
&mut provider_request_headers,
|
||||
input.behavior.upstream_is_stream,
|
||||
);
|
||||
}
|
||||
if input.behavior.is_claude_code_transport || input.behavior.is_claude_code {
|
||||
claude_code_profile.apply_beta_policy(&mut provider_request_headers, input.api_operation);
|
||||
}
|
||||
if matches!(
|
||||
input.api_operation,
|
||||
Some(aether_ai_formats::ApiOperation::ClaudeCountTokens)
|
||||
) {
|
||||
provider_request_headers.insert("accept".to_string(), "application/json".to_string());
|
||||
} else if input.behavior.upstream_is_stream {
|
||||
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
|
||||
force_identity_accept_encoding(&mut provider_request_headers);
|
||||
}
|
||||
Some(provider_request_headers)
|
||||
}
|
||||
@@ -599,7 +654,7 @@ pub fn same_format_provider_transport_unsupported_reason(
|
||||
local_kiro_request_transport_unsupported_reason_with_network(transport)
|
||||
} else if behavior.is_antigravity {
|
||||
None
|
||||
} else if behavior.is_claude_code {
|
||||
} else if behavior.is_claude_code_transport {
|
||||
local_claude_code_transport_unsupported_reason_with_network(transport, api_format)
|
||||
} else if behavior.is_vertex {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network(transport)
|
||||
@@ -639,7 +694,7 @@ pub fn same_format_provider_transport_unsupported_reason_for_trace(
|
||||
},
|
||||
);
|
||||
if !behavior.is_antigravity
|
||||
&& !behavior.is_claude_code
|
||||
&& !behavior.is_claude_code_transport
|
||||
&& !behavior.is_gemini_cli
|
||||
&& !behavior.is_vertex
|
||||
&& !behavior.is_kiro
|
||||
@@ -825,6 +880,183 @@ mod tests {
|
||||
assert_eq!(behavior.report_kind, "gemini_cli_sync_finalize");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_compatibility_profile_controls_legacy_claude_code_edits() {
|
||||
let mut native = sample_transport("custom");
|
||||
native.endpoint.api_format = "claude:messages".to_string();
|
||||
let native_behavior = classify_same_format_provider_request_behavior(
|
||||
&native,
|
||||
SameFormatProviderRequestBehaviorParams {
|
||||
require_streaming: false,
|
||||
provider_api_format: "claude:messages",
|
||||
report_kind: "claude_chat_sync_success",
|
||||
},
|
||||
);
|
||||
assert!(!native_behavior.is_claude_code);
|
||||
assert!(!native_behavior.is_claude_code_transport);
|
||||
assert_eq!(
|
||||
native_behavior.anthropic_compatibility_profile,
|
||||
AnthropicCompatibilityProfile::NativeTransparent
|
||||
);
|
||||
|
||||
native.endpoint.config = Some(json!({
|
||||
"anthropic": {"compatibility_profile": "claude_code_legacy"}
|
||||
}));
|
||||
let compat_behavior = classify_same_format_provider_request_behavior(
|
||||
&native,
|
||||
SameFormatProviderRequestBehaviorParams {
|
||||
require_streaming: false,
|
||||
provider_api_format: "claude:messages",
|
||||
report_kind: "claude_chat_sync_success",
|
||||
},
|
||||
);
|
||||
assert!(compat_behavior.is_claude_code);
|
||||
assert!(!compat_behavior.is_claude_code_transport);
|
||||
assert_eq!(
|
||||
compat_behavior.anthropic_compatibility_profile,
|
||||
AnthropicCompatibilityProfile::ClaudeCodeLegacy
|
||||
);
|
||||
|
||||
let request_body = json!({
|
||||
"model": "claude-client",
|
||||
"thinking": {"type": "enabled"},
|
||||
"context_management": {"edits": [{"type": "client_strategy"}]},
|
||||
"system": [{
|
||||
"type": "text",
|
||||
"text": "x-anthropic-billing-header: cc_version=9.9.9.abc; cc_entrypoint=cli;"
|
||||
}],
|
||||
"messages": [{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "text", "text": "visible"},
|
||||
{"type": "thinking", "thinking": "unsigned", "signature": ""}
|
||||
]
|
||||
}]
|
||||
});
|
||||
let build_body = |behavior: SameFormatProviderRequestBehavior| {
|
||||
build_same_format_provider_request_body(SameFormatProviderRequestBodyInput {
|
||||
body_json: &request_body,
|
||||
mapped_model: "claude-upstream",
|
||||
client_api_format: "claude:messages",
|
||||
provider_api_format: "claude:messages",
|
||||
source_model: Some("claude-client"),
|
||||
family: SameFormatProviderFamily::Standard,
|
||||
body_rules: None,
|
||||
request_headers: None,
|
||||
upstream_is_stream: false,
|
||||
force_body_stream_field: false,
|
||||
kiro_auth_config: None,
|
||||
is_claude_code: behavior.is_claude_code,
|
||||
enable_model_directives: false,
|
||||
})
|
||||
.expect("body should build")
|
||||
};
|
||||
assert_eq!(
|
||||
build_body(native_behavior)["messages"][0]["content"]
|
||||
.as_array()
|
||||
.map(Vec::len),
|
||||
Some(2),
|
||||
"native Anthropic requests must remain untouched"
|
||||
);
|
||||
assert_eq!(
|
||||
build_body(native_behavior)["system"][0]["text"],
|
||||
request_body["system"][0]["text"],
|
||||
"native transparent requests must not rewrite billing identity"
|
||||
);
|
||||
assert_eq!(
|
||||
build_body(native_behavior)["context_management"],
|
||||
request_body["context_management"],
|
||||
"native transparent requests must not apply Claude Code body gates"
|
||||
);
|
||||
assert_eq!(
|
||||
build_body(compat_behavior)["messages"][0]["content"]
|
||||
.as_array()
|
||||
.map(Vec::len),
|
||||
Some(1),
|
||||
"legacy compatibility may sanitize invalid thinking blocks"
|
||||
);
|
||||
assert_eq!(
|
||||
build_body(compat_behavior)["system"][0]["text"],
|
||||
"x-anthropic-billing-header: cc_version=2.1.161.abc; cc_entrypoint=cli;"
|
||||
);
|
||||
|
||||
let mut legacy = sample_transport("claude_code");
|
||||
legacy.endpoint.api_format = "claude:messages".to_string();
|
||||
legacy.endpoint.config = Some(json!({
|
||||
"anthropic": {"compatibility_profile": "native_transparent"}
|
||||
}));
|
||||
let transparent_legacy_transport = classify_same_format_provider_request_behavior(
|
||||
&legacy,
|
||||
SameFormatProviderRequestBehaviorParams {
|
||||
require_streaming: false,
|
||||
provider_api_format: "claude:messages",
|
||||
report_kind: "claude_chat_sync_success",
|
||||
},
|
||||
);
|
||||
assert!(!transparent_legacy_transport.is_claude_code);
|
||||
assert!(transparent_legacy_transport.is_claude_code_transport);
|
||||
|
||||
let provider_request_body = json!({"model": "claude-upstream"});
|
||||
let empty_headers = http::HeaderMap::new();
|
||||
let empty_extra_headers = BTreeMap::new();
|
||||
let build_headers =
|
||||
|behavior: SameFormatProviderRequestBehavior,
|
||||
api_operation: Option<aether_ai_formats::ApiOperation>| {
|
||||
build_same_format_provider_headers(SameFormatProviderHeadersInput {
|
||||
headers: &empty_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &request_body,
|
||||
header_rules: None,
|
||||
behavior,
|
||||
api_operation,
|
||||
auth_header: Some("x-api-key"),
|
||||
auth_value: Some("upstream-secret"),
|
||||
extra_headers: &empty_extra_headers,
|
||||
kiro_auth_config: None,
|
||||
kiro_machine_id: None,
|
||||
})
|
||||
.expect("headers should build")
|
||||
};
|
||||
assert!(
|
||||
build_headers(compat_behavior, None).get("x-app").is_none(),
|
||||
"compatibility profile must not impersonate the Claude Code transport"
|
||||
);
|
||||
assert_eq!(
|
||||
build_headers(transparent_legacy_transport, None)
|
||||
.get("x-app")
|
||||
.map(String::as_str),
|
||||
Some("cli"),
|
||||
"Claude Code transport headers must survive a transparent body profile"
|
||||
);
|
||||
assert!(
|
||||
build_headers(
|
||||
native_behavior,
|
||||
Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
)
|
||||
.get("anthropic-beta")
|
||||
.is_none(),
|
||||
"native transparent token counting must not inject compatibility betas"
|
||||
);
|
||||
assert!(
|
||||
build_headers(
|
||||
compat_behavior,
|
||||
Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
)["anthropic-beta"]
|
||||
.split(',')
|
||||
.any(|token| token == "token-counting-2024-11-01"),
|
||||
"legacy compatibility token counting requires the token-counting beta"
|
||||
);
|
||||
assert!(
|
||||
build_headers(
|
||||
transparent_legacy_transport,
|
||||
Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
)["anthropic-beta"]
|
||||
.split(',')
|
||||
.any(|token| token == "token-counting-2024-11-01"),
|
||||
"Claude Code transport token counting requires the token-counting beta"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn same_format_behavior_resolves_endpoint_stream_policy() {
|
||||
let mut force_stream = sample_transport("openai");
|
||||
@@ -908,6 +1140,53 @@ mod tests {
|
||||
assert!(!search_behavior.force_body_stream_field);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn count_tokens_overrides_endpoint_stream_policy_and_removes_stream_field() {
|
||||
let mut transport = sample_transport("custom");
|
||||
transport.endpoint.api_format = "claude:messages".to_string();
|
||||
transport.endpoint.config = Some(json!({
|
||||
"upstream_stream_policy": "force_stream"
|
||||
}));
|
||||
let behavior = classify_same_format_provider_request_behavior_for_operation(
|
||||
&transport,
|
||||
SameFormatProviderRequestBehaviorParams {
|
||||
require_streaming: false,
|
||||
provider_api_format: "claude:messages",
|
||||
report_kind: "claude_count_tokens_sync_success",
|
||||
},
|
||||
Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
);
|
||||
|
||||
assert!(!behavior.upstream_is_stream);
|
||||
assert!(!behavior.force_body_stream_field);
|
||||
|
||||
let mut body = json!({"model": "claude-sonnet-4", "messages": [], "stream": true});
|
||||
assert!(enforce_same_format_provider_api_operation_body_policy(
|
||||
&mut body,
|
||||
Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
));
|
||||
assert!(body.get("stream").is_none());
|
||||
|
||||
let headers = build_same_format_provider_headers(SameFormatProviderHeadersInput {
|
||||
headers: &http::HeaderMap::new(),
|
||||
provider_request_body: &body,
|
||||
original_request_body: &body,
|
||||
header_rules: None,
|
||||
behavior,
|
||||
api_operation: Some(aether_ai_formats::ApiOperation::ClaudeCountTokens),
|
||||
auth_header: Some("x-api-key"),
|
||||
auth_value: Some("secret"),
|
||||
extra_headers: &BTreeMap::new(),
|
||||
kiro_auth_config: None,
|
||||
kiro_machine_id: None,
|
||||
})
|
||||
.expect("headers should build");
|
||||
assert_eq!(
|
||||
headers.get("accept").map(String::as_str),
|
||||
Some("application/json")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn same_format_behavior_preserves_hard_streaming_constraint() {
|
||||
let mut kiro = sample_transport("kiro");
|
||||
@@ -1857,8 +2136,13 @@ mod tests {
|
||||
fn builds_same_format_headers_with_auth_and_stream_accept() {
|
||||
let provider_request_body = json!({"model": "upstream-model"});
|
||||
let original_request_body = json!({"model": "client-model"});
|
||||
let mut request_headers = http::HeaderMap::new();
|
||||
request_headers.insert(
|
||||
http::header::ACCEPT_ENCODING,
|
||||
http::HeaderValue::from_static("gzip, br"),
|
||||
);
|
||||
let headers = build_same_format_provider_headers(SameFormatProviderHeadersInput {
|
||||
headers: &http::HeaderMap::new(),
|
||||
headers: &request_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &original_request_body,
|
||||
header_rules: None,
|
||||
@@ -1866,16 +2150,18 @@ mod tests {
|
||||
is_antigravity: false,
|
||||
is_gemini_cli: false,
|
||||
is_claude_code: false,
|
||||
is_claude_code_transport: false,
|
||||
anthropic_compatibility_profile: AnthropicCompatibilityProfile::NativeTransparent,
|
||||
is_vertex: false,
|
||||
is_kiro: false,
|
||||
upstream_is_stream: true,
|
||||
force_body_stream_field: false,
|
||||
report_kind: "openai_chat_stream_success",
|
||||
},
|
||||
api_operation: None,
|
||||
auth_header: Some("x-api-key"),
|
||||
auth_value: Some("secret"),
|
||||
extra_headers: &BTreeMap::new(),
|
||||
key_fingerprint: None,
|
||||
kiro_auth_config: None,
|
||||
kiro_machine_id: None,
|
||||
})
|
||||
@@ -1890,5 +2176,89 @@ mod tests {
|
||||
headers.get("accept").map(String::as_str),
|
||||
Some("text/event-stream")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("accept-encoding").map(String::as_str),
|
||||
Some("identity")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn same_format_headers_cannot_restore_credentials_or_internal_headers() {
|
||||
let provider_request_body = json!({"model": "upstream-model"});
|
||||
let original_request_body = json!({"model": "client-model"});
|
||||
let mut request_headers = http::HeaderMap::new();
|
||||
for (name, value) in [
|
||||
("authorization", "Bearer client"),
|
||||
("api-key", "client-api-key"),
|
||||
("x-api-key", "client-x-api-key"),
|
||||
("cookie", "session=client"),
|
||||
("proxy-authorization", "Basic client"),
|
||||
("x-aether-auth-user-id", "user-private"),
|
||||
("x-aether-auth-api-key-id", "key-private"),
|
||||
("x-aether-auth-balance-remaining", "12.34"),
|
||||
("x-aether-gateway", "gateway-internal"),
|
||||
] {
|
||||
request_headers.insert(
|
||||
http::HeaderName::from_bytes(name.as_bytes()).expect("valid header name"),
|
||||
http::HeaderValue::from_str(value).expect("valid header value"),
|
||||
);
|
||||
}
|
||||
let behavior = SameFormatProviderRequestBehavior {
|
||||
is_antigravity: false,
|
||||
is_gemini_cli: false,
|
||||
is_claude_code: false,
|
||||
is_claude_code_transport: false,
|
||||
anthropic_compatibility_profile: AnthropicCompatibilityProfile::NativeTransparent,
|
||||
is_vertex: false,
|
||||
is_kiro: false,
|
||||
upstream_is_stream: false,
|
||||
force_body_stream_field: false,
|
||||
report_kind: "claude_chat_sync_success",
|
||||
};
|
||||
let header_rules = json!([
|
||||
{"action": "set", "key": "cookie", "value": "session=rule"},
|
||||
{"action": "set", "key": "authorization", "value": "Bearer rule"},
|
||||
{"action": "set", "key": "x-aether-auth-user-id", "value": "user-rule"}
|
||||
]);
|
||||
let extra_headers = BTreeMap::from([
|
||||
("api-key".to_string(), "extra-api-key".to_string()),
|
||||
("authorization".to_string(), "Bearer extra".to_string()),
|
||||
(
|
||||
"x-aether-auth-balance-remaining".to_string(),
|
||||
"99.99".to_string(),
|
||||
),
|
||||
("anthropic-beta".to_string(), "custom-beta".to_string()),
|
||||
]);
|
||||
|
||||
let headers = build_same_format_provider_headers(SameFormatProviderHeadersInput {
|
||||
headers: &request_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &original_request_body,
|
||||
header_rules: Some(&header_rules),
|
||||
behavior,
|
||||
api_operation: None,
|
||||
auth_header: Some("x-api-key"),
|
||||
auth_value: Some("upstream-secret"),
|
||||
extra_headers: &extra_headers,
|
||||
kiro_auth_config: None,
|
||||
kiro_machine_id: None,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("x-api-key").map(String::as_str),
|
||||
Some("upstream-secret")
|
||||
);
|
||||
for stripped in ["authorization", "api-key", "cookie", "proxy-authorization"] {
|
||||
assert!(!headers.contains_key(stripped), "should strip {stripped}");
|
||||
}
|
||||
assert!(
|
||||
headers.keys().all(|name| !name.starts_with("x-aether-")),
|
||||
"Aether-owned headers must never leave provider egress: {headers:?}"
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("anthropic-beta").map(String::as_str),
|
||||
Some("custom-beta")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,6 +3,25 @@ use std::collections::BTreeMap;
|
||||
use url::form_urlencoded;
|
||||
use url::Url;
|
||||
|
||||
pub(crate) const GATEWAY_CREDENTIAL_QUERY_KEYS: &[&str] = &["key"];
|
||||
|
||||
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());
|
||||
let mut retained = false;
|
||||
for (key, value) in form_urlencoded::parse(query.as_bytes()) {
|
||||
if GATEWAY_CREDENTIAL_QUERY_KEYS
|
||||
.iter()
|
||||
.any(|blocked| key.eq_ignore_ascii_case(blocked))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
serializer.append_pair(&key, &value);
|
||||
retained = true;
|
||||
}
|
||||
retained.then(|| serializer.finish())
|
||||
}
|
||||
|
||||
pub fn build_openai_chat_url(upstream_base_url: &str, query: Option<&str>) -> String {
|
||||
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed = trimmed.trim_end_matches('/');
|
||||
@@ -72,10 +91,61 @@ fn openai_image_base_includes_operation_path(base_url: &str) -> bool {
|
||||
}
|
||||
|
||||
pub fn build_claude_messages_url(upstream_base_url: &str, query: Option<&str>) -> String {
|
||||
build_claude_messages_operation_url(upstream_base_url, "", query)
|
||||
}
|
||||
|
||||
pub(crate) fn build_claude_count_tokens_url(
|
||||
upstream_base_url: &str,
|
||||
query: Option<&str>,
|
||||
) -> String {
|
||||
build_claude_messages_operation_url(upstream_base_url, "/count_tokens", query)
|
||||
}
|
||||
|
||||
fn build_claude_messages_operation_url(
|
||||
upstream_base_url: &str,
|
||||
operation_suffix: &str,
|
||||
query: Option<&str>,
|
||||
) -> String {
|
||||
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed = trimmed.trim_end_matches('/');
|
||||
let mut url = format!("{trimmed}/messages");
|
||||
append_merged_query(&mut url, base_query, None, query, &[]);
|
||||
let parsed_url = Url::parse(trimmed).ok();
|
||||
let parsed_path = parsed_url
|
||||
.as_ref()
|
||||
.map(|url| url.path().trim_matches('/'))
|
||||
.unwrap_or_default();
|
||||
let base_includes_count_tokens =
|
||||
parsed_path == "messages/count_tokens" || parsed_path.ends_with("/messages/count_tokens");
|
||||
let messages_base = if base_includes_count_tokens {
|
||||
trimmed
|
||||
.strip_suffix("/count_tokens")
|
||||
.unwrap_or(trimmed)
|
||||
.to_string()
|
||||
} else if parsed_path.is_empty()
|
||||
&& parsed_url
|
||||
.as_ref()
|
||||
.and_then(Url::host_str)
|
||||
.is_some_and(|host| host.eq_ignore_ascii_case("api.anthropic.com"))
|
||||
{
|
||||
format!("{trimmed}/v1/messages")
|
||||
} else if parsed_path.is_empty() {
|
||||
format!("{trimmed}/messages")
|
||||
} else if parsed_path.rsplit('/').next() == Some("messages") {
|
||||
trimmed.to_string()
|
||||
} else {
|
||||
format!("{trimmed}/messages")
|
||||
};
|
||||
let mut url = if base_includes_count_tokens && operation_suffix == "/count_tokens" {
|
||||
trimmed.to_string()
|
||||
} else {
|
||||
format!("{messages_base}{operation_suffix}")
|
||||
};
|
||||
append_merged_query(
|
||||
&mut url,
|
||||
base_query,
|
||||
None,
|
||||
query,
|
||||
GATEWAY_CREDENTIAL_QUERY_KEYS,
|
||||
);
|
||||
url
|
||||
}
|
||||
|
||||
@@ -358,9 +428,10 @@ fn append_merged_query(
|
||||
base_query: Option<&str>,
|
||||
path_query: Option<&str>,
|
||||
request_query: Option<&str>,
|
||||
blocked_keys: &[&str],
|
||||
blocked_request_keys: &[&str],
|
||||
) {
|
||||
let Some(query) = merge_query_layers(base_query, path_query, request_query, blocked_keys)
|
||||
let Some(query) =
|
||||
merge_query_layers(base_query, path_query, request_query, blocked_request_keys)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
@@ -376,9 +447,9 @@ fn merge_query_layers(
|
||||
base_query: Option<&str>,
|
||||
path_query: Option<&str>,
|
||||
request_query: Option<&str>,
|
||||
blocked_keys: &[&str],
|
||||
blocked_request_keys: &[&str],
|
||||
) -> Option<String> {
|
||||
if blocked_keys.is_empty()
|
||||
if blocked_request_keys.is_empty()
|
||||
&& path_query.is_none()
|
||||
&& base_query.is_none()
|
||||
&& request_query
|
||||
@@ -392,9 +463,9 @@ fn merge_query_layers(
|
||||
}
|
||||
|
||||
let mut merged = BTreeMap::new();
|
||||
for source in [base_query, path_query, request_query] {
|
||||
merge_query_string(&mut merged, source, blocked_keys);
|
||||
}
|
||||
merge_query_string(&mut merged, base_query, &[]);
|
||||
merge_query_string(&mut merged, path_query, &[]);
|
||||
merge_query_string(&mut merged, request_query, blocked_request_keys);
|
||||
if merged.is_empty() {
|
||||
return None;
|
||||
}
|
||||
@@ -429,11 +500,11 @@ fn merge_query_string(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_bigmodel_coding_models_url, build_claude_messages_url, build_gemini_content_url,
|
||||
build_gemini_files_passthrough_url, build_gemini_video_predict_long_running_url,
|
||||
build_openai_chat_url, build_openai_compatible_models_url, build_openai_image_url,
|
||||
build_openai_responses_url, build_openai_search_url, build_passthrough_path_url,
|
||||
normalize_gemini_content_action_path,
|
||||
build_bigmodel_coding_models_url, build_claude_count_tokens_url, build_claude_messages_url,
|
||||
build_gemini_content_url, build_gemini_files_passthrough_url,
|
||||
build_gemini_video_predict_long_running_url, build_openai_chat_url,
|
||||
build_openai_compatible_models_url, build_openai_image_url, build_openai_responses_url,
|
||||
build_openai_search_url, build_passthrough_path_url, normalize_gemini_content_action_path,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -534,6 +605,54 @@ mod tests {
|
||||
build_claude_messages_url("https://api.anthropic.example", None),
|
||||
"https://api.anthropic.example/messages"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_messages_url("https://api.anthropic.com", None),
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_messages_url(
|
||||
"https://proxy.example.com/anthropic?key=base-secret&tenant=base",
|
||||
Some("KEY=request-secret&trace=1")
|
||||
),
|
||||
"https://proxy.example.com/anthropic/messages?key=base-secret&tenant=base&trace=1"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_count_tokens_url_preserves_configured_query_and_strips_request_credentials() {
|
||||
assert_eq!(
|
||||
build_claude_count_tokens_url(
|
||||
"https://proxy.example.com/anthropic?key=base-secret&tenant=base",
|
||||
Some("key=request-secret&trace=1")
|
||||
),
|
||||
"https://proxy.example.com/anthropic/messages/count_tokens?key=base-secret&tenant=base&trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_count_tokens_url("https://api.anthropic.com", None),
|
||||
"https://api.anthropic.com/v1/messages/count_tokens"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_count_tokens_url("https://api.anthropic.example", None),
|
||||
"https://api.anthropic.example/messages/count_tokens"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_count_tokens_url(
|
||||
"https://proxy.example.com/anthropic/messages",
|
||||
Some("trace=1")
|
||||
),
|
||||
"https://proxy.example.com/anthropic/messages/count_tokens?trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_count_tokens_url(
|
||||
"https://proxy.example.com/v1/messages/count_tokens?key=base-secret",
|
||||
Some("key=request-secret&trace=1")
|
||||
),
|
||||
"https://proxy.example.com/v1/messages/count_tokens?key=base-secret&trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_messages_url("https://proxy.example.com/v1/messages/count_tokens", None),
|
||||
"https://proxy.example.com/v1/messages"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -9,7 +9,7 @@ use rsa::pkcs8::DecodePrivateKey;
|
||||
use rsa::signature::{SignatureEncoding, Signer};
|
||||
use rsa::RsaPrivateKey;
|
||||
use serde_json::{json, Value};
|
||||
use sha2::Sha256;
|
||||
use sha2::{Digest, Sha256};
|
||||
use url::form_urlencoded;
|
||||
|
||||
use super::super::oauth_refresh::{
|
||||
@@ -153,7 +153,7 @@ impl LocalOAuthRefreshAdapter for VertexServiceAccountRefreshAdapter {
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
if !entry
|
||||
@@ -162,6 +162,9 @@ impl LocalOAuthRefreshAdapter for VertexServiceAccountRefreshAdapter {
|
||||
{
|
||||
return None;
|
||||
}
|
||||
if !vertex_service_account_cached_entry_matches_transport(transport, entry) {
|
||||
return None;
|
||||
}
|
||||
if service_account_token_expires_soon(entry.expires_at_unix_secs) {
|
||||
return None;
|
||||
}
|
||||
@@ -194,6 +197,34 @@ impl LocalOAuthRefreshAdapter for VertexServiceAccountRefreshAdapter {
|
||||
.is_none()
|
||||
}
|
||||
|
||||
fn refresh_fingerprint(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
let source_fingerprint = vertex_service_account_credential_fingerprint(transport)?;
|
||||
Some(
|
||||
entry
|
||||
.filter(|entry| {
|
||||
vertex_service_account_cached_entry_matches_transport(transport, entry)
|
||||
})
|
||||
.map(|entry| {
|
||||
let mut digest = Sha256::new();
|
||||
digest.update(source_fingerprint.as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(entry.auth_header_value.as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(entry.expires_at_unix_secs.unwrap_or_default().to_be_bytes());
|
||||
format!("{:x}", digest.finalize())
|
||||
})
|
||||
.unwrap_or(source_fingerprint),
|
||||
)
|
||||
}
|
||||
|
||||
fn shares_refresh_through_transport_persistence(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
@@ -259,11 +290,48 @@ impl LocalOAuthRefreshAdapter for VertexServiceAccountRefreshAdapter {
|
||||
"project_id": auth_config.project_id,
|
||||
"client_email": auth_config.client_email,
|
||||
})),
|
||||
source_fingerprint: None,
|
||||
source_fingerprint: vertex_service_account_credential_fingerprint(transport),
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
fn vertex_service_account_credential_fingerprint(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<String> {
|
||||
supports_local_vertex_service_account_auth_resolution(transport).then(|| {
|
||||
let provider_type = transport.provider.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
|
||||
let auth_config = transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.unwrap_or_default();
|
||||
let mut digest = Sha256::new();
|
||||
for field in [
|
||||
provider_type.as_bytes(),
|
||||
auth_type.as_bytes(),
|
||||
auth_config.as_bytes(),
|
||||
transport.key.decrypted_api_key.as_bytes(),
|
||||
] {
|
||||
digest.update((field.len() as u64).to_be_bytes());
|
||||
digest.update(field);
|
||||
}
|
||||
format!("{:x}", digest.finalize())
|
||||
})
|
||||
}
|
||||
|
||||
fn vertex_service_account_cached_entry_matches_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> bool {
|
||||
entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE)
|
||||
&& vertex_service_account_credential_fingerprint(transport)
|
||||
.as_deref()
|
||||
.is_some_and(|fingerprint| entry.source_fingerprint.as_deref() == Some(fingerprint))
|
||||
}
|
||||
|
||||
pub fn build_vertex_service_account_assertion(
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
now_unix_secs: u64,
|
||||
@@ -324,6 +392,9 @@ fn body_excerpt(value: &str) -> String {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::super::oauth_refresh::{
|
||||
CachedOAuthEntry, LocalOAuthRefreshAdapter, LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
@@ -335,7 +406,10 @@ mod tests {
|
||||
use super::{
|
||||
decode_vertex_service_account_private_key, parse_vertex_service_account_auth_config,
|
||||
resolve_local_vertex_api_key_query_auth,
|
||||
supports_local_vertex_service_account_auth_resolution, VERTEX_API_KEY_QUERY_PARAM,
|
||||
supports_local_vertex_service_account_auth_resolution,
|
||||
vertex_service_account_credential_fingerprint, VertexServiceAccountRefreshAdapter,
|
||||
VERTEX_API_KEY_QUERY_PARAM, VERTEX_SERVICE_ACCOUNT_AUTH_HEADER,
|
||||
VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
@@ -395,6 +469,59 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_service_account_transport(private_key: &str) -> GatewayProviderTransportSnapshot {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
serde_json::json!({
|
||||
"client_email": "[email protected]",
|
||||
"private_key": private_key,
|
||||
"project_id": "demo-project"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
transport
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cached_token_is_bound_to_vertex_service_account_generation() {
|
||||
let source_transport = sample_service_account_transport("SOURCE-PRIVATE-KEY");
|
||||
let source_fingerprint = vertex_service_account_credential_fingerprint(&source_transport)
|
||||
.expect("source service account should have a fingerprint");
|
||||
let entry = CachedOAuthEntry {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE.to_string(),
|
||||
auth_header_name: VERTEX_SERVICE_ACCOUNT_AUTH_HEADER.to_string(),
|
||||
auth_header_value: "Bearer source-access-token".to_string(),
|
||||
expires_at_unix_secs: Some(u64::MAX),
|
||||
metadata: None,
|
||||
source_fingerprint: Some(source_fingerprint),
|
||||
};
|
||||
let adapter = VertexServiceAccountRefreshAdapter;
|
||||
|
||||
assert!(!adapter.shares_refresh_through_transport_persistence());
|
||||
assert_eq!(
|
||||
adapter.resolve_cached(&source_transport, &entry),
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: VERTEX_SERVICE_ACCOUNT_AUTH_HEADER.to_string(),
|
||||
value: "Bearer source-access-token".to_string(),
|
||||
})
|
||||
);
|
||||
assert_ne!(
|
||||
adapter.refresh_fingerprint(&source_transport, None),
|
||||
adapter.refresh_fingerprint(&source_transport, Some(&entry))
|
||||
);
|
||||
|
||||
let replacement_transport = sample_service_account_transport("ADMIN-PRIVATE-KEY");
|
||||
assert!(adapter
|
||||
.resolve_cached(&replacement_transport, &entry)
|
||||
.is_none());
|
||||
assert_eq!(
|
||||
adapter.refresh_fingerprint(&replacement_transport, Some(&entry)),
|
||||
vertex_service_account_credential_fingerprint(&replacement_transport)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_query_auth_for_vertex_api_key_subset() {
|
||||
let auth = resolve_local_vertex_api_key_query_auth(&sample_transport())
|
||||
|
||||
Reference in New Issue
Block a user