Skip to content
Open
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
1 change: 1 addition & 0 deletions frontend/rust-lib/Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

3 changes: 3 additions & 0 deletions frontend/rust-lib/flowy-ai/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,9 @@ base64 = "0.21.5"
futures-util = "0.3.30"
flowy-storage-pub = { workspace = true }
ollama-rs.workspace = true
# Already compiled as a non-optional dependency of the langchain-rust fork; used directly here for
# the OpenAI-compatible backend because langchain's OpenAiEmbedder cannot request `dimensions`.
async-openai = "0.28.1"
schemars = "0.8.22"
twox-hash = { version = "2.1.0", features = ["xxhash64"] }
async-trait.workspace = true
Expand Down
49 changes: 34 additions & 15 deletions frontend/rust-lib/flowy-ai/src/embeddings/context.rs
Original file line number Diff line number Diff line change
@@ -1,15 +1,17 @@
use crate::embeddings::scheduler::EmbeddingScheduler;
use crate::local_ai::client::AIClient;
use arc_swap::ArcSwapOption;
use flowy_error::{ErrorCode, FlowyError, FlowyResult};
use flowy_sqlite_vec::db::VectorSqliteDB;
use lib_infra::util::get_operating_system;
use ollama_rs::Ollama;
use std::path::PathBuf;
use std::sync::{Arc, OnceLock};
use tracing::{error, info, warn};

