Skip to main content

rmcp_traces/
http.rs

1use std::fmt;
2
3use ::http::{HeaderMap, HeaderValue};
4use rmcp::model::Meta;
5
6use crate::trace_context::{
7    parse_traceparent_value, validate_baggage_value, validate_tracestate_value,
8};
9use crate::{
10    TraceLimits, TraceParseError, TraceSummary, TraceTrust, BAGGAGE_KEY, TRACEPARENT_KEY,
11    TRACESTATE_KEY,
12};
13
14#[derive(Clone, Copy, Debug, PartialEq, Eq)]
15pub struct HttpTracePolicy {
16    pub trust: TraceTrust,
17    pub limits: TraceLimits,
18    pub include_baggage: bool,
19}
20
21impl Default for HttpTracePolicy {
22    fn default() -> Self {
23        Self {
24            trust: TraceTrust::Untrusted,
25            limits: TraceLimits::default(),
26            include_baggage: false,
27        }
28    }
29}
30
31pub struct HttpTraceExtraction {
32    pub meta: Meta,
33    pub summary: TraceSummary,
34}
35
36impl fmt::Debug for HttpTraceExtraction {
37    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
38        let meta_keys = [TRACEPARENT_KEY, TRACESTATE_KEY, BAGGAGE_KEY]
39            .into_iter()
40            .filter(|key| self.meta.get(*key).is_some())
41            .collect::<Vec<_>>();
42        f.debug_struct("HttpTraceExtraction")
43            .field("meta_keys", &meta_keys)
44            .field("summary", &self.summary)
45            .finish()
46    }
47}
48
49pub fn extract_http_trace(headers: &HeaderMap, policy: HttpTracePolicy) -> HttpTraceExtraction {
50    let mut meta = Meta::new();
51    let traceparent_value = match single_header_value(headers, TRACEPARENT_KEY) {
52        Ok(Some(value)) => value,
53        Ok(None) => {
54            return HttpTraceExtraction {
55                meta,
56                summary: TraceSummary::absent_with_trust(policy.trust),
57            };
58        }
59        Err(error) => {
60            let mut summary = TraceSummary::absent_with_trust(policy.trust);
61            summary.record_invalid_reason(&error);
62            return HttpTraceExtraction { meta, summary };
63        }
64    };
65
66    let traceparent = match parse_traceparent_value(traceparent_value, policy.limits) {
67        Ok(traceparent) => traceparent,
68        Err(error) => {
69            let mut summary = TraceSummary::absent_with_trust(policy.trust);
70            summary.record_invalid_reason(&error);
71            return HttpTraceExtraction { meta, summary };
72        }
73    };
74
75    meta.set_traceparent(traceparent_value);
76    let mut summary = TraceSummary::from_valid_traceparent(&traceparent, policy.trust);
77
78    match bounded_join_headers(headers, TRACESTATE_KEY, policy.limits.max_tracestate_len) {
79        Ok(Some(tracestate)) => match validate_tracestate_value(Some(&tracestate)) {
80            Ok(()) => {
81                meta.set_tracestate(tracestate);
82                summary.set_has_tracestate_for_http();
83            }
84            Err(error) => summary.record_invalid_reason(&error),
85        },
86        Ok(None) => {}
87        Err(error) => summary.record_invalid_reason(&error),
88    }
89
90    if policy.include_baggage {
91        match bounded_join_headers(headers, BAGGAGE_KEY, policy.limits.max_baggage_len) {
92            Ok(Some(baggage)) => {
93                match validate_baggage_value(Some(&baggage), policy.limits.max_baggage_members) {
94                    Ok(()) => {
95                        summary.set_baggage_counts_for_http(Some(&baggage));
96                        meta.set_baggage(baggage);
97                    }
98                    Err(error) => summary.record_invalid_reason(&error),
99                }
100            }
101            Ok(None) => {}
102            Err(error) => summary.record_invalid_reason(&error),
103        }
104    }
105
106    HttpTraceExtraction { meta, summary }
107}
108
109fn single_header_value<'a>(
110    headers: &'a HeaderMap,
111    field: &'static str,
112) -> Result<Option<&'a str>, TraceParseError> {
113    let values = headers.get_all(field);
114    match values.iter().take(2).count() {
115        0 => Ok(None),
116        1 => {
117            let value = values
118                .iter()
119                .next()
120                .expect("count confirmed one header value");
121            header_to_str(value, field).map(Some)
122        }
123        _ => Err(TraceParseError::MultipleHeaderValues { field }),
124    }
125}
126
127fn bounded_join_headers(
128    headers: &HeaderMap,
129    field: &'static str,
130    max: usize,
131) -> Result<Option<String>, TraceParseError> {
132    let mut joined = String::new();
133    let mut saw_value = false;
134    for value in headers.get_all(field).iter() {
135        let value = header_to_str(value, field)?;
136        let separator_len = usize::from(saw_value);
137        let required_len = joined.len() + separator_len + value.len();
138        if required_len > max {
139            return Err(TraceParseError::ValueTooLong {
140                field,
141                actual: max + 1,
142                max,
143            });
144        }
145        if saw_value {
146            joined.push(',');
147        }
148        joined.push_str(value);
149        saw_value = true;
150    }
151    Ok(saw_value.then_some(joined))
152}
153
154fn header_to_str<'a>(
155    value: &'a HeaderValue,
156    field: &'static str,
157) -> Result<&'a str, TraceParseError> {
158    value
159        .to_str()
160        .map_err(|_| TraceParseError::InvalidHeaderValue { field })
161}