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}