refactor: 拆分 gateway 单体为独立 crate,新增 systemd 部署方案

将 gateway 内部的 model-fetch、provider-transport、scheduler-core、
usage-runtime、video-tasks-core 模块提取为独立 crate;重构 gateway
内部模块结构(state/router/cache/data/query 等);移除大量遗留模块
文件;新增 systemd 二进制部署骨架及相关文档;更新前端 usage 相关
API 和组件。
This commit is contained in:
fawney19
2026-04-05 20:23:16 +08:00
parent cbc811f6ce
commit 763ff03a7b
777 changed files with 42659 additions and 21469 deletions

View File

@@ -0,0 +1,28 @@
[package]
name = "aether-provider-transport"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
description = "Provider transport core extracted from aether-gateway"
[dependencies]
aether-contracts.workspace = true
aether-crypto.workspace = true
aether-data.workspace = true
aether-video-tasks-core.workspace = true
async-trait.workspace = true
http.workspace = true
regex.workspace = true
reqwest.workspace = true
serde.workspace = true
serde_json.workspace = true
sha2.workspace = true
thiserror.workspace = true
tokio.workspace = true
tracing.workspace = true
url.workspace = true
uuid.workspace = true
[dev-dependencies]
axum = { version = "0.8", features = ["ws"] }

View File

@@ -0,0 +1,23 @@
mod auth;
mod policy;
mod request;
mod url;
pub use auth::{
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, 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,
};

View File

@@ -0,0 +1,213 @@
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";
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 Some(project_id) = find_string_by_paths(
&auth_config,
&[
&["project_id"],
&["projectId"],
&["project", "id"],
&["project", "project_id"],
&["project", "projectId"],
&["antigravity", "project_id"],
&["antigravity", "projectId"],
&["metadata", "project_id"],
&["metadata", "projectId"],
],
) else {
return AntigravityRequestAuthSupport::Unsupported(
AntigravityRequestAuthUnsupportedReason::MissingProjectId,
);
};
let client_version = find_string_by_paths(
&auth_config,
&[
&["client_version"],
&["clientVersion"],
&["antigravity", "client_version"],
&["antigravity", "clientVersion"],
&["metadata", "client_version"],
&["metadata", "clientVersion"],
],
);
let session_id = find_string_by_paths(
&auth_config,
&[
&["session_id"],
&["sessionId"],
&["antigravity", "session_id"],
&["antigravity", "sessionId"],
&["metadata", "session_id"],
&["metadata", "sessionId"],
],
);
AntigravityRequestAuthSupport::Supported(AntigravityRequestAuth {
project_id,
client_version,
session_id,
})
}
pub fn build_antigravity_static_identity_headers(
auth: &AntigravityRequestAuth,
) -> 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),
),
]);
if let Some(client_version) = auth
.client_version
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
headers.insert(String::from("x-client-version"), client_version.to_string());
}
if let Some(session_id) = auth
.session_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
headers.insert(String::from("x-vscode-sessionid"), session_id.to_string());
}
headers
}
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
}
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"
)
}

View File

@@ -0,0 +1,113 @@
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,
};
#[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 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 !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(ANTIGRAVITY_PROVIDER_TYPE)
{
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::WrongProviderType,
);
}
let endpoint_format = transport.endpoint.api_format.trim();
if !endpoint_format.eq_ignore_ascii_case("gemini:chat")
&& !endpoint_format.eq_ignore_ascii_case("gemini:cli")
{
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 transport.endpoint.header_rules.is_some() {
return AntigravityRequestSideSupport::Unsupported(
AntigravityRequestSideUnsupportedReason::UnsupportedHeaderRules,
);
}
if transport.endpoint.body_rules.is_some() {
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 })
}

View File

@@ -0,0 +1,120 @@
use serde_json::{Map, Value};
use super::auth::{AntigravityRequestAuth, ANTIGRAVITY_REQUEST_USER_AGENT};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AntigravityEnvelopeRequestType {
Agent,
EndpointTest,
}
impl AntigravityEnvelopeRequestType {
fn as_str(self) -> &'static str {
match self {
Self::Agent => "agent",
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,
ComplexEnvelopeTransform,
}
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") {
return Err(AntigravityRequestEnvelopeUnsupportedReason::MissingContents);
}
if contains_blocked_request_features(request_body) {
return Err(AntigravityRequestEnvelopeUnsupportedReason::ComplexEnvelopeTransform);
}
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,
);
};
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 contains_blocked_request_features(value: &Value) -> bool {
match value {
Value::Object(map) => map.iter().any(|(key, inner)| {
is_blocked_request_key(key.as_str()) || contains_blocked_request_features(inner)
}),
Value::Array(items) => items.iter().any(contains_blocked_request_features),
_ => false,
}
}
fn is_blocked_request_key(key: &str) -> bool {
matches!(
key.trim(),
"systemInstruction"
| "system_instruction"
| "tools"
| "toolConfig"
| "tool_config"
| "thinkingConfig"
| "thinking_config"
| "imageConfig"
| "image_config"
| "functionCall"
| "function_call"
| "functionResponse"
| "function_response"
)
}

View File

@@ -0,0 +1,69 @@
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") {
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)
}

View File

@@ -0,0 +1,144 @@
use std::collections::BTreeMap;
use super::headers::should_skip_upstream_passthrough_header;
use super::snapshot::GatewayProviderTransportSnapshot;
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 value = value.trim();
if value.is_empty() {
continue;
}
out.insert(key, value.to_string());
}
for (key, value) in extra_headers {
let normalized_key = key.to_ascii_lowercase();
let value = value.trim();
if value.is_empty() {
continue;
}
out.insert(normalized_key, value.to_string());
}
out
}
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_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());
}
}
pub fn resolve_local_openai_chat_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
let auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
if !matches!(auth_type.as_str(), "api_key" | "bearer") {
return None;
}
let secret = transport.key.decrypted_api_key.trim();
if secret.is_empty() {
return None;
}
Some(("authorization".to_string(), format!("Bearer {secret}")))
}
pub fn resolve_local_standard_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
let auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
let secret = transport.key.decrypted_api_key.trim();
if secret.is_empty() {
return None;
}
match auth_type.as_str() {
"api_key" => Some(("x-api-key".to_string(), secret.to_string())),
"bearer" => Some(("authorization".to_string(), format!("Bearer {secret}"))),
_ => None,
}
}
pub fn resolve_local_gemini_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
let auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
let secret = transport.key.decrypted_api_key.trim();
if secret.is_empty() {
return None;
}
match auth_type.as_str() {
"api_key" => Some(("x-goog-api-key".to_string(), secret.to_string())),
"bearer" => Some(("authorization".to_string(), format!("Bearer {secret}"))),
_ => None,
}
}

View File

@@ -0,0 +1,600 @@
use std::collections::BTreeMap;
use serde_json::Value;
use url::form_urlencoded;
const UNSAFE_AUTH_CONFIG_HEADER_NAMES: &[&str] = &[
"api-key",
"authorization",
"content-length",
"content-type",
"cookie",
"host",
"proxy-authorization",
"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",
"email",
"model_regions",
"organizations",
"plan_type",
"project_id",
"provider_type",
"region",
"tier",
"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 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)?;
Ok(LocalAuthConfigSafeSubset {
headers,
query,
path,
})
}
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_string_map(headers, value, normalize_auth_config_header_name)?
}
"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 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, 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":{"authorization":"Bearer x"}}"#),
),
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 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": "user@example.com",
"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": "user@example.com",
"plan_type": "plus",
"workspace_name": "demo"
}"#,
),
),
LocalAuthConfigAbsorption::Unsupported
);
}
}

View File

@@ -0,0 +1,129 @@
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,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: Some(1),
proxy: None,
fingerprint: 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(
&current, &refreshed
));
refreshed.key.decrypted_api_key = "sk-updated".to_string();
assert!(provider_transport_snapshot_looks_refreshed(
&current, &refreshed
));
let mut refreshed = current.clone();
refreshed.key.decrypted_auth_config = Some("{\"token\":\"y\"}".to_string());
assert!(provider_transport_snapshot_looks_refreshed(
&current, &refreshed
));
let mut refreshed = current.clone();
refreshed.key.expires_at_unix_secs = Some(2);
assert!(provider_transport_snapshot_looks_refreshed(
&current, &refreshed
));
}
}

View File

@@ -0,0 +1,9 @@
mod auth;
mod policy;
mod request;
mod url;
pub use auth::supports_local_claude_code_auth;
pub use policy::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;

View File

@@ -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()
|| supports_local_oauth_request_auth_resolution(transport)
}

View File

@@ -0,0 +1,53 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::super::{
body_rules_are_locally_supported, header_rules_are_locally_supported,
resolve_transport_tls_profile, supports_local_oauth_request_auth_resolution,
transport_proxy_is_locally_supported,
};
use super::auth::supports_local_claude_code_auth;
pub fn supports_local_claude_code_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> 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("claude_code")
{
return false;
}
if !transport
.endpoint
.api_format
.trim()
.eq_ignore_ascii_case(api_format.trim())
{
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 !supports_local_claude_code_auth(transport) {
return false;
}
if transport.key.decrypted_auth_config.is_some()
&& !supports_local_oauth_request_auth_resolution(transport)
{
return false;
}
if !transport_proxy_is_locally_supported(transport) {
return false;
}
if transport.key.fingerprint.is_some() && resolve_transport_tls_profile(transport).is_none() {
return false;
}
true
}

View File

@@ -0,0 +1,330 @@
use std::collections::{BTreeMap, BTreeSet};
use serde_json::{Map, Value};
use super::super::auth::build_openai_passthrough_headers;
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"];
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"),
);
out.insert("accept".to_string(), DEFAULT_ACCEPT.to_string());
out.insert(
"anthropic-version".to_string(),
DEFAULT_ANTHROPIC_VERSION.to_string(),
);
out.insert(
"anthropic-beta".to_string(),
merge_anthropic_beta_tokens(out.get("anthropic-beta").map(String::as_str)),
);
out.insert("x-stainless-lang".to_string(), "js".to_string());
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".to_string(), "node".to_string());
out.insert(
"x-stainless-runtime-version".to_string(),
"v24.13.0".to_string(),
);
out.insert("x-stainless-retry-count".to_string(), "0".to_string());
out.insert("x-stainless-timeout".to_string(), "600".to_string());
out.insert("x-app".to_string(), "cli".to_string());
out.insert(
"anthropic-dangerous-direct-browser-access".to_string(),
"true".to_string(),
);
if stream {
out.insert(
"x-stainless-helper-method".to_string(),
STREAM_HELPER_METHOD.to_string(),
);
} else {
out.remove("x-stainless-helper-method");
}
if let Some(fingerprint) = fingerprint.and_then(Value::as_object) {
override_header_from_fingerprint(
&mut out,
fingerprint,
"stainless_package_version",
"x-stainless-package-version",
);
override_header_from_fingerprint(&mut out, fingerprint, "stainless_os", "x-stainless-os");
override_header_from_fingerprint(
&mut out,
fingerprint,
"stainless_arch",
"x-stainless-arch",
);
override_header_from_fingerprint(
&mut out,
fingerprint,
"stainless_runtime_version",
"x-stainless-runtime-version",
);
override_header_from_fingerprint(
&mut out,
fingerprint,
"stainless_timeout",
"x-stainless-timeout",
);
override_header_from_fingerprint(&mut out, fingerprint, "user_agent", "user-agent");
}
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());
}
}
fn override_header_from_fingerprint(
headers: &mut BTreeMap<String, String>,
fingerprint: &Map<String, Value>,
fingerprint_key: &str,
header_key: &str,
) {
let Some(value) = fingerprint
.get(fingerprint_key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return;
};
headers.insert(header_key.to_string(), value.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_merge_required_betas_and_stream_helper() {
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!({
"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"}
])
);
}
}

View File

@@ -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"
);
}
}

View File

