use std::sync::Arc;
use jammi_ai::session::InferenceSession;
use jammi_ai::wire::{
asof_join_from_proto, assemble_context_request_from_proto, assemble_context_to_proto,
build_neighbor_graph_from_proto, propagate_request_from_proto, recompute_from_proto,
recompute_report_to_proto,
};
use jammi_db::error::JammiError;
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::{
AsofJoinRequest, AssembleContextRequest, AssembleContextResponse, BuildNeighborGraphRequest,
PropagateEmbeddingsRequest, RecomputeReport as ProtoRecomputeReport, RecomputeRequest,
};
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, outcome) = scoped(&self.session, tenant, || async {
self.session
.build_neighbor_graph(
&args.source_id,
args.embedding_table.as_deref(),
&args.params,
args.cache,
)
.await
})
.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 propagate_embeddings(
&self,
request: Request<PropagateEmbeddingsRequest>,
) -> Result<Response<ResultTable>, Status> {
let tenant = session_tenant_traced(&request);
let (req, cache) = propagate_request_from_proto(request.into_inner())?;
let (record, outcome) = scoped(&self.session, tenant, || async {
self.session.propagate_embeddings(&req, cache).await
})
.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 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)?))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn asof_join(
&self,
request: Request<AsofJoinRequest>,
) -> Result<Response<ResultTable>, Status> {
let tenant = session_tenant_traced(&request);
let args = asof_join_from_proto(request.into_inner())?;
let record = scoped(&self.session, tenant, || async {
self.session
.asof_join(&args.spine, &args.facts, &args.spec)
.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 recompute(
&self,
request: Request<RecomputeRequest>,
) -> Result<Response<ProtoRecomputeReport>, Status> {
let tenant = session_tenant_traced(&request);
let args = recompute_from_proto(request.into_inner())?;
let report = scoped(&self.session, tenant, || async {
let record = self
.session
.catalog()
.get_result_table(&args.table)
.await?
.ok_or_else(|| {
JammiError::Catalog(format!("Result table '{}' not found", args.table))
})?;
self.session.recompute(&record, args.cascade).await
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(recompute_report_to_proto(report)))
}
}