diff --git a/.github/workflows/dockerized-test.yml b/.github/workflows/dockerized-test.yml new file mode 100644 index 00000000..c75d5f9f --- /dev/null +++ b/.github/workflows/dockerized-test.yml @@ -0,0 +1,31 @@ +name: dockerized-test + +permissions: + contents: read + +on: + push: + branches: [main] + pull_request: + branches: ['*'] + workflow_dispatch: + +jobs: + dockerized-test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v5 + + - name: Build the test image + run: | + docker build . \ + -t local/test \ + -f Dockerfile.test \ + --build-arg BASE_IMAGE=public.ecr.aws/lambda/provided:al2023 + + - name: Run dockerized suites + uses: aws/containerized-test-runner-for-aws-lambda@0863dd17b5fc19585250a2405c0f939a77b4f397 # main + with: + suiteFileArray: '["./test/dockerized/suites/*.json"]' + dockerImageName: 'local/test' + taskFolder: './test/dockerized/tasks' diff --git a/Dockerfile.test b/Dockerfile.test new file mode 100644 index 00000000..e0d9d11a --- /dev/null +++ b/Dockerfile.test @@ -0,0 +1,13 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# SPDX-License-Identifier: Apache-2.0 + +ARG BASE_IMAGE=public.ecr.aws/lambda/provided:al2023 + +FROM golang:1.26 AS build +WORKDIR /src +COPY . . +RUN CGO_ENABLED=0 GOOS=linux go build -o /bootstrap ./test/dockerized/tasks + +FROM ${BASE_IMAGE} +COPY --from=build /bootstrap ${LAMBDA_RUNTIME_DIR}/bootstrap +CMD [ "bootstrap" ] diff --git a/lambda/invoke_loop.go b/lambda/invoke_loop.go index b6857f99..01e766ab 100644 --- a/lambda/invoke_loop.go +++ b/lambda/invoke_loop.go @@ -55,7 +55,7 @@ func handleInvoke(invoke *invoke, handler *handlerOptions) error { InvokedFunctionArn: invoke.headers.Get(headerInvokedFunctionARN), TenantID: invoke.headers.Get(headerTenantID), } - if err := parseClientContext(invoke, &lc.ClientContext); err != nil { + if err := parseClientContext(invoke, &lc); err != nil { return reportFailure(invoke, lambdaErrorResponse(err)) } if err := parseCognitoIdentity(invoke, &lc.Identity); err != nil { @@ -147,13 +147,19 @@ func parseCognitoIdentity(invoke *invoke, out *lambdacontext.CognitoIdentity) er return nil } -func parseClientContext(invoke *invoke, out *lambdacontext.ClientContext) error { +func parseClientContext(invoke *invoke, lc *lambdacontext.LambdaContext) error { clientContextJSON := invoke.headers.Get(headerClientContext) - if clientContextJSON != "" { - if err := json.Unmarshal([]byte(clientContextJSON), out); err != nil { - return fmt.Errorf("failed to unmarshal client context json: %v", err) - } + if clientContextJSON == "" { + return nil + } + if err := json.Unmarshal([]byte(clientContextJSON), &lc.ClientContext); err != nil { + return fmt.Errorf("failed to unmarshal client context json: %v", err) } + // Extract the allowlisted W3C trace-context fields from the raw header. + // clientContext.w3c is not part of the ClientContext struct, so it is + // dropped from lc.ClientContext on unmarshal and is only reachable via + // lc.W3C(). + lc.ExtractW3C([]byte(clientContextJSON)) return nil } diff --git a/lambdacontext/context.go b/lambdacontext/context.go index f3e70399..ee981f3f 100644 --- a/lambdacontext/context.go +++ b/lambdacontext/context.go @@ -10,12 +10,27 @@ package lambdacontext import ( + "bytes" "context" "encoding/json" "os" "strconv" ) +// w3cAllowedFields is the allowlist of W3C trace-context fields that may be +// surfaced through LambdaContext.W3C(). Any other key carried on +// clientContext.w3c is ignored, and any allowlisted key whose value is not a +// JSON string is dropped. +var w3cAllowedFields = [...]string{"traceparent", "tracestate", "baggage"} + +// W3CAllowedFields returns the allowlist of W3C trace-context fields that may be +// surfaced through LambdaContext.W3C() (traceparent, tracestate, baggage). +func W3CAllowedFields() []string { + out := make([]string, len(w3cAllowedFields)) + copy(out, w3cAllowedFields[:]) + return out +} + // LogGroupName is the name of the log group that contains the log streams of the current Lambda Function var LogGroupName string @@ -111,6 +126,65 @@ type LambdaContext struct { Identity CognitoIdentity ClientContext ClientContext TenantID string `json:",omitempty"` + w3c map[string]string +} + +// W3C returns the W3C trace-context fields (see W3CAllowedFields) that were +// carried on clientContext.w3c at invoke time. The returned map is a fresh copy +// on every call, so mutating it never affects the context. It is never nil; an +// invoke that carried no W3C trace-context yields an empty map. +// +// The w3c key is not part of ClientContext, so it is deliberately never +// surfaced through lc.ClientContext — W3C() is the only accessor. +func (lc *LambdaContext) W3C() map[string]string { + out := make(map[string]string, len(lc.w3c)) + for k, v := range lc.w3c { + out[k] = v + } + return out +} + +func (lc *LambdaContext) ExtractW3C(clientContextJSON []byte) { + lc.w3c = extractW3CFields(clientContextJSON) +} + +func extractW3CFields(clientContextJSON []byte) map[string]string { + fields := map[string]string{} + if len(clientContextJSON) == 0 { + return fields + } + + var envelope struct { + W3C json.RawMessage `json:"w3c"` + } + if err := json.Unmarshal(clientContextJSON, &envelope); err != nil || len(envelope.W3C) == 0 { + return fields + } + + // w3c must be a JSON object; a string, array, number, etc. yields empty. + var raw map[string]json.RawMessage + if err := json.Unmarshal(envelope.W3C, &raw); err != nil { + return fields + } + + for _, key := range w3cAllowedFields { + value, ok := raw[key] + if !ok { + continue + } + // Only keep values that are genuine JSON strings. A JSON string always + // begins with a double-quote, so this rejects numbers, null, objects + // and arrays without a second unmarshal attempt. + trimmed := bytes.TrimSpace(value) + if len(trimmed) == 0 || trimmed[0] != '"' { + continue + } + var s string + if err := json.Unmarshal(trimmed, &s); err == nil { + fields[key] = s + } + } + return fields } // An unexported type to be used as the key for types in this package. diff --git a/lambdacontext/context_w3c_test.go b/lambdacontext/context_w3c_test.go new file mode 100644 index 00000000..6ac5e034 --- /dev/null +++ b/lambdacontext/context_w3c_test.go @@ -0,0 +1,101 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package lambdacontext + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestLambdaContextW3C(t *testing.T) { + traceparent := "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01" + + cases := []struct { + name string + clientContext string + expected map[string]string + }{ + { + name: "empty when no client context", + clientContext: "", + expected: map[string]string{}, + }, + { + name: "empty when client context has no w3c key", + clientContext: `{"custom":{"value":"test"}}`, + expected: map[string]string{}, + }, + { + name: "baggage only", + clientContext: `{"w3c":{"baggage":"userId=alice"}}`, + expected: map[string]string{"baggage": "userId=alice"}, + }, + { + name: "all three allowlisted fields", + clientContext: `{"custom":{"value":"test"},"w3c":{"traceparent":"` + traceparent + `","tracestate":"rojo=00f067aa0ba902b7","baggage":"userId=alice"}}`, + expected: map[string]string{ + "traceparent": traceparent, + "tracestate": "rojo=00f067aa0ba902b7", + "baggage": "userId=alice", + }, + }, + { + name: "allowlist drops non-allowlisted keys", + clientContext: `{"w3c":{"baggage":"keep=me","unknownField":"nope","x-custom-trace":"nope"}}`, + expected: map[string]string{"baggage": "keep=me"}, + }, + { + name: "drops allowlisted fields with non-string values", + clientContext: `{"w3c":{"traceparent":42,"tracestate":null,"baggage":{"nested":"no"}}}`, + expected: map[string]string{}, + }, + { + name: "non-object w3c treated as empty", + clientContext: `{"w3c":"not-an-object"}`, + expected: map[string]string{}, + }, + { + name: "array w3c treated as empty", + clientContext: `{"w3c":["baggage=abc"]}`, + expected: map[string]string{}, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + lc := &LambdaContext{} + lc.ExtractW3C([]byte(tc.clientContext)) + assert.Equal(t, tc.expected, lc.W3C()) + }) + } +} + +func TestLambdaContextW3CReturnsFreshCopy(t *testing.T) { + lc := &LambdaContext{} + lc.ExtractW3C([]byte(`{"w3c":{"baggage":"abc"}}`)) + + // Mutating a returned map must not affect the context's internal state. + first := lc.W3C() + first["baggage"] = "tampered" + first["injected"] = "nope" + + assert.Equal(t, map[string]string{"baggage": "abc"}, lc.W3C()) +} + +func TestLambdaContextW3CNeverNil(t *testing.T) { + // A zero-value context (ExtractW3C never called) still yields a usable, + // non-nil empty map. + lc := &LambdaContext{} + assert.Equal(t, map[string]string{}, lc.W3C()) +} + +func TestW3CAllowedFieldsIsImmutable(t *testing.T) { + assert.Equal(t, []string{"traceparent", "tracestate", "baggage"}, W3CAllowedFields()) + got := W3CAllowedFields() + got[0] = "tampered" + got = append(got, "injected") + assert.NotEqual(t, got, W3CAllowedFields()) + assert.Equal(t, []string{"traceparent", "tracestate", "baggage"}, W3CAllowedFields()) +} diff --git a/test/dockerized/suites/ctx.json b/test/dockerized/suites/ctx.json new file mode 100644 index 00000000..a80afd77 --- /dev/null +++ b/test/dockerized/suites/ctx.json @@ -0,0 +1,46 @@ +{ + "tests": [ + { + "name": "client_context_siblings_are_echoed_and_w3c_is_stripped", + "handler": "w3c.getW3cAndSource", + "request": {}, + "clientContext": { + "custom": { "value": "test" }, + "w3c": { + "traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01", + "baggage": "userId=alice" + } + }, + "assertions": [ + { + "response": { + "w3c": { + "traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01", + "baggage": "userId=alice" + }, + "clientContextCustom": { "value": "test" }, + "clientContextHasW3c": false + } + } + ] + }, + + { + "name": "client_context_is_echoed_when_no_w3c_key", + "handler": "w3c.echoClientContext", + "request": {}, + "clientContext": { + "custom": { "value": "hello" }, + "env": { "stage": "beta" } + }, + "assertions": [ + { + "response": { + "custom": { "value": "hello" }, + "env": { "stage": "beta" } + } + } + ] + } + ] +} diff --git a/test/dockerized/suites/w3c.json b/test/dockerized/suites/w3c.json new file mode 100644 index 00000000..2072fdf5 --- /dev/null +++ b/test/dockerized/suites/w3c.json @@ -0,0 +1,114 @@ +{ + "tests": [ + { + "name": "w3c_returns_empty_when_no_client_context_header", + "handler": "w3c.getW3c", + "request": {}, + "assertions": [ + { "response": {} } + ] + }, + + { + "name": "w3c_returns_empty_when_client_context_has_no_w3c_key", + "handler": "w3c.getW3c", + "request": {}, + "clientContext": { + "custom": { "value": "test" } + }, + "assertions": [ + { "response": {} } + ] + }, + + { + "name": "w3c_returns_baggage_only", + "handler": "w3c.getW3c", + "request": {}, + "clientContext": { + "w3c": { "baggage": "userId=alice" } + }, + "assertions": [ + { "response": { "baggage": "userId=alice" } } + ] + }, + + { + "name": "w3c_returns_all_three_allowlisted_fields", + "handler": "w3c.getW3c", + "request": {}, + "clientContext": { + "w3c": { + "traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01", + "tracestate": "rojo=00f067aa0ba902b7", + "baggage": "userId=alice" + } + }, + "assertions": [ + { + "response": { + "traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01", + "tracestate": "rojo=00f067aa0ba902b7", + "baggage": "userId=alice" + } + } + ] + }, + + { + "name": "w3c_allowlist_drops_non_allowlisted_keys", + "handler": "w3c.getW3c", + "request": {}, + "clientContext": { + "w3c": { + "baggage": "keep=me", + "unknownField": "should-not-appear", + "x-custom-trace": "should-not-appear" + } + }, + "assertions": [ + { "response": { "baggage": "keep=me" } } + ] + }, + + { + "name": "w3c_drops_allowlisted_fields_with_non_string_values", + "handler": "w3c.getW3c", + "request": {}, + "clientContext": { + "w3c": { + "traceparent": 42, + "tracestate": null, + "baggage": { "nested": "no" } + } + }, + "assertions": [ + { "response": {} } + ] + }, + + { + "name": "w3c_treats_non_object_as_empty", + "handler": "w3c.getW3c", + "request": {}, + "clientContext": { + "w3c": "not-an-object" + }, + "assertions": [ + { "response": {} } + ] + }, + + { + "name": "w3c_treats_array_as_empty", + "handler": "w3c.getW3c", + "request": {}, + "clientContext": { + "w3c": ["baggage=abc"] + }, + "assertions": [ + { "response": {} } + ] + } + ] +} diff --git a/test/dockerized/tasks/main.go b/test/dockerized/tasks/main.go new file mode 100644 index 00000000..174b6c28 --- /dev/null +++ b/test/dockerized/tasks/main.go @@ -0,0 +1,47 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "context" + "encoding/json" + "fmt" + "os" + + "github.com/aws/aws-lambda-go/lambda" + "github.com/aws/aws-lambda-go/lambdacontext" +) + +func dispatch(ctx context.Context, _ json.RawMessage) (interface{}, error) { + lc, ok := lambdacontext.FromContext(ctx) + if !ok { + return nil, fmt.Errorf("no lambda context on invoke context") + } + + switch handler := os.Getenv("_HANDLER"); handler { + case "w3c.getW3c": + return lc.W3C(), nil + + case "w3c.getW3cAndSource": + _, hasW3c := lc.ClientContext.Custom["w3c"] + return map[string]interface{}{ + "w3c": lc.W3C(), + "clientContextCustom": lc.ClientContext.Custom, + "clientContextHasW3c": hasW3c, + }, nil + + case "w3c.echoClientContext": + return map[string]interface{}{ + "custom": lc.ClientContext.Custom, + "env": lc.ClientContext.Env, + }, nil + + default: + return nil, fmt.Errorf("unknown handler: %q", handler) + } +} + +func main() { + lambda.Start(dispatch) +}