@@ -0,0 +1,444 @@
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use serde_json::{json, Value};
use super::oauth_refresh::{
CachedOAuthEntry, LocalOAuthRefreshAdapter, LocalOAuthRefreshError,
LocalResolvedOAuthRequestAuth,
};
use super::snapshot::GatewayProviderTransportSnapshot;
const AUTH_HEADER_NAME: &str = "authorization";
const OAUTH_REFRESH_SKEW_SECS: u64 = 120;
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct GenericOAuthTemplate {
provider_type: &'static str,
token_url: &'static str,
client_id: &'static str,
client_secret: &'static str,
scopes: &'static [&'static str],
uses_json_payload: bool,
}
const GENERIC_OAUTH_TEMPLATES: &[GenericOAuthTemplate] = &[
GenericOAuthTemplate {
provider_type: "claude_code",
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"],
uses_json_payload: true,
},
GenericOAuthTemplate {
provider_type: "codex",
token_url: "https://auth.openai.com/oauth/token",
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
client_secret: "",
scopes: &["openid", "email", "profile", "offline_access"],
uses_json_payload: false,
},
GenericOAuthTemplate {
provider_type: "gemini_cli",
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",
],
uses_json_payload: false,
},
GenericOAuthTemplate {
provider_type: "antigravity",
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",
],
uses_json_payload: false,
},
];
pub fn supports_local_generic_oauth_request_auth_resolution(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
&& template_for_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 token_url_for_template(&self, template: GenericOAuthTemplate) -> String {
self.token_url_overrides
.get(template.provider_type)
.cloned()
.unwrap_or_else(|| template.token_url.to_string())
}
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 base_auth_config(
&self,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Option<Value> {
entry
.and_then(|cached| Self::auth_config_from_entry(transport, cached))
.or_else(|| Self::auth_config_from_transport(transport))
}
fn resolve_direct_header(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<LocalResolvedOAuthRequestAuth> {
if !supports_local_generic_oauth_request_auth_resolution(transport) {
return None;
}
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(
&self,
template: GenericOAuthTemplate,
access_token: &str,
metadata: Value,
expires_at_unix_secs: Option<u64>,
) -> CachedOAuthEntry {
CachedOAuthEntry {
provider_type: template.provider_type.to_string(),
auth_header_name: AUTH_HEADER_NAME.to_string(),
auth_header_value: format!("Bearer {access_token}"),
expires_at_unix_secs,
metadata: Some(metadata),
}
}
}
#[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 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,
client: &reqwest::Client,
transport: &GatewayProviderTransportSnapshot,
entry: Option<&CachedOAuthEntry>,
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
let Some(template) = template_for_provider_type(transport.provider.provider_type.as_str())
else {
return Ok(None);
};
let mut metadata = self
.base_auth_config(transport, entry)
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
let Some(refresh_token) = metadata.get("refresh_token").and_then(non_empty_string) else {
return Ok(None);
};
let token_url = self.token_url_for_template(template);
let scope = (!template.scopes.is_empty()).then(|| template.scopes.join(" "));
let request = client.post(token_url);
let response = if template.uses_json_payload {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
Value::String("refresh_token".to_string()),
),
(
"client_id".to_string(),
Value::String(template.client_id.to_string()),
),
(
"refresh_token".to_string(),
Value::String(refresh_token.clone()),
),
]);
if let Some(scope) = scope.as_ref() {
body.insert("scope".to_string(), Value::String(scope.clone()));
}
request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&Value::Object(body))
.send()
.await
} else {
let mut form = vec![
("grant_type", "refresh_token".to_string()),
("client_id", template.client_id.to_string()),
("refresh_token", refresh_token.clone()),
];
if let Some(scope) = scope.as_ref() {
form.push(("scope", scope.clone()));
}
if !template.client_secret.trim().is_empty() {
form.push(("client_secret", template.client_secret.to_string()));
}
request
.header("Content-Type", "application/x-www-form-urlencoded")
.header("Accept", "application/json")
.form(&form)
.send()
.await
}
.map_err(|source| LocalOAuthRefreshError::Transport {
provider_type: template.provider_type,
source,
})?;
let status = response.status();
let body = response
.text()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
provider_type: template.provider_type,
source,
})?;
if !status.is_success() {
return Err(LocalOAuthRefreshError::HttpStatus {
provider_type: template.provider_type,
status_code: status.as_u16(),
body_excerpt: truncate_body(&body),
});
}
let payload: Value =
serde_json::from_str(&body).map_err(|_| LocalOAuthRefreshError::InvalidResponse {
provider_type: template.provider_type,
message: "generic oauth refresh returned non-json body".to_string(),
})?;
let Some(access_token) = payload.get("access_token").and_then(non_empty_string) else {
return Err(LocalOAuthRefreshError::InvalidResponse {
provider_type: template.provider_type,
message: "generic oauth refresh returned empty access_token".to_string(),
});
};
let expires_at_unix_secs = resolve_expires_at(payload.get("expires_in"));
metadata.insert(
"provider_type".to_string(),
Value::String(template.provider_type.to_string()),
);
metadata.insert("updated_at".to_string(), json!(current_unix_secs()));
if let Some(refresh_token) = payload.get("refresh_token").and_then(non_empty_string) {
metadata.insert("refresh_token".to_string(), Value::String(refresh_token));
}
if let Some(token_type) = payload.get("token_type").and_then(non_empty_string) {
metadata.insert("token_type".to_string(), Value::String(token_type));
}
if let Some(scope) = payload.get("scope").and_then(non_empty_string) {
metadata.insert("scope".to_string(), Value::String(scope));
}
match expires_at_unix_secs {
Some(expires_at_unix_secs) => {
metadata.insert("expires_at".to_string(), json!(expires_at_unix_secs));
}
None => {
metadata.remove("expires_at");
}
}
Ok(Some(self.build_cached_entry(
template,
access_token.as_str(),
Value::Object(metadata),
expires_at_unix_secs,
)))
}
}
fn template_for_provider_type(provider_type: &str) -> Option<GenericOAuthTemplate> {
let normalized = provider_type.trim();
GENERIC_OAUTH_TEMPLATES
.iter()
.find(|template| normalized.eq_ignore_ascii_case(template.provider_type))
.copied()
}
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_soon(auth_config: Option<&Value>) -> bool {
expires_at_requires_refresh(
auth_config
.and_then(|value| value.as_object())
.and_then(|object| object.get("expires_at"))
.and_then(|value| parse_u64_value(Some(value))),
)
}
fn expires_at_requires_refresh(expires_at_unix_secs: Option<u64>) -> bool {
expires_at_unix_secs
.map(|expires_at_unix_secs| {
current_unix_secs() >= expires_at_unix_secs.saturating_sub(OAUTH_REFRESH_SKEW_SECS)
})
.unwrap_or(false)
}
fn resolve_expires_at(expires_in: Option<&Value>) -> Option<u64> {
parse_u64_value(expires_in).map(|expires_in| current_unix_secs().saturating_add(expires_in))
}
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 current_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|value| value.as_secs())
.unwrap_or_default()
}
fn truncate_body(body: &str) -> String {
let body = body.trim();
if body.is_empty() {
return String::from("-");
}
body.chars().take(500).collect()
}

View File

@@ -0,0 +1,41 @@
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-control-execute-fallback"
| "x-aether-rate-limit-preflight"
)
}
pub fn should_skip_upstream_passthrough_header(name: &str) -> bool {
matches!(
name.to_ascii_lowercase().as_str(),
"authorization"
| "x-api-key"
| "x-goog-api-key"
| "host"
| "content-length"
| "transfer-encoding"
| "connection"
| "accept-encoding"
| "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)
}

View File

@@ -0,0 +1,31 @@
mod auth;
mod converter;
mod credentials;
mod headers;
mod policy;
mod refresh;
mod request;
mod url;
pub use auth::{
build_kiro_request_auth_from_config, resolve_local_kiro_bearer_auth,
resolve_local_kiro_request_auth, supports_local_kiro_auth_prerequisites,
supports_local_kiro_request_auth_resolution, KiroBearerAuth, KiroRequestAuth, KIRO_AUTH_HEADER,
PROVIDER_TYPE,
};
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, AWS_EVENTSTREAM_CONTENT_TYPE};
pub use policy::{
supports_local_kiro_request_transport, supports_local_kiro_request_transport_with_network,
};
pub use refresh::KiroOAuthRefreshAdapter;
pub use request::{
apply_local_body_rules, apply_local_header_rules, body_rules_are_locally_supported,
build_kiro_provider_headers, build_kiro_provider_request_body,
header_rules_are_locally_supported, supports_local_kiro_request_shape,
};
pub use url::{
build_kiro_generate_assistant_response_url, resolve_kiro_base_url,
GENERATE_ASSISTANT_RESPONSE_PATH, KIRO_ENVELOPE_NAME,
};

View File

@@ -0,0 +1,311 @@
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 build_kiro_request_auth_from_config(
auth_config: KiroAuthConfig,
fallback_secret: Option<&str>,
) -> Option<KiroRequestAuth> {
let fallback_secret = fallback_secret
.map(str::trim)
.filter(|value| !value.is_empty() && *value != "__placeholder__");
let token = auth_config
.cached_access_token()
.filter(|_| !auth_config.cached_access_token_requires_refresh(120))
.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 !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(PROVIDER_TYPE)
{
return None;
}
if transport.key.decrypted_auth_config.is_some() {
return None;
}
if !transport
.key
.auth_type
.trim()
.eq_ignore_ascii_case("bearer")
{
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 !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(PROVIDER_TYPE)
{
return None;
}
if !transport
.key
.auth_type
.trim()
.eq_ignore_ascii_case("bearer")
{
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| {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(PROVIDER_TYPE)
&& transport
.key
.auth_type
.trim()
.eq_ignore_ascii_case("bearer")
&& auth_config.can_refresh_access_token()
})
}
#[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:cli".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:cli".to_string()]),
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: "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_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",
"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 falls_back_to_decrypted_api_key_when_cached_access_token_is_expired() {
let mut transport = sample_transport();
transport.key.decrypted_api_key = "live-upstream-token".to_string();
transport.key.decrypted_auth_config = Some(
r#"{
"access_token":"expired-token",
"expires_at": 1,
"refresh_token":"rrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrrr"
}"#
.to_string(),
);
let auth = resolve_local_kiro_request_auth(&transport)
.expect("request auth should fall back to decrypted api key");
assert_eq!(auth.value, "Bearer live-upstream-token");
}
#[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));
}
}

View File

@@ -0,0 +1,716 @@
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)
);
}
}

View File

