use crate::{CacheKey, CachePolicy, CacheStorage, PutHandle, StoredEntry, tee::TeeingReader};
use futures_lite::{AsyncReadExt, AsyncWrite, AsyncWriteExt};
use std::{
fmt::{self, Debug, Formatter},
io,
pin::Pin,
task::{Context, Poll},
};
use trillium_http::{Body, Headers};
use trillium_server_common::{Runtime, RuntimeTrait};
pub struct TieredStorage<Hot, Cold> {
hot: Hot,
cold: Cold,
runtime: Runtime,
}
impl<Hot, Cold> TieredStorage<Hot, Cold> {
pub fn new(hot: Hot, cold: Cold, runtime: impl RuntimeTrait) -> Self {
Self {
hot,
cold,
runtime: runtime.into(),
}
}
pub fn hot(&self) -> &Hot {
&self.hot
}
pub fn cold(&self) -> &Cold {
&self.cold
}
}
impl<Hot: Clone, Cold: Clone> Clone for TieredStorage<Hot, Cold> {
fn clone(&self) -> Self {
Self {
hot: self.hot.clone(),
cold: self.cold.clone(),
runtime: self.runtime.clone(),
}
}
}
impl<Hot: CacheStorage, Cold: CacheStorage> Debug for TieredStorage<Hot, Cold> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("TieredStorage")
.field("hot", &self.hot)
.field("cold", &self.cold)
.finish_non_exhaustive()
}
}
impl<Hot, Cold> CacheStorage for TieredStorage<Hot, Cold>
where
Hot: CacheStorage + Clone,
Cold: CacheStorage + Clone,
{
type StoredEntry = TieredEntry<Hot, Cold>;
type PutHandle = TieredPutHandle<Hot, Cold>;
async fn get(&self, key: &CacheKey) -> Vec<Self::StoredEntry> {
let hot = self.hot.get(key).await;
if !hot.is_empty() {
return hot.into_iter().map(TieredEntry::Hot).collect();
}
self.cold
.get(key)
.await
.into_iter()
.map(|entry| TieredEntry::Cold {
entry,
hot: self.hot.clone(),
key: key.clone(),
})
.collect()
}
async fn put(&self, key: CacheKey, policy: CachePolicy) -> io::Result<Self::PutHandle> {
let hot = self.hot.put(key.clone(), policy.clone()).await?;
Ok(TieredPutHandle {
hot,
hot_store: self.hot.clone(),
cold: self.cold.clone(),
runtime: self.runtime.clone(),
key,
policy,
})
}
async fn invalidate(&self, key: &CacheKey) {
self.hot.invalidate(key).await;
self.cold.invalidate(key).await;
}
}
pub enum TieredEntry<Hot: CacheStorage, Cold: CacheStorage> {
Hot(Hot::StoredEntry),
Cold {
entry: Cold::StoredEntry,
hot: Hot,
key: CacheKey,
},
}
impl<Hot, Cold> Clone for TieredEntry<Hot, Cold>
where
Hot: CacheStorage + Clone,
Cold: CacheStorage,
{
fn clone(&self) -> Self {
match self {
Self::Hot(entry) => Self::Hot(entry.clone()),
Self::Cold { entry, hot, key } => Self::Cold {
entry: entry.clone(),
hot: hot.clone(),
key: key.clone(),
},
}
}
}
impl<Hot: CacheStorage, Cold: CacheStorage> Debug for TieredEntry<Hot, Cold> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
match self {
Self::Hot(entry) => f.debug_tuple("Hot").field(entry).finish(),
Self::Cold { entry, key, .. } => f
.debug_struct("Cold")
.field("entry", entry)
.field("key", key)
.finish_non_exhaustive(),
}
}
}
impl<Hot, Cold> StoredEntry for TieredEntry<Hot, Cold>
where
Hot: CacheStorage + Clone,
Cold: CacheStorage,
{
fn policy(&self) -> &CachePolicy {
match self {
Self::Hot(entry) => entry.policy(),
Self::Cold { entry, .. } => entry.policy(),
}
}
async fn refresh_policy(&mut self, new_policy: CachePolicy) -> io::Result<()> {
match self {
Self::Hot(entry) => entry.refresh_policy(new_policy).await,
Self::Cold { entry, .. } => entry.refresh_policy(new_policy).await,
}
}
async fn open(self) -> io::Result<Body> {
match self {
Self::Hot(entry) => entry.open().await,
Self::Cold { entry, hot, key } => {
let policy = entry.policy().clone();
let cold_body = entry.open().await?;
let len = cold_body.len();
match hot.put(key, policy).await {
Ok(put_handle) => {
let tee = TeeingReader::new(cold_body, put_handle, u64::MAX);
Ok(Body::new_with_trailers(tee, len))
}
Err(e) => {
log::warn!("cache: promotion put failed: {e}, serving cold entry only");
Ok(cold_body)
}
}
}
}
}
}
pub struct TieredPutHandle<Hot: CacheStorage, Cold: CacheStorage> {
hot: Hot::PutHandle,
hot_store: Hot,
cold: Cold,
runtime: Runtime,
key: CacheKey,
policy: CachePolicy,
}
impl<Hot: CacheStorage, Cold: CacheStorage> Debug for TieredPutHandle<Hot, Cold> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("TieredPutHandle")
.field("key", &self.key)
.finish_non_exhaustive()
}
}
impl<Hot: CacheStorage, Cold: CacheStorage> Unpin for TieredPutHandle<Hot, Cold> {}
impl<Hot: CacheStorage, Cold: CacheStorage> AsyncWrite for TieredPutHandle<Hot, Cold> {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut self.get_mut().hot).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().hot).poll_flush(cx)
}
fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().hot).poll_close(cx)
}
}
impl<Hot, Cold> PutHandle for TieredPutHandle<Hot, Cold>
where
Hot: CacheStorage + Clone,
Cold: CacheStorage,
{
async fn finalize(self, trailers: Option<Headers>) -> io::Result<()> {
let Self {
hot,
hot_store,
cold,
runtime,
key,
policy,
} = self;
hot.finalize(trailers).await?;
let log_key = key.clone();
let _detached = runtime.spawn(async move {
if let Err(e) = flush_to_cold(hot_store, cold, key, policy).await {
log::warn!("cache: tiered background flush to cold failed for {log_key}: {e}");
}
});
Ok(())
}
}
async fn flush_to_cold<Hot, Cold>(
hot_store: Hot,
cold: Cold,
key: CacheKey,
policy: CachePolicy,
) -> io::Result<()>
where
Hot: CacheStorage,
Cold: CacheStorage,
{
let Some(entry) = hot_store
.get(&key)
.await
.into_iter()
.find(|entry| entry.policy().same_variant_as(&policy))
else {
return Ok(());
};
let mut body = entry.open().await?;
let mut put = cold.put(key, policy).await?;
let mut buf = [0u8; 8192];
loop {
let n = body.read(&mut buf).await?;
if n == 0 {
break;
}
put.write_all(&buf[..n]).await?;
}
let trailers = body.trailers();
put.finalize(trailers).await
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{InMemoryStorage, test_helpers::*};
use std::time::{Duration, SystemTime};
use trillium_http::{KnownHeaderName::*, Method, Status};
use trillium_testing::{TestResult, harness, runtime, test};
fn key() -> CacheKey {
CacheKey::new(Method::Get, "http://example.com/".parse().unwrap())
}
fn tiered() -> TieredStorage<InMemoryStorage, InMemoryStorage> {
TieredStorage::new(InMemoryStorage::new(), InMemoryStorage::new(), runtime())
}
async fn store_into(storage: &impl CacheStorage, key: CacheKey, body: &[u8]) {
let conn = exchange(
Method::Get,
&[],
Status::Ok,
&[(CacheControl, "max-age=600")],
);
let policy = policy_from(&conn, SystemTime::now(), private_cache());
let mut handle = storage.put(key, policy).await.unwrap();
handle.write_all(body).await.unwrap();
handle.finalize(None).await.unwrap();
}
async fn read_body(entry: impl StoredEntry) -> Vec<u8> {
let mut body = entry.open().await.unwrap();
let mut buf = Vec::new();
body.read_to_end(&mut buf).await.unwrap();
buf
}
async fn cold_settles<Hot, Cold>(
storage: &TieredStorage<Hot, Cold>,
key: &CacheKey,
) -> Vec<Cold::StoredEntry>
where
Hot: CacheStorage + Clone,
Cold: CacheStorage + Clone,
{
for _ in 0..200 {
let entries = storage.cold().get(key).await;
if !entries.is_empty() {
return entries;
}
storage.runtime.delay(Duration::from_millis(5)).await;
}
panic!("cold tier never populated");
}
#[test(harness)]
async fn hot_populated_synchronously_cold_written_back() -> TestResult {
let storage = tiered();
store_into(&storage, key(), b"hello").await;
let entries = storage.get(&key()).await;
assert_eq!(entries.len(), 1);
assert!(matches!(entries[0], TieredEntry::Hot(_)));
assert_eq!(read_body(entries[0].clone()).await, b"hello");
let cold = cold_settles(&storage, &key()).await;
assert_eq!(cold.len(), 1);
assert_eq!(read_body(cold[0].clone()).await, b"hello");
Ok(())
}
#[test(harness)]
async fn cold_hit_promotes_into_hot() -> TestResult {
let storage = tiered();
store_into(storage.cold(), key(), b"promoted").await;
assert!(storage.hot().get(&key()).await.is_empty());
let entries = storage.get(&key()).await;
assert_eq!(entries.len(), 1);
assert!(matches!(entries[0], TieredEntry::Cold { .. }));
assert_eq!(read_body(entries[0].clone()).await, b"promoted");
let hot = storage.hot().get(&key()).await;
assert_eq!(hot.len(), 1);
assert_eq!(read_body(hot[0].clone()).await, b"promoted");
Ok(())
}
#[test(harness)]
async fn invalidate_clears_both_tiers() -> TestResult {
let storage = tiered();
store_into(&storage, key(), b"x").await;
cold_settles(&storage, &key()).await;
storage.invalidate(&key()).await;
assert!(storage.get(&key()).await.is_empty());
assert!(storage.hot().get(&key()).await.is_empty());
assert!(storage.cold().get(&key()).await.is_empty());
Ok(())
}
#[test(harness)]
async fn drop_put_handle_without_finalize_stores_nothing() -> TestResult {
let storage = tiered();
let conn = exchange(
Method::Get,
&[],
Status::Ok,
&[(CacheControl, "max-age=600")],
);
let policy = policy_from(&conn, SystemTime::now(), private_cache());
let mut handle = storage.put(key(), policy).await.unwrap();
handle.write_all(b"partial").await.unwrap();
drop(handle);
assert!(storage.hot().get(&key()).await.is_empty());
assert!(storage.cold().get(&key()).await.is_empty());
Ok(())
}
#[cfg(feature = "fs")]
#[test(harness)]
async fn memory_over_filesystem_promotes_from_disk() -> TestResult {
use crate::FileSystemStorage;
let dir = tempfile::tempdir().unwrap();
{
let storage = TieredStorage::new(
InMemoryStorage::new(),
FileSystemStorage::new(dir.path()),
runtime(),
);
store_into(&storage, key(), b"on-disk").await;
cold_settles(&storage, &key()).await;
}
let reopened = TieredStorage::new(
InMemoryStorage::new(),
FileSystemStorage::new(dir.path()),
runtime(),
);
assert!(reopened.hot().get(&key()).await.is_empty());
let entries = reopened.get(&key()).await;
assert_eq!(entries.len(), 1);
assert!(matches!(entries[0], TieredEntry::Cold { .. }));
assert_eq!(read_body(entries[0].clone()).await, b"on-disk");
let hot = reopened.hot().get(&key()).await;
assert_eq!(hot.len(), 1);
assert_eq!(read_body(hot[0].clone()).await, b"on-disk");
Ok(())
}
}