use std::sync::Arc;
use jammi_ai::session::InferenceSession;
use jammi_ai::Session;
use jammi_wire::{calibration_shape_from_proto, cohorts_from_proto, EvalTaskFromWire};
use tonic::{Request, Response, Status};
use crate::grpc::proto::eval as pb;
use crate::grpc::proto::eval::eval_service_server::EvalService;
use crate::grpc::wire::{map_engine_error, require_nonempty, scoped, session_tenant};
pub struct EvalServer {
session: Arc<InferenceSession>,
}
impl EvalServer {
pub fn new(session: Arc<InferenceSession>) -> Self {
Self { session }
}
fn local(&self) -> Session {
Session::new(Arc::clone(&self.session))
}
}
#[tonic::async_trait]
impl EvalService for EvalServer {
async fn eval_embeddings(
&self,
request: Request<pb::EvalEmbeddingsRequest>,
) -> Result<Response<pb::EmbeddingEvalReport>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.source_id, "source_id")?;
require_nonempty(&req.golden_source, "golden_source")?;
let embedding_table = optional_str(&req.embedding_table);
let cohorts = cohorts_from_proto(req.cohorts);
let session = self.local();
let report = scoped(&self.session, tenant, || {
session.eval_embeddings(
&req.source_id,
embedding_table,
&req.golden_source,
req.k as usize,
&cohorts,
)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(report.into()))
}
async fn eval_per_query(
&self,
request: Request<pb::EvalPerQueryRequest>,
) -> Result<Response<pb::EvalPerQueryResponse>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.eval_run_id, "eval_run_id")?;
let session = self.local();
let records = scoped(&self.session, tenant, || {
session.eval_per_query(&req.eval_run_id)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(pb::EvalPerQueryResponse {
records: records.into_iter().map(Into::into).collect(),
}))
}
async fn eval_inference(
&self,
request: Request<pb::EvalInferenceRequest>,
) -> Result<Response<pb::InferenceEvalReport>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.model_id, "model_id")?;
require_nonempty(&req.source_id, "source_id")?;
require_nonempty(&req.golden_source, "golden_source")?;
require_nonempty(&req.label_column, "label_column")?;
if req.columns.is_empty() {
return Err(Status::invalid_argument("columns is required"));
}
let task = EvalTaskFromWire::try_from(req.task)?.0;
let session = self.local();
let report = scoped(&self.session, tenant, || {
session.eval_inference(
&req.model_id,
&req.source_id,
&req.columns,
task,
&req.golden_source,
&req.label_column,
)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(report.into()))
}
async fn eval_compare(
&self,
request: Request<pb::EvalCompareRequest>,
) -> Result<Response<pb::CompareEvalReport>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.source_id, "source_id")?;
require_nonempty(&req.golden_source, "golden_source")?;
if req.embedding_tables.len() < 2 {
return Err(Status::invalid_argument(
"embedding_tables requires at least two tables",
));
}
let session = self.local();
let report = scoped(&self.session, tenant, || {
session.eval_compare(
&req.embedding_tables,
&req.source_id,
&req.golden_source,
req.k as usize,
)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(report.into()))
}
async fn eval_calibration(
&self,
request: Request<pb::EvalCalibrationRequest>,
) -> Result<Response<pb::CalibrationEvalReport>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.source_id, "source_id")?;
require_nonempty(&req.golden_source, "golden_source")?;
let shape = calibration_shape_from_proto(req.shape)?;
let cohorts = cohorts_from_proto(req.cohorts);
let report = scoped(&self.session, tenant, || async {
self.session
.eval_calibration(&req.source_id, &req.golden_source, shape, &cohorts)
.await
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(report.into()))
}
}
fn optional_str(s: &str) -> Option<&str> {
if s.is_empty() {
None
} else {
Some(s)
}
}