use tonic::transport::Channel;
use tonic::{Request, Status};
use crate::Error;
use crate::eval::PlaintextEvaluator;
use crate::eval::fhe::ClientContext;
use crate::model::{Objective, WeirwoodTree};
use super::rpc::inference_service_client::InferenceServiceClient;
use super::rpc::{InitSessionRequest, PredictRequest};
use super::{deserialize_score, serialize_feature, serialize_server_context};
const MAX_GRPC_MESSAGE_BYTES: usize = 512 * 1024 * 1024;
pub struct WeirwoodClient {
grpc: InferenceServiceClient<Channel>,
fhe: ClientContext,
session_id: String,
}
impl WeirwoodClient {
pub async fn connect(dst: impl Into<String>) -> Result<Self, Error> {
let endpoint = tonic::transport::Endpoint::from_shared(dst.into())
.map_err(|e| Error::Other(format!("invalid server endpoint: {e}")))?;
let channel = endpoint
.connect()
.await
.map_err(|e| Error::Other(format!("failed to connect to inference server: {e}")))?;
let mut grpc = InferenceServiceClient::new(channel)
.max_decoding_message_size(MAX_GRPC_MESSAGE_BYTES)
.max_encoding_message_size(MAX_GRPC_MESSAGE_BYTES);
let fhe = ClientContext::generate()?;
let server_key_bytes = serialize_server_context(&fhe.server_context())?;
let req = Request::new(InitSessionRequest {
server_key: server_key_bytes,
});
let resp = grpc
.init_session(req)
.await
.map_err(status_to_error)?
.into_inner();
Ok(Self {
grpc,
fhe,
session_id: resp.session_id,
})
}
pub async fn predict_proba(
&mut self,
model: &WeirwoodTree,
features: &[f32],
) -> Result<f32, Error> {
let raw = self.predict_raw(features).await?;
Ok(match &model.objective {
Objective::BinaryLogistic => sigmoid(raw),
Objective::RegSquaredError => raw,
Objective::MultiSoftmax { num_class } => panic!(
"predict_proba returns a single f32; multi:softmax with num_class={} \
produces a vector — use predict_multiclass_proba instead",
num_class
),
Objective::Other(_) => raw,
})
}
pub async fn predict_multiclass_proba(
&mut self,
model: &WeirwoodTree,
features: &[f32],
) -> Result<Vec<f32>, Error> {
let _ = self.predict_raw(features).await?;
Ok(PlaintextEvaluator.predict_multiclass_proba(model, features))
}
pub async fn predict_raw(&mut self, features: &[f32]) -> Result<f32, Error> {
let encrypted = self.fhe.encrypt(features);
let mut feature_bytes = Vec::with_capacity(encrypted.len());
for feat in &encrypted {
feature_bytes.push(serialize_feature(feat)?);
}
let req = Request::new(PredictRequest {
session_id: self.session_id.clone(),
features: feature_bytes,
});
let resp = self
.grpc
.predict(req)
.await
.map_err(status_to_error)?
.into_inner();
let encrypted_score = deserialize_score(&resp.encrypted_score)?;
Ok(self.fhe.decrypt_score(&encrypted_score))
}
pub fn session_id(&self) -> &str {
&self.session_id
}
}
fn status_to_error(status: Status) -> Error {
Error::Other(format!("gRPC error: {status}"))
}
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}