use std::collections::BTreeMap;
use std::fmt;
use std::sync::Arc;
use std::time::{Duration, Instant};
use serde_json::Value;
use crate::answers::{decode_models, ModelMetadata, SystemOneResponse};
use crate::call_options::{CallOptions, ResolvedCall};
use crate::config::{
resolve_api_key, resolve_base_url, resolve_default_model, validate_timeout, LogLevel,
};
use crate::errors::{api_error_message, request_id_of, ApiError, ApiErrorKind, Error};
use crate::logging;
use crate::request::SystemOneRequest;
use crate::retry::RetryPolicy;
pub(crate) struct RawResponse {
pub status: u16,
pub body: String,
pub request_id: Option<String>,
pub endpoint: String,
}
#[derive(Debug, Clone, Default)]
pub struct RequestOptions {
pub(crate) options: CallOptions,
}
impl RequestOptions {
pub fn new() -> Self {
Self::default()
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.options.timeout = Some(timeout);
self
}
pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
self.options.retry_policy = Some(policy);
self
}
pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.options.headers.insert(name.into(), value.into());
self
}
pub fn extra_body(mut self, name: impl Into<String>, value: Value) -> Self {
self.options
.extra_body
.get_or_insert_with(serde_json::Map::new)
.insert(name.into(), value);
self
}
}
struct Inner {
api_key: String,
base_url: String,
default_model: String,
timeout: Duration,
retry_policy: RetryPolicy,
default_headers: BTreeMap<String, String>,
log_level: LogLevel,
http: reqwest::Client,
}
#[derive(Clone)]
pub struct Client {
inner: Arc<Inner>,
}
impl fmt::Debug for Client {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Client")
.field("base_url", &self.inner.base_url)
.field("default_model", &self.inner.default_model)
.field("timeout", &self.inner.timeout)
.field("retry_policy", &self.inner.retry_policy)
.field("default_headers", &self.inner.default_headers)
.field("log_level", &self.inner.log_level)
.finish_non_exhaustive()
}
}
#[derive(Default)]
pub struct ClientBuilder {
api_key: Option<String>,
base_url: Option<String>,
default_model: Option<String>,
timeout: Option<Duration>,
retry_policy: Option<RetryPolicy>,
default_headers: BTreeMap<String, String>,
log_level: Option<LogLevel>,
http_client: Option<reqwest::Client>,
}
impl fmt::Debug for ClientBuilder {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ClientBuilder")
.field("base_url", &self.base_url)
.field("default_model", &self.default_model)
.field("timeout", &self.timeout)
.field("retry_policy", &self.retry_policy)
.field("default_headers", &self.default_headers)
.field("log_level", &self.log_level)
.field(
"http_client",
&self.http_client.as_ref().map(|_| "reqwest::Client"),
)
.finish()
}
}
impl ClientBuilder {
pub fn api_key(mut self, api_key: impl Into<String>) -> Self {
self.api_key = Some(api_key.into());
self
}
pub fn base_url(mut self, base_url: impl Into<String>) -> Self {
self.base_url = Some(base_url.into());
self
}
pub fn default_model(mut self, model: impl Into<String>) -> Self {
self.default_model = Some(model.into());
self
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout);
self
}
pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
self.retry_policy = Some(policy);
self
}
pub fn default_header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.default_headers.insert(name.into(), value.into());
self
}
pub fn log_level(mut self, level: LogLevel) -> Self {
self.log_level = Some(level);
self
}
pub fn http_client(mut self, http: reqwest::Client) -> Self {
self.http_client = Some(http);
self
}
pub fn build(self) -> Result<Client, Error> {
let api_key = resolve_api_key(self.api_key.as_deref())?;
let base_url = resolve_base_url(self.base_url.as_deref());
validate_base_url(&base_url)?;
let default_model = resolve_default_model(self.default_model.as_deref());
let timeout = match self.timeout {
Some(timeout) => validate_timeout(timeout)?,
None => crate::DEFAULT_TIMEOUT,
};
let retry_policy = self.retry_policy.unwrap_or_default();
retry_policy.validate()?;
let log_level = LogLevel::resolve(self.log_level);
let http = self.http_client.unwrap_or_default();
Ok(Client {
inner: Arc::new(Inner {
api_key,
base_url,
default_model,
timeout,
retry_policy,
default_headers: self.default_headers,
log_level,
http,
}),
})
}
}
impl Client {
pub fn builder() -> ClientBuilder {
ClientBuilder::default()
}
pub fn from_env() -> Result<Self, Error> {
Self::builder().build()
}
pub fn base_url(&self) -> &str {
&self.inner.base_url
}
pub fn default_model(&self) -> &str {
&self.inner.default_model
}
pub fn timeout(&self) -> Duration {
self.inner.timeout
}
pub fn retry_policy(&self) -> &RetryPolicy {
&self.inner.retry_policy
}
pub fn log_level(&self) -> LogLevel {
self.inner.log_level
}
pub async fn system_one(&self, request: SystemOneRequest) -> Result<SystemOneResponse, Error> {
self.system_one_with(request, RequestOptions::new()).await
}
pub async fn system_one_with(
&self,
request: SystemOneRequest,
options: RequestOptions,
) -> Result<SystemOneResponse, Error> {
request.validate()?;
let call = self.resolve_call(options.options)?;
let model = request
.model
.clone()
.unwrap_or_else(|| self.inner.default_model.clone());
let mut body = serde_json::Map::new();
body.insert("state".into(), request.state.clone());
body.insert("model".into(), Value::String(model));
body.insert(
"questions".into(),
serde_json::to_value(&request.questions).map_err(|error| {
Error::InvalidRequest(format!("questions failed to serialize: {error}"))
})?,
);
if let Some(extra) = call.extra_body.clone() {
for (name, value) in extra {
body.insert(name, value);
}
}
let path = "/v1/systemone";
let url = format!("{}{}", self.inner.base_url, path);
let raw = self
.execute(
reqwest::Method::POST,
&url,
path,
Some(&Value::Object(body)),
&call,
)
.await?;
let response = SystemOneResponse::decode(&raw)?;
let overridden = call
.extra_body
.as_ref()
.is_some_and(|extra| extra.iter().any(|(name, _)| name == "questions"));
if !overridden {
crate::answers::check_complete(&raw, &response.answers, &request.questions)?;
}
Ok(response)
}
pub async fn list_models(&self) -> Result<Vec<ModelMetadata>, Error> {
self.list_models_with(RequestOptions::new()).await
}
pub async fn list_models_with(
&self,
options: RequestOptions,
) -> Result<Vec<ModelMetadata>, Error> {
let call = self.resolve_call(options.options)?;
let path = "/v1/models";
let url = format!("{}{}", self.inner.base_url, path);
let raw = self
.execute(reqwest::Method::GET, &url, path, None, &call)
.await?;
decode_models(&raw)
}
fn resolve_call(&self, options: CallOptions) -> Result<ResolvedCall, Error> {
let per_call = |error: Error| match error {
Error::Config(message) => Error::InvalidRequest(format!("per-call option: {message}")),
other => other,
};
let timeout = match options.timeout {
Some(timeout) => crate::config::validate_timeout(timeout).map_err(per_call)?,
None => self.inner.timeout,
};
let retry_policy = match options.retry_policy {
Some(policy) => {
policy.validate().map_err(per_call)?;
policy
}
None => self.inner.retry_policy.clone(),
};
Ok(ResolvedCall {
timeout,
retry_policy,
headers: options.headers,
extra_body: options.extra_body,
})
}
async fn execute(
&self,
method: reqwest::Method,
url: &str,
path: &str,
body: Option<&Value>,
call: &ResolvedCall,
) -> Result<RawResponse, Error> {
let endpoint = format!("{method} {url}");
let mut headers: BTreeMap<String, String> = BTreeMap::new();
for (name, value) in self.inner.default_headers.iter().chain(call.headers.iter()) {
headers.insert(name.to_ascii_lowercase(), value.clone());
}
headers.remove("x-typesafe-retry-count");
if body.is_none() {
headers.remove("content-type");
}
headers.insert(
"authorization".into(),
format!("Bearer {}", self.inner.api_key),
);
headers.insert("accept".into(), "application/json".into());
if body.is_some() {
headers.insert("content-type".into(), "application/json".into());
}
headers.insert("user-agent".into(), sdk_version_header());
headers.insert("x-typesafe-sdk".into(), sdk_version_header());
headers.insert("x-typesafe-runtime".into(), runtime_header());
let policy = call.retry_policy.clone();
let total_budget = policy.total_budget.filter(|budget| !budget.is_zero());
let call_start = Instant::now();
let mut last_error: Option<Error> = None;
for retry_number in 0..=policy.max_retries {
let mut attempt_headers = headers.clone();
if retry_number > 0 {
attempt_headers.insert("x-typesafe-retry-count".into(), retry_number.to_string());
}
let reqwest_headers = build_header_map(&attempt_headers)?;
if self.inner.log_level >= LogLevel::Debug {
logging::log_request_debug(method.as_str(), path, &reqwest_headers, body);
}
let started = Instant::now();
let mut attempt = self
.inner
.http
.request(method.clone(), url)
.headers(reqwest_headers.clone())
.timeout(call.timeout);
if let Some(body) = body {
attempt = attempt.body(body.to_string());
}
let result = attempt.send().await;
let response = match result {
Ok(response) => response,
Err(error) => {
let error: Error = if error.is_timeout() {
Error::Timeout {
timeout: call.timeout,
}
} else {
Error::Connection {
message: format!("Connection error: {error}"),
source: Some(Box::new(error)),
}
};
if self.inner.log_level >= LogLevel::Info {
logging::log_attempt_error(method.as_str(), path, &error.to_string());
}
last_error = Some(error);
if !should_retry(&policy, &last_error, retry_number) {
return Err(last_error.unwrap());
}
let delay = policy.delay_for_retry(retry_number, None);
if !delay_within_budget(&delay, &total_budget, call_start) {
return Err(last_error.unwrap());
}
if self.inner.log_level >= LogLevel::Info {
logging::log_retry_scheduled(
method.as_str(),
path,
delay,
retry_number + 1,
&last_error.as_ref().unwrap().to_string(),
);
}
tokio::time::sleep(delay).await;
continue;
}
};
let status = response.status().as_u16();
let response_headers = response.headers().clone();
let request_id = request_id_of(&response_headers);
let text = match response.text().await {
Ok(text) => text,
Err(error) => {
let error = if error.is_timeout() {
Error::Timeout {
timeout: call.timeout,
}
} else {
Error::Connection {
message: format!("Connection error: {error}"),
source: Some(Box::new(error)),
}
};
if self.inner.log_level >= LogLevel::Info {
logging::log_attempt_error(method.as_str(), path, &error.to_string());
}
last_error = Some(error);
if !should_retry(&policy, &last_error, retry_number) {
return Err(last_error.unwrap());
}
let delay = policy.delay_for_retry(retry_number, None);
if !delay_within_budget(&delay, &total_budget, call_start) {
return Err(last_error.unwrap());
}
if self.inner.log_level >= LogLevel::Info {
logging::log_retry_scheduled(
method.as_str(),
path,
delay,
retry_number + 1,
&last_error.as_ref().unwrap().to_string(),
);
}
tokio::time::sleep(delay).await;
continue;
}
};
let duration = started.elapsed();
if self.inner.log_level >= LogLevel::Info {
logging::log_attempt_info(
method.as_str(),
path,
status,
duration,
request_id.as_deref(),
);
}
if self.inner.log_level >= LogLevel::Debug {
logging::log_response_debug(
method.as_str(),
path,
status,
&response_headers,
&text,
);
}
if (200..300).contains(&status) {
return Ok(RawResponse {
status,
body: text,
request_id,
endpoint,
});
}
let parsed_body = parse_body(&text);
let retry_after = policy.parse_retry_after(&response_headers);
let api_error = ApiError {
status,
kind: ApiErrorKind::from_status(status),
message: api_error_message(&parsed_body),
body: parsed_body,
headers: response_headers,
request_id,
endpoint: endpoint.clone(),
retry_after,
};
let error = Error::from(api_error);
last_error = Some(error);
if !should_retry(&policy, &last_error, retry_number) {
return Err(last_error.unwrap());
}
let delay =
policy.delay_for_retry(retry_number, last_error.as_ref().and_then(Error::as_api));
if !delay_within_budget(&delay, &total_budget, call_start) {
return Err(last_error.unwrap());
}
if self.inner.log_level >= LogLevel::Info {
logging::log_retry_scheduled(
method.as_str(),
path,
delay,
retry_number + 1,
&last_error.as_ref().unwrap().to_string(),
);
}
tokio::time::sleep(delay).await;
}
Err(last_error.unwrap_or_else(|| Error::Connection {
message: "Connection error: exhausted retries.".into(),
source: None,
}))
}
}
fn should_retry(policy: &RetryPolicy, last_error: &Option<Error>, retry_number: u32) -> bool {
if retry_number >= policy.max_retries {
return false;
}
last_error
.as_ref()
.map(|error| policy.is_retryable(error))
.unwrap_or(false)
}
fn delay_within_budget(
delay: &Duration,
total_budget: &Option<Duration>,
call_start: Instant,
) -> bool {
match total_budget {
None => true,
Some(budget) => {
let elapsed = call_start.elapsed();
elapsed.saturating_add(*delay) < *budget
}
}
}
fn parse_body(text: &str) -> Option<Value> {
if text.is_empty() {
return None;
}
match serde_json::from_str(text) {
Ok(value) => Some(value),
Err(_) => Some(Value::String(text.to_owned())),
}
}
fn build_header_map(
headers: &BTreeMap<String, String>,
) -> Result<reqwest::header::HeaderMap, Error> {
let mut map = reqwest::header::HeaderMap::new();
for (name, value) in headers {
let name: reqwest::header::HeaderName = name
.parse()
.map_err(|error| Error::Config(format!("invalid header name {name:?}: {error}")))?;
let value = reqwest::header::HeaderValue::from_str(value).map_err(|error| {
Error::Config(format!("invalid value for header {name}: {error}"))
})?;
map.insert(name, value);
}
Ok(map)
}
fn validate_base_url(base_url: &str) -> Result<(), Error> {
let invalid = || {
Error::Config(format!("invalid base URL {base_url:?}: expected an absolute http(s) URL with no query or fragment"))
};
let url = reqwest::Url::parse(base_url).map_err(|_| invalid())?;
let scheme_ok = matches!(url.scheme(), "http" | "https");
if !scheme_ok || url.host_str().is_none() || url.query().is_some() || url.fragment().is_some() {
return Err(invalid());
}
Ok(())
}
fn sdk_version_header() -> String {
format!("typesafe-client-rust/{}", env!("CARGO_PKG_VERSION"))
}
fn runtime_header() -> String {
format!(
"rust ({}; {})",
std::env::consts::OS,
std::env::consts::ARCH
)
}