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