use dataflow_rs::datalogic_rs;
use std::sync::Arc;
use serde_json::Value;
use super::ChannelRuntimeConfig;
use crate::channel::cache_namespace;
use crate::connector::cache_backend::CacheBackend;
use crate::metrics;
use sha2::{Digest, Sha256};
pub(super) fn resolve_key_field<'a>(data: &'a Value, field: &str) -> Option<&'a Value> {
fn walk<'a>(mut cur: &'a Value, path: &str) -> Option<&'a Value> {
for segment in path.split('.') {
if segment.is_empty() {
return None;
}
cur = cur.get(segment)?;
}
Some(cur)
}
if let Some(v) = data.get(field) {
return Some(v);
}
if !field.contains('.') {
return None;
}
walk(data, field).or_else(|| field.strip_prefix("data.").and_then(|p| walk(data, p)))
}
pub(super) fn compute_cache_key(
channel: &str,
data: &Value,
metadata: &Value,
cache_cfg: &crate::channel::ChannelCacheConfig,
key_logic: Option<&datalogic_rs::Logic>,
datalogic: &datalogic_rs::Engine,
) -> Option<String> {
let mut h = Sha256::new();
fn feed(h: &mut Sha256, bytes: &[u8]) {
h.update((bytes.len() as u64).to_be_bytes());
h.update(bytes);
}
feed(
&mut h,
metadata
.get("http_method")
.and_then(Value::as_str)
.unwrap_or("")
.as_bytes(),
);
feed_object_sorted(&mut h, metadata.get("params"));
feed_object_sorted(&mut h, metadata.get("query"));
if let Some(compiled) = key_logic {
let context = serde_json::json!({ "data": data, "metadata": metadata });
let key = datalogic
.session()
.eval_into::<Value, _>(compiled, &context)
.ok()?;
if key.is_null() {
return None;
}
feed(&mut h, &serde_json::to_vec(&key).unwrap_or_default());
} else if let Some(ref fields) = cache_cfg.cache_key_fields {
let mut resolved = 0usize;
for f in fields {
feed(&mut h, f.as_bytes());
match resolve_key_field(data, f) {
Some(v) => {
resolved += 1;
h.update([1u8]);
feed(&mut h, &serde_json::to_vec(v).unwrap_or_default());
}
None => h.update([0u8]),
}
}
if resolved == 0 {
return None;
}
} else {
feed(&mut h, &serde_json::to_vec(data).unwrap_or_default());
};
let digest = h.finalize();
Some(format!("cache:{channel}:{}", hex::encode(&digest[..16])))
}
pub(super) fn feed_object_sorted(h: &mut Sha256, v: Option<&Value>) {
let Some(Value::Object(map)) = v else {
h.update([0u8]);
return;
};
h.update([1u8]);
h.update((map.len() as u64).to_be_bytes());
let mut keys: Vec<&String> = map.keys().collect();
keys.sort_unstable();
for k in keys {
h.update((k.len() as u64).to_be_bytes());
h.update(k.as_bytes());
let bytes = serde_json::to_vec(&map[k.as_str()]).unwrap_or_default();
h.update((bytes.len() as u64).to_be_bytes());
h.update(&bytes);
}
}
pub struct CacheStoreCtx {
key: String,
backend: Arc<dyn CacheBackend>,
ttl_secs: u64,
versions: Option<Vec<i64>>,
_flight: Option<Box<FlightGuard>>,
}
impl CacheStoreCtx {
pub async fn store(&self, body: &str) -> Result<(), crate::errors::OrionError> {
match &self.versions {
None => self.backend.set_ex(&self.key, body, self.ttl_secs).await,
Some(versions) => {
let stored = cache_namespace::encode_entry(versions, body);
self.backend.set_ex(&self.key, &stored, self.ttl_secs).await
}
}
}
fn miss(
key: String,
backend: &Arc<dyn CacheBackend>,
ttl_secs: u64,
versions: Option<Vec<i64>>,
) -> CacheLookup {
CacheLookup::Miss(Some(Self {
key,
backend: backend.clone(),
ttl_secs,
versions,
_flight: None,
}))
}
}
pub(super) enum CacheLookup {
Hit(String),
Miss(Option<CacheStoreCtx>),
}
pub(super) async fn check_response_cache(
channel: &str,
data: &Value,
metadata: &Value,
channel_config: &Option<Arc<ChannelRuntimeConfig>>,
datalogic: &datalogic_rs::Engine,
) -> CacheLookup {
let Some(cfg) = channel_config else {
return CacheLookup::Miss(None);
};
let Some(ref cache_cfg) = cfg.parsed_config.cache else {
return CacheLookup::Miss(None);
};
if !cache_cfg.enabled {
return CacheLookup::Miss(None);
}
let Some(ref cache) = cfg.response_cache else {
return CacheLookup::Miss(None);
};
let Some(key) = compute_cache_key(
channel,
data,
metadata,
cache_cfg,
cfg.cache_key_logic.as_ref(),
datalogic,
) else {
tracing::warn!(
channel = %channel,
fields = ?cache_cfg.cache_key_fields,
has_key_logic = cfg.cache_key_logic.is_some(),
"No cache key resolved against the request; bypassing the response cache. \
Field names are literal payload keys or dotted paths (`user.id`, or \
`data.user_id` for a top-level `user_id`)."
);
return CacheLookup::Miss(None);
};
let ttl_secs = cache_cfg.ttl_secs.unwrap_or(300);
let namespaces = cache_cfg.namespaces.as_deref().unwrap_or(&[]);
let key = if namespaces.is_empty() {
key
} else {
cache_namespace::entry_key(&key)
};
let first = lookup_once(channel, cache, key, ttl_secs, namespaces).await;
let Some(flights) = cfg.cache_flights.as_ref() else {
return first;
};
let CacheLookup::Miss(Some(mut ctx)) = first else {
return first;
};
match flights.join(&ctx.key) {
Flight::Leader(guard) => {
ctx._flight = Some(Box::new(guard));
CacheLookup::Miss(Some(ctx))
}
Flight::Follower(mut done) => {
let wait = cfg
.parsed_config
.timeout_ms
.map_or(MAX_COALESCE_WAIT, |ms| {
std::time::Duration::from_millis(ms).min(MAX_COALESCE_WAIT)
});
let _ = tokio::time::timeout(wait, done.changed()).await;
let retry = lookup_once(channel, cache, ctx.key, ttl_secs, namespaces).await;
if matches!(retry, CacheLookup::Hit(_)) {
metrics::record_cache_coalesced(channel);
}
retry
}
}
}
const MAX_COALESCE_WAIT: std::time::Duration = std::time::Duration::from_secs(5);
async fn lookup_once(
channel: &str,
cache: &Arc<dyn CacheBackend>,
key: String,
ttl_secs: u64,
namespaces: &[String],
) -> CacheLookup {
if namespaces.is_empty() {
return match cache.get(&key).await {
Ok(Some(cached)) => {
metrics::record_cache_hit(channel);
CacheLookup::Hit(cached)
}
_ => {
metrics::record_cache_miss(channel);
CacheStoreCtx::miss(key, cache, ttl_secs, None)
}
};
}
let mut keys: Vec<String> = namespaces
.iter()
.map(|ns| cache_namespace::version_key(ns))
.collect();
keys.push(key);
let result = cache.get_many(&keys).await;
let key = keys.pop().unwrap_or_default();
let mut values = match result {
Ok(values) if values.len() == keys.len() + 1 => values,
other => {
if let Err(e) = other {
tracing::debug!(channel = %channel, error = %e, "Response-cache lookup failed");
}
metrics::record_cache_miss(channel);
return CacheLookup::Miss(None);
}
};
let entry = values.pop().flatten();
let versions: Option<Vec<i64>> = values
.iter()
.map(|raw| cache_namespace::parse_version(raw.as_deref()))
.collect();
let Some(versions) = versions else {
tracing::warn!(
channel = %channel,
"A response-cache namespace counter holds something other than an integer; \
bypassing the response cache"
);
metrics::record_cache_miss(channel);
return CacheLookup::Miss(None);
};
if let Some(body) = entry.and_then(|stored| cache_namespace::decode_entry(stored, &versions)) {
metrics::record_cache_hit(channel);
return CacheLookup::Hit(body);
}
metrics::record_cache_miss(channel);
CacheStoreCtx::miss(key, cache, ttl_secs, Some(versions))
}
#[derive(Default)]
pub struct CacheFlights(dashmap::DashMap<String, tokio::sync::watch::Receiver<()>>);
pub struct FlightGuard {
flights: Arc<CacheFlights>,
key: String,
_done: tokio::sync::watch::Sender<()>,
}
impl Drop for FlightGuard {
fn drop(&mut self) {
self.flights.0.remove(&self.key);
}
}
enum Flight {
Leader(FlightGuard),
Follower(tokio::sync::watch::Receiver<()>),
}
impl CacheFlights {
fn join(self: &Arc<Self>, key: &str) -> Flight {
use dashmap::mapref::entry::Entry;
if let Some(flight) = self.0.get(key) {
return Flight::Follower(flight.clone());
}
match self.0.entry(key.to_string()) {
Entry::Occupied(flight) => Flight::Follower(flight.get().clone()),
Entry::Vacant(slot) => {
let (done, waiting) = tokio::sync::watch::channel(());
slot.insert(waiting);
Flight::Leader(FlightGuard {
flights: self.clone(),
key: key.to_string(),
_done: done,
})
}
}
}
#[cfg(test)]
pub(crate) fn in_flight(&self) -> usize {
self.0.len()
}
}