@@ -0,0 +1,436 @@
use serde_json::Value;
use sha2::{Digest, Sha256};
use std::time::{SystemTime, UNIX_EPOCH};
pub const DEFAULT_REGION: &str = "us-east-1";
pub const DEFAULT_KIRO_VERSION: &str = "0.8.0";
pub const DEFAULT_NODE_VERSION: &str = "22.21.1";
pub const DEFAULT_SYSTEM_VERSION: &str = "other#unknown";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct KiroAuthConfig {
pub auth_method: Option<String>,
pub refresh_token: Option<String>,
pub expires_at: Option<u64>,
pub profile_arn: Option<String>,
pub region: Option<String>,
pub auth_region: Option<String>,
pub api_region: Option<String>,
pub client_id: Option<String>,
pub client_secret: Option<String>,
pub machine_id: Option<String>,
pub kiro_version: Option<String>,
pub system_version: Option<String>,
pub node_version: Option<String>,
pub access_token: Option<String>,
}
impl KiroAuthConfig {
pub fn from_raw_json(raw: Option<&str>) -> Option<Self> {
let raw = raw?.trim();
if raw.is_empty() {
return None;
}
let parsed: Value = serde_json::from_str(raw).ok()?;
Self::from_json_value(&parsed)
}
pub fn from_json_value(raw: &Value) -> Option<Self> {
let object = raw.as_object()?;
Some(Self {
auth_method: get_nonempty_string(
object,
&["auth_method", "authMethod", "auth_type", "authType"],
)
.map(|value| normalize_auth_method(&value)),
refresh_token: get_nonempty_string(object, &["refresh_token", "refreshToken"]),
expires_at: get_epoch_seconds(object.get("expires_at"))
.or_else(|| get_epoch_seconds(object.get("expiresAt"))),
profile_arn: get_nonempty_string(object, &["profile_arn", "profileArn"]),
region: get_nonempty_string(object, &["region"]),
auth_region: get_nonempty_string(object, &["auth_region", "authRegion"]),
api_region: get_nonempty_string(object, &["api_region", "apiRegion"]),
client_id: get_nonempty_string(object, &["client_id", "clientId"]),
client_secret: get_nonempty_string(object, &["client_secret", "clientSecret"]),
machine_id: get_nonempty_string(object, &["machine_id", "machineId"]),
kiro_version: get_nonempty_string(object, &["kiro_version", "kiroVersion"]),
system_version: get_nonempty_string(object, &["system_version", "systemVersion"]),
node_version: get_nonempty_string(object, &["node_version", "nodeVersion"]),
access_token: get_nonempty_string(object, &["access_token", "accessToken"]),
})
}
pub fn to_json_value(&self) -> Value {
let mut object = serde_json::Map::new();
insert_optional_string(&mut object, "auth_method", self.auth_method.as_deref());
insert_optional_string(&mut object, "refresh_token", self.refresh_token.as_deref());
if let Some(expires_at) = self.expires_at {
object.insert("expires_at".to_string(), Value::from(expires_at));
}
insert_optional_string(&mut object, "profile_arn", self.profile_arn.as_deref());
insert_optional_string(&mut object, "region", self.region.as_deref());
insert_optional_string(&mut object, "auth_region", self.auth_region.as_deref());
insert_optional_string(&mut object, "api_region", self.api_region.as_deref());
insert_optional_string(&mut object, "client_id", self.client_id.as_deref());
insert_optional_string(&mut object, "client_secret", self.client_secret.as_deref());
insert_optional_string(&mut object, "machine_id", self.machine_id.as_deref());
insert_optional_string(&mut object, "kiro_version", self.kiro_version.as_deref());
insert_optional_string(
&mut object,
"system_version",
self.system_version.as_deref(),
);
insert_optional_string(&mut object, "node_version", self.node_version.as_deref());
insert_optional_string(&mut object, "access_token", self.access_token.as_deref());
Value::Object(object)
}
pub fn effective_api_region(&self) -> &str {
self.api_region
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_REGION)
}
pub fn effective_auth_region(&self) -> &str {
self.auth_region
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.or_else(|| {
self.region
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
})
.unwrap_or(DEFAULT_REGION)
}
pub fn effective_kiro_version(&self) -> &str {
self.kiro_version
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_KIRO_VERSION)
}
pub fn effective_system_version(&self) -> &str {
self.system_version
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_SYSTEM_VERSION)
}
pub fn effective_node_version(&self) -> &str {
self.node_version
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_NODE_VERSION)
}
pub fn cached_access_token(&self) -> Option<&str> {
self.access_token
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
}
pub fn cached_access_token_requires_refresh(&self, skew_seconds: u64) -> bool {
let Some(expires_at) = self.expires_at else {
return false;
};
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|value| value.as_secs())
.unwrap_or_default();
now >= expires_at.saturating_sub(skew_seconds)
}
pub fn is_idc_auth(&self) -> bool {
let explicit_method = self
.auth_method
.as_deref()
.map(normalize_auth_method)
.unwrap_or_else(|| "social".to_string());
if explicit_method != "social" {
return explicit_method == "idc";
}
self.client_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.is_some()
&& self
.client_secret
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.is_some()
}
pub fn profile_arn_for_payload(&self) -> Option<&str> {
if self.is_idc_auth() {
return None;
}
self.profile_arn
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
}
pub fn can_refresh_access_token(&self) -> bool {
let refresh_token = self
.refresh_token
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|value| value.len() >= 100 && !value.contains("..."));
if refresh_token.is_none() {
return false;
}
if !self.is_idc_auth() {
return true;
}
self.client_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.is_some()
&& self
.client_secret
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.is_some()
}
}
pub fn normalize_machine_id(raw: &str) -> Option<String> {
let raw = raw.trim();
if raw.is_empty() {
return None;
}
if raw.len() == 64 && raw.bytes().all(|byte| byte.is_ascii_hexdigit()) {
return Some(raw.to_ascii_lowercase());
}
if raw.len() == 36
&& raw.chars().enumerate().all(|(idx, ch)| match idx {
8 | 13 | 18 | 23 => ch == '-',
_ => ch.is_ascii_hexdigit(),
})
{
let normalized = raw.replace('-', "").to_ascii_lowercase();
return Some(format!("{normalized}{normalized}"));
}
None
}
pub fn generate_machine_id(
auth_config: &KiroAuthConfig,
fallback_secret: Option<&str>,
) -> Option<String> {
if let Some(machine_id) = auth_config
.machine_id
.as_deref()
.and_then(normalize_machine_id)
{
return Some(machine_id);
}
let seed = auth_config
.refresh_token
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.or_else(|| {
fallback_secret
.map(str::trim)
.filter(|value| !value.is_empty())
})?;
let mut hasher = Sha256::new();
hasher.update(b"KotlinNativeAPI/");
hasher.update(seed.as_bytes());
Some(format!("{:x}", hasher.finalize()))
}
fn get_nonempty_string(object: &serde_json::Map<String, Value>, keys: &[&str]) -> Option<String> {
keys.iter()
.find_map(|key| object.get(*key))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn insert_optional_string(
object: &mut serde_json::Map<String, Value>,
key: &str,
value: Option<&str>,
) {
let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) else {
return;
};
object.insert(key.to_string(), Value::String(value.to_string()));
}
fn get_epoch_seconds(value: Option<&Value>) -> Option<u64> {
match value? {
Value::Number(number) => number.as_u64().or_else(|| {
number
.as_i64()
.and_then(|value| (value >= 0).then_some(value as u64))
}),
Value::String(text) => text.trim().parse::<u64>().ok(),
_ => None,
}
}
fn normalize_auth_method(raw: &str) -> String {
let value = raw.trim().to_ascii_lowercase();
match value.as_str() {
"" => "social".to_string(),
"builder-id"
| "builder_id"
| "builderid"
| "device"
| "device-auth"
| "device_authorization"
| "iam"
| "identity-center"
| "identity_center"
| "identitycenter"
| "idc" => "idc".to_string(),
_ => value,
}
}
#[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 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.profile_arn_for_payload().is_none());
}
#[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);
}
}

View File

@@ -0,0 +1,122 @@
use std::collections::BTreeMap;
use uuid::Uuid;
use super::credentials::KiroAuthConfig;
pub const AWS_EVENTSTREAM_CONTENT_TYPE: &str = "application/vnd.amazon.eventstream";
const AWS_SDK_JS_MAIN_VERSION: &str = "1.0.27";
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(),
),
])
}
#[cfg(test)]
mod tests {
use super::super::credentials::KiroAuthConfig;
use super::{build_generate_assistant_headers, AWS_EVENTSTREAM_CONTENT_TYPE};
#[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")
);
}
}

View File

@@ -0,0 +1,133 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::super::{resolve_transport_tls_profile, transport_proxy_is_locally_supported};
use super::{supports_local_kiro_request_auth_resolution, supports_local_kiro_request_shape};
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
.endpoint
.api_format
.trim()
.eq_ignore_ascii_case("claude:cli")
{
return false;
}
if !supports_local_kiro_request_auth_resolution(transport) {
return false;
}
if !supports_local_kiro_request_shape(
transport.endpoint.header_rules.as_ref(),
transport.endpoint.body_rules.as_ref(),
) {
return false;
}
true
}
pub fn supports_local_kiro_request_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
supports_local_kiro_request_transport(transport)
&& transport_proxy_is_locally_supported(transport)
&& (transport.key.fingerprint.is_none()
|| resolve_transport_tls_profile(transport).is_some())
}
#[cfg(test)]
mod tests {
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
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:cli".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:cli".to_string()]),
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: 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
));
}
}

View File

@@ -0,0 +1,702 @@
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use serde_json::{json, Value};
use super::super::oauth_refresh::{
CachedOAuthEntry, LocalOAuthRefreshAdapter, LocalOAuthRefreshError,
LocalResolvedOAuthRequestAuth,
};
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::auth::{
build_kiro_request_auth_from_config, resolve_local_kiro_request_auth, PROVIDER_TYPE,
};
use super::credentials::{generate_machine_id, KiroAuthConfig};
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,
client: &reqwest::Client,
auth_config: &KiroAuthConfig,
) -> Result<KiroAuthConfig, LocalOAuthRefreshError> {
if auth_config.is_idc_auth() {
self.refresh_idc_token(client, auth_config).await
} else {
self.refresh_social_token(client, auth_config).await
}
}
fn social_refresh_url(&self, auth_config: &KiroAuthConfig) -> String {
if let Some(base_url) = self
.social_refresh_base_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
return format!("{}/refreshToken", base_url.trim_end_matches('/'));
}
let region = auth_config.effective_auth_region();
format!("https://prod.{region}.auth.desktop.kiro.dev/refreshToken")
}
fn idc_refresh_url(&self, auth_config: &KiroAuthConfig) -> String {
if let Some(base_url) = self
.idc_refresh_base_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
return format!("{}/token", base_url.trim_end_matches('/'));
}
let region = auth_config.effective_auth_region();
format!("https://oidc.{region}.amazonaws.com/token")
}
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()),
})
}
async fn refresh_social_token(
&self,
client: &reqwest::Client,
auth_config: &KiroAuthConfig,
) -> Result<KiroAuthConfig, LocalOAuthRefreshError> {
let url = self.social_refresh_url(auth_config);
let host = reqwest::Url::parse(&url)
.ok()
.and_then(|value| value.host_str().map(ToOwned::to_owned))
.unwrap_or_else(|| {
format!(
"prod.{}.auth.desktop.kiro.dev",
auth_config.effective_auth_region()
)
});
let machine_id = generate_machine_id(auth_config, None).ok_or_else(|| {
LocalOAuthRefreshError::InvalidResponse {
provider_type: PROVIDER_TYPE,
message: "missing machine_id seed for social refresh".to_string(),
}
})?;
let kiro_version = auth_config.effective_kiro_version();
let user_agent = build_kiro_ide_tag(kiro_version, &machine_id);
let response = client
.post(url)
.header("User-Agent", user_agent)
.header("Host", host)
.header("Accept", "application/json, text/plain, */*")
.header("Content-Type", "application/json")
.header("Connection", "close")
.header("Accept-Encoding", "gzip, compress, deflate, br")
.json(&json!({
"refreshToken": auth_config
.refresh_token
.as_deref()
.map(str::trim)
.unwrap_or_default()
}))
.send()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
provider_type: PROVIDER_TYPE,
source,
})?;
let status = response.status();
let body = response
.text()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
provider_type: PROVIDER_TYPE,
source,
})?;
if !status.is_success() {
return Err(LocalOAuthRefreshError::HttpStatus {
provider_type: PROVIDER_TYPE,
status_code: status.as_u16(),
body_excerpt: truncate_body(&body),
});
}
let payload: Value =
serde_json::from_str(&body).map_err(|_| LocalOAuthRefreshError::InvalidResponse {
provider_type: PROVIDER_TYPE,
message: "social refresh returned non-json body".to_string(),
})?;
let access_token = payload
.get("accessToken")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| LocalOAuthRefreshError::InvalidResponse {
provider_type: PROVIDER_TYPE,
message: "social refresh returned empty accessToken".to_string(),
})?;
let mut refreshed = auth_config.clone();
refreshed.access_token = Some(access_token.to_string());
refreshed.expires_at = Some(resolve_expires_at(&payload));
if refreshed
.machine_id
.as_deref()
.map(str::trim)
.is_none_or(|value| value.is_empty())
{
refreshed.machine_id = Some(machine_id);
}
if let Some(refresh_token) = payload
.get("refreshToken")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
refreshed.refresh_token = Some(refresh_token.to_string());
}
if let Some(profile_arn) = payload
.get("profileArn")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
refreshed.profile_arn = Some(profile_arn.to_string());
}
Ok(refreshed)
}
async fn refresh_idc_token(
&self,
client: &reqwest::Client,
auth_config: &KiroAuthConfig,
) -> Result<KiroAuthConfig, LocalOAuthRefreshError> {
let url = self.idc_refresh_url(auth_config);
let host = reqwest::Url::parse(&url)
.ok()
.and_then(|value| value.host_str().map(ToOwned::to_owned))
.unwrap_or_else(|| {
format!("oidc.{}.amazonaws.com", auth_config.effective_auth_region())
});
let response = client
.post(url)
.header("Content-Type", "application/json")
.header("Host", host)
.header("x-amz-user-agent", IDC_AMZ_USER_AGENT)
.header("User-Agent", "node")
.header("Accept", "*/*")
.json(&json!({
"clientId": auth_config
.client_id
.as_deref()
.map(str::trim)
.unwrap_or_default(),
"clientSecret": auth_config
.client_secret
.as_deref()
.map(str::trim)
.unwrap_or_default(),
"refreshToken": auth_config
.refresh_token
.as_deref()
.map(str::trim)
.unwrap_or_default(),
"grantType": "refresh_token"
}))
.send()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
provider_type: PROVIDER_TYPE,
source,
})?;
let status = response.status();
let body = response
.text()
.await
.map_err(|source| LocalOAuthRefreshError::Transport {
provider_type: PROVIDER_TYPE,
source,
})?;
if !status.is_success() {
return Err(LocalOAuthRefreshError::HttpStatus {
provider_type: PROVIDER_TYPE,
status_code: status.as_u16(),
body_excerpt: truncate_body(&body),
});
}
let payload: Value =
serde_json::from_str(&body).map_err(|_| LocalOAuthRefreshError::InvalidResponse {
provider_type: PROVIDER_TYPE,
message: "idc refresh returned non-json body".to_string(),
})?;
let access_token = payload
.get("accessToken")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| LocalOAuthRefreshError::InvalidResponse {
provider_type: PROVIDER_TYPE,
message: "idc refresh returned empty accessToken".to_string(),
})?;
let mut refreshed = auth_config.clone();
refreshed.access_token = Some(access_token.to_string());
refreshed.expires_at = Some(resolve_expires_at(&payload));
if refreshed
.machine_id
.as_deref()
.map(str::trim)
.is_none_or(|value| value.is_empty())
{
refreshed.machine_id = generate_machine_id(auth_config, None);
}
if let Some(refresh_token) = payload
.get("refreshToken")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
refreshed.refresh_token = Some(refresh_token.to_string());
}
Ok(refreshed)
}
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,
client: &reqwest::Client,
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 = if auth_config.is_idc_auth() {
self.refresh_idc_token(client, &auth_config).await?
} else {
self.refresh_social_token(client, &auth_config).await?
};
Ok(Self::build_cached_entry(&refreshed))
}
}
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 resolve_expires_at(payload: &Value) -> u64 {
let expires_in = payload
.get("expiresIn")
.and_then(|value| {
value
.as_u64()
.or_else(|| value.as_str()?.parse::<u64>().ok())
})
.unwrap_or(3600);
current_unix_secs().saturating_add(expires_in)
}
fn current_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|value| value.as_secs())
.unwrap_or_default()
}
fn truncate_body(body: &str) -> String {
let body = body.trim();
if body.is_empty() {
return String::from("-");
}
body.chars().take(500).collect()
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use super::super::super::oauth_refresh::{
LocalOAuthRefreshAdapter, LocalResolvedOAuthRequestAuth,
};
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:cli".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:cli".to_string()]),
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: 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 entry = adapter
.refresh(&reqwest::Client::new(), &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 entry = adapter
.refresh(&reqwest::Client::new(), &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();
}
}

