rig-core 0.44.0

An opinionated library for building LLM powered applications.
Documentation
//! Text and image embedding models, responses, and input identifiers.
//!
//! ```no_run
//! use rig_core::DynModel;
//! use rig_core::operation::Embedding;
//!
//! # async fn example(model: DynModel<Embedding>) -> Result<(), Box<dyn std::error::Error>> {
//! let embedding = model.embed_text("A document").await?;
//! # let _ = embedding;
//! # Ok(())
//! # }
//! ```

use crate::completion::Usage;
use crate::error::ProviderError;
use serde::{Deserialize, Serialize};

impl<W, T> crate::driver::Model<W, T>
where
    W: crate::wire::Wire<Op = crate::operation::Embedding>,
    T: crate::driver::Transport<W>,
{
    /// Embed one text, returning the last vector or an error if none is returned.
    pub async fn embed_text(&self, text: &str) -> Result<Embedding, ProviderError> {
        last_embedding(self.call(vec![text.to_owned()]).await?)
    }
}

impl crate::driver::DynModel<crate::operation::Embedding> {
    /// Embed one text, returning the last vector or an error if none is returned.
    pub async fn embed_text(&self, text: &str) -> Result<Embedding, ProviderError> {
        last_embedding(self.call(vec![text.to_owned()]).await?)
    }
}

/// The last vector of a one-text batch, or the empty-reply error.
fn last_embedding(response: EmbeddingResponse) -> Result<Embedding, ProviderError> {
    let mut embeddings = response.embeddings;
    embeddings.pop().ok_or_else(|| {
        ProviderError::Response(
            "embedding provider returned an empty response for embed_text".to_string(),
        )
    })
}

/// Text or image embeddings and normalized provider metadata.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingResponse {
    /// The embeddings returned by the provider, one per input, in input order.
    pub embeddings: Vec<Embedding>,
    /// Token usage for this request; every counter is `None` when the
    /// provider reported none (see [`Usage`]).
    #[serde(default)]
    pub usage: Usage,
    /// Stable descriptor name of the provider that produced this response,
    /// for example `"openai"`. Always populated.
    pub provider: String,
    /// Provider-reported model identifier, when the wire response named one.
    #[serde(default)]
    pub model: Option<String>,
    /// Provider-assigned response-scoped identifier, when reported.
    #[serde(default, skip_serializing_if = "Option::is_none")]
    pub response_id: Option<String>,
    /// Transport request identifier from HTTP response headers, or `None` when absent.
    #[serde(default, skip_serializing_if = "Option::is_none")]
    pub provider_request_id: Option<String>,
    /// Provider response payload, or null when no raw payload was attached.
    #[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
    pub raw: serde_json::Value,
}

impl EmbeddingResponse {
    /// A response carrying `embeddings`. The driver writes the provider, the
    /// transport request id and the reply document; decoders set what the
    /// provider reported.
    pub fn new(embeddings: Vec<Embedding>) -> Self {
        Self {
            embeddings,
            usage: Usage::default(),
            provider: String::new(),
            model: None,
            response_id: None,
            provider_request_id: None,
            raw: serde_json::Value::Null,
        }
    }

    /// A response carrying bare `vectors` in input order. The operation's
    /// fold pairs each with the input it belongs to, since a reply is not
    /// trusted to echo its inputs back.
    pub(crate) fn from_vectors(vectors: impl IntoIterator<Item = Vec<f64>>) -> Self {
        Self::new(
            vectors
                .into_iter()
                .map(|vec| Embedding {
                    document: String::new(),
                    vec,
                })
                .collect(),
        )
    }
}

/// A document identifier and its vector. Equality compares only the document,
/// not vector values.
#[derive(Clone, Default, Deserialize, Serialize, Debug)]
pub struct Embedding {
    /// The text that was embedded, or a non-sensitive input identifier for
    /// non-text embeddings. Used for debugging and equality.
    pub document: String,
    /// The embedding vector
    pub vec: Vec<f64>,
}

impl PartialEq for Embedding {
    fn eq(&self, other: &Self) -> bool {
        self.document == other.document
    }
}

impl Eq for Embedding {}

#[cfg(test)]
mod provider_response_tests;

/// The media type of an encoded image, sniffed from its magic bytes.
///
/// The image-embedding wires need it twice: once to reject a format the
/// provider does not accept, and once to name the vector's input.
pub fn image_media_type(bytes: &[u8]) -> Option<&'static str> {
    if bytes.starts_with(b"\x89PNG\r\n\x1a\n") {
        Some("image/png")
    } else if bytes.starts_with(b"\xff\xd8\xff") {
        Some("image/jpeg")
    } else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
        Some("image/gif")
    } else if bytes.starts_with(b"RIFF") && bytes.get(8..12) == Some(b"WEBP".as_slice()) {
        Some("image/webp")
    } else {
        None
    }
}

/// Identifies image bytes by media type and a URL-safe, unpadded SHA-256 digest,
/// without retaining the image or a reversible encoding. Unknown formats use
/// `application/octet-stream`; this function does not validate provider support.
pub fn image_document(bytes: &[u8]) -> String {
    use base64::Engine as _;
    use sha2::Digest as _;
    let media_type = image_media_type(bytes).unwrap_or("application/octet-stream");
    let digest = sha2::Sha256::digest(bytes);
    format!(
        "{media_type};sha256={}",
        base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest)
    )
}