openai-interface 0.10.0

A low-level Rust interface for the OpenAI API
Documentation
use serde::Serialize;
use url::Url;

use crate::{
    errors::OapiError,
    rest::post::{Post, PostNoStream},
};

/// Creates an embedding vector representing the input text.
#[derive(Debug, Serialize, Default, Clone)]
pub struct EmbeddingRequest {
    /// Input text to embed, encoded as a string, an array of strings, an
    /// array of tokens, or an array of token arrays.
    ///
    /// The input must not exceed the max input tokens for the model (8192
    /// tokens for all embedding models) and cannot be an empty string.
    pub input: EmbeddingInput,
    /// ID of the model to use, e.g. `text-embedding-v4`.
    pub model: String,
    /// The number of dimensions the resulting output embeddings should have.
    ///
    /// Only supported in `text-embedding-3` and later models.
    #[serde(skip_serializing_if = "Option::is_none")]
    pub dimensions: Option<u32>,
    /// The format to return the embeddings in. Can be either `float` or
    /// `base64`.
    ///
    /// Note that some OpenAI-compatible providers (e.g. Qwen) only support
    /// `float`.
    #[serde(skip_serializing_if = "Option::is_none")]
    pub encoding_format: Option<EncodingFormat>,
    /// A unique identifier representing your end-user, which can help OpenAI
    /// to monitor and detect abuse.
    #[serde(skip_serializing_if = "Option::is_none")]
    pub user: Option<String>,
    /// Add additional JSON properties to the request.
    #[serde(flatten, skip_serializing_if = "Option::is_none")]
    pub extra_body: Option<serde_json::Map<String, serde_json::Value>>,
}

/// The input text to embed.
#[derive(Debug, Serialize, Clone)]
#[serde(untagged)]
pub enum EmbeddingInput {
    /// A single string to embed.
    String(String),
    /// An array of strings to embed.
    StringArray(Vec<String>),
    /// A single array of tokens to embed.
    Tokens(Vec<u32>),
    /// An array of token arrays to embed.
    TokenArray(Vec<Vec<u32>>),
}

impl Default for EmbeddingInput {
    fn default() -> Self {
        Self::String(String::new())
    }
}

/// The format to return the embeddings in.
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "lowercase")]
pub enum EncodingFormat {
    Float,
    Base64,
}

impl EmbeddingRequest {
    pub fn is_streaming(&self) -> bool {
        false
    }
}

impl Post for EmbeddingRequest {
    fn is_streaming(&self) -> bool {
        false
    }

    /// Builds the URL for the request.
    ///
    /// `base_url` should be like <https://api.openai.com/v1>
    fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
        let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
        url.path_segments_mut()
            .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
            .push("embeddings");

        Ok(url.to_string())
    }
}

impl PostNoStream for EmbeddingRequest {
    type Response = super::response::CreateEmbeddingResponse;
}

#[cfg(test)]
mod tests {
    use super::*;

    /// Serializes a simple string input request.
    #[test]
    fn string_input_serialization() {
        let request = EmbeddingRequest {
            input: EmbeddingInput::String("Hello".to_string()),
            model: "text-embedding-v4".to_string(),
            dimensions: Some(1024),
            encoding_format: Some(EncodingFormat::Float),
            ..Default::default()
        };

        let json = serde_json::to_string(&request).unwrap();
        assert!(json.contains(r#""input":"Hello""#), "json: {json}");
        assert!(
            json.contains(r#""model":"text-embedding-v4""#),
            "json: {json}"
        );
        assert!(json.contains(r#""dimensions":1024"#), "json: {json}");
        assert!(
            json.contains(r#""encoding_format":"float""#),
            "json: {json}"
        );
    }

    /// Serializes an array of strings and a token array input.
    #[test]
    fn array_input_serialization() {
        let request = EmbeddingRequest {
            input: EmbeddingInput::StringArray(vec!["Hello".to_string(), "World".to_string()]),
            model: "text-embedding-3-small".to_string(),
            ..Default::default()
        };
        let json = serde_json::to_string(&request).unwrap();
        assert!(
            json.contains(r#""input":["Hello","World"]"#),
            "json: {json}"
        );

        let request = EmbeddingRequest {
            input: EmbeddingInput::TokenArray(vec![vec![1234, 5678]]),
            model: "text-embedding-3-small".to_string(),
            ..Default::default()
        };
        let json = serde_json::to_string(&request).unwrap();
        assert!(json.contains(r#""input":[[1234,5678]]"#), "json: {json}");
    }

    /// A flat token id array is sent as `input: [..]`.
    #[test]
    fn flat_token_array_serialization() {
        let request = EmbeddingRequest {
            input: EmbeddingInput::Tokens(vec![1234, 5678]),
            model: "text-embedding-3-small".to_string(),
            ..Default::default()
        };
        let json = serde_json::to_string(&request).unwrap();
        assert!(json.contains(r#""input":[1234,5678]"#), "json: {json}");
    }

    #[test]
    fn test_build_url() {
        let request = EmbeddingRequest::default();
        let url = request.build_url("https://api.openai.com/v1/").unwrap();
        assert_eq!(url, "https://api.openai.com/v1/embeddings");
    }
}