View File

@@ -0,0 +1,284 @@
use std::collections::BTreeMap;
use serde_json::{json, Value};
pub use super::super::rules::{
apply_local_body_rules, apply_local_header_rules, 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>,
) -> 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(&mut provider_request_body, body_rules, Some(body_json)) {
return None;
}
Some(provider_request_body)
}
pub fn build_kiro_provider_headers(
headers: &http::HeaderMap,
provider_request_body: &Value,
original_request_body: &Value,
header_rules: Option<&Value>,
auth_header: &str,
auth_value: &str,
auth_config: &KiroAuthConfig,
machine_id: &str,
) -> 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() {
continue;
}
out.insert(key, value.to_string());
}
if !apply_local_header_rules(
&mut out,
header_rules,
&[auth_header, "content-type"],
provider_request_body,
Some(original_request_body),
) {
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,
};
#[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"}
])),
)
.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(
&http::HeaderMap::new(),
&json!({"conversationState": {}}),
&json!({"messages": []}),
Some(&json!([
{"action":"set","key":"accept","value":"text/plain"},
{"action":"set","key":"x-endpoint-tag","value":"kiro-local"}
])),
"authorization",
"Bearer cached-token",
&auth_config,
"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,
)
.expect("payload should build");
assert!(payload.get("profileArn").is_none());
}
}

View File

@@ -0,0 +1,71 @@
use super::super::url::build_passthrough_path_url;
use super::credentials::DEFAULT_REGION;
pub const GENERATE_ASSISTANT_RESPONSE_PATH: &str = "/generateAssistantResponse";
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,
&[],
)
}
#[cfg(test)]
mod tests {
use super::{
build_kiro_generate_assistant_response_url, resolve_kiro_base_url,
GENERATE_ASSISTANT_RESPONSE_PATH, KIRO_ENVELOPE_NAME,
};
#[test]
fn exposes_kiro_request_constants() {
assert_eq!(
GENERATE_ASSISTANT_RESPONSE_PATH,
"/generateAssistantResponse"
);
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"
);
}
}

View File

@@ -0,0 +1,50 @@
pub mod antigravity;
pub mod auth;
mod auth_config;
mod cache;
pub mod claude_code;
mod generic_oauth;
mod headers;
pub mod kiro;
mod network;
pub mod oauth_refresh;
pub mod policy;
pub mod provider_types;
pub mod rules;
pub mod snapshot;
pub mod url;
pub mod vertex;
mod video;
pub use auth::{build_passthrough_headers, ensure_upstream_auth_header};
pub use cache::{provider_transport_snapshot_looks_refreshed, ProviderTransportSnapshotCacheKey};
pub use generic_oauth::{
supports_local_generic_oauth_request_auth_resolution, GenericOAuthRefreshAdapter,
};
pub use headers::{should_skip_request_header, should_skip_upstream_passthrough_header};
pub use network::{
resolve_transport_execution_timeouts, resolve_transport_proxy_snapshot,
resolve_transport_proxy_snapshot_with_tunnel_affinity, resolve_transport_tls_profile,
transport_proxy_is_locally_supported, TransportTunnelAffinityLookup,
TransportTunnelAttachmentOwner,
};
pub use oauth_refresh::{
supports_local_oauth_request_auth_resolution, CachedOAuthEntry, LocalOAuthRefreshCoordinator,
LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth,
};
pub use policy::{
supports_local_gemini_transport, supports_local_gemini_transport_with_network,
supports_local_standard_transport,
};
pub use rules::{
apply_local_body_rules, apply_local_header_rules, body_rules_are_locally_supported,
header_rules_are_locally_supported,
};
pub use snapshot::{
read_provider_transport_snapshot, GatewayProviderTransportSnapshot,
ProviderTransportSnapshotSource,
};
pub use video::{
reconstruct_local_video_task_snapshot, resolve_local_video_task_transport,
VideoTaskTransportSnapshotLookup,
};

View File

@@ -0,0 +1,396 @@
use aether_contracts::{ExecutionTimeouts, ProxySnapshot};
use async_trait::async_trait;
use serde_json::{json, Map, Value};
use tracing::warn;
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";
#[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> {
let total_ms = transport
.provider
.request_timeout_secs
.filter(|value| value.is_finite() && *value > 0.0)
.map(|value| (value * 1000.0).round() as u64);
let first_byte_ms = transport
.provider
.stream_first_byte_timeout_secs
.filter(|value| value.is_finite() && *value > 0.0)
.map(|value| (value * 1000.0).round() as u64);
if total_ms.is_none() && first_byte_ms.is_none() {
return None;
}
Some(ExecutionTimeouts {
total_ms,
first_byte_ms,
..ExecutionTimeouts::default()
})
}
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_tls_profile(
transport: &GatewayProviderTransportSnapshot,
) -> Option<String> {
transport
.key
.fingerprint
.as_ref()
.and_then(|value| value.get("tls_profile"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn effective_proxy_config(transport: &GatewayProviderTransportSnapshot) -> Option<&Value> {
for candidate in [
transport.key.proxy.as_ref(),
transport.endpoint.proxy.as_ref(),
transport.provider.proxy.as_ref(),
]
.into_iter()
.flatten()
{
if proxy_enabled(candidate) {
return Some(candidate);
}
}
None
}
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_proxy_snapshot, resolve_transport_proxy_snapshot_with_tunnel_affinity,
resolve_transport_tls_profile, transport_proxy_is_locally_supported,
TransportTunnelAffinityLookup, TransportTunnelAttachmentOwner,
};
#[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,
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!({"tls_profile":"chrome_136"})),
decrypted_api_key: "sk-test".to_string(),
decrypted_auth_config: None,
},
}
}
#[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_tls_profile_from_key_fingerprint() {
assert_eq!(
resolve_transport_tls_profile(&sample_transport()).as_deref(),
Some("chrome_136")
);
assert!(transport_proxy_is_locally_supported(&sample_transport()));
}
}

View File

@@ -0,0 +1,470 @@
use std::collections::BTreeMap;
use std::fmt;
use std::sync::Arc;
use aether_data::redis::{RedisLockKey, RedisLockRunner};
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;
#[derive(Debug, Clone, PartialEq, Eq)]
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, 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 returned invalid response: {message}")]
InvalidResponse {
provider_type: &'static str,
message: String,
},
}
#[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,
client: &reqwest::Client,
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(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 resolve_with_result(
&self,
client: &reqwest::Client,
transport: &GatewayProviderTransportSnapshot,
distributed_lock: Option<&RedisLockRunner>,
distributed_owner: Option<&str>,
) -> 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 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 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() => {
let lock_key = RedisLockKey(format!("provider_oauth_refresh_lock:{key_id}"));
match lock
.try_acquire(
&lock_key,
owner,
Some(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,
};
let refresh_result = adapter
.refresh(client, transport, cached_entry.as_ref())
.await;
if let (Some(lock), Some(lease)) = (distributed_lock, distributed_lease.as_ref()) {
if let Err(err) = 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);
};
self.insert_cached_entry(key_id, refreshed_entry.clone())
.await;
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_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, LocalOAuthRefreshAdapter, LocalOAuthRefreshCoordinator,
LocalOAuthRefreshError, LocalOAuthResolution, LocalResolvedOAuthRequestAuth,
};
use async_trait::async_trait;
use std::sync::Arc;
#[derive(Debug)]
struct TestAdapter {
refresh_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,
_client: &reqwest::Client,
_transport: &GatewayProviderTransportSnapshot,
_entry: Option<&CachedOAuthEntry>,
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
self.refresh_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:cli".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,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "__placeholder__".to_string(),
decrypted_auth_config: 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 coordinator =
LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(TestAdapter {
refresh_hits: Arc::clone(&refresh_hits),
})]);
let transport = sample_transport();
let client = reqwest::Client::new();
let first = coordinator
.resolve_with_result(&client, &transport, None, None)
.await
.expect("first resolve should succeed");
let second = coordinator
.resolve_with_result(&client, &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,
})
);
}
}

View File

@@ -0,0 +1,137 @@
use super::provider_types::{
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_tls_profile, supports_local_oauth_request_auth_resolution,
transport_proxy_is_locally_supported,
};
pub fn supports_local_openai_chat_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return false;
}
if !transport
.endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:chat")
{
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 transport.key.decrypted_auth_config.is_some()
&& !supports_local_oauth_request_auth_resolution(transport)
{
return false;
}
if !transport_proxy_is_locally_supported(transport) {
return false;
}
if transport.key.fingerprint.is_some() && resolve_transport_tls_profile(transport).is_none() {
return false;
}
if !provider_type_supports_local_openai_chat_transport(&transport.provider.provider_type) {
return false;
}
true
}
pub fn supports_local_standard_transport(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
supports_local_same_format_transport(transport, api_format, false)
}
pub fn supports_local_gemini_transport(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
supports_local_same_format_transport(transport, api_format, false)
}
pub fn supports_local_standard_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
supports_local_same_format_transport(transport, api_format, true)
}
pub fn supports_local_gemini_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
supports_local_same_format_transport(transport, api_format, true)
}
fn supports_local_same_format_transport(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
allow_network_passthrough: bool,
) -> bool {
if !transport.provider.is_active || !transport.endpoint.is_active || !transport.key.is_active {
return false;
}
if !transport
.endpoint
.api_format
.trim()
.eq_ignore_ascii_case(api_format.trim())
{
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 transport.key.decrypted_auth_config.is_some()
&& !supports_local_oauth_request_auth_resolution(transport)
{
return false;
}
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 false;
}
if allow_network_passthrough {
if !transport_proxy_is_locally_supported(transport) {
return false;
}
if transport.key.fingerprint.is_some() && resolve_transport_tls_profile(transport).is_none()
{
return false;
}
} else if transport.provider.proxy.is_some()
|| transport.endpoint.proxy.is_some()
|| transport.key.proxy.is_some()
|| transport
.key
.fingerprint
.as_ref()
.and_then(|value| value.get("tls_profile"))
.and_then(|value| value.as_str())
.is_some_and(|value| !value.trim().is_empty())
{
return false;
}
if !provider_type_supports_local_same_format_transport(&transport.provider.provider_type) {
return false;
}
true
}

View File

@@ -0,0 +1,139 @@
#[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,
}
pub fn provider_type_is_fixed(provider_type: &str) -> bool {
matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"claude_code" | "kiro" | "codex" | "gemini_cli" | "antigravity" | "vertex_ai"
)
}
pub fn provider_type_enables_format_conversion_by_default(provider_type: &str) -> bool {
matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"claude_code" | "kiro" | "codex" | "antigravity" | "vertex_ai"
)
}
pub fn fixed_provider_template(
provider_type: &str,
) -> Option<(&'static str, &'static [&'static str])> {
match provider_type.trim().to_ascii_lowercase().as_str() {
"claude_code" => Some(("https://api.anthropic.com", &["claude:cli"])),
"codex" => Some((
"https://chatgpt.com/backend-api/codex",
&["openai:cli", "openai:compact"],
)),
"kiro" => Some(("https://q.{region}.amazonaws.com", &["claude:cli"])),
"gemini_cli" => Some(("https://cloudcode-pa.googleapis.com", &["gemini:cli"])),
"vertex_ai" => Some((
"https://aiplatform.googleapis.com",
&["gemini:chat", "claude:chat"],
)),
"antigravity" => Some(("https://cloudcode-pa.googleapis.com", &["gemini:chat"])),
_ => None,
}
}
pub fn provider_type_supports_model_fetch(provider_type: &str) -> bool {
!matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"vertex_ai" | "antigravity" | "codex" | "kiro" | "claude_code"
)
}
pub fn provider_type_supports_local_openai_chat_transport(provider_type: &str) -> bool {
!matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"antigravity" | "claude_code" | "codex" | "gemini_cli" | "kiro" | "vertex_ai"
)
}
pub fn provider_type_supports_local_same_format_transport(provider_type: &str) -> bool {
!matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"antigravity" | "claude_code" | "kiro" | "vertex_ai"
)
}
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,
}),
"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,
}),
_ => None,
}
}
pub const ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES: &[&str] =
&["claude_code", "codex", "gemini_cli", "antigravity"];

