diff --git a/README.md b/README.md index 38578c2..adf5eca 100644 --- a/README.md +++ b/README.md @@ -72,4 +72,25 @@ All timestamps should be RFC3339 in UTC, a subset of ISO8601. For example: `2020 ## Tracing Traces consist of a `request_id`, a `trace_id`, and the `trace_start` timestamp. A `trace_id` should begin when an "event" starts in our system (i.e, a customer request comes in, a cron job starts, etc) and travels across services. A `request_id` is a unique identifier for a single request: i.e, within the bounds of a single service. A trace may outlive a request, but a request will always be part of a trace. +### HTTP header contract +The trace crosses service boundaries as HTTP headers: + +| Header | Meaning | Required for interop | +|--------|---------|----------------------| +| `X-Trace-ID` | stable id for the whole trace | yes | +| `X-Request-ID` | id for a single request; a fresh one is minted per outbound sub-request | yes | +| `X-Trace-Start` | RFC3339 timestamp of trace origin | no (metadata) | +| `X-Trace-Source` | service that originated the trace | no (metadata) | +| `X-Request-Source` | service that originated this request | no (metadata) | + +Only the first two are required; the rest are metadata. Every hop applies **propagate-or-generate** to both required ids: reuse the valid inbound value, otherwise mint one. One `trace_id` spans the whole trace. A fresh `request_id` is minted per **outbound** sub-request (`ClientMiddleware`); an inbound edge honors a valid caller-supplied `X-Request-ID` (so a caller that already tagged its request keeps that id through the hop), and mints one only when it is absent or invalid. + +### Go usage (`trace` package) +- **Inbound:** `trc := trace.FromHeaderOrNew(r.Header)` then `ctx = trace.CtxWith(ctx, trc)` (or use `trace.ServerMiddleware`). Store the `Trace` on the request context so downstream code and logs pick it up. +- **Outbound:** wrap your client transport with `trace.ClientMiddleware(transport)`. It reads the `Trace` off the request context, preserves the `TraceID`, and mints a fresh `RequestID` per call. +- **Response echo:** `trace.SaveToHeader(w.Header(), trc)` so the caller (and a browser, via CORS `Access-Control-Expose-Headers`) can read the id back. Set it *before* you start streaming a response — SSE handlers flush headers early. + +### Inbound validation +`FromHeaderOrNew` treats inbound header values as untrusted input. An `X-Trace-ID` / `X-Request-ID` that is empty, longer than 200 bytes, or contains any byte outside `[A-Za-z0-9._-]` (e.g. CR/LF, control bytes, spaces) is rejected and a fresh id is generated in its place — preventing header/log injection (CWE-93). The `X-Trace-Source` / `X-Request-Source` metadata values are validated the same way and dropped (blanked) rather than regenerated when invalid. The Python port (`py/trace.py`) applies the identical length cap and charset. This means a value that survives validation in one service is accepted unchanged by the next. + diff --git a/py/trace.py b/py/trace.py index 9170d23..053b28d 100644 --- a/py/trace.py +++ b/py/trace.py @@ -1,6 +1,6 @@ +import re from dataclasses import dataclass from datetime import datetime, timezone -from typing import Optional from uuid7 import uuid7 @@ -9,6 +9,58 @@ def as_rfc3339(dt: datetime) -> str: return dt.astimezone(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%f")[:-4] + "Z" +# An inbound id may be at most _MAX_ID_LEN chars and contain only [A-Za-z0-9._-]. +# Anything else (empty, over-long, or carrying CR/LF, control, or non-ascii bytes) +# is untrusted: an id is regenerated and a source is dropped, so a poisoned inbound +# header can never be re-emitted by save_to_headers to corrupt a downstream request +# or a log line. Mirrors validID in the Go trace package. +_MAX_ID_LEN = 200 +_VALID_ID = re.compile(r"\A[A-Za-z0-9._-]+\Z") + +# A source is a service name, not an id, so it is bounded far tighter than an id. +# Mirrors maxSourceLen in the Go trace package. +_MAX_SOURCE_LEN = 64 + + +def _valid_id(s: str) -> bool: + return bool(s) and len(s) <= _MAX_ID_LEN and _VALID_ID.match(s) is not None + + +def _resolve_id(s: str) -> str: + return s if _valid_id(s) else uuid7() + + +def _valid_source(s: str) -> bool: + # Empty is the unset / "unknown" case; otherwise a charset-valid slug no longer + # than _MAX_SOURCE_LEN. Mirrors Go's validSource. + return s == "" or (len(s) <= _MAX_SOURCE_LEN and _valid_id(s)) + + +def _resolve_source(s: str) -> str: + # Match Go's resolveSource: keep a valid source, drop anything else to "" so + # save_to_headers can never re-emit unsafe or oversized bytes. + return s if _valid_source(s) else "" + + +def _resolve_start(s: str, now: datetime) -> str: + # Parse an inbound RFC3339 timestamp and re-emit it canonically, falling back + # to `now` when the header is absent, malformed, or in the future. Clamping a + # future value to now silently (like a malformed one) mirrors Go's + # FromHeaderOrNew and stops a caller from choosing the trace's start time. + # Re-formatting through as_rfc3339 also means the stored value can never carry + # CR/LF (or any other injected bytes) into save_to_headers. + if s: + try: + parsed = datetime.fromisoformat(s.replace("Z", "+00:00")) + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=timezone.utc) + if parsed <= now: + return as_rfc3339(parsed) + except ValueError: + pass + return as_rfc3339(now) + + @dataclass class Trace: """A trace object that can be used to track a request through multiple services. @@ -25,17 +77,24 @@ class Trace: def from_headers(headers: dict[str, str]) -> "Trace": """get a trace from a dictionary of headers, or create a new one if it doesn't exist. this will over-write the global trace. + + Inbound header values are untrusted: an absent or malformed X-Trace-ID / + X-Request-ID is replaced with a fresh uuid, a malformed source is dropped + to "", and a malformed X-Trace-Start falls back to now (see _valid_id / + _resolve_start). There is no X-Request-Start header — the request timing + starts when the server receives the request, matching Go's SaveToHeader. """ global _trace - now = as_rfc3339(datetime.now()) + now_dt = datetime.now(timezone.utc) + now = as_rfc3339(now_dt) t = Trace( - request_id=headers.get("X-Request-ID", uuid7()), - request_source=headers.get("X-Request-Source", "unknown"), - request_start=headers.get("X-Request-Start", now), - trace_id=headers.get("X-Trace-ID", uuid7()), - trace_source=headers.get("X-Trace-Source", "unknown"), - trace_start=headers.get("X-Trace-Start", now), + request_id=_resolve_id(headers.get("X-Request-ID", "")), + request_source=_resolve_source(headers.get("X-Request-Source", "")), + request_start=now, + trace_id=_resolve_id(headers.get("X-Trace-ID", "")), + trace_source=_resolve_source(headers.get("X-Trace-Source", "")), + trace_start=_resolve_start(headers.get("X-Trace-Start", ""), now_dt), ) _trace = t return t @@ -54,9 +113,9 @@ def current(cls) -> "Trace": return _trace @staticmethod - def new(): + def new() -> "Trace": """start a fresh trace and return it, overwriting the global trace if it exists.""" - now = as_rfc3339(datetime.datetime.now()) + now = as_rfc3339(datetime.now(timezone.utc)) global _trace t = Trace( request_id=uuid7(), @@ -71,13 +130,21 @@ def new(): def save_to_headers(self, headers: dict[str, str]) -> None: """save the trace to a dictionary of headers in preparation for an HTTP request. - This creates a new request_id and sets the request_start time to now so that the next service in the chain can add its own trace information. + + Writes the same five headers as the Go trace.SaveToHeader, and mints a fresh + X-Request-ID: this is a new request within the same trace, so the trace_id + persists across the hop while the request_id identifies this sub-request. + + Every stored field is already validated at construction (ids charset-checked, + source dropped to "" if invalid, trace_start re-formatted through as_rfc3339), + so no value written here can carry CR/LF into an outbound header. """ - headers["X-Request-ID"] = uuid7(), - headers["X-Request-Source"] = self.request_source - headers["X-Request-Start"] = as_rfc3339(datetime.now()) headers["X-Trace-ID"] = self.trace_id + headers["X-Request-ID"] = uuid7() + headers["X-Trace-Start"] = self.trace_start + headers["X-Trace-Source"] = self.trace_source + headers["X-Request-Source"] = self.request_source """the current trace, if any. this is only valid in a truly single-threaded environment.""" -_trace: Optional[Trace] = None +_trace: "Trace | None" = None diff --git a/py/trace_test.py b/py/trace_test.py new file mode 100644 index 0000000..92b18f2 --- /dev/null +++ b/py/trace_test.py @@ -0,0 +1,193 @@ +import unittest +from datetime import datetime, timedelta, timezone +from trace import Trace, _resolve_id, _resolve_source, _valid_id + + +class TestValidId(unittest.TestCase): + def test_cases(self): + cases = [ + {"description": "canonical uuid", "input": "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", "expected": True}, + {"description": "service-prefixed id", "input": "trace_018f3a1c2b4d", "expected": True}, + {"description": "dotted id", "input": "svc.abc.123", "expected": True}, + {"description": "max length accepted", "input": "a" * 200, "expected": True}, + {"description": "empty rejected", "input": "", "expected": False}, + {"description": "one over max rejected", "input": "a" * 201, "expected": False}, + {"description": "CRLF rejected", "input": "abc\r\nX-Evil: 1", "expected": False}, + {"description": "newline rejected", "input": "abc\ndef", "expected": False}, + {"description": "space rejected", "input": "abc def", "expected": False}, + {"description": "control byte rejected", "input": "abc\x00def", "expected": False}, + {"description": "non-ascii rejected", "input": "abcé", "expected": False}, + {"description": "slash rejected", "input": "a/b", "expected": False}, + ] + for c in cases: + with self.subTest(c["description"]): + self.assertEqual(_valid_id(c["input"]), c["expected"]) + + +class TestResolveId(unittest.TestCase): + def test_cases(self): + cases = [ + {"description": "valid id propagated", "input": "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", "propagate": True}, + {"description": "absent id generated", "input": "", "propagate": False}, + {"description": "poisoned id regenerated", "input": "abc\r\nX-Evil: 1", "propagate": False}, + {"description": "over-long id regenerated", "input": "a" * 201, "propagate": False}, + ] + for c in cases: + with self.subTest(c["description"]): + got = _resolve_id(c["input"]) + if c["propagate"]: + self.assertEqual(got, c["input"]) + else: + self.assertNotEqual(got, c["input"]) + self.assertTrue(_valid_id(got), f"generated id {got!r} is not valid") + + +class TestResolveSource(unittest.TestCase): + def test_cases(self): + cases = [ + {"description": "empty source stays empty", "input": "", "expected": ""}, + {"description": "valid source kept", "input": "runpod-graphql", "expected": "runpod-graphql"}, + {"description": "source at max length kept", "input": "a" * 64, "expected": "a" * 64}, + {"description": "poisoned source dropped to empty", "input": "svc\r\nX-Evil: 1", "expected": ""}, + {"description": "spaced source dropped to empty", "input": "not a slug", "expected": ""}, + {"description": "over-length source dropped to empty", "input": "a" * 65, "expected": ""}, + {"description": "charset-valid but over-length source dropped", "input": "a" * 200, "expected": ""}, + ] + for c in cases: + with self.subTest(c["description"]): + self.assertEqual(_resolve_source(c["input"]), c["expected"]) + + +class TestFromHeaders(unittest.TestCase): + def test_cases(self): + good_trace = "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31" + good_req = "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061" + cases = [ + { + "description": "valid ids propagated", + "headers": {"X-Trace-ID": good_trace, "X-Request-ID": good_req, "X-Trace-Source": "main-ui"}, + "want_trace": good_trace, + "want_request": good_req, + "want_trace_source": "main-ui", + }, + { + "description": "absent ids generated", + "headers": {}, + "want_trace": None, + "want_request": None, + "want_trace_source": "", + }, + { + "description": "poisoned trace id regenerated, valid request id kept", + "headers": {"X-Trace-ID": "abc\r\nX-Evil: 1", "X-Request-ID": good_req}, + "want_trace": None, + "want_request": good_req, + "want_trace_source": "", + }, + { + "description": "poisoned source dropped to empty", + "headers": {"X-Trace-ID": good_trace, "X-Request-ID": good_req, "X-Trace-Source": "svc\r\nX-Evil"}, + "want_trace": good_trace, + "want_request": good_req, + "want_trace_source": "", + }, + ] + for c in cases: + with self.subTest(c["description"]): + t = Trace.from_headers(dict(c["headers"])) + if c["want_trace"] is None: + self.assertTrue(_valid_id(t.trace_id)) + self.assertNotIn("\r", t.trace_id) + else: + self.assertEqual(t.trace_id, c["want_trace"]) + if c["want_request"] is None: + self.assertTrue(_valid_id(t.request_id)) + else: + self.assertEqual(t.request_id, c["want_request"]) + self.assertEqual(t.trace_source, c["want_trace_source"]) + + +class TestSaveToHeaders(unittest.TestCase): + def test_writes_go_five_headers_and_fresh_request_id(self): + t = Trace.from_headers({"X-Trace-ID": "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31"}) + original_request_id = t.request_id + headers: dict[str, str] = {} + t.save_to_headers(headers) + + for key in ("X-Trace-ID", "X-Request-ID", "X-Trace-Start", "X-Trace-Source", "X-Request-Source"): + self.assertIn(key, headers) + + self.assertEqual(headers["X-Trace-ID"], t.trace_id) + # Regression: X-Request-ID must be a plain string, not a tuple (trailing-comma bug). + self.assertIsInstance(headers["X-Request-ID"], str) + # A fresh sub-request id, distinct from the inbound one. + self.assertNotEqual(headers["X-Request-ID"], original_request_id) + + def test_emits_no_unsafe_bytes_from_poisoned_trace(self): + # Every inbound header poisoned, including X-Trace-Start — the timestamp is + # the header that previously leaked CR/LF straight through save_to_headers. + cases = [ + {"description": "poisoned trace id", "header": "X-Trace-ID", "value": "abc\r\nX-Evil: 1"}, + {"description": "poisoned trace source", "header": "X-Trace-Source", "value": "svc\r\nX-Evil: 2"}, + {"description": "poisoned request id", "header": "X-Request-ID", "value": "req\r\nX-Evil: 3"}, + {"description": "poisoned trace start", "header": "X-Trace-Start", "value": "2020-01-01T00:00:00Z\r\nX-Evil: 4"}, + ] + for c in cases: + with self.subTest(c["description"]): + t = Trace.from_headers({c["header"]: c["value"]}) + headers: dict[str, str] = {} + t.save_to_headers(headers) + for key, value in headers.items(): + self.assertNotIn("\r", value, f"{key} carries CR") + self.assertNotIn("\n", value, f"{key} carries LF") + + def test_malformed_trace_start_falls_back_to_now(self): + cases = [ + {"description": "non-RFC3339 falls back", "value": "not-a-date"}, + {"description": "CRLF-poisoned falls back", "value": "2020-01-01T00:00:00Z\r\nX-Evil: 1"}, + ] + for c in cases: + with self.subTest(c["description"]): + t = Trace.from_headers({"X-Trace-ID": "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", "X-Trace-Start": c["value"]}) + # The stored value re-parses as an RFC3339 timestamp (the fallback + # to now), never the poisoned input. + self.assertNotIn("\r", t.trace_start) + parsed = datetime.fromisoformat(t.trace_start.replace("Z", "+00:00")) + delta = abs((datetime.now(timezone.utc) - parsed).total_seconds()) + self.assertLess(delta, 5, f"trace_start {t.trace_start!r} is not ~now") + + def test_valid_trace_start_is_preserved(self): + t = Trace.from_headers( + {"X-Trace-ID": "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", "X-Trace-Start": "2020-01-01T00:00:00Z"} + ) + parsed = datetime.fromisoformat(t.trace_start.replace("Z", "+00:00")) + self.assertEqual((parsed.year, parsed.month, parsed.day), (2020, 1, 1)) + + def test_future_trace_start_is_clamped_to_now(self): + # A caller-supplied future timestamp is clamped to now (matching Go), so it + # never drives the trace's start time. + future = (datetime.now(timezone.utc) + timedelta(hours=1)).strftime("%Y-%m-%dT%H:%M:%SZ") + t = Trace.from_headers( + {"X-Trace-ID": "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", "X-Trace-Start": future} + ) + parsed = datetime.fromisoformat(t.trace_start.replace("Z", "+00:00")) + delta = abs((datetime.now(timezone.utc) - parsed).total_seconds()) + self.assertLess(delta, 5, f"future trace_start {t.trace_start!r} was not clamped to now") + + +class TestNewRegression(unittest.TestCase): + def test_new_does_not_raise(self): + # Regression: new() used datetime.datetime.now() and raised AttributeError. + t = Trace.new() + self.assertTrue(_valid_id(t.trace_id)) + self.assertTrue(_valid_id(t.request_id)) + + def test_new_generates_unique_ids(self): + seen = set() + for _ in range(1000): + seen.add(Trace.new().trace_id) + self.assertEqual(len(seen), 1000) + + +if __name__ == "__main__": + unittest.main() diff --git a/trace/trace.go b/trace/trace.go index 7dde0aa..b56ec93 100644 --- a/trace/trace.go +++ b/trace/trace.go @@ -2,7 +2,6 @@ package trace import ( "context" - "log/slog" "net/http" "time" @@ -102,15 +101,25 @@ func FromCtxOrNew(ctx context.Context) Trace { return t } +// Trace headers carried across service boundaries. Only X-Trace-ID and +// X-Request-ID are required for interop; the rest are metadata. +const ( + headerTraceID = "X-Trace-ID" + headerRequestID = "X-Request-ID" + headerTraceStart = "X-Trace-Start" + headerTraceSource = "X-Trace-Source" + headerRequestSource = "X-Request-Source" +) + // Save a Trace into the given header, over-writing the X-Trace-ID, X-Request-ID, and X-Trace-Start headers. // Note that there is no RequestStart header: the request timing starts when the server receives the request. // This is in contrast to the TraceStart header, which is the time the trace was created and persists across service boundaries. func SaveToHeader(h http.Header, t Trace) { - h.Set("X-Trace-ID", t.TraceID) - h.Set("X-Request-ID", t.RequestID) - h.Set("X-Trace-Start", t.TraceStart.Format(time.RFC3339)) - h.Set("X-Trace-Source", t.TraceSource) - h.Set("X-Request-Source", t.RequestSource) + h.Set(headerTraceID, t.TraceID) + h.Set(headerRequestID, t.RequestID) + h.Set(headerTraceStart, t.TraceStart.Format(time.RFC3339)) + h.Set(headerTraceSource, t.TraceSource) + h.Set(headerRequestSource, t.RequestSource) } // uuid generates a new UUID, preferring V7 over V4, but falling back to V4 if V7 is not available. @@ -123,35 +132,89 @@ func newuuid() string { } // FromHeaderOrNew returns a Trace from the given header, if it exists, and creates a new one if it doesn't. +// +// Inbound header values are untrusted. A TraceID/RequestID that is empty, over +// maxIDLen, or carries bytes outside [A-Za-z0-9._-] is replaced with a fresh +// uuid, and an out-of-charset TraceSource/RequestSource is dropped to "". This +// keeps every value SaveToHeader later re-emits safe to write into an HTTP +// header: a CR/LF-bearing id would otherwise make the outbound request (or the +// response echo) fail, or smuggle a header into a lenient downstream. func FromHeaderOrNew(h http.Header) Trace { now := time.Now().UTC() - var traceStart time.Time - var err error - if traceStart, err = time.Parse(time.RFC3339, h.Get("X-Trace-Start")); err != nil { - traceStart = now - } - - if traceStart.After(now) { - slog.Warn("trace start is in the future", slog.Time("trace_start", traceStart), slog.Time("now", now)) - traceStart = now + // X-Trace-Start is caller-controlled and untrusted. Parse it, but fall back to + // now for absent, malformed, or future values. A future timestamp is clamped + // silently, exactly like a malformed one, rather than logged: warning per + // request would let a caller drive this service's log volume by choosing the + // header value. + traceStart := now + if raw := h.Get(headerTraceStart); raw != "" { + if parsed, err := time.Parse(time.RFC3339, raw); err == nil && !parsed.After(now) { + traceStart = parsed + } } return Trace{ - TraceID: orelse(h.Get("X-Trace-ID"), newuuid), - RequestID: orelse(h.Get("X-Request-ID"), newuuid), + TraceID: resolveID(h.Get(headerTraceID)), + RequestID: resolveID(h.Get(headerRequestID)), TraceStart: traceStart, RequestStart: now, - TraceSource: h.Get("X-Trace-Source"), - RequestSource: h.Get("X-Request-Source"), + TraceSource: resolveSource(h.Get(headerTraceSource)), + RequestSource: resolveSource(h.Get(headerRequestSource)), + } +} + +// maxIDLen caps how long an inbound trace/request id may be before we treat it +// as untrustworthy and mint a fresh one. +const maxIDLen = 200 + +// maxSourceLen caps an inbound X-Trace-Source / X-Request-Source. A source is a +// service name (e.g. "runpod-graphql"), not an id, so it is bounded far tighter +// than maxIDLen: there is no legitimate reason to accept a 200-byte source, and +// the value is logged verbatim. +const maxSourceLen = 64 + +// validID reports whether s is safe to accept verbatim from an untrusted +// inbound header: non-empty, within maxIDLen, and limited to characters that +// cannot corrupt a log line or an HTTP header value (no CR/LF, control bytes, +// spaces, or non-ASCII). The scan indexes bytes and allocates nothing. +func validID(s string) bool { + if s == "" || len(s) > maxIDLen { + return false + } + for i := 0; i < len(s); i++ { + c := s[i] + switch { + case c >= 'a' && c <= 'z', c >= 'A' && c <= 'Z', c >= '0' && c <= '9': + case c == '-', c == '_', c == '.': + default: + return false + } + } + return true +} + +// resolveID returns id when it passes validID, otherwise a fresh uuid. +func resolveID(id string) string { + if validID(id) { + return id } + return newuuid() +} + +// validSource reports whether an inbound source header is safe to keep: empty +// (the unset / "unknown" case) or a charset-valid slug no longer than +// maxSourceLen. It is tighter than validID because a source is a service name, +// not an id. +func validSource(s string) bool { + return s == "" || (len(s) <= maxSourceLen && validID(s)) } -// return a if it's non-zero, otherwise call f and return its result. -func orelse[T comparable](a T, f func() T) T { - var zero T - if a == zero { - return f() +// resolveSource returns src unchanged when it passes validSource, and drops any +// other value to "" so SaveToHeader can never re-emit unsafe or oversized bytes. +func resolveSource(src string) string { + if validSource(src) { + return src } - return a + return "" } diff --git a/trace/trace_test.go b/trace/trace_test.go new file mode 100644 index 0000000..3c650fb --- /dev/null +++ b/trace/trace_test.go @@ -0,0 +1,414 @@ +package trace + +import ( + "context" + "net/http" + "strings" + "testing" + "time" +) + +// Shared fixtures, hoisted so the sample ids and source aren't repeated literals +// across the tables (keeps goconst quiet and the intent obvious). +const ( + sampleTraceID = "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31" + sampleRequestID = "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061" + sampleSource = "runpod-graphql" +) + +func TestValidID(t *testing.T) { + tests := []struct { + description string + input string + expected bool + }{ + {description: "canonical uuid v4", input: sampleTraceID, expected: true}, + {description: "uuid v7 style", input: sampleRequestID, expected: true}, + {description: "service-prefixed id with underscore", input: "trace_018f3a1c2b4d7e8f", expected: true}, + {description: "dotted id", input: "svc.abc.123", expected: true}, + {description: "all allowed classes", input: "AZaz09-_.", expected: true}, + {description: "exactly maxIDLen chars", input: strings.Repeat("a", maxIDLen), expected: true}, + {description: "empty string is rejected", input: "", expected: false}, + {description: "one over maxIDLen is rejected", input: strings.Repeat("a", maxIDLen+1), expected: false}, + {description: "embedded CRLF (header/log injection) is rejected", input: "abc\r\nX-Evil: 1", expected: false}, + {description: "bare newline is rejected", input: "abc\ndef", expected: false}, + {description: "space is rejected", input: "abc def", expected: false}, + {description: "control byte is rejected", input: "abc\x00def", expected: false}, + {description: "non-ascii is rejected", input: "abcé", expected: false}, + {description: "slash is rejected", input: "abc/def", expected: false}, + {description: "colon is rejected", input: "abc:def", expected: false}, + } + for _, tt := range tests { + t.Run(tt.description, func(t *testing.T) { + if got := validID(tt.input); got != tt.expected { + t.Errorf("validID(%q) = %v, want %v", tt.input, got, tt.expected) + } + }) + } +} + +func TestResolveID(t *testing.T) { + tests := []struct { + description string + input string + wantPropagate bool // true: returned verbatim; false: freshly generated + }{ + {description: "valid id is propagated verbatim", input: sampleTraceID, wantPropagate: true}, + {description: "absent id is generated", input: "", wantPropagate: false}, + {description: "oversized id is regenerated", input: strings.Repeat("a", maxIDLen+1), wantPropagate: false}, + {description: "CRLF id is regenerated", input: "abc\r\nInjected: 1", wantPropagate: false}, + {description: "junk id is regenerated", input: "abc def", wantPropagate: false}, + } + for _, tt := range tests { + t.Run(tt.description, func(t *testing.T) { + got := resolveID(tt.input) + if tt.wantPropagate { + if got != tt.input { + t.Errorf("resolveID(%q) = %q, want it propagated verbatim", tt.input, got) + } + return + } + if got == tt.input { + t.Errorf("resolveID(%q) returned the input; expected a fresh id", tt.input) + } + if !validID(got) { + t.Errorf("resolveID(%q) generated an invalid id %q", tt.input, got) + } + }) + } +} + +func TestResolveSource(t *testing.T) { + tests := []struct { + description string + input string + expected string + }{ + {description: "empty source stays empty", input: "", expected: ""}, + {description: "valid service name is kept", input: sampleSource, expected: sampleSource}, + {description: "source at maxSourceLen is kept", input: strings.Repeat("a", maxSourceLen), expected: strings.Repeat("a", maxSourceLen)}, + {description: "CRLF source is dropped to empty", input: "svc\r\nX-Evil: 1", expected: ""}, + {description: "spaced source is dropped to empty", input: "not a slug", expected: ""}, + {description: "over-length source is dropped to empty", input: strings.Repeat("a", maxSourceLen+1), expected: ""}, + {description: "charset-valid but over-length source is dropped", input: strings.Repeat("a", maxIDLen), expected: ""}, + } + for _, tt := range tests { + t.Run(tt.description, func(t *testing.T) { + if got := resolveSource(tt.input); got != tt.expected { + t.Errorf("resolveSource(%q) = %q, want %q", tt.input, got, tt.expected) + } + }) + } +} + +func TestFromHeaderOrNew(t *testing.T) { + const goodTrace = sampleTraceID + const goodReq = sampleRequestID + + tests := []struct { + description string + headers map[string]string + wantTraceID string // "" => expect a generated valid id + wantRequestID string // "" => expect a generated valid id + wantTraceSource string + wantRequestSource string + }{ + { + description: "all ids present and valid are propagated", + headers: map[string]string{headerTraceID: goodTrace, headerRequestID: goodReq, headerTraceSource: "main-ui", headerRequestSource: "hapi"}, + wantTraceID: goodTrace, + wantRequestID: goodReq, + wantTraceSource: "main-ui", + wantRequestSource: "hapi", + }, + { + description: "no headers generates both ids", + headers: map[string]string{}, + wantTraceID: "", + wantRequestID: "", + }, + { + description: "poisoned trace id is regenerated, valid request id kept", + headers: map[string]string{headerTraceID: "abc\r\nX-Evil: 1", headerRequestID: goodReq}, + wantTraceID: "", + wantRequestID: goodReq, + }, + { + description: "poisoned source is dropped to empty", + headers: map[string]string{headerTraceID: goodTrace, headerRequestID: goodReq, headerTraceSource: "svc\r\nX-Evil: 1"}, + wantTraceID: goodTrace, + wantRequestID: goodReq, + wantTraceSource: "", + }, + } + for _, tt := range tests { + t.Run(tt.description, func(t *testing.T) { + h := http.Header{} + for k, v := range tt.headers { + h.Set(k, v) + } + got := FromHeaderOrNew(h) + + assertID := func(name, want, actual string) { + if want == "" { + if actual == tt.headers[name] && tt.headers[name] != "" { + t.Errorf("%s = %q, expected a regenerated id", name, actual) + } + if !validID(actual) { + t.Errorf("%s = %q is not a valid generated id", name, actual) + } + return + } + if actual != want { + t.Errorf("%s = %q, want %q", name, actual, want) + } + } + assertID(headerTraceID, tt.wantTraceID, got.TraceID) + assertID(headerRequestID, tt.wantRequestID, got.RequestID) + if got.TraceSource != tt.wantTraceSource { + t.Errorf("TraceSource = %q, want %q", got.TraceSource, tt.wantTraceSource) + } + if got.RequestSource != tt.wantRequestSource { + t.Errorf("RequestSource = %q, want %q", got.RequestSource, tt.wantRequestSource) + } + }) + } +} + +func TestFromHeaderOrNewTraceStart(t *testing.T) { + now := time.Now().UTC() + tests := []struct { + description string + header string + expected func(got time.Time) bool + }{ + { + description: "absent trace-start defaults to ~now", + header: "", + expected: func(got time.Time) bool { + return !got.After(now.Add(time.Second)) && !got.Before(now.Add(-time.Second)) + }, + }, + { + description: "valid past trace-start is preserved", + header: "2020-01-01T00:00:00Z", + expected: func(got time.Time) bool { + want, _ := time.Parse(time.RFC3339, "2020-01-01T00:00:00Z") + return got.Equal(want) + }, + }, + { + description: "invalid trace-start falls back to ~now", + header: "not-a-timestamp", + expected: func(got time.Time) bool { return !got.Before(now.Add(-time.Second)) }, + }, + { + description: "future trace-start is clamped to ~now", + header: now.Add(time.Hour).Format(time.RFC3339), + expected: func(got time.Time) bool { return !got.After(now.Add(time.Second)) }, + }, + } + for _, tt := range tests { + t.Run(tt.description, func(t *testing.T) { + h := http.Header{} + if tt.header != "" { + h.Set(headerTraceStart, tt.header) + } + got := FromHeaderOrNew(h) + if !tt.expected(got.TraceStart) { + t.Errorf("TraceStart = %v not within expectation for %q", got.TraceStart, tt.header) + } + }) + } +} + +func TestSaveToHeaderRoundTrip(t *testing.T) { + orig := Trace{ + TraceID: sampleTraceID, + RequestID: sampleRequestID, + TraceSource: sampleSource, + RequestSource: sampleSource, + TraceStart: time.Now().UTC().Truncate(time.Second), + } + h := http.Header{} + SaveToHeader(h, orig) + + got := FromHeaderOrNew(h) + if got.TraceID != orig.TraceID || got.RequestID != orig.RequestID { + t.Errorf("round-trip ids: got trace=%q req=%q, want trace=%q req=%q", got.TraceID, got.RequestID, orig.TraceID, orig.RequestID) + } + if got.TraceSource != orig.TraceSource { + t.Errorf("round-trip TraceSource = %q, want %q", got.TraceSource, orig.TraceSource) + } + if !got.TraceStart.Equal(orig.TraceStart) { + t.Errorf("round-trip TraceStart = %v, want %v", got.TraceStart, orig.TraceStart) + } +} + +func TestSaveToHeaderEmitsNoUnsafeBytes(t *testing.T) { + // A Trace built from a poisoned inbound header must never re-emit CR/LF (or + // other control bytes) through SaveToHeader — otherwise the outbound request + // fails or smuggles a header downstream. + poisoned := http.Header{} + poisoned.Set(headerTraceID, "abc\r\nX-Evil: 1") + poisoned.Set(headerRequestID, "req\r\nX-Evil: 2") + poisoned.Set(headerTraceSource, "svc\r\nX-Evil: 3") + trc := FromHeaderOrNew(poisoned) + + out := http.Header{} + SaveToHeader(out, trc) + for key, values := range out { + for _, v := range values { + if strings.ContainsAny(v, "\r\n\x00") { + t.Errorf("header %q carries unsafe bytes after SaveToHeader: %q", key, v) + } + } + } +} + +type capturingRT struct{ req *http.Request } + +func (c *capturingRT) RoundTrip(r *http.Request) (*http.Response, error) { + c.req = r + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header)}, nil +} + +func TestClientMiddleware(t *testing.T) { + t.Run("reuses trace id and mints a new request id when ctx has a trace", func(t *testing.T) { + parent := Trace{TraceID: "trace-abc", RequestID: "req-parent", TraceSource: "svc"} + cap := &capturingRT{} + rt := ClientMiddleware(cap) + + req, _ := http.NewRequest(http.MethodGet, "http://example.invalid/", nil) + req = req.WithContext(CtxWith(req.Context(), parent)) + resp, err := rt.RoundTrip(req) + if err != nil { + t.Fatalf("RoundTrip: %v", err) + } + resp.Body.Close() + if got := cap.req.Header.Get(headerTraceID); got != "trace-abc" { + t.Errorf("X-Trace-ID = %q, want the parent trace id", got) + } + if got := cap.req.Header.Get(headerRequestID); got == "" || got == "req-parent" { + t.Errorf("X-Request-ID = %q, want a fresh sub-request id", got) + } + }) + + t.Run("creates a fresh trace when ctx has none", func(t *testing.T) { + cap := &capturingRT{} + rt := ClientMiddleware(cap) + req, _ := http.NewRequest(http.MethodGet, "http://example.invalid/", nil) + resp, err := rt.RoundTrip(req) + if err != nil { + t.Fatalf("RoundTrip: %v", err) + } + resp.Body.Close() + if got := cap.req.Header.Get(headerTraceID); !validID(got) { + t.Errorf("X-Trace-ID = %q, want a generated valid id", got) + } + }) +} + +func TestServerMiddlewarePutsTraceInContext(t *testing.T) { + var seen Trace + var ok bool + h := ServerMiddleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + seen, ok = FromCtx(r.Context()) + })) + req, _ := http.NewRequest(http.MethodGet, "http://example.invalid/", nil) + req.Header.Set(headerTraceID, "trace-from-header") + h.ServeHTTP(nil, req) + if !ok { + t.Fatal("expected a Trace in the request context") + } + if seen.TraceID != "trace-from-header" { + t.Errorf("TraceID = %q, want %q", seen.TraceID, "trace-from-header") + } +} + +func TestNewGeneratesUniqueIDs(t *testing.T) { + const n = 10_000 + traces := make(map[string]struct{}, n) + reqs := make(map[string]struct{}, n) + for i := 0; i < n; i++ { + tr := New() + if _, dup := traces[tr.TraceID]; dup { + t.Fatalf("duplicate TraceID at i=%d: %q", i, tr.TraceID) + } + if _, dup := reqs[tr.RequestID]; dup { + t.Fatalf("duplicate RequestID at i=%d: %q", i, tr.RequestID) + } + traces[tr.TraceID] = struct{}{} + reqs[tr.RequestID] = struct{}{} + } +} + +func TestValidIDIsAllocationFree(t *testing.T) { + input := sampleTraceID + if allocs := testing.AllocsPerRun(1000, func() { _ = validID(input) }); allocs != 0 { + t.Errorf("validID allocated %v times per run, want 0", allocs) + } +} + +func BenchmarkFromHeaderOrNew(b *testing.B) { + valid := http.Header{} + valid.Set(headerTraceID, sampleTraceID) + valid.Set(headerRequestID, sampleRequestID) + + invalid := http.Header{} + invalid.Set(headerTraceID, "abc\r\nX-Evil: 1") + + absent := http.Header{} + + cases := map[string]http.Header{"present": valid, "absent": absent, "invalid": invalid} + for name, h := range cases { + b.Run(name, func(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + _ = FromHeaderOrNew(h) + } + }) + } +} + +func BenchmarkValidID(b *testing.B) { + input := sampleTraceID + b.ReportAllocs() + for i := 0; i < b.N; i++ { + _ = validID(input) + } +} + +func BenchmarkNew(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + _ = New() + } +} + +func BenchmarkNewUUID(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + _ = newuuid() + } +} + +func BenchmarkSaveToHeader(b *testing.B) { + trc := New() + h := http.Header{} + b.ReportAllocs() + for i := 0; i < b.N; i++ { + SaveToHeader(h, trc) + } +} + +func BenchmarkClientMiddlewareRoundTrip(b *testing.B) { + rt := ClientMiddleware(&capturingRT{}) + req, _ := http.NewRequest(http.MethodGet, "http://example.invalid/", nil) + req = req.WithContext(CtxWith(context.Background(), New())) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + resp, _ := rt.RoundTrip(req) + resp.Body.Close() + } +}