diff --git a/crates/umem_ai/src/lib.rs b/crates/umem_ai/src/lib.rs index d82d464..2c8cb28 100644 --- a/crates/umem_ai/src/lib.rs +++ b/crates/umem_ai/src/lib.rs @@ -1,5 +1,6 @@ // TODO: remove this allow once the module is fully implemented #![allow(dead_code)] +mod model_impl; mod providers; mod response_generators; mod utils; @@ -7,14 +8,13 @@ mod utils; use anyhow::Result; use async_trait::async_trait; use lazy_static::lazy_static; +pub use model_impl::*; pub use providers::*; pub use response_generators::*; use schemars::JsonSchema; use serde::{Serialize, de::DeserializeOwned}; use std::sync::Arc; use thiserror::Error; -use tokio::sync::OnceCell; -use umem_config::CONFIG; pub type HashMap = rustc_hash::FxHashMap; @@ -33,8 +33,6 @@ pub struct LanguageModel { pub model_name: String, } -static LANGUAGE_MODEL: OnceCell> = OnceCell::const_new(); - impl LanguageModel { fn new(provider: Arc, model_name: String) -> Self { Self { @@ -42,48 +40,6 @@ impl LanguageModel { model_name, } } - - pub async fn get_model() -> Result, LanguageModelError> { - LANGUAGE_MODEL - .get_or_try_init(|| async { - match CONFIG.language_model.provider.clone() { - umem_config::Provider::OpenAI(open_ai) => { - let openai_provider = OpenAIProvider::builder() - .api_key(open_ai.api_key) - .base_url(open_ai.base_url) - .default_headers(open_ai.default_headers.unwrap_or_default()) - .project(open_ai.project) - .organization(open_ai.organization) - .build(); - - let provider = Arc::new(AIProvider::from(openai_provider)); - - Ok(Arc::new(LanguageModel { - provider, - model_name: CONFIG.language_model.model.clone(), - })) - } - umem_config::Provider::AmazonBedrock(config) => { - let provider = AmazonBedrockProviderBuilder::default() - .region(config.region) - .access_key_id(config.key_id) - .secret_access_key(config.access_key) - .build() - .await - .map_err(|e| AIProviderError::ProviderBuilderError(e.into()))?; - - let provider = Arc::new(AIProvider::from(provider)); - - Ok(Arc::new(LanguageModel { - provider, - model_name: CONFIG.language_model.model.clone(), - })) - } - } - }) - .await - .cloned() - } } #[derive(Error, Debug)] @@ -98,8 +54,6 @@ pub struct RerankingModel { pub model_name: String, } -static RERANKING_MODEL: OnceCell> = OnceCell::const_new(); - impl RerankingModel { fn new(provider: Arc, model_name: String) -> Self { Self { @@ -107,48 +61,6 @@ impl RerankingModel { model_name, } } - - pub async fn get_model() -> Result, LanguageModelError> { - RERANKING_MODEL - .get_or_try_init(|| async { - match CONFIG.reranking_model.provider.clone() { - umem_config::Provider::OpenAI(open_ai) => { - let openai_provider = OpenAIProvider::builder() - .api_key(open_ai.api_key) - .base_url(open_ai.base_url) - .default_headers(open_ai.default_headers.unwrap_or_default()) - .project(open_ai.project) - .organization(open_ai.organization) - .build(); - - let provider = Arc::new(AIProvider::from(openai_provider)); - - Ok(Arc::new(RerankingModel { - provider, - model_name: CONFIG.reranking_model.model.clone(), - })) - } - umem_config::Provider::AmazonBedrock(config) => { - let provider = AmazonBedrockProviderBuilder::default() - .region(config.region) - .access_key_id(config.key_id) - .secret_access_key(config.access_key) - .build() - .await - .map_err(|e| AIProviderError::ProviderBuilderError(e.into()))?; - - let provider = Arc::new(AIProvider::from(provider)); - - Ok(Arc::new(RerankingModel { - provider, - model_name: CONFIG.reranking_model.model.clone(), - })) - } - } - }) - .await - .cloned() - } } #[derive(Error, Debug)] @@ -163,8 +75,6 @@ pub struct EmbeddingModel { pub model_name: String, } -static EMBEDDING_MODEL: OnceCell> = OnceCell::const_new(); - impl EmbeddingModel { fn new(provider: Arc, model_name: String) -> Self { Self { @@ -172,48 +82,6 @@ impl EmbeddingModel { model_name, } } - - pub async fn get_model() -> Result, LanguageModelError> { - EMBEDDING_MODEL - .get_or_try_init(|| async { - match CONFIG.embedding_model.provider.clone() { - umem_config::Provider::OpenAI(open_ai) => { - let openai_provider = OpenAIProvider::builder() - .api_key(open_ai.api_key) - .base_url(open_ai.base_url) - .default_headers(open_ai.default_headers.unwrap_or_default()) - .project(open_ai.project) - .organization(open_ai.organization) - .build(); - - let provider = Arc::new(AIProvider::from(openai_provider)); - - Ok(Arc::new(EmbeddingModel { - provider, - model_name: CONFIG.embedding_model.model.clone(), - })) - } - umem_config::Provider::AmazonBedrock(config) => { - let provider = AmazonBedrockProviderBuilder::default() - .region(config.region) - .access_key_id(config.key_id) - .secret_access_key(config.access_key) - .build() - .await - .map_err(|e| AIProviderError::ProviderBuilderError(e.into()))?; - - let provider = Arc::new(AIProvider::from(provider)); - - Ok(Arc::new(EmbeddingModel { - provider, - model_name: CONFIG.embedding_model.model.clone(), - })) - } - } - }) - .await - .cloned() - } } #[derive(Debug)] diff --git a/crates/umem_ai/src/model_impl/embedding_model.rs b/crates/umem_ai/src/model_impl/embedding_model.rs new file mode 100644 index 0000000..a6a7aef --- /dev/null +++ b/crates/umem_ai/src/model_impl/embedding_model.rs @@ -0,0 +1,53 @@ +use crate::{ + AIProvider, AIProviderError, AmazonBedrockProviderBuilder, EmbeddingModel, EmbeddingModelError, + OpenAIProvider, +}; +use std::sync::Arc; +use tokio::sync::OnceCell; +use umem_config::CONFIG; + +pub static EMBEDDING_MODEL: OnceCell> = OnceCell::const_new(); + +impl EmbeddingModel { + pub async fn get_model() -> Result, EmbeddingModelError> { + EMBEDDING_MODEL + .get_or_try_init(|| async { + match CONFIG.embedding_model.provider.clone() { + umem_config::Provider::OpenAI(open_ai) => { + let openai_provider = OpenAIProvider::builder() + .api_key(open_ai.api_key) + .base_url(open_ai.base_url) + .default_headers(open_ai.default_headers.unwrap_or_default()) + .project(open_ai.project) + .organization(open_ai.organization) + .build(); + + let provider = Arc::new(AIProvider::from(openai_provider)); + + Ok(Arc::new(EmbeddingModel { + provider, + model_name: CONFIG.embedding_model.model.clone(), + })) + } + umem_config::Provider::AmazonBedrock(config) => { + let provider = AmazonBedrockProviderBuilder::default() + .region(config.region) + .access_key_id(config.key_id) + .secret_access_key(config.access_key) + .build() + .await + .map_err(|e| AIProviderError::ProviderBuilderError(e.into()))?; + + let provider = Arc::new(AIProvider::from(provider)); + + Ok(Arc::new(EmbeddingModel { + provider, + model_name: CONFIG.embedding_model.model.clone(), + })) + } + } + }) + .await + .cloned() + } +} diff --git a/crates/umem_ai/src/model_impl/language_model.rs b/crates/umem_ai/src/model_impl/language_model.rs new file mode 100644 index 0000000..48a4e0f --- /dev/null +++ b/crates/umem_ai/src/model_impl/language_model.rs @@ -0,0 +1,53 @@ +use crate::{ + AIProvider, AIProviderError, AmazonBedrockProviderBuilder, LanguageModel, LanguageModelError, + OpenAIProvider, +}; +use std::sync::Arc; +use tokio::sync::OnceCell; +use umem_config::CONFIG; + +pub static LANGUAGE_MODEL: OnceCell> = OnceCell::const_new(); + +impl LanguageModel { + pub async fn get_model() -> Result, LanguageModelError> { + LANGUAGE_MODEL + .get_or_try_init(|| async { + match CONFIG.language_model.provider.clone() { + umem_config::Provider::OpenAI(open_ai) => { + let openai_provider = OpenAIProvider::builder() + .api_key(open_ai.api_key) + .base_url(open_ai.base_url) + .default_headers(open_ai.default_headers.unwrap_or_default()) + .project(open_ai.project) + .organization(open_ai.organization) + .build(); + + let provider = Arc::new(AIProvider::from(openai_provider)); + + Ok(Arc::new(LanguageModel { + provider, + model_name: CONFIG.language_model.model.clone(), + })) + } + umem_config::Provider::AmazonBedrock(config) => { + let provider = AmazonBedrockProviderBuilder::default() + .region(config.region) + .access_key_id(config.key_id) + .secret_access_key(config.access_key) + .build() + .await + .map_err(|e| AIProviderError::ProviderBuilderError(e.into()))?; + + let provider = Arc::new(AIProvider::from(provider)); + + Ok(Arc::new(LanguageModel { + provider, + model_name: CONFIG.language_model.model.clone(), + })) + } + } + }) + .await + .cloned() + } +} diff --git a/crates/umem_ai/src/model_impl/mod.rs b/crates/umem_ai/src/model_impl/mod.rs new file mode 100644 index 0000000..4b838df --- /dev/null +++ b/crates/umem_ai/src/model_impl/mod.rs @@ -0,0 +1,7 @@ +mod embedding_model; +mod language_model; +mod reranking_model; + +pub use embedding_model::*; +pub use language_model::*; +pub use reranking_model::*; diff --git a/crates/umem_ai/src/model_impl/reranking_model.rs b/crates/umem_ai/src/model_impl/reranking_model.rs new file mode 100644 index 0000000..81d4caf --- /dev/null +++ b/crates/umem_ai/src/model_impl/reranking_model.rs @@ -0,0 +1,53 @@ +use crate::{ + AIProvider, AIProviderError, AmazonBedrockProviderBuilder, OpenAIProvider, RerankingModel, + RerankingModelError, +}; +use std::sync::Arc; +use tokio::sync::OnceCell; +use umem_config::CONFIG; + +pub static RERANKING_MODEL: OnceCell> = OnceCell::const_new(); + +impl RerankingModel { + pub async fn get_model() -> Result, RerankingModelError> { + RERANKING_MODEL + .get_or_try_init(|| async { + match CONFIG.reranking_model.provider.clone() { + umem_config::Provider::OpenAI(open_ai) => { + let openai_provider = OpenAIProvider::builder() + .api_key(open_ai.api_key) + .base_url(open_ai.base_url) + .default_headers(open_ai.default_headers.unwrap_or_default()) + .project(open_ai.project) + .organization(open_ai.organization) + .build(); + + let provider = Arc::new(AIProvider::from(openai_provider)); + + Ok(Arc::new(RerankingModel { + provider, + model_name: CONFIG.reranking_model.model.clone(), + })) + } + umem_config::Provider::AmazonBedrock(config) => { + let provider = AmazonBedrockProviderBuilder::default() + .region(config.region) + .access_key_id(config.key_id) + .secret_access_key(config.access_key) + .build() + .await + .map_err(|e| AIProviderError::ProviderBuilderError(e.into()))?; + + let provider = Arc::new(AIProvider::from(provider)); + + Ok(Arc::new(RerankingModel { + provider, + model_name: CONFIG.reranking_model.model.clone(), + })) + } + } + }) + .await + .cloned() + } +}