pub mod anthropic;
pub mod error;
pub mod functions;
#[cfg(feature = "bert")]
pub mod huggingface;
mod inference;
pub mod openai;
pub mod streaming;
use self::{
anthropic::builder::AnthropicCompletionModel, error::CompletionResult, functions::Function,
inference::CompletionRequestBuilder, openai::builder::OpenAiCompletionModel,
streaming::ProviderStreamHandler,
};
use crate::agents::memory::MessageStack;
use anyhow::anyhow;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::fmt::Debug;
use tracing::{info, warn};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum CompletionProvider {
OpenAi(OpenAiCompletionModel),
Anthropic(AnthropicCompletionModel),
}
impl From<OpenAiCompletionModel> for CompletionProvider {
fn from(value: OpenAiCompletionModel) -> Self {
Self::OpenAi(value)
}
}
impl From<AnthropicCompletionModel> for CompletionProvider {
fn from(value: AnthropicCompletionModel) -> Self {
Self::Anthropic(value)
}
}
impl CompletionProvider {
fn inner_builder(&self) -> Box<&dyn CompletionRequestBuilder> {
match &self {
Self::OpenAi(b) => return Box::new(b),
Self::Anthropic(b) => return Box::new(b),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompletionModel {
pub provider: CompletionProvider,
pub params: ModelParameters,
pub api_key: String,
#[serde(skip)]
client: Client,
}
impl Eq for CompletionModel {}
impl PartialEq for CompletionModel {
fn eq(&self, other: &Self) -> bool {
self.provider == other.provider
&& self.params == other.params
&& self.api_key == other.api_key
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ModelParameters {
pub total_token_count: u32,
pub temperature: Option<u8>,
pub frequency_penalty: Option<i8>,
pub max_tokens: Option<u32>,
pub n: Option<u32>,
pub presence_penalty: Option<i8>,
}
impl Default for ModelParameters {
fn default() -> Self {
Self {
total_token_count: 0,
temperature: Some(70),
frequency_penalty: None,
max_tokens: None,
n: Some(1),
presence_penalty: None,
}
}
}
impl ModelParameters {
fn temperature(&self) -> Result<f32, anyhow::Error> {
Ok((self.temperature.ok_or(anyhow!("No temperature"))? / 100) as f32)
}
}
impl CompletionModel {
pub fn new(
m: impl Into<CompletionProvider>,
params: ModelParameters,
api_key: &str,
) -> CompletionModel {
let client = Client::new();
Self {
provider: m.into(),
params,
client,
api_key: api_key.to_owned(),
}
}
pub fn default_openai(api_key: &str) -> CompletionModel {
let provider = CompletionProvider::OpenAi(OpenAiCompletionModel::default());
let client = reqwest::Client::new();
CompletionModel {
provider,
params: ModelParameters::default(),
api_key: api_key.to_owned(),
client,
}
}
pub fn default_anthropic(api_key: &str) -> CompletionModel {
let provider = CompletionProvider::Anthropic(AnthropicCompletionModel::default());
let client = reqwest::Client::new();
CompletionModel {
provider,
params: ModelParameters::default(),
api_key: api_key.to_owned(),
client,
}
}
#[tracing::instrument(name = "io completion", skip_all)]
pub(crate) async fn get_io_completion(
&self,
messages: &MessageStack,
) -> CompletionResult<String> {
let builder = self.provider.inner_builder();
let headers = builder.headers(&self.api_key);
let url = builder.url_str();
let req = builder.into_io_req(messages, &self.params)?;
let json_req = req.as_json()?;
info!(
"\nSending request:\n{:?}\nto: {}\nwith headers: {:?}\n",
json_req, url, headers
);
let response = self
.client
.post(url)
.headers(headers)
.json(&json_req)
.send()
.await?;
match req.process_response(response).await {
Ok(r) => return Ok(TryInto::<String>::try_into(r)?),
Err(err) => {
warn!("Error getting Io completion: {:?}", err);
Err(err)
}
}
}
#[tracing::instrument(name = "streamed completion", skip_all)]
pub(crate) async fn get_stream_completion(
&self,
messages: &MessageStack,
) -> CompletionResult<ProviderStreamHandler> {
let builder = self.provider.inner_builder();
let headers = builder.headers(&self.api_key);
let url = builder.url_str();
let req = builder.into_stream_req(messages, &self.params)?;
let json_req = req.as_json()?;
info!(
"\nSending request:\n{:?}\nto: {}\nwith headers: {:?}\n",
json_req, url, headers
);
let response = self
.client
.post(url)
.headers(headers)
.json(&json_req)
.send()
.await?;
match req.process_response(response).await {
Ok(r) => return Ok(TryInto::<ProviderStreamHandler>::try_into(r)?),
Err(err) => {
warn!("Error getting streamed Io completion: {:?}", err);
Err(err.into())
}
}
}
#[tracing::instrument(name = "function completion", skip_all)]
pub(crate) async fn get_fn_completion(
&self,
messages: &MessageStack,
function: Function,
) -> CompletionResult<Value> {
let builder = self.provider.inner_builder();
let headers = builder.headers(&self.api_key);
let url = builder.url_str();
let req = builder.serialize_function(messages, function)?;
info!(
"\nSending request:\n{:?}\nto: {}\nwith headers: {:?}\n",
req, url, headers
);
let response = self
.client
.post(url)
.headers(headers)
.json(&req)
.send()
.await?;
let json = response.json().await?;
info!("Got response: {json:#?}");
match builder.process_function_response(json) {
Ok(r) => return Ok(r),
Err(err) => {
warn!("Error getting function completion: {:?}", err);
Err(err.into())
}
}
}
}