mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor ai serving modules and crates
This commit is contained in:
@@ -9,8 +9,9 @@ pub use auth::{
|
||||
ANTIGRAVITY_PROVIDER_TYPE, ANTIGRAVITY_REQUEST_USER_AGENT,
|
||||
};
|
||||
pub use policy::{
|
||||
classify_local_antigravity_request_support, AntigravityRequestSideSpec,
|
||||
AntigravityRequestSideSupport, AntigravityRequestSideUnsupportedReason,
|
||||
classify_local_antigravity_request_support, is_antigravity_provider_transport,
|
||||
AntigravityRequestSideSpec, AntigravityRequestSideSupport,
|
||||
AntigravityRequestSideUnsupportedReason,
|
||||
};
|
||||
pub use request::{
|
||||
build_antigravity_safe_v1internal_request, classify_antigravity_safe_request_body,
|
||||
|
||||
@@ -35,6 +35,14 @@ pub enum AntigravityRequestSideUnsupportedReason {
|
||||
UnsupportedEnvelope(AntigravityRequestEnvelopeUnsupportedReason),
|
||||
}
|
||||
|
||||
pub fn is_antigravity_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(ANTIGRAVITY_PROVIDER_TYPE)
|
||||
}
|
||||
|
||||
pub fn classify_local_antigravity_request_support(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request_body: &Value,
|
||||
@@ -45,12 +53,7 @@ pub fn classify_local_antigravity_request_support(
|
||||
AntigravityRequestSideUnsupportedReason::InactiveTransport,
|
||||
);
|
||||
}
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(ANTIGRAVITY_PROVIDER_TYPE)
|
||||
{
|
||||
if !is_antigravity_provider_transport(transport) {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::WrongProviderType,
|
||||
);
|
||||
|
||||
590
crates/aether-provider-transport/src/conversion.rs
Normal file
590
crates/aether-provider-transport/src/conversion.rs
Normal file
@@ -0,0 +1,590 @@
|
||||
use aether_ai_formats::matrix::{
|
||||
request_conversion_kind, request_conversion_requires_enable_flag, RequestConversionKind,
|
||||
};
|
||||
use aether_ai_formats::normalize_api_format_alias;
|
||||
|
||||
use crate::auth::{
|
||||
resolve_local_gemini_auth, resolve_local_openai_bearer_auth, resolve_local_standard_auth,
|
||||
};
|
||||
use crate::kiro::{
|
||||
is_kiro_claude_messages_transport, local_kiro_request_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::policy::{
|
||||
local_gemini_transport_unsupported_reason_with_network,
|
||||
local_openai_chat_transport_unsupported_reason,
|
||||
local_standard_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::vertex::{
|
||||
is_vertex_api_key_transport_context,
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
|
||||
resolve_local_vertex_api_key_query_auth, VERTEX_API_KEY_QUERY_PARAM,
|
||||
};
|
||||
use crate::GatewayProviderTransportSnapshot;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct CandidateTransportPolicyFacts<'a> {
|
||||
pub endpoint_api_format: &'a str,
|
||||
pub global_model_name: &'a str,
|
||||
pub selected_provider_model_name: &'a str,
|
||||
pub mapping_matched_model: Option<&'a str>,
|
||||
}
|
||||
|
||||
pub fn request_conversion_enabled_for_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
client_api_format: &str,
|
||||
provider_api_format: &str,
|
||||
) -> bool {
|
||||
let client_api_format = normalize_api_format_alias(client_api_format);
|
||||
let provider_api_format = normalize_api_format_alias(provider_api_format);
|
||||
if client_api_format == provider_api_format {
|
||||
return true;
|
||||
}
|
||||
if request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str()).is_none() {
|
||||
return false;
|
||||
}
|
||||
if !request_conversion_requires_enable_flag(
|
||||
client_api_format.as_str(),
|
||||
provider_api_format.as_str(),
|
||||
) {
|
||||
return true;
|
||||
}
|
||||
transport.provider.enable_format_conversion
|
||||
|| endpoint_accepts_client_api_format(transport, client_api_format.as_str())
|
||||
}
|
||||
|
||||
pub fn request_pair_allowed_for_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
client_api_format: &str,
|
||||
provider_api_format: &str,
|
||||
) -> bool {
|
||||
let client_api_format = normalize_api_format_alias(client_api_format);
|
||||
let provider_api_format = normalize_api_format_alias(provider_api_format);
|
||||
if client_api_format == provider_api_format {
|
||||
return true;
|
||||
}
|
||||
if request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str()).is_none() {
|
||||
return false;
|
||||
}
|
||||
if is_kiro_claude_messages_transport(transport, &provider_api_format) {
|
||||
return request_conversion_enabled_for_transport(
|
||||
transport,
|
||||
client_api_format.as_str(),
|
||||
provider_api_format.as_str(),
|
||||
) && local_kiro_request_transport_unsupported_reason_with_network(transport)
|
||||
.is_none();
|
||||
}
|
||||
request_conversion_enabled_for_transport(
|
||||
transport,
|
||||
client_api_format.as_str(),
|
||||
provider_api_format.as_str(),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn request_conversion_transport_supported(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
kind: RequestConversionKind,
|
||||
) -> bool {
|
||||
request_conversion_transport_unsupported_reason(transport, kind).is_none()
|
||||
}
|
||||
|
||||
pub fn request_conversion_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
_kind: RequestConversionKind,
|
||||
) -> Option<&'static str> {
|
||||
if is_kiro_claude_messages_transport(transport, &transport.endpoint.api_format) {
|
||||
return local_kiro_request_transport_unsupported_reason_with_network(transport);
|
||||
}
|
||||
|
||||
match normalize_api_format_alias(&transport.endpoint.api_format).as_str() {
|
||||
"openai:chat" => local_openai_chat_transport_unsupported_reason(transport),
|
||||
"openai:responses" | "openai:responses:compact" => {
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
transport.endpoint.api_format.trim(),
|
||||
)
|
||||
}
|
||||
"claude:messages" => {
|
||||
local_standard_transport_unsupported_reason_with_network(transport, "claude:messages")
|
||||
}
|
||||
"gemini:generate_content" if is_vertex_api_key_transport_context(transport) => {
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(transport)
|
||||
}
|
||||
"gemini:generate_content" => local_gemini_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
"gemini:generate_content",
|
||||
),
|
||||
_ => Some("transport_api_format_unsupported"),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn request_conversion_direct_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
_kind: RequestConversionKind,
|
||||
) -> Option<(String, String)> {
|
||||
match normalize_api_format_alias(&transport.endpoint.api_format).as_str() {
|
||||
"openai:chat" | "openai:responses" | "openai:responses:compact" => {
|
||||
resolve_local_openai_bearer_auth(transport)
|
||||
}
|
||||
"gemini:generate_content" => {
|
||||
if is_vertex_api_key_transport_context(transport) {
|
||||
resolve_local_vertex_api_key_query_auth(transport)
|
||||
.map(|auth| (VERTEX_API_KEY_QUERY_PARAM.to_string(), auth.value))
|
||||
} else {
|
||||
resolve_local_gemini_auth(transport)
|
||||
}
|
||||
}
|
||||
"claude:messages" => resolve_local_standard_auth(transport),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn candidate_common_transport_skip_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
candidate: CandidateTransportPolicyFacts<'_>,
|
||||
requested_model: Option<&str>,
|
||||
) -> Option<&'static str> {
|
||||
let requested_model = requested_model.unwrap_or_default();
|
||||
|
||||
if !transport.provider.is_active {
|
||||
return Some("provider_inactive");
|
||||
}
|
||||
if !transport.endpoint.is_active {
|
||||
return Some("endpoint_inactive");
|
||||
}
|
||||
if !transport.key.is_active {
|
||||
return Some("key_inactive");
|
||||
}
|
||||
|
||||
let endpoint_api_format = transport.endpoint.api_format.trim();
|
||||
if !aether_ai_formats::api_format_alias_matches(
|
||||
candidate.endpoint_api_format,
|
||||
endpoint_api_format,
|
||||
) {
|
||||
return Some("endpoint_api_format_changed");
|
||||
}
|
||||
|
||||
if !transport_key_supports_api_format(transport, endpoint_api_format) {
|
||||
return Some("key_api_format_disabled");
|
||||
}
|
||||
if !transport_key_allows_candidate_model(transport, requested_model, candidate) {
|
||||
return Some("key_model_disabled");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub fn candidate_transport_pair_skip_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
normalized_client_api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
let endpoint_api_format = transport.endpoint.api_format.trim();
|
||||
if aether_ai_formats::api_format_alias_matches(
|
||||
endpoint_api_format,
|
||||
normalized_client_api_format,
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
|
||||
if let Some(skip_reason) =
|
||||
disabled_format_conversion_skip_reason(transport, normalized_client_api_format)
|
||||
{
|
||||
return Some(skip_reason);
|
||||
}
|
||||
|
||||
if !request_pair_allowed_for_transport(
|
||||
transport,
|
||||
normalized_client_api_format,
|
||||
endpoint_api_format,
|
||||
) {
|
||||
return Some("transport_unsupported");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
fn disabled_format_conversion_skip_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
normalized_client_api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
let endpoint_api_format = transport.endpoint.api_format.trim();
|
||||
if aether_ai_formats::api_format_alias_matches(
|
||||
endpoint_api_format,
|
||||
normalized_client_api_format,
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
|
||||
request_conversion_kind(normalized_client_api_format, endpoint_api_format)?;
|
||||
|
||||
if request_conversion_requires_enable_flag(normalized_client_api_format, endpoint_api_format)
|
||||
&& !request_conversion_enabled_for_transport(
|
||||
transport,
|
||||
normalized_client_api_format,
|
||||
endpoint_api_format,
|
||||
)
|
||||
{
|
||||
return Some("format_conversion_disabled");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
fn transport_key_supports_api_format(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
endpoint_api_format: &str,
|
||||
) -> bool {
|
||||
let inherits_provider_api_formats =
|
||||
crate::provider_types::fixed_provider_key_inherits_api_formats(
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.key.auth_type.as_str(),
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
if inherits_provider_api_formats {
|
||||
return true;
|
||||
}
|
||||
|
||||
match transport.key.api_formats.as_deref() {
|
||||
None => true,
|
||||
Some(formats) => formats
|
||||
.iter()
|
||||
.any(|value| aether_ai_formats::api_format_alias_matches(value, endpoint_api_format)),
|
||||
}
|
||||
}
|
||||
|
||||
fn transport_key_allows_candidate_model(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
requested_model: &str,
|
||||
candidate: CandidateTransportPolicyFacts<'_>,
|
||||
) -> bool {
|
||||
let Some(allowed_models) = transport.key.allowed_models.as_deref() else {
|
||||
return true;
|
||||
};
|
||||
|
||||
let requested_model = requested_model.trim();
|
||||
let global_model_name = candidate.global_model_name.trim();
|
||||
let selected_provider_model_name = candidate.selected_provider_model_name.trim();
|
||||
let mapping_matched_model = candidate
|
||||
.mapping_matched_model
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
for allowed_model in allowed_models.iter().map(String::as_str).map(str::trim) {
|
||||
if allowed_model.is_empty() {
|
||||
continue;
|
||||
}
|
||||
if allowed_model == requested_model
|
||||
|| allowed_model == global_model_name
|
||||
|| allowed_model == selected_provider_model_name
|
||||
|| mapping_matched_model.is_some_and(|value| value == allowed_model)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
|
||||
fn endpoint_accepts_client_api_format(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
client_api_format: &str,
|
||||
) -> bool {
|
||||
let Some(config) = transport
|
||||
.endpoint
|
||||
.format_acceptance_config
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
if !config
|
||||
.get("enabled")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
if config
|
||||
.get("reject_formats")
|
||||
.is_some_and(|value| json_format_list_contains(value, client_api_format))
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
match config.get("accept_formats") {
|
||||
Some(value) => json_format_list_contains(value, client_api_format),
|
||||
None => true,
|
||||
}
|
||||
}
|
||||
|
||||
fn json_format_list_contains(value: &serde_json::Value, api_format: &str) -> bool {
|
||||
let Some(items) = value.as_array() else {
|
||||
return false;
|
||||
};
|
||||
items.iter().any(|item| {
|
||||
item.as_str().is_some_and(|candidate| {
|
||||
aether_ai_formats::api_format_alias_matches(candidate, api_format)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
candidate_common_transport_skip_reason, candidate_transport_pair_skip_reason,
|
||||
request_conversion_direct_auth, request_conversion_enabled_for_transport,
|
||||
request_conversion_transport_supported, request_pair_allowed_for_transport,
|
||||
CandidateTransportPolicyFacts,
|
||||
};
|
||||
use aether_ai_formats::matrix::RequestConversionKind;
|
||||
use serde_json::json;
|
||||
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn transport_snapshot(
|
||||
provider_type: &str,
|
||||
api_format: &str,
|
||||
auth_type: &str,
|
||||
enable_format_conversion: bool,
|
||||
format_acceptance_config: Option<serde_json::Value>,
|
||||
) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: format!("provider-{provider_type}"),
|
||||
name: provider_type.to_string(),
|
||||
provider_type: provider_type.to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: format!("provider-{provider_type}"),
|
||||
api_format: api_format.to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: if provider_type == "kiro" {
|
||||
"https://q.{region}.amazonaws.com".to_string()
|
||||
} else if provider_type == "vertex_ai" {
|
||||
"https://aiplatform.googleapis.com".to_string()
|
||||
} else {
|
||||
"https://api.example.com".to_string()
|
||||
},
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: format!("provider-{provider_type}"),
|
||||
name: "key".to_string(),
|
||||
auth_type: auth_type.to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec![api_format.to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn candidate_facts(api_format: &str) -> CandidateTransportPolicyFacts<'_> {
|
||||
CandidateTransportPolicyFacts {
|
||||
endpoint_api_format: api_format,
|
||||
global_model_name: "gpt-4.1",
|
||||
selected_provider_model_name: "provider-gpt-4.1",
|
||||
mapping_matched_model: Some("gpt-4.1-mini"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn conversion_helpers_follow_transport_api_format() {
|
||||
let transport = transport_snapshot("openai", "openai:chat", "bearer", true, None);
|
||||
|
||||
assert!(request_conversion_transport_supported(
|
||||
&transport,
|
||||
RequestConversionKind::ToOpenAIChat
|
||||
));
|
||||
assert_eq!(
|
||||
request_conversion_direct_auth(&transport, RequestConversionKind::ToOpenAIChat),
|
||||
Some(("authorization".to_string(), "Bearer secret".to_string()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn endpoint_level_format_acceptance_enables_cross_format_pair_without_provider_flag() {
|
||||
let transport = transport_snapshot(
|
||||
"custom",
|
||||
"openai:responses",
|
||||
"bearer",
|
||||
false,
|
||||
Some(json!({
|
||||
"enabled": true,
|
||||
"accept_formats": ["claude:messages"],
|
||||
})),
|
||||
);
|
||||
|
||||
assert!(request_conversion_enabled_for_transport(
|
||||
&transport,
|
||||
"claude:messages",
|
||||
"openai:responses"
|
||||
));
|
||||
assert!(request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"claude:messages",
|
||||
"openai:responses"
|
||||
));
|
||||
assert!(!request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"gemini:generate_content",
|
||||
"openai:responses"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn endpoint_reject_formats_override_endpoint_cross_format_enablement() {
|
||||
let transport = transport_snapshot(
|
||||
"custom",
|
||||
"openai:responses",
|
||||
"bearer",
|
||||
false,
|
||||
Some(json!({
|
||||
"enabled": true,
|
||||
"reject_formats": ["claude:messages"],
|
||||
})),
|
||||
);
|
||||
|
||||
assert!(!request_conversion_enabled_for_transport(
|
||||
&transport,
|
||||
"claude:messages",
|
||||
"openai:responses"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_gemini_transport_supports_cross_format_conversion_with_query_auth() {
|
||||
let transport = transport_snapshot(
|
||||
"vertex_ai",
|
||||
"gemini:generate_content",
|
||||
"api_key",
|
||||
true,
|
||||
None,
|
||||
);
|
||||
|
||||
assert!(request_conversion_transport_supported(
|
||||
&transport,
|
||||
RequestConversionKind::ToGeminiStandard
|
||||
));
|
||||
assert_eq!(
|
||||
request_conversion_direct_auth(&transport, RequestConversionKind::ToGeminiStandard),
|
||||
Some(("key".to_string(), "secret".to_string()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_claude_messages_transport_supports_cross_format_conversion_via_envelope() {
|
||||
let transport = transport_snapshot("kiro", "claude:messages", "bearer", true, None);
|
||||
|
||||
assert!(request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"openai:chat",
|
||||
"claude:messages"
|
||||
));
|
||||
assert!(request_conversion_transport_supported(
|
||||
&transport,
|
||||
RequestConversionKind::ToClaudeStandard
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn candidate_common_transport_policy_checks_active_state_format_and_allowed_models() {
|
||||
let mut transport = transport_snapshot("custom", "openai:chat", "bearer", true, None);
|
||||
transport.key.allowed_models = Some(vec!["gpt-4.1-mini".to_string()]);
|
||||
|
||||
assert_eq!(
|
||||
candidate_common_transport_skip_reason(
|
||||
&transport,
|
||||
candidate_facts("openai:chat"),
|
||||
Some("gpt-4.1")
|
||||
),
|
||||
None
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
candidate_common_transport_skip_reason(
|
||||
&transport,
|
||||
candidate_facts("claude:messages"),
|
||||
Some("gpt-4.1")
|
||||
),
|
||||
Some("endpoint_api_format_changed")
|
||||
);
|
||||
|
||||
transport.key.allowed_models = Some(vec!["other-model".to_string()]);
|
||||
assert_eq!(
|
||||
candidate_common_transport_skip_reason(
|
||||
&transport,
|
||||
candidate_facts("openai:chat"),
|
||||
Some("gpt-4.1")
|
||||
),
|
||||
Some("key_model_disabled")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn fixed_provider_oauth_keys_inherit_endpoint_api_formats_for_candidate_policy() {
|
||||
let mut transport = transport_snapshot("codex", "openai:responses", "oauth", true, None);
|
||||
transport.key.api_formats = Some(vec!["openai:image".to_string()]);
|
||||
|
||||
assert_eq!(
|
||||
candidate_common_transport_skip_reason(
|
||||
&transport,
|
||||
candidate_facts("openai:responses"),
|
||||
None,
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn candidate_transport_pair_policy_reports_disabled_conversion_and_unsupported_pairs() {
|
||||
let transport = transport_snapshot("custom", "openai:responses", "bearer", false, None);
|
||||
|
||||
assert_eq!(
|
||||
candidate_transport_pair_skip_reason(&transport, "openai:chat"),
|
||||
Some("format_conversion_disabled")
|
||||
);
|
||||
assert_eq!(
|
||||
candidate_transport_pair_skip_reason(&transport, "gemini:video"),
|
||||
Some("transport_unsupported")
|
||||
);
|
||||
|
||||
let enabled_transport =
|
||||
transport_snapshot("custom", "openai:responses", "bearer", true, None);
|
||||
assert_eq!(
|
||||
candidate_transport_pair_skip_reason(&enabled_transport, "openai:chat"),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
418
crates/aether-provider-transport/src/diagnostics.rs
Normal file
418
crates/aether-provider-transport/src/diagnostics.rs
Normal file
@@ -0,0 +1,418 @@
|
||||
use aether_ai_formats::matrix::{request_conversion_kind, request_conversion_requires_enable_flag};
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use serde_json::{json, Map, Value};
|
||||
|
||||
use crate::conversion::{
|
||||
request_conversion_enabled_for_transport, request_conversion_transport_unsupported_reason,
|
||||
request_pair_allowed_for_transport,
|
||||
};
|
||||
use crate::network::{resolve_transport_tls_profile, transport_proxy_is_locally_supported};
|
||||
use crate::policy::{
|
||||
local_gemini_transport_unsupported_reason_with_network,
|
||||
local_openai_chat_transport_unsupported_reason,
|
||||
local_standard_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::rules::{body_rules_are_locally_supported, header_rules_are_locally_supported};
|
||||
use crate::same_format_provider::same_format_provider_transport_unsupported_reason_for_trace;
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
pub fn build_request_trace_proxy_value(
|
||||
transport: Option<&GatewayProviderTransportSnapshot>,
|
||||
resolved_proxy: Option<&ProxySnapshot>,
|
||||
) -> Option<Value> {
|
||||
let resolved_proxy = resolved_proxy?;
|
||||
let mut object = Map::new();
|
||||
|
||||
if let Some(node_id) = trimmed_non_empty(resolved_proxy.node_id.as_deref()) {
|
||||
object.insert("node_id".to_string(), Value::String(node_id));
|
||||
}
|
||||
if let Some(node_name) = trimmed_non_empty(resolved_proxy.label.as_deref()) {
|
||||
object.insert("node_name".to_string(), Value::String(node_name));
|
||||
}
|
||||
if let Some(url) = sanitize_trace_proxy_url(resolved_proxy.url.as_deref()) {
|
||||
object.insert("url".to_string(), Value::String(url));
|
||||
}
|
||||
if let Some(source) = resolve_request_trace_proxy_source(transport, true) {
|
||||
object.insert("source".to_string(), Value::String(source.to_string()));
|
||||
}
|
||||
|
||||
(!object.is_empty()).then_some(Value::Object(object))
|
||||
}
|
||||
|
||||
pub fn append_transport_diagnostics_to_value(
|
||||
value: Value,
|
||||
transport: Option<&GatewayProviderTransportSnapshot>,
|
||||
client_api_format: &str,
|
||||
provider_api_format: &str,
|
||||
) -> Value {
|
||||
let Value::Object(mut object) = value else {
|
||||
return value;
|
||||
};
|
||||
object.insert(
|
||||
"transport_diagnostics".to_string(),
|
||||
transport
|
||||
.map(|transport| {
|
||||
build_transport_diagnostics(transport, client_api_format, provider_api_format)
|
||||
})
|
||||
.unwrap_or_else(|| json!({ "transport_snapshot_available": false })),
|
||||
);
|
||||
Value::Object(object)
|
||||
}
|
||||
|
||||
pub fn build_transport_diagnostics(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
client_api_format: &str,
|
||||
provider_api_format: &str,
|
||||
) -> Value {
|
||||
let resolved_tls_profile = resolve_transport_tls_profile(transport);
|
||||
let configured_tls_profile = transport
|
||||
.key
|
||||
.fingerprint
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get("tls_profile"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null);
|
||||
let has_oauth_config = transport.key.decrypted_auth_config.is_some();
|
||||
let oauth_resolution_supported =
|
||||
!has_oauth_config || crate::supports_local_oauth_request_auth_resolution(transport);
|
||||
let request_transport_unsupported_reason = resolve_request_transport_unsupported_reason(
|
||||
transport,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
);
|
||||
|
||||
json!({
|
||||
"transport_snapshot_available": true,
|
||||
"provider_type": transport.provider.provider_type,
|
||||
"provider_is_active": transport.provider.is_active,
|
||||
"endpoint_is_active": transport.endpoint.is_active,
|
||||
"key_is_active": transport.key.is_active,
|
||||
"provider_enable_format_conversion": transport.provider.enable_format_conversion,
|
||||
"provider_keep_priority_on_conversion": transport.provider.keep_priority_on_conversion,
|
||||
"endpoint_format_acceptance_config": transport.endpoint.format_acceptance_config,
|
||||
"endpoint_custom_path": transport.endpoint.custom_path,
|
||||
"header_rules": transport.endpoint.header_rules,
|
||||
"header_rules_supported": header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()),
|
||||
"body_rules": transport.endpoint.body_rules,
|
||||
"body_rules_supported": body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()),
|
||||
"proxy": {
|
||||
"locally_supported": transport_proxy_is_locally_supported(transport),
|
||||
"provider": summarize_proxy_config(transport.provider.proxy.as_ref()),
|
||||
"endpoint": summarize_proxy_config(transport.endpoint.proxy.as_ref()),
|
||||
"key": summarize_proxy_config(transport.key.proxy.as_ref()),
|
||||
},
|
||||
"auth": {
|
||||
"key_auth_type": transport.key.auth_type,
|
||||
"has_oauth_config": has_oauth_config,
|
||||
"oauth_request_auth_resolution_supported": oauth_resolution_supported,
|
||||
},
|
||||
"fingerprint": transport.key.fingerprint,
|
||||
"configured_tls_profile": configured_tls_profile,
|
||||
"resolved_tls_profile": resolved_tls_profile,
|
||||
"request_pair": {
|
||||
"client_api_format": client_api_format,
|
||||
"provider_api_format": provider_api_format,
|
||||
"requires_conversion_enable_flag": request_conversion_requires_enable_flag(
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
),
|
||||
"conversion_enabled": request_conversion_enabled_for_transport(
|
||||
transport,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
),
|
||||
"pair_allowed": request_pair_allowed_for_transport(
|
||||
transport,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
),
|
||||
"transport_unsupported_reason": request_transport_unsupported_reason,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
fn summarize_proxy_config(proxy: Option<&Value>) -> Value {
|
||||
let Some(object) = proxy.and_then(Value::as_object) else {
|
||||
return Value::Null;
|
||||
};
|
||||
let has_url = object
|
||||
.get("url")
|
||||
.or_else(|| object.get("proxy_url"))
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty());
|
||||
json!({
|
||||
"enabled": object.get("enabled").cloned().unwrap_or(Value::Null),
|
||||
"mode": object.get("mode").cloned().unwrap_or(Value::Null),
|
||||
"node_id": object.get("node_id").cloned().unwrap_or(Value::Null),
|
||||
"label": object.get("label").cloned().unwrap_or(Value::Null),
|
||||
"has_url": has_url,
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_request_trace_proxy_source(
|
||||
transport: Option<&GatewayProviderTransportSnapshot>,
|
||||
has_resolved_proxy: bool,
|
||||
) -> Option<&'static str> {
|
||||
let transport = transport?;
|
||||
if transport_has_explicit_proxy(transport.key.proxy.as_ref()) {
|
||||
return Some("key");
|
||||
}
|
||||
if transport_has_explicit_proxy(transport.endpoint.proxy.as_ref()) {
|
||||
return Some("endpoint");
|
||||
}
|
||||
if transport_has_explicit_proxy(transport.provider.proxy.as_ref()) {
|
||||
return Some("provider");
|
||||
}
|
||||
has_resolved_proxy.then_some("system")
|
||||
}
|
||||
|
||||
fn transport_has_explicit_proxy(proxy: Option<&Value>) -> bool {
|
||||
let Some(object) = proxy.and_then(Value::as_object) else {
|
||||
return false;
|
||||
};
|
||||
let enabled = object
|
||||
.get("enabled")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(true);
|
||||
if !enabled {
|
||||
return false;
|
||||
}
|
||||
|
||||
object
|
||||
.get("node_id")
|
||||
.or_else(|| object.get("url"))
|
||||
.or_else(|| object.get("proxy_url"))
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty())
|
||||
}
|
||||
|
||||
fn sanitize_trace_proxy_url(url: Option<&str>) -> Option<String> {
|
||||
let raw = url.map(str::trim).filter(|value| !value.is_empty())?;
|
||||
let parsed = url::Url::parse(raw).ok()?;
|
||||
let scheme = parsed.scheme().trim();
|
||||
let host = parsed.host_str()?.trim();
|
||||
if scheme.is_empty() || host.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut safe = format!("{scheme}://{host}");
|
||||
if let Some(port) = parsed.port() {
|
||||
safe.push(':');
|
||||
safe.push_str(port.to_string().as_str());
|
||||
}
|
||||
Some(safe)
|
||||
}
|
||||
|
||||
fn trimmed_non_empty(value: Option<&str>) -> Option<String> {
|
||||
value
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn resolve_request_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
client_api_format: &str,
|
||||
provider_api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
let client_api_format = client_api_format.trim().to_ascii_lowercase();
|
||||
let provider_api_format = provider_api_format.trim().to_ascii_lowercase();
|
||||
if client_api_format == provider_api_format {
|
||||
if let Some(skip_reason) = same_format_provider_transport_unsupported_reason_for_trace(
|
||||
transport,
|
||||
provider_api_format.as_str(),
|
||||
) {
|
||||
return Some(skip_reason);
|
||||
}
|
||||
return match provider_api_format.as_str() {
|
||||
"openai:chat" => local_openai_chat_transport_unsupported_reason(transport),
|
||||
"gemini:generate_content" => local_gemini_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
provider_api_format.as_str(),
|
||||
),
|
||||
_ => local_standard_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
provider_api_format.as_str(),
|
||||
),
|
||||
};
|
||||
}
|
||||
match request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str()) {
|
||||
Some(kind) => request_conversion_transport_unsupported_reason(transport, kind),
|
||||
None => Some("transport_api_format_unsupported"),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{build_request_trace_proxy_value, build_transport_diagnostics};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "RightCode".to_string(),
|
||||
provider_type: "codex".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: true,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: Some(json!({"enabled": true, "mode": "node", "node_id": "proxy-node-1"})),
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "openai:responses".to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: "https://example.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: Some("/v1/responses".to_string()),
|
||||
config: None,
|
||||
format_acceptance_config: Some(json!({
|
||||
"enabled": true,
|
||||
"accept_formats": ["claude:messages"]
|
||||
})),
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "codex".to_string(),
|
||||
auth_type: "oauth".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: Some(json!({
|
||||
"tls_profile": "chrome_136",
|
||||
"user_agent": "Mozilla/5.0"
|
||||
})),
|
||||
decrypted_api_key: "sk-test".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_claude_code_transport_without_auth() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-cc-1".to_string(),
|
||||
name: "NekoCode".to_string(),
|
||||
provider_type: "claude_code".to_string(),
|
||||
website: Some("https://nekocode.ai".to_string()),
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: true,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-cc-1".to_string(),
|
||||
provider_id: "provider-cc-1".to_string(),
|
||||
api_format: "claude:messages".to_string(),
|
||||
api_family: Some("claude".to_string()),
|
||||
endpoint_kind: Some("cli".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://api.anthropic.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-cc-1".to_string(),
|
||||
provider_id: "provider-cc-1".to_string(),
|
||||
name: "CC".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec!["claude:messages".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_diagnostics_include_format_and_network_policy() {
|
||||
let diagnostics =
|
||||
build_transport_diagnostics(&sample_transport(), "claude:messages", "openai:responses");
|
||||
|
||||
assert_eq!(diagnostics["provider_type"], "codex");
|
||||
assert_eq!(diagnostics["fingerprint"]["tls_profile"], "chrome_136");
|
||||
assert_eq!(diagnostics["resolved_tls_profile"], "chrome_136");
|
||||
assert_eq!(
|
||||
diagnostics["request_pair"]["conversion_enabled"],
|
||||
Value::Bool(true)
|
||||
);
|
||||
assert!(diagnostics["request_pair"]["transport_unsupported_reason"].is_null());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_diagnostics_use_provider_private_same_format_reason() {
|
||||
let diagnostics = build_transport_diagnostics(
|
||||
&sample_claude_code_transport_without_auth(),
|
||||
"claude:messages",
|
||||
"claude:messages",
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
diagnostics["request_pair"]["transport_unsupported_reason"],
|
||||
Value::String("transport_auth_unavailable".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_trace_proxy_value_sanitizes_url_and_marks_config_source() {
|
||||
let transport = sample_transport();
|
||||
let proxy = ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
mode: Some("node".to_string()),
|
||||
node_id: Some("proxy-node-1".to_string()),
|
||||
label: Some("Primary proxy".to_string()),
|
||||
url: Some("https://user:pass@example.test:9443/path".to_string()),
|
||||
extra: None,
|
||||
};
|
||||
|
||||
let value = build_request_trace_proxy_value(Some(&transport), Some(&proxy)).unwrap();
|
||||
|
||||
assert_eq!(value["node_id"], "proxy-node-1");
|
||||
assert_eq!(value["node_name"], "Primary proxy");
|
||||
assert_eq!(value["url"], "https://example.test:9443");
|
||||
assert_eq!(value["source"], "provider");
|
||||
}
|
||||
}
|
||||
269
crates/aether-provider-transport/src/gemini_files.rs
Normal file
269
crates/aether-provider-transport/src/gemini_files.rs
Normal file
@@ -0,0 +1,269 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use crate::auth::{build_passthrough_headers_with_auth, resolve_local_gemini_auth};
|
||||
use crate::policy::local_gemini_transport_unsupported_reason_with_network;
|
||||
use crate::rules::{apply_local_body_rules, apply_local_header_rules};
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::url::build_gemini_files_passthrough_url;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GeminiFilesRequestBodyError {
|
||||
BodyRulesUnsupportedForBinaryUpload,
|
||||
BodyRulesApplyFailed,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct GeminiFilesRequestBodyParts {
|
||||
pub provider_request_body: Option<Value>,
|
||||
pub provider_request_body_base64: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct GeminiFilesHeadersInput<'a> {
|
||||
pub headers: &'a http::HeaderMap,
|
||||
pub auth_header: &'a str,
|
||||
pub auth_value: &'a str,
|
||||
pub header_rules: Option<&'a Value>,
|
||||
pub provider_request_body: Option<&'a Value>,
|
||||
pub provider_request_body_base64: Option<&'a str>,
|
||||
pub original_request_body_json: &'a Value,
|
||||
pub original_body_is_empty: bool,
|
||||
}
|
||||
|
||||
pub fn gemini_files_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
local_gemini_transport_unsupported_reason_with_network(transport, api_format)
|
||||
}
|
||||
|
||||
pub fn resolve_gemini_files_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<(String, String)> {
|
||||
resolve_local_gemini_auth(transport)
|
||||
}
|
||||
|
||||
pub fn build_gemini_files_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request_path: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let custom_path = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
let passthrough_path = custom_path.unwrap_or(request_path);
|
||||
build_gemini_files_passthrough_url(
|
||||
&transport.endpoint.base_url,
|
||||
passthrough_path,
|
||||
request_query,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_gemini_files_request_body(
|
||||
body_json: &Value,
|
||||
body_base64: Option<&str>,
|
||||
body_is_empty: bool,
|
||||
is_upload: bool,
|
||||
body_rules: Option<&Value>,
|
||||
) -> Result<GeminiFilesRequestBodyParts, GeminiFilesRequestBodyError> {
|
||||
let mut provider_request_body = if is_upload && !body_is_empty && body_base64.is_none() {
|
||||
Some(body_json.clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let provider_request_body_base64 = if is_upload {
|
||||
body_base64
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if provider_request_body_base64.is_some() && body_rules.is_some() {
|
||||
return Err(GeminiFilesRequestBodyError::BodyRulesUnsupportedForBinaryUpload);
|
||||
}
|
||||
if let Some(body) = provider_request_body.as_mut() {
|
||||
if !apply_local_body_rules(body, body_rules, Some(body_json)) {
|
||||
return Err(GeminiFilesRequestBodyError::BodyRulesApplyFailed);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(GeminiFilesRequestBodyParts {
|
||||
provider_request_body,
|
||||
provider_request_body_base64,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn build_gemini_files_headers(
|
||||
input: GeminiFilesHeadersInput<'_>,
|
||||
) -> Option<BTreeMap<String, String>> {
|
||||
let mut provider_request_headers = build_passthrough_headers_with_auth(
|
||||
input.headers,
|
||||
input.auth_header,
|
||||
input.auth_value,
|
||||
&BTreeMap::new(),
|
||||
);
|
||||
let null_original_request_body = Value::Null;
|
||||
let base64_original_request_body = input
|
||||
.provider_request_body_base64
|
||||
.map(|body_bytes_b64| json!({ "body_bytes_b64": body_bytes_b64 }));
|
||||
let original_request_body = base64_original_request_body
|
||||
.as_ref()
|
||||
.or_else(|| (!input.original_body_is_empty).then_some(input.original_request_body_json))
|
||||
.unwrap_or(&null_original_request_body);
|
||||
if !apply_local_header_rules(
|
||||
&mut provider_request_headers,
|
||||
input.header_rules,
|
||||
&[input.auth_header, "content-type"],
|
||||
input.provider_request_body.unwrap_or(original_request_body),
|
||||
Some(original_request_body),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
Some(provider_request_headers)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Provider".to_string(),
|
||||
provider_type: "gemini".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "gemini:files".to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: "https://generativelanguage.googleapis.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_gemini_files_url_from_request_path_and_strips_client_key() {
|
||||
let url = build_gemini_files_upstream_url(
|
||||
&sample_transport(),
|
||||
"/v1beta/files/demo",
|
||||
Some("key=client-key&alt=json"),
|
||||
)
|
||||
.expect("url should build");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://generativelanguage.googleapis.com/v1beta/files/demo?alt=json"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_json_upload_body_and_applies_rules() {
|
||||
let body = build_gemini_files_request_body(
|
||||
&json!({"display_name": "demo"}),
|
||||
None,
|
||||
false,
|
||||
true,
|
||||
Some(&json!([
|
||||
{"action":"set","path":"metadata.source","value":"local"}
|
||||
])),
|
||||
)
|
||||
.expect("body should build");
|
||||
|
||||
assert_eq!(
|
||||
body.provider_request_body
|
||||
.as_ref()
|
||||
.and_then(|value| value.pointer("/metadata/source")),
|
||||
Some(&json!("local"))
|
||||
);
|
||||
assert!(body.provider_request_body_base64.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_body_rules_for_binary_upload() {
|
||||
assert_eq!(
|
||||
build_gemini_files_request_body(
|
||||
&json!({}),
|
||||
Some("YWJj"),
|
||||
false,
|
||||
true,
|
||||
Some(&json!([{"action":"set","path":"x","value":1}]))
|
||||
),
|
||||
Err(GeminiFilesRequestBodyError::BodyRulesUnsupportedForBinaryUpload)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_headers_using_binary_body_as_original_context() {
|
||||
let headers = build_gemini_files_headers(GeminiFilesHeadersInput {
|
||||
headers: &http::HeaderMap::new(),
|
||||
auth_header: "x-goog-api-key",
|
||||
auth_value: "secret",
|
||||
header_rules: Some(&json!([
|
||||
{"action":"set","key":"x-upload-mode","value":"binary"}
|
||||
])),
|
||||
provider_request_body: None,
|
||||
provider_request_body_base64: Some("YWJj"),
|
||||
original_request_body_json: &json!({}),
|
||||
original_body_is_empty: false,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("x-goog-api-key").map(String::as_str),
|
||||
Some("secret")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-upload-mode").map(String::as_str),
|
||||
Some("binary")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -8,10 +8,10 @@ mod request;
|
||||
mod url;
|
||||
|
||||
pub use auth::{
|
||||
build_kiro_request_auth_from_config, resolve_local_kiro_bearer_auth,
|
||||
resolve_local_kiro_request_auth, supports_local_kiro_auth_prerequisites,
|
||||
supports_local_kiro_request_auth_resolution, KiroBearerAuth, KiroRequestAuth, KIRO_AUTH_HEADER,
|
||||
PROVIDER_TYPE,
|
||||
build_kiro_request_auth_from_config, is_kiro_claude_messages_transport,
|
||||
is_kiro_provider_transport, resolve_local_kiro_bearer_auth, resolve_local_kiro_request_auth,
|
||||
supports_local_kiro_auth_prerequisites, supports_local_kiro_request_auth_resolution,
|
||||
KiroBearerAuth, KiroRequestAuth, KIRO_AUTH_HEADER, PROVIDER_TYPE,
|
||||
};
|
||||
pub use converter::convert_claude_messages_to_conversation_state;
|
||||
pub use credentials::{generate_machine_id, normalize_machine_id, KiroAuthConfig};
|
||||
|
||||
@@ -18,6 +18,22 @@ pub struct KiroRequestAuth {
|
||||
pub machine_id: String,
|
||||
}
|
||||
|
||||
pub fn is_kiro_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(PROVIDER_TYPE)
|
||||
}
|
||||
|
||||
pub fn is_kiro_claude_messages_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> bool {
|
||||
is_kiro_provider_transport(transport)
|
||||
&& aether_ai_formats::api_format_alias_matches(provider_api_format, "claude:messages")
|
||||
}
|
||||
|
||||
pub fn build_kiro_request_auth_from_config(
|
||||
auth_config: KiroAuthConfig,
|
||||
fallback_secret: Option<&str>,
|
||||
@@ -47,12 +63,7 @@ pub fn build_kiro_request_auth_from_config(
|
||||
pub fn resolve_local_kiro_bearer_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<KiroBearerAuth> {
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(PROVIDER_TYPE)
|
||||
{
|
||||
if !is_kiro_provider_transport(transport) {
|
||||
return None;
|
||||
}
|
||||
if transport.key.decrypted_auth_config.is_some() {
|
||||
@@ -82,12 +93,7 @@ pub fn supports_local_kiro_auth_prerequisites(
|
||||
pub fn resolve_local_kiro_request_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<KiroRequestAuth> {
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(PROVIDER_TYPE)
|
||||
{
|
||||
if !is_kiro_provider_transport(transport) {
|
||||
return None;
|
||||
}
|
||||
if !kiro_auth_type_supported(transport.key.auth_type.as_str()) {
|
||||
@@ -127,11 +133,7 @@ pub fn supports_local_kiro_request_auth_resolution(
|
||||
resolve_local_kiro_request_auth(transport).is_some()
|
||||
|| KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
|
||||
.is_some_and(|auth_config| {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(PROVIDER_TYPE)
|
||||
is_kiro_provider_transport(transport)
|
||||
&& kiro_auth_type_supported(transport.key.auth_type.as_str())
|
||||
&& auth_config.can_refresh_access_token()
|
||||
})
|
||||
|
||||
@@ -3,16 +3,22 @@ pub mod auth;
|
||||
mod auth_config;
|
||||
mod cache;
|
||||
pub mod claude_code;
|
||||
pub mod conversion;
|
||||
mod diagnostics;
|
||||
mod gemini_files;
|
||||
mod generic_oauth;
|
||||
mod headers;
|
||||
pub mod kiro;
|
||||
mod network;
|
||||
pub mod oauth_refresh;
|
||||
mod openai_image;
|
||||
pub mod policy;
|
||||
pub mod provider_types;
|
||||
mod request_url;
|
||||
pub mod rules;
|
||||
pub mod same_format_provider;
|
||||
pub mod snapshot;
|
||||
mod standard;
|
||||
pub mod url;
|
||||
pub mod vertex;
|
||||
mod video;
|
||||
@@ -20,6 +26,21 @@ mod video;
|
||||
pub use aether_oauth as oauth;
|
||||
pub use auth::{build_passthrough_headers, ensure_upstream_auth_header};
|
||||
pub use cache::{provider_transport_snapshot_looks_refreshed, ProviderTransportSnapshotCacheKey};
|
||||
pub use conversion::{
|
||||
candidate_common_transport_skip_reason, candidate_transport_pair_skip_reason,
|
||||
request_conversion_direct_auth, request_conversion_enabled_for_transport,
|
||||
request_conversion_transport_supported, request_conversion_transport_unsupported_reason,
|
||||
request_pair_allowed_for_transport, CandidateTransportPolicyFacts,
|
||||
};
|
||||
pub use diagnostics::{
|
||||
append_transport_diagnostics_to_value, build_request_trace_proxy_value,
|
||||
build_transport_diagnostics,
|
||||
};
|
||||
pub use gemini_files::{
|
||||
build_gemini_files_headers, build_gemini_files_request_body, build_gemini_files_upstream_url,
|
||||
gemini_files_transport_unsupported_reason, resolve_gemini_files_auth, GeminiFilesHeadersInput,
|
||||
GeminiFilesRequestBodyError, GeminiFilesRequestBodyParts,
|
||||
};
|
||||
pub use generic_oauth::{
|
||||
supports_local_generic_oauth_request_auth_resolution, GenericOAuthRefreshAdapter,
|
||||
};
|
||||
@@ -35,6 +56,11 @@ pub use oauth_refresh::{
|
||||
LocalOAuthHttpRequest, LocalOAuthHttpResponse, LocalOAuthRefreshCoordinator,
|
||||
LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
|
||||
};
|
||||
pub use openai_image::{
|
||||
build_openai_image_headers, build_openai_image_upstream_url,
|
||||
openai_image_transport_unsupported_reason, resolve_openai_image_auth,
|
||||
ProviderOpenAiImageHeadersInput,
|
||||
};
|
||||
pub use policy::{
|
||||
local_gemini_transport_unsupported_reason,
|
||||
local_gemini_transport_unsupported_reason_with_network,
|
||||
@@ -42,17 +68,42 @@ pub use policy::{
|
||||
local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport,
|
||||
supports_local_gemini_transport_with_network, supports_local_standard_transport,
|
||||
};
|
||||
pub use request_url::{build_transport_request_url, TransportRequestUrlParams};
|
||||
pub use request_url::{
|
||||
build_cross_format_openai_chat_upstream_url, build_cross_format_openai_responses_upstream_url,
|
||||
build_kiro_cross_format_upstream_url, build_local_openai_chat_upstream_url,
|
||||
build_local_openai_responses_upstream_url, build_transport_request_url,
|
||||
TransportRequestUrlParams,
|
||||
};
|
||||
pub use rules::{
|
||||
apply_local_body_rules, apply_local_header_rules, body_rules_are_locally_supported,
|
||||
body_rules_handle_path, header_rules_are_locally_supported,
|
||||
};
|
||||
pub use same_format_provider::{
|
||||
build_same_format_provider_headers, build_same_format_provider_request_body,
|
||||
build_same_format_provider_upstream_url, classify_same_format_provider_request_behavior,
|
||||
resolve_same_format_provider_direct_auth, same_format_provider_transport_supported,
|
||||
same_format_provider_transport_unsupported_reason,
|
||||
same_format_provider_transport_unsupported_reason_for_trace,
|
||||
should_try_same_format_provider_oauth_auth, SameFormatProviderFamily,
|
||||
SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior,
|
||||
SameFormatProviderRequestBehaviorParams, SameFormatProviderRequestBodyInput,
|
||||
SameFormatProviderUpstreamUrlParams,
|
||||
};
|
||||
pub use snapshot::{
|
||||
read_provider_transport_snapshot, GatewayProviderTransportSnapshot,
|
||||
ProviderTransportSnapshotSource,
|
||||
};
|
||||
pub use standard::{
|
||||
apply_standard_provider_request_body_rules, build_standard_plan_fallback_headers,
|
||||
build_standard_plan_fallback_openai_chat_url,
|
||||
build_standard_plan_fallback_openai_responses_url, build_standard_provider_request_headers,
|
||||
StandardPlanFallbackAcceptPolicy, StandardPlanFallbackHeadersInput,
|
||||
StandardProviderRequestHeaders, StandardProviderRequestHeadersInput,
|
||||
};
|
||||
pub use vertex::{is_vertex_api_key_transport_context, uses_vertex_api_key_query_auth};
|
||||
pub use video::{
|
||||
build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url,
|
||||
reconstruct_local_video_task_snapshot, resolve_local_video_task_transport,
|
||||
VideoTaskTransportSnapshotLookup,
|
||||
resolve_video_create_auth, video_create_transport_unsupported_reason,
|
||||
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, VideoTaskTransportSnapshotLookup,
|
||||
};
|
||||
|
||||
167
crates/aether-provider-transport/src/openai_image.rs
Normal file
167
crates/aether-provider-transport/src/openai_image.rs
Normal file
@@ -0,0 +1,167 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::auth::{build_passthrough_headers_with_auth, resolve_local_openai_bearer_auth};
|
||||
use crate::policy::local_standard_transport_unsupported_reason_with_network;
|
||||
use crate::rules::apply_local_header_rules;
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::url::build_openai_responses_url;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct ProviderOpenAiImageHeadersInput<'a> {
|
||||
pub headers: &'a http::HeaderMap,
|
||||
pub auth_header: &'a str,
|
||||
pub auth_value: &'a str,
|
||||
pub header_rules: Option<&'a Value>,
|
||||
pub provider_request_body: &'a Value,
|
||||
pub original_request_body: &'a Value,
|
||||
}
|
||||
|
||||
pub fn openai_image_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
local_standard_transport_unsupported_reason_with_network(transport, api_format)
|
||||
}
|
||||
|
||||
pub fn resolve_openai_image_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<(String, String)> {
|
||||
resolve_local_openai_bearer_auth(transport)
|
||||
}
|
||||
|
||||
pub fn build_openai_image_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request_query: Option<&str>,
|
||||
) -> String {
|
||||
build_openai_responses_url(&transport.endpoint.base_url, request_query, false)
|
||||
}
|
||||
|
||||
pub fn build_openai_image_headers(
|
||||
input: ProviderOpenAiImageHeadersInput<'_>,
|
||||
) -> Option<BTreeMap<String, String>> {
|
||||
let mut provider_request_headers = build_passthrough_headers_with_auth(
|
||||
input.headers,
|
||||
input.auth_header,
|
||||
input.auth_value,
|
||||
&BTreeMap::new(),
|
||||
);
|
||||
provider_request_headers.insert("content-type".to_string(), "application/json".to_string());
|
||||
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
|
||||
if !apply_local_header_rules(
|
||||
&mut provider_request_headers,
|
||||
input.header_rules,
|
||||
&[input.auth_header, "content-type", "accept"],
|
||||
input.provider_request_body,
|
||||
Some(input.original_request_body),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
Some(provider_request_headers)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use http::HeaderMap;
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
build_openai_image_headers, build_openai_image_upstream_url,
|
||||
ProviderOpenAiImageHeadersInput,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Provider".to_string(),
|
||||
provider_type: "codex".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "openai:image".to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: "https://api.openai.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "bearer".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_openai_image_url_on_responses_surface() {
|
||||
let url = build_openai_image_upstream_url(&sample_transport(), Some("trace=1"));
|
||||
|
||||
assert_eq!(url, "https://api.openai.com/v1/responses?trace=1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_json_eventstream_headers_and_applies_rules() {
|
||||
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
|
||||
headers: &HeaderMap::new(),
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer secret",
|
||||
header_rules: Some(&json!([
|
||||
{"action":"set","key":"x-image-route","value":"codex"}
|
||||
])),
|
||||
provider_request_body: &json!({"model":"gpt-5.4-mini"}),
|
||||
original_request_body: &json!({"prompt":"draw"}),
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("authorization"),
|
||||
Some(&"Bearer secret".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("content-type"),
|
||||
Some(&"application/json".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("accept"),
|
||||
Some(&"text/event-stream".to_string())
|
||||
);
|
||||
assert_eq!(headers.get("x-image-route"), Some(&"codex".to_string()));
|
||||
}
|
||||
}
|
||||
@@ -158,6 +158,22 @@ pub fn provider_type_is_fixed(provider_type: &str) -> bool {
|
||||
)
|
||||
}
|
||||
|
||||
pub fn fixed_provider_key_inherits_api_formats(
|
||||
provider_type: &str,
|
||||
auth_type: &str,
|
||||
decrypted_auth_config: Option<&str>,
|
||||
) -> bool {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = auth_type.trim().to_ascii_lowercase();
|
||||
provider_type_is_fixed(&provider_type)
|
||||
&& (auth_type == "oauth"
|
||||
|| provider_type == "kiro"
|
||||
&& auth_type == "bearer"
|
||||
&& decrypted_auth_config
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty()))
|
||||
}
|
||||
|
||||
pub fn provider_type_enables_format_conversion_by_default(provider_type: &str) -> bool {
|
||||
matches!(
|
||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||
@@ -284,8 +300,8 @@ pub const ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES: &[&str] =
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
fixed_provider_endpoint_template_by_api_format, fixed_provider_template,
|
||||
FixedProviderEndpointConfigValue,
|
||||
fixed_provider_endpoint_template_by_api_format, fixed_provider_key_inherits_api_formats,
|
||||
fixed_provider_template, FixedProviderEndpointConfigValue,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -321,4 +337,22 @@ mod tests {
|
||||
)]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() {
|
||||
assert!(fixed_provider_key_inherits_api_formats(
|
||||
"codex", "oauth", None
|
||||
));
|
||||
assert!(fixed_provider_key_inherits_api_formats(
|
||||
"kiro",
|
||||
"bearer",
|
||||
Some("{}")
|
||||
));
|
||||
assert!(!fixed_provider_key_inherits_api_formats(
|
||||
"kiro", "bearer", None
|
||||
));
|
||||
assert!(!fixed_provider_key_inherits_api_formats(
|
||||
"custom", "oauth", None
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,10 @@ use std::sync::OnceLock;
|
||||
use regex::Regex;
|
||||
use url::form_urlencoded;
|
||||
|
||||
use crate::antigravity::{build_antigravity_v1internal_url, AntigravityRequestUrlAction};
|
||||
use crate::antigravity::{
|
||||
build_antigravity_v1internal_url, is_antigravity_provider_transport,
|
||||
AntigravityRequestUrlAction,
|
||||
};
|
||||
use crate::claude_code::build_claude_code_messages_url;
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::url::{
|
||||
@@ -95,6 +98,105 @@ pub fn build_transport_request_url(
|
||||
))
|
||||
}
|
||||
|
||||
pub fn build_local_openai_chat_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
build_transport_request_url(
|
||||
transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "openai:chat",
|
||||
mapped_model: None,
|
||||
upstream_is_stream: false,
|
||||
request_query,
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_cross_format_openai_chat_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
mapped_model: &str,
|
||||
provider_api_format: &str,
|
||||
upstream_is_stream: bool,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
aether_ai_formats::request_conversion_kind("openai:chat", provider_api_format)?;
|
||||
build_transport_request_url(
|
||||
transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format,
|
||||
mapped_model: Some(mapped_model),
|
||||
upstream_is_stream,
|
||||
request_query,
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_local_openai_responses_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
compact: bool,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let provider_api_format = if compact {
|
||||
"openai:responses:compact"
|
||||
} else {
|
||||
"openai:responses"
|
||||
};
|
||||
build_transport_request_url(
|
||||
transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format,
|
||||
mapped_model: None,
|
||||
upstream_is_stream: false,
|
||||
request_query,
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_cross_format_openai_responses_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
mapped_model: &str,
|
||||
client_api_format: &str,
|
||||
provider_api_format: &str,
|
||||
upstream_is_stream: bool,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
aether_ai_formats::request_conversion_kind(client_api_format, provider_api_format)?;
|
||||
build_transport_request_url(
|
||||
transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format,
|
||||
mapped_model: Some(mapped_model),
|
||||
upstream_is_stream,
|
||||
request_query,
|
||||
kiro_api_region: None,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_kiro_cross_format_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
mapped_model: &str,
|
||||
provider_api_format: &str,
|
||||
upstream_is_stream: bool,
|
||||
request_query: Option<&str>,
|
||||
api_region: &str,
|
||||
) -> Option<String> {
|
||||
build_transport_request_url(
|
||||
transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format,
|
||||
mapped_model: Some(mapped_model),
|
||||
upstream_is_stream,
|
||||
request_query,
|
||||
kiro_api_region: Some(api_region),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn build_transport_hook_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
params: TransportRequestUrlParams<'_>,
|
||||
@@ -135,12 +237,7 @@ fn build_transport_hook_url(
|
||||
}
|
||||
}
|
||||
|
||||
if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
{
|
||||
if is_antigravity_provider_transport(transport) {
|
||||
let query = params.request_query.map(|raw| {
|
||||
form_urlencoded::parse(raw.as_bytes())
|
||||
.into_owned()
|
||||
@@ -255,7 +352,10 @@ fn custom_path_template_regex() -> &'static Regex {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{build_transport_request_url, TransportRequestUrlParams};
|
||||
use super::{
|
||||
build_kiro_cross_format_upstream_url, build_transport_request_url,
|
||||
TransportRequestUrlParams,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
@@ -422,4 +522,29 @@ mod tests {
|
||||
|
||||
assert_eq!(url, "https://api.example.com/v1/messages/{model}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_cross_format_helper_uses_region_specific_generate_assistant_url() {
|
||||
let transport = sample_transport(
|
||||
"kiro",
|
||||
"claude:messages",
|
||||
"https://codewhisperer.{region}.amazonaws.com/",
|
||||
None,
|
||||
);
|
||||
|
||||
let url = build_kiro_cross_format_upstream_url(
|
||||
&transport,
|
||||
"claude-sonnet-4",
|
||||
"claude:messages",
|
||||
true,
|
||||
Some("conversationId=abc"),
|
||||
"us-west-2",
|
||||
)
|
||||
.expect("kiro url");
|
||||
|
||||
assert!(url.starts_with(
|
||||
"https://codewhisperer.us-west-2.amazonaws.com/generateAssistantResponse"
|
||||
));
|
||||
assert!(url.contains("conversationId=abc"));
|
||||
}
|
||||
}
|
||||
|
||||
525
crates/aether-provider-transport/src/same_format_provider.rs
Normal file
525
crates/aether-provider-transport/src/same_format_provider.rs
Normal file
@@ -0,0 +1,525 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
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_standard_auth,
|
||||
};
|
||||
use crate::claude_code::build_claude_code_passthrough_headers;
|
||||
use crate::claude_code::local_claude_code_transport_unsupported_reason_with_network;
|
||||
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,
|
||||
KiroProviderHeadersInput,
|
||||
};
|
||||
use crate::policy::{
|
||||
local_gemini_transport_unsupported_reason_with_network,
|
||||
local_standard_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::rules::{apply_local_body_rules, apply_local_header_rules};
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::vertex::{
|
||||
is_vertex_api_key_transport_context,
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
|
||||
};
|
||||
use crate::{build_transport_request_url, ensure_upstream_auth_header, TransportRequestUrlParams};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SameFormatProviderFamily {
|
||||
Standard,
|
||||
Gemini,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct SameFormatProviderRequestBehaviorParams {
|
||||
pub require_streaming: bool,
|
||||
pub report_kind: &'static str,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct SameFormatProviderRequestBehavior {
|
||||
pub is_antigravity: bool,
|
||||
pub is_claude_code: bool,
|
||||
pub is_vertex: bool,
|
||||
pub is_kiro: bool,
|
||||
pub upstream_is_stream: bool,
|
||||
pub report_kind: &'static str,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct SameFormatProviderRequestBodyInput<'a> {
|
||||
pub body_json: &'a Value,
|
||||
pub mapped_model: &'a str,
|
||||
pub family: SameFormatProviderFamily,
|
||||
pub body_rules: Option<&'a Value>,
|
||||
pub upstream_is_stream: bool,
|
||||
pub kiro_auth_config: Option<&'a KiroAuthConfig>,
|
||||
pub is_claude_code: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct SameFormatProviderUpstreamUrlParams<'a> {
|
||||
pub provider_api_format: &'a str,
|
||||
pub mapped_model: &'a str,
|
||||
pub upstream_is_stream: bool,
|
||||
pub request_query: Option<&'a str>,
|
||||
pub kiro_api_region: Option<&'a str>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct SameFormatProviderHeadersInput<'a> {
|
||||
pub headers: &'a http::HeaderMap,
|
||||
pub provider_request_body: &'a Value,
|
||||
pub original_request_body: &'a Value,
|
||||
pub header_rules: Option<&'a Value>,
|
||||
pub behavior: SameFormatProviderRequestBehavior,
|
||||
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>,
|
||||
}
|
||||
|
||||
pub fn classify_same_format_provider_request_behavior(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
params: SameFormatProviderRequestBehaviorParams,
|
||||
) -> SameFormatProviderRequestBehavior {
|
||||
let is_antigravity = is_antigravity_provider_transport(transport);
|
||||
let is_claude_code = transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("claude_code");
|
||||
let is_vertex = is_vertex_api_key_transport_context(transport);
|
||||
let is_kiro = is_kiro_provider_transport(transport);
|
||||
let upstream_is_stream = is_kiro || is_antigravity || params.require_streaming;
|
||||
let report_kind = if is_kiro && !params.require_streaming {
|
||||
"claude_cli_sync_finalize"
|
||||
} else if is_antigravity && !params.require_streaming {
|
||||
match params.report_kind {
|
||||
"gemini_chat_sync_success" => "gemini_chat_sync_finalize",
|
||||
"gemini_cli_sync_success" => "gemini_cli_sync_finalize",
|
||||
_ => params.report_kind,
|
||||
}
|
||||
} else {
|
||||
params.report_kind
|
||||
};
|
||||
|
||||
SameFormatProviderRequestBehavior {
|
||||
is_antigravity,
|
||||
is_claude_code,
|
||||
is_vertex,
|
||||
is_kiro,
|
||||
upstream_is_stream,
|
||||
report_kind,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_same_format_provider_request_body(
|
||||
input: SameFormatProviderRequestBodyInput<'_>,
|
||||
) -> Option<Value> {
|
||||
if let Some(kiro_auth_config) = input.kiro_auth_config {
|
||||
return build_kiro_provider_request_body(
|
||||
input.body_json,
|
||||
input.mapped_model,
|
||||
kiro_auth_config,
|
||||
input.body_rules,
|
||||
);
|
||||
}
|
||||
|
||||
let request_body_object = input.body_json.as_object()?;
|
||||
let mut provider_request_body = serde_json::Map::from_iter(
|
||||
request_body_object
|
||||
.iter()
|
||||
.map(|(key, value)| (key.clone(), value.clone())),
|
||||
);
|
||||
match input.family {
|
||||
SameFormatProviderFamily::Standard => {
|
||||
provider_request_body.insert(
|
||||
"model".to_string(),
|
||||
Value::String(input.mapped_model.to_string()),
|
||||
);
|
||||
if input.upstream_is_stream {
|
||||
provider_request_body.insert("stream".to_string(), Value::Bool(true));
|
||||
}
|
||||
}
|
||||
SameFormatProviderFamily::Gemini => {
|
||||
provider_request_body.remove("model");
|
||||
}
|
||||
}
|
||||
let mut provider_request_body = Value::Object(provider_request_body);
|
||||
if input.is_claude_code {
|
||||
crate::claude_code::sanitize_claude_code_request_body(&mut provider_request_body);
|
||||
}
|
||||
if !apply_local_body_rules(
|
||||
&mut provider_request_body,
|
||||
input.body_rules,
|
||||
Some(input.body_json),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
Some(provider_request_body)
|
||||
}
|
||||
|
||||
pub fn build_same_format_provider_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
params: SameFormatProviderUpstreamUrlParams<'_>,
|
||||
) -> Option<String> {
|
||||
build_transport_request_url(
|
||||
transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: params.provider_api_format,
|
||||
mapped_model: Some(params.mapped_model),
|
||||
upstream_is_stream: params.upstream_is_stream,
|
||||
request_query: params.request_query,
|
||||
kiro_api_region: params.kiro_api_region,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_same_format_provider_headers(
|
||||
input: SameFormatProviderHeadersInput<'_>,
|
||||
) -> Option<BTreeMap<String, String>> {
|
||||
if let Some(kiro_auth_config) = input.kiro_auth_config {
|
||||
return build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
headers: input.headers,
|
||||
provider_request_body: input.provider_request_body,
|
||||
original_request_body: input.original_request_body,
|
||||
header_rules: input.header_rules,
|
||||
auth_header: input.auth_header.unwrap_or_default(),
|
||||
auth_value: input.auth_value.unwrap_or_default(),
|
||||
auth_config: kiro_auth_config,
|
||||
machine_id: input.kiro_machine_id.unwrap_or_default(),
|
||||
});
|
||||
}
|
||||
|
||||
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 {
|
||||
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(
|
||||
input.headers,
|
||||
input.extra_headers,
|
||||
Some("application/json"),
|
||||
)
|
||||
} else {
|
||||
build_complete_passthrough_headers_with_auth(
|
||||
input.headers,
|
||||
auth_header,
|
||||
auth_value,
|
||||
input.extra_headers,
|
||||
Some("application/json"),
|
||||
)
|
||||
};
|
||||
|
||||
let protected_headers = input
|
||||
.auth_header
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.map(|value| vec![value, "content-type"])
|
||||
.unwrap_or_else(|| vec!["content-type"]);
|
||||
if !apply_local_header_rules(
|
||||
&mut provider_request_headers,
|
||||
input.header_rules,
|
||||
&protected_headers,
|
||||
input.provider_request_body,
|
||||
Some(input.original_request_body),
|
||||
) {
|
||||
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);
|
||||
}
|
||||
if input.behavior.upstream_is_stream {
|
||||
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
|
||||
}
|
||||
Some(provider_request_headers)
|
||||
}
|
||||
|
||||
pub fn same_format_provider_transport_supported(
|
||||
behavior: &SameFormatProviderRequestBehavior,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
family: SameFormatProviderFamily,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
same_format_provider_transport_unsupported_reason(behavior, transport, family, api_format)
|
||||
.is_none()
|
||||
}
|
||||
|
||||
pub fn same_format_provider_transport_unsupported_reason(
|
||||
behavior: &SameFormatProviderRequestBehavior,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
family: SameFormatProviderFamily,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
if behavior.is_kiro {
|
||||
local_kiro_request_transport_unsupported_reason_with_network(transport)
|
||||
} else if behavior.is_antigravity {
|
||||
None
|
||||
} else if behavior.is_claude_code {
|
||||
local_claude_code_transport_unsupported_reason_with_network(transport, api_format)
|
||||
} else if behavior.is_vertex {
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(transport)
|
||||
} else {
|
||||
match family {
|
||||
SameFormatProviderFamily::Standard => {
|
||||
local_standard_transport_unsupported_reason_with_network(transport, api_format)
|
||||
}
|
||||
SameFormatProviderFamily::Gemini => {
|
||||
local_gemini_transport_unsupported_reason_with_network(transport, api_format)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn same_format_provider_transport_unsupported_reason_for_trace(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
let normalized_api_format =
|
||||
match aether_ai_formats::normalize_api_format_alias(provider_api_format).as_str() {
|
||||
"openai:chat" => "openai:chat",
|
||||
"openai:responses" => "openai:responses",
|
||||
"openai:responses:compact" => "openai:responses:compact",
|
||||
"claude:messages" => "claude:messages",
|
||||
"gemini:generate_content" => "gemini:generate_content",
|
||||
_ => return Some("transport_api_format_unsupported"),
|
||||
};
|
||||
let behavior = classify_same_format_provider_request_behavior(
|
||||
transport,
|
||||
SameFormatProviderRequestBehaviorParams {
|
||||
require_streaming: false,
|
||||
report_kind: "trace_candidate_metadata",
|
||||
},
|
||||
);
|
||||
if !behavior.is_antigravity
|
||||
&& !behavior.is_claude_code
|
||||
&& !behavior.is_vertex
|
||||
&& !behavior.is_kiro
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let family = if normalized_api_format.starts_with("gemini:") {
|
||||
SameFormatProviderFamily::Gemini
|
||||
} else {
|
||||
SameFormatProviderFamily::Standard
|
||||
};
|
||||
same_format_provider_transport_unsupported_reason(
|
||||
&behavior,
|
||||
transport,
|
||||
family,
|
||||
normalized_api_format,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn should_try_same_format_provider_oauth_auth(
|
||||
behavior: &SameFormatProviderRequestBehavior,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
family: SameFormatProviderFamily,
|
||||
) -> bool {
|
||||
behavior.is_kiro
|
||||
|| matches!(family, SameFormatProviderFamily::Standard)
|
||||
&& resolve_local_standard_auth(transport).is_none()
|
||||
|| matches!(family, SameFormatProviderFamily::Gemini)
|
||||
&& !behavior.is_vertex
|
||||
&& resolve_local_gemini_auth(transport).is_none()
|
||||
}
|
||||
|
||||
pub fn resolve_same_format_provider_direct_auth(
|
||||
behavior: &SameFormatProviderRequestBehavior,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
family: SameFormatProviderFamily,
|
||||
) -> Option<(String, String)> {
|
||||
if behavior.is_vertex {
|
||||
None
|
||||
} else {
|
||||
match family {
|
||||
SameFormatProviderFamily::Standard => resolve_local_standard_auth(transport),
|
||||
SameFormatProviderFamily::Gemini => resolve_local_gemini_auth(transport),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::*;
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider,
|
||||
};
|
||||
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: true,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "openai:chat".to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: "https://api.example.test".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_streaming_and_report_kind_for_provider_private_transports() {
|
||||
let kiro = sample_transport("kiro");
|
||||
let behavior = classify_same_format_provider_request_behavior(
|
||||
&kiro,
|
||||
SameFormatProviderRequestBehaviorParams {
|
||||
require_streaming: false,
|
||||
report_kind: "claude_chat_sync_success",
|
||||
},
|
||||
);
|
||||
|
||||
assert!(behavior.is_kiro);
|
||||
assert!(behavior.upstream_is_stream);
|
||||
assert_eq!(behavior.report_kind, "claude_cli_sync_finalize");
|
||||
|
||||
let antigravity = sample_transport("antigravity");
|
||||
let behavior = classify_same_format_provider_request_behavior(
|
||||
&antigravity,
|
||||
SameFormatProviderRequestBehaviorParams {
|
||||
require_streaming: false,
|
||||
report_kind: "gemini_chat_sync_success",
|
||||
},
|
||||
);
|
||||
|
||||
assert!(behavior.is_antigravity);
|
||||
assert!(behavior.upstream_is_stream);
|
||||
assert_eq!(behavior.report_kind, "gemini_chat_sync_finalize");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_direct_auth_except_vertex() {
|
||||
let transport = sample_transport("openai");
|
||||
let behavior = classify_same_format_provider_request_behavior(
|
||||
&transport,
|
||||
SameFormatProviderRequestBehaviorParams {
|
||||
require_streaming: false,
|
||||
report_kind: "openai_chat_sync_success",
|
||||
},
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
resolve_same_format_provider_direct_auth(
|
||||
&behavior,
|
||||
&transport,
|
||||
SameFormatProviderFamily::Standard,
|
||||
),
|
||||
Some(("x-api-key".to_string(), "secret".to_string()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_same_format_standard_body_with_mapped_model_and_stream_flag() {
|
||||
let body = build_same_format_provider_request_body(SameFormatProviderRequestBodyInput {
|
||||
body_json: &json!({
|
||||
"model": "client-model",
|
||||
"messages": [{"role": "user", "content": "hello"}]
|
||||
}),
|
||||
mapped_model: "upstream-model",
|
||||
family: SameFormatProviderFamily::Standard,
|
||||
body_rules: None,
|
||||
upstream_is_stream: true,
|
||||
kiro_auth_config: None,
|
||||
is_claude_code: false,
|
||||
})
|
||||
.expect("body should build");
|
||||
|
||||
assert_eq!(body.get("model"), Some(&json!("upstream-model")));
|
||||
assert_eq!(body.get("stream"), Some(&json!(true)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
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 headers = build_same_format_provider_headers(SameFormatProviderHeadersInput {
|
||||
headers: &http::HeaderMap::new(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &original_request_body,
|
||||
header_rules: None,
|
||||
behavior: SameFormatProviderRequestBehavior {
|
||||
is_antigravity: false,
|
||||
is_claude_code: false,
|
||||
is_vertex: false,
|
||||
is_kiro: false,
|
||||
upstream_is_stream: true,
|
||||
report_kind: "openai_chat_stream_success",
|
||||
},
|
||||
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,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(headers.get("x-api-key").map(String::as_str), Some("secret"));
|
||||
assert_eq!(
|
||||
headers.get("content-type").map(String::as_str),
|
||||
Some("application/json")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("accept").map(String::as_str),
|
||||
Some("text/event-stream")
|
||||
);
|
||||
}
|
||||
}
|
||||
454
crates/aether-provider-transport/src/standard.rs
Normal file
454
crates/aether-provider-transport/src/standard.rs
Normal file
@@ -0,0 +1,454 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::auth::{
|
||||
build_claude_passthrough_headers, build_complete_passthrough_headers_with_auth,
|
||||
build_openai_passthrough_headers, build_passthrough_headers, ensure_upstream_auth_header,
|
||||
};
|
||||
use crate::rules::{apply_local_body_rules, apply_local_header_rules};
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::url::{build_openai_chat_url, build_openai_responses_url};
|
||||
use crate::vertex::uses_vertex_api_key_query_auth;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct StandardProviderRequestHeadersInput<'a> {
|
||||
pub transport: &'a GatewayProviderTransportSnapshot,
|
||||
pub provider_api_format: &'a str,
|
||||
pub same_format: bool,
|
||||
pub headers: &'a http::HeaderMap,
|
||||
pub auth_header: &'a str,
|
||||
pub auth_value: &'a str,
|
||||
pub extra_headers: &'a BTreeMap<String, String>,
|
||||
pub header_rules: Option<&'a Value>,
|
||||
pub provider_request_body: &'a Value,
|
||||
pub original_request_body: &'a Value,
|
||||
pub upstream_is_stream: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct StandardProviderRequestHeaders {
|
||||
pub headers: BTreeMap<String, String>,
|
||||
pub auth_header: String,
|
||||
pub auth_value: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum StandardPlanFallbackAcceptPolicy {
|
||||
None,
|
||||
TextEventStreamIfStreaming,
|
||||
TextEventStreamRequired,
|
||||
ProviderEventStreamIfMissing,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct StandardPlanFallbackHeadersInput<'a> {
|
||||
pub request_headers: &'a http::HeaderMap,
|
||||
pub existing_provider_request_headers: BTreeMap<String, String>,
|
||||
pub auth_header: Option<&'a str>,
|
||||
pub auth_value: Option<&'a str>,
|
||||
pub extra_headers: &'a BTreeMap<String, String>,
|
||||
pub content_type: Option<&'a str>,
|
||||
pub provider_api_format: &'a str,
|
||||
pub client_api_format: &'a str,
|
||||
pub upstream_is_stream: bool,
|
||||
pub build_from_request_when_empty: bool,
|
||||
pub accept_policy: StandardPlanFallbackAcceptPolicy,
|
||||
}
|
||||
|
||||
pub fn build_standard_plan_fallback_openai_chat_url(
|
||||
upstream_base_url: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> String {
|
||||
build_openai_chat_url(upstream_base_url, request_query)
|
||||
}
|
||||
|
||||
pub fn build_standard_plan_fallback_openai_responses_url(
|
||||
upstream_base_url: &str,
|
||||
request_query: Option<&str>,
|
||||
compact: bool,
|
||||
) -> String {
|
||||
build_openai_responses_url(upstream_base_url, request_query, compact)
|
||||
}
|
||||
|
||||
pub fn build_standard_plan_fallback_headers(
|
||||
input: StandardPlanFallbackHeadersInput<'_>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let auth_pair = input.auth_header.zip(input.auth_value);
|
||||
let mut headers = if !input.existing_provider_request_headers.is_empty() {
|
||||
input.existing_provider_request_headers
|
||||
} else if input.build_from_request_when_empty {
|
||||
match auth_pair {
|
||||
Some((auth_header, auth_value))
|
||||
if input.provider_api_format == input.client_api_format =>
|
||||
{
|
||||
build_complete_passthrough_headers_with_auth(
|
||||
input.request_headers,
|
||||
auth_header,
|
||||
auth_value,
|
||||
input.extra_headers,
|
||||
input.content_type,
|
||||
)
|
||||
}
|
||||
Some((auth_header, auth_value)) if input.provider_api_format.starts_with("claude:") => {
|
||||
build_claude_passthrough_headers(
|
||||
input.request_headers,
|
||||
auth_header,
|
||||
auth_value,
|
||||
input.extra_headers,
|
||||
input.content_type,
|
||||
)
|
||||
}
|
||||
Some((auth_header, auth_value)) => build_openai_passthrough_headers(
|
||||
input.request_headers,
|
||||
auth_header,
|
||||
auth_value,
|
||||
input.extra_headers,
|
||||
input.content_type,
|
||||
),
|
||||
None => build_passthrough_headers(
|
||||
input.request_headers,
|
||||
input.extra_headers,
|
||||
input.content_type,
|
||||
),
|
||||
}
|
||||
} else {
|
||||
input.existing_provider_request_headers
|
||||
};
|
||||
|
||||
if let Some((auth_header, auth_value)) = auth_pair {
|
||||
ensure_upstream_auth_header(&mut headers, auth_header, auth_value);
|
||||
}
|
||||
|
||||
match input.accept_policy {
|
||||
StandardPlanFallbackAcceptPolicy::None => {}
|
||||
StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreaming => {
|
||||
if input.upstream_is_stream {
|
||||
headers
|
||||
.entry("accept".to_string())
|
||||
.or_insert_with(|| "text/event-stream".to_string());
|
||||
}
|
||||
}
|
||||
StandardPlanFallbackAcceptPolicy::TextEventStreamRequired => {
|
||||
headers.insert("accept".to_string(), "text/event-stream".to_string());
|
||||
}
|
||||
StandardPlanFallbackAcceptPolicy::ProviderEventStreamIfMissing => {
|
||||
headers
|
||||
.entry("accept".to_string())
|
||||
.or_insert_with(|| "application/vnd.amazon.eventstream".to_string());
|
||||
}
|
||||
}
|
||||
|
||||
headers
|
||||
}
|
||||
|
||||
pub fn apply_standard_provider_request_body_rules(
|
||||
mut provider_request_body: Value,
|
||||
body_rules: Option<&Value>,
|
||||
original_request_body: &Value,
|
||||
) -> Option<Value> {
|
||||
if !apply_local_body_rules(
|
||||
&mut provider_request_body,
|
||||
body_rules,
|
||||
Some(original_request_body),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
Some(provider_request_body)
|
||||
}
|
||||
|
||||
pub fn build_standard_provider_request_headers(
|
||||
input: StandardProviderRequestHeadersInput<'_>,
|
||||
) -> Option<StandardProviderRequestHeaders> {
|
||||
let uses_vertex_query_auth =
|
||||
uses_vertex_api_key_query_auth(input.transport, input.provider_api_format);
|
||||
let mut headers = if input.same_format {
|
||||
build_complete_passthrough_headers_with_auth(
|
||||
input.headers,
|
||||
input.auth_header,
|
||||
input.auth_value,
|
||||
input.extra_headers,
|
||||
Some("application/json"),
|
||||
)
|
||||
} else if input.provider_api_format.starts_with("claude:") {
|
||||
build_claude_passthrough_headers(
|
||||
input.headers,
|
||||
input.auth_header,
|
||||
input.auth_value,
|
||||
input.extra_headers,
|
||||
Some("application/json"),
|
||||
)
|
||||
} else {
|
||||
build_openai_passthrough_headers(
|
||||
input.headers,
|
||||
input.auth_header,
|
||||
input.auth_value,
|
||||
input.extra_headers,
|
||||
Some("application/json"),
|
||||
)
|
||||
};
|
||||
|
||||
let protected_headers = if uses_vertex_query_auth {
|
||||
&["content-type"][..]
|
||||
} else {
|
||||
&[input.auth_header, "content-type"][..]
|
||||
};
|
||||
if !apply_local_header_rules(
|
||||
&mut headers,
|
||||
input.header_rules,
|
||||
protected_headers,
|
||||
input.provider_request_body,
|
||||
Some(input.original_request_body),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let (auth_header, auth_value) = if uses_vertex_query_auth {
|
||||
headers.remove("x-goog-api-key");
|
||||
(String::new(), String::new())
|
||||
} else {
|
||||
ensure_upstream_auth_header(&mut headers, input.auth_header, input.auth_value);
|
||||
(input.auth_header.to_string(), input.auth_value.to_string())
|
||||
};
|
||||
if input.upstream_is_stream {
|
||||
headers
|
||||
.entry("accept".to_string())
|
||||
.or_insert_with(|| "text/event-stream".to_string());
|
||||
}
|
||||
|
||||
Some(StandardProviderRequestHeaders {
|
||||
headers,
|
||||
auth_header,
|
||||
auth_value,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use http::HeaderMap;
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
apply_standard_provider_request_body_rules, build_standard_plan_fallback_headers,
|
||||
build_standard_plan_fallback_openai_chat_url,
|
||||
build_standard_plan_fallback_openai_responses_url, build_standard_provider_request_headers,
|
||||
StandardPlanFallbackAcceptPolicy, StandardPlanFallbackHeadersInput,
|
||||
StandardProviderRequestHeadersInput,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport(api_format: &str) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Provider".to_string(),
|
||||
provider_type: "openai".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: true,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: api_format.to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: "https://api.example.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "bearer".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_same_format_headers_with_complete_passthrough_and_stream_accept() {
|
||||
let mut request_headers = HeaderMap::new();
|
||||
request_headers.insert("x-client", "demo".parse().expect("header"));
|
||||
let transport = sample_transport("openai:chat");
|
||||
let resolved =
|
||||
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
|
||||
transport: &transport,
|
||||
provider_api_format: "openai:chat",
|
||||
same_format: true,
|
||||
headers: &request_headers,
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer secret",
|
||||
extra_headers: &BTreeMap::new(),
|
||||
header_rules: None,
|
||||
provider_request_body: &json!({"model":"gpt-5"}),
|
||||
original_request_body: &json!({"model":"gpt-5"}),
|
||||
upstream_is_stream: true,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(resolved.auth_header, "authorization");
|
||||
assert_eq!(resolved.auth_value, "Bearer secret");
|
||||
assert_eq!(
|
||||
resolved.headers.get("authorization"),
|
||||
Some(&"Bearer secret".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
resolved.headers.get("accept"),
|
||||
Some(&"text/event-stream".to_string())
|
||||
);
|
||||
assert_eq!(resolved.headers.get("x-client"), Some(&"demo".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn applies_header_rules_after_base_headers() {
|
||||
let transport = sample_transport("claude:messages");
|
||||
let resolved =
|
||||
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
|
||||
transport: &transport,
|
||||
provider_api_format: "claude:messages",
|
||||
same_format: false,
|
||||
headers: &HeaderMap::new(),
|
||||
auth_header: "x-api-key",
|
||||
auth_value: "secret",
|
||||
extra_headers: &BTreeMap::new(),
|
||||
header_rules: Some(&json!([
|
||||
{"action":"set","key":"x-route","value":"standard"}
|
||||
])),
|
||||
provider_request_body: &json!({"model":"claude"}),
|
||||
original_request_body: &json!({"model":"claude"}),
|
||||
upstream_is_stream: false,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
resolved.headers.get("x-api-key"),
|
||||
Some(&"secret".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
resolved.headers.get("x-route"),
|
||||
Some(&"standard".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
resolved.headers.get("anthropic-version"),
|
||||
Some(&"2023-06-01".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn applies_standard_body_rules_to_surface_built_body() {
|
||||
let body = apply_standard_provider_request_body_rules(
|
||||
json!({"model":"gpt-5"}),
|
||||
Some(&json!([
|
||||
{"action":"set","path":"metadata.source","value":"standard"}
|
||||
])),
|
||||
&json!({"model":"client"}),
|
||||
)
|
||||
.expect("body rules should apply");
|
||||
|
||||
assert_eq!(body["metadata"]["source"], json!("standard"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_plan_fallback_headers_from_request_when_enabled() {
|
||||
let mut request_headers = HeaderMap::new();
|
||||
request_headers.insert("x-client", "demo".parse().expect("header"));
|
||||
|
||||
let headers = build_standard_plan_fallback_headers(StandardPlanFallbackHeadersInput {
|
||||
request_headers: &request_headers,
|
||||
existing_provider_request_headers: BTreeMap::new(),
|
||||
auth_header: Some("authorization"),
|
||||
auth_value: Some("Bearer secret"),
|
||||
extra_headers: &BTreeMap::new(),
|
||||
content_type: Some("application/json"),
|
||||
provider_api_format: "openai:chat",
|
||||
client_api_format: "openai:chat",
|
||||
upstream_is_stream: true,
|
||||
build_from_request_when_empty: true,
|
||||
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreaming,
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
headers.get("authorization"),
|
||||
Some(&"Bearer secret".to_string())
|
||||
);
|
||||
assert_eq!(headers.get("x-client"), Some(&"demo".to_string()));
|
||||
assert_eq!(
|
||||
headers.get("accept"),
|
||||
Some(&"text/event-stream".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn plan_fallback_headers_preserve_empty_existing_mode() {
|
||||
let headers = build_standard_plan_fallback_headers(StandardPlanFallbackHeadersInput {
|
||||
request_headers: &HeaderMap::new(),
|
||||
existing_provider_request_headers: BTreeMap::new(),
|
||||
auth_header: Some("authorization"),
|
||||
auth_value: Some("Bearer secret"),
|
||||
extra_headers: &BTreeMap::new(),
|
||||
content_type: Some("application/json"),
|
||||
provider_api_format: "openai:responses",
|
||||
client_api_format: "openai:responses",
|
||||
upstream_is_stream: true,
|
||||
build_from_request_when_empty: false,
|
||||
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreaming,
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
headers.get("authorization"),
|
||||
Some(&"Bearer secret".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("accept"),
|
||||
Some(&"text/event-stream".to_string())
|
||||
);
|
||||
assert!(!headers.contains_key("content-type"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn plan_fallback_url_helpers_route_openai_surfaces() {
|
||||
assert_eq!(
|
||||
build_standard_plan_fallback_openai_chat_url("https://api.example.com", Some("x=1")),
|
||||
"https://api.example.com/v1/chat/completions?x=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_standard_plan_fallback_openai_responses_url(
|
||||
"https://api.example.com",
|
||||
Some("x=1"),
|
||||
true,
|
||||
),
|
||||
"https://api.example.com/v1/responses/compact?x=1"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -1,13 +1,41 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_data_contracts::repository::video_tasks::StoredVideoTask;
|
||||
use aether_video_tasks_core::{
|
||||
LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::auth::{resolve_local_gemini_auth, resolve_local_openai_bearer_auth};
|
||||
use super::auth::{
|
||||
build_passthrough_headers_with_auth, resolve_local_gemini_auth,
|
||||
resolve_local_openai_bearer_auth,
|
||||
};
|
||||
use super::network::resolve_transport_execution_timeouts;
|
||||
use super::policy::{supports_local_gemini_transport, supports_local_standard_transport};
|
||||
use super::policy::{
|
||||
local_gemini_transport_unsupported_reason_with_network,
|
||||
local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport,
|
||||
supports_local_standard_transport,
|
||||
};
|
||||
use super::rules::{apply_local_body_rules, apply_local_header_rules};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::url::{build_gemini_video_predict_long_running_url, build_passthrough_path_url};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ProviderVideoCreateFamily {
|
||||
OpenAi,
|
||||
Gemini,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct ProviderVideoCreateHeadersInput<'a> {
|
||||
pub headers: &'a http::HeaderMap,
|
||||
pub auth_header: &'a str,
|
||||
pub auth_value: &'a str,
|
||||
pub header_rules: Option<&'a Value>,
|
||||
pub provider_request_body: &'a Value,
|
||||
pub original_request_body: &'a Value,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait VideoTaskTransportSnapshotLookup: Send + Sync {
|
||||
@@ -59,6 +87,115 @@ pub fn resolve_local_video_task_transport(
|
||||
))
|
||||
}
|
||||
|
||||
pub fn video_create_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
family: ProviderVideoCreateFamily,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
match family {
|
||||
ProviderVideoCreateFamily::OpenAi => {
|
||||
local_standard_transport_unsupported_reason_with_network(transport, api_format)
|
||||
}
|
||||
ProviderVideoCreateFamily::Gemini => {
|
||||
local_gemini_transport_unsupported_reason_with_network(transport, api_format)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_video_create_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
family: ProviderVideoCreateFamily,
|
||||
) -> Option<(String, String)> {
|
||||
match family {
|
||||
ProviderVideoCreateFamily::OpenAi => resolve_local_openai_bearer_auth(transport),
|
||||
ProviderVideoCreateFamily::Gemini => resolve_local_gemini_auth(transport),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_video_create_request_body(
|
||||
body_json: &Value,
|
||||
family: ProviderVideoCreateFamily,
|
||||
mapped_model: &str,
|
||||
body_rules: Option<&Value>,
|
||||
) -> Option<Value> {
|
||||
let mut provider_request_body = match family {
|
||||
ProviderVideoCreateFamily::OpenAi => {
|
||||
let mut provider_request_body = body_json.as_object().cloned().unwrap_or_default();
|
||||
provider_request_body
|
||||
.insert("model".to_string(), Value::String(mapped_model.to_string()));
|
||||
Value::Object(provider_request_body)
|
||||
}
|
||||
ProviderVideoCreateFamily::Gemini => body_json.clone(),
|
||||
};
|
||||
if !apply_local_body_rules(&mut provider_request_body, body_rules, Some(body_json)) {
|
||||
return None;
|
||||
}
|
||||
Some(provider_request_body)
|
||||
}
|
||||
|
||||
pub fn build_video_create_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request_path: &str,
|
||||
request_query: Option<&str>,
|
||||
mapped_model: &str,
|
||||
family: ProviderVideoCreateFamily,
|
||||
) -> Option<String> {
|
||||
let custom_path = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
if let Some(path) = custom_path {
|
||||
let blocked_keys = match family {
|
||||
ProviderVideoCreateFamily::OpenAi => &[][..],
|
||||
ProviderVideoCreateFamily::Gemini => &["key"][..],
|
||||
};
|
||||
return build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
path,
|
||||
request_query,
|
||||
blocked_keys,
|
||||
);
|
||||
}
|
||||
|
||||
match family {
|
||||
ProviderVideoCreateFamily::OpenAi => build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_path,
|
||||
request_query,
|
||||
&[],
|
||||
),
|
||||
ProviderVideoCreateFamily::Gemini => build_gemini_video_predict_long_running_url(
|
||||
&transport.endpoint.base_url,
|
||||
mapped_model,
|
||||
request_query,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_video_create_headers(
|
||||
input: ProviderVideoCreateHeadersInput<'_>,
|
||||
) -> Option<BTreeMap<String, String>> {
|
||||
let mut provider_request_headers = build_passthrough_headers_with_auth(
|
||||
input.headers,
|
||||
input.auth_header,
|
||||
input.auth_value,
|
||||
&BTreeMap::new(),
|
||||
);
|
||||
if !apply_local_header_rules(
|
||||
&mut provider_request_headers,
|
||||
input.header_rules,
|
||||
&[input.auth_header, "content-type"],
|
||||
input.provider_request_body,
|
||||
Some(input.original_request_body),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
Some(provider_request_headers)
|
||||
}
|
||||
|
||||
pub async fn reconstruct_local_video_task_snapshot(
|
||||
lookup: &dyn VideoTaskTransportSnapshotLookup,
|
||||
task: &StoredVideoTask,
|
||||
@@ -109,8 +246,10 @@ mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
reconstruct_local_video_task_snapshot, resolve_local_video_task_transport,
|
||||
VideoTaskTransportSnapshotLookup,
|
||||
build_video_create_headers, build_video_create_request_body,
|
||||
build_video_create_upstream_url, reconstruct_local_video_task_snapshot,
|
||||
resolve_local_video_task_transport, ProviderVideoCreateFamily,
|
||||
ProviderVideoCreateHeadersInput, VideoTaskTransportSnapshotLookup,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
@@ -262,6 +401,64 @@ mod tests {
|
||||
assert_eq!(transport.endpoint_id, "endpoint-1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_openai_video_create_request_body_with_mapped_model() {
|
||||
let body = build_video_create_request_body(
|
||||
&json!({"prompt": "make a clip", "model": "client-model"}),
|
||||
ProviderVideoCreateFamily::OpenAi,
|
||||
"upstream-video-model",
|
||||
None,
|
||||
)
|
||||
.expect("body should build");
|
||||
|
||||
assert_eq!(body.get("prompt"), Some(&json!("make a clip")));
|
||||
assert_eq!(body.get("model"), Some(&json!("upstream-video-model")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_gemini_video_create_url_and_removes_client_key_query() {
|
||||
let transport = sample_transport("gemini:video", "api_key");
|
||||
let url = build_video_create_upstream_url(
|
||||
&transport,
|
||||
"/v1beta/models/client-model:predictLongRunning",
|
||||
Some("key=client-key&trace=1"),
|
||||
"veo-upstream",
|
||||
ProviderVideoCreateFamily::Gemini,
|
||||
)
|
||||
.expect("url should build");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://example.com/v1beta/models/veo-upstream:predictLongRunning?trace=1"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_video_create_headers_with_auth_and_rules() {
|
||||
let provider_request_body = json!({"prompt": "make a clip"});
|
||||
let original_request_body = provider_request_body.clone();
|
||||
let headers = build_video_create_headers(ProviderVideoCreateHeadersInput {
|
||||
headers: &http::HeaderMap::new(),
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer secret",
|
||||
header_rules: Some(&json!([
|
||||
{"action":"set","key":"x-provider-tag","value":"video"}
|
||||
])),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &original_request_body,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer secret")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-provider-tag").map(String::as_str),
|
||||
Some("video")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_mismatched_video_transport_format() {
|
||||
let transport = sample_transport("openai:chat", "bearer");
|
||||
|
||||
Reference in New Issue
Block a user