diff --git a/lib/crates/fabro-llm/src/providers/bedrock/eventstream.rs b/lib/crates/fabro-llm/src/providers/bedrock/eventstream.rs new file mode 100644 index 000000000..70cd5b5e8 --- /dev/null +++ b/lib/crates/fabro-llm/src/providers/bedrock/eventstream.rs @@ -0,0 +1,218 @@ +//! Decoder for Bedrock's `application/vnd.amazon.eventstream` streaming +//! responses. +//! +//! ConverseStream wraps each event in a binary event-stream frame: the event +//! name (`messageStart`, `contentBlockDelta`, `metadata`, ...) travels in the +//! frame's `:event-type` header and the payload is that event's JSON +//! directly. (The base64 `{"bytes": ...}` wrapping belongs to +//! `InvokeModelWithResponseStream`'s `PayloadPart` and does not apply here.) +//! Exception and error frames are surfaced as stream errors. + +use aws_smithy_eventstream::frame::{DecodedFrame, MessageFrameDecoder}; +use aws_smithy_types::event_stream::Message; +use aws_smithy_types::str_bytes::StrBytes; +use bytes::BytesMut; + +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. +pub(crate) struct DecodedEvent { + pub event_type: String, + pub payload: String, +} + +/// Incremental decoder over event-stream bytes. +pub(crate) struct FrameDecoder { + inner: MessageFrameDecoder, + buffer: BytesMut, +} + +impl FrameDecoder { + pub(crate) fn new() -> Self { + Self { + inner: MessageFrameDecoder::new(), + buffer: BytesMut::new(), + } + } + + /// Feed newly received bytes and return any complete events decoded from + /// them. Bedrock exception and error frames are surfaced as errors. + pub(crate) fn push(&mut self, bytes: &[u8]) -> Result, Error> { + self.buffer.extend_from_slice(bytes); + let mut events = Vec::new(); + loop { + // `decode_frame` advances `self.buffer` and retains partial-frame + // state internally, so repeated calls over a growing buffer work. + let frame = self.inner.decode_frame(&mut self.buffer).map_err(|e| { + Error::stream_error( + format!("bedrock event-stream decode: {e}"), + std::io::Error::other(e.to_string()), + ) + })?; + match frame { + DecodedFrame::Complete(message) => { + if let Some(event) = Self::message_to_event(&message)? { + events.push(event); + } + } + DecodedFrame::Incomplete => break, + } + } + Ok(events) + } + + /// Classify one event-stream message. + /// + /// `event` frames yield their `:event-type` name and JSON payload; + /// `exception` frames (modeled AWS errors such as `throttlingException`, + /// arriving in-band after HTTP 200) and `error` frames (unmodeled) are + /// turned into errors. Frames without an event type are skipped. + fn message_to_event(message: &Message) -> Result, Error> { + match header_str(message, ":message-type") { + Some("exception") => { + let kind = header_str(message, ":exception-type").unwrap_or("unknown"); + let body = String::from_utf8_lossy(message.payload()); + Err(Error::stream_error( + format!("bedrock stream exception ({kind}): {body}"), + std::io::Error::other("bedrock event-stream exception frame"), + )) + } + Some("error") => { + let code = header_str(message, ":error-code").unwrap_or("unknown"); + let detail = header_str(message, ":error-message").unwrap_or(""); + Err(Error::stream_error( + format!("bedrock stream error ({code}): {detail}"), + std::io::Error::other("bedrock event-stream error frame"), + )) + } + _ => { + let Some(event_type) = header_str(message, ":event-type") else { + return Ok(None); + }; + Ok(Some(DecodedEvent { + event_type: event_type.to_string(), + payload: String::from_utf8_lossy(message.payload()).into_owned(), + })) + } + } + } +} + +/// Read a string-valued event-stream header by name. +fn header_str<'a>(message: &'a Message, name: &str) -> Option<&'a str> { + message + .headers() + .iter() + .find(|header| header.name().as_str() == name) + .and_then(|header| header.value().as_string().ok()) + .map(StrBytes::as_str) +} + +#[cfg(test)] +pub(crate) mod tests { + use aws_smithy_eventstream::frame::write_message_to; + use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; + + use super::*; + + /// Build one ConverseStream event frame: event name in `:event-type`, + /// payload = the event JSON directly. + fn encode_event_frame(event_type: &str, payload_json: &str) -> Vec { + let message = Message::new(payload_json.as_bytes().to_vec()) + .add_header(Header::new( + ":message-type", + HeaderValue::String("event".into()), + )) + .add_header(Header::new( + ":event-type", + HeaderValue::String(event_type.into()), + )) + .add_header(Header::new( + ":content-type", + HeaderValue::String("application/json".into()), + )); + let mut buf = Vec::new(); + write_message_to(&message, &mut buf).unwrap(); + buf + } + + /// Build a full streaming body from `(event_type, payload_json)` pairs. + pub(crate) fn build_stream_body(events: &[(&str, &str)]) -> Vec { + let mut body = Vec::new(); + for (event_type, payload) in events { + body.extend_from_slice(&encode_event_frame(event_type, payload)); + } + body + } + + #[test] + fn decodes_event_frame_to_typed_payload() { + let frame = encode_event_frame( + "contentBlockDelta", + r#"{"delta":{"text":"hi"},"contentBlockIndex":0}"#, + ); + let mut decoder = FrameDecoder::new(); + let events = decoder.push(&frame).unwrap(); + assert_eq!(events.len(), 1); + assert_eq!(events[0].event_type, "contentBlockDelta"); + let payload: serde_json::Value = serde_json::from_str(&events[0].payload).unwrap(); + assert_eq!(payload["delta"]["text"], "hi"); + } + + #[test] + fn reassembles_frame_split_across_pushes() { + let frame = encode_event_frame("messageStop", r#"{"stopReason":"end_turn"}"#); + let split = frame.len() / 2; + let mut decoder = FrameDecoder::new(); + assert!(decoder.push(&frame[..split]).unwrap().is_empty()); + let events = decoder.push(&frame[split..]).unwrap(); + assert_eq!(events.len(), 1); + assert_eq!(events[0].event_type, "messageStop"); + } + + #[test] + fn exception_frame_surfaces_as_error() { + let message = Message::new(br#"{"message":"Too many requests"}"#.to_vec()) + .add_header(Header::new( + ":message-type", + HeaderValue::String("exception".into()), + )) + .add_header(Header::new( + ":exception-type", + HeaderValue::String("throttlingException".into()), + )); + let mut buf = Vec::new(); + write_message_to(&message, &mut buf).unwrap(); + + let mut decoder = FrameDecoder::new(); + let err = decoder.push(&buf).unwrap_err(); + let rendered = err.to_string(); + assert!(rendered.contains("throttlingException"), "{rendered}"); + assert!(rendered.contains("Too many requests"), "{rendered}"); + } + + #[test] + fn unmodeled_error_frame_surfaces_as_error() { + let message = Message::new(Vec::new()) + .add_header(Header::new( + ":message-type", + HeaderValue::String("error".into()), + )) + .add_header(Header::new( + ":error-code", + HeaderValue::String("InternalError".into()), + )) + .add_header(Header::new( + ":error-message", + HeaderValue::String("stream broke".into()), + )); + let mut buf = Vec::new(); + write_message_to(&message, &mut buf).unwrap(); + + let mut decoder = FrameDecoder::new(); + let err = decoder.push(&buf).unwrap_err(); + let rendered = err.to_string(); + assert!(rendered.contains("InternalError"), "{rendered}"); + } +} diff --git a/lib/crates/fabro-llm/src/providers/bedrock/mod.rs b/lib/crates/fabro-llm/src/providers/bedrock/mod.rs new file mode 100644 index 000000000..833e49e19 --- /dev/null +++ b/lib/crates/fabro-llm/src/providers/bedrock/mod.rs @@ -0,0 +1,111 @@ +//! Amazon Bedrock transport primitives: SigV4 signing, AWS event-stream +//! decoding, auth-mode selection, and region derivation. +//! +//! The adapter that composes these over the `bedrock_converse` codec lands +//! later in this series; until then the pieces carry ahead-of-use allows. + +pub(crate) mod eventstream; +pub(crate) mod sigv4; + +use tokio::sync::OnceCell; + +use crate::error::Error; + +/// 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), + /// SigV4 signing. The signer (holding the AWS default credential chain) + /// 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), +} + +/// Derive the AWS region from a Bedrock runtime endpoint URL. +/// +/// The region is a SigV4 signing parameter, so it is parsed from the +/// 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!( + "bedrock base_url '{base_url}' is not a recognized Bedrock runtime endpoint \ + (expected https://bedrock-runtime[-fips]..amazonaws.com[.cn])" + ), + 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); + let rest = host + .strip_prefix("bedrock-runtime-fips.") + .or_else(|| host.strip_prefix("bedrock-runtime.")) + .ok_or_else(invalid)?; + let region = rest + .strip_suffix(".amazonaws.com.cn") + .or_else(|| rest.strip_suffix(".amazonaws.com")) + .ok_or_else(invalid)?; + let valid = !region.is_empty() + && region + .bytes() + .all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'-'); + if valid { + Ok(region.to_string()) + } else { + Err(invalid()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn region_parses_from_standard_endpoint() { + assert_eq!( + region_from_base_url("https://bedrock-runtime.eu-west-1.amazonaws.com").unwrap(), + "eu-west-1" + ); + } + + #[test] + fn region_parses_from_fips_endpoint() { + assert_eq!( + region_from_base_url("https://bedrock-runtime-fips.us-gov-west-1.amazonaws.com") + .unwrap(), + "us-gov-west-1" + ); + } + + #[test] + fn region_parses_from_china_endpoint() { + assert_eq!( + region_from_base_url("https://bedrock-runtime.cn-north-1.amazonaws.com.cn").unwrap(), + "cn-north-1" + ); + } + + #[test] + fn region_rejects_non_bedrock_hosts() { + for url in [ + "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}"); + } + } +} diff --git a/lib/crates/fabro-llm/src/providers/bedrock/sigv4.rs b/lib/crates/fabro-llm/src/providers/bedrock/sigv4.rs new file mode 100644 index 000000000..1f327bfc0 --- /dev/null +++ b/lib/crates/fabro-llm/src/providers/bedrock/sigv4.rs @@ -0,0 +1,247 @@ +//! AWS Signature Version 4 signing for Bedrock requests. +//! +//! Wraps the `aws-sigv4` crate to compute the `Authorization`, `x-amz-date`, +//! and (for temporary credentials) `x-amz-security-token` headers for a fully +//! built request. The headers are then attached to the shared `fabro-http` +//! request builder, so signed Bedrock requests still flow through the same +//! retry/redaction/transport layers as every other adapter. + +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use aws_credential_types::Credentials; +use aws_credential_types::provider::SharedCredentialsProvider; +use aws_sigv4::http_request::{SignableBody, SignableRequest, SigningSettings, sign}; +use aws_sigv4::sign::v4; +use aws_smithy_runtime_api::client::identity::Identity; + +use crate::error::Error; + +/// Service name used in the SigV4 credential scope for Bedrock runtime calls. +pub(crate) const SERVICE: &str = "bedrock"; + +/// Where the signer's credentials come from. +enum CredentialSource { + /// Fixed credentials (tests / explicitly supplied keys). + Static(Credentials), + /// The AWS default provider chain. Credentials are resolved per request + /// so expiring session credentials (STS, IRSA, instance roles) refresh + /// through the chain's identity cache instead of being snapshotted once + /// at startup. + Chain(SharedCredentialsProvider), +} + +/// Signs HTTP requests for AWS services with SigV4. +pub(crate) struct Sigv4Signer { + credentials: CredentialSource, +} + +impl Sigv4Signer { + /// Build a signer from static keys. Test-only: production paths resolve + /// credentials through the AWS chain. + #[cfg(test)] + pub(crate) fn from_static( + access_key_id: &str, + secret_access_key: &str, + session_token: Option, + ) -> Self { + Self { + credentials: CredentialSource::Static(Credentials::from_keys( + access_key_id, + secret_access_key, + session_token, + )), + } + } + + /// Build a signer over the standard AWS provider chain (environment, + /// IRSA/web identity, EC2/ECS instance profile, SSO, assume-role). The + /// chain is resolved once; the credentials it yields are fetched per + /// signing call so they stay fresh over long-lived adapters. + pub(crate) async fn from_default_chain() -> Result { + let config = aws_config::defaults(aws_config::BehaviorVersion::latest()) + .load() + .await; + let provider = config + .credentials_provider() + .ok_or_else(|| Error::Configuration { + message: "no AWS credentials provider found in the default chain".to_string(), + source: None, + })?; + Ok(Self { + credentials: CredentialSource::Chain(provider), + }) + } + + /// The credentials to sign the next request with. + async fn current_credentials(&self) -> Result { + use aws_credential_types::provider::ProvideCredentials; + + match &self.credentials { + CredentialSource::Static(credentials) => Ok(credentials.clone()), + CredentialSource::Chain(provider) => { + provider + .provide_credentials() + .await + .map_err(|e| Error::Configuration { + message: format!("failed to resolve AWS credentials: {e}"), + source: None, + }) + } + } + } + + /// Compute the SigV4 headers for a request: `Authorization`, `x-amz-date`, + /// and `x-amz-security-token` when the credentials carry a session token. + fn signed_headers( + credentials: &Credentials, + region: &str, + service: &str, + method: &str, + url: &str, + body: &[u8], + epoch_secs: u64, + ) -> Result, Error> { + let identity: Identity = credentials.clone().into(); + let signing_params = v4::SigningParams::builder() + .identity(&identity) + .region(region) + .name(service) + .time(UNIX_EPOCH + Duration::from_secs(epoch_secs)) + .settings(SigningSettings::default()) + .build() + .map_err(|e| Error::Configuration { + message: format!("sigv4 params: {e}"), + source: None, + })? + .into(); + + let signable = + SignableRequest::new(method, url, std::iter::empty(), SignableBody::Bytes(body)) + .map_err(|e| Error::Configuration { + message: format!("sigv4 signable request: {e}"), + source: None, + })?; + + let (instructions, _signature) = sign(signable, &signing_params) + .map_err(|e| Error::Configuration { + message: format!("sigv4 signing failed: {e}"), + source: None, + })? + .into_parts(); + + Ok(instructions + .headers() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect()) + } + + /// Apply SigV4 signed headers to a `fabro-http` request builder for a + /// `POST` to `url` carrying `body`. + pub(crate) async fn sign_post( + &self, + mut req: fabro_http::RequestBuilder, + region: &str, + url: &str, + body: &[u8], + ) -> Result { + let credentials = self.current_credentials().await?; + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map_err(|e| Error::Configuration { + message: format!("system clock before epoch: {e}"), + source: None, + })? + .as_secs(); + for (name, value) in + Self::signed_headers(&credentials, region, SERVICE, "POST", url, body, now)? + { + req = req.header(name, value); + } + Ok(req.body(body.to_vec())) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + // Fixed credentials + time produce a deterministic Authorization header. + // The expected value is locked below after the first green run so the test + // guards against accidental changes to the signing logic. + const ACCESS_KEY: &str = "AKIDEXAMPLE"; + const SECRET_KEY: &str = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"; + const FIXED_EPOCH: u64 = 1_716_960_000; // 2024-05-29T04:00:00Z + const URL: &str = "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-sonnet-4-6/converse"; + + fn static_credentials(signer: &Sigv4Signer) -> Credentials { + match &signer.credentials { + CredentialSource::Static(credentials) => credentials.clone(), + CredentialSource::Chain(_) => panic!("test signer should hold static credentials"), + } + } + + fn auth_header(headers: &[(String, String)]) -> &str { + headers + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("authorization")) + .map(|(_, value)| value.as_str()) + .expect("authorization header must be present") + } + + fn sign_fixed(signer: &Sigv4Signer, body: &[u8]) -> Vec<(String, String)> { + Sigv4Signer::signed_headers( + &static_credentials(signer), + "us-east-1", + SERVICE, + "POST", + URL, + body, + FIXED_EPOCH, + ) + .unwrap() + } + + #[test] + fn produces_authorization_and_date_headers() { + let signer = Sigv4Signer::from_static(ACCESS_KEY, SECRET_KEY, None); + let headers = sign_fixed(&signer, br#"{"messages":[]}"#); + + assert!( + headers + .iter() + .any(|(n, _)| n.eq_ignore_ascii_case("authorization")) + ); + assert!( + headers + .iter() + .any(|(n, _)| n.eq_ignore_ascii_case("x-amz-date")) + ); + let auth = auth_header(&headers); + assert!(auth.starts_with("AWS4-HMAC-SHA256 ")); + assert!(auth.contains("Credential=AKIDEXAMPLE/20240529/us-east-1/bedrock/aws4_request")); + assert!(auth.contains("SignedHeaders=")); + assert!(auth.contains("Signature=")); + } + + #[test] + fn deterministic_signature_is_stable() { + let signer = Sigv4Signer::from_static(ACCESS_KEY, SECRET_KEY, None); + // Same inputs must yield an identical signature (regression lock). + assert_eq!( + auth_header(&sign_fixed(&signer, br#"{"messages":[]}"#)), + auth_header(&sign_fixed(&signer, br#"{"messages":[]}"#)), + ); + } + + #[test] + fn session_token_adds_security_token_header() { + let signer = + Sigv4Signer::from_static(ACCESS_KEY, SECRET_KEY, Some("session-tok".to_string())); + let headers = sign_fixed(&signer, b"{}"); + assert!( + headers + .iter() + .any(|(n, v)| n.eq_ignore_ascii_case("x-amz-security-token") && v == "session-tok") + ); + } +} diff --git a/lib/crates/fabro-llm/src/providers/mod.rs b/lib/crates/fabro-llm/src/providers/mod.rs index ecc2ee96c..d0c72788d 100644 --- a/lib/crates/fabro-llm/src/providers/mod.rs +++ b/lib/crates/fabro-llm/src/providers/mod.rs @@ -1,4 +1,5 @@ pub mod anthropic; +pub(crate) mod bedrock; pub mod common; pub mod fabro_server; pub mod gemini; @@ -6,6 +7,7 @@ pub mod openai; pub mod openai_compatible; pub use anthropic::Adapter as AnthropicAdapter; +pub use bedrock::Adapter as BedrockAdapter; pub use fabro_server::Adapter as FabroServerAdapter; pub use gemini::Adapter as GeminiAdapter; pub use openai::Adapter as OpenAiAdapter;