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};
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 "stop": request.stop,
80 "seed": request.seed,
81 "frequency_penalty": request.frequency_penalty,
82 "presence_penalty": request.presence_penalty,
83 "output_format": request.output_format,
84 "thinking": request.thinking,
85 "anthropic_beta_features": request.anthropic_beta_features,
86 "native_tools": request.native_tools,
87 "tool_choice": request.tool_choice,
88 "cache": request.cache,
89 "timeout": request.timeout,
90 "stream": request.stream,
91 "provider_overrides": request.provider_overrides,
92 "prefill": request.prefill,
93 "mock_scope": request.mock_scope,
94 });
95 let mut hasher = std::collections::hash_map::DefaultHasher::new();
96 serde_json::to_string(&canonical)
97 .unwrap_or_default()
98 .hash(&mut hasher);
99 format!("{:016x}", hasher.finish())
100}
101
102pub(crate) struct PredicateEvaluationGuard;
103
104impl PredicateEvaluationGuard {
105 pub fn finish(self) -> PredicateEvaluationCapture {
106 finish_predicate_evaluation()
107 }
108}
109
110impl Drop for PredicateEvaluationGuard {
111 fn drop(&mut self) {
112 ACTIVE_PREDICATE_EVALUATION.with(|slot| {
113 *slot.borrow_mut() = None;
114 });
115 }
116}
117
118pub(crate) fn start_predicate_evaluation(
119 budget: TriggerPredicateBudget,
120 replay_entries: Vec<PredicateCacheEntry>,
121) -> PredicateEvaluationGuard {
122 ACTIVE_PREDICATE_EVALUATION.with(|slot| {
123 *slot.borrow_mut() = Some(PredicateEvaluationState {
124 budget,
125 replay_cache: replay_entries
126 .into_iter()
127 .map(|entry| (entry.request_hash, entry.result))
128 .collect(),
129 ..Default::default()
130 });
131 });
132 PredicateEvaluationGuard
133}
134
135fn finish_predicate_evaluation() -> PredicateEvaluationCapture {
136 ACTIVE_PREDICATE_EVALUATION.with(|slot| {
137 let Some(state) = slot.borrow_mut().take() else {
138 return PredicateEvaluationCapture::default();
139 };
140 PredicateEvaluationCapture {
141 entries: state
142 .entries
143 .into_iter()
144 .map(|(request_hash, result)| PredicateCacheEntry {
145 request_hash,
146 result,
147 })
148 .collect(),
149 total_tokens: state.total_tokens,
150 total_cost_usd: state.total_cost_usd,
151 cached: state.cached,
152 budget_exceeded: state.budget_exceeded,
153 }
154 })
155}
156
157pub(crate) fn lookup_cached_result(request: &LlmRequestPayload) -> Option<LlmResult> {
158 ACTIVE_PREDICATE_EVALUATION.with(|slot| {
159 let mut borrowed = slot.borrow_mut();
160 let state = borrowed.as_mut()?;
161 if state.budget_exceeded {
162 return None;
163 }
164 let hash = request_hash(request);
165 let cached = state.replay_cache.get(&hash).cloned().or_else(|| {
166 request_cache()
167 .lock()
168 .ok()
169 .and_then(|cache| cache.get(&hash).cloned())
170 });
171 if let Some(result) = cached {
172 state.cached = true;
173 state.entries.insert(hash, result.clone());
174 return Some(result);
175 }
176 None
177 })
178}
179
180pub(crate) fn note_result(request: &LlmRequestPayload, result: &LlmResult) {
181 ACTIVE_PREDICATE_EVALUATION.with(|slot| {
182 let mut borrowed = slot.borrow_mut();
183 let Some(state) = borrowed.as_mut() else {
184 return;
185 };
186 let hash = request_hash(request);
187 state.entries.insert(hash.clone(), result.clone());
188 if let Ok(mut cache) = request_cache().lock() {
189 cache.insert(hash, result.clone());
190 }
191 let usage = result.usage();
192 let call_tokens = usage
193 .input_tokens
194 .saturating_add(usage.output_tokens)
195 .max(0) as u64;
196 state.total_tokens = state.total_tokens.saturating_add(call_tokens);
197 state.total_cost_usd += usage.cost_usd.unwrap_or(0.0);
198 if state
199 .budget
200 .tokens_max
201 .is_some_and(|limit| state.total_tokens > limit)
202 {
203 state.budget_exceeded = true;
204 }
205 if state
206 .budget
207 .max_cost_usd
208 .is_some_and(|limit| state.total_cost_usd > limit)
209 {
210 state.budget_exceeded = true;
211 }
212 });
213}