ruchat 0.1.2

ollama/chroma command-line AI chat tool
Documentation
use crate::chroma::create_client;
use crate::error::RuChatError;
use crate::ollama::model::get_name;
use chromadb::collection::{ChromaCollection, CollectionEntries};
use chromadb::embeddings::EmbeddingFunction;
use clap::Parser;
use log::warn;
use ollama_rs::Ollama;
use ollama_rs::generation::embeddings::request::GenerateEmbeddingsRequest;
use serde_json::{Map, Value};

/// Command-line arguments for embedding data into a Chroma database.
///
/// This struct defines the arguments required to perform an embedding
/// operation in a Chroma database, including model details, prompt,
/// and database connection information.
#[derive(Parser, Debug, Clone, PartialEq)]
pub struct EmbedArgs {
    /// The model to use for generating embeddings.
    #[clap(short, long, default_value = "nomic-embed-text:latest")]
    pub(crate) model: String,

    /// The prompt to embed.
    #[clap(short, long)]
    pub(crate) prompt: String,

    /// Chroma database server address and port.
    #[clap(short = 'C', long, default_value = "http://localhost:8000")]
    pub(crate) chroma_server: String,

    /// Chroma database name.
    #[clap(short = 'd', long, default_value = "default")]
    pub(crate) chroma_database: String,

    /// Chroma token for authentication.
    #[clap(short = 't', long)]
    pub(crate) chroma_token: Option<String>,

    /// Chroma database collection name.
    #[clap(short, long, default_value = "default")]
    pub(crate) collection: String,

    /// Chroma collection metadata, comma separated key:value pairs.
    #[clap(short, long, default_value = "version:0.01")]
    pub(crate) collection_metadata: Option<String>,

    /// Chroma entries metadata, comma separated key:value pairs.
    #[clap(short, long, default_value = "version:0.01")]
    pub(crate) entries_metadata: Option<String>,
}

/// Parses metadata from a string of comma-separated key:value pairs.
///
/// # Parameters
///
/// - `arg_metadata`: An optional string containing metadata.
///
/// # Returns
///
/// A `Result` containing an optional map of metadata or a `RuChatError`.
fn get_metadata(arg_metadata: &Option<String>) -> Result<Option<Map<String, Value>>, RuChatError> {
    if arg_metadata.is_none() {
        return Ok(None);
    }
    let mut metadata = Map::new();
    if let Some(md) = arg_metadata {
        for s in md.split(',') {
            match s.split_once(':') {
                Some((k, v)) => _ = metadata.insert(k.to_string(), v.into()),
                None => return Err(RuChatError::InvalidMetadata(s.to_string())),
            }
        }
    }
    Ok(Some(metadata))
}

/// Embeds data into a Chroma database.
///
/// This function connects to a Chroma database using the provided
/// arguments, generates embeddings for the specified prompt, and
/// stores the embeddings in the database.
///
/// # Parameters
///
/// - `ollama`: The Ollama client for generating embeddings.
/// - `args`: The command-line arguments for the embedding operation.
///
/// # Returns
///
/// A `Result` indicating success or failure.
pub(crate) async fn embed(ollama: Ollama, args: &EmbedArgs) -> Result<(), RuChatError> {
    let model_name = get_name(&ollama, &args.model).await?;
    if !model_name.contains("embed") {
        warn!("Model {} might not be an embeddings model", model_name);
    }
    let entries_metadata = get_metadata(&args.entries_metadata)?;

    let request = GenerateEmbeddingsRequest::new(model_name, vec![args.prompt.as_str()].into());
    let client = create_client(
        args.chroma_token.as_deref(),
        &args.chroma_server,
        &args.chroma_database,
    )
    .await?;
    let res = ollama.generate_embeddings(request).await?;

    let collection_metadata = get_metadata(&args.collection_metadata)?;

    eprintln!("Collection name: {}", args.collection);
    // XXX error here.
    let collection: ChromaCollection = client
        .get_or_create_collection(&args.collection, collection_metadata)
        .await?;

    let id = collection.id().to_string();
    eprintln!("Collection Name: {}", collection.name());
    eprintln!("Collection ID: {}", id);
    eprintln!("Collection Metadata: {:?}", collection.metadata());
    eprintln!("Collection Count: {}", collection.count().await?);

    let collection_entries = CollectionEntries {
        ids: vec![id.as_str()],
        embeddings: Some(res.embeddings),
        metadatas: entries_metadata.map(|md| vec![md]),
        documents: Some(vec![&args.prompt]),
    };
    // The function to use to compute the embeddings. If None, embeddings must be provided.
    let embedding_function: Option<Box<dyn EmbeddingFunction>> = None;

    let result: Value = collection
        .upsert(collection_entries, embedding_function)
        .await?;
    eprintln!("{:?}", result);
    Ok(())
}

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

    #[test]
    fn test_get_metadata_valid() {
        let metadata_str = Some("key1:value1,key2:value2".to_string());
        let result = get_metadata(&metadata_str);
        assert!(result.is_ok());
        let metadata = result.unwrap().unwrap();
        assert_eq!(metadata["key1"], "value1");
        assert_eq!(metadata["key2"], "value2");
    }

    #[test]
    fn test_get_metadata_invalid() {
        let metadata_str = Some("key1value1".to_string());
        let result = get_metadata(&metadata_str);
        assert!(result.is_err());
    }
}