View File

@@ -0,0 +1,820 @@
use std::collections::{BTreeMap, HashSet};
use regex::Regex;
use serde_json::{Map, Value};
const ORIGINAL_PLACEHOLDER: &str = "{{$original}}";
const CONDITION_SOURCES: &[&str] = &["current", "original"];
const CONDITION_TYPE_VALUES: &[&str] = &["string", "number", "boolean", "array", "object", "null"];
#[derive(Debug, Clone, PartialEq, Eq)]
enum BodyPathSegment {
Key(String),
Index(isize),
}
pub fn header_rules_are_locally_supported(rules: Option<&Value>) -> bool {
let Some(rules) = rules else {
return true;
};
let Some(rules) = rules.as_array() else {
return false;
};
rules.iter().all(|rule| {
let Some(rule) = rule.as_object() else {
return false;
};
if rule
.get("condition")
.is_some_and(|value| !value.is_null() && !condition_is_locally_supported(value))
{
return false;
}
match rule
.get("action")
.and_then(Value::as_str)
.map(str::trim)
.map(str::to_ascii_lowercase)
.as_deref()
{
Some("set") => {
rule.get("key")
.and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty())
&& rule.get("value").is_some_and(Value::is_string)
}
Some("drop") => rule
.get("key")
.and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty()),
Some("rename") => {
rule.get("from")
.and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty())
&& rule
.get("to")
.and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty())
}
_ => false,
}
})
}
pub fn apply_local_header_rules(
headers: &mut BTreeMap<String, String>,
rules: Option<&Value>,
protected_keys: &[&str],
body: &Value,
original_body: Option<&Value>,
) -> bool {
let Some(rules) = rules else {
return true;
};
let Some(rules) = rules.as_array() else {
return false;
};
let protected_keys: HashSet<String> = protected_keys
.iter()
.map(|value| value.trim().to_ascii_lowercase())
.collect();
for rule in rules {
let Some(rule) = rule.as_object() else {
return false;
};
if let Some(condition) = rule.get("condition").filter(|value| !value.is_null()) {
if !condition_is_locally_supported(condition) {
return false;
}
if !evaluate_local_condition(body, condition, original_body) {
continue;
}
}
match rule
.get("action")
.and_then(Value::as_str)
.map(str::trim)
.map(str::to_ascii_lowercase)
.as_deref()
{
Some("set") => {
let Some(key) = rule.get("key").and_then(Value::as_str).map(str::trim) else {
return false;
};
let Some(value) = rule.get("value").and_then(Value::as_str) else {
return false;
};
let key = key.to_ascii_lowercase();
if !protected_keys.contains(&key) {
headers.insert(key, value.to_string());
}
}
Some("drop") => {
let Some(key) = rule.get("key").and_then(Value::as_str).map(str::trim) else {
return false;
};
let key = key.to_ascii_lowercase();
if !protected_keys.contains(&key) {
headers.remove(&key);
}
}
Some("rename") => {
let Some(from) = rule.get("from").and_then(Value::as_str).map(str::trim) else {
return false;
};
let Some(to) = rule.get("to").and_then(Value::as_str).map(str::trim) else {
return false;
};
let from = from.to_ascii_lowercase();
let to = to.to_ascii_lowercase();
if protected_keys.contains(&from) || protected_keys.contains(&to) {
continue;
}
if let Some(value) = headers.remove(&from) {
headers.insert(to, value);
}
}
_ => return false,
}
}
true
}
pub fn body_rules_are_locally_supported(rules: Option<&Value>) -> bool {
let Some(rules) = rules else {
return true;
};
let Some(rules) = rules.as_array() else {
return false;
};
rules.iter().all(|rule| {
let Some(rule) = rule.as_object() else {
return false;
};
if rule
.get("condition")
.is_some_and(|value| !value.is_null() && !condition_is_locally_supported(value))
{
return false;
}
match rule
.get("action")
.and_then(Value::as_str)
.map(str::trim)
.map(str::to_ascii_lowercase)
.as_deref()
{
Some("set") => {
rule.get("path")
.and_then(Value::as_str)
.and_then(parse_body_path)
.is_some()
&& !rule.get("value").is_some_and(contains_original_placeholder)
}
Some("drop") => rule
.get("path")
.and_then(Value::as_str)
.and_then(parse_body_path)
.is_some(),
Some("rename") => {
rule.get("from")
.and_then(Value::as_str)
.and_then(parse_body_path)
.is_some()
&& rule
.get("to")
.and_then(Value::as_str)
.and_then(parse_body_path)
.is_some()
}
_ => false,
}
})
}
pub fn apply_local_body_rules(
body: &mut Value,
rules: Option<&Value>,
original_body: Option<&Value>,
) -> bool {
let Some(rules) = rules else {
return true;
};
let Some(rules) = rules.as_array() else {
return false;
};
for rule in rules {
let Some(rule) = rule.as_object() else {
return false;
};
if let Some(condition) = rule.get("condition").filter(|value| !value.is_null()) {
if !condition_is_locally_supported(condition) {
return false;
}
if !evaluate_local_condition(body, condition, original_body) {
continue;
}
}
match rule
.get("action")
.and_then(Value::as_str)
.map(str::trim)
.map(str::to_ascii_lowercase)
.as_deref()
{
Some("set") => {
let Some(path) = rule
.get("path")
.and_then(Value::as_str)
.and_then(parse_body_path)
else {
return false;
};
let value = rule.get("value").cloned().unwrap_or(Value::Null);
if contains_original_placeholder(&value) {
return false;
}
let _ = set_nested_value(body, &path, value);
}
Some("drop") => {
let Some(path) = rule
.get("path")
.and_then(Value::as_str)
.and_then(parse_body_path)
else {
return false;
};
let _ = delete_nested_value(body, &path);
}
Some("rename") => {
let Some(from) = rule
.get("from")
.and_then(Value::as_str)
.and_then(parse_body_path)
else {
return false;
};
let Some(to) = rule
.get("to")
.and_then(Value::as_str)
.and_then(parse_body_path)
else {
return false;
};
let _ = rename_nested_value(body, &from, &to);
}
_ => return false,
}
}
true
}
fn condition_is_locally_supported(condition: &Value) -> bool {
let Some(condition) = condition.as_object() else {
return false;
};
if let Some(children) = condition.get("all").and_then(Value::as_array) {
return !children.is_empty() && children.iter().all(condition_is_locally_supported);
}
if let Some(children) = condition.get("any").and_then(Value::as_array) {
return !children.is_empty() && children.iter().all(condition_is_locally_supported);
}
let source = condition
.get("source")
.and_then(Value::as_str)
.map(str::trim)
.unwrap_or("current");
if !CONDITION_SOURCES.contains(&source) {
return false;
}
let Some(op) = condition.get("op").and_then(Value::as_str).map(str::trim) else {
return false;
};
let Some(path) = condition
.get("path")
.and_then(Value::as_str)
.map(str::trim)
.and_then(parse_body_path)
else {
return false;
};
if path.is_empty() {
return false;
}
match op {
"exists" | "not_exists" | "eq" | "neq" => true,
"gt" | "lt" | "gte" | "lte" => condition
.get("value")
.is_some_and(|value| value.as_f64().is_some() && !value.is_boolean()),
"starts_with" | "ends_with" | "matches" => condition
.get("value")
.and_then(Value::as_str)
.is_some_and(|value| {
if op == "matches" {
Regex::new(value).is_ok()
} else {
true
}
}),
"contains" => condition.get("value").is_some(),
"in" => condition.get("value").is_some_and(Value::is_array),
"type_is" => condition
.get("value")
.and_then(Value::as_str)
.is_some_and(|value| CONDITION_TYPE_VALUES.contains(&value)),
_ => false,
}
}
fn evaluate_local_condition(
body: &Value,
condition: &Value,
original_body: Option<&Value>,
) -> bool {
let Some(condition) = condition.as_object() else {
return false;
};
if let Some(children) = condition.get("all").and_then(Value::as_array) {
return !children.is_empty()
&& children
.iter()
.all(|child| evaluate_local_condition(body, child, original_body));
}
if let Some(children) = condition.get("any").and_then(Value::as_array) {
return !children.is_empty()
&& children
.iter()
.any(|child| evaluate_local_condition(body, child, original_body));
}
let source = condition
.get("source")
.and_then(Value::as_str)
.map(str::trim)
.unwrap_or("current");
let target = if source.eq_ignore_ascii_case("original") {
original_body.unwrap_or(body)
} else {
body
};
let Some(op) = condition.get("op").and_then(Value::as_str).map(str::trim) else {
return false;
};
let Some(path) = condition
.get("path")
.and_then(Value::as_str)
.map(str::trim)
.and_then(parse_body_path)
else {
return false;
};
let current_value = get_nested_value(target, &path);
if op == "exists" {
return current_value.is_some();
}
if op == "not_exists" {
return current_value.is_none();
}
let Some(current_value) = current_value else {
return false;
};
let expected = condition.get("value");
match op {
"eq" => expected == Some(&current_value),
"neq" => expected != Some(&current_value),
"gt" | "lt" | "gte" | "lte" => {
let Some(current) = json_number(&current_value) else {
return false;
};
let Some(expected) = expected.and_then(json_number) else {
return false;
};
match op {
"gt" => current > expected,
"lt" => current < expected,
"gte" => current >= expected,
"lte" => current <= expected,
_ => false,
}
}
"starts_with" => current_value
.as_str()
.zip(expected.and_then(Value::as_str))
.is_some_and(|(current, expected)| current.starts_with(expected)),
"ends_with" => current_value
.as_str()
.zip(expected.and_then(Value::as_str))
.is_some_and(|(current, expected)| current.ends_with(expected)),
"contains" => match (current_value, expected) {
(Value::String(current), Some(Value::String(expected))) => current.contains(expected),
(Value::Array(current), Some(expected)) => {
current.iter().any(|value| value == expected)
}
_ => false,
},
"matches" => current_value
.as_str()
.zip(expected.and_then(Value::as_str))
.is_some_and(|(current, expected)| {
Regex::new(expected)
.map(|pattern| pattern.is_match(current))
.unwrap_or(false)
}),
"in" => expected
.and_then(Value::as_array)
.is_some_and(|values| values.iter().any(|value| value == &current_value)),
"type_is" => expected
.and_then(Value::as_str)
.is_some_and(|expected| match expected {
"string" => current_value.is_string(),
"number" => current_value.as_f64().is_some() && !current_value.is_boolean(),
"boolean" => current_value.is_boolean(),
"array" => current_value.is_array(),
"object" => current_value.is_object(),
"null" => current_value.is_null(),
_ => false,
}),
_ => false,
}
}
fn json_number(value: &Value) -> Option<f64> {
value.as_f64().filter(|_| !value.is_boolean())
}
fn parse_body_path(path: &str) -> Option<Vec<BodyPathSegment>> {
let raw = path.trim();
if raw.is_empty() {
return None;
}
let chars: Vec<char> = raw.chars().collect();
let mut parts = Vec::new();
let mut current = String::new();
let mut expect_key = true;
let mut index = 0usize;
while index < chars.len() {
let ch = chars[index];
if ch == '\\' && chars.get(index + 1).copied() == Some('.') {
current.push('.');
expect_key = false;
index += 2;
continue;
}
if ch == '.' {
if !current.is_empty() {
parts.push(BodyPathSegment::Key(std::mem::take(&mut current)));
} else if expect_key {
return None;
}
expect_key = true;
index += 1;
continue;
}
if ch == '[' {
if !current.is_empty() {
parts.push(BodyPathSegment::Key(std::mem::take(&mut current)));
}
let mut close_index = index + 1;
while close_index < chars.len() && chars[close_index] != ']' {
close_index += 1;
}
if close_index >= chars.len() {
return None;
}
let inner = chars[index + 1..close_index]
.iter()
.collect::<String>()
.trim()
.to_string();
if inner.is_empty() || inner == "*" {
return None;
}
let Ok(index_value) = inner.parse::<isize>() else {
return None;
};
parts.push(BodyPathSegment::Index(index_value));
expect_key = false;
index = close_index + 1;
continue;
}
current.push(ch);
expect_key = false;
index += 1;
}
if !current.is_empty() {
parts.push(BodyPathSegment::Key(current));
} else if expect_key {
return None;
}
(!parts.is_empty()).then_some(parts)
}
fn contains_original_placeholder(value: &Value) -> bool {
match value {
Value::String(value) => value.contains(ORIGINAL_PLACEHOLDER),
Value::Array(items) => items.iter().any(contains_original_placeholder),
Value::Object(items) => items.values().any(contains_original_placeholder),
_ => false,
}
}
fn resolve_index(len: usize, index: isize) -> Option<usize> {
if index >= 0 {
((index as usize) < len).then_some(index as usize)
} else {
let resolved = len as isize + index;
(resolved >= 0).then_some(resolved as usize)
}
}
fn get_nested_value(value: &Value, path: &[BodyPathSegment]) -> Option<Value> {
let mut current = value;
for segment in path {
match segment {
BodyPathSegment::Key(key) => {
current = current.as_object()?.get(key)?;
}
BodyPathSegment::Index(index) => {
let values = current.as_array()?;
let resolved = resolve_index(values.len(), *index)?;
current = values.get(resolved)?;
}
}
}
Some(current.clone())
}
fn get_existing_child_mut<'a>(
current: &'a mut Value,
segment: &BodyPathSegment,
) -> Option<&'a mut Value> {
match segment {
BodyPathSegment::Key(key) => current.as_object_mut()?.get_mut(key),
BodyPathSegment::Index(index) => {
let values = current.as_array_mut()?;
let resolved = resolve_index(values.len(), *index)?;
values.get_mut(resolved)
}
}
}
fn set_nested_value(current: &mut Value, path: &[BodyPathSegment], value: Value) -> bool {
let Some((last, parents)) = path.split_last() else {
return false;
};
let mut current = current;
for (offset, segment) in parents.iter().enumerate() {
let next = &path[offset + 1];
match segment {
BodyPathSegment::Key(key) => {
let Some(object) = current.as_object_mut() else {
return false;
};
match next {
BodyPathSegment::Key(_) => {
let child = object
.entry(key.clone())
.or_insert_with(|| Value::Object(Map::new()));
if !child.is_object() {
*child = Value::Object(Map::new());
}
current = child;
}
BodyPathSegment::Index(_) => {
let Some(child) = object.get_mut(key) else {
return false;
};
if !child.is_array() {
return false;
}
current = child;
}
}
}
BodyPathSegment::Index(index) => {
let Some(values) = current.as_array_mut() else {
return false;
};
let Some(resolved) = resolve_index(values.len(), *index) else {
return false;
};
current = &mut values[resolved];
}
}
}
match last {
BodyPathSegment::Key(key) => {
let Some(object) = current.as_object_mut() else {
return false;
};
object.insert(key.clone(), value);
true
}
BodyPathSegment::Index(index) => {
let Some(values) = current.as_array_mut() else {
return false;
};
let Some(resolved) = resolve_index(values.len(), *index) else {
return false;
};
values[resolved] = value;
true
}
}
}
fn delete_nested_value(current: &mut Value, path: &[BodyPathSegment]) -> bool {
let Some((last, parents)) = path.split_last() else {
return false;
};
let mut current = current;
for segment in parents {
let Some(child) = get_existing_child_mut(current, segment) else {
return false;
};
current = child;
}
match last {
BodyPathSegment::Key(key) => current
.as_object_mut()
.and_then(|object| object.remove(key))
.is_some(),
BodyPathSegment::Index(index) => {
let Some(values) = current.as_array_mut() else {
return false;
};
let Some(resolved) = resolve_index(values.len(), *index) else {
return false;
};
values.remove(resolved);
true
}
}
}
fn rename_nested_value(
current: &mut Value,
from: &[BodyPathSegment],
to: &[BodyPathSegment],
) -> bool {
if from == to {
return get_nested_value(current, from).is_some();
}
let Some(value) = get_nested_value(current, from) else {
return false;
};
if !set_nested_value(current, to, value) {
return false;
}
delete_nested_value(current, from)
}
#[cfg(test)]
mod tests {
use super::{
apply_local_body_rules, apply_local_header_rules, body_rules_are_locally_supported,
header_rules_are_locally_supported,
};
#[test]
fn header_rules_allow_simple_set_drop_and_rename() {
let rules = serde_json::json!([
{"action":"set","key":"x-added","value":"1"},
{"action":"drop","key":"x-drop"},
{"action":"rename","from":"x-old","to":"x-new"}
]);
assert!(header_rules_are_locally_supported(Some(&rules)));
let mut headers = std::collections::BTreeMap::from([
("x-drop".to_string(), "drop-me".to_string()),
("x-old".to_string(), "old-value".to_string()),
("authorization".to_string(), "Bearer keep".to_string()),
]);
assert!(apply_local_header_rules(
&mut headers,
Some(&rules),
&["authorization", "content-type"],
&serde_json::json!({}),
None,
));
assert_eq!(headers.get("x-added").map(String::as_str), Some("1"));
assert!(!headers.contains_key("x-drop"));
assert_eq!(headers.get("x-new").map(String::as_str), Some("old-value"));
assert_eq!(
headers.get("authorization").map(String::as_str),
Some("Bearer keep")
);
}
#[test]
fn header_rules_allow_simple_conditions() {
let rules = serde_json::json!([
{"action":"set","key":"x-added","value":"1","condition":{"path":"metadata.mode","op":"eq","value":"safe"}},
{"action":"set","key":"x-from-original","value":"1","condition":{"path":"metadata.client","op":"exists","source":"original"}}
]);
assert!(header_rules_are_locally_supported(Some(&rules)));
let mut headers = std::collections::BTreeMap::new();
assert!(apply_local_header_rules(
&mut headers,
Some(&rules),
&[],
&serde_json::json!({"metadata":{"mode":"safe"}}),
Some(&serde_json::json!({"metadata":{"client":"desktop"}})),
));
assert_eq!(headers.get("x-added").map(String::as_str), Some("1"));
assert_eq!(
headers.get("x-from-original").map(String::as_str),
Some("1")
);
}
#[test]
fn body_rules_allow_simple_nested_set_drop_and_rename() {
let rules = serde_json::json!([
{"action":"set","path":"metadata.mode","value":"safe"},
{"action":"drop","path":"tools[1]"},
{"action":"rename","from":"messages[0].content","to":"messages[0].text"}
]);
assert!(body_rules_are_locally_supported(Some(&rules)));
let mut body = serde_json::json!({
"messages": [{"content":"hello"}],
"tools": [{"name":"a"},{"name":"b"}]
});
assert!(apply_local_body_rules(&mut body, Some(&rules), None));
assert_eq!(body["metadata"]["mode"], "safe");
assert_eq!(body["tools"], serde_json::json!([{"name":"a"}]));
assert_eq!(body["messages"][0]["text"], "hello");
assert!(body["messages"][0].get("content").is_none());
}
#[test]
fn body_rules_allow_simple_conditions() {
let rules = serde_json::json!([
{"action":"set","path":"instructions","value":"You are GPT-5.","condition":{"path":"instructions","op":"not_exists"}},
{"action":"set","path":"metadata.origin","value":"desktop","condition":{"path":"metadata.client","op":"exists","source":"original"}}
]);
assert!(body_rules_are_locally_supported(Some(&rules)));
let original = serde_json::json!({
"metadata": {
"client": "desktop"
}
});
let mut body = serde_json::json!({
"metadata": {
"mode": "safe"
}
});
assert!(apply_local_body_rules(
&mut body,
Some(&rules),
Some(&original)
));
assert_eq!(body["instructions"], "You are GPT-5.");
assert_eq!(body["metadata"]["origin"], "desktop");
}
#[test]
fn body_rules_reject_placeholder_and_wildcard_paths() {
let placeholder_rules =
serde_json::json!([{"action":"set","path":"model","value":"{{$original}}"}]);
assert!(!body_rules_are_locally_supported(Some(&placeholder_rules)));
let wildcard_rules = serde_json::json!([{"action":"drop","path":"items[*].value"}]);
assert!(!body_rules_are_locally_supported(Some(&wildcard_rules)));
}
}

