use crate::error::{Result, XbergError};
use crate::types::Chunk;
#[allow(dead_code)]
pub(crate) fn apply_sparse_vectors(chunks: &mut [Chunk], vectors: Vec<crate::SparseEmbedding>) -> Result<()> {
if chunks.len() != vectors.len() {
return Err(XbergError::Validation {
message: format!(
"Sparse-embedding generation returned {got} vectors for {expected} chunks; refusing to \
attach vectors because a positional zip would misalign them with the wrong chunks",
got = vectors.len(),
expected = chunks.len(),
),
source: None,
});
}
for (chunk, vector) in chunks.iter_mut().zip(vectors) {
chunk.sparse_embedding = Some(vector);
}
Ok(())
}
#[cfg(feature = "sparse-embeddings")]
pub(crate) fn generate_sparse_vectors_for_chunks(
chunks: &mut [Chunk],
config: &crate::core::config::SparseEmbeddingConfig,
) -> Result<()> {
if chunks.is_empty() {
return Ok(());
}
let texts: Vec<&str> = chunks.iter().map(|c| c.content.as_str()).collect();
let vectors = crate::sparse_embeddings::embed_sparse(&texts, config)?;
apply_sparse_vectors(chunks, vectors)
}
#[allow(dead_code)]
pub(crate) fn apply_late_interaction_vectors(
chunks: &mut [Chunk],
vectors: Vec<crate::MultiVectorEmbedding>,
) -> Result<()> {
if chunks.len() != vectors.len() {
return Err(XbergError::Validation {
message: format!(
"Late-interaction embedding generation returned {got} vectors for {expected} chunks; \
refusing to attach vectors because a positional zip would misalign them with the wrong \
chunks",
got = vectors.len(),
expected = chunks.len(),
),
source: None,
});
}
for (chunk, vector) in chunks.iter_mut().zip(vectors) {
chunk.late_interaction = Some(vector);
}
Ok(())
}
#[cfg(feature = "late-interaction")]
pub(crate) fn generate_late_interaction_vectors_for_chunks(
chunks: &mut [Chunk],
config: &crate::core::config::LateInteractionConfig,
) -> Result<()> {
if chunks.is_empty() {
return Ok(());
}
let texts: Vec<&str> = chunks.iter().map(|c| c.content.as_str()).collect();
let vectors = crate::late_interaction::embed_multi_vector(&texts, config, false)?;
apply_late_interaction_vectors(chunks, vectors)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{ChunkMetadata, ChunkType};
fn test_chunk(content: &str) -> Chunk {
Chunk {
content: content.to_string(),
chunk_type: ChunkType::default(),
embedding: None,
sparse_embedding: None,
late_interaction: None,
metadata: ChunkMetadata {
byte_start: 0,
byte_end: content.len(),
token_count: None,
chunk_index: 0,
total_chunks: 1,
first_page: None,
last_page: None,
heading_context: None,
heading_path: Vec::new(),
image_indices: Vec::new(),
node_ids: Vec::new(),
page_spans: Vec::new(),
classifications: Vec::new(),
},
}
}
#[test]
fn should_attach_exact_sparse_vector_to_each_chunk_when_applied() {
let mut chunks = vec![test_chunk("alpha"), test_chunk("beta")];
let vectors = vec![
crate::SparseEmbedding {
indices: vec![1, 4],
values: vec![0.9, 0.1],
},
crate::SparseEmbedding {
indices: vec![2],
values: vec![0.5],
},
];
apply_sparse_vectors(&mut chunks, vectors).expect("lengths match");
assert_eq!(chunks[0].sparse_embedding.as_ref().unwrap().indices, vec![1, 4]);
assert_eq!(chunks[0].sparse_embedding.as_ref().unwrap().values, vec![0.9, 0.1]);
assert_eq!(chunks[1].sparse_embedding.as_ref().unwrap().indices, vec![2]);
assert_eq!(chunks[1].sparse_embedding.as_ref().unwrap().values, vec![0.5]);
assert!(
chunks[0].late_interaction.is_none(),
"late_interaction must stay untouched"
);
}
#[test]
fn should_leave_sparse_embedding_none_when_apply_is_never_called() {
let chunks = [test_chunk("alpha"), test_chunk("beta")];
assert!(chunks.iter().all(|c| c.sparse_embedding.is_none()));
}
#[test]
fn should_error_and_leave_chunks_untouched_when_sparse_vector_count_mismatches() {
let mut chunks = vec![test_chunk("alpha"), test_chunk("beta")];
let vectors = vec![crate::SparseEmbedding {
indices: vec![1],
values: vec![1.0],
}];
let err = apply_sparse_vectors(&mut chunks, vectors).expect_err("length mismatch must error");
assert!(err.to_string().contains("1 vectors for 2 chunks"), "got: {err}");
assert!(
chunks.iter().all(|c| c.sparse_embedding.is_none()),
"chunks must be untouched on error"
);
}
#[test]
fn should_attach_exact_late_interaction_vector_to_each_chunk_when_applied() {
let mut chunks = vec![test_chunk("alpha")];
let vectors = vec![crate::MultiVectorEmbedding {
num_tokens: 2,
dim: 2,
data: vec![0.1, 0.2, 0.3, 0.4],
}];
apply_late_interaction_vectors(&mut chunks, vectors).expect("lengths match");
let late = chunks[0].late_interaction.as_ref().unwrap();
assert_eq!(late.num_tokens, 2);
assert_eq!(late.dim, 2);
assert_eq!(late.data, vec![0.1, 0.2, 0.3, 0.4]);
assert!(
chunks[0].sparse_embedding.is_none(),
"sparse_embedding must stay untouched"
);
}
#[test]
fn should_leave_late_interaction_none_when_apply_is_never_called() {
let chunks = [test_chunk("alpha")];
assert!(chunks.iter().all(|c| c.late_interaction.is_none()));
}
#[test]
fn should_error_and_leave_chunks_untouched_when_late_interaction_vector_count_mismatches() {
let mut chunks = vec![test_chunk("alpha"), test_chunk("beta")];
let vectors = vec![];
let err = apply_late_interaction_vectors(&mut chunks, vectors).expect_err("length mismatch must error");
assert!(err.to_string().contains("0 vectors for 2 chunks"), "got: {err}");
assert!(
chunks.iter().all(|c| c.late_interaction.is_none()),
"chunks must be untouched on error"
);
}
}