mod local;
pub use local::LocalMockTransport;
use core::fmt;
use core::sync::atomic::{AtomicUsize, Ordering};
use cloud_sdk::Method;
use cloud_sdk::authentication::{
AsyncAuthenticatedTransport, AuthenticatedRequest, BlockingAuthenticatedTransport,
};
use cloud_sdk::transport::{
AsyncResponseStaging, AsyncTransport, BlockingTransport, BoundTransport, EndpointIdentity,
EndpointIdentityError, HeaderSensitivity, RequestHeaders, RequestTarget, ResponseAttempt,
ResponseCompletion, ResponseContentType, ResponseHeaders, ResponseMetadata, ResponseWriter,
TransportRequest,
};
use crate::{FixtureBodyError, ResponseFixture};
#[derive(Clone, Copy)]
pub struct ExpectedRequest<'a> {
method: Method,
target: RequestTarget<'a>,
body: &'a [u8],
headers: RequestHeaders<'a>,
}
impl<'a> ExpectedRequest<'a> {
#[must_use]
pub const fn new(method: Method, target: RequestTarget<'a>) -> Self {
Self {
method,
target,
body: &[],
headers: RequestHeaders::EMPTY,
}
}
#[must_use]
pub const fn with_body(mut self, body: &'a [u8]) -> Self {
self.body = body;
self
}
#[must_use]
pub const fn with_headers(mut self, headers: RequestHeaders<'a>) -> Self {
self.headers = headers;
self
}
const fn method(self) -> Method {
self.method
}
const fn target(self) -> RequestTarget<'a> {
self.target
}
const fn body(self) -> &'a [u8] {
self.body
}
const fn headers(self) -> RequestHeaders<'a> {
self.headers
}
}
impl fmt::Debug for ExpectedRequest<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ExpectedRequest")
.field("method", &self.method)
.field("target", &"[redacted]")
.field("body", &"[redacted]")
.field("headers", &self.headers)
.finish()
}
}
#[derive(Debug)]
pub struct MockExchange<'a> {
request: ExpectedRequest<'a>,
response: ResponseFixture<'a>,
}
impl<'a> MockExchange<'a> {
#[must_use]
pub const fn new(request: ExpectedRequest<'a>, response: ResponseFixture<'a>) -> Self {
Self { request, response }
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum MockError {
Exhausted,
MethodMismatch,
TargetMismatch,
BodyMismatch,
HeadersMismatch,
ResponseBufferTooSmall,
CursorOverflow,
ConcurrentRequest,
InvalidFixtureMetadata,
ResponseWriterRejected,
}
impl_static_error!(MockError,
Self::Exhausted => "mock transport has no expected exchange remaining",
Self::MethodMismatch => "mock request method differs from expectation",
Self::TargetMismatch => "mock request target differs from expectation",
Self::BodyMismatch => "mock request body differs from expectation",
Self::HeadersMismatch => "mock request headers differ from expectation",
Self::ResponseBufferTooSmall => "mock response buffer is too small",
Self::CursorOverflow => "mock transport cursor overflowed",
Self::ConcurrentRequest => "mock transport cursor changed concurrently",
Self::InvalidFixtureMetadata => "mock fixture metadata is invalid",
Self::ResponseWriterRejected => "mock response writer rejected output",
);
pub struct MockTransport<'a> {
exchanges: &'a [MockExchange<'a>],
cursor: AtomicUsize,
endpoint: Option<EndpointIdentity<'a>>,
}
impl<'a> MockTransport<'a> {
#[must_use]
pub const fn new(exchanges: &'a [MockExchange<'a>]) -> Self {
Self {
exchanges,
cursor: AtomicUsize::new(0),
endpoint: None,
}
}
#[must_use]
pub const fn with_endpoint(mut self, endpoint: EndpointIdentity<'a>) -> Self {
self.endpoint = Some(endpoint);
self
}
#[must_use]
pub fn remaining(&self) -> usize {
self.exchanges
.len()
.saturating_sub(self.cursor.load(Ordering::Acquire))
}
#[must_use]
pub fn is_complete(&self) -> bool {
self.remaining() == 0
}
fn send_inner(
&self,
request: TransportRequest<'_>,
response: &mut ResponseWriter<'_>,
) -> Result<(), MockError> {
if response.is_committed() {
return Err(MockError::ResponseWriterRejected);
}
let mut response = response
.begin_attempt()
.map_err(|_| MockError::ResponseWriterRejected)?;
let completion = self.stage_inner(request, &mut response)?;
response
.commit_completion(completion)
.map_err(|_| MockError::ResponseWriterRejected)
}
fn stage_inner<'buffer>(
&self,
request: TransportRequest<'_>,
response: &mut impl MockResponseSink<'buffer>,
) -> Result<ResponseCompletion, MockError> {
let cursor = self.cursor.load(Ordering::Acquire);
let exchange = self.exchanges.get(cursor).ok_or(MockError::Exhausted)?;
if request.method() != exchange.request.method() {
return Err(MockError::MethodMismatch);
}
if request.target() != exchange.request.target() {
return Err(MockError::TargetMismatch);
}
if request.body() != exchange.request.body() {
return Err(MockError::BodyMismatch);
}
if !request_headers_match(request.headers(), exchange.request.headers()) {
return Err(MockError::HeadersMismatch);
}
let next_cursor = cursor.checked_add(1).ok_or(MockError::CursorOverflow)?;
let completion = stage_response(&exchange.response, response)?;
self.cursor
.compare_exchange(cursor, next_cursor, Ordering::AcqRel, Ordering::Acquire)
.map_err(|_| MockError::ConcurrentRequest)?;
Ok(completion)
}
}
pub(crate) fn stage_response<'buffer>(
fixture: &ResponseFixture<'_>,
response: &mut impl MockResponseSink<'buffer>,
) -> Result<ResponseCompletion, MockError> {
let _content_type = fixture
.content_type()
.map(ResponseContentType::new)
.transpose()
.map_err(|_| MockError::InvalidFixtureMetadata)?;
let rate_limit = fixture
.rate_limit()
.map(|value| value.into_rate_limit())
.transpose()
.map_err(|_| MockError::InvalidFixtureMetadata)?;
{
let response_headers = response
.headers_mut()
.map_err(|_| MockError::ResponseWriterRejected)?;
if let Some(source) = fixture.headers() {
for header in source.iter() {
response_headers
.try_push(header.name(), header.value(), header.sensitivity())
.map_err(|_| MockError::InvalidFixtureMetadata)?;
}
}
if let Some(value) = fixture.content_type() {
response_headers
.try_push("content-type", value.as_bytes(), HeaderSensitivity::Public)
.map_err(|_| MockError::InvalidFixtureMetadata)?;
}
if let Some(value) = rate_limit {
push_rate_limit_headers(response_headers, value)
.map_err(|_| MockError::InvalidFixtureMetadata)?;
}
}
let body_len = fixture
.body()
.write_to(
response
.body_mut()
.map_err(|_| MockError::ResponseWriterRejected)?,
)
.map_err(|error| match error {
FixtureBodyError::OutputTooSmall | FixtureBodyError::TooLarge => {
MockError::ResponseBufferTooSmall
}
})?;
let mut metadata = ResponseMetadata::EMPTY;
if let Some(value) = rate_limit {
metadata = metadata.with_rate_limit(value);
}
Ok(ResponseCompletion::new(
fixture.status(),
body_len,
metadata,
))
}
pub(crate) trait MockResponseSink<'buffer> {
fn body_mut(&mut self) -> Result<&mut [u8], MockError>;
fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError>;
}
impl<'buffer> MockResponseSink<'buffer> for ResponseAttempt<'_, 'buffer> {
fn body_mut(&mut self) -> Result<&mut [u8], MockError> {
ResponseAttempt::body_mut(self).map_err(|_| MockError::ResponseWriterRejected)
}
fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError> {
ResponseAttempt::headers_mut(self).map_err(|_| MockError::ResponseWriterRejected)
}
}
impl<'buffer> MockResponseSink<'buffer> for AsyncResponseStaging<'_, 'buffer> {
fn body_mut(&mut self) -> Result<&mut [u8], MockError> {
AsyncResponseStaging::body_mut(self).map_err(|_| MockError::ResponseWriterRejected)
}
fn headers_mut(&mut self) -> Result<&mut ResponseHeaders<'buffer>, MockError> {
AsyncResponseStaging::headers_mut(self).map_err(|_| MockError::ResponseWriterRejected)
}
}
fn request_headers_match(actual: RequestHeaders<'_>, expected: RequestHeaders<'_>) -> bool {
let actual = actual.as_slice();
let expected = expected.as_slice();
actual.len() == expected.len()
&& actual.iter().zip(expected).all(|(actual, expected)| {
actual.name() == expected.name()
&& actual.value().as_str().as_bytes() == expected.value().as_str().as_bytes()
&& actual.sensitivity() == expected.sensitivity()
})
}
fn push_rate_limit_headers(
headers: &mut ResponseHeaders,
rate_limit: cloud_sdk::rate_limit::RateLimit,
) -> Result<(), ()> {
let mut storage = [0_u8; 20];
for (name, value) in [
("ratelimit-limit", rate_limit.limit()),
("ratelimit-remaining", rate_limit.remaining()),
("ratelimit-reset", rate_limit.reset_epoch_seconds()),
] {
let text = write_decimal(value, &mut storage).ok_or(())?;
headers
.try_push(name, text.as_bytes(), HeaderSensitivity::Public)
.map_err(|_| ())?;
}
Ok(())
}
fn write_decimal(value: u64, output: &mut [u8; 20]) -> Option<&str> {
let mut value = value;
let mut cursor = output.len();
loop {
cursor = cursor.checked_sub(1)?;
let digit = u8::try_from(value % 10).ok()?;
*output.get_mut(cursor)? = b'0'.checked_add(digit)?;
value /= 10;
if value == 0 {
break;
}
}
core::str::from_utf8(output.get(cursor..)?).ok()
}
impl BlockingTransport for MockTransport<'_> {
type Error = MockError;
fn send(
&self,
request: TransportRequest<'_>,
response: &mut ResponseWriter<'_>,
) -> Result<(), Self::Error> {
self.send_inner(request, response)
}
}
impl BlockingAuthenticatedTransport for MockTransport<'_> {
type Error = MockError;
fn send_authenticated(
&self,
request: AuthenticatedRequest<'_, '_>,
response: &mut ResponseWriter<'_>,
) -> Result<(), Self::Error> {
self.send_inner(request.transport_request(), response)
}
}
impl AsyncTransport for MockTransport<'_> {
type Error = MockError;
async fn send<'transport, 'request, 'writer, 'buffer>(
&'transport self,
request: TransportRequest<'request>,
mut response: AsyncResponseStaging<'writer, 'buffer>,
) -> Result<ResponseCompletion, Self::Error>
where
'transport: 'writer,
'request: 'writer,
'buffer: 'writer,
{
self.stage_inner(request, &mut response)
}
}
impl AsyncAuthenticatedTransport for MockTransport<'_> {
type Error = MockError;
async fn send_authenticated<'transport, 'request, 'policy, 'writer, 'buffer>(
&'transport self,
request: AuthenticatedRequest<'request, 'policy>,
mut response: AsyncResponseStaging<'writer, 'buffer>,
) -> Result<ResponseCompletion, Self::Error>
where
'transport: 'writer,
'request: 'writer,
'policy: 'writer,
'buffer: 'writer,
{
self.stage_inner(request.transport_request(), &mut response)
}
}
impl BoundTransport for MockTransport<'_> {
fn endpoint_identity(&self) -> Result<EndpointIdentity<'_>, EndpointIdentityError> {
self.endpoint.ok_or(EndpointIdentityError::UnboundTransport)
}
}
impl fmt::Debug for MockTransport<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("MockTransport")
.field("remaining", &self.remaining())
.finish_non_exhaustive()
}
}