Skip to main content

lean_ctx/core/eval_ab/
model.rs

1//! Pinned, reproducible model adapter (#234).
2//!
3//! The harness talks to a model through one small [`ModelRunner`] trait. Two real
4//! implementations ship:
5//!
6//! * [`OpenAiRunner`] — a synchronous OpenAI-compatible chat client (OpenAI, Azure OpenAI,
7//!   vLLM, llama.cpp, Ollama's `/v1` …). Decoding is pinned (`temperature = 0`, fixed `seed`)
8//!   so a compliant provider is as deterministic as it can be.
9//! * [`RecordedRunner`] — strict replay of responses previously captured from a real provider.
10//!   Missing keys are a hard error (never a silent fallback), which is what makes CI runs and
11//!   the determinism digest byte-stable across machines.
12//!
13//! [`RecordingRunner`] wraps any real runner and captures every response so an operator can
14//! produce a replay file once with `eval ab --record`. Secrets (API keys) never enter a
15//! fingerprint, recording, or report.
16
17use std::cell::RefCell;
18use std::collections::BTreeMap;
19use std::path::Path;
20
21use anyhow::{Context, Result, anyhow, bail};
22use serde::{Deserialize, Serialize};
23
24use super::sha256_hex;
25
26/// Provider label for the OpenAI-compatible chat runner.
27pub const PROVIDER_OPENAI: &str = "openai-compatible";
28/// Provider label for the strict replay runner.
29pub const PROVIDER_RECORDED: &str = "recorded";
30
31/// Decoding parameters pinned for reproducibility.
32#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
33pub struct ModelParams {
34    /// Provider model identifier, e.g. `gpt-4o-mini` or `qwen2.5-coder:7b`.
35    pub model: String,
36    pub temperature: f64,
37    pub top_p: f64,
38    pub max_tokens: u32,
39    /// Best-effort decoding seed (forwarded to providers that honour it).
40    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/// Identifies exactly which model + params produced a set of answers. Embedded in the report
56/// and the determinism digest so a third party knows precisely what was run.
57#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
58pub struct ModelFingerprint {
59    /// `openai-compatible` or `recorded`.
60    pub provider: String,
61    /// Base URL or recording path — never contains credentials.
62    pub endpoint: String,
63    pub params: ModelParams,
64}
65
66impl ModelFingerprint {
67    /// Stable hex digest of the fingerprint (canonical JSON over sorted keys via serde).
68    pub fn digest(&self) -> String {
69        let canonical = serde_json::to_vec(self).unwrap_or_default();
70        sha256_hex(&canonical)
71    }
72}
73
74/// A single chat request (one system + one user turn). Deliberately minimal so the recording
75/// key is stable across runs.
76#[derive(Debug, Clone, PartialEq)]
77pub struct ModelRequest {
78    pub system: String,
79    pub user: String,
80}
81
82impl ModelRequest {
83    /// Content-addressed replay key: hex SHA-256 over `system` and `user`.
84    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/// The model's answer plus a content digest for auditing.
94#[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    /// Hex SHA-256 of the answer text (goes into the determinism digest).
105    pub fn digest(&self) -> String {
106        sha256_hex(self.text.as_bytes())
107    }
108}
109
110/// Anything that can turn a request into a response under a fixed fingerprint.
111pub trait ModelRunner {
112    /// The pinned identity of this runner (model, params, provider).
113    fn fingerprint(&self) -> &ModelFingerprint;
114    /// Executes one request. Implementations must be deterministic given the same fingerprint.
115    fn run(&self, req: &ModelRequest) -> Result<ModelResponse>;
116}
117
118// ---------------------------------------------------------------------------
119// OpenAI-compatible runner (real HTTP)
120// ---------------------------------------------------------------------------
121
122/// Synchronous OpenAI-compatible chat client. Credentials are held only in memory.
123pub struct OpenAiRunner {
124    fingerprint: ModelFingerprint,
125    api_key: String,
126}
127
128impl OpenAiRunner {
129    /// Builds a runner against `base_url` (e.g. `https://api.openai.com/v1`).
130    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    /// Builds a runner from the standard environment:
147    /// `LEAN_CTX_EVAL_MODEL_URL`, `LEAN_CTX_EVAL_MODEL_KEY`, `LEAN_CTX_EVAL_MODEL`.
148    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// ---------------------------------------------------------------------------
252// Recording + strict replay
253// ---------------------------------------------------------------------------
254
255/// A captured set of real responses, keyed by [`ModelRequest::key`]. Serialized as JSON with a
256/// `BTreeMap` so the on-disk form is byte-stable.
257#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
258pub struct Recording {
259    /// `lean-ctx.eval-recording`.
260    pub kind: String,
261    /// The fingerprint that produced these responses (used to label the replay runner).
262    pub fingerprint: ModelFingerprint,
263    /// `request key -> response`.
264    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    /// Loads + validates a recording file.
279    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    /// Writes the recording as pretty JSON (creating parent dirs).
291    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
302/// Strict replay runner: every request must hit a recorded entry, or it errors. This guarantees
303/// the run is fully deterministic and machine-independent.
304pub struct RecordedRunner {
305    recording: Recording,
306}
307
308impl RecordedRunner {
309    pub fn new(recording: Recording) -> Self {
310        Self { recording }
311    }
312
313    /// Loads a replay runner from a recording file.
314    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
333/// Wraps a real runner and captures every response so it can be replayed later. Pass-through:
334/// the wrapped runner's answers are returned unchanged.
335pub 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    /// Consumes the wrapper and returns everything captured so far.
350    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}