diff --git a/lib/crates/fabro-llm/src/adapter_registry.rs b/lib/crates/fabro-llm/src/adapter_registry.rs index a14c3083c..f37f80fa9 100644 --- a/lib/crates/fabro-llm/src/adapter_registry.rs +++ b/lib/crates/fabro-llm/src/adapter_registry.rs @@ -207,23 +207,6 @@ fn build_openai_compatible(config: AdapterConfig) -> Result Result, Error> { - Err(Error::Configuration { - message: format!( - "provider '{}': the bedrock adapter is not yet wired", - config.provider_id - ), - source: None, - }) -} - /// Return the factory for a known adapter kind. #[must_use] pub fn factory_for(adapter_kind: AdapterKind) -> AdapterFactory { @@ -232,7 +215,7 @@ pub fn factory_for(adapter_kind: AdapterKind) -> AdapterFactory { AdapterKind::OpenAi => build_openai, AdapterKind::Gemini => build_gemini, AdapterKind::OpenAiCompatible => build_openai_compatible, - AdapterKind::Bedrock => build_bedrock, + AdapterKind::Bedrock => providers::bedrock::build, } } diff --git a/lib/crates/fabro-llm/src/codec/bedrock_converse/decode.rs b/lib/crates/fabro-llm/src/codec/bedrock_converse/decode.rs new file mode 100644 index 000000000..117c96a88 --- /dev/null +++ b/lib/crates/fabro-llm/src/codec/bedrock_converse/decode.rs @@ -0,0 +1,283 @@ +//! Response decoding: Converse body → canonical `Response`. + +use serde_json::Value; + +use crate::codec::CodecCtx; +use crate::error::{Error, error_from_status_code}; +use crate::types::{ + ContentPart, FinishReason, Message, RateLimitInfo, Response, Role, ThinkingData, TokenCounts, + ToolCall, +}; + +/// Map a non-2xx Bedrock runtime response to an `Error`, pulling the human +/// reason out of AWS's error envelope. Bedrock uses several shapes for the +/// same field — top-level `message` (SigV4 path) and `Message` (API-key +/// path), occasionally nested `error.message` — and tags the type in +/// `__type`. The generic codec parser only reads `error.message`, so without +/// this these surface as "Unknown error". +pub(super) fn bedrock_error( + status: u16, + body: &str, + provider: &str, + retry_after: Option, +) -> Error { + let raw: Option = serde_json::from_str(body).ok(); + let message = raw + .as_ref() + .and_then(extract_error_message) + .unwrap_or_else(|| { + if body.trim().is_empty() { + "Unknown error".to_string() + } else { + body.to_string() + } + }); + // `__type` is often an ARN-ish `prefix#ThrottlingException`; keep the tail. + let code = raw + .as_ref() + .and_then(|v| { + v.get("__type") + .or_else(|| v.get("code")) + .and_then(Value::as_str) + }) + .map(|t| t.rsplit('#').next().unwrap_or(t).to_string()); + error_from_status_code( + status, + message, + provider.to_string(), + code, + raw, + retry_after, + ) +} + +fn extract_error_message(v: &Value) -> Option { + v.get("message") + .and_then(Value::as_str) + .or_else(|| v.get("Message").and_then(Value::as_str)) + .or_else(|| { + v.get("error") + .and_then(|e| e.get("message")) + .and_then(Value::as_str) + }) + .map(String::from) +} + +pub(super) fn decode_response( + body: &str, + ctx: &CodecCtx<'_>, + rate_limit: Option, +) -> Result { + let raw: Value = serde_json::from_str(body) + .map_err(|e| Error::network(format!("failed to parse converse response: {e}"), e))?; + + let content_parts = raw + .pointer("/output/message/content") + .and_then(Value::as_array) + .map(|blocks| blocks.iter().filter_map(decode_content_block).collect()) + .unwrap_or_default(); + + let finish_reason = map_stop_reason(raw.get("stopReason").and_then(Value::as_str)); + let usage = token_counts_from_usage(raw.get("usage")); + + Ok(Response { + // Converse responses carry no id; synthesize one like the gemini + // codec does so downstream consumers always see a non-empty id. + id: uuid::Uuid::new_v4().to_string(), + model: ctx.request.model.clone(), + provider: ctx.provider_name.to_string(), + message: Message { + role: Role::Assistant, + content: content_parts, + name: None, + tool_call_id: None, + }, + finish_reason, + usage, + raw: Some(raw), + warnings: vec![], + rate_limit, + cost_usd: None, + cost_source: None, + }) +} + +/// Decode one Converse content block into a canonical part. Unknown block +/// kinds are skipped (the union grows: `citationsContent`, `searchResult`, +/// `video`, ...). +pub(super) fn decode_content_block(block: &Value) -> Option { + if let Some(text) = block.get("text").and_then(Value::as_str) { + if text.is_empty() { + return None; + } + return Some(ContentPart::text(text)); + } + if let Some(tool_use) = block.get("toolUse") { + let id = tool_use.get("toolUseId").and_then(Value::as_str)?; + let name = tool_use.get("name").and_then(Value::as_str)?; + let input = tool_use.get("input").cloned().unwrap_or(Value::Null); + return Some(ContentPart::ToolCall(ToolCall::new(id, name, input))); + } + if let Some(reasoning) = block.get("reasoningContent") { + if let Some(text_block) = reasoning.get("reasoningText") { + return Some(ContentPart::Thinking(ThinkingData { + text: text_block + .get("text") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(), + signature: text_block + .get("signature") + .and_then(Value::as_str) + .map(str::to_string), + redacted: false, + })); + } + if let Some(redacted) = reasoning.get("redactedContent").and_then(Value::as_str) { + return Some(ContentPart::Thinking(ThinkingData { + text: redacted.to_string(), + signature: None, + redacted: true, + })); + } + } + None +} + +/// Map a Converse `stopReason` onto the canonical finish vocabulary. +pub(super) fn map_stop_reason(reason: Option<&str>) -> FinishReason { + match reason { + None | Some("end_turn" | "stop_sequence") => FinishReason::Stop, + Some("max_tokens" | "model_context_window_exceeded") => FinishReason::Length, + Some("tool_use") => FinishReason::ToolCalls, + Some("guardrail_intervened" | "content_filtered") => FinishReason::ContentFilter, + Some(other) => FinishReason::Other(other.to_string()), + } +} + +/// Converse usage maps directly onto the disjoint buckets: `inputTokens` +/// already excludes cached tokens (documented), so no subtraction applies. +pub(super) fn token_counts_from_usage(usage: Option<&Value>) -> TokenCounts { + let Some(usage) = usage else { + return TokenCounts::default(); + }; + let count = |key: &str| usage.get(key).and_then(Value::as_i64).unwrap_or(0); + TokenCounts { + input_tokens: count("inputTokens"), + output_tokens: count("outputTokens"), + reasoning_tokens: 0, + cache_read_tokens: count("cacheReadInputTokens"), + cache_write_tokens: count("cacheWriteInputTokens"), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn stop_reasons_map_to_canonical_vocabulary() { + assert_eq!(map_stop_reason(Some("end_turn")), FinishReason::Stop); + assert_eq!(map_stop_reason(Some("stop_sequence")), FinishReason::Stop); + assert_eq!(map_stop_reason(Some("max_tokens")), FinishReason::Length); + assert_eq!( + map_stop_reason(Some("model_context_window_exceeded")), + FinishReason::Length + ); + assert_eq!(map_stop_reason(Some("tool_use")), FinishReason::ToolCalls); + assert_eq!( + map_stop_reason(Some("guardrail_intervened")), + FinishReason::ContentFilter + ); + assert_eq!( + map_stop_reason(Some("content_filtered")), + FinishReason::ContentFilter + ); + assert_eq!( + map_stop_reason(Some("malformed_tool_use")), + FinishReason::Other("malformed_tool_use".to_string()) + ); + assert_eq!(map_stop_reason(None), FinishReason::Stop); + } + + #[test] + fn usage_maps_without_subtraction() { + let usage = serde_json::json!({ + "inputTokens": 30, + "outputTokens": 628, + "totalTokens": 658, + "cacheReadInputTokens": 1024, + "cacheWriteInputTokens": 512, + }); + let counts = token_counts_from_usage(Some(&usage)); + assert_eq!(counts.input_tokens, 30); + assert_eq!(counts.output_tokens, 628); + assert_eq!(counts.cache_read_tokens, 1024); + assert_eq!(counts.cache_write_tokens, 512); + assert_eq!(counts.reasoning_tokens, 0); + } + + #[test] + fn bedrock_error_extracts_aws_message_shapes() { + // SigV4 path: top-level lowercase `message`. + let sigv4 = bedrock_error( + 403, + r#"{"message":"Model access is denied due to IAM ..."}"#, + "bedrock", + None, + ); + assert!( + sigv4.to_string().contains("Model access is denied"), + "{sigv4}" + ); + + // API-key path: top-level capitalized `Message`. + let api_key = bedrock_error( + 403, + r#"{"Message":"Authentication failed: Please make sure your API Key is valid."}"#, + "bedrock", + None, + ); + assert!( + api_key.to_string().contains("Authentication failed"), + "{api_key}" + ); + + // `__type` becomes the error code (tail after `#`). + let typed = bedrock_error( + 429, + r#"{"__type":"com.amazon.coral.service#ThrottlingException","message":"slow down"}"#, + "bedrock", + None, + ); + let Error::Provider { detail, .. } = &typed else { + panic!("expected provider error: {typed}"); + }; + assert_eq!(detail.error_code.as_deref(), Some("ThrottlingException")); + + // Garbage body falls back rather than panicking. + let opaque = bedrock_error(500, "not json", "bedrock", None); + assert!(opaque.to_string().contains("not json"), "{opaque}"); + } + + #[test] + fn unknown_content_blocks_are_skipped() { + assert!(decode_content_block(&serde_json::json!({"citationsContent": {}})).is_none()); + assert!(decode_content_block(&serde_json::json!({"text": ""})).is_none()); + } + + #[test] + fn reasoning_text_block_round_trips_signature() { + let block = serde_json::json!({ + "reasoningContent": { + "reasoningText": { "text": "thinking...", "signature": "sig-1" } + } + }); + let Some(ContentPart::Thinking(thinking)) = decode_content_block(&block) else { + panic!("expected thinking part"); + }; + assert_eq!(thinking.text, "thinking..."); + assert_eq!(thinking.signature.as_deref(), Some("sig-1")); + assert!(!thinking.redacted); + } +} diff --git a/lib/crates/fabro-llm/src/codec/bedrock_converse/encode.rs b/lib/crates/fabro-llm/src/codec/bedrock_converse/encode.rs new file mode 100644 index 000000000..ee1272c25 --- /dev/null +++ b/lib/crates/fabro-llm/src/codec/bedrock_converse/encode.rs @@ -0,0 +1,491 @@ +//! Request encoding: canonical `Request` → Converse envelope. + +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::error::Error; +use crate::types::{ContentPart, Message, Request, Role, ToolChoice}; + +pub(super) fn encode(ctx: &CodecCtx<'_>, stream: bool) -> Result { + let request = ctx.request; + if request.response_format.is_some() { + return Err(Error::Configuration { + message: format!( + "provider '{}' does not support response_format yet (Bedrock Converse \ + structured output is a named follow-up)", + ctx.provider_name + ), + source: None, + }); + } + + let caching = supports_prompt_cache(ctx); + let (system, conversation) = extract_system_prompt(&request.messages); + + let mut body = Map::new(); + + if let Some(system) = system { + let mut blocks = vec![json!({ "text": system })]; + if caching { + blocks.push(cache_point()); + } + body.insert("system".to_string(), Value::Array(blocks)); + } + + let mut messages = Vec::new(); + for message in conversation { + if let Some(value) = encode_message(message) { + messages.push(value); + } + } + if caching { + apply_cache_point_to_conversation_prefix(&mut messages); + } + body.insert("messages".to_string(), Value::Array(messages)); + + let mut inference = Map::new(); + if let Some(max_tokens) = request.max_tokens { + inference.insert("maxTokens".to_string(), json!(max_tokens)); + } + if let Some(temperature) = request.temperature { + inference.insert("temperature".to_string(), json!(temperature)); + } + if let Some(top_p) = request.top_p { + inference.insert("topP".to_string(), json!(top_p)); + } + if let Some(stop) = &request.stop_sequences { + if !stop.is_empty() { + inference.insert("stopSequences".to_string(), json!(stop)); + } + } + if !inference.is_empty() { + body.insert("inferenceConfig".to_string(), Value::Object(inference)); + } + + if let Some(tool_config) = encode_tool_config(request, caching) { + body.insert("toolConfig".to_string(), tool_config); + } + + let mut body = Value::Object(body); + merge_provider_options( + &mut body, + request.provider_options.as_ref(), + ctx.provider_name, + ); + + let action = if stream { + "converse-stream" + } else { + "converse" + }; + Ok(EncodedRequest { + body, + endpoint: format!("/model/{}/{action}", ctx.deployment_id), + headers: Vec::new(), + }) +} + +fn supports_prompt_cache(ctx: &CodecCtx<'_>) -> bool { + ctx.model.is_some_and(|m| m.features.prompt_cache) +} + +fn cache_point() -> Value { + json!({ "cachePoint": { "type": "default" } }) +} + +/// Encode one conversation message. Tool-role messages carry their results in +/// user-role messages (Converse has no tool role). Returns `None` when no +/// block survives translation. +fn encode_message(message: &Message) -> Option { + let role = match message.role { + Role::Assistant => "assistant", + // Tool results ride in user messages on the Converse wire. + _ => "user", + }; + + let mut blocks: Vec = message + .content + .iter() + .filter_map(encode_content_part) + .collect(); + + // Tool-role messages whose result lives on the message rather than in a + // ToolResult part. + if blocks.is_empty() && message.role == Role::Tool { + if let Some(tool_call_id) = &message.tool_call_id { + let text = message.text(); + blocks.push(json!({ + "toolResult": { + "toolUseId": tool_call_id, + "content": [{ "text": text }], + } + })); + } + } + + if blocks.is_empty() { + return None; + } + Some(json!({ "role": role, "content": blocks })) +} + +fn encode_content_part(part: &ContentPart) -> Option { + match part { + ContentPart::Text(text) => { + if text.is_empty() { + None + } else { + Some(json!({ "text": text })) + } + } + // Converse has no URL sources; the adapter's attachment resolution + // inlines file-backed parts ahead of encoding, and URL-only parts are + // dropped (the established drop-don't-fail attachment contract). + ContentPart::Image(image) => { + let bytes = image.data.as_ref()?; + Some(json!({ + "image": { + "format": media_format(image.media_type.as_deref(), "png"), + "source": { "bytes": BASE64.encode(bytes) }, + } + })) + } + ContentPart::Document(document) => { + let bytes = document.data.as_ref()?; + Some(json!({ + "document": { + "format": media_format(document.media_type.as_deref(), "pdf"), + "name": document.file_name.as_deref().unwrap_or("document"), + "source": { "bytes": BASE64.encode(bytes) }, + } + })) + } + ContentPart::ToolCall(tool_call) => Some(json!({ + "toolUse": { + "toolUseId": tool_call.id, + "name": tool_call.name, + "input": tool_call.arguments, + } + })), + ContentPart::ToolResult(result) => { + let content = match &result.content { + Value::String(text) => json!([{ "text": text }]), + other => json!([{ "json": other }]), + }; + let mut block = Map::new(); + block.insert("toolUseId".to_string(), json!(result.tool_call_id)); + block.insert("content".to_string(), content); + if result.is_error { + block.insert("status".to_string(), json!("error")); + } + Some(json!({ "toolResult": Value::Object(block) })) + } + ContentPart::Thinking(thinking) => { + if thinking.redacted { + Some(json!({ + "reasoningContent": { "redactedContent": thinking.text } + })) + } else { + let mut text_block = Map::new(); + text_block.insert("text".to_string(), json!(thinking.text)); + if let Some(signature) = &thinking.signature { + // Echoed back unmodified — Bedrock validates it. + text_block.insert("signature".to_string(), json!(signature)); + } + Some(json!({ + "reasoningContent": { "reasoningText": Value::Object(text_block) } + })) + } + } + // Audio input and opaque foreign parts have no Converse encoding. + ContentPart::Audio(_) | ContentPart::Other { .. } => None, + } +} + +/// `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() +} + +fn encode_tool_config(request: &Request, caching: bool) -> Option { + let tools = request.tools.as_ref()?; + if tools.is_empty() { + return None; + } + // `tool_choice: none` is rejected at the adapter's validate_request; + // defensively drop the toolConfig if it slips through. + if request.tool_choice == Some(ToolChoice::None) { + return None; + } + + let mut entries: Vec = tools + .iter() + .map(|tool| { + json!({ + "toolSpec": { + "name": tool.name, + "description": tool.description, + "inputSchema": { "json": tool.parameters }, + } + }) + }) + .collect(); + if caching { + entries.push(cache_point()); + } + + let mut config = Map::new(); + config.insert("tools".to_string(), Value::Array(entries)); + match &request.tool_choice { + Some(ToolChoice::Required) => { + config.insert("toolChoice".to_string(), json!({ "any": {} })); + } + Some(ToolChoice::Named { tool_name }) => { + config.insert( + "toolChoice".to_string(), + json!({ "tool": { "name": tool_name } }), + ); + } + // Auto is the wire default; ToolChoice::None dropped the config above. + Some(ToolChoice::Auto | ToolChoice::None) | None => {} + } + Some(Value::Object(config)) +} + +/// Mirror the anthropic codec's conversation-prefix cache placement: a +/// `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 = 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 target = user_indices[user_indices.len() - 2]; + if let Some(content) = messages[target] + .get_mut("content") + .and_then(Value::as_array_mut) + { + content.push(cache_point()); + } +} + +/// Merge `provider_options.` keys into the top level of the +/// body (the same adapter-name-keyed contract as the openai_compatible +/// 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()); + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + use crate::codec::CodecParams; + use crate::types::{ + ResponseFormat, ResponseFormatType, ThinkingData, ToolDefinition, ToolResult, + }; + + fn base_request(model: &str) -> Request { + Request { + model: model.to_string(), + messages: vec![Message::user("Hello")], + provider: Some("bedrock".to_string()), + tools: None, + tool_choice: None, + response_format: None, + temperature: Some(0.5), + top_p: None, + max_tokens: Some(256), + stop_sequences: None, + reasoning_effort: None, + speed: None, + metadata: None, + provider_options: None, + } + } + + fn encode_with(request: &Request) -> EncodedRequest { + let params = CodecParams::default(); + let ctx = CodecCtx { + request, + provider_name: "bedrock", + deployment_id: "us.anthropic.claude-sonnet-4-6", + model: None, + params: ¶ms, + }; + encode(&ctx, false).unwrap() + } + + #[test] + fn endpoint_carries_model_and_action() { + let request = base_request("claude"); + let params = CodecParams::default(); + let ctx = CodecCtx { + request: &request, + provider_name: "bedrock", + deployment_id: "us.anthropic.claude-sonnet-4-6", + model: None, + params: ¶ms, + }; + assert_eq!( + encode(&ctx, false).unwrap().endpoint, + "/model/us.anthropic.claude-sonnet-4-6/converse" + ); + assert_eq!( + encode(&ctx, true).unwrap().endpoint, + "/model/us.anthropic.claude-sonnet-4-6/converse-stream" + ); + } + + #[test] + fn system_messages_become_top_level_system_blocks() { + let mut request = base_request("claude"); + request.messages = vec![Message::system("Be brief"), Message::user("Hi")]; + let encoded = encode_with(&request); + assert_eq!(encoded.body["system"][0]["text"], "Be brief"); + assert_eq!(encoded.body["messages"][0]["role"], "user"); + assert_eq!(encoded.body["messages"][0]["content"][0]["text"], "Hi"); + } + + #[test] + fn inference_config_uses_camel_case() { + let encoded = encode_with(&base_request("claude")); + assert_eq!(encoded.body["inferenceConfig"]["maxTokens"], 256); + assert_eq!(encoded.body["inferenceConfig"]["temperature"], 0.5); + } + + #[test] + fn tools_encode_as_tool_specs_with_choice() { + let mut request = base_request("claude"); + request.tools = Some(vec![ToolDefinition::function( + "search", + "Search things", + json!({"type": "object"}), + )]); + request.tool_choice = Some(ToolChoice::named("search")); + let encoded = encode_with(&request); + let spec = &encoded.body["toolConfig"]["tools"][0]["toolSpec"]; + assert_eq!(spec["name"], "search"); + assert_eq!(spec["inputSchema"]["json"]["type"], "object"); + assert_eq!( + encoded.body["toolConfig"]["toolChoice"]["tool"]["name"], + "search" + ); + } + + #[test] + fn tool_results_ride_in_user_messages() { + let mut request = base_request("claude"); + request.messages = vec![Message { + role: Role::Tool, + content: vec![ContentPart::ToolResult(ToolResult { + tool_call_id: "tool-1".to_string(), + content: json!("42"), + is_error: false, + image_data: None, + image_media_type: None, + })], + name: None, + tool_call_id: Some("tool-1".to_string()), + }]; + let encoded = encode_with(&request); + let message = &encoded.body["messages"][0]; + assert_eq!(message["role"], "user"); + assert_eq!(message["content"][0]["toolResult"]["toolUseId"], "tool-1"); + assert_eq!( + message["content"][0]["toolResult"]["content"][0]["text"], + "42" + ); + } + + #[test] + fn thinking_parts_restructure_into_reasoning_text_blocks() { + let mut request = base_request("claude"); + request.messages = vec![Message { + role: Role::Assistant, + content: vec![ContentPart::Thinking(ThinkingData { + text: "prior thoughts".to_string(), + signature: Some("sig-1".to_string()), + redacted: false, + })], + name: None, + tool_call_id: None, + }]; + let encoded = encode_with(&request); + let block = &encoded.body["messages"][0]["content"][0]["reasoningContent"]["reasoningText"]; + assert_eq!(block["text"], "prior thoughts"); + assert_eq!(block["signature"], "sig-1"); + } + + #[test] + fn provider_options_merge_top_level() { + let mut request = base_request("claude"); + request.provider_options = Some(json!({ + "bedrock": { + "additionalModelRequestFields": {"top_k": 200}, + "serviceTier": {"type": "flex"} + } + })); + let encoded = encode_with(&request); + assert_eq!(encoded.body["additionalModelRequestFields"]["top_k"], 200); + assert_eq!(encoded.body["serviceTier"]["type"], "flex"); + } + + #[test] + fn response_format_is_rejected() { + let mut request = base_request("claude"); + request.response_format = Some(ResponseFormat { + kind: ResponseFormatType::JsonSchema, + json_schema: Some(json!({"type": "object"})), + strict: false, + }); + let params = CodecParams::default(); + let ctx = CodecCtx { + request: &request, + provider_name: "bedrock", + deployment_id: "m", + model: None, + params: ¶ms, + }; + assert!(encode(&ctx, false).is_err()); + } + + #[test] + fn cache_points_follow_the_anthropic_placement() { + let mut messages = vec![ + json!({"role": "user", "content": [{"text": "turn 1"}]}), + json!({"role": "assistant", "content": [{"text": "reply 1"}]}), + json!({"role": "user", "content": [{"text": "turn 2"}]}), + ]; + apply_cache_point_to_conversation_prefix(&mut messages); + // Second-to-last user message gains the cachePoint. + assert!(messages[0]["content"][1].get("cachePoint").is_some()); + assert_eq!(messages[2]["content"].as_array().unwrap().len(), 1); + } +} diff --git a/lib/crates/fabro-llm/src/codec/bedrock_converse/mod.rs b/lib/crates/fabro-llm/src/codec/bedrock_converse/mod.rs new file mode 100644 index 000000000..12b315094 --- /dev/null +++ b/lib/crates/fabro-llm/src/codec/bedrock_converse/mod.rs @@ -0,0 +1,58 @@ +//! The Amazon Bedrock Converse codec. +//! +//! Pure translation: no HTTP, auth, signing, or event-stream framing — the +//! Bedrock adapter shell owns those. Converse is Bedrock's model-agnostic +//! envelope (AWS translates it to each hosted family's native dialect +//! server-side), which is what makes this one codec serve Claude, Nova, +//! Llama, Mistral, DeepSeek, Qwen, Kimi, GLM, MiniMax, Nemotron, and +//! gpt-oss alike. The codec fully forms its endpoints (model-in-path, +//! `/converse` vs `/converse-stream`), mirrors the anthropic codec's prompt +//! cache placement with `cachePoint` blocks, and round-trips +//! `reasoningContent` thinking signatures unmodified. + +mod decode; +mod encode; +mod stream; + +use crate::codec::{Codec, CodecCtx, EncodedRequest, StreamDecoder}; +use crate::error::Error; +use crate::types::{RateLimitInfo, Response}; + +/// Codec for the Bedrock Converse wire dialect. +pub(crate) struct BedrockConverse; + +impl Codec for BedrockConverse { + fn encode(&self, ctx: &CodecCtx<'_>, stream: bool) -> Result { + encode::encode(ctx, stream) + } + + fn decode_response( + &self, + body: &str, + ctx: &CodecCtx<'_>, + rate_limit: Option, + ) -> Result { + decode::decode_response(body, ctx, rate_limit) + } + + fn stream_decoder( + &self, + ctx: &CodecCtx<'_>, + rate_limit: Option, + ) -> Box { + Box::new(stream::ConverseStreamDecoder::new(ctx, rate_limit)) + } + + /// Bedrock error bodies are AWS-shaped (top-level `message`/`Message`, + /// `__type`), which the default parser misses — extract them so failures + /// surface the real reason instead of "Unknown error". + fn decode_error( + &self, + status: u16, + body: &str, + ctx: &CodecCtx<'_>, + retry_after: Option, + ) -> Error { + decode::bedrock_error(status, body, ctx.provider_name, retry_after) + } +} diff --git a/lib/crates/fabro-llm/src/codec/bedrock_converse/stream.rs b/lib/crates/fabro-llm/src/codec/bedrock_converse/stream.rs new file mode 100644 index 000000000..c0008e819 --- /dev/null +++ b/lib/crates/fabro-llm/src/codec/bedrock_converse/stream.rs @@ -0,0 +1,489 @@ +//! Streaming decoder: ConverseStream events → canonical `StreamEvent`s. +//! +//! Event names arrive in the transport's `RawEvent::event` (the frame's +//! `:event-type` header); payloads are the event JSON. The documented +//! sequence is `messageStart` → per content block (`contentBlockStart` +//! [tool use only] → `contentBlockDelta`* → `contentBlockStop`) → +//! `messageStop{stopReason}` → `metadata{usage}`. Usage arrives ONLY in the +//! terminal `metadata` event, which is also where the final `Finish` is +//! synthesized. + +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::error::Error; +use crate::types::{ + ContentPart, FinishReason, Message, RateLimitInfo, Response, Role, StreamEvent, ThinkingData, + TokenCounts, ToolCall, +}; + +/// Per-content-block accumulation state, keyed by `contentBlockIndex`. +enum BlockState { + Text(String), + Reasoning { + text: String, + signature: Option, + redacted: Option, + }, + ToolUse { + id: String, + name: String, + input: String, + }, +} + +/// Accumulated state while decoding one ConverseStream response. +pub(super) struct ConverseStreamDecoder { + provider_name: String, + model: String, + blocks: BTreeMap, + /// Completed blocks in arrival order, for the final response message. + parts: Vec, + finish_reason: FinishReason, + usage: TokenCounts, + text_started: bool, + finished: bool, + rate_limit: Option, +} + +impl ConverseStreamDecoder { + pub(super) fn new(ctx: &CodecCtx<'_>, rate_limit: Option) -> Self { + Self { + provider_name: ctx.provider_name.to_string(), + model: ctx.request.model.clone(), + blocks: BTreeMap::new(), + parts: Vec::new(), + finish_reason: FinishReason::Stop, + usage: TokenCounts::default(), + text_started: false, + finished: false, + rate_limit, + } + } + + fn block_index(payload: &Value) -> u64 { + payload + .get("contentBlockIndex") + .and_then(Value::as_u64) + .unwrap_or(0) + } + + fn on_block_start(&mut self, payload: &Value) -> Vec { + let index = Self::block_index(payload); + if let Some(tool_use) = payload.pointer("/start/toolUse") { + let id = tool_use + .get("toolUseId") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(); + let name = tool_use + .get("name") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(); + let started = ToolCall::new(&id, &name, Value::Null); + self.blocks.insert(index, BlockState::ToolUse { + id, + name, + input: String::new(), + }); + return vec![StreamEvent::ToolCallStart { tool_call: started }]; + } + Vec::new() + } + + fn on_block_delta(&mut self, payload: &Value) -> Vec { + let index = Self::block_index(payload); + let Some(delta) = payload.get("delta") else { + return Vec::new(); + }; + + if let Some(text) = delta.get("text").and_then(Value::as_str) { + if text.is_empty() { + return Vec::new(); + } + let mut events = Vec::new(); + if !self.text_started { + self.text_started = true; + events.push(StreamEvent::TextStart { text_id: None }); + } + match self + .blocks + .entry(index) + .or_insert_with(|| BlockState::Text(String::new())) + { + BlockState::Text(buffer) => buffer.push_str(text), + // A text delta against a non-text block: tolerate by ignoring + // the mismatch rather than corrupting tool/reasoning state. + _ => return events, + } + events.push(StreamEvent::text_delta(text, None)); + return events; + } + + if let Some(input) = delta.pointer("/toolUse/input").and_then(Value::as_str) { + if let Some(BlockState::ToolUse { + id, + name, + input: buffer, + }) = self.blocks.get_mut(&index) + { + buffer.push_str(input); + let partial = ToolCall::new(id.as_str(), name.as_str(), Value::Null); + return vec![StreamEvent::ToolCallDelta { tool_call: partial }]; + } + return Vec::new(); + } + + if let Some(reasoning) = delta.get("reasoningContent") { + let entry = self + .blocks + .entry(index) + .or_insert_with(|| BlockState::Reasoning { + text: String::new(), + signature: None, + redacted: None, + }); + let BlockState::Reasoning { + text, + signature, + redacted, + } = entry + else { + return Vec::new(); + }; + let mut events = Vec::new(); + if text.is_empty() && signature.is_none() && redacted.is_none() { + events.push(StreamEvent::ReasoningStart); + } + // Streaming reasoning deltas carry text/signature as FLAT union + // members (unlike the nested request-side reasoningText block). + if let Some(fragment) = reasoning.get("text").and_then(Value::as_str) { + text.push_str(fragment); + events.push(StreamEvent::ReasoningDelta { + delta: fragment.to_string(), + }); + } + if let Some(sig) = reasoning.get("signature").and_then(Value::as_str) { + *signature = Some(sig.to_string()); + } + if let Some(blob) = reasoning.get("redactedContent").and_then(Value::as_str) { + *redacted = Some(blob.to_string()); + } + return events; + } + + Vec::new() + } + + fn on_block_stop(&mut self, payload: &Value) -> Vec { + let index = Self::block_index(payload); + let Some(block) = self.blocks.remove(&index) else { + return Vec::new(); + }; + match block { + BlockState::Text(text) => { + let mut events = Vec::new(); + if self.text_started { + self.text_started = false; + events.push(StreamEvent::TextEnd { text_id: None }); + } + if !text.is_empty() { + self.parts.push(ContentPart::text(&text)); + } + events + } + BlockState::Reasoning { + text, + signature, + redacted, + } => { + let part = if let Some(blob) = redacted { + ThinkingData { + text: blob, + signature: None, + redacted: true, + } + } else { + ThinkingData { + text, + signature, + redacted: false, + } + }; + self.parts.push(ContentPart::Thinking(part)); + vec![StreamEvent::ReasoningEnd] + } + BlockState::ToolUse { id, name, input } => { + let arguments = serde_json::from_str(&input).unwrap_or(Value::Null); + let mut tool_call = ToolCall::new(&id, &name, arguments); + tool_call.raw_arguments = Some(input); + self.parts.push(ContentPart::ToolCall(tool_call.clone())); + vec![StreamEvent::ToolCallEnd { tool_call }] + } + } + } + + /// Build the final `Finish` from accumulated state. + fn finish_event(&mut self) -> StreamEvent { + self.finished = true; + // Flush any blocks that never saw a contentBlockStop. + let dangling: Vec = self.blocks.keys().copied().collect(); + for index in dangling { + let _ = self.on_block_stop(&serde_json::json!({ "contentBlockIndex": index })); + } + + let response = Response { + id: uuid::Uuid::new_v4().to_string(), + model: self.model.clone(), + provider: self.provider_name.clone(), + message: Message { + role: Role::Assistant, + content: std::mem::take(&mut self.parts), + name: None, + tool_call_id: None, + }, + finish_reason: self.finish_reason.clone(), + usage: self.usage.clone(), + raw: None, + warnings: vec![], + rate_limit: self.rate_limit.clone(), + cost_usd: None, + cost_source: None, + }; + StreamEvent::finish(self.finish_reason.clone(), self.usage.clone(), response) + } +} + +impl StreamDecoder for ConverseStreamDecoder { + fn on_event(&mut self, ev: RawEvent<'_>) -> Result, Error> { + let Some(event_type) = ev.event else { + return Ok(Vec::new()); + }; + let payload: Value = serde_json::from_str(ev.data) + .map_err(|e| Error::stream_error(format!("converse stream event json: {e}"), e))?; + + Ok(match event_type { + "messageStart" => vec![StreamEvent::StreamStart], + "contentBlockStart" => self.on_block_start(&payload), + "contentBlockDelta" => self.on_block_delta(&payload), + "contentBlockStop" => self.on_block_stop(&payload), + "messageStop" => { + self.finish_reason = + map_stop_reason(payload.get("stopReason").and_then(Value::as_str)); + Vec::new() + } + "metadata" => { + self.usage = token_counts_from_usage(payload.get("usage")); + vec![self.finish_event()] + } + // Tolerate unknown event types — the union grows. + _ => Vec::new(), + }) + } + + /// Byte-stream end: `metadata` is the documented terminus, but if the + /// stream ends without one, synthesize the `Finish` from accumulated + /// state so callers still receive a response (mirrors the gemini + /// decoder's unconditional synthesis). + fn finish(&mut self) -> Vec { + if self.finished { + return Vec::new(); + } + vec![self.finish_event()] + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::codec::CodecParams; + use crate::types::{Message as RequestMessage, Request}; + + fn decoder() -> ConverseStreamDecoder { + let request = Request { + model: "us.anthropic.claude-sonnet-4-6".to_string(), + messages: vec![RequestMessage::user("hi")], + provider: Some("bedrock".to_string()), + tools: None, + tool_choice: None, + response_format: None, + temperature: None, + top_p: None, + max_tokens: None, + stop_sequences: None, + reasoning_effort: None, + speed: None, + metadata: None, + provider_options: None, + }; + let params = CodecParams::default(); + let ctx = CodecCtx { + request: &request, + provider_name: "bedrock", + deployment_id: "us.anthropic.claude-sonnet-4-6", + model: None, + params: ¶ms, + }; + ConverseStreamDecoder::new(&ctx, None) + } + + fn feed(decoder: &mut ConverseStreamDecoder, event: &str, data: &str) -> Vec { + decoder + .on_event(RawEvent { + event: Some(event), + data, + }) + .unwrap() + } + + #[test] + fn text_happy_path_finishes_on_metadata() { + let mut d = decoder(); + assert!(matches!( + feed(&mut d, "messageStart", r#"{"role":"assistant"}"#)[0], + StreamEvent::StreamStart + )); + let events = feed( + &mut d, + "contentBlockDelta", + r#"{"delta":{"text":"Hel"},"contentBlockIndex":0}"#, + ); + assert!(matches!(events[0], StreamEvent::TextStart { .. })); + assert!(matches!(events[1], StreamEvent::TextDelta { .. })); + feed( + &mut d, + "contentBlockDelta", + r#"{"delta":{"text":"lo"},"contentBlockIndex":0}"#, + ); + let stop = feed(&mut d, "contentBlockStop", r#"{"contentBlockIndex":0}"#); + assert!(matches!(stop[0], StreamEvent::TextEnd { .. })); + assert!(feed(&mut d, "messageStop", r#"{"stopReason":"end_turn"}"#).is_empty()); + + let finish = feed( + &mut d, + "metadata", + r#"{"usage":{"inputTokens":12,"outputTokens":5,"totalTokens":17}}"#, + ); + let StreamEvent::Finish { + finish_reason, + usage, + response, + } = &finish[0] + else { + panic!("expected Finish"); + }; + assert_eq!(*finish_reason, FinishReason::Stop); + assert_eq!(usage.input_tokens, 12); + assert_eq!(response.text(), "Hello"); + assert_eq!(response.provider, "bedrock"); + // Byte-stream end after metadata adds nothing. + assert!(d.finish().is_empty()); + } + + #[test] + fn tool_use_accumulates_string_input_fragments() { + let mut d = decoder(); + feed(&mut d, "messageStart", r#"{"role":"assistant"}"#); + let start = feed( + &mut d, + "contentBlockStart", + r#"{"start":{"toolUse":{"toolUseId":"tool-1","name":"search"}},"contentBlockIndex":0}"#, + ); + assert!(matches!(start[0], StreamEvent::ToolCallStart { .. })); + feed( + &mut d, + "contentBlockDelta", + r#"{"delta":{"toolUse":{"input":"{\"que"}},"contentBlockIndex":0}"#, + ); + feed( + &mut d, + "contentBlockDelta", + r#"{"delta":{"toolUse":{"input":"ry\":\"foo\"}"}},"contentBlockIndex":0}"#, + ); + let stop = feed(&mut d, "contentBlockStop", r#"{"contentBlockIndex":0}"#); + let StreamEvent::ToolCallEnd { tool_call } = &stop[0] else { + panic!("expected ToolCallEnd"); + }; + assert_eq!(tool_call.id, "tool-1"); + assert_eq!(tool_call.arguments["query"], "foo"); + + feed(&mut d, "messageStop", r#"{"stopReason":"tool_use"}"#); + let finish = feed( + &mut d, + "metadata", + r#"{"usage":{"inputTokens":1,"outputTokens":1}}"#, + ); + let StreamEvent::Finish { finish_reason, .. } = &finish[0] else { + panic!("expected Finish"); + }; + assert_eq!(*finish_reason, FinishReason::ToolCalls); + } + + #[test] + fn reasoning_deltas_round_trip_signature() { + let mut d = decoder(); + let events = feed( + &mut d, + "contentBlockDelta", + r#"{"delta":{"reasoningContent":{"text":"thinking"}},"contentBlockIndex":0}"#, + ); + assert!(matches!(events[0], StreamEvent::ReasoningStart)); + assert!(matches!(events[1], StreamEvent::ReasoningDelta { .. })); + feed( + &mut d, + "contentBlockDelta", + r#"{"delta":{"reasoningContent":{"signature":"sig-9"}},"contentBlockIndex":0}"#, + ); + let stop = feed(&mut d, "contentBlockStop", r#"{"contentBlockIndex":0}"#); + assert!(matches!(stop[0], StreamEvent::ReasoningEnd)); + + let finish = feed( + &mut d, + "metadata", + r#"{"usage":{"inputTokens":1,"outputTokens":1}}"#, + ); + let StreamEvent::Finish { response, .. } = &finish[0] else { + panic!("expected Finish"); + }; + let ContentPart::Thinking(thinking) = &response.message.content[0] else { + panic!("expected thinking part"); + }; + assert_eq!(thinking.text, "thinking"); + assert_eq!(thinking.signature.as_deref(), Some("sig-9")); + } + + #[test] + fn stream_end_without_metadata_synthesizes_finish() { + let mut d = decoder(); + feed( + &mut d, + "contentBlockDelta", + r#"{"delta":{"text":"partial"},"contentBlockIndex":0}"#, + ); + let events = d.finish(); + let StreamEvent::Finish { response, .. } = &events[0] else { + panic!("expected synthesized Finish"); + }; + assert_eq!(response.text(), "partial"); + // Synthesis happens once. + assert!(d.finish().is_empty()); + } + + #[test] + fn unknown_events_are_tolerated() { + let mut d = decoder(); + assert!(feed(&mut d, "futureEventKind", r#"{"anything":1}"#).is_empty()); + assert!( + d.on_event(RawEvent { + event: None, + data: "{}", + }) + .unwrap() + .is_empty() + ); + } +} diff --git a/lib/crates/fabro-llm/src/codec/mod.rs b/lib/crates/fabro-llm/src/codec/mod.rs index e8d8d6e2a..5e770794c 100644 --- a/lib/crates/fabro-llm/src/codec/mod.rs +++ b/lib/crates/fabro-llm/src/codec/mod.rs @@ -11,6 +11,7 @@ //! methods, never extend the contract. pub(crate) mod anthropic_messages; +pub(crate) mod bedrock_converse; pub(crate) mod gemini_generate; pub(crate) mod openai_compatible; pub(crate) mod openai_responses; diff --git a/lib/crates/fabro-llm/src/providers/bedrock/eventstream.rs b/lib/crates/fabro-llm/src/providers/bedrock/eventstream.rs index 70cd5b5e8..bb4015eff 100644 --- a/lib/crates/fabro-llm/src/providers/bedrock/eventstream.rs +++ b/lib/crates/fabro-llm/src/providers/bedrock/eventstream.rs @@ -17,6 +17,7 @@ use crate::error::Error; /// One decoded ConverseStream event: the `:event-type` header value plus the /// frame's JSON payload, ready to feed a stream decoder. +#[derive(Debug)] pub(crate) struct DecodedEvent { pub event_type: String, pub payload: String, @@ -126,7 +127,7 @@ pub(crate) mod tests { )) .add_header(Header::new( ":event-type", - HeaderValue::String(event_type.into()), + HeaderValue::String(event_type.to_string().into()), )) .add_header(Header::new( ":content-type", diff --git a/lib/crates/fabro-llm/src/providers/bedrock/mod.rs b/lib/crates/fabro-llm/src/providers/bedrock/mod.rs index 833e49e19..c1bf7fb46 100644 --- a/lib/crates/fabro-llm/src/providers/bedrock/mod.rs +++ b/lib/crates/fabro-llm/src/providers/bedrock/mod.rs @@ -1,21 +1,37 @@ -//! Amazon Bedrock transport primitives: SigV4 signing, AWS event-stream -//! decoding, auth-mode selection, and region derivation. +//! Provider adapter for Amazon Bedrock (Converse/ConverseStream). //! -//! The adapter that composes these over the `bedrock_converse` codec lands -//! later in this series; until then the pieces carry ahead-of-use allows. +//! A thin transport shell over the `bedrock_converse` codec: it owns auth +//! (SigV4 signing or a bearer Bedrock API key), the region derivation, and +//! the AWS event-stream byte loop. All wire translation lives in the codec; +//! one codec serves every Converse-capable family because AWS translates the +//! envelope server-side. pub(crate) mod eventstream; pub(crate) mod sigv4; -use tokio::sync::OnceCell; +use std::collections::VecDeque; +use std::sync::Arc; +use std::time::Duration; +use eventstream::FrameDecoder; +use fabro_auth::ApiKeyHeader; +use fabro_model::Catalog; +use futures::stream; +use sigv4::Sigv4Signer; +use tokio::sync::OnceCell; +use tokio::time; + +use crate::adapter_registry::AdapterConfig; +use crate::attachments::{self, AttachmentPolicy}; +use crate::codec::bedrock_converse::BedrockConverse; +use crate::codec::{Codec, CodecCtx, CodecParams, EncodedRequest, RawEvent, StreamDecoder}; use crate::error::Error; +use crate::provider::{self, ProviderAdapter, StreamEventStream}; +use crate::providers::common::{self as common}; +use crate::transport::{self, HttpTransport}; +use crate::types::{AdapterTimeout, Request, Response, StreamEvent}; /// How the adapter authenticates to Bedrock. -#[expect( - dead_code, - reason = "Consumed by the Bedrock adapter later in this series." -)] pub(crate) enum BedrockAuth { /// Bedrock API key, sent as an `Authorization: Bearer` token. ApiKey(String), @@ -23,7 +39,329 @@ pub(crate) enum BedrockAuth { /// is resolved on first use and cached; the chain itself re-resolves /// expiring credentials per request. Tests pre-seed the cell with a /// static signer. - Sigv4(OnceCell), + Sigv4(OnceCell), +} + +/// Build a boxed Bedrock adapter from a resolved [`AdapterConfig`]. +/// +/// Kept in this module (rather than the generic adapter registry) so that +/// Bedrock-specific construction stays encapsulated here. The auth mode is +/// implied by the resolved credential: an `aws_sigv4` credential signs with +/// the AWS chain; a static token is sent as a bearer API key. +pub(crate) fn build(config: AdapterConfig) -> Result, Error> { + let base_url = config + .base_url + .clone() + .ok_or_else(|| Error::Configuration { + message: format!( + "bedrock provider '{}' requires a base_url (the Bedrock runtime endpoint)", + config.provider_id + ), + source: None, + })?; + 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)? + } + None => { + return Err(Error::Configuration { + message: format!( + "bedrock provider '{}' has no resolved credential (configure `aws_sigv4` or \ + an API key)", + config.provider_id + ), + source: None, + }); + } + }; + let mut adapter = adapter.with_name(config.provider_id); + if let Some(catalog) = config.catalog { + adapter = adapter.with_catalog(catalog); + } + Ok(Arc::new(adapter)) +} + +/// Provider adapter for Amazon Bedrock. +pub struct Adapter { + pub(crate) http: HttpTransport, + provider_name: String, + region: String, + auth: BedrockAuth, + catalog: Option>, +} + +impl Adapter { + /// Construct an adapter that authenticates with a Bedrock API key. + /// `base_url` is the Bedrock runtime endpoint; the signing region is + /// parsed from it. + pub fn new_api_key( + token: impl Into, + base_url: impl Into, + ) -> Result { + Self::with_auth(base_url, BedrockAuth::ApiKey(token.into())) + } + + /// Construct a SigV4 adapter. Credentials resolve lazily from the AWS + /// default chain on the first request, so construction stays synchronous. + pub fn new_sigv4(base_url: impl Into) -> Result { + Self::with_auth(base_url, BedrockAuth::Sigv4(OnceCell::new())) + } + + fn with_auth(base_url: impl Into, auth: BedrockAuth) -> Result { + let base_url = base_url.into(); + let region = region_from_base_url(&base_url)?; + Ok(Self { + http: HttpTransport::new_optional(None, base_url), + provider_name: "bedrock".to_string(), + region, + auth, + catalog: None, + }) + } + + #[must_use] + pub fn with_name(mut self, name: impl Into) -> Self { + self.provider_name = name.into(); + self + } + + #[must_use] + pub fn with_catalog(mut self, catalog: Arc) -> Self { + self.catalog = Some(catalog); + self + } + + #[must_use] + pub fn with_timeout(self, timeout: AdapterTimeout) -> Self { + Self { + http: self.http.with_timeout(timeout), + ..self + } + } + + fn codec_ctx<'a>( + &'a self, + request: &'a Request, + deployment_id: &'a str, + params: &'a CodecParams, + ) -> CodecCtx<'a> { + CodecCtx { + request, + provider_name: &self.provider_name, + deployment_id, + model: common::catalog_model(self.catalog.as_deref(), &request.model), + params, + } + } + + /// Resolve file-backed attachments to inline data first: Converse takes + /// inline image and document bytes (no URL sources). + async fn resolve_request<'a>(&self, request: &'a Request) -> std::borrow::Cow<'a, Request> { + let policy = AttachmentPolicy { + images: true, + documents: true, + audio: false, + }; + attachments::resolve(request, policy).await + } + + /// Build the signed/bearer HTTP request for an encoded Converse call. + async fn build_http_request( + &self, + encoded: &EncodedRequest, + stream: bool, + ) -> Result { + let url = format!("{}{}", self.http.base_url, encoded.endpoint); + let body = serde_json::to_vec(&encoded.body).map_err(|e| Error::Configuration { + message: format!("failed to serialize converse request: {e}"), + source: None, + })?; + + let mut req = self.http.client.post(&url); + for (key, value) in &self.http.default_headers { + req = req.header(key, value); + } + + req = match &self.auth { + BedrockAuth::ApiKey(token) => req.bearer_auth(token).body(body), + BedrockAuth::Sigv4(cell) => { + let signer = cell + .get_or_try_init(Sigv4Signer::from_default_chain) + .await?; + signer.sign_post(req, &self.region, &url, &body).await? + } + }; + + req = req.header("content-type", "application/json"); + if stream { + req = req.header("accept", "application/vnd.amazon.eventstream"); + } + if let Some(t) = self.http.request_timeout { + if !stream { + req = req.timeout(t); + } + } + Ok(req) + } +} + +#[async_trait::async_trait] +impl ProviderAdapter for Adapter { + fn name(&self) -> &str { + &self.provider_name + } + + async fn complete(&self, request: &Request) -> Result { + self.validate_request(request)?; + + let resolved = self.resolve_request(request).await; + let codec = BedrockConverse; + let deployment_id = common::api_model_id(self.catalog.as_deref(), &resolved.model); + let params = CodecParams::default(); + let ctx = self.codec_ctx(&resolved, &deployment_id, ¶ms); + + let encoded = codec.encode(&ctx, false)?; + let req = self.build_http_request(&encoded, false).await?; + transport::complete_via_http(req, &codec, &ctx).await + } + + async fn stream(&self, request: &Request) -> Result { + self.validate_request(request)?; + + let resolved = self.resolve_request(request).await; + let codec = BedrockConverse; + let deployment_id = common::api_model_id(self.catalog.as_deref(), &resolved.model); + let params = CodecParams::default(); + let ctx = self.codec_ctx(&resolved, &deployment_id, ¶ms); + + let encoded = codec.encode(&ctx, true)?; + let req = self.build_http_request(&encoded, true).await?; + + let http_resp = req + .send() + .await + .map_err(|e| Error::network(e.to_string(), e))?; + let status = http_resp.status(); + if !status.is_success() { + let retry_after = transport::parse_retry_after(http_resp.headers()); + let body = http_resp + .text() + .await + .map_err(|e| Error::network(e.to_string(), e))?; + return Err(codec.decode_error(status.as_u16(), &body, &ctx, retry_after)); + } + + let rate_limit = transport::parse_rate_limit_headers(http_resp.headers()); + let decoder = codec.stream_decoder(&ctx, rate_limit); + Ok(decode_eventstream( + http_resp, + decoder, + self.http.stream_read_timeout, + )) + } + + fn supports_tool_choice(&self, mode: &str) -> bool { + // Converse has no `none` tool choice on the wire. + matches!(mode, "auto" | "required" | "named") + } + + fn validate_request(&self, request: &Request) -> Result<(), Error> { + if let Some(tool_choice) = &request.tool_choice { + provider::validate_tool_choice(self, tool_choice)?; + } + Ok(()) + } +} + +/// State driving the event-stream byte loop: the codec's decoder plus the +/// frame decoder, with a buffer that flattens batched events. +struct EventStreamLoop { + response: fabro_http::Response, + frames: FrameDecoder, + decoder: Box, + pending: VecDeque, + done: bool, + /// `finish()` already drained. + finished: bool, + timeout: Option, +} + +/// Drive `decoder` over the AWS event-stream byte stream of `response`: the +/// event-stream sibling of the transport's shared SSE loop, anticipated by +/// the transport consolidation notes. +fn decode_eventstream( + response: fabro_http::Response, + decoder: Box, + timeout: Option, +) -> StreamEventStream { + let out = stream::unfold( + EventStreamLoop { + response, + frames: FrameDecoder::new(), + decoder, + pending: VecDeque::new(), + done: false, + finished: false, + timeout, + }, + move |mut state| async move { + loop { + if let Some(event) = state.pending.pop_front() { + return Some((Ok(event), state)); + } + + if state.done { + if state.finished { + return None; + } + state.finished = true; + state.pending.extend(state.decoder.finish()); + if state.pending.is_empty() { + return None; + } + continue; + } + + let chunk_result = match state.timeout { + Some(timeout) => time::timeout(timeout, state.response.chunk()).await, + None => Ok(state.response.chunk().await), + }; + match chunk_result { + Ok(Ok(Some(bytes))) => { + let frames = match state.frames.push(&bytes) { + Ok(frames) => frames, + Err(e) => return Some((Err(e), state)), + }; + for frame in frames { + let raw = RawEvent { + event: Some(frame.event_type.as_str()), + data: frame.payload.as_str(), + }; + match state.decoder.on_event(raw) { + Ok(events) => state.pending.extend(events), + Err(e) => return Some((Err(e), state)), + } + } + } + Ok(Ok(None)) => state.done = true, + Ok(Err(e)) => { + return Some((Err(Error::stream_error(e.to_string(), e)), state)); + } + Err(_) => { + return Some(( + Err(Error::Stream { + message: "stream read timed out waiting for next event".to_string(), + source: None, + }), + state, + )); + } + } + } + }, + ); + Box::pin(out) } /// Derive the AWS region from a Bedrock runtime endpoint URL. @@ -32,10 +370,6 @@ pub(crate) enum BedrockAuth { /// configured base URL rather than carried as a separate AWS-specific config /// field. It is validated as `[a-z0-9-]` since it ultimately appears in a /// signed request. -#[expect( - dead_code, - reason = "Consumed by the Bedrock adapter later in this series." -)] fn region_from_base_url(base_url: &str) -> Result { let invalid = || Error::Configuration { message: format!( @@ -70,7 +404,43 @@ fn region_from_base_url(base_url: &str) -> Result { #[cfg(test)] mod tests { + use futures::StreamExt; + use httpmock::prelude::*; + use super::*; + use crate::types::{FinishReason, Message}; + + fn make_request(model: &str) -> Request { + Request { + model: model.to_string(), + messages: vec![Message::user("Hello")], + provider: Some("bedrock".to_string()), + tools: None, + tool_choice: None, + response_format: None, + temperature: None, + top_p: None, + max_tokens: Some(64), + stop_sequences: None, + reasoning_effort: None, + speed: None, + metadata: None, + provider_options: None, + } + } + + /// Adapter pointed at httpmock: region parsing only applies to real + /// bedrock-runtime URLs, so the test constructor sets the region field + /// directly. + fn test_adapter(server: &MockServer) -> Adapter { + Adapter { + http: HttpTransport::new_optional(None, server.base_url()), + provider_name: "bedrock".to_string(), + region: "us-east-1".to_string(), + auth: BedrockAuth::ApiKey("test-bedrock-key".to_string()), + catalog: None, + } + } #[test] fn region_parses_from_standard_endpoint() { @@ -108,4 +478,135 @@ mod tests { assert!(region_from_base_url(url).is_err(), "{url}"); } } + + #[tokio::test] + async fn complete_posts_converse_body_with_bearer_auth() { + let server = MockServer::start(); + let mock = server.mock(|when, then| { + when.method(POST) + .path("/model/us.anthropic.claude-sonnet-4-6/converse") + .header("authorization", "Bearer test-bedrock-key") + .json_body_includes( + r#"{"messages":[{"role":"user","content":[{"text":"Hello"}]}],"inferenceConfig":{"maxTokens":64}}"#, + ); + then.status(200) + .header("content-type", "application/json") + .json_body(serde_json::json!({ + "output": {"message": {"role": "assistant", "content": [{"text": "Hi!"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 8, "outputTokens": 2, "totalTokens": 10} + })); + }); + + let adapter = test_adapter(&server); + let response = adapter + .complete(&make_request("us.anthropic.claude-sonnet-4-6")) + .await + .unwrap(); + + mock.assert(); + assert_eq!(response.text(), "Hi!"); + assert_eq!(response.finish_reason, FinishReason::Stop); + assert_eq!(response.usage.input_tokens, 8); + assert_eq!(response.provider, "bedrock"); + } + + #[tokio::test] + async fn complete_signs_with_sigv4_when_configured() { + let server = MockServer::start(); + let mock = server.mock(|when, then| { + when.method(POST) + .path("/model/m/converse") + .header_exists("authorization") + .header_exists("x-amz-date"); + 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 mut adapter = test_adapter(&server); + let cell = OnceCell::new(); + cell.set(Sigv4Signer::from_static("AKIDEXAMPLE", "secret", None)) + .ok(); + adapter.auth = BedrockAuth::Sigv4(cell); + + let response = adapter.complete(&make_request("m")).await.unwrap(); + mock.assert(); + assert_eq!(response.text(), "ok"); + } + + #[tokio::test] + async fn stream_decodes_eventstream_frames() { + let server = MockServer::start(); + let body = eventstream::tests::build_stream_body(&[ + ("messageStart", r#"{"role":"assistant"}"#), + ( + "contentBlockDelta", + r#"{"delta":{"text":"Hel"},"contentBlockIndex":0}"#, + ), + ( + "contentBlockDelta", + r#"{"delta":{"text":"lo"},"contentBlockIndex":0}"#, + ), + ("contentBlockStop", r#"{"contentBlockIndex":0}"#), + ("messageStop", r#"{"stopReason":"end_turn"}"#), + ( + "metadata", + r#"{"usage":{"inputTokens":9,"outputTokens":3,"totalTokens":12}}"#, + ), + ]); + server.mock(|when, then| { + when.method(POST) + .path("/model/m/converse-stream") + .header("accept", "application/vnd.amazon.eventstream"); + then.status(200) + .header("content-type", "application/vnd.amazon.eventstream") + .body(body); + }); + + let adapter = test_adapter(&server); + let mut stream = adapter.stream(&make_request("m")).await.unwrap(); + + let mut text = String::new(); + let mut finish: Option = None; + while let Some(event) = stream.next().await { + match event.unwrap() { + StreamEvent::TextDelta { delta, .. } => text.push_str(&delta), + StreamEvent::Finish { response, .. } => finish = Some(*response), + _ => {} + } + } + assert_eq!(text, "Hello"); + let response = finish.expect("stream should finish"); + assert_eq!(response.text(), "Hello"); + assert_eq!(response.usage.input_tokens, 9); + } + + #[tokio::test] + async fn stream_surfaces_http_error_before_bytes() { + let server = MockServer::start(); + server.mock(|when, then| { + when.method(POST).path("/model/m/converse-stream"); + then.status(429) + .json_body(serde_json::json!({"message": "Too many requests"})); + }); + + let adapter = test_adapter(&server); + let Err(err) = adapter.stream(&make_request("m")).await else { + panic!("expected an HTTP error before any stream bytes"); + }; + assert_eq!(err.status_code(), Some(429)); + } + + #[test] + fn tool_choice_none_is_rejected() { + let server = MockServer::start(); + let adapter = test_adapter(&server); + assert!(!adapter.supports_tool_choice("none")); + assert!(adapter.supports_tool_choice("auto")); + } } diff --git a/lib/crates/fabro-llm/src/providers/bedrock/sigv4.rs b/lib/crates/fabro-llm/src/providers/bedrock/sigv4.rs index 1f327bfc0..08d89c1a8 100644 --- a/lib/crates/fabro-llm/src/providers/bedrock/sigv4.rs +++ b/lib/crates/fabro-llm/src/providers/bedrock/sigv4.rs @@ -22,6 +22,7 @@ pub(crate) const SERVICE: &str = "bedrock"; /// Where the signer's credentials come from. enum CredentialSource { /// Fixed credentials (tests / explicitly supplied keys). + #[cfg(test)] Static(Credentials), /// The AWS default provider chain. Credentials are resolved per request /// so expiring session credentials (STS, IRSA, instance roles) refresh @@ -77,6 +78,7 @@ impl Sigv4Signer { use aws_credential_types::provider::ProvideCredentials; match &self.credentials { + #[cfg(test)] CredentialSource::Static(credentials) => Ok(credentials.clone()), CredentialSource::Chain(provider) => { provider