use std::borrow::Cow;
use std::io::{Read, Seek, SeekFrom, Write};
use std::path::Path;
use std::time::{Duration, Instant};
use crc32fast::Hasher;
use futures::FutureExt;
use reqwest::{Request, Response};
use rkyv::util::AlignedVec;
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use tracing::{Instrument, Span, debug, info_span, instrument, trace, warn};
use zerocopy::byteorder::little_endian::{U32, U64};
use zerocopy::{FromBytes, Immutable, IntoBytes, KnownLayout, Unaligned};
use uv_cache::{CacheEntry, Freshness};
use uv_fastid::Id;
use uv_fs::write_atomic;
use uv_redacted::DisplaySafeUrl;
use crate::base_client::CertificateSource;
use crate::httpcache::{
AfterResponse, ArchivedCachePolicy, BeforeRequest, CachePolicy, CachePolicyBuilder,
};
use crate::{BaseClient, Error, ErrorKind, OwnedArchive, ProblemDetails, RetryState};
pub(crate) trait Cacheable: Sized {
type Target: Send + 'static;
fn from_aligned_bytes(bytes: AlignedVec) -> Result<Self::Target, Error>;
fn to_bytes(&self) -> Result<Cow<'_, [u8]>, Error>;
fn into_target(self) -> Self::Target;
}
#[derive(Debug, Deserialize, Serialize)]
#[serde(transparent)]
struct SerdeCacheable<T> {
inner: T,
}
impl<T: Serialize + DeserializeOwned + Send + 'static> Cacheable for SerdeCacheable<T> {
type Target = T;
fn from_aligned_bytes(bytes: AlignedVec) -> Result<T, Error> {
Ok(rmp_serde::from_slice::<T>(&bytes).map_err(ErrorKind::Decode)?)
}
fn to_bytes(&self) -> Result<Cow<'_, [u8]>, Error> {
Ok(Cow::from(
rmp_serde::to_vec(&self.inner).map_err(ErrorKind::Encode)?,
))
}
fn into_target(self) -> Self::Target {
self.inner
}
}
impl<A> Cacheable for OwnedArchive<A>
where
A: rkyv::Archive + for<'a> rkyv::Serialize<crate::rkyvutil::Serializer<'a>> + Send + 'static,
A::Archived: rkyv::Portable
+ rkyv::Deserialize<A, crate::rkyvutil::Deserializer>
+ for<'a> rkyv::bytecheck::CheckBytes<crate::rkyvutil::Validator<'a>>,
{
type Target = Self;
fn from_aligned_bytes(bytes: AlignedVec) -> Result<Self, Error> {
Self::new(bytes)
}
fn to_bytes(&self) -> Result<Cow<'_, [u8]>, Error> {
Ok(Cow::from(Self::as_bytes(self)))
}
fn into_target(self) -> Self::Target {
self
}
}
#[derive(Debug)]
pub enum CachedClientError<CallbackError: std::error::Error + 'static> {
Client(Error),
Callback {
retries: u32,
err: CallbackError,
duration: Duration,
},
}
impl<CallbackError: std::error::Error + 'static> CachedClientError<CallbackError> {
fn with_retries(self, retries: u32) -> Self {
match self {
Self::Client(err) => Self::Client(err.with_retries(retries)),
Self::Callback {
retries: _,
err,
duration,
} => Self::Callback {
retries,
err,
duration,
},
}
}
fn retries(&self) -> u32 {
match self {
Self::Client(err) => err.retries(),
Self::Callback { retries, .. } => *retries,
}
}
fn error(&self) -> &(dyn std::error::Error + 'static) {
match self {
Self::Client(err) => err,
Self::Callback { err, .. } => err,
}
}
}
impl<CallbackError: std::error::Error + 'static> From<Error> for CachedClientError<CallbackError> {
fn from(error: Error) -> Self {
Self::Client(error)
}
}
impl<CallbackError: std::error::Error + 'static> From<ErrorKind>
for CachedClientError<CallbackError>
{
fn from(error: ErrorKind) -> Self {
Self::Client(error.into())
}
}
impl<E: Into<Self> + std::error::Error + 'static> From<CachedClientError<E>> for Error {
fn from(error: CachedClientError<E>) -> Self {
match error {
CachedClientError::Client(error) => error,
CachedClientError::Callback {
retries,
err,
duration,
} => Self::new(err.into().into_kind(), retries, duration),
}
}
}
#[derive(Debug, Clone)]
pub enum CacheControl {
None,
MustRevalidate,
AllowStale,
Override(http::HeaderValue),
}
impl From<Freshness> for CacheControl {
fn from(value: Freshness) -> Self {
match value {
Freshness::Fresh => Self::None,
Freshness::Stale => Self::MustRevalidate,
Freshness::Missing => Self::None,
}
}
}
#[derive(Debug, Clone)]
pub struct CachedClient(BaseClient);
impl CachedClient {
pub fn new(client: BaseClient) -> Self {
Self(client)
}
pub fn uncached(&self) -> &BaseClient {
&self.0
}
pub(crate) fn certificate_source(&self) -> CertificateSource {
self.0.certificate_source()
}
#[instrument(skip_all)]
async fn get_cacheable<
Payload: Cacheable + 'static,
CallBackError: std::error::Error + 'static,
Callback: AsyncFnOnce(Response) -> Result<Payload, CallBackError>,
>(
&self,
req: Request,
cache_entry: &CacheEntry,
cache_control: CacheControl,
response_callback: Callback,
) -> Result<Payload::Target, CachedClientError<CallBackError>> {
let start = Instant::now();
if matches!(cache_control, CacheControl::AllowStale) {
let (req, cached) = self
.read_and_decode_stale_cache::<Payload>(req, cache_entry)
.await;
match cached {
Ok(Some(payload)) => return Ok(payload),
Ok(None) => warn!(
"Cached response doesn't match current request for: {}",
DisplaySafeUrl::from_url(req.url().clone())
),
Err(err) if err.is_file_not_exists() => {
trace!(
"No cache entry exists for `{}`",
cache_entry.path().display()
);
}
Err(err) => {
warn!(
"Broken cache entry at `{}`, removing: {err}",
cache_entry.path().display()
);
let _ = fs_err::tokio::remove_file(&cache_entry.path()).await;
}
}
let (response, cache_policy) = self.fresh_request(req, cache_control).await?;
return self
.run_response_callback(
cache_entry,
cache_policy,
start,
response,
response_callback,
)
.await;
}
let fresh_req = req.try_clone().expect("HTTP request must be cloneable");
let (req, cached) = self
.read_cache::<Payload>(req, cache_entry, cache_control.clone())
.await;
let cached_response = match cached {
Some(CachedEntry::Fresh(payload)) => {
return match payload {
Ok(payload) => Ok(payload),
Err(err) => {
warn!(
"Broken fresh cache entry (for payload) at `{}`, removing: {err}",
cache_entry.path().display()
);
self.resend_and_heal_cache(
fresh_req,
cache_entry,
cache_control,
response_callback,
)
.await
}
};
}
Some(CachedEntry::Stale {
cached,
new_cache_policy_builder,
}) => {
self.send_cached_handle_stale(
req,
cache_control.clone(),
cached,
*new_cache_policy_builder,
)
.boxed_local()
.await?
}
None => {
debug!(
"No cache entry for: {}",
DisplaySafeUrl::from_url(req.url().clone())
);
let (response, cache_policy) =
self.fresh_request(req, cache_control.clone()).await?;
CachedResponse::ModifiedOrNew {
response,
cache_policy,
}
}
};
match cached_response {
CachedResponse::NotModified { cached, new_policy } => {
let refresh_cache =
info_span!("refresh_cache", file = %cache_entry.path().display());
async {
let path = cache_entry.path().to_path_buf();
let span = Span::current();
let payload = tokio::task::spawn_blocking(move || {
span.in_scope(|| {
cached.refresh_policy(&path, &new_policy)?;
Ok::<_, Error>(Payload::from_aligned_bytes(cached.into_data()))
})
})
.await
.expect("cache refresh task panicked")?;
match payload {
Ok(payload) => Ok(payload),
Err(err) => {
warn!(
"Broken fresh cache entry after revalidation \
(for payload) at `{}`, removing: {err}",
cache_entry.path().display()
);
self.resend_and_heal_cache(
fresh_req,
cache_entry,
cache_control.clone(),
response_callback,
)
.await
}
}
}
.instrument(refresh_cache)
.await
}
CachedResponse::ModifiedOrNew {
response,
cache_policy,
} => {
if response.status() == http::StatusCode::NOT_MODIFIED {
warn!(
"Server returned unusable 304 for: {}",
DisplaySafeUrl::from_url(fresh_req.url().clone())
);
self.resend_and_heal_cache(
fresh_req,
cache_entry,
cache_control,
response_callback,
)
.await
} else {
self.run_response_callback(
cache_entry,
cache_policy,
start,
response,
response_callback,
)
.await
}
}
}
}
async fn skip_cache<
Payload: Serialize + DeserializeOwned + Send + 'static,
CallBackError: std::error::Error + 'static,
Callback: AsyncFnOnce(Response) -> Result<Payload, CallBackError>,
>(
&self,
req: Request,
cache_entry: &CacheEntry,
cache_control: CacheControl,
response_callback: Callback,
) -> Result<Payload, CachedClientError<CallBackError>> {
let start = Instant::now();
let (response, cache_policy) = self.fresh_request(req, cache_control).await?;
let payload = self
.run_response_callback(cache_entry, cache_policy, start, response, async |resp| {
let payload = response_callback(resp).await?;
Ok(SerdeCacheable { inner: payload })
})
.await?;
Ok(payload)
}
async fn resend_and_heal_cache<
Payload: Cacheable,
CallBackError: std::error::Error + 'static,
Callback: AsyncFnOnce(Response) -> Result<Payload, CallBackError>,
>(
&self,
req: Request,
cache_entry: &CacheEntry,
cache_control: CacheControl,
response_callback: Callback,
) -> Result<Payload::Target, CachedClientError<CallBackError>> {
let _ = fs_err::tokio::remove_file(&cache_entry.path()).await;
let start = Instant::now();
let (response, cache_policy) = self.fresh_request(req, cache_control).await?;
self.run_response_callback(
cache_entry,
cache_policy,
start,
response,
response_callback,
)
.await
}
async fn run_response_callback<
Payload: Cacheable,
CallBackError: std::error::Error + 'static,
Callback: AsyncFnOnce(Response) -> Result<Payload, CallBackError>,
>(
&self,
cache_entry: &CacheEntry,
cache_policy: Option<Box<CachePolicy>>,
start: Instant,
response: Response,
response_callback: Callback,
) -> Result<Payload::Target, CachedClientError<CallBackError>> {
let new_cache = info_span!("new_cache", file = %cache_entry.path().display());
let data = response_callback(response)
.boxed_local()
.await
.map_err(|err| CachedClientError::Callback {
retries: 0,
err,
duration: start.elapsed(),
})?;
let Some(cache_policy) = cache_policy else {
return Ok(data.into_target());
};
async {
fs_err::tokio::create_dir_all(cache_entry.dir())
.await
.map_err(ErrorKind::CacheWrite)?;
let data_with_cache_policy_bytes =
DataWithCachePolicy::serialize(&cache_policy, &data.to_bytes()?)?;
write_atomic(cache_entry.path(), data_with_cache_policy_bytes)
.await
.map_err(ErrorKind::CacheWrite)?;
Ok(data.into_target())
}
.instrument(new_cache)
.await
}
#[instrument(name = "read_and_parse_cache", skip_all, fields(file = %cache_entry.path().display()))]
async fn read_cache<Payload: Cacheable + 'static>(
&self,
mut req: Request,
cache_entry: &CacheEntry,
cache_control: CacheControl,
) -> (Request, Option<CachedEntry<Payload::Target>>) {
let path = cache_entry.path().to_path_buf();
let span = Span::current();
let (req, cached) = self
.0
.cache_read_runtime()
.spawn_blocking(move || {
span.in_scope(|| {
let cached = DataWithCachePolicy::from_path_sync(&path).map(|cached| {
if let CacheControl::MustRevalidate = cache_control {
req.headers_mut().insert(
http::header::CACHE_CONTROL,
http::HeaderValue::from_static("no-cache"),
);
}
let url = DisplaySafeUrl::from_url(req.url().clone());
match cached.cache_policy().before_request(&mut req) {
BeforeRequest::Fresh => {
debug!("Found fresh response for: {url}");
Some(CachedEntry::Fresh(Payload::from_aligned_bytes(
cached.into_data(),
)))
}
BeforeRequest::Stale(new_cache_policy_builder) => {
debug!("Found stale response for: {url}");
Some(CachedEntry::Stale {
cached,
new_cache_policy_builder: Box::new(new_cache_policy_builder),
})
}
BeforeRequest::NoMatch => {
warn!("Cached response doesn't match current request for: {url}");
None
}
}
});
(req, cached)
})
})
.await
.expect("cache read and payload decoding task panicked");
let cached = match cached {
Ok(cached) => cached,
Err(err) => {
if err.is_file_not_exists() {
trace!(
"No cache entry exists for `{}`",
cache_entry.path().display()
);
} else {
warn!(
"Broken cache policy entry at `{}`, removing: {err}",
cache_entry.path().display()
);
let _ = fs_err::tokio::remove_file(&cache_entry.path()).await;
}
None
}
};
(req, cached)
}
#[instrument(name = "read_and_decode_stale_cache", skip_all, fields(file = %cache_entry.path().display()))]
async fn read_and_decode_stale_cache<Payload: Cacheable + 'static>(
&self,
req: Request,
cache_entry: &CacheEntry,
) -> (Request, Result<Option<Payload::Target>, Error>) {
let path = cache_entry.path().to_path_buf();
self.0
.cache_read_runtime()
.spawn_blocking(move || {
let cached = DataWithCachePolicy::from_path_sync(&path).and_then(|cached| {
if cached.cache_policy().matches_stale_request(&req) {
Payload::from_aligned_bytes(cached.into_data()).map(Some)
} else {
Ok(None)
}
});
(req, cached)
})
.await
.expect("cache read and payload decoding task panicked")
}
async fn send_cached_handle_stale(
&self,
req: Request,
cache_control: CacheControl,
cached: DataWithCachePolicy,
new_cache_policy_builder: CachePolicyBuilder,
) -> Result<CachedResponse, Error> {
let url = DisplaySafeUrl::from_url(req.url().clone());
debug!("Sending revalidation request for: {url}");
let start = Instant::now();
let mut response = self
.0
.execute(req)
.instrument(info_span!("revalidation_request", url = %url))
.await
.map_err(|err| {
Error::from_reqwest_middleware(url.clone(), err, start, self.certificate_source())
})?;
trace!(
"Received response for revalidation request with status {} for: {}",
response.status(),
url
);
let retry_count = response
.extensions()
.get::<reqwest_retry::RetryCount>()
.map(|retries| retries.value());
if let Err(status_error) = response.error_for_status_ref() {
let problem_details = ProblemDetails::try_from_response(response).await;
return Err(Error::new(
ErrorKind::from_reqwest_with_problem_details(url, status_error, problem_details),
retry_count.unwrap_or_default(),
start.elapsed(),
));
}
if let CacheControl::Override(header) = &cache_control {
response
.headers_mut()
.insert(http::header::CACHE_CONTROL, header.clone());
}
match cached
.cache_policy()
.after_response(new_cache_policy_builder, &response)
{
AfterResponse::NotModified(new_policy) => {
debug!("Found not-modified response for: {url}");
Ok(CachedResponse::NotModified {
cached,
new_policy: Box::new(new_policy),
})
}
AfterResponse::Modified(new_policy) => {
debug!("Found modified response for: {url}");
Ok(CachedResponse::ModifiedOrNew {
response,
cache_policy: new_policy
.to_archived()
.is_storable()
.then(|| Box::new(new_policy)),
})
}
}
}
#[instrument(skip_all, fields(url = %DisplaySafeUrl::from_url(req.url().clone())))]
async fn fresh_request(
&self,
req: Request,
cache_control: CacheControl,
) -> Result<(Response, Option<Box<CachePolicy>>), Error> {
let url = DisplaySafeUrl::from_url(req.url().clone());
debug!("Sending fresh {} request for: {}", req.method(), url);
let cache_policy_builder = CachePolicyBuilder::new(&req);
let start = Instant::now();
let mut response = self.0.execute(req).await.map_err(|err| {
Error::from_reqwest_middleware(url.clone(), err, start, self.certificate_source())
})?;
trace!(
"Received response for fresh request with status {} for: {}",
response.status(),
url
);
if let CacheControl::Override(header) = &cache_control {
response
.headers_mut()
.insert(http::header::CACHE_CONTROL, header.clone());
}
let retry_count = response
.extensions()
.get::<reqwest_retry::RetryCount>()
.map(|retries| retries.value());
if let Err(status_error) = response.error_for_status_ref() {
let problem_details = ProblemDetails::try_from_response(response).await;
return Err(Error::new(
ErrorKind::from_reqwest_with_problem_details(url, status_error, problem_details),
retry_count.unwrap_or_default(),
start.elapsed(),
));
}
let cache_policy = cache_policy_builder.build(&response);
let cache_policy = if cache_policy.to_archived().is_storable() {
Some(Box::new(cache_policy))
} else {
None
};
Ok((response, cache_policy))
}
#[instrument(skip_all)]
pub async fn get_serde_with_retry<
Payload: Serialize + DeserializeOwned + Send + 'static,
CallBackError: std::error::Error + 'static,
Callback: AsyncFn(Response, &mut RetryState) -> Result<Payload, CallBackError>,
>(
&self,
req: Request,
cache_entry: &CacheEntry,
cache_control: CacheControl,
response_callback: Callback,
) -> Result<Payload, CachedClientError<CallBackError>> {
let payload = self
.get_cacheable_with_retry(
req,
cache_entry,
cache_control,
async |resp, retry_state| {
let payload = response_callback(resp, retry_state).await?;
Ok(SerdeCacheable { inner: payload })
},
)
.await?;
Ok(payload)
}
#[instrument(skip_all)]
pub(crate) async fn get_cacheable_with_retry<
Payload: Cacheable + 'static,
CallBackError: std::error::Error + 'static,
Callback: AsyncFn(Response, &mut RetryState) -> Result<Payload, CallBackError>,
>(
&self,
req: Request,
cache_entry: &CacheEntry,
cache_control: CacheControl,
response_callback: Callback,
) -> Result<Payload::Target, CachedClientError<CallBackError>> {
let mut retry_state = RetryState::start(self.uncached().retry_policy(), req.url().clone());
loop {
let fresh_req = req.try_clone().expect("HTTP request must be cloneable");
let result = self
.get_cacheable(
fresh_req,
cache_entry,
cache_control.clone(),
async |response| {
retry_state
.handle_response(response, &response_callback)
.await
},
)
.await;
match result {
Ok(ok) => return Ok(ok),
Err(err)
if let Some(backoff) = retry_state.should_retry(err.error(), err.retries()) =>
{
retry_state.sleep_backoff(backoff).await;
}
Err(err) => return Err(err.with_retries(retry_state.total_retries())),
}
}
}
pub async fn skip_cache_with_retry<
Payload: Serialize + DeserializeOwned + Send + 'static,
CallBackError: std::error::Error + 'static,
Callback: AsyncFn(Response, &mut RetryState) -> Result<Payload, CallBackError>,
>(
&self,
req: Request,
cache_entry: &CacheEntry,
cache_control: CacheControl,
response_callback: Callback,
) -> Result<Payload, CachedClientError<CallBackError>> {
let mut retry_state = RetryState::start(self.uncached().retry_policy(), req.url().clone());
loop {
let fresh_req = req.try_clone().expect("HTTP request must be cloneable");
let result = self
.skip_cache(
fresh_req,
cache_entry,
cache_control.clone(),
async |response| {
retry_state
.handle_response(response, &response_callback)
.await
},
)
.await;
match result {
Ok(ok) => return Ok(ok),
Err(err)
if let Some(backoff) = retry_state.should_retry(err.error(), err.retries()) =>
{
retry_state.sleep_backoff(backoff).await;
}
Err(err) => return Err(err.with_retries(retry_state.total_retries())),
}
}
}
}
#[derive(Debug)]
enum CachedEntry<Payload> {
Fresh(Result<Payload, Error>),
Stale {
cached: DataWithCachePolicy,
new_cache_policy_builder: Box<CachePolicyBuilder>,
},
}
#[derive(Debug)]
enum CachedResponse {
NotModified {
cached: DataWithCachePolicy,
new_policy: Box<CachePolicy>,
},
ModifiedOrNew {
response: Response,
cache_policy: Option<Box<CachePolicy>>,
},
}
#[derive(Debug)]
pub struct DataWithCachePolicy {
bytes: AlignedVec,
data_len: usize,
policy_start: usize,
entry_id: [u8; 16],
}
#[derive(Debug, FromBytes, Immutable, IntoBytes, KnownLayout, Unaligned)]
#[repr(C)]
struct CachePolicyFooter {
checksum: U32,
data_len: U64,
}
impl DataWithCachePolicy {
#[instrument]
fn from_path_sync(path: &Path) -> Result<Self, Error> {
let mut file = fs_err::File::open(path).map_err(ErrorKind::Io)?;
let file_size = file.metadata().map_err(ErrorKind::Io)?.len();
let file_size = usize::try_from(file_size)
.ok()
.filter(|&file_size| file_size <= AlignedVec::<16>::MAX_CAPACITY)
.ok_or_else(|| {
ErrorKind::Io(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"cache entry file size of {file_size} bytes exceeds the maximum supported \
size of {} bytes",
AlignedVec::<16>::MAX_CAPACITY,
),
))
})?;
let mut aligned_bytes = AlignedVec::with_capacity(file_size);
aligned_bytes.resize(file_size, 0);
file.read_exact(&mut aligned_bytes).map_err(ErrorKind::Io)?;
Self::from_aligned_bytes(aligned_bytes)
}
pub fn from_reader(mut rdr: impl std::io::Read) -> Result<Self, Error> {
let mut aligned_bytes = AlignedVec::new();
aligned_bytes
.extend_from_reader(&mut rdr)
.map_err(ErrorKind::Io)?;
Self::from_aligned_bytes(aligned_bytes)
}
fn from_aligned_bytes(mut bytes: AlignedVec) -> Result<Self, Error> {
let (contents, footer) = CachePolicyFooter::ref_from_suffix(&bytes).map_err(|_| {
ErrorKind::ArchiveRead("HTTP cache entry is shorter than its trailer".to_owned())
})?;
let data_len = usize::try_from(footer.data_len.get())
.map_err(|_| ErrorKind::ArchiveRead("invalid HTTP cache payload length".to_owned()))?;
let tail = contents.get(data_len..).ok_or_else(|| {
ErrorKind::ArchiveRead("invalid HTTP cache payload length".to_owned())
})?;
let padding_len = Self::padding_len(data_len);
let aligned_tail = tail.get(padding_len..).ok_or_else(|| {
ErrorKind::ArchiveRead("invalid HTTP cache payload length".to_owned())
})?;
let (entry_id, policy_bytes) = <[u8; 16]>::ref_from_prefix(aligned_tail)
.map_err(|_| ErrorKind::ArchiveRead("invalid HTTP cache payload length".to_owned()))?;
let mut hasher = Hasher::new();
hasher.update(tail);
hasher.update(footer.data_len.as_bytes());
if hasher.finalize() != footer.checksum.get() {
return Err(
ErrorKind::ArchiveRead("HTTP cache policy checksum mismatch".to_owned()).into(),
);
}
rkyv::access::<ArchivedCachePolicy, rkyv::rancor::Error>(policy_bytes)
.map_err(|err| ErrorKind::ArchiveRead(err.to_string()))?;
let entry_id = *entry_id;
let policy_start = contents.len() - policy_bytes.len();
let contents_len = contents.len();
bytes.resize(contents_len, 0);
Ok(Self {
bytes,
data_len,
policy_start,
entry_id,
})
}
fn serialize(policy: &CachePolicy, data: &[u8]) -> Result<Vec<u8>, Error> {
let entry_id = *Id::secure()
.as_bytes()
.as_array::<16>()
.expect("IDs have 16 bytes");
let tail = Self::serialize_policy(policy, &entry_id, data.len())?;
let entry_id_start = data.len() + Self::padding_len(data.len());
let mut bytes = Vec::with_capacity(entry_id_start + entry_id.len() + tail.len());
bytes.extend_from_slice(data);
bytes.resize(entry_id_start, 0);
bytes.extend_from_slice(&entry_id);
bytes.extend_from_slice(&tail);
Ok(bytes)
}
fn serialize_policy(
policy: &CachePolicy,
entry_id: &[u8; 16],
data_len: usize,
) -> Result<Vec<u8>, Error> {
let policy = OwnedArchive::from_unarchived(policy)?;
let policy = OwnedArchive::as_bytes(&policy);
let padding_len = Self::padding_len(data_len);
let data_len = U64::new(
u64::try_from(data_len).map_err(|err| ErrorKind::ArchiveWrite(err.to_string()))?,
);
let mut hasher = Hasher::new();
hasher.update(&[0; AlignedVec::<16>::ALIGNMENT][..padding_len]);
hasher.update(entry_id);
hasher.update(policy);
hasher.update(data_len.as_bytes());
let footer = CachePolicyFooter {
checksum: U32::new(hasher.finalize()),
data_len,
};
let mut bytes = Vec::with_capacity(policy.len() + size_of::<CachePolicyFooter>());
bytes.extend_from_slice(policy);
bytes.extend_from_slice(footer.as_bytes());
Ok(bytes)
}
pub fn into_data(mut self) -> AlignedVec {
self.bytes.resize(self.data_len, 0);
self.bytes
}
fn cache_policy(&self) -> &ArchivedCachePolicy {
#[expect(unsafe_code)]
unsafe {
rkyv::access_unchecked::<ArchivedCachePolicy>(&self.bytes[self.policy_start..])
}
}
fn padding_len(data_len: usize) -> usize {
(AlignedVec::<16>::ALIGNMENT - data_len % AlignedVec::<16>::ALIGNMENT)
% AlignedVec::<16>::ALIGNMENT
}
fn refresh_policy(&self, path: &Path, policy: &CachePolicy) -> Result<(), Error> {
let mut file = match fs_err::OpenOptions::new().read(true).write(true).open(path) {
Ok(file) => file,
Err(err) if err.kind() == std::io::ErrorKind::NotFound => return Ok(()),
Err(err)
if err.kind() == std::io::ErrorKind::PermissionDenied
|| err.kind() == std::io::ErrorKind::ReadOnlyFilesystem =>
{
debug!("Skipping cache policy update at {}: {err}", path.display());
return Ok(());
}
Err(err) => return Err(ErrorKind::CacheWrite(err).into()),
};
file.seek(SeekFrom::Start(
(self.policy_start - self.entry_id.len()) as u64,
))
.map_err(ErrorKind::CacheWrite)?;
let mut entry_id = [0; 16];
match file.read_exact(&mut entry_id) {
Ok(()) => {}
Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(()),
Err(err) => return Err(ErrorKind::CacheWrite(err).into()),
}
if entry_id != self.entry_id {
return Ok(());
}
let tail = Self::serialize_policy(policy, &entry_id, self.data_len)?;
file.write_all(&tail).map_err(ErrorKind::CacheWrite)?;
file.set_len((self.policy_start + tail.len()) as u64)
.map_err(ErrorKind::CacheWrite)?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use anyhow::Result;
use rkyv::util::AlignedVec;
use zerocopy::IntoBytes;
use zerocopy::byteorder::little_endian::{U32, U64};
use super::{
ArchivedCachePolicy, CachePolicy, CachePolicyBuilder, CachePolicyFooter,
DataWithCachePolicy, Hasher, OwnedArchive,
};
fn policy(vary: &str) -> Result<CachePolicy> {
let request = reqwest::Request::new(http::Method::GET, "https://example.com/".parse()?);
let response = reqwest::Response::from(
http::Response::builder()
.header("cache-control", "public, max-age=3600")
.header("vary", vary)
.body(Vec::new())?,
);
Ok(CachePolicyBuilder::new(&request).build(&response))
}
#[test]
fn payload_and_policy_share_an_aligned_buffer() -> Result<()> {
let policy = policy("x-small")?;
for data_len in 0..32 {
let payload = vec![42; data_len];
let serialized = DataWithCachePolicy::serialize(&policy, &payload)?;
let mut bytes = AlignedVec::with_capacity(serialized.len());
bytes.extend_from_slice(&serialized);
let allocation = bytes.as_ptr();
let cached = DataWithCachePolicy::from_aligned_bytes(bytes)?;
assert_eq!(cached.policy_start % AlignedVec::<16>::ALIGNMENT, 0);
assert_eq!(
&cached.bytes[cached.policy_start..],
OwnedArchive::as_bytes(&policy.to_archived()),
);
let data = cached.into_data();
assert_eq!(data.as_ptr(), allocation);
assert_eq!(data.as_slice(), payload);
}
Ok(())
}
#[test]
fn policy_pointers_cannot_reference_the_payload() -> Result<()> {
let policy = policy(&format!("x-{}", "large".repeat(20)))?;
let serialized = DataWithCachePolicy::serialize(&policy, b"")?;
let mut bytes = AlignedVec::new();
bytes.extend_from_slice(&serialized);
let footer_start = bytes.len() - size_of::<CachePolicyFooter>();
let data_len = U64::new(16);
let mut hasher = Hasher::new();
hasher.update(&bytes[usize::try_from(data_len.get())?..footer_start]);
hasher.update(data_len.as_bytes());
let footer = CachePolicyFooter {
checksum: U32::new(hasher.finalize()),
data_len,
};
bytes[footer_start..].copy_from_slice(footer.as_bytes());
assert!(
rkyv::access::<ArchivedCachePolicy, rkyv::rancor::Error>(&bytes[..footer_start])
.is_ok()
);
assert!(DataWithCachePolicy::from_aligned_bytes(bytes).is_err());
Ok(())
}
#[test]
fn policy_can_grow_and_shrink_in_place() -> Result<()> {
let directory = tempfile::tempdir()?;
let path = directory.path().join("response");
let small = policy("x-small")?;
let large = policy(&format!("x-{}", "large".repeat(1000)))?;
let payload = vec![42; 1024 * 1024 + 7];
let original = DataWithCachePolicy::serialize(&small, &payload)?;
fs_err::write(&path, &original)?;
let cached = DataWithCachePolicy::from_path_sync(&path)?;
cached.refresh_policy(&path, &large)?;
assert!(fs_err::metadata(&path)?.len() > original.len() as u64);
let grown = DataWithCachePolicy::from_path_sync(&path)?;
assert_eq!(&grown.bytes[..grown.data_len], payload);
assert_eq!(grown.entry_id, cached.entry_id);
assert_eq!(
&grown.bytes[grown.policy_start..],
OwnedArchive::as_bytes(&large.to_archived()),
);
grown.refresh_policy(&path, &small)?;
assert_eq!(fs_err::read(&path)?, original);
Ok(())
}
#[test]
fn interrupted_policy_writes_never_accept_a_mixed_policy() -> Result<()> {
let small = policy("x-small")?;
let large = policy(&format!("x-{}", "large".repeat(20)))?;
for (before, after) in [(&small, &large), (&large, &small)] {
let original = DataWithCachePolicy::serialize(before, b"payload")?;
let cached = DataWithCachePolicy::from_reader(original.as_slice())?;
let tail =
DataWithCachePolicy::serialize_policy(after, &cached.entry_id, cached.data_len)?;
let offset = cached.policy_start;
let before = before.to_archived();
let after = after.to_archived();
for written in 0..=tail.len() {
let mut interrupted = original.clone();
interrupted.resize(interrupted.len().max(offset + written), 0);
interrupted[offset..offset + written].copy_from_slice(&tail[..written]);
if let Ok(read) = DataWithCachePolicy::from_reader(interrupted.as_slice()) {
assert_eq!(&read.bytes[..read.data_len], b"payload");
let policy = &read.bytes[read.policy_start..];
assert!(
policy == OwnedArchive::as_bytes(&before)
|| policy == OwnedArchive::as_bytes(&after)
);
}
}
}
Ok(())
}
#[test]
fn cache_trailer_corruption_is_rejected() -> Result<()> {
let bytes = DataWithCachePolicy::serialize(&policy("x-small")?, b"payload")?;
for index in b"payload".len()..bytes.len() {
let mut corrupted = bytes.clone();
corrupted[index] ^= 1;
assert!(DataWithCachePolicy::from_reader(corrupted.as_slice()).is_err());
}
Ok(())
}
}