use std::{error::Error, fmt};
use axum::{
http::StatusCode,
response::{IntoResponse, Response},
};
use crate::RequestId;
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct TraceContext {
trace_id: String,
parent_id: String,
flags: u8,
traceparent: String,
tracestate: Option<String>,
}
impl TraceContext {
pub(crate) fn new(trace_id: String, parent_id: String, flags: u8, traceparent: String) -> Self {
Self {
trace_id,
parent_id,
flags,
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 fn traceparent(&self) -> &str {
&self.traceparent
}
#[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("");
}
}