1use 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
93struct 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 ×tamp,
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 = ×tamp[..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(®ion_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 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}