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::eval::Evaluator;
use crate::model::WeirwoodTree;
use super::client::{EncryptedScore, SCALE};
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 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() {
let scaled: i32 = (node.leaf_value * SCALE)
.round()
.clamp(i32::MIN as f32, i32::MAX as f32) as i32;
FheInt32::encrypt_trivial(scaled)
} else {
let threshold: i32 = (node.split_threshold * SCALE).round() as i32;
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 base_scaled = (weirwood_tree.base_score * SCALE)
.round()
.clamp(i32::MIN as f32, i32::MAX as f32) as i32;
let mut total: FheInt32 = FheInt32::encrypt_trivial(base_scaled);
for score in tree_scores {
total += score;
}
total
}
}