pub struct EmbedContext {
ollama: ArcSwapOption<Ollama>,
client: ArcSwapOption<AIClient>,
/// Embedding model name for the active client. Provider-specific, so it travels with the client.
embedding_model: ArcSwapOption<String>,
vector_db: ArcSwapOption<VectorSqliteDB>,
scheduler: ArcSwapOption<EmbeddingScheduler>,
}
Expand All @@ -19,7 +21,8 @@ impl EmbedContext {
static INSTANCE: OnceLock<Arc<EmbedContext>> = OnceLock::new();
INSTANCE.get_or_init(|| {
Arc::new(EmbedContext {
ollama: ArcSwapOption::empty(),
client: ArcSwapOption::empty(),
embedding_model: ArcSwapOption::empty(),
vector_db: ArcSwapOption::empty(),
scheduler: Default::default(),
})
Expand Down Expand Up @@ -50,19 +53,31 @@ impl EmbedContext {
}
}

pub fn set_ollama(&self, ollama: Option<Arc<Ollama>>) {
if let Some(ollama) = ollama {
if let Some(o) = self.ollama.load().as_ref() {
if o.uri() == ollama.uri() {
info!("[Embedding] Ollama does not change");
return;
}
pub fn set_client(&self, client: Option<Arc<AIClient>>, embedding_model: &str) {
if let Some(client) = client {
let unchanged = self
.client
.load()
.as_ref()
.is_some_and(|c| c.same_as(&client))
&& self
.embedding_model
.load()
.as_ref()
.is_some_and(|m| m.as_str() == embedding_model);
if unchanged {
info!("[Embedding] ai client does not change");
return;
}

self.ollama.store(Some(ollama));
self.client.store(Some(client));
self
.embedding_model
.store(Some(Arc::new(embedding_model.to_string())));
self.try_create_scheduler();
} else {
self.ollama.store(None);
self.client.store(None);
self.embedding_model.store(None);
if let Some(s) = self.scheduler.swap(None) {
info!("[Embedding] Stopping scheduler");
let _ = s.stop_tx.send(());
Expand All @@ -80,17 +95,21 @@ impl EmbedContext {
}

fn try_create_scheduler(&self) {
if let (Some(ollama), Some(vector_db)) = (self.ollama.load_full(), self.vector_db.load_full()) {
if let (Some(client), Some(model), Some(vector_db)) = (
self.client.load_full(),
self.embedding_model.load_full(),
self.vector_db.load_full(),
) {
info!("[Embedding] Creating scheduler");
match EmbeddingScheduler::new(ollama, vector_db) {
match EmbeddingScheduler::new(client, model.to_string(), vector_db) {
Ok(s) => {
info!("[Embedding] create scheduler successfully");
self.scheduler.store(Some(s));
},
Err(err) => error!("[Embedding] Failed to create scheduler: {}", err),
}
} else {
info!("[Embedding] Ollama or vector db is not initialized, remove embedding scheduler");
info!("[Embedding] ai client or vector db is not initialized, remove embedding scheduler");
self.scheduler.store(None);
}
}
Expand Down
15 changes: 5 additions & 10 deletions frontend/rust-lib/flowy-ai/src/embeddings/document_indexer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@ use crate::embeddings::indexer::{EmbeddingModel, Indexer};
use flowy_ai_pub::entities::{EmbeddedChunk, SOURCE, SOURCE_ID, SOURCE_NAME};
use flowy_error::FlowyError;
use lib_infra::async_trait::async_trait;
use ollama_rs::generation::embeddings::request::{EmbeddingsInput, GenerateEmbeddingsRequest};
use serde_json::json;
use text_splitter::{ChunkConfig, TextSplitter};
use tracing::{debug, error, trace, warn};
Expand Down Expand Up @@ -54,25 +53,21 @@ impl Indexer for DocumentIndexer {
contents.push(chunks[i].content.as_ref().unwrap().to_owned());
}

let request = GenerateEmbeddingsRequest::new(
embedder.model().name().to_string(),
EmbeddingsInput::Multiple(contents),
);
let resp = embedder.embed(request).await?;
if resp.embeddings.len() != valid_indices.len() {
let embeddings = embedder.embed(contents).await?;
if embeddings.len() != valid_indices.len() {
error!(
"[Embedding] requested {} embeddings, received {} embeddings",
valid_indices.len(),
resp.embeddings.len()
embeddings.len()
);
return Err(FlowyError::internal().with_context(format!(
"Mismatch in number of embeddings requested and received: {} vs {}",
valid_indices.len(),
resp.embeddings.len()
embeddings.len()
)));
}

for (index, embedding) in resp.embeddings.into_iter().enumerate() {
for (index, embedding) in embeddings.into_iter().enumerate() {
let chunk_idx = valid_indices[index];
chunks[chunk_idx].embeddings = Some(embedding);
}
Expand Down
45 changes: 17 additions & 28 deletions frontend/rust-lib/flowy-ai/src/embeddings/embedder.rs
Original file line number Diff line number Diff line change
@@ -1,41 +1,30 @@
use crate::embeddings::indexer::EmbeddingModel;
use crate::local_ai::client::AIClient;
use flowy_error::FlowyResult;
use ollama_rs::Ollama;
use ollama_rs::generation::embeddings::GenerateEmbeddingsResponse;
use ollama_rs::generation::embeddings::request::GenerateEmbeddingsRequest;
use std::sync::Arc;

/// Pairs the configured backend with the embedding model to use.
///
/// The model name lives here because it is provider-specific (`nomic-embed-text:latest` for Ollama
/// versus `text-embedding-3-small` for OpenAI-compatible endpoints) and callers should not need to
/// know which backend is active. Regardless of provider the vectors are always
/// [`crate::local_ai::client::EMBEDDING_DIMENSION`] long.
#[derive(Debug, Clone)]
pub enum Embedder {
Ollama(OllamaEmbedder),
pub struct Embedder {
client: Arc<AIClient>,
model: String,
}

impl Embedder {
pub async fn embed(
&self,
request: GenerateEmbeddingsRequest,
) -> FlowyResult<GenerateEmbeddingsResponse> {
match self {
Embedder::Ollama(ollama) => ollama.embed(request).await,
}
pub fn new(client: Arc<AIClient>, model: String) -> Self {
Self { client, model }
}

pub fn model(&self) -> EmbeddingModel {
EmbeddingModel::NomicEmbedText
/// Returns one vector per input string, in the same order.
pub async fn embed(&self, input: Vec<String>) -> FlowyResult<Vec<Vec<f32>>> {
self.client.embed(&self.model, input).await
}
}

#[derive(Debug, Clone)]
pub struct OllamaEmbedder {
pub ollama: Arc<Ollama>,
}

impl OllamaEmbedder {
pub async fn embed(
&self,
request: GenerateEmbeddingsRequest,
) -> FlowyResult<GenerateEmbeddingsResponse> {
let resp = self.ollama.generate_embeddings(request).await?;
Ok(resp)
pub fn model_name(&self) -> &str {
&self.model
}
}
37 changes: 17 additions & 20 deletions frontend/rust-lib/flowy-ai/src/embeddings/scheduler.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
use crate::embeddings::embedder::{Embedder, OllamaEmbedder};
use crate::embeddings::indexer::IndexerProvider;
use crate::embeddings::embedder::Embedder;
use crate::embeddings::indexer::{EmbeddingModel, IndexerProvider};
use crate::local_ai::client::AIClient;
use crate::search::summary::{LLMDocument, summarize_documents};
use flowy_ai_pub::cloud::search_dto::{
SearchContentType, SearchDocumentResponseItem, SearchResult, SearchSummaryResult, Summary,
Expand All @@ -8,8 +9,6 @@ use flowy_ai_pub::entities::{EmbeddingRecord, UnindexedCollab, UnindexedData};
use flowy_error::{ErrorCode, FlowyError, FlowyResult};
use flowy_sqlite::internal::derives::multiconnection::chrono::Utc;
use flowy_sqlite_vec::db::VectorSqliteDB;
use ollama_rs::Ollama;
use ollama_rs::generation::embeddings::request::{EmbeddingsInput, GenerateEmbeddingsRequest};
use std::sync::{Arc, Weak};
use tokio::select;
use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
Expand All @@ -23,14 +22,16 @@ pub struct EmbeddingScheduler {
indexer_provider: Arc<IndexerProvider>,
write_embedding_tx: UnboundedSender<EmbeddingRecord>,
generate_embedding_tx: mpsc::Sender<UnindexedCollab>,
ollama: Arc<Ollama>,
client: Arc<AIClient>,
embedding_model: String,
vector_db: Arc<VectorSqliteDB>,
pub(crate) stop_tx: tokio::sync::broadcast::Sender<()>,
}

impl EmbeddingScheduler {
pub fn new(
ollama: Arc<Ollama>,
client: Arc<AIClient>,
embedding_model: String,
vector_db: Arc<VectorSqliteDB>,
) -> FlowyResult<Arc<EmbeddingScheduler>> {
let indexer_provider = IndexerProvider::new();
Expand All @@ -42,7 +43,8 @@ impl EmbeddingScheduler {
indexer_provider,
write_embedding_tx,
generate_embedding_tx,
ollama,
client,
embedding_model,
vector_db,
stop_tx,
});
Expand All @@ -67,10 +69,10 @@ impl EmbeddingScheduler {
}

pub(crate) fn create_embedder(&self) -> Result<Embedder, FlowyError> {
let embedder = Embedder::Ollama(OllamaEmbedder {
ollama: self.ollama.clone(),
});
Ok(embedder)
Ok(Embedder::new(
self.client.clone(),
self.embedding_model.clone(),
))
}

pub async fn index_collab(&self, data: UnindexedCollab) -> FlowyResult<()> {
Expand Down Expand Up @@ -99,13 +101,8 @@ impl EmbeddingScheduler {
query: &str,
) -> FlowyResult<Vec<SearchDocumentResponseItem>> {
let embedder = self.create_embedder()?;
let request = GenerateEmbeddingsRequest::new(
embedder.model().name().to_string(),
EmbeddingsInput::Single(query.to_string()),
);

let resp = embedder.embed(request).await?;
match resp.embeddings.first() {
let embeddings = embedder.embed(vec![query.to_string()]).await?;
match embeddings.first() {
None => Ok(vec![]),
Some(query_embed) => {
let result = self
Expand Down Expand Up @@ -155,7 +152,7 @@ impl EmbeddingScheduler {
})
.collect::<Vec<_>>();

let resp = summarize_documents(&self.ollama, question, model_name, docs)
let resp = summarize_documents(&self.client, question, model_name, docs)
.await
.map_err(|err| {
error!("[Embedding] Failed to generate summary: {}", err);
Expand Down Expand Up @@ -295,7 +292,7 @@ async fn spawn_generate_embeddings(
match indexer.create_embedded_chunks_from_text(
record.object_id,
paragraphs,
embedder.model(),
EmbeddingModel::NomicEmbedText,
) {
Ok(mut chunks) => {
if let Some(fragment_ids) = existing_embeddings.get(&record.object_id) {
Expand Down
33 changes: 16 additions & 17 deletions frontend/rust-lib/flowy-ai/src/embeddings/store.rs
Original file line number Diff line number Diff line change
@@ -1,18 +1,17 @@
use crate::embeddings::document_indexer::split_text_into_chunks;
use crate::embeddings::embedder::{Embedder, OllamaEmbedder};
use crate::embeddings::embedder::Embedder;
use crate::embeddings::indexer::{EmbeddingModel, IndexerProvider};
use crate::local_ai::chat::retriever::MultipleSourceRetrieverStore;
use crate::local_ai::client::AIClient;
use async_trait::async_trait;
use flowy_ai_pub::cloud::CollabType;
use flowy_ai_pub::entities::{RAG_IDS, SOURCE_ID};
use flowy_error::{FlowyError, FlowyResult};
use flowy_sqlite_vec::db::VectorSqliteDB;
use flowy_sqlite_vec::entities::{EmbeddedContent, SqliteEmbeddedDocument};
use futures::stream::{self, StreamExt};
use langchain_rust::llm::client::OllamaClient;
use langchain_rust::schemas::Document;
use langchain_rust::vectorstore::{VecStoreOptions, VectorStore};
use ollama_rs::generation::embeddings::request::{EmbeddingsInput, GenerateEmbeddingsRequest};
use serde_json::Value;
use std::collections::HashMap;
use std::error::Error;
Expand All @@ -22,28 +21,33 @@ use uuid::Uuid;

#[derive(Clone)]
pub struct SqliteVectorStore {
ollama: Weak<OllamaClient>,
client: Weak<AIClient>,
embedding_model: String,
vector_db: Weak<VectorSqliteDB>,
indexer_provider: Arc<IndexerProvider>,
}

impl SqliteVectorStore {
pub fn new(ollama: Weak<OllamaClient>, vector_db: Weak<VectorSqliteDB>) -> Self {
pub fn new(
client: Weak<AIClient>,
embedding_model: String,
vector_db: Weak<VectorSqliteDB>,
) -> Self {
Self {
ollama,
client,
embedding_model,
vector_db,
indexer_provider: IndexerProvider::new(),
}
}

pub(crate) fn create_embedder(&self) -> Result<Embedder, FlowyError> {
let ollama = self
.ollama
let client = self
.client
.upgrade()
.ok_or_else(|| FlowyError::internal().with_context("Ollama reference was dropped"))?;
.ok_or_else(|| FlowyError::internal().with_context("AI client reference was dropped"))?;

let embedder = Embedder::Ollama(OllamaEmbedder { ollama });
Ok(embedder)
Ok(Embedder::new(client, self.embedding_model.clone()))
}

pub(crate) async fn select_all_embedded_documents(
Expand Down Expand Up @@ -107,12 +111,7 @@ impl MultipleSourceRetrieverStore for SqliteVectorStore {

// Create embedder and generate embedding for query
let embedder = self.create_embedder()?;
let request = GenerateEmbeddingsRequest::new(
embedder.model().name().to_string(),
EmbeddingsInput::Single(query.to_string()),
);

let embedding = embedder.embed(request).await?.embeddings;
let embedding = embedder.embed(vec![query.to_string()]).await?;
if embedding.is_empty() {
return Ok(Vec::new());
}
Expand Down
Loading
Loading