saddle_observability/
trace.rs1use std::{error::Error, fmt};
2
3use saddle_core::TraceId;
4
5#[derive(Clone, Copy, Debug, Eq, PartialEq)]
7pub enum InboundTrace {
8 Inherited,
10 Created,
12 ReplacedInvalid,
14}
15
16impl InboundTrace {
17 pub(crate) const fn as_str(self) -> &'static str {
18 match self {
19 Self::Inherited => "inherited",
20 Self::Created => "created",
21 Self::ReplacedInvalid => "replaced_invalid",
22 }
23 }
24}
25
26#[derive(Clone, Copy, Debug, Eq, PartialEq)]
28pub enum TraceIdError {
29 InvalidLength,
30 InvalidHex,
31 Zero,
32}
33
34impl fmt::Display for TraceIdError {
35 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
36 formatter.write_str(match self {
37 Self::InvalidLength => "trace id must contain exactly 32 hexadecimal characters",
38 Self::InvalidHex => "trace id contains a non-hexadecimal character",
39 Self::Zero => "trace id must not be zero",
40 })
41 }
42}
43
44impl Error for TraceIdError {}
45
46pub(crate) fn select_trace_id(
47 value: Option<&str>,
48 generated: impl FnOnce() -> TraceId,
49) -> (TraceId, InboundTrace) {
50 match value {
51 Some(value) => match trace_id_from_hex(value) {
52 Ok(trace_id) => (trace_id, InboundTrace::Inherited),
53 Err(_) => (generated(), InboundTrace::ReplacedInvalid),
54 },
55 None => (generated(), InboundTrace::Created),
56 }
57}
58
59pub fn trace_id_from_hex(value: &str) -> Result<TraceId, TraceIdError> {
61 if value.len() != 32 {
62 return Err(TraceIdError::InvalidLength);
63 }
64 if !value.bytes().all(|byte| byte.is_ascii_hexdigit()) {
65 return Err(TraceIdError::InvalidHex);
66 }
67
68 let value = u128::from_str_radix(value, 16).map_err(|_| TraceIdError::InvalidHex)?;
69 if value == 0 {
70 return Err(TraceIdError::Zero);
71 }
72 Ok(TraceId::from_u128(value))
73}
74
75#[cfg(test)]
76mod tests {
77 use super::*;
78
79 #[test]
80 fn accepts_only_non_zero_fixed_width_hex_trace_ids() {
81 let expected = "00112233445566778899aabbccddeeff";
82 assert_eq!(trace_id_from_hex(expected).unwrap().to_string(), expected);
83 assert_eq!(trace_id_from_hex("abcd"), Err(TraceIdError::InvalidLength));
84 assert_eq!(
85 trace_id_from_hex("00112233445566778899aabbccddeefg"),
86 Err(TraceIdError::InvalidHex)
87 );
88 assert_eq!(
89 trace_id_from_hex("00000000000000000000000000000000"),
90 Err(TraceIdError::Zero)
91 );
92 }
93}