use std::sync::Arc;
use jammi_ai::session::InferenceSession;
use jammi_ai::Session;
use tonic::{Request, Response, Status};
use std::collections::HashMap;
use arrow::array::{Array, Float32Array, RecordBatch, StringArray};
use arrow::util::display::{ArrayFormatter, FormatOptions};
use crate::grpc::proto::embedding::embedding_service_server::EmbeddingService;
use crate::grpc::proto::embedding::{
EncodeQueryRequest, EncodeQueryResponse, GenerateEmbeddingsRequest, ImportEmbeddingsRequest,
ResultTable, SearchHit, SearchRequest, SearchResponse,
};
use crate::grpc::wire::{map_engine_error, scoped, session_tenant_traced};
pub struct EmbeddingServer {
session: Arc<InferenceSession>,
}
impl EmbeddingServer {
pub fn new(session: Arc<InferenceSession>) -> Self {
Self { session }
}
fn local(&self) -> Session {
Session::new(Arc::clone(&self.session))
}
}
#[tonic::async_trait]
impl EmbeddingService for EmbeddingServer {
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn generate_embeddings(
&self,
request: Request<GenerateEmbeddingsRequest>,
) -> Result<Response<ResultTable>, Status> {
let tenant = session_tenant_traced(&request);
let args = jammi_ai::wire::generate_embeddings_from_proto(request.into_inner())?;
let session = self.local();
let (record, outcome) = scoped(&self.session, tenant, || {
session.generate_embeddings(
&args.source_id,
&args.model_id,
&args.columns,
&args.key_column,
args.modality,
args.cache,
)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(jammi_wire::result_table_with_outcome(
record,
jammi_ai::wire::cache_outcome_to_proto(&outcome),
)))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn import_embeddings(
&self,
request: Request<ImportEmbeddingsRequest>,
) -> Result<Response<ResultTable>, Status> {
let tenant = session_tenant_traced(&request);
let args = jammi_ai::wire::import_embeddings_from_proto(request.into_inner())?;
let session = self.local();
let record = scoped(&self.session, tenant, || {
session.import_embeddings(
&args.source_id,
&args.model_id,
&args.vectors_url,
&args.key_column,
&args.text_columns,
args.dimensions,
)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(jammi_wire::result_table_with_outcome(
record,
jammi_ai::wire::cache_outcome_to_proto(&jammi_db::store::CacheOutcome::Computed),
)))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn encode_query(
&self,
request: Request<EncodeQueryRequest>,
) -> Result<Response<EncodeQueryResponse>, Status> {
let tenant = session_tenant_traced(&request);
let args = jammi_ai::wire::encode_query_from_proto(request.into_inner())?;
let session = self.local();
let embedding = scoped(&self.session, tenant, || {
session.encode_query(&args.model_id, args.input, args.modality)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(EncodeQueryResponse { embedding }))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn search(
&self,
request: Request<SearchRequest>,
) -> Result<Response<SearchResponse>, Status> {
let tenant = session_tenant_traced(&request);
let mut request = jammi_ai::wire::search_from_proto(request.into_inner())?;
let select = std::mem::take(&mut request.select);
request.select = search_select(&select);
let session = self.local();
let batches = scoped(&self.session, tenant, || session.search(request))
.await
.map_err(map_engine_error)?;
let hits = batches_to_hits(&batches, &select)?;
Ok(Response::new(SearchResponse { hits }))
}
}
fn search_select(select: &[String]) -> Vec<String> {
if select.is_empty() {
return Vec::new();
}
let mut columns: Vec<String> = vec!["_row_id".to_string(), "similarity".to_string()];
for name in select {
if name != "_row_id" && name != "similarity" {
columns.push(name.clone());
}
}
columns
}
fn batches_to_hits(batches: &[RecordBatch], select: &[String]) -> Result<Vec<SearchHit>, Status> {
let mut hits = Vec::new();
let format = FormatOptions::default();
for batch in batches {
let keys = column_as::<StringArray>(batch, "_row_id")?;
let scores = column_as::<Float32Array>(batch, "similarity")?;
let formatters: Vec<(String, ArrayFormatter)> = select
.iter()
.map(|name| {
let array = batch.column_by_name(name).ok_or_else(|| {
Status::invalid_argument(format!("select column '{name}' not in results"))
})?;
let formatter = ArrayFormatter::try_new(array.as_ref(), &format)
.map_err(|e| Status::internal(format!("format column '{name}': {e}")))?;
Ok((name.clone(), formatter))
})
.collect::<Result<_, Status>>()?;
for row in 0..batch.num_rows() {
let columns: HashMap<String, String> = formatters
.iter()
.map(|(name, fmt)| (name.clone(), fmt.value(row).to_string()))
.collect();
hits.push(SearchHit {
key: keys.value(row).to_string(),
score: scores.value(row),
columns,
});
}
}
Ok(hits)
}
fn column_as<'a, A: Array + 'static>(batch: &'a RecordBatch, name: &str) -> Result<&'a A, Status> {
batch
.column_by_name(name)
.ok_or_else(|| Status::internal(format!("search result missing '{name}' column")))?
.as_any()
.downcast_ref::<A>()
.ok_or_else(|| {
Status::internal(format!("search result '{name}' column has unexpected type"))
})
}