use std::sync::Arc;
use frankensearch_core::generation::EmbeddingIdentityBundleV1;
use frankensearch_core::traits::{IdentityBoundEmbedding, ModelCategory, ModelTier, SearchFuture};
use super::checkpoint;
use crate::{Cx, Embedder, SearchError, SearchResult};
pub(super) fn wrap(
inner: Arc<dyn Embedder>,
identity: EmbeddingIdentityBundleV1,
) -> Arc<dyn Embedder> {
Arc::new(SplittingEmbedder { inner, identity })
}
struct SplittingEmbedder {
inner: Arc<dyn Embedder>,
identity: EmbeddingIdentityBundleV1,
}
impl SplittingEmbedder {
fn identity_error() -> SearchError {
SearchError::UnverifiableRemoteSpace {
producer: "native_ann.builder.batch".to_owned(),
reason: "batch inference no longer matches the admitted producer contract".to_owned(),
}
}
fn admit_producer(&self) -> SearchResult<()> {
let current = self.inner.identity().map_err(|_| Self::identity_error())?;
if current != &self.identity
|| usize::try_from(self.identity.space.dimension).ok() != Some(self.inner.dimension())
{
return Err(Self::identity_error());
}
Ok(())
}
fn admit_response(&self, response: &IdentityBoundEmbedding) -> SearchResult<()> {
if response.identity != self.identity {
return Err(Self::identity_error());
}
response.validate()
}
async fn split_batch(
&self,
cx: &Cx,
texts: &[&str],
) -> SearchResult<Vec<IdentityBoundEmbedding>> {
checkpoint(cx, "native_ann.builder.before_batch")?;
self.admit_producer()?;
let mut pending = Vec::new();
pending.push(0..texts.len());
let mut accepted = Vec::with_capacity(texts.len());
while let Some(range) = pending.pop() {
checkpoint(cx, "native_ann.builder.before_batch")?;
self.admit_producer()?;
let outcome = self
.inner
.embed_batch_bound(cx, &texts[range.clone()])
.await;
checkpoint(cx, "native_ann.builder.after_batch")?;
let outcome = match outcome {
Err(error @ SearchError::Cancelled { .. }) => return Err(error),
outcome => outcome,
};
self.admit_producer()?;
let mut response = match outcome {
Ok(response) => response,
Err(SearchError::EmbeddingFailed { .. }) if range.len() > 1 => {
let middle = range.start + range.len() / 2;
tracing::debug!(
batch_size = range.len(),
left_size = middle - range.start,
"splitting rejected bound embedding batch with the same producer"
);
pending.push(middle..range.end);
pending.push(range.start..middle);
continue;
}
Err(error) => return Err(error),
};
if response.len() != range.len() {
return Err(SearchError::InvalidConfig {
field: "native_ann.builder.batch_cardinality".to_owned(),
value: response.len().to_string(),
reason: format!("expected exactly {} bound outputs", range.len()),
});
}
for value in &response {
checkpoint(cx, "native_ann.builder.admit_output")?;
self.admit_response(value)?;
}
accepted.append(&mut response);
}
checkpoint(cx, "native_ann.builder.batch_complete")?;
self.admit_producer()?;
Ok(accepted)
}
}
impl Embedder for SplittingEmbedder {
fn embed<'a>(&'a self, cx: &'a Cx, text: &'a str) -> SearchFuture<'a, Vec<f32>> {
self.inner.embed(cx, text)
}
fn embed_batch<'a>(
&'a self,
cx: &'a Cx,
texts: &'a [&'a str],
) -> SearchFuture<'a, Vec<Vec<f32>>> {
self.inner.embed_batch(cx, texts)
}
fn embed_bound<'a>(
&'a self,
cx: &'a Cx,
text: &'a str,
) -> SearchFuture<'a, IdentityBoundEmbedding> {
self.inner.embed_bound(cx, text)
}
fn embed_batch_bound<'a>(
&'a self,
cx: &'a Cx,
texts: &'a [&'a str],
) -> SearchFuture<'a, Vec<IdentityBoundEmbedding>> {
Box::pin(self.split_batch(cx, texts))
}
fn identity(&self) -> SearchResult<&EmbeddingIdentityBundleV1> {
self.inner.identity()
}
fn bound_batch_is_native(&self) -> bool {
self.inner.bound_batch_is_native()
}
fn dimension(&self) -> usize {
self.inner.dimension()
}
fn id(&self) -> &str {
self.inner.id()
}
fn model_name(&self) -> &str {
self.inner.model_name()
}
fn is_ready(&self) -> bool {
self.inner.is_ready()
}
fn is_semantic(&self) -> bool {
self.inner.is_semantic()
}
fn category(&self) -> ModelCategory {
self.inner.category()
}
fn tier(&self) -> ModelTier {
self.inner.tier()
}
fn supports_mrl(&self) -> bool {
self.inner.supports_mrl()
}
}
#[cfg(test)]
mod tests;