mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
refactor(workspace): enforce layered crate boundaries
This commit is contained in:
@@ -0,0 +1,406 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
pub const ANTIGRAVITY_PROVIDER_TYPE: &str = "antigravity";
|
||||
pub const ANTIGRAVITY_REQUEST_USER_AGENT: &str =
|
||||
"antigravity/cli/1.0.16 (aidev_client; os_type=linux; arch=arm64; auth_method=consumer)";
|
||||
const ANTIGRAVITY_CLIENT_NAME: &str = "antigravity";
|
||||
const ANTIGRAVITY_GOOG_API_CLIENT: &str = "gl-node/18.18.2 fire/0.8.6 grpc/1.10.x";
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct AntigravityRequestAuth {
|
||||
pub project_id: String,
|
||||
pub client_version: Option<String>,
|
||||
pub session_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum AntigravityRequestAuthSupport {
|
||||
Supported(AntigravityRequestAuth),
|
||||
Unsupported(AntigravityRequestAuthUnsupportedReason),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum AntigravityRequestAuthUnsupportedReason {
|
||||
WrongProviderType,
|
||||
MissingAuthConfig,
|
||||
InvalidAuthConfigJson,
|
||||
ComplexDynamicAuthConfig,
|
||||
MissingProjectId,
|
||||
}
|
||||
|
||||
pub fn resolve_local_antigravity_request_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> AntigravityRequestAuthSupport {
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(ANTIGRAVITY_PROVIDER_TYPE)
|
||||
{
|
||||
return AntigravityRequestAuthSupport::Unsupported(
|
||||
AntigravityRequestAuthUnsupportedReason::WrongProviderType,
|
||||
);
|
||||
}
|
||||
|
||||
let Some(raw_auth_config) = transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return AntigravityRequestAuthSupport::Unsupported(
|
||||
AntigravityRequestAuthUnsupportedReason::MissingAuthConfig,
|
||||
);
|
||||
};
|
||||
|
||||
let Ok(auth_config) = serde_json::from_str::<Value>(raw_auth_config) else {
|
||||
return AntigravityRequestAuthSupport::Unsupported(
|
||||
AntigravityRequestAuthUnsupportedReason::InvalidAuthConfigJson,
|
||||
);
|
||||
};
|
||||
|
||||
if contains_blocked_auth_fields(&auth_config) {
|
||||
return AntigravityRequestAuthSupport::Unsupported(
|
||||
AntigravityRequestAuthUnsupportedReason::ComplexDynamicAuthConfig,
|
||||
);
|
||||
}
|
||||
|
||||
let upstream_metadata = transport.key.upstream_metadata.as_ref();
|
||||
let Some(project_id) = find_antigravity_string(
|
||||
upstream_metadata,
|
||||
&auth_config,
|
||||
ANTIGRAVITY_PROJECT_ID_PATHS,
|
||||
) else {
|
||||
return AntigravityRequestAuthSupport::Unsupported(
|
||||
AntigravityRequestAuthUnsupportedReason::MissingProjectId,
|
||||
);
|
||||
};
|
||||
|
||||
let client_version = find_antigravity_string(
|
||||
upstream_metadata,
|
||||
&auth_config,
|
||||
ANTIGRAVITY_CLIENT_VERSION_PATHS,
|
||||
);
|
||||
let session_id = find_antigravity_string(
|
||||
upstream_metadata,
|
||||
&auth_config,
|
||||
ANTIGRAVITY_SESSION_ID_PATHS,
|
||||
);
|
||||
|
||||
AntigravityRequestAuthSupport::Supported(AntigravityRequestAuth {
|
||||
project_id,
|
||||
client_version,
|
||||
session_id,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn build_antigravity_static_identity_headers(
|
||||
auth: &AntigravityRequestAuth,
|
||||
) -> BTreeMap<String, String> {
|
||||
build_antigravity_static_client_headers(
|
||||
auth.client_version.as_deref(),
|
||||
auth.session_id.as_deref(),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_antigravity_static_client_headers(
|
||||
client_version: Option<&str>,
|
||||
session_id: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut headers = BTreeMap::from([
|
||||
(
|
||||
String::from("x-client-name"),
|
||||
String::from(ANTIGRAVITY_CLIENT_NAME),
|
||||
),
|
||||
(
|
||||
String::from("x-goog-api-client"),
|
||||
String::from(ANTIGRAVITY_GOOG_API_CLIENT),
|
||||
),
|
||||
(
|
||||
String::from("user-agent"),
|
||||
String::from(ANTIGRAVITY_REQUEST_USER_AGENT),
|
||||
),
|
||||
]);
|
||||
|
||||
if let Some(client_version) = client_version
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
headers.insert(String::from("x-client-version"), client_version.to_string());
|
||||
}
|
||||
if let Some(session_id) = session_id.map(str::trim).filter(|value| !value.is_empty()) {
|
||||
headers.insert(String::from("x-vscode-sessionid"), session_id.to_string());
|
||||
}
|
||||
|
||||
headers
|
||||
}
|
||||
|
||||
const ANTIGRAVITY_PROJECT_ID_PATHS: &[&[&str]] = &[
|
||||
&["project_id"],
|
||||
&["projectId"],
|
||||
&["project", "id"],
|
||||
&["project", "project_id"],
|
||||
&["project", "projectId"],
|
||||
&["cloudaicompanionProject"],
|
||||
&["cloudaicompanionProject", "id"],
|
||||
&["cloudAiCompanionProject"],
|
||||
&["cloudAiCompanionProject", "id"],
|
||||
&["antigravity", "project_id"],
|
||||
&["antigravity", "projectId"],
|
||||
&["antigravity", "project", "id"],
|
||||
&["antigravity", "cloudaicompanionProject"],
|
||||
&["antigravity", "cloudaicompanionProject", "id"],
|
||||
&["antigravity", "cloudAiCompanionProject"],
|
||||
&["antigravity", "cloudAiCompanionProject", "id"],
|
||||
&["metadata", "project_id"],
|
||||
&["metadata", "projectId"],
|
||||
&["metadata", "cloudaicompanionProject"],
|
||||
&["metadata", "cloudaicompanionProject", "id"],
|
||||
&["metadata", "cloudAiCompanionProject"],
|
||||
&["metadata", "cloudAiCompanionProject", "id"],
|
||||
];
|
||||
|
||||
const ANTIGRAVITY_CLIENT_VERSION_PATHS: &[&[&str]] = &[
|
||||
&["client_version"],
|
||||
&["clientVersion"],
|
||||
&["antigravity", "client_version"],
|
||||
&["antigravity", "clientVersion"],
|
||||
&["metadata", "client_version"],
|
||||
&["metadata", "clientVersion"],
|
||||
];
|
||||
|
||||
const ANTIGRAVITY_SESSION_ID_PATHS: &[&[&str]] = &[
|
||||
&["session_id"],
|
||||
&["sessionId"],
|
||||
&["antigravity", "session_id"],
|
||||
&["antigravity", "sessionId"],
|
||||
&["metadata", "session_id"],
|
||||
&["metadata", "sessionId"],
|
||||
];
|
||||
|
||||
fn find_antigravity_string(
|
||||
upstream_metadata: Option<&Value>,
|
||||
auth_config: &Value,
|
||||
paths: &[&[&str]],
|
||||
) -> Option<String> {
|
||||
upstream_metadata
|
||||
.and_then(|metadata| find_string_by_paths(metadata, paths))
|
||||
.or_else(|| find_string_by_paths(auth_config, paths))
|
||||
}
|
||||
|
||||
fn find_string_by_paths(value: &Value, paths: &[&[&str]]) -> Option<String> {
|
||||
for path in paths {
|
||||
let mut current = value;
|
||||
let mut matched = true;
|
||||
for segment in *path {
|
||||
let Some(next) = current.get(*segment) else {
|
||||
matched = false;
|
||||
break;
|
||||
};
|
||||
current = next;
|
||||
}
|
||||
if !matched {
|
||||
continue;
|
||||
}
|
||||
if let Some(string) = current
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|item| !item.is_empty())
|
||||
{
|
||||
return Some(string.to_string());
|
||||
}
|
||||
if let Some(string) = current
|
||||
.as_object()
|
||||
.and_then(|object| {
|
||||
object
|
||||
.get("id")
|
||||
.or_else(|| object.get("project_id"))
|
||||
.or_else(|| object.get("projectId"))
|
||||
})
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|item| !item.is_empty())
|
||||
{
|
||||
return Some(string.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
fn contains_blocked_auth_fields(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::Object(map) => map.iter().any(|(key, inner)| {
|
||||
is_blocked_auth_key(key.as_str()) || contains_blocked_auth_fields(inner)
|
||||
}),
|
||||
Value::Array(items) => items.iter().any(contains_blocked_auth_fields),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_blocked_auth_key(key: &str) -> bool {
|
||||
matches!(
|
||||
key.trim().to_ascii_lowercase().as_str(),
|
||||
"private_key"
|
||||
| "privateKey"
|
||||
| "private_key_id"
|
||||
| "privateKeyId"
|
||||
| "service_account"
|
||||
| "serviceAccount"
|
||||
| "service_account_json"
|
||||
| "serviceAccountJson"
|
||||
| "service_account_key"
|
||||
| "serviceAccountKey"
|
||||
| "credential_source"
|
||||
| "credentialSource"
|
||||
| "token_url"
|
||||
| "tokenUrl"
|
||||
| "auth_uri"
|
||||
| "authUri"
|
||||
| "subject"
|
||||
| "audience"
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
build_antigravity_static_client_headers, resolve_local_antigravity_request_auth,
|
||||
AntigravityRequestAuth, AntigravityRequestAuthSupport, ANTIGRAVITY_REQUEST_USER_AGENT,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport(auth_config: &str) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Antigravity".to_string(),
|
||||
provider_type: "antigravity".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: true,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "gemini:generate_content".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("generate_content".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://daily-cloudcode-pa.googleapis.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "oauth".to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec!["gemini:generate_content".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: Some(auth_config.to_string()),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_cloudaicompanion_project_object_from_auth_config() {
|
||||
let transport = sample_transport(
|
||||
r#"{
|
||||
"provider_type":"antigravity",
|
||||
"refresh_token":"rt",
|
||||
"cloudaicompanionProject":{"id":"project-from-auth-config"}
|
||||
}"#,
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
resolve_local_antigravity_request_auth(&transport),
|
||||
AntigravityRequestAuthSupport::Supported(AntigravityRequestAuth {
|
||||
project_id: "project-from-auth-config".to_string(),
|
||||
client_version: None,
|
||||
session_id: None,
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_identity_from_antigravity_upstream_metadata() {
|
||||
let mut transport = sample_transport(
|
||||
r#"{
|
||||
"provider_type":"antigravity",
|
||||
"refresh_token":"rt"
|
||||
}"#,
|
||||
);
|
||||
transport.key.upstream_metadata = Some(json!({
|
||||
"antigravity": {
|
||||
"project_id": "project-from-metadata",
|
||||
"client_version": "1.99.0",
|
||||
"session_id": "session-from-metadata"
|
||||
}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
resolve_local_antigravity_request_auth(&transport),
|
||||
AntigravityRequestAuthSupport::Supported(AntigravityRequestAuth {
|
||||
project_id: "project-from-metadata".to_string(),
|
||||
client_version: Some("1.99.0".to_string()),
|
||||
session_id: Some("session-from-metadata".to_string()),
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn static_client_headers_use_native_antigravity_cli_user_agent() {
|
||||
let headers = build_antigravity_static_client_headers(Some("1.0.16"), Some("session-abc"));
|
||||
|
||||
assert_eq!(
|
||||
headers.get("user-agent").map(String::as_str),
|
||||
Some(ANTIGRAVITY_REQUEST_USER_AGENT)
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-client-name").map(String::as_str),
|
||||
Some("antigravity")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-client-version").map(String::as_str),
|
||||
Some("1.0.16")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-vscode-sessionid").map(String::as_str),
|
||||
Some("session-abc")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
mod auth;
|
||||
mod policy;
|
||||
mod request;
|
||||
mod url;
|
||||
|
||||
pub use auth::{
|
||||
build_antigravity_static_client_headers, build_antigravity_static_identity_headers,
|
||||
resolve_local_antigravity_request_auth, AntigravityRequestAuth, AntigravityRequestAuthSupport,
|
||||
AntigravityRequestAuthUnsupportedReason, ANTIGRAVITY_PROVIDER_TYPE,
|
||||
ANTIGRAVITY_REQUEST_USER_AGENT,
|
||||
};
|
||||
pub use policy::{
|
||||
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,
|
||||
AntigravityEnvelopeRequestType, AntigravityRequestEnvelopeSupport,
|
||||
AntigravityRequestEnvelopeUnsupportedReason,
|
||||
};
|
||||
pub use url::{
|
||||
build_antigravity_v1internal_url, AntigravityRequestUrlAction,
|
||||
ANTIGRAVITY_V1INTERNAL_PATH_TEMPLATE,
|
||||
};
|
||||
@@ -0,0 +1,116 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::auth::{
|
||||
resolve_local_antigravity_request_auth, AntigravityRequestAuth, AntigravityRequestAuthSupport,
|
||||
AntigravityRequestAuthUnsupportedReason, ANTIGRAVITY_PROVIDER_TYPE,
|
||||
};
|
||||
use super::request::{
|
||||
classify_antigravity_safe_request_body, AntigravityEnvelopeRequestType,
|
||||
AntigravityRequestEnvelopeUnsupportedReason,
|
||||
};
|
||||
use crate::rules::{body_rules_have_enabled_rules, header_rules_have_enabled_rules};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct AntigravityRequestSideSpec {
|
||||
pub auth: AntigravityRequestAuth,
|
||||
pub request_type: AntigravityEnvelopeRequestType,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum AntigravityRequestSideSupport {
|
||||
Supported(AntigravityRequestSideSpec),
|
||||
Unsupported(AntigravityRequestSideUnsupportedReason),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum AntigravityRequestSideUnsupportedReason {
|
||||
InactiveTransport,
|
||||
WrongProviderType,
|
||||
UnsupportedApiFormat,
|
||||
UnsupportedCustomPath,
|
||||
UnsupportedHeaderRules,
|
||||
UnsupportedBodyRules,
|
||||
UnsupportedNetworkConfig,
|
||||
UnsupportedAuth(AntigravityRequestAuthUnsupportedReason),
|
||||
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,
|
||||
request_type: AntigravityEnvelopeRequestType,
|
||||
) -> AntigravityRequestSideSupport {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::InactiveTransport,
|
||||
);
|
||||
}
|
||||
if !is_antigravity_provider_transport(transport) {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::WrongProviderType,
|
||||
);
|
||||
}
|
||||
|
||||
let endpoint_format =
|
||||
aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
|
||||
if endpoint_format != "gemini:generate_content" {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedApiFormat,
|
||||
);
|
||||
}
|
||||
if transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.is_some_and(|value| !value.trim().is_empty())
|
||||
{
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedCustomPath,
|
||||
);
|
||||
}
|
||||
if header_rules_have_enabled_rules(transport.endpoint.header_rules.as_ref()) {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedHeaderRules,
|
||||
);
|
||||
}
|
||||
if body_rules_have_enabled_rules(transport.endpoint.body_rules.as_ref()) {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedBodyRules,
|
||||
);
|
||||
}
|
||||
if transport.provider.proxy.is_some()
|
||||
|| transport.endpoint.proxy.is_some()
|
||||
|| transport.key.proxy.is_some()
|
||||
|| transport.key.fingerprint.is_some()
|
||||
{
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedNetworkConfig,
|
||||
);
|
||||
}
|
||||
|
||||
let auth = match resolve_local_antigravity_request_auth(transport) {
|
||||
AntigravityRequestAuthSupport::Supported(auth) => auth,
|
||||
AntigravityRequestAuthSupport::Unsupported(reason) => {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedAuth(reason),
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(reason) = classify_antigravity_safe_request_body(request_body) {
|
||||
return AntigravityRequestSideSupport::Unsupported(
|
||||
AntigravityRequestSideUnsupportedReason::UnsupportedEnvelope(reason),
|
||||
);
|
||||
}
|
||||
|
||||
AntigravityRequestSideSupport::Supported(AntigravityRequestSideSpec { auth, request_type })
|
||||
}
|
||||
@@ -0,0 +1,367 @@
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::auth::{AntigravityRequestAuth, ANTIGRAVITY_REQUEST_USER_AGENT};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum AntigravityEnvelopeRequestType {
|
||||
Agent,
|
||||
Checkpoint,
|
||||
EndpointTest,
|
||||
}
|
||||
|
||||
impl AntigravityEnvelopeRequestType {
|
||||
fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Agent => "agent",
|
||||
Self::Checkpoint => "checkpoint",
|
||||
Self::EndpointTest => "endpoint_test",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum AntigravityRequestEnvelopeSupport {
|
||||
Supported(Value),
|
||||
Unsupported(AntigravityRequestEnvelopeUnsupportedReason),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum AntigravityRequestEnvelopeUnsupportedReason {
|
||||
NonObjectBody,
|
||||
MissingContents,
|
||||
MissingRequestId,
|
||||
MissingModel,
|
||||
}
|
||||
|
||||
pub fn classify_antigravity_safe_request_body(
|
||||
request_body: &Value,
|
||||
) -> Result<(), AntigravityRequestEnvelopeUnsupportedReason> {
|
||||
let Value::Object(map) = request_body else {
|
||||
return Err(AntigravityRequestEnvelopeUnsupportedReason::NonObjectBody);
|
||||
};
|
||||
if !map.contains_key("contents") && existing_v1internal_request_object(map).is_none() {
|
||||
return Err(AntigravityRequestEnvelopeUnsupportedReason::MissingContents);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn build_antigravity_safe_v1internal_request(
|
||||
auth: &AntigravityRequestAuth,
|
||||
request_id: &str,
|
||||
model: &str,
|
||||
request_body: &Value,
|
||||
request_type: AntigravityEnvelopeRequestType,
|
||||
) -> AntigravityRequestEnvelopeSupport {
|
||||
if request_id.trim().is_empty() {
|
||||
return AntigravityRequestEnvelopeSupport::Unsupported(
|
||||
AntigravityRequestEnvelopeUnsupportedReason::MissingRequestId,
|
||||
);
|
||||
}
|
||||
if model.trim().is_empty() {
|
||||
return AntigravityRequestEnvelopeSupport::Unsupported(
|
||||
AntigravityRequestEnvelopeUnsupportedReason::MissingModel,
|
||||
);
|
||||
}
|
||||
if let Err(reason) = classify_antigravity_safe_request_body(request_body) {
|
||||
return AntigravityRequestEnvelopeSupport::Unsupported(reason);
|
||||
}
|
||||
|
||||
let Value::Object(source) = request_body else {
|
||||
return AntigravityRequestEnvelopeSupport::Unsupported(
|
||||
AntigravityRequestEnvelopeUnsupportedReason::NonObjectBody,
|
||||
);
|
||||
};
|
||||
|
||||
if let Some(existing_request) = existing_v1internal_request_object(source) {
|
||||
let mut inner_request: Map<String, Value> = existing_request.clone();
|
||||
inner_request.remove("model");
|
||||
inner_request.remove("safetySettings");
|
||||
inner_request.remove("safety_settings");
|
||||
let request_id = non_empty_string_field(source, "requestId").unwrap_or(request_id);
|
||||
let user_agent =
|
||||
non_empty_string_field(source, "userAgent").unwrap_or(ANTIGRAVITY_REQUEST_USER_AGENT);
|
||||
let request_type =
|
||||
existing_v1internal_request_type(source).unwrap_or_else(|| request_type.as_str());
|
||||
|
||||
return AntigravityRequestEnvelopeSupport::Supported(serde_json::json!({
|
||||
"project": auth.project_id,
|
||||
"requestId": request_id,
|
||||
"request": Value::Object(inner_request),
|
||||
"model": model,
|
||||
"userAgent": user_agent,
|
||||
"requestType": request_type,
|
||||
}));
|
||||
}
|
||||
|
||||
let mut inner_request: Map<String, Value> = source.clone();
|
||||
inner_request.remove("model");
|
||||
inner_request.remove("safetySettings");
|
||||
inner_request.remove("safety_settings");
|
||||
|
||||
AntigravityRequestEnvelopeSupport::Supported(serde_json::json!({
|
||||
"project": auth.project_id,
|
||||
"requestId": request_id,
|
||||
"request": Value::Object(inner_request),
|
||||
"model": model,
|
||||
"userAgent": ANTIGRAVITY_REQUEST_USER_AGENT,
|
||||
"requestType": request_type.as_str(),
|
||||
}))
|
||||
}
|
||||
|
||||
fn existing_v1internal_request_object(source: &Map<String, Value>) -> Option<&Map<String, Value>> {
|
||||
source
|
||||
.get("request")
|
||||
.and_then(Value::as_object)
|
||||
.filter(|request| request.contains_key("contents"))
|
||||
}
|
||||
|
||||
fn non_empty_string_field<'a>(source: &'a Map<String, Value>, key: &str) -> Option<&'a str> {
|
||||
source
|
||||
.get(key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn existing_v1internal_request_type(source: &Map<String, Value>) -> Option<&str> {
|
||||
match non_empty_string_field(source, "requestType")? {
|
||||
"agent" => Some("agent"),
|
||||
"checkpoint" => Some("checkpoint"),
|
||||
"endpoint_test" => Some("endpoint_test"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
build_antigravity_safe_v1internal_request, classify_antigravity_safe_request_body,
|
||||
AntigravityEnvelopeRequestType, AntigravityRequestAuth, AntigravityRequestEnvelopeSupport,
|
||||
};
|
||||
use crate::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT;
|
||||
|
||||
fn sample_auth() -> AntigravityRequestAuth {
|
||||
AntigravityRequestAuth {
|
||||
project_id: "project-ant-123".to_string(),
|
||||
client_version: None,
|
||||
session_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn real_agent_request_preserves_antigravity_agent_fields() {
|
||||
let request_body = json!({
|
||||
"model": "client-side-model-should-not-be-nested",
|
||||
"contents": [
|
||||
{
|
||||
"role": "user",
|
||||
"parts": [
|
||||
{ "text": "Reply with OK only." }
|
||||
]
|
||||
}
|
||||
],
|
||||
"systemInstruction": {
|
||||
"role": "user",
|
||||
"parts": [
|
||||
{ "text": "Antigravity agent system prompt" }
|
||||
]
|
||||
},
|
||||
"generationConfig": {
|
||||
"maxOutputTokens": 8192,
|
||||
"thinkingConfig": {
|
||||
"includeThoughts": true,
|
||||
"thinkingBudget": 4000
|
||||
}
|
||||
},
|
||||
"toolConfig": {
|
||||
"functionCallingConfig": {
|
||||
"mode": "VALIDATED"
|
||||
}
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"functionDeclarations": [
|
||||
{
|
||||
"name": "run_command",
|
||||
"description": "Run a command",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"cmd": { "type": "string" }
|
||||
},
|
||||
"required": ["cmd"]
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"labels": {
|
||||
"trajectory_id": "trajectory-123",
|
||||
"used_claude": "false"
|
||||
},
|
||||
"sessionId": "session-ant-123",
|
||||
"safetySettings": [
|
||||
{ "category": "HARM_CATEGORY_UNSPECIFIED" }
|
||||
]
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
classify_antigravity_safe_request_body(&request_body),
|
||||
Ok(())
|
||||
);
|
||||
|
||||
let envelope = match build_antigravity_safe_v1internal_request(
|
||||
&sample_auth(),
|
||||
"request-ant-agent-123",
|
||||
"gemini-3.5-flash-low",
|
||||
&request_body,
|
||||
AntigravityEnvelopeRequestType::Agent,
|
||||
) {
|
||||
AntigravityRequestEnvelopeSupport::Supported(envelope) => envelope,
|
||||
AntigravityRequestEnvelopeSupport::Unsupported(reason) => {
|
||||
panic!("real agent envelope should be supported: {reason:?}")
|
||||
}
|
||||
};
|
||||
|
||||
assert_eq!(envelope["project"], "project-ant-123");
|
||||
assert_eq!(envelope["requestId"], "request-ant-agent-123");
|
||||
assert_eq!(envelope["model"], "gemini-3.5-flash-low");
|
||||
assert_eq!(envelope["userAgent"], ANTIGRAVITY_REQUEST_USER_AGENT);
|
||||
assert_eq!(envelope["requestType"], "agent");
|
||||
assert!(envelope["request"].get("model").is_none());
|
||||
assert!(envelope["request"].get("safetySettings").is_none());
|
||||
assert_eq!(
|
||||
envelope["request"]["systemInstruction"]["parts"][0]["text"],
|
||||
"Antigravity agent system prompt"
|
||||
);
|
||||
assert_eq!(
|
||||
envelope["request"]["generationConfig"]["thinkingConfig"]["thinkingBudget"],
|
||||
4000
|
||||
);
|
||||
assert_eq!(
|
||||
envelope["request"]["toolConfig"]["functionCallingConfig"]["mode"],
|
||||
"VALIDATED"
|
||||
);
|
||||
assert_eq!(
|
||||
envelope["request"]["tools"][0]["functionDeclarations"][0]["name"],
|
||||
"run_command"
|
||||
);
|
||||
assert_eq!(
|
||||
envelope["request"]["labels"]["trajectory_id"],
|
||||
"trajectory-123"
|
||||
);
|
||||
assert_eq!(envelope["request"]["sessionId"], "session-ant-123");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn checkpoint_request_type_builds_checkpoint_envelope() {
|
||||
let request_body = json!({
|
||||
"contents": [
|
||||
{
|
||||
"role": "user",
|
||||
"parts": [
|
||||
{ "text": "checkpoint context" }
|
||||
]
|
||||
}
|
||||
],
|
||||
"generationConfig": {
|
||||
"maxOutputTokens": 8192,
|
||||
"thinkingConfig": {
|
||||
"includeThoughts": true,
|
||||
"thinkingBudget": 4000
|
||||
}
|
||||
},
|
||||
"toolConfig": {
|
||||
"functionCallingConfig": {
|
||||
"mode": "NONE"
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
let envelope = match build_antigravity_safe_v1internal_request(
|
||||
&sample_auth(),
|
||||
"request-ant-checkpoint-123",
|
||||
"gemini-3.5-flash-low",
|
||||
&request_body,
|
||||
AntigravityEnvelopeRequestType::Checkpoint,
|
||||
) {
|
||||
AntigravityRequestEnvelopeSupport::Supported(envelope) => envelope,
|
||||
AntigravityRequestEnvelopeSupport::Unsupported(reason) => {
|
||||
panic!("checkpoint envelope should be supported: {reason:?}")
|
||||
}
|
||||
};
|
||||
|
||||
assert_eq!(envelope["requestType"], "checkpoint");
|
||||
assert_eq!(
|
||||
envelope["request"]["toolConfig"]["functionCallingConfig"]["mode"],
|
||||
"NONE"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn existing_v1internal_envelope_is_not_double_wrapped() {
|
||||
let request_body = json!({
|
||||
"project": "client-side-project",
|
||||
"requestId": "client-request-id-123",
|
||||
"model": "gemini-3.5-flash-low",
|
||||
"userAgent": "antigravity",
|
||||
"requestType": "checkpoint",
|
||||
"request": {
|
||||
"contents": [
|
||||
{
|
||||
"role": "user",
|
||||
"parts": [
|
||||
{ "text": "checkpoint context" }
|
||||
]
|
||||
}
|
||||
],
|
||||
"generationConfig": {
|
||||
"thinkingConfig": {
|
||||
"includeThoughts": true
|
||||
}
|
||||
},
|
||||
"toolConfig": {
|
||||
"functionCallingConfig": {
|
||||
"mode": "NONE"
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
classify_antigravity_safe_request_body(&request_body),
|
||||
Ok(())
|
||||
);
|
||||
|
||||
let envelope = match build_antigravity_safe_v1internal_request(
|
||||
&sample_auth(),
|
||||
"trace-request-id-should-not-overwrite-client-id",
|
||||
"mapped-antigravity-model",
|
||||
&request_body,
|
||||
AntigravityEnvelopeRequestType::Agent,
|
||||
) {
|
||||
AntigravityRequestEnvelopeSupport::Supported(envelope) => envelope,
|
||||
AntigravityRequestEnvelopeSupport::Unsupported(reason) => {
|
||||
panic!("existing v1internal envelope should be supported: {reason:?}")
|
||||
}
|
||||
};
|
||||
|
||||
assert_eq!(envelope["project"], "project-ant-123");
|
||||
assert_eq!(envelope["requestId"], "client-request-id-123");
|
||||
assert_eq!(envelope["model"], "mapped-antigravity-model");
|
||||
assert_eq!(envelope["userAgent"], "antigravity");
|
||||
assert_eq!(envelope["requestType"], "checkpoint");
|
||||
assert!(envelope["request"].get("request").is_none());
|
||||
assert_eq!(
|
||||
envelope["request"]["contents"][0]["parts"][0]["text"],
|
||||
"checkpoint context"
|
||||
);
|
||||
assert_eq!(
|
||||
envelope["request"]["toolConfig"]["functionCallingConfig"]["mode"],
|
||||
"NONE"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use url::form_urlencoded;
|
||||
|
||||
pub const ANTIGRAVITY_V1INTERNAL_PATH_TEMPLATE: &str = "/v1internal:{action}";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum AntigravityRequestUrlAction {
|
||||
GenerateContent,
|
||||
StreamGenerateContent,
|
||||
}
|
||||
|
||||
impl AntigravityRequestUrlAction {
|
||||
fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::GenerateContent => "generateContent",
|
||||
Self::StreamGenerateContent => "streamGenerateContent",
|
||||
}
|
||||
}
|
||||
|
||||
fn is_stream(self) -> bool {
|
||||
matches!(self, Self::StreamGenerateContent)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_antigravity_v1internal_url(
|
||||
base_url: &str,
|
||||
action: AntigravityRequestUrlAction,
|
||||
query: Option<&BTreeMap<String, String>>,
|
||||
) -> Option<String> {
|
||||
let trimmed_base = base_url.trim();
|
||||
if trimmed_base.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let path = ANTIGRAVITY_V1INTERNAL_PATH_TEMPLATE.replace("{action}", action.as_str());
|
||||
let mut url = format!("{}{}", trimmed_base.trim_end_matches('/'), path);
|
||||
|
||||
let mut params = BTreeMap::new();
|
||||
if let Some(query) = query {
|
||||
for (key, value) in query {
|
||||
let key = key.trim();
|
||||
let value = value.trim();
|
||||
if key.is_empty()
|
||||
|| value.is_empty()
|
||||
|| key.eq_ignore_ascii_case("beta")
|
||||
|| key.eq_ignore_ascii_case("key")
|
||||
{
|
||||
continue;
|
||||
}
|
||||
params.insert(key.to_string(), value.to_string());
|
||||
}
|
||||
}
|
||||
if action.is_stream() {
|
||||
params
|
||||
.entry(String::from("alt"))
|
||||
.or_insert_with(|| String::from("sse"));
|
||||
}
|
||||
|
||||
if !params.is_empty() {
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in params {
|
||||
serializer.append_pair(key.as_str(), value.as_str());
|
||||
}
|
||||
let query_string = serializer.finish();
|
||||
if !query_string.is_empty() {
|
||||
url.push('?');
|
||||
url.push_str(&query_string);
|
||||
}
|
||||
}
|
||||
|
||||
Some(url)
|
||||
}
|
||||
@@ -0,0 +1,599 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::headers::{
|
||||
normalize_upstream_accept_encoding, should_skip_upstream_complete_passthrough_header,
|
||||
should_skip_upstream_passthrough_header,
|
||||
};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
const DEFAULT_ANTHROPIC_VERSION: &str = "2023-06-01";
|
||||
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
|
||||
|
||||
fn collect_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = BTreeMap::new();
|
||||
for (name, value) in headers.iter() {
|
||||
let Ok(value) = value.to_str() else {
|
||||
continue;
|
||||
};
|
||||
let key = name.as_str().to_ascii_lowercase();
|
||||
if should_skip_upstream_passthrough_header(&key) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = normalize_passthrough_header_value(&key, value) else {
|
||||
continue;
|
||||
};
|
||||
out.insert(key, value);
|
||||
}
|
||||
|
||||
for (key, value) in extra_headers {
|
||||
let normalized_key = key.to_ascii_lowercase();
|
||||
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
|
||||
continue;
|
||||
};
|
||||
out.insert(normalized_key, value);
|
||||
}
|
||||
|
||||
out
|
||||
}
|
||||
|
||||
fn collect_complete_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = BTreeMap::new();
|
||||
for (name, value) in headers.iter() {
|
||||
let Ok(value) = value.to_str() else {
|
||||
continue;
|
||||
};
|
||||
let key = name.as_str().to_ascii_lowercase();
|
||||
if should_skip_upstream_complete_passthrough_header(&key) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = normalize_passthrough_header_value(&key, value) else {
|
||||
continue;
|
||||
};
|
||||
out.insert(key, value);
|
||||
}
|
||||
|
||||
for (key, value) in extra_headers {
|
||||
let normalized_key = key.to_ascii_lowercase();
|
||||
let Some(value) = normalize_passthrough_header_value(&normalized_key, value) else {
|
||||
continue;
|
||||
};
|
||||
out.insert(normalized_key, value);
|
||||
}
|
||||
|
||||
out
|
||||
}
|
||||
|
||||
fn normalize_passthrough_header_value(key: &str, value: &str) -> Option<String> {
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
if key.eq_ignore_ascii_case("accept-encoding") {
|
||||
return normalize_upstream_accept_encoding(value);
|
||||
}
|
||||
|
||||
Some(value.to_string())
|
||||
}
|
||||
|
||||
pub fn build_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = collect_passthrough_headers(headers, extra_headers);
|
||||
out.entry("content-type".to_string()).or_insert_with(|| {
|
||||
content_type
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or("application/json")
|
||||
.trim()
|
||||
.to_string()
|
||||
});
|
||||
out.remove("content-length");
|
||||
out
|
||||
}
|
||||
|
||||
pub fn build_openai_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
auth_header: &str,
|
||||
auth_value: &str,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = build_passthrough_headers(headers, extra_headers, content_type);
|
||||
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
|
||||
out
|
||||
}
|
||||
|
||||
pub fn build_complete_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = collect_complete_passthrough_headers(headers, extra_headers);
|
||||
out.entry("content-type".to_string()).or_insert_with(|| {
|
||||
content_type
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or("application/json")
|
||||
.trim()
|
||||
.to_string()
|
||||
});
|
||||
out.remove("content-length");
|
||||
out
|
||||
}
|
||||
|
||||
pub fn build_complete_passthrough_headers_with_auth(
|
||||
headers: &http::HeaderMap,
|
||||
auth_header: &str,
|
||||
auth_value: &str,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = build_complete_passthrough_headers(headers, extra_headers, content_type);
|
||||
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
|
||||
out
|
||||
}
|
||||
|
||||
pub fn build_claude_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
auth_header: &str,
|
||||
auth_value: &str,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
content_type: Option<&str>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = build_openai_passthrough_headers(
|
||||
headers,
|
||||
auth_header,
|
||||
auth_value,
|
||||
extra_headers,
|
||||
content_type,
|
||||
);
|
||||
|
||||
for (name, value) in headers.iter() {
|
||||
let Ok(value) = value.to_str() else {
|
||||
continue;
|
||||
};
|
||||
let key = name.as_str().to_ascii_lowercase();
|
||||
let value = value.trim();
|
||||
if value.is_empty() || !should_restore_claude_passthrough_header(&key) {
|
||||
continue;
|
||||
}
|
||||
|
||||
if key == "anthropic-beta" {
|
||||
let merged = merge_comma_header_values(out.get(&key).map(String::as_str), Some(value));
|
||||
if let Some(merged) = merged {
|
||||
out.insert(key, merged);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
out.entry(key).or_insert_with(|| value.to_string());
|
||||
}
|
||||
|
||||
out.entry("anthropic-version".to_string())
|
||||
.or_insert_with(|| DEFAULT_ANTHROPIC_VERSION.to_string());
|
||||
out
|
||||
}
|
||||
|
||||
pub fn build_passthrough_headers_with_auth(
|
||||
headers: &http::HeaderMap,
|
||||
auth_header: &str,
|
||||
auth_value: &str,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = collect_passthrough_headers(headers, extra_headers);
|
||||
ensure_upstream_auth_header(&mut out, auth_header, auth_value);
|
||||
out.remove("content-length");
|
||||
out
|
||||
}
|
||||
|
||||
pub fn ensure_upstream_auth_header(
|
||||
headers: &mut BTreeMap<String, String>,
|
||||
auth_header: &str,
|
||||
auth_value: &str,
|
||||
) {
|
||||
let header_name = auth_header.trim().to_ascii_lowercase();
|
||||
let header_value = auth_value.trim();
|
||||
if header_name.is_empty() || header_value.is_empty() {
|
||||
return;
|
||||
}
|
||||
|
||||
if headers
|
||||
.get(&header_name)
|
||||
.map(|value| value.trim().is_empty())
|
||||
.unwrap_or(true)
|
||||
{
|
||||
headers.insert(header_name, header_value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
fn should_restore_claude_passthrough_header(name: &str) -> bool {
|
||||
name.starts_with("anthropic-") || name.starts_with("x-stainless-") || name == "x-app"
|
||||
}
|
||||
|
||||
fn merge_comma_header_values(left: Option<&str>, right: Option<&str>) -> Option<String> {
|
||||
let mut merged = Vec::new();
|
||||
|
||||
for raw in [left, right].into_iter().flatten() {
|
||||
for token in raw.split(',') {
|
||||
let token = token.trim();
|
||||
if token.is_empty() || merged.iter().any(|existing: &String| existing == token) {
|
||||
continue;
|
||||
}
|
||||
merged.push(token.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
if merged.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(merged.join(","))
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_local_openai_bearer_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<(String, String)> {
|
||||
let auth_type = resolve_local_auth_type_for_transport_format(transport);
|
||||
if !matches!(auth_type.as_str(), "api_key" | "bearer") {
|
||||
return None;
|
||||
}
|
||||
let secret = resolved_local_secret(transport)?;
|
||||
|
||||
Some(("authorization".to_string(), bearer_auth_value(secret)))
|
||||
}
|
||||
|
||||
pub fn resolve_local_standard_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<(String, String)> {
|
||||
let auth_type = resolve_local_auth_type_for_transport_format(transport);
|
||||
let secret = resolved_local_secret(transport)?;
|
||||
|
||||
match auth_type.as_str() {
|
||||
"api_key" => Some(("x-api-key".to_string(), secret.to_string())),
|
||||
"bearer" => Some(("authorization".to_string(), bearer_auth_value(secret))),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_local_gemini_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<(String, String)> {
|
||||
let auth_type = resolve_local_auth_type_for_transport_format(transport);
|
||||
let secret = resolved_local_secret(transport)?;
|
||||
|
||||
match auth_type.as_str() {
|
||||
"api_key" => Some(("x-goog-api-key".to_string(), secret.to_string())),
|
||||
"bearer" => Some(("authorization".to_string(), bearer_auth_value(secret))),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_local_auth_type_for_transport_format(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> String {
|
||||
let default_auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
|
||||
let api_format = aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
|
||||
let Some(overrides) = transport
|
||||
.key
|
||||
.auth_type_by_format
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
else {
|
||||
return default_auth_type;
|
||||
};
|
||||
|
||||
overrides
|
||||
.get(&api_format)
|
||||
.or_else(|| overrides.get(transport.endpoint.api_format.trim()))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.map(str::to_ascii_lowercase)
|
||||
.filter(|value| matches!(value.as_str(), "api_key" | "bearer"))
|
||||
.unwrap_or(default_auth_type)
|
||||
}
|
||||
|
||||
fn resolved_local_secret(transport: &GatewayProviderTransportSnapshot) -> Option<&str> {
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if !secret.is_empty() && secret != PLACEHOLDER_API_KEY {
|
||||
Some(secret)
|
||||
} else if transport.key.decrypted_auth_config.is_some() {
|
||||
None
|
||||
} else {
|
||||
Some("")
|
||||
}
|
||||
}
|
||||
|
||||
fn bearer_auth_value(secret: &str) -> String {
|
||||
if secret.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!("Bearer {secret}")
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_claude_passthrough_headers, build_complete_passthrough_headers_with_auth,
|
||||
build_openai_passthrough_headers, resolve_local_openai_bearer_auth,
|
||||
resolve_local_standard_auth,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "provider".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "claude:messages".to_string(),
|
||||
api_family: Some("claude".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://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: "bearer".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_passthrough_headers_restore_stripped_anthropic_headers() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert(
|
||||
"anthropic-beta",
|
||||
http::HeaderValue::from_static("prompt-caching-2024-07-31,context-1m-2025-08-07"),
|
||||
);
|
||||
headers.insert(
|
||||
"x-stainless-runtime-version",
|
||||
http::HeaderValue::from_static("v22.14.0"),
|
||||
);
|
||||
headers.insert("x-app", http::HeaderValue::from_static("cli"));
|
||||
|
||||
let built = build_claude_passthrough_headers(
|
||||
&headers,
|
||||
"x-api-key",
|
||||
"sk-upstream-claude",
|
||||
&BTreeMap::from([("anthropic-beta".to_string(), "custom-beta".to_string())]),
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
built.get("anthropic-version").map(String::as_str),
|
||||
Some("2023-06-01")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("anthropic-beta").map(String::as_str),
|
||||
Some("custom-beta,prompt-caching-2024-07-31,context-1m-2025-08-07")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("x-stainless-runtime-version").map(String::as_str),
|
||||
Some("v22.14.0")
|
||||
);
|
||||
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
|
||||
assert_eq!(
|
||||
built.get("x-api-key").map(String::as_str),
|
||||
Some("sk-upstream-claude")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn passthrough_headers_preserve_supported_response_compression() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert(
|
||||
http::header::ACCEPT_ENCODING,
|
||||
http::HeaderValue::from_static("gzip, br"),
|
||||
);
|
||||
|
||||
let built = build_openai_passthrough_headers(
|
||||
&headers,
|
||||
"authorization",
|
||||
"Bearer upstream",
|
||||
&BTreeMap::new(),
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
built.get("accept-encoding").map(String::as_str),
|
||||
Some("gzip")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_passthrough_headers_preserve_explicit_anthropic_version_override() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert(
|
||||
"anthropic-version",
|
||||
http::HeaderValue::from_static("2024-01-01"),
|
||||
);
|
||||
|
||||
let built = build_claude_passthrough_headers(
|
||||
&headers,
|
||||
"authorization",
|
||||
"Bearer upstream-token",
|
||||
&BTreeMap::new(),
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
built.get("anthropic-version").map(String::as_str),
|
||||
Some("2024-01-01")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn complete_passthrough_headers_preserve_business_headers() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert(
|
||||
"anthropic-beta",
|
||||
http::HeaderValue::from_static("prompt-caching-2024-07-31"),
|
||||
);
|
||||
headers.insert(
|
||||
"x-stainless-runtime-version",
|
||||
http::HeaderValue::from_static("v24.0.0"),
|
||||
);
|
||||
headers.insert("x-app", http::HeaderValue::from_static("cli"));
|
||||
headers.insert(
|
||||
"authorization",
|
||||
http::HeaderValue::from_static("Bearer client-token"),
|
||||
);
|
||||
|
||||
let built = build_complete_passthrough_headers_with_auth(
|
||||
&headers,
|
||||
"x-api-key",
|
||||
"sk-upstream",
|
||||
&BTreeMap::new(),
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
built.get("anthropic-beta").map(String::as_str),
|
||||
Some("prompt-caching-2024-07-31")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("x-stainless-runtime-version").map(String::as_str),
|
||||
Some("v24.0.0")
|
||||
);
|
||||
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
|
||||
assert_eq!(built.get("authorization"), None);
|
||||
assert_eq!(
|
||||
built.get("x-api-key").map(String::as_str),
|
||||
Some("sk-upstream")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_standard_auth_keeps_header_shape_for_placeholder_secret() {
|
||||
assert_eq!(
|
||||
resolve_local_standard_auth(&sample_transport()),
|
||||
Some(("authorization".to_string(), String::new()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_standard_auth_keeps_header_shape_for_empty_secret() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "api_key".to_string();
|
||||
transport.key.decrypted_api_key = String::new();
|
||||
|
||||
assert_eq!(
|
||||
resolve_local_standard_auth(&transport),
|
||||
Some(("x-api-key".to_string(), String::new()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_standard_auth_defers_to_auth_config_when_raw_secret_is_empty() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_auth_config =
|
||||
Some(r#"{"access_token":"cached-token"}"#.to_string());
|
||||
|
||||
assert!(resolve_local_standard_auth(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_openai_bearer_auth_maps_api_key_to_bearer_authorization() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "api_key".to_string();
|
||||
transport.key.decrypted_api_key = "sk-openai".to_string();
|
||||
|
||||
assert_eq!(
|
||||
resolve_local_openai_bearer_auth(&transport),
|
||||
Some(("authorization".to_string(), "Bearer sk-openai".to_string(),))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_openai_bearer_auth_preserves_bearer_header_shape() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "bearer".to_string();
|
||||
transport.key.decrypted_api_key = "sk-openai".to_string();
|
||||
|
||||
assert_eq!(
|
||||
resolve_local_openai_bearer_auth(&transport),
|
||||
Some(("authorization".to_string(), "Bearer sk-openai".to_string(),))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_standard_auth_uses_format_auth_type_override() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "api_key".to_string();
|
||||
transport.key.auth_type_by_format = Some(serde_json::json!({
|
||||
"claude:messages": "bearer"
|
||||
}));
|
||||
transport.key.decrypted_api_key = "sk-claude".to_string();
|
||||
|
||||
assert_eq!(
|
||||
resolve_local_standard_auth(&transport),
|
||||
Some(("authorization".to_string(), "Bearer sk-claude".to_string(),))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_gemini_auth_falls_back_to_default_when_other_format_is_overridden() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "gemini:generate_content".to_string();
|
||||
transport.key.auth_type = "api_key".to_string();
|
||||
transport.key.auth_type_by_format = Some(serde_json::json!({
|
||||
"claude:messages": "bearer"
|
||||
}));
|
||||
transport.key.decrypted_api_key = "sk-gemini".to_string();
|
||||
|
||||
assert_eq!(
|
||||
super::resolve_local_gemini_auth(&transport),
|
||||
Some(("x-goog-api-key".to_string(), "sk-gemini".to_string(),))
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,784 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::Value;
|
||||
use url::form_urlencoded;
|
||||
|
||||
const UNSAFE_AUTH_CONFIG_HEADER_NAMES: &[&str] = &["content-length", "host", "proxy-authorization"];
|
||||
const RUNTIME_ONLY_AUTH_CONFIG_HEADER_NAMES: &[&str] = &[
|
||||
"api-key",
|
||||
"authorization",
|
||||
"content-type",
|
||||
"cookie",
|
||||
"x-api-key",
|
||||
"x-goog-api-key",
|
||||
];
|
||||
const UNSAFE_AUTH_CONFIG_QUERY_NAMES: &[&str] = &[
|
||||
"access_token",
|
||||
"api_key",
|
||||
"apikey",
|
||||
"authorization",
|
||||
"key",
|
||||
"token",
|
||||
];
|
||||
const SENSITIVE_AUTH_CONFIG_KEYS: &[&str] = &[
|
||||
"access_token",
|
||||
"api_key",
|
||||
"apikey",
|
||||
"authorization",
|
||||
"client_email",
|
||||
"client_id",
|
||||
"client_secret",
|
||||
"expires_at",
|
||||
"id_token",
|
||||
"key",
|
||||
"private_key",
|
||||
"refresh_token",
|
||||
"service_account",
|
||||
"token",
|
||||
"token_uri",
|
||||
];
|
||||
const IGNORABLE_AUTH_CONFIG_METADATA_KEYS: &[&str] = &[
|
||||
"account_id",
|
||||
"account_name",
|
||||
"account_user_id",
|
||||
"auth_method",
|
||||
"access_token_import_temporary",
|
||||
"email",
|
||||
"expires_at",
|
||||
"is_fedramp",
|
||||
"model_regions",
|
||||
"organizations",
|
||||
"plan_type",
|
||||
"project_id",
|
||||
"provider_type",
|
||||
"refresh_token_import_error",
|
||||
"region",
|
||||
"scope",
|
||||
"tier",
|
||||
"token_type",
|
||||
"updated_at",
|
||||
"user_id",
|
||||
"workspace_id",
|
||||
"workspace_name",
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct LocalAuthConfigSafeSubset {
|
||||
pub headers: BTreeMap<String, String>,
|
||||
pub query: BTreeMap<String, String>,
|
||||
pub path: Option<String>,
|
||||
}
|
||||
|
||||
impl LocalAuthConfigSafeSubset {
|
||||
fn is_empty(&self) -> bool {
|
||||
self.headers.is_empty() && self.query.is_empty() && self.path.is_none()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum LocalAuthConfigAbsorption {
|
||||
Missing,
|
||||
Unsupported,
|
||||
Absorbed {
|
||||
base_url: String,
|
||||
header_rules: Option<Value>,
|
||||
custom_path: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
pub fn apply_local_auth_config_header_overrides(
|
||||
headers: &mut BTreeMap<String, String>,
|
||||
raw_auth_config: Option<&str>,
|
||||
) {
|
||||
let Some(raw_auth_config) = raw_auth_config
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return;
|
||||
};
|
||||
let Ok(parsed) = serde_json::from_str::<Value>(raw_auth_config) else {
|
||||
return;
|
||||
};
|
||||
let Some(object) = parsed.as_object() else {
|
||||
return;
|
||||
};
|
||||
|
||||
let mut overrides = BTreeMap::new();
|
||||
collect_auth_config_header_overrides(object, &mut overrides);
|
||||
for (key, value) in overrides {
|
||||
headers.insert(key, value);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn absorb_local_auth_config_safe_subset(
|
||||
base_url: &str,
|
||||
header_rules: Option<Value>,
|
||||
custom_path: Option<String>,
|
||||
raw_auth_config: Option<&str>,
|
||||
) -> LocalAuthConfigAbsorption {
|
||||
let Some(raw_auth_config) = raw_auth_config
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return LocalAuthConfigAbsorption::Missing;
|
||||
};
|
||||
|
||||
let subset = match parse_local_auth_config_safe_subset(raw_auth_config) {
|
||||
Ok(subset) => subset,
|
||||
Err(()) => return LocalAuthConfigAbsorption::Unsupported,
|
||||
};
|
||||
if subset.is_empty() {
|
||||
return LocalAuthConfigAbsorption::Unsupported;
|
||||
}
|
||||
let header_rules = match merge_auth_config_header_rules(header_rules, &subset.headers) {
|
||||
Some(rules) => rules,
|
||||
None => return LocalAuthConfigAbsorption::Unsupported,
|
||||
};
|
||||
let base_url = match merge_auth_config_base_url(base_url, &subset.query) {
|
||||
Some(value) => value,
|
||||
None => return LocalAuthConfigAbsorption::Unsupported,
|
||||
};
|
||||
let custom_path = match merge_auth_config_custom_path(custom_path, subset.path) {
|
||||
Some(path) => path,
|
||||
None => return LocalAuthConfigAbsorption::Unsupported,
|
||||
};
|
||||
|
||||
LocalAuthConfigAbsorption::Absorbed {
|
||||
base_url,
|
||||
header_rules,
|
||||
custom_path,
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_local_auth_config_safe_subset(raw: &str) -> Result<LocalAuthConfigSafeSubset, ()> {
|
||||
let parsed: Value = serde_json::from_str(raw).map_err(|_| ())?;
|
||||
let object = parsed.as_object().ok_or(())?;
|
||||
|
||||
let mut headers = BTreeMap::new();
|
||||
let mut query = BTreeMap::new();
|
||||
let mut path = None;
|
||||
|
||||
parse_local_auth_config_object(object, &mut headers, &mut query, &mut path, true)?;
|
||||
if headers
|
||||
.keys()
|
||||
.any(|key| RUNTIME_ONLY_AUTH_CONFIG_HEADER_NAMES.contains(&key.as_str()))
|
||||
{
|
||||
return Err(());
|
||||
}
|
||||
|
||||
Ok(LocalAuthConfigSafeSubset {
|
||||
headers,
|
||||
query,
|
||||
path,
|
||||
})
|
||||
}
|
||||
|
||||
fn collect_auth_config_header_overrides(
|
||||
object: &serde_json::Map<String, Value>,
|
||||
out: &mut BTreeMap<String, String>,
|
||||
) {
|
||||
for (key, value) in object {
|
||||
let normalized = key.trim().to_ascii_lowercase();
|
||||
match normalized.as_str() {
|
||||
"headers" | "extra_headers" | "extraheaders" => {
|
||||
merge_header_string_map_lenient(out, value);
|
||||
}
|
||||
"transport" | "request" => {
|
||||
if let Some(nested) = value.as_object() {
|
||||
collect_auth_config_header_overrides(nested, out);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_local_auth_config_object(
|
||||
object: &serde_json::Map<String, Value>,
|
||||
headers: &mut BTreeMap<String, String>,
|
||||
query: &mut BTreeMap<String, String>,
|
||||
path: &mut Option<String>,
|
||||
allow_metadata: bool,
|
||||
) -> Result<(), ()> {
|
||||
for (key, value) in object {
|
||||
let normalized = key.trim().to_ascii_lowercase();
|
||||
match normalized.as_str() {
|
||||
"headers" | "extra_headers" | "extraheaders" => {
|
||||
merge_header_string_map(headers, value)?
|
||||
}
|
||||
"query" | "query_params" | "queryparams" => {
|
||||
merge_string_map(query, value, normalize_auth_config_query_key)?
|
||||
}
|
||||
"path" | "custom_path" => {
|
||||
let value = value.as_str().ok_or(())?;
|
||||
let normalized = normalize_auth_config_path(value).ok_or(())?;
|
||||
*path = Some(normalized);
|
||||
}
|
||||
"custompath" => {
|
||||
let value = value.as_str().ok_or(())?;
|
||||
let normalized = normalize_auth_config_path(value).ok_or(())?;
|
||||
*path = Some(normalized);
|
||||
}
|
||||
"transport" | "request" => {
|
||||
let nested = value.as_object().ok_or(())?;
|
||||
parse_local_auth_config_object(nested, headers, query, path, false)?;
|
||||
}
|
||||
_ if allow_metadata && is_ignorable_auth_config_metadata_key(&normalized) => {}
|
||||
_ if allow_metadata && is_sensitive_auth_config_key(&normalized) => return Err(()),
|
||||
_ => return Err(()),
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn merge_string_map(
|
||||
out: &mut BTreeMap<String, String>,
|
||||
value: &Value,
|
||||
normalize_key: fn(&str) -> Option<String>,
|
||||
) -> Result<(), ()> {
|
||||
let object = value.as_object().ok_or(())?;
|
||||
for (raw_key, raw_value) in object {
|
||||
let key = normalize_key(raw_key).ok_or(())?;
|
||||
let value = parse_static_auth_config_value(raw_value).ok_or(())?;
|
||||
out.insert(key, value);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn parse_static_auth_config_value(value: &Value) -> Option<String> {
|
||||
match value {
|
||||
Value::String(raw) => {
|
||||
let normalized = raw.trim();
|
||||
if normalized.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(normalized.to_string())
|
||||
}
|
||||
}
|
||||
Value::Number(raw) => Some(raw.to_string()),
|
||||
Value::Bool(raw) => Some(raw.to_string()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_header_string_map(out: &mut BTreeMap<String, String>, value: &Value) -> Result<(), ()> {
|
||||
let object = value.as_object().ok_or(())?;
|
||||
for (raw_key, raw_value) in object {
|
||||
let key = normalize_auth_config_header_name(raw_key).ok_or(())?;
|
||||
let value = parse_static_auth_config_header_value(raw_value).ok_or(())?;
|
||||
out.insert(key, value);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn merge_header_string_map_lenient(out: &mut BTreeMap<String, String>, value: &Value) {
|
||||
let Some(object) = value.as_object() else {
|
||||
return;
|
||||
};
|
||||
for (raw_key, raw_value) in object {
|
||||
let Some(key) = normalize_auth_config_header_name(raw_key) else {
|
||||
continue;
|
||||
};
|
||||
let Some(value) = parse_static_auth_config_header_value(raw_value) else {
|
||||
continue;
|
||||
};
|
||||
out.insert(key, value);
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_static_auth_config_header_value(value: &Value) -> Option<String> {
|
||||
let value = parse_static_auth_config_value(value)?;
|
||||
http::header::HeaderValue::from_str(&value)
|
||||
.is_ok()
|
||||
.then_some(value)
|
||||
}
|
||||
|
||||
fn normalize_auth_config_header_name(raw: &str) -> Option<String> {
|
||||
let value = raw.trim().to_ascii_lowercase();
|
||||
if value.is_empty()
|
||||
|| value.chars().any(|char| char.is_ascii_control())
|
||||
|| UNSAFE_AUTH_CONFIG_HEADER_NAMES.contains(&value.as_str())
|
||||
{
|
||||
return None;
|
||||
}
|
||||
http::header::HeaderName::from_bytes(value.as_bytes())
|
||||
.ok()
|
||||
.map(|name| name.as_str().to_string())
|
||||
}
|
||||
|
||||
fn normalize_auth_config_query_key(raw: &str) -> Option<String> {
|
||||
let value = raw.trim();
|
||||
if value.is_empty()
|
||||
|| value.chars().any(|char| matches!(char, '&' | '=' | '#'))
|
||||
|| value.chars().any(|char| char.is_ascii_control())
|
||||
|| UNSAFE_AUTH_CONFIG_QUERY_NAMES
|
||||
.iter()
|
||||
.any(|blocked| value.eq_ignore_ascii_case(blocked))
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(value.to_string())
|
||||
}
|
||||
|
||||
fn is_sensitive_auth_config_key(key: &str) -> bool {
|
||||
SENSITIVE_AUTH_CONFIG_KEYS
|
||||
.iter()
|
||||
.any(|blocked| key.eq_ignore_ascii_case(blocked))
|
||||
}
|
||||
|
||||
fn is_ignorable_auth_config_metadata_key(key: &str) -> bool {
|
||||
IGNORABLE_AUTH_CONFIG_METADATA_KEYS
|
||||
.iter()
|
||||
.any(|allowed| key.eq_ignore_ascii_case(allowed))
|
||||
}
|
||||
|
||||
fn normalize_auth_config_path(raw: &str) -> Option<String> {
|
||||
let value = raw.trim();
|
||||
if value.is_empty()
|
||||
|| !value.starts_with('/')
|
||||
|| value.contains("://")
|
||||
|| value
|
||||
.chars()
|
||||
.any(|char| matches!(char, '{' | '}' | '$' | '#'))
|
||||
|| value.chars().any(|char| char.is_ascii_control())
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(value.to_string())
|
||||
}
|
||||
|
||||
fn merge_auth_config_header_rules(
|
||||
existing_rules: Option<Value>,
|
||||
headers: &BTreeMap<String, String>,
|
||||
) -> Option<Option<Value>> {
|
||||
if headers.is_empty() {
|
||||
return Some(existing_rules);
|
||||
}
|
||||
|
||||
let mut merged = match existing_rules {
|
||||
Some(Value::Array(items)) => items,
|
||||
Some(_) => return None,
|
||||
None => Vec::new(),
|
||||
};
|
||||
for (key, value) in headers {
|
||||
merged.push(serde_json::json!({
|
||||
"action": "set",
|
||||
"key": key,
|
||||
"value": value,
|
||||
}));
|
||||
}
|
||||
Some(Some(Value::Array(merged)))
|
||||
}
|
||||
|
||||
fn merge_auth_config_custom_path(
|
||||
existing_custom_path: Option<String>,
|
||||
path_override: Option<String>,
|
||||
) -> Option<Option<String>> {
|
||||
let base_path = path_override.or(existing_custom_path);
|
||||
let Some(base_path) = base_path else {
|
||||
return Some(None);
|
||||
};
|
||||
|
||||
let (path_only, query) = split_path_and_query(&base_path)?;
|
||||
if query.is_empty() {
|
||||
return Some(Some(path_only));
|
||||
}
|
||||
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in query {
|
||||
serializer.append_pair(&key, &value);
|
||||
}
|
||||
Some(Some(format!("{path_only}?{}", serializer.finish())))
|
||||
}
|
||||
|
||||
fn merge_auth_config_base_url(base_url: &str, query: &BTreeMap<String, String>) -> Option<String> {
|
||||
if query.is_empty() {
|
||||
return Some(base_url.to_string());
|
||||
}
|
||||
|
||||
let raw_base_url = base_url.trim();
|
||||
let had_implicit_root = raw_base_url
|
||||
.split_once("://")
|
||||
.map(|(_, rest)| {
|
||||
let authority = rest.split_once('?').map(|(head, _)| head).unwrap_or(rest);
|
||||
!authority.contains('/')
|
||||
})
|
||||
.unwrap_or(false);
|
||||
|
||||
let mut url = url::Url::parse(raw_base_url).ok()?;
|
||||
let mut merged = BTreeMap::new();
|
||||
for (key, value) in url.query_pairs() {
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
merged.insert(key.into_owned(), value.to_string());
|
||||
}
|
||||
for (key, value) in query {
|
||||
merged.insert(key.clone(), value.clone());
|
||||
}
|
||||
|
||||
if merged.is_empty() {
|
||||
url.set_query(None);
|
||||
return Some(url.to_string());
|
||||
}
|
||||
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in merged {
|
||||
serializer.append_pair(&key, &value);
|
||||
}
|
||||
url.set_query(Some(&serializer.finish()));
|
||||
|
||||
let mut normalized = url.to_string();
|
||||
if had_implicit_root {
|
||||
normalized = normalized.replacen("/?", "?", 1);
|
||||
}
|
||||
Some(normalized)
|
||||
}
|
||||
|
||||
fn split_path_and_query(path: &str) -> Option<(String, BTreeMap<String, String>)> {
|
||||
let normalized = normalize_auth_config_path(path)?;
|
||||
let (path_only, query_part) = if let Some((path, query)) = normalized.split_once('?') {
|
||||
(path.to_string(), Some(query))
|
||||
} else {
|
||||
(normalized, None)
|
||||
};
|
||||
if path_only.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut query = BTreeMap::new();
|
||||
if let Some(query_part) = query_part.filter(|value| !value.trim().is_empty()) {
|
||||
for (key, value) in form_urlencoded::parse(query_part.as_bytes()) {
|
||||
let key = normalize_auth_config_query_key(key.as_ref())?;
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
query.insert(key, value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
Some((path_only, query))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
absorb_local_auth_config_safe_subset, apply_local_auth_config_header_overrides,
|
||||
LocalAuthConfigAbsorption,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn absorbs_static_headers_and_query_into_existing_transport_fields() {
|
||||
let result = absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
Some(json!([{"action":"set","key":"x-base","value":"1"}])),
|
||||
None,
|
||||
Some(
|
||||
r#"{
|
||||
"headers": {"x-account-id": "acc-1"},
|
||||
"query": {"tenant": "demo"}
|
||||
}"#,
|
||||
),
|
||||
);
|
||||
|
||||
let LocalAuthConfigAbsorption::Absorbed {
|
||||
base_url,
|
||||
header_rules,
|
||||
custom_path,
|
||||
} = result
|
||||
else {
|
||||
panic!("auth_config should be absorbed");
|
||||
};
|
||||
|
||||
assert_eq!(base_url, "https://api.openai.example/v1?tenant=demo");
|
||||
assert_eq!(
|
||||
header_rules,
|
||||
Some(json!([
|
||||
{"action":"set","key":"x-base","value":"1"},
|
||||
{"action":"set","key":"x-account-id","value":"acc-1"}
|
||||
]))
|
||||
);
|
||||
assert_eq!(custom_path, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn absorbs_path_override_and_query_aliases() {
|
||||
let result = absorb_local_auth_config_safe_subset(
|
||||
"https://generativelanguage.googleapis.com/v1beta",
|
||||
None,
|
||||
Some("/v1beta/models/original:generateContent".to_string()),
|
||||
Some(
|
||||
r#"{
|
||||
"extra_headers": {"x-tenant": "demo"},
|
||||
"query_params": {"alt": "sse"},
|
||||
"custom_path": "/v1beta/models/gemini-2.5-pro:streamGenerateContent"
|
||||
}"#,
|
||||
),
|
||||
);
|
||||
|
||||
let LocalAuthConfigAbsorption::Absorbed {
|
||||
base_url,
|
||||
header_rules,
|
||||
custom_path,
|
||||
} = result
|
||||
else {
|
||||
panic!("auth_config should be absorbed");
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
base_url,
|
||||
"https://generativelanguage.googleapis.com/v1beta?alt=sse"
|
||||
);
|
||||
assert_eq!(
|
||||
header_rules,
|
||||
Some(json!([{"action":"set","key":"x-tenant","value":"demo"}]))
|
||||
);
|
||||
assert_eq!(
|
||||
custom_path.as_deref(),
|
||||
Some("/v1beta/models/gemini-2.5-pro:streamGenerateContent")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unknown_keys_and_reserved_headers() {
|
||||
assert_eq!(
|
||||
absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
None,
|
||||
Some(r#"{"provider_type":"custom"}"#),
|
||||
),
|
||||
LocalAuthConfigAbsorption::Unsupported
|
||||
);
|
||||
assert_eq!(
|
||||
absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
None,
|
||||
Some(r#"{"headers":{"host":"api.example.test"}}"#),
|
||||
),
|
||||
LocalAuthConfigAbsorption::Unsupported
|
||||
);
|
||||
assert_eq!(
|
||||
absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
None,
|
||||
Some(r#"{"query":{"key":"secret"}}"#),
|
||||
),
|
||||
LocalAuthConfigAbsorption::Unsupported
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn applies_header_overrides_even_when_auth_config_has_refresh_token() {
|
||||
let mut headers = std::collections::BTreeMap::from([
|
||||
(
|
||||
"authorization".to_string(),
|
||||
"Bearer direct-token".to_string(),
|
||||
),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
]);
|
||||
|
||||
apply_local_auth_config_header_overrides(
|
||||
&mut headers,
|
||||
Some(
|
||||
r#"{
|
||||
"refresh_token": "rt-1",
|
||||
"headers": {
|
||||
"authorization": "Bearer imported-session",
|
||||
"content-type": "text/plain",
|
||||
"host": "blocked.example"
|
||||
}
|
||||
}"#,
|
||||
),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
headers.get("authorization"),
|
||||
Some(&"Bearer imported-session".to_string())
|
||||
);
|
||||
assert_eq!(headers.get("content-type"), Some(&"text/plain".to_string()));
|
||||
assert!(!headers.contains_key("host"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ignores_invalid_auth_config_header_values_when_applying_overrides() {
|
||||
let mut headers = std::collections::BTreeMap::new();
|
||||
|
||||
apply_local_auth_config_header_overrides(
|
||||
&mut headers,
|
||||
Some(r#"{"headers":{"authorization":"Bearer ok","x-bad":"line\nbreak"}}"#),
|
||||
);
|
||||
|
||||
assert_eq!(headers.get("authorization"), Some(&"Bearer ok".to_string()));
|
||||
assert!(!headers.contains_key("x-bad"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_imported_authorization_headers_for_runtime_override() {
|
||||
assert_eq!(
|
||||
absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
None,
|
||||
Some(
|
||||
r#"{
|
||||
"provider_type": "codex",
|
||||
"access_token_import_temporary": true,
|
||||
"headers": {
|
||||
"authorization": "Bearer imported-session",
|
||||
"chatgpt-account-id": "acct-1"
|
||||
}
|
||||
}"#,
|
||||
),
|
||||
),
|
||||
LocalAuthConfigAbsorption::Unsupported
|
||||
);
|
||||
|
||||
let mut headers = std::collections::BTreeMap::from([(
|
||||
"authorization".to_string(),
|
||||
"Bearer direct-token".to_string(),
|
||||
)]);
|
||||
apply_local_auth_config_header_overrides(
|
||||
&mut headers,
|
||||
Some(
|
||||
r#"{
|
||||
"provider_type": "codex",
|
||||
"access_token_import_temporary": true,
|
||||
"headers": {
|
||||
"authorization": "Bearer imported-session",
|
||||
"chatgpt-account-id": "acct-1"
|
||||
}
|
||||
}"#,
|
||||
),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
headers.get("authorization"),
|
||||
Some(&"Bearer imported-session".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("chatgpt-account-id"),
|
||||
Some(&"acct-1".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn absorbs_query_only_configs_into_base_url_for_dynamic_path_formats() {
|
||||
let result = absorb_local_auth_config_safe_subset(
|
||||
"https://generativelanguage.googleapis.com/v1beta",
|
||||
None,
|
||||
None,
|
||||
Some(r#"{"query":{"alt":"sse"}}"#),
|
||||
);
|
||||
let LocalAuthConfigAbsorption::Absorbed {
|
||||
base_url,
|
||||
header_rules,
|
||||
custom_path,
|
||||
} = result
|
||||
else {
|
||||
panic!("query-only auth_config should be absorbed");
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
base_url,
|
||||
"https://generativelanguage.googleapis.com/v1beta?alt=sse"
|
||||
);
|
||||
assert_eq!(header_rules, None);
|
||||
assert_eq!(custom_path, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn absorbs_camel_case_transport_keys_with_ignorable_metadata() {
|
||||
let result = absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
None,
|
||||
Some(
|
||||
r#"{
|
||||
"email": "[email protected]",
|
||||
"plan_type": "plus",
|
||||
"request": {
|
||||
"extraHeaders": {"x-org-id": "org-1"},
|
||||
"queryParams": {"tenant": "demo", "retry": 2, "stream": true},
|
||||
"customPath": "/v1/responses"
|
||||
}
|
||||
}"#,
|
||||
),
|
||||
);
|
||||
let LocalAuthConfigAbsorption::Absorbed {
|
||||
base_url,
|
||||
header_rules,
|
||||
custom_path,
|
||||
} = result
|
||||
else {
|
||||
panic!("camelCase auth_config should be absorbed");
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
base_url,
|
||||
"https://api.openai.example/v1?retry=2&stream=true&tenant=demo"
|
||||
);
|
||||
assert_eq!(
|
||||
header_rules,
|
||||
Some(json!([{"action":"set","key":"x-org-id","value":"org-1"}]))
|
||||
);
|
||||
assert_eq!(custom_path.as_deref(), Some("/v1/responses"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_sensitive_oauth_fields_even_with_transport_subset() {
|
||||
assert_eq!(
|
||||
absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
None,
|
||||
Some(
|
||||
r#"{
|
||||
"headers": {"x-org-id": "org-1"},
|
||||
"refresh_token": "rt-1"
|
||||
}"#,
|
||||
),
|
||||
),
|
||||
LocalAuthConfigAbsorption::Unsupported
|
||||
);
|
||||
assert_eq!(
|
||||
absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
None,
|
||||
Some(
|
||||
r#"{
|
||||
"query": {"tenant": "demo"},
|
||||
"access_token": "at-1"
|
||||
}"#,
|
||||
),
|
||||
),
|
||||
LocalAuthConfigAbsorption::Unsupported
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_metadata_only_auth_config_without_transport_subset() {
|
||||
assert_eq!(
|
||||
absorb_local_auth_config_safe_subset(
|
||||
"https://api.openai.example/v1",
|
||||
None,
|
||||
None,
|
||||
Some(
|
||||
r#"{
|
||||
"email": "[email protected]",
|
||||
"plan_type": "plus",
|
||||
"workspace_name": "demo"
|
||||
}"#,
|
||||
),
|
||||
),
|
||||
LocalAuthConfigAbsorption::Unsupported
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
||||
pub struct ProviderTransportSnapshotCacheKey {
|
||||
provider_id: String,
|
||||
endpoint_id: String,
|
||||
key_id: String,
|
||||
}
|
||||
|
||||
impl ProviderTransportSnapshotCacheKey {
|
||||
pub fn new(provider_id: &str, endpoint_id: &str, key_id: &str) -> Option<Self> {
|
||||
let provider_id = provider_id.trim();
|
||||
let endpoint_id = endpoint_id.trim();
|
||||
let key_id = key_id.trim();
|
||||
if provider_id.is_empty() || endpoint_id.is_empty() || key_id.is_empty() {
|
||||
return None;
|
||||
}
|
||||
Some(Self {
|
||||
provider_id: provider_id.to_string(),
|
||||
endpoint_id: endpoint_id.to_string(),
|
||||
key_id: key_id.to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn provider_transport_snapshot_looks_refreshed(
|
||||
current: &GatewayProviderTransportSnapshot,
|
||||
refreshed: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
current.key.decrypted_api_key != refreshed.key.decrypted_api_key
|
||||
|| current.key.decrypted_auth_config != refreshed.key.decrypted_auth_config
|
||||
|| current.key.expires_at_unix_secs != refreshed.key.expires_at_unix_secs
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{provider_transport_snapshot_looks_refreshed, ProviderTransportSnapshotCacheKey};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_snapshot() -> 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: 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".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: 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,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: Some(1),
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "sk-test".to_string(),
|
||||
decrypted_auth_config: Some("{\"token\":\"x\"}".to_string()),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cache_key_requires_non_empty_segments() {
|
||||
assert!(ProviderTransportSnapshotCacheKey::new("provider", "endpoint", "key").is_some());
|
||||
assert!(ProviderTransportSnapshotCacheKey::new("", "endpoint", "key").is_none());
|
||||
assert!(ProviderTransportSnapshotCacheKey::new("provider", " ", "key").is_none());
|
||||
assert!(ProviderTransportSnapshotCacheKey::new("provider", "endpoint", "").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn refresh_detection_tracks_key_material_and_expiry() {
|
||||
let current = sample_snapshot();
|
||||
let mut refreshed = current.clone();
|
||||
assert!(!provider_transport_snapshot_looks_refreshed(
|
||||
¤t, &refreshed
|
||||
));
|
||||
|
||||
refreshed.key.decrypted_api_key = "sk-updated".to_string();
|
||||
assert!(provider_transport_snapshot_looks_refreshed(
|
||||
¤t, &refreshed
|
||||
));
|
||||
|
||||
let mut refreshed = current.clone();
|
||||
refreshed.key.decrypted_auth_config = Some("{\"token\":\"y\"}".to_string());
|
||||
assert!(provider_transport_snapshot_looks_refreshed(
|
||||
¤t, &refreshed
|
||||
));
|
||||
|
||||
let mut refreshed = current.clone();
|
||||
refreshed.key.expires_at_unix_secs = Some(2);
|
||||
assert!(provider_transport_snapshot_looks_refreshed(
|
||||
¤t, &refreshed
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
use super::super::auth::resolve_local_standard_auth;
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::super::supports_local_oauth_request_auth_resolution;
|
||||
|
||||
pub fn supports_local_claude_code_auth(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
resolve_local_standard_auth(transport).is_some_and(|(_, value)| !value.trim().is_empty())
|
||||
|| supports_local_oauth_request_auth_resolution(transport)
|
||||
}
|
||||
@@ -0,0 +1,386 @@
|
||||
use aether_contracts::{
|
||||
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
use sha2::{Digest, Sha256};
|
||||
use uuid::Uuid;
|
||||
|
||||
// Chrome impersonate profiles
|
||||
const CHROME_IMPERSONATE_PROFILES: &[&str] = &[
|
||||
"chrome110",
|
||||
"chrome116",
|
||||
"chrome119",
|
||||
"chrome120",
|
||||
"chrome123",
|
||||
"chrome124",
|
||||
"chrome131",
|
||||
"chrome133",
|
||||
];
|
||||
|
||||
const CHROME_VERSIONS: &[(&str, &str)] = &[
|
||||
("chrome110", "110.0.5481.177"),
|
||||
("chrome116", "116.0.5845.188"),
|
||||
("chrome119", "119.0.6045.214"),
|
||||
("chrome120", "120.0.6099.216"),
|
||||
("chrome123", "123.0.6312.122"),
|
||||
("chrome124", "124.0.6367.243"),
|
||||
("chrome131", "131.0.6778.265"),
|
||||
("chrome133", "133.0.6943.142"),
|
||||
];
|
||||
|
||||
// (os, arch, platform_token, platform_info)
|
||||
const PLATFORM_VARIANTS: &[(&str, &str, &str, &str)] = &[
|
||||
("Linux", "x64", "X11; Linux x86_64", "Linux x86_64"),
|
||||
("Linux", "arm64", "X11; Linux arm64", "Linux arm64"),
|
||||
(
|
||||
"Windows",
|
||||
"x64",
|
||||
"Windows NT 10.0; Win64; x64",
|
||||
"Windows x64",
|
||||
),
|
||||
(
|
||||
"MacOS",
|
||||
"x64",
|
||||
"Macintosh; Intel Mac OS X 10_15_7",
|
||||
"Darwin x64",
|
||||
),
|
||||
(
|
||||
"MacOS",
|
||||
"arm64",
|
||||
"Macintosh; ARM Mac OS X 14_0_0",
|
||||
"Darwin arm64",
|
||||
),
|
||||
];
|
||||
|
||||
const STAINLESS_PACKAGE_VERSIONS: &[&str] = &["0.68.0", "0.69.0", "0.70.0", "0.71.0"];
|
||||
const NODE_VERSIONS: &[&str] = &["v20.18.1", "v22.12.0", "v22.14.0", "v24.13.0"];
|
||||
const ELECTRON_VERSIONS: &[&str] = &["35.5.1", "36.7.1", "37.3.0", "38.7.0", "39.2.3"];
|
||||
const STAINLESS_TIMEOUTS: &[&str] = &["600", "900"];
|
||||
const CLAUDE_CODE_TRANSPORT_PROFILE_ID: &str = "claude_code_nodejs";
|
||||
|
||||
/// Deterministic hash-based index picker, compatible with Python implementation.
|
||||
/// Each `slot` produces a different selection from the same seed.
|
||||
struct SeededPicker {
|
||||
seed_bytes: [u8; 32],
|
||||
}
|
||||
|
||||
impl SeededPicker {
|
||||
fn new(seed: &str) -> Self {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(seed.as_bytes());
|
||||
Self {
|
||||
seed_bytes: hasher.finalize().into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Pick an index from `[0, len)` using a specific slot.
|
||||
/// Different slots produce independent-looking selections from the same seed.
|
||||
fn pick(&self, slot: u8, len: usize) -> usize {
|
||||
if len == 0 {
|
||||
return 0;
|
||||
}
|
||||
// Hash seed_bytes + slot to get a new digest, take first 8 bytes as u64
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(self.seed_bytes);
|
||||
hasher.update([slot]);
|
||||
let hash = hasher.finalize();
|
||||
let value = u64::from_be_bytes(hash[..8].try_into().unwrap());
|
||||
(value % len as u64) as usize
|
||||
}
|
||||
}
|
||||
|
||||
fn chrome_version_for_profile(profile: &str) -> &'static str {
|
||||
for (p, v) in CHROME_VERSIONS {
|
||||
if p.eq_ignore_ascii_case(profile) {
|
||||
return v;
|
||||
}
|
||||
}
|
||||
"120.0.6099.216"
|
||||
}
|
||||
|
||||
fn build_user_agent(platform_token: &str, chrome_version: &str, electron_version: &str) -> String {
|
||||
format!(
|
||||
"Mozilla/5.0 ({platform_token}) AppleWebKit/537.36 (KHTML, like Gecko) \
|
||||
Chrome/{chrome_version} Electron/{electron_version} Safari/537.36"
|
||||
)
|
||||
}
|
||||
|
||||
fn resolve_platform_token(os: &str, arch: &str) -> &'static str {
|
||||
let os_lower = os.to_ascii_lowercase();
|
||||
let arch_lower = arch.to_ascii_lowercase();
|
||||
|
||||
if os_lower.starts_with("win") {
|
||||
return "Windows NT 10.0; Win64; x64";
|
||||
}
|
||||
if matches!(os_lower.as_str(), "darwin" | "mac" | "macos") {
|
||||
return if matches!(arch_lower.as_str(), "arm64" | "aarch64") {
|
||||
"Macintosh; ARM Mac OS X 14_0_0"
|
||||
} else {
|
||||
"Macintosh; Intel Mac OS X 10_15_7"
|
||||
};
|
||||
}
|
||||
if matches!(arch_lower.as_str(), "arm64" | "aarch64") {
|
||||
"X11; Linux arm64"
|
||||
} else {
|
||||
"X11; Linux x86_64"
|
||||
}
|
||||
}
|
||||
|
||||
/// Generate a complete Claude Code transport fingerprint from a seed.
|
||||
pub fn generate_fingerprint(seed: &str) -> Value {
|
||||
wrap_header_fingerprint(generate_header_fingerprint(seed))
|
||||
}
|
||||
|
||||
fn generate_header_fingerprint(seed: &str) -> Value {
|
||||
let picker = SeededPicker::new(seed);
|
||||
|
||||
let impersonate =
|
||||
CHROME_IMPERSONATE_PROFILES[picker.pick(0, CHROME_IMPERSONATE_PROFILES.len())];
|
||||
let chrome_version = chrome_version_for_profile(impersonate);
|
||||
let node_version = NODE_VERSIONS[picker.pick(1, NODE_VERSIONS.len())];
|
||||
let electron_version = ELECTRON_VERSIONS[picker.pick(2, ELECTRON_VERSIONS.len())];
|
||||
let platform = PLATFORM_VARIANTS[picker.pick(3, PLATFORM_VARIANTS.len())];
|
||||
let (stainless_os, stainless_arch, platform_token, platform_info) = platform;
|
||||
let stainless_package_version =
|
||||
STAINLESS_PACKAGE_VERSIONS[picker.pick(4, STAINLESS_PACKAGE_VERSIONS.len())];
|
||||
let stainless_timeout = STAINLESS_TIMEOUTS[picker.pick(5, STAINLESS_TIMEOUTS.len())];
|
||||
|
||||
let vscode_session_id = Uuid::new_v5(
|
||||
&Uuid::NAMESPACE_URL,
|
||||
format!("aether:fingerprint:{seed}").as_bytes(),
|
||||
)
|
||||
.simple()
|
||||
.to_string();
|
||||
|
||||
let user_agent = build_user_agent(platform_token, chrome_version, electron_version);
|
||||
|
||||
serde_json::json!({
|
||||
"impersonate": impersonate,
|
||||
"stainless_package_version": stainless_package_version,
|
||||
"stainless_os": stainless_os,
|
||||
"stainless_arch": stainless_arch,
|
||||
"stainless_runtime_version": node_version,
|
||||
"stainless_timeout": stainless_timeout,
|
||||
"node_version": node_version,
|
||||
"chrome_version": chrome_version,
|
||||
"electron_version": electron_version,
|
||||
"vscode_session_id": vscode_session_id,
|
||||
"platform_info": platform_info,
|
||||
"user_agent": user_agent,
|
||||
})
|
||||
}
|
||||
|
||||
fn wrap_header_fingerprint(header_fingerprint: Value) -> Value {
|
||||
serde_json::json!({
|
||||
"transport_profile": {
|
||||
"profile_id": CLAUDE_CODE_TRANSPORT_PROFILE_ID,
|
||||
"backend": TRANSPORT_BACKEND_REQWEST_RUSTLS,
|
||||
"http_mode": TRANSPORT_HTTP_MODE_AUTO,
|
||||
"pool_scope": TRANSPORT_POOL_SCOPE_KEY,
|
||||
"header_fingerprint": header_fingerprint,
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
pub fn header_fingerprint_from_fingerprint(fingerprint: &Value) -> Option<&Map<String, Value>> {
|
||||
fingerprint
|
||||
.get("transport_profile")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|profile| profile.get("header_fingerprint"))
|
||||
.and_then(Value::as_object)
|
||||
}
|
||||
|
||||
/// Generate a random (non-deterministic) fingerprint.
|
||||
pub fn generate_random_fingerprint() -> Value {
|
||||
let random_seed = Uuid::new_v4().to_string();
|
||||
generate_fingerprint(&random_seed)
|
||||
}
|
||||
|
||||
/// Sanitize an existing fingerprint JSON, filling missing fields with
|
||||
/// deterministic fallbacks derived from `key_id`.
|
||||
pub fn sanitize_fingerprint(raw: &Value, key_id: &str) -> Value {
|
||||
let generated = generate_header_fingerprint(key_id);
|
||||
let gen_map = generated.as_object().unwrap();
|
||||
let raw_map = header_fingerprint_from_fingerprint(raw);
|
||||
|
||||
let mut out = Map::new();
|
||||
|
||||
// Start with generated values, then overlay non-empty raw values
|
||||
for (key, gen_value) in gen_map {
|
||||
let value = raw_map
|
||||
.and_then(|raw_map| raw_map.get(key))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|v| !v.is_empty())
|
||||
.map(|v| Value::String(v.to_string()))
|
||||
.unwrap_or_else(|| gen_value.clone());
|
||||
out.insert(key.clone(), value);
|
||||
}
|
||||
|
||||
// Normalize impersonate to known profile
|
||||
let impersonate = out
|
||||
.get("impersonate")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default()
|
||||
.to_ascii_lowercase();
|
||||
let is_known = CHROME_IMPERSONATE_PROFILES
|
||||
.iter()
|
||||
.any(|p| p.eq_ignore_ascii_case(&impersonate));
|
||||
if !is_known {
|
||||
out.insert("impersonate".to_string(), gen_map["impersonate"].clone());
|
||||
}
|
||||
|
||||
// Ensure chrome_version matches impersonate profile
|
||||
let profile = out
|
||||
.get("impersonate")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let chrome_version = chrome_version_for_profile(profile);
|
||||
out.insert(
|
||||
"chrome_version".to_string(),
|
||||
Value::String(chrome_version.to_string()),
|
||||
);
|
||||
|
||||
// Rebuild user_agent if missing
|
||||
let has_ua = out
|
||||
.get("user_agent")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|v| !v.is_empty());
|
||||
if !has_ua {
|
||||
let os = out
|
||||
.get("stainless_os")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("Linux");
|
||||
let arch = out
|
||||
.get("stainless_arch")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("x64");
|
||||
let electron = out
|
||||
.get("electron_version")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("38.7.0");
|
||||
let platform_token = resolve_platform_token(os, arch);
|
||||
out.insert(
|
||||
"user_agent".to_string(),
|
||||
Value::String(build_user_agent(platform_token, chrome_version, electron)),
|
||||
);
|
||||
}
|
||||
|
||||
wrap_header_fingerprint(Value::Object(out))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn deterministic_generation_from_seed() {
|
||||
let fp1 = generate_fingerprint("key-abc-123");
|
||||
let fp2 = generate_fingerprint("key-abc-123");
|
||||
assert_eq!(fp1, fp2, "same seed should produce identical fingerprint");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn different_seeds_produce_different_fingerprints() {
|
||||
let fp1 = generate_fingerprint("key-1");
|
||||
let fp2 = generate_fingerprint("key-2");
|
||||
// At least one field should differ (statistically near-certain)
|
||||
assert_ne!(fp1, fp2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_fingerprint_has_all_fields() {
|
||||
let fp = generate_fingerprint("test-key");
|
||||
let expected_keys = [
|
||||
"impersonate",
|
||||
"stainless_package_version",
|
||||
"stainless_os",
|
||||
"stainless_arch",
|
||||
"stainless_runtime_version",
|
||||
"stainless_timeout",
|
||||
"node_version",
|
||||
"chrome_version",
|
||||
"electron_version",
|
||||
"vscode_session_id",
|
||||
"platform_info",
|
||||
"user_agent",
|
||||
];
|
||||
let map = header_fingerprint_from_fingerprint(&fp).unwrap();
|
||||
for key in expected_keys {
|
||||
assert!(map.contains_key(key), "missing field: {key}");
|
||||
let value = map[key].as_str().unwrap();
|
||||
assert!(!value.is_empty(), "empty field: {key}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitize_preserves_user_overrides() {
|
||||
let raw = serde_json::json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "claude_code_nodejs",
|
||||
"header_fingerprint": {
|
||||
"stainless_os": "MacOS",
|
||||
"stainless_arch": "arm64",
|
||||
"stainless_timeout": "900",
|
||||
"user_agent": "Custom-Agent/1.0"
|
||||
}
|
||||
}
|
||||
});
|
||||
let sanitized = sanitize_fingerprint(&raw, "test-key");
|
||||
let map = header_fingerprint_from_fingerprint(&sanitized).unwrap();
|
||||
assert_eq!(map["stainless_os"].as_str(), Some("MacOS"));
|
||||
assert_eq!(map["stainless_arch"].as_str(), Some("arm64"));
|
||||
assert_eq!(map["stainless_timeout"].as_str(), Some("900"));
|
||||
assert_eq!(map["user_agent"].as_str(), Some("Custom-Agent/1.0"));
|
||||
// Other fields should be filled from generation
|
||||
assert!(map.contains_key("impersonate"));
|
||||
assert!(map.contains_key("stainless_package_version"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitize_fills_missing_fields_from_seed() {
|
||||
let raw = serde_json::json!({});
|
||||
let sanitized = sanitize_fingerprint(&raw, "test-key");
|
||||
let generated = generate_header_fingerprint("test-key");
|
||||
// All fields should match generated since raw is empty
|
||||
let s = header_fingerprint_from_fingerprint(&sanitized).unwrap();
|
||||
let g = generated.as_object().unwrap();
|
||||
for key in g.keys() {
|
||||
assert!(s.contains_key(key), "sanitized missing key: {key}");
|
||||
assert!(
|
||||
!s[key].as_str().unwrap().is_empty(),
|
||||
"sanitized empty key: {key}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitize_normalizes_unknown_impersonate_profile() {
|
||||
let raw = serde_json::json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "claude_code_nodejs",
|
||||
"header_fingerprint": {
|
||||
"impersonate": "firefox99"
|
||||
}
|
||||
}
|
||||
});
|
||||
let sanitized = sanitize_fingerprint(&raw, "test-key");
|
||||
let profile = header_fingerprint_from_fingerprint(&sanitized).unwrap()["impersonate"]
|
||||
.as_str()
|
||||
.unwrap();
|
||||
assert!(
|
||||
CHROME_IMPERSONATE_PROFILES
|
||||
.iter()
|
||||
.any(|p| p.eq_ignore_ascii_case(profile)),
|
||||
"should normalize to known profile, got: {profile}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn random_fingerprint_differs_each_call() {
|
||||
let fp1 = generate_random_fingerprint();
|
||||
let fp2 = generate_random_fingerprint();
|
||||
assert_ne!(fp1, fp2);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
mod auth;
|
||||
mod fingerprint;
|
||||
mod policy;
|
||||
mod request;
|
||||
mod url;
|
||||
|
||||
pub use auth::supports_local_claude_code_auth;
|
||||
pub use fingerprint::{
|
||||
generate_fingerprint, generate_random_fingerprint, header_fingerprint_from_fingerprint,
|
||||
sanitize_fingerprint,
|
||||
};
|
||||
pub use policy::{
|
||||
local_claude_code_transport_unsupported_reason_with_network,
|
||||
supports_local_claude_code_transport_with_network,
|
||||
};
|
||||
pub use request::{build_claude_code_passthrough_headers, sanitize_claude_code_request_body};
|
||||
pub use url::build_claude_code_messages_url;
|
||||
@@ -0,0 +1,162 @@
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::super::{
|
||||
body_rules_are_locally_supported, header_rules_are_locally_supported,
|
||||
resolve_transport_profile, supports_local_oauth_request_auth_resolution,
|
||||
transport_profile_is_configured, transport_proxy_is_locally_supported,
|
||||
};
|
||||
use super::auth::supports_local_claude_code_auth;
|
||||
|
||||
pub fn local_claude_code_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return if !transport.provider.is_active {
|
||||
Some("provider_inactive")
|
||||
} else if !transport.endpoint.is_active {
|
||||
Some("endpoint_inactive")
|
||||
} else {
|
||||
Some("key_inactive")
|
||||
};
|
||||
}
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("claude_code")
|
||||
{
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
if !transport
|
||||
.endpoint
|
||||
.api_format
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(api_format.trim())
|
||||
{
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
|
||||
return Some("transport_header_rules_unsupported");
|
||||
}
|
||||
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
if transport.key.decrypted_auth_config.is_some()
|
||||
&& !supports_local_oauth_request_auth_resolution(transport)
|
||||
{
|
||||
return Some("transport_oauth_resolution_unsupported");
|
||||
}
|
||||
if !supports_local_claude_code_auth(transport) {
|
||||
return Some("transport_auth_unavailable");
|
||||
}
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return Some("transport_proxy_unsupported");
|
||||
}
|
||||
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return Some("transport_profile_unsupported");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub fn supports_local_claude_code_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
local_claude_code_transport_unsupported_reason_with_network(transport, api_format).is_none()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::{
|
||||
local_claude_code_transport_unsupported_reason_with_network,
|
||||
supports_local_claude_code_transport_with_network,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Claude Code".to_string(),
|
||||
provider_type: "claude_code".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "claude:messages".to_string(),
|
||||
api_family: Some("claude".to_string()),
|
||||
endpoint_kind: Some("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-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "bearer".to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec!["claude:messages".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: Some(json!({"transport_profile":"chrome_136"})),
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "sk-ant-123".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_claude_code_transport_when_auth_and_profile_are_valid() {
|
||||
assert!(supports_local_claude_code_transport_with_network(
|
||||
&sample_transport(),
|
||||
"claude:messages"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reports_auth_unavailable_for_claude_code_without_local_auth() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "api_key".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
|
||||
assert_eq!(
|
||||
local_claude_code_transport_unsupported_reason_with_network(
|
||||
&transport,
|
||||
"claude:messages"
|
||||
),
|
||||
Some("transport_auth_unavailable")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::super::auth::build_openai_passthrough_headers;
|
||||
use super::fingerprint::header_fingerprint_from_fingerprint;
|
||||
|
||||
const DEFAULT_ANTHROPIC_VERSION: &str = "2023-06-01";
|
||||
const DEFAULT_ACCEPT: &str = "application/json";
|
||||
const STREAM_HELPER_METHOD: &str = "stream";
|
||||
const DUMMY_THINKING_SIGNATURE: &str = "skip_thought_signature_validator";
|
||||
const REQUIRED_BETA_TOKENS: &[&str] = &[
|
||||
"claude-code-20250219",
|
||||
"oauth-2025-04-20",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
];
|
||||
const EXCLUDED_BETA_TOKENS: &[&str] = &["context-1m-2025-08-07"];
|
||||
|
||||
/// Fingerprint field -> HTTP header mapping.
|
||||
/// Every stainless / identity dimension that can vary per-key is listed here.
|
||||
const FINGERPRINT_HEADER_MAP: &[(&str, &str)] = &[
|
||||
("stainless_package_version", "x-stainless-package-version"),
|
||||
("stainless_os", "x-stainless-os"),
|
||||
("stainless_arch", "x-stainless-arch"),
|
||||
("stainless_runtime_version", "x-stainless-runtime-version"),
|
||||
("stainless_timeout", "x-stainless-timeout"),
|
||||
("user_agent", "user-agent"),
|
||||
];
|
||||
|
||||
pub fn build_claude_code_passthrough_headers(
|
||||
headers: &http::HeaderMap,
|
||||
auth_header: &str,
|
||||
auth_value: &str,
|
||||
extra_headers: &BTreeMap<String, String>,
|
||||
stream: bool,
|
||||
fingerprint: Option<&Value>,
|
||||
) -> BTreeMap<String, String> {
|
||||
let mut out = build_openai_passthrough_headers(
|
||||
headers,
|
||||
auth_header,
|
||||
auth_value,
|
||||
extra_headers,
|
||||
Some("application/json"),
|
||||
);
|
||||
|
||||
// -- Anthropic protocol headers --
|
||||
out.insert("accept".to_string(), DEFAULT_ACCEPT.to_string());
|
||||
out.insert(
|
||||
"anthropic-version".to_string(),
|
||||
DEFAULT_ANTHROPIC_VERSION.to_string(),
|
||||
);
|
||||
// Read incoming anthropic-beta directly from the original HeaderMap because the
|
||||
// upstream passthrough filter now strips `anthropic-*` headers to avoid leaking
|
||||
// them to non-Anthropic upstreams.
|
||||
let incoming_anthropic_beta = headers
|
||||
.get("anthropic-beta")
|
||||
.and_then(|value| value.to_str().ok());
|
||||
out.insert(
|
||||
"anthropic-beta".to_string(),
|
||||
merge_anthropic_beta_tokens(incoming_anthropic_beta),
|
||||
);
|
||||
out.insert(
|
||||
"anthropic-dangerous-direct-browser-access".to_string(),
|
||||
"true".to_string(),
|
||||
);
|
||||
out.insert("x-app".to_string(), "cli".to_string());
|
||||
|
||||
// -- Stainless SDK identity headers --
|
||||
// Fixed values: these don't vary per fingerprint.
|
||||
out.insert("x-stainless-lang".to_string(), "js".to_string());
|
||||
out.insert("x-stainless-runtime".to_string(), "node".to_string());
|
||||
out.insert("x-stainless-retry-count".to_string(), "0".to_string());
|
||||
|
||||
// Defaults for fingerprint-overridable fields (used when no fingerprint is present).
|
||||
out.insert(
|
||||
"x-stainless-package-version".to_string(),
|
||||
"0.70.0".to_string(),
|
||||
);
|
||||
out.insert("x-stainless-os".to_string(), "Linux".to_string());
|
||||
out.insert("x-stainless-arch".to_string(), "arm64".to_string());
|
||||
out.insert(
|
||||
"x-stainless-runtime-version".to_string(),
|
||||
"v24.13.0".to_string(),
|
||||
);
|
||||
out.insert("x-stainless-timeout".to_string(), "600".to_string());
|
||||
|
||||
if stream {
|
||||
out.insert(
|
||||
"x-stainless-helper-method".to_string(),
|
||||
STREAM_HELPER_METHOD.to_string(),
|
||||
);
|
||||
} else {
|
||||
out.remove("x-stainless-helper-method");
|
||||
}
|
||||
|
||||
// Override from the formal transport profile header fingerprint.
|
||||
if let Some(fp) = fingerprint.and_then(header_fingerprint_from_fingerprint) {
|
||||
for &(fp_key, header_key) in FINGERPRINT_HEADER_MAP {
|
||||
if let Some(value) = fp
|
||||
.get(fp_key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|v| !v.is_empty())
|
||||
{
|
||||
out.insert(header_key.to_string(), value.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
out
|
||||
}
|
||||
|
||||
pub fn sanitize_claude_code_request_body(body: &mut Value) {
|
||||
let Some(body_object) = body.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
let thinking_enabled = body_object
|
||||
.get("thinking")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|thinking| thinking.get("type"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| matches!(value.to_ascii_lowercase().as_str(), "enabled" | "adaptive"));
|
||||
|
||||
let Some(messages) = body_object
|
||||
.get_mut("messages")
|
||||
.and_then(Value::as_array_mut)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
|
||||
for message in messages {
|
||||
let Some(message_object) = message.as_object_mut() else {
|
||||
continue;
|
||||
};
|
||||
let role = message_object
|
||||
.get("role")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
let Some(content) = message_object
|
||||
.get_mut("content")
|
||||
.and_then(Value::as_array_mut)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let mut filtered = Vec::with_capacity(content.len());
|
||||
for block in std::mem::take(content) {
|
||||
let Value::Object(block_object) = block else {
|
||||
filtered.push(block);
|
||||
continue;
|
||||
};
|
||||
if keep_claude_code_block(&block_object, &role, thinking_enabled) {
|
||||
filtered.push(Value::Object(block_object));
|
||||
}
|
||||
}
|
||||
*content = filtered;
|
||||
}
|
||||
}
|
||||
|
||||
fn keep_claude_code_block(
|
||||
block_object: &Map<String, Value>,
|
||||
role: &str,
|
||||
thinking_enabled: bool,
|
||||
) -> bool {
|
||||
let block_type = block_object
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default();
|
||||
if matches!(block_type, "thinking" | "redacted_thinking") {
|
||||
let signature = block_object
|
||||
.get("signature")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default();
|
||||
return thinking_enabled
|
||||
&& role.eq_ignore_ascii_case("assistant")
|
||||
&& !signature.is_empty()
|
||||
&& signature != DUMMY_THINKING_SIGNATURE;
|
||||
}
|
||||
if block_type.is_empty() && block_object.contains_key("thinking") {
|
||||
return false;
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
fn merge_anthropic_beta_tokens(incoming: Option<&str>) -> String {
|
||||
let mut seen = BTreeSet::new();
|
||||
let mut merged = Vec::new();
|
||||
|
||||
for token in REQUIRED_BETA_TOKENS {
|
||||
append_beta_token(&mut seen, &mut merged, token);
|
||||
}
|
||||
for token in incoming.unwrap_or_default().split(',') {
|
||||
let token = token.trim();
|
||||
if EXCLUDED_BETA_TOKENS
|
||||
.iter()
|
||||
.any(|excluded| token.eq_ignore_ascii_case(excluded))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
append_beta_token(&mut seen, &mut merged, token);
|
||||
}
|
||||
|
||||
merged.join(",")
|
||||
}
|
||||
|
||||
fn append_beta_token(seen: &mut BTreeSet<String>, merged: &mut Vec<String>, token: &str) {
|
||||
let normalized = token.trim();
|
||||
if normalized.is_empty() {
|
||||
return;
|
||||
}
|
||||
let key = normalized.to_ascii_lowercase();
|
||||
if seen.insert(key) {
|
||||
merged.push(normalized.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{build_claude_code_passthrough_headers, sanitize_claude_code_request_body};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn claude_code_headers_use_transport_profile_header_fingerprint_and_merge_required_betas() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert(
|
||||
"anthropic-beta",
|
||||
http::HeaderValue::from_static("context-1m-2025-08-07,custom-beta"),
|
||||
);
|
||||
headers.insert(
|
||||
"user-agent",
|
||||
http::HeaderValue::from_static("Claude-Code/Test"),
|
||||
);
|
||||
let built = build_claude_code_passthrough_headers(
|
||||
&headers,
|
||||
"authorization",
|
||||
"Bearer upstream-token",
|
||||
&BTreeMap::new(),
|
||||
true,
|
||||
Some(&json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "claude_code_nodejs",
|
||||
"header_fingerprint": {
|
||||
"user_agent":"Claude-Code/9.9",
|
||||
"stainless_package_version":"1.0.5",
|
||||
"stainless_runtime_version":"v22.12.0",
|
||||
"stainless_timeout":"900"
|
||||
}
|
||||
}
|
||||
})),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
built.get("anthropic-beta").map(String::as_str),
|
||||
Some(
|
||||
"claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14,custom-beta"
|
||||
)
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("anthropic-version").map(String::as_str),
|
||||
Some("2023-06-01")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("accept").map(String::as_str),
|
||||
Some("application/json")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("x-stainless-helper-method").map(String::as_str),
|
||||
Some("stream")
|
||||
);
|
||||
assert_eq!(built.get("x-app").map(String::as_str), Some("cli"));
|
||||
assert_eq!(
|
||||
built.get("x-stainless-package-version").map(String::as_str),
|
||||
Some("1.0.5")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("x-stainless-runtime-version").map(String::as_str),
|
||||
Some("v22.12.0")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("x-stainless-timeout").map(String::as_str),
|
||||
Some("900")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("user-agent").map(String::as_str),
|
||||
Some("Claude-Code/9.9")
|
||||
);
|
||||
assert_eq!(
|
||||
built.get("authorization").map(String::as_str),
|
||||
Some("Bearer upstream-token")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_code_body_sanitizer_drops_invalid_thinking_blocks() {
|
||||
let mut body = json!({
|
||||
"thinking": {"type":"enabled"},
|
||||
"messages": [{
|
||||
"role":"assistant",
|
||||
"content":[
|
||||
{"type":"thinking","thinking":"keep","signature":"sig_valid"},
|
||||
{"type":"thinking","thinking":"drop-empty","signature":""},
|
||||
{"type":"redacted_thinking","data":"keep-redacted","signature":"sig_redacted"},
|
||||
{"type":"redacted_thinking","data":"drop-no-signature"},
|
||||
{"thinking":"drop-no-type"},
|
||||
{"type":"text","text":"ok"}
|
||||
]
|
||||
}]
|
||||
});
|
||||
|
||||
sanitize_claude_code_request_body(&mut body);
|
||||
|
||||
assert_eq!(
|
||||
body["messages"][0]["content"],
|
||||
json!([
|
||||
{"type":"thinking","thinking":"keep","signature":"sig_valid"},
|
||||
{"type":"redacted_thinking","data":"keep-redacted","signature":"sig_redacted"},
|
||||
{"type":"text","text":"ok"}
|
||||
])
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use url::form_urlencoded;
|
||||
|
||||
pub fn build_claude_code_messages_url(upstream_base_url: &str, query: Option<&str>) -> String {
|
||||
let (trimmed_base_url, base_query) = split_query(upstream_base_url.trim());
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
let mut url =
|
||||
if trimmed_base_url.ends_with("/v1/messages") || trimmed_base_url.ends_with("/messages") {
|
||||
trimmed_base_url.to_string()
|
||||
} else if trimmed_base_url.ends_with("/v1") {
|
||||
format!("{trimmed_base_url}/messages")
|
||||
} else {
|
||||
format!("{trimmed_base_url}/v1/messages")
|
||||
};
|
||||
append_merged_query(&mut url, base_query, query);
|
||||
url
|
||||
}
|
||||
|
||||
fn split_query(value: &str) -> (&str, Option<&str>) {
|
||||
value
|
||||
.split_once('?')
|
||||
.map(|(base, query)| (base, Some(query)))
|
||||
.unwrap_or((value, None))
|
||||
}
|
||||
|
||||
fn append_merged_query(url: &mut String, base_query: Option<&str>, request_query: Option<&str>) {
|
||||
let Some(query) = merge_query_layers(base_query, request_query) else {
|
||||
return;
|
||||
};
|
||||
if url.contains('?') {
|
||||
url.push('&');
|
||||
} else {
|
||||
url.push('?');
|
||||
}
|
||||
url.push_str(&query);
|
||||
}
|
||||
|
||||
fn merge_query_layers(base_query: Option<&str>, request_query: Option<&str>) -> Option<String> {
|
||||
let mut merged = BTreeMap::new();
|
||||
for source in [base_query, request_query] {
|
||||
let Some(source) = source.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||
continue;
|
||||
};
|
||||
for (key, value) in form_urlencoded::parse(source.as_bytes()) {
|
||||
merged.insert(key.into_owned(), value.into_owned());
|
||||
}
|
||||
}
|
||||
if merged.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in merged {
|
||||
serializer.append_pair(&key, &value);
|
||||
}
|
||||
Some(serializer.finish())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::build_claude_code_messages_url;
|
||||
|
||||
#[test]
|
||||
fn keeps_existing_messages_suffix_without_duplication() {
|
||||
assert_eq!(
|
||||
build_claude_code_messages_url("https://api.anthropic.com/v1/messages", None),
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn appends_messages_and_merges_query() {
|
||||
assert_eq!(
|
||||
build_claude_code_messages_url(
|
||||
"https://api.anthropic.com/v1?beta=true",
|
||||
Some("foo=bar"),
|
||||
),
|
||||
"https://api.anthropic.com/v1/messages?beta=true&foo=bar"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,864 @@
|
||||
use aether_ai_formats::formats::matrix::{
|
||||
api_data_format_id, request_conversion_kind, request_conversion_requires_enable_flag,
|
||||
RequestConversionKind,
|
||||
};
|
||||
use aether_ai_formats::normalize_api_format_alias;
|
||||
|
||||
use crate::antigravity::is_antigravity_provider_transport;
|
||||
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, is_vertex_transport_context,
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network,
|
||||
resolve_local_vertex_api_key_query_auth, VERTEX_API_KEY_QUERY_PARAM,
|
||||
};
|
||||
use crate::windsurf::{
|
||||
is_windsurf_provider_transport,
|
||||
local_windsurf_request_transport_unsupported_reason_with_network,
|
||||
resolve_windsurf_cascade_auth,
|
||||
};
|
||||
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;
|
||||
}
|
||||
let conversion_kind =
|
||||
request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str());
|
||||
if conversion_kind.is_none()
|
||||
&& !same_data_format_transport_pair(
|
||||
client_api_format.as_str(),
|
||||
provider_api_format.as_str(),
|
||||
)
|
||||
{
|
||||
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;
|
||||
}
|
||||
let conversion_kind =
|
||||
request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str());
|
||||
if conversion_kind.is_none()
|
||||
&& !same_data_format_transport_pair(
|
||||
client_api_format.as_str(),
|
||||
provider_api_format.as_str(),
|
||||
)
|
||||
{
|
||||
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(),
|
||||
)
|
||||
}
|
||||
|
||||
fn same_data_format_transport_pair(client_api_format: &str, provider_api_format: &str) -> bool {
|
||||
if aether_ai_formats::api_format_alias_matches(client_api_format, provider_api_format) {
|
||||
return false;
|
||||
}
|
||||
matches!(
|
||||
(
|
||||
api_data_format_id(client_api_format),
|
||||
api_data_format_id(provider_api_format)
|
||||
),
|
||||
(Some("embedding"), Some("embedding")) | (Some("rerank"), Some("rerank"))
|
||||
)
|
||||
}
|
||||
|
||||
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);
|
||||
}
|
||||
if is_windsurf_provider_transport(transport)
|
||||
&& normalize_api_format_alias(&transport.endpoint.api_format) == "openai:chat"
|
||||
{
|
||||
return local_windsurf_request_transport_unsupported_reason_with_network(transport);
|
||||
}
|
||||
if is_antigravity_provider_transport(transport)
|
||||
&& normalize_api_format_alias(&transport.endpoint.api_format) == "gemini:generate_content"
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
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_transport_context(transport) => {
|
||||
local_vertex_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_pair_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
client_api_format: &str,
|
||||
provider_api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
let client_api_format = normalize_api_format_alias(client_api_format);
|
||||
let provider_api_format = normalize_api_format_alias(provider_api_format);
|
||||
|
||||
if let Some(kind) =
|
||||
request_conversion_kind(client_api_format.as_str(), provider_api_format.as_str())
|
||||
{
|
||||
return request_conversion_transport_unsupported_reason(transport, kind);
|
||||
}
|
||||
|
||||
if !same_data_format_transport_pair(client_api_format.as_str(), provider_api_format.as_str()) {
|
||||
return Some("transport_api_format_unsupported");
|
||||
}
|
||||
|
||||
match provider_api_format.as_str() {
|
||||
"gemini:embedding" => {
|
||||
if is_vertex_transport_context(transport) {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network(transport)
|
||||
} else {
|
||||
local_gemini_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
"gemini:embedding",
|
||||
)
|
||||
}
|
||||
}
|
||||
"openai:embedding"
|
||||
| "jina:embedding"
|
||||
| "doubao:embedding"
|
||||
| "aliyun:multimodal_embedding"
|
||||
| "openai:rerank"
|
||||
| "jina:rerank" => local_standard_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
provider_api_format.as_str(),
|
||||
),
|
||||
_ => Some("transport_api_format_unsupported"),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn request_conversion_direct_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
_kind: RequestConversionKind,
|
||||
) -> Option<(String, String)> {
|
||||
request_direct_auth_for_provider_format(transport, transport.endpoint.api_format.as_str())
|
||||
}
|
||||
|
||||
pub fn request_pair_direct_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> Option<(String, String)> {
|
||||
request_direct_auth_for_provider_format(transport, provider_api_format)
|
||||
}
|
||||
|
||||
fn request_direct_auth_for_provider_format(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> Option<(String, String)> {
|
||||
match normalize_api_format_alias(provider_api_format).as_str() {
|
||||
"openai:chat" if is_windsurf_provider_transport(transport) => {
|
||||
resolve_windsurf_cascade_auth(transport)
|
||||
}
|
||||
"openai:chat"
|
||||
| "openai:responses"
|
||||
| "openai:responses:compact"
|
||||
| "openai:embedding"
|
||||
| "jina:embedding"
|
||||
| "doubao:embedding"
|
||||
| "aliyun:multimodal_embedding"
|
||||
| "openai:rerank"
|
||||
| "jina:rerank" => resolve_local_openai_bearer_auth(transport),
|
||||
"gemini:generate_content" | "gemini:embedding" => {
|
||||
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_permission_covers(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());
|
||||
let requested_base_model = aether_ai_formats::model_directive_base_model(requested_model);
|
||||
|
||||
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
|
||||
|| requested_base_model
|
||||
.as_deref()
|
||||
.is_some_and(|base_model| allowed_model == base_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,
|
||||
request_pair_direct_auth, CandidateTransportPolicyFacts,
|
||||
};
|
||||
use aether_ai_formats::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,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
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 cross_format_pair_requires_provider_or_endpoint_enablement() {
|
||||
let disabled = transport_snapshot("custom", "openai:responses", "bearer", false, None);
|
||||
assert!(!request_conversion_enabled_for_transport(
|
||||
&disabled,
|
||||
"claude:messages",
|
||||
"openai:responses",
|
||||
));
|
||||
assert_eq!(
|
||||
candidate_transport_pair_skip_reason(&disabled, "claude:messages"),
|
||||
Some("format_conversion_disabled")
|
||||
);
|
||||
|
||||
let provider_enabled =
|
||||
transport_snapshot("custom", "openai:responses", "bearer", true, None);
|
||||
assert!(request_conversion_enabled_for_transport(
|
||||
&provider_enabled,
|
||||
"claude:messages",
|
||||
"openai:responses",
|
||||
));
|
||||
assert_eq!(
|
||||
candidate_transport_pair_skip_reason(&provider_enabled, "claude:messages"),
|
||||
None
|
||||
);
|
||||
|
||||
let endpoint_enabled = transport_snapshot(
|
||||
"custom",
|
||||
"openai:responses",
|
||||
"bearer",
|
||||
false,
|
||||
Some(json!({
|
||||
"enabled": true,
|
||||
"accept_formats": ["claude:messages"],
|
||||
})),
|
||||
);
|
||||
assert!(request_conversion_enabled_for_transport(
|
||||
&endpoint_enabled,
|
||||
"claude:messages",
|
||||
"openai:responses",
|
||||
));
|
||||
assert_eq!(
|
||||
candidate_transport_pair_skip_reason(&endpoint_enabled, "claude:messages"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[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 antigravity_gemini_transport_supports_standard_cross_format_conversion_via_envelope() {
|
||||
let transport = transport_snapshot(
|
||||
"antigravity",
|
||||
"gemini:generate_content",
|
||||
"oauth",
|
||||
true,
|
||||
None,
|
||||
);
|
||||
|
||||
assert!(request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"openai:chat",
|
||||
"gemini:generate_content"
|
||||
));
|
||||
assert!(request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"openai:responses",
|
||||
"gemini:generate_content"
|
||||
));
|
||||
assert!(request_conversion_transport_supported(
|
||||
&transport,
|
||||
RequestConversionKind::ToGeminiStandard
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_gemini_embedding_transport_supports_openai_embedding_conversion() {
|
||||
let transport = transport_snapshot("vertex_ai", "gemini:embedding", "api_key", true, None);
|
||||
|
||||
assert!(request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"openai:embedding",
|
||||
"gemini:embedding"
|
||||
));
|
||||
assert_eq!(
|
||||
request_pair_direct_auth(&transport, "gemini:embedding"),
|
||||
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 windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() {
|
||||
let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None);
|
||||
transport.key.decrypted_api_key = "devin-session-token$abc".to_string();
|
||||
transport.key.decrypted_auth_config = Some(r#"{"provider_type":"windsurf"}"#.to_string());
|
||||
|
||||
assert!(request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"claude:messages",
|
||||
"openai:chat"
|
||||
));
|
||||
assert!(request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"openai:responses",
|
||||
"openai:chat"
|
||||
));
|
||||
assert!(request_conversion_transport_supported(
|
||||
&transport,
|
||||
RequestConversionKind::ToOpenAIChat
|
||||
));
|
||||
assert_eq!(
|
||||
request_conversion_direct_auth(&transport, RequestConversionKind::ToOpenAIChat),
|
||||
Some((
|
||||
"authorization".to_string(),
|
||||
"Bearer devin-session-token$abc".to_string()
|
||||
))
|
||||
);
|
||||
}
|
||||
|
||||
#[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 candidate_common_transport_policy_allows_model_directive_base_model() {
|
||||
let mut transport = transport_snapshot("custom", "openai:responses", "bearer", true, None);
|
||||
transport.key.allowed_models = Some(vec!["gpt-5.5".to_string()]);
|
||||
|
||||
assert_eq!(
|
||||
candidate_common_transport_skip_reason(
|
||||
&transport,
|
||||
CandidateTransportPolicyFacts {
|
||||
endpoint_api_format: "openai:responses",
|
||||
global_model_name: "gpt-5",
|
||||
selected_provider_model_name: "provider-gpt-5",
|
||||
mapping_matched_model: None,
|
||||
},
|
||||
Some("gpt-5.5-xhigh"),
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[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 responses_key_permission_covers_search_without_changing_endpoint_identity() {
|
||||
let mut transport = transport_snapshot("custom", "openai:search", "bearer", true, None);
|
||||
transport.key.api_formats = Some(vec!["openai:responses".to_string()]);
|
||||
|
||||
assert_eq!(
|
||||
candidate_common_transport_skip_reason(
|
||||
&transport,
|
||||
candidate_facts("openai:search"),
|
||||
None,
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
candidate_common_transport_skip_reason(
|
||||
&transport,
|
||||
candidate_facts("openai:responses"),
|
||||
None,
|
||||
),
|
||||
Some("endpoint_api_format_changed")
|
||||
);
|
||||
}
|
||||
|
||||
#[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
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,511 @@
|
||||
use aether_ai_formats::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::grok::grok_browser_resolved_transport_profile_from_auth_config;
|
||||
use crate::network::{
|
||||
resolve_transport_profile, resolve_transport_profile_id, 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_transport_profile_id = resolve_transport_profile_id(transport);
|
||||
let resolved_transport_profile = resolve_transport_profile(transport)
|
||||
.and_then(|profile| serde_json::to_value(profile).ok())
|
||||
.unwrap_or(Value::Null);
|
||||
let configured_key_transport_profile = transport
|
||||
.key
|
||||
.fingerprint
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get("transport_profile"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null);
|
||||
let configured_provider_transport_profile = transport
|
||||
.provider
|
||||
.config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("fingerprint"))
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get("transport_profile"))
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null);
|
||||
let configured_legacy_grok_transport_profile = if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok")
|
||||
{
|
||||
transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.and_then(|auth_config| {
|
||||
grok_browser_resolved_transport_profile_from_auth_config(
|
||||
&auth_config,
|
||||
"grok_auth_config",
|
||||
)
|
||||
.and_then(|profile| serde_json::to_value(profile).ok())
|
||||
})
|
||||
.unwrap_or(Value::Null)
|
||||
} else {
|
||||
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_key_transport_profile": configured_key_transport_profile,
|
||||
"configured_provider_transport_profile": configured_provider_transport_profile,
|
||||
"configured_legacy_grok_transport_profile": configured_legacy_grok_transport_profile,
|
||||
"resolved_transport_profile_id": resolved_transport_profile_id,
|
||||
"resolved_transport_profile": resolved_transport_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,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: Some(json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "chrome_136",
|
||||
"header_fingerprint": {
|
||||
"user_agent": "Mozilla/5.0"
|
||||
}
|
||||
}
|
||||
})),
|
||||
upstream_metadata: None,
|
||||
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,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: 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"]["transport_profile"]["profile_id"],
|
||||
"chrome_136"
|
||||
);
|
||||
assert_eq!(diagnostics["resolved_transport_profile_id"], "chrome_136");
|
||||
assert_eq!(
|
||||
diagnostics["request_pair"]["conversion_enabled"],
|
||||
Value::Bool(true)
|
||||
);
|
||||
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())
|
||||
);
|
||||
}
|
||||
|
||||
fn sample_grok_transport_with_legacy_user_agent() -> GatewayProviderTransportSnapshot {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token",
|
||||
"user_agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/137.0.0.0 Safari/537.36"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
transport
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_diagnostics_include_legacy_grok_transport_profile() {
|
||||
let diagnostics = build_transport_diagnostics(
|
||||
&sample_grok_transport_with_legacy_user_agent(),
|
||||
"openai:chat",
|
||||
"openai:chat",
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
diagnostics["configured_legacy_grok_transport_profile"]["profile_id"],
|
||||
"chrome137"
|
||||
);
|
||||
assert_eq!(
|
||||
diagnostics["resolved_transport_profile"]["profile_id"],
|
||||
"chrome137"
|
||||
);
|
||||
assert_eq!(
|
||||
diagnostics["resolved_transport_profile"]["backend"],
|
||||
"browser_wreq"
|
||||
);
|
||||
}
|
||||
|
||||
#[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:[email protected]: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");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,272 @@
|
||||
use serde_json::Value;
|
||||
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
pub const GEMINI_CLI_PROVIDER_TYPE: &str = "gemini_cli";
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
pub struct GeminiCliRequestAuth {
|
||||
pub project_id: Option<String>,
|
||||
pub session_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GeminiCliRequestAuthSupport {
|
||||
Supported(GeminiCliRequestAuth),
|
||||
Unsupported(GeminiCliRequestAuthUnsupportedReason),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GeminiCliRequestAuthUnsupportedReason {
|
||||
WrongProviderType,
|
||||
InvalidAuthConfigJson,
|
||||
}
|
||||
|
||||
pub fn is_gemini_cli_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(GEMINI_CLI_PROVIDER_TYPE)
|
||||
}
|
||||
|
||||
pub fn resolve_local_gemini_cli_request_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> GeminiCliRequestAuthSupport {
|
||||
if !is_gemini_cli_provider_transport(transport) {
|
||||
return GeminiCliRequestAuthSupport::Unsupported(
|
||||
GeminiCliRequestAuthUnsupportedReason::WrongProviderType,
|
||||
);
|
||||
}
|
||||
|
||||
let metadata = transport.key.upstream_metadata.as_ref();
|
||||
let metadata_project_id = metadata.and_then(resolve_project_id_from_value);
|
||||
let metadata_session_id = metadata.and_then(resolve_session_id_from_value);
|
||||
|
||||
let Some(raw_auth_config) = transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return GeminiCliRequestAuthSupport::Supported(GeminiCliRequestAuth {
|
||||
project_id: metadata_project_id,
|
||||
session_id: metadata_session_id,
|
||||
});
|
||||
};
|
||||
|
||||
let Ok(auth_config) = serde_json::from_str::<Value>(raw_auth_config) else {
|
||||
return GeminiCliRequestAuthSupport::Unsupported(
|
||||
GeminiCliRequestAuthUnsupportedReason::InvalidAuthConfigJson,
|
||||
);
|
||||
};
|
||||
|
||||
GeminiCliRequestAuthSupport::Supported(GeminiCliRequestAuth {
|
||||
project_id: metadata_project_id.or_else(|| resolve_project_id_from_value(&auth_config)),
|
||||
session_id: metadata_session_id.or_else(|| resolve_session_id_from_value(&auth_config)),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn resolve_gemini_cli_project_id(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<String> {
|
||||
match resolve_local_gemini_cli_request_auth(transport) {
|
||||
GeminiCliRequestAuthSupport::Supported(auth) => auth.project_id,
|
||||
GeminiCliRequestAuthSupport::Unsupported(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_project_id_from_value(value: &Value) -> Option<String> {
|
||||
find_string_by_paths(
|
||||
value,
|
||||
&[
|
||||
&["project"],
|
||||
&["project_id"],
|
||||
&["projectId"],
|
||||
&["project", "id"],
|
||||
&["project", "project_id"],
|
||||
&["project", "projectId"],
|
||||
&["cloudaicompanionProject"],
|
||||
&["cloudaicompanionProject", "id"],
|
||||
&["cloudaicompanion_project"],
|
||||
&["cloudaicompanion_project", "id"],
|
||||
&["gemini_cli", "project"],
|
||||
&["gemini_cli", "project_id"],
|
||||
&["gemini_cli", "projectId"],
|
||||
&["gemini_cli", "cloudaicompanionProject"],
|
||||
&["gemini_cli", "cloudaicompanionProject", "id"],
|
||||
&["geminiCli", "project"],
|
||||
&["geminiCli", "project_id"],
|
||||
&["geminiCli", "projectId"],
|
||||
&["metadata", "project"],
|
||||
&["metadata", "project_id"],
|
||||
&["metadata", "projectId"],
|
||||
],
|
||||
)
|
||||
}
|
||||
|
||||
fn resolve_session_id_from_value(value: &Value) -> Option<String> {
|
||||
find_string_by_paths(
|
||||
value,
|
||||
&[
|
||||
&["session_id"],
|
||||
&["sessionId"],
|
||||
&["gemini_cli", "session_id"],
|
||||
&["gemini_cli", "sessionId"],
|
||||
&["geminiCli", "session_id"],
|
||||
&["geminiCli", "sessionId"],
|
||||
&["metadata", "session_id"],
|
||||
&["metadata", "sessionId"],
|
||||
],
|
||||
)
|
||||
}
|
||||
|
||||
fn find_string_by_paths(value: &Value, paths: &[&[&str]]) -> Option<String> {
|
||||
for path in paths {
|
||||
let mut current = value;
|
||||
let mut matched = true;
|
||||
for segment in *path {
|
||||
let Some(next) = current.get(*segment) else {
|
||||
matched = false;
|
||||
break;
|
||||
};
|
||||
current = next;
|
||||
}
|
||||
if !matched {
|
||||
continue;
|
||||
}
|
||||
if let Some(string) = current
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|item| !item.is_empty())
|
||||
{
|
||||
return Some(string.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
resolve_gemini_cli_project_id, resolve_local_gemini_cli_request_auth, GeminiCliRequestAuth,
|
||||
GeminiCliRequestAuthSupport,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Gemini CLI".to_string(),
|
||||
provider_type: "gemini_cli".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: true,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "gemini:generate_content".to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: "https://cloudcode-pa.googleapis.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "oauth".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__oauth__".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extracts_project_and_session_metadata_when_available() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"cloudaicompanionProject": {"id": "project-123"},
|
||||
"metadata": {"sessionId": "session-123"}
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
resolve_local_gemini_cli_request_auth(&transport),
|
||||
GeminiCliRequestAuthSupport::Supported(GeminiCliRequestAuth {
|
||||
project_id: Some("project-123".to_string()),
|
||||
session_id: Some("session-123".to_string()),
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn project_id_prefers_upstream_metadata_then_auth_config() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.upstream_metadata = Some(serde_json::json!({
|
||||
"gemini_cli": {
|
||||
"project_id": "metadata-project"
|
||||
}
|
||||
}));
|
||||
transport.key.decrypted_auth_config = Some(r#"{"project_id":"auth-project"}"#.to_string());
|
||||
|
||||
assert_eq!(
|
||||
resolve_gemini_cli_project_id(&transport).as_deref(),
|
||||
Some("metadata-project")
|
||||
);
|
||||
|
||||
transport.key.upstream_metadata = None;
|
||||
assert_eq!(
|
||||
resolve_gemini_cli_project_id(&transport).as_deref(),
|
||||
Some("auth-project")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_metadata_still_supports_request_envelope() {
|
||||
let transport = sample_transport();
|
||||
|
||||
assert_eq!(
|
||||
resolve_local_gemini_cli_request_auth(&transport),
|
||||
GeminiCliRequestAuthSupport::Supported(GeminiCliRequestAuth {
|
||||
project_id: None,
|
||||
session_id: None,
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
mod auth;
|
||||
mod policy;
|
||||
mod request;
|
||||
mod url;
|
||||
|
||||
pub use auth::{
|
||||
is_gemini_cli_provider_transport, resolve_gemini_cli_project_id,
|
||||
resolve_local_gemini_cli_request_auth, GeminiCliRequestAuth, GeminiCliRequestAuthSupport,
|
||||
GeminiCliRequestAuthUnsupportedReason, GEMINI_CLI_PROVIDER_TYPE,
|
||||
};
|
||||
pub use policy::gemini_cli_v1internal_requires_upstream_streaming;
|
||||
pub use request::{
|
||||
build_gemini_cli_v1internal_request, classify_gemini_cli_v1internal_request_body,
|
||||
GeminiCliRequestEnvelopeSupport, GeminiCliRequestEnvelopeUnsupportedReason,
|
||||
};
|
||||
pub use url::{
|
||||
build_gemini_cli_v1internal_url, GeminiCliRequestUrlAction,
|
||||
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT,
|
||||
GEMINI_CLI_V1INTERNAL_PATH_TEMPLATE,
|
||||
};
|
||||
|
||||
pub const GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME: &str = "gemini_cli:v1internal";
|
||||
@@ -0,0 +1,29 @@
|
||||
pub fn gemini_cli_v1internal_requires_upstream_streaming(
|
||||
provider_api_format: &str,
|
||||
client_requires_streaming: bool,
|
||||
) -> bool {
|
||||
client_requires_streaming
|
||||
&& aether_ai_formats::normalize_api_format_alias(provider_api_format)
|
||||
== "gemini:generate_content"
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::gemini_cli_v1internal_requires_upstream_streaming;
|
||||
|
||||
#[test]
|
||||
fn v1internal_generate_content_requires_upstream_streaming_for_stream_clients() {
|
||||
assert!(gemini_cli_v1internal_requires_upstream_streaming(
|
||||
"gemini:generate_content",
|
||||
true,
|
||||
));
|
||||
assert!(!gemini_cli_v1internal_requires_upstream_streaming(
|
||||
"gemini:generate_content",
|
||||
false,
|
||||
));
|
||||
assert!(!gemini_cli_v1internal_requires_upstream_streaming(
|
||||
"openai:chat",
|
||||
true,
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,243 @@
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::auth::GeminiCliRequestAuth;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum GeminiCliRequestEnvelopeSupport {
|
||||
Supported(Value),
|
||||
Unsupported(GeminiCliRequestEnvelopeUnsupportedReason),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GeminiCliRequestEnvelopeUnsupportedReason {
|
||||
NonObjectBody,
|
||||
MissingContents,
|
||||
MissingUserPromptId,
|
||||
MissingModel,
|
||||
}
|
||||
|
||||
pub fn classify_gemini_cli_v1internal_request_body(
|
||||
request_body: &Value,
|
||||
) -> Result<(), GeminiCliRequestEnvelopeUnsupportedReason> {
|
||||
let Value::Object(map) = request_body else {
|
||||
return Err(GeminiCliRequestEnvelopeUnsupportedReason::NonObjectBody);
|
||||
};
|
||||
if !map.contains_key("contents") && existing_v1internal_request_object(map).is_none() {
|
||||
return Err(GeminiCliRequestEnvelopeUnsupportedReason::MissingContents);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn build_gemini_cli_v1internal_request(
|
||||
auth: &GeminiCliRequestAuth,
|
||||
user_prompt_id: &str,
|
||||
model: &str,
|
||||
request_body: &Value,
|
||||
) -> GeminiCliRequestEnvelopeSupport {
|
||||
if user_prompt_id.trim().is_empty() {
|
||||
return GeminiCliRequestEnvelopeSupport::Unsupported(
|
||||
GeminiCliRequestEnvelopeUnsupportedReason::MissingUserPromptId,
|
||||
);
|
||||
}
|
||||
if model.trim().is_empty() {
|
||||
return GeminiCliRequestEnvelopeSupport::Unsupported(
|
||||
GeminiCliRequestEnvelopeUnsupportedReason::MissingModel,
|
||||
);
|
||||
}
|
||||
if let Err(reason) = classify_gemini_cli_v1internal_request_body(request_body) {
|
||||
return GeminiCliRequestEnvelopeSupport::Unsupported(reason);
|
||||
}
|
||||
|
||||
let Value::Object(source) = request_body else {
|
||||
return GeminiCliRequestEnvelopeSupport::Unsupported(
|
||||
GeminiCliRequestEnvelopeUnsupportedReason::NonObjectBody,
|
||||
);
|
||||
};
|
||||
|
||||
let existing_request = existing_v1internal_request_object(source);
|
||||
let mut inner_request: Map<String, Value> =
|
||||
existing_request.cloned().unwrap_or_else(|| source.clone());
|
||||
sanitize_inner_request(&mut inner_request);
|
||||
maybe_insert_session_id(&mut inner_request, auth.session_id.as_deref());
|
||||
|
||||
let project = non_empty_string_field(source, "project")
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| {
|
||||
auth.project_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
});
|
||||
let user_prompt_id = non_empty_string_field(source, "user_prompt_id")
|
||||
.or_else(|| non_empty_string_field(source, "userPromptId"))
|
||||
.unwrap_or(user_prompt_id)
|
||||
.trim()
|
||||
.to_string();
|
||||
|
||||
let mut envelope = Map::new();
|
||||
envelope.insert("model".to_string(), Value::String(model.trim().to_string()));
|
||||
if let Some(project) = project {
|
||||
envelope.insert("project".to_string(), Value::String(project));
|
||||
}
|
||||
envelope.insert("user_prompt_id".to_string(), Value::String(user_prompt_id));
|
||||
envelope.insert("request".to_string(), Value::Object(inner_request));
|
||||
|
||||
GeminiCliRequestEnvelopeSupport::Supported(Value::Object(envelope))
|
||||
}
|
||||
|
||||
fn sanitize_inner_request(inner_request: &mut Map<String, Value>) {
|
||||
inner_request.remove("model");
|
||||
inner_request.remove("stream");
|
||||
}
|
||||
|
||||
fn maybe_insert_session_id(inner_request: &mut Map<String, Value>, session_id: Option<&str>) {
|
||||
let Some(session_id) = session_id.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||
return;
|
||||
};
|
||||
if inner_request.contains_key("session_id") || inner_request.contains_key("sessionId") {
|
||||
return;
|
||||
}
|
||||
inner_request.insert(
|
||||
"session_id".to_string(),
|
||||
Value::String(session_id.to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
fn existing_v1internal_request_object(source: &Map<String, Value>) -> Option<&Map<String, Value>> {
|
||||
source
|
||||
.get("request")
|
||||
.and_then(Value::as_object)
|
||||
.filter(|request| request.contains_key("contents"))
|
||||
}
|
||||
|
||||
fn non_empty_string_field<'a>(source: &'a Map<String, Value>, key: &str) -> Option<&'a str> {
|
||||
source
|
||||
.get(key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
build_gemini_cli_v1internal_request, classify_gemini_cli_v1internal_request_body,
|
||||
GeminiCliRequestEnvelopeSupport,
|
||||
};
|
||||
use crate::gemini_cli::GeminiCliRequestAuth;
|
||||
|
||||
fn sample_auth() -> GeminiCliRequestAuth {
|
||||
GeminiCliRequestAuth {
|
||||
project_id: Some("project-123".to_string()),
|
||||
session_id: Some("session-123".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn wraps_generate_content_body_in_gemini_cli_v1internal_envelope() {
|
||||
let request_body = json!({
|
||||
"model": "client-model",
|
||||
"contents": [
|
||||
{"role": "user", "parts": [{"text": "hello"}]}
|
||||
],
|
||||
"stream": true,
|
||||
"generationConfig": {"temperature": 0.2}
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
classify_gemini_cli_v1internal_request_body(&request_body),
|
||||
Ok(())
|
||||
);
|
||||
assert_eq!(
|
||||
build_gemini_cli_v1internal_request(
|
||||
&sample_auth(),
|
||||
"trace-123",
|
||||
"gemini-2.5-pro",
|
||||
&request_body,
|
||||
),
|
||||
GeminiCliRequestEnvelopeSupport::Supported(json!({
|
||||
"model": "gemini-2.5-pro",
|
||||
"project": "project-123",
|
||||
"user_prompt_id": "trace-123",
|
||||
"request": {
|
||||
"contents": [
|
||||
{"role": "user", "parts": [{"text": "hello"}]}
|
||||
],
|
||||
"generationConfig": {"temperature": 0.2},
|
||||
"session_id": "session-123"
|
||||
}
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_existing_v1internal_request_shape_without_antigravity_fields() {
|
||||
let request_body = json!({
|
||||
"model": "old-model",
|
||||
"project": "project-from-body",
|
||||
"user_prompt_id": "prompt-from-body",
|
||||
"request": {
|
||||
"model": "nested-client-model",
|
||||
"contents": [
|
||||
{"role": "user", "parts": [{"text": "hello"}]}
|
||||
],
|
||||
"stream": false,
|
||||
"labels": {"source": "test"}
|
||||
},
|
||||
"userAgent": "antigravity",
|
||||
"requestType": "agent"
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
build_gemini_cli_v1internal_request(
|
||||
&sample_auth(),
|
||||
"trace-123",
|
||||
"gemini-2.5-pro",
|
||||
&request_body,
|
||||
),
|
||||
GeminiCliRequestEnvelopeSupport::Supported(json!({
|
||||
"model": "gemini-2.5-pro",
|
||||
"project": "project-from-body",
|
||||
"user_prompt_id": "prompt-from-body",
|
||||
"request": {
|
||||
"contents": [
|
||||
{"role": "user", "parts": [{"text": "hello"}]}
|
||||
],
|
||||
"labels": {"source": "test"},
|
||||
"session_id": "session-123"
|
||||
}
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn omits_optional_project_and_session_when_metadata_is_absent() {
|
||||
let request_body = json!({
|
||||
"contents": [
|
||||
{"role": "user", "parts": [{"text": "hello"}]}
|
||||
]
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
build_gemini_cli_v1internal_request(
|
||||
&GeminiCliRequestAuth::default(),
|
||||
"trace-123",
|
||||
"gemini-2.5-pro",
|
||||
&request_body,
|
||||
),
|
||||
GeminiCliRequestEnvelopeSupport::Supported(json!({
|
||||
"model": "gemini-2.5-pro",
|
||||
"user_prompt_id": "trace-123",
|
||||
"request": {
|
||||
"contents": [
|
||||
{"role": "user", "parts": [{"text": "hello"}]}
|
||||
]
|
||||
}
|
||||
}))
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use url::form_urlencoded;
|
||||
|
||||
pub const GEMINI_CLI_USER_AGENT: &str = "GeminiCLI/0.1.5 (Windows; AMD64)";
|
||||
pub const GEMINI_CLI_V1INTERNAL_PATH_TEMPLATE: &str = "/v1internal:{action}";
|
||||
pub const GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH: &str = "/v1internal:retrieveUserQuota";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum GeminiCliRequestUrlAction {
|
||||
GenerateContent,
|
||||
StreamGenerateContent,
|
||||
}
|
||||
|
||||
impl GeminiCliRequestUrlAction {
|
||||
fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::GenerateContent => "generateContent",
|
||||
Self::StreamGenerateContent => "streamGenerateContent",
|
||||
}
|
||||
}
|
||||
|
||||
fn is_stream(self) -> bool {
|
||||
matches!(self, Self::StreamGenerateContent)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_gemini_cli_v1internal_url(
|
||||
base_url: &str,
|
||||
action: GeminiCliRequestUrlAction,
|
||||
query: Option<&BTreeMap<String, String>>,
|
||||
) -> Option<String> {
|
||||
let trimmed_base = base_url.trim();
|
||||
if trimmed_base.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let path = GEMINI_CLI_V1INTERNAL_PATH_TEMPLATE.replace("{action}", action.as_str());
|
||||
let mut url = format!("{}{}", trimmed_base.trim_end_matches('/'), path);
|
||||
|
||||
let mut params = BTreeMap::new();
|
||||
if let Some(query) = query {
|
||||
for (key, value) in query {
|
||||
let key = key.trim();
|
||||
let value = value.trim();
|
||||
if key.is_empty()
|
||||
|| value.is_empty()
|
||||
|| key.eq_ignore_ascii_case("beta")
|
||||
|| key.eq_ignore_ascii_case("key")
|
||||
{
|
||||
continue;
|
||||
}
|
||||
params.insert(key.to_string(), value.to_string());
|
||||
}
|
||||
}
|
||||
if action.is_stream() {
|
||||
params
|
||||
.entry(String::from("alt"))
|
||||
.or_insert_with(|| String::from("sse"));
|
||||
}
|
||||
|
||||
if !params.is_empty() {
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in params {
|
||||
serializer.append_pair(key.as_str(), value.as_str());
|
||||
}
|
||||
let query_string = serializer.finish();
|
||||
if !query_string.is_empty() {
|
||||
url.push('?');
|
||||
url.push_str(&query_string);
|
||||
}
|
||||
}
|
||||
|
||||
Some(url)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::{
|
||||
build_gemini_cli_v1internal_url, GeminiCliRequestUrlAction,
|
||||
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn builds_gemini_cli_stream_url_with_alt_sse() {
|
||||
let query = BTreeMap::from([
|
||||
("foo".to_string(), "bar".to_string()),
|
||||
("key".to_string(), "blocked".to_string()),
|
||||
]);
|
||||
|
||||
assert_eq!(
|
||||
build_gemini_cli_v1internal_url(
|
||||
"https://cloudcode-pa.googleapis.com/",
|
||||
GeminiCliRequestUrlAction::StreamGenerateContent,
|
||||
Some(&query),
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://cloudcode-pa.googleapis.com/v1internal:streamGenerateContent?alt=sse&foo=bar")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_gemini_cli_sync_url_without_stream_query() {
|
||||
assert_eq!(
|
||||
build_gemini_cli_v1internal_url(
|
||||
"https://cloudcode-pa.googleapis.com",
|
||||
GeminiCliRequestUrlAction::GenerateContent,
|
||||
None,
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://cloudcode-pa.googleapis.com/v1internal:generateContent")
|
||||
);
|
||||
assert_eq!(
|
||||
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH,
|
||||
"/v1internal:retrieveUserQuota"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,283 @@
|
||||
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_with_request_headers, apply_local_header_rules_with_request_headers,
|
||||
body_rules_have_enabled_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>,
|
||||
request_headers: Option<&http::HeaderMap>,
|
||||
) -> 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_have_enabled_rules(body_rules) {
|
||||
return Err(GeminiFilesRequestBodyError::BodyRulesUnsupportedForBinaryUpload);
|
||||
}
|
||||
if let Some(body) = provider_request_body.as_mut() {
|
||||
if !apply_local_body_rules_with_request_headers(
|
||||
body,
|
||||
body_rules,
|
||||
Some(body_json),
|
||||
request_headers,
|
||||
) {
|
||||
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_with_request_headers(
|
||||
&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),
|
||||
Some(input.headers),
|
||||
) {
|
||||
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,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn 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"}
|
||||
])),
|
||||
None,
|
||||
)
|
||||
.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}])),
|
||||
None,
|
||||
),
|
||||
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")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,497 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_oauth::provider::providers::{
|
||||
GenericProviderOAuthAdapter, GENERIC_PROVIDER_OAUTH_TEMPLATES,
|
||||
};
|
||||
use aether_oauth::provider::{ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthTokenSet};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::oauth_refresh::{
|
||||
oauth_error_to_local_refresh_error, provider_oauth_transport_context_from_snapshot,
|
||||
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter, LocalOAuthRefreshError,
|
||||
LocalResolvedOAuthRequestAuth, ProviderOAuthLocalHttpExecutor,
|
||||
};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
const AUTH_HEADER_NAME: &str = "authorization";
|
||||
const OAUTH_REFRESH_SKEW_SECS: u64 = 120;
|
||||
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
|
||||
|
||||
pub fn supports_local_generic_oauth_request_auth_resolution(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
&& generic_provider_type(transport.provider.provider_type.as_str()).is_some()
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct GenericOAuthRefreshAdapter {
|
||||
token_url_overrides: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl GenericOAuthRefreshAdapter {
|
||||
pub fn with_token_url_for_tests(
|
||||
mut self,
|
||||
provider_type: &str,
|
||||
token_url: impl Into<String>,
|
||||
) -> Self {
|
||||
self.token_url_overrides
|
||||
.insert(provider_type.trim().to_ascii_lowercase(), token_url.into());
|
||||
self
|
||||
}
|
||||
|
||||
fn adapter_for_provider_type(
|
||||
&self,
|
||||
provider_type: &'static str,
|
||||
) -> Option<GenericProviderOAuthAdapter> {
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?;
|
||||
if let Some(token_url) = self.token_url_overrides.get(provider_type) {
|
||||
return Some(adapter.with_token_url_override(token_url.clone()));
|
||||
}
|
||||
Some(adapter)
|
||||
}
|
||||
|
||||
fn auth_config_from_transport(transport: &GatewayProviderTransportSnapshot) -> Option<Value> {
|
||||
transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||
}
|
||||
|
||||
fn auth_config_from_entry(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<Value> {
|
||||
entry
|
||||
.metadata
|
||||
.as_ref()
|
||||
.filter(|_| {
|
||||
entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(transport.provider.provider_type.as_str())
|
||||
})
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn auth_config_updated_at(auth_config: &Value) -> Option<u64> {
|
||||
auth_config
|
||||
.as_object()
|
||||
.and_then(|object| object.get("updated_at"))
|
||||
.and_then(|value| parse_u64_value(Some(value)))
|
||||
}
|
||||
|
||||
fn base_auth_config(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<Value> {
|
||||
let cached = entry.and_then(|cached| Self::auth_config_from_entry(transport, cached));
|
||||
let transport_auth = Self::auth_config_from_transport(transport);
|
||||
|
||||
match (cached, transport_auth) {
|
||||
(Some(cached), Some(transport_auth)) => {
|
||||
let cached_updated_at = Self::auth_config_updated_at(&cached);
|
||||
let transport_updated_at = Self::auth_config_updated_at(&transport_auth);
|
||||
if transport_updated_at > cached_updated_at {
|
||||
Some(transport_auth)
|
||||
} else {
|
||||
Some(cached)
|
||||
}
|
||||
}
|
||||
(Some(cached), None) => Some(cached),
|
||||
(None, Some(transport_auth)) => Some(transport_auth),
|
||||
(None, None) => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_direct_header(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
if !supports_local_generic_oauth_request_auth_resolution(transport) {
|
||||
return None;
|
||||
}
|
||||
|
||||
if let Some(value) =
|
||||
auth_config_authorization_header(transport.key.decrypted_auth_config.as_deref())
|
||||
{
|
||||
return Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: AUTH_HEADER_NAME.to_string(),
|
||||
value,
|
||||
});
|
||||
}
|
||||
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if secret.is_empty() || secret == PLACEHOLDER_API_KEY {
|
||||
return None;
|
||||
}
|
||||
|
||||
let auth_config = Self::auth_config_from_transport(transport);
|
||||
let refreshable = auth_config
|
||||
.as_ref()
|
||||
.and_then(refresh_token_from_auth_config)
|
||||
.is_some();
|
||||
if refreshable && auth_config_expires_soon(auth_config.as_ref()) {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: AUTH_HEADER_NAME.to_string(),
|
||||
value: format!("Bearer {secret}"),
|
||||
})
|
||||
}
|
||||
|
||||
fn build_cached_entry(
|
||||
provider_type: &'static str,
|
||||
refreshed: ProviderOAuthTokenSet,
|
||||
) -> CachedOAuthEntry {
|
||||
CachedOAuthEntry {
|
||||
provider_type: provider_type.to_string(),
|
||||
auth_header_name: AUTH_HEADER_NAME.to_string(),
|
||||
auth_header_value: refreshed.token_set.bearer_header_value(),
|
||||
expires_at_unix_secs: refreshed.token_set.expires_at_unix_secs,
|
||||
metadata: Some(refreshed.auth_config),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
"generic_oauth"
|
||||
}
|
||||
|
||||
fn supports(&self, transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
supports_local_generic_oauth_request_auth_resolution(transport)
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
if !entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(transport.provider.provider_type.as_str())
|
||||
{
|
||||
return None;
|
||||
}
|
||||
if let Some(value) =
|
||||
auth_config_authorization_header(transport.key.decrypted_auth_config.as_deref())
|
||||
{
|
||||
return Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: AUTH_HEADER_NAME.to_string(),
|
||||
value,
|
||||
});
|
||||
}
|
||||
if expires_at_requires_refresh(entry.expires_at_unix_secs) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let name = entry.auth_header_name.trim();
|
||||
let value = entry.auth_header_value.trim();
|
||||
if name.is_empty() || value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: name.to_ascii_lowercase(),
|
||||
value: value.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
self.resolve_direct_header(transport)
|
||||
}
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool {
|
||||
if !supports_local_generic_oauth_request_auth_resolution(transport) {
|
||||
return false;
|
||||
}
|
||||
if entry
|
||||
.and_then(|cached| self.resolve_cached(transport, cached))
|
||||
.is_some()
|
||||
|| self.resolve_direct_header(transport).is_some()
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
self.base_auth_config(transport, entry)
|
||||
.as_ref()
|
||||
.and_then(refresh_token_from_auth_config)
|
||||
.is_some()
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
|
||||
let Some(provider_type) = generic_provider_type(transport.provider.provider_type.as_str())
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(auth_config) = self.base_auth_config(transport, entry) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(refresh_token) = refresh_token_from_auth_config(&auth_config) else {
|
||||
tracing::warn!(
|
||||
key_id = %transport.key.id,
|
||||
provider_id = %transport.provider.id,
|
||||
provider_type,
|
||||
"gateway generic oauth refresh skipped because auth_config has no refresh_token"
|
||||
);
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(adapter) = self.adapter_for_provider_type(provider_type) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
key_id = %transport.key.id,
|
||||
provider_id = %transport.provider.id,
|
||||
endpoint_id = %transport.endpoint.id,
|
||||
provider_type,
|
||||
request_refresh_token_len = refresh_token.len(),
|
||||
"gateway generic oauth refresh delegated to provider oauth adapter"
|
||||
);
|
||||
|
||||
let oauth_executor =
|
||||
ProviderOAuthLocalHttpExecutor::new(provider_type, transport, executor);
|
||||
let ctx = provider_oauth_transport_context_from_snapshot(transport);
|
||||
let account = ProviderOAuthAccount {
|
||||
provider_type: provider_type.to_string(),
|
||||
access_token: current_access_token(transport, entry).unwrap_or_default(),
|
||||
expires_at_unix_secs: auth_config_expires_at(&auth_config),
|
||||
auth_config,
|
||||
identity: BTreeMap::new(),
|
||||
};
|
||||
let refreshed = adapter
|
||||
.refresh(&oauth_executor, &ctx, &account)
|
||||
.await
|
||||
.map_err(|error| oauth_error_to_local_refresh_error(provider_type, error))?;
|
||||
|
||||
tracing::info!(
|
||||
key_id = %transport.key.id,
|
||||
provider_id = %transport.provider.id,
|
||||
endpoint_id = %transport.endpoint.id,
|
||||
provider_type,
|
||||
expires_at_unix_secs = ?refreshed.token_set.expires_at_unix_secs,
|
||||
response_has_refresh_token = refreshed.token_set.refresh_token.is_some(),
|
||||
"gateway generic oauth refresh succeeded"
|
||||
);
|
||||
|
||||
Ok(Some(Self::build_cached_entry(provider_type, refreshed)))
|
||||
}
|
||||
}
|
||||
|
||||
fn generic_provider_type(provider_type: &str) -> Option<&'static str> {
|
||||
let normalized = provider_type.trim();
|
||||
GENERIC_PROVIDER_OAUTH_TEMPLATES
|
||||
.iter()
|
||||
.find(|template| normalized.eq_ignore_ascii_case(template.provider_type))
|
||||
.map(|template| template.provider_type)
|
||||
}
|
||||
|
||||
fn refresh_token_from_auth_config(auth_config: &Value) -> Option<String> {
|
||||
auth_config
|
||||
.as_object()
|
||||
.and_then(|object| object.get("refresh_token"))
|
||||
.and_then(non_empty_string)
|
||||
}
|
||||
|
||||
fn auth_config_expires_at(auth_config: &Value) -> Option<u64> {
|
||||
auth_config
|
||||
.as_object()
|
||||
.and_then(|object| object.get("expires_at"))
|
||||
.and_then(|value| parse_u64_value(Some(value)))
|
||||
}
|
||||
|
||||
fn auth_config_expires_soon(auth_config: Option<&Value>) -> bool {
|
||||
expires_at_requires_refresh(auth_config.and_then(auth_config_expires_at))
|
||||
}
|
||||
|
||||
fn expires_at_requires_refresh(expires_at_unix_secs: Option<u64>) -> bool {
|
||||
expires_at_unix_secs
|
||||
.map(|expires_at_unix_secs| {
|
||||
aether_oauth::core::current_unix_secs()
|
||||
>= expires_at_unix_secs.saturating_sub(OAUTH_REFRESH_SKEW_SECS)
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn parse_u64_value(value: Option<&Value>) -> Option<u64> {
|
||||
match value? {
|
||||
Value::Number(number) => number.as_u64(),
|
||||
Value::String(string) => string.trim().parse::<u64>().ok(),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn non_empty_string(value: &Value) -> Option<String> {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn auth_config_authorization_header(raw_auth_config: Option<&str>) -> Option<String> {
|
||||
let mut headers = BTreeMap::new();
|
||||
crate::auth_config::apply_local_auth_config_header_overrides(&mut headers, raw_auth_config);
|
||||
headers
|
||||
.remove(AUTH_HEADER_NAME)
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn current_access_token(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
entry
|
||||
.and_then(|entry| {
|
||||
entry
|
||||
.auth_header_value
|
||||
.trim()
|
||||
.strip_prefix("Bearer ")
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
.or_else(|| {
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
(!secret.is_empty() && secret != PLACEHOLDER_API_KEY).then(|| secret.to_string())
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::oauth_refresh::{
|
||||
CachedOAuthEntry, LocalOAuthRefreshAdapter, LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::GenericOAuthRefreshAdapter;
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Codex".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:responses".to_string(),
|
||||
api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("responses".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://chatgpt.com/backend-api/codex".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: "OAuth headers".to_string(),
|
||||
auth_type: "oauth".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: Some(
|
||||
json!({
|
||||
"provider_type": "codex",
|
||||
"access_token_import_temporary": true,
|
||||
"headers": {
|
||||
"Authorization": "Bearer imported-session",
|
||||
"Host": "blocked.example"
|
||||
}
|
||||
})
|
||||
.to_string(),
|
||||
),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_imported_authorization_header_without_api_key_secret() {
|
||||
let adapter = GenericOAuthRefreshAdapter::default();
|
||||
let auth = adapter
|
||||
.resolve_without_refresh(&sample_transport())
|
||||
.expect("auth_config authorization header should resolve");
|
||||
|
||||
assert_eq!(
|
||||
auth,
|
||||
LocalResolvedOAuthRequestAuth::Header {
|
||||
name: "authorization".to_string(),
|
||||
value: "Bearer imported-session".to_string(),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_config_authorization_header_overrides_cached_oauth_entry() {
|
||||
let adapter = GenericOAuthRefreshAdapter::default();
|
||||
let entry = CachedOAuthEntry {
|
||||
provider_type: "codex".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "Bearer refreshed-access-token".to_string(),
|
||||
expires_at_unix_secs: Some(u64::MAX),
|
||||
metadata: None,
|
||||
};
|
||||
let auth = adapter
|
||||
.resolve_cached(&sample_transport(), &entry)
|
||||
.expect("auth_config authorization header should override cache");
|
||||
|
||||
assert_eq!(
|
||||
auth,
|
||||
LocalResolvedOAuthRequestAuth::Header {
|
||||
name: "authorization".to_string(),
|
||||
value: "Bearer imported-session".to_string(),
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,304 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_contracts::USAGE_SERVER_NOW_UNIX_MS_HEADER;
|
||||
|
||||
pub fn should_skip_request_header(name: &str) -> bool {
|
||||
let normalized = name.to_ascii_lowercase();
|
||||
matches!(
|
||||
normalized.as_str(),
|
||||
"connection"
|
||||
| "keep-alive"
|
||||
| "proxy-authenticate"
|
||||
| "proxy-authorization"
|
||||
| "proxy-connection"
|
||||
| "te"
|
||||
| "trailer"
|
||||
| "transfer-encoding"
|
||||
| "upgrade"
|
||||
| "x-aether-execution-path"
|
||||
| "x-aether-dependency-reason"
|
||||
| "x-aether-execution-loop-guard"
|
||||
| "x-aether-control-execute-fallback"
|
||||
| "x-aether-rate-limit-preflight"
|
||||
| USAGE_SERVER_NOW_UNIX_MS_HEADER
|
||||
)
|
||||
}
|
||||
|
||||
pub fn should_skip_upstream_passthrough_header(name: &str) -> bool {
|
||||
let lower = name.to_ascii_lowercase();
|
||||
// Anthropic SDK (stainless) client metadata and Anthropic-specific headers
|
||||
// (anthropic-version / anthropic-beta / anthropic-dangerous-direct-browser-access / ...).
|
||||
// These are only meaningful when the upstream is Anthropic. The Claude Code
|
||||
// adapter re-reads what it needs directly from the original HeaderMap and
|
||||
// injects its own values *after* this filter runs, so stripping both prefixes
|
||||
// here prevents leakage to any other upstream (OpenAI/Gemini/Codex/...).
|
||||
if lower.starts_with("x-stainless-") || lower.starts_with("anthropic-") {
|
||||
return true;
|
||||
}
|
||||
matches!(
|
||||
lower.as_str(),
|
||||
"authorization"
|
||||
| "x-api-key"
|
||||
| "x-goog-api-key"
|
||||
| "host"
|
||||
| "content-length"
|
||||
| "transfer-encoding"
|
||||
| "connection"
|
||||
| "content-encoding"
|
||||
| "x-real-ip"
|
||||
| "x-real-proto"
|
||||
| "x-forwarded-for"
|
||||
| "x-forwarded-proto"
|
||||
| "x-forwarded-scheme"
|
||||
| "x-forwarded-host"
|
||||
| "x-forwarded-port"
|
||||
// Claude CLI client identifier; re-injected by the Claude Code adapter
|
||||
// when the upstream is Anthropic, filtered for everybody else.
|
||||
| "x-app"
|
||||
) || should_skip_request_header(name)
|
||||
}
|
||||
|
||||
pub(crate) fn should_skip_upstream_complete_passthrough_header(name: &str) -> bool {
|
||||
let lower = name.to_ascii_lowercase();
|
||||
matches!(
|
||||
lower.as_str(),
|
||||
"authorization"
|
||||
| "x-api-key"
|
||||
| "x-goog-api-key"
|
||||
| "host"
|
||||
| "content-length"
|
||||
| "transfer-encoding"
|
||||
| "connection"
|
||||
| "content-encoding"
|
||||
| "x-real-ip"
|
||||
| "x-real-proto"
|
||||
| "x-forwarded-for"
|
||||
| "x-forwarded-proto"
|
||||
| "x-forwarded-scheme"
|
||||
| "x-forwarded-host"
|
||||
| "x-forwarded-port"
|
||||
) || should_skip_request_header(name)
|
||||
}
|
||||
|
||||
pub fn normalize_upstream_accept_encoding(value: &str) -> Option<String> {
|
||||
let mut accepted = Vec::new();
|
||||
let mut wildcard_allowed = false;
|
||||
let mut gzip_disabled = false;
|
||||
let mut deflate_disabled = false;
|
||||
let mut identity_disabled = false;
|
||||
|
||||
for item in value.split(',') {
|
||||
let Some((token, normalized_item, enabled)) = parse_accept_encoding_item(item) else {
|
||||
continue;
|
||||
};
|
||||
match token.as_str() {
|
||||
"gzip" if enabled => accepted.push(normalized_item),
|
||||
"gzip" => gzip_disabled = true,
|
||||
"deflate" if enabled => accepted.push(normalized_item),
|
||||
"deflate" => deflate_disabled = true,
|
||||
"identity" if enabled => accepted.push(normalized_item),
|
||||
"identity" => identity_disabled = true,
|
||||
"*" if enabled => wildcard_allowed = true,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if !accepted.is_empty() {
|
||||
return Some(accepted.join(", "));
|
||||
}
|
||||
|
||||
if wildcard_allowed && !gzip_disabled {
|
||||
Some("gzip".to_string())
|
||||
} else if wildcard_allowed && !deflate_disabled {
|
||||
Some("deflate".to_string())
|
||||
} else if wildcard_allowed && !identity_disabled {
|
||||
Some("identity".to_string())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_accept_encoding_item(raw_item: &str) -> Option<(String, String, bool)> {
|
||||
let mut parts = raw_item.trim().split(';');
|
||||
let token = parts.next()?.trim().to_ascii_lowercase();
|
||||
if token.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut enabled = true;
|
||||
let mut normalized = token.clone();
|
||||
for raw_param in parts {
|
||||
let param = raw_param.trim();
|
||||
if param.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let Some((name, value)) = param.split_once('=') else {
|
||||
continue;
|
||||
};
|
||||
if name.trim().eq_ignore_ascii_case("q") {
|
||||
let value = value.trim();
|
||||
if q_value_is_zero(value) {
|
||||
enabled = false;
|
||||
continue;
|
||||
}
|
||||
normalized.push_str(";q=");
|
||||
normalized.push_str(value);
|
||||
}
|
||||
}
|
||||
|
||||
Some((token, normalized, enabled))
|
||||
}
|
||||
|
||||
fn q_value_is_zero(value: &str) -> bool {
|
||||
value
|
||||
.trim_matches('"')
|
||||
.parse::<f32>()
|
||||
.is_ok_and(|q| q <= 0.0)
|
||||
}
|
||||
|
||||
pub fn force_identity_accept_encoding(headers: &mut BTreeMap<String, String>) {
|
||||
if let Some(existing_key) = headers
|
||||
.keys()
|
||||
.find(|key| key.eq_ignore_ascii_case("accept-encoding"))
|
||||
.cloned()
|
||||
{
|
||||
headers.remove(&existing_key);
|
||||
}
|
||||
headers.insert("accept-encoding".to_string(), "identity".to_string());
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
force_identity_accept_encoding, normalize_upstream_accept_encoding,
|
||||
should_skip_request_header, should_skip_upstream_complete_passthrough_header,
|
||||
should_skip_upstream_passthrough_header,
|
||||
};
|
||||
use aether_contracts::USAGE_SERVER_NOW_UNIX_MS_HEADER;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn strips_all_stainless_headers() {
|
||||
let stainless = [
|
||||
"x-stainless-arch",
|
||||
"x-stainless-lang",
|
||||
"x-stainless-os",
|
||||
"x-stainless-package-version",
|
||||
"x-stainless-retry-count",
|
||||
"x-stainless-runtime",
|
||||
"x-stainless-runtime-version",
|
||||
"x-stainless-timeout",
|
||||
"x-stainless-helper-method",
|
||||
"X-Stainless-Arch",
|
||||
"X-STAINLESS-FUTURE-HEADER",
|
||||
];
|
||||
for h in stainless {
|
||||
assert!(
|
||||
should_skip_upstream_passthrough_header(h),
|
||||
"should skip {h}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accept_encoding_is_not_classified_as_hop_by_hop_passthrough_skip() {
|
||||
assert!(!should_skip_upstream_passthrough_header("accept-encoding"));
|
||||
assert!(!should_skip_upstream_complete_passthrough_header(
|
||||
"accept-encoding"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalizes_accept_encoding_to_supported_upstream_codecs() {
|
||||
assert_eq!(
|
||||
normalize_upstream_accept_encoding("gzip, br").as_deref(),
|
||||
Some("gzip")
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_upstream_accept_encoding("br, deflate").as_deref(),
|
||||
Some("deflate")
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_upstream_accept_encoding("identity").as_deref(),
|
||||
Some("identity")
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_upstream_accept_encoding("gzip;q=0.5, br").as_deref(),
|
||||
Some("gzip;q=0.5")
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_upstream_accept_encoding("gzip;q=0, br, deflate").as_deref(),
|
||||
Some("deflate")
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_upstream_accept_encoding("gzip;q=0, br").as_deref(),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_upstream_accept_encoding("*").as_deref(),
|
||||
Some("gzip")
|
||||
);
|
||||
assert_eq!(normalize_upstream_accept_encoding("br"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn force_identity_accept_encoding_replaces_existing_casing() {
|
||||
let mut headers = BTreeMap::from([("Accept-Encoding".to_string(), "gzip".to_string())]);
|
||||
|
||||
force_identity_accept_encoding(&mut headers);
|
||||
|
||||
assert_eq!(
|
||||
headers.get("accept-encoding").map(String::as_str),
|
||||
Some("identity")
|
||||
);
|
||||
assert!(!headers.contains_key("Accept-Encoding"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strips_anthropic_and_claude_cli_identity_headers() {
|
||||
let anthropic = [
|
||||
"anthropic-version",
|
||||
"anthropic-beta",
|
||||
"anthropic-dangerous-direct-browser-access",
|
||||
"Anthropic-Version",
|
||||
"ANTHROPIC-FUTURE-HEADER",
|
||||
"x-app",
|
||||
"X-App",
|
||||
];
|
||||
for h in anthropic {
|
||||
assert!(
|
||||
should_skip_upstream_passthrough_header(h),
|
||||
"should skip {h}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strips_usage_server_time_header_from_provider_requests() {
|
||||
for h in [
|
||||
USAGE_SERVER_NOW_UNIX_MS_HEADER,
|
||||
"X-Aether-Server-Now-Unix-Ms",
|
||||
] {
|
||||
assert!(should_skip_request_header(h), "should skip {h}");
|
||||
assert!(
|
||||
should_skip_upstream_passthrough_header(h),
|
||||
"should skip passthrough {h}"
|
||||
);
|
||||
assert!(
|
||||
should_skip_upstream_complete_passthrough_header(h),
|
||||
"should skip complete passthrough {h}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn allows_normal_headers_through() {
|
||||
let allowed = ["user-agent", "accept", "content-type", "x-custom-header"];
|
||||
for h in allowed {
|
||||
assert!(
|
||||
!should_skip_upstream_passthrough_header(h),
|
||||
"should allow {h}"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,353 @@
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::credentials::{generate_machine_id, KiroAuthConfig};
|
||||
|
||||
pub const PROVIDER_TYPE: &str = "kiro";
|
||||
pub const KIRO_AUTH_HEADER: &str = "authorization";
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct KiroBearerAuth {
|
||||
pub name: &'static str,
|
||||
pub value: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct KiroRequestAuth {
|
||||
pub name: &'static str,
|
||||
pub value: String,
|
||||
pub auth_config: KiroAuthConfig,
|
||||
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>,
|
||||
) -> Option<KiroRequestAuth> {
|
||||
let cached_token_needs_refresh = auth_config.cached_access_token_requires_refresh(120);
|
||||
let fallback_secret = fallback_secret
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty() && *value != "__placeholder__");
|
||||
if cached_token_needs_refresh && auth_config.can_refresh_access_token() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let token = auth_config
|
||||
.cached_access_token()
|
||||
.filter(|_| !cached_token_needs_refresh)
|
||||
.or(fallback_secret)?;
|
||||
let machine_id = generate_machine_id(&auth_config, Some(token))?;
|
||||
|
||||
Some(KiroRequestAuth {
|
||||
name: KIRO_AUTH_HEADER,
|
||||
value: format!("Bearer {token}"),
|
||||
auth_config,
|
||||
machine_id,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn resolve_local_kiro_bearer_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<KiroBearerAuth> {
|
||||
if !is_kiro_provider_transport(transport) {
|
||||
return None;
|
||||
}
|
||||
if transport.key.decrypted_auth_config.is_some() {
|
||||
return None;
|
||||
}
|
||||
if !kiro_auth_type_supported(transport.key.auth_type.as_str()) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if secret.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(KiroBearerAuth {
|
||||
name: KIRO_AUTH_HEADER,
|
||||
value: format!("Bearer {secret}"),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn supports_local_kiro_auth_prerequisites(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
resolve_local_kiro_bearer_auth(transport).is_some()
|
||||
}
|
||||
|
||||
pub fn resolve_local_kiro_request_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<KiroRequestAuth> {
|
||||
if !is_kiro_provider_transport(transport) {
|
||||
return None;
|
||||
}
|
||||
if !kiro_auth_type_supported(transport.key.auth_type.as_str()) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let auth_config = KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
|
||||
.unwrap_or(KiroAuthConfig {
|
||||
auth_method: None,
|
||||
refresh_token: None,
|
||||
expires_at: None,
|
||||
profile_arn: None,
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: None,
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: None,
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: None,
|
||||
});
|
||||
let fallback_secret = transport
|
||||
.key
|
||||
.decrypted_api_key
|
||||
.trim()
|
||||
.strip_prefix("__placeholder__")
|
||||
.map(|_| "")
|
||||
.unwrap_or(transport.key.decrypted_api_key.trim());
|
||||
build_kiro_request_auth_from_config(auth_config, Some(fallback_secret))
|
||||
}
|
||||
|
||||
pub fn supports_local_kiro_request_auth_resolution(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
resolve_local_kiro_request_auth(transport).is_some()
|
||||
|| KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
|
||||
.is_some_and(|auth_config| {
|
||||
is_kiro_provider_transport(transport)
|
||||
&& kiro_auth_type_supported(transport.key.auth_type.as_str())
|
||||
&& auth_config.can_refresh_access_token()
|
||||
})
|
||||
}
|
||||
|
||||
fn kiro_auth_type_supported(auth_type: &str) -> bool {
|
||||
matches!(
|
||||
auth_type.trim().to_ascii_lowercase().as_str(),
|
||||
"bearer" | "oauth"
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::{
|
||||
resolve_local_kiro_bearer_auth, resolve_local_kiro_request_auth,
|
||||
supports_local_kiro_auth_prerequisites, supports_local_kiro_request_auth_resolution,
|
||||
KIRO_AUTH_HEADER,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Kiro".to_string(),
|
||||
provider_type: "kiro".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "claude:messages".to_string(),
|
||||
api_family: Some("claude".to_string()),
|
||||
endpoint_kind: Some("cli".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://kiro.example".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: Some(vec!["claude:messages".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "upstream-key".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_bearer_auth_for_known_kiro_subset() {
|
||||
let auth = resolve_local_kiro_bearer_auth(&sample_transport())
|
||||
.expect("kiro bearer auth should resolve");
|
||||
assert_eq!(auth.name, KIRO_AUTH_HEADER);
|
||||
assert_eq!(auth.value, "Bearer upstream-key");
|
||||
assert!(supports_local_kiro_auth_prerequisites(&sample_transport()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_auth_config_subset() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_auth_config = Some("{\"mode\":\"custom\"}".to_string());
|
||||
assert!(resolve_local_kiro_bearer_auth(&transport).is_none());
|
||||
assert!(!supports_local_kiro_auth_prerequisites(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_non_bearer_subset() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "api_key".to_string();
|
||||
assert!(resolve_local_kiro_bearer_auth(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_request_auth_when_legacy_oauth_auth_type_is_used() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "oauth".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"access_token":"cached-token",
|
||||
"expires_at":4102444800,
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
|
||||
"machine_id":"123e4567-e89b-12d3-a456-426614174000",
|
||||
"api_region":"us-west-2"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let auth = resolve_local_kiro_request_auth(&transport)
|
||||
.expect("request auth should resolve from legacy oauth auth_type");
|
||||
assert_eq!(auth.value, "Bearer cached-token");
|
||||
assert!(supports_local_kiro_request_auth_resolution(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_request_auth_from_cached_access_token() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"access_token":"cached-token",
|
||||
"expires_at":4102444800,
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
|
||||
"machine_id":"123e4567-e89b-12d3-a456-426614174000",
|
||||
"api_region":"us-west-2"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let auth = resolve_local_kiro_request_auth(&transport)
|
||||
.expect("request auth should resolve from cached token");
|
||||
assert_eq!(auth.name, KIRO_AUTH_HEADER);
|
||||
assert_eq!(auth.value, "Bearer cached-token");
|
||||
assert_eq!(auth.auth_config.effective_api_region(), "us-west-2");
|
||||
assert_eq!(
|
||||
auth.machine_id,
|
||||
"123e4567e89b12d3a456426614174000123e4567e89b12d3a456426614174000"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn skips_expired_cached_access_token_without_fallback_secret() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"access_token":"expired-token",
|
||||
"expires_at": 1,
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(resolve_local_kiro_request_auth(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn skips_refreshable_cached_access_token_without_expiry() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"access_token":"cached-token-without-expiry",
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(resolve_local_kiro_request_auth(&transport).is_none());
|
||||
assert!(supports_local_kiro_request_auth_resolution(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn refreshable_expired_cached_access_token_does_not_fallback_to_decrypted_api_key() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "stale-upstream-token".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"access_token":"expired-token",
|
||||
"expires_at": 1,
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(resolve_local_kiro_request_auth(&transport).is_none());
|
||||
assert!(supports_local_kiro_request_auth_resolution(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_refresh_only_resolution_without_cached_access_token() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(resolve_local_kiro_request_auth(&transport).is_none());
|
||||
assert!(supports_local_kiro_request_auth_resolution(&transport));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,714 @@
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use serde_json::{json, Map, Value};
|
||||
use tracing::warn;
|
||||
use uuid::Uuid;
|
||||
|
||||
const SYSTEM_CHUNKED_POLICY: &str = "When the Write or Edit tool has content size limits, always comply silently. Never suggest bypassing these limits via alternative tools. Never ask the user whether to switch approaches. Complete all chunked operations without commentary.";
|
||||
const WRITE_TOOL_DESCRIPTION_SUFFIX: &str = "- IMPORTANT: If the content to write exceeds 150 lines, you MUST only write the first 50 lines using this tool, then use `Edit` tool to append the remaining content in chunks of no more than 50 lines each. If needed, leave a unique placeholder to help append content. Do NOT attempt to write all content at once.";
|
||||
const EDIT_TOOL_DESCRIPTION_SUFFIX: &str = "- IMPORTANT: If the `new_string` content exceeds 50 lines, you MUST split it into multiple Edit calls, each replacing no more than 50 lines at a time. If used to append content, leave a unique placeholder to help append content. On the final chunk, do NOT include the placeholder.";
|
||||
|
||||
pub fn convert_claude_messages_to_conversation_state(
|
||||
request_body: &Value,
|
||||
model: &str,
|
||||
) -> Option<Value> {
|
||||
let model_id = model.trim();
|
||||
if model_id.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let messages = request_body.get("messages")?.as_array()?;
|
||||
if messages.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let conversation_id = request_body
|
||||
.get("metadata")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| {
|
||||
metadata
|
||||
.get("user_id")
|
||||
.or_else(|| metadata.get("userId"))
|
||||
.and_then(Value::as_str)
|
||||
})
|
||||
.and_then(extract_session_id)
|
||||
.unwrap_or_else(|| Uuid::new_v4().to_string());
|
||||
let agent_continuation_id = Uuid::new_v4().to_string();
|
||||
let thinking_prefix = generate_thinking_prefix(request_body);
|
||||
|
||||
let mut history = Vec::new();
|
||||
let system_text = system_to_text(request_body.get("system"));
|
||||
if !system_text.is_empty() {
|
||||
history.push(json!({
|
||||
"userInputMessage": {
|
||||
"content": format!("{system_text}\n{SYSTEM_CHUNKED_POLICY}"),
|
||||
"modelId": model_id,
|
||||
"origin": "AI_EDITOR"
|
||||
}
|
||||
}));
|
||||
history.push(json!({
|
||||
"assistantResponseMessage": {
|
||||
"content": "I will follow these instructions."
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
let last_is_assistant = messages
|
||||
.last()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|message| message.get("role"))
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|role| role == "assistant");
|
||||
let history_end_index = if last_is_assistant {
|
||||
messages.len()
|
||||
} else {
|
||||
messages.len().saturating_sub(1)
|
||||
};
|
||||
|
||||
let mut user_buffer = Vec::new();
|
||||
for message in &messages[..history_end_index] {
|
||||
let Some(message) = message.as_object() else {
|
||||
continue;
|
||||
};
|
||||
match message.get("role").and_then(Value::as_str) {
|
||||
Some("user") => user_buffer.push(message),
|
||||
Some("assistant") => {
|
||||
if let Some(user_item) = flush_user_buffer(&mut user_buffer, model_id) {
|
||||
history.push(user_item);
|
||||
} else if history.is_empty()
|
||||
|| history
|
||||
.last()
|
||||
.and_then(Value::as_object)
|
||||
.is_some_and(|item| item.contains_key("assistantResponseMessage"))
|
||||
{
|
||||
history.push(json!({
|
||||
"userInputMessage": {
|
||||
"content": "Continue.",
|
||||
"modelId": model_id,
|
||||
"origin": "AI_EDITOR"
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
if let Some(assistant_item) = convert_assistant_message(message) {
|
||||
history.push(json!({"assistantResponseMessage": assistant_item}));
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(tail_user) = flush_user_buffer(&mut user_buffer, model_id) {
|
||||
history.push(tail_user);
|
||||
history.push(json!({"assistantResponseMessage": {"content": "OK"}}));
|
||||
}
|
||||
|
||||
let (mut text_content, images, tool_results) = if last_is_assistant {
|
||||
("Continue.".to_string(), Vec::new(), Vec::new())
|
||||
} else {
|
||||
let last = messages.last()?.as_object()?;
|
||||
if last.get("role").and_then(Value::as_str) != Some("user") {
|
||||
return None;
|
||||
}
|
||||
process_message_content(last.get("content"))
|
||||
};
|
||||
|
||||
let mut tools = convert_tools(request_body.get("tools"));
|
||||
let mut history_tool_names = BTreeSet::new();
|
||||
let mut history_tool_result_ids = BTreeSet::new();
|
||||
let mut history_tool_use_ids = BTreeSet::new();
|
||||
|
||||
for item in &history {
|
||||
let Some(item) = item.as_object() else {
|
||||
continue;
|
||||
};
|
||||
if let Some(user_input) = item.get("userInputMessage").and_then(Value::as_object) {
|
||||
if let Some(results) = user_input
|
||||
.get("userInputMessageContext")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|ctx| ctx.get("toolResults"))
|
||||
.and_then(Value::as_array)
|
||||
{
|
||||
for result in results {
|
||||
if let Some(tool_use_id) = result
|
||||
.as_object()
|
||||
.and_then(|result| result.get("toolUseId"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
history_tool_result_ids.insert(tool_use_id.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if let Some(assistant) = item
|
||||
.get("assistantResponseMessage")
|
||||
.and_then(Value::as_object)
|
||||
{
|
||||
if let Some(tool_uses) = assistant.get("toolUses").and_then(Value::as_array) {
|
||||
for tool_use in tool_uses {
|
||||
let Some(tool_use) = tool_use.as_object() else {
|
||||
continue;
|
||||
};
|
||||
if let Some(name) = tool_use
|
||||
.get("name")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
history_tool_names.insert(name.to_string());
|
||||
}
|
||||
if let Some(tool_use_id) = tool_use
|
||||
.get("toolUseId")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
history_tool_use_ids.insert(tool_use_id.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let existing_tool_names = tools
|
||||
.iter()
|
||||
.filter_map(|tool| {
|
||||
tool.get("toolSpecification")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|spec| spec.get("name"))
|
||||
.and_then(Value::as_str)
|
||||
.map(|name| name.to_ascii_lowercase())
|
||||
})
|
||||
.collect::<BTreeSet<_>>();
|
||||
for tool_name in history_tool_names {
|
||||
if !existing_tool_names.contains(&tool_name.to_ascii_lowercase()) {
|
||||
tools.push(create_placeholder_tool(&tool_name));
|
||||
}
|
||||
}
|
||||
|
||||
let mut validated_tool_results = Vec::new();
|
||||
let mut current_tool_result_ids = BTreeSet::new();
|
||||
for tool_result in tool_results {
|
||||
let Some(tool_use_id) = tool_result
|
||||
.get("toolUseId")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if !history_tool_use_ids.contains(tool_use_id)
|
||||
|| history_tool_result_ids.contains(tool_use_id)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
current_tool_result_ids.insert(tool_use_id.to_string());
|
||||
validated_tool_results.push(tool_result);
|
||||
}
|
||||
|
||||
let orphaned_tool_use_ids = history_tool_use_ids
|
||||
.difference(&history_tool_result_ids)
|
||||
.filter(|tool_use_id| !current_tool_result_ids.contains(*tool_use_id))
|
||||
.cloned()
|
||||
.collect::<BTreeSet<_>>();
|
||||
if !orphaned_tool_use_ids.is_empty() {
|
||||
warn!(
|
||||
"kiro: removing {} orphaned tool_use(s) from history",
|
||||
orphaned_tool_use_ids.len()
|
||||
);
|
||||
for item in &mut history {
|
||||
let Some(item) = item.as_object_mut() else {
|
||||
continue;
|
||||
};
|
||||
let Some(assistant) = item
|
||||
.get_mut("assistantResponseMessage")
|
||||
.and_then(Value::as_object_mut)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let Some(tool_uses) = assistant.get_mut("toolUses").and_then(Value::as_array_mut)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
tool_uses.retain(|tool_use| {
|
||||
!tool_use
|
||||
.get("toolUseId")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|tool_use_id| orphaned_tool_use_ids.contains(tool_use_id))
|
||||
});
|
||||
if tool_uses.is_empty() {
|
||||
assistant.remove("toolUses");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut user_context = Map::new();
|
||||
if !tools.is_empty() {
|
||||
user_context.insert("tools".to_string(), Value::Array(tools));
|
||||
}
|
||||
if !validated_tool_results.is_empty() {
|
||||
user_context.insert(
|
||||
"toolResults".to_string(),
|
||||
Value::Array(validated_tool_results),
|
||||
);
|
||||
}
|
||||
if let Some(thinking_prefix) = thinking_prefix.as_deref() {
|
||||
if !has_thinking_tags(&text_content) {
|
||||
text_content = format!("{thinking_prefix}\n{text_content}");
|
||||
}
|
||||
}
|
||||
|
||||
let mut user_input = Map::new();
|
||||
user_input.insert(
|
||||
"userInputMessageContext".to_string(),
|
||||
Value::Object(user_context),
|
||||
);
|
||||
user_input.insert("content".to_string(), Value::String(text_content));
|
||||
user_input.insert("modelId".to_string(), Value::String(model_id.to_string()));
|
||||
user_input.insert("origin".to_string(), Value::String("AI_EDITOR".to_string()));
|
||||
if !images.is_empty() {
|
||||
user_input.insert("images".to_string(), Value::Array(images));
|
||||
}
|
||||
|
||||
Some(json!({
|
||||
"agentContinuationId": agent_continuation_id,
|
||||
"agentTaskType": "vibe",
|
||||
"chatTriggerType": "MANUAL",
|
||||
"currentMessage": {
|
||||
"userInputMessage": Value::Object(user_input)
|
||||
},
|
||||
"conversationId": conversation_id,
|
||||
"history": history,
|
||||
}))
|
||||
}
|
||||
|
||||
fn extract_session_id(user_id: &str) -> Option<String> {
|
||||
let position = user_id.find("session_")?;
|
||||
let candidate = user_id.get(position + "session_".len()..position + "session_".len() + 36)?;
|
||||
(candidate.matches('-').count() == 4).then(|| candidate.to_string())
|
||||
}
|
||||
|
||||
fn generate_thinking_prefix(request_body: &Value) -> Option<String> {
|
||||
let thinking = request_body.get("thinking")?.as_object()?;
|
||||
match thinking.get("type").and_then(Value::as_str).map(str::trim) {
|
||||
Some("enabled") => {
|
||||
let budget_tokens = thinking
|
||||
.get("budget_tokens")
|
||||
.and_then(Value::as_i64)
|
||||
.unwrap_or_default();
|
||||
Some(format!(
|
||||
"<thinking_mode>enabled</thinking_mode><max_thinking_length>{budget_tokens}</max_thinking_length>"
|
||||
))
|
||||
}
|
||||
Some("adaptive") => {
|
||||
let effort = request_body
|
||||
.get("output_config")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|cfg| cfg.get("effort"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("high");
|
||||
Some(format!(
|
||||
"<thinking_mode>adaptive</thinking_mode><thinking_effort>{effort}</thinking_effort>"
|
||||
))
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn has_thinking_tags(content: &str) -> bool {
|
||||
content.contains("<thinking_mode>") || content.contains("<max_thinking_length>")
|
||||
}
|
||||
|
||||
fn system_to_text(system: Option<&Value>) -> String {
|
||||
match system {
|
||||
Some(Value::String(text)) => text.clone(),
|
||||
Some(Value::Array(items)) => items
|
||||
.iter()
|
||||
.filter_map(|item| {
|
||||
item.as_object()
|
||||
.and_then(|item| item.get("text"))
|
||||
.and_then(Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"),
|
||||
_ => String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn flush_user_buffer(user_buffer: &mut Vec<&Map<String, Value>>, model_id: &str) -> Option<Value> {
|
||||
if user_buffer.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut parts = Vec::new();
|
||||
let mut images = Vec::new();
|
||||
let mut tool_results = Vec::new();
|
||||
for message in user_buffer.drain(..) {
|
||||
let (text, mut message_images, mut message_tool_results) =
|
||||
process_message_content(message.get("content"));
|
||||
if !text.is_empty() {
|
||||
parts.push(text);
|
||||
}
|
||||
images.append(&mut message_images);
|
||||
tool_results.append(&mut message_tool_results);
|
||||
}
|
||||
|
||||
let mut payload = Map::new();
|
||||
payload.insert("content".to_string(), Value::String(parts.join("\n")));
|
||||
payload.insert("modelId".to_string(), Value::String(model_id.to_string()));
|
||||
payload.insert("origin".to_string(), Value::String("AI_EDITOR".to_string()));
|
||||
if !images.is_empty() {
|
||||
payload.insert("images".to_string(), Value::Array(images));
|
||||
}
|
||||
if !tool_results.is_empty() {
|
||||
payload.insert(
|
||||
"userInputMessageContext".to_string(),
|
||||
json!({"toolResults": tool_results}),
|
||||
);
|
||||
}
|
||||
|
||||
Some(json!({"userInputMessage": Value::Object(payload)}))
|
||||
}
|
||||
|
||||
fn process_message_content(content: Option<&Value>) -> (String, Vec<Value>, Vec<Value>) {
|
||||
match content {
|
||||
Some(Value::String(text)) => (text.clone(), Vec::new(), Vec::new()),
|
||||
Some(Value::Array(blocks)) => {
|
||||
let mut text_parts = Vec::new();
|
||||
let mut images = Vec::new();
|
||||
let mut tool_results = Vec::new();
|
||||
|
||||
for block in blocks {
|
||||
let Some(block) = block.as_object() else {
|
||||
continue;
|
||||
};
|
||||
match block
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default()
|
||||
{
|
||||
"text" => {
|
||||
if let Some(text) = block.get("text").and_then(Value::as_str) {
|
||||
text_parts.push(text.to_string());
|
||||
}
|
||||
}
|
||||
"image" => {
|
||||
let Some(source) = block.get("source").and_then(Value::as_object) else {
|
||||
continue;
|
||||
};
|
||||
let Some(format) = source
|
||||
.get("media_type")
|
||||
.or_else(|| source.get("mediaType"))
|
||||
.and_then(Value::as_str)
|
||||
.and_then(image_format)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let Some(bytes) = source.get("data").and_then(Value::as_str) else {
|
||||
continue;
|
||||
};
|
||||
images.push(json!({
|
||||
"format": format,
|
||||
"source": {"bytes": bytes}
|
||||
}));
|
||||
}
|
||||
"tool_result" => {
|
||||
let Some(tool_use_id) = block
|
||||
.get("tool_use_id")
|
||||
.or_else(|| block.get("toolUseId"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let text = match block.get("content") {
|
||||
Some(Value::String(text)) => text.clone(),
|
||||
Some(Value::Array(items)) => items
|
||||
.iter()
|
||||
.filter_map(|item| {
|
||||
item.as_object()
|
||||
.filter(|item| {
|
||||
item.get("type").and_then(Value::as_str) == Some("text")
|
||||
})
|
||||
.and_then(|item| item.get("text"))
|
||||
.and_then(Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"),
|
||||
Some(other) => {
|
||||
serde_json::to_string(other).unwrap_or_else(|_| other.to_string())
|
||||
}
|
||||
None => String::new(),
|
||||
};
|
||||
let is_error = block
|
||||
.get("is_error")
|
||||
.or_else(|| block.get("isError"))
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
tool_results.push(json!({
|
||||
"toolUseId": tool_use_id,
|
||||
"content": [{"text": text}],
|
||||
"status": if is_error { "error" } else { "success" },
|
||||
"isError": is_error,
|
||||
}));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
(text_parts.join(""), images, tool_results)
|
||||
}
|
||||
_ => (String::new(), Vec::new(), Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn image_format(media_type: &str) -> Option<&'static str> {
|
||||
let (prefix, suffix) = media_type.split_once('/')?;
|
||||
if prefix != "image" {
|
||||
return None;
|
||||
}
|
||||
match suffix.trim().to_ascii_lowercase().as_str() {
|
||||
"jpeg" => Some("jpeg"),
|
||||
"png" => Some("png"),
|
||||
"gif" => Some("gif"),
|
||||
"webp" => Some("webp"),
|
||||
"jpg" => Some("jpeg"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn clean_tool_schema(value: &Value) -> Value {
|
||||
match value {
|
||||
Value::Object(object) => {
|
||||
let mut out = Map::new();
|
||||
for (key, inner) in object {
|
||||
if key == "additionalProperties" {
|
||||
continue;
|
||||
}
|
||||
if key == "required" && inner.as_array().is_some_and(|items| items.is_empty()) {
|
||||
continue;
|
||||
}
|
||||
out.insert(key.clone(), clean_tool_schema(inner));
|
||||
}
|
||||
Value::Object(out)
|
||||
}
|
||||
Value::Array(items) => Value::Array(items.iter().map(clean_tool_schema).collect()),
|
||||
_ => value.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_tools(tools: Option<&Value>) -> Vec<Value> {
|
||||
let Some(tools) = tools.and_then(Value::as_array) else {
|
||||
return Vec::new();
|
||||
};
|
||||
|
||||
tools
|
||||
.iter()
|
||||
.filter_map(|tool| {
|
||||
let tool = tool.as_object()?;
|
||||
let name = tool.get("name")?.as_str()?.trim();
|
||||
if name.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let mut description = tool
|
||||
.get("description")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
let suffix = match name {
|
||||
"Write" => Some(WRITE_TOOL_DESCRIPTION_SUFFIX),
|
||||
"Edit" => Some(EDIT_TOOL_DESCRIPTION_SUFFIX),
|
||||
_ => None,
|
||||
};
|
||||
if let Some(suffix) = suffix {
|
||||
description = if description.is_empty() {
|
||||
suffix.to_string()
|
||||
} else {
|
||||
format!("{description}\n{suffix}")
|
||||
};
|
||||
}
|
||||
if description.len() > 10_000 {
|
||||
description.truncate(10_000);
|
||||
}
|
||||
let input_schema = tool
|
||||
.get("input_schema")
|
||||
.or_else(|| tool.get("inputSchema"))
|
||||
.filter(|value| value.is_object())
|
||||
.map(clean_tool_schema)
|
||||
.unwrap_or_else(|| json!({}));
|
||||
|
||||
Some(json!({
|
||||
"toolSpecification": {
|
||||
"name": name,
|
||||
"description": description,
|
||||
"inputSchema": {
|
||||
"json": input_schema
|
||||
}
|
||||
}
|
||||
}))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn create_placeholder_tool(name: &str) -> Value {
|
||||
json!({
|
||||
"toolSpecification": {
|
||||
"name": name,
|
||||
"description": "Tool used in conversation history",
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn convert_assistant_message(message: &Map<String, Value>) -> Option<Value> {
|
||||
let content = message.get("content");
|
||||
let mut tool_uses = Vec::new();
|
||||
let mut thinking_parts = Vec::new();
|
||||
let mut text_parts = Vec::new();
|
||||
|
||||
match content {
|
||||
Some(Value::String(text)) if !text.is_empty() => {
|
||||
text_parts.push(text.clone());
|
||||
}
|
||||
Some(Value::Array(blocks)) => {
|
||||
for block in blocks {
|
||||
let Some(block) = block.as_object() else {
|
||||
continue;
|
||||
};
|
||||
match block
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default()
|
||||
{
|
||||
"thinking" => {
|
||||
if let Some(thinking) = block.get("thinking").and_then(Value::as_str) {
|
||||
if !thinking.is_empty() {
|
||||
thinking_parts.push(thinking.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
"text" => {
|
||||
if let Some(text) = block.get("text").and_then(Value::as_str) {
|
||||
if !text.is_empty() {
|
||||
text_parts.push(text.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
"tool_use" => {
|
||||
let Some(tool_use_id) = block
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let Some(name) = block
|
||||
.get("name")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let input = block
|
||||
.get("input")
|
||||
.filter(|value| value.is_object())
|
||||
.cloned()
|
||||
.unwrap_or_else(|| json!({}));
|
||||
tool_uses.push(json!({
|
||||
"toolUseId": tool_use_id,
|
||||
"name": name,
|
||||
"input": input
|
||||
}));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
let thinking_str = thinking_parts.join("");
|
||||
let text_str = text_parts.join("");
|
||||
let mut content_str = if thinking_str.is_empty() {
|
||||
text_str
|
||||
} else if text_str.is_empty() {
|
||||
format!("<thinking>{thinking_str}</thinking>")
|
||||
} else {
|
||||
format!("<thinking>{thinking_str}</thinking>\n\n{text_str}")
|
||||
};
|
||||
|
||||
if content_str.is_empty() && !tool_uses.is_empty() {
|
||||
content_str = " ".to_string();
|
||||
}
|
||||
if content_str.is_empty() && tool_uses.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut out = Map::new();
|
||||
out.insert("content".to_string(), Value::String(content_str));
|
||||
if !tool_uses.is_empty() {
|
||||
out.insert("toolUses".to_string(), Value::Array(tool_uses));
|
||||
}
|
||||
Some(Value::Object(out))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::convert_claude_messages_to_conversation_state;
|
||||
|
||||
#[test]
|
||||
fn converts_simple_claude_request_into_conversation_state() {
|
||||
let conversation_state = convert_claude_messages_to_conversation_state(
|
||||
&json!({
|
||||
"messages": [
|
||||
{"role":"user","content":"hello"}
|
||||
],
|
||||
"thinking": {"type": "enabled", "budget_tokens": 128},
|
||||
"tools": [
|
||||
{"name":"Write","description":"write file","input_schema":{"type":"object","properties":{},"required":[]}}
|
||||
]
|
||||
}),
|
||||
"claude-sonnet-4-upstream",
|
||||
)
|
||||
.expect("conversation state should build");
|
||||
|
||||
assert_eq!(
|
||||
conversation_state
|
||||
.get("currentMessage")
|
||||
.and_then(|value| value.get("userInputMessage"))
|
||||
.and_then(|value| value.get("content"))
|
||||
.and_then(|value| value.as_str()),
|
||||
Some(
|
||||
"<thinking_mode>enabled</thinking_mode><max_thinking_length>128</max_thinking_length>\nhello"
|
||||
)
|
||||
);
|
||||
assert_eq!(
|
||||
conversation_state
|
||||
.get("currentMessage")
|
||||
.and_then(|value| value.get("userInputMessage"))
|
||||
.and_then(|value| value.get("userInputMessageContext"))
|
||||
.and_then(|value| value.get("tools"))
|
||||
.and_then(|value| value.as_array())
|
||||
.map(Vec::len),
|
||||
Some(1)
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
pub use aether_oauth::provider::providers::{
|
||||
generate_kiro_machine_id as generate_machine_id,
|
||||
normalize_kiro_machine_id as normalize_machine_id, KiroAuthConfig, DEFAULT_REGION,
|
||||
};
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{generate_machine_id, normalize_machine_id, KiroAuthConfig, DEFAULT_REGION};
|
||||
|
||||
#[test]
|
||||
fn normalizes_uuid_machine_id() {
|
||||
assert_eq!(
|
||||
normalize_machine_id("123e4567-e89b-12d3-a456-426614174000").as_deref(),
|
||||
Some("123e4567e89b12d3a456426614174000123e4567e89b12d3a456426614174000")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn hashes_refresh_token_into_machine_id() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: None,
|
||||
refresh_token: Some("r".repeat(128)),
|
||||
expires_at: None,
|
||||
profile_arn: None,
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: None,
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: None,
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: None,
|
||||
};
|
||||
|
||||
let machine_id = generate_machine_id(&auth_config, None).expect("machine id should exist");
|
||||
assert_eq!(machine_id.len(), 64);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_auth_config_aliases() {
|
||||
let auth_config = KiroAuthConfig::from_raw_json(Some(
|
||||
r#"{
|
||||
"authMethod":"identity_center",
|
||||
"refreshToken":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
|
||||
"expires_at": 4102444800,
|
||||
"profileArn":"arn:aws:bedrock:demo",
|
||||
"apiRegion":"us-west-2",
|
||||
"clientId":"cid",
|
||||
"clientSecret":"secret",
|
||||
"machineId":"123e4567-e89b-12d3-a456-426614174000",
|
||||
"kiroVersion":"1.2.3",
|
||||
"systemVersion":"darwin#24.6.0",
|
||||
"nodeVersion":"22.21.1",
|
||||
"accessToken":"cached-token"
|
||||
}"#,
|
||||
))
|
||||
.expect("auth config should parse");
|
||||
|
||||
assert_eq!(auth_config.auth_method.as_deref(), Some("idc"));
|
||||
assert_eq!(
|
||||
auth_config.refresh_token.as_deref(),
|
||||
Some(
|
||||
"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
|
||||
)
|
||||
);
|
||||
assert_eq!(auth_config.expires_at, Some(4_102_444_800));
|
||||
assert_eq!(
|
||||
auth_config.profile_arn.as_deref(),
|
||||
Some("arn:aws:bedrock:demo")
|
||||
);
|
||||
assert_eq!(auth_config.client_id.as_deref(), Some("cid"));
|
||||
assert_eq!(auth_config.client_secret.as_deref(), Some("secret"));
|
||||
assert_eq!(auth_config.effective_api_region(), "us-west-2");
|
||||
assert_eq!(auth_config.effective_kiro_version(), "1.2.3");
|
||||
assert_eq!(auth_config.effective_system_version(), "darwin#24.6.0");
|
||||
assert_eq!(auth_config.effective_node_version(), "22.21.1");
|
||||
assert_eq!(auth_config.access_token.as_deref(), Some("cached-token"));
|
||||
assert!(auth_config.is_idc_auth());
|
||||
assert!(auth_config.profile_arn_for_payload().is_none());
|
||||
assert_eq!(auth_config.effective_auth_region(), "us-east-1");
|
||||
assert!(auth_config.can_refresh_access_token());
|
||||
assert_eq!(DEFAULT_REGION, "us-east-1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_external_idp_auth_method_for_header_selection() {
|
||||
let auth_config = KiroAuthConfig::from_raw_json(Some(
|
||||
r#"{
|
||||
"authMethod":"external_idp",
|
||||
"refreshToken":"rt-1",
|
||||
"clientId":"cid",
|
||||
"clientSecret":"secret",
|
||||
"profileArn":"arn:aws:bedrock:demo"
|
||||
}"#,
|
||||
))
|
||||
.expect("auth config should parse");
|
||||
|
||||
assert_eq!(auth_config.auth_method.as_deref(), Some("external_idp"));
|
||||
assert!(auth_config.is_idc_auth());
|
||||
assert!(auth_config.uses_external_idp_token_type());
|
||||
assert!(auth_config.profile_arn_for_payload().is_none());
|
||||
assert_eq!(
|
||||
auth_config.profile_arn_for_mcp(),
|
||||
Some("arn:aws:bedrock:demo")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_idc_when_client_credentials_exist() {
|
||||
let auth_config = KiroAuthConfig::from_raw_json(Some(
|
||||
r#"{
|
||||
"refreshToken":"rt-1",
|
||||
"clientId":"cid",
|
||||
"clientSecret":"secret",
|
||||
"profileArn":"arn:aws:bedrock:demo"
|
||||
}"#,
|
||||
))
|
||||
.expect("auth config should parse");
|
||||
|
||||
assert!(auth_config.is_idc_auth());
|
||||
assert!(!auth_config.uses_external_idp_token_type());
|
||||
assert!(auth_config.profile_arn_for_payload().is_none());
|
||||
assert_eq!(
|
||||
auth_config.profile_arn_for_mcp(),
|
||||
Some("arn:aws:bedrock:demo")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn round_trips_json_value() {
|
||||
let auth_config = KiroAuthConfig::from_raw_json(Some(
|
||||
r#"{
|
||||
"auth_method":"social",
|
||||
"refreshToken":"rt-1....................................................................................................",
|
||||
"expires_at": 4102444800,
|
||||
"profileArn":"arn:aws:bedrock:demo",
|
||||
"region":"eu-north-1",
|
||||
"apiRegion":"us-west-2",
|
||||
"machineId":"123e4567-e89b-12d3-a456-426614174000",
|
||||
"kiroVersion":"1.2.3",
|
||||
"systemVersion":"darwin#24.6.0",
|
||||
"nodeVersion":"22.21.1",
|
||||
"accessToken":"cached-token"
|
||||
}"#,
|
||||
))
|
||||
.expect("auth config should parse");
|
||||
|
||||
let value = auth_config.to_json_value();
|
||||
let reparsed = KiroAuthConfig::from_json_value(&value).expect("auth config should reparse");
|
||||
assert_eq!(reparsed, auth_config);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,334 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use uuid::Uuid;
|
||||
|
||||
use super::credentials::KiroAuthConfig;
|
||||
|
||||
pub const AWS_EVENTSTREAM_CONTENT_TYPE: &str = "application/vnd.amazon.eventstream";
|
||||
pub const KIRO_PROFILE_ARN_HEADER: &str = "x-amzn-kiro-profile-arn";
|
||||
pub const KIRO_TOKEN_TYPE_HEADER: &str = "TokenType";
|
||||
pub const KIRO_EXTERNAL_IDP_TOKEN_TYPE: &str = "EXTERNAL_IDP";
|
||||
const AWS_SDK_JS_MAIN_VERSION: &str = "1.0.27";
|
||||
const AWS_SDK_JS_LIST_MODELS_VERSION: &str = "1.0.0";
|
||||
const CODEWHISPERER_OPTOUT: &str = "true";
|
||||
const KIRO_AGENT_MODE: &str = "vibe";
|
||||
|
||||
fn build_kiro_ide_tag(kiro_version: &str, machine_id: &str) -> String {
|
||||
if machine_id.trim().is_empty() {
|
||||
format!("KiroIDE-{kiro_version}")
|
||||
} else {
|
||||
format!("KiroIDE-{kiro_version}-{machine_id}")
|
||||
}
|
||||
}
|
||||
|
||||
fn build_x_amz_user_agent_main(kiro_version: &str, machine_id: &str) -> String {
|
||||
format!(
|
||||
"aws-sdk-js/{AWS_SDK_JS_MAIN_VERSION} {}",
|
||||
build_kiro_ide_tag(kiro_version, machine_id)
|
||||
)
|
||||
}
|
||||
|
||||
fn build_user_agent_main(
|
||||
system_version: &str,
|
||||
node_version: &str,
|
||||
kiro_version: &str,
|
||||
machine_id: &str,
|
||||
) -> String {
|
||||
format!(
|
||||
"aws-sdk-js/{AWS_SDK_JS_MAIN_VERSION} ua/2.1 os/{system_version} lang/js md/nodejs#{node_version} api/codewhispererstreaming#{AWS_SDK_JS_MAIN_VERSION} m/E {}",
|
||||
build_kiro_ide_tag(kiro_version, machine_id)
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_generate_assistant_headers(
|
||||
auth_config: &KiroAuthConfig,
|
||||
machine_id: &str,
|
||||
) -> BTreeMap<String, String> {
|
||||
let kiro_version = auth_config.effective_kiro_version();
|
||||
let system_version = auth_config.effective_system_version();
|
||||
let node_version = auth_config.effective_node_version();
|
||||
let region = auth_config.effective_api_region();
|
||||
let host = format!("q.{region}.amazonaws.com");
|
||||
|
||||
BTreeMap::from([
|
||||
(
|
||||
"accept".to_string(),
|
||||
AWS_EVENTSTREAM_CONTENT_TYPE.to_string(),
|
||||
),
|
||||
(
|
||||
"amz-sdk-invocation-id".to_string(),
|
||||
Uuid::new_v4().to_string(),
|
||||
),
|
||||
(
|
||||
"amz-sdk-request".to_string(),
|
||||
"attempt=1; max=3".to_string(),
|
||||
),
|
||||
("connection".to_string(), "close".to_string()),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
("host".to_string(), host),
|
||||
(
|
||||
"user-agent".to_string(),
|
||||
build_user_agent_main(system_version, node_version, kiro_version, machine_id),
|
||||
),
|
||||
(
|
||||
"x-amz-user-agent".to_string(),
|
||||
build_x_amz_user_agent_main(kiro_version, machine_id),
|
||||
),
|
||||
(
|
||||
"x-amzn-codewhisperer-optout".to_string(),
|
||||
CODEWHISPERER_OPTOUT.to_string(),
|
||||
),
|
||||
(
|
||||
"x-amzn-kiro-agent-mode".to_string(),
|
||||
KIRO_AGENT_MODE.to_string(),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
||||
pub fn build_mcp_headers(
|
||||
auth_config: &KiroAuthConfig,
|
||||
machine_id: &str,
|
||||
) -> BTreeMap<String, String> {
|
||||
let kiro_version = auth_config.effective_kiro_version();
|
||||
let system_version = auth_config.effective_system_version();
|
||||
let node_version = auth_config.effective_node_version();
|
||||
let region = auth_config.effective_api_region();
|
||||
let host = format!("q.{region}.amazonaws.com");
|
||||
|
||||
let mut headers = BTreeMap::from([
|
||||
("accept".to_string(), "application/json".to_string()),
|
||||
(
|
||||
"amz-sdk-invocation-id".to_string(),
|
||||
Uuid::new_v4().to_string(),
|
||||
),
|
||||
(
|
||||
"amz-sdk-request".to_string(),
|
||||
"attempt=1; max=3".to_string(),
|
||||
),
|
||||
("connection".to_string(), "close".to_string()),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
("host".to_string(), host),
|
||||
(
|
||||
"user-agent".to_string(),
|
||||
build_user_agent_main(system_version, node_version, kiro_version, machine_id),
|
||||
),
|
||||
(
|
||||
"x-amz-user-agent".to_string(),
|
||||
build_x_amz_user_agent_main(kiro_version, machine_id),
|
||||
),
|
||||
(
|
||||
"x-amzn-codewhisperer-optout".to_string(),
|
||||
CODEWHISPERER_OPTOUT.to_string(),
|
||||
),
|
||||
]);
|
||||
if let Some(profile_arn) = auth_config.profile_arn_for_mcp() {
|
||||
headers.insert(KIRO_PROFILE_ARN_HEADER.to_string(), profile_arn.to_string());
|
||||
}
|
||||
headers
|
||||
}
|
||||
|
||||
pub fn build_list_available_models_headers(
|
||||
auth_config: &KiroAuthConfig,
|
||||
machine_id: &str,
|
||||
) -> BTreeMap<String, String> {
|
||||
let kiro_version = auth_config.effective_kiro_version();
|
||||
let system_version = auth_config.effective_system_version();
|
||||
let node_version = auth_config.effective_node_version();
|
||||
let region = auth_config.effective_api_region();
|
||||
let host = format!("q.{region}.amazonaws.com");
|
||||
let ide_tag = build_kiro_ide_tag(kiro_version, machine_id);
|
||||
|
||||
BTreeMap::from([
|
||||
("accept".to_string(), "application/json".to_string()),
|
||||
(
|
||||
"amz-sdk-invocation-id".to_string(),
|
||||
Uuid::new_v4().to_string(),
|
||||
),
|
||||
(
|
||||
"amz-sdk-request".to_string(),
|
||||
"attempt=1; max=1".to_string(),
|
||||
),
|
||||
("connection".to_string(), "close".to_string()),
|
||||
("host".to_string(), host),
|
||||
(
|
||||
"user-agent".to_string(),
|
||||
format!(
|
||||
"aws-sdk-js/{AWS_SDK_JS_LIST_MODELS_VERSION} ua/2.1 os/{system_version} lang/js md/nodejs#{node_version} api/codewhispererruntime#1.0.0 m/N,E {ide_tag}"
|
||||
),
|
||||
),
|
||||
(
|
||||
"x-amz-user-agent".to_string(),
|
||||
format!("aws-sdk-js/{AWS_SDK_JS_LIST_MODELS_VERSION} {ide_tag}"),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::credentials::KiroAuthConfig;
|
||||
use super::{
|
||||
build_generate_assistant_headers, build_list_available_models_headers, build_mcp_headers,
|
||||
AWS_EVENTSTREAM_CONTENT_TYPE, KIRO_PROFILE_ARN_HEADER, KIRO_TOKEN_TYPE_HEADER,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn builds_generate_assistant_headers_for_region() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: None,
|
||||
refresh_token: None,
|
||||
expires_at: None,
|
||||
profile_arn: None,
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: Some("us-west-2".to_string()),
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: None,
|
||||
kiro_version: Some("1.2.3".to_string()),
|
||||
system_version: Some("darwin#24.6.0".to_string()),
|
||||
node_version: Some("22.21.1".to_string()),
|
||||
access_token: None,
|
||||
};
|
||||
|
||||
let headers = build_generate_assistant_headers(&auth_config, "machine-123");
|
||||
assert_eq!(
|
||||
headers.get("accept").map(String::as_str),
|
||||
Some(AWS_EVENTSTREAM_CONTENT_TYPE)
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("host").map(String::as_str),
|
||||
Some("q.us-west-2.amazonaws.com")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-amzn-kiro-agent-mode").map(String::as_str),
|
||||
Some("vibe")
|
||||
);
|
||||
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_list_available_models_headers() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: Some("social".to_string()),
|
||||
refresh_token: None,
|
||||
expires_at: None,
|
||||
profile_arn: Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo".into()),
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: Some("us-east-1".to_string()),
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: None,
|
||||
kiro_version: Some("0.12.155".to_string()),
|
||||
system_version: Some("darwin#24.6.0".to_string()),
|
||||
node_version: Some("22.21.1".to_string()),
|
||||
access_token: None,
|
||||
};
|
||||
|
||||
let headers = build_list_available_models_headers(&auth_config, "machine-123");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("host").map(String::as_str),
|
||||
Some("q.us-east-1.amazonaws.com")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("amz-sdk-request").map(String::as_str),
|
||||
Some("attempt=1; max=1")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-amz-user-agent").map(String::as_str),
|
||||
Some("aws-sdk-js/1.0.0 KiroIDE-0.12.155-machine-123")
|
||||
);
|
||||
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_generate_assistant_headers_without_external_idp_token_type() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: Some("idc".to_string()),
|
||||
refresh_token: None,
|
||||
expires_at: None,
|
||||
profile_arn: Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo".into()),
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: Some("us-east-1".to_string()),
|
||||
client_id: Some("client-id".to_string()),
|
||||
client_secret: Some("client-secret".to_string()),
|
||||
machine_id: None,
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: None,
|
||||
};
|
||||
|
||||
let headers = build_generate_assistant_headers(&auth_config, "machine-123");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("accept").map(String::as_str),
|
||||
Some(AWS_EVENTSTREAM_CONTENT_TYPE)
|
||||
);
|
||||
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_mcp_headers_with_profile_arn_for_social_auth() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: Some("social".to_string()),
|
||||
refresh_token: None,
|
||||
expires_at: None,
|
||||
profile_arn: Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo".into()),
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: Some("us-east-1".to_string()),
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: None,
|
||||
kiro_version: Some("0.3.210".to_string()),
|
||||
system_version: Some("darwin#24.6.0".to_string()),
|
||||
node_version: Some("22.21.1".to_string()),
|
||||
access_token: None,
|
||||
};
|
||||
|
||||
let headers = build_mcp_headers(&auth_config, "machine-123");
|
||||
assert_eq!(
|
||||
headers.get("accept").map(String::as_str),
|
||||
Some("application/json")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get(KIRO_PROFILE_ARN_HEADER).map(String::as_str),
|
||||
Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo")
|
||||
);
|
||||
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_mcp_headers_with_profile_arn_for_idc_auth() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: Some("idc".to_string()),
|
||||
refresh_token: None,
|
||||
expires_at: None,
|
||||
profile_arn: Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo".into()),
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: Some("us-west-2".to_string()),
|
||||
client_id: Some("client-id".to_string()),
|
||||
client_secret: Some("client-secret".to_string()),
|
||||
machine_id: None,
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: None,
|
||||
};
|
||||
|
||||
let headers = build_mcp_headers(&auth_config, "machine-123");
|
||||
assert_eq!(
|
||||
headers.get("host").map(String::as_str),
|
||||
Some("q.us-west-2.amazonaws.com")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get(KIRO_PROFILE_ARN_HEADER).map(String::as_str),
|
||||
Some("arn:aws:codewhisperer:us-east-1:123456789012:profile/demo")
|
||||
);
|
||||
assert!(!headers.contains_key(KIRO_TOKEN_TYPE_HEADER));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
mod auth;
|
||||
mod converter;
|
||||
mod credentials;
|
||||
mod headers;
|
||||
mod policy;
|
||||
mod refresh;
|
||||
mod request;
|
||||
mod url;
|
||||
|
||||
use crate::provider_types::{
|
||||
ProviderApiFormatInheritance, ProviderLocalEmbeddingSupport, ProviderRuntimePolicy,
|
||||
};
|
||||
|
||||
pub const RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrConfiguredBearer,
|
||||
enable_format_conversion_by_default: true,
|
||||
allow_auth_channel_mismatch_by_default: true,
|
||||
oauth_is_bearer_like: true,
|
||||
supports_model_fetch: true,
|
||||
supports_local_openai_chat_transport: false,
|
||||
supports_local_same_format_transport: false,
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::None,
|
||||
};
|
||||
|
||||
pub use auth::{
|
||||
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};
|
||||
pub use headers::{
|
||||
build_generate_assistant_headers, build_list_available_models_headers, build_mcp_headers,
|
||||
AWS_EVENTSTREAM_CONTENT_TYPE, KIRO_EXTERNAL_IDP_TOKEN_TYPE, KIRO_PROFILE_ARN_HEADER,
|
||||
KIRO_TOKEN_TYPE_HEADER,
|
||||
};
|
||||
pub use policy::{
|
||||
local_kiro_request_transport_unsupported_reason_with_network,
|
||||
supports_local_kiro_request_transport, supports_local_kiro_request_transport_with_network,
|
||||
};
|
||||
pub use refresh::KiroOAuthRefreshAdapter;
|
||||
pub use request::{
|
||||
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
|
||||
body_rules_are_locally_supported, build_kiro_provider_headers,
|
||||
build_kiro_provider_request_body, header_rules_are_locally_supported,
|
||||
supports_local_kiro_request_shape, KiroProviderHeadersInput,
|
||||
};
|
||||
pub use url::{
|
||||
build_kiro_generate_assistant_response_url, build_kiro_list_available_models_url,
|
||||
build_kiro_mcp_url, build_kiro_mcp_url_from_resolved_url, resolve_kiro_base_url,
|
||||
GENERATE_ASSISTANT_RESPONSE_PATH, KIRO_ENVELOPE_NAME, LIST_AVAILABLE_MODELS_PATH, MCP_PATH,
|
||||
MCP_STREAM_PATH,
|
||||
};
|
||||
@@ -0,0 +1,203 @@
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::super::{
|
||||
resolve_transport_profile, transport_profile_is_configured,
|
||||
transport_proxy_is_locally_supported,
|
||||
};
|
||||
use super::{
|
||||
body_rules_are_locally_supported, header_rules_are_locally_supported,
|
||||
supports_local_kiro_request_auth_resolution, supports_local_kiro_request_shape, PROVIDER_TYPE,
|
||||
};
|
||||
|
||||
pub fn local_kiro_request_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return if !transport.provider.is_active {
|
||||
Some("provider_inactive")
|
||||
} else if !transport.endpoint.is_active {
|
||||
Some("endpoint_inactive")
|
||||
} else {
|
||||
Some("key_inactive")
|
||||
};
|
||||
}
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(PROVIDER_TYPE)
|
||||
{
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
if !aether_ai_formats::api_format_alias_matches(
|
||||
&transport.endpoint.api_format,
|
||||
"claude:messages",
|
||||
) {
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
|
||||
return Some("transport_header_rules_unsupported");
|
||||
}
|
||||
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
if transport.key.decrypted_auth_config.is_some()
|
||||
&& !supports_local_kiro_request_auth_resolution(transport)
|
||||
{
|
||||
return Some("transport_oauth_resolution_unsupported");
|
||||
}
|
||||
if !supports_local_kiro_request_auth_resolution(transport) {
|
||||
return Some("transport_auth_unavailable");
|
||||
}
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return Some("transport_proxy_unsupported");
|
||||
}
|
||||
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return Some("transport_profile_unsupported");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub fn supports_local_kiro_request_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return false;
|
||||
}
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(PROVIDER_TYPE)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
if !aether_ai_formats::api_format_alias_matches(
|
||||
&transport.endpoint.api_format,
|
||||
"claude:messages",
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
supports_local_kiro_request_shape(
|
||||
transport.endpoint.header_rules.as_ref(),
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
) && supports_local_kiro_request_auth_resolution(transport)
|
||||
}
|
||||
|
||||
pub fn supports_local_kiro_request_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
local_kiro_request_transport_unsupported_reason_with_network(transport).is_none()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::{
|
||||
local_kiro_request_transport_unsupported_reason_with_network,
|
||||
supports_local_kiro_request_transport, supports_local_kiro_request_transport_with_network,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Kiro".to_string(),
|
||||
provider_type: "kiro".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "claude:messages".to_string(),
|
||||
api_family: Some("claude".to_string()),
|
||||
endpoint_kind: Some("messages".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://kiro.example".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: Some(vec!["claude:messages".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: Some(
|
||||
r#"{
|
||||
"access_token":"cached-token",
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
|
||||
"machine_id":"123e4567-e89b-12d3-a456-426614174000"
|
||||
}"#
|
||||
.to_string(),
|
||||
),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_kiro_request_transport_when_cached_access_token_exists() {
|
||||
assert!(supports_local_kiro_request_transport(&sample_transport()));
|
||||
assert!(supports_local_kiro_request_transport_with_network(
|
||||
&sample_transport()
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_kiro_request_transport_when_refresh_only_auth_exists() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(supports_local_kiro_request_transport(&transport));
|
||||
assert!(supports_local_kiro_request_transport_with_network(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reports_auth_unavailable_when_kiro_key_cannot_resolve_request_auth() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "api_key".to_string();
|
||||
transport.key.decrypted_auth_config = None;
|
||||
|
||||
assert_eq!(
|
||||
local_kiro_request_transport_unsupported_reason_with_network(&transport),
|
||||
Some("transport_auth_unavailable")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,446 @@
|
||||
use aether_oauth::provider::providers::KiroProviderOAuthAdapter as CoreKiroProviderOAuthAdapter;
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::super::oauth_refresh::{
|
||||
oauth_error_to_local_refresh_error, provider_oauth_transport_context_from_snapshot,
|
||||
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter, LocalOAuthRefreshError,
|
||||
LocalResolvedOAuthRequestAuth, ProviderOAuthLocalHttpExecutor,
|
||||
};
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::auth::{
|
||||
build_kiro_request_auth_from_config, resolve_local_kiro_request_auth, PROVIDER_TYPE,
|
||||
};
|
||||
use super::credentials::KiroAuthConfig;
|
||||
|
||||
#[cfg(test)]
|
||||
const IDC_AMZ_USER_AGENT: &str = "aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct KiroOAuthRefreshAdapter {
|
||||
social_refresh_base_url: Option<String>,
|
||||
idc_refresh_base_url: Option<String>,
|
||||
}
|
||||
|
||||
impl KiroOAuthRefreshAdapter {
|
||||
pub fn with_refresh_base_urls(
|
||||
mut self,
|
||||
social_refresh_base_url: Option<String>,
|
||||
idc_refresh_base_url: Option<String>,
|
||||
) -> Self {
|
||||
self.social_refresh_base_url = social_refresh_base_url;
|
||||
self.idc_refresh_base_url = idc_refresh_base_url;
|
||||
self
|
||||
}
|
||||
|
||||
pub async fn refresh_auth_config(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
auth_config: &KiroAuthConfig,
|
||||
) -> Result<KiroAuthConfig, LocalOAuthRefreshError> {
|
||||
let adapter = CoreKiroProviderOAuthAdapter::default().with_refresh_base_urls(
|
||||
self.social_refresh_base_url.clone(),
|
||||
self.idc_refresh_base_url.clone(),
|
||||
);
|
||||
let oauth_executor =
|
||||
ProviderOAuthLocalHttpExecutor::new(PROVIDER_TYPE, transport, executor);
|
||||
let ctx = provider_oauth_transport_context_from_snapshot(transport);
|
||||
adapter
|
||||
.refresh_auth_config(&oauth_executor, &ctx, auth_config)
|
||||
.await
|
||||
.map_err(|error| oauth_error_to_local_refresh_error(PROVIDER_TYPE, error))
|
||||
}
|
||||
|
||||
fn auth_config_from_entry(entry: &CachedOAuthEntry) -> Option<KiroAuthConfig> {
|
||||
entry
|
||||
.metadata
|
||||
.as_ref()
|
||||
.filter(|_| entry.provider_type.eq_ignore_ascii_case(PROVIDER_TYPE))
|
||||
.and_then(KiroAuthConfig::from_json_value)
|
||||
}
|
||||
|
||||
fn base_auth_config(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<KiroAuthConfig> {
|
||||
entry.and_then(Self::auth_config_from_entry).or_else(|| {
|
||||
KiroAuthConfig::from_raw_json(transport.key.decrypted_auth_config.as_deref())
|
||||
})
|
||||
}
|
||||
|
||||
fn build_cached_entry(auth_config: &KiroAuthConfig) -> Option<CachedOAuthEntry> {
|
||||
let request_auth = build_kiro_request_auth_from_config(auth_config.clone(), None)?;
|
||||
Some(CachedOAuthEntry {
|
||||
provider_type: PROVIDER_TYPE.to_string(),
|
||||
auth_header_name: request_auth.name.to_string(),
|
||||
auth_header_value: request_auth.value,
|
||||
expires_at_unix_secs: auth_config.expires_at,
|
||||
metadata: Some(auth_config.to_json_value()),
|
||||
})
|
||||
}
|
||||
|
||||
fn refreshable_auth_config(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<KiroAuthConfig> {
|
||||
let auth_config = self.base_auth_config(transport, entry)?;
|
||||
auth_config
|
||||
.can_refresh_access_token()
|
||||
.then_some(auth_config)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for KiroOAuthRefreshAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
PROVIDER_TYPE
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
let auth_config = Self::auth_config_from_entry(entry)?;
|
||||
let request_auth = build_kiro_request_auth_from_config(auth_config, None)?;
|
||||
Some(LocalResolvedOAuthRequestAuth::Kiro(request_auth))
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
resolve_local_kiro_request_auth(transport).map(LocalResolvedOAuthRequestAuth::Kiro)
|
||||
}
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool {
|
||||
entry
|
||||
.and_then(|cached| self.resolve_cached(transport, cached))
|
||||
.is_none()
|
||||
&& self.resolve_without_refresh(transport).is_none()
|
||||
&& self.refreshable_auth_config(transport, entry).is_some()
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
|
||||
let Some(auth_config) = self.refreshable_auth_config(transport, entry) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let refreshed = self
|
||||
.refresh_auth_config(executor, transport, &auth_config)
|
||||
.await?;
|
||||
Ok(Self::build_cached_entry(&refreshed))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use super::super::super::oauth_refresh::{
|
||||
LocalOAuthRefreshAdapter, LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
|
||||
};
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::{KiroOAuthRefreshAdapter, IDC_AMZ_USER_AGENT};
|
||||
use axum::body::to_bytes;
|
||||
use axum::extract::Request;
|
||||
use axum::response::IntoResponse;
|
||||
use axum::routing::any;
|
||||
use axum::{Json, Router};
|
||||
use http::StatusCode;
|
||||
use serde_json::{json, Value};
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct SeenRefreshRequest {
|
||||
body: Value,
|
||||
authorization: String,
|
||||
host: String,
|
||||
user_agent: String,
|
||||
x_amz_user_agent: String,
|
||||
}
|
||||
|
||||
fn sample_transport(raw_auth_config: &str) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Kiro".to_string(),
|
||||
provider_type: "kiro".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "claude:messages".to_string(),
|
||||
api_family: Some("claude".to_string()),
|
||||
endpoint_kind: Some("cli".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://kiro.example".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: Some(vec!["claude:messages".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: Some(raw_auth_config.to_string()),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
async fn start_server(app: Router) -> (String, JoinHandle<()>) {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener
|
||||
.local_addr()
|
||||
.expect("listener should expose local addr");
|
||||
let handle = tokio::spawn(async move {
|
||||
axum::serve(listener, app).await.expect("server should run");
|
||||
});
|
||||
(format!("http://{addr}"), handle)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn refreshes_social_token_via_adapter() {
|
||||
let seen_request = Arc::new(Mutex::new(None::<SeenRefreshRequest>));
|
||||
let seen_request_clone = Arc::clone(&seen_request);
|
||||
let server = Router::new().route(
|
||||
"/refreshToken",
|
||||
any(move |request: Request| {
|
||||
let seen_request_inner = Arc::clone(&seen_request_clone);
|
||||
async move {
|
||||
let (parts, body) = request.into_parts();
|
||||
let raw_body = to_bytes(body, usize::MAX).await.expect("body should read");
|
||||
let body: Value =
|
||||
serde_json::from_slice(&raw_body).expect("body should parse as json");
|
||||
*seen_request_inner.lock().expect("mutex should lock") =
|
||||
Some(SeenRefreshRequest {
|
||||
body,
|
||||
authorization: parts
|
||||
.headers
|
||||
.get("authorization")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
host: parts
|
||||
.headers
|
||||
.get("host")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
user_agent: parts
|
||||
.headers
|
||||
.get("user-agent")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
x_amz_user_agent: parts
|
||||
.headers
|
||||
.get("x-amz-user-agent")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
});
|
||||
(
|
||||
StatusCode::OK,
|
||||
Json(json!({
|
||||
"accessToken": "cached-kiro-access-token",
|
||||
"refreshToken": "s".repeat(120),
|
||||
"expiresIn": 3600,
|
||||
"profileArn": "arn:aws:bedrock:demo"
|
||||
})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}),
|
||||
);
|
||||
let (server_url, server_handle) = start_server(server).await;
|
||||
let adapter =
|
||||
KiroOAuthRefreshAdapter::default().with_refresh_base_urls(Some(server_url), None);
|
||||
let transport = sample_transport(
|
||||
r#"{
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
|
||||
"machine_id":"123e4567-e89b-12d3-a456-426614174000",
|
||||
"kiro_version":"1.2.3"
|
||||
}"#,
|
||||
);
|
||||
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
|
||||
|
||||
let entry = adapter
|
||||
.refresh(&executor, &transport, None)
|
||||
.await
|
||||
.expect("refresh should succeed")
|
||||
.expect("cached entry should exist");
|
||||
let resolved = adapter
|
||||
.resolve_cached(&transport, &entry)
|
||||
.expect("cached entry should resolve");
|
||||
let seen_request = seen_request
|
||||
.lock()
|
||||
.expect("mutex should lock")
|
||||
.clone()
|
||||
.expect("refresh request should be captured");
|
||||
|
||||
assert_eq!(seen_request.body["refreshToken"], json!("r".repeat(120)));
|
||||
assert_eq!(seen_request.authorization, "");
|
||||
assert!(!seen_request.user_agent.is_empty());
|
||||
assert_eq!(seen_request.x_amz_user_agent, "");
|
||||
assert!(!seen_request.host.trim().is_empty());
|
||||
match resolved {
|
||||
LocalResolvedOAuthRequestAuth::Kiro(auth) => {
|
||||
assert_eq!(auth.value, "Bearer cached-kiro-access-token");
|
||||
assert_eq!(
|
||||
auth.auth_config.profile_arn.as_deref(),
|
||||
Some("arn:aws:bedrock:demo")
|
||||
);
|
||||
assert!(auth.auth_config.expires_at.is_some());
|
||||
}
|
||||
other => panic!("unexpected resolved auth: {other:?}"),
|
||||
}
|
||||
|
||||
server_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn refreshes_idc_token_via_adapter() {
|
||||
let seen_request = Arc::new(Mutex::new(None::<SeenRefreshRequest>));
|
||||
let seen_request_clone = Arc::clone(&seen_request);
|
||||
let server = Router::new().route(
|
||||
"/token",
|
||||
any(move |request: Request| {
|
||||
let seen_request_inner = Arc::clone(&seen_request_clone);
|
||||
async move {
|
||||
let (parts, body) = request.into_parts();
|
||||
let raw_body = to_bytes(body, usize::MAX).await.expect("body should read");
|
||||
let body: Value =
|
||||
serde_json::from_slice(&raw_body).expect("body should parse as json");
|
||||
*seen_request_inner.lock().expect("mutex should lock") =
|
||||
Some(SeenRefreshRequest {
|
||||
body,
|
||||
authorization: parts
|
||||
.headers
|
||||
.get("authorization")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
host: parts
|
||||
.headers
|
||||
.get("host")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
user_agent: parts
|
||||
.headers
|
||||
.get("user-agent")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
x_amz_user_agent: parts
|
||||
.headers
|
||||
.get("x-amz-user-agent")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
});
|
||||
(
|
||||
StatusCode::OK,
|
||||
Json(json!({
|
||||
"accessToken": "cached-idc-access-token",
|
||||
"refreshToken": "i".repeat(120),
|
||||
"expiresIn": 1800
|
||||
})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}),
|
||||
);
|
||||
let (server_url, server_handle) = start_server(server).await;
|
||||
let adapter =
|
||||
KiroOAuthRefreshAdapter::default().with_refresh_base_urls(None, Some(server_url));
|
||||
let transport = sample_transport(
|
||||
r#"{
|
||||
"auth_method":"identity_center",
|
||||
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr",
|
||||
"client_id":"cid",
|
||||
"client_secret":"secret",
|
||||
"profile_arn":"arn:aws:bedrock:demo"
|
||||
}"#,
|
||||
);
|
||||
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
|
||||
|
||||
let entry = adapter
|
||||
.refresh(&executor, &transport, None)
|
||||
.await
|
||||
.expect("refresh should succeed")
|
||||
.expect("cached entry should exist");
|
||||
let resolved = adapter
|
||||
.resolve_cached(&transport, &entry)
|
||||
.expect("cached entry should resolve");
|
||||
let seen_request = seen_request
|
||||
.lock()
|
||||
.expect("mutex should lock")
|
||||
.clone()
|
||||
.expect("refresh request should be captured");
|
||||
|
||||
assert_eq!(
|
||||
seen_request.body["grantType"].as_str(),
|
||||
Some("refresh_token")
|
||||
);
|
||||
assert_eq!(seen_request.body["clientId"].as_str(), Some("cid"));
|
||||
assert_eq!(seen_request.user_agent, "node");
|
||||
assert_eq!(seen_request.x_amz_user_agent, IDC_AMZ_USER_AGENT);
|
||||
assert!(!seen_request.host.trim().is_empty());
|
||||
match resolved {
|
||||
LocalResolvedOAuthRequestAuth::Kiro(auth) => {
|
||||
assert_eq!(auth.value, "Bearer cached-idc-access-token");
|
||||
assert!(auth.auth_config.profile_arn_for_payload().is_none());
|
||||
assert!(auth.auth_config.expires_at.is_some());
|
||||
}
|
||||
other => panic!("unexpected resolved auth: {other:?}"),
|
||||
}
|
||||
|
||||
server_handle.abort();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,309 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub use super::super::rules::{
|
||||
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
|
||||
body_rules_are_locally_supported, header_rules_are_locally_supported,
|
||||
};
|
||||
use super::super::should_skip_upstream_passthrough_header;
|
||||
use super::converter::convert_claude_messages_to_conversation_state;
|
||||
use super::credentials::KiroAuthConfig;
|
||||
use super::headers::build_generate_assistant_headers;
|
||||
|
||||
pub fn supports_local_kiro_request_shape(
|
||||
header_rules: Option<&Value>,
|
||||
body_rules: Option<&Value>,
|
||||
) -> bool {
|
||||
header_rules_are_locally_supported(header_rules) && body_rules_are_locally_supported(body_rules)
|
||||
}
|
||||
|
||||
pub fn build_kiro_provider_request_body(
|
||||
body_json: &Value,
|
||||
mapped_model: &str,
|
||||
auth_config: &KiroAuthConfig,
|
||||
body_rules: Option<&Value>,
|
||||
request_headers: Option<&http::HeaderMap>,
|
||||
) -> Option<Value> {
|
||||
let conversation_state =
|
||||
convert_claude_messages_to_conversation_state(body_json, mapped_model)?;
|
||||
let mut provider_request_body = json!({
|
||||
"conversationState": conversation_state
|
||||
});
|
||||
|
||||
let mut inference_config = serde_json::Map::new();
|
||||
if let Some(max_tokens) = body_json
|
||||
.get("max_tokens")
|
||||
.and_then(|value| {
|
||||
value
|
||||
.as_i64()
|
||||
.or_else(|| value.as_u64().map(|value| value as i64))
|
||||
})
|
||||
.filter(|value| *value > 0)
|
||||
{
|
||||
inference_config.insert("maxTokens".to_string(), Value::from(max_tokens));
|
||||
}
|
||||
if let Some(temperature) = body_json
|
||||
.get("temperature")
|
||||
.and_then(Value::as_f64)
|
||||
.filter(|value| *value >= 0.0)
|
||||
{
|
||||
inference_config.insert("temperature".to_string(), Value::from(temperature));
|
||||
}
|
||||
if let Some(top_p) = body_json
|
||||
.get("top_p")
|
||||
.and_then(Value::as_f64)
|
||||
.filter(|value| *value > 0.0)
|
||||
{
|
||||
inference_config.insert("topP".to_string(), Value::from(top_p));
|
||||
}
|
||||
if !inference_config.is_empty() {
|
||||
provider_request_body.as_object_mut()?.insert(
|
||||
"inferenceConfig".to_string(),
|
||||
Value::Object(inference_config),
|
||||
);
|
||||
}
|
||||
|
||||
if let Some(profile_arn) = auth_config.profile_arn_for_payload() {
|
||||
provider_request_body.as_object_mut()?.insert(
|
||||
"profileArn".to_string(),
|
||||
Value::String(profile_arn.to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
if !apply_local_body_rules_with_request_headers(
|
||||
&mut provider_request_body,
|
||||
body_rules,
|
||||
Some(body_json),
|
||||
request_headers,
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(provider_request_body)
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug)]
|
||||
pub struct KiroProviderHeadersInput<'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 auth_header: &'a str,
|
||||
pub auth_value: &'a str,
|
||||
pub auth_config: &'a KiroAuthConfig,
|
||||
pub machine_id: &'a str,
|
||||
}
|
||||
|
||||
pub fn build_kiro_provider_headers(
|
||||
input: KiroProviderHeadersInput<'_>,
|
||||
) -> Option<BTreeMap<String, String>> {
|
||||
let KiroProviderHeadersInput {
|
||||
headers,
|
||||
provider_request_body,
|
||||
original_request_body,
|
||||
header_rules,
|
||||
auth_header,
|
||||
auth_value,
|
||||
auth_config,
|
||||
machine_id,
|
||||
} = input;
|
||||
|
||||
let mut out = BTreeMap::new();
|
||||
for (name, value) in headers {
|
||||
let Ok(value) = value.to_str() else {
|
||||
continue;
|
||||
};
|
||||
let key = name.as_str().to_ascii_lowercase();
|
||||
if should_skip_upstream_passthrough_header(&key) {
|
||||
continue;
|
||||
}
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
continue;
|
||||
}
|
||||
out.insert(key, value.to_string());
|
||||
}
|
||||
|
||||
if !apply_local_header_rules_with_request_headers(
|
||||
&mut out,
|
||||
header_rules,
|
||||
&[auth_header, "content-type"],
|
||||
provider_request_body,
|
||||
Some(original_request_body),
|
||||
Some(headers),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
|
||||
for (key, value) in build_generate_assistant_headers(auth_config, machine_id) {
|
||||
out.insert(key, value);
|
||||
}
|
||||
out.insert(
|
||||
auth_header.trim().to_ascii_lowercase(),
|
||||
auth_value.trim().to_string(),
|
||||
);
|
||||
out.entry("content-type".to_string())
|
||||
.or_insert_with(|| "application/json".to_string());
|
||||
out.remove("content-length");
|
||||
Some(out)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::credentials::KiroAuthConfig;
|
||||
use super::{
|
||||
build_kiro_provider_headers, build_kiro_provider_request_body,
|
||||
supports_local_kiro_request_shape, KiroProviderHeadersInput,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn supports_empty_local_request_shape() {
|
||||
assert!(supports_local_kiro_request_shape(None, None));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unsupported_rule_shape() {
|
||||
assert!(!supports_local_kiro_request_shape(
|
||||
Some(&json!({"action":"set"})),
|
||||
None
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_simple_header_and_body_rules() {
|
||||
assert!(supports_local_kiro_request_shape(
|
||||
Some(&json!([{"action":"set","key":"x-provider-extra","value":"1"}])),
|
||||
Some(&json!([{"action":"set","path":"debugTag","value":true}]))
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn wraps_claude_request_into_kiro_payload_before_body_rules() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: None,
|
||||
refresh_token: Some("r".repeat(128)),
|
||||
expires_at: None,
|
||||
profile_arn: Some("arn:aws:bedrock:demo".to_string()),
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: Some("us-east-1".to_string()),
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: Some("123e4567-e89b-12d3-a456-426614174000".to_string()),
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: Some("cached-token".to_string()),
|
||||
};
|
||||
|
||||
let payload = build_kiro_provider_request_body(
|
||||
&json!({
|
||||
"messages": [{"role":"user","content":"hello"}],
|
||||
"max_tokens": 64
|
||||
}),
|
||||
"claude-sonnet-4-upstream",
|
||||
&auth_config,
|
||||
Some(&json!([
|
||||
{"action":"set","path":"debugTag","value":"kiro-local"}
|
||||
])),
|
||||
None,
|
||||
)
|
||||
.expect("payload should build");
|
||||
|
||||
assert!(payload.get("conversationState").is_some());
|
||||
assert_eq!(
|
||||
payload
|
||||
.get("inferenceConfig")
|
||||
.and_then(|value| value.get("maxTokens")),
|
||||
Some(&json!(64))
|
||||
);
|
||||
assert_eq!(
|
||||
payload.get("profileArn"),
|
||||
Some(&json!("arn:aws:bedrock:demo"))
|
||||
);
|
||||
assert_eq!(payload.get("debugTag"), Some(&json!("kiro-local")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn applies_header_rules_before_kiro_extra_headers() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: None,
|
||||
refresh_token: Some("r".repeat(128)),
|
||||
expires_at: None,
|
||||
profile_arn: None,
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: Some("us-east-1".to_string()),
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: None,
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: Some("cached-token".to_string()),
|
||||
};
|
||||
let headers = build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
headers: &http::HeaderMap::new(),
|
||||
provider_request_body: &json!({"conversationState": {}}),
|
||||
original_request_body: &json!({"messages": []}),
|
||||
header_rules: Some(&json!([
|
||||
{"action":"set","key":"accept","value":"text/plain"},
|
||||
{"action":"set","key":"x-endpoint-tag","value":"kiro-local"}
|
||||
])),
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer cached-token",
|
||||
auth_config: &auth_config,
|
||||
machine_id: "machine-123",
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("accept").map(String::as_str),
|
||||
Some("application/vnd.amazon.eventstream")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer cached-token")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("x-endpoint-tag").map(String::as_str),
|
||||
Some("kiro-local")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn omits_profile_arn_for_idc_auth() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: Some("identity_center".to_string()),
|
||||
refresh_token: Some("r".repeat(128)),
|
||||
expires_at: None,
|
||||
profile_arn: Some("arn:aws:bedrock:demo".to_string()),
|
||||
region: None,
|
||||
auth_region: None,
|
||||
api_region: Some("us-east-1".to_string()),
|
||||
client_id: Some("cid".to_string()),
|
||||
client_secret: Some("secret".to_string()),
|
||||
machine_id: None,
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: Some("cached-token".to_string()),
|
||||
};
|
||||
|
||||
let payload = build_kiro_provider_request_body(
|
||||
&json!({
|
||||
"messages": [{"role":"user","content":"hello"}]
|
||||
}),
|
||||
"claude-sonnet-4-upstream",
|
||||
&auth_config,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("payload should build");
|
||||
|
||||
assert!(payload.get("profileArn").is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
use super::super::url::build_passthrough_path_url;
|
||||
use super::credentials::DEFAULT_REGION;
|
||||
|
||||
pub const GENERATE_ASSISTANT_RESPONSE_PATH: &str = "/generateAssistantResponse";
|
||||
pub const LIST_AVAILABLE_MODELS_PATH: &str = "/ListAvailableModels";
|
||||
pub const MCP_PATH: &str = "/mcp";
|
||||
pub const MCP_STREAM_PATH: &str = "/mcp/stream";
|
||||
pub const KIRO_ENVELOPE_NAME: &str = "kiro:generateAssistantResponse";
|
||||
|
||||
pub fn resolve_kiro_base_url(upstream_base_url: &str, api_region: Option<&str>) -> String {
|
||||
let region = api_region
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or(DEFAULT_REGION);
|
||||
upstream_base_url
|
||||
.trim()
|
||||
.replace("{region}", region)
|
||||
.trim_end_matches('/')
|
||||
.to_string()
|
||||
}
|
||||
|
||||
pub fn build_kiro_generate_assistant_response_url(
|
||||
upstream_base_url: &str,
|
||||
query: Option<&str>,
|
||||
api_region: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let upstream_base_url = resolve_kiro_base_url(upstream_base_url, api_region);
|
||||
build_passthrough_path_url(
|
||||
upstream_base_url.as_str(),
|
||||
GENERATE_ASSISTANT_RESPONSE_PATH,
|
||||
query,
|
||||
&[],
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_kiro_mcp_url(upstream_base_url: &str, api_region: Option<&str>) -> Option<String> {
|
||||
let upstream_base_url = resolve_kiro_base_url(upstream_base_url, api_region);
|
||||
build_passthrough_path_url(upstream_base_url.as_str(), MCP_PATH, None, &[])
|
||||
}
|
||||
|
||||
pub fn build_kiro_list_available_models_url(
|
||||
upstream_base_url: &str,
|
||||
api_region: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let mut serializer = url::form_urlencoded::Serializer::new(String::new());
|
||||
serializer.append_pair("origin", "AI_EDITOR");
|
||||
let query = serializer.finish();
|
||||
let upstream_base_url = resolve_kiro_base_url(upstream_base_url, api_region);
|
||||
build_passthrough_path_url(
|
||||
upstream_base_url.as_str(),
|
||||
LIST_AVAILABLE_MODELS_PATH,
|
||||
Some(&query),
|
||||
&[],
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_kiro_mcp_url_from_resolved_url(resolved_url: &str) -> Option<String> {
|
||||
let mut parsed = url::Url::parse(resolved_url).ok()?;
|
||||
parsed.set_path(MCP_PATH);
|
||||
parsed.set_query(None);
|
||||
Some(parsed.to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_kiro_generate_assistant_response_url, build_kiro_list_available_models_url,
|
||||
build_kiro_mcp_url, build_kiro_mcp_url_from_resolved_url, resolve_kiro_base_url,
|
||||
GENERATE_ASSISTANT_RESPONSE_PATH, KIRO_ENVELOPE_NAME, LIST_AVAILABLE_MODELS_PATH, MCP_PATH,
|
||||
MCP_STREAM_PATH,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn exposes_kiro_request_constants() {
|
||||
assert_eq!(
|
||||
GENERATE_ASSISTANT_RESPONSE_PATH,
|
||||
"/generateAssistantResponse"
|
||||
);
|
||||
assert_eq!(LIST_AVAILABLE_MODELS_PATH, "/ListAvailableModels");
|
||||
assert_eq!(MCP_PATH, "/mcp");
|
||||
assert_eq!(MCP_STREAM_PATH, "/mcp/stream");
|
||||
assert_eq!(KIRO_ENVELOPE_NAME, "kiro:generateAssistantResponse");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_generate_assistant_response_url() {
|
||||
assert_eq!(
|
||||
build_kiro_generate_assistant_response_url(
|
||||
"https://kiro.{region}.example?tenant=demo",
|
||||
Some("stream=true"),
|
||||
Some("us-west-2")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://kiro.us-west-2.example/generateAssistantResponse?stream=true&tenant=demo"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_region_placeholder_in_base_url() {
|
||||
assert_eq!(
|
||||
resolve_kiro_base_url("https://kiro.{region}.example/", Some("us-west-2")),
|
||||
"https://kiro.us-west-2.example"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_mcp_url_for_latest_kiro_endpoint() {
|
||||
assert_eq!(
|
||||
build_kiro_mcp_url("https://q.{region}.amazonaws.com", Some("eu-west-1")).as_deref(),
|
||||
Some("https://q.eu-west-1.amazonaws.com/mcp")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_list_available_models_url_for_profile() {
|
||||
assert_eq!(
|
||||
build_kiro_list_available_models_url(
|
||||
"https://q.{region}.amazonaws.com",
|
||||
Some("us-west-2")
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://q.us-west-2.amazonaws.com/ListAvailableModels?origin=AI_EDITOR")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rewrites_generate_assistant_url_to_mcp_url() {
|
||||
assert_eq!(
|
||||
build_kiro_mcp_url_from_resolved_url(
|
||||
"https://q.us-east-1.amazonaws.com/generateAssistantResponse?beta=true"
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://q.us-east-1.amazonaws.com/mcp")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
pub mod antigravity;
|
||||
pub mod auth;
|
||||
mod auth_config;
|
||||
mod cache;
|
||||
pub mod claude_code;
|
||||
pub mod conversion;
|
||||
mod diagnostics;
|
||||
pub mod gemini_cli;
|
||||
mod gemini_files;
|
||||
mod generic_oauth;
|
||||
pub mod grok;
|
||||
mod headers;
|
||||
pub mod kiro;
|
||||
mod network;
|
||||
pub mod oauth_refresh;
|
||||
mod openai_image;
|
||||
pub mod policy;
|
||||
pub mod provider_types;
|
||||
mod request_body;
|
||||
mod request_url;
|
||||
pub mod rules;
|
||||
pub mod same_format_provider;
|
||||
pub mod snapshot;
|
||||
mod standard;
|
||||
pub mod url;
|
||||
pub mod vertex;
|
||||
mod video;
|
||||
pub mod windsurf;
|
||||
|
||||
pub use aether_oauth as oauth;
|
||||
pub use auth::{build_passthrough_headers, ensure_upstream_auth_header};
|
||||
pub use auth_config::apply_local_auth_config_header_overrides;
|
||||
pub use cache::{provider_transport_snapshot_looks_refreshed, ProviderTransportSnapshotCacheKey};
|
||||
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, request_pair_direct_auth,
|
||||
request_pair_transport_unsupported_reason, CandidateTransportPolicyFacts,
|
||||
};
|
||||
pub use diagnostics::{
|
||||
append_transport_diagnostics_to_value, build_request_trace_proxy_value,
|
||||
build_transport_diagnostics,
|
||||
};
|
||||
pub use gemini_cli::{
|
||||
build_gemini_cli_v1internal_request, build_gemini_cli_v1internal_url,
|
||||
classify_gemini_cli_v1internal_request_body, gemini_cli_v1internal_requires_upstream_streaming,
|
||||
is_gemini_cli_provider_transport, resolve_gemini_cli_project_id,
|
||||
resolve_local_gemini_cli_request_auth, GeminiCliRequestAuth, GeminiCliRequestAuthSupport,
|
||||
GeminiCliRequestAuthUnsupportedReason, GeminiCliRequestEnvelopeSupport,
|
||||
GeminiCliRequestEnvelopeUnsupportedReason, GeminiCliRequestUrlAction, GEMINI_CLI_PROVIDER_TYPE,
|
||||
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT,
|
||||
GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME, GEMINI_CLI_V1INTERNAL_PATH_TEMPLATE,
|
||||
};
|
||||
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,
|
||||
};
|
||||
pub use grok::{
|
||||
build_grok_app_chat_body, build_grok_browser_headers, build_grok_upstream_url, grok_base_url,
|
||||
grok_browser_profile_id_from_user_agent,
|
||||
grok_browser_profile_metadata_from_resolved_transport_profile,
|
||||
grok_browser_resolved_transport_profile,
|
||||
grok_browser_resolved_transport_profile_from_auth_config,
|
||||
grok_browser_transport_fingerprint_from_auth_config, is_grok_provider_transport,
|
||||
resolve_grok_session_auth, GrokBrowserProfileMetadata, GrokHeaderInput, GROK_CHAT_PATH,
|
||||
GROK_DEFAULT_BASE_URL, GROK_DEFAULT_BROWSER_PROFILE, GROK_DEFAULT_USER_AGENT,
|
||||
GROK_INTERNAL_HEADER, GROK_RATE_LIMITS_PATH,
|
||||
};
|
||||
pub use headers::{should_skip_request_header, should_skip_upstream_passthrough_header};
|
||||
pub use network::{
|
||||
resolve_transport_execution_timeouts, resolve_transport_profile, resolve_transport_profile_id,
|
||||
resolve_transport_proxy_snapshot, resolve_transport_proxy_snapshot_with_tunnel_affinity,
|
||||
transport_profile_is_configured, transport_proxy_is_locally_supported,
|
||||
TransportTunnelAffinityLookup, TransportTunnelAttachmentOwner,
|
||||
};
|
||||
pub use oauth_refresh::{
|
||||
supports_local_oauth_request_auth_resolution, CachedOAuthEntry, LocalOAuthHttpExecutor,
|
||||
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,
|
||||
local_openai_chat_transport_unsupported_reason, local_standard_transport_unsupported_reason,
|
||||
local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport,
|
||||
supports_local_gemini_transport_with_network, supports_local_standard_transport,
|
||||
};
|
||||
pub use request_body::{
|
||||
apply_transport_request_body_semantics, TransportRequestBodySemanticsError,
|
||||
};
|
||||
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,
|
||||
build_transport_request_url_for_request_body, gemini_embedding_request_body_uses_batch,
|
||||
TransportRequestUrlParams,
|
||||
};
|
||||
pub use rules::{
|
||||
apply_local_body_rules, apply_local_body_rules_with_request_headers, apply_local_header_rules,
|
||||
apply_local_header_rules_with_request_headers, body_rules_are_locally_supported,
|
||||
body_rules_handle_path, body_rules_have_enabled_rules, header_rules_are_locally_supported,
|
||||
header_rules_have_enabled_rules,
|
||||
};
|
||||
pub use same_format_provider::{
|
||||
build_same_format_provider_headers, build_same_format_provider_request_body,
|
||||
build_same_format_provider_request_body_with_compatibility_report,
|
||||
build_same_format_provider_upstream_url, classify_same_format_provider_request_behavior,
|
||||
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, SameFormatProviderCompatibilityEdit,
|
||||
SameFormatProviderCompatibilityEditAction, SameFormatProviderFamily,
|
||||
SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior,
|
||||
SameFormatProviderRequestBehaviorParams, SameFormatProviderRequestBodyInput,
|
||||
SameFormatProviderRequestBodyOutput, SameFormatProviderUpstreamUrlParams,
|
||||
};
|
||||
pub use snapshot::{
|
||||
read_provider_transport_snapshot, GatewayProviderTransportSnapshot,
|
||||
ProviderTransportSnapshotSource,
|
||||
};
|
||||
pub use standard::{
|
||||
apply_standard_provider_request_body_rules,
|
||||
apply_standard_provider_request_body_rules_with_request_headers,
|
||||
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, is_vertex_service_account_transport_context,
|
||||
is_vertex_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,
|
||||
resolve_video_create_auth, video_create_transport_unsupported_reason,
|
||||
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, VideoTaskTransportSnapshotLookup,
|
||||
};
|
||||
pub use windsurf::{
|
||||
build_windsurf_cascade_headers, build_windsurf_cascade_request_body,
|
||||
build_windsurf_cascade_upstream_url, is_windsurf_provider_transport,
|
||||
local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH,
|
||||
WINDSURF_ENVELOPE_NAME,
|
||||
};
|
||||
@@ -0,0 +1,848 @@
|
||||
use aether_contracts::{
|
||||
ExecutionTimeouts, ProxySnapshot, ResolvedTransportProfile, TRANSPORT_BACKEND_REQWEST_RUSTLS,
|
||||
TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::{json, Map, Value};
|
||||
use tracing::warn;
|
||||
|
||||
use crate::grok::grok_browser_resolved_transport_profile_from_auth_config;
|
||||
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
const TUNNEL_BASE_URL_EXTRA_KEY: &str = "tunnel_base_url";
|
||||
const TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY: &str = "tunnel_owner_instance_id";
|
||||
const TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY: &str = "tunnel_owner_observed_at_unix_secs";
|
||||
const DEFAULT_PROVIDER_STREAM_FIRST_BYTE_TIMEOUT_SECS: f64 = 30.0;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct TransportTunnelAttachmentOwner {
|
||||
pub gateway_instance_id: String,
|
||||
pub relay_base_url: String,
|
||||
pub observed_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait TransportTunnelAffinityLookup: Send + Sync {
|
||||
async fn lookup_tunnel_attachment_owner(
|
||||
&self,
|
||||
node_id: &str,
|
||||
) -> Result<Option<TransportTunnelAttachmentOwner>, String>;
|
||||
}
|
||||
|
||||
pub fn resolve_transport_execution_timeouts(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ExecutionTimeouts> {
|
||||
Some(ExecutionTimeouts {
|
||||
total_ms: transport
|
||||
.provider
|
||||
.request_timeout_secs
|
||||
.filter(|value| value.is_finite() && *value > 0.0)
|
||||
.map(timeout_secs_to_ms),
|
||||
first_byte_ms: Some(timeout_secs_to_ms(
|
||||
transport
|
||||
.provider
|
||||
.stream_first_byte_timeout_secs
|
||||
.filter(|value| value.is_finite() && *value > 0.0)
|
||||
.unwrap_or(DEFAULT_PROVIDER_STREAM_FIRST_BYTE_TIMEOUT_SECS),
|
||||
)),
|
||||
..ExecutionTimeouts::default()
|
||||
})
|
||||
}
|
||||
|
||||
fn timeout_secs_to_ms(secs: f64) -> u64 {
|
||||
((secs * 1000.0).round() as u64).max(1)
|
||||
}
|
||||
|
||||
pub fn resolve_transport_proxy_snapshot(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ProxySnapshot> {
|
||||
let raw = effective_proxy_config(transport)?;
|
||||
proxy_snapshot_from_value(raw)
|
||||
}
|
||||
|
||||
pub async fn resolve_transport_proxy_snapshot_with_tunnel_affinity(
|
||||
lookup: &dyn TransportTunnelAffinityLookup,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ProxySnapshot> {
|
||||
let mut snapshot = resolve_transport_proxy_snapshot(transport)?;
|
||||
let Some(node_id) = snapshot
|
||||
.node_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Some(snapshot);
|
||||
};
|
||||
|
||||
let owner = match lookup.lookup_tunnel_attachment_owner(node_id).await {
|
||||
Ok(owner) => owner,
|
||||
Err(error) => {
|
||||
warn!(error = %error, node_id = node_id, "failed to load tunnel attachment owner");
|
||||
None
|
||||
}
|
||||
};
|
||||
let Some(owner) = owner else {
|
||||
return Some(snapshot);
|
||||
};
|
||||
|
||||
let mut extra = snapshot
|
||||
.extra
|
||||
.take()
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_default();
|
||||
let configured_tunnel_base_url = extra
|
||||
.get(TUNNEL_BASE_URL_EXTRA_KEY)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
if configured_tunnel_base_url.is_none() {
|
||||
extra.insert(
|
||||
TUNNEL_BASE_URL_EXTRA_KEY.to_string(),
|
||||
Value::String(owner.relay_base_url.clone()),
|
||||
);
|
||||
}
|
||||
extra.insert(
|
||||
TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY.to_string(),
|
||||
Value::String(owner.gateway_instance_id),
|
||||
);
|
||||
extra.insert(
|
||||
TUNNEL_OWNER_OBSERVED_AT_EXTRA_KEY.to_string(),
|
||||
json!(owner.observed_at_unix_secs),
|
||||
);
|
||||
snapshot.extra = Some(Value::Object(extra));
|
||||
Some(snapshot)
|
||||
}
|
||||
|
||||
pub fn transport_proxy_is_locally_supported(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
let has_configured_proxy = transport.provider.proxy.is_some()
|
||||
|| transport.endpoint.proxy.is_some()
|
||||
|| transport.key.proxy.is_some();
|
||||
if !has_configured_proxy {
|
||||
return true;
|
||||
}
|
||||
|
||||
let Some(snapshot) = resolve_transport_proxy_snapshot(transport) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if snapshot.enabled == Some(false) {
|
||||
return true;
|
||||
}
|
||||
|
||||
snapshot
|
||||
.url
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
|| snapshot
|
||||
.node_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub fn resolve_transport_profile_id(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<String> {
|
||||
resolve_transport_profile(transport).map(|profile| profile.profile_id)
|
||||
}
|
||||
|
||||
pub fn resolve_transport_profile(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
resolve_transport_profile_from_fingerprint(transport.key.fingerprint.as_ref()).or_else(|| {
|
||||
resolve_transport_profile_from_provider_config(transport.provider.config.as_ref())
|
||||
.or_else(|| resolve_grok_browser_transport_profile(transport))
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_grok_browser_transport_profile(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
let auth_config = transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok())?;
|
||||
let object = auth_config.as_object()?;
|
||||
let has_session = json_string_field(object, "sso_token")
|
||||
.or_else(|| json_string_field(object, "access_token"))
|
||||
.or_else(|| json_string_field(object, "token"))
|
||||
.is_some();
|
||||
if !has_session {
|
||||
return None;
|
||||
}
|
||||
grok_browser_resolved_transport_profile_from_auth_config(object, "grok_auth_config")
|
||||
}
|
||||
|
||||
fn resolve_transport_profile_from_provider_config(
|
||||
config: Option<&Value>,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
let fingerprint = config?.get("fingerprint");
|
||||
resolve_transport_profile_from_fingerprint(fingerprint)
|
||||
}
|
||||
|
||||
fn resolve_transport_profile_from_fingerprint(
|
||||
fingerprint: Option<&Value>,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
let fingerprint = fingerprint?;
|
||||
fingerprint
|
||||
.get("transport_profile")
|
||||
.and_then(parse_transport_profile_value)
|
||||
}
|
||||
|
||||
pub fn transport_profile_is_configured(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport_profile_configured_in_fingerprint(transport.key.fingerprint.as_ref())
|
||||
|| transport_profile_configured_in_provider_config(transport.provider.config.as_ref())
|
||||
}
|
||||
|
||||
fn transport_profile_configured_in_provider_config(config: Option<&Value>) -> bool {
|
||||
let fingerprint = config.and_then(|value| value.get("fingerprint"));
|
||||
transport_profile_configured_in_fingerprint(fingerprint)
|
||||
}
|
||||
|
||||
fn transport_profile_configured_in_fingerprint(fingerprint: Option<&Value>) -> bool {
|
||||
fingerprint
|
||||
.and_then(|value| value.get("transport_profile"))
|
||||
.is_some_and(|value| !value.is_null())
|
||||
}
|
||||
|
||||
fn parse_transport_profile_value(value: &Value) -> Option<ResolvedTransportProfile> {
|
||||
if let Some(profile_id) = value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return Some(ResolvedTransportProfile {
|
||||
profile_id: profile_id.to_string(),
|
||||
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: None,
|
||||
});
|
||||
}
|
||||
|
||||
let object = value.as_object()?;
|
||||
let profile_id = json_string_field(object, "profile_id")
|
||||
.or_else(|| json_string_field(object, "id"))
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())?;
|
||||
let backend = json_string_field(object, "backend")
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or_else(|| TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string());
|
||||
let http_mode = json_string_field(object, "http_mode")
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or_else(|| TRANSPORT_HTTP_MODE_AUTO.to_string());
|
||||
let pool_scope = json_string_field(object, "pool_scope")
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or_else(|| TRANSPORT_POOL_SCOPE_KEY.to_string());
|
||||
let header_fingerprint = object.get("header_fingerprint").cloned();
|
||||
let extra = object.get("extra").cloned();
|
||||
|
||||
Some(ResolvedTransportProfile {
|
||||
profile_id,
|
||||
backend,
|
||||
http_mode,
|
||||
pool_scope,
|
||||
header_fingerprint,
|
||||
extra,
|
||||
})
|
||||
}
|
||||
|
||||
fn effective_proxy_config(transport: &GatewayProviderTransportSnapshot) -> Option<&Value> {
|
||||
[
|
||||
transport.key.proxy.as_ref(),
|
||||
transport.endpoint.proxy.as_ref(),
|
||||
transport.provider.proxy.as_ref(),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.find(|candidate| proxy_enabled(candidate))
|
||||
}
|
||||
|
||||
fn proxy_enabled(value: &Value) -> bool {
|
||||
value
|
||||
.as_object()
|
||||
.and_then(|object| object.get("enabled"))
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
fn proxy_snapshot_from_value(value: &Value) -> Option<ProxySnapshot> {
|
||||
let object = value.as_object()?;
|
||||
let enabled = object.get("enabled").and_then(Value::as_bool);
|
||||
let mode = json_string_field(object, "mode");
|
||||
let node_id = json_string_field(object, "node_id");
|
||||
let label = json_string_field(object, "label");
|
||||
let url = json_string_field(object, "url").or_else(|| json_string_field(object, "proxy_url"));
|
||||
|
||||
let mut extra = Map::new();
|
||||
for (key, value) in object {
|
||||
if matches!(
|
||||
key.as_str(),
|
||||
"enabled" | "mode" | "node_id" | "label" | "url" | "proxy_url"
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
extra.insert(key.clone(), value.clone());
|
||||
}
|
||||
|
||||
Some(ProxySnapshot {
|
||||
enabled,
|
||||
mode,
|
||||
node_id,
|
||||
label,
|
||||
url,
|
||||
extra: if extra.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(Value::Object(extra))
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
fn json_string_field(object: &Map<String, Value>, key: &str) -> Option<String> {
|
||||
object
|
||||
.get(key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::{
|
||||
resolve_transport_execution_timeouts, resolve_transport_profile,
|
||||
resolve_transport_profile_id, resolve_transport_proxy_snapshot,
|
||||
resolve_transport_proxy_snapshot_with_tunnel_affinity, transport_profile_is_configured,
|
||||
transport_proxy_is_locally_supported, TransportTunnelAffinityLookup,
|
||||
TransportTunnelAttachmentOwner,
|
||||
};
|
||||
use aether_contracts::TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE;
|
||||
|
||||
#[derive(Default)]
|
||||
struct TestTunnelAffinityLookup {
|
||||
owners: BTreeMap<String, TransportTunnelAttachmentOwner>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl TransportTunnelAffinityLookup for TestTunnelAffinityLookup {
|
||||
async fn lookup_tunnel_attachment_owner(
|
||||
&self,
|
||||
node_id: &str,
|
||||
) -> Result<Option<TransportTunnelAttachmentOwner>, String> {
|
||||
Ok(self.owners.get(node_id).cloned())
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_lookup() -> TestTunnelAffinityLookup {
|
||||
let mut owners = BTreeMap::new();
|
||||
owners.insert(
|
||||
"proxy-node-1".to_string(),
|
||||
TransportTunnelAttachmentOwner {
|
||||
gateway_instance_id: "gateway-b".to_string(),
|
||||
relay_base_url: "http://gateway-b.internal".to_string(),
|
||||
observed_at_unix_secs: 4_102_444_800u64,
|
||||
},
|
||||
);
|
||||
TestTunnelAffinityLookup { owners }
|
||||
}
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "provider".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: Some(json!({"url":"http://provider-proxy:8080"})),
|
||||
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: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://api.openai.example".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: Some(json!({"enabled":false,"url":"http://endpoint-proxy:8080"})),
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: Some(json!({"node_id":"proxy-node-1","kind":"manual"})),
|
||||
fingerprint: Some(json!({"transport_profile":"chrome_136"})),
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "sk-test".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_execution_timeouts_use_provider_defaults_when_unset() {
|
||||
let transport = sample_transport();
|
||||
|
||||
let timeouts = resolve_transport_execution_timeouts(&transport)
|
||||
.expect("default provider timeouts should resolve");
|
||||
|
||||
assert_eq!(timeouts.total_ms, None);
|
||||
assert_eq!(timeouts.first_byte_ms, Some(30_000));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_execution_timeouts_preserve_configured_values_independently() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.request_timeout_secs = Some(12.0);
|
||||
|
||||
let timeouts = resolve_transport_execution_timeouts(&transport)
|
||||
.expect("provider timeouts should resolve");
|
||||
|
||||
assert_eq!(timeouts.total_ms, Some(12_000));
|
||||
assert_eq!(timeouts.first_byte_ms, Some(30_000));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_execution_timeouts_preserve_the_configurable_maximum() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.request_timeout_secs =
|
||||
Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS as f64);
|
||||
|
||||
let timeouts = resolve_transport_execution_timeouts(&transport)
|
||||
.expect("provider timeouts should resolve");
|
||||
|
||||
assert_eq!(
|
||||
timeouts.total_ms,
|
||||
Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_MS)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_execution_timeouts_preserve_configured_first_byte_value() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.stream_first_byte_timeout_secs = Some(7.5);
|
||||
|
||||
let timeouts = resolve_transport_execution_timeouts(&transport)
|
||||
.expect("provider timeouts should resolve");
|
||||
|
||||
assert_eq!(timeouts.total_ms, None);
|
||||
assert_eq!(timeouts.first_byte_ms, Some(7_500));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_transport_proxy_with_key_precedence() {
|
||||
let snapshot = resolve_transport_proxy_snapshot(&sample_transport())
|
||||
.expect("proxy snapshot should resolve");
|
||||
assert_eq!(snapshot.node_id.as_deref(), Some("proxy-node-1"));
|
||||
assert_eq!(snapshot.url, None);
|
||||
assert_eq!(snapshot.extra, Some(json!({"kind":"manual"})));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enriches_transport_proxy_snapshot_with_tunnel_owner_hint() {
|
||||
let state = sample_lookup();
|
||||
|
||||
let snapshot =
|
||||
resolve_transport_proxy_snapshot_with_tunnel_affinity(&state, &sample_transport())
|
||||
.await
|
||||
.expect("proxy snapshot should resolve");
|
||||
|
||||
assert_eq!(snapshot.node_id.as_deref(), Some("proxy-node-1"));
|
||||
assert_eq!(
|
||||
snapshot
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("tunnel_base_url"))
|
||||
.and_then(Value::as_str),
|
||||
Some("http://gateway-b.internal")
|
||||
);
|
||||
assert_eq!(
|
||||
snapshot
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("tunnel_owner_instance_id"))
|
||||
.and_then(Value::as_str),
|
||||
Some("gateway-b")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn preserves_explicit_tunnel_base_url_when_owner_hint_exists() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.proxy = Some(json!({
|
||||
"node_id": "proxy-node-1",
|
||||
"kind": "manual",
|
||||
"tunnel_base_url": "http://configured-gateway.internal",
|
||||
}));
|
||||
let state = sample_lookup();
|
||||
|
||||
let snapshot = resolve_transport_proxy_snapshot_with_tunnel_affinity(&state, &transport)
|
||||
.await
|
||||
.expect("proxy snapshot should resolve");
|
||||
|
||||
assert_eq!(
|
||||
snapshot
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("tunnel_base_url"))
|
||||
.and_then(Value::as_str),
|
||||
Some("http://configured-gateway.internal")
|
||||
);
|
||||
assert_eq!(
|
||||
snapshot
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("tunnel_owner_instance_id"))
|
||||
.and_then(Value::as_str),
|
||||
Some("gateway-b")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_transport_profile_id_from_key_fingerprint() {
|
||||
assert_eq!(
|
||||
resolve_transport_profile_id(&sample_transport()).as_deref(),
|
||||
Some("chrome_136")
|
||||
);
|
||||
assert!(transport_proxy_is_locally_supported(&sample_transport()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_transport_profile_from_key_fingerprint_before_provider_default() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.config = Some(json!({
|
||||
"fingerprint": {"transport_profile": "provider_profile"}
|
||||
}));
|
||||
transport.key.fingerprint = Some(json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "key_profile",
|
||||
"backend": "reqwest_rustls",
|
||||
"http_mode": "http1_only"
|
||||
}
|
||||
}));
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "key_profile");
|
||||
assert_eq!(profile.backend, "reqwest_rustls");
|
||||
assert_eq!(profile.http_mode, "http1_only");
|
||||
assert_eq!(profile.pool_scope, "key");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_transport_profile_from_provider_default() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = Some(json!({
|
||||
"fingerprint": {"transport_profile": "provider_profile"}
|
||||
}));
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "provider_profile");
|
||||
assert_eq!(profile.backend, "reqwest_rustls");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn maps_string_transport_profile_to_resolved_profile() {
|
||||
let profile = resolve_transport_profile(&sample_transport()).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "chrome_136");
|
||||
assert_eq!(profile.backend, "reqwest_rustls");
|
||||
assert_eq!(profile.http_mode, "auto");
|
||||
assert_eq!(profile.pool_scope, "key");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_h2c_prior_knowledge_transport_profile() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.fingerprint = Some(json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "mock-h2c",
|
||||
"backend": "reqwest_rustls",
|
||||
"http_mode": "h2c_prior_knowledge"
|
||||
}
|
||||
}));
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "mock-h2c");
|
||||
assert_eq!(profile.http_mode, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_no_transport_profile_without_fingerprint_configuration() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
|
||||
assert!(resolve_transport_profile(&transport).is_none());
|
||||
assert!(!transport_profile_is_configured(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_grok_browser_transport_profile_from_session_auth_config() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token",
|
||||
"browser_profile": "chrome136",
|
||||
"cf_clearance": "clearance"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "chrome136");
|
||||
assert_eq!(profile.backend, "browser_wreq");
|
||||
assert_eq!(profile.http_mode, "auto");
|
||||
assert_eq!(profile.pool_scope, "key");
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("browser_profile"))
|
||||
.and_then(Value::as_str),
|
||||
Some("chrome136")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_grok_browser_transport_profile_default_from_session_auth_config() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "chrome136");
|
||||
assert_eq!(profile.backend, "browser_wreq");
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("source"))
|
||||
.and_then(Value::as_str),
|
||||
Some("grok_auth_config")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_grok_browser_transport_profile_normalizes_auth_config_alias() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token",
|
||||
"browser_profile": "Chrome-137"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "chrome137");
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("browser_profile"))
|
||||
.and_then(Value::as_str),
|
||||
Some("chrome137")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_grok_browser_transport_profile_from_legacy_user_agent() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token",
|
||||
"user_agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/137.0.0.0 Safari/537.36"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "chrome137");
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("browser_profile"))
|
||||
.and_then(Value::as_str),
|
||||
Some("chrome137")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn key_fingerprint_wins_over_grok_auth_config_fallback() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.provider.config = None;
|
||||
transport.key.fingerprint = Some(json!({
|
||||
"transport_profile": {
|
||||
"profile_id": "chrome136",
|
||||
"backend": "browser_wreq",
|
||||
"extra": {"browser_profile": "chrome136", "source": "key"}
|
||||
}
|
||||
}));
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token",
|
||||
"browser_profile": "chrome137"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "chrome136");
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("source"))
|
||||
.and_then(Value::as_str),
|
||||
Some("key")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_fingerprint_wins_over_grok_auth_config_fallback() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = Some(json!({
|
||||
"fingerprint": {
|
||||
"transport_profile": {
|
||||
"profile_id": "chrome136",
|
||||
"backend": "browser_wreq",
|
||||
"extra": {"browser_profile": "chrome136", "source": "provider"}
|
||||
}
|
||||
}
|
||||
}));
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token",
|
||||
"browser_profile": "chrome137"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let profile = resolve_transport_profile(&transport).expect("profile");
|
||||
|
||||
assert_eq!(profile.profile_id, "chrome136");
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("source"))
|
||||
.and_then(Value::as_str),
|
||||
Some("provider")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unsupported_grok_auth_config_browser_profile() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token",
|
||||
"browser_profile": "safari999"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(resolve_transport_profile(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unsupported_grok_auth_config_user_agent_profile() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.fingerprint = None;
|
||||
transport.provider.config = None;
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"sso_token": "sso-token",
|
||||
"user_agent": "Mozilla/5.0 Version/18.0 Safari/605.1.15"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(resolve_transport_profile(&transport).is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,791 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::fmt;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_oauth::core::OAuthError;
|
||||
use aether_oauth::network::{
|
||||
OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse, OAuthNetworkContext,
|
||||
};
|
||||
use aether_oauth::provider::ProviderOAuthTransportContext;
|
||||
use aether_runtime_state::RuntimeState;
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use super::generic_oauth::supports_local_generic_oauth_request_auth_resolution;
|
||||
pub use super::generic_oauth::GenericOAuthRefreshAdapter;
|
||||
use super::kiro::{
|
||||
supports_local_kiro_request_auth_resolution, KiroOAuthRefreshAdapter, KiroRequestAuth,
|
||||
};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::vertex::{
|
||||
supports_local_vertex_service_account_auth_resolution, VertexServiceAccountRefreshAdapter,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[allow(clippy::large_enum_variant)]
|
||||
pub enum LocalResolvedOAuthRequestAuth {
|
||||
#[allow(dead_code)]
|
||||
Header {
|
||||
name: String,
|
||||
value: String,
|
||||
},
|
||||
Kiro(KiroRequestAuth),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct LocalOAuthResolution {
|
||||
pub auth: Option<LocalResolvedOAuthRequestAuth>,
|
||||
pub refreshed_entry: Option<CachedOAuthEntry>,
|
||||
pub refresh_in_flight: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct CachedOAuthEntry {
|
||||
pub provider_type: String,
|
||||
pub auth_header_name: String,
|
||||
pub auth_header_value: String,
|
||||
pub expires_at_unix_secs: Option<u64>,
|
||||
pub metadata: Option<Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct LocalOAuthHttpRequest {
|
||||
pub request_id: &'static str,
|
||||
pub method: reqwest::Method,
|
||||
pub url: String,
|
||||
pub headers: BTreeMap<String, String>,
|
||||
pub json_body: Option<Value>,
|
||||
pub body_bytes: Option<Vec<u8>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct LocalOAuthHttpResponse {
|
||||
pub status_code: u16,
|
||||
pub body_text: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum LocalOAuthRefreshError {
|
||||
#[error("{provider_type} oauth refresh request failed: {source}")]
|
||||
Transport {
|
||||
provider_type: &'static str,
|
||||
#[source]
|
||||
source: reqwest::Error,
|
||||
},
|
||||
#[error("{provider_type} oauth refresh returned HTTP {status_code}: {body_excerpt}")]
|
||||
HttpStatus {
|
||||
provider_type: &'static str,
|
||||
status_code: u16,
|
||||
body_excerpt: String,
|
||||
},
|
||||
#[error("{provider_type} oauth refresh transport failed: {message}")]
|
||||
TransportMessage {
|
||||
provider_type: &'static str,
|
||||
message: String,
|
||||
},
|
||||
#[error("{provider_type} oauth refresh returned invalid response: {message}")]
|
||||
InvalidResponse {
|
||||
provider_type: &'static str,
|
||||
message: String,
|
||||
},
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait LocalOAuthHttpExecutor: Send + Sync {
|
||||
async fn execute(
|
||||
&self,
|
||||
provider_type: &'static str,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request: &LocalOAuthHttpRequest,
|
||||
) -> Result<LocalOAuthHttpResponse, LocalOAuthRefreshError>;
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ReqwestLocalOAuthHttpExecutor {
|
||||
client: reqwest::Client,
|
||||
}
|
||||
|
||||
impl ReqwestLocalOAuthHttpExecutor {
|
||||
pub fn new(client: reqwest::Client) -> Self {
|
||||
Self { client }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthHttpExecutor for ReqwestLocalOAuthHttpExecutor {
|
||||
async fn execute(
|
||||
&self,
|
||||
provider_type: &'static str,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
request: &LocalOAuthHttpRequest,
|
||||
) -> Result<LocalOAuthHttpResponse, LocalOAuthRefreshError> {
|
||||
let mut builder = self
|
||||
.client
|
||||
.request(request.method.clone(), request.url.as_str());
|
||||
for (name, value) in &request.headers {
|
||||
builder = builder.header(name, value);
|
||||
}
|
||||
if let Some(json_body) = request.json_body.as_ref() {
|
||||
builder = builder.json(json_body);
|
||||
} else if let Some(body_bytes) = request.body_bytes.as_ref() {
|
||||
builder = builder.body(body_bytes.clone());
|
||||
}
|
||||
|
||||
let response =
|
||||
builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|source| LocalOAuthRefreshError::Transport {
|
||||
provider_type,
|
||||
source,
|
||||
})?;
|
||||
let status_code = response.status().as_u16();
|
||||
let body_text =
|
||||
response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|source| LocalOAuthRefreshError::Transport {
|
||||
provider_type,
|
||||
source,
|
||||
})?;
|
||||
Ok(LocalOAuthHttpResponse {
|
||||
status_code,
|
||||
body_text,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) struct ProviderOAuthLocalHttpExecutor<'a> {
|
||||
provider_type: &'static str,
|
||||
transport: &'a GatewayProviderTransportSnapshot,
|
||||
inner: &'a dyn LocalOAuthHttpExecutor,
|
||||
}
|
||||
|
||||
impl<'a> ProviderOAuthLocalHttpExecutor<'a> {
|
||||
pub(crate) fn new(
|
||||
provider_type: &'static str,
|
||||
transport: &'a GatewayProviderTransportSnapshot,
|
||||
inner: &'a dyn LocalOAuthHttpExecutor,
|
||||
) -> Self {
|
||||
Self {
|
||||
provider_type,
|
||||
transport,
|
||||
inner,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthHttpExecutor for ProviderOAuthLocalHttpExecutor<'_> {
|
||||
async fn execute(&self, request: OAuthHttpRequest) -> Result<OAuthHttpResponse, OAuthError> {
|
||||
let response = self
|
||||
.inner
|
||||
.execute(
|
||||
self.provider_type,
|
||||
self.transport,
|
||||
&LocalOAuthHttpRequest {
|
||||
request_id: "provider-oauth:local-refresh-token",
|
||||
method: request.method,
|
||||
url: request.url,
|
||||
headers: request.headers,
|
||||
json_body: request.json_body,
|
||||
body_bytes: request.body_bytes,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.map_err(local_refresh_error_to_oauth_error)?;
|
||||
let json_body = serde_json::from_str::<Value>(&response.body_text).ok();
|
||||
Ok(OAuthHttpResponse {
|
||||
status_code: response.status_code,
|
||||
body_text: response.body_text,
|
||||
json_body,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn provider_oauth_transport_context_from_snapshot(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> ProviderOAuthTransportContext {
|
||||
ProviderOAuthTransportContext {
|
||||
provider_id: transport.provider.id.clone(),
|
||||
provider_type: transport.provider.provider_type.clone(),
|
||||
endpoint_id: Some(transport.endpoint.id.clone()),
|
||||
key_id: Some(transport.key.id.clone()),
|
||||
auth_type: Some(transport.key.auth_type.clone()),
|
||||
decrypted_api_key: Some(transport.key.decrypted_api_key.clone()),
|
||||
decrypted_auth_config: transport.key.decrypted_auth_config.clone(),
|
||||
provider_config: transport.provider.config.clone(),
|
||||
endpoint_config: transport.endpoint.config.clone(),
|
||||
key_config: None,
|
||||
network: OAuthNetworkContext::provider_operation(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn oauth_error_to_local_refresh_error(
|
||||
provider_type: &'static str,
|
||||
error: OAuthError,
|
||||
) -> LocalOAuthRefreshError {
|
||||
match error {
|
||||
OAuthError::HttpStatus {
|
||||
status_code,
|
||||
body_excerpt,
|
||||
} => LocalOAuthRefreshError::HttpStatus {
|
||||
provider_type,
|
||||
status_code,
|
||||
body_excerpt,
|
||||
},
|
||||
OAuthError::Transport(message) => LocalOAuthRefreshError::TransportMessage {
|
||||
provider_type,
|
||||
message,
|
||||
},
|
||||
OAuthError::InvalidRequest(message)
|
||||
| OAuthError::InvalidResponse(message)
|
||||
| OAuthError::Storage(message)
|
||||
| OAuthError::UnsupportedProvider(message) => LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type,
|
||||
message,
|
||||
},
|
||||
OAuthError::InvalidState => LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type,
|
||||
message: "oauth state is invalid or expired".to_string(),
|
||||
},
|
||||
OAuthError::EncryptionUnavailable => LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type,
|
||||
message: "oauth encryption unavailable".to_string(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn local_refresh_error_to_oauth_error(error: LocalOAuthRefreshError) -> OAuthError {
|
||||
match error {
|
||||
LocalOAuthRefreshError::Transport { source, .. } => {
|
||||
OAuthError::Transport(source.to_string())
|
||||
}
|
||||
LocalOAuthRefreshError::TransportMessage { message, .. } => OAuthError::Transport(message),
|
||||
LocalOAuthRefreshError::HttpStatus {
|
||||
status_code,
|
||||
body_excerpt,
|
||||
..
|
||||
} => OAuthError::HttpStatus {
|
||||
status_code,
|
||||
body_excerpt,
|
||||
},
|
||||
LocalOAuthRefreshError::InvalidResponse { message, .. } => {
|
||||
OAuthError::InvalidResponse(message)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait LocalOAuthRefreshAdapter: Send + Sync {
|
||||
fn provider_type(&self) -> &'static str;
|
||||
|
||||
fn supports(&self, transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(self.provider_type())
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth>;
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth>;
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool;
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError>;
|
||||
}
|
||||
|
||||
pub struct LocalOAuthRefreshCoordinator {
|
||||
adapters: Vec<Arc<dyn LocalOAuthRefreshAdapter>>,
|
||||
cache: Mutex<BTreeMap<String, CachedOAuthEntry>>,
|
||||
key_locks: Mutex<BTreeMap<String, Arc<Mutex<()>>>>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LocalOAuthRefreshCoordinator {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("LocalOAuthRefreshCoordinator")
|
||||
.field("adapter_count", &self.adapters.len())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for LocalOAuthRefreshCoordinator {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOAuthRefreshCoordinator {
|
||||
const DISTRIBUTED_REFRESH_LOCK_TTL_MS: u64 = 30_000;
|
||||
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
adapters: vec![
|
||||
Arc::new(KiroOAuthRefreshAdapter::default()),
|
||||
Arc::new(VertexServiceAccountRefreshAdapter),
|
||||
Arc::new(GenericOAuthRefreshAdapter::default()),
|
||||
],
|
||||
cache: Mutex::new(BTreeMap::new()),
|
||||
key_locks: Mutex::new(BTreeMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn lock_for_key(&self, key_id: &str) -> Arc<Mutex<()>> {
|
||||
let mut key_locks = self.key_locks.lock().await;
|
||||
key_locks
|
||||
.entry(key_id.to_string())
|
||||
.or_insert_with(|| Arc::new(Mutex::new(())))
|
||||
.clone()
|
||||
}
|
||||
|
||||
async fn cached_entry(&self, key_id: &str) -> Option<CachedOAuthEntry> {
|
||||
self.cache.lock().await.get(key_id).cloned()
|
||||
}
|
||||
|
||||
async fn insert_cached_entry(&self, key_id: &str, entry: CachedOAuthEntry) {
|
||||
self.cache.lock().await.insert(key_id.to_string(), entry);
|
||||
}
|
||||
|
||||
pub async fn store_cached_entry(&self, key_id: &str, entry: CachedOAuthEntry) {
|
||||
self.insert_cached_entry(key_id, entry).await;
|
||||
}
|
||||
|
||||
pub async fn invalidate_cached_entry(&self, key_id: &str) -> bool {
|
||||
self.cache.lock().await.remove(key_id).is_some()
|
||||
}
|
||||
|
||||
pub async fn resolve_with_result(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
distributed_lock: Option<&RuntimeState>,
|
||||
distributed_owner: Option<&str>,
|
||||
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
|
||||
self.resolve_with_result_mode(
|
||||
executor,
|
||||
transport,
|
||||
distributed_lock,
|
||||
distributed_owner,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn force_refresh_with_result(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
distributed_lock: Option<&RuntimeState>,
|
||||
distributed_owner: Option<&str>,
|
||||
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
|
||||
self.resolve_with_result_mode(
|
||||
executor,
|
||||
transport,
|
||||
distributed_lock,
|
||||
distributed_owner,
|
||||
true,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn resolve_with_result_mode(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
distributed_lock: Option<&RuntimeState>,
|
||||
distributed_owner: Option<&str>,
|
||||
force_refresh: bool,
|
||||
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
|
||||
let Some(adapter) = self
|
||||
.adapters
|
||||
.iter()
|
||||
.find(|adapter| adapter.supports(transport))
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let key_id = transport.key.id.trim();
|
||||
|
||||
let cached_entry = if key_id.is_empty() {
|
||||
None
|
||||
} else {
|
||||
self.cached_entry(key_id).await
|
||||
};
|
||||
if !force_refresh {
|
||||
if let Some(auth) = cached_entry
|
||||
.as_ref()
|
||||
.and_then(|entry| adapter.resolve_cached(transport, entry))
|
||||
{
|
||||
return Ok(Some(LocalOAuthResolution::resolved(auth, None)));
|
||||
}
|
||||
if let Some(auth) = adapter.resolve_without_refresh(transport) {
|
||||
return Ok(Some(LocalOAuthResolution::resolved(auth, None)));
|
||||
}
|
||||
if !adapter.should_refresh(transport, cached_entry.as_ref()) {
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
if key_id.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let key_lock = self.lock_for_key(key_id).await;
|
||||
let _key_guard = key_lock.lock().await;
|
||||
|
||||
let cached_entry = self.cached_entry(key_id).await;
|
||||
if !force_refresh {
|
||||
if let Some(auth) = cached_entry
|
||||
.as_ref()
|
||||
.and_then(|entry| adapter.resolve_cached(transport, entry))
|
||||
{
|
||||
return Ok(Some(LocalOAuthResolution::resolved(auth, None)));
|
||||
}
|
||||
if let Some(auth) = adapter.resolve_without_refresh(transport) {
|
||||
return Ok(Some(LocalOAuthResolution::resolved(auth, None)));
|
||||
}
|
||||
if !adapter.should_refresh(transport, cached_entry.as_ref()) {
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
|
||||
let distributed_lease = match (distributed_lock, distributed_owner) {
|
||||
(Some(lock), Some(owner)) if !owner.trim().is_empty() => {
|
||||
match lock
|
||||
.lock_try_acquire(
|
||||
&format!("provider_oauth_refresh_lock:{key_id}"),
|
||||
owner,
|
||||
std::time::Duration::from_millis(Self::DISTRIBUTED_REFRESH_LOCK_TTL_MS),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(lease)) => Some(lease),
|
||||
Ok(None) => return Ok(Some(LocalOAuthResolution::refresh_in_flight())),
|
||||
Err(err) => {
|
||||
tracing::warn!(
|
||||
key_id = %key_id,
|
||||
provider_type = adapter.provider_type(),
|
||||
error = ?err,
|
||||
"gateway local oauth refresh distributed lock unavailable"
|
||||
);
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
|
||||
// Forced refresh still needs the latest rotated refresh_token as input.
|
||||
// Otherwise a second overlapping refresh can acquire the lock after the
|
||||
// first one completes, then immediately retry with the stale token that
|
||||
// came from the original transport snapshot.
|
||||
let refresh_entry = cached_entry.as_ref();
|
||||
let refresh_result = adapter.refresh(executor, transport, refresh_entry).await;
|
||||
if let (Some(lock), Some(lease)) = (distributed_lock, distributed_lease.as_ref()) {
|
||||
if let Err(err) = lock.lock_release(lease).await {
|
||||
tracing::warn!(
|
||||
key_id = %key_id,
|
||||
provider_type = adapter.provider_type(),
|
||||
error = ?err,
|
||||
"gateway local oauth refresh distributed lock release failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
let Some(refreshed_entry) = refresh_result? else {
|
||||
return Ok(None);
|
||||
};
|
||||
Ok(adapter
|
||||
.resolve_cached(transport, &refreshed_entry)
|
||||
.map(|auth| LocalOAuthResolution::resolved(auth, Some(refreshed_entry))))
|
||||
}
|
||||
|
||||
pub fn with_adapters_for_tests(adapters: Vec<Arc<dyn LocalOAuthRefreshAdapter>>) -> Self {
|
||||
Self {
|
||||
adapters,
|
||||
cache: Mutex::new(BTreeMap::new()),
|
||||
key_locks: Mutex::new(BTreeMap::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOAuthResolution {
|
||||
fn resolved(
|
||||
auth: LocalResolvedOAuthRequestAuth,
|
||||
refreshed_entry: Option<CachedOAuthEntry>,
|
||||
) -> Self {
|
||||
Self {
|
||||
auth: Some(auth),
|
||||
refreshed_entry,
|
||||
refresh_in_flight: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn refresh_in_flight() -> Self {
|
||||
Self {
|
||||
auth: None,
|
||||
refreshed_entry: None,
|
||||
refresh_in_flight: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn supports_local_oauth_request_auth_resolution(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
supports_local_kiro_request_auth_resolution(transport)
|
||||
|| supports_local_vertex_service_account_auth_resolution(transport)
|
||||
|| supports_local_generic_oauth_request_auth_resolution(transport)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::{
|
||||
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter,
|
||||
LocalOAuthRefreshCoordinator, LocalOAuthRefreshError, LocalOAuthResolution,
|
||||
LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use std::sync::Arc;
|
||||
|
||||
#[derive(Debug)]
|
||||
struct TestAdapter {
|
||||
refresh_hits: Arc<AtomicUsize>,
|
||||
refresh_with_entry_hits: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for TestAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
"test-oauth"
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
(entry.provider_type == "test-oauth").then(|| LocalResolvedOAuthRequestAuth::Header {
|
||||
name: entry.auth_header_name.clone(),
|
||||
value: entry.auth_header_value.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
(!secret.is_empty() && secret != "__placeholder__").then(|| {
|
||||
LocalResolvedOAuthRequestAuth::Header {
|
||||
name: "authorization".to_string(),
|
||||
value: format!("Bearer {secret}"),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool {
|
||||
entry.is_none() && transport.key.decrypted_api_key.trim() == "__placeholder__"
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
_executor: &dyn LocalOAuthHttpExecutor,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
|
||||
self.refresh_hits.fetch_add(1, Ordering::SeqCst);
|
||||
if entry.is_some() {
|
||||
self.refresh_with_entry_hits.fetch_add(1, Ordering::SeqCst);
|
||||
}
|
||||
Ok(Some(CachedOAuthEntry {
|
||||
provider_type: "test-oauth".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "Bearer refreshed-token".to_string(),
|
||||
expires_at_unix_secs: Some(4_102_444_800),
|
||||
metadata: None,
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "test".to_string(),
|
||||
provider_type: "test-oauth".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "claude:messages".to_string(),
|
||||
api_family: Some("claude".to_string()),
|
||||
endpoint_kind: Some("cli".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://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: "bearer".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "__placeholder__".to_string(),
|
||||
decrypted_auth_config: Some("{\"refresh_token\":\"rt-1\"}".to_string()),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn coordinator_reuses_runtime_cached_refresh_result() {
|
||||
let refresh_hits = Arc::new(AtomicUsize::new(0));
|
||||
let refresh_with_entry_hits = Arc::new(AtomicUsize::new(0));
|
||||
let coordinator =
|
||||
LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(TestAdapter {
|
||||
refresh_hits: Arc::clone(&refresh_hits),
|
||||
refresh_with_entry_hits: Arc::clone(&refresh_with_entry_hits),
|
||||
})]);
|
||||
let transport = sample_transport();
|
||||
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
|
||||
|
||||
let first = coordinator
|
||||
.resolve_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.expect("first resolve should succeed");
|
||||
coordinator
|
||||
.insert_cached_entry(
|
||||
transport.key.id.as_str(),
|
||||
first
|
||||
.as_ref()
|
||||
.and_then(|result| result.refreshed_entry.clone())
|
||||
.expect("first resolve should provide cached entry"),
|
||||
)
|
||||
.await;
|
||||
let second = coordinator
|
||||
.resolve_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.expect("second resolve should succeed");
|
||||
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(
|
||||
first,
|
||||
Some(LocalOAuthResolution {
|
||||
auth: Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: "authorization".to_string(),
|
||||
value: "Bearer refreshed-token".to_string(),
|
||||
}),
|
||||
refreshed_entry: Some(CachedOAuthEntry {
|
||||
provider_type: "test-oauth".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "Bearer refreshed-token".to_string(),
|
||||
expires_at_unix_secs: Some(4_102_444_800),
|
||||
metadata: None,
|
||||
}),
|
||||
refresh_in_flight: false,
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
second,
|
||||
Some(LocalOAuthResolution {
|
||||
auth: Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: "authorization".to_string(),
|
||||
value: "Bearer refreshed-token".to_string(),
|
||||
}),
|
||||
refreshed_entry: None,
|
||||
refresh_in_flight: false,
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn coordinator_force_refresh_bypasses_runtime_cache() {
|
||||
let refresh_hits = Arc::new(AtomicUsize::new(0));
|
||||
let refresh_with_entry_hits = Arc::new(AtomicUsize::new(0));
|
||||
let coordinator =
|
||||
LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(TestAdapter {
|
||||
refresh_hits: Arc::clone(&refresh_hits),
|
||||
refresh_with_entry_hits: Arc::clone(&refresh_with_entry_hits),
|
||||
})]);
|
||||
let transport = sample_transport();
|
||||
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
|
||||
|
||||
let first = coordinator
|
||||
.resolve_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.expect("initial resolve should succeed");
|
||||
coordinator
|
||||
.insert_cached_entry(
|
||||
transport.key.id.as_str(),
|
||||
first
|
||||
.as_ref()
|
||||
.and_then(|result| result.refreshed_entry.clone())
|
||||
.expect("first resolve should provide cached entry"),
|
||||
)
|
||||
.await;
|
||||
let forced = coordinator
|
||||
.force_refresh_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.expect("forced refresh should succeed");
|
||||
|
||||
assert!(first.and_then(|result| result.refreshed_entry).is_some());
|
||||
assert!(forced.and_then(|result| result.refreshed_entry).is_some());
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 2);
|
||||
assert_eq!(refresh_with_entry_hits.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,419 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::auth::{build_passthrough_headers_with_auth, resolve_local_openai_bearer_auth};
|
||||
use crate::grok::{is_grok_provider_transport, resolve_grok_session_auth};
|
||||
use crate::policy::local_standard_transport_unsupported_reason_with_network;
|
||||
use crate::rules::apply_local_header_rules_with_request_headers;
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::url::build_openai_image_url;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct ProviderOpenAiImageHeadersInput<'a> {
|
||||
pub transport: &'a GatewayProviderTransportSnapshot,
|
||||
pub headers: &'a http::HeaderMap,
|
||||
pub auth_header: &'a str,
|
||||
pub auth_value: &'a str,
|
||||
pub accept: Option<&'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> {
|
||||
let reason = local_standard_transport_unsupported_reason_with_network(transport, api_format);
|
||||
if is_dedicated_openai_image_provider(transport)
|
||||
&& matches!(
|
||||
reason,
|
||||
Some("transport_provider_type_unsupported")
|
||||
| Some("transport_oauth_resolution_unsupported")
|
||||
)
|
||||
{
|
||||
return None;
|
||||
}
|
||||
reason
|
||||
}
|
||||
|
||||
fn is_dedicated_openai_image_provider(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("chatgpt_web")
|
||||
|| transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
|| is_grok_provider_transport(transport)
|
||||
}
|
||||
|
||||
pub fn resolve_openai_image_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<(String, String)> {
|
||||
if is_grok_provider_transport(transport) {
|
||||
return resolve_grok_session_auth(transport);
|
||||
}
|
||||
resolve_local_openai_bearer_auth(transport)
|
||||
}
|
||||
|
||||
pub fn build_openai_image_upstream_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request_path: Option<&str>,
|
||||
request_query: Option<&str>,
|
||||
) -> String {
|
||||
build_openai_image_url(&transport.endpoint.base_url, request_path, request_query)
|
||||
}
|
||||
|
||||
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());
|
||||
if let Some(accept) = input.accept {
|
||||
provider_request_headers.insert("accept".to_string(), accept.to_string());
|
||||
} else {
|
||||
provider_request_headers.remove("accept");
|
||||
}
|
||||
crate::apply_local_auth_config_header_overrides(
|
||||
&mut provider_request_headers,
|
||||
input.transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
if !apply_local_header_rules_with_request_headers(
|
||||
&mut provider_request_headers,
|
||||
input.header_rules,
|
||||
&[input.auth_header, "content-type", "accept"],
|
||||
input.provider_request_body,
|
||||
Some(input.original_request_body),
|
||||
Some(input.headers),
|
||||
) {
|
||||
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,
|
||||
openai_image_transport_unsupported_reason, 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/v1".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,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_openai_image_url_uses_images_surface() {
|
||||
let url = build_openai_image_upstream_url(
|
||||
&sample_transport(),
|
||||
Some("/v1/images/generations"),
|
||||
Some("trace=1"),
|
||||
);
|
||||
|
||||
assert_eq!(url, "https://api.openai.com/v1/images/generations?trace=1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_openai_image_edit_url_uses_images_edit_surface() {
|
||||
let url = build_openai_image_upstream_url(
|
||||
&sample_transport(),
|
||||
Some("/v1/images/edits"),
|
||||
Some("trace=1"),
|
||||
);
|
||||
|
||||
assert_eq!(url, "https://api.openai.com/v1/images/edits?trace=1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standard_openai_image_url_uses_images_surface() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "openai".to_string();
|
||||
|
||||
let url = build_openai_image_upstream_url(
|
||||
&transport,
|
||||
Some("/v1/images/generations"),
|
||||
Some("trace=1"),
|
||||
);
|
||||
|
||||
assert_eq!(url, "https://api.openai.com/v1/images/generations?trace=1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standard_openai_image_url_preserves_edit_surface() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "openai".to_string();
|
||||
|
||||
let url =
|
||||
build_openai_image_upstream_url(&transport, Some("/v1/images/edits"), Some("trace=1"));
|
||||
|
||||
assert_eq!(url, "https://api.openai.com/v1/images/edits?trace=1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chatgpt_web_is_supported_by_dedicated_openai_image_transport_policy() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "chatgpt_web".to_string();
|
||||
|
||||
assert_eq!(
|
||||
openai_image_transport_unsupported_reason(&transport, "openai:image"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_is_supported_by_dedicated_openai_image_transport_policy() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
|
||||
assert_eq!(
|
||||
openai_image_transport_unsupported_reason(&transport, "openai:image"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_is_supported_by_dedicated_openai_image_transport_policy() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "codex".to_string();
|
||||
transport.key.auth_type = "oauth".to_string();
|
||||
transport.key.decrypted_auth_config = Some(json!({"access_token":"token"}).to_string());
|
||||
|
||||
assert_eq!(
|
||||
openai_image_transport_unsupported_reason(&transport, "openai:image"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_oauth_session_is_supported_by_dedicated_openai_image_transport_policy() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "grok".to_string();
|
||||
transport.key.auth_type = "oauth".to_string();
|
||||
transport.key.decrypted_api_key = String::new();
|
||||
transport.key.decrypted_auth_config = Some(json!({"sso_token":"abc"}).to_string());
|
||||
|
||||
assert_eq!(
|
||||
openai_image_transport_unsupported_reason(&transport, "openai:image"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chatgpt_web_oauth_is_supported_by_dedicated_openai_image_transport_policy() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "chatgpt_web".to_string();
|
||||
transport.key.auth_type = "oauth".to_string();
|
||||
transport.key.decrypted_api_key = String::new();
|
||||
transport.key.decrypted_auth_config = Some(json!({"access_token":"token"}).to_string());
|
||||
|
||||
assert_eq!(
|
||||
openai_image_transport_unsupported_reason(&transport, "openai:image"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_json_eventstream_headers_and_applies_rules() {
|
||||
let transport = sample_transport();
|
||||
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
|
||||
transport: &transport,
|
||||
headers: &HeaderMap::new(),
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer secret",
|
||||
accept: Some("text/event-stream"),
|
||||
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()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_config_headers_override_default_authorization() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"refresh_token": "rt-1",
|
||||
"headers": {
|
||||
"authorization": "Bearer imported-session",
|
||||
"chatgpt-account-id": "acct-1"
|
||||
}
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
|
||||
transport: &transport,
|
||||
headers: &HeaderMap::new(),
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer refreshed-access-token",
|
||||
accept: Some("text/event-stream"),
|
||||
header_rules: None,
|
||||
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 imported-session".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("chatgpt-account-id"),
|
||||
Some(&"acct-1".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standard_openai_compatible_image_headers_can_request_json() {
|
||||
let transport = sample_transport();
|
||||
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
|
||||
transport: &transport,
|
||||
headers: &HeaderMap::new(),
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer secret",
|
||||
accept: Some("application/json"),
|
||||
header_rules: None,
|
||||
provider_request_body: &json!({
|
||||
"model": "upstream-image-model",
|
||||
"prompt": "draw a city",
|
||||
}),
|
||||
original_request_body: &json!({"prompt":"draw a city"}),
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(headers.get("accept"), Some(&"application/json".to_string()));
|
||||
assert_eq!(
|
||||
headers.get("content-type"),
|
||||
Some(&"application/json".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_images_omits_explicit_accept_header() {
|
||||
let transport = sample_transport();
|
||||
let mut request_headers = HeaderMap::new();
|
||||
request_headers.insert("accept", "text/event-stream".parse().expect("valid header"));
|
||||
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
|
||||
transport: &transport,
|
||||
headers: &request_headers,
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer secret",
|
||||
accept: None,
|
||||
header_rules: Some(&json!([
|
||||
{"action":"set","key":"accept","value":"application/json"}
|
||||
])),
|
||||
provider_request_body: &json!({
|
||||
"model": "gpt-image-2",
|
||||
"prompt": "draw a city",
|
||||
}),
|
||||
original_request_body: &json!({"prompt":"draw a city"}),
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert!(!headers.contains_key("accept"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standard_openai_compatible_image_url_supports_aether_api_root() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "custom".to_string();
|
||||
transport.endpoint.base_url = "https://upstream-aether.example/v1".to_string();
|
||||
|
||||
let url = build_openai_image_upstream_url(
|
||||
&transport,
|
||||
Some("/v1/images/generations"),
|
||||
Some("trace=1"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://upstream-aether.example/v1/images/generations?trace=1"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,395 @@
|
||||
use super::provider_types::{
|
||||
provider_type_supports_local_embedding_transport,
|
||||
provider_type_supports_local_openai_chat_transport,
|
||||
provider_type_supports_local_same_format_transport,
|
||||
};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::{
|
||||
body_rules_are_locally_supported, header_rules_are_locally_supported,
|
||||
resolve_transport_profile, supports_local_oauth_request_auth_resolution,
|
||||
transport_profile_is_configured, transport_proxy_is_locally_supported,
|
||||
};
|
||||
|
||||
pub fn supports_local_openai_chat_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
local_openai_chat_transport_unsupported_reason(transport).is_none()
|
||||
}
|
||||
|
||||
pub fn supports_local_standard_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
local_standard_transport_unsupported_reason(transport, api_format).is_none()
|
||||
}
|
||||
|
||||
pub fn supports_local_gemini_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
local_gemini_transport_unsupported_reason(transport, api_format).is_none()
|
||||
}
|
||||
|
||||
pub fn supports_local_standard_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
local_standard_transport_unsupported_reason_with_network(transport, api_format).is_none()
|
||||
}
|
||||
|
||||
pub fn supports_local_gemini_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
local_gemini_transport_unsupported_reason_with_network(transport, api_format).is_none()
|
||||
}
|
||||
|
||||
pub fn local_openai_chat_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
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");
|
||||
}
|
||||
if !transport
|
||||
.endpoint
|
||||
.api_format
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("openai:chat")
|
||||
{
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
|
||||
return Some("transport_header_rules_unsupported");
|
||||
}
|
||||
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
if transport.key.decrypted_auth_config.is_some()
|
||||
&& !supports_local_oauth_request_auth_resolution(transport)
|
||||
{
|
||||
return Some("transport_oauth_resolution_unsupported");
|
||||
}
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return Some("transport_proxy_unsupported");
|
||||
}
|
||||
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return Some("transport_profile_unsupported");
|
||||
}
|
||||
if !provider_type_supports_local_openai_chat_transport(&transport.provider.provider_type) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub fn local_standard_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
local_same_format_transport_unsupported_reason(
|
||||
transport,
|
||||
api_format,
|
||||
false,
|
||||
provider_type_supports_local_same_format_transport,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn local_gemini_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
local_same_format_transport_unsupported_reason(
|
||||
transport,
|
||||
api_format,
|
||||
false,
|
||||
provider_type_supports_local_same_format_transport,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn local_standard_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
local_same_format_transport_unsupported_reason(
|
||||
transport,
|
||||
api_format,
|
||||
true,
|
||||
provider_type_supports_local_same_format_transport,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn local_gemini_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
local_same_format_transport_unsupported_reason(
|
||||
transport,
|
||||
api_format,
|
||||
true,
|
||||
provider_type_supports_local_same_format_transport,
|
||||
)
|
||||
}
|
||||
|
||||
fn local_same_format_transport_unsupported_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
allow_network_passthrough: bool,
|
||||
provider_type_supported: fn(&str) -> bool,
|
||||
) -> Option<&'static str> {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return if !transport.provider.is_active {
|
||||
Some("provider_inactive")
|
||||
} else if !transport.endpoint.is_active {
|
||||
Some("endpoint_inactive")
|
||||
} else {
|
||||
Some("key_inactive")
|
||||
};
|
||||
}
|
||||
if !same_api_format(&transport.endpoint.api_format, api_format) {
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
|
||||
return Some("transport_header_rules_unsupported");
|
||||
}
|
||||
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
if transport.key.decrypted_auth_config.is_some()
|
||||
&& !supports_local_oauth_request_auth_resolution(transport)
|
||||
{
|
||||
return Some("transport_oauth_resolution_unsupported");
|
||||
}
|
||||
let has_custom_path = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.is_some_and(|value| !value.trim().is_empty());
|
||||
if has_custom_path && !allow_network_passthrough {
|
||||
return Some("transport_custom_path_unsupported");
|
||||
}
|
||||
if allow_network_passthrough {
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return Some("transport_proxy_unsupported");
|
||||
}
|
||||
if transport_profile_is_configured(transport)
|
||||
&& resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return Some("transport_profile_unsupported");
|
||||
}
|
||||
} else if transport.provider.proxy.is_some()
|
||||
|| transport.endpoint.proxy.is_some()
|
||||
|| transport.key.proxy.is_some()
|
||||
|| transport_profile_is_configured(transport)
|
||||
{
|
||||
return Some("transport_proxy_or_profile_unsupported");
|
||||
}
|
||||
|
||||
if !provider_type_supported(&transport.provider.provider_type) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
if aether_ai_formats::is_embedding_api_format(api_format) {
|
||||
if !endpoint_kind_allows_embedding(transport.endpoint.endpoint_kind.as_deref()) {
|
||||
return Some("transport_endpoint_kind_unsupported");
|
||||
}
|
||||
if !provider_type_supports_local_embedding_transport(
|
||||
&transport.provider.provider_type,
|
||||
api_format,
|
||||
) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
}
|
||||
if aether_ai_formats::is_rerank_api_format(api_format) {
|
||||
if !endpoint_kind_allows_rerank(transport.endpoint.endpoint_kind.as_deref()) {
|
||||
return Some("transport_endpoint_kind_unsupported");
|
||||
}
|
||||
if !provider_type_supports_local_embedding_transport(
|
||||
&transport.provider.provider_type,
|
||||
api_format,
|
||||
) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
fn endpoint_kind_allows_embedding(endpoint_kind: Option<&str>) -> bool {
|
||||
endpoint_kind
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| {
|
||||
matches!(
|
||||
value.to_ascii_lowercase().as_str(),
|
||||
"embedding" | "embeddings" | "multimodal_embedding" | "multimodal_embeddings"
|
||||
)
|
||||
})
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
fn endpoint_kind_allows_rerank(endpoint_kind: Option<&str>) -> bool {
|
||||
endpoint_kind
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| matches!(value.to_ascii_lowercase().as_str(), "rerank" | "reranking"))
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
fn same_api_format(left: &str, right: &str) -> bool {
|
||||
aether_ai_formats::api_format_alias_matches(left, right)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::local_standard_transport_unsupported_reason_with_network;
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport(
|
||||
provider_type: &str,
|
||||
api_format: &str,
|
||||
endpoint_kind: Option<&str>,
|
||||
) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "provider".to_string(),
|
||||
provider_type: provider_type.to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: None,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: api_format.to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: endpoint_kind.map(ToOwned::to_owned),
|
||||
is_active: true,
|
||||
base_url: "https://provider.example".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "api_key".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "sk-test".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unsupported_embedding_provider_format_pairs() {
|
||||
let openai_on_gemini = sample_transport("openai", "gemini:embedding", Some("embedding"));
|
||||
let gemini_on_openai = sample_transport("gemini", "openai:embedding", Some("embedding"));
|
||||
let chat_marked_embedding = sample_transport("openai", "openai:embedding", Some("chat"));
|
||||
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&openai_on_gemini,
|
||||
"gemini:embedding"
|
||||
),
|
||||
Some("transport_provider_type_unsupported")
|
||||
);
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&gemini_on_openai,
|
||||
"openai:embedding"
|
||||
),
|
||||
Some("transport_provider_type_unsupported")
|
||||
);
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&chat_marked_embedding,
|
||||
"openai:embedding"
|
||||
),
|
||||
Some("transport_endpoint_kind_unsupported")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_supported_embedding_provider_format_pairs() {
|
||||
for (provider_type, api_format) in [
|
||||
("openai", "openai:embedding"),
|
||||
("gemini", "gemini:embedding"),
|
||||
("google", "gemini:embedding"),
|
||||
("jina", "jina:embedding"),
|
||||
("doubao", "doubao:embedding"),
|
||||
("volcengine", "doubao:embedding"),
|
||||
("aliyun", "aliyun:multimodal_embedding"),
|
||||
("dashscope", "aliyun:multimodal_embedding"),
|
||||
("custom", "openai:embedding"),
|
||||
("custom", "gemini:embedding"),
|
||||
("custom", "jina:embedding"),
|
||||
("custom", "doubao:embedding"),
|
||||
("custom", "aliyun:multimodal_embedding"),
|
||||
] {
|
||||
let transport = sample_transport(provider_type, api_format, Some("embedding"));
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(&transport, api_format),
|
||||
None,
|
||||
"{provider_type} should support {api_format}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_policy_accepts_embedding_endpoint_kind_aliases_only() {
|
||||
for endpoint_kind in [None, Some(""), Some(" embedding "), Some("EMBEDDINGS")] {
|
||||
let transport = sample_transport("openai", "openai:embedding", endpoint_kind);
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&transport,
|
||||
"openai:embedding"
|
||||
),
|
||||
None,
|
||||
"endpoint kind {endpoint_kind:?} should be accepted"
|
||||
);
|
||||
}
|
||||
|
||||
for endpoint_kind in [Some("chat"), Some("responses"), Some("image")] {
|
||||
let transport = sample_transport("openai", "openai:embedding", endpoint_kind);
|
||||
assert_eq!(
|
||||
local_standard_transport_unsupported_reason_with_network(
|
||||
&transport,
|
||||
"openai:embedding"
|
||||
),
|
||||
Some("transport_endpoint_kind_unsupported"),
|
||||
"endpoint kind {endpoint_kind:?} should fail closed"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,961 @@
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct ProviderOAuthTemplate {
|
||||
pub provider_type: &'static str,
|
||||
pub display_name: &'static str,
|
||||
pub authorize_url: &'static str,
|
||||
pub token_url: &'static str,
|
||||
pub client_id: &'static str,
|
||||
pub client_secret: &'static str,
|
||||
pub scopes: &'static [&'static str],
|
||||
pub redirect_uri: &'static str,
|
||||
pub use_pkce: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum FixedProviderEndpointConfigValue {
|
||||
String(&'static str),
|
||||
Bool(bool),
|
||||
I64(i64),
|
||||
}
|
||||
|
||||
impl FixedProviderEndpointConfigValue {
|
||||
pub fn to_json_value(self) -> serde_json::Value {
|
||||
match self {
|
||||
Self::String(value) => serde_json::Value::String(value.to_string()),
|
||||
Self::Bool(value) => serde_json::Value::Bool(value),
|
||||
Self::I64(value) => serde_json::Value::Number(value.into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct FixedProviderEndpointConfigDefault {
|
||||
pub key: &'static str,
|
||||
pub value: FixedProviderEndpointConfigValue,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct FixedProviderEndpointTemplate {
|
||||
pub item_key: &'static str,
|
||||
pub api_format: &'static str,
|
||||
pub custom_path: Option<&'static str>,
|
||||
pub config_defaults: &'static [FixedProviderEndpointConfigDefault],
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ProviderApiFormatInheritance {
|
||||
None,
|
||||
OAuth,
|
||||
OAuthOrBearer,
|
||||
OAuthOrServiceAccount,
|
||||
OAuthOrConfiguredBearer,
|
||||
}
|
||||
|
||||
impl ProviderApiFormatInheritance {
|
||||
pub fn key_inherits_api_formats(
|
||||
self,
|
||||
auth_type: &str,
|
||||
decrypted_auth_config: Option<&str>,
|
||||
) -> bool {
|
||||
let auth_type = auth_type.trim().to_ascii_lowercase();
|
||||
match self {
|
||||
Self::None => false,
|
||||
Self::OAuth => auth_type == "oauth",
|
||||
Self::OAuthOrBearer => auth_type == "oauth" || auth_type == "bearer",
|
||||
Self::OAuthOrServiceAccount => {
|
||||
auth_type == "oauth" || auth_type == "service_account" || auth_type == "vertex_ai"
|
||||
}
|
||||
Self::OAuthOrConfiguredBearer => {
|
||||
auth_type == "oauth"
|
||||
|| auth_type == "bearer"
|
||||
&& decrypted_auth_config
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ProviderLocalEmbeddingSupport {
|
||||
None,
|
||||
AnyKnown,
|
||||
OpenAi,
|
||||
Gemini,
|
||||
Jina,
|
||||
Doubao,
|
||||
Aliyun,
|
||||
}
|
||||
|
||||
impl ProviderLocalEmbeddingSupport {
|
||||
pub fn supports_api_format(self, api_format: &str) -> bool {
|
||||
let api_format = aether_ai_formats::normalize_api_format_alias(api_format);
|
||||
match self {
|
||||
Self::None => false,
|
||||
Self::AnyKnown => matches!(
|
||||
api_format.as_str(),
|
||||
"openai:embedding"
|
||||
| "openai:rerank"
|
||||
| "gemini:embedding"
|
||||
| "jina:embedding"
|
||||
| "jina:rerank"
|
||||
| "doubao:embedding"
|
||||
| "aliyun:multimodal_embedding"
|
||||
),
|
||||
Self::OpenAi => matches!(api_format.as_str(), "openai:embedding" | "openai:rerank"),
|
||||
Self::Gemini => api_format == "gemini:embedding",
|
||||
Self::Jina => matches!(api_format.as_str(), "jina:embedding" | "jina:rerank"),
|
||||
Self::Doubao => api_format == "doubao:embedding",
|
||||
Self::Aliyun => api_format == "aliyun:multimodal_embedding",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct ProviderRuntimePolicy {
|
||||
pub fixed_provider: bool,
|
||||
pub api_format_inheritance: ProviderApiFormatInheritance,
|
||||
pub enable_format_conversion_by_default: bool,
|
||||
pub allow_auth_channel_mismatch_by_default: bool,
|
||||
pub oauth_is_bearer_like: bool,
|
||||
pub supports_model_fetch: bool,
|
||||
pub supports_local_openai_chat_transport: bool,
|
||||
pub supports_local_same_format_transport: bool,
|
||||
pub local_embedding_support: ProviderLocalEmbeddingSupport,
|
||||
}
|
||||
|
||||
impl ProviderRuntimePolicy {
|
||||
pub const fn standard() -> Self {
|
||||
Self {
|
||||
fixed_provider: false,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::None,
|
||||
enable_format_conversion_by_default: false,
|
||||
allow_auth_channel_mismatch_by_default: false,
|
||||
oauth_is_bearer_like: false,
|
||||
supports_model_fetch: true,
|
||||
supports_local_openai_chat_transport: true,
|
||||
supports_local_same_format_transport: true,
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn key_inherits_api_formats(
|
||||
self,
|
||||
auth_type: &str,
|
||||
decrypted_auth_config: Option<&str>,
|
||||
) -> bool {
|
||||
self.api_format_inheritance
|
||||
.key_inherits_api_formats(auth_type, decrypted_auth_config)
|
||||
}
|
||||
|
||||
pub fn supports_local_embedding_transport(self, api_format: &str) -> bool {
|
||||
self.local_embedding_support.supports_api_format(api_format)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct FixedProviderTemplate {
|
||||
pub provider_type: &'static str,
|
||||
pub version: u32,
|
||||
pub base_url: &'static str,
|
||||
pub endpoints: &'static [FixedProviderEndpointTemplate],
|
||||
pub runtime_policy: ProviderRuntimePolicy,
|
||||
}
|
||||
|
||||
const EMPTY_ENDPOINT_CONFIG_DEFAULTS: &[FixedProviderEndpointConfigDefault] = &[];
|
||||
const FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS: &[FixedProviderEndpointConfigDefault] =
|
||||
&[FixedProviderEndpointConfigDefault {
|
||||
key: "upstream_stream_policy",
|
||||
value: FixedProviderEndpointConfigValue::String("force_stream"),
|
||||
}];
|
||||
const AUTO_STREAM_ENDPOINT_CONFIG_DEFAULTS: &[FixedProviderEndpointConfigDefault] =
|
||||
&[FixedProviderEndpointConfigDefault {
|
||||
key: "upstream_stream_policy",
|
||||
value: FixedProviderEndpointConfigValue::String("auto"),
|
||||
}];
|
||||
|
||||
const STANDARD_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy::standard();
|
||||
const CUSTOM_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::AnyKnown,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const OPENAI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::OpenAi,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const GEMINI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::Gemini,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const JINA_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::Jina,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const DOUBAO_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::Doubao,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const ALIYUN_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::Aliyun,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
|
||||
const CLAUDE_CODE_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
|
||||
enable_format_conversion_by_default: true,
|
||||
oauth_is_bearer_like: true,
|
||||
supports_model_fetch: false,
|
||||
supports_local_openai_chat_transport: false,
|
||||
supports_local_same_format_transport: false,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const CODEX_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
|
||||
enable_format_conversion_by_default: true,
|
||||
supports_model_fetch: false,
|
||||
supports_local_openai_chat_transport: false,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const CHATGPT_WEB_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrBearer,
|
||||
enable_format_conversion_by_default: true,
|
||||
oauth_is_bearer_like: true,
|
||||
supports_model_fetch: false,
|
||||
supports_local_openai_chat_transport: false,
|
||||
supports_local_same_format_transport: false,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const GEMINI_CLI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
|
||||
oauth_is_bearer_like: true,
|
||||
supports_local_openai_chat_transport: false,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const VERTEX_AI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrServiceAccount,
|
||||
enable_format_conversion_by_default: true,
|
||||
supports_model_fetch: false,
|
||||
supports_local_openai_chat_transport: false,
|
||||
supports_local_same_format_transport: false,
|
||||
local_embedding_support: ProviderLocalEmbeddingSupport::Gemini,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const ANTIGRAVITY_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
|
||||
enable_format_conversion_by_default: true,
|
||||
oauth_is_bearer_like: true,
|
||||
supports_model_fetch: false,
|
||||
supports_local_openai_chat_transport: false,
|
||||
supports_local_same_format_transport: false,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
const GROK_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuth,
|
||||
enable_format_conversion_by_default: true,
|
||||
supports_model_fetch: false,
|
||||
supports_local_openai_chat_transport: false,
|
||||
supports_local_same_format_transport: false,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
|
||||
const WINDSURF_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrBearer,
|
||||
enable_format_conversion_by_default: true,
|
||||
oauth_is_bearer_like: true,
|
||||
supports_model_fetch: false,
|
||||
supports_local_openai_chat_transport: false,
|
||||
supports_local_same_format_transport: false,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
|
||||
const CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "claude_code",
|
||||
version: 1,
|
||||
base_url: "https://api.anthropic.com",
|
||||
endpoints: &[FixedProviderEndpointTemplate {
|
||||
item_key: "claude:messages",
|
||||
api_format: "claude:messages",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
}],
|
||||
runtime_policy: CLAUDE_CODE_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const CODEX_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "codex",
|
||||
version: 1,
|
||||
base_url: "https://chatgpt.com/backend-api/codex",
|
||||
endpoints: &[
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:responses",
|
||||
api_format: "openai:responses",
|
||||
custom_path: None,
|
||||
config_defaults: FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:responses:compact",
|
||||
api_format: "openai:responses:compact",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:search",
|
||||
api_format: "openai:search",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:image",
|
||||
api_format: "openai:image",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
],
|
||||
runtime_policy: CODEX_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const CHATGPT_WEB_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "chatgpt_web",
|
||||
version: 1,
|
||||
base_url: "https://chatgpt.com",
|
||||
endpoints: &[FixedProviderEndpointTemplate {
|
||||
item_key: "openai:image",
|
||||
api_format: "openai:image",
|
||||
custom_path: None,
|
||||
config_defaults: FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS,
|
||||
}],
|
||||
runtime_policy: CHATGPT_WEB_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const KIRO_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "kiro",
|
||||
version: 1,
|
||||
base_url: "https://q.{region}.amazonaws.com",
|
||||
endpoints: &[FixedProviderEndpointTemplate {
|
||||
item_key: "claude:messages",
|
||||
api_format: "claude:messages",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
}],
|
||||
runtime_policy: crate::kiro::RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const GEMINI_CLI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "gemini_cli",
|
||||
version: 3,
|
||||
base_url: "https://cloudcode-pa.googleapis.com",
|
||||
endpoints: &[FixedProviderEndpointTemplate {
|
||||
item_key: "gemini:generate_content",
|
||||
api_format: "gemini:generate_content",
|
||||
custom_path: Some("/v1internal:{action}"),
|
||||
config_defaults: AUTO_STREAM_ENDPOINT_CONFIG_DEFAULTS,
|
||||
}],
|
||||
runtime_policy: GEMINI_CLI_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const VERTEX_AI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "vertex_ai",
|
||||
version: 1,
|
||||
base_url: "https://aiplatform.googleapis.com",
|
||||
endpoints: &[
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "gemini:generate_content",
|
||||
api_format: "gemini:generate_content",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "gemini:embedding",
|
||||
api_format: "gemini:embedding",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "claude:messages",
|
||||
api_format: "claude:messages",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
],
|
||||
runtime_policy: VERTEX_AI_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "antigravity",
|
||||
version: 2,
|
||||
base_url: "https://daily-cloudcode-pa.googleapis.com",
|
||||
endpoints: &[FixedProviderEndpointTemplate {
|
||||
item_key: "gemini:generate_content",
|
||||
api_format: "gemini:generate_content",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
}],
|
||||
runtime_policy: ANTIGRAVITY_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const GROK_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "grok",
|
||||
version: 1,
|
||||
base_url: "https://grok.com",
|
||||
endpoints: &[
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:chat",
|
||||
api_format: "openai:chat",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:responses",
|
||||
api_format: "openai:responses",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "claude:messages",
|
||||
api_format: "claude:messages",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:image",
|
||||
api_format: "openai:image",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
],
|
||||
runtime_policy: GROK_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const WINDSURF_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "windsurf",
|
||||
version: 1,
|
||||
base_url: "https://server.codeium.com",
|
||||
endpoints: &[FixedProviderEndpointTemplate {
|
||||
item_key: "openai:chat",
|
||||
api_format: "openai:chat",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
}],
|
||||
runtime_policy: WINDSURF_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
pub fn provider_type_is_fixed(provider_type: &str) -> bool {
|
||||
provider_runtime_policy(provider_type).fixed_provider
|
||||
}
|
||||
|
||||
pub fn fixed_provider_key_inherits_api_formats(
|
||||
provider_type: &str,
|
||||
auth_type: &str,
|
||||
decrypted_auth_config: Option<&str>,
|
||||
) -> bool {
|
||||
provider_runtime_policy(provider_type)
|
||||
.key_inherits_api_formats(auth_type, decrypted_auth_config)
|
||||
}
|
||||
|
||||
pub fn provider_type_enables_format_conversion_by_default(provider_type: &str) -> bool {
|
||||
provider_runtime_policy(provider_type).enable_format_conversion_by_default
|
||||
}
|
||||
|
||||
pub fn provider_type_allows_auth_channel_mismatch_by_default(provider_type: &str) -> bool {
|
||||
provider_runtime_policy(provider_type).allow_auth_channel_mismatch_by_default
|
||||
}
|
||||
|
||||
pub fn provider_type_oauth_is_bearer_like(provider_type: &str) -> bool {
|
||||
provider_runtime_policy(provider_type).oauth_is_bearer_like
|
||||
}
|
||||
|
||||
pub fn provider_runtime_policy(provider_type: &str) -> ProviderRuntimePolicy {
|
||||
if let Some(template) = fixed_provider_template(provider_type) {
|
||||
return template.runtime_policy;
|
||||
}
|
||||
|
||||
match provider_type.trim().to_ascii_lowercase().as_str() {
|
||||
"custom" => CUSTOM_RUNTIME_POLICY,
|
||||
"openai" => OPENAI_RUNTIME_POLICY,
|
||||
"gemini" | "google" => GEMINI_RUNTIME_POLICY,
|
||||
"jina" => JINA_RUNTIME_POLICY,
|
||||
"doubao" | "volcengine" => DOUBAO_RUNTIME_POLICY,
|
||||
"aliyun" | "dashscope" => ALIYUN_RUNTIME_POLICY,
|
||||
_ => STANDARD_RUNTIME_POLICY,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn fixed_provider_template(provider_type: &str) -> Option<&'static FixedProviderTemplate> {
|
||||
match provider_type.trim().to_ascii_lowercase().as_str() {
|
||||
"claude_code" => Some(&CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE),
|
||||
"codex" => Some(&CODEX_FIXED_PROVIDER_TEMPLATE),
|
||||
"chatgpt_web" => Some(&CHATGPT_WEB_FIXED_PROVIDER_TEMPLATE),
|
||||
"kiro" => Some(&KIRO_FIXED_PROVIDER_TEMPLATE),
|
||||
"grok" => Some(&GROK_FIXED_PROVIDER_TEMPLATE),
|
||||
"gemini_cli" => Some(&GEMINI_CLI_FIXED_PROVIDER_TEMPLATE),
|
||||
"vertex_ai" => Some(&VERTEX_AI_FIXED_PROVIDER_TEMPLATE),
|
||||
"antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE),
|
||||
"windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn fixed_provider_endpoint_template_by_api_format(
|
||||
provider_type: &str,
|
||||
api_format: &str,
|
||||
) -> Option<&'static FixedProviderEndpointTemplate> {
|
||||
let normalized = aether_ai_formats::normalize_api_format_alias(api_format);
|
||||
fixed_provider_template(provider_type)?
|
||||
.endpoints
|
||||
.iter()
|
||||
.find(|item| item.api_format.eq_ignore_ascii_case(&normalized))
|
||||
}
|
||||
|
||||
pub fn provider_type_supports_model_fetch(provider_type: &str) -> bool {
|
||||
provider_runtime_policy(provider_type).supports_model_fetch
|
||||
}
|
||||
|
||||
pub fn provider_type_supports_local_openai_chat_transport(provider_type: &str) -> bool {
|
||||
provider_runtime_policy(provider_type).supports_local_openai_chat_transport
|
||||
}
|
||||
|
||||
pub fn provider_type_supports_local_same_format_transport(provider_type: &str) -> bool {
|
||||
provider_runtime_policy(provider_type).supports_local_same_format_transport
|
||||
}
|
||||
|
||||
pub fn provider_type_supports_local_embedding_transport(
|
||||
provider_type: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
provider_runtime_policy(provider_type).supports_local_embedding_transport(api_format)
|
||||
}
|
||||
|
||||
pub fn is_codex_cli_backend_url(url: &str) -> bool {
|
||||
let url = url.trim().to_ascii_lowercase();
|
||||
url.contains("/codex") && (url.contains("/backend-api/") || url.contains("/backendapi/"))
|
||||
}
|
||||
|
||||
pub fn provider_type_is_fixed_for_admin_oauth(provider_type: &str) -> bool {
|
||||
provider_type_is_fixed(provider_type)
|
||||
}
|
||||
|
||||
pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<ProviderOAuthTemplate> {
|
||||
match provider_type.trim().to_ascii_lowercase().as_str() {
|
||||
"claude_code" => Some(ProviderOAuthTemplate {
|
||||
provider_type: "claude_code",
|
||||
display_name: "ClaudeCode",
|
||||
authorize_url: "https://claude.ai/oauth/authorize",
|
||||
token_url: "https://console.anthropic.com/v1/oauth/token",
|
||||
client_id: "9d1c250a-e61b-44d9-88ed-5944d1962f5e",
|
||||
client_secret: "",
|
||||
scopes: &["org:create_api_key", "user:profile", "user:inference"],
|
||||
redirect_uri: "http://localhost:54545/callback",
|
||||
use_pkce: true,
|
||||
}),
|
||||
"codex" => Some(ProviderOAuthTemplate {
|
||||
provider_type: "codex",
|
||||
display_name: "Codex",
|
||||
authorize_url: "https://auth.openai.com/oauth/authorize",
|
||||
token_url: "https://auth.openai.com/oauth/token",
|
||||
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
|
||||
client_secret: "",
|
||||
scopes: &["openid", "email", "profile", "offline_access"],
|
||||
redirect_uri: "http://localhost:1455/auth/callback",
|
||||
use_pkce: true,
|
||||
}),
|
||||
"chatgpt_web" => Some(ProviderOAuthTemplate {
|
||||
provider_type: "chatgpt_web",
|
||||
display_name: "ChatGPT Web",
|
||||
authorize_url: "https://auth.openai.com/oauth/authorize",
|
||||
token_url: "https://auth.openai.com/oauth/token",
|
||||
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
|
||||
client_secret: "",
|
||||
scopes: &["openid", "email", "profile", "offline_access"],
|
||||
redirect_uri: "http://localhost:1455/auth/callback",
|
||||
use_pkce: true,
|
||||
}),
|
||||
"gemini_cli" => Some(ProviderOAuthTemplate {
|
||||
provider_type: "gemini_cli",
|
||||
display_name: "GeminiCli",
|
||||
authorize_url: "https://accounts.google.com/o/oauth2/v2/auth",
|
||||
token_url: "https://oauth2.googleapis.com/token",
|
||||
client_id: "681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com",
|
||||
client_secret: "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",
|
||||
scopes: &[
|
||||
"https://www.googleapis.com/auth/cloud-platform",
|
||||
"https://www.googleapis.com/auth/userinfo.email",
|
||||
"https://www.googleapis.com/auth/userinfo.profile",
|
||||
],
|
||||
redirect_uri: "http://localhost:8085/oauth2callback",
|
||||
use_pkce: false,
|
||||
}),
|
||||
"antigravity" => Some(ProviderOAuthTemplate {
|
||||
provider_type: "antigravity",
|
||||
display_name: "Antigravity",
|
||||
authorize_url: "https://accounts.google.com/o/oauth2/v2/auth",
|
||||
token_url: "https://oauth2.googleapis.com/token",
|
||||
client_id: "1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com",
|
||||
client_secret: "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",
|
||||
scopes: &[
|
||||
"https://www.googleapis.com/auth/cloud-platform",
|
||||
"https://www.googleapis.com/auth/userinfo.email",
|
||||
"https://www.googleapis.com/auth/userinfo.profile",
|
||||
"https://www.googleapis.com/auth/cclog",
|
||||
"https://www.googleapis.com/auth/experimentsandconfigs",
|
||||
],
|
||||
redirect_uri: "http://localhost:51121/oauth2callback",
|
||||
use_pkce: true,
|
||||
}),
|
||||
"windsurf" => Some(ProviderOAuthTemplate {
|
||||
provider_type: "windsurf",
|
||||
display_name: "Windsurf",
|
||||
authorize_url: "https://windsurf.com/windsurf/signin",
|
||||
token_url: "https://register.windsurf.com/exa.seat_management_pb.SeatManagementService/RegisterUser",
|
||||
client_id: "3GUryQ7ldAeKEuD2obYnppsnmj58eP5u",
|
||||
client_secret: "",
|
||||
scopes: &[],
|
||||
redirect_uri: "show-auth-token",
|
||||
use_pkce: false,
|
||||
}),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub const ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES: &[&str] = &[
|
||||
"claude_code",
|
||||
"codex",
|
||||
"chatgpt_web",
|
||||
"gemini_cli",
|
||||
"antigravity",
|
||||
"windsurf",
|
||||
];
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
fixed_provider_endpoint_template_by_api_format, fixed_provider_key_inherits_api_formats,
|
||||
fixed_provider_template, provider_runtime_policy, provider_type_admin_oauth_template,
|
||||
provider_type_allows_auth_channel_mismatch_by_default, provider_type_oauth_is_bearer_like,
|
||||
provider_type_supports_local_embedding_transport,
|
||||
provider_type_supports_local_same_format_transport, FixedProviderEndpointConfigValue,
|
||||
ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn codex_fixed_provider_template_includes_codex_companion_endpoints() {
|
||||
let template = fixed_provider_template("codex").expect("codex template should exist");
|
||||
assert_eq!(template.base_url, "https://chatgpt.com/backend-api/codex");
|
||||
assert_eq!(template.version, 1);
|
||||
assert_eq!(
|
||||
template
|
||||
.endpoints
|
||||
.iter()
|
||||
.map(|item| item.api_format)
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"openai:responses",
|
||||
"openai:responses:compact",
|
||||
"openai:search",
|
||||
"openai:image"
|
||||
]
|
||||
);
|
||||
|
||||
let image_template =
|
||||
fixed_provider_endpoint_template_by_api_format("codex", "openai:image")
|
||||
.expect("codex image endpoint should exist");
|
||||
assert!(image_template.config_defaults.is_empty());
|
||||
|
||||
let search_template =
|
||||
fixed_provider_endpoint_template_by_api_format("codex", "openai:search")
|
||||
.expect("codex search endpoint should exist");
|
||||
assert!(search_template.config_defaults.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chatgpt_web_fixed_provider_template_only_exposes_openai_image() {
|
||||
let template =
|
||||
fixed_provider_template("chatgpt_web").expect("chatgpt_web template should exist");
|
||||
assert_eq!(template.base_url, "https://chatgpt.com");
|
||||
assert_eq!(template.version, 1);
|
||||
assert_eq!(
|
||||
template
|
||||
.endpoints
|
||||
.iter()
|
||||
.map(|item| item.api_format)
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["openai:image"]
|
||||
);
|
||||
|
||||
let image_template =
|
||||
fixed_provider_endpoint_template_by_api_format("chatgpt_web", "openai:image")
|
||||
.expect("chatgpt_web image endpoint should exist");
|
||||
assert_eq!(
|
||||
image_template
|
||||
.config_defaults
|
||||
.iter()
|
||||
.map(|item| (item.key, item.value))
|
||||
.collect::<Vec<_>>(),
|
||||
vec![(
|
||||
"upstream_stream_policy",
|
||||
FixedProviderEndpointConfigValue::String("force_stream")
|
||||
)]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_fixed_provider_template_exposes_chat_responses_messages_and_image() {
|
||||
let template = fixed_provider_template("grok").expect("grok template should exist");
|
||||
assert_eq!(template.base_url, "https://grok.com");
|
||||
assert_eq!(template.version, 1);
|
||||
assert_eq!(
|
||||
template
|
||||
.endpoints
|
||||
.iter()
|
||||
.map(|item| item.api_format)
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"openai:chat",
|
||||
"openai:responses",
|
||||
"claude:messages",
|
||||
"openai:image"
|
||||
]
|
||||
);
|
||||
assert!(!template.runtime_policy.supports_model_fetch);
|
||||
assert!(!template.runtime_policy.supports_local_openai_chat_transport);
|
||||
assert!(!template.runtime_policy.supports_local_same_format_transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gemini_cli_fixed_provider_template_uses_v1internal_endpoint_path() {
|
||||
let template =
|
||||
fixed_provider_template("gemini_cli").expect("gemini_cli template should exist");
|
||||
assert_eq!(template.base_url, "https://cloudcode-pa.googleapis.com");
|
||||
assert_eq!(template.version, 3);
|
||||
|
||||
let endpoint =
|
||||
fixed_provider_endpoint_template_by_api_format("gemini_cli", "gemini:generate_content")
|
||||
.expect("gemini_cli generateContent endpoint should exist");
|
||||
assert_eq!(endpoint.custom_path, Some("/v1internal:{action}"));
|
||||
assert_eq!(
|
||||
endpoint
|
||||
.config_defaults
|
||||
.iter()
|
||||
.map(|item| (item.key, item.value))
|
||||
.collect::<Vec<_>>(),
|
||||
vec![(
|
||||
"upstream_stream_policy",
|
||||
FixedProviderEndpointConfigValue::String("auto")
|
||||
)]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn antigravity_fixed_provider_template_uses_daily_cloudcode_endpoint() {
|
||||
let template =
|
||||
fixed_provider_template("antigravity").expect("antigravity template should exist");
|
||||
assert_eq!(
|
||||
template.base_url,
|
||||
"https://daily-cloudcode-pa.googleapis.com"
|
||||
);
|
||||
assert_eq!(template.version, 2);
|
||||
|
||||
let endpoint = fixed_provider_endpoint_template_by_api_format(
|
||||
"antigravity",
|
||||
"gemini:generate_content",
|
||||
)
|
||||
.expect("antigravity generateContent endpoint should exist");
|
||||
assert_eq!(endpoint.custom_path, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_fixed_provider_template_exposes_openai_chat() {
|
||||
let template = fixed_provider_template("windsurf").expect("windsurf template should exist");
|
||||
assert_eq!(template.provider_type, "windsurf");
|
||||
assert_eq!(template.base_url, "https://server.codeium.com");
|
||||
assert_eq!(template.version, 1);
|
||||
assert_eq!(
|
||||
template
|
||||
.endpoints
|
||||
.iter()
|
||||
.map(|item| item.api_format)
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["openai:chat"]
|
||||
);
|
||||
assert!(
|
||||
fixed_provider_endpoint_template_by_api_format("windsurf", "openai:chat").is_some()
|
||||
);
|
||||
|
||||
let policy = provider_runtime_policy("windsurf");
|
||||
assert!(policy.fixed_provider);
|
||||
assert!(policy.enable_format_conversion_by_default);
|
||||
assert!(policy.oauth_is_bearer_like);
|
||||
assert!(!policy.supports_model_fetch);
|
||||
assert!(!policy.supports_local_same_format_transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_admin_oauth_template_is_advertised() {
|
||||
let template =
|
||||
provider_type_admin_oauth_template("windsurf").expect("windsurf oauth template");
|
||||
|
||||
assert_eq!(template.provider_type, "windsurf");
|
||||
assert_eq!(template.display_name, "Windsurf");
|
||||
assert_eq!(
|
||||
template.authorize_url,
|
||||
"https://windsurf.com/windsurf/signin"
|
||||
);
|
||||
assert_eq!(template.redirect_uri, "show-auth-token");
|
||||
assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf"));
|
||||
}
|
||||
|
||||
#[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(
|
||||
"chatgpt_web",
|
||||
"oauth",
|
||||
None
|
||||
));
|
||||
assert!(fixed_provider_key_inherits_api_formats(
|
||||
"kiro",
|
||||
"bearer",
|
||||
Some("{}")
|
||||
));
|
||||
assert!(fixed_provider_key_inherits_api_formats(
|
||||
"vertex_ai",
|
||||
"service_account",
|
||||
None
|
||||
));
|
||||
assert!(!fixed_provider_key_inherits_api_formats(
|
||||
"kiro", "bearer", None
|
||||
));
|
||||
assert!(!fixed_provider_key_inherits_api_formats(
|
||||
"custom", "oauth", None
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_allows_auth_channel_mismatch_by_default() {
|
||||
let policy = provider_runtime_policy("kiro");
|
||||
assert!(policy.fixed_provider);
|
||||
assert!(policy.enable_format_conversion_by_default);
|
||||
assert!(policy.oauth_is_bearer_like);
|
||||
assert!(policy.supports_model_fetch);
|
||||
assert!(!policy.supports_local_openai_chat_transport);
|
||||
assert!(!policy.supports_local_same_format_transport);
|
||||
assert!(policy.key_inherits_api_formats("oauth", None));
|
||||
assert!(policy.key_inherits_api_formats("bearer", Some("{}")));
|
||||
assert!(!policy.key_inherits_api_formats("bearer", None));
|
||||
|
||||
assert!(provider_type_allows_auth_channel_mismatch_by_default(
|
||||
"kiro"
|
||||
));
|
||||
assert!(provider_type_allows_auth_channel_mismatch_by_default(
|
||||
" KIRO "
|
||||
));
|
||||
assert!(!provider_type_allows_auth_channel_mismatch_by_default(
|
||||
"claude_code"
|
||||
));
|
||||
assert!(!provider_type_allows_auth_channel_mismatch_by_default(
|
||||
"custom"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_policy_preserves_other_fixed_provider_behavior() {
|
||||
let codex = provider_runtime_policy("codex");
|
||||
assert!(codex.fixed_provider);
|
||||
assert!(codex.enable_format_conversion_by_default);
|
||||
assert!(!codex.oauth_is_bearer_like);
|
||||
assert!(!codex.supports_model_fetch);
|
||||
assert!(!codex.supports_local_openai_chat_transport);
|
||||
assert!(codex.supports_local_same_format_transport);
|
||||
|
||||
let gemini_cli = provider_runtime_policy("gemini_cli");
|
||||
assert!(gemini_cli.fixed_provider);
|
||||
assert!(!gemini_cli.enable_format_conversion_by_default);
|
||||
assert!(provider_type_oauth_is_bearer_like("gemini_cli"));
|
||||
assert!(gemini_cli.supports_model_fetch);
|
||||
assert!(!gemini_cli.supports_local_openai_chat_transport);
|
||||
assert!(gemini_cli.supports_local_same_format_transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chatgpt_web_does_not_use_generic_same_format_transport() {
|
||||
assert!(!provider_type_supports_local_same_format_transport(
|
||||
"chatgpt_web"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_type_supports_only_matching_embedding_formats() {
|
||||
for (provider_type, api_format) in [
|
||||
("openai", "openai:embedding"),
|
||||
("custom", "openai:embedding"),
|
||||
("gemini", "gemini:embedding"),
|
||||
("google", "gemini:embedding"),
|
||||
("vertex_ai", "gemini:embedding"),
|
||||
("jina", "jina:embedding"),
|
||||
("doubao", "doubao:embedding"),
|
||||
("volcengine", "doubao:embedding"),
|
||||
("aliyun", "aliyun:multimodal_embedding"),
|
||||
("dashscope", "aliyun:multimodal_embedding"),
|
||||
] {
|
||||
assert!(
|
||||
provider_type_supports_local_embedding_transport(provider_type, api_format),
|
||||
"{provider_type} should support {api_format}"
|
||||
);
|
||||
}
|
||||
|
||||
for (provider_type, api_format) in [
|
||||
("openai", "gemini:embedding"),
|
||||
("gemini", "openai:embedding"),
|
||||
("vertex_ai", "openai:embedding"),
|
||||
("jina", "doubao:embedding"),
|
||||
("doubao", "jina:embedding"),
|
||||
("aliyun", "openai:embedding"),
|
||||
("openai", "aliyun:multimodal_embedding"),
|
||||
("claude_code", "openai:embedding"),
|
||||
("openai", "openai:chat"),
|
||||
] {
|
||||
assert!(
|
||||
!provider_type_supports_local_embedding_transport(provider_type, api_format),
|
||||
"{provider_type} should not support {api_format}"
|
||||
);
|
||||
}
|
||||
|
||||
assert!(provider_type_supports_local_embedding_transport(
|
||||
" Google ",
|
||||
"GEMINI:EMBEDDING"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_fixed_provider_template_includes_gemini_embedding_endpoint() {
|
||||
let template =
|
||||
fixed_provider_template("vertex_ai").expect("vertex_ai template should exist");
|
||||
|
||||
assert_eq!(
|
||||
template
|
||||
.endpoints
|
||||
.iter()
|
||||
.map(|item| item.api_format)
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"gemini:generate_content",
|
||||
"gemini:embedding",
|
||||
"claude:messages",
|
||||
]
|
||||
);
|
||||
|
||||
assert!(
|
||||
fixed_provider_endpoint_template_by_api_format("vertex_ai", "gemini:embedding")
|
||||
.is_some()
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,416 @@
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::vertex::is_vertex_transport_context;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct TransportRequestBodySemanticsError {
|
||||
message: &'static str,
|
||||
}
|
||||
|
||||
impl TransportRequestBodySemanticsError {
|
||||
const fn new(message: &'static str) -> Self {
|
||||
Self { message }
|
||||
}
|
||||
|
||||
pub const fn message(&self) -> &'static str {
|
||||
self.message
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for TransportRequestBodySemanticsError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.write_str(self.message)
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for TransportRequestBodySemanticsError {}
|
||||
|
||||
pub fn apply_transport_request_body_semantics(
|
||||
provider_request_body: &mut Value,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> Result<(), TransportRequestBodySemanticsError> {
|
||||
let provider_api_format = aether_ai_formats::normalize_api_format_alias(provider_api_format);
|
||||
if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) {
|
||||
apply_vertex_gemini_embedding_body_semantics(provider_request_body)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn apply_vertex_gemini_embedding_body_semantics(
|
||||
provider_request_body: &mut Value,
|
||||
) -> Result<(), TransportRequestBodySemanticsError> {
|
||||
let object = provider_request_body.as_object_mut().ok_or_else(|| {
|
||||
TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding request body must be a JSON object",
|
||||
)
|
||||
})?;
|
||||
|
||||
if object.contains_key("instances") {
|
||||
validate_existing_vertex_predict_body(object)?;
|
||||
object.remove("model");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let next = build_vertex_predict_body_from_gemini_embedding_object(object)?;
|
||||
*object = next;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn build_vertex_predict_body_from_gemini_embedding_object(
|
||||
object: &Map<String, Value>,
|
||||
) -> Result<Map<String, Value>, TransportRequestBodySemanticsError> {
|
||||
if let Some(requests) = object.get("requests") {
|
||||
if object.keys().any(|key| key != "requests") {
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding batch body cannot mix requests with other top-level fields",
|
||||
));
|
||||
}
|
||||
let request_items = requests.as_array().ok_or_else(|| {
|
||||
TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding requests must be an array",
|
||||
)
|
||||
})?;
|
||||
let request_objects = request_items
|
||||
.iter()
|
||||
.map(Value::as_object)
|
||||
.collect::<Option<Vec<_>>>()
|
||||
.ok_or_else(|| {
|
||||
TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding requests must be an array of objects",
|
||||
)
|
||||
})?;
|
||||
if request_objects.is_empty() {
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding requests must contain at least one item",
|
||||
));
|
||||
}
|
||||
return build_vertex_predict_body_from_gemini_embedding_items(&request_objects);
|
||||
}
|
||||
|
||||
build_vertex_predict_body_from_gemini_embedding_items(&[object])
|
||||
}
|
||||
|
||||
fn build_vertex_predict_body_from_gemini_embedding_items(
|
||||
items: &[&Map<String, Value>],
|
||||
) -> Result<Map<String, Value>, TransportRequestBodySemanticsError> {
|
||||
if items.iter().any(|item| {
|
||||
item.keys().any(|key| {
|
||||
!matches!(
|
||||
key.as_str(),
|
||||
"model"
|
||||
| "content"
|
||||
| "taskType"
|
||||
| "title"
|
||||
| "outputDimensionality"
|
||||
| "autoTruncate"
|
||||
)
|
||||
})
|
||||
}) {
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding body contains fields that cannot be mapped to predict instances",
|
||||
));
|
||||
}
|
||||
|
||||
let instances = items
|
||||
.iter()
|
||||
.map(|item| build_vertex_predict_instance(item))
|
||||
.collect::<Option<Vec<_>>>()
|
||||
.ok_or_else(|| {
|
||||
TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding body must contain text content parts",
|
||||
)
|
||||
})?;
|
||||
if instances.is_empty() {
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding body must contain at least one instance",
|
||||
));
|
||||
}
|
||||
|
||||
let mut output = Map::new();
|
||||
output.insert("instances".to_string(), Value::Array(instances));
|
||||
|
||||
let mut parameters = Map::new();
|
||||
insert_shared_parameter(items, &mut parameters, "outputDimensionality")?;
|
||||
insert_shared_parameter(items, &mut parameters, "autoTruncate")?;
|
||||
if !parameters.is_empty() {
|
||||
output.insert("parameters".to_string(), Value::Object(parameters));
|
||||
}
|
||||
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
fn validate_existing_vertex_predict_body(
|
||||
object: &Map<String, Value>,
|
||||
) -> Result<(), TransportRequestBodySemanticsError> {
|
||||
if object
|
||||
.keys()
|
||||
.any(|key| !matches!(key.as_str(), "model" | "instances" | "parameters"))
|
||||
{
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding predict body contains unsupported top-level fields",
|
||||
));
|
||||
}
|
||||
let Some(instances) = object.get("instances").and_then(Value::as_array) else {
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding predict body must contain an instances array",
|
||||
));
|
||||
};
|
||||
if instances.is_empty() {
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding predict body must contain at least one instance",
|
||||
));
|
||||
}
|
||||
if object
|
||||
.get("parameters")
|
||||
.is_some_and(|parameters| !parameters.is_object())
|
||||
{
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding predict parameters must be an object",
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn build_vertex_predict_instance(item: &Map<String, Value>) -> Option<Value> {
|
||||
let content = gemini_embedding_content_text(item.get("content")?)?;
|
||||
let mut instance = Map::new();
|
||||
instance.insert("content".to_string(), Value::String(content));
|
||||
if let Some(task_type) = item.get("taskType") {
|
||||
instance.insert(
|
||||
"task_type".to_string(),
|
||||
Value::String(task_type.as_str()?.to_string()),
|
||||
);
|
||||
}
|
||||
if let Some(title) = item.get("title") {
|
||||
instance.insert(
|
||||
"title".to_string(),
|
||||
Value::String(title.as_str()?.to_string()),
|
||||
);
|
||||
}
|
||||
Some(Value::Object(instance))
|
||||
}
|
||||
|
||||
fn gemini_embedding_content_text(content: &Value) -> Option<String> {
|
||||
let parts = content
|
||||
.as_object()?
|
||||
.get("parts")?
|
||||
.as_array()?
|
||||
.iter()
|
||||
.filter_map(|part| part.as_object()?.get("text")?.as_str())
|
||||
.filter(|text| !text.trim().is_empty())
|
||||
.collect::<Vec<_>>();
|
||||
if parts.is_empty() {
|
||||
return None;
|
||||
}
|
||||
Some(parts.join(""))
|
||||
}
|
||||
|
||||
fn insert_shared_parameter(
|
||||
items: &[&Map<String, Value>],
|
||||
parameters: &mut Map<String, Value>,
|
||||
key: &str,
|
||||
) -> Result<(), TransportRequestBodySemanticsError> {
|
||||
let mut value: Option<Value> = None;
|
||||
for item in items {
|
||||
let Some(next) = item.get(key) else {
|
||||
continue;
|
||||
};
|
||||
match &value {
|
||||
Some(current) if current != next => {
|
||||
return Err(TransportRequestBodySemanticsError::new(
|
||||
"Vertex Gemini embedding batch items must use the same shared parameters",
|
||||
));
|
||||
}
|
||||
None => value = Some(next.clone()),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
if let Some(value) = value {
|
||||
parameters.insert(key.to_string(), value);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::apply_transport_request_body_semantics;
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport(provider_type: &str, base_url: &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: "gemini:embedding".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("embedding".to_string()),
|
||||
is_active: true,
|
||||
base_url: base_url.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: Some(vec!["gemini:embedding".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_gemini_embedding_single_body_uses_predict_contract() {
|
||||
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
|
||||
let mut body = json!({
|
||||
"model": "gemini-embedding-2",
|
||||
"content": {"parts": [{"text": "hello"}]},
|
||||
"taskType": "RETRIEVAL_QUERY",
|
||||
"outputDimensionality": 768
|
||||
});
|
||||
|
||||
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
|
||||
.expect("body semantics should apply");
|
||||
|
||||
assert!(body.get("model").is_none());
|
||||
assert!(body.get("content").is_none());
|
||||
assert_eq!(body["instances"][0]["content"], "hello");
|
||||
assert_eq!(body["instances"][0]["task_type"], "RETRIEVAL_QUERY");
|
||||
assert_eq!(body["parameters"]["outputDimensionality"], 768);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gemini_api_embedding_single_body_keeps_model_for_developer_api() {
|
||||
let transport =
|
||||
sample_transport("gemini", "https://generativelanguage.googleapis.com/v1beta");
|
||||
let mut body = json!({
|
||||
"model": "gemini-embedding-2",
|
||||
"content": {"parts": [{"text": "hello"}]}
|
||||
});
|
||||
|
||||
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
|
||||
.expect("developer API body should pass through");
|
||||
|
||||
assert_eq!(body["model"], "gemini-embedding-2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_gemini_embedding_batch_body_uses_predict_instances() {
|
||||
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
|
||||
let mut body = json!({
|
||||
"requests": [
|
||||
{
|
||||
"model": "models/gemini-embedding-2",
|
||||
"content": {"parts": [{"text": "hello"}]}
|
||||
},
|
||||
{
|
||||
"model": "models/gemini-embedding-2",
|
||||
"content": {"parts": [{"text": "world"}]}
|
||||
}
|
||||
]
|
||||
});
|
||||
|
||||
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
|
||||
.expect("batch body semantics should apply");
|
||||
|
||||
assert!(body.get("requests").is_none());
|
||||
assert_eq!(body["instances"][0]["content"], "hello");
|
||||
assert_eq!(body["instances"][1]["content"], "world");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_gemini_embedding_existing_predict_body_removes_duplicate_model() {
|
||||
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
|
||||
let mut body = json!({
|
||||
"model": "gemini-embedding-2",
|
||||
"instances": [
|
||||
{"content": "hello", "task_type": "RETRIEVAL_QUERY"}
|
||||
],
|
||||
"parameters": {
|
||||
"outputDimensionality": 768
|
||||
}
|
||||
});
|
||||
|
||||
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
|
||||
.expect("existing predict body should be accepted");
|
||||
|
||||
assert!(body.get("model").is_none());
|
||||
assert_eq!(body["instances"][0]["content"], "hello");
|
||||
assert_eq!(body["parameters"]["outputDimensionality"], 768);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_gemini_embedding_existing_predict_body_rejects_unconsumed_fields() {
|
||||
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
|
||||
let mut body = json!({
|
||||
"model": "gemini-embedding-2",
|
||||
"instances": [
|
||||
{"content": "hello"}
|
||||
],
|
||||
"input": "this field would not be consumed by Vertex predict"
|
||||
});
|
||||
|
||||
let error =
|
||||
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
|
||||
.expect_err("predict body must not carry unconsumed OpenAI fields");
|
||||
|
||||
assert!(error.message().contains("unsupported top-level fields"));
|
||||
assert!(body.get("model").is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_gemini_embedding_rejects_unconverted_openai_body() {
|
||||
let transport = sample_transport("vertex_ai", "https://aiplatform.googleapis.com");
|
||||
let mut body = json!({
|
||||
"model": "gemini-embedding-2",
|
||||
"input": "hello"
|
||||
});
|
||||
|
||||
let error =
|
||||
apply_transport_request_body_semantics(&mut body, &transport, "gemini:embedding")
|
||||
.expect_err("OpenAI embedding body must not be sent to Vertex native predict");
|
||||
|
||||
assert!(error.message().contains("cannot be mapped"));
|
||||
assert!(body.get("input").is_some());
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1 @@
|
||||
pub use aether_ai_formats::provider_compat::proxy::rules::*;
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,249 @@
|
||||
use aether_crypto::{decrypt_python_fernet_ciphertext, looks_like_python_fernet_ciphertext};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use super::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey, GatewayProviderTransportProvider,
|
||||
};
|
||||
|
||||
pub(super) fn map_provider(
|
||||
provider: StoredProviderCatalogProvider,
|
||||
) -> GatewayProviderTransportProvider {
|
||||
GatewayProviderTransportProvider {
|
||||
id: provider.id,
|
||||
name: provider.name,
|
||||
provider_type: provider.provider_type,
|
||||
website: provider.website,
|
||||
is_active: provider.is_active,
|
||||
keep_priority_on_conversion: provider.keep_priority_on_conversion,
|
||||
enable_format_conversion: provider.enable_format_conversion,
|
||||
concurrent_limit: provider.concurrent_limit,
|
||||
max_retries: provider.max_retries,
|
||||
proxy: normalize_optional_json(provider.proxy),
|
||||
request_timeout_secs: provider.request_timeout_secs,
|
||||
stream_first_byte_timeout_secs: provider.stream_first_byte_timeout_secs,
|
||||
config: normalize_optional_json(provider.config),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn map_endpoint(
|
||||
endpoint: StoredProviderCatalogEndpoint,
|
||||
) -> GatewayProviderTransportEndpoint {
|
||||
GatewayProviderTransportEndpoint {
|
||||
id: endpoint.id,
|
||||
provider_id: endpoint.provider_id,
|
||||
api_format: endpoint.api_format,
|
||||
api_family: endpoint.api_family,
|
||||
endpoint_kind: endpoint.endpoint_kind,
|
||||
is_active: endpoint.is_active,
|
||||
base_url: endpoint.base_url,
|
||||
header_rules: normalize_optional_json(endpoint.header_rules),
|
||||
body_rules: normalize_optional_json(endpoint.body_rules),
|
||||
max_retries: endpoint.max_retries,
|
||||
custom_path: endpoint.custom_path,
|
||||
config: normalize_optional_json(endpoint.config),
|
||||
format_acceptance_config: normalize_optional_json(endpoint.format_acceptance_config),
|
||||
proxy: normalize_optional_json(endpoint.proxy),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn map_key(
|
||||
key: StoredProviderCatalogKey,
|
||||
encryption_key: &str,
|
||||
fallback_encryption_keys: &[String],
|
||||
) -> Result<GatewayProviderTransportKey, DataLayerError> {
|
||||
let decrypted_api_key = key
|
||||
.encrypted_api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|ciphertext| {
|
||||
decrypt_secret(
|
||||
encryption_key,
|
||||
fallback_encryption_keys,
|
||||
ciphertext,
|
||||
"provider_api_keys.api_key",
|
||||
)
|
||||
})
|
||||
.transpose()?
|
||||
.unwrap_or_default();
|
||||
let decrypted_auth_config = key
|
||||
.encrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|ciphertext| {
|
||||
decrypt_secret(
|
||||
encryption_key,
|
||||
fallback_encryption_keys,
|
||||
ciphertext,
|
||||
"provider_api_keys.auth_config",
|
||||
)
|
||||
})
|
||||
.transpose()?;
|
||||
|
||||
Ok(GatewayProviderTransportKey {
|
||||
id: key.id,
|
||||
provider_id: key.provider_id,
|
||||
name: key.name,
|
||||
auth_type: key.auth_type,
|
||||
is_active: key.is_active,
|
||||
api_formats: normalize_string_list(
|
||||
normalize_optional_json(key.api_formats),
|
||||
"provider_api_keys.api_formats",
|
||||
)?,
|
||||
auth_type_by_format: normalize_optional_json(key.auth_type_by_format),
|
||||
allow_auth_channel_mismatch_formats: normalize_optional_json(
|
||||
key.allow_auth_channel_mismatch_formats,
|
||||
),
|
||||
allowed_models: normalize_string_list(
|
||||
normalize_optional_json(key.allowed_models),
|
||||
"provider_api_keys.allowed_models",
|
||||
)?,
|
||||
capabilities: normalize_optional_json(key.capabilities),
|
||||
rate_multipliers: normalize_optional_json(key.rate_multipliers),
|
||||
global_priority_by_format: normalize_optional_json(key.global_priority_by_format),
|
||||
expires_at_unix_secs: key.expires_at_unix_secs,
|
||||
proxy: normalize_optional_json(key.proxy),
|
||||
fingerprint: normalize_optional_json(key.fingerprint),
|
||||
upstream_metadata: normalize_optional_json(key.upstream_metadata),
|
||||
decrypted_api_key,
|
||||
decrypted_auth_config,
|
||||
})
|
||||
}
|
||||
|
||||
fn normalize_optional_json(value: Option<serde_json::Value>) -> Option<serde_json::Value> {
|
||||
match value {
|
||||
Some(serde_json::Value::Null) | None => None,
|
||||
Some(value) => Some(value),
|
||||
}
|
||||
}
|
||||
|
||||
fn decrypt_secret(
|
||||
encryption_key: &str,
|
||||
fallback_encryption_keys: &[String],
|
||||
ciphertext: &str,
|
||||
field_name: &str,
|
||||
) -> Result<String, DataLayerError> {
|
||||
if should_use_plaintext_secret(ciphertext, field_name) {
|
||||
return Ok(ciphertext.trim().to_string());
|
||||
}
|
||||
|
||||
match decrypt_python_fernet_ciphertext(encryption_key, ciphertext) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(error) => {
|
||||
for fallback_encryption_key in fallback_encryption_keys {
|
||||
if let Ok(value) =
|
||||
decrypt_python_fernet_ciphertext(fallback_encryption_key, ciphertext)
|
||||
{
|
||||
return Ok(value);
|
||||
}
|
||||
}
|
||||
Err(DataLayerError::UnexpectedValue(format!(
|
||||
"failed to decrypt {field_name}: {error}"
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn fallback_encryption_keys(primary_encryption_key: &str) -> Vec<String> {
|
||||
let mut keys = Vec::new();
|
||||
for env_key in ["AETHER_GATEWAY_DATA_ENCRYPTION_KEY", "ENCRYPTION_KEY"] {
|
||||
let Ok(value) = std::env::var(env_key) else {
|
||||
continue;
|
||||
};
|
||||
let value = value.trim();
|
||||
if value.is_empty()
|
||||
|| value == primary_encryption_key
|
||||
|| keys.iter().any(|existing| existing == value)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
keys.push(value.to_string());
|
||||
}
|
||||
keys
|
||||
}
|
||||
|
||||
fn should_use_plaintext_secret(ciphertext: &str, field_name: &str) -> bool {
|
||||
let ciphertext = ciphertext.trim();
|
||||
if ciphertext.is_empty() {
|
||||
return false;
|
||||
}
|
||||
|
||||
match field_name {
|
||||
"provider_api_keys.api_key" => {
|
||||
if ciphertext.starts_with('{') || ciphertext.starts_with('[') {
|
||||
return false;
|
||||
}
|
||||
!looks_like_python_fernet_ciphertext(ciphertext)
|
||||
}
|
||||
"provider_api_keys.auth_config" => {
|
||||
if ciphertext.starts_with('{') || ciphertext.starts_with('[') {
|
||||
return true;
|
||||
}
|
||||
false
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_string_list(
|
||||
raw: Option<serde_json::Value>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
let Some(raw) = raw else {
|
||||
return Ok(None);
|
||||
};
|
||||
normalize_string_list_value(&raw, field_name)
|
||||
}
|
||||
|
||||
fn normalize_string_list_value(
|
||||
raw: &serde_json::Value,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
match raw {
|
||||
serde_json::Value::Null => Ok(None),
|
||||
serde_json::Value::Array(items) => normalize_string_list_array(items, field_name).map(Some),
|
||||
serde_json::Value::String(raw) => normalize_embedded_string_list(raw, field_name),
|
||||
_ => Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} is not a JSON array"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_embedded_string_list(
|
||||
raw: &str,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
|
||||
return normalize_string_list_value(&decoded, field_name);
|
||||
}
|
||||
|
||||
Ok(Some(vec![raw.to_string()]))
|
||||
}
|
||||
|
||||
fn normalize_string_list_array(
|
||||
items: &[serde_json::Value],
|
||||
field_name: &str,
|
||||
) -> Result<Vec<String>, DataLayerError> {
|
||||
let mut values = Vec::with_capacity(items.len());
|
||||
for item in items {
|
||||
let Some(value) = item.as_str() else {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains a non-string item"
|
||||
)));
|
||||
};
|
||||
let value = value.trim();
|
||||
if !value.is_empty() {
|
||||
values.push(value.to_string());
|
||||
}
|
||||
}
|
||||
Ok(values)
|
||||
}
|
||||
@@ -0,0 +1,703 @@
|
||||
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::headers::force_identity_accept_encoding;
|
||||
use crate::rules::{
|
||||
apply_local_body_rules, apply_local_body_rules_with_request_headers,
|
||||
apply_local_header_rules_with_request_headers,
|
||||
};
|
||||
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,
|
||||
TextEventStreamIfStreamingOrWildcard,
|
||||
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::TextEventStreamIfStreamingOrWildcard => {
|
||||
if input.upstream_is_stream {
|
||||
set_accept_if_missing_or_wildcard(&mut headers, "text/event-stream");
|
||||
}
|
||||
}
|
||||
StandardPlanFallbackAcceptPolicy::TextEventStreamRequired => {
|
||||
headers.insert("accept".to_string(), "text/event-stream".to_string());
|
||||
}
|
||||
StandardPlanFallbackAcceptPolicy::ProviderEventStreamIfMissing => {
|
||||
set_accept_if_missing_or_wildcard(&mut headers, "application/vnd.amazon.eventstream");
|
||||
}
|
||||
}
|
||||
|
||||
if input.upstream_is_stream {
|
||||
force_identity_accept_encoding(&mut headers);
|
||||
}
|
||||
|
||||
headers
|
||||
}
|
||||
|
||||
fn set_accept_if_missing_or_wildcard(headers: &mut BTreeMap<String, String>, value: &str) {
|
||||
let Some(existing_key) = headers
|
||||
.keys()
|
||||
.find(|key| key.eq_ignore_ascii_case("accept"))
|
||||
.cloned()
|
||||
else {
|
||||
headers.insert("accept".to_string(), value.to_string());
|
||||
return;
|
||||
};
|
||||
|
||||
if headers
|
||||
.get(&existing_key)
|
||||
.is_some_and(|existing_value| accept_is_wildcard_only(existing_value))
|
||||
{
|
||||
headers.insert(existing_key, value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
fn accept_is_wildcard_only(value: &str) -> bool {
|
||||
let mut saw_value = false;
|
||||
for raw_part in value.split(',') {
|
||||
let media_type = raw_part.trim().split(';').next().unwrap_or_default().trim();
|
||||
if media_type.is_empty() {
|
||||
continue;
|
||||
}
|
||||
saw_value = true;
|
||||
if media_type != "*/*" {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
saw_value
|
||||
}
|
||||
|
||||
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 apply_standard_provider_request_body_rules_with_request_headers(
|
||||
mut provider_request_body: Value,
|
||||
body_rules: Option<&Value>,
|
||||
original_request_body: &Value,
|
||||
request_headers: &http::HeaderMap,
|
||||
) -> Option<Value> {
|
||||
if !apply_local_body_rules_with_request_headers(
|
||||
&mut provider_request_body,
|
||||
body_rules,
|
||||
Some(original_request_body),
|
||||
Some(request_headers),
|
||||
) {
|
||||
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"),
|
||||
)
|
||||
};
|
||||
|
||||
crate::apply_local_auth_config_header_overrides(
|
||||
&mut headers,
|
||||
input.transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
let protected_headers = if uses_vertex_query_auth {
|
||||
&["content-type"][..]
|
||||
} else {
|
||||
&[input.auth_header, "content-type"][..]
|
||||
};
|
||||
if !apply_local_header_rules_with_request_headers(
|
||||
&mut headers,
|
||||
input.header_rules,
|
||||
protected_headers,
|
||||
input.provider_request_body,
|
||||
Some(input.original_request_body),
|
||||
Some(input.headers),
|
||||
) {
|
||||
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());
|
||||
force_identity_accept_encoding(&mut headers);
|
||||
}
|
||||
|
||||
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,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn 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"));
|
||||
request_headers.insert(
|
||||
http::header::ACCEPT_ENCODING,
|
||||
"gzip, br".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("accept-encoding"),
|
||||
Some(&"identity".to_string())
|
||||
);
|
||||
assert_eq!(resolved.headers.get("x-client"), Some(&"demo".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_sync_headers_preserving_supported_accept_encoding() {
|
||||
let mut request_headers = HeaderMap::new();
|
||||
request_headers.insert(
|
||||
http::header::ACCEPT_ENCODING,
|
||||
"gzip, br".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: false,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
resolved.headers.get("accept-encoding"),
|
||||
Some(&"gzip".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 endpoint_header_rules_do_not_override_protected_authorization() {
|
||||
let mut transport = sample_transport("openai:responses");
|
||||
transport.endpoint.header_rules = Some(json!([
|
||||
{"action":"set","key":"authorization","value":"Bearer imported-session"}
|
||||
]));
|
||||
|
||||
let resolved =
|
||||
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
|
||||
transport: &transport,
|
||||
provider_api_format: "openai:responses",
|
||||
same_format: true,
|
||||
headers: &HeaderMap::new(),
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer jwt-access-token",
|
||||
extra_headers: &BTreeMap::new(),
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &json!({"model":"gpt-test"}),
|
||||
original_request_body: &json!({"model":"gpt-test"}),
|
||||
upstream_is_stream: false,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
resolved.headers.get("authorization"),
|
||||
Some(&"Bearer jwt-access-token".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_config_headers_can_override_authorization_when_refresh_token_is_present() {
|
||||
let mut transport = sample_transport("openai:responses");
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"refresh_token": "rt-1",
|
||||
"headers": {
|
||||
"authorization": "Bearer imported-session"
|
||||
}
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
let resolved =
|
||||
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
|
||||
transport: &transport,
|
||||
provider_api_format: "openai:responses",
|
||||
same_format: true,
|
||||
headers: &HeaderMap::new(),
|
||||
auth_header: "authorization",
|
||||
auth_value: "Bearer refreshed-access-token",
|
||||
extra_headers: &BTreeMap::new(),
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &json!({"model":"gpt-test"}),
|
||||
original_request_body: &json!({"model":"gpt-test"}),
|
||||
upstream_is_stream: false,
|
||||
})
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
resolved.headers.get("authorization"),
|
||||
Some(&"Bearer imported-session".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 stream_fallback_headers_treat_wildcard_accept_as_absent() {
|
||||
let mut request_headers = HeaderMap::new();
|
||||
request_headers.insert(http::header::ACCEPT, "*/*".parse().expect("header"));
|
||||
request_headers.insert(
|
||||
http::header::ACCEPT_ENCODING,
|
||||
"gzip, br".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::TextEventStreamIfStreamingOrWildcard,
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
headers.get("accept"),
|
||||
Some(&"text/event-stream".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("accept-encoding"),
|
||||
Some(&"identity".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_fallback_headers_preserve_wildcard_in_missing_only_mode() {
|
||||
let mut request_headers = HeaderMap::new();
|
||||
request_headers.insert(http::header::ACCEPT, "*/*".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: "gemini:generate_content",
|
||||
client_api_format: "openai:responses",
|
||||
upstream_is_stream: true,
|
||||
build_from_request_when_empty: true,
|
||||
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreaming,
|
||||
});
|
||||
|
||||
assert_eq!(headers.get("accept"), Some(&"*/*".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_fallback_headers_preserve_explicit_accept() {
|
||||
let mut existing_headers = BTreeMap::new();
|
||||
existing_headers.insert("accept".to_string(), "application/json".to_string());
|
||||
|
||||
let headers = build_standard_plan_fallback_headers(StandardPlanFallbackHeadersInput {
|
||||
request_headers: &HeaderMap::new(),
|
||||
existing_provider_request_headers: existing_headers,
|
||||
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: false,
|
||||
accept_policy: StandardPlanFallbackAcceptPolicy::TextEventStreamIfStreamingOrWildcard,
|
||||
});
|
||||
|
||||
assert_eq!(headers.get("accept"), Some(&"application/json".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/v1", 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/v1",
|
||||
Some("x=1"),
|
||||
true,
|
||||
),
|
||||
"https://api.example.com/v1/responses/compact?x=1"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,752 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use url::form_urlencoded;
|
||||
use url::Url;
|
||||
|
||||
pub fn build_openai_chat_url(upstream_base_url: &str, query: Option<&str>) -> String {
|
||||
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed = trimmed.trim_end_matches('/');
|
||||
let mut url = format!("{trimmed}/chat/completions");
|
||||
append_merged_query(&mut url, base_query, None, query, &[]);
|
||||
url
|
||||
}
|
||||
|
||||
pub fn build_openai_responses_url(
|
||||
upstream_base_url: &str,
|
||||
query: Option<&str>,
|
||||
compact: bool,
|
||||
) -> String {
|
||||
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed = trimmed.trim_end_matches('/');
|
||||
let suffix = if compact {
|
||||
"/responses/compact"
|
||||
} else {
|
||||
"/responses"
|
||||
};
|
||||
let mut url = format!("{trimmed}{suffix}");
|
||||
append_merged_query(&mut url, base_query, None, query, &[]);
|
||||
url
|
||||
}
|
||||
|
||||
pub fn build_openai_search_url(upstream_base_url: &str, query: Option<&str>) -> String {
|
||||
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed = trimmed.trim_end_matches('/');
|
||||
let mut url = format!("{trimmed}/alpha/search");
|
||||
append_merged_query(&mut url, base_query, None, query, &[]);
|
||||
url
|
||||
}
|
||||
|
||||
pub fn build_openai_image_url(
|
||||
upstream_base_url: &str,
|
||||
request_path: Option<&str>,
|
||||
query: Option<&str>,
|
||||
) -> String {
|
||||
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed = trimmed.trim_end_matches('/');
|
||||
let suffix = openai_image_path_suffix(request_path);
|
||||
let mut url = if openai_image_base_includes_operation_path(trimmed) {
|
||||
trimmed.to_string()
|
||||
} else {
|
||||
format!("{trimmed}{suffix}")
|
||||
};
|
||||
append_merged_query(&mut url, base_query, None, query, &[]);
|
||||
url
|
||||
}
|
||||
|
||||
fn openai_image_path_suffix(request_path: Option<&str>) -> &'static str {
|
||||
match request_path
|
||||
.map(str::trim)
|
||||
.map(|value| value.trim_end_matches('/'))
|
||||
{
|
||||
Some("/v1/images/edits") | Some("/images/edits") => "/images/edits",
|
||||
_ => "/images/generations",
|
||||
}
|
||||
}
|
||||
|
||||
fn openai_image_base_includes_operation_path(base_url: &str) -> bool {
|
||||
let path = Url::parse(base_url)
|
||||
.ok()
|
||||
.map(|url| url.path().trim_end_matches('/').to_string())
|
||||
.unwrap_or_else(|| base_url.trim_end_matches('/').to_string());
|
||||
path.ends_with("/images/generations") || path.ends_with("/images/edits")
|
||||
}
|
||||
|
||||
pub fn build_claude_messages_url(upstream_base_url: &str, query: Option<&str>) -> String {
|
||||
let (trimmed, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed = trimmed.trim_end_matches('/');
|
||||
let mut url = format!("{trimmed}/messages");
|
||||
append_merged_query(&mut url, base_query, None, query, &[]);
|
||||
url
|
||||
}
|
||||
|
||||
pub fn build_gemini_content_url(
|
||||
upstream_base_url: &str,
|
||||
model: &str,
|
||||
stream: bool,
|
||||
query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
let trimmed_model = model.trim();
|
||||
if trimmed_base_url.is_empty() || trimmed_model.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let operation = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
let mut url = if trimmed_base_url.ends_with("/v1") || trimmed_base_url.ends_with("/v1beta") {
|
||||
format!("{trimmed_base_url}/models/{trimmed_model}:{operation}")
|
||||
} else if gemini_content_base_url_contains_model_path(trimmed_base_url) {
|
||||
let trimmed_base_url = strip_gemini_content_action(trimmed_base_url);
|
||||
format!("{trimmed_base_url}:{operation}")
|
||||
} else {
|
||||
format!("{trimmed_base_url}/v1beta/models/{trimmed_model}:{operation}")
|
||||
};
|
||||
append_merged_query(&mut url, base_query, None, query, &["key"]);
|
||||
Some(url)
|
||||
}
|
||||
|
||||
pub fn normalize_gemini_content_action_path(path: &str, stream: bool) -> String {
|
||||
let trimmed = path.trim();
|
||||
let (path, query) = split_path_query(trimmed);
|
||||
let action = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
let normalized = strip_gemini_content_action(path);
|
||||
let normalized = if normalized.len() == path.len() {
|
||||
path.to_string()
|
||||
} else {
|
||||
format!("{normalized}:{action}")
|
||||
};
|
||||
match query {
|
||||
Some(query) => format!("{normalized}?{query}"),
|
||||
None => normalized,
|
||||
}
|
||||
}
|
||||
|
||||
fn strip_gemini_content_action(value: &str) -> &str {
|
||||
value
|
||||
.strip_suffix(":streamGenerateContent")
|
||||
.or_else(|| value.strip_suffix(":generateContent"))
|
||||
.unwrap_or(value)
|
||||
}
|
||||
|
||||
fn gemini_content_base_url_contains_model_path(value: &str) -> bool {
|
||||
value.contains("/v1/models/") || value.contains("/v1beta/models/")
|
||||
}
|
||||
|
||||
pub fn build_gemini_video_predict_long_running_url(
|
||||
upstream_base_url: &str,
|
||||
model: &str,
|
||||
query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
let trimmed_model = model.trim();
|
||||
if trimmed_base_url.is_empty() || trimmed_model.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut url = if trimmed_base_url.ends_with("/v1") || trimmed_base_url.ends_with("/v1beta") {
|
||||
format!("{trimmed_base_url}/models/{trimmed_model}:predictLongRunning")
|
||||
} else if gemini_content_base_url_contains_model_path(trimmed_base_url) {
|
||||
format!("{trimmed_base_url}:predictLongRunning")
|
||||
} else {
|
||||
format!("{trimmed_base_url}/v1beta/models/{trimmed_model}:predictLongRunning")
|
||||
};
|
||||
append_merged_query(&mut url, base_query, None, query, &["key"]);
|
||||
Some(url)
|
||||
}
|
||||
|
||||
pub fn build_passthrough_path_url(
|
||||
upstream_base_url: &str,
|
||||
path: &str,
|
||||
query: Option<&str>,
|
||||
blocked_keys: &[&str],
|
||||
) -> Option<String> {
|
||||
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
let trimmed_path = path.trim();
|
||||
if trimmed_base_url.is_empty() || trimmed_path.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let (trimmed_path, path_query) = split_path_query(trimmed_path);
|
||||
let normalized_base_url =
|
||||
if trimmed_base_url.ends_with("/v1beta") && trimmed_path.starts_with("/v1beta") {
|
||||
trimmed_base_url.trim_end_matches("/v1beta")
|
||||
} else if trimmed_base_url.ends_with("/v1") && trimmed_path.starts_with("/v1/") {
|
||||
trimmed_base_url.trim_end_matches("/v1")
|
||||
} else {
|
||||
trimmed_base_url
|
||||
};
|
||||
|
||||
let mut url = format!("{normalized_base_url}{trimmed_path}");
|
||||
append_merged_query(&mut url, base_query, path_query, query, blocked_keys);
|
||||
Some(url)
|
||||
}
|
||||
|
||||
pub fn build_bigmodel_coding_models_url(upstream_base_url: &str) -> Option<String> {
|
||||
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
if trimmed_base_url.is_empty() || !bigmodel_coding_models_base_is_supported(trimmed_base_url) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let path = Url::parse(trimmed_base_url)
|
||||
.ok()
|
||||
.map(|url| url.path().trim_end_matches('/').to_string())
|
||||
.unwrap_or_else(|| trimmed_base_url.trim_end_matches('/').to_string());
|
||||
let mut url = if path.ends_with("/models") {
|
||||
trimmed_base_url.to_string()
|
||||
} else {
|
||||
format!("{trimmed_base_url}/models")
|
||||
};
|
||||
append_merged_query(&mut url, base_query, None, None, &[]);
|
||||
Some(url)
|
||||
}
|
||||
|
||||
pub fn build_openai_compatible_models_url(upstream_base_url: &str) -> Option<String> {
|
||||
if let Some(url) = build_bigmodel_coding_models_url(upstream_base_url) {
|
||||
return Some(url);
|
||||
}
|
||||
|
||||
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
if trimmed_base_url.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut url = if trimmed_base_url.ends_with("/models") {
|
||||
trimmed_base_url.to_string()
|
||||
} else {
|
||||
format!("{trimmed_base_url}/models")
|
||||
};
|
||||
append_merged_query(&mut url, base_query, None, None, &[]);
|
||||
Some(url)
|
||||
}
|
||||
|
||||
pub fn build_gemini_files_passthrough_url(
|
||||
upstream_base_url: &str,
|
||||
path: &str,
|
||||
query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let (trimmed_base_url, base_query) = split_base_url_query(upstream_base_url);
|
||||
let trimmed_base_url = trimmed_base_url.trim_end_matches('/');
|
||||
let trimmed_path = path.trim();
|
||||
if trimmed_base_url.is_empty() || trimmed_path.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let (trimmed_path, path_query) = split_path_query(trimmed_path);
|
||||
let normalized_base_url = if trimmed_base_url.ends_with("/v1beta")
|
||||
&& (trimmed_path.starts_with("/v1beta/") || trimmed_path.starts_with("/upload/v1beta/"))
|
||||
{
|
||||
trimmed_base_url.trim_end_matches("/v1beta")
|
||||
} else {
|
||||
trimmed_base_url
|
||||
};
|
||||
|
||||
let mut url = format!("{normalized_base_url}{trimmed_path}");
|
||||
append_merged_query(&mut url, base_query, path_query, query, &["key"]);
|
||||
Some(url)
|
||||
}
|
||||
|
||||
fn split_base_url_query(base_url: &str) -> (&str, Option<&str>) {
|
||||
let trimmed = base_url.trim();
|
||||
trimmed
|
||||
.split_once('?')
|
||||
.map(|(base, query)| (base, Some(query)))
|
||||
.unwrap_or((trimmed, None))
|
||||
}
|
||||
|
||||
pub(crate) fn google_openai_compat_base_includes_api_root(base_url: &str) -> bool {
|
||||
let Ok(parsed) = Url::parse(base_url.trim()) else {
|
||||
return false;
|
||||
};
|
||||
let Some(host) = parsed.host_str().map(|value| value.to_ascii_lowercase()) else {
|
||||
return false;
|
||||
};
|
||||
let path = parsed.path().trim_end_matches('/');
|
||||
|
||||
if host == "generativelanguage.googleapis.com" {
|
||||
return path == "/v1beta/openai" || path == "/v1/openai";
|
||||
}
|
||||
|
||||
if looks_like_vertex_ai_host(&host) {
|
||||
return path.ends_with("/endpoints/openapi");
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
|
||||
pub fn openai_compatible_base_includes_api_root(base_url: &str) -> bool {
|
||||
let trimmed = base_url.trim().trim_end_matches('/');
|
||||
trimmed.ends_with("/v1")
|
||||
|| google_openai_compat_base_includes_api_root(trimmed)
|
||||
|| bigmodel_coding_base_includes_api_root(trimmed)
|
||||
|| openai_compatible_base_includes_unversioned_api_root(trimmed)
|
||||
}
|
||||
|
||||
pub fn v1_compatible_base_includes_api_root(base_url: &str) -> bool {
|
||||
let trimmed = base_url.trim().trim_end_matches('/');
|
||||
trimmed.ends_with("/v1") || openai_compatible_base_includes_unversioned_api_root(trimmed)
|
||||
}
|
||||
|
||||
pub fn openai_compatible_base_includes_unversioned_api_root(base_url: &str) -> bool {
|
||||
let trimmed = base_url.trim().trim_end_matches('/');
|
||||
let path = Url::parse(trimmed)
|
||||
.ok()
|
||||
.map(|url| url.path().trim_end_matches('/').to_ascii_lowercase())
|
||||
.unwrap_or_else(|| {
|
||||
trimmed
|
||||
.split_once('/')
|
||||
.map(|(_, path)| format!("/{path}"))
|
||||
.unwrap_or_default()
|
||||
.trim_end_matches('/')
|
||||
.to_ascii_lowercase()
|
||||
});
|
||||
!path.is_empty()
|
||||
}
|
||||
|
||||
fn bigmodel_coding_base_includes_api_root(base_url: &str) -> bool {
|
||||
let Ok(parsed) = Url::parse(base_url.trim()) else {
|
||||
return false;
|
||||
};
|
||||
let Some(host) = parsed.host_str().map(|value| value.to_ascii_lowercase()) else {
|
||||
return false;
|
||||
};
|
||||
host == "open.bigmodel.cn" && parsed.path().trim_end_matches('/') == "/api/coding/paas/v4"
|
||||
}
|
||||
|
||||
fn bigmodel_coding_models_base_is_supported(base_url: &str) -> bool {
|
||||
let Ok(parsed) = Url::parse(base_url.trim()) else {
|
||||
return false;
|
||||
};
|
||||
let Some(host) = parsed.host_str().map(|value| value.to_ascii_lowercase()) else {
|
||||
return false;
|
||||
};
|
||||
if host != "open.bigmodel.cn" {
|
||||
return false;
|
||||
}
|
||||
matches!(
|
||||
parsed.path().trim_end_matches('/'),
|
||||
"/api/coding/paas/v4" | "/api/coding/paas/v4/models"
|
||||
)
|
||||
}
|
||||
|
||||
fn looks_like_vertex_ai_host(host: &str) -> bool {
|
||||
const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com";
|
||||
host == VERTEX_AI_HOST
|
||||
|| host.ends_with(&format!(".{VERTEX_AI_HOST}"))
|
||||
|| host.ends_with(&format!("-{VERTEX_AI_HOST}"))
|
||||
}
|
||||
|
||||
fn split_path_query(path: &str) -> (&str, Option<&str>) {
|
||||
path.split_once('?')
|
||||
.map(|(path, query)| (path, Some(query)))
|
||||
.unwrap_or((path, None))
|
||||
}
|
||||
|
||||
fn append_merged_query(
|
||||
url: &mut String,
|
||||
base_query: Option<&str>,
|
||||
path_query: Option<&str>,
|
||||
request_query: Option<&str>,
|
||||
blocked_keys: &[&str],
|
||||
) {
|
||||
let Some(query) = merge_query_layers(base_query, path_query, request_query, blocked_keys)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if url.contains('?') {
|
||||
url.push('&');
|
||||
} else {
|
||||
url.push('?');
|
||||
}
|
||||
url.push_str(&query);
|
||||
}
|
||||
|
||||
fn merge_query_layers(
|
||||
base_query: Option<&str>,
|
||||
path_query: Option<&str>,
|
||||
request_query: Option<&str>,
|
||||
blocked_keys: &[&str],
|
||||
) -> Option<String> {
|
||||
if blocked_keys.is_empty()
|
||||
&& path_query.is_none()
|
||||
&& base_query.is_none()
|
||||
&& request_query
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
{
|
||||
return request_query
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
}
|
||||
|
||||
let mut merged = BTreeMap::new();
|
||||
for source in [base_query, path_query, request_query] {
|
||||
merge_query_string(&mut merged, source, blocked_keys);
|
||||
}
|
||||
if merged.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in merged {
|
||||
serializer.append_pair(&key, &value);
|
||||
}
|
||||
Some(serializer.finish())
|
||||
}
|
||||
|
||||
fn merge_query_string(
|
||||
out: &mut BTreeMap<String, String>,
|
||||
query: Option<&str>,
|
||||
blocked_keys: &[&str],
|
||||
) {
|
||||
let Some(query) = query.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||
return;
|
||||
};
|
||||
|
||||
for (key, value) in form_urlencoded::parse(query.as_bytes()) {
|
||||
if blocked_keys
|
||||
.iter()
|
||||
.any(|blocked| key.as_ref().eq_ignore_ascii_case(blocked))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
out.insert(key.into_owned(), value.into_owned());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_bigmodel_coding_models_url, build_claude_messages_url, build_gemini_content_url,
|
||||
build_gemini_files_passthrough_url, build_gemini_video_predict_long_running_url,
|
||||
build_openai_chat_url, build_openai_compatible_models_url, build_openai_image_url,
|
||||
build_openai_responses_url, build_openai_search_url, build_passthrough_path_url,
|
||||
normalize_gemini_content_action_path,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn merges_base_url_query_for_same_format_urls() {
|
||||
assert_eq!(
|
||||
build_openai_chat_url(
|
||||
"https://api.openai.example/v1?tenant=demo",
|
||||
Some("mode=fast&tenant=override")
|
||||
),
|
||||
"https://api.openai.example/v1/chat/completions?mode=fast&tenant=override"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_chat_url_preserves_google_openai_compat_roots() {
|
||||
assert_eq!(
|
||||
build_openai_chat_url(
|
||||
"https://generativelanguage.googleapis.com/v1beta/openai",
|
||||
Some("trace=1")
|
||||
),
|
||||
"https://generativelanguage.googleapis.com/v1beta/openai/chat/completions?trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_chat_url(
|
||||
"https://aiplatform.googleapis.com/v1/projects/project-1/locations/global/endpoints/openapi",
|
||||
None,
|
||||
),
|
||||
"https://aiplatform.googleapis.com/v1/projects/project-1/locations/global/endpoints/openapi/chat/completions"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_urls_preserve_bigmodel_coding_api_root() {
|
||||
assert_eq!(
|
||||
build_openai_chat_url(
|
||||
"https://open.bigmodel.cn/api/coding/paas/v4",
|
||||
Some("trace=1")
|
||||
),
|
||||
"https://open.bigmodel.cn/api/coding/paas/v4/chat/completions?trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_responses_url("https://open.bigmodel.cn/api/coding/paas/v4", None, false),
|
||||
"https://open.bigmodel.cn/api/coding/paas/v4/responses"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_urls_preserve_unversioned_api_root() {
|
||||
assert_eq!(
|
||||
build_openai_chat_url("https://proxy.example.com/api", Some("trace=1")),
|
||||
"https://proxy.example.com/api/chat/completions?trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_chat_url("https://proxy.example.com/openai", None),
|
||||
"https://proxy.example.com/openai/chat/completions"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_chat_url("https://proxy.example.com", None),
|
||||
"https://proxy.example.com/chat/completions"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_chat_url("https://api.deepseek.com", None),
|
||||
"https://api.deepseek.com/chat/completions"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_responses_url("https://proxy.example.com/api", None, false),
|
||||
"https://proxy.example.com/api/responses"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_responses_url("https://api.deepseek.com", None, false),
|
||||
"https://api.deepseek.com/responses"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_image_url(
|
||||
"https://proxy.example.com/api",
|
||||
Some("/v1/images/generations"),
|
||||
None
|
||||
),
|
||||
"https://proxy.example.com/api/images/generations"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_messages_url_preserves_v1_and_unversioned_api_roots() {
|
||||
assert_eq!(
|
||||
build_claude_messages_url("https://api.anthropic.example/v1", Some("trace=1")),
|
||||
"https://api.anthropic.example/v1/messages?trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_messages_url("https://proxy.example.com/api", None),
|
||||
"https://proxy.example.com/api/messages"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_messages_url("https://proxy.example.com/anthropic", None),
|
||||
"https://proxy.example.com/anthropic/messages"
|
||||
);
|
||||
assert_eq!(
|
||||
build_claude_messages_url("https://api.anthropic.example", None),
|
||||
"https://api.anthropic.example/messages"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn bigmodel_coding_models_url_uses_models_resource() {
|
||||
assert_eq!(
|
||||
build_bigmodel_coding_models_url(
|
||||
"https://open.bigmodel.cn/api/coding/paas/v4?tenant=demo"
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://open.bigmodel.cn/api/coding/paas/v4/models?tenant=demo")
|
||||
);
|
||||
assert_eq!(
|
||||
build_bigmodel_coding_models_url("https://open.bigmodel.cn/api/coding/paas/v4/models")
|
||||
.as_deref(),
|
||||
Some("https://open.bigmodel.cn/api/coding/paas/v4/models")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_compatible_models_url_preserves_unversioned_api_root() {
|
||||
assert_eq!(
|
||||
build_openai_compatible_models_url("https://proxy.example.com/api?tenant=demo")
|
||||
.as_deref(),
|
||||
Some("https://proxy.example.com/api/models?tenant=demo")
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_compatible_models_url("https://proxy.example.com/openai").as_deref(),
|
||||
Some("https://proxy.example.com/openai/models")
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_compatible_models_url("https://proxy.example.com").as_deref(),
|
||||
Some("https://proxy.example.com/models")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_responses_url_preserves_codex_path_prefix() {
|
||||
assert_eq!(
|
||||
build_openai_responses_url("https://tiger.bookapi.cc/codex", None, false),
|
||||
"https://tiger.bookapi.cc/codex/responses"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_responses_url("https://tiger.bookapi.cc/codex?tenant=demo", None, true),
|
||||
"https://tiger.bookapi.cc/codex/responses/compact?tenant=demo"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_search_url_preserves_api_and_codex_roots() {
|
||||
assert_eq!(
|
||||
build_openai_search_url(
|
||||
"https://api.openai.com/v1?tenant=base",
|
||||
Some("trace=1&tenant=request")
|
||||
),
|
||||
"https://api.openai.com/v1/alpha/search?tenant=request&trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_search_url("https://chatgpt.com/backend-api/codex/", None),
|
||||
"https://chatgpt.com/backend-api/codex/alpha/search"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_image_url_uses_images_surface() {
|
||||
assert_eq!(
|
||||
build_openai_image_url(
|
||||
"https://api.openai.example/v1?tenant=demo",
|
||||
Some("/v1/images/generations"),
|
||||
Some("trace=1")
|
||||
),
|
||||
"https://api.openai.example/v1/images/generations?tenant=demo&trace=1"
|
||||
);
|
||||
assert_eq!(
|
||||
build_openai_image_url(
|
||||
"https://api.openai.example/v1",
|
||||
Some("/v1/images/edits"),
|
||||
None
|
||||
),
|
||||
"https://api.openai.example/v1/images/edits"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merges_base_url_query_for_dynamic_gemini_content_urls() {
|
||||
assert_eq!(
|
||||
build_gemini_content_url(
|
||||
"https://generativelanguage.googleapis.com/v1beta?alt=sse",
|
||||
"gemini-2.5-pro",
|
||||
true,
|
||||
Some("foo=bar&key=secret")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:streamGenerateContent?alt=sse&foo=bar"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gemini_content_urls_rewrite_existing_base_action_for_stream_mode() {
|
||||
assert_eq!(
|
||||
build_gemini_content_url(
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent",
|
||||
"ignored-model",
|
||||
true,
|
||||
Some("foo=bar")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:streamGenerateContent?foo=bar"
|
||||
)
|
||||
);
|
||||
assert_eq!(
|
||||
build_gemini_content_url(
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:streamGenerateContent",
|
||||
"ignored-model",
|
||||
false,
|
||||
Some("foo=bar")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent?foo=bar"
|
||||
)
|
||||
);
|
||||
assert_eq!(
|
||||
build_gemini_content_url(
|
||||
"https://generativelanguage.googleapis.com/v1/models/gemini-2.5-pro:generateContent",
|
||||
"ignored-model",
|
||||
true,
|
||||
Some("foo=bar")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://generativelanguage.googleapis.com/v1/models/gemini-2.5-pro:streamGenerateContent?foo=bar"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalizes_gemini_content_action_in_custom_paths() {
|
||||
assert_eq!(
|
||||
normalize_gemini_content_action_path(
|
||||
"/v1beta/models/gemini-2.5-pro:generateContent?alt=sse",
|
||||
true
|
||||
),
|
||||
"/v1beta/models/gemini-2.5-pro:streamGenerateContent?alt=sse"
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_gemini_content_action_path(
|
||||
"/v1beta/models/gemini-2.5-pro:streamGenerateContent",
|
||||
false
|
||||
),
|
||||
"/v1beta/models/gemini-2.5-pro:generateContent"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merges_base_path_and_request_query_for_passthrough_paths() {
|
||||
assert_eq!(
|
||||
build_passthrough_path_url(
|
||||
"https://api.openai.example/v1?tenant=demo",
|
||||
"/videos/generations?variant=video",
|
||||
Some("size=1024"),
|
||||
&[]
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://api.openai.example/v1/videos/generations?size=1024&tenant=demo&variant=video"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn passthrough_path_does_not_duplicate_openai_v1_root() {
|
||||
assert_eq!(
|
||||
build_passthrough_path_url(
|
||||
"https://api.openai.example/v1?tenant=demo",
|
||||
"/v1/chat/completions?variant=chat",
|
||||
Some("trace=1"),
|
||||
&[]
|
||||
)
|
||||
.as_deref(),
|
||||
Some("https://api.openai.example/v1/chat/completions?tenant=demo&trace=1&variant=chat")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merges_base_url_query_for_gemini_files_passthrough_urls() {
|
||||
assert_eq!(
|
||||
build_gemini_files_passthrough_url(
|
||||
"https://generativelanguage.googleapis.com/v1beta?alt=media",
|
||||
"/upload/v1beta/files?uploadType=resumable",
|
||||
Some("key=secret&pageSize=10")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://generativelanguage.googleapis.com/upload/v1beta/files?alt=media&pageSize=10&uploadType=resumable"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merges_base_url_query_for_gemini_video_urls() {
|
||||
assert_eq!(
|
||||
build_gemini_video_predict_long_running_url(
|
||||
"https://generativelanguage.googleapis.com/v1beta?alt=sse",
|
||||
"veo-3.0-generate-preview",
|
||||
Some("foo=bar&key=secret")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/veo-3.0-generate-preview:predictLongRunning?alt=sse&foo=bar"
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,485 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
||||
use base64::Engine as _;
|
||||
use rsa::pkcs1::DecodeRsaPrivateKey;
|
||||
use rsa::pkcs1v15::SigningKey;
|
||||
use rsa::pkcs8::DecodePrivateKey;
|
||||
use rsa::signature::{SignatureEncoding, Signer};
|
||||
use rsa::RsaPrivateKey;
|
||||
use serde_json::{json, Value};
|
||||
use sha2::Sha256;
|
||||
use url::form_urlencoded;
|
||||
|
||||
use super::super::oauth_refresh::{
|
||||
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthHttpRequest, LocalOAuthRefreshAdapter,
|
||||
LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
pub const VERTEX_API_KEY_QUERY_PARAM: &str = "key";
|
||||
pub const VERTEX_SERVICE_ACCOUNT_AUTH_HEADER: &str = "authorization";
|
||||
pub const VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE: &str = "vertex_ai";
|
||||
pub const GOOGLE_OAUTH_TOKEN_URL: &str = "https://oauth2.googleapis.com/token";
|
||||
const GOOGLE_CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
|
||||
const SERVICE_ACCOUNT_REFRESH_SKEW_SECS: u64 = 120;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct VertexApiKeyQueryAuth {
|
||||
pub name: &'static str,
|
||||
pub value: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct VertexServiceAccountAuthConfig {
|
||||
pub client_email: String,
|
||||
pub private_key: String,
|
||||
pub project_id: String,
|
||||
pub token_uri: String,
|
||||
pub region: Option<String>,
|
||||
pub model_regions: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
pub fn resolve_local_vertex_api_key_query_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<VertexApiKeyQueryAuth> {
|
||||
if !super::is_vertex_api_key_transport_context(transport) {
|
||||
return None;
|
||||
}
|
||||
|
||||
if transport.key.decrypted_auth_config.is_some() {
|
||||
return None;
|
||||
}
|
||||
|
||||
if !transport
|
||||
.key
|
||||
.auth_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("api_key")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if secret.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(VertexApiKeyQueryAuth {
|
||||
name: VERTEX_API_KEY_QUERY_PARAM,
|
||||
value: secret.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn resolve_local_vertex_service_account_auth_config(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<VertexServiceAccountAuthConfig> {
|
||||
if !super::is_vertex_service_account_transport_context(transport) {
|
||||
return None;
|
||||
}
|
||||
parse_vertex_service_account_auth_config(transport.key.decrypted_auth_config.as_deref())
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_service_account_auth_resolution(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
resolve_local_vertex_service_account_auth_config(transport).is_some()
|
||||
}
|
||||
|
||||
pub fn parse_vertex_service_account_auth_config(
|
||||
raw: Option<&str>,
|
||||
) -> Option<VertexServiceAccountAuthConfig> {
|
||||
let raw = raw.map(str::trim).filter(|value| !value.is_empty())?;
|
||||
let value: Value = serde_json::from_str(raw).ok()?;
|
||||
parse_vertex_service_account_auth_config_value(&value)
|
||||
}
|
||||
|
||||
fn parse_vertex_service_account_auth_config_value(
|
||||
value: &Value,
|
||||
) -> Option<VertexServiceAccountAuthConfig> {
|
||||
let client_email = json_string(value.get("client_email"))?;
|
||||
let private_key = json_string(value.get("private_key"))?;
|
||||
let project_id = json_string(value.get("project_id"))?;
|
||||
let token_uri =
|
||||
json_string(value.get("token_uri")).unwrap_or_else(|| GOOGLE_OAUTH_TOKEN_URL.to_string());
|
||||
let region = json_string(value.get("region"));
|
||||
let model_regions = value
|
||||
.get("model_regions")
|
||||
.and_then(Value::as_object)
|
||||
.map(|items| {
|
||||
items
|
||||
.iter()
|
||||
.filter_map(|(model, region)| {
|
||||
let model = model.trim();
|
||||
let region = region.as_str()?.trim();
|
||||
(!model.is_empty() && !region.is_empty())
|
||||
.then(|| (model.to_string(), region.to_string()))
|
||||
})
|
||||
.collect::<BTreeMap<_, _>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
Some(VertexServiceAccountAuthConfig {
|
||||
client_email,
|
||||
private_key,
|
||||
project_id,
|
||||
token_uri,
|
||||
region,
|
||||
model_regions,
|
||||
})
|
||||
}
|
||||
|
||||
fn json_string(value: Option<&Value>) -> Option<String> {
|
||||
value
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct VertexServiceAccountRefreshAdapter;
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for VertexServiceAccountRefreshAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE
|
||||
}
|
||||
|
||||
fn supports(&self, transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
supports_local_vertex_service_account_auth_resolution(transport)
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
if !entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE)
|
||||
{
|
||||
return None;
|
||||
}
|
||||
if service_account_token_expires_soon(entry.expires_at_unix_secs) {
|
||||
return None;
|
||||
}
|
||||
let name = entry.auth_header_name.trim();
|
||||
let value = entry.auth_header_value.trim();
|
||||
if name.is_empty() || value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: name.to_ascii_lowercase(),
|
||||
value: value.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
None
|
||||
}
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool {
|
||||
supports_local_vertex_service_account_auth_resolution(transport)
|
||||
&& entry
|
||||
.and_then(|cached| self.resolve_cached(transport, cached))
|
||||
.is_none()
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
_entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
|
||||
let Some(auth_config) = resolve_local_vertex_service_account_auth_config(transport) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let now = aether_oauth::core::current_unix_secs();
|
||||
let assertion = build_vertex_service_account_assertion(&auth_config, now)?;
|
||||
let body = form_urlencoded::Serializer::new(String::new())
|
||||
.append_pair("grant_type", "urn:ietf:params:oauth:grant-type:jwt-bearer")
|
||||
.append_pair("assertion", &assertion)
|
||||
.finish();
|
||||
let response = executor
|
||||
.execute(
|
||||
VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
transport,
|
||||
&LocalOAuthHttpRequest {
|
||||
request_id: "vertex_ai:service-account-token",
|
||||
method: reqwest::Method::POST,
|
||||
url: auth_config.token_uri.clone(),
|
||||
headers: BTreeMap::from([(
|
||||
"content-type".to_string(),
|
||||
"application/x-www-form-urlencoded".to_string(),
|
||||
)]),
|
||||
json_body: None,
|
||||
body_bytes: Some(body.into_bytes()),
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if response.status_code != 200 {
|
||||
return Err(LocalOAuthRefreshError::HttpStatus {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
status_code: response.status_code,
|
||||
body_excerpt: body_excerpt(&response.body_text),
|
||||
});
|
||||
}
|
||||
let body_json: Value = serde_json::from_str(&response.body_text).map_err(|err| {
|
||||
LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
message: format!("vertex service account token response is not JSON: {err}"),
|
||||
}
|
||||
})?;
|
||||
let access_token = json_string(body_json.get("access_token")).ok_or_else(|| {
|
||||
LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
message: "vertex service account token response missing access_token".to_string(),
|
||||
}
|
||||
})?;
|
||||
let expires_in = body_json
|
||||
.get("expires_in")
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(3600);
|
||||
|
||||
Ok(Some(CachedOAuthEntry {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE.to_string(),
|
||||
auth_header_name: VERTEX_SERVICE_ACCOUNT_AUTH_HEADER.to_string(),
|
||||
auth_header_value: format!("Bearer {access_token}"),
|
||||
expires_at_unix_secs: Some(now.saturating_add(expires_in)),
|
||||
metadata: Some(json!({
|
||||
"project_id": auth_config.project_id,
|
||||
"client_email": auth_config.client_email,
|
||||
})),
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_vertex_service_account_assertion(
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<String, LocalOAuthRefreshError> {
|
||||
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#);
|
||||
let payload = URL_SAFE_NO_PAD.encode(
|
||||
serde_json::to_string(&json!({
|
||||
"iss": auth_config.client_email,
|
||||
"sub": auth_config.client_email,
|
||||
"scope": GOOGLE_CLOUD_PLATFORM_SCOPE,
|
||||
"aud": auth_config.token_uri,
|
||||
"iat": now_unix_secs,
|
||||
"exp": now_unix_secs.saturating_add(3600),
|
||||
}))
|
||||
.map_err(|err| LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
message: format!("vertex service account jwt payload encode failed: {err}"),
|
||||
})?,
|
||||
);
|
||||
let message = format!("{header}.{payload}");
|
||||
let private_key = decode_vertex_service_account_private_key(auth_config.private_key.as_str())?;
|
||||
let signing_key = SigningKey::<Sha256>::new(private_key);
|
||||
let signature = signing_key.sign(message.as_bytes());
|
||||
Ok(format!(
|
||||
"{message}.{}",
|
||||
URL_SAFE_NO_PAD.encode(signature.to_bytes())
|
||||
))
|
||||
}
|
||||
|
||||
fn decode_vertex_service_account_private_key(
|
||||
private_key_pem: &str,
|
||||
) -> Result<RsaPrivateKey, LocalOAuthRefreshError> {
|
||||
match RsaPrivateKey::from_pkcs8_pem(private_key_pem) {
|
||||
Ok(private_key) => Ok(private_key),
|
||||
Err(pkcs8_err) => RsaPrivateKey::from_pkcs1_pem(private_key_pem).map_err(|pkcs1_err| {
|
||||
LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type: VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
message: format!(
|
||||
"vertex service account private_key parse failed: pkcs8: {pkcs8_err}; pkcs1: {pkcs1_err}"
|
||||
),
|
||||
}
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn service_account_token_expires_soon(expires_at_unix_secs: Option<u64>) -> bool {
|
||||
expires_at_unix_secs
|
||||
.map(|expires_at_unix_secs| {
|
||||
aether_oauth::core::current_unix_secs()
|
||||
>= expires_at_unix_secs.saturating_sub(SERVICE_ACCOUNT_REFRESH_SKEW_SECS)
|
||||
})
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
fn body_excerpt(value: &str) -> String {
|
||||
value.chars().take(500).collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use rsa::pkcs1::{EncodeRsaPrivateKey, LineEnding};
|
||||
use rsa::rand_core::OsRng;
|
||||
use rsa::RsaPrivateKey;
|
||||
|
||||
use super::{
|
||||
decode_vertex_service_account_private_key, parse_vertex_service_account_auth_config,
|
||||
resolve_local_vertex_api_key_query_auth,
|
||||
supports_local_vertex_service_account_auth_resolution, VERTEX_API_KEY_QUERY_PARAM,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Vertex".to_string(),
|
||||
provider_type: "vertex_ai".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:generate_content".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://aiplatform.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: Some(vec!["gemini:generate_content".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "vertex-secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_query_auth_for_vertex_api_key_subset() {
|
||||
let auth = resolve_local_vertex_api_key_query_auth(&sample_transport())
|
||||
.expect("vertex api key query auth should resolve");
|
||||
assert_eq!(auth.name, VERTEX_API_KEY_QUERY_PARAM);
|
||||
assert_eq!(auth.value, "vertex-secret");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_non_api_key_transport() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
assert!(resolve_local_vertex_api_key_query_auth(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_vertex_auth_config_transport() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_auth_config = Some("{\"project_id\":\"demo-project\"}".to_string());
|
||||
assert!(resolve_local_vertex_api_key_query_auth(&transport).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_query_auth_for_custom_aiplatform_transport() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "custom".to_string();
|
||||
transport.endpoint.api_format = "gemini:generate_content".to_string();
|
||||
|
||||
let auth = resolve_local_vertex_api_key_query_auth(&transport)
|
||||
.expect("custom aiplatform transport should resolve");
|
||||
assert_eq!(auth.value, "vertex-secret");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_service_account_auth_config() {
|
||||
let config = parse_vertex_service_account_auth_config(Some(
|
||||
r#"{
|
||||
"client_email":"[email protected]",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project",
|
||||
"region":"global",
|
||||
"model_regions":{"gemini-2.0-flash":"us-central1"}
|
||||
}"#,
|
||||
))
|
||||
.expect("service account config should parse");
|
||||
|
||||
assert_eq!(config.client_email, "[email protected]");
|
||||
assert_eq!(config.project_id, "demo-project");
|
||||
assert_eq!(config.region.as_deref(), Some("global"));
|
||||
assert_eq!(
|
||||
config
|
||||
.model_regions
|
||||
.get("gemini-2.0-flash")
|
||||
.map(String::as_str),
|
||||
Some("us-central1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_service_account_auth_resolution() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"client_email":"[email protected]",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(supports_local_vertex_service_account_auth_resolution(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decodes_pkcs1_service_account_private_key() {
|
||||
let mut rng = OsRng;
|
||||
let private_key = RsaPrivateKey::new(&mut rng, 1024)
|
||||
.expect("test RSA private key should generate")
|
||||
.to_pkcs1_pem(LineEnding::LF)
|
||||
.expect("test RSA private key should encode as PKCS#1 PEM");
|
||||
|
||||
decode_vertex_service_account_private_key(private_key.as_str())
|
||||
.expect("PKCS#1 private key should decode");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
use url::Url;
|
||||
|
||||
use super::super::auth::resolve_local_auth_type_for_transport_format;
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
const VERTEX_AI_HOST: &str = "aiplatform.googleapis.com";
|
||||
|
||||
pub fn looks_like_vertex_ai_host(base_url: &str) -> bool {
|
||||
let trimmed = base_url.trim();
|
||||
if trimmed.is_empty() {
|
||||
return false;
|
||||
}
|
||||
|
||||
let Ok(parsed) = Url::parse(trimmed) else {
|
||||
return false;
|
||||
};
|
||||
let Some(host) = parsed
|
||||
.host_str()
|
||||
.map(|value| value.trim().to_ascii_lowercase())
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
|
||||
host == VERTEX_AI_HOST
|
||||
|| host.ends_with(&format!(".{VERTEX_AI_HOST}"))
|
||||
|| host.ends_with(&format!("-{VERTEX_AI_HOST}"))
|
||||
}
|
||||
|
||||
pub fn is_vertex_api_key_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
if is_vertex_provider_type(transport) {
|
||||
return resolve_local_auth_type_for_transport_format(transport)
|
||||
.eq_ignore_ascii_case("api_key");
|
||||
}
|
||||
|
||||
if !is_vertex_host_format_context(transport) {
|
||||
return false;
|
||||
}
|
||||
|
||||
resolve_local_auth_type_for_transport_format(transport).eq_ignore_ascii_case("api_key")
|
||||
}
|
||||
|
||||
pub fn is_vertex_service_account_transport_context(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
if !is_vertex_provider_type(transport) && !is_vertex_host_format_context(transport) {
|
||||
return false;
|
||||
}
|
||||
|
||||
matches!(
|
||||
resolve_local_auth_type_for_transport_format(transport).as_str(),
|
||||
"service_account" | "vertex_ai"
|
||||
)
|
||||
}
|
||||
|
||||
pub fn is_vertex_transport_context(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
is_vertex_api_key_transport_context(transport)
|
||||
|| is_vertex_service_account_transport_context(transport)
|
||||
}
|
||||
|
||||
pub fn uses_vertex_api_key_query_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> bool {
|
||||
is_vertex_api_key_transport_context(transport)
|
||||
&& provider_api_format
|
||||
.trim()
|
||||
.to_ascii_lowercase()
|
||||
.starts_with("gemini:")
|
||||
}
|
||||
|
||||
fn is_vertex_provider_type(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
|
||||
}
|
||||
|
||||
fn is_vertex_host_format_context(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
if !looks_like_vertex_ai_host(&transport.endpoint.base_url) {
|
||||
return false;
|
||||
}
|
||||
|
||||
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
||||
endpoint_api_format.starts_with("gemini:")
|
||||
|| endpoint_api_format.starts_with("claude:")
|
||||
|| (endpoint_api_format.starts_with("openai:")
|
||||
&& looks_like_vertex_openai_compat_base(&transport.endpoint.base_url))
|
||||
}
|
||||
|
||||
fn looks_like_vertex_openai_compat_base(base_url: &str) -> bool {
|
||||
let Ok(parsed) = Url::parse(base_url.trim()) else {
|
||||
return false;
|
||||
};
|
||||
parsed
|
||||
.path()
|
||||
.trim_end_matches('/')
|
||||
.ends_with("/endpoints/openapi")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
is_vertex_api_key_transport_context, is_vertex_service_account_transport_context,
|
||||
is_vertex_transport_context, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Vertex".to_string(),
|
||||
provider_type: "custom".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:generate_content".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("cli".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://aiplatform.googleapis.com".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: Some("/v1/publishers/google/models/{model}:{action}".to_string()),
|
||||
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: Some(vec!["gemini:generate_content".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "vertex-secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_vertex_host() {
|
||||
assert!(looks_like_vertex_ai_host(
|
||||
"https://aiplatform.googleapis.com"
|
||||
));
|
||||
assert!(looks_like_vertex_ai_host(
|
||||
"https://us-central1-aiplatform.googleapis.com"
|
||||
));
|
||||
assert!(!looks_like_vertex_ai_host("https://example.com"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_vertex_api_key_context_for_custom_aiplatform_transport() {
|
||||
assert!(is_vertex_api_key_transport_context(&sample_transport()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_non_api_key_custom_aiplatform_transport() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "bearer".to_string();
|
||||
assert!(!is_vertex_api_key_transport_context(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_vertex_service_account_context_for_fixed_provider() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "vertex_ai".to_string();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
|
||||
assert!(is_vertex_service_account_transport_context(&transport));
|
||||
assert!(is_vertex_transport_context(&transport));
|
||||
assert!(!is_vertex_api_key_transport_context(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_vertex_query_auth_usage_for_gemini_formats() {
|
||||
let transport = sample_transport();
|
||||
assert!(uses_vertex_api_key_query_auth(
|
||||
&transport,
|
||||
"gemini:generate_content"
|
||||
));
|
||||
assert!(!uses_vertex_api_key_query_auth(
|
||||
&transport,
|
||||
"claude:messages"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_vertex_service_account_context_for_openai_compat_endpoint_root() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "openai:chat".to_string();
|
||||
transport.endpoint.base_url =
|
||||
"https://aiplatform.googleapis.com/v1/projects/project-1/locations/global/endpoints/openapi"
|
||||
.to_string();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
|
||||
assert!(is_vertex_service_account_transport_context(&transport));
|
||||
assert!(is_vertex_transport_context(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn does_not_infer_vertex_context_for_generic_openai_format_on_aiplatform_root() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "openai:chat".to_string();
|
||||
transport.endpoint.base_url = "https://aiplatform.googleapis.com".to_string();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
|
||||
assert!(!is_vertex_service_account_transport_context(&transport));
|
||||
assert!(!is_vertex_transport_context(&transport));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
mod auth;
|
||||
mod context;
|
||||
mod policy;
|
||||
mod url;
|
||||
|
||||
pub use auth::{
|
||||
parse_vertex_service_account_auth_config, resolve_local_vertex_api_key_query_auth,
|
||||
resolve_local_vertex_service_account_auth_config,
|
||||
supports_local_vertex_service_account_auth_resolution, VertexApiKeyQueryAuth,
|
||||
VertexServiceAccountAuthConfig, VertexServiceAccountRefreshAdapter, VERTEX_API_KEY_QUERY_PARAM,
|
||||
VERTEX_SERVICE_ACCOUNT_AUTH_HEADER, VERTEX_SERVICE_ACCOUNT_PROVIDER_TYPE,
|
||||
};
|
||||
pub use context::{
|
||||
is_vertex_api_key_transport_context, is_vertex_service_account_transport_context,
|
||||
is_vertex_transport_context, looks_like_vertex_ai_host, uses_vertex_api_key_query_auth,
|
||||
};
|
||||
pub use policy::{
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network,
|
||||
supports_local_vertex_api_key_gemini_transport,
|
||||
supports_local_vertex_api_key_gemini_transport_with_network,
|
||||
supports_local_vertex_api_key_imagen_transport,
|
||||
supports_local_vertex_api_key_imagen_transport_with_network,
|
||||
supports_local_vertex_gemini_transport_with_network,
|
||||
};
|
||||
pub use url::{
|
||||
build_vertex_api_key_gemini_content_url, build_vertex_api_key_gemini_embedding_url,
|
||||
build_vertex_api_key_imagen_content_url, build_vertex_service_account_gemini_content_url,
|
||||
build_vertex_service_account_gemini_embedding_url, resolve_vertex_service_account_region,
|
||||
VERTEX_API_KEY_BASE_URL,
|
||||
};
|
||||
|
||||
pub const PROVIDER_TYPE: &str = "vertex_ai";
|
||||
@@ -0,0 +1,356 @@
|
||||
use super::super::snapshot::GatewayProviderTransportSnapshot;
|
||||
use super::super::{
|
||||
body_rules_are_locally_supported, header_rules_are_locally_supported,
|
||||
resolve_transport_profile, supports_local_oauth_request_auth_resolution,
|
||||
transport_profile_is_configured, transport_proxy_is_locally_supported,
|
||||
};
|
||||
use super::auth::{
|
||||
resolve_local_vertex_api_key_query_auth, supports_local_vertex_service_account_auth_resolution,
|
||||
};
|
||||
|
||||
fn is_vertex_transport_family(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
|
||||
|| super::looks_like_vertex_ai_host(&transport.endpoint.base_url)
|
||||
}
|
||||
|
||||
pub fn local_vertex_api_key_gemini_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network_impl(transport, true)
|
||||
}
|
||||
|
||||
pub fn local_vertex_gemini_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network_impl(transport, false)
|
||||
}
|
||||
|
||||
fn local_vertex_gemini_transport_unsupported_reason_with_network_impl(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
require_api_key: bool,
|
||||
) -> Option<&'static str> {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return if !transport.provider.is_active {
|
||||
Some("provider_inactive")
|
||||
} else if !transport.endpoint.is_active {
|
||||
Some("endpoint_inactive")
|
||||
} else {
|
||||
Some("key_inactive")
|
||||
};
|
||||
}
|
||||
let endpoint_api_format =
|
||||
aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
|
||||
if !matches!(
|
||||
endpoint_api_format.as_str(),
|
||||
"gemini:generate_content" | "gemini:embedding"
|
||||
) {
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if !is_vertex_transport_family(transport) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
|
||||
return Some("transport_header_rules_unsupported");
|
||||
}
|
||||
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
let has_api_key_auth = resolve_local_vertex_api_key_query_auth(transport).is_some();
|
||||
let has_service_account_auth = supports_local_vertex_service_account_auth_resolution(transport)
|
||||
&& supports_local_oauth_request_auth_resolution(transport);
|
||||
if require_api_key {
|
||||
if !has_api_key_auth {
|
||||
return Some("transport_auth_unavailable");
|
||||
}
|
||||
} else if !has_api_key_auth && !has_service_account_auth {
|
||||
return Some("transport_auth_unavailable");
|
||||
}
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return Some("transport_proxy_unsupported");
|
||||
}
|
||||
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return Some("transport_profile_unsupported");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_api_key_gemini_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
supports_local_vertex_api_key_same_format_transport(
|
||||
transport,
|
||||
&["gemini:generate_content"],
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_api_key_gemini_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(transport).is_none()
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_gemini_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network(transport).is_none()
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_api_key_imagen_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
supports_local_vertex_api_key_same_format_transport(
|
||||
transport,
|
||||
&["gemini:generate_content"],
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn supports_local_vertex_api_key_imagen_transport_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
supports_local_vertex_api_key_same_format_transport(
|
||||
transport,
|
||||
&["gemini:generate_content"],
|
||||
true,
|
||||
)
|
||||
}
|
||||
|
||||
fn supports_local_vertex_api_key_same_format_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_formats: &[&str],
|
||||
allow_network_passthrough: bool,
|
||||
) -> bool {
|
||||
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
|
||||
return false;
|
||||
}
|
||||
let endpoint_api_format =
|
||||
aether_ai_formats::normalize_api_format_alias(&transport.endpoint.api_format);
|
||||
if !api_formats
|
||||
.iter()
|
||||
.any(|api_format| endpoint_api_format.eq_ignore_ascii_case(api_format))
|
||||
{
|
||||
return false;
|
||||
}
|
||||
if !super::is_vertex_api_key_transport_context(transport) {
|
||||
return false;
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref())
|
||||
|| !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref())
|
||||
{
|
||||
return false;
|
||||
}
|
||||
if resolve_local_vertex_api_key_query_auth(transport).is_none() {
|
||||
return false;
|
||||
}
|
||||
|
||||
let has_custom_path = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.is_some_and(|value: &str| !value.trim().is_empty());
|
||||
if has_custom_path && !allow_network_passthrough {
|
||||
return false;
|
||||
}
|
||||
|
||||
if allow_network_passthrough {
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return false;
|
||||
}
|
||||
if transport_profile_is_configured(transport)
|
||||
&& resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return false;
|
||||
}
|
||||
} else if transport.provider.proxy.is_some()
|
||||
|| transport.endpoint.proxy.is_some()
|
||||
|| transport.key.proxy.is_some()
|
||||
|| transport_profile_is_configured(transport)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
true
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
use super::{
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network,
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network,
|
||||
supports_local_vertex_api_key_gemini_transport,
|
||||
supports_local_vertex_api_key_gemini_transport_with_network,
|
||||
supports_local_vertex_gemini_transport_with_network,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Vertex".to_string(),
|
||||
provider_type: "vertex_ai".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:generate_content".to_string(),
|
||||
api_family: Some("gemini".to_string()),
|
||||
endpoint_kind: Some("generate_content".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://aiplatform.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: Some(vec!["gemini:generate_content".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "vertex-secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_api_key_same_format_subset() {
|
||||
assert!(supports_local_vertex_api_key_gemini_transport(
|
||||
&sample_transport()
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_api_key_gemini_generate_content_subset() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "gemini:generate_content".to_string();
|
||||
assert!(supports_local_vertex_api_key_gemini_transport(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_custom_aiplatform_gemini_generate_content_subset() {
|
||||
let mut transport = sample_transport();
|
||||
transport.provider.provider_type = "custom".to_string();
|
||||
transport.endpoint.api_format = "gemini:generate_content".to_string();
|
||||
assert!(supports_local_vertex_api_key_gemini_transport(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_vertex_service_account_from_api_key_subset() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_auth_config = Some("{\"project_id\":\"demo-project\"}".to_string());
|
||||
assert!(!supports_local_vertex_api_key_gemini_transport(&transport));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_service_account_gemini_transport_with_network() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"client_email":"[email protected]",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(!supports_local_vertex_api_key_gemini_transport_with_network(&transport));
|
||||
assert!(supports_local_vertex_gemini_transport_with_network(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn supports_vertex_service_account_gemini_embedding_transport_with_network() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.api_format = "gemini:embedding".to_string();
|
||||
transport.endpoint.endpoint_kind = Some("embedding".to_string());
|
||||
transport.key.api_formats = Some(vec!["gemini:embedding".to_string()]);
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
r#"{
|
||||
"client_email":"[email protected]",
|
||||
"private_key":"TEST-PRIVATE-KEY",
|
||||
"project_id":"demo-project"
|
||||
}"#
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(supports_local_vertex_gemini_transport_with_network(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn allows_network_passthrough_for_custom_path_with_local_proxy_support() {
|
||||
let mut transport = sample_transport();
|
||||
transport.endpoint.custom_path =
|
||||
Some("/v1/publishers/google/models/gemini-2.5-pro:generateContent".to_string());
|
||||
transport.key.proxy = Some(json!({"url":"http://proxy.example:8080"}));
|
||||
transport.key.fingerprint = Some(json!({"transport_profile":"chrome_136"}));
|
||||
assert!(!supports_local_vertex_api_key_gemini_transport(&transport));
|
||||
assert!(supports_local_vertex_api_key_gemini_transport_with_network(
|
||||
&transport
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reports_auth_unavailable_when_vertex_api_key_query_auth_is_missing() {
|
||||
let mut transport = sample_transport();
|
||||
transport.key.auth_type = "service_account".to_string();
|
||||
transport.key.decrypted_api_key = "__placeholder__".to_string();
|
||||
|
||||
assert_eq!(
|
||||
local_vertex_api_key_gemini_transport_unsupported_reason_with_network(&transport),
|
||||
Some("transport_auth_unavailable")
|
||||
);
|
||||
assert_eq!(
|
||||
local_vertex_gemini_transport_unsupported_reason_with_network(&transport),
|
||||
Some("transport_auth_unavailable")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use url::form_urlencoded;
|
||||
|
||||
use super::super::url::build_passthrough_path_url;
|
||||
use super::auth::VertexServiceAccountAuthConfig;
|
||||
|
||||
pub const VERTEX_API_KEY_BASE_URL: &str = "https://aiplatform.googleapis.com";
|
||||
|
||||
pub fn build_vertex_api_key_gemini_content_url(
|
||||
model: &str,
|
||||
stream: bool,
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let action = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
build_vertex_api_key_google_model_url(model, action, stream, api_key, request_query)
|
||||
}
|
||||
|
||||
pub fn build_vertex_api_key_imagen_content_url(
|
||||
model: &str,
|
||||
stream: bool,
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let action = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
build_vertex_api_key_google_model_url(model, action, stream, api_key, request_query)
|
||||
}
|
||||
|
||||
pub fn build_vertex_api_key_gemini_embedding_url(
|
||||
model: &str,
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
build_vertex_api_key_google_model_url(model, "predict", false, api_key, request_query)
|
||||
}
|
||||
|
||||
pub fn build_vertex_service_account_gemini_content_url(
|
||||
model: &str,
|
||||
stream: bool,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let action = if stream {
|
||||
"streamGenerateContent"
|
||||
} else {
|
||||
"generateContent"
|
||||
};
|
||||
build_vertex_service_account_google_model_url(model, action, stream, auth_config, request_query)
|
||||
}
|
||||
|
||||
pub fn build_vertex_service_account_gemini_embedding_url(
|
||||
model: &str,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
build_vertex_service_account_google_model_url(
|
||||
model,
|
||||
"predict",
|
||||
false,
|
||||
auth_config,
|
||||
request_query,
|
||||
)
|
||||
}
|
||||
|
||||
fn build_vertex_api_key_google_model_url(
|
||||
model: &str,
|
||||
action: &str,
|
||||
stream: bool,
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let trimmed_model = model.trim();
|
||||
let trimmed_action = action.trim();
|
||||
let trimmed_api_key = api_key.trim();
|
||||
if trimmed_model.is_empty() || trimmed_action.is_empty() || trimmed_api_key.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let path = format!("/v1/publishers/google/models/{trimmed_model}:{trimmed_action}");
|
||||
let merged_query = build_vertex_api_key_query(trimmed_api_key, request_query, stream);
|
||||
build_passthrough_path_url(VERTEX_API_KEY_BASE_URL, &path, merged_query.as_deref(), &[])
|
||||
}
|
||||
|
||||
fn build_vertex_service_account_google_model_url(
|
||||
model: &str,
|
||||
action: &str,
|
||||
stream: bool,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
request_query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
let trimmed_model = model.trim();
|
||||
let trimmed_action = action.trim();
|
||||
let project_id = auth_config.project_id.trim();
|
||||
if trimmed_model.is_empty() || trimmed_action.is_empty() || project_id.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let region = resolve_vertex_service_account_region(trimmed_model, auth_config);
|
||||
let base_url = if region == "global" {
|
||||
VERTEX_API_KEY_BASE_URL.to_string()
|
||||
} else {
|
||||
format!("https://{region}-aiplatform.googleapis.com")
|
||||
};
|
||||
let path = format!(
|
||||
"/v1/projects/{project_id}/locations/{region}/publishers/google/models/{trimmed_model}:{trimmed_action}"
|
||||
);
|
||||
let merged_query = build_vertex_service_account_query(request_query, stream);
|
||||
build_passthrough_path_url(&base_url, &path, merged_query.as_deref(), &[])
|
||||
}
|
||||
|
||||
pub fn resolve_vertex_service_account_region(
|
||||
model: &str,
|
||||
auth_config: &VertexServiceAccountAuthConfig,
|
||||
) -> String {
|
||||
let trimmed_model = model.trim();
|
||||
if let Some(region) = auth_config
|
||||
.model_regions
|
||||
.get(trimmed_model)
|
||||
.map(String::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return region.to_string();
|
||||
}
|
||||
if let Some(region) = default_vertex_model_region(trimmed_model) {
|
||||
return region.to_string();
|
||||
}
|
||||
if let Some(region) = auth_config
|
||||
.region
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return region.to_string();
|
||||
}
|
||||
"global".to_string()
|
||||
}
|
||||
|
||||
fn default_vertex_model_region(model: &str) -> Option<&'static str> {
|
||||
if model.starts_with("gemini-3.") || model == "gemini-3-pro-image-preview" {
|
||||
return Some("global");
|
||||
}
|
||||
if matches!(
|
||||
model,
|
||||
"gemini-2.0-flash"
|
||||
| "gemini-2.0-flash-exp"
|
||||
| "gemini-2.0-flash-001"
|
||||
| "gemini-2.0-pro-exp"
|
||||
| "gemini-2.0-flash-exp-image-generation"
|
||||
| "gemini-1.5-pro"
|
||||
| "gemini-1.5-pro-001"
|
||||
| "gemini-1.5-pro-002"
|
||||
| "gemini-1.5-flash"
|
||||
| "gemini-1.5-flash-001"
|
||||
| "gemini-1.5-flash-002"
|
||||
| "imagen-3.0-generate-001"
|
||||
| "imagen-3.0-fast-generate-001"
|
||||
) {
|
||||
return Some("us-central1");
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn build_vertex_api_key_query(
|
||||
api_key: &str,
|
||||
request_query: Option<&str>,
|
||||
stream: bool,
|
||||
) -> Option<String> {
|
||||
let mut merged = BTreeMap::new();
|
||||
merge_query_string(&mut merged, request_query);
|
||||
merged.remove("beta");
|
||||
merged.insert("key".to_string(), api_key.to_string());
|
||||
if stream {
|
||||
merged
|
||||
.entry("alt".to_string())
|
||||
.or_insert_with(|| "sse".to_string());
|
||||
}
|
||||
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in merged {
|
||||
serializer.append_pair(&key, &value);
|
||||
}
|
||||
let query = serializer.finish();
|
||||
if query.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(query)
|
||||
}
|
||||
}
|
||||
|
||||
fn build_vertex_service_account_query(request_query: Option<&str>, stream: bool) -> Option<String> {
|
||||
let mut merged = BTreeMap::new();
|
||||
merge_query_string(&mut merged, request_query);
|
||||
merged.remove("beta");
|
||||
merged.remove("key");
|
||||
if stream {
|
||||
merged
|
||||
.entry("alt".to_string())
|
||||
.or_insert_with(|| "sse".to_string());
|
||||
}
|
||||
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
for (key, value) in merged {
|
||||
serializer.append_pair(&key, &value);
|
||||
}
|
||||
let query = serializer.finish();
|
||||
if query.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(query)
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_query_string(out: &mut BTreeMap<String, String>, query: Option<&str>) {
|
||||
let Some(query) = query.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||
return;
|
||||
};
|
||||
|
||||
for (key, value) in form_urlencoded::parse(query.as_bytes()) {
|
||||
out.insert(key.into_owned(), value.into_owned());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::{
|
||||
build_vertex_api_key_gemini_content_url, build_vertex_api_key_imagen_content_url,
|
||||
build_vertex_service_account_gemini_content_url,
|
||||
};
|
||||
use crate::vertex::VertexServiceAccountAuthConfig;
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_gemini_api_key_stream_url() {
|
||||
assert_eq!(
|
||||
build_vertex_api_key_gemini_content_url(
|
||||
"gemini-2.5-pro",
|
||||
true,
|
||||
"vertex-secret",
|
||||
Some("foo=bar&beta=v1")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://aiplatform.googleapis.com/v1/publishers/google/models/gemini-2.5-pro:streamGenerateContent?alt=sse&foo=bar&key=vertex-secret"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_imagen_api_key_sync_url() {
|
||||
assert_eq!(
|
||||
build_vertex_api_key_imagen_content_url(
|
||||
"imagen-3.0-generate-001",
|
||||
false,
|
||||
"vertex-secret",
|
||||
Some("view=full")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://aiplatform.googleapis.com/v1/publishers/google/models/imagen-3.0-generate-001:generateContent?key=vertex-secret&view=full"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_service_account_gemini_sync_url() {
|
||||
let auth_config = VertexServiceAccountAuthConfig {
|
||||
client_email: "[email protected]".to_string(),
|
||||
private_key: "not-used".to_string(),
|
||||
project_id: "demo-project".to_string(),
|
||||
token_uri: "https://oauth2.googleapis.com/token".to_string(),
|
||||
region: None,
|
||||
model_regions: BTreeMap::new(),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
build_vertex_service_account_gemini_content_url(
|
||||
"gemini-3.1-pro-preview",
|
||||
false,
|
||||
&auth_config,
|
||||
Some("foo=bar&beta=1&key=client-key")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://aiplatform.googleapis.com/v1/projects/demo-project/locations/global/publishers/google/models/gemini-3.1-pro-preview:generateContent?foo=bar"
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_vertex_service_account_gemini_stream_url_with_model_region_override() {
|
||||
let auth_config = VertexServiceAccountAuthConfig {
|
||||
client_email: "[email protected]".to_string(),
|
||||
private_key: "not-used".to_string(),
|
||||
project_id: "demo-project".to_string(),
|
||||
token_uri: "https://oauth2.googleapis.com/token".to_string(),
|
||||
region: Some("global".to_string()),
|
||||
model_regions: BTreeMap::from([(
|
||||
"gemini-2.0-flash".to_string(),
|
||||
"us-central1".to_string(),
|
||||
)]),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
build_vertex_service_account_gemini_content_url(
|
||||
"gemini-2.0-flash",
|
||||
true,
|
||||
&auth_config,
|
||||
Some("foo=bar")
|
||||
)
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://us-central1-aiplatform.googleapis.com/v1/projects/demo-project/locations/us-central1/publishers/google/models/gemini-2.0-flash:streamGenerateContent?alt=sse&foo=bar"
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,520 @@
|
||||
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::{
|
||||
build_passthrough_headers_with_auth, resolve_local_gemini_auth,
|
||||
resolve_local_openai_bearer_auth,
|
||||
};
|
||||
use super::network::{resolve_transport_execution_timeouts, resolve_transport_profile};
|
||||
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_with_request_headers, apply_local_header_rules_with_request_headers,
|
||||
};
|
||||
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 {
|
||||
async fn read_video_task_provider_transport_snapshot(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
endpoint_id: &str,
|
||||
key_id: &str,
|
||||
) -> Result<Option<GatewayProviderTransportSnapshot>, String>;
|
||||
}
|
||||
|
||||
pub fn resolve_local_video_task_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
model_name: Option<String>,
|
||||
) -> Option<LocalVideoTaskTransport> {
|
||||
let api_format = api_format.trim();
|
||||
let (auth_header, auth_value) = match api_format {
|
||||
"openai:video" => {
|
||||
if !supports_local_standard_transport(transport, api_format) {
|
||||
return None;
|
||||
}
|
||||
resolve_local_openai_bearer_auth(transport)?
|
||||
}
|
||||
"gemini:video" => {
|
||||
if !supports_local_gemini_transport(transport, api_format) {
|
||||
return None;
|
||||
}
|
||||
resolve_local_gemini_auth(transport)?
|
||||
}
|
||||
_ => return None,
|
||||
};
|
||||
|
||||
Some(LocalVideoTaskTransport::from_bridge_input(
|
||||
LocalVideoTaskTransportBridgeInput {
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
provider_name: Some(transport.provider.name.clone()),
|
||||
provider_id: transport.provider.id.clone(),
|
||||
endpoint_id: transport.endpoint.id.clone(),
|
||||
key_id: transport.key.id.clone(),
|
||||
auth_header,
|
||||
auth_value,
|
||||
content_type: Some("application/json".to_string()),
|
||||
model_name,
|
||||
proxy: None,
|
||||
transport_profile: resolve_transport_profile(transport),
|
||||
timeouts: resolve_transport_execution_timeouts(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>,
|
||||
request_headers: Option<&http::HeaderMap>,
|
||||
) -> 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_with_request_headers(
|
||||
&mut provider_request_body,
|
||||
body_rules,
|
||||
Some(body_json),
|
||||
request_headers,
|
||||
) {
|
||||
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,
|
||||
openai_video_api_root_request_path(request_path),
|
||||
request_query,
|
||||
&[],
|
||||
),
|
||||
ProviderVideoCreateFamily::Gemini => build_gemini_video_predict_long_running_url(
|
||||
&transport.endpoint.base_url,
|
||||
mapped_model,
|
||||
request_query,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
fn openai_video_api_root_request_path(request_path: &str) -> &str {
|
||||
if request_path.starts_with("/v1/") {
|
||||
&request_path[3..]
|
||||
} else {
|
||||
request_path
|
||||
}
|
||||
}
|
||||
|
||||
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_with_request_headers(
|
||||
&mut provider_request_headers,
|
||||
input.header_rules,
|
||||
&[input.auth_header, "content-type"],
|
||||
input.provider_request_body,
|
||||
Some(input.original_request_body),
|
||||
Some(input.headers),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
Some(provider_request_headers)
|
||||
}
|
||||
|
||||
pub async fn reconstruct_local_video_task_snapshot(
|
||||
lookup: &dyn VideoTaskTransportSnapshotLookup,
|
||||
task: &StoredVideoTask,
|
||||
) -> Result<Option<LocalVideoTaskSnapshot>, String> {
|
||||
let provider_api_format = task
|
||||
.provider_api_format
|
||||
.as_deref()
|
||||
.unwrap_or_default()
|
||||
.trim();
|
||||
if !matches!(provider_api_format, "openai:video" | "gemini:video") {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(provider_id) = task.provider_id.as_deref() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(endpoint_id) = task.endpoint_id.as_deref() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(key_id) = task.key_id.as_deref() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some(transport) = lookup
|
||||
.read_video_task_provider_transport_snapshot(provider_id, endpoint_id, key_id)
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some(local_transport) =
|
||||
resolve_local_video_task_transport(&transport, provider_api_format, task.model.clone())
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
Ok(LocalVideoTaskSnapshot::from_stored_task_with_transport(
|
||||
task,
|
||||
local_transport,
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskStatus};
|
||||
use aether_video_tasks_core::LocalVideoTaskSnapshot;
|
||||
use async_trait::async_trait;
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
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,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
fn sample_transport(api_format: &str, auth_type: &str) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Provider One".to_string(),
|
||||
provider_type: "openai".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: Some(30.0),
|
||||
stream_first_byte_timeout_secs: Some(5.0),
|
||||
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://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: auth_type.to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_stored_video_task() -> StoredVideoTask {
|
||||
StoredVideoTask {
|
||||
id: "task-1".to_string(),
|
||||
short_id: Some("short-1".to_string()),
|
||||
request_id: "request-1".to_string(),
|
||||
user_id: Some("user-1".to_string()),
|
||||
api_key_id: Some("api-key-1".to_string()),
|
||||
username: Some("user".to_string()),
|
||||
api_key_name: Some("key".to_string()),
|
||||
external_task_id: Some("upstream-task-1".to_string()),
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
endpoint_id: Some("endpoint-1".to_string()),
|
||||
key_id: Some("key-1".to_string()),
|
||||
client_api_format: Some("openai:video".to_string()),
|
||||
provider_api_format: Some("openai:video".to_string()),
|
||||
format_converted: false,
|
||||
model: Some("sora".to_string()),
|
||||
prompt: Some("generate".to_string()),
|
||||
original_request_body: Some(json!({"prompt": "generate"})),
|
||||
duration_seconds: None,
|
||||
resolution: None,
|
||||
aspect_ratio: None,
|
||||
size: Some("1024x1024".to_string()),
|
||||
status: VideoTaskStatus::Submitted,
|
||||
progress_percent: 0,
|
||||
progress_message: None,
|
||||
retry_count: 0,
|
||||
poll_interval_seconds: 10,
|
||||
next_poll_at_unix_secs: None,
|
||||
poll_count: 0,
|
||||
max_poll_count: 360,
|
||||
created_at_unix_ms: 1,
|
||||
submitted_at_unix_secs: Some(1),
|
||||
completed_at_unix_secs: None,
|
||||
updated_at_unix_secs: 1,
|
||||
error_code: None,
|
||||
error_message: None,
|
||||
video_url: None,
|
||||
request_metadata: None,
|
||||
}
|
||||
}
|
||||
|
||||
struct TestLookup(Option<GatewayProviderTransportSnapshot>);
|
||||
|
||||
#[async_trait]
|
||||
impl VideoTaskTransportSnapshotLookup for TestLookup {
|
||||
async fn read_video_task_provider_transport_snapshot(
|
||||
&self,
|
||||
_provider_id: &str,
|
||||
_endpoint_id: &str,
|
||||
_key_id: &str,
|
||||
) -> Result<Option<GatewayProviderTransportSnapshot>, String> {
|
||||
Ok(self.0.clone())
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_openai_video_transport() {
|
||||
let transport = resolve_local_video_task_transport(
|
||||
&sample_transport("openai:video", "api_key"),
|
||||
"openai:video",
|
||||
Some("sora".to_string()),
|
||||
)
|
||||
.expect("transport");
|
||||
|
||||
assert_eq!(
|
||||
transport.headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer secret")
|
||||
);
|
||||
assert_eq!(transport.model_name.as_deref(), Some("sora"));
|
||||
assert_eq!(transport.provider_id, "provider-1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_gemini_video_transport() {
|
||||
let transport = resolve_local_video_task_transport(
|
||||
&sample_transport("gemini:video", "api_key"),
|
||||
"gemini:video",
|
||||
Some("veo".to_string()),
|
||||
)
|
||||
.expect("transport");
|
||||
|
||||
assert_eq!(
|
||||
transport.headers.get("x-goog-api-key").map(String::as_str),
|
||||
Some("secret")
|
||||
);
|
||||
assert_eq!(transport.model_name.as_deref(), Some("veo"));
|
||||
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,
|
||||
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_openai_video_create_url_from_api_root_base() {
|
||||
let mut transport = sample_transport("openai:video", "bearer");
|
||||
transport.endpoint.base_url = "https://api.openai.example/v1".to_string();
|
||||
let url = build_video_create_upstream_url(
|
||||
&transport,
|
||||
"/v1/videos",
|
||||
Some("trace=1"),
|
||||
"sora-upstream",
|
||||
ProviderVideoCreateFamily::OpenAi,
|
||||
)
|
||||
.expect("url should build");
|
||||
|
||||
assert_eq!(url, "https://api.openai.example/v1/videos?trace=1");
|
||||
}
|
||||
|
||||
#[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");
|
||||
assert!(resolve_local_video_task_transport(&transport, "openai:video", None).is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reconstructs_openai_video_snapshot_via_lookup_trait() {
|
||||
let lookup = TestLookup(Some(sample_transport("openai:video", "bearer")));
|
||||
let snapshot = reconstruct_local_video_task_snapshot(&lookup, &sample_stored_video_task())
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("snapshot");
|
||||
|
||||
match snapshot {
|
||||
LocalVideoTaskSnapshot::OpenAi(seed) => {
|
||||
assert_eq!(seed.transport.provider_id, "provider-1");
|
||||
assert_eq!(seed.transport.model_name.as_deref(), Some("sora"));
|
||||
}
|
||||
LocalVideoTaskSnapshot::Gemini(_) => panic!("expected openai snapshot"),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,582 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::{json, Value};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::rules::{
|
||||
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
|
||||
body_rules_are_locally_supported, header_rules_are_locally_supported,
|
||||
};
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
use crate::url::build_passthrough_path_url;
|
||||
use crate::{
|
||||
resolve_transport_profile, should_skip_upstream_passthrough_header,
|
||||
supports_local_oauth_request_auth_resolution, transport_profile_is_configured,
|
||||
transport_proxy_is_locally_supported,
|
||||
};
|
||||
|
||||
pub mod cascade;
|
||||
pub mod models;
|
||||
pub mod proto;
|
||||
|
||||
pub const PROVIDER_TYPE: &str = "windsurf";
|
||||
pub const WINDSURF_ENVELOPE_NAME: &str = "windsurf:GetChatMessage";
|
||||
pub const GET_CHAT_MESSAGE_PATH: &str = "/exa.api_server_pb.ApiServerService/GetChatMessage";
|
||||
const DEFAULT_IDE_VERSION: &str = "1.9600.41";
|
||||
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
|
||||
|
||||
pub fn is_windsurf_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(PROVIDER_TYPE)
|
||||
}
|
||||
|
||||
pub fn local_windsurf_request_transport_unsupported_reason_with_network(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
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");
|
||||
}
|
||||
if !is_windsurf_provider_transport(transport) {
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
if !transport
|
||||
.endpoint
|
||||
.api_format
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("openai:chat")
|
||||
{
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if !header_rules_are_locally_supported(transport.endpoint.header_rules.as_ref()) {
|
||||
return Some("transport_header_rules_unsupported");
|
||||
}
|
||||
if !body_rules_are_locally_supported(transport.endpoint.body_rules.as_ref()) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
if transport.key.decrypted_auth_config.is_some()
|
||||
&& !supports_local_oauth_request_auth_resolution(transport)
|
||||
&& !supports_local_windsurf_request_auth_resolution(transport)
|
||||
{
|
||||
return Some("transport_oauth_resolution_unsupported");
|
||||
}
|
||||
if !transport_proxy_is_locally_supported(transport) {
|
||||
return Some("transport_proxy_unsupported");
|
||||
}
|
||||
if transport_profile_is_configured(transport) && resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return Some("transport_profile_unsupported");
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub fn supports_local_windsurf_request_auth_resolution(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
resolve_windsurf_cascade_auth(transport).is_some()
|
||||
}
|
||||
|
||||
pub fn resolve_windsurf_cascade_auth(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<(String, String)> {
|
||||
if !is_windsurf_provider_transport(transport) {
|
||||
return None;
|
||||
}
|
||||
let auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
|
||||
if !matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer") {
|
||||
return None;
|
||||
}
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if secret.is_empty() || secret == PLACEHOLDER_API_KEY {
|
||||
return None;
|
||||
}
|
||||
Some(("authorization".to_string(), format!("Bearer {secret}")))
|
||||
}
|
||||
|
||||
pub fn build_windsurf_cascade_upstream_url(
|
||||
upstream_base_url: &str,
|
||||
query: Option<&str>,
|
||||
) -> Option<String> {
|
||||
build_passthrough_path_url(upstream_base_url, GET_CHAT_MESSAGE_PATH, query, &[])
|
||||
}
|
||||
|
||||
pub fn build_windsurf_cascade_request_body(
|
||||
body_json: &Value,
|
||||
mapped_model: &str,
|
||||
auth_value: &str,
|
||||
body_rules: Option<&Value>,
|
||||
request_headers: Option<&http::HeaderMap>,
|
||||
upstream_is_stream: bool,
|
||||
) -> Option<Value> {
|
||||
let mapped_model = mapped_model.trim();
|
||||
if mapped_model.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let messages = body_json.get("messages")?.as_array()?.clone();
|
||||
if messages.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let conversation_id =
|
||||
extract_conversation_id(body_json).unwrap_or_else(|| Uuid::new_v4().to_string());
|
||||
let message_text =
|
||||
latest_message_snapshot_text(&messages).unwrap_or_else(|| "Continue.".to_string());
|
||||
let mut provider_request_body = json!({
|
||||
"metadata": windsurf_metadata_from_auth(auth_value),
|
||||
"model": mapped_model,
|
||||
"modelName": mapped_model,
|
||||
"stream": upstream_is_stream,
|
||||
"conversationId": conversation_id,
|
||||
"message": message_text,
|
||||
"messages": messages,
|
||||
});
|
||||
|
||||
if let Some(max_tokens) = body_json
|
||||
.get("max_tokens")
|
||||
.or_else(|| body_json.get("maxTokens"))
|
||||
{
|
||||
provider_request_body
|
||||
.as_object_mut()?
|
||||
.insert("maxTokens".to_string(), max_tokens.clone());
|
||||
}
|
||||
if let Some(temperature) = body_json.get("temperature") {
|
||||
provider_request_body
|
||||
.as_object_mut()?
|
||||
.insert("temperature".to_string(), temperature.clone());
|
||||
}
|
||||
if let Some(top_p) = body_json.get("top_p").or_else(|| body_json.get("topP")) {
|
||||
provider_request_body
|
||||
.as_object_mut()?
|
||||
.insert("topP".to_string(), top_p.clone());
|
||||
}
|
||||
for field in [
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"toolChoice",
|
||||
"parallel_tool_calls",
|
||||
"response_format",
|
||||
] {
|
||||
if let Some(value) = body_json.get(field) {
|
||||
provider_request_body
|
||||
.as_object_mut()?
|
||||
.insert(field.to_string(), value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
if !apply_local_body_rules_with_request_headers(
|
||||
&mut provider_request_body,
|
||||
body_rules,
|
||||
Some(body_json),
|
||||
request_headers,
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
Some(provider_request_body)
|
||||
}
|
||||
|
||||
pub fn build_windsurf_cascade_headers(
|
||||
headers: &http::HeaderMap,
|
||||
provider_request_body: &Value,
|
||||
original_request_body: &Value,
|
||||
header_rules: Option<&Value>,
|
||||
auth_header: &str,
|
||||
auth_value: &str,
|
||||
_upstream_is_stream: bool,
|
||||
) -> Option<BTreeMap<String, String>> {
|
||||
let mut out = BTreeMap::new();
|
||||
for (name, value) in headers {
|
||||
let Ok(value) = value.to_str() else {
|
||||
continue;
|
||||
};
|
||||
let key = name.as_str().to_ascii_lowercase();
|
||||
if should_skip_upstream_passthrough_header(&key) {
|
||||
continue;
|
||||
}
|
||||
let value = value.trim();
|
||||
if !value.is_empty() {
|
||||
out.insert(key, value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
let auth_header = auth_header.trim().to_ascii_lowercase();
|
||||
if !apply_local_header_rules_with_request_headers(
|
||||
&mut out,
|
||||
header_rules,
|
||||
&[
|
||||
auth_header.as_str(),
|
||||
"content-type",
|
||||
"connect-protocol-version",
|
||||
],
|
||||
provider_request_body,
|
||||
Some(original_request_body),
|
||||
Some(headers),
|
||||
) {
|
||||
return None;
|
||||
}
|
||||
|
||||
out.insert(
|
||||
"content-type".to_string(),
|
||||
"application/connect+json".to_string(),
|
||||
);
|
||||
out.insert("connect-protocol-version".to_string(), "1".to_string());
|
||||
out.insert(
|
||||
"user-agent".to_string(),
|
||||
format!("windsurf/{DEFAULT_IDE_VERSION}"),
|
||||
);
|
||||
out.insert("accept".to_string(), "application/connect+json".to_string());
|
||||
if !auth_header.is_empty() {
|
||||
out.insert(auth_header, auth_value.trim().to_string());
|
||||
}
|
||||
out.remove("content-length");
|
||||
Some(out)
|
||||
}
|
||||
|
||||
fn windsurf_metadata_from_auth(auth_value: &str) -> Value {
|
||||
json!({
|
||||
"apiKey": auth_secret_from_header_value(auth_value),
|
||||
"ideName": "windsurf",
|
||||
"ideVersion": DEFAULT_IDE_VERSION,
|
||||
"extensionName": "windsurf",
|
||||
"extensionVersion": DEFAULT_IDE_VERSION,
|
||||
"locale": "en",
|
||||
})
|
||||
}
|
||||
|
||||
fn auth_secret_from_header_value(auth_value: &str) -> String {
|
||||
let value = auth_value.trim();
|
||||
value
|
||||
.strip_prefix("Bearer ")
|
||||
.or_else(|| value.strip_prefix("bearer "))
|
||||
.unwrap_or(value)
|
||||
.trim()
|
||||
.to_string()
|
||||
}
|
||||
|
||||
fn extract_conversation_id(body_json: &Value) -> Option<String> {
|
||||
let object = body_json.as_object()?;
|
||||
string_value(object.get("conversation_id"))
|
||||
.or_else(|| string_value(object.get("conversationId")))
|
||||
.or_else(|| string_value(object.get("session_id")))
|
||||
.or_else(|| string_value(object.get("sessionId")))
|
||||
.or_else(|| {
|
||||
object
|
||||
.get("metadata")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| {
|
||||
string_value(metadata.get("conversation_id"))
|
||||
.or_else(|| string_value(metadata.get("conversationId")))
|
||||
.or_else(|| string_value(metadata.get("session_id")))
|
||||
.or_else(|| string_value(metadata.get("sessionId")))
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn string_value(value: Option<&Value>) -> Option<String> {
|
||||
value
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn latest_message_snapshot_text(messages: &[Value]) -> Option<String> {
|
||||
messages
|
||||
.iter()
|
||||
.rev()
|
||||
.filter_map(Value::as_object)
|
||||
.find_map(|message| {
|
||||
let role = message.get("role").and_then(Value::as_str)?;
|
||||
match role {
|
||||
"user" => openai_content_to_text(message.get("content"))
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty()),
|
||||
"tool" => {
|
||||
let content = openai_content_to_text(message.get("content"))
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())?;
|
||||
let tool_call_id = message
|
||||
.get("tool_call_id")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("unknown");
|
||||
Some(format!(
|
||||
"<tool_result tool_call_id=\"{}\">\n{content}\n</tool_result>",
|
||||
escape_xml_attr(tool_call_id)
|
||||
))
|
||||
}
|
||||
"assistant" => openai_content_to_text(message.get("content"))
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty()),
|
||||
_ => None,
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn openai_content_to_text(value: Option<&Value>) -> Option<String> {
|
||||
match value? {
|
||||
Value::String(text) => Some(text.clone()),
|
||||
Value::Array(items) => {
|
||||
let parts = items
|
||||
.iter()
|
||||
.filter_map(|item| {
|
||||
item.as_object()
|
||||
.and_then(|object| object.get("text"))
|
||||
.and_then(Value::as_str)
|
||||
})
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.collect::<Vec<_>>();
|
||||
(!parts.is_empty()).then(|| parts.join("\n"))
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn escape_xml_attr(value: &str) -> String {
|
||||
value
|
||||
.replace('&', "&")
|
||||
.replace('"', """)
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use http::HeaderMap;
|
||||
use serde_json::json;
|
||||
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
|
||||
use super::{
|
||||
build_windsurf_cascade_headers, build_windsurf_cascade_request_body,
|
||||
build_windsurf_cascade_upstream_url,
|
||||
local_windsurf_request_transport_unsupported_reason_with_network,
|
||||
resolve_windsurf_cascade_auth, GET_CHAT_MESSAGE_PATH,
|
||||
};
|
||||
|
||||
fn sample_windsurf_transport(auth_type: &str) -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: "provider-windsurf".to_string(),
|
||||
name: "Windsurf".to_string(),
|
||||
provider_type: "windsurf".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-windsurf-chat".to_string(),
|
||||
provider_id: "provider-windsurf".to_string(),
|
||||
api_format: "openai:chat".to_string(),
|
||||
api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://server.codeium.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-windsurf".to_string(),
|
||||
provider_id: "provider-windsurf".to_string(),
|
||||
name: "[email protected]".to_string(),
|
||||
auth_type: auth_type.to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
decrypted_api_key: "devin-session-token$abc".to_string(),
|
||||
decrypted_auth_config: Some(r#"{"provider_type":"windsurf"}"#.to_string()),
|
||||
upstream_metadata: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_windsurf_cascade_url() {
|
||||
assert_eq!(
|
||||
build_windsurf_cascade_upstream_url("https://server.codeium.com", Some("debug=1"))
|
||||
.as_deref(),
|
||||
Some(
|
||||
"https://server.codeium.com/exa.api_server_pb.ApiServerService/GetChatMessage?debug=1"
|
||||
)
|
||||
);
|
||||
assert!(GET_CHAT_MESSAGE_PATH.ends_with("/GetChatMessage"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_cascade_request_body_with_metadata_and_messages() {
|
||||
let body = build_windsurf_cascade_request_body(
|
||||
&json!({
|
||||
"model": "gpt-5",
|
||||
"conversation_id": "conv-1",
|
||||
"messages": [
|
||||
{"role": "system", "content": "brief"},
|
||||
{"role": "user", "content": [{"type": "text", "text": "hello"}]}
|
||||
],
|
||||
"max_tokens": 128
|
||||
}),
|
||||
"windsurf-model",
|
||||
"Bearer devin-session-token$abc",
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("body should build");
|
||||
|
||||
assert_eq!(body["metadata"]["apiKey"], json!("devin-session-token$abc"));
|
||||
assert_eq!(body["modelName"], json!("windsurf-model"));
|
||||
assert_eq!(body["stream"], json!(true));
|
||||
assert_eq!(body["conversationId"], json!("conv-1"));
|
||||
assert_eq!(body["message"], json!("hello"));
|
||||
assert_eq!(body["maxTokens"], json!(128));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_openai_tool_fields_for_native_windsurf_runtime() {
|
||||
let body = build_windsurf_cascade_request_body(
|
||||
&json!({
|
||||
"model": "gpt-5-5-low",
|
||||
"messages": [
|
||||
{"role": "user", "content": "read Cargo.toml"}
|
||||
],
|
||||
"tools": [{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "Read",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"file_path": {"type": "string"}}
|
||||
}
|
||||
}
|
||||
}],
|
||||
"tool_choice": "required",
|
||||
"parallel_tool_calls": false,
|
||||
"response_format": {"type": "json_object"}
|
||||
}),
|
||||
"gpt-5-5-low",
|
||||
"Bearer devin-session-token$abc",
|
||||
None,
|
||||
None,
|
||||
false,
|
||||
)
|
||||
.expect("body should build");
|
||||
|
||||
assert_eq!(body["tools"][0]["function"]["name"], json!("Read"));
|
||||
assert_eq!(body["tool_choice"], json!("required"));
|
||||
assert_eq!(body["parallel_tool_calls"], json!(false));
|
||||
assert_eq!(body["response_format"]["type"], json!("json_object"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_cascade_request_body_message_snapshot_from_latest_tool_result() {
|
||||
let body = build_windsurf_cascade_request_body(
|
||||
&json!({
|
||||
"model": "gpt-5-5-low",
|
||||
"messages": [
|
||||
{"role": "user", "content": "read Cargo.toml"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": null,
|
||||
"tool_calls": [{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "Read", "arguments": "{\"file_path\":\"Cargo.toml\"}"}
|
||||
}]
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "workspace Cargo.toml content"}
|
||||
]
|
||||
}),
|
||||
"gpt-5-5-low",
|
||||
"Bearer devin-session-token$abc",
|
||||
None,
|
||||
None,
|
||||
false,
|
||||
)
|
||||
.expect("body should build");
|
||||
|
||||
assert!(body["message"]
|
||||
.as_str()
|
||||
.expect("message should be a string")
|
||||
.contains(r#"<tool_result tool_call_id="call_1">"#));
|
||||
assert!(body["message"]
|
||||
.as_str()
|
||||
.expect("message should be a string")
|
||||
.contains("workspace Cargo.toml content"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_cascade_headers_with_connect_protocol_and_auth() {
|
||||
let headers = build_windsurf_cascade_headers(
|
||||
&HeaderMap::new(),
|
||||
&json!({"metadata": {"apiKey": "secret"}}),
|
||||
&json!({"messages": []}),
|
||||
None,
|
||||
"authorization",
|
||||
"Bearer secret",
|
||||
false,
|
||||
)
|
||||
.expect("headers should build");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("connect-protocol-version").map(String::as_str),
|
||||
Some("1")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("authorization").map(String::as_str),
|
||||
Some("Bearer secret")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("content-type").map(String::as_str),
|
||||
Some("application/connect+json")
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("accept").map(String::as_str),
|
||||
Some("application/connect+json")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_windsurf_transport_resolves_direct_bearer_auth() {
|
||||
let transport = sample_windsurf_transport("oauth");
|
||||
|
||||
assert_eq!(
|
||||
local_windsurf_request_transport_unsupported_reason_with_network(&transport),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_windsurf_cascade_auth(&transport),
|
||||
Some((
|
||||
"authorization".to_string(),
|
||||
"Bearer devin-session-token$abc".to_string()
|
||||
))
|
||||
);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,384 @@
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct WindsurfModel {
|
||||
pub canonical_name: &'static str,
|
||||
pub enum_value: u32,
|
||||
pub model_uid: Option<&'static str>,
|
||||
pub credit_multiplier: f32,
|
||||
pub provider: &'static str,
|
||||
pub deprecated: bool,
|
||||
}
|
||||
|
||||
#[rustfmt::skip]
|
||||
const MODELS: &[WindsurfModel] = &[
|
||||
WindsurfModel { canonical_name: "claude-3.5-sonnet", enum_value: 166, model_uid: None, credit_multiplier: 2.0, provider: "anthropic", deprecated: true },
|
||||
WindsurfModel { canonical_name: "claude-3.7-sonnet", enum_value: 226, model_uid: None, credit_multiplier: 2.0, provider: "anthropic", deprecated: true },
|
||||
WindsurfModel { canonical_name: "claude-3.7-sonnet-thinking", enum_value: 227, model_uid: None, credit_multiplier: 3.0, provider: "anthropic", deprecated: true },
|
||||
WindsurfModel { canonical_name: "claude-4-sonnet", enum_value: 281, model_uid: Some("MODEL_CLAUDE_4_SONNET"), credit_multiplier: 2.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4-sonnet-thinking", enum_value: 282, model_uid: Some("MODEL_CLAUDE_4_SONNET_THINKING"), credit_multiplier: 3.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4-opus", enum_value: 290, model_uid: Some("MODEL_CLAUDE_4_OPUS"), credit_multiplier: 4.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4-opus-thinking", enum_value: 291, model_uid: Some("MODEL_CLAUDE_4_OPUS_THINKING"), credit_multiplier: 5.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4.1-opus", enum_value: 328, model_uid: Some("MODEL_CLAUDE_4_1_OPUS"), credit_multiplier: 4.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4.1-opus-thinking", enum_value: 329, model_uid: Some("MODEL_CLAUDE_4_1_OPUS_THINKING"), credit_multiplier: 5.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4.5-haiku", enum_value: 0, model_uid: Some("MODEL_PRIVATE_11"), credit_multiplier: 1.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4.5-sonnet", enum_value: 353, model_uid: Some("MODEL_PRIVATE_2"), credit_multiplier: 2.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4.5-sonnet-thinking", enum_value: 354, model_uid: Some("MODEL_PRIVATE_3"), credit_multiplier: 3.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4.5-opus", enum_value: 391, model_uid: Some("MODEL_CLAUDE_4_5_OPUS"), credit_multiplier: 4.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-4.5-opus-thinking", enum_value: 392, model_uid: Some("MODEL_CLAUDE_4_5_OPUS_THINKING"), credit_multiplier: 5.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-sonnet-4.6", enum_value: 0, model_uid: Some("claude-sonnet-4-6"), credit_multiplier: 4.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-sonnet-4.6-thinking", enum_value: 0, model_uid: Some("claude-sonnet-4-6-thinking"), credit_multiplier: 6.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-sonnet-4.6-1m", enum_value: 0, model_uid: Some("claude-sonnet-4-6-1m"), credit_multiplier: 12.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-sonnet-4.6-thinking-1m", enum_value: 0, model_uid: Some("claude-sonnet-4-6-thinking-1m"), credit_multiplier: 16.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4.6", enum_value: 0, model_uid: Some("claude-opus-4-6"), credit_multiplier: 6.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4.6-thinking", enum_value: 0, model_uid: Some("claude-opus-4-6-thinking"), credit_multiplier: 8.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4-7-medium", enum_value: 0, model_uid: Some("claude-opus-4-7-medium"), credit_multiplier: 8.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4-7-low", enum_value: 0, model_uid: Some("claude-opus-4-7-low"), credit_multiplier: 6.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4-7-high", enum_value: 0, model_uid: Some("claude-opus-4-7-high"), credit_multiplier: 10.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4-7-xhigh", enum_value: 0, model_uid: Some("claude-opus-4-7-xhigh"), credit_multiplier: 12.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4-7-medium-thinking", enum_value: 0, model_uid: Some("claude-opus-4-7-medium-thinking"), credit_multiplier: 10.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4-7-high-thinking", enum_value: 0, model_uid: Some("claude-opus-4-7-high-thinking"), credit_multiplier: 12.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4-7-xhigh-thinking", enum_value: 0, model_uid: Some("claude-opus-4-7-xhigh-thinking"), credit_multiplier: 16.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "claude-opus-4-7-max", enum_value: 0, model_uid: Some("claude-opus-4-7-max"), credit_multiplier: 16.0, provider: "anthropic", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-4o", enum_value: 109, model_uid: Some("MODEL_CHAT_GPT_4O_2024_08_06"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-4o-mini", enum_value: 113, model_uid: None, credit_multiplier: 0.5, provider: "openai", deprecated: true },
|
||||
WindsurfModel { canonical_name: "gpt-4.1", enum_value: 259, model_uid: Some("MODEL_CHAT_GPT_4_1_2025_04_14"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-4.1-mini", enum_value: 260, model_uid: None, credit_multiplier: 0.5, provider: "openai", deprecated: true },
|
||||
WindsurfModel { canonical_name: "gpt-4.1-nano", enum_value: 261, model_uid: None, credit_multiplier: 0.25, provider: "openai", deprecated: true },
|
||||
WindsurfModel { canonical_name: "gpt-5", enum_value: 340, model_uid: Some("MODEL_PRIVATE_6"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5-medium", enum_value: 0, model_uid: Some("MODEL_PRIVATE_7"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5-high", enum_value: 0, model_uid: Some("MODEL_PRIVATE_8"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5-mini", enum_value: 337, model_uid: None, credit_multiplier: 0.25, provider: "openai", deprecated: true },
|
||||
WindsurfModel { canonical_name: "gpt-5-codex", enum_value: 346, model_uid: Some("MODEL_CHAT_GPT_5_CODEX"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1", enum_value: 0, model_uid: Some("MODEL_PRIVATE_12"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-low", enum_value: 0, model_uid: Some("MODEL_PRIVATE_13"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-medium", enum_value: 0, model_uid: Some("MODEL_PRIVATE_14"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-high", enum_value: 0, model_uid: Some("MODEL_PRIVATE_15"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-fast", enum_value: 0, model_uid: Some("MODEL_PRIVATE_20"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-low-fast", enum_value: 0, model_uid: Some("MODEL_PRIVATE_21"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-medium-fast", enum_value: 0, model_uid: Some("MODEL_PRIVATE_22"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-high-fast", enum_value: 0, model_uid: Some("MODEL_PRIVATE_23"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-codex-low", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_LOW"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-codex-medium", enum_value: 0, model_uid: Some("MODEL_PRIVATE_9"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-codex-mini-low", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_MINI_LOW"), credit_multiplier: 0.25, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-codex-mini", enum_value: 0, model_uid: Some("MODEL_PRIVATE_19"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-codex-max-low", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_MAX_LOW"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-codex-max-medium", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_MAX_MEDIUM"), credit_multiplier: 1.25, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.1-codex-max-high", enum_value: 0, model_uid: Some("MODEL_GPT_5_1_CODEX_MAX_HIGH"), credit_multiplier: 1.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2", enum_value: 401, model_uid: Some("MODEL_GPT_5_2_MEDIUM"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-none", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_NONE"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-low", enum_value: 400, model_uid: Some("MODEL_GPT_5_2_LOW"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-high", enum_value: 402, model_uid: Some("MODEL_GPT_5_2_HIGH"), credit_multiplier: 3.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-xhigh", enum_value: 403, model_uid: Some("MODEL_GPT_5_2_XHIGH"), credit_multiplier: 8.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-none-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_NONE_PRIORITY"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-low-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_LOW_PRIORITY"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-medium-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_MEDIUM_PRIORITY"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-high-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_HIGH_PRIORITY"), credit_multiplier: 6.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-xhigh-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_XHIGH_PRIORITY"), credit_multiplier: 16.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-codex-low", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_LOW"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-codex-medium", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_MEDIUM"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-codex-high", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_HIGH"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-codex-xhigh", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_XHIGH"), credit_multiplier: 3.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-codex-low-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_LOW_PRIORITY"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-codex-medium-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_MEDIUM_PRIORITY"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-codex-high-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_HIGH_PRIORITY"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.2-codex-xhigh-fast", enum_value: 0, model_uid: Some("MODEL_GPT_5_2_CODEX_XHIGH_PRIORITY"), credit_multiplier: 6.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.3-codex", enum_value: 0, model_uid: Some("gpt-5-3-codex-medium"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-none", enum_value: 0, model_uid: Some("gpt-5-4-none"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-low", enum_value: 0, model_uid: Some("gpt-5-4-low"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-medium", enum_value: 0, model_uid: Some("gpt-5-4-medium"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-high", enum_value: 0, model_uid: Some("gpt-5-4-high"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-xhigh", enum_value: 0, model_uid: Some("gpt-5-4-xhigh"), credit_multiplier: 8.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-mini-low", enum_value: 0, model_uid: Some("gpt-5-4-mini-low"), credit_multiplier: 1.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-mini-medium", enum_value: 0, model_uid: Some("gpt-5-4-mini-medium"), credit_multiplier: 1.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-mini-high", enum_value: 0, model_uid: Some("gpt-5-4-mini-high"), credit_multiplier: 4.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.4-mini-xhigh", enum_value: 0, model_uid: Some("gpt-5-4-mini-xhigh"), credit_multiplier: 12.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5", enum_value: 0, model_uid: Some("gpt-5-5-medium"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-none", enum_value: 0, model_uid: Some("gpt-5-5-none"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-low", enum_value: 0, model_uid: Some("gpt-5-5-low"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-medium", enum_value: 0, model_uid: Some("gpt-5-5-medium"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-high", enum_value: 0, model_uid: Some("gpt-5-5-high"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-xhigh", enum_value: 0, model_uid: Some("gpt-5-5-xhigh"), credit_multiplier: 8.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-none-fast", enum_value: 0, model_uid: Some("gpt-5-5-none-priority"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-low-fast", enum_value: 0, model_uid: Some("gpt-5-5-low-priority"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-medium-fast", enum_value: 0, model_uid: Some("gpt-5-5-medium-priority"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-high-fast", enum_value: 0, model_uid: Some("gpt-5-5-high-priority"), credit_multiplier: 8.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.5-xhigh-fast", enum_value: 0, model_uid: Some("gpt-5-5-xhigh-priority"), credit_multiplier: 16.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.3-codex-low", enum_value: 0, model_uid: Some("gpt-5-3-codex-low"), credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.3-codex-high", enum_value: 0, model_uid: Some("gpt-5-3-codex-high"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.3-codex-xhigh", enum_value: 0, model_uid: Some("gpt-5-3-codex-xhigh"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.3-codex-low-fast", enum_value: 0, model_uid: Some("gpt-5-3-codex-low-priority"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.3-codex-medium-fast", enum_value: 0, model_uid: Some("gpt-5-3-codex-medium-priority"), credit_multiplier: 2.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.3-codex-high-fast", enum_value: 0, model_uid: Some("gpt-5-3-codex-high-priority"), credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-5.3-codex-xhigh-fast", enum_value: 0, model_uid: Some("gpt-5-3-codex-xhigh-priority"), credit_multiplier: 6.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gpt-oss-120b", enum_value: 0, model_uid: Some("MODEL_GPT_OSS_120B"), credit_multiplier: 0.25, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "o3-mini", enum_value: 207, model_uid: None, credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "o3", enum_value: 218, model_uid: Some("MODEL_CHAT_O3"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "o3-high", enum_value: 0, model_uid: Some("MODEL_CHAT_O3_HIGH"), credit_multiplier: 1.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "o3-pro", enum_value: 294, model_uid: None, credit_multiplier: 4.0, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "o4-mini", enum_value: 264, model_uid: None, credit_multiplier: 0.5, provider: "openai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-2.5-pro", enum_value: 246, model_uid: Some("MODEL_GOOGLE_GEMINI_2_5_PRO"), credit_multiplier: 1.0, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-2.5-flash", enum_value: 312, model_uid: Some("MODEL_GOOGLE_GEMINI_2_5_FLASH"), credit_multiplier: 0.5, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-3.0-pro", enum_value: 412, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_PRO_LOW"), credit_multiplier: 1.0, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-3.0-flash-minimal", enum_value: 0, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_FLASH_MINIMAL"), credit_multiplier: 0.75, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-3.0-flash-low", enum_value: 0, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_FLASH_LOW"), credit_multiplier: 1.0, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-3.0-flash", enum_value: 415, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_FLASH_MEDIUM"), credit_multiplier: 1.0, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-3.0-flash-high", enum_value: 0, model_uid: Some("MODEL_GOOGLE_GEMINI_3_0_FLASH_HIGH"), credit_multiplier: 1.75, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-3.1-pro-low", enum_value: 0, model_uid: Some("gemini-3-1-pro-low"), credit_multiplier: 1.0, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "gemini-3.1-pro-high", enum_value: 0, model_uid: Some("gemini-3-1-pro-high"), credit_multiplier: 2.0, provider: "google", deprecated: false },
|
||||
WindsurfModel { canonical_name: "deepseek-v3", enum_value: 205, model_uid: None, credit_multiplier: 0.5, provider: "deepseek", deprecated: true },
|
||||
WindsurfModel { canonical_name: "deepseek-v3-2", enum_value: 409, model_uid: None, credit_multiplier: 0.5, provider: "deepseek", deprecated: true },
|
||||
WindsurfModel { canonical_name: "deepseek-r1", enum_value: 206, model_uid: None, credit_multiplier: 1.0, provider: "deepseek", deprecated: true },
|
||||
WindsurfModel { canonical_name: "grok-3", enum_value: 217, model_uid: Some("MODEL_XAI_GROK_3"), credit_multiplier: 1.0, provider: "xai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "grok-3-mini", enum_value: 234, model_uid: None, credit_multiplier: 0.5, provider: "xai", deprecated: true },
|
||||
WindsurfModel { canonical_name: "grok-3-mini-thinking", enum_value: 0, model_uid: Some("MODEL_XAI_GROK_3_MINI_REASONING"), credit_multiplier: 0.125, provider: "xai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "grok-code-fast-1", enum_value: 0, model_uid: Some("MODEL_PRIVATE_4"), credit_multiplier: 0.5, provider: "xai", deprecated: false },
|
||||
WindsurfModel { canonical_name: "qwen-3", enum_value: 324, model_uid: None, credit_multiplier: 0.5, provider: "alibaba", deprecated: true },
|
||||
WindsurfModel { canonical_name: "kimi-k2", enum_value: 323, model_uid: Some("MODEL_KIMI_K2"), credit_multiplier: 0.5, provider: "moonshot", deprecated: false },
|
||||
WindsurfModel { canonical_name: "kimi-k2-thinking", enum_value: 394, model_uid: Some("MODEL_KIMI_K2_THINKING"), credit_multiplier: 1.0, provider: "moonshot", deprecated: false },
|
||||
WindsurfModel { canonical_name: "kimi-k2.5", enum_value: 0, model_uid: Some("kimi-k2-5"), credit_multiplier: 1.0, provider: "moonshot", deprecated: false },
|
||||
WindsurfModel { canonical_name: "kimi-k2-6", enum_value: 0, model_uid: Some("kimi-k2-6"), credit_multiplier: 1.0, provider: "moonshot", deprecated: false },
|
||||
WindsurfModel { canonical_name: "glm-4.7", enum_value: 417, model_uid: Some("MODEL_GLM_4_7"), credit_multiplier: 0.25, provider: "zhipu", deprecated: false },
|
||||
WindsurfModel { canonical_name: "glm-4.7-fast", enum_value: 418, model_uid: Some("MODEL_GLM_4_7_FAST"), credit_multiplier: 0.5, provider: "zhipu", deprecated: false },
|
||||
WindsurfModel { canonical_name: "glm-5", enum_value: 0, model_uid: Some("glm-5"), credit_multiplier: 1.5, provider: "zhipu", deprecated: false },
|
||||
WindsurfModel { canonical_name: "glm-5.1", enum_value: 0, model_uid: Some("glm-5-1"), credit_multiplier: 1.5, provider: "zhipu", deprecated: false },
|
||||
WindsurfModel { canonical_name: "minimax-m2.5", enum_value: 419, model_uid: Some("MODEL_MINIMAX_M2_1"), credit_multiplier: 1.0, provider: "minimax", deprecated: false },
|
||||
WindsurfModel { canonical_name: "swe-1.5", enum_value: 377, model_uid: Some("MODEL_SWE_1_5_SLOW"), credit_multiplier: 0.5, provider: "windsurf", deprecated: false },
|
||||
WindsurfModel { canonical_name: "swe-1.5-fast", enum_value: 359, model_uid: Some("MODEL_SWE_1_5"), credit_multiplier: 0.5, provider: "windsurf", deprecated: false },
|
||||
WindsurfModel { canonical_name: "swe-1.5-thinking", enum_value: 369, model_uid: Some("MODEL_SWE_1_5_THINKING"), credit_multiplier: 0.75, provider: "windsurf", deprecated: false },
|
||||
WindsurfModel { canonical_name: "swe-1.6", enum_value: 420, model_uid: Some("MODEL_SWE_1_6"), credit_multiplier: 0.5, provider: "windsurf", deprecated: false },
|
||||
WindsurfModel { canonical_name: "swe-1.6-fast", enum_value: 421, model_uid: Some("MODEL_SWE_1_6_FAST"), credit_multiplier: 0.5, provider: "windsurf", deprecated: false },
|
||||
WindsurfModel { canonical_name: "adaptive", enum_value: 0, model_uid: Some("adaptive"), credit_multiplier: 1.0, provider: "windsurf", deprecated: true },
|
||||
WindsurfModel { canonical_name: "arena-fast", enum_value: 0, model_uid: Some("arena-fast"), credit_multiplier: 0.5, provider: "windsurf", deprecated: true },
|
||||
WindsurfModel { canonical_name: "arena-smart", enum_value: 0, model_uid: Some("arena-smart"), credit_multiplier: 1.0, provider: "windsurf", deprecated: true },
|
||||
];
|
||||
|
||||
#[rustfmt::skip]
|
||||
const ALIASES: &[(&str, &str)] = &[
|
||||
("claude-3-5-haiku-20241022", "claude-4.5-haiku"),
|
||||
("claude-3-5-haiku-latest", "claude-4.5-haiku"),
|
||||
("claude-3-5-sonnet-20240620", "claude-3.5-sonnet"),
|
||||
("claude-3-5-sonnet-20241022", "claude-3.5-sonnet"),
|
||||
("claude-3-5-sonnet-latest", "claude-3.5-sonnet"),
|
||||
("claude-3-7-sonnet-20250219", "claude-3.7-sonnet"),
|
||||
("claude-3-7-sonnet-latest", "claude-3.7-sonnet"),
|
||||
("claude-4.6", "claude-sonnet-4.6"),
|
||||
("claude-4.6-1m", "claude-sonnet-4.6-1m"),
|
||||
("claude-4.6-thinking", "claude-sonnet-4.6-thinking"),
|
||||
("claude-4.6-thinking-1m", "claude-sonnet-4.6-thinking-1m"),
|
||||
("claude-haiku-3-5", "claude-4.5-haiku"),
|
||||
("claude-haiku-3-5-latest", "claude-4.5-haiku"),
|
||||
("claude-haiku-4-5", "claude-4.5-haiku"),
|
||||
("claude-haiku-4-5-20251001", "claude-4.5-haiku"),
|
||||
("claude-haiku-4-5-latest", "claude-4.5-haiku"),
|
||||
("claude-haiku-4.5", "claude-4.5-haiku"),
|
||||
("claude-haiku-4.5-latest", "claude-4.5-haiku"),
|
||||
("claude-opus-4-0", "claude-4-opus"),
|
||||
("claude-opus-4-1", "claude-4.1-opus"),
|
||||
("claude-opus-4-1-20250805", "claude-4.1-opus"),
|
||||
("claude-opus-4-20250514", "claude-4-opus"),
|
||||
("claude-opus-4-5", "claude-4.5-opus"),
|
||||
("claude-opus-4-5-20251101", "claude-4.5-opus"),
|
||||
("claude-opus-4-5-latest", "claude-4.5-opus"),
|
||||
("claude-opus-4-6", "claude-opus-4.6"),
|
||||
("claude-opus-4-6-thinking", "claude-opus-4.6-thinking"),
|
||||
("claude-opus-4-7", "claude-opus-4-7-medium"),
|
||||
("claude-opus-4-7-latest", "claude-opus-4-7-medium"),
|
||||
("claude-opus-4-7-thinking", "claude-opus-4-7-medium-thinking"),
|
||||
("claude-opus-4.5", "claude-4.5-opus"),
|
||||
("claude-opus-4.5-thinking", "claude-4.5-opus-thinking"),
|
||||
("claude-opus-4.7", "claude-opus-4-7-medium"),
|
||||
("claude-opus-4.7-high", "claude-opus-4-7-high"),
|
||||
("claude-opus-4.7-high-thinking", "claude-opus-4-7-high-thinking"),
|
||||
("claude-opus-4.7-low", "claude-opus-4-7-low"),
|
||||
("claude-opus-4.7-max", "claude-opus-4-7-max"),
|
||||
("claude-opus-4.7-medium", "claude-opus-4-7-medium"),
|
||||
("claude-opus-4.7-medium-thinking", "claude-opus-4-7-medium-thinking"),
|
||||
("claude-opus-4.7-thinking", "claude-opus-4-7-medium-thinking"),
|
||||
("claude-opus-4.7-xhigh", "claude-opus-4-7-xhigh"),
|
||||
("claude-opus-4.7-xhigh-thinking", "claude-opus-4-7-xhigh-thinking"),
|
||||
("claude-sonnet-4-0", "claude-4-sonnet"),
|
||||
("claude-sonnet-4-20250514", "claude-4-sonnet"),
|
||||
("claude-sonnet-4-5", "claude-4.5-sonnet"),
|
||||
("claude-sonnet-4-5-20250929", "claude-4.5-sonnet"),
|
||||
("claude-sonnet-4-5-latest", "claude-4.5-sonnet"),
|
||||
("claude-sonnet-4-6", "claude-sonnet-4.6"),
|
||||
("claude-sonnet-4-6-1m", "claude-sonnet-4.6-1m"),
|
||||
("claude-sonnet-4-6-thinking", "claude-sonnet-4.6-thinking"),
|
||||
("claude-sonnet-4-6-thinking-1m", "claude-sonnet-4.6-thinking-1m"),
|
||||
("claude-sonnet-4.5", "claude-4.5-sonnet"),
|
||||
("claude-sonnet-4.5-thinking", "claude-4.5-sonnet-thinking"),
|
||||
("gpt-4.1-2025-04-14", "gpt-4.1"),
|
||||
("gpt-4.1-mini-2025-04-14", "gpt-4.1-mini"),
|
||||
("gpt-4.1-nano-2025-04-14", "gpt-4.1-nano"),
|
||||
("gpt-4o-2024-05-13", "gpt-4o"),
|
||||
("gpt-4o-2024-08-06", "gpt-4o"),
|
||||
("gpt-4o-2024-11-20", "gpt-4o"),
|
||||
("gpt-4o-mini-2024-07-18", "gpt-4o-mini"),
|
||||
("gpt-5-2-codex-medium", "gpt-5.2-codex-medium"),
|
||||
("gpt-5-2-medium", "gpt-5.2"),
|
||||
("gpt-5-2025-08-07", "gpt-5"),
|
||||
("gpt-5-3-codex-high", "gpt-5.3-codex-high"),
|
||||
("gpt-5-3-codex-high-priority", "gpt-5.3-codex-high-fast"),
|
||||
("gpt-5-3-codex-low", "gpt-5.3-codex-low"),
|
||||
("gpt-5-3-codex-low-priority", "gpt-5.3-codex-low-fast"),
|
||||
("gpt-5-3-codex-medium", "gpt-5.3-codex"),
|
||||
("gpt-5-3-codex-medium-priority", "gpt-5.3-codex-medium-fast"),
|
||||
("gpt-5-3-codex-xhigh", "gpt-5.3-codex-xhigh"),
|
||||
("gpt-5-3-codex-xhigh-priority", "gpt-5.3-codex-xhigh-fast"),
|
||||
("gpt-5-4-high", "gpt-5.4-high"),
|
||||
("gpt-5-4-low", "gpt-5.4-low"),
|
||||
("gpt-5-4-medium", "gpt-5.4-medium"),
|
||||
("gpt-5-4-mini-high", "gpt-5.4-mini-high"),
|
||||
("gpt-5-4-mini-low", "gpt-5.4-mini-low"),
|
||||
("gpt-5-4-mini-medium", "gpt-5.4-mini-medium"),
|
||||
("gpt-5-4-mini-xhigh", "gpt-5.4-mini-xhigh"),
|
||||
("gpt-5-4-none", "gpt-5.4-none"),
|
||||
("gpt-5-4-xhigh", "gpt-5.4-xhigh"),
|
||||
("gpt-5-5", "gpt-5.5-medium"),
|
||||
("gpt-5-5-high", "gpt-5.5-high"),
|
||||
("gpt-5-5-high-priority", "gpt-5.5-high-fast"),
|
||||
("gpt-5-5-low", "gpt-5.5-low"),
|
||||
("gpt-5-5-low-priority", "gpt-5.5-low-fast"),
|
||||
("gpt-5-5-medium", "gpt-5.5-medium"),
|
||||
("gpt-5-5-medium-priority", "gpt-5.5-medium-fast"),
|
||||
("gpt-5-5-none", "gpt-5.5-none"),
|
||||
("gpt-5-5-none-priority", "gpt-5.5-none-fast"),
|
||||
("gpt-5-5-xhigh", "gpt-5.5-xhigh"),
|
||||
("gpt-5-5-xhigh-priority", "gpt-5.5-xhigh-fast"),
|
||||
("gpt-5-pro-2025-10-06", "gpt-5-high"),
|
||||
("gpt-5.2-codex", "gpt-5.2-codex-medium"),
|
||||
("gpt-5.2-medium", "gpt-5.2"),
|
||||
("gpt-5.3-codex-medium", "gpt-5.3-codex"),
|
||||
("gpt-5.4", "gpt-5.4-medium"),
|
||||
("gpt-5.5", "gpt-5.5-medium"),
|
||||
("haiku-4.5", "claude-4.5-haiku"),
|
||||
("kimi-k2-5", "kimi-k2.5"),
|
||||
("minimax-m2-5", "minimax-m2.5"),
|
||||
("model_claude_4_5_sonnet", "claude-4.5-sonnet"),
|
||||
("model_claude_4_5_sonnet_thinking", "claude-4.5-sonnet-thinking"),
|
||||
("o4.7", "claude-opus-4-7-medium"),
|
||||
("opus-4", "claude-4-opus"),
|
||||
("opus-4-7", "claude-opus-4-7-medium"),
|
||||
("opus-4.1", "claude-4.1-opus"),
|
||||
("opus-4.6", "claude-opus-4.6"),
|
||||
("opus-4.6-thinking", "claude-opus-4.6-thinking"),
|
||||
("opus-4.7", "claude-opus-4-7-medium"),
|
||||
("opus-4.7-thinking", "claude-opus-4-7-medium-thinking"),
|
||||
("sonnet-3.5", "claude-3.5-sonnet"),
|
||||
("sonnet-3.7", "claude-3.7-sonnet"),
|
||||
("sonnet-4", "claude-4-sonnet"),
|
||||
("sonnet-4.5", "claude-4.5-sonnet"),
|
||||
("sonnet-4.5-thinking", "claude-4.5-sonnet-thinking"),
|
||||
("sonnet-4.6", "claude-sonnet-4.6"),
|
||||
("sonnet-4.6-1m", "claude-sonnet-4.6-1m"),
|
||||
("sonnet-4.6-thinking", "claude-sonnet-4.6-thinking"),
|
||||
("swe-1-6", "swe-1.6"),
|
||||
("swe-1-6-fast", "swe-1.6-fast"),
|
||||
("ws-haiku", "claude-4.5-haiku"),
|
||||
("ws-opus", "claude-opus-4.6"),
|
||||
("ws-opus-thinking", "claude-opus-4.6-thinking"),
|
||||
("ws-sonnet", "claude-sonnet-4.6"),
|
||||
("ws-sonnet-thinking", "claude-sonnet-4.6-thinking"),
|
||||
];
|
||||
|
||||
pub fn windsurf_models() -> &'static [WindsurfModel] {
|
||||
MODELS
|
||||
}
|
||||
|
||||
pub fn resolve_windsurf_model(name: &str) -> Option<WindsurfModel> {
|
||||
let normalized = name.trim().to_ascii_lowercase();
|
||||
if normalized.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let canonical = ALIASES
|
||||
.iter()
|
||||
.find_map(|(alias, canonical)| (*alias == normalized).then_some(*canonical))
|
||||
.unwrap_or(normalized.as_str());
|
||||
MODELS
|
||||
.iter()
|
||||
.find(|model| model_matches(model, canonical))
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn model_matches(model: &WindsurfModel, value: &str) -> bool {
|
||||
model.canonical_name.eq_ignore_ascii_case(value)
|
||||
|| model
|
||||
.model_uid
|
||||
.is_some_and(|uid| uid.eq_ignore_ascii_case(value))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn resolves_gpt55_cloud_alias_to_windsurf_model_uid() {
|
||||
let model = resolve_windsurf_model("gpt-5-5-low").expect("model should resolve");
|
||||
|
||||
assert_eq!(model.canonical_name, "gpt-5.5-low");
|
||||
assert_eq!(model.model_uid, Some("gpt-5-5-low"));
|
||||
assert_eq!(model.enum_value, 0);
|
||||
assert_eq!(model.credit_multiplier, 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_claude_opus_47_bare_alias_to_medium() {
|
||||
let model = resolve_windsurf_model("claude-opus-4.7").expect("model should resolve");
|
||||
|
||||
assert_eq!(model.canonical_name, "claude-opus-4-7-medium");
|
||||
assert_eq!(model.model_uid, Some("claude-opus-4-7-medium"));
|
||||
assert_eq!(model.enum_value, 0);
|
||||
assert_eq!(model.credit_multiplier, 8.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_priority_alias_to_fast_variant() {
|
||||
let model = resolve_windsurf_model("gpt-5-5-low-priority").expect("model should resolve");
|
||||
|
||||
assert_eq!(model.canonical_name, "gpt-5.5-low-fast");
|
||||
assert_eq!(model.model_uid, Some("gpt-5-5-low-priority"));
|
||||
assert_eq!(model.credit_multiplier, 2.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_full_gpt55_effort_ladder_and_priority_aliases() {
|
||||
let none = resolve_windsurf_model("gpt-5-5-none").expect("none should resolve");
|
||||
assert_eq!(none.canonical_name, "gpt-5.5-none");
|
||||
assert_eq!(none.model_uid, Some("gpt-5-5-none"));
|
||||
assert_eq!(none.credit_multiplier, 1.0);
|
||||
|
||||
let high = resolve_windsurf_model("gpt-5.5-high").expect("high should resolve");
|
||||
assert_eq!(high.canonical_name, "gpt-5.5-high");
|
||||
assert_eq!(high.model_uid, Some("gpt-5-5-high"));
|
||||
assert_eq!(high.credit_multiplier, 4.0);
|
||||
|
||||
let xhigh_fast = resolve_windsurf_model("gpt-5-5-xhigh-priority")
|
||||
.expect("xhigh priority should resolve");
|
||||
assert_eq!(xhigh_fast.canonical_name, "gpt-5.5-xhigh-fast");
|
||||
assert_eq!(xhigh_fast.model_uid, Some("gpt-5-5-xhigh-priority"));
|
||||
assert_eq!(xhigh_fast.credit_multiplier, 16.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_windsurfapi_catalog_aliases_beyond_gpt55() {
|
||||
let gpt52_medium = resolve_windsurf_model("gpt-5.2-medium").expect("gpt-5.2 medium alias");
|
||||
assert_eq!(gpt52_medium.canonical_name, "gpt-5.2");
|
||||
assert_eq!(gpt52_medium.model_uid, Some("MODEL_GPT_5_2_MEDIUM"));
|
||||
|
||||
let haiku = resolve_windsurf_model("claude-haiku-4-5-20251001").expect("dated haiku alias");
|
||||
assert_eq!(haiku.canonical_name, "claude-4.5-haiku");
|
||||
assert_eq!(haiku.model_uid, Some("MODEL_PRIVATE_11"));
|
||||
|
||||
let uid = resolve_windsurf_model("MODEL_GPT_5_2_LOW").expect("model uid alias");
|
||||
assert_eq!(uid.canonical_name, "gpt-5.2-low");
|
||||
assert_eq!(uid.enum_value, 400);
|
||||
|
||||
let cursor = resolve_windsurf_model("ws-opus").expect("cursor alias");
|
||||
assert_eq!(cursor.canonical_name, "claude-opus-4.6");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_catalog_covers_current_windsurfapi_model_set() {
|
||||
assert_eq!(MODELS.len(), 139);
|
||||
assert!(ALIASES.len() >= 100);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
use std::fmt;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum WireType {
|
||||
Varint = 0,
|
||||
Fixed64 = 1,
|
||||
Len = 2,
|
||||
Fixed32 = 5,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum FieldValue {
|
||||
Varint(u64),
|
||||
Bytes(Vec<u8>),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct Field {
|
||||
pub number: u32,
|
||||
pub wire_type: WireType,
|
||||
pub value: FieldValue,
|
||||
}
|
||||
|
||||
impl Field {
|
||||
pub fn bytes(&self) -> &[u8] {
|
||||
match &self.value {
|
||||
FieldValue::Bytes(bytes) => bytes.as_slice(),
|
||||
FieldValue::Varint(_) => &[],
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct ProtoError {
|
||||
message: String,
|
||||
}
|
||||
|
||||
impl ProtoError {
|
||||
fn new(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for ProtoError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.write_str(&self.message)
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for ProtoError {}
|
||||
|
||||
pub fn encode_varint(value: u64) -> Vec<u8> {
|
||||
let mut value = value;
|
||||
let mut out = Vec::new();
|
||||
loop {
|
||||
let mut byte = (value & 0x7f) as u8;
|
||||
value >>= 7;
|
||||
if value != 0 {
|
||||
byte |= 0x80;
|
||||
}
|
||||
out.push(byte);
|
||||
if value == 0 {
|
||||
break;
|
||||
}
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
fn decode_varint(buf: &[u8], offset: usize) -> Result<(u64, usize), ProtoError> {
|
||||
let mut value = 0u64;
|
||||
let mut shift = 0u32;
|
||||
let mut pos = offset;
|
||||
while pos < buf.len() {
|
||||
let byte = buf[pos];
|
||||
pos += 1;
|
||||
value |= u64::from(byte & 0x7f) << shift;
|
||||
if byte & 0x80 == 0 {
|
||||
return Ok((value, pos - offset));
|
||||
}
|
||||
shift += 7;
|
||||
if shift >= 64 {
|
||||
return Err(ProtoError::new("varint overflow"));
|
||||
}
|
||||
}
|
||||
Err(ProtoError::new("truncated varint"))
|
||||
}
|
||||
|
||||
fn tag(field: u32, wire_type: WireType) -> Vec<u8> {
|
||||
encode_varint((u64::from(field) << 3) | wire_type as u64)
|
||||
}
|
||||
|
||||
pub fn write_varint_field(field: u32, value: u64) -> Vec<u8> {
|
||||
let mut out = tag(field, WireType::Varint);
|
||||
out.extend(encode_varint(value));
|
||||
out
|
||||
}
|
||||
|
||||
pub fn write_string_field(field: u32, value: &str) -> Vec<u8> {
|
||||
let bytes = value.as_bytes();
|
||||
let mut out = tag(field, WireType::Len);
|
||||
out.extend(encode_varint(bytes.len() as u64));
|
||||
out.extend(bytes);
|
||||
out
|
||||
}
|
||||
|
||||
pub fn write_message_field(field: u32, value: &[u8]) -> Vec<u8> {
|
||||
if value.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
let mut out = tag(field, WireType::Len);
|
||||
out.extend(encode_varint(value.len() as u64));
|
||||
out.extend(value);
|
||||
out
|
||||
}
|
||||
|
||||
pub fn write_bool_field(field: u32, value: bool) -> Vec<u8> {
|
||||
if value {
|
||||
write_varint_field(field, 1)
|
||||
} else {
|
||||
Vec::new()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parse_fields(buf: &[u8]) -> Result<Vec<Field>, ProtoError> {
|
||||
let mut fields = Vec::new();
|
||||
let mut pos = 0usize;
|
||||
while pos < buf.len() {
|
||||
let (tag, tag_len) = decode_varint(buf, pos)?;
|
||||
pos += tag_len;
|
||||
let number = (tag >> 3) as u32;
|
||||
let wire_type = match tag & 0x07 {
|
||||
0 => WireType::Varint,
|
||||
1 => WireType::Fixed64,
|
||||
2 => WireType::Len,
|
||||
5 => WireType::Fixed32,
|
||||
other => {
|
||||
return Err(ProtoError::new(format!(
|
||||
"unknown wire type {other} at offset {pos}"
|
||||
)))
|
||||
}
|
||||
};
|
||||
let value = match wire_type {
|
||||
WireType::Varint => {
|
||||
let (value, value_len) = decode_varint(buf, pos)?;
|
||||
pos += value_len;
|
||||
FieldValue::Varint(value)
|
||||
}
|
||||
WireType::Len => {
|
||||
let (len, len_len) = decode_varint(buf, pos)?;
|
||||
pos += len_len;
|
||||
let len = usize::try_from(len)
|
||||
.map_err(|_| ProtoError::new("length-delimited field too large"))?;
|
||||
if pos + len > buf.len() {
|
||||
return Err(ProtoError::new(format!(
|
||||
"truncated len-delimited field {number} at offset {pos}"
|
||||
)));
|
||||
}
|
||||
let bytes = buf[pos..pos + len].to_vec();
|
||||
pos += len;
|
||||
FieldValue::Bytes(bytes)
|
||||
}
|
||||
WireType::Fixed64 => {
|
||||
if pos + 8 > buf.len() {
|
||||
return Err(ProtoError::new(format!("truncated fixed64 field {number}")));
|
||||
}
|
||||
let bytes = buf[pos..pos + 8].to_vec();
|
||||
pos += 8;
|
||||
FieldValue::Bytes(bytes)
|
||||
}
|
||||
WireType::Fixed32 => {
|
||||
if pos + 4 > buf.len() {
|
||||
return Err(ProtoError::new(format!("truncated fixed32 field {number}")));
|
||||
}
|
||||
let bytes = buf[pos..pos + 4].to_vec();
|
||||
pos += 4;
|
||||
FieldValue::Bytes(bytes)
|
||||
}
|
||||
};
|
||||
fields.push(Field {
|
||||
number,
|
||||
wire_type,
|
||||
value,
|
||||
});
|
||||
}
|
||||
Ok(fields)
|
||||
}
|
||||
|
||||
pub fn get_field(fields: &[Field], number: u32, wire_type: Option<WireType>) -> Option<&Field> {
|
||||
fields
|
||||
.iter()
|
||||
.find(|field| field.number == number && wire_type.is_none_or(|ty| field.wire_type == ty))
|
||||
}
|
||||
|
||||
pub fn get_all_fields(fields: &[Field], number: u32) -> Vec<&Field> {
|
||||
fields
|
||||
.iter()
|
||||
.filter(|field| field.number == number)
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn get_varint(fields: &[Field], number: u32) -> Option<u64> {
|
||||
match get_field(fields, number, Some(WireType::Varint))?.value {
|
||||
FieldValue::Varint(value) => Some(value),
|
||||
FieldValue::Bytes(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_string(fields: &[Field], number: u32) -> Option<String> {
|
||||
let field = get_field(fields, number, Some(WireType::Len))?;
|
||||
String::from_utf8(field.bytes().to_vec()).ok()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn encodes_varint_and_string_fields_like_windsurfapi() {
|
||||
assert_eq!(encode_varint(300), vec![0xac, 0x02]);
|
||||
assert_eq!(
|
||||
write_string_field(3, "abc"),
|
||||
vec![0x1a, 0x03, b'a', b'b', b'c']
|
||||
);
|
||||
assert_eq!(write_bool_field(2, false), Vec::<u8>::new());
|
||||
assert_eq!(write_bool_field(2, true), vec![0x10, 0x01]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_repeated_len_delimited_fields() {
|
||||
let mut bytes = Vec::new();
|
||||
bytes.extend(write_string_field(1, "alpha"));
|
||||
bytes.extend(write_string_field(1, "beta"));
|
||||
bytes.extend(write_varint_field(2, 42));
|
||||
|
||||
let fields = parse_fields(&bytes).expect("fields should parse");
|
||||
assert_eq!(get_all_fields(&fields, 1).len(), 2);
|
||||
assert_eq!(get_string(&fields, 1).as_deref(), Some("alpha"));
|
||||
assert_eq!(get_varint(&fields, 2), Some(42));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_truncated_len_delimited_field() {
|
||||
let err = parse_fields(&[0x0a, 0x05, b'a']).expect_err("must reject truncated field");
|
||||
assert!(err.to_string().contains("truncated"));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user