rig-core 0.42.0

An opinionated library for building LLM powered applications.
Documentation
use crate::{
    Embed,
    client::{self, BearerAuth, DebugExt, Provider},
    embeddings::EmbeddingsBuilder,
    http_client::HttpClientExt,
    wasm_compat::*,
};

use super::{CompletionModel, EmbeddingModel, ImageEmbeddingModel};
use serde::Deserialize;

// ================================================================
// Main Cohere Client
// ================================================================

#[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 {
    /// Provider error message; tolerant of `{"message": "..."}`,
    /// `{"error": "..."}`, nested `{"error": {"message": ...}}`, and bodies
    /// carrying both keys. Used for logging only — the raw body is preserved
    /// on the returned error.
    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))
    }

    /// Note: default embedding dimension of 0 will be used if model is not known.
    /// If this is the case, it's better to use function `embedding_model_with_ndims`
    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)
    }

    /// Create a Cohere `embed-english-v3.0` model for embedding PNG, JPEG,
    /// WebP, or GIF bytes.
    ///
    /// Images must be at least 2×2 pixels and no larger than 5 MB.
    /// Cohere accepts one image per request, so
    /// [`crate::embeddings::ImageEmbeddingModel::embed_images`] sends batches
    /// as ordered individual requests.
    pub fn image_embedding_model(&self) -> ImageEmbeddingModel<T> {
        ImageEmbeddingModel::new(self.clone())
    }

    /// Create an embedding model with the given name and the number of dimensions in the embedding generated by the model.
    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");
    }
}