use std::{error::Error, fmt};
use axum::{
http::StatusCode,
response::{IntoResponse, Response},
};
use crate::{RequestId, TraceContextLevel};
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct TraceContext {
version: u8,
trace_id: String,
parent_id: String,
flags: u8,
level: TraceContextLevel,
traceparent: Box<[u8]>,
tracestate: Option<String>,
}
impl TraceContext {
pub(crate) fn new(
version: u8,
trace_id: String,
parent_id: String,
flags: u8,
level: TraceContextLevel,
traceparent: Box<[u8]>,
) -> Self {
Self {
version,
trace_id,
parent_id,
flags,
level,
traceparent,
tracestate: None,
}
}
pub(crate) fn with_tracestate(mut self, tracestate: Option<String>) -> Self {
self.tracestate = tracestate;
self
}
#[must_use]
pub fn trace_id(&self) -> &str {
&self.trace_id
}
#[must_use]
pub fn parent_id(&self) -> &str {
&self.parent_id
}
#[must_use]
pub const fn flags(&self) -> u8 {
self.flags
}
#[must_use]
pub const fn sampled(&self) -> bool {
self.flags & 1 == 1
}
#[must_use]
pub const fn trace_context_level(&self) -> TraceContextLevel {
self.level
}
#[must_use]
pub const fn trace_id_random(&self) -> Option<bool> {
match (self.level, self.version) {
(TraceContextLevel::Level2, 0) => Some(self.flags & 2 == 2),
(TraceContextLevel::Level1 | TraceContextLevel::Level2, _) => None,
}
}
#[must_use]
pub fn traceparent_bytes(&self) -> &[u8] {
&self.traceparent
}
#[must_use]
pub fn traceparent(&self) -> Option<&str> {
std::str::from_utf8(&self.traceparent).ok()
}
#[must_use]
pub fn tracestate(&self) -> Option<&str> {
self.tracestate.as_deref()
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct RequestContext {
request_id: RequestId,
correlation_id: String,
trace_context: Option<TraceContext>,
}
impl<S> axum::extract::FromRequestParts<S> for RequestContext
where
S: Send + Sync,
{
type Rejection = MissingRequestContext;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<Self>()
.cloned()
.ok_or(MissingRequestContext)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub struct MissingRequestContext;
impl fmt::Display for MissingRequestContext {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("request context unavailable")
}
}
impl Error for MissingRequestContext {}
impl IntoResponse for MissingRequestContext {
fn into_response(self) -> Response {
(StatusCode::INTERNAL_SERVER_ERROR, self.to_string()).into_response()
}
}
impl RequestContext {
pub(crate) fn new(request_id: RequestId, trace_context: Option<TraceContext>) -> Self {
let correlation_id = trace_context.as_ref().map_or_else(
|| request_id.as_str().to_owned(),
|trace| trace.trace_id.clone(),
);
Self {
request_id,
correlation_id,
trace_context,
}
}
#[must_use]
pub const fn request_id(&self) -> &RequestId {
&self.request_id
}
#[must_use]
pub fn correlation_id(&self) -> &str {
&self.correlation_id
}
#[must_use]
pub const fn trace_context(&self) -> Option<&TraceContext> {
self.trace_context.as_ref()
}
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub struct OperationId(&'static str);
impl OperationId {
#[track_caller]
#[must_use]
pub const fn from_static(value: &'static str) -> Self {
assert!(!value.is_empty(), "operation ID must not be empty");
Self(value)
}
#[must_use]
pub fn as_str(&self) -> &str {
self.0
}
}
impl AsRef<str> for OperationId {
fn as_ref(&self) -> &str {
self.as_str()
}
}
impl fmt::Display for OperationId {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(self.as_str())
}
}
#[cfg(test)]
mod tests {
use super::OperationId;
#[test]
fn operation_id_preserves_nonempty_static_values() {
const OPERATION: OperationId = OperationId::from_static("create-item");
assert_eq!(OPERATION.as_str(), "create-item");
assert_eq!(OPERATION.as_ref(), "create-item");
assert_eq!(OPERATION.to_string(), "create-item");
}
#[test]
#[should_panic(expected = "operation ID must not be empty")]
fn operation_id_rejects_empty_static_values() {
let _ = OperationId::from_static("");
}
#[test]
fn operation_id_preserves_application_static_controls() {
let operation = OperationId::from_static("create-item\nvariant");
assert_eq!(operation.as_str(), "create-item\nvariant");
}
}