1use std::cell::RefCell;
18use std::collections::BTreeMap;
19use std::path::Path;
20
21use anyhow::{anyhow, bail, Context, Result};
22use serde::{Deserialize, Serialize};
23
24use super::sha256_hex;
25
26pub const PROVIDER_OPENAI: &str = "openai-compatible";
28pub const PROVIDER_RECORDED: &str = "recorded";
30
31#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
33pub struct ModelParams {
34 pub model: String,
36 pub temperature: f64,
37 pub top_p: f64,
38 pub max_tokens: u32,
39 pub seed: u64,
41}
42
43impl Default for ModelParams {
44 fn default() -> Self {
45 Self {
46 model: String::new(),
47 temperature: 0.0,
48 top_p: 1.0,
49 max_tokens: 1024,
50 seed: 7,
51 }
52 }
53}
54
55#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
58pub struct ModelFingerprint {
59 pub provider: String,
61 pub endpoint: String,
63 pub params: ModelParams,
64}
65
66impl ModelFingerprint {
67 pub fn digest(&self) -> String {
69 let canonical = serde_json::to_vec(self).unwrap_or_default();
70 sha256_hex(&canonical)
71 }
72}
73
74#[derive(Debug, Clone, PartialEq)]
77pub struct ModelRequest {
78 pub system: String,
79 pub user: String,
80}
81
82impl ModelRequest {
83 pub fn key(&self) -> String {
85 let mut joined = Vec::with_capacity(self.system.len() + self.user.len() + 1);
86 joined.extend_from_slice(self.system.as_bytes());
87 joined.push(0);
88 joined.extend_from_slice(self.user.as_bytes());
89 sha256_hex(&joined)
90 }
91}
92
93#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
95pub struct ModelResponse {
96 pub text: String,
97}
98
99impl ModelResponse {
100 pub fn new(text: impl Into<String>) -> Self {
101 Self { text: text.into() }
102 }
103
104 pub fn digest(&self) -> String {
106 sha256_hex(self.text.as_bytes())
107 }
108}
109
110pub trait ModelRunner {
112 fn fingerprint(&self) -> &ModelFingerprint;
114 fn run(&self, req: &ModelRequest) -> Result<ModelResponse>;
116}
117
118pub struct OpenAiRunner {
124 fingerprint: ModelFingerprint,
125 api_key: String,
126}
127
128impl OpenAiRunner {
129 pub fn new(
131 base_url: impl Into<String>,
132 api_key: impl Into<String>,
133 params: ModelParams,
134 ) -> Self {
135 let endpoint = base_url.into().trim_end_matches('/').to_string();
136 Self {
137 fingerprint: ModelFingerprint {
138 provider: PROVIDER_OPENAI.to_string(),
139 endpoint,
140 params,
141 },
142 api_key: api_key.into(),
143 }
144 }
145
146 pub fn from_env() -> Result<Self> {
149 let base = std::env::var("LEAN_CTX_EVAL_MODEL_URL")
150 .context("LEAN_CTX_EVAL_MODEL_URL not set (OpenAI-compatible base URL)")?;
151 let key = std::env::var("LEAN_CTX_EVAL_MODEL_KEY").unwrap_or_default();
152 let model = std::env::var("LEAN_CTX_EVAL_MODEL")
153 .context("LEAN_CTX_EVAL_MODEL not set (provider model id)")?;
154 let seed = std::env::var("LEAN_CTX_EVAL_SEED")
155 .ok()
156 .and_then(|s| s.parse().ok())
157 .unwrap_or(ModelParams::default().seed);
158 let params = ModelParams {
159 model,
160 seed,
161 ..ModelParams::default()
162 };
163 Ok(Self::new(base, key, params))
164 }
165}
166
167#[derive(Serialize)]
168struct ChatPayload<'a> {
169 model: &'a str,
170 temperature: f64,
171 top_p: f64,
172 max_tokens: u32,
173 seed: u64,
174 messages: Vec<ChatMessage<'a>>,
175}
176
177#[derive(Serialize)]
178struct ChatMessage<'a> {
179 role: &'a str,
180 content: &'a str,
181}
182
183#[derive(Deserialize)]
184struct ChatResponse {
185 choices: Vec<ChatChoice>,
186}
187
188#[derive(Deserialize)]
189struct ChatChoice {
190 message: ChatChoiceMessage,
191}
192
193#[derive(Deserialize)]
194struct ChatChoiceMessage {
195 content: String,
196}
197
198impl ModelRunner for OpenAiRunner {
199 fn fingerprint(&self) -> &ModelFingerprint {
200 &self.fingerprint
201 }
202
203 fn run(&self, req: &ModelRequest) -> Result<ModelResponse> {
204 let p = &self.fingerprint.params;
205 let payload = ChatPayload {
206 model: &p.model,
207 temperature: p.temperature,
208 top_p: p.top_p,
209 max_tokens: p.max_tokens,
210 seed: p.seed,
211 messages: vec![
212 ChatMessage {
213 role: "system",
214 content: &req.system,
215 },
216 ChatMessage {
217 role: "user",
218 content: &req.user,
219 },
220 ],
221 };
222 let url = format!("{}/chat/completions", self.fingerprint.endpoint);
223 let body = serde_json::to_vec(&payload).context("serialize chat payload")?;
224 let mut request = ureq::post(&url).header("Content-Type", "application/json");
225 if !self.api_key.is_empty() {
226 request = request.header("Authorization", &format!("Bearer {}", self.api_key));
227 }
228 let resp = request
229 .send(&body[..])
230 .map_err(|e| anyhow!("model request to {url} failed: {e}"))?;
231 let status = resp.status().as_u16();
232 let text = resp
233 .into_body()
234 .read_to_string()
235 .map_err(|e| anyhow!("reading model response failed: {e}"))?;
236 if status != 200 {
237 bail!("model endpoint returned HTTP {status}: {text}");
238 }
239 let parsed: ChatResponse = serde_json::from_str(&text)
240 .with_context(|| format!("parsing model response: {text}"))?;
241 let answer = parsed
242 .choices
243 .into_iter()
244 .next()
245 .map(|c| c.message.content)
246 .ok_or_else(|| anyhow!("model response contained no choices"))?;
247 Ok(ModelResponse::new(answer))
248 }
249}
250
251#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
258pub struct Recording {
259 pub kind: String,
261 pub fingerprint: ModelFingerprint,
263 pub entries: BTreeMap<String, ModelResponse>,
265}
266
267const RECORDING_KIND: &str = "lean-ctx.eval-recording";
268
269impl Recording {
270 pub fn new(fingerprint: ModelFingerprint) -> Self {
271 Self {
272 kind: RECORDING_KIND.to_string(),
273 fingerprint,
274 entries: BTreeMap::new(),
275 }
276 }
277
278 pub fn load(path: &Path) -> Result<Self> {
280 let raw = std::fs::read_to_string(path)
281 .with_context(|| format!("reading recording {}", path.display()))?;
282 let rec: Recording = serde_json::from_str(&raw)
283 .with_context(|| format!("parsing recording {}", path.display()))?;
284 if rec.kind != RECORDING_KIND {
285 bail!("not a {RECORDING_KIND} file (kind = {:?})", rec.kind);
286 }
287 Ok(rec)
288 }
289
290 pub fn save(&self, path: &Path) -> Result<()> {
292 if let Some(parent) = path.parent() {
293 std::fs::create_dir_all(parent).ok();
294 }
295 let json = serde_json::to_string_pretty(self).context("serialize recording")?;
296 std::fs::write(path, json)
297 .with_context(|| format!("writing recording {}", path.display()))?;
298 Ok(())
299 }
300}
301
302pub struct RecordedRunner {
305 recording: Recording,
306}
307
308impl RecordedRunner {
309 pub fn new(recording: Recording) -> Self {
310 Self { recording }
311 }
312
313 pub fn from_file(path: &Path) -> Result<Self> {
315 Ok(Self::new(Recording::load(path)?))
316 }
317}
318
319impl ModelRunner for RecordedRunner {
320 fn fingerprint(&self) -> &ModelFingerprint {
321 &self.recording.fingerprint
322 }
323
324 fn run(&self, req: &ModelRequest) -> Result<ModelResponse> {
325 self.recording
326 .entries
327 .get(&req.key())
328 .cloned()
329 .ok_or_else(|| anyhow!("no recorded response for request key {}", req.key()))
330 }
331}
332
333pub struct RecordingRunner<R: ModelRunner> {
336 inner: R,
337 recording: RefCell<Recording>,
338}
339
340impl<R: ModelRunner> RecordingRunner<R> {
341 pub fn new(inner: R) -> Self {
342 let recording = Recording::new(inner.fingerprint().clone());
343 Self {
344 inner,
345 recording: RefCell::new(recording),
346 }
347 }
348
349 pub fn into_recording(self) -> Recording {
351 self.recording.into_inner()
352 }
353}
354
355impl<R: ModelRunner> ModelRunner for RecordingRunner<R> {
356 fn fingerprint(&self) -> &ModelFingerprint {
357 self.inner.fingerprint()
358 }
359
360 fn run(&self, req: &ModelRequest) -> Result<ModelResponse> {
361 let resp = self.inner.run(req)?;
362 self.recording
363 .borrow_mut()
364 .entries
365 .insert(req.key(), resp.clone());
366 Ok(resp)
367 }
368}
369
370#[cfg(test)]
371mod tests {
372 use super::*;
373
374 fn fp() -> ModelFingerprint {
375 ModelFingerprint {
376 provider: PROVIDER_RECORDED.to_string(),
377 endpoint: "rec".into(),
378 params: ModelParams {
379 model: "test-model".into(),
380 ..ModelParams::default()
381 },
382 }
383 }
384
385 #[test]
386 fn request_key_is_stable_and_order_sensitive() {
387 let a = ModelRequest {
388 system: "s".into(),
389 user: "u".into(),
390 };
391 let b = ModelRequest {
392 system: "s".into(),
393 user: "u".into(),
394 };
395 assert_eq!(a.key(), b.key());
396 let swapped = ModelRequest {
397 system: "u".into(),
398 user: "s".into(),
399 };
400 assert_ne!(a.key(), swapped.key());
401 }
402
403 #[test]
404 fn fingerprint_digest_changes_with_params() {
405 let mut f1 = fp();
406 let d1 = f1.digest();
407 f1.params.seed = 99;
408 assert_ne!(d1, f1.digest());
409 }
410
411 #[test]
412 fn recorded_runner_replays_and_errors_on_miss() {
413 let mut rec = Recording::new(fp());
414 let req = ModelRequest {
415 system: "sys".into(),
416 user: "hi".into(),
417 };
418 rec.entries.insert(req.key(), ModelResponse::new("hello"));
419 let runner = RecordedRunner::new(rec);
420 assert_eq!(runner.run(&req).unwrap().text, "hello");
421
422 let miss = ModelRequest {
423 system: "sys".into(),
424 user: "missing".into(),
425 };
426 assert!(runner.run(&miss).is_err(), "unknown key must hard-error");
427 }
428
429 #[test]
430 fn recording_runner_captures_then_replays() {
431 let mut seed = Recording::new(fp());
432 let req = ModelRequest {
433 system: "sys".into(),
434 user: "q".into(),
435 };
436 seed.entries.insert(req.key(), ModelResponse::new("answer"));
437 let recorder = RecordingRunner::new(RecordedRunner::new(seed));
438 let _ = recorder.run(&req).unwrap();
439 let captured = recorder.into_recording();
440 assert_eq!(captured.entries.get(&req.key()).unwrap().text, "answer");
441 }
442}