1use crate::{SpanId, TraceFlags, TraceId};
2use std::collections::VecDeque;
3use std::hash::Hash;
4use std::str::FromStr;
5use thiserror::Error;
6
7#[derive(Clone, Debug, Default, Eq, PartialEq, Hash)]
15pub struct TraceState(Option<VecDeque<(String, String)>>);
16
17const MAX_LIST_MEMBERS: usize = 32;
18
19impl TraceState {
20 pub const NONE: TraceState = TraceState(None);
22
23 fn valid_key(key: &str) -> bool {
27 if key.len() > 256 {
28 return false;
29 }
30
31 let allowed_special = |b: u8| b == b'_' || b == b'-' || b == b'*' || b == b'/';
32 let mut vendor_start = None;
33 for (i, &b) in key.as_bytes().iter().enumerate() {
34 if !(b.is_ascii_lowercase() || b.is_ascii_digit() || allowed_special(b) || b == b'@') {
35 return false;
36 }
37
38 if i == 0 && (!b.is_ascii_lowercase() && !b.is_ascii_digit()) {
39 return false;
40 } else if b == b'@' {
41 if vendor_start.is_some() || i + 14 < key.len() {
42 return false;
43 }
44 vendor_start = Some(i);
45 } else if let Some(start) = vendor_start {
46 if i == start + 1 && !(b.is_ascii_lowercase() || b.is_ascii_digit()) {
47 return false;
48 }
49 }
50 }
51
52 true
53 }
54
55 fn valid_value(value: &str) -> bool {
59 if value.len() > 256 {
60 return false;
61 }
62
63 !(value.contains(',') || value.contains('='))
64 }
65
66 pub fn from_key_value<T, K, V>(trace_state: T) -> TraceStateResult<Self>
85 where
86 T: IntoIterator<Item = (K, V)>,
87 K: ToString,
88 V: ToString,
89 {
90 let ordered_data = trace_state
91 .into_iter()
92 .take(MAX_LIST_MEMBERS)
93 .map(|(key, value)| {
94 let (key, value) = (key.to_string(), value.to_string());
95 if !TraceState::valid_key(key.as_str()) {
96 return Err(TraceStateError::Key(key));
97 }
98 if !TraceState::valid_value(value.as_str()) {
99 return Err(TraceStateError::Value(value));
100 }
101
102 Ok((key, value))
103 })
104 .collect::<Result<VecDeque<_>, TraceStateError>>()?;
105
106 if ordered_data.is_empty() {
107 Ok(TraceState(None))
108 } else {
109 Ok(TraceState(Some(ordered_data)))
110 }
111 }
112
113 pub fn get(&self, key: &str) -> Option<&str> {
115 self.0.as_ref().and_then(|kvs| {
116 kvs.iter().find_map(|item| {
117 if item.0.as_str() == key {
118 Some(item.1.as_str())
119 } else {
120 None
121 }
122 })
123 })
124 }
125
126 pub fn insert<K, V>(&self, key: K, value: V) -> TraceStateResult<TraceState>
136 where
137 K: Into<String>,
138 V: Into<String>,
139 {
140 let (key, value) = (key.into(), value.into());
141 if !TraceState::valid_key(key.as_str()) {
142 return Err(TraceStateError::Key(key));
143 }
144 if !TraceState::valid_value(value.as_str()) {
145 return Err(TraceStateError::Value(value));
146 }
147
148 let mut trace_state = self.delete_from_deque(&key);
149 let kvs = trace_state.0.get_or_insert(VecDeque::with_capacity(1));
150 if kvs.len() >= MAX_LIST_MEMBERS {
151 kvs.truncate(MAX_LIST_MEMBERS - 1);
152 }
153
154 kvs.push_front((key, value));
155
156 Ok(trace_state)
157 }
158
159 pub fn delete<K: Into<String>>(&self, key: K) -> TraceStateResult<TraceState> {
167 let key = key.into();
168 if !TraceState::valid_key(key.as_str()) {
169 return Err(TraceStateError::Key(key));
170 }
171
172 Ok(self.delete_from_deque(&key))
173 }
174
175 fn delete_from_deque(&self, key: &str) -> TraceState {
177 let mut owned = self.clone();
178 if let Some(kvs) = owned.0.as_mut() {
179 if let Some(index) = kvs.iter().position(|x| x.0 == key) {
180 kvs.remove(index);
181 }
182 }
183 owned
184 }
185
186 pub fn header(&self) -> String {
189 self.header_delimited("=", ",")
190 }
191
192 pub fn header_delimited(&self, entry_delimiter: &str, list_delimiter: &str) -> String {
194 self.0
195 .as_ref()
196 .map(|kvs| {
197 kvs.iter()
198 .map(|(key, value)| format!("{key}{entry_delimiter}{value}"))
199 .collect::<Vec<String>>()
200 .join(list_delimiter)
201 })
202 .unwrap_or_default()
203 }
204}
205
206impl FromStr for TraceState {
207 type Err = TraceStateError;
208
209 fn from_str(s: &str) -> Result<Self, Self::Err> {
210 let mut key_value_pairs: Vec<(String, String)> = Vec::new();
211
212 for list_member in s.split_terminator(',').take(MAX_LIST_MEMBERS) {
213 match list_member.find('=') {
214 None => return Err(TraceStateError::List(list_member.to_string())),
215 Some(separator_index) => {
216 let (key, value) = list_member.split_at(separator_index);
217 key_value_pairs
218 .push((key.to_string(), value.trim_start_matches('=').to_string()));
219 }
220 }
221 }
222
223 TraceState::from_key_value(key_value_pairs)
224 }
225}
226
227#[derive(Debug)]
229pub struct TraceStateIter<'a> {
230 inner: Option<std::collections::vec_deque::Iter<'a, (String, String)>>,
231}
232
233impl<'a> Iterator for TraceStateIter<'a> {
234 type Item = (&'a str, &'a str);
235
236 fn next(&mut self) -> Option<Self::Item> {
237 self.inner
238 .as_mut()?
239 .next()
240 .map(|(key, value)| (key.as_str(), value.as_str()))
241 }
242
243 fn size_hint(&self) -> (usize, Option<usize>) {
244 match &self.inner {
245 Some(iter) => iter.size_hint(),
246 None => (0, Some(0)),
247 }
248 }
249}
250
251impl ExactSizeIterator for TraceStateIter<'_> {
252 fn len(&self) -> usize {
253 match &self.inner {
254 Some(iter) => iter.len(),
255 None => 0,
256 }
257 }
258}
259
260impl<'a> IntoIterator for &'a TraceState {
261 type Item = (&'a str, &'a str);
262 type IntoIter = TraceStateIter<'a>;
263
264 fn into_iter(self) -> Self::IntoIter {
265 TraceStateIter {
266 inner: self.0.as_ref().map(|deque| deque.iter()),
267 }
268 }
269}
270
271type TraceStateResult<T> = Result<T, TraceStateError>;
273
274#[derive(Error, Debug)]
276#[non_exhaustive]
277pub enum TraceStateError {
278 #[error("{0} is not a valid key in TraceState, see https://www.w3.org/TR/trace-context-2/#key for more details")]
282 Key(String),
283
284 #[error("{0} is not a valid value in TraceState, see https://www.w3.org/TR/trace-context-2/#value for more details")]
288 Value(String),
289
290 #[error("{0} is not a valid list member in TraceState, see https://www.w3.org/TR/trace-context-2/#list for more details")]
294 List(String),
295}
296
297#[derive(Clone, Debug, PartialEq, Hash, Eq)]
307pub struct SpanContext {
308 trace_id: TraceId,
309 span_id: SpanId,
310 trace_flags: TraceFlags,
311 is_remote: bool,
312 trace_state: TraceState,
313}
314
315impl SpanContext {
316 pub const NONE: SpanContext = SpanContext {
318 trace_id: TraceId::INVALID,
319 span_id: SpanId::INVALID,
320 trace_flags: TraceFlags::NOT_SAMPLED,
321 is_remote: false,
322 trace_state: TraceState::NONE,
323 };
324
325 pub fn empty_context() -> Self {
327 SpanContext::NONE
328 }
329
330 pub fn new(
332 trace_id: TraceId,
333 span_id: SpanId,
334 trace_flags: TraceFlags,
335 is_remote: bool,
336 trace_state: TraceState,
337 ) -> Self {
338 SpanContext {
339 trace_id,
340 span_id,
341 trace_flags,
342 is_remote,
343 trace_state,
344 }
345 }
346
347 pub fn trace_id(&self) -> TraceId {
349 self.trace_id
350 }
351
352 pub fn span_id(&self) -> SpanId {
354 self.span_id
355 }
356
357 pub fn trace_flags(&self) -> TraceFlags {
362 self.trace_flags
363 }
364
365 pub fn is_valid(&self) -> bool {
368 self.trace_id != TraceId::INVALID && self.span_id != SpanId::INVALID
369 }
370
371 pub fn is_remote(&self) -> bool {
373 self.is_remote
374 }
375
376 pub fn is_sampled(&self) -> bool {
380 self.trace_flags.is_sampled()
381 }
382
383 pub fn trace_state(&self) -> &TraceState {
385 &self.trace_state
386 }
387}
388
389#[cfg(test)]
390mod tests {
391 use super::*;
392 use crate::{trace::TraceContextExt, Context};
393
394 #[rustfmt::skip]
395 fn trace_state_test_data() -> Vec<(TraceState, &'static str, &'static str)> {
396 vec![
397 (TraceState::from_key_value(vec![("foo", "bar")]).unwrap(), "foo=bar", "foo"),
398 (TraceState::from_key_value(vec![("foo", ""), ("apple", "banana")]).unwrap(), "foo=,apple=banana", "apple"),
399 (TraceState::from_key_value(vec![("foo", "bar"), ("apple", "banana")]).unwrap(), "foo=bar,apple=banana", "apple"),
400 ]
401 }
402
403 #[test]
404 fn test_trace_state() {
405 for test_case in trace_state_test_data() {
406 assert_eq!(test_case.0.clone().header(), test_case.1);
407
408 let new_key = format!("{}-{}", test_case.0.get(test_case.2).unwrap(), "test");
409
410 let updated_trace_state = test_case.0.insert(test_case.2, new_key.clone());
411 assert!(updated_trace_state.is_ok());
412 let updated_trace_state = updated_trace_state.unwrap();
413
414 let updated = format!("{}={}", test_case.2, new_key);
415
416 let index = updated_trace_state.clone().header().find(&updated);
417
418 assert!(index.is_some());
419 assert_eq!(index.unwrap(), 0);
420
421 let deleted_trace_state = updated_trace_state.delete(test_case.2.to_string());
422 assert!(deleted_trace_state.is_ok());
423
424 let deleted_trace_state = deleted_trace_state.unwrap();
425
426 assert!(deleted_trace_state.get(test_case.2).is_none());
427 }
428 }
429
430 #[test]
431 fn test_trace_state_key() {
432 let test_data: Vec<(&'static str, bool)> = vec![
433 ("123", true),
434 ("bar", true),
435 ("foo@bar", true),
436 ("foo@0123456789abcdef", false),
437 ("foo@012345678", true),
438 ("FOO@BAR", false),
439 ("你好", false),
440 ];
441
442 for (key, expected) in test_data {
443 assert_eq!(TraceState::valid_key(key), expected, "test key: {key:?}");
444 }
445 }
446
447 #[test]
448 fn test_trace_state_insert() {
449 let trace_state = TraceState::from_key_value(vec![("foo", "bar")]).unwrap();
450 let inserted_trace_state = trace_state.insert("testkey", "testvalue").unwrap();
451 assert!(trace_state.get("testkey").is_none()); assert_eq!(inserted_trace_state.get("testkey").unwrap(), "testvalue"); }
454
455 #[test]
456 fn test_context_span_debug() {
457 let cx = Context::current();
458 assert_eq!(
459 format!("{cx:?}"),
460 "Context { span: \"None\", entries count: 0, suppress_telemetry: false }"
461 );
462 let cx = Context::current().with_remote_span_context(SpanContext::NONE);
463 assert_eq!(
464 format!("{cx:?}"),
465 "Context { \
466 span: SpanContext { \
467 trace_id: 00000000000000000000000000000000, \
468 span_id: 0000000000000000, \
469 trace_flags: TraceFlags(0), \
470 is_remote: false, \
471 trace_state: TraceState(None) \
472 }, \
473 entries count: 1, suppress_telemetry: false \
474 }"
475 );
476 }
477
478 #[test]
479 fn test_tracestate_iter_empty() {
480 let ts = TraceState::NONE;
481 let mut iter = ts.into_iter();
482 assert_eq!(iter.next(), None);
483 assert_eq!(iter.size_hint(), (0, Some(0)));
484 assert_eq!(iter.len(), 0);
485 }
486
487 #[test]
488 fn test_tracestate_iter_single() {
489 let ts = TraceState::from_key_value(vec![("foo", "bar")]).unwrap();
490 let mut iter = ts.into_iter();
491 assert_eq!(iter.next(), Some(("foo", "bar")));
492 assert_eq!(iter.next(), None);
493 assert_eq!(iter.size_hint(), (0, Some(0)));
494 }
495
496 #[test]
497 fn test_tracestate_iter_multiple() {
498 let ts = TraceState::from_key_value(vec![("foo", "bar"), ("apple", "banana")]).unwrap();
499 let mut iter = ts.into_iter();
500 assert_eq!(iter.next(), Some(("foo", "bar")));
501 assert_eq!(iter.next(), Some(("apple", "banana")));
502 assert_eq!(iter.next(), None);
503 }
504
505 #[test]
506 fn test_tracestate_iter_size_hint_and_len() {
507 let ts = TraceState::from_key_value(vec![("foo", "bar"), ("apple", "banana")]).unwrap();
508 let iter = ts.into_iter();
509 assert_eq!(iter.size_hint(), (2, Some(2)));
510 assert_eq!(iter.len(), 2);
511 }
512
513 #[test]
514 fn test_tracestate_from_str_keeps_at_most_32_list_members() {
515 let header = (0..64)
516 .map(|i| format!("key{i}=value{i}"))
517 .collect::<Vec<_>>()
518 .join(",");
519 let trace_state = TraceState::from_str(&header).unwrap();
520
521 assert_eq!(trace_state.into_iter().count(), 32);
522 assert_eq!(trace_state.get("key0"), Some("value0"));
523 assert_eq!(trace_state.get("key31"), Some("value31"));
524 assert_eq!(trace_state.get("key32"), None);
525
526 let expected_header = (0..32)
527 .map(|i| format!("key{i}=value{i}"))
528 .collect::<Vec<_>>()
529 .join(",");
530 assert_eq!(trace_state.header(), expected_header);
531 }
532 #[test]
533 fn test_tracestate_from_str_ignores_invalid_members_beyond_the_limit() {
534 let mut members: Vec<String> = (0..32).map(|i| format!("key{i}=value{i}")).collect();
535 members.push("invalid-no-separator".to_string());
536 let trace_state = TraceState::from_str(&members.join(",")).unwrap();
537
538 assert_eq!(trace_state.into_iter().count(), 32);
539 }
540 #[test]
541 fn test_tracestate_from_str_rejects_invalid_members_within_the_limit() {
542 let mut members: Vec<String> = (0..31).map(|i| format!("key{i}=value{i}")).collect();
543 members.push("invalid-no-separator".to_string());
544
545 assert!(TraceState::from_str(&members.join(",")).is_err());
546 }
547 #[test]
548 fn test_tracestate_from_key_value_keeps_at_most_32_list_members() {
549 let kvs: Vec<(String, String)> = (0..64)
550 .map(|i| (format!("key{i}"), format!("value{i}")))
551 .collect();
552 let trace_state = TraceState::from_key_value(kvs).unwrap();
553
554 assert_eq!(trace_state.into_iter().count(), 32);
555 assert_eq!(trace_state.get("key0"), Some("value0"));
556 assert_eq!(trace_state.get("key32"), None);
557 }
558 #[test]
559 fn test_tracestate_insert_keeps_at_most_32_list_members() {
560 let mut trace_state = TraceState::default();
561 for i in 0..64 {
562 trace_state = trace_state
563 .insert(format!("key{i}"), format!("value{i}"))
564 .unwrap();
565 }
566
567 assert_eq!(trace_state.into_iter().count(), 32);
568 assert_eq!(trace_state.get("key63"), Some("value63"));
569 assert_eq!(trace_state.get("key32"), Some("value32"));
570 assert_eq!(trace_state.get("key31"), None);
571 }
572 #[test]
573 fn test_tracestate_insert_of_an_existing_key_evicts_nothing() {
574 let mut trace_state = TraceState::default();
575 for i in 0..32 {
576 trace_state = trace_state
577 .insert(format!("key{i}"), format!("value{i}"))
578 .unwrap();
579 }
580 let updated = trace_state.insert("key20", "updated").unwrap();
581
582 assert_eq!(updated.into_iter().count(), 32);
583 assert_eq!(updated.get("key20"), Some("updated"));
584 assert_eq!(updated.get("key0"), Some("value0"));
585 }
586}