use crate::{embedding::EmbeddingProvider, parsing::Document, JobHandle}; use gpui::executor::Background; use parking_lot::Mutex; use smol::channel; use std::{mem, ops::Range, path::Path, sync::Arc, time::SystemTime}; #[derive(Clone)] pub struct FileToEmbed { pub worktree_id: i64, pub path: Arc, pub mtime: SystemTime, pub documents: Vec, pub job_handle: JobHandle, } impl std::fmt::Debug for FileToEmbed { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("FileToEmbed") .field("worktree_id", &self.worktree_id) .field("path", &self.path) .field("mtime", &self.mtime) .field("document", &self.documents) .finish_non_exhaustive() } } impl PartialEq for FileToEmbed { fn eq(&self, other: &Self) -> bool { self.worktree_id == other.worktree_id && self.path == other.path && self.mtime == other.mtime && self.documents == other.documents } } pub struct EmbeddingQueue { embedding_provider: Arc, pending_batch: Vec, executor: Arc, pending_batch_token_count: usize, finished_files_tx: channel::Sender, finished_files_rx: channel::Receiver, } #[derive(Clone)] pub struct FileToEmbedFragment { file: Arc>, document_range: Range, } impl EmbeddingQueue { pub fn new(embedding_provider: Arc, executor: Arc) -> Self { let (finished_files_tx, finished_files_rx) = channel::unbounded(); Self { embedding_provider, executor, pending_batch: Vec::new(), pending_batch_token_count: 0, finished_files_tx, finished_files_rx, } } pub fn push(&mut self, file: FileToEmbed) { if file.documents.is_empty() { self.finished_files_tx.try_send(file).unwrap(); return; } let file = Arc::new(Mutex::new(file)); self.pending_batch.push(FileToEmbedFragment { file: file.clone(), document_range: 0..0, }); let mut fragment_range = &mut self.pending_batch.last_mut().unwrap().document_range; let mut saved_tokens = 0; for (ix, document) in file.lock().documents.iter().enumerate() { let document_token_count = if document.embedding.is_none() { document.token_count } else { saved_tokens += document.token_count; 0 }; let next_token_count = self.pending_batch_token_count + document_token_count; if next_token_count > self.embedding_provider.max_tokens_per_batch() { let range_end = fragment_range.end; self.flush(); self.pending_batch.push(FileToEmbedFragment { file: file.clone(), document_range: range_end..range_end, }); fragment_range = &mut self.pending_batch.last_mut().unwrap().document_range; } fragment_range.end = ix + 1; self.pending_batch_token_count += document_token_count; } log::trace!("Saved Tokens: {:?}", saved_tokens); } pub fn flush(&mut self) { let batch = mem::take(&mut self.pending_batch); self.pending_batch_token_count = 0; if batch.is_empty() { return; } let finished_files_tx = self.finished_files_tx.clone(); let embedding_provider = self.embedding_provider.clone(); self.executor.spawn(async move { let mut spans = Vec::new(); let mut document_count = 0; for fragment in &batch { let file = fragment.file.lock(); document_count += file.documents[fragment.document_range.clone()].len(); spans.extend( { file.documents[fragment.document_range.clone()] .iter().filter(|d| d.embedding.is_none()) .map(|d| d.content.clone()) } ); } log::trace!("Documents Length: {:?}", document_count); log::trace!("Span Length: {:?}", spans.clone().len()); // If spans is 0, just send the fragment to the finished files if its the last one. if spans.len() == 0 { for fragment in batch.clone() { if let Some(file) = Arc::into_inner(fragment.file) { finished_files_tx.try_send(file.into_inner()).unwrap(); } } return; }; match embedding_provider.embed_batch(spans).await { Ok(embeddings) => { let mut embeddings = embeddings.into_iter(); for fragment in batch { for document in &mut fragment.file.lock().documents[fragment.document_range.clone()].iter_mut().filter(|d| d.embedding.is_none()) { if let Some(embedding) = embeddings.next() { document.embedding = Some(embedding); } else { // log::error!("number of embeddings returned different from number of documents"); } } if let Some(file) = Arc::into_inner(fragment.file) { finished_files_tx.try_send(file.into_inner()).unwrap(); } } } Err(error) => { log::error!("{:?}", error); } } }) .detach(); } pub fn finished_files(&self) -> channel::Receiver { self.finished_files_rx.clone() } }