Skip to main content

harn_vm/llm/
trigger_predicate.rs

1use std::cell::RefCell;
2use std::collections::{BTreeMap, HashMap};
3use std::sync::{Mutex, OnceLock};
4use std::time::Duration;
5
6use serde::{Deserialize, Serialize};
7
8use super::api::{LlmRequestPayload, LlmResult};
9
10#[derive(Clone, Debug, Default, Serialize, Deserialize, PartialEq)]
11pub struct TriggerPredicateBudget {
12    pub max_cost_usd: Option<f64>,
13    pub tokens_max: Option<u64>,
14    pub timeout_ms: Option<u64>,
15}
16
17impl TriggerPredicateBudget {
18    pub fn timeout(&self) -> Option<Duration> {
19        self.timeout_ms.map(Duration::from_millis)
20    }
21}
22
23#[derive(Clone, Debug, Serialize, Deserialize)]
24pub(crate) struct PredicateCacheEntry {
25    pub request_hash: String,
26    pub(crate) result: LlmResult,
27}
28
29#[derive(Clone, Debug, Default)]
30pub struct PredicateEvaluationCapture {
31    pub entries: Vec<PredicateCacheEntry>,
32    pub total_tokens: u64,
33    pub total_cost_usd: f64,
34    pub cached: bool,
35    pub budget_exceeded: bool,
36}
37
38#[derive(Clone, Debug, Default)]
39struct PredicateEvaluationState {
40    budget: TriggerPredicateBudget,
41    replay_cache: HashMap<String, LlmResult>,
42    entries: BTreeMap<String, LlmResult>,
43    total_tokens: u64,
44    total_cost_usd: f64,
45    cached: bool,
46    budget_exceeded: bool,
47}
48
49thread_local! {
50    static ACTIVE_PREDICATE_EVALUATION: RefCell<Option<PredicateEvaluationState>> = const { RefCell::new(None) };
51}
52
53fn request_cache() -> &'static Mutex<HashMap<String, LlmResult>> {
54    static CACHE: OnceLock<Mutex<HashMap<String, LlmResult>>> = OnceLock::new();
55    CACHE.get_or_init(|| Mutex::new(HashMap::new()))
56}
57
58pub(crate) fn reset_trigger_predicate_state() {
59    ACTIVE_PREDICATE_EVALUATION.with(|slot| {
60        *slot.borrow_mut() = None;
61    });
62    if let Ok(mut cache) = request_cache().lock() {
63        cache.clear();
64    }
65}
66
67pub(crate) fn request_hash(request: &LlmRequestPayload) -> String {
68    use std::hash::{Hash, Hasher};
69
70    let canonical = serde_json::json!({
71        "provider": request.provider,
72        "model": request.model,
73        "messages": request.messages,
74        "system": request.system,
75        "max_tokens": request.max_tokens,
76        "temperature": request.temperature,
77        "top_p": request.top_p,
78        "top_k": request.top_k,
79        "logprobs": request.logprobs,
80        "logit_bias": request.logit_bias,
81        "min_p": request.min_p,
82        "repetition_penalty": request.repetition_penalty,
83        "prediction": request.prediction,
84        "verbosity": request.verbosity,
85        "mirostat": request.mirostat,
86        "stop": request.stop,
87        "seed": request.seed,
88        "frequency_penalty": request.frequency_penalty,
89        "presence_penalty": request.presence_penalty,
90        "parallel_tool_calls": request.parallel_tool_calls,
91        "output_format": request.output_format,
92        "thinking": request.thinking,
93        "anthropic_beta_features": request.anthropic_beta_features,
94        "native_tools": request.native_tools,
95        "tool_choice": request.tool_choice,
96        "cache": request.cache,
97        "timeout": request.timeout,
98        "stream": request.stream,
99        "provider_overrides": request.provider_overrides,
100        "prefill": request.prefill,
101        "mock_scope": request.mock_scope,
102    });
103    let mut hasher = std::collections::hash_map::DefaultHasher::new();
104    serde_json::to_string(&canonical)
105        .unwrap_or_default()
106        .hash(&mut hasher);
107    format!("{:016x}", hasher.finish())
108}
109
110pub(crate) struct PredicateEvaluationGuard;
111
112impl PredicateEvaluationGuard {
113    pub fn finish(self) -> PredicateEvaluationCapture {
114        finish_predicate_evaluation()
115    }
116}
117
118impl Drop for PredicateEvaluationGuard {
119    fn drop(&mut self) {
120        ACTIVE_PREDICATE_EVALUATION.with(|slot| {
121            *slot.borrow_mut() = None;
122        });
123    }
124}
125
126pub(crate) fn start_predicate_evaluation(
127    budget: TriggerPredicateBudget,
128    replay_entries: Vec<PredicateCacheEntry>,
129) -> PredicateEvaluationGuard {
130    ACTIVE_PREDICATE_EVALUATION.with(|slot| {
131        *slot.borrow_mut() = Some(PredicateEvaluationState {
132            budget,
133            replay_cache: replay_entries
134                .into_iter()
135                .map(|entry| (entry.request_hash, entry.result))
136                .collect(),
137            ..Default::default()
138        });
139    });
140    PredicateEvaluationGuard
141}
142
143fn finish_predicate_evaluation() -> PredicateEvaluationCapture {
144    ACTIVE_PREDICATE_EVALUATION.with(|slot| {
145        let Some(state) = slot.borrow_mut().take() else {
146            return PredicateEvaluationCapture::default();
147        };
148        PredicateEvaluationCapture {
149            entries: state
150                .entries
151                .into_iter()
152                .map(|(request_hash, result)| PredicateCacheEntry {
153                    request_hash,
154                    result,
155                })
156                .collect(),
157            total_tokens: state.total_tokens,
158            total_cost_usd: state.total_cost_usd,
159            cached: state.cached,
160            budget_exceeded: state.budget_exceeded,
161        }
162    })
163}
164
165pub(crate) fn lookup_cached_result(request: &LlmRequestPayload) -> Option<LlmResult> {
166    ACTIVE_PREDICATE_EVALUATION.with(|slot| {
167        let mut borrowed = slot.borrow_mut();
168        let state = borrowed.as_mut()?;
169        if state.budget_exceeded {
170            return None;
171        }
172        let hash = request_hash(request);
173        let cached = state.replay_cache.get(&hash).cloned().or_else(|| {
174            request_cache()
175                .lock()
176                .ok()
177                .and_then(|cache| cache.get(&hash).cloned())
178        });
179        if let Some(result) = cached {
180            state.cached = true;
181            state.entries.insert(hash, result.clone());
182            return Some(result);
183        }
184        None
185    })
186}
187
188pub(crate) fn note_result(request: &LlmRequestPayload, result: &LlmResult) {
189    ACTIVE_PREDICATE_EVALUATION.with(|slot| {
190        let mut borrowed = slot.borrow_mut();
191        let Some(state) = borrowed.as_mut() else {
192            return;
193        };
194        let hash = request_hash(request);
195        state.entries.insert(hash.clone(), result.clone());
196        if let Ok(mut cache) = request_cache().lock() {
197            cache.insert(hash, result.clone());
198        }
199        let usage = result.usage();
200        let call_tokens = usage
201            .input_tokens
202            .saturating_add(usage.output_tokens)
203            .max(0) as u64;
204        state.total_tokens = state.total_tokens.saturating_add(call_tokens);
205        state.total_cost_usd += usage.cost_usd.unwrap_or(0.0);
206        if state
207            .budget
208            .tokens_max
209            .is_some_and(|limit| state.total_tokens > limit)
210        {
211            state.budget_exceeded = true;
212        }
213        if state
214            .budget
215            .max_cost_usd
216            .is_some_and(|limit| state.total_cost_usd > limit)
217        {
218            state.budget_exceeded = true;
219        }
220    });
221}