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