use std::sync::Arc;
use jammi_ai::fine_tune::spec::TrainingSpec;
use jammi_ai::fine_tune::training_job::TrainingJob;
use jammi_ai::session::InferenceSession;
use jammi_ai::wire::training_spec_from_proto;
use tonic::{Request, Response, Status};
use crate::grpc::proto::training as pb;
use crate::grpc::proto::training::training_service_server::TrainingService;
use crate::grpc::wire::{map_engine_error, require_nonempty, scoped, session_tenant_traced};
pub struct TrainingServer {
session: Arc<InferenceSession>,
}
impl TrainingServer {
pub fn new(session: Arc<InferenceSession>) -> Self {
Self { session }
}
async fn submit(&self, spec: TrainingSpec) -> Result<TrainingJob, jammi_db::error::JammiError> {
self.session.run_training_spec(spec).await
}
}
#[tonic::async_trait]
impl TrainingService for TrainingServer {
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn start_training(
&self,
request: Request<pb::StartTrainingRequest>,
) -> Result<Response<pb::StartTrainingResponse>, Status> {
let tenant = session_tenant_traced(&request);
let spec = training_spec_from_proto(request.into_inner())?;
let job = scoped(&self.session, tenant, || self.submit(spec))
.await
.map_err(map_engine_error)?;
Ok(Response::new(pb::StartTrainingResponse {
job_id: job.job_id,
model_id: job.model_id,
}))
}
#[tracing::instrument(skip(self, request), fields(tenant_id = tracing::field::Empty))]
async fn training_status(
&self,
request: Request<pb::TrainingStatusRequest>,
) -> Result<Response<pb::TrainingStatusResponse>, Status> {
let tenant = session_tenant_traced(&request);
let req = request.into_inner();
require_nonempty(&req.job_id, "job_id")?;
let record = scoped(&self.session, tenant, || {
self.session.catalog().get_training_job(&req.job_id)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(pb::TrainingStatusResponse {
status: record.status,
model_id: record.output_model_id.unwrap_or_default(),
error: record.error_message.unwrap_or_default(),
}))
}
}