use crate::ErrorResult;
use crate::doc::DocumentableDTO;
use crate::utils::request_parser::{MultipartBody, UploadedFile};
use crate::validation::Validate;
use serde_json::Value;
use std::any::Any;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use uuid::Uuid;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CorrelationFlow {
Once,
Start,
Continue,
End,
}
impl std::str::FromStr for CorrelationFlow {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.trim().to_uppercase().as_str() {
"ONCE" => Ok(Self::Once),
"START" => Ok(Self::Start),
"CONTINUE" => Ok(Self::Continue),
"END" => Ok(Self::End),
other => Err(format!("unknown flow {other}")),
}
}
}
impl std::fmt::Display for CorrelationFlow {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let s = match self {
Self::Once => "ONCE",
Self::Start => "START",
Self::Continue => "CONTINUE",
Self::End => "END",
};
write!(f, "{s}")
}
}
#[derive(Debug, Clone)]
pub struct CorrelationContext {
inner: Arc<Mutex<Inner>>,
}
#[derive(Debug)]
struct Inner {
correlation_id: String,
request_id: String,
flow: CorrelationFlow,
user_id: Option<String>,
body: Option<Vec<u8>>,
headers: http::HeaderMap,
query_params: HashMap<String, String>,
multipart: Option<MultipartBody>,
data: HashMap<String, Box<dyn Any + Send + Sync>>,
path_params: HashMap<String, String>,
pagination_cursor: Option<String>,
pagination_limit: usize,
}
impl CorrelationContext {
pub fn new() -> Self {
let corr = Uuid::new_v4().to_string();
let req = hex::encode(rand::random::<[u8; 8]>());
Self {
inner: Arc::new(Mutex::new(Inner {
correlation_id: corr,
request_id: req,
flow: CorrelationFlow::Once,
user_id: None,
body: None,
headers: http::HeaderMap::new(),
query_params: HashMap::new(),
path_params: HashMap::new(),
multipart: None,
data: HashMap::new(),
pagination_cursor: None,
pagination_limit: 15,
})),
}
}
pub fn with_ids(correlation_id: &str, request_id: &str) -> Self {
let ctx = Self::new();
{
let mut inner = ctx.inner.lock().unwrap();
inner.correlation_id = correlation_id.to_string();
inner.request_id = request_id.to_string();
}
ctx
}
pub(crate) fn set_body(&self, body: Vec<u8>) {
self.inner.lock().unwrap().body = Some(body);
}
pub(crate) fn set_headers(&self, headers: http::HeaderMap) {
self.inner.lock().unwrap().headers = headers;
}
pub(crate) fn set_params(&self, params: HashMap<String, String>) {
self.inner.lock().unwrap().path_params = params;
}
pub(crate) fn set_query_params(&self, params: HashMap<String, String>) {
self.inner.lock().unwrap().query_params = params;
}
pub(crate) fn set_multipart(&self, multipart: MultipartBody) {
self.inner.lock().unwrap().multipart = Some(multipart);
}
pub fn is_multipart(&self) -> bool {
self.inner.lock().unwrap().multipart.is_some()
}
pub fn multipart(&self) -> Option<MultipartBody> {
self.inner.lock().unwrap().multipart.clone()
}
pub fn form_field(&self, name: &str) -> Option<String> {
self.inner
.lock()
.unwrap()
.multipart
.as_ref()
.and_then(|mp| mp.field(name))
.map(|s| s.to_string())
}
pub fn form_field_all(&self, name: &str) -> Vec<String> {
self.inner
.lock()
.unwrap()
.multipart
.as_ref()
.and_then(|mp| mp.fields.get(name).cloned())
.unwrap_or_default()
}
pub fn files(&self) -> Vec<UploadedFile> {
self.inner
.lock()
.unwrap()
.multipart
.as_ref()
.map(|mp| mp.files.clone())
.unwrap_or_default()
}
pub fn files_for(&self, field: &str) -> Vec<UploadedFile> {
self.files()
.into_iter()
.filter(|f| f.field_name == field)
.collect()
}
pub fn body_bytes(&self) -> Option<Vec<u8>> {
self.inner.lock().unwrap().body.clone()
}
pub fn body_string(&self) -> Result<String, ErrorResult> {
if let Some(payload) = self.form_field("body") {
return Ok(payload);
}
let body_bytes_opt = self.inner.lock().unwrap().body.clone();
let Some(body_bytes) = body_bytes_opt else {
return Err(ErrorResult::bad_request("no body"));
};
String::from_utf8(body_bytes).map_err(|_| ErrorResult::bad_request("invalid body"))
}
pub fn body<T>(&self) -> Result<T, ErrorResult>
where
T: DocumentableDTO + Validate,
{
if let Some(mp) = self.inner.lock().unwrap().multipart.clone() {
return self.multipart_body(&mp);
}
let body_bytes_opt = self.inner.lock().unwrap().body.clone();
let Some(body_bytes) = body_bytes_opt else {
return Err(ErrorResult::bad_request("no body"));
};
let parsed_body: Option<T> = serde_json::from_slice(&body_bytes).ok().or_else(|| {
serde_urlencoded::from_bytes::<Value>(&body_bytes)
.ok()
.and_then(|v| serde_json::from_value(v).ok())
});
let Some(parsed_body) = parsed_body else {
return Err(ErrorResult::bad_request("invalid body"));
};
if let Err(validated) = parsed_body.validate() {
return Err(ErrorResult::new(
validated.message,
Some(Value::String(validated.field)),
400,
));
}
Ok(parsed_body.clone())
}
fn multipart_body<T>(&self, mp: &MultipartBody) -> Result<T, ErrorResult>
where
T: DocumentableDTO + Validate,
{
if let Some(payload) = mp.field("body") {
let parsed: T = serde_json::from_str(payload)
.map_err(|_| ErrorResult::bad_request("invalid body"))?;
return self.validated(parsed);
}
let mut map = serde_json::Map::new();
for (k, vs) in &mp.fields {
let v = if vs.len() == 1 {
serde_json::Value::String(vs[0].clone())
} else {
serde_json::Value::Array(
vs.iter().cloned().map(serde_json::Value::String).collect(),
)
};
map.insert(k.clone(), v);
}
if map.is_empty() {
return Err(ErrorResult::bad_request("no body"));
}
let parsed: T = serde_json::from_value(serde_json::Value::Object(map))
.map_err(|_| ErrorResult::bad_request("invalid body"))?;
self.validated(parsed)
}
fn validated<T>(&self, parsed: T) -> Result<T, ErrorResult>
where
T: DocumentableDTO + Validate,
{
if let Err(e) = parsed.validate() {
return Err(ErrorResult::bad_request(e.message));
}
Ok(parsed)
}
pub fn headers(&self) -> http::HeaderMap {
self.inner.lock().unwrap().headers.clone()
}
pub fn header(&self, name: &str) -> Option<String> {
self.inner
.lock()
.unwrap()
.headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
}
pub fn query_params(&self) -> HashMap<String, String> {
self.inner.lock().unwrap().query_params.clone()
}
pub fn query_param(&self, name: &str) -> Option<String> {
self.inner.lock().unwrap().query_params.get(name).cloned()
}
pub fn query_param_or(&self, name: &str, default: &str) -> String {
self.query_param(name)
.unwrap_or_else(|| default.to_string())
}
pub fn path_params(&self) -> HashMap<String, String> {
self.inner.lock().unwrap().path_params.clone()
}
pub fn path_param(&self, name: &str) -> Option<String> {
self.inner.lock().unwrap().path_params.get(name).cloned()
}
pub fn path_param_or(&self, name: &str, default: &str) -> String {
self.query_param(name)
.unwrap_or_else(|| default.to_string())
}
pub fn correlation_id(&self) -> String {
self.inner.lock().unwrap().correlation_id.clone()
}
pub fn request_id(&self) -> String {
self.inner.lock().unwrap().request_id.clone()
}
pub fn set_request_id(&self, id: &str) {
self.inner.lock().unwrap().request_id = id.to_string();
}
pub fn set_correlation_id(&self, id: &str) {
self.inner.lock().unwrap().correlation_id = id.to_string();
}
pub fn flow(&self) -> CorrelationFlow {
self.inner.lock().unwrap().flow
}
pub fn set_flow(&self, flow: CorrelationFlow) {
self.inner.lock().unwrap().flow = flow;
}
pub fn with_flow(self, flow: CorrelationFlow) -> Self {
self.set_flow(flow);
self
}
pub fn with_correlation_id(self, id: &str) -> Self {
self.set_correlation_id(id);
self
}
pub fn user_id(&self) -> Option<String> {
self.inner.lock().unwrap().user_id.clone()
}
pub fn set_user_id(&self, id: Option<String>) {
self.inner.lock().unwrap().user_id = id;
}
pub fn set<T>(&self, key: &str, value: T)
where
T: Any + Send + Sync,
{
self.inner
.lock()
.unwrap()
.data
.insert(key.to_string(), Box::new(value));
}
pub fn get<T>(&self, key: &str) -> Option<T>
where
T: Any + Clone,
{
self.inner
.lock()
.unwrap()
.data
.get(key)
.and_then(|value| value.downcast_ref::<T>())
.cloned()
}
pub fn set_string(&self, key: &str, value: impl Into<String>) {
self.set(key, value.into());
}
pub fn get_string(&self, key: &str) -> Option<String> {
self.get(key)
}
pub fn set_bool(&self, key: &str, value: bool) {
self.set(key, value);
}
pub fn get_bool(&self, key: &str) -> Option<bool> {
self.get(key)
}
pub fn set_number(&self, key: &str, value: f64) {
self.set(key, value);
}
pub fn get_number(&self, key: &str) -> Option<f64> {
self.get(key)
}
pub fn pagination_cursor(&self) -> Option<String> {
self.inner.lock().unwrap().pagination_cursor.clone()
}
pub fn pagination_limit(&self) -> usize {
self.inner.lock().unwrap().pagination_limit
}
pub fn set_pagination(&self, cursor: Option<String>, limit: usize) {
let mut inner = self.inner.lock().unwrap();
inner.pagination_cursor = cursor;
inner.pagination_limit = limit;
}
pub fn client_ip(
headers: &http::HeaderMap,
remote_addr: Option<std::net::SocketAddr>,
) -> String {
if let Some(v) = headers.get("x-forwarded-for").and_then(|h| h.to_str().ok()) {
if let Some(first) = v.split(',').next() {
let ip = first.trim();
if !ip.is_empty() {
return ip.to_string();
}
}
}
if let Some(v) = headers.get("x-real-ip").and_then(|h| h.to_str().ok()) {
return v.to_string();
}
remote_addr
.map(|a| a.ip().to_string())
.unwrap_or_else(|| "unknown".to_string())
}
}
impl Default for CorrelationContext {
fn default() -> Self {
Self::new()
}
}
tokio::task_local! {
pub static CORRELATION_CTX: CorrelationContext;
}
#[cfg(test)]
mod tests {
use super::CorrelationContext;
#[test]
fn stores_and_returns_owned_typed_values() {
let context = CorrelationContext::new();
context.set("count", 42_u32);
assert_eq!(context.get::<u32>("count"), Some(42));
assert_eq!(context.get::<String>("count"), None);
assert_eq!(context.get::<u32>("missing"), None);
}
#[test]
fn convenience_accessors_store_and_return_values() {
let context = CorrelationContext::new();
context.set_string("name", "Ada");
context.set_bool("enabled", true);
context.set_number("ratio", 1.5);
assert_eq!(context.get_string("name"), Some("Ada".to_string()));
assert_eq!(context.get_bool("enabled"), Some(true));
assert_eq!(context.get_number("ratio"), Some(1.5));
}
}