Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
136 changes: 2 additions & 134 deletions crates/umem_ai/src/lib.rs
Original file line number Diff line number Diff line change
@@ -1,20 +1,20 @@
// TODO: remove this allow once the module is fully implemented
#![allow(dead_code)]
mod model_impl;
mod providers;
mod response_generators;
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<K, V> = rustc_hash::FxHashMap<K, V>;

Expand All @@ -33,57 +33,13 @@ pub struct LanguageModel {
pub model_name: String,
}

static LANGUAGE_MODEL: OnceCell<Arc<LanguageModel>> = OnceCell::const_new();

impl LanguageModel {
fn new(provider: Arc<AIProvider>, model_name: String) -> Self {
Self {
provider,
model_name,
}
}

pub async fn get_model() -> Result<Arc<LanguageModel>, 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)]
Expand All @@ -98,57 +54,13 @@ pub struct RerankingModel {
pub model_name: String,
}

static RERANKING_MODEL: OnceCell<Arc<RerankingModel>> = OnceCell::const_new();

impl RerankingModel {
fn new(provider: Arc<AIProvider>, model_name: String) -> Self {
Self {
provider,
model_name,
}
}

pub async fn get_model() -> Result<Arc<RerankingModel>, 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)]
Expand All @@ -163,57 +75,13 @@ pub struct EmbeddingModel {
pub model_name: String,
}

static EMBEDDING_MODEL: OnceCell<Arc<EmbeddingModel>> = OnceCell::const_new();

impl EmbeddingModel {
fn new(provider: Arc<AIProvider>, model_name: String) -> Self {
Self {
provider,
model_name,
}
}

pub async fn get_model() -> Result<Arc<EmbeddingModel>, 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)]
Expand Down
53 changes: 53 additions & 0 deletions crates/umem_ai/src/model_impl/embedding_model.rs
Original file line number Diff line number Diff line change
@@ -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<Arc<EmbeddingModel>> = OnceCell::const_new();

impl EmbeddingModel {
pub async fn get_model() -> Result<Arc<EmbeddingModel>, 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()
}
}
53 changes: 53 additions & 0 deletions crates/umem_ai/src/model_impl/language_model.rs
Original file line number Diff line number Diff line change
@@ -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<Arc<LanguageModel>> = OnceCell::const_new();

impl LanguageModel {
pub async fn get_model() -> Result<Arc<LanguageModel>, 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()
}
}
7 changes: 7 additions & 0 deletions crates/umem_ai/src/model_impl/mod.rs
Original file line number Diff line number Diff line change
@@ -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::*;
53 changes: 53 additions & 0 deletions crates/umem_ai/src/model_impl/reranking_model.rs
Original file line number Diff line number Diff line change
@@ -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<Arc<RerankingModel>> = OnceCell::const_new();

impl RerankingModel {
pub async fn get_model() -> Result<Arc<RerankingModel>, 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()
}
}