use std::sync::Arc;
use jammi_ai::session::InferenceSession;
use jammi_ai::wire::{
assemble_context_request_from_proto, assemble_context_to_proto,
build_neighbor_graph_from_proto, propagate_request_from_proto,
};
use tonic::{Request, Response, Status};
use crate::grpc::proto::embedding::ResultTable;
use crate::grpc::proto::pipeline::pipeline_service_server::PipelineService;
use crate::grpc::proto::pipeline::{
AssembleContextRequest, AssembleContextResponse, BuildNeighborGraphRequest,
PropagateEmbeddingsRequest,
};
use crate::grpc::wire::{map_engine_error, scoped, session_tenant_traced};
pub struct PipelineServer {
session: Arc<InferenceSession>,
}
impl PipelineServer {
pub fn new(session: Arc<InferenceSession>) -> Self {
Self { session }
}
}
#[tonic::async_trait]
impl PipelineService for PipelineServer {
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn build_neighbor_graph(
&self,
request: Request<BuildNeighborGraphRequest>,
) -> Result<Response<ResultTable>, Status> {
let tenant = session_tenant_traced(&request);
let args = build_neighbor_graph_from_proto(request.into_inner())?;
let record = scoped(&self.session, tenant, || async {
self.session
.build_neighbor_graph(
&args.source_id,
args.embedding_table.as_deref(),
&args.params,
)
.await
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(record.into()))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn propagate_embeddings(
&self,
request: Request<PropagateEmbeddingsRequest>,
) -> Result<Response<ResultTable>, Status> {
let tenant = session_tenant_traced(&request);
let req = propagate_request_from_proto(request.into_inner())?;
let record = scoped(&self.session, tenant, || async {
self.session.propagate_embeddings(&req).await
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(record.into()))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn assemble_context(
&self,
request: Request<AssembleContextRequest>,
) -> Result<Response<AssembleContextResponse>, Status> {
let tenant = session_tenant_traced(&request);
let req = assemble_context_request_from_proto(request.into_inner())?;
let context = scoped(&self.session, tenant, || async {
self.session.assemble_context(&req).await
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(assemble_context_to_proto(context)?))
}
}