fix(bedrock): simplify adapter cleanup

This commit is contained in:
Scott Werner 2026-06-16 11:24:20 -04:00
parent c2cf1ccbf7
commit e0a4e6bc1a
11 changed files with 222 additions and 85 deletions

View file

@ -7,7 +7,7 @@
use super::SYNTHETIC_TOOL_NAME;
use super::decode::{convert_synthetic_tool_to_text, map_finish_reason, refusal_error};
use crate::codec::{RawEvent, StreamDecoder};
use crate::codec::{RawEvent, StreamDecoder, parse_tool_arguments_or_empty};
use crate::error::{Error, ProviderErrorDetail, ProviderErrorKind};
use crate::types::{
ContentPart, FinishReason, Message, RateLimitInfo, Response, Role, StreamEvent, ThinkingData,
@ -253,8 +253,7 @@ impl SseAccumulator {
}
Some(ContentBlockKind::ToolUse { id, name }) => {
let raw_args = std::mem::take(&mut self.current_tool_args);
let arguments =
serde_json::from_str(&raw_args).unwrap_or_else(|_| serde_json::json!({}));
let arguments = parse_tool_arguments_or_empty(&raw_args);
let mut tool_call = ToolCall::new(id, name, arguments);
tool_call.raw_arguments = Some(raw_args);
self.content_parts

View file

@ -4,7 +4,7 @@ use base64::Engine;
use base64::engine::general_purpose::STANDARD as BASE64;
use serde_json::{Map, Value, json};
use crate::codec::{CodecCtx, EncodedRequest, extract_system_prompt};
use crate::codec::{CodecCtx, EncodedRequest, extract_system_prompt, merge_named_provider_options};
use crate::error::Error;
use crate::types::{ContentPart, Message, Request, Role, ToolChoice};
@ -226,13 +226,28 @@ fn encode_content_part(part: &ContentPart) -> Option<Value> {
}
}
/// `image/png` → `png`; missing/odd media types fall back to `default`.
fn media_format(media_type: Option<&str>, default: &str) -> String {
media_type
.and_then(|m| m.split('/').next_back())
.filter(|s| !s.is_empty())
.unwrap_or(default)
.to_string()
/// Convert common MIME types into Bedrock's media `format` enum values.
fn media_format<'a>(media_type: Option<&str>, default: &'a str) -> &'a str {
match media_type {
Some("image/png") => "png",
Some("image/jpeg" | "image/jpg") => "jpeg",
Some("image/gif") => "gif",
Some("image/webp") => "webp",
Some("application/pdf") => "pdf",
Some("text/plain") => "txt",
Some("text/markdown") => "md",
Some("text/html") => "html",
Some("text/csv") => "csv",
Some(
"application/msword"
| "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
) => "docx",
Some(
"application/vnd.ms-excel"
| "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
) => "xlsx",
_ => default,
}
}
fn encode_tool_config(request: &Request, caching: bool) -> Option<Value> {
@ -302,18 +317,18 @@ fn tool_input_schema(parameters: &Value) -> Value {
/// `cachePoint` at the end of the second-to-last user message, so the prior
/// turns stay cached while the newest turn streams.
fn apply_cache_point_to_conversation_prefix(messages: &mut [Value]) {
let user_indices: Vec<usize> = messages
.iter()
.enumerate()
.filter(|(_, m)| m.get("role").and_then(Value::as_str) == Some("user"))
.map(|(i, _)| i)
.collect();
if user_indices.len() < 2 {
return;
let mut previous_user = None;
let mut last_user = None;
for (index, message) in messages.iter().enumerate() {
if message.get("role").and_then(Value::as_str) == Some("user") {
previous_user = last_user;
last_user = Some(index);
}
}
let target = user_indices[user_indices.len() - 2];
let Some(target) = previous_user else {
return;
};
if let Some(content) = messages[target]
.get_mut("content")
.and_then(Value::as_array_mut)
@ -327,18 +342,7 @@ fn apply_cache_point_to_conversation_prefix(messages: &mut [Value]) {
/// codec). This is the passthrough for `additionalModelRequestFields`,
/// `guardrailConfig`, `serviceTier`, and other Converse extensions.
fn merge_provider_options(body: &mut Value, provider_options: Option<&Value>, provider_name: &str) {
let Some(opts) = provider_options.and_then(|opts| opts.get(provider_name)) else {
return;
};
let Some(body_map) = body.as_object_mut() else {
return;
};
let Some(opts_map) = opts.as_object() else {
return;
};
for (key, value) in opts_map {
body_map.insert(key.clone(), value.clone());
}
merge_named_provider_options(body, provider_options, provider_name);
}
#[cfg(test)]
@ -537,6 +541,14 @@ mod tests {
assert_eq!(block["signature"], "sig-1");
}
#[test]
fn media_format_maps_common_mime_types_to_bedrock_formats() {
assert_eq!(media_format(Some("image/jpeg"), "png"), "jpeg");
assert_eq!(media_format(Some("text/plain"), "pdf"), "txt");
assert_eq!(media_format(Some("text/markdown"), "pdf"), "md");
assert_eq!(media_format(Some("application/octet-stream"), "pdf"), "pdf");
}
#[test]
fn provider_options_merge_top_level() {
let mut request = base_request("claude");

View file

@ -13,7 +13,7 @@ use std::collections::BTreeMap;
use serde_json::Value;
use super::decode::{map_stop_reason, token_counts_from_usage};
use crate::codec::{CodecCtx, RawEvent, StreamDecoder};
use crate::codec::{CodecCtx, RawEvent, StreamDecoder, parse_tool_arguments_or_empty};
use crate::error::Error;
use crate::types::{
ContentPart, FinishReason, Message, RateLimitInfo, Response, Role, StreamEvent, ThinkingData,
@ -222,11 +222,7 @@ impl ConverseStreamDecoder {
// the buffer empty; canonically that is an empty object, not
// null (matching the anthropic/openai codecs, and what Bedrock
// wants back on re-encode).
let arguments = if input.trim().is_empty() {
serde_json::json!({})
} else {
serde_json::from_str(&input).unwrap_or(Value::Null)
};
let arguments = parse_tool_arguments_or_empty(&input);
let mut tool_call = ToolCall::new(&id, &name, arguments);
tool_call.raw_arguments = Some(input);
self.parts.push(ContentPart::ToolCall(tool_call.clone()));

View file

@ -21,6 +21,35 @@ use fabro_model::Model;
use crate::error::{Error, error_from_status_code};
use crate::types::{Message, RateLimitInfo, Request, Response, Role, StreamEvent};
/// Parse a streamed/generated tool-argument JSON string, defaulting malformed
/// or absent arguments to the canonical no-argument object.
pub(crate) fn parse_tool_arguments_or_empty(raw_arguments: &str) -> serde_json::Value {
serde_json::from_str(raw_arguments).unwrap_or_else(|_| serde_json::json!({}))
}
/// Merge `provider_options.<provider_name>` fields into an encoded request
/// body. Used by codecs whose provider-options namespace is adapter-name keyed
/// rather than a single fixed provider.
pub(crate) fn merge_named_provider_options(
body: &mut serde_json::Value,
provider_options: Option<&serde_json::Value>,
provider_name: &str,
) {
let Some(opts) = provider_options.and_then(|opts| opts.get(provider_name)) else {
return;
};
let Some(body_map) = body.as_object_mut() else {
return;
};
let Some(opts_map) = opts.as_object() else {
return;
};
for (key, value) in opts_map {
body_map.insert(key.clone(), value.clone());
}
}
/// Per-request context. Borrowed — the codec reads what it needs and returns.
pub(crate) struct CodecCtx<'a> {
/// The canonical request being translated. Decoders read it too

View file

@ -2,7 +2,7 @@
use super::translate;
use super::wire::ApiRequest;
use crate::codec::{CodecCtx, EncodedRequest};
use crate::codec::{CodecCtx, EncodedRequest, merge_named_provider_options};
/// Build the Chat Completions request for `ctx.request`. `stream` toggles the
/// `stream` body field. The body is assembled as a `serde_json::Value` so
@ -63,19 +63,7 @@ pub(super) fn merge_provider_options(
provider_options: Option<&serde_json::Value>,
provider_name: &str,
) {
let Some(opts) = provider_options.and_then(|opts| opts.get(provider_name)) else {
return;
};
let Some(body_map) = body.as_object_mut() else {
return;
};
let Some(opts_map) = opts.as_object() else {
return;
};
for (key, value) in opts_map {
body_map.insert(key.clone(), value.clone());
}
merge_named_provider_options(body, provider_options, provider_name);
}
#[cfg(test)]

View file

@ -3,7 +3,7 @@
use serde::Deserialize;
use super::wire::{ApiResponse, ApiUsage, InputTokensResponse};
use crate::codec::CodecCtx;
use crate::codec::{CodecCtx, parse_tool_arguments_or_empty};
use crate::error::Error;
use crate::types::{
ContentPart, FinishReason, Message, RateLimitInfo, Response, Role, TokenCounts, ToolCall,
@ -75,7 +75,7 @@ pub(super) fn tool_call_from_item(item: &serde_json::Value, custom: bool) -> Too
.get("arguments")
.and_then(serde_json::Value::as_str)
.unwrap_or("{}");
let arguments = serde_json::from_str(args_str).unwrap_or_else(|_| serde_json::json!({}));
let arguments = parse_tool_arguments_or_empty(args_str);
let mut tc = ToolCall::new(call_id, name, arguments);
tc.raw_arguments = Some(args_str.to_string());
tc

View file

@ -9,7 +9,7 @@
pub(crate) mod eventstream;
pub(crate) mod sigv4;
use std::collections::VecDeque;
use std::collections::{HashMap, VecDeque};
use std::sync::Arc;
use std::time::Duration;
@ -22,6 +22,8 @@ use tokio::sync::OnceCell;
use tokio::time;
use crate::adapter_registry::AdapterConfig;
#[cfg(test)]
use crate::adapter_registry::AdapterKindOptions;
use crate::attachments::{self, AttachmentPolicy};
use crate::codec::bedrock_converse::BedrockConverse;
use crate::codec::{Codec, CodecCtx, CodecParams, EncodedRequest, RawEvent, StreamDecoder};
@ -61,8 +63,16 @@ pub(crate) fn build(config: AdapterConfig) -> Result<Arc<dyn ProviderAdapter>, E
})?;
let adapter = match config.auth_header {
Some(ApiKeyHeader::AwsSigv4) => Adapter::new_sigv4(base_url)?,
Some(ApiKeyHeader::Bearer(token) | ApiKeyHeader::Custom { value: token, .. }) => {
Adapter::new_api_key(token, base_url)?
Some(ApiKeyHeader::Bearer(token)) => Adapter::new_api_key(token, base_url)?,
Some(ApiKeyHeader::Custom { name, .. }) => {
return Err(Error::Configuration {
message: format!(
"bedrock provider '{}' does not support custom auth header '{}' (use bearer \
credentials or aws_sigv4)",
config.provider_id, name
),
source: None,
});
}
None => {
return Err(Error::Configuration {
@ -76,6 +86,9 @@ pub(crate) fn build(config: AdapterConfig) -> Result<Arc<dyn ProviderAdapter>, E
}
};
let mut adapter = adapter.with_name(config.provider_id);
if !config.extra_headers.is_empty() {
adapter = adapter.with_default_headers(config.extra_headers);
}
if let Some(catalog) = config.catalog {
adapter = adapter.with_catalog(catalog);
}
@ -132,6 +145,12 @@ impl Adapter {
self
}
#[must_use]
pub fn with_default_headers(mut self, headers: HashMap<String, String>) -> Self {
self.http = self.http.with_default_headers(headers);
self
}
#[must_use]
pub fn with_timeout(self, timeout: AdapterTimeout) -> Self {
Self {
@ -182,6 +201,9 @@ impl Adapter {
for (key, value) in &self.http.default_headers {
req = req.header(key, value);
}
for (key, value) in &encoded.headers {
req = req.header(key, value);
}
req = match &self.auth {
BedrockAuth::ApiKey(token) => req.bearer_auth(token).body(body),
@ -189,7 +211,7 @@ impl Adapter {
let signer = cell
.get_or_try_init(Sigv4Signer::from_default_chain)
.await?;
signer.sign_post(req, &self.region, &url, &body).await?
signer.sign_post(req, &self.region, &url, body).await?
}
};
@ -378,11 +400,12 @@ fn region_from_base_url(base_url: &str) -> Result<String, Error> {
),
source: None,
};
let host = base_url
.strip_prefix("https://")
.or_else(|| base_url.strip_prefix("http://"))
.unwrap_or(base_url);
let host = host.split('/').next().unwrap_or(host);
#[expect(
clippy::disallowed_types,
reason = "Bedrock region derivation needs URL host parsing; the raw URL is not logged or rendered."
)]
let parsed = fabro_http::Url::parse(base_url).map_err(|_| invalid())?;
let host = parsed.host_str().ok_or_else(invalid)?;
let rest = host
.strip_prefix("bedrock-runtime-fips.")
.or_else(|| host.strip_prefix("bedrock-runtime."))
@ -473,12 +496,19 @@ mod tests {
"https://example.com",
"https://bedrock.us-east-1.amazonaws.com",
"https://bedrock-runtime.amazonaws.com",
"https://bedrock-runtime.UPPER.amazonaws.com",
] {
assert!(region_from_base_url(url).is_err(), "{url}");
}
}
#[test]
fn region_normalizes_hostname_case() {
assert_eq!(
region_from_base_url("https://bedrock-runtime.US-EAST-1.amazonaws.com").unwrap(),
"us-east-1"
);
}
#[tokio::test]
async fn complete_posts_converse_body_with_bearer_auth() {
let server = MockServer::start();
@ -511,6 +541,55 @@ mod tests {
assert_eq!(response.provider, "bedrock");
}
#[tokio::test]
async fn complete_applies_default_headers() {
let server = MockServer::start();
let mock = server.mock(|when, then| {
when.method(POST)
.path("/model/m/converse")
.header("x-fabro-test", "present");
then.status(200)
.header("content-type", "application/json")
.json_body(serde_json::json!({
"output": {"message": {"role": "assistant", "content": [{"text": "ok"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 1, "outputTokens": 1, "totalTokens": 2}
}));
});
let adapter = test_adapter(&server).with_default_headers(HashMap::from([(
"x-fabro-test".to_string(),
"present".to_string(),
)]));
let response = adapter.complete(&make_request("m")).await.unwrap();
mock.assert();
assert_eq!(response.text(), "ok");
}
#[test]
fn factory_rejects_custom_auth_header() {
let result = build(AdapterConfig {
provider_id: "bedrock".to_string(),
auth_header: Some(ApiKeyHeader::Custom {
name: "x-api-key".to_string(),
value: "secret".to_string(),
}),
base_url: Some("https://bedrock-runtime.us-east-1.amazonaws.com".to_string()),
extra_headers: HashMap::new(),
kind_options: AdapterKindOptions::None,
catalog: None,
});
let Err(err) = result else {
panic!("expected custom auth header to be rejected");
};
assert!(
err.to_string()
.contains("does not support custom auth header")
);
}
#[tokio::test]
async fn complete_signs_with_sigv4_when_configured() {
let server = MockServer::start();

View file

@ -144,7 +144,7 @@ impl Sigv4Signer {
mut req: fabro_http::RequestBuilder,
region: &str,
url: &str,
body: &[u8],
body: Vec<u8>,
) -> Result<fabro_http::RequestBuilder, Error> {
let credentials = self.current_credentials().await?;
let now = SystemTime::now()
@ -155,11 +155,11 @@ impl Sigv4Signer {
})?
.as_secs();
for (name, value) in
Self::signed_headers(&credentials, region, SERVICE, "POST", url, body, now)?
Self::signed_headers(&credentials, region, SERVICE, "POST", url, &body, now)?
{
req = req.header(name, value);
}
Ok(req.body(body.to_vec()))
Ok(req.body(body))
}
}

View file

@ -176,9 +176,8 @@ pub struct CostRates {
/// `Vault`/`Env` reference a stored secret resolved to an auth header.
/// `AwsSigv4` is an opaque source: the credential comes from the AWS default
/// credential chain and the request is SigV4-signed rather than carrying a
/// static secret. Folding this into the credential list (instead of a separate
/// `scheme` field) keeps the "where do credentials come from" decision in one
/// place and makes invalid combinations unrepresentable.
/// static secret. It is only valid on Bedrock providers, which catalog
/// validation enforces before adapter construction.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(into = "String", try_from = "String")]
pub enum CredentialRef {
@ -616,6 +615,13 @@ pub enum CatalogBuildError {
},
#[error("provider '{provider}' API-key auth must declare at least one credential")]
EmptyApiKeyCredentials { provider: ProviderId },
#[error(
"provider '{provider}' uses aws_sigv4 credentials, but adapter '{adapter}' does not support SigV4"
)]
UnsupportedAwsSigv4Credential {
provider: ProviderId,
adapter: AdapterKind,
},
#[error("provider identifier '{identifier}' is declared by both '{first}' and '{second}'")]
DuplicateProviderIdentifier {
identifier: String,
@ -1355,7 +1361,7 @@ fn build_providers(
let codec = resolve_provider_codec(&provider_id, adapter, settings.codec)?;
let agent_profile = settings.agent_profile.unwrap_or(defaults.agent_profile);
let auth = settings.auth.clone();
validate_provider_auth(&provider_id, auth.as_ref())?;
validate_provider_auth(&provider_id, adapter, auth.as_ref())?;
providers.push(CatalogProvider {
id: provider_id,
@ -1441,6 +1447,7 @@ fn resolve_model_codec(
fn validate_provider_auth(
provider: &ProviderId,
adapter: AdapterKind,
auth: Option<&ProviderAuthConfig>,
) -> Result<(), CatalogBuildError> {
match auth {
@ -1449,6 +1456,18 @@ fn validate_provider_auth(
provider: provider.clone(),
})
}
Some(auth)
if adapter != AdapterKind::Bedrock
&& auth
.credentials
.iter()
.any(|credential| matches!(credential, CredentialRef::AwsSigv4)) =>
{
Err(CatalogBuildError::UnsupportedAwsSigv4Credential {
provider: provider.clone(),
adapter,
})
}
_ => Ok(()),
}
}
@ -3806,6 +3825,23 @@ credentials = []
CatalogBuildError::EmptyApiKeyCredentials { provider }
if provider == ProviderId::new("test")
));
let sigv4_on_openai = minimal_settings(
r#"
[providers.test]
display_name = "Test"
adapter = "openai"
agent_profile = "openai"
[providers.test.auth]
credentials = ["aws_sigv4"]
"#,
);
assert!(matches!(
Catalog::from_settings(&sigv4_on_openai).unwrap_err(),
CatalogBuildError::UnsupportedAwsSigv4Credential { provider, adapter }
if provider == ProviderId::new("test") && adapter == AdapterKind::OpenAi
));
}
#[test]

View file

@ -41,9 +41,9 @@ credentials = [
# ---------- Anthropic Claude ----------
#
# Claude bills Anthropic-style cache reads/writes, so these rows override
# the provider's billing default. Claude Fable 5 is deliberately absent:
# its Bedrock deployment pins sampling parameters (temperature must be
# unset) that the Converse route does not gate yet — a named follow-up.
# the provider's billing default. Claude Fable 5 appears at the end of this
# file because its Bedrock deployment pins sampling parameters and requires an
# extra data-sharing opt-in.
[models."us.anthropic.claude-sonnet-4-6"]
provider = "bedrock"

View file

@ -23,18 +23,17 @@ const WORKER_ENV_ALLOWLIST: &[&str] = &[
// AWS chain on every request so STS/SSO/IRSA sessions can refresh, which
// means the chain's *inputs* must survive `env_clear()` in the worker, not
// a snapshot taken at launch. We pass the identity surface only (static
// keys, session token, the Bedrock bearer key under either accepted name,
// profile/region selectors,
// and the web-identity/ECS role vars); HOME already carries the shared
// keys, session token, profile/region selectors, and the web-identity/ECS
// role vars); HOME already carries the shared
// `~/.aws` config + SSO cache. Endpoint/metadata overrides
// (AWS_ENDPOINT_*, AWS_METADATA_ENDPOINT, AWS_IMDSV1_FALLBACK) are
// deliberately excluded — they belong to the server's S3 path, not to the
// worker's outbound model calls.
// worker's outbound model calls. Bedrock bearer API keys are optional LLM
// provider secrets, so server workers read them through the vault rather
// than inheriting process env.
EnvVars::AWS_ACCESS_KEY_ID,
EnvVars::AWS_SECRET_ACCESS_KEY,
EnvVars::AWS_SESSION_TOKEN,
EnvVars::AWS_BEARER_TOKEN_BEDROCK,
EnvVars::BEDROCK_API_KEY,
EnvVars::AWS_PROFILE,
EnvVars::AWS_REGION,
EnvVars::AWS_DEFAULT_REGION,
@ -122,6 +121,7 @@ mod tests {
("AWS_SECRET_ACCESS_KEY".to_string(), "secret".to_string()),
("AWS_SESSION_TOKEN".to_string(), "session".to_string()),
("AWS_BEARER_TOKEN_BEDROCK".to_string(), "bearer".to_string()),
("BEDROCK_API_KEY".to_string(), "alias-bearer".to_string()),
("AWS_REGION".to_string(), "us-east-2".to_string()),
("SESSION_SECRET".to_string(), "leak".to_string()),
("FABRO_JWT_PRIVATE_KEY".to_string(), "leak".to_string()),
@ -169,14 +169,12 @@ mod tests {
actual.get("AWS_SESSION_TOKEN").map(String::as_str),
Some("session")
);
assert_eq!(
actual.get("AWS_BEARER_TOKEN_BEDROCK").map(String::as_str),
Some("bearer")
);
assert_eq!(
actual.get("AWS_REGION").map(String::as_str),
Some("us-east-2")
);
assert!(!actual.contains_key("AWS_BEARER_TOKEN_BEDROCK"));
assert!(!actual.contains_key("BEDROCK_API_KEY"));
assert!(!actual.contains_key("FABRO_LOG_DESTINATION"));
assert_eq!(
actual.get("FABRO_DEV_TOKEN").map(String::as_str),