use core::fmt;
use super::RequestIdPolicy;
use crate::rate_limit::RateLimit;
use crate::transport::{
MediaType, ResponseBuffer, ResponseContentType, ResponseDecodeWorkspace, ResponseWriterError,
RetainedMetadataError, RetainedResponseMetadata, StatusCode, TransportResponse,
};
#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub enum ResponseBodyPolicy {
Required,
Optional,
Forbidden,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ContentTypePolicy {
Required(&'static [MediaType<'static>]),
Optional(&'static [MediaType<'static>]),
Forbidden,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ResponsePolicyValidationError {
MissingSuccessStatus,
NonSuccessStatus,
DuplicateSuccessStatus,
MissingAcceptedMediaType,
DuplicateAcceptedMediaType,
RequiredBodyHasZeroLimit,
ForbiddenBodyHasNonzeroLimit,
ForbiddenBodyAllowsContentType,
RequiredBodyForbidsContentType,
}
impl_static_error!(ResponsePolicyValidationError,
Self::MissingSuccessStatus => "response policy has no success status",
Self::NonSuccessStatus => "response policy contains a non-success status",
Self::DuplicateSuccessStatus => "response policy contains duplicate statuses",
Self::MissingAcceptedMediaType => "response policy has no accepted media type",
Self::DuplicateAcceptedMediaType => "response policy contains duplicate media types",
Self::RequiredBodyHasZeroLimit => "required response body has a zero limit",
Self::ForbiddenBodyHasNonzeroLimit => "forbidden response body has a nonzero limit",
Self::ForbiddenBodyAllowsContentType => "forbidden response body allows a content type",
Self::RequiredBodyForbidsContentType => "required response body forbids its content type",
);
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ResponsePolicyError {
UnexpectedStatus,
BodyTooLarge,
MissingBody,
ForbiddenBody,
MissingContentType,
InvalidContentType,
UnexpectedContentType,
ForbiddenContentType,
UncommittedResponse,
InvalidRequestId,
}
impl_static_error!(ResponsePolicyError,
Self::UnexpectedStatus => "response status is not expected",
Self::BodyTooLarge => "response body exceeds the operation limit",
Self::MissingBody => "required response body is missing",
Self::ForbiddenBody => "response body is forbidden",
Self::MissingContentType => "required response content type is missing",
Self::InvalidContentType => "response content type is invalid",
Self::UnexpectedContentType => "response content type is not accepted",
Self::ForbiddenContentType => "response content type is forbidden",
Self::UncommittedResponse => "response writer is not committed",
Self::InvalidRequestId => "response request identifier is invalid",
);
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ResponsePolicy {
success_statuses: &'static [StatusCode],
content_type: ContentTypePolicy,
body: ResponseBodyPolicy,
max_body_bytes: usize,
}
impl ResponsePolicy {
pub fn new(
success_statuses: &'static [StatusCode],
content_type: ContentTypePolicy,
body: ResponseBodyPolicy,
max_body_bytes: usize,
) -> Result<Self, ResponsePolicyValidationError> {
validate_statuses(success_statuses)?;
validate_media_types(content_type)?;
match (body, content_type, max_body_bytes) {
(ResponseBodyPolicy::Required, _, 0) => {
return Err(ResponsePolicyValidationError::RequiredBodyHasZeroLimit);
}
(ResponseBodyPolicy::Forbidden, _, limit) if limit != 0 => {
return Err(ResponsePolicyValidationError::ForbiddenBodyHasNonzeroLimit);
}
(
ResponseBodyPolicy::Forbidden,
ContentTypePolicy::Required(_) | ContentTypePolicy::Optional(_),
_,
) => {
return Err(ResponsePolicyValidationError::ForbiddenBodyAllowsContentType);
}
(ResponseBodyPolicy::Required, ContentTypePolicy::Forbidden, _) => {
return Err(ResponsePolicyValidationError::RequiredBodyForbidsContentType);
}
_ => {}
}
Ok(Self {
success_statuses,
content_type,
body,
max_body_bytes,
})
}
#[must_use]
pub const fn success_statuses(self) -> &'static [StatusCode] {
self.success_statuses
}
#[must_use]
pub const fn content_type_policy(self) -> ContentTypePolicy {
self.content_type
}
#[must_use]
pub const fn body_policy(self) -> ResponseBodyPolicy {
self.body
}
#[must_use]
pub const fn max_body_bytes(self) -> usize {
self.max_body_bytes
}
pub fn validate<'buffer>(
self,
mut writer: ResponseBuffer<'buffer>,
request_id_policy: RequestIdPolicy,
) -> Result<CheckedResponseGuard<'buffer>, ResponsePolicyError> {
apply_request_id_policy(&mut writer, request_id_policy)?;
let snapshot = {
let response = writer.response().map_err(map_writer_error)?;
self.validate_view(response, request_id_policy)?
};
Ok(CheckedResponseGuard {
writer,
snapshot,
workspace: ResponseDecodeWorkspace::new(),
})
}
fn validate_view(
self,
response: TransportResponse<'_, '_>,
request_id_policy: RequestIdPolicy,
) -> Result<CheckedResponseSnapshot, ResponsePolicyError> {
if !self.success_statuses.contains(&response.status()) {
return Err(ResponsePolicyError::UnexpectedStatus);
}
match self.body {
ResponseBodyPolicy::Forbidden if !response.body().is_empty() => {
return Err(ResponsePolicyError::ForbiddenBody);
}
_ => {}
}
if response.body().len() > self.max_body_bytes {
return Err(ResponsePolicyError::BodyTooLarge);
}
if matches!(self.body, ResponseBodyPolicy::Required) && response.body().is_empty() {
return Err(ResponsePolicyError::MissingBody);
}
let content_type = response
.content_type()
.map_err(|_| ResponsePolicyError::InvalidContentType)?;
validate_content_type(self.content_type, content_type)?;
Ok(CheckedResponseSnapshot {
status: response.status(),
body_len: response.body().len(),
rate_limit: response.rate_limit(),
request_id_policy,
})
}
}
pub(crate) fn apply_request_id_policy(
writer: &mut ResponseBuffer<'_>,
request_id_policy: RequestIdPolicy,
) -> Result<(), ResponsePolicyError> {
writer.response().map_err(map_writer_error)?;
writer
.apply_request_id_policy(request_id_policy)
.map_err(|_| ResponsePolicyError::InvalidRequestId)
}
#[derive(Clone, Copy)]
pub struct CheckedResponse<'body> {
status: StatusCode,
body: &'body [u8],
content_type: Option<ResponseContentType<'body>>,
rate_limit: Option<RateLimit>,
request_id: Option<&'body [u8]>,
request_id_policy: RequestIdPolicy,
}
impl<'body> CheckedResponse<'body> {
#[must_use]
pub const fn status(&self) -> StatusCode {
self.status
}
#[must_use]
pub const fn body(&self) -> &[u8] {
self.body
}
#[must_use]
pub const fn content_type(&self) -> Option<ResponseContentType<'body>> {
self.content_type
}
#[must_use]
pub const fn rate_limit(&self) -> Option<RateLimit> {
self.rate_limit
}
#[must_use]
pub const fn request_id_policy(&self) -> RequestIdPolicy {
self.request_id_policy
}
pub fn with_request_id<R>(&self, inspect: impl FnOnce(Option<&[u8]>) -> R) -> R {
inspect(self.request_id)
}
}
#[derive(Clone, Copy)]
struct CheckedResponseSnapshot {
status: StatusCode,
body_len: usize,
rate_limit: Option<RateLimit>,
request_id_policy: RequestIdPolicy,
}
pub struct CheckedResponseGuard<'buffer> {
writer: ResponseBuffer<'buffer>,
snapshot: CheckedResponseSnapshot,
workspace: ResponseDecodeWorkspace,
}
impl CheckedResponseGuard<'_> {
#[must_use]
pub const fn status(&self) -> StatusCode {
self.snapshot.status
}
#[must_use]
pub fn content_type(&self) -> Option<ResponseContentType<'_>> {
self.writer
.response()
.ok()
.and_then(|response| response.content_type().ok().flatten())
}
#[must_use]
pub const fn rate_limit(&self) -> Option<RateLimit> {
self.snapshot.rate_limit
}
pub fn with_borrowed<R>(
&self,
inspect: impl for<'response> FnOnce(CheckedResponse<'response>) -> R,
) -> R {
inspect(self.checked_response())
}
pub fn decode_owned<R, E>(
self,
decode: impl for<'response> FnOnce(CheckedResponse<'response>) -> Result<R, E>,
) -> Result<R, E> {
self.decode_owned_with_workspace(|response, _workspace| decode(response))
}
pub fn decode_owned_with_workspace<R, E>(
mut self,
decode: impl for<'response> FnOnce(
CheckedResponse<'response>,
&mut ResponseDecodeWorkspace,
) -> Result<R, E>,
) -> Result<R, E> {
let result = {
let Self {
writer,
snapshot,
workspace,
} = &mut self;
let response = checked_response(writer, *snapshot);
decode(response, workspace)
};
drop(self);
result
}
pub fn retain_metadata_into<'destination>(
&mut self,
destination: &'destination mut [u8],
request_id_limit: usize,
) -> Result<RetainedResponseMetadata<'destination>, RetainedMetadataError> {
if self.snapshot.request_id_policy != RequestIdPolicy::Retain {
return Err(RetainedMetadataError::RetentionForbidden);
}
self.writer.retain_request_id(destination, request_id_limit)
}
fn checked_response(&self) -> CheckedResponse<'_> {
checked_response(&self.writer, self.snapshot)
}
}
impl fmt::Debug for CheckedResponseGuard<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CheckedResponseGuard")
.field("status", &self.status())
.field("body_len", &self.snapshot.body_len)
.field("body", &"[redacted]")
.field("content_type", &self.content_type())
.field("rate_limit", &self.rate_limit())
.field("request_id", &"[redacted]")
.finish()
}
}
impl fmt::Debug for CheckedResponse<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CheckedResponse")
.field("status", &self.status())
.field("body_len", &self.body().len())
.field("body", &"[redacted]")
.field("content_type", &self.content_type())
.field("rate_limit", &self.rate_limit())
.field("request_id", &"[redacted]")
.finish()
}
}
fn checked_response<'response>(
writer: &'response ResponseBuffer<'_>,
snapshot: CheckedResponseSnapshot,
) -> CheckedResponse<'response> {
CheckedResponse {
status: snapshot.status,
body: writer.initialized_body(snapshot.body_len),
content_type: writer
.response()
.ok()
.and_then(|response| response.content_type().ok().flatten()),
rate_limit: snapshot.rate_limit,
request_id: writer.request_id(),
request_id_policy: snapshot.request_id_policy,
}
}
fn validate_statuses(statuses: &[StatusCode]) -> Result<(), ResponsePolicyValidationError> {
if statuses.is_empty() {
return Err(ResponsePolicyValidationError::MissingSuccessStatus);
}
for (index, status) in statuses.iter().enumerate() {
if !status.is_success() {
return Err(ResponsePolicyValidationError::NonSuccessStatus);
}
if statuses
.get(..index)
.is_some_and(|seen| seen.contains(status))
{
return Err(ResponsePolicyValidationError::DuplicateSuccessStatus);
}
}
Ok(())
}
fn validate_media_types(policy: ContentTypePolicy) -> Result<(), ResponsePolicyValidationError> {
let media_types = match policy {
ContentTypePolicy::Required(values) | ContentTypePolicy::Optional(values) => values,
ContentTypePolicy::Forbidden => return Ok(()),
};
if media_types.is_empty() {
return Err(ResponsePolicyValidationError::MissingAcceptedMediaType);
}
for (index, media_type) in media_types.iter().enumerate() {
if media_types.get(..index).is_some_and(|seen| {
seen.iter()
.any(|candidate| candidate.as_str().eq_ignore_ascii_case(media_type.as_str()))
}) {
return Err(ResponsePolicyValidationError::DuplicateAcceptedMediaType);
}
}
Ok(())
}
fn validate_content_type(
policy: ContentTypePolicy,
actual: Option<ResponseContentType<'_>>,
) -> Result<(), ResponsePolicyError> {
match (policy, actual) {
(ContentTypePolicy::Required(_), None) => Err(ResponsePolicyError::MissingContentType),
(ContentTypePolicy::Forbidden, Some(_)) => Err(ResponsePolicyError::ForbiddenContentType),
(ContentTypePolicy::Forbidden | ContentTypePolicy::Optional(_), None) => Ok(()),
(
ContentTypePolicy::Required(accepted) | ContentTypePolicy::Optional(accepted),
Some(actual),
) => {
if accepted
.iter()
.any(|media_type| actual.matches(*media_type))
{
Ok(())
} else {
Err(ResponsePolicyError::UnexpectedContentType)
}
}
}
}
const fn map_writer_error(error: ResponseWriterError) -> ResponsePolicyError {
match error {
ResponseWriterError::NotCommitted
| ResponseWriterError::AlreadyCommitted
| ResponseWriterError::InitializedLengthTooLarge => {
ResponsePolicyError::UncommittedResponse
}
}
}