From d6e87d01f4b6f9df2d2081aa2a88a163d5e768f8 Mon Sep 17 00:00:00 2001 From: "randomizedcoder dave.seddon.ca@gmail.com" Date: Mon, 21 Sep 2026 17:05:07 -0700 Subject: [PATCH 1/5] feat: add inbound header validation and tests to trace package Validate inbound X-Trace-ID/X-Request-ID (length + charset, regenerate on failure) to prevent header/log injection; fix Python Trace.new() and save_to_headers tuple bug and bring it to five-header parity with Go; add table-driven Go + Python tests and benchmarks. Part of the end-to-end X-Trace-ID tracing effort. Co-Authored-By: Claude Opus 4.8 --- README.md | 21 +++ docs/x-trace-id-plan.md | 73 ++++++++ py/trace.py | 52 ++++-- py/trace_test.py | 147 +++++++++++++++ trace/trace.go | 69 +++++-- trace/trace_test.go | 398 ++++++++++++++++++++++++++++++++++++++++ 6 files changed, 735 insertions(+), 25 deletions(-) create mode 100644 docs/x-trace-id-plan.md create mode 100644 py/trace_test.py create mode 100644 trace/trace_test.go diff --git a/README.md b/README.md index 38578c2..1a20c58 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` | fresh id per request/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**: reuse a valid inbound `X-Trace-ID`, otherwise mint one. One `trace_id` spans the whole trace; a fresh `request_id` is minted per sub-request. + +### 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/docs/x-trace-id-plan.md b/docs/x-trace-id-plan.md new file mode 100644 index 0000000..72e3928 --- /dev/null +++ b/docs/x-trace-id-plan.md @@ -0,0 +1,73 @@ +# rplog — X-Trace-ID hardening plan (trace package) + +## Context + +`rplog/trace` is the shared library that defines RunPod's cross-service trace contract +(`X-Trace-ID` / `X-Request-ID` + `X-Trace-Start` / `X-Trace-Source` / `X-Request-Source`). +The `host` daemon already depends on it; `hapi` and `proxy` are adopting it, and the RunPod +GraphQL API (TypeScript) is being taught to emit and propagate the same headers. As part of that +end-to-end tracing work, the library needed: + +1. **Inbound-header validation.** `FromHeaderOrNew` trusted the raw `X-Trace-ID` / `X-Request-ID` + verbatim. A CR/LF-bearing value would flow into a `Trace`, then `SaveToHeader` would re-emit it + on an outbound request or the response echo — failing the request (`net/http` rejects invalid + header values) or smuggling a header into a lenient downstream. +2. **A defect the Python port carried** (`py/trace.py`): `Trace.new()` raised `AttributeError` + (`datetime.datetime.now()` on a directly-imported `datetime`), and `save_to_headers` stored a + **tuple** for `X-Request-ID` (trailing comma) and omitted `X-Trace-Source` / `X-Trace-Start`. +3. **No tests** for the Go `trace` package. + +Scope: Go + Python. A JS `trace` module is a separate follow-up (`js/` has none today). + +## Changes (done) + +### Go — `trace/trace.go` +- Added `validID` (allocation-free byte scan: non-empty, ≤ `maxIDLen` = 200, charset + `[A-Za-z0-9._-]`), `resolveID` (regenerate on invalid), `resolveSource` (drop invalid to `""`). +- `FromHeaderOrNew` now resolves the two IDs and both sources through those helpers, so a `Trace` + can never hold bytes that `SaveToHeader` would unsafely re-emit. +- Guarded the `X-Trace-Start` parse: the common no-header path skips `time.Parse` entirely + (previously it parsed `""` on every request and fell back). +- Removed the now-unused `orelse` generic. + +### Python — `py/trace.py` +- Fixed `Trace.new()` (`datetime.now()`), fixed the `save_to_headers` tuple bug, and made it emit + the same five headers as Go with a freshly minted `X-Request-ID`. +- Added `_valid_id` / `_resolve_id` / `_resolve_source` mirroring the Go validation. + +### Tests +- `trace/trace_test.go` — table-driven (`{description, input, expected}`) covering `validID`, + `resolveID`, `resolveSource`, `FromHeaderOrNew` (present / absent / poisoned / source-drop), + trace-start (absent / past / invalid / future-clamp), `SaveToHeader` round-trip + no-unsafe-bytes, + `ClientMiddleware` (reuse trace id, mint new request id) and `ServerMiddleware`, `New()` + uniqueness (10k), and `AllocsPerRun` proving `validID` is 0-alloc. Plus benchmarks. +- `py/trace_test.py` — `unittest` table-driven mirror, incl. regression tests for both fixed bugs. + +## Benchmark results (baseline; `go test -bench=. -benchmem`) + +``` +BenchmarkFromHeaderOrNew/present 1057 ns/op 32 B/op 2 allocs/op +BenchmarkFromHeaderOrNew/absent 1899 ns/op 160 B/op 6 allocs/op +BenchmarkFromHeaderOrNew/invalid 2220 ns/op 160 B/op 6 allocs/op +BenchmarkValidID 128 ns/op 0 B/op 0 allocs/op +BenchmarkNew 1359 ns/op 128 B/op 4 allocs/op +BenchmarkNewUUID 492 ns/op 64 B/op 2 allocs/op +BenchmarkSaveToHeader 1694 ns/op 136 B/op 8 allocs/op +BenchmarkClientMiddlewareRoundTrip 3982 ns/op 872 B/op 15 allocs/op +``` + +Review: `validID` meets the 0-alloc requirement. The guarded `time.Parse` keeps the hot inbound +"present" path at 2 allocs / ~1µs; the 6-alloc paths are the uuid-generation cases (unavoidable). +All costs are negligible against real network I/O — no further tuning is justified. + +## Remaining / follow-ups + +- Tag a new module version (e.g. `v0.1.2`) and bump `require github.com/runpod/rplog` in + `host` / `hapi` / `proxy` to pick up validation. Backward compatible — the only visible change is + that a hostile inbound id is sanitized to a fresh uuid. +- Add a JS `trace` module (`js/trace.ts`) so TS/JS services can converge on the shared library + (RunPod currently uses its own monorepo `@runpod/rplog` helper). +- ~~Update the top-level README to document the validation and response-echo usage.~~ Done: + the README `## Tracing` section now carries the HTTP header contract table, the Go + inbound/outbound/response-echo usage, and the inbound-validation behavior. (GO_README stays + the minimal slog-logger quickstart; the header contract lives in README.) diff --git a/py/trace.py b/py/trace.py index 9170d23..5464b1a 100644 --- a/py/trace.py +++ b/py/trace.py @@ -1,3 +1,4 @@ +import re from dataclasses import dataclass from datetime import datetime, timezone from typing import Optional @@ -9,6 +10,27 @@ 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") + + +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 _resolve_source(s: str) -> str: + return s if _valid_id(s) else "unknown" + + @dataclass class Trace: """A trace object that can be used to track a request through multiple services. @@ -25,17 +47,21 @@ 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, and a malformed source is + dropped to "unknown" (see _valid_id). """ global _trace now = as_rfc3339(datetime.now()) 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=headers.get("X-Request-Start") or now, + trace_id=_resolve_id(headers.get("X-Trace-ID", "")), + trace_source=_resolve_source(headers.get("X-Trace-Source", "")), + trace_start=headers.get("X-Trace-Start") or now, ) _trace = t return t @@ -56,7 +82,7 @@ def current(cls) -> "Trace": @staticmethod def new(): """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()) global _trace t = Trace( request_id=uuid7(), @@ -71,12 +97,16 @@ 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. """ - 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.""" diff --git a/py/trace_test.py b/py/trace_test.py new file mode 100644 index 0000000..ae1c44b --- /dev/null +++ b/py/trace_test.py @@ -0,0 +1,147 @@ +import unittest + +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 becomes unknown", "input": "", "expected": "unknown"}, + {"description": "valid source kept", "input": "runpod-graphql", "expected": "runpod-graphql"}, + {"description": "poisoned source becomes unknown", "input": "svc\r\nX-Evil: 1", "expected": "unknown"}, + {"description": "spaced source becomes unknown", "input": "not a slug", "expected": "unknown"}, + ] + 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": "unknown", + }, + { + "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": "unknown", + }, + { + "description": "poisoned source dropped to unknown", + "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": "unknown", + }, + ] + 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): + t = Trace.from_headers({"X-Trace-ID": "abc\r\nX-Evil: 1", "X-Trace-Source": "svc\r\nX-Evil: 2"}) + 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") + + +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..235a682 100644 --- a/trace/trace.go +++ b/trace/trace.go @@ -123,13 +123,21 @@ 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 + traceStart := now + if raw := h.Get("X-Trace-Start"); raw != "" { + if parsed, err := time.Parse(time.RFC3339, raw); err == nil { + traceStart = parsed + } } if traceStart.After(now) { @@ -138,20 +146,53 @@ func FromHeaderOrNew(h http.Header) Trace { } return Trace{ - TraceID: orelse(h.Get("X-Trace-ID"), newuuid), - RequestID: orelse(h.Get("X-Request-ID"), newuuid), + TraceID: resolveID(h.Get("X-Trace-ID")), + RequestID: resolveID(h.Get("X-Request-ID")), TraceStart: traceStart, RequestStart: now, - TraceSource: h.Get("X-Trace-Source"), - RequestSource: h.Get("X-Request-Source"), + TraceSource: resolveSource(h.Get("X-Trace-Source")), + RequestSource: resolveSource(h.Get("X-Request-Source")), + } +} + +// 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 + +// 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() } -// 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 empty (the "unknown" case) or when +// it passes validID, and drops any other value to "" so SaveToHeader can never +// re-emit unsafe bytes. +func resolveSource(src string) string { + if src == "" || validID(src) { + return src } - return a + return "" } diff --git a/trace/trace_test.go b/trace/trace_test.go new file mode 100644 index 0000000..3e76f05 --- /dev/null +++ b/trace/trace_test.go @@ -0,0 +1,398 @@ +package trace + +import ( + "context" + "net/http" + "strings" + "testing" + "time" +) + +func TestValidID(t *testing.T) { + tests := []struct { + description string + input string + expected bool + }{ + {description: "canonical uuid v4", input: "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", expected: true}, + {description: "uuid v7 style", input: "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061", 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: "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", 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 is kept as unknown", input: "", expected: ""}, + {description: "valid service name is kept", input: "runpod-graphql", expected: "runpod-graphql"}, + {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: ""}, + } + 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 = "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31" + const goodReq = "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061" + + 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{"X-Trace-ID": goodTrace, "X-Request-ID": goodReq, "X-Trace-Source": "main-ui", "X-Request-Source": "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{"X-Trace-ID": "abc\r\nX-Evil: 1", "X-Request-ID": goodReq}, + wantTraceID: "", + wantRequestID: goodReq, + }, + { + description: "poisoned source is dropped to empty", + headers: map[string]string{"X-Trace-ID": goodTrace, "X-Request-ID": goodReq, "X-Trace-Source": "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("X-Trace-ID", tt.wantTraceID, got.TraceID) + assertID("X-Request-ID", 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("X-Trace-Start", 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: "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", + RequestID: "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061", + TraceSource: "runpod-graphql", + RequestSource: "runpod-graphql", + 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("X-Trace-ID", "abc\r\nX-Evil: 1") + poisoned.Set("X-Request-ID", "req\r\nX-Evil: 2") + poisoned.Set("X-Trace-Source", "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)) + if _, err := rt.RoundTrip(req); err != nil { + t.Fatalf("RoundTrip: %v", err) + } + if got := cap.req.Header.Get("X-Trace-ID"); got != "trace-abc" { + t.Errorf("X-Trace-ID = %q, want the parent trace id", got) + } + if got := cap.req.Header.Get("X-Request-ID"); 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) + if _, err := rt.RoundTrip(req); err != nil { + t.Fatalf("RoundTrip: %v", err) + } + if got := cap.req.Header.Get("X-Trace-ID"); !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("X-Trace-ID", "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 := "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31" + 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("X-Trace-ID", "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31") + valid.Set("X-Request-ID", "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061") + + invalid := http.Header{} + invalid.Set("X-Trace-ID", "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 := "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31" + 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++ { + _, _ = rt.RoundTrip(req) + } +} From d2accaeaa9970ec051f8bbc7088f1d998ec0c39f Mon Sep 17 00:00:00 2001 From: "randomizedcoder dave.seddon.ca@gmail.com" Date: Mon, 21 Sep 2026 17:56:03 -0700 Subject: [PATCH 2/5] =?UTF-8?q?fix:=20address=20self-review=20=E2=80=94=20?= =?UTF-8?q?validate=20Python=20X-Trace-Start,=20parity,=20and=20lint?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - py: parse/validate inbound X-Trace-Start (RFC3339, fall back to now) and drop the non-contract X-Request-Start read, closing a CWE-93 header-injection vector where a poisoned timestamp was re-emitted verbatim by save_to_headers - py: _resolve_source drops empty/invalid to "" to match Go's resolveSource; datetime.now(timezone.utc) for correct UTC; new() return annotation; Optional -> | None - py tests: add poisoned/invalid X-Trace-Start rows + valid-preserved row - go: extract X-Trace-* header names to consts (SaveToHeader/FromHeaderOrNew); close response bodies in trace tests; reuse fixture consts; fix an empty-source test description - docs: clarify request_id is minted fresh per outbound sub-request; the inbound edge honors a valid caller-supplied X-Request-ID Co-Authored-By: Claude Opus 4.8 --- README.md | 4 +-- py/trace.py | 42 +++++++++++++++++++------ py/trace_test.py | 60 ++++++++++++++++++++++++++--------- trace/trace.go | 30 ++++++++++++------ trace/trace_test.go | 77 ++++++++++++++++++++++++++------------------- 5 files changed, 145 insertions(+), 68 deletions(-) diff --git a/README.md b/README.md index 1a20c58..adf5eca 100644 --- a/README.md +++ b/README.md @@ -78,12 +78,12 @@ 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` | fresh id per request/sub-request | 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**: reuse a valid inbound `X-Trace-ID`, otherwise mint one. One `trace_id` spans the whole trace; a fresh `request_id` is minted per sub-request. +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. diff --git a/py/trace.py b/py/trace.py index 5464b1a..fd49a33 100644 --- a/py/trace.py +++ b/py/trace.py @@ -1,7 +1,6 @@ import re from dataclasses import dataclass from datetime import datetime, timezone -from typing import Optional from uuid7 import uuid7 @@ -28,7 +27,23 @@ def _resolve_id(s: str) -> str: def _resolve_source(s: str) -> str: - return s if _valid_id(s) else "unknown" + # Match Go's resolveSource: keep a valid source, drop anything else (including + # empty) to "" so save_to_headers can never re-emit unsafe bytes. + return s if _valid_id(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 or malformed. Re-formatting through + # as_rfc3339 means the stored value can never carry CR/LF (or any other + # injected bytes) into save_to_headers — the emit-side guard that mirrors Go's + # time.Parse of X-Trace-Start. + if s: + try: + return as_rfc3339(datetime.fromisoformat(s.replace("Z", "+00:00"))) + except ValueError: + pass + return as_rfc3339(now) @dataclass @@ -49,19 +64,22 @@ def from_headers(headers: dict[str, str]) -> "Trace": 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, and a malformed source is - dropped to "unknown" (see _valid_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=_resolve_id(headers.get("X-Request-ID", "")), request_source=_resolve_source(headers.get("X-Request-Source", "")), - request_start=headers.get("X-Request-Start") or now, + request_start=now, trace_id=_resolve_id(headers.get("X-Trace-ID", "")), trace_source=_resolve_source(headers.get("X-Trace-Source", "")), - trace_start=headers.get("X-Trace-Start") or now, + trace_start=_resolve_start(headers.get("X-Trace-Start", ""), now_dt), ) _trace = t return t @@ -80,9 +98,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.now()) + now = as_rfc3339(datetime.now(timezone.utc)) global _trace t = Trace( request_id=uuid7(), @@ -101,6 +119,10 @@ def save_to_headers(self, headers: dict[str, str]) -> None: 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-Trace-ID"] = self.trace_id headers["X-Request-ID"] = uuid7() @@ -110,4 +132,4 @@ def save_to_headers(self, headers: dict[str, str]) -> None: """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 index ae1c44b..b8e99dd 100644 --- a/py/trace_test.py +++ b/py/trace_test.py @@ -1,5 +1,5 @@ import unittest - +from datetime import datetime, timezone from trace import Trace, _resolve_id, _resolve_source, _valid_id @@ -45,10 +45,10 @@ def test_cases(self): class TestResolveSource(unittest.TestCase): def test_cases(self): cases = [ - {"description": "empty source becomes unknown", "input": "", "expected": "unknown"}, + {"description": "empty source stays empty", "input": "", "expected": ""}, {"description": "valid source kept", "input": "runpod-graphql", "expected": "runpod-graphql"}, - {"description": "poisoned source becomes unknown", "input": "svc\r\nX-Evil: 1", "expected": "unknown"}, - {"description": "spaced source becomes unknown", "input": "not a slug", "expected": "unknown"}, + {"description": "poisoned source dropped to empty", "input": "svc\r\nX-Evil: 1", "expected": ""}, + {"description": "spaced source dropped to empty", "input": "not a slug", "expected": ""}, ] for c in cases: with self.subTest(c["description"]): @@ -72,21 +72,21 @@ def test_cases(self): "headers": {}, "want_trace": None, "want_request": None, - "want_trace_source": "unknown", + "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": "unknown", + "want_trace_source": "", }, { - "description": "poisoned source dropped to unknown", + "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": "unknown", + "want_trace_source": "", }, ] for c in cases: @@ -121,12 +121,44 @@ def test_writes_go_five_headers_and_fresh_request_id(self): self.assertNotEqual(headers["X-Request-ID"], original_request_id) def test_emits_no_unsafe_bytes_from_poisoned_trace(self): - t = Trace.from_headers({"X-Trace-ID": "abc\r\nX-Evil: 1", "X-Trace-Source": "svc\r\nX-Evil: 2"}) - 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") + # 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)) class TestNewRegression(unittest.TestCase): diff --git a/trace/trace.go b/trace/trace.go index 235a682..ed549eb 100644 --- a/trace/trace.go +++ b/trace/trace.go @@ -102,15 +102,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. @@ -134,7 +144,7 @@ func FromHeaderOrNew(h http.Header) Trace { now := time.Now().UTC() traceStart := now - if raw := h.Get("X-Trace-Start"); raw != "" { + if raw := h.Get(headerTraceStart); raw != "" { if parsed, err := time.Parse(time.RFC3339, raw); err == nil { traceStart = parsed } @@ -146,12 +156,12 @@ func FromHeaderOrNew(h http.Header) Trace { } return Trace{ - TraceID: resolveID(h.Get("X-Trace-ID")), - RequestID: resolveID(h.Get("X-Request-ID")), + TraceID: resolveID(h.Get(headerTraceID)), + RequestID: resolveID(h.Get(headerRequestID)), TraceStart: traceStart, RequestStart: now, - TraceSource: resolveSource(h.Get("X-Trace-Source")), - RequestSource: resolveSource(h.Get("X-Request-Source")), + TraceSource: resolveSource(h.Get(headerTraceSource)), + RequestSource: resolveSource(h.Get(headerRequestSource)), } } diff --git a/trace/trace_test.go b/trace/trace_test.go index 3e76f05..5c7788c 100644 --- a/trace/trace_test.go +++ b/trace/trace_test.go @@ -8,14 +8,22 @@ import ( "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: "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", expected: true}, - {description: "uuid v7 style", input: "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061", expected: true}, + {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}, @@ -45,7 +53,7 @@ func TestResolveID(t *testing.T) { input string wantPropagate bool // true: returned verbatim; false: freshly generated }{ - {description: "valid id is propagated verbatim", input: "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", wantPropagate: true}, + {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}, @@ -76,8 +84,8 @@ func TestResolveSource(t *testing.T) { input string expected string }{ - {description: "empty source is kept as unknown", input: "", expected: ""}, - {description: "valid service name is kept", input: "runpod-graphql", expected: "runpod-graphql"}, + {description: "empty source stays empty", input: "", expected: ""}, + {description: "valid service name is kept", input: sampleSource, expected: sampleSource}, {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: ""}, } @@ -91,8 +99,8 @@ func TestResolveSource(t *testing.T) { } func TestFromHeaderOrNew(t *testing.T) { - const goodTrace = "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31" - const goodReq = "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061" + const goodTrace = sampleTraceID + const goodReq = sampleRequestID tests := []struct { description string @@ -104,7 +112,7 @@ func TestFromHeaderOrNew(t *testing.T) { }{ { description: "all ids present and valid are propagated", - headers: map[string]string{"X-Trace-ID": goodTrace, "X-Request-ID": goodReq, "X-Trace-Source": "main-ui", "X-Request-Source": "hapi"}, + headers: map[string]string{headerTraceID: goodTrace, headerRequestID: goodReq, headerTraceSource: "main-ui", headerRequestSource: "hapi"}, wantTraceID: goodTrace, wantRequestID: goodReq, wantTraceSource: "main-ui", @@ -118,13 +126,13 @@ func TestFromHeaderOrNew(t *testing.T) { }, { description: "poisoned trace id is regenerated, valid request id kept", - headers: map[string]string{"X-Trace-ID": "abc\r\nX-Evil: 1", "X-Request-ID": goodReq}, + 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{"X-Trace-ID": goodTrace, "X-Request-ID": goodReq, "X-Trace-Source": "svc\r\nX-Evil: 1"}, + headers: map[string]string{headerTraceID: goodTrace, headerRequestID: goodReq, headerTraceSource: "svc\r\nX-Evil: 1"}, wantTraceID: goodTrace, wantRequestID: goodReq, wantTraceSource: "", @@ -152,8 +160,8 @@ func TestFromHeaderOrNew(t *testing.T) { t.Errorf("%s = %q, want %q", name, actual, want) } } - assertID("X-Trace-ID", tt.wantTraceID, got.TraceID) - assertID("X-Request-ID", tt.wantRequestID, got.RequestID) + 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) } @@ -201,7 +209,7 @@ func TestFromHeaderOrNewTraceStart(t *testing.T) { t.Run(tt.description, func(t *testing.T) { h := http.Header{} if tt.header != "" { - h.Set("X-Trace-Start", tt.header) + h.Set(headerTraceStart, tt.header) } got := FromHeaderOrNew(h) if !tt.expected(got.TraceStart) { @@ -213,10 +221,10 @@ func TestFromHeaderOrNewTraceStart(t *testing.T) { func TestSaveToHeaderRoundTrip(t *testing.T) { orig := Trace{ - TraceID: "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31", - RequestID: "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061", - TraceSource: "runpod-graphql", - RequestSource: "runpod-graphql", + TraceID: sampleTraceID, + RequestID: sampleRequestID, + TraceSource: sampleSource, + RequestSource: sampleSource, TraceStart: time.Now().UTC().Truncate(time.Second), } h := http.Header{} @@ -239,9 +247,9 @@ func TestSaveToHeaderEmitsNoUnsafeBytes(t *testing.T) { // other control bytes) through SaveToHeader — otherwise the outbound request // fails or smuggles a header downstream. poisoned := http.Header{} - poisoned.Set("X-Trace-ID", "abc\r\nX-Evil: 1") - poisoned.Set("X-Request-ID", "req\r\nX-Evil: 2") - poisoned.Set("X-Trace-Source", "svc\r\nX-Evil: 3") + 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{} @@ -270,13 +278,15 @@ func TestClientMiddleware(t *testing.T) { req, _ := http.NewRequest(http.MethodGet, "http://example.invalid/", nil) req = req.WithContext(CtxWith(req.Context(), parent)) - if _, err := rt.RoundTrip(req); err != nil { + resp, err := rt.RoundTrip(req) + if err != nil { t.Fatalf("RoundTrip: %v", err) } - if got := cap.req.Header.Get("X-Trace-ID"); got != "trace-abc" { + 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("X-Request-ID"); got == "" || got == "req-parent" { + if got := cap.req.Header.Get(headerRequestID); got == "" || got == "req-parent" { t.Errorf("X-Request-ID = %q, want a fresh sub-request id", got) } }) @@ -285,10 +295,12 @@ func TestClientMiddleware(t *testing.T) { cap := &capturingRT{} rt := ClientMiddleware(cap) req, _ := http.NewRequest(http.MethodGet, "http://example.invalid/", nil) - if _, err := rt.RoundTrip(req); err != nil { + resp, err := rt.RoundTrip(req) + if err != nil { t.Fatalf("RoundTrip: %v", err) } - if got := cap.req.Header.Get("X-Trace-ID"); !validID(got) { + resp.Body.Close() + if got := cap.req.Header.Get(headerTraceID); !validID(got) { t.Errorf("X-Trace-ID = %q, want a generated valid id", got) } }) @@ -301,7 +313,7 @@ func TestServerMiddlewarePutsTraceInContext(t *testing.T) { seen, ok = FromCtx(r.Context()) })) req, _ := http.NewRequest(http.MethodGet, "http://example.invalid/", nil) - req.Header.Set("X-Trace-ID", "trace-from-header") + req.Header.Set(headerTraceID, "trace-from-header") h.ServeHTTP(nil, req) if !ok { t.Fatal("expected a Trace in the request context") @@ -329,7 +341,7 @@ func TestNewGeneratesUniqueIDs(t *testing.T) { } func TestValidIDIsAllocationFree(t *testing.T) { - input := "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31" + input := sampleTraceID if allocs := testing.AllocsPerRun(1000, func() { _ = validID(input) }); allocs != 0 { t.Errorf("validID allocated %v times per run, want 0", allocs) } @@ -337,11 +349,11 @@ func TestValidIDIsAllocationFree(t *testing.T) { func BenchmarkFromHeaderOrNew(b *testing.B) { valid := http.Header{} - valid.Set("X-Trace-ID", "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31") - valid.Set("X-Request-ID", "018f3a1c-2b4d-7e8f-9a0b-1c2d3e4f5061") + valid.Set(headerTraceID, sampleTraceID) + valid.Set(headerRequestID, sampleRequestID) invalid := http.Header{} - invalid.Set("X-Trace-ID", "abc\r\nX-Evil: 1") + invalid.Set(headerTraceID, "abc\r\nX-Evil: 1") absent := http.Header{} @@ -357,7 +369,7 @@ func BenchmarkFromHeaderOrNew(b *testing.B) { } func BenchmarkValidID(b *testing.B) { - input := "5f9c2e6a-1b3d-4c8e-9a0f-2b7c6d5e4f31" + input := sampleTraceID b.ReportAllocs() for i := 0; i < b.N; i++ { _ = validID(input) @@ -393,6 +405,7 @@ func BenchmarkClientMiddlewareRoundTrip(b *testing.B) { req = req.WithContext(CtxWith(context.Background(), New())) b.ReportAllocs() for i := 0; i < b.N; i++ { - _, _ = rt.RoundTrip(req) + resp, _ := rt.RoundTrip(req) + resp.Body.Close() } } From c96445133529fbe4cd70afbd1d537de69b420a00 Mon Sep 17 00:00:00 2001 From: "randomizedcoder dave.seddon.ca@gmail.com" Date: Mon, 21 Sep 2026 18:33:07 -0700 Subject: [PATCH 3/5] chore: drop internal planning doc from PR (planning-only, not shipped) Co-Authored-By: Claude Opus 4.8 --- docs/x-trace-id-plan.md | 73 ----------------------------------------- 1 file changed, 73 deletions(-) delete mode 100644 docs/x-trace-id-plan.md diff --git a/docs/x-trace-id-plan.md b/docs/x-trace-id-plan.md deleted file mode 100644 index 72e3928..0000000 --- a/docs/x-trace-id-plan.md +++ /dev/null @@ -1,73 +0,0 @@ -# rplog — X-Trace-ID hardening plan (trace package) - -## Context - -`rplog/trace` is the shared library that defines RunPod's cross-service trace contract -(`X-Trace-ID` / `X-Request-ID` + `X-Trace-Start` / `X-Trace-Source` / `X-Request-Source`). -The `host` daemon already depends on it; `hapi` and `proxy` are adopting it, and the RunPod -GraphQL API (TypeScript) is being taught to emit and propagate the same headers. As part of that -end-to-end tracing work, the library needed: - -1. **Inbound-header validation.** `FromHeaderOrNew` trusted the raw `X-Trace-ID` / `X-Request-ID` - verbatim. A CR/LF-bearing value would flow into a `Trace`, then `SaveToHeader` would re-emit it - on an outbound request or the response echo — failing the request (`net/http` rejects invalid - header values) or smuggling a header into a lenient downstream. -2. **A defect the Python port carried** (`py/trace.py`): `Trace.new()` raised `AttributeError` - (`datetime.datetime.now()` on a directly-imported `datetime`), and `save_to_headers` stored a - **tuple** for `X-Request-ID` (trailing comma) and omitted `X-Trace-Source` / `X-Trace-Start`. -3. **No tests** for the Go `trace` package. - -Scope: Go + Python. A JS `trace` module is a separate follow-up (`js/` has none today). - -## Changes (done) - -### Go — `trace/trace.go` -- Added `validID` (allocation-free byte scan: non-empty, ≤ `maxIDLen` = 200, charset - `[A-Za-z0-9._-]`), `resolveID` (regenerate on invalid), `resolveSource` (drop invalid to `""`). -- `FromHeaderOrNew` now resolves the two IDs and both sources through those helpers, so a `Trace` - can never hold bytes that `SaveToHeader` would unsafely re-emit. -- Guarded the `X-Trace-Start` parse: the common no-header path skips `time.Parse` entirely - (previously it parsed `""` on every request and fell back). -- Removed the now-unused `orelse` generic. - -### Python — `py/trace.py` -- Fixed `Trace.new()` (`datetime.now()`), fixed the `save_to_headers` tuple bug, and made it emit - the same five headers as Go with a freshly minted `X-Request-ID`. -- Added `_valid_id` / `_resolve_id` / `_resolve_source` mirroring the Go validation. - -### Tests -- `trace/trace_test.go` — table-driven (`{description, input, expected}`) covering `validID`, - `resolveID`, `resolveSource`, `FromHeaderOrNew` (present / absent / poisoned / source-drop), - trace-start (absent / past / invalid / future-clamp), `SaveToHeader` round-trip + no-unsafe-bytes, - `ClientMiddleware` (reuse trace id, mint new request id) and `ServerMiddleware`, `New()` - uniqueness (10k), and `AllocsPerRun` proving `validID` is 0-alloc. Plus benchmarks. -- `py/trace_test.py` — `unittest` table-driven mirror, incl. regression tests for both fixed bugs. - -## Benchmark results (baseline; `go test -bench=. -benchmem`) - -``` -BenchmarkFromHeaderOrNew/present 1057 ns/op 32 B/op 2 allocs/op -BenchmarkFromHeaderOrNew/absent 1899 ns/op 160 B/op 6 allocs/op -BenchmarkFromHeaderOrNew/invalid 2220 ns/op 160 B/op 6 allocs/op -BenchmarkValidID 128 ns/op 0 B/op 0 allocs/op -BenchmarkNew 1359 ns/op 128 B/op 4 allocs/op -BenchmarkNewUUID 492 ns/op 64 B/op 2 allocs/op -BenchmarkSaveToHeader 1694 ns/op 136 B/op 8 allocs/op -BenchmarkClientMiddlewareRoundTrip 3982 ns/op 872 B/op 15 allocs/op -``` - -Review: `validID` meets the 0-alloc requirement. The guarded `time.Parse` keeps the hot inbound -"present" path at 2 allocs / ~1µs; the 6-alloc paths are the uuid-generation cases (unavoidable). -All costs are negligible against real network I/O — no further tuning is justified. - -## Remaining / follow-ups - -- Tag a new module version (e.g. `v0.1.2`) and bump `require github.com/runpod/rplog` in - `host` / `hapi` / `proxy` to pick up validation. Backward compatible — the only visible change is - that a hostile inbound id is sanitized to a fresh uuid. -- Add a JS `trace` module (`js/trace.ts`) so TS/JS services can converge on the shared library - (RunPod currently uses its own monorepo `@runpod/rplog` helper). -- ~~Update the top-level README to document the validation and response-echo usage.~~ Done: - the README `## Tracing` section now carries the HTTP header contract table, the Go - inbound/outbound/response-echo usage, and the inbound-validation behavior. (GO_README stays - the minimal slog-logger quickstart; the header contract lives in README.) From cabcbaa5e925fed680a4d90977c336691373f0ef Mon Sep 17 00:00:00 2001 From: "randomizedcoder dave.seddon.ca@gmail.com" Date: Tue, 22 Sep 2026 12:01:29 -0700 Subject: [PATCH 4/5] fix(trace): clamp a future X-Trace-Start silently instead of warning FromHeaderOrNew logged an slog.Warn once per request whenever an inbound X-Trace-Start was in the future. That value is caller-controlled, so an untrusted client could drive a consuming service's log volume by sending a future timestamp on every request. A malformed start already falls back to now with no log; the future case is now handled the same way, folded into the parse guard. Align the Python resolver, which previously neither clamped nor warned on a future start (it propagated it verbatim): it now clamps to now to match Go, covered by a regression test. This removes the need for each public-facing consumer to strip inbound X-Trace-Start defensively. Co-Authored-By: Claude Opus 4.8 --- py/trace.py | 15 ++++++++++----- py/trace_test.py | 13 ++++++++++++- trace/trace.go | 13 ++++++------- 3 files changed, 28 insertions(+), 13 deletions(-) diff --git a/py/trace.py b/py/trace.py index fd49a33..82c08cb 100644 --- a/py/trace.py +++ b/py/trace.py @@ -34,13 +34,18 @@ def _resolve_source(s: str) -> str: 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 or malformed. Re-formatting through - # as_rfc3339 means the stored value can never carry CR/LF (or any other - # injected bytes) into save_to_headers — the emit-side guard that mirrors Go's - # time.Parse of X-Trace-Start. + # 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: - return as_rfc3339(datetime.fromisoformat(s.replace("Z", "+00:00"))) + 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) diff --git a/py/trace_test.py b/py/trace_test.py index b8e99dd..f66b71e 100644 --- a/py/trace_test.py +++ b/py/trace_test.py @@ -1,5 +1,5 @@ import unittest -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from trace import Trace, _resolve_id, _resolve_source, _valid_id @@ -160,6 +160,17 @@ def test_valid_trace_start_is_preserved(self): 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): diff --git a/trace/trace.go b/trace/trace.go index ed549eb..786d860 100644 --- a/trace/trace.go +++ b/trace/trace.go @@ -2,7 +2,6 @@ package trace import ( "context" - "log/slog" "net/http" "time" @@ -143,18 +142,18 @@ func newuuid() string { func FromHeaderOrNew(h http.Header) Trace { now := time.Now().UTC() + // 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 { + if parsed, err := time.Parse(time.RFC3339, raw); err == nil && !parsed.After(now) { traceStart = parsed } } - if traceStart.After(now) { - slog.Warn("trace start is in the future", slog.Time("trace_start", traceStart), slog.Time("now", now)) - traceStart = now - } - return Trace{ TraceID: resolveID(h.Get(headerTraceID)), RequestID: resolveID(h.Get(headerRequestID)), From 251cfb8d07cee4bd6a780bfbe4213d02fafe25df Mon Sep 17 00:00:00 2001 From: "randomizedcoder dave.seddon.ca@gmail.com" Date: Thu, 24 Sep 2026 15:49:23 -0700 Subject: [PATCH 5/5] feat(trace): bound inbound trace/request source to maxSourceLen (64) X-Trace-Source / X-Request-Source is a service name, not an id, so it no longer inherits the 200-byte id budget: an inbound source over 64 bytes is dropped to "" alongside the existing charset check. Adds validSource in Go and _valid_source in Python (parity), with boundary + over-length table rows in both test suites. Co-Authored-By: Claude Opus 4.8 --- py/trace.py | 16 +++++++++++++--- py/trace_test.py | 3 +++ trace/trace.go | 21 +++++++++++++++++---- trace/trace_test.go | 3 +++ 4 files changed, 36 insertions(+), 7 deletions(-) diff --git a/py/trace.py b/py/trace.py index 82c08cb..053b28d 100644 --- a/py/trace.py +++ b/py/trace.py @@ -17,6 +17,10 @@ def as_rfc3339(dt: datetime) -> str: _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 @@ -26,10 +30,16 @@ 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 (including - # empty) to "" so save_to_headers can never re-emit unsafe bytes. - return s if _valid_id(s) else "" + # 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: diff --git a/py/trace_test.py b/py/trace_test.py index f66b71e..92b18f2 100644 --- a/py/trace_test.py +++ b/py/trace_test.py @@ -47,8 +47,11 @@ 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"]): diff --git a/trace/trace.go b/trace/trace.go index 786d860..b56ec93 100644 --- a/trace/trace.go +++ b/trace/trace.go @@ -168,6 +168,12 @@ func FromHeaderOrNew(h http.Header) Trace { // 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, @@ -196,11 +202,18 @@ func resolveID(id string) string { return newuuid() } -// resolveSource returns src unchanged when empty (the "unknown" case) or when -// it passes validID, and drops any other value to "" so SaveToHeader can never -// re-emit unsafe bytes. +// 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)) +} + +// 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 src == "" || validID(src) { + if validSource(src) { return src } return "" diff --git a/trace/trace_test.go b/trace/trace_test.go index 5c7788c..3c650fb 100644 --- a/trace/trace_test.go +++ b/trace/trace_test.go @@ -86,8 +86,11 @@ func TestResolveSource(t *testing.T) { }{ {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) {