Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.


97 changes: 82 additions & 15 deletions py/trace.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import re
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Optional

from uuid7 import uuid7

Expand All @@ -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.
Expand All @@ -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", "")),

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

from_headers now defaults an absent X-Request-Source / X-Trace-Source to "" (via _resolve_source(headers.get(..., ""))), instead of the repo-wide "unknown" sentinel this replaced. That sentinel is still used by new() in this same file, and by trace/trace.go (thisServiceName defaults to "unknown"), so new() and from_headers now disagree on the default for no source given.

Any consumer that logs, queries, or dashboards on trace_source == "unknown" for untagged inbound requests will silently start seeing "". The PR frames the Python changes as bug fixes plus mirroring the Go validation, not as removing the "unknown" default, so this reads as an undocumented contract change.

Is this deliberate cross-language alignment? If so, please also update new() to match and call it out in the PR description. If not, keep "unknown" as the fallback here.

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
Expand All @@ -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(),
Expand All @@ -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
193 changes: 193 additions & 0 deletions py/trace_test.py
Original file line number Diff line number Diff line change
@@ -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()
Loading
Loading