use std::cell::Cell;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use rayon::prelude::*;
use tfhe::FheInt32;
use tfhe::ServerKey;
use tfhe::prelude::*;
use crate::Error;
use crate::eval::Evaluator;
use crate::model::WeirwoodTree;
use super::client::{EncryptedScore, encode_fixed_point};
use super::server::ServerContext;
static EVALUATOR_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
thread_local! {
static INSTALLED_EVALUATOR_ID: Cell<Option<u64>> = const { Cell::new(None) };
}
pub struct FheEvaluator {
id: u64,
server_key: Arc<ServerKey>,
thread_pool: rayon::ThreadPool,
}
impl FheEvaluator {
pub fn try_new(model: &WeirwoodTree, ctx: ServerContext) -> Result<Self, Error> {
let warnings = model.validate_for_fhe();
if !warnings.is_empty() {
let summary = warnings
.iter()
.map(|w| w.to_string())
.collect::<Vec<_>>()
.join("; ");
return Err(Error::Format(format!(
"model is unsafe for FHE evaluation ({} issue(s)): {summary}",
warnings.len()
)));
}
Ok(Self::new(ctx))
}
pub fn new(ctx: ServerContext) -> Self {
let id = EVALUATOR_ID_COUNTER.fetch_add(1, Ordering::Relaxed);
let server_key = Arc::new(ctx.server_key);
let server_key_for_workers = Arc::clone(&server_key);
let thread_pool = rayon::ThreadPoolBuilder::new()
.start_handler(move |_| {
tfhe::set_server_key((*server_key_for_workers).clone());
})
.build()
.expect("failed to build Rayon thread pool for FHE evaluation");
FheEvaluator {
id,
server_key,
thread_pool,
}
}
fn ensure_key_installed(&self) {
INSTALLED_EVALUATOR_ID.with(|cell| {
if cell.get() != Some(self.id) {
tfhe::set_server_key((*self.server_key).clone());
cell.set(Some(self.id));
}
});
}
}
fn eval_node(tree: &crate::model::Tree, node_idx: usize, features: &[FheInt32]) -> FheInt32 {
let node: &crate::model::Node = &tree.nodes[node_idx];
if node.is_leaf() {
return FheInt32::encrypt_trivial(encode_fixed_point(node.leaf_value));
}
let threshold: i32 = encode_fixed_point(node.split_threshold);
let go_left: tfhe::FheBool = features[node.split_feature as usize].le(threshold);
let left_score: FheInt32 = eval_node(tree, node.left_child as usize, features);
let right_score: FheInt32 = eval_node(tree, node.right_child as usize, features);
go_left.if_then_else(&left_score, &right_score)
}
impl Evaluator for FheEvaluator {
type Input = [FheInt32];
type Output = EncryptedScore;
fn predict(
&self,
weirwood_tree: &WeirwoodTree,
encrypted_features: &[FheInt32],
) -> EncryptedScore {
self.ensure_key_installed();
let tree_scores: Vec<FheInt32> = self.thread_pool.install(|| {
weirwood_tree
.trees
.par_iter()
.map(|tree| eval_node(tree, 0, encrypted_features))
.collect()
});
let mut total: FheInt32 =
FheInt32::encrypt_trivial(encode_fixed_point(weirwood_tree.base_score));
for score in tree_scores {
total += score;
}
total
}
}