use std::time::Duration;
use http::header::{HeaderName, HeaderValue};
use serde::Serialize;
use serde_json::Value;
use crate::error::{Error, Result};
use crate::question::Questions;
use crate::response::{ListModelsResponse, SystemOneResponse};
use crate::retry::RetryPolicy;
#[derive(Debug)]
pub struct Client {
inner: crate::Client,
rt: tokio::runtime::Runtime,
}
impl Client {
pub fn new(inner: crate::Client) -> Result<Self> {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|e| Error::Config(format!("could not start runtime: {e}")))?;
Ok(Self { inner, rt })
}
pub fn from_env() -> Result<Self> {
Self::new(crate::Client::from_env()?)
}
pub fn system_one<S: Serialize>(
&self,
state: S,
questions: impl Into<Questions>,
) -> SystemOneRequest<'_> {
SystemOneRequest {
rt: &self.rt,
req: self.inner.system_one(state, questions),
}
}
pub fn list_models(&self) -> ListModelsRequest<'_> {
ListModelsRequest {
rt: &self.rt,
req: self.inner.models().list(),
}
}
}
macro_rules! forward {
($($name:ident($($arg:ident: $ty:ty),*)),* $(,)?) => {$(
#[doc = concat!("See the async request's `", stringify!($name), "`.")]
pub fn $name(mut self, $($arg: $ty),*) -> Self {
self.req = self.req.$name($($arg),*);
self
}
)*};
}
#[must_use = "call .send()"]
#[derive(Debug)]
pub struct SystemOneRequest<'a> {
rt: &'a tokio::runtime::Runtime,
req: crate::SystemOneRequest,
}
impl SystemOneRequest<'_> {
forward!(
retry(policy: RetryPolicy),
timeout(timeout: Duration),
header(name: HeaderName, value: HeaderValue),
);
pub fn model(mut self, model: impl Into<String>) -> Self {
self.req = self.req.model(model);
self
}
pub fn extra_body(mut self, key: impl Into<String>, value: impl Into<Value>) -> Self {
self.req = self.req.extra_body(key, value);
self
}
pub fn send(self) -> Result<SystemOneResponse> {
self.rt.block_on(self.req.send())
}
}
#[must_use = "call .send()"]
#[derive(Debug)]
pub struct ListModelsRequest<'a> {
rt: &'a tokio::runtime::Runtime,
req: crate::ListModelsRequest,
}
impl ListModelsRequest<'_> {
forward!(
retry(policy: RetryPolicy),
timeout(timeout: Duration),
header(name: HeaderName, value: HeaderValue),
);
pub fn send(self) -> Result<ListModelsResponse> {
self.rt.block_on(self.req.send())
}
}