use crate::{
Embed,
client::{self, BearerAuth, DebugExt, Provider},
embeddings::EmbeddingsBuilder,
http_client::HttpClientExt,
wasm_compat::*,
};
use super::{CompletionModel, EmbeddingModel, ImageEmbeddingModel};
use serde::Deserialize;
#[derive(Debug, Default, Clone, Copy)]
pub struct CohereExt;
#[derive(Debug, Default, Clone, Copy)]
pub struct CohereBuilder;
type CohereApiKey = BearerAuth;
pub type Client<H = reqwest::Client> = client::Client<CohereExt, H>;
pub type ClientBuilder<H = crate::markers::Missing> =
client::ClientBuilder<CohereBuilder, CohereApiKey, H>;
impl Provider for CohereExt {
type Builder = CohereBuilder;
const VERIFY_PATH: &'static str = "/models";
}
client::impl_capabilities!(
CohereExt,
completion = CompletionModel<H>,
embeddings = EmbeddingModel<H>,
);
impl DebugExt for CohereExt {}
client::impl_default_provider_builder!(
CohereBuilder => CohereExt,
api_key = CohereApiKey,
base_url = "https://api.cohere.ai",
);
client::impl_provider_client!(Client, input = CohereApiKey, api_key_env = "COHERE_API_KEY",);
#[derive(Debug)]
pub struct ApiErrorResponse {
pub message: String,
}
impl<'de> Deserialize<'de> for ApiErrorResponse {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
Ok(Self {
message: crate::providers::internal::envelope::error_message(deserializer)?,
})
}
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
pub enum ApiResponse<T> {
Ok(T),
Err(ApiErrorResponse),
}
impl<T> Client<T>
where
T: HttpClientExt + Clone + WasmCompatSend + WasmCompatSync + 'static,
{
pub fn embeddings<D: Embed>(
&self,
model: impl Into<String>,
input_type: &str,
) -> EmbeddingsBuilder<EmbeddingModel<T>, D> {
EmbeddingsBuilder::new(self.embedding_model(model, input_type))
}
pub fn embedding_model(&self, model: impl Into<String>, input_type: &str) -> EmbeddingModel<T> {
let model = model.into();
let ndims = super::model_dimensions_from_identifier(&model).unwrap_or_default();
EmbeddingModel::new(self.clone(), model, input_type, ndims)
}
pub fn image_embedding_model(&self) -> ImageEmbeddingModel<T> {
ImageEmbeddingModel::new(self.clone())
}
pub fn embedding_model_with_ndims(
&self,
model: impl Into<String>,
input_type: &str,
ndims: usize,
) -> EmbeddingModel<T> {
EmbeddingModel::new(self.clone(), model, input_type, ndims)
}
}
#[cfg(test)]
mod tests {
#[test]
fn test_client_initialization() {
let _client =
crate::providers::cohere::Client::new("dummy-key").expect("Client::new() failed");
let _client_from_builder = crate::providers::cohere::Client::builder()
.api_key("dummy-key")
.build()
.expect("Client::builder() failed");
}
}