View File

@@ -0,0 +1,985 @@
use aether_data::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data::DataLayerError;
use async_trait::async_trait;
use super::auth_config::{absorb_local_auth_config_safe_subset, LocalAuthConfigAbsorption};
#[path = "snapshot_mapping.rs"]
mod snapshot_mapping;
use self::snapshot_mapping::{fallback_encryption_keys, map_endpoint, map_key, map_provider};
#[derive(Debug, Clone, PartialEq, serde::Serialize)]
pub struct GatewayProviderTransportSnapshot {
pub provider: GatewayProviderTransportProvider,
pub endpoint: GatewayProviderTransportEndpoint,
pub key: GatewayProviderTransportKey,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize)]
pub struct GatewayProviderTransportProvider {
pub id: String,
pub name: String,
pub provider_type: String,
pub website: Option<String>,
pub is_active: bool,
pub keep_priority_on_conversion: bool,
pub enable_format_conversion: bool,
pub concurrent_limit: Option<i32>,
pub max_retries: Option<i32>,
pub proxy: Option<serde_json::Value>,
pub request_timeout_secs: Option<f64>,
pub stream_first_byte_timeout_secs: Option<f64>,
pub config: Option<serde_json::Value>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize)]
pub struct GatewayProviderTransportEndpoint {
pub id: String,
pub provider_id: String,
pub api_format: String,
pub api_family: Option<String>,
pub endpoint_kind: Option<String>,
pub is_active: bool,
pub base_url: String,
pub header_rules: Option<serde_json::Value>,
pub body_rules: Option<serde_json::Value>,
pub max_retries: Option<i32>,
pub custom_path: Option<String>,
pub config: Option<serde_json::Value>,
pub format_acceptance_config: Option<serde_json::Value>,
pub proxy: Option<serde_json::Value>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize)]
pub struct GatewayProviderTransportKey {
pub id: String,
pub provider_id: String,
pub name: String,
pub auth_type: String,
pub is_active: bool,
pub api_formats: Option<Vec<String>>,
pub allowed_models: Option<Vec<String>>,
pub capabilities: Option<serde_json::Value>,
pub rate_multipliers: Option<serde_json::Value>,
pub global_priority_by_format: Option<serde_json::Value>,
pub expires_at_unix_secs: Option<u64>,
pub proxy: Option<serde_json::Value>,
pub fingerprint: Option<serde_json::Value>,
pub decrypted_api_key: String,
pub decrypted_auth_config: Option<String>,
}
#[async_trait]
pub trait ProviderTransportSnapshotSource: Send + Sync {
fn encryption_key(&self) -> Option<&str>;
async fn list_provider_catalog_providers_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError>;
async fn list_provider_catalog_endpoints_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError>;
async fn list_provider_catalog_keys_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError>;
}
pub async fn read_provider_transport_snapshot(
state: &dyn ProviderTransportSnapshotSource,
provider_id: &str,
endpoint_id: &str,
key_id: &str,
) -> Result<Option<GatewayProviderTransportSnapshot>, DataLayerError> {
let Some(encryption_key) = state.encryption_key() else {
return Ok(None);
};
let fallback_encryption_keys = fallback_encryption_keys(encryption_key);
let providers = state
.list_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await?;
let endpoints = state
.list_provider_catalog_endpoints_by_ids(&[endpoint_id.to_string()])
.await?;
let keys = state
.list_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await?;
let Some(provider) = providers.into_iter().next() else {
return Ok(None);
};
let Some(endpoint) = endpoints.into_iter().next() else {
return Ok(None);
};
let Some(key) = keys.into_iter().next() else {
return Ok(None);
};
if endpoint.provider_id != provider.id {
return Err(DataLayerError::UnexpectedValue(format!(
"provider_endpoints.provider_id mismatch: expected {}, got {}",
provider.id, endpoint.provider_id
)));
}
if key.provider_id != provider.id {
return Err(DataLayerError::UnexpectedValue(format!(
"provider_api_keys.provider_id mismatch: expected {}, got {}",
provider.id, key.provider_id
)));
}
let provider = map_provider(provider);
let mut endpoint = map_endpoint(endpoint);
let mut key = map_key(key, encryption_key, &fallback_encryption_keys)?;
if let LocalAuthConfigAbsorption::Absorbed {
base_url,
header_rules,
custom_path,
} = absorb_local_auth_config_safe_subset(
&endpoint.base_url,
endpoint.header_rules.clone(),
endpoint.custom_path.clone(),
key.decrypted_auth_config.as_deref(),
) {
endpoint.base_url = base_url;
endpoint.header_rules = header_rules;
endpoint.custom_path = custom_path;
key.decrypted_auth_config = None;
}
Ok(Some(GatewayProviderTransportSnapshot {
provider,
endpoint,
key,
}))
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data::repository::provider_catalog::{
InMemoryProviderCatalogReadRepository, ProviderCatalogReadRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data::DataLayerError;
use async_trait::async_trait;
use super::super::policy::{
supports_local_openai_chat_transport, supports_local_standard_transport_with_network,
};
use super::{
map_key, read_provider_transport_snapshot, GatewayProviderTransportSnapshot,
ProviderTransportSnapshotSource,
};
struct TestSnapshotSource {
repository: Arc<InMemoryProviderCatalogReadRepository>,
encryption_key: Option<String>,
}
impl TestSnapshotSource {
fn new(
repository: Arc<InMemoryProviderCatalogReadRepository>,
encryption_key: impl Into<Option<String>>,
) -> Self {
Self {
repository,
encryption_key: encryption_key.into(),
}
}
}
#[async_trait]
impl ProviderTransportSnapshotSource for TestSnapshotSource {
fn encryption_key(&self) -> Option<&str> {
self.encryption_key.as_deref()
}
async fn list_provider_catalog_providers_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
self.repository.list_providers_by_ids(ids).await
}
async fn list_provider_catalog_endpoints_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
self.repository.list_endpoints_by_ids(ids).await
}
async fn list_provider_catalog_keys_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
self.repository.list_keys_by_ids(ids).await
}
}
fn sample_provider() -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
"provider-1".to_string(),
"OpenAI".to_string(),
Some("https://openai.com".to_string()),
"custom".to_string(),
)
.expect("provider should build")
.with_transport_fields(
true,
false,
true,
Some(32),
Some(3),
Some(serde_json::json!({"url":"http://provider-proxy"})),
Some(20.0),
Some(8.0),
Some(serde_json::json!({"region":"global"})),
)
}
fn sample_endpoint() -> StoredProviderCatalogEndpoint {
StoredProviderCatalogEndpoint::new(
"endpoint-1".to_string(),
"provider-1".to_string(),
"openai:chat".to_string(),
Some("openai".to_string()),
Some("chat".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://api.openai.com".to_string(),
Some(serde_json::json!([{"action":"set","key":"x-test","value":"1"}])),
Some(serde_json::json!([{"action":"drop","path":"stream"}])),
Some(2),
Some("/v1/chat/completions".to_string()),
Some(serde_json::json!({"api_version":"v1"})),
Some(serde_json::json!({"allow":["openai:chat"]})),
Some(serde_json::json!({"url":"http://endpoint-proxy"})),
)
.expect("endpoint transport fields should build")
}
fn sample_key() -> StoredProviderCatalogKey {
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-live-openai")
.expect("api key ciphertext should build");
let encrypted_auth_config = encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
"{\"refresh_token\":\"rt-1\",\"project\":\"demo\"}",
)
.expect("auth config ciphertext should build");
StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"prod-key".to_string(),
"api_key".to_string(),
Some(serde_json::json!({"cache_1h": true})),
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!(["openai:chat", "openai:cli"])),
encrypted_api_key,
Some(encrypted_auth_config),
Some(serde_json::json!({"openai:chat": 0.8})),
Some(serde_json::json!({"openai:chat": 1})),
Some(serde_json::json!(["gpt-4.1", "gpt-4.1-mini"])),
Some(1_800_000_000),
Some(serde_json::json!({"node_id":"proxy-node-1"})),
Some(serde_json::json!({"tls_profile":"chrome_136"})),
)
.expect("key transport fields should build")
}
fn read_state() -> TestSnapshotSource {
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
vec![sample_endpoint()],
vec![sample_key()],
));
TestSnapshotSource::new(repository, Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()))
}
#[tokio::test]
async fn reads_decrypted_provider_transport_snapshot() {
let state = read_state();
let snapshot =
read_provider_transport_snapshot(&state, "provider-1", "endpoint-1", "key-1")
.await
.expect("snapshot should read")
.expect("snapshot should exist");
assert_eq!(
snapshot,
GatewayProviderTransportSnapshot {
provider: super::GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "OpenAI".to_string(),
provider_type: "custom".to_string(),
website: Some("https://openai.com".to_string()),
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: Some(32),
max_retries: Some(3),
proxy: Some(serde_json::json!({"url":"http://provider-proxy"})),
request_timeout_secs: Some(20.0),
stream_first_byte_timeout_secs: Some(8.0),
config: Some(serde_json::json!({"region":"global"})),
},
endpoint: super::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.com".to_string(),
header_rules: Some(
serde_json::json!([{"action":"set","key":"x-test","value":"1"}]),
),
body_rules: Some(serde_json::json!([{"action":"drop","path":"stream"}])),
max_retries: Some(2),
custom_path: Some("/v1/chat/completions".to_string()),
config: Some(serde_json::json!({"api_version":"v1"})),
format_acceptance_config: Some(serde_json::json!({"allow":["openai:chat"]}),),
proxy: Some(serde_json::json!({"url":"http://endpoint-proxy"})),
},
key: super::GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "prod-key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: Some(vec!["openai:chat".to_string(), "openai:cli".to_string(),]),
allowed_models: Some(vec!["gpt-4.1".to_string(), "gpt-4.1-mini".to_string(),]),
capabilities: Some(serde_json::json!({"cache_1h": true})),
rate_multipliers: Some(serde_json::json!({"openai:chat": 0.8})),
global_priority_by_format: Some(serde_json::json!({"openai:chat": 1})),
expires_at_unix_secs: Some(1_800_000_000),
proxy: Some(serde_json::json!({"node_id":"proxy-node-1"})),
fingerprint: Some(serde_json::json!({"tls_profile":"chrome_136"})),
decrypted_api_key: "sk-live-openai".to_string(),
decrypted_auth_config: Some(
"{\"refresh_token\":\"rt-1\",\"project\":\"demo\"}".to_string(),
),
},
}
);
}
#[tokio::test]
async fn returns_none_when_encryption_key_is_not_configured() {
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
vec![sample_endpoint()],
vec![sample_key()],
));
let state = TestSnapshotSource::new(repository, None);
let snapshot =
read_provider_transport_snapshot(&state, "provider-1", "endpoint-1", "key-1")
.await
.expect("snapshot read should not error");
assert!(snapshot.is_none());
}
#[tokio::test]
async fn absorbs_safe_auth_config_into_local_transport_fields() {
let provider = sample_provider();
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-safe-1".to_string(),
"provider-1".to_string(),
"openai:chat".to_string(),
Some("openai".to_string()),
Some("chat".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://api.openai.com".to_string(),
Some(serde_json::json!([{"action":"set","key":"x-test","value":"1"}])),
None,
Some(2),
None,
None,
None,
None,
)
.expect("endpoint transport fields should build");
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-live-openai")
.expect("api key ciphertext should build");
let encrypted_auth_config = encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"headers":{"x-account-id":"acc-1"},"query":{"tenant":"demo"}}"#,
)
.expect("auth config ciphertext should build");
let key = StoredProviderCatalogKey::new(
"key-safe-1".to_string(),
"provider-1".to_string(),
"safe-key".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!(["openai:chat"])),
encrypted_api_key,
Some(encrypted_auth_config),
None,
None,
None,
None,
None,
None,
)
.expect("key transport fields should build");
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let state =
TestSnapshotSource::new(repository, Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()));
let snapshot =
read_provider_transport_snapshot(&state, "provider-1", "endpoint-safe-1", "key-safe-1")
.await
.expect("snapshot read should succeed")
.expect("snapshot should exist");
assert_eq!(snapshot.key.decrypted_auth_config, None);
assert_eq!(
snapshot.endpoint.base_url,
"https://api.openai.com?tenant=demo"
);
assert_eq!(snapshot.endpoint.custom_path.as_deref(), None);
assert_eq!(
snapshot.endpoint.header_rules,
Some(serde_json::json!([
{"action":"set","key":"x-test","value":"1"},
{"action":"set","key":"x-account-id","value":"acc-1"}
]))
);
assert!(supports_local_openai_chat_transport(&snapshot));
}
#[tokio::test]
async fn accepts_plaintext_legacy_key_material() {
let provider = sample_provider();
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-legacy-1".to_string(),
"provider-1".to_string(),
"openai:chat".to_string(),
Some("openai".to_string()),
Some("chat".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://api.openai.com".to_string(),
Some(serde_json::json!([{"action":"set","key":"x-test","value":"1"}])),
None,
Some(2),
None,
None,
None,
None,
)
.expect("endpoint transport fields should build");
let key = StoredProviderCatalogKey::new(
"key-legacy-1".to_string(),
"provider-1".to_string(),
"legacy-key".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!(["openai:chat"])),
"sk-plaintext-openai".to_string(),
Some(r#"{"headers":{"x-account-id":"acc-legacy"}}"#.to_string()),
None,
None,
None,
None,
None,
None,
)
.expect("key transport fields should build");
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let state =
TestSnapshotSource::new(repository, Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()));
let snapshot = read_provider_transport_snapshot(
&state,
"provider-1",
"endpoint-legacy-1",
"key-legacy-1",
)
.await
.expect("snapshot read should succeed")
.expect("snapshot should exist");
assert_eq!(snapshot.key.decrypted_api_key, "sk-plaintext-openai");
assert_eq!(snapshot.key.decrypted_auth_config, None);
assert_eq!(
snapshot.endpoint.header_rules,
Some(serde_json::json!([
{"action":"set","key":"x-test","value":"1"},
{"action":"set","key":"x-account-id","value":"acc-legacy"}
]))
);
}
#[tokio::test]
async fn rejects_fernet_shaped_key_material_when_encryption_key_is_wrong() {
let key = sample_key();
let error = map_key(key, "wrong-encryption-key", &[])
.expect_err("snapshot read should fail for Fernet-shaped data with wrong key");
assert!(matches!(error, DataLayerError::UnexpectedValue(message)
if message.contains("failed to decrypt provider_api_keys.api_key")));
}
#[test]
fn decrypts_fernet_shaped_key_material_with_fallback_encryption_key() {
let key = sample_key();
let mapped = map_key(
key,
"wrong-encryption-key",
&[DEVELOPMENT_ENCRYPTION_KEY.to_string()],
)
.expect("fallback key should decrypt");
assert_eq!(mapped.decrypted_api_key, "sk-live-openai");
assert_eq!(
mapped.decrypted_auth_config.as_deref(),
Some("{\"refresh_token\":\"rt-1\",\"project\":\"demo\"}")
);
}
#[test]
fn accepts_stringified_allowed_models_in_transport_key() {
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-live-openai")
.expect("api key ciphertext should build");
let key = StoredProviderCatalogKey::new(
"key-compat-1".to_string(),
"provider-1".to_string(),
"compat-key".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!(["openai:chat"])),
encrypted_api_key,
None,
None,
None,
Some(serde_json::json!("[\"gpt-5.2\", \"gpt-5\"]")),
None,
None,
None,
)
.expect("key transport fields should build");
let mapped =
map_key(key, DEVELOPMENT_ENCRYPTION_KEY, &[]).expect("stringified list should parse");
assert_eq!(
mapped.allowed_models,
Some(vec!["gpt-5.2".to_string(), "gpt-5".to_string()])
);
}
#[test]
fn accepts_single_string_api_format_in_transport_key() {
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-live-openai")
.expect("api key ciphertext should build");
let key = StoredProviderCatalogKey::new(
"key-compat-2".to_string(),
"provider-1".to_string(),
"compat-key".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!("openai:chat")),
encrypted_api_key,
None,
None,
None,
Some(serde_json::json!("gpt-5.2")),
None,
None,
None,
)
.expect("key transport fields should build");
let mapped =
map_key(key, DEVELOPMENT_ENCRYPTION_KEY, &[]).expect("single string should parse");
assert_eq!(mapped.api_formats, Some(vec!["openai:chat".to_string()]));
assert_eq!(mapped.allowed_models, Some(vec!["gpt-5.2".to_string()]));
}
#[tokio::test]
async fn keeps_unsupported_auth_config_blocking_local_transport() {
let provider = sample_provider();
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-safe-2".to_string(),
"provider-1".to_string(),
"openai:cli".to_string(),
Some("openai".to_string()),
Some("cli".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://api.openai.com".to_string(),
None,
None,
Some(2),
None,
None,
None,
None,
)
.expect("endpoint transport fields should build");
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-live-openai")
.expect("api key ciphertext should build");
let encrypted_auth_config = encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"refresh_token":"rt-1","project":"demo"}"#,
)
.expect("auth config ciphertext should build");
let key = StoredProviderCatalogKey::new(
"key-safe-2".to_string(),
"provider-1".to_string(),
"unsafe-key".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!(["openai:cli"])),
encrypted_api_key,
Some(encrypted_auth_config),
None,
None,
None,
None,
None,
None,
)
.expect("key transport fields should build");
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let state =
TestSnapshotSource::new(repository, Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()));
let snapshot =
read_provider_transport_snapshot(&state, "provider-1", "endpoint-safe-2", "key-safe-2")
.await
.expect("snapshot read should succeed")
.expect("snapshot should exist");
assert_eq!(
snapshot.key.decrypted_auth_config.as_deref(),
Some(r#"{"refresh_token":"rt-1","project":"demo"}"#)
);
assert!(!supports_local_standard_transport_with_network(
&snapshot,
"openai:cli"
));
}
#[tokio::test]
async fn absorbs_query_only_auth_config_into_gemini_base_url() {
let provider = sample_provider();
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-safe-3".to_string(),
"provider-1".to_string(),
"gemini:chat".to_string(),
Some("gemini".to_string()),
Some("chat".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://generativelanguage.googleapis.com/v1beta".to_string(),
None,
None,
Some(2),
None,
None,
None,
None,
)
.expect("endpoint transport fields should build");
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-live-openai")
.expect("api key ciphertext should build");
let encrypted_auth_config = encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"query":{"alt":"sse"}}"#,
)
.expect("auth config ciphertext should build");
let key = StoredProviderCatalogKey::new(
"key-safe-3".to_string(),
"provider-1".to_string(),
"safe-key".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!(["gemini:chat"])),
encrypted_api_key,
Some(encrypted_auth_config),
None,
None,
None,
None,
None,
None,
)
.expect("key transport fields should build");
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let state =
TestSnapshotSource::new(repository, Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()));
let snapshot =
read_provider_transport_snapshot(&state, "provider-1", "endpoint-safe-3", "key-safe-3")
.await
.expect("snapshot read should succeed")
.expect("snapshot should exist");
assert_eq!(
snapshot.endpoint.base_url,
"https://generativelanguage.googleapis.com/v1beta?alt=sse"
);
assert_eq!(snapshot.endpoint.custom_path, None);
assert_eq!(snapshot.key.decrypted_auth_config, None);
}
#[tokio::test]
async fn absorbs_transport_subset_when_metadata_is_present() {
let provider = sample_provider();
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-safe-4".to_string(),
"provider-1".to_string(),
"openai:cli".to_string(),
Some("openai".to_string()),
Some("cli".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://api.openai.com/v1".to_string(),
Some(serde_json::json!([{"action":"set","key":"x-base","value":"1"}])),
None,
Some(2),
None,
None,
None,
None,
)
.expect("endpoint transport fields should build");
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-live-openai")
.expect("api key ciphertext should build");
let encrypted_auth_config = encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{
"email":"user@example.com",
"plan_type":"plus",
"transport":{
"extraHeaders":{"x-org-id":"org-1"},
"queryParams":{"tenant":"demo","retry":2}
}
}"#,
)
.expect("auth config ciphertext should build");
let key = StoredProviderCatalogKey::new(
"key-safe-4".to_string(),
"provider-1".to_string(),
"safe-key".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::json!(["openai:cli"])),
encrypted_api_key,
Some(encrypted_auth_config),
None,
None,
None,
None,
None,
None,
)
.expect("key transport fields should build");
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let state =
TestSnapshotSource::new(repository, Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()));
let snapshot =
read_provider_transport_snapshot(&state, "provider-1", "endpoint-safe-4", "key-safe-4")
.await
.expect("snapshot read should succeed")
.expect("snapshot should exist");
assert_eq!(snapshot.key.decrypted_auth_config, None);
assert_eq!(
snapshot.endpoint.base_url,
"https://api.openai.com/v1?retry=2&tenant=demo"
);
assert_eq!(
snapshot.endpoint.header_rules,
Some(serde_json::json!([
{"action":"set","key":"x-base","value":"1"},
{"action":"set","key":"x-org-id","value":"org-1"}
]))
);
assert!(supports_local_standard_transport_with_network(
&snapshot,
"openai:cli"
));
}
#[tokio::test]
async fn normalizes_json_null_transport_fields_before_local_support_checks() {
let provider = sample_provider().with_transport_fields(
true,
false,
false,
None,
Some(2),
Some(serde_json::Value::Null),
Some(20.0),
Some(8.0),
Some(serde_json::Value::Null),
);
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-null-json".to_string(),
"provider-1".to_string(),
"openai:chat".to_string(),
Some("openai".to_string()),
Some("chat".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://api.openai.com".to_string(),
Some(serde_json::Value::Null),
Some(serde_json::Value::Null),
Some(2),
None,
Some(serde_json::Value::Null),
Some(serde_json::Value::Null),
Some(serde_json::Value::Null),
)
.expect("endpoint transport fields should build");
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "sk-live-openai")
.expect("api key ciphertext should build");
let key = StoredProviderCatalogKey::new(
"key-null-json".to_string(),
"provider-1".to_string(),
"safe-key".to_string(),
"api_key".to_string(),
Some(serde_json::Value::Null),
true,
)
.expect("key should build")
.with_transport_fields(
Some(serde_json::Value::Null),
encrypted_api_key,
None,
Some(serde_json::Value::Null),
Some(serde_json::Value::Null),
Some(serde_json::Value::Null),
None,
Some(serde_json::Value::Null),
Some(serde_json::Value::Null),
)
.expect("key transport fields should build");
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let state =
TestSnapshotSource::new(repository, Some(DEVELOPMENT_ENCRYPTION_KEY.to_string()));
let snapshot = read_provider_transport_snapshot(
&state,
"provider-1",
"endpoint-null-json",
"key-null-json",
)
.await
.expect("snapshot read should succeed")
.expect("snapshot should exist");
assert_eq!(snapshot.provider.proxy, None);
assert_eq!(snapshot.provider.config, None);
assert_eq!(snapshot.endpoint.header_rules, None);
assert_eq!(snapshot.endpoint.body_rules, None);
assert_eq!(snapshot.endpoint.config, None);
assert_eq!(snapshot.endpoint.format_acceptance_config, None);
assert_eq!(snapshot.endpoint.proxy, None);
assert_eq!(snapshot.key.api_formats, None);
assert_eq!(snapshot.key.allowed_models, None);
assert_eq!(snapshot.key.capabilities, None);
assert_eq!(snapshot.key.rate_multipliers, None);
assert_eq!(snapshot.key.global_priority_by_format, None);
assert_eq!(snapshot.key.proxy, None);
assert_eq!(snapshot.key.fingerprint, None);
assert!(supports_local_openai_chat_transport(&snapshot));
}
}

