Source code for taimoe.platform.proxy.parse
"""Server-side parsers for the headers produced by ``build_taimoe_headers``.
Backend audit middleware (and tests) should depend on this module so the
producer/consumer contract for ``X-Taimoe-*`` and ``traceparent`` lives in
exactly one place.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Mapping
from taimoe.platform._ids import is_valid_span_id, is_valid_trace_id
from .headers import (
HEADER_AGENT_ID,
HEADER_PARENT_AGENT_ID,
HEADER_PARENT_SPAN_ID,
HEADER_RUNTIME_ID,
HEADER_SESSION_ID,
HEADER_SPAN_ID,
HEADER_STEP,
HEADER_TRACE_ID,
HEADER_TRACEPARENT,
HEADER_USER_ID,
)
[docs]
@dataclass(frozen=True)
class TraceParent:
"""Parsed W3C ``traceparent`` value."""
version: str
trace_id: str
span_id: str
sampled: bool
[docs]
@dataclass(frozen=True)
class TaimoeRequestContext:
"""Result of parsing a request's Taimoe-native headers."""
trace_id: str | None
span_id: str | None
parent_span_id: str | None
runtime_id: str | None
agent_id: str | None
parent_agent_id: str | None
step: str | None
session_id: str | None
user_id: str | None
def _get(headers: Mapping[str, str], name: str) -> str | None:
"""Case-insensitive header lookup that tolerates either dict or Mapping input."""
if name in headers:
return headers[name]
lowered = name.lower()
for key, value in headers.items():
if key.lower() == lowered:
return value
return None
[docs]
def parse_traceparent(value: str) -> TraceParent | None:
"""Parse a W3C ``traceparent`` value, returning ``None`` if malformed.
Format: ``version-traceid-spanid-flags`` (all hex).
Per spec, unknown future versions MUST still parse the first three fields.
"""
if not isinstance(value, str):
return None
parts = value.strip().split("-")
if len(parts) < 4:
return None
version, trace_id, span_id, flags = parts[0], parts[1], parts[2], parts[3]
if len(version) != 2 or not all(c in "0123456789abcdef" for c in version.lower()):
return None
if not is_valid_trace_id(trace_id.lower()):
return None
if not is_valid_span_id(span_id.lower()):
return None
try:
flag_bits = int(flags, 16)
except ValueError:
return None
return TraceParent(
version=version.lower(),
trace_id=trace_id.lower(),
span_id=span_id.lower(),
sampled=bool(flag_bits & 0x01),
)