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