use std::sync::Arc;
use jammi_ai::session::InferenceSession;
use jammi_ai::{Modality, SearchQuery, SearchRequest as SessionSearch, Session};
use jammi_wire::ProtoQueryInput;
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::{
search_request::Query as ProtoQuery, EncodeQueryRequest, EncodeQueryResponse,
GenerateEmbeddingsRequest, ResultTable, SearchHit, SearchRequest, SearchResponse,
};
use crate::grpc::wire::{map_engine_error, require_nonempty, scoped, session_tenant};
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 {
async fn generate_embeddings(
&self,
request: Request<GenerateEmbeddingsRequest>,
) -> Result<Response<ResultTable>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.source_id, "source_id")?;
require_nonempty(&req.model_id, "model_id")?;
require_nonempty(&req.key_column, "key_column")?;
if req.columns.is_empty() {
return Err(Status::invalid_argument("columns is required"));
}
let modality = Modality::try_from(req.modality)?;
let session = self.local();
let record = scoped(&self.session, tenant, || {
session.generate_embeddings(
&req.source_id,
&req.model_id,
&req.columns,
&req.key_column,
modality,
)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(ResultTable::from(record)))
}
async fn encode_query(
&self,
request: Request<EncodeQueryRequest>,
) -> Result<Response<EncodeQueryResponse>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.model_id, "model_id")?;
let modality = Modality::try_from(req.modality)?;
let input = ProtoQueryInput {
input: req.input,
modality,
}
.try_into()?;
let session = self.local();
let embedding = scoped(&self.session, tenant, || {
session.encode_query(&req.model_id, input, modality)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(EncodeQueryResponse { embedding }))
}
async fn search(
&self,
request: Request<SearchRequest>,
) -> Result<Response<SearchResponse>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.source_id, "source_id")?;
let query = match req.query.ok_or_else(|| {
Status::invalid_argument("query (query_vector or row_key) is required")
})? {
ProtoQuery::QueryVector(v) => SearchQuery::Vector(v.values),
ProtoQuery::RowKey(key) => SearchQuery::RowKey(key),
};
let select = req.select;
let request = SessionSearch {
source_id: req.source_id,
query,
k: req.k as usize,
filter: req.filter,
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"))
})
}