Files
Aether/crates/aether-provider-transport/src/kiro/request.rs

285 lines
8.7 KiB
Rust
Raw Normal View History

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());
}
}