use crate::cache::ArticleAvailability;
use crate::cache::ttl::CacheTier;
use crate::protocol::{
RequestCacheEntryMetadata, RequestCacheStatus, RequestContext, RequestKind,
RequestResponseMetadata, ResponseWireLen, StatusCode,
};
use crate::router::BackendSelector;
use crate::session::{ClientSession, precheck};
use crate::types::{BackendId, BackendToClientBytes, MessageId};
use anyhow::Result;
use std::sync::Arc;
use tokio::io::{AsyncWrite, AsyncWriteExt};
use tracing::debug;
pub(super) enum CacheLookupResult {
Hit,
PartialHit,
Miss,
}
impl ClientSession {
pub(super) async fn try_serve_from_cache<W>(
&self,
request: &mut RequestContext,
router: &Arc<BackendSelector>,
client_write: &mut W,
backend_to_client_bytes: &mut BackendToClientBytes,
) -> Result<CacheLookupResult>
where
W: AsyncWrite + Unpin,
{
let Some(msg_id_for_lookup) = request.message_id() else {
request.record_cache_status(RequestCacheStatus::Miss);
return Ok(CacheLookupResult::Miss);
};
debug!(
"Client {} checking cache for {}",
self.client_addr, msg_id_for_lookup
);
let Some(cached) = self.cache.get_request_message_id(msg_id_for_lookup).await else {
debug!("Cache MISS for message-ID: {}", msg_id_for_lookup);
request.record_cache_status(RequestCacheStatus::Miss);
return Ok(CacheLookupResult::Miss);
};
let availability = cached.availability();
request.record_cache_entry_metadata(cache_entry_metadata(&cached, &availability));
debug!(
"Client {} cache HIT for {} (cache_articles={})",
self.client_addr,
request.message_id().unwrap_or("<invalid>"),
self.cache_articles
);
if !self.cache_articles {
if self.adaptive_precheck && request.is_stat() {
precheck::spawn_background_precheck(
self.precheck_deps(router),
request.clone(),
MessageId::from_borrowed(
request
.message_id()
.expect("cached request still has message id"),
)
.expect("cached request has validated message id")
.to_owned(),
);
}
request.record_cache_status(RequestCacheStatus::PartialHit);
return Ok(CacheLookupResult::PartialHit);
}
if !request.is_stat() && !cached.is_complete_article() {
debug!(
"Client {} cache entry for {} has no complete payload (payload_len={}), fetching full article",
self.client_addr,
request.message_id().unwrap_or("<invalid>"),
cached.payload_len().get()
);
request.record_cache_status(RequestCacheStatus::PartialHit);
return Ok(CacheLookupResult::PartialHit);
}
let request_kind = request.kind();
let msg_id_for_write = request
.message_id()
.expect("cached request still has message id");
let Some(write) =
write_cached_article_response(client_write, &cached, request_kind, msg_id_for_write)
.await?
else {
let status_code = cached.status_code().as_u16();
debug!(
"Client {} cached response (code={}) can't serve request kind {:?}",
self.client_addr, status_code, request_kind
);
request.record_cache_status(RequestCacheStatus::PartialHit);
return Ok(CacheLookupResult::PartialHit);
};
*backend_to_client_bytes = backend_to_client_bytes.add(write.wire_len.get());
request.record_cache_response(write.metadata());
Ok(CacheLookupResult::Hit)
}
pub(super) fn spawn_cache_upsert_buffer(
&self,
msg_id: &crate::types::MessageId<'_>,
buffer: crate::cache::CacheIngestResponse,
backend: BackendId,
tier: CacheTier,
) {
if !self.cache.stores_payload_responses() {
return;
}
let cache_clone = self.cache.clone();
let msg_id_owned = msg_id.to_owned();
tokio::spawn(async move {
cache_clone
.upsert_ingest(msg_id_owned, buffer, backend, tier)
.await;
});
}
pub(super) fn spawn_cache_upsert_availability(
&self,
msg_id: &crate::types::MessageId<'_>,
status_code: StatusCode,
backend: BackendId,
tier: CacheTier,
) {
if !self.cache.records_backend_has_status() {
return;
}
let cache_clone = self.cache.clone();
let msg_id_owned = msg_id.to_owned();
tokio::spawn(async move {
cache_clone
.record_backend_has_status(msg_id_owned, status_code, backend, tier)
.await;
});
}
pub(super) fn tier_for_backend(&self, backend_id: BackendId) -> CacheTier {
self.router
.as_ref()
.and_then(|r| r.get_tier(backend_id))
.unwrap_or(0)
.into()
}
pub(super) const fn precheck_deps<'a>(
&'a self,
router: &'a Arc<BackendSelector>,
) -> precheck::PrecheckDeps<'a> {
precheck::PrecheckDeps {
router,
cache: &self.cache,
buffer_pool: &self.buffer_pool,
metrics: &self.metrics,
cache_articles: self.cache_articles,
}
}
}
fn cache_entry_metadata(
cached: &crate::cache::CachedArticle,
availability: &ArticleAvailability,
) -> RequestCacheEntryMetadata {
cached.request_cache_metadata(availability)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct CachedResponseWrite {
pub status: StatusCode,
pub wire_len: ResponseWireLen,
}
impl CachedResponseWrite {
#[must_use]
pub const fn metadata(self) -> RequestResponseMetadata {
RequestResponseMetadata::new(self.status, self.wire_len)
}
}
pub(super) async fn write_cached_article_response<W>(
client_write: &mut W,
cached: &crate::cache::CachedArticle,
request_kind: RequestKind,
message_id: &str,
) -> std::io::Result<Option<CachedResponseWrite>>
where
W: AsyncWrite + Unpin,
{
let Some(response) = cached.cached_response_for(request_kind, message_id) else {
return Ok(None);
};
let wire_len = response.wire_len();
response.write_to(client_write).await?;
client_write.flush().await?;
Ok(Some(CachedResponseWrite {
status: response.status(),
wire_len,
}))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::auth::AuthHandler;
use crate::cache::UnifiedCache;
use crate::metrics::MetricsCollector;
use crate::pool::{BufferPool, DeadpoolConnectionProvider};
use crate::protocol::{
RequestCacheArticleNumber, RequestCachePayloadKind, RequestCacheTimestampMillis,
};
use crate::types::{BufferSize, ClientAddress, ServerName};
use std::net::SocketAddr;
use std::time::Duration;
use tokio::io::AsyncReadExt;
use tokio::net::{TcpListener, TcpStream};
fn test_session() -> ClientSession {
let addr: SocketAddr = "127.0.0.1:0".parse().expect("valid address");
ClientSession::builder(
ClientAddress::from(addr),
BufferPool::new(BufferSize::try_new(1024).expect("valid buffer size"), 1),
Arc::new(AuthHandler::new(None, None).expect("auth disabled")),
MetricsCollector::new(1),
)
.with_cache(Arc::new(UnifiedCache::memory(
1024,
Duration::from_secs(60),
)))
.with_cache_articles(true)
.build()
}
async fn tcp_write_pair() -> (TcpStream, TcpStream) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind listener");
let addr = listener.local_addr().expect("listener address");
let connect = TcpStream::connect(addr);
let accept = listener.accept();
let (client, server) = tokio::join!(connect, accept);
(
client.expect("connect client"),
server.expect("accept client").0,
)
}
fn request_context(line: &[u8]) -> RequestContext {
RequestContext::parse(line).expect("valid request line")
}
#[tokio::test]
async fn cache_miss_is_recorded_on_request_context() {
let session = test_session();
let router = Arc::new(BackendSelector::new());
let mut metrics = BackendToClientBytes::zero();
let (mut client, _server) = tcp_write_pair().await;
let (_read, mut write) = client.split();
let mut request = request_context(b"ARTICLE <missing@example>\r\n");
let result = session
.try_serve_from_cache(&mut request, &router, &mut write, &mut metrics)
.await
.expect("lookup succeeds");
assert!(matches!(result, CacheLookupResult::Miss));
assert_eq!(request.cache_status(), Some(RequestCacheStatus::Miss));
assert_eq!(metrics, BackendToClientBytes::zero());
}
#[tokio::test]
async fn request_without_message_id_records_cache_miss() {
let session = test_session();
let router = Arc::new(BackendSelector::new());
let mut metrics = BackendToClientBytes::zero();
let (mut client, _server) = tcp_write_pair().await;
let (_read, mut write) = client.split();
let mut request = request_context(b"DATE\r\n");
let result = session
.try_serve_from_cache(&mut request, &router, &mut write, &mut metrics)
.await
.expect("lookup succeeds");
assert!(matches!(result, CacheLookupResult::Miss));
assert_eq!(request.cache_status(), Some(RequestCacheStatus::Miss));
assert_eq!(metrics, BackendToClientBytes::zero());
}
#[tokio::test]
async fn cache_hit_records_response_metadata_on_request_context() {
let session = test_session();
let msg_id = MessageId::new("<hit@example>".to_string()).expect("valid message id");
let expected = b"220 0 <hit@example>\r\nHeader: v\r\n\r\nBody\r\n.\r\n";
session
.cache
.upsert_ingest(
msg_id.clone(),
expected.to_vec(),
BackendId::from_index(0),
0.into(),
)
.await;
let router = Arc::new(BackendSelector::new());
let mut metrics = BackendToClientBytes::zero();
let (mut client, mut server) = tcp_write_pair().await;
let (_read, mut write) = client.split();
let mut request = request_context(b"ARTICLE <hit@example>\r\n");
let result = session
.try_serve_from_cache(&mut request, &router, &mut write, &mut metrics)
.await
.expect("lookup succeeds");
let mut written = vec![0; expected.len()];
server
.read_exact(&mut written)
.await
.expect("cached response written");
assert!(matches!(result, CacheLookupResult::Hit));
assert_eq!(written, expected);
assert_eq!(request.cache_status(), Some(RequestCacheStatus::Hit));
assert_eq!(request.backend_id(), None);
assert_eq!(request.response_status(), Some(StatusCode::new(220)));
assert_eq!(
request
.cache_availability()
.expect("cache hit records cache metadata")
.missing_bits(),
0
);
assert!(!request.cache_records_backend_has_article(BackendId::from_index(0)));
assert_eq!(
request.cache_article_number(),
Some(RequestCacheArticleNumber::new(0))
);
assert_eq!(
request.response_wire_len(),
Some(ResponseWireLen::new(expected.len()))
);
assert_eq!(metrics, BackendToClientBytes::zero().add(expected.len()));
}
#[tokio::test]
async fn partial_cache_hit_records_availability_on_request_context() {
let session = test_session();
let msg_id = MessageId::new("<partial@example>".to_string()).expect("valid message id");
session
.cache
.record_backend_missing(msg_id.clone(), BackendId::from_index(0))
.await;
let expected_timestamp = session
.cache
.get(&msg_id)
.await
.expect("cached availability entry")
.inserted_at();
let mut router = BackendSelector::new();
router.add_backend(
ServerName::try_new("partial-backend".to_string()).expect("server name"),
DeadpoolConnectionProvider::new(
"127.0.0.1".to_string(),
119,
"partial-backend".to_string(),
1,
None,
None,
),
0,
);
let router = Arc::new(router);
let mut metrics = BackendToClientBytes::zero();
let (mut client, _server) = tcp_write_pair().await;
let (_read, mut write) = client.split();
let mut request = request_context(b"ARTICLE <partial@example>\r\n");
let result = session
.try_serve_from_cache(&mut request, &router, &mut write, &mut metrics)
.await
.expect("lookup succeeds");
assert!(matches!(result, CacheLookupResult::PartialHit));
assert_eq!(request.cache_status(), Some(RequestCacheStatus::PartialHit));
assert_eq!(request.cache_entry_status(), Some(StatusCode::new(430)));
assert_eq!(
request.cache_entry_tier(),
Some(crate::protocol::RequestCacheTier::new(0))
);
assert_eq!(
request.cache_entry_timestamp(),
Some(RequestCacheTimestampMillis::new(expected_timestamp.get()))
);
assert_eq!(
request.cache_payload_kind(),
Some(RequestCachePayloadKind::Missing)
);
assert_eq!(request.cache_article_number(), None);
assert_eq!(
request.cache_availability().map(|availability| (
availability.missing_bits(),
availability.backend_has_article(BackendId::from_index(0))
)),
Some((0b0000_0001, false))
);
assert_eq!(metrics, BackendToClientBytes::zero());
}
}