use std::sync::Arc;
use jammi_ai::session::InferenceSession;
use jammi_ai::wire::{infer_result_to_proto, model_task_from_proto};
use jammi_ai::{LocalSession, Session};
use tonic::{Request, Response, Status};
use crate::grpc::proto::inference::inference_service_server::InferenceService;
use crate::grpc::proto::inference::{InferRequest, InferResponse};
use crate::grpc::wire::{map_engine_error, require_nonempty, scoped, session_tenant};
pub struct InferenceServer {
session: Arc<InferenceSession>,
}
impl InferenceServer {
pub fn new(session: Arc<InferenceSession>) -> Self {
Self { session }
}
fn local(&self) -> Session {
Session::Local(LocalSession::new(Arc::clone(&self.session)))
}
}
#[tonic::async_trait]
impl InferenceService for InferenceServer {
async fn infer(
&self,
request: Request<InferRequest>,
) -> Result<Response<InferResponse>, 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 task = model_task_from_proto(req.task)?;
let session = self.local();
let batches = scoped(&self.session, tenant, || {
session.infer(
&req.source_id,
&req.model_id,
task,
&req.columns,
&req.key_column,
)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(InferResponse {
result: Some(infer_result_to_proto(batches)?),
}))
}
}