use std::{borrow::Cow, fmt, marker::PhantomData, time::Duration};
use bytes::Bytes;
use http::Method;
use serde::Serialize;
use crate::{
client::Client,
codec::{self, EncodeError},
config::ZERO_TIMEOUT,
de::{self, AnswerSet},
error::Error,
question::{PreparedQuestions, upsert},
response::{Answers, SystemOneResponse},
retry::{self, RetryPolicy},
text,
transport::{self, Exchange, HttpService},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub(crate) enum Deadline {
#[default]
Client,
After(Duration),
Never,
}
impl Deadline {
pub(crate) fn resolve(self, client: Option<Duration>) -> Result<Option<Duration>, Error> {
match self {
Self::Client => Ok(client),
Self::After(timeout) if timeout.is_zero() => Err(Error::invalid_request(ZERO_TIMEOUT)),
Self::After(timeout) => Ok(Some(timeout)),
Self::Never => Ok(None),
}
}
}
#[derive(Clone, Default)]
pub(crate) struct CallHeaders<'a>(Vec<(Cow<'a, str>, Cow<'a, str>)>);
impl<'a> CallHeaders<'a> {
pub(crate) fn push(&mut self, name: Cow<'a, str>, value: Cow<'a, str>) {
self.0.push((name, value));
}
pub(crate) fn parse(
&self,
with_body: bool,
) -> Result<Vec<(http::HeaderName, http::HeaderValue)>, Error> {
transport::call_headers(
self.0.iter().map(|(name, value)| (name.as_ref(), value.as_ref())),
with_body,
)
}
}
impl fmt::Debug for CallHeaders<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.debug_list().entries(self.0.iter().map(|(name, _)| name)).finish()
}
}
const STATE: &str = "state";
const MODEL: &str = "model";
const QUESTIONS: &str = "questions";
type ExtraMember<'a> = (Cow<'a, str>, Result<Vec<u8>, EncodeError>);
#[must_use = "a request does nothing until it is sent"]
pub struct SystemOne<'a, S, T: ?Sized, A = Answers> {
client: &'a Client<S>,
state: &'a T,
questions: &'a PreparedQuestions,
model: Option<Cow<'a, str>>,
deadline: Deadline,
headers: CallHeaders<'a>,
extra: Vec<ExtraMember<'a>>,
retry: Option<RetryPolicy>,
answers: PhantomData<fn() -> A>,
}
impl<'a, S, T> SystemOne<'a, S, T>
where
T: ?Sized,
{
pub(crate) fn new(
client: &'a Client<S>,
state: &'a T,
questions: &'a PreparedQuestions,
) -> Self {
Self {
client,
state,
questions,
model: None,
deadline: Deadline::Client,
headers: CallHeaders::default(),
extra: Vec::new(),
retry: None,
answers: PhantomData,
}
}
}
impl<'a, S, T, A> SystemOne<'a, S, T, A>
where
T: ?Sized,
{
pub fn model(mut self, model: impl Into<Cow<'a, str>>) -> Self {
self.model = Some(model.into());
self
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.deadline = Deadline::After(timeout);
self
}
pub fn no_timeout(mut self) -> Self {
self.deadline = Deadline::Never;
self
}
pub fn header(mut self, name: impl Into<Cow<'a, str>>, value: impl Into<Cow<'a, str>>) -> Self {
self.headers.push(name.into(), value.into());
self
}
pub fn retry(mut self, policy: RetryPolicy) -> Self {
self.retry = Some(policy);
self
}
pub fn extra_body<V>(mut self, name: impl Into<Cow<'a, str>>, value: &V) -> Self
where
V: Serialize + ?Sized,
{
let mut encoded = Vec::new();
let value = codec::encode_into(&mut encoded, value).map(|()| encoded);
upsert(&mut self.extra, name.into(), value);
self
}
pub fn typed<B>(self) -> SystemOne<'a, S, T, B>
where
B: AnswerSet,
{
SystemOne {
client: self.client,
state: self.state,
questions: self.questions,
model: self.model,
deadline: self.deadline,
headers: self.headers,
extra: self.extra,
retry: self.retry,
answers: PhantomData,
}
}
}
impl<S, T, A> SystemOne<'_, S, T, A>
where
S: HttpService,
T: Serialize + ?Sized,
A: AnswerSet,
{
pub async fn send(self) -> Result<SystemOneResponse<A>, Error> {
let shared = self.client.shared();
let deadline = self.deadline.resolve(shared.config.timeout())?;
let headers = self.headers.parse(true)?;
let mut body = self.encode()?;
let uri = shared.config.endpoints().system_one();
let exchange = Exchange {
method: &Method::POST,
uri,
base_headers: &shared.post_headers,
call_headers: &headers,
deadline,
max_response_bytes: shared.config.max_response_bytes(),
};
let policy = self.retry.as_ref().unwrap_or(&shared.retry);
let retain = policy.can_retry();
let asked =
de::AnswerContext::new(self.questions.len()).with_levels(self.questions.max_levels());
retry::run(policy, &Method::POST, uri, move |retry| {
let body = if retain { body.clone() } else { std::mem::take(&mut body) };
async move {
let (status, headers, body) =
transport::attempt(&shared.service, exchange, retry, Some(body)).await?;
de::decode_system_one_with(body, status, headers, asked, Some((&Method::POST, uri)))
}
})
.await
}
fn encode(&self) -> Result<Bytes, Error> {
for (name, value) in &self.extra {
if let Err(error) = value {
let part = format!("the extra member {}", text::quoted(name));
return Err(encode_failure(&part, error));
}
}
let extra = |wanted: &str| {
self.extra.iter().find_map(|(name, value)| match value {
Ok(bytes) if name == wanted => Some(bytes.as_slice()),
_ => None,
})
};
let mut state_is_json_content = true;
let body = codec::encode_body(|buffer| {
buffer.extend_from_slice(br#"{"state":"#);
match extra(STATE) {
Some(bytes) => buffer.extend_from_slice(bytes),
None => {
let start = buffer.len();
codec::encode_into(buffer, self.state)?;
if !matches!(buffer.get(start), Some(b'"' | b'{' | b'[')) {
state_is_json_content = false;
return Ok(());
}
}
}
buffer.extend_from_slice(br#","model":"#);
match (extra(MODEL), &self.model) {
(Some(bytes), _) => buffer.extend_from_slice(bytes),
(None, Some(model)) => codec::write_json_string(buffer, model),
(None, None) => buffer.extend_from_slice(&self.client.shared().model_json),
}
buffer.extend_from_slice(br#","questions":"#);
buffer.extend_from_slice(extra(QUESTIONS).unwrap_or(self.questions.as_bytes()));
for (name, value) in &self.extra {
if let (Ok(bytes), false) = (value, [STATE, MODEL, QUESTIONS].contains(&&**name)) {
buffer.push(b',');
codec::write_json_string(buffer, name);
buffer.push(b':');
buffer.extend_from_slice(bytes);
}
}
buffer.push(b'}');
Ok(())
});
let body = body.map_err(|error| encode_failure("the state", &error))?;
if !state_is_json_content {
return Err(Error::invalid_request(
"The state must be a JSON string, object or array; \
it encoded as a number, a boolean or null.",
));
}
Ok(body)
}
}
fn encode_failure(part: &str, error: &EncodeError) -> Error {
Error::invalid_request(format!(
"The request body could not be encoded as JSON: {part}: {}",
text::bounded(&error.message(), text::MAX_MESSAGE_CHARS)
))
}
impl<S, T, A> fmt::Debug for SystemOne<'_, S, T, A>
where
T: ?Sized,
{
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut shown = formatter.debug_struct("SystemOne");
shown
.field("questions", &self.questions.len())
.field("model", &self.model)
.field("deadline", &self.deadline)
.field("headers", &self.headers)
.field("extra_body", &self.extra.iter().map(|(name, _)| name).collect::<Vec<_>>());
if let Some(retry) = &self.retry {
shown.field("retry", retry);
}
shown.finish_non_exhaustive()
}
}
#[cfg(test)]
#[path = "request_tests.rs"]
mod tests;