diff --git a/.github/workflows/run-integration-test.yml b/.github/workflows/run-integration-test.yml index d43ed9c71..ff93a51f2 100644 --- a/.github/workflows/run-integration-test.yml +++ b/.github/workflows/run-integration-test.yml @@ -54,13 +54,16 @@ jobs: sam deploy --stack-name "${stackName}" --parameter-overrides "ParameterKey=SecretToken,ParameterValue=${{ secrets.SECRET_TOKEN }}" "ParameterKey=LambdaRole,ParameterValue=${{ secrets.AWS_LAMBDA_ROLE }}" --no-confirm-changeset --no-progressbar > disable_output TEST_ENDPOINT=$(sam list stack-outputs --stack-name "${stackName}" --output json | jq -r '.[] | select(.OutputKey=="HelloApiEndpoint") | .OutputValue') TENANT_ID_TEST_FUNCTION=$(sam list stack-outputs --stack-name "${stackName}" --output json | jq -r '.[] | select(.OutputKey=="TenantIdTestFunction") | .OutputValue') + W3C_TEST_FUNCTION=$(sam list stack-outputs --stack-name "${stackName}" --output json | jq -r '.[] | select(.OutputKey=="W3CTestFunction") | .OutputValue') echo "TEST_ENDPOINT=$TEST_ENDPOINT" >> "$GITHUB_OUTPUT" echo "TENANT_ID_TEST_FUNCTION=$TENANT_ID_TEST_FUNCTION" >> "$GITHUB_OUTPUT" + echo "W3C_TEST_FUNCTION=$W3C_TEST_FUNCTION" >> "$GITHUB_OUTPUT" - name: run test env: SECRET_TOKEN: ${{ secrets.SECRET_TOKEN }} TEST_ENDPOINT: ${{ steps.deploy_stack.outputs.TEST_ENDPOINT }} TENANT_ID_TEST_FUNCTION: ${{ steps.deploy_stack.outputs.TENANT_ID_TEST_FUNCTION }} + W3C_TEST_FUNCTION: ${{ steps.deploy_stack.outputs.W3C_TEST_FUNCTION }} run: cd lambda-integration-tests && cargo test - name: cleanup if: always() && steps.deploy_stack.outputs.STACK_NAME diff --git a/lambda-integration-tests/Cargo.toml b/lambda-integration-tests/Cargo.toml index e2497333e..7a237f721 100644 --- a/lambda-integration-tests/Cargo.toml +++ b/lambda-integration-tests/Cargo.toml @@ -21,6 +21,7 @@ tracing = "0.1" [dev-dependencies] reqwest = { version = "0.13.1", features = ["blocking"] } +base64 = "0.22.1" [features] catch-all-fields = ["aws_lambda_events/catch-all-fields"] @@ -36,3 +37,7 @@ path = "src/authorizer.rs" [[bin]] name = "tenant-id-test" path = "src/tenant_id_test.rs" + +[[bin]] +name = "w3c-test" +path = "src/w3c_test.rs" diff --git a/lambda-integration-tests/src/w3c_test.rs b/lambda-integration-tests/src/w3c_test.rs new file mode 100644 index 000000000..dd348cc1f --- /dev/null +++ b/lambda-integration-tests/src/w3c_test.rs @@ -0,0 +1,34 @@ +use lambda_runtime::{service_fn, Error, LambdaEvent}; +use serde_json::{json, Value}; + +async fn function_handler(event: LambdaEvent) -> Result { + let (_event, context) = event.into_parts(); + + let w3c_fields = context.w3c(); + tracing::info!("w3c fields observed on context: {:?}", w3c_fields); + + let client_context_has_custom = context + .client_context + .as_ref() + .map(|cc| !cc.custom.is_empty()) + .unwrap_or(false); + + let response = json!({ + "statusCode": 200, + "body": json!({ + "message": "W3C test successful", + "request_id": context.request_id, + "w3c": w3c_fields, + "has_client_context": context.client_context.is_some(), + "client_context_has_custom": client_context_has_custom, + }).to_string() + }); + + Ok(response) +} + +#[tokio::main] +async fn main() -> Result<(), Error> { + lambda_runtime::tracing::init_default_subscriber(); + lambda_runtime::run(service_fn(function_handler)).await +} diff --git a/lambda-integration-tests/template.yaml b/lambda-integration-tests/template.yaml index 9ba694e32..1a5612bd5 100644 --- a/lambda-integration-tests/template.yaml +++ b/lambda-integration-tests/template.yaml @@ -53,6 +53,18 @@ Resources: Runtime: provided.al2023 Role: !Ref LambdaRole + W3CTestFunction: + Type: AWS::Serverless::Function + Metadata: + BuildMethod: rust-cargolambda + BuildProperties: + Binary: w3c-test + Properties: + CodeUri: ./ + Handler: bootstrap + Runtime: provided.al2023 + Role: !Ref LambdaRole + AuthorizerFunction: Type: AWS::Serverless::Function Metadata: @@ -74,4 +86,7 @@ Outputs: Value: !Sub "https://${API}.execute-api.${AWS::Region}.amazonaws.com/integ-test/hello/" TenantIdTestFunction: Description: "Tenant ID test function name" - Value: !Ref TenantIdTestFunction \ No newline at end of file + Value: !Ref TenantIdTestFunction + W3CTestFunction: + Description: "W3C trace-context test function name" + Value: !Ref W3CTestFunction \ No newline at end of file diff --git a/lambda-integration-tests/tests/w3c_prod_test.rs b/lambda-integration-tests/tests/w3c_prod_test.rs new file mode 100644 index 000000000..75b28be40 --- /dev/null +++ b/lambda-integration-tests/tests/w3c_prod_test.rs @@ -0,0 +1,107 @@ +use base64::prelude::*; +use serde_json::json; + +fn function_name() -> String { + std::env::var("W3C_TEST_FUNCTION").expect("W3C_TEST_FUNCTION environment variable not set") +} + +/// Invoke the deployed Lambda function and return the parsed body JSON. +fn invoke(payload: &serde_json::Value, client_context_json: Option<&serde_json::Value>) -> serde_json::Value { + let response_path = format!( + "/tmp/w3c_response_{}.json", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos() + ); + + let mut args: Vec = vec![ + "lambda".into(), + "invoke".into(), + "--function-name".into(), + function_name(), + "--payload".into(), + payload.to_string(), + "--cli-binary-format".into(), + "raw-in-base64-out".into(), + ]; + + if let Some(cc) = client_context_json { + let encoded = BASE64_STANDARD.encode(cc.to_string()); + args.push("--client-context".into()); + args.push(encoded); + } + + args.push(response_path.clone()); + + let output = std::process::Command::new("aws") + .args(&args) + .output() + .expect("Failed to invoke Lambda function"); + + assert!( + output.status.success(), + "Lambda invocation failed: {}", + String::from_utf8_lossy(&output.stderr) + ); + + let response = std::fs::read_to_string(&response_path).expect("Failed to read response file"); + let response_json: serde_json::Value = serde_json::from_str(&response).expect("Failed to parse response JSON"); + + assert_eq!( + response_json["statusCode"], + 200, + "handler returned non-200: {}", + serde_json::to_string_pretty(&response_json).unwrap() + ); + + let body: serde_json::Value = + serde_json::from_str(response_json["body"].as_str().expect("Body should be a string")) + .expect("Failed to parse body JSON"); + + let _ = std::fs::remove_file(&response_path); + body +} + +#[test] +fn test_w3c_propagation_in_production() { + // 1. No client context — `context.w3c()` must be empty. + let body = invoke(&json!({ "test": "w3c_no_client_context" }), None); + assert_eq!(body["message"], "W3C test successful"); + assert_eq!(body["has_client_context"], false); + assert_eq!( + body["w3c"], + json!({}), + "w3c should be empty when no clientContext header is present, got: {}", + body["w3c"] + ); + + // 2. Client context carrying `w3c` + a sibling `custom` field — `w3c()` + // must surface all three allowlisted fields, and the sibling `custom` + // must still be reachable on `context.client_context` (the `w3c` key + // was stripped during Context construction, not the whole object). + let client_context = json!({ + "custom": { "source": "integ-test" }, + "w3c": { + "traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01", + "tracestate": "rojo=00f067aa0ba902b7", + "baggage": "userId=alice" + } + }); + + let body = invoke(&json!({ "test": "w3c_with_client_context" }), Some(&client_context)); + assert_eq!(body["message"], "W3C test successful"); + assert_eq!(body["has_client_context"], true); + assert_eq!( + body["client_context_has_custom"], true, + "sibling clientContext fields must still be reachable after w3c is stripped" + ); + assert_eq!( + body["w3c"]["traceparent"], + "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01" + ); + assert_eq!(body["w3c"]["tracestate"], "rojo=00f067aa0ba902b7"); + assert_eq!(body["w3c"]["baggage"], "userId=alice"); + + println!("✅ W3C trace-context propagation test passed"); +} diff --git a/lambda-runtime/src/constants.rs b/lambda-runtime/src/constants.rs index 98c789d0f..6e123890f 100644 --- a/lambda-runtime/src/constants.rs +++ b/lambda-runtime/src/constants.rs @@ -7,3 +7,4 @@ pub(crate) const LAMBDA_RUNTIME_CLIENT_CONTEXT: &str = "lambda-runtime-client-co pub(crate) const LAMBDA_RUNTIME_COGNITO_IDENTITY: &str = "lambda-runtime-cognito-identity"; pub(crate) const LAMBDA_RUNTIME_TENANT_ID: &str = "lambda-runtime-aws-tenant-id"; pub(crate) const LAMBDA_RUNTIME_INVOCATION_ID: &str = "lambda-runtime-invocation-id"; +pub(crate) const W3C_ALLOWED_FIELDS: &[&str] = &["traceparent", "tracestate", "baggage"]; diff --git a/lambda-runtime/src/types.rs b/lambda-runtime/src/types.rs index edd932402..aa78504a7 100644 --- a/lambda-runtime/src/types.rs +++ b/lambda-runtime/src/types.rs @@ -2,7 +2,7 @@ use crate::{ constants::{ LAMBDA_RUNTIME_CLIENT_CONTEXT, LAMBDA_RUNTIME_COGNITO_IDENTITY, LAMBDA_RUNTIME_DEADLINE_MS, LAMBDA_RUNTIME_INVOKED_FUNCTION_ARN, LAMBDA_RUNTIME_REQUEST_ID, LAMBDA_RUNTIME_TENANT_ID, - LAMBDA_RUNTIME_TRACE_ID, + LAMBDA_RUNTIME_TRACE_ID, W3C_ALLOWED_FIELDS, }, Error, RefConfig, }; @@ -92,6 +92,10 @@ pub struct Context { /// Includes information such as the function name, memory allocation, /// version, and log streams. pub env_config: RefConfig, + /// Allowlisted W3C trace-context fields (`traceparent`, `tracestate`, + /// `baggage`) carried on `clientContext.w3c` at invoke time. + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub(crate) w3c_fields: HashMap, } impl Default for Context { @@ -105,6 +109,7 @@ impl Default for Context { identity: None, tenant_id: None, env_config: std::sync::Arc::new(crate::Config::default()), + w3c_fields: HashMap::new(), } } } @@ -113,15 +118,22 @@ impl Context { /// Create a new [Context] struct based on the function configuration /// and the incoming request data. pub fn new(request_id: &str, env_config: RefConfig, headers: &HeaderMap) -> Result { - let client_context: Option = if let Some(value) = headers.get(LAMBDA_RUNTIME_CLIENT_CONTEXT) { - let raw = value.to_str()?; - if raw.is_empty() { - None + let mut client_context_value: Option = + if let Some(value) = headers.get(LAMBDA_RUNTIME_CLIENT_CONTEXT) { + let raw = value.to_str()?; + if raw.is_empty() { + None + } else { + Some(serde_json::from_str(raw)?) + } } else { - Some(serde_json::from_str(raw)?) - } - } else { - None + None + }; + + let w3c_fields = Self::extract_and_strip_w3c(client_context_value.as_mut()); + let client_context: Option = match client_context_value { + Some(v) => Some(serde_json::from_value(v)?), + None => None, }; let identity: Option = if let Some(value) = headers.get(LAMBDA_RUNTIME_COGNITO_IDENTITY) { @@ -158,6 +170,7 @@ impl Context { .get(LAMBDA_RUNTIME_TENANT_ID) .map(|v| String::from_utf8_lossy(v.as_bytes()).to_string()), env_config, + w3c_fields, }; Ok(ctx) @@ -167,6 +180,37 @@ impl Context { pub fn deadline(&self) -> SystemTime { SystemTime::UNIX_EPOCH + Duration::from_millis(self.deadline) } + + /// Return the W3C trace-context fields (`traceparent`, `tracestate`, + /// `baggage`) carried on `clientContext.w3c` at invoke time. + pub fn w3c(&self) -> HashMap { + self.w3c_fields.clone() + } + + /// Pop `w3c` out of the parsed `clientContext` and return a normalized + /// map of the allowlisted string fields (see `W3C_ALLOWED_FIELDS`). + fn extract_and_strip_w3c(client_context: Option<&mut serde_json::Value>) -> HashMap { + let Some(client_context) = client_context else { + return HashMap::new(); + }; + let Some(obj) = client_context.as_object_mut() else { + return HashMap::new(); + }; + let Some(raw_w3c) = obj.remove("w3c") else { + return HashMap::new(); + }; + let Some(w3c_obj) = raw_w3c.as_object() else { + return HashMap::new(); + }; + + let mut fields = HashMap::new(); + for key in W3C_ALLOWED_FIELDS { + if let Some(serde_json::Value::String(s)) = w3c_obj.get(*key) { + fields.insert((*key).to_string(), s.clone()); + } + } + fields + } } /// Extract the invocation request id from the incoming request. @@ -543,4 +587,141 @@ mod test { let context = Context::new("id", config, &headers).unwrap(); assert_eq!(context.tenant_id, None); } + + // ----- W3C trace-context (`context.w3c()`) tests --------------------------------- + // + // These mirror the Python RIC tests at + // tests/test_lambda_context.py::TestLambdaContextW3C + // and the Node.js RIC tests at + // src/context/context-builder.test.ts::describe("w3c", ...) + // Consolidated: each test covers one distinct behavior rather than one + // input shape. + + fn w3c_headers_with_client_context(client_context: &serde_json::Value) -> HeaderMap { + let mut headers = HeaderMap::new(); + headers.insert("lambda-runtime-aws-request-id", HeaderValue::from_static("my-id")); + headers.insert("lambda-runtime-deadline-ms", HeaderValue::from_static("123")); + headers.insert( + "lambda-runtime-client-context", + HeaderValue::from_str(&serde_json::to_string(client_context).unwrap()).unwrap(), + ); + headers + } + + #[test] + fn w3c_is_empty_when_no_w3c_is_present() { + let config = Arc::new(Config::default()); + + let mut no_header = HeaderMap::new(); + no_header.insert("lambda-runtime-aws-request-id", HeaderValue::from_static("my-id")); + no_header.insert("lambda-runtime-deadline-ms", HeaderValue::from_static("123")); + let ctx = Context::new("id", config.clone(), &no_header).unwrap(); + assert!(ctx.w3c().is_empty()); + + let no_w3c = w3c_headers_with_client_context(&serde_json::json!({ + "custom": { "value": "test" } + })); + let ctx = Context::new("id", config, &no_w3c).unwrap(); + assert!(ctx.w3c().is_empty()); + let cc = ctx.client_context.expect("client_context should be set"); + assert_eq!(cc.custom.get("value"), Some(&"test".to_string())); + } + + #[test] + fn w3c_returns_all_allowlisted_fields_and_strips_source() { + let config = Arc::new(Config::default()); + let headers = w3c_headers_with_client_context(&serde_json::json!({ + "custom": { "value": "test" }, + "w3c": { + "traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01", + "tracestate": "rojo=00f067aa0ba902b7", + "baggage": "userId=alice" + } + })); + + let ctx = Context::new("id", config, &headers).unwrap(); + let fields = ctx.w3c(); + assert_eq!(fields.len(), 3); + assert_eq!( + fields.get("traceparent"), + Some(&"00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01".to_string()) + ); + assert_eq!(fields.get("tracestate"), Some(&"rojo=00f067aa0ba902b7".to_string())); + assert_eq!(fields.get("baggage"), Some(&"userId=alice".to_string())); + + let cc = ctx.client_context.expect("client_context should be set"); + assert_eq!(cc.custom.get("value"), Some(&"test".to_string())); + } + + #[test] + fn w3c_drops_non_string_and_absent_allowlisted_keys() { + let config = Arc::new(Config::default()); + let headers = w3c_headers_with_client_context(&serde_json::json!({ + "w3c": { + "baggage": "abc", + "traceparent": 42, + "tracestate": null, + } + })); + + let ctx = Context::new("id", config, &headers).unwrap(); + let fields = ctx.w3c(); + assert_eq!(fields.len(), 1); + assert_eq!(fields.get("baggage"), Some(&"abc".to_string())); + assert!(!fields.contains_key("traceparent")); + assert!(!fields.contains_key("tracestate")); + } + + #[test] + fn w3c_treats_non_object_as_empty_and_still_strips_source() { + let config = Arc::new(Config::default()); + for bad_w3c in [serde_json::json!("not-an-object"), serde_json::json!(["baggage=abc"])] { + let headers = w3c_headers_with_client_context(&serde_json::json!({ + "w3c": bad_w3c, + "custom": { "k": "v" } + })); + let ctx = Context::new("id", config.clone(), &headers).unwrap(); + assert!(ctx.w3c().is_empty()); + let cc = ctx.client_context.expect("client_context should be set"); + assert_eq!(cc.custom.get("k"), Some(&"v".to_string())); + } + } + + #[test] + fn w3c_drops_non_allowlisted_keys() { + let config = Arc::new(Config::default()); + let headers = w3c_headers_with_client_context(&serde_json::json!({ + "w3c": { + "baggage": "keep=me", + "unknownField": "should-not-appear", + "x-custom-trace": "should-not-appear", + "__proto__": "should-not-appear", + "constructor": "should-not-appear", + "toString": "should-not-appear" + } + })); + + let ctx = Context::new("id", config, &headers).unwrap(); + let fields = ctx.w3c(); + assert_eq!(fields.len(), 1); + assert_eq!(fields.get("baggage"), Some(&"keep=me".to_string())); + } + + #[test] + fn w3c_returned_map_is_a_defensive_copy() { + let config = Arc::new(Config::default()); + let headers = w3c_headers_with_client_context(&serde_json::json!({ + "w3c": { "baggage": "abc" } + })); + + let ctx = Context::new("id", config, &headers).unwrap(); + let mut first = ctx.w3c(); + first.insert("baggage".to_string(), "tampered".to_string()); + first.insert("injected".to_string(), "nope".to_string()); + + let second = ctx.w3c(); + assert_eq!(second.len(), 1); + assert_eq!(second.get("baggage"), Some(&"abc".to_string())); + assert!(!second.contains_key("injected")); + } }