use std::sync::Arc;
use jammi_ai::fine_tune::FineTuneConfig;
use jammi_ai::local_session::FineTuneJobId;
use jammi_ai::session::InferenceSession;
use jammi_ai::wire::{method_from_proto, model_task_from_proto};
use jammi_ai::{LocalSession, Session};
use tonic::{Request, Response, Status};
use crate::grpc::proto::fine_tune as pb;
use crate::grpc::proto::fine_tune::fine_tune_service_server::FineTuneService;
use crate::grpc::wire::{map_engine_error, require_nonempty, scoped, session_tenant};
pub struct FineTuneServer {
session: Arc<InferenceSession>,
}
impl FineTuneServer {
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 FineTuneService for FineTuneServer {
async fn start_fine_tune(
&self,
request: Request<pb::StartFineTuneRequest>,
) -> Result<Response<pb::StartFineTuneResponse>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.source_id, "source_id")?;
require_nonempty(&req.base_model, "base_model")?;
if req.columns.is_empty() {
return Err(Status::invalid_argument("columns is required"));
}
let method = method_from_proto(req.method)?;
let task = model_task_from_proto(req.task)?;
let config = req.config.map(FineTuneConfig::try_from).transpose()?;
let session = self.local();
let job_id = scoped(&self.session, tenant, || {
session.fine_tune(
&req.source_id,
&req.base_model,
&req.columns,
method,
task,
config,
)
})
.await
.map_err(map_engine_error)?;
Ok(Response::new(pb::StartFineTuneResponse {
job_id: job_id.0,
}))
}
async fn fine_tune_status(
&self,
request: Request<pb::FineTuneStatusRequest>,
) -> Result<Response<pb::FineTuneStatusResponse>, Status> {
let tenant = session_tenant(&request);
let req = request.into_inner();
require_nonempty(&req.job_id, "job_id")?;
let id = FineTuneJobId(req.job_id);
let session = self.local();
let status = scoped(&self.session, tenant, || session.fine_tune_status(&id))
.await
.map_err(map_engine_error)?;
Ok(Response::new(pb::FineTuneStatusResponse { status }))
}
}