Skip to main content

lean_ctx/proxy/
bedrock.rs

1//! Amazon Bedrock Runtime request validation and AWS Signature Version 4.
2//!
3//! Credentials are read only from the standard AWS environment variables at
4//! request time. The final body bytes (after any bounded proxy transform) are
5//! the bytes covered by the payload digest and signature.
6
7use std::collections::BTreeMap;
8
9use axum::http::{
10    HeaderMap, HeaderName, HeaderValue, Method, Request, StatusCode, header, request::Parts,
11};
12use chrono::Utc;
13use futures::{Stream, StreamExt};
14use hmac::{Hmac, KeyInit, Mac};
15use serde_json::Value;
16use sha2::{Digest, Sha256};
17
18use crate::core::config::ResolvedProvider;
19
20const SERVICE: &str = "bedrock";
21const MAX_BODY_BYTES: usize = 25_000_000;
22const MAX_CREDENTIAL_BYTES: usize = 4 * 1024;
23const ACCESS_KEY_ENV: &str = "AWS_ACCESS_KEY_ID";
24const SECRET_KEY_ENV: &str = "AWS_SECRET_ACCESS_KEY";
25const SESSION_TOKEN_ENV: &str = "AWS_SESSION_TOKEN";
26const BEDROCK_REQUEST_HEADERS: &[&str] = &[
27    "x-amz-date",
28    "x-amz-content-sha256",
29    "x-amz-security-token",
30    "x-amzn-bedrock-accept",
31    "x-amzn-bedrock-trace",
32    "x-amzn-bedrock-guardrailidentifier",
33    "x-amzn-bedrock-guardrailversion",
34    "x-amzn-bedrock-guardrailtrace",
35    "x-amzn-bedrock-performanceconfig-latency",
36    "x-amzn-bedrock-service-tier",
37    "x-amzn-bedrock-request-metadata",
38];
39
40pub(super) fn is_bedrock_request_header(name: &str) -> bool {
41    BEDROCK_REQUEST_HEADERS.contains(&name)
42}
43
44pub(super) fn is_bedrock_response_header(name: &str) -> bool {
45    name.starts_with("x-amzn-bedrock-") || matches!(name, "x-amzn-requestid" | "x-amzn-errortype")
46}
47
48pub(super) fn response_is_sse(headers: &HeaderMap) -> bool {
49    headers
50        .get("content-type")
51        .and_then(|value| value.to_str().ok())
52        .is_some_and(|value| value.to_ascii_lowercase().contains("text/event-stream"))
53}
54
55#[allow(clippy::needless_pass_by_value)]
56pub(super) fn passthrough_request_body(
57    parsed: Value,
58    original_size: usize,
59) -> (Vec<u8>, usize, usize) {
60    let body = serde_json::to_vec(&parsed).unwrap_or_default();
61    (body, original_size, original_size)
62}
63
64pub(super) fn final_request_body(
65    provider_label: &str,
66    parts: &Parts,
67    raw: &[u8],
68    transformed: Vec<u8>,
69) -> Vec<u8> {
70    if provider_label == "Bedrock" && !parts.headers.contains_key("content-encoding") {
71        raw.to_vec()
72    } else {
73        transformed
74    }
75}
76
77pub(super) fn finalize_request(
78    provider_label: &str,
79    parts: &mut Parts,
80    raw: &[u8],
81    transformed: Vec<u8>,
82    limit: usize,
83    url: &str,
84) -> Result<Vec<u8>, StatusCode> {
85    let body = final_request_body(provider_label, parts, raw, transformed);
86    if body.len() > limit {
87        return Err(StatusCode::PAYLOAD_TOO_LARGE);
88    }
89    sign_request_from_environment(parts, url, &body)?;
90    Ok(body)
91}
92
93/// AWS event-stream frame decoder used only for usage observation. Bytes still
94/// flow through unchanged; JSON payloads are fed to the Anthropic-shaped usage
95/// scanner, while malformed/unknown frames fail closed for metering only.
96struct EventStreamScanner {
97    buffered: Vec<u8>,
98    scanner: crate::proxy::usage::Scanner,
99}
100
101impl EventStreamScanner {
102    fn new(scanner: crate::proxy::usage::Scanner) -> Self {
103        Self {
104            buffered: Vec::new(),
105            scanner,
106        }
107    }
108
109    fn feed(&mut self, chunk: &[u8]) {
110        self.buffered.extend_from_slice(chunk);
111        loop {
112            if self.buffered.len() < 12 {
113                return;
114            }
115            let total = u32::from_be_bytes(
116                self.buffered[0..4]
117                    .try_into()
118                    .expect("buffered length >= 12 guarantees 4-byte prefix"),
119            ) as usize;
120            let headers = u32::from_be_bytes(
121                self.buffered[4..8]
122                    .try_into()
123                    .expect("buffered length >= 12 guarantees 4-byte header count"),
124            ) as usize;
125            if !(16..=8 * 1024 * 1024).contains(&total) || headers > total - 16 {
126                self.buffered.clear();
127                return;
128            }
129            if self.buffered.len() < total {
130                return;
131            }
132            let frame: Vec<u8> = self.buffered.drain(..total).collect();
133            if crc32(&frame[..8])
134                != u32::from_be_bytes(
135                    frame[8..12]
136                        .try_into()
137                        .expect("frame total >= 16 guarantees 4-byte header CRC"),
138                )
139                || crc32(&frame[..total - 4])
140                    != u32::from_be_bytes(
141                        frame[total - 4..]
142                            .try_into()
143                            .expect("frame total >= 16 guarantees 4-byte trailer CRC"),
144                    )
145            {
146                continue;
147            }
148            let payload_end = total - 4;
149            let payload = &frame[12 + headers..payload_end];
150            let Ok(value) = serde_json::from_slice::<Value>(payload) else {
151                continue;
152            };
153            if let Some(metrics) = value.get("amazon-bedrock-invocationMetrics") {
154                let mapped = serde_json::json!({
155                    "usage": {
156                        "input_tokens": metrics.get("inputTokenCount").and_then(Value::as_u64).unwrap_or(0),
157                        "output_tokens": metrics.get("outputTokenCount").and_then(Value::as_u64).unwrap_or(0),
158                    }
159                });
160                self.scanner
161                    .feed_body(&serde_json::to_vec(&mapped).unwrap_or_default());
162            } else {
163                self.scanner.feed_body(payload);
164            }
165        }
166    }
167
168    fn finalize(self) -> Option<crate::proxy::usage::RealUsage> {
169        self.scanner.finalize()
170    }
171}
172
173fn crc32(bytes: &[u8]) -> u32 {
174    let mut crc = u32::MAX;
175    for byte in bytes {
176        crc ^= u32::from(*byte);
177        for _ in 0..8 {
178            crc = if crc & 1 == 1 {
179                (crc >> 1) ^ 0xedb88320
180            } else {
181                crc >> 1
182            };
183        }
184    }
185    !crc
186}
187
188pub(super) fn tee_eventstream<S, B, E>(
189    inner: S,
190    scanner: crate::proxy::usage::Scanner,
191) -> impl Stream<Item = Result<B, E>> + Send + 'static
192where
193    S: Stream<Item = Result<B, E>> + Send + Unpin + 'static,
194    B: AsRef<[u8]> + Send + 'static,
195    E: Send + 'static,
196{
197    futures::stream::unfold(
198        (inner, Some(EventStreamScanner::new(scanner))),
199        |(mut inner, mut scanner)| async move {
200            match inner.next().await {
201                Some(Ok(chunk)) => {
202                    if let Some(s) = scanner.as_mut() {
203                        s.feed(chunk.as_ref());
204                    }
205                    Some((Ok(chunk), (inner, scanner)))
206                }
207                Some(err) => Some((err, (inner, scanner))),
208                None => {
209                    if let Some(s) = scanner.take()
210                        && let Some(usage) = s.finalize()
211                    {
212                        crate::proxy::usage_meter::record(&usage);
213                    }
214                    None
215                }
216            }
217        },
218    )
219}
220
221pub(super) fn build_stream_body<S>(
222    inner: S,
223    scanner: crate::proxy::usage::Scanner,
224    is_sse: bool,
225    eventstream: bool,
226    xlat: bool,
227) -> axum::body::Body
228where
229    S: Stream<Item = Result<axum::body::Bytes, reqwest::Error>> + Send + Unpin + 'static,
230{
231    if is_sse {
232        return super::forward::xlat_stream_body(
233            super::sse_keepalive::keepalive_stream(Box::pin(crate::proxy::usage::tee_stream(
234                inner, scanner,
235            ))),
236            xlat,
237        );
238    }
239    if eventstream {
240        return super::forward::xlat_stream_body(Box::pin(tee_eventstream(inner, scanner)), xlat);
241    }
242    super::forward::xlat_stream_body(
243        Box::pin(crate::proxy::usage::tee_stream(inner, scanner)),
244        xlat,
245    )
246}
247
248#[derive(Clone)]
249pub(super) struct SigningContext {
250    pub(super) region: String,
251}
252
253impl std::fmt::Debug for SigningContext {
254    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
255        formatter
256            .debug_struct("SigningContext")
257            .field("region", &self.region)
258            .finish_non_exhaustive()
259    }
260}
261
262#[derive(Debug, PartialEq, Eq)]
263enum SigningError {
264    MissingCredential,
265    InvalidCredential,
266    InvalidUrl,
267    InvalidHeader,
268    InvalidTimestamp,
269}
270
271struct Credentials {
272    access_key: String,
273    secret_key: String,
274    session_token: Option<String>,
275}
276
277impl Credentials {
278    fn from_environment() -> Result<Self, SigningError> {
279        Ok(Self {
280            access_key: required_credential(ACCESS_KEY_ENV)?,
281            secret_key: required_credential(SECRET_KEY_ENV)?,
282            session_token: optional_credential(SESSION_TOKEN_ENV)?,
283        })
284    }
285}
286
287fn required_credential(name: &str) -> Result<String, SigningError> {
288    optional_credential(name)?.ok_or(SigningError::MissingCredential)
289}
290
291fn optional_credential(name: &str) -> Result<Option<String>, SigningError> {
292    let value = match std::env::var(name) {
293        Ok(value) => value,
294        Err(std::env::VarError::NotPresent) => return Ok(None),
295        Err(std::env::VarError::NotUnicode(_)) => return Err(SigningError::InvalidCredential),
296    };
297    let value = value.trim();
298    if value.is_empty() {
299        return Ok(None);
300    }
301    if value.len() > MAX_CREDENTIAL_BYTES || value.bytes().any(|byte| byte.is_ascii_control()) {
302        return Err(SigningError::InvalidCredential);
303    }
304    Ok(Some(value.to_string()))
305}
306
307pub(super) fn attach_signing_context(
308    provider: &ResolvedProvider,
309    request: &mut Request<axum::body::Body>,
310) -> Result<(), StatusCode> {
311    let Some(region) = provider
312        .aws_region
313        .as_deref()
314        .filter(|value| !value.is_empty())
315    else {
316        return Err(StatusCode::BAD_GATEWAY);
317    };
318    Credentials::from_environment().map_err(|_| StatusCode::BAD_GATEWAY)?;
319    strip_untrusted_signing_headers(request.headers_mut());
320    request.extensions_mut().insert(SigningContext {
321        region: region.to_string(),
322    });
323    Ok(())
324}
325
326pub(super) fn request_body_limit(parts: &Parts) -> Option<usize> {
327    parts
328        .extensions
329        .get::<SigningContext>()
330        .map(|_| MAX_BODY_BYTES)
331}
332
333pub(super) fn validate_invoke_request<B>(request: &Request<B>) -> Result<(), StatusCode> {
334    if request.method() != Method::POST || request.uri().query().is_some() {
335        return Err(StatusCode::METHOD_NOT_ALLOWED);
336    }
337    let Some(rest) = request.uri().path().strip_prefix("/model/") else {
338        return Err(StatusCode::NOT_FOUND);
339    };
340    let Some((model, operation)) = rest.rsplit_once('/') else {
341        return Err(StatusCode::NOT_FOUND);
342    };
343    if model.is_empty()
344        || model.len() > 2_048
345        || model
346            .split('/')
347            .any(|segment| segment.is_empty() || segment == "..")
348        || !matches!(operation, "invoke" | "invoke-with-response-stream")
349    {
350        return Err(StatusCode::NOT_FOUND);
351    }
352    Ok(())
353}
354
355pub(super) fn sign_request_from_environment(
356    parts: &mut Parts,
357    url: &str,
358    body: &[u8],
359) -> Result<(), StatusCode> {
360    let Some(context) = parts.extensions.get::<SigningContext>().cloned() else {
361        return Ok(());
362    };
363    let credentials = Credentials::from_environment().map_err(|_| StatusCode::BAD_GATEWAY)?;
364    let timestamp = Utc::now().format("%Y%m%dT%H%M%SZ").to_string();
365    if !parts.headers.contains_key(header::CONTENT_TYPE) {
366        parts.headers.insert(
367            header::CONTENT_TYPE,
368            HeaderValue::from_static("application/json"),
369        );
370    }
371    sign_headers_at(
372        &parts.method,
373        url,
374        &mut parts.headers,
375        body,
376        &credentials,
377        &context.region,
378        SERVICE,
379        &timestamp,
380    )
381    .map_err(|_| StatusCode::BAD_GATEWAY)
382}
383
384fn strip_untrusted_signing_headers(headers: &mut HeaderMap) {
385    let names = headers
386        .keys()
387        .filter(|name| {
388            name.as_str().eq_ignore_ascii_case("authorization")
389                || name.as_str().starts_with("x-amz-")
390        })
391        .cloned()
392        .collect::<Vec<_>>();
393    for name in names {
394        headers.remove(name);
395    }
396}
397
398#[allow(clippy::too_many_arguments)]
399fn sign_headers_at(
400    method: &Method,
401    url: &str,
402    headers: &mut HeaderMap,
403    body: &[u8],
404    credentials: &Credentials,
405    region: &str,
406    service: &str,
407    timestamp: &str,
408) -> Result<(), SigningError> {
409    if timestamp.len() != 16
410        || timestamp.as_bytes().get(8) != Some(&b'T')
411        || timestamp.as_bytes().last() != Some(&b'Z')
412        || !timestamp
413            .bytes()
414            .enumerate()
415            .all(|(index, byte)| matches!(index, 8 | 15) || byte.is_ascii_digit())
416    {
417        return Err(SigningError::InvalidTimestamp);
418    }
419    let parsed = reqwest::Url::parse(url).map_err(|_| SigningError::InvalidUrl)?;
420    let host = canonical_host(&parsed)?;
421    strip_untrusted_signing_headers(headers);
422    let payload_hash = sha256_hex(body);
423    insert_header(headers, "x-amz-content-sha256", &payload_hash)?;
424    insert_header(headers, "x-amz-date", timestamp)?;
425    if let Some(token) = credentials.session_token.as_deref() {
426        insert_header(headers, "x-amz-security-token", token)?;
427    }
428
429    let (canonical_headers, signed_headers) = canonical_headers(headers, &host)?;
430    let canonical_request = format!(
431        "{}\n{}\n{}\n{}\n{}\n{}",
432        method.as_str(),
433        canonical_uri(parsed.path()),
434        canonical_query(parsed.query().unwrap_or_default()),
435        canonical_headers,
436        signed_headers,
437        payload_hash,
438    );
439    let date = &timestamp[..8];
440    let scope = format!("{date}/{region}/{service}/aws4_request");
441    let string_to_sign = format!(
442        "AWS4-HMAC-SHA256\n{timestamp}\n{scope}\n{}",
443        sha256_hex(canonical_request.as_bytes())
444    );
445    let date_key = hmac_sha256(
446        format!("AWS4{}", credentials.secret_key).as_bytes(),
447        date.as_bytes(),
448    );
449    let region_key = hmac_sha256(&date_key, region.as_bytes());
450    let service_key = hmac_sha256(&region_key, service.as_bytes());
451    let signing_key = hmac_sha256(&service_key, b"aws4_request");
452    let signature = hex(&hmac_sha256(&signing_key, string_to_sign.as_bytes()));
453    let authorization = format!(
454        "AWS4-HMAC-SHA256 Credential={}/{scope}, SignedHeaders={signed_headers}, Signature={signature}",
455        credentials.access_key
456    );
457    insert_header(headers, header::AUTHORIZATION.as_str(), &authorization)?;
458    Ok(())
459}
460
461fn insert_header(headers: &mut HeaderMap, name: &str, value: &str) -> Result<(), SigningError> {
462    let name = HeaderName::from_bytes(name.as_bytes()).map_err(|_| SigningError::InvalidHeader)?;
463    let value = HeaderValue::from_str(value).map_err(|_| SigningError::InvalidHeader)?;
464    headers.insert(name, value);
465    Ok(())
466}
467
468fn canonical_host(url: &reqwest::Url) -> Result<String, SigningError> {
469    let host = url.host_str().ok_or(SigningError::InvalidUrl)?;
470    let default_port = match url.scheme() {
471        "http" => Some(80),
472        "https" => Some(443),
473        _ => None,
474    };
475    let port = url.port().filter(|value| Some(*value) != default_port);
476    Ok(port.map_or_else(|| host.to_string(), |port| format!("{host}:{port}")))
477}
478
479fn canonical_headers(headers: &HeaderMap, host: &str) -> Result<(String, String), SigningError> {
480    let mut values = BTreeMap::<String, Vec<String>>::new();
481    values.insert("host".into(), vec![host.to_string()]);
482    for (name, value) in headers {
483        let name = name.as_str().to_ascii_lowercase();
484        if name != "content-type"
485            && name != "x-amz-content-sha256"
486            && name != "x-amz-date"
487            && name != "x-amz-security-token"
488            && !name.starts_with("x-amzn-")
489        {
490            continue;
491        }
492        let value = value.to_str().map_err(|_| SigningError::InvalidHeader)?;
493        values.entry(name).or_default().push(collapse_spaces(value));
494    }
495    let signed_headers = values.keys().cloned().collect::<Vec<_>>().join(";");
496    let canonical = values
497        .into_iter()
498        .fold(String::new(), |mut output, (name, values)| {
499            use std::fmt::Write as _;
500            let _ = writeln!(output, "{name}:{}", values.join(","));
501            output
502        });
503    Ok((canonical, signed_headers))
504}
505
506fn collapse_spaces(value: &str) -> String {
507    value.split_ascii_whitespace().collect::<Vec<_>>().join(" ")
508}
509
510fn canonical_uri(path: &str) -> String {
511    if path.is_empty() {
512        return "/".into();
513    }
514    aws_encode(path.as_bytes(), true)
515}
516
517fn canonical_query(query: &str) -> String {
518    let mut pairs = query
519        .split('&')
520        .filter(|pair| !pair.is_empty())
521        .map(|pair| {
522            let (name, value) = pair.split_once('=').unwrap_or((pair, ""));
523            (
524                aws_encode(name.as_bytes(), false),
525                aws_encode(value.as_bytes(), false),
526            )
527        })
528        .collect::<Vec<_>>();
529    pairs.sort();
530    pairs
531        .into_iter()
532        .map(|(name, value)| format!("{name}={value}"))
533        .collect::<Vec<_>>()
534        .join("&")
535}
536
537fn aws_encode(input: &[u8], preserve_slash: bool) -> String {
538    const HEX: &[u8; 16] = b"0123456789ABCDEF";
539    let mut encoded = String::with_capacity(input.len());
540    let mut index = 0;
541    while index < input.len() {
542        let byte = input[index];
543        if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') {
544            encoded.push(char::from(byte));
545        } else if preserve_slash && byte == b'/' {
546            encoded.push('/');
547        } else if byte == b'%'
548            && index + 2 < input.len()
549            && input[index + 1].is_ascii_hexdigit()
550            && input[index + 2].is_ascii_hexdigit()
551        {
552            encoded.push('%');
553            encoded.push(char::from(input[index + 1]).to_ascii_uppercase());
554            encoded.push(char::from(input[index + 2]).to_ascii_uppercase());
555            index += 2;
556        } else {
557            encoded.push('%');
558            encoded.push(char::from(HEX[(byte >> 4) as usize]));
559            encoded.push(char::from(HEX[(byte & 0x0f) as usize]));
560        }
561        index += 1;
562    }
563    encoded
564}
565
566fn sha256_hex(value: &[u8]) -> String {
567    hex(&Sha256::digest(value))
568}
569
570fn hmac_sha256(key: &[u8], value: &[u8]) -> Vec<u8> {
571    let mut mac = Hmac::<Sha256>::new_from_slice(key).expect("HMAC accepts any key length");
572    mac.update(value);
573    mac.finalize().into_bytes().to_vec()
574}
575
576fn hex(value: &[u8]) -> String {
577    const HEX: &[u8; 16] = b"0123456789abcdef";
578    let mut encoded = String::with_capacity(value.len() * 2);
579    for byte in value {
580        encoded.push(char::from(HEX[(byte >> 4) as usize]));
581        encoded.push(char::from(HEX[(byte & 0x0f) as usize]));
582    }
583    encoded
584}
585
586#[cfg(test)]
587mod tests {
588    use super::*;
589
590    #[test]
591    fn bedrock_sigv4_is_deterministic_for_fixed_request() {
592        let credentials = Credentials {
593            access_key: "AKIDEXAMPLE".into(),
594            secret_key: "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".into(),
595            session_token: None,
596        };
597        let mut headers = HeaderMap::new();
598        headers.insert("content-type", HeaderValue::from_static("application/json"));
599        let mut second = headers.clone();
600        let body = br#"{"prompt":"hello"}"#;
601        sign_headers_at(
602            &Method::POST,
603            "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/invoke",
604            &mut headers,
605            body,
606            &credentials,
607            "us-east-1",
608            SERVICE,
609            "20240301T000000Z",
610        )
611        .unwrap();
612        sign_headers_at(
613            &Method::POST,
614            "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/invoke",
615            &mut second,
616            body,
617            &credentials,
618            "us-east-1",
619            SERVICE,
620            "20240301T000000Z",
621        )
622        .unwrap();
623        assert_eq!(
624            headers[header::AUTHORIZATION],
625            second[header::AUTHORIZATION]
626        );
627        assert!(
628            headers[header::AUTHORIZATION]
629                .to_str()
630                .unwrap()
631                .contains("/20240301/us-east-1/bedrock/aws4_request")
632        );
633        assert_eq!(headers["x-amz-content-sha256"], sha256_hex(body));
634    }
635
636    #[test]
637    fn bedrock_sigv4_matches_fixed_reference_vector() {
638        // Provenance: AWS General Reference, "Signing AWS API requests with
639        // Signature Version 4" —
640        // https://docs.aws.amazon.com/general/latest/gr/signature-version-4.html.
641        // The published canonical/HMAC procedure is adapted here to the
642        // Bedrock Runtime host and service scope; expected Authorization stays
643        // a fixed, independently recomputed vector (no live AWS call).
644        let credentials = Credentials {
645            access_key: "AKIDEXAMPLE".into(),
646            secret_key: "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".into(),
647            session_token: None,
648        };
649        let mut headers = HeaderMap::new();
650        headers.insert("content-type", HeaderValue::from_static("application/json"));
651        let body = br#"{"prompt":"hello"}"#;
652        sign_headers_at(
653            &Method::POST,
654            "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/invoke",
655            &mut headers,
656            body,
657            &credentials,
658            "us-east-1",
659            SERVICE,
660            "20240301T000000Z",
661        )
662        .unwrap();
663        assert_eq!(
664            headers[header::AUTHORIZATION],
665            "AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20240301/us-east-1/bedrock/aws4_request, SignedHeaders=content-type;host;x-amz-content-sha256;x-amz-date, Signature=9b82ca1f601090dcdb8213bb724217d852eb1efa80c8ef496cbc1d8c95c89ff8"
666        );
667    }
668
669    #[test]
670    fn bedrock_sigv4_signs_mock_session_credentials() {
671        let credentials = Credentials {
672            access_key: "AKIDEXAMPLE".into(),
673            secret_key: "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".into(),
674            session_token: Some("session-token".into()),
675        };
676        let mut headers = HeaderMap::new();
677        sign_headers_at(
678            &Method::POST,
679            "https://bedrock-runtime.us-east-1.amazonaws.com/model/demo/invoke",
680            &mut headers,
681            b"{}",
682            &credentials,
683            "us-east-1",
684            SERVICE,
685            "20240301T000000Z",
686        )
687        .unwrap();
688        assert_eq!(headers["x-amz-security-token"], "session-token");
689        assert!(
690            headers[header::AUTHORIZATION].to_str().unwrap().contains(
691                "SignedHeaders=host;x-amz-content-sha256;x-amz-date;x-amz-security-token"
692            )
693        );
694    }
695
696    #[test]
697    fn eventstream_sigv4_binds_exact_binary_payload() {
698        let credentials = Credentials {
699            access_key: "AKIDEXAMPLE".into(),
700            secret_key: "secret".into(),
701            session_token: None,
702        };
703        let body = [0u8, 1, 2, 0xff, 0x7f];
704        let mut headers = HeaderMap::new();
705        headers.insert(
706            header::CONTENT_TYPE,
707            HeaderValue::from_static("application/vnd.amazon.eventstream"),
708        );
709        sign_headers_at(
710            &Method::POST,
711            "https://bedrock-runtime.us-east-1.amazonaws.com/model/demo/invoke-with-response-stream",
712            &mut headers,
713            &body,
714            &credentials,
715            "us-east-1",
716            SERVICE,
717            "20240301T000000Z",
718        )
719        .unwrap();
720        assert_eq!(headers["x-amz-content-sha256"], sha256_hex(&body));
721        assert!(
722            headers[header::AUTHORIZATION]
723                .to_str()
724                .unwrap()
725                .contains("SignedHeaders=content-type;host;x-amz-content-sha256;x-amz-date")
726        );
727    }
728
729    #[test]
730    fn bedrock_request_metadata_is_signed_when_forwarded() {
731        let credentials = Credentials {
732            access_key: "AKIDEXAMPLE".into(),
733            secret_key: "secret".into(),
734            session_token: None,
735        };
736        let mut headers = HeaderMap::new();
737        headers.insert(
738            "x-amzn-bedrock-request-metadata",
739            HeaderValue::from_static("{\"team\":\"platform\"}"),
740        );
741        sign_headers_at(
742            &Method::POST,
743            "https://bedrock-runtime.us-east-1.amazonaws.com/model/demo/invoke",
744            &mut headers,
745            b"{}",
746            &credentials,
747            "us-east-1",
748            SERVICE,
749            "20240301T000000Z",
750        )
751        .unwrap();
752        assert!(
753            headers[header::AUTHORIZATION]
754                .to_str()
755                .unwrap()
756                .contains("x-amzn-bedrock-request-metadata")
757        );
758    }
759
760    #[test]
761    fn binary_eventstream_never_uses_sse_keepalive() {
762        let mut headers = HeaderMap::new();
763        headers.insert(
764            "content-type",
765            "application/vnd.amazon.eventstream".parse().unwrap(),
766        );
767        assert!(!response_is_sse(&headers));
768        headers.insert(
769            "content-type",
770            "text/event-stream; charset=utf-8".parse().unwrap(),
771        );
772        assert!(response_is_sse(&headers));
773        assert!(is_bedrock_response_header("x-amzn-requestid"));
774        assert!(!is_bedrock_response_header("x-amzn-secret"));
775    }
776
777    #[test]
778    fn identity_body_preserves_provider_bytes() {
779        let parts = Request::post("/model/demo/invoke")
780            .body(())
781            .unwrap()
782            .into_parts()
783            .0;
784        let raw = br#"{"b":1,"a":2}"#;
785        assert_eq!(
786            final_request_body("Bedrock", &parts, raw, br#"{"a":2,"b":1}"#.to_vec()),
787            raw
788        );
789    }
790
791    fn event_frame(payload: &[u8]) -> Vec<u8> {
792        let total = 16 + payload.len() as u32;
793        let mut frame = Vec::with_capacity(total as usize);
794        frame.extend_from_slice(&total.to_be_bytes());
795        frame.extend_from_slice(&0u32.to_be_bytes());
796        frame.extend_from_slice(&crc32(&frame).to_be_bytes());
797        frame.extend_from_slice(payload);
798        let crc = crc32(&frame);
799        frame.extend_from_slice(&crc.to_be_bytes());
800        frame
801    }
802
803    #[test]
804    fn framed_eventstream_decodes_metrics_across_chunks() {
805        let frame = event_frame(
806            br#"{"amazon-bedrock-invocationMetrics":{"inputTokenCount":17,"outputTokenCount":9}}"#,
807        );
808        let mut stream = EventStreamScanner::new(crate::proxy::usage::Scanner::new(
809            crate::proxy::usage::Provider::Anthropic,
810            None,
811        ));
812        stream.feed(&frame[..7]);
813        stream.feed(&frame[7..]);
814        let usage = stream.finalize().expect("Bedrock metrics extracted");
815        assert_eq!(usage.input_tokens, 17);
816        assert_eq!(usage.output_tokens, 9);
817    }
818
819    #[tokio::test]
820    async fn framed_eventstream_tee_replays_exact_bytes() {
821        let frame = event_frame(
822            br#"{"amazon-bedrock-invocationMetrics":{"inputTokenCount":3,"outputTokenCount":2}}"#,
823        );
824        let split = frame.len() / 2;
825        let chunks = vec![
826            Ok::<_, std::convert::Infallible>(frame[..split].to_vec()),
827            Ok(frame[split..].to_vec()),
828        ];
829        let stream = tee_eventstream(
830            futures::stream::iter(chunks),
831            crate::proxy::usage::Scanner::new(crate::proxy::usage::Provider::Anthropic, None),
832        );
833        let replayed = stream
834            .collect::<Vec<_>>()
835            .await
836            .into_iter()
837            .map(Result::unwrap)
838            .collect::<Vec<Vec<u8>>>()
839            .concat();
840        assert_eq!(replayed, frame);
841    }
842
843    #[test]
844    fn invoke_validation_rejects_queries_and_unknown_operations() {
845        for path in [
846            "/model/anthropic.claude-v2/invoke",
847            "/model/arn%3Aaws%3Abedrock%3Aus-east-1%3A123%3Aprofile%2Fdemo/invoke-with-response-stream",
848        ] {
849            validate_invoke_request(&Request::post(path).body(()).unwrap()).unwrap();
850        }
851        for path in [
852            "/model/anthropic.claude-v2/converse",
853            "/model/anthropic.claude-v2/invoke?unsigned=true",
854            "/model/../invoke",
855        ] {
856            assert!(validate_invoke_request(&Request::post(path).body(()).unwrap()).is_err());
857        }
858    }
859
860    #[test]
861    fn incoming_signing_headers_are_removed_before_new_signature() {
862        let mut headers = HeaderMap::new();
863        headers.insert("authorization", HeaderValue::from_static("caller"));
864        headers.insert("x-amz-date", HeaderValue::from_static("old"));
865        headers.insert("x-amz-target", HeaderValue::from_static("old-target"));
866        strip_untrusted_signing_headers(&mut headers);
867        assert!(headers.is_empty());
868    }
869
870    #[test]
871    fn body_digest_binds_exact_final_bytes() {
872        let credentials = Credentials {
873            access_key: "AKIDEXAMPLE".into(),
874            secret_key: "secret".into(),
875            session_token: None,
876        };
877        let mut first = HeaderMap::new();
878        let mut second = HeaderMap::new();
879        sign_headers_at(
880            &Method::POST,
881            "https://bedrock-runtime.us-east-1.amazonaws.com/model/demo/invoke",
882            &mut first,
883            b"final-body-a",
884            &credentials,
885            "us-east-1",
886            SERVICE,
887            "20240301T000000Z",
888        )
889        .unwrap();
890        sign_headers_at(
891            &Method::POST,
892            "https://bedrock-runtime.us-east-1.amazonaws.com/model/demo/invoke",
893            &mut second,
894            b"final-body-b",
895            &credentials,
896            "us-east-1",
897            SERVICE,
898            "20240301T000000Z",
899        )
900        .unwrap();
901        assert_ne!(
902            first["x-amz-content-sha256"],
903            second["x-amz-content-sha256"]
904        );
905        assert_ne!(first[header::AUTHORIZATION], second[header::AUTHORIZATION]);
906    }
907
908    #[test]
909    fn bedrock_body_transform_is_semantic_passthrough() {
910        let value = serde_json::json!({
911            "messages": [{"role": "user", "content": "keep me"}],
912            "temperature": 0.2,
913        });
914        let original = serde_json::to_vec(&value).unwrap();
915        let (body, original_size, compressed_size) =
916            passthrough_request_body(value.clone(), original.len());
917        assert_eq!(serde_json::from_slice::<Value>(&body).unwrap(), value);
918        assert_eq!(original_size, compressed_size);
919    }
920
921    #[test]
922    fn missing_credentials_fail_closed() {
923        let _lock = crate::core::data_dir::test_env_lock();
924        crate::test_env::remove_var(ACCESS_KEY_ENV);
925        crate::test_env::remove_var(SECRET_KEY_ENV);
926        assert!(matches!(
927            Credentials::from_environment(),
928            Err(SigningError::MissingCredential)
929        ));
930    }
931
932    #[test]
933    fn malformed_credentials_fail_closed() {
934        let _lock = crate::core::data_dir::test_env_lock();
935        crate::test_env::set_var(ACCESS_KEY_ENV, "AKID\nEXAMPLE");
936        crate::test_env::set_var(SECRET_KEY_ENV, "secret");
937        assert!(matches!(
938            Credentials::from_environment(),
939            Err(SigningError::InvalidCredential)
940        ));
941
942        crate::test_env::set_var(ACCESS_KEY_ENV, "A".repeat(MAX_CREDENTIAL_BYTES + 1));
943        assert!(matches!(
944            Credentials::from_environment(),
945            Err(SigningError::InvalidCredential)
946        ));
947        crate::test_env::remove_var(ACCESS_KEY_ENV);
948        crate::test_env::remove_var(SECRET_KEY_ENV);
949    }
950
951    #[test]
952    fn malformed_signing_inputs_return_typed_errors() {
953        let credentials = Credentials {
954            access_key: "AKIDEXAMPLE".into(),
955            secret_key: "secret".into(),
956            session_token: None,
957        };
958        let mut headers = HeaderMap::new();
959        assert_eq!(
960            sign_headers_at(
961                &Method::POST,
962                "not a URL",
963                &mut headers,
964                b"{}",
965                &credentials,
966                "us-east-1",
967                SERVICE,
968                "20240301T000000Z",
969            ),
970            Err(SigningError::InvalidUrl)
971        );
972        assert_eq!(
973            sign_headers_at(
974                &Method::POST,
975                "https://bedrock-runtime.us-east-1.amazonaws.com/model/demo/invoke",
976                &mut headers,
977                b"{}",
978                &credentials,
979                "us-east-1",
980                SERVICE,
981                "20240301T000000",
982            ),
983            Err(SigningError::InvalidTimestamp)
984        );
985    }
986
987    #[test]
988    fn malformed_eventstream_frames_fail_closed() {
989        let mut frame = event_frame(
990            br#"{"amazon-bedrock-invocationMetrics":{"inputTokenCount":1,"outputTokenCount":1}}"#,
991        );
992        frame[8] ^= 1;
993        let mut stream = EventStreamScanner::new(crate::proxy::usage::Scanner::new(
994            crate::proxy::usage::Provider::Anthropic,
995            None,
996        ));
997        stream.feed(&frame);
998        assert!(stream.finalize().is_none());
999
1000        let mut stream = EventStreamScanner::new(crate::proxy::usage::Scanner::new(
1001            crate::proxy::usage::Provider::Anthropic,
1002            None,
1003        ));
1004        stream.feed(&[0, 0, 0, 8, 0, 0, 0, 0, 0, 0, 0, 0]);
1005        assert!(stream.finalize().is_none());
1006
1007        let mut stream = EventStreamScanner::new(crate::proxy::usage::Scanner::new(
1008            crate::proxy::usage::Provider::Anthropic,
1009            None,
1010        ));
1011        stream.feed(&event_frame(b"not-json"));
1012        assert!(stream.finalize().is_none());
1013    }
1014
1015    #[tokio::test]
1016    async fn stalled_eventstream_can_be_bounded_by_request_timeout() {
1017        use futures::pin_mut;
1018        let stream = tee_eventstream(
1019            futures::stream::pending::<Result<Vec<u8>, std::convert::Infallible>>(),
1020            crate::proxy::usage::Scanner::new(crate::proxy::usage::Provider::Anthropic, None),
1021        );
1022        pin_mut!(stream);
1023        let result = tokio::time::timeout(std::time::Duration::ZERO, stream.next()).await;
1024        assert!(result.is_err());
1025    }
1026}