use std::sync::Arc;
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;
use crate::rubric::Rubric;
#[derive(Debug, Clone)]
pub struct Client {
inner: crate::Client,
rt: Arc<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_caused("could not start the async runtime", e))?;
Ok(Self {
inner,
rt: Arc::new(rt),
})
}
pub fn from_env() -> Result<Self> {
Self::new(crate::Client::from_env()?)
}
pub fn default_model(&self) -> &str {
self.inner.default_model()
}
pub fn models(&self) -> Models<'_> {
Models { client: self }
}
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 ask<R: Rubric>(&self, state: impl Serialize) -> AskRequest<'_, R> {
AskRequest {
rt: &self.rt,
req: self.inner.ask(state),
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct Models<'a> {
client: &'a Client,
}
impl<'a> Models<'a> {
pub fn list(&self) -> ListModelsRequest<'a> {
ListModelsRequest {
rt: &self.client.rt,
req: self.client.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 AskRequest<'a, R> {
rt: &'a tokio::runtime::Runtime,
req: crate::AskRequest<R>,
}
impl<R: Rubric> AskRequest<'_, R> {
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<R> {
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())
}
}