View File

@@ -0,0 +1,229 @@
use aether_crypto::{decrypt_python_fernet_ciphertext, looks_like_python_fernet_ciphertext};
use aether_data::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data::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 = decrypt_secret(
encryption_key,
fallback_encryption_keys,
&key.encrypted_api_key,
"provider_api_keys.api_key",
)?;
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",
)?,
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),
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> {
match decrypt_python_fernet_ciphertext(encryption_key, ciphertext) {
Ok(value) => Ok(value),
Err(_error) if should_use_plaintext_secret(ciphertext, field_name) => {
Ok(ciphertext.trim().to_string())
}
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;
}
if looks_like_python_fernet_ciphertext(ciphertext) {
return false;
}
match field_name {
"provider_api_keys.api_key" => !ciphertext.starts_with('{') && !ciphertext.starts_with('['),
"provider_api_keys.auth_config" => {
ciphertext.starts_with('{') || ciphertext.starts_with('[')
}
_ => 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)
}

View File

@@ -0,0 +1,321 @@
use std::collections::BTreeMap;
use super::provider_types::is_codex_cli_backend_url;
use url::form_urlencoded;
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 = if trimmed.ends_with("/v1") {
format!("{trimmed}/chat/completions")
} else {
format!("{trimmed}/v1/chat/completions")
};
append_merged_query(&mut url, base_query, None, query, &[]);
url
}
pub fn build_openai_cli_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 = if is_codex_cli_backend_url(trimmed)
|| trimmed.ends_with("/codex")
|| trimmed.ends_with("/v1")
{
format!("{trimmed}{suffix}")
} else {
format!("{trimmed}/v1{suffix}")
};
append_merged_query(&mut url, base_query, None, query, &[]);
url
}
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 = if trimmed.ends_with("/v1") {
format!("{trimmed}/messages")
} else {
format!("{trimmed}/v1/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("/v1beta") {
format!("{trimmed_base_url}/models/{trimmed_model}:{operation}")
} else if trimmed_base_url.contains("/v1beta/models/") {
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 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("/v1beta") {
format!("{trimmed_base_url}/models/{trimmed_model}:predictLongRunning")
} else if trimmed_base_url.contains("/v1beta/models/") {
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 {
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_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))
}
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_gemini_content_url, build_gemini_files_passthrough_url,
build_gemini_video_predict_long_running_url, build_openai_chat_url,
build_passthrough_path_url,
};
#[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 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 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 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"
)
);
}
}

View File

@@ -0,0 +1,19 @@
mod auth;
mod policy;
mod url;
pub use auth::{
resolve_local_vertex_api_key_query_auth, VertexApiKeyQueryAuth, VERTEX_API_KEY_QUERY_PARAM,
};
pub use policy::{
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,
};
pub use url::{
build_vertex_api_key_gemini_content_url, build_vertex_api_key_imagen_content_url,
VERTEX_API_KEY_BASE_URL,
};
pub const PROVIDER_TYPE: &str = "vertex_ai";

View File

@@ -0,0 +1,130 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
pub const VERTEX_API_KEY_QUERY_PARAM: &str = "key";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VertexApiKeyQueryAuth {
pub name: &'static str,
pub value: String,
}
pub fn resolve_local_vertex_api_key_query_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<VertexApiKeyQueryAuth> {
if !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
{
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(),
})
}
#[cfg(test)]
mod tests {
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{resolve_local_vertex_api_key_query_auth, 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:chat".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:chat".to_string()]),
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: "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());
}
}

View File

@@ -0,0 +1,204 @@
use super::super::snapshot::GatewayProviderTransportSnapshot;
use super::super::{
body_rules_are_locally_supported, header_rules_are_locally_supported,
resolve_transport_tls_profile, transport_proxy_is_locally_supported,
};
use super::auth::resolve_local_vertex_api_key_query_auth;
pub fn supports_local_vertex_api_key_gemini_transport(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
supports_local_vertex_api_key_same_format_transport(
transport,
&["gemini:chat", "gemini:cli"],
false,
)
}
pub fn supports_local_vertex_api_key_gemini_transport_with_network(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
supports_local_vertex_api_key_same_format_transport(
transport,
&["gemini:chat", "gemini:cli"],
true,
)
}
pub fn supports_local_vertex_api_key_imagen_transport(
transport: &GatewayProviderTransportSnapshot,
) -> bool {
supports_local_vertex_api_key_same_format_transport(transport, &["gemini:chat"], 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:chat"], 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;
}
if !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(super::PROVIDER_TYPE)
{
return false;
}
let endpoint_api_format = transport.endpoint.api_format.trim();
if !api_formats
.iter()
.any(|api_format| endpoint_api_format.eq_ignore_ascii_case(api_format))
{
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());
let has_tls_profile = resolve_transport_tls_profile(transport)
.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.key.fingerprint.is_some() && resolve_transport_tls_profile(transport).is_none()
{
return false;
}
} else if transport.provider.proxy.is_some()
|| transport.endpoint.proxy.is_some()
|| transport.key.proxy.is_some()
|| has_tls_profile
{
return false;
}
true
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::super::super::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use super::{
supports_local_vertex_api_key_gemini_transport,
supports_local_vertex_api_key_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:chat".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:chat".to_string()]),
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: "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_cli_subset() {
let mut transport = sample_transport();
transport.endpoint.api_format = "gemini:cli".to_string();
assert!(supports_local_vertex_api_key_gemini_transport(&transport));
}
#[test]
fn rejects_vertex_service_account_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 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!({"tls_profile":"chrome_136"}));
assert!(!supports_local_vertex_api_key_gemini_transport(&transport));
assert!(supports_local_vertex_api_key_gemini_transport_with_network(
&transport
));
}
}

View File

@@ -0,0 +1,121 @@
use std::collections::BTreeMap;
use url::form_urlencoded;
use super::super::url::build_passthrough_path_url;
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> {
build_vertex_api_key_google_model_url(model, 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> {
build_vertex_api_key_google_model_url(model, stream, api_key, request_query)
}
fn build_vertex_api_key_google_model_url(
model: &str,
stream: bool,
api_key: &str,
request_query: Option<&str>,
) -> Option<String> {
let trimmed_model = model.trim();
let trimmed_api_key = api_key.trim();
if trimmed_model.is_empty() || trimmed_api_key.is_empty() {
return None;
}
let action = if stream {
"streamGenerateContent"
} else {
"generateContent"
};
let path = format!("/v1/publishers/google/models/{trimmed_model}:{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_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 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 super::{build_vertex_api_key_gemini_content_url, build_vertex_api_key_imagen_content_url};
#[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"
)
);
}
}

View File

@@ -0,0 +1,285 @@
use aether_data::repository::video_tasks::StoredVideoTask;
use aether_video_tasks_core::{
LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput,
};
use async_trait::async_trait;
use super::auth::{resolve_local_gemini_auth, resolve_local_standard_auth};
use super::network::resolve_transport_execution_timeouts;
use super::policy::{supports_local_gemini_transport, supports_local_standard_transport};
use super::snapshot::GatewayProviderTransportSnapshot;
#[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_standard_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,
tls_profile: None,
timeouts: resolve_transport_execution_timeouts(transport),
},
))
}
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::repository::video_tasks::{StoredVideoTask, VideoTaskStatus};
use aether_video_tasks_core::LocalVideoTaskSnapshot;
use async_trait::async_trait;
use serde_json::json;
use super::{
reconstruct_local_video_task_snapshot, resolve_local_video_task_transport,
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,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
}
}
fn 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_secs: 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", "bearer"),
"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 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"),
}
}
}