Skip to main content

el_engine_candle/
lib.rs

1//! `el-engine-candle` — inference engine adapter over **Candle** (ADR-002),
2//! implementing [`el_runtime::InferenceEngine`] / `RuntimeAcl`.
3//!
4//! Consumers supply their own model file; see [`CandleEngine::from_path`] and
5//! [`CandleEngine::from_bytes`].  For tests that need a working engine without
6//! a model asset, use [`CandleEngine::toy`].
7//!
8//! Expected GGUF tensor names:
9//! - `token_embd.weight`  — embedding table  `[vocab, dim]`
10//! - `output.weight` or `lm_head.weight` — lm-head  `[vocab, dim]`  (standard Llama layout)
11//!
12//! Float logits are quantised to integer milli-logits at the ACL boundary, so
13//! Candle's `Tensor`/`Device` types never cross into the domain.
14
15#![forbid(unsafe_code)]
16
17use candle_core::{Device, Tensor};
18use el_core::{
19    ChatMessage, ChatRequest, ChatResponse, ChatRole, ChatToken, EdgeError, LlmProvider, Result,
20    SessionConfig, SessionId, Token,
21};
22use el_provenance::LoadPermit;
23use el_runtime::{InferenceEngine, InferenceSession, Ports};
24
25/// Candle-backed inference engine.
26pub struct CandleEngine {
27    embed: Tensor,
28    w_out: Tensor,
29    vocab: usize,
30    eos: Token,
31}
32
33impl CandleEngine {
34    /// Build a deterministic toy model on the CPU — no model file required.
35    ///
36    /// Uses fixed synthetic weights so tests are deterministic.
37    pub fn toy(vocab: usize, dim: usize, eos: Token) -> Result<Self> {
38        let device = Device::Cpu;
39
40        let embed_data: Vec<f32> = (0..vocab * dim)
41            .map(|k| {
42                let (i, j) = (k / dim, k % dim);
43                (((i + j) % 7) as f32) * 0.1
44            })
45            .collect();
46        let wout_data: Vec<f32> = (0..dim * vocab)
47            .map(|k| {
48                let (a, b) = (k / vocab, k % vocab);
49                ((((a * 31 + b * 17) % 13) as f32) * 0.1) - 0.6
50            })
51            .collect();
52
53        let embed = Tensor::from_vec(embed_data, (vocab, dim), &device)
54            .map_err(|_| EdgeError::Engine("candle: embed tensor build failed"))?;
55        let w_out = Tensor::from_vec(wout_data, (dim, vocab), &device)
56            .map_err(|_| EdgeError::Engine("candle: w_out tensor build failed"))?;
57
58        Ok(Self {
59            embed,
60            w_out,
61            vocab,
62            eos,
63        })
64    }
65
66    /// Load `token_embd.weight` and `output.weight` from a consumer-supplied GGUF file.
67    ///
68    /// # Limitations
69    /// This engine's forward pass is `embed[last_token] · w_out` — a single linear
70    /// projection.  Only these two tensors are used; transformer blocks, attention,
71    /// RoPE, and norms present in the GGUF are ignored.  Logits will not match a
72    /// real Llama/Mistral/etc. forward.  This is the ADR-002 engine-seam proof; for
73    /// a full transformer forward implement a separate [`InferenceEngine`] using
74    /// `candle-transformers`.
75    pub fn from_path(path: impl AsRef<std::path::Path>, eos: Token) -> Result<Self> {
76        let file = std::fs::File::open(path.as_ref())
77            .map_err(|_| EdgeError::Engine("model file not found or not readable"))?;
78        Self::load_gguf(&mut std::io::BufReader::new(file), eos)
79    }
80
81    /// Load from raw bytes (WASM / memory-mapped scenarios).
82    ///
83    /// Same limitations as [`Self::from_path`]: only `token_embd.weight` and
84    /// `output.weight` are used; the forward is `embed[last] · w_out`.
85    pub fn from_bytes(data: &[u8], eos: Token) -> Result<Self> {
86        Self::load_gguf(&mut std::io::Cursor::new(data), eos)
87    }
88
89    fn load_gguf<R: std::io::Read + std::io::Seek>(reader: &mut R, eos: Token) -> Result<Self> {
90        use candle_core::quantized::gguf_file;
91
92        let content = gguf_file::Content::read(reader)
93            .map_err(|_| EdgeError::Engine("GGUF: invalid or unrecognised file"))?;
94        let device = Device::Cpu;
95
96        let embed = content
97            .tensor(reader, "token_embd.weight", &device)
98            .map_err(|_| EdgeError::Engine("GGUF: missing 'token_embd.weight'"))?
99            .dequantize(&device)
100            .map_err(|_| EdgeError::Engine("GGUF: cannot dequantize embed tensor"))?;
101
102        let (vocab, dim) = match embed.shape().dims() {
103            [v, d] => (*v, *d),
104            _ => return Err(EdgeError::Engine("GGUF: 'token_embd.weight' must be 2-D")),
105        };
106
107        let raw_w_q = match content.tensor(reader, "output.weight", &device) {
108            Ok(t) => t,
109            Err(_) => content
110                .tensor(reader, "lm_head.weight", &device)
111                .map_err(|_| {
112                    EdgeError::Engine("GGUF: missing 'output.weight' / 'lm_head.weight'")
113                })?,
114        };
115        let raw_w = raw_w_q
116            .dequantize(&device)
117            .map_err(|_| EdgeError::Engine("GGUF: cannot dequantize output weight"))?;
118
119        // Standard GGUF / Llama convention: output.weight is [vocab, dim].
120        // We need [dim, vocab] so that embed_row [1,dim] × w_out [dim,vocab] → logits [1,vocab].
121        let w_out = match raw_w.shape().dims() {
122            [v, _d] if *v == vocab => raw_w
123                .t()
124                .map_err(|_| EdgeError::Engine("GGUF: failed to transpose output weight"))?,
125            _ => raw_w,
126        };
127
128        // Validate that the output weight's inner dimension matches the embedding dimension.
129        // A mismatch would silently produce all-zero logits at inference time.
130        match w_out.shape().dims() {
131            [d, v] if *d == dim && *v == vocab => {}
132            _ => return Err(EdgeError::Engine(
133                "GGUF: output weight shape incompatible with embed dim — expected [dim, vocab] after transpose",
134            )),
135        }
136
137        Ok(Self {
138            embed,
139            w_out,
140            vocab,
141            eos,
142        })
143    }
144
145    /// One real Candle forward: `embed[last] · w_out` → length-`vocab` logits.
146    fn forward(&self, last: usize) -> candle_core::Result<Vec<f32>> {
147        let row = self.embed.narrow(0, last, 1)?; // [1, dim]
148        let logits = row.matmul(&self.w_out)?; // [1, vocab]
149        Ok(logits.to_vec2::<f32>()?.remove(0))
150    }
151}
152
153impl InferenceEngine for CandleEngine {
154    fn prefill(&mut self, tokens: &[Token]) -> Result<u32> {
155        Ok(tokens.len() as u32)
156    }
157
158    fn next_logits(&mut self, committed: &[Token]) -> Vec<i32> {
159        let last = committed
160            .last()
161            .copied()
162            .unwrap_or(0)
163            .min(self.vocab as u32 - 1) as usize;
164        match self.forward(last) {
165            Ok(logits) => logits.iter().map(|x| (x * 1000.0).round() as i32).collect(),
166            Err(_) => vec![0; self.vocab],
167        }
168    }
169
170    fn eos_token(&self) -> Token {
171        self.eos
172    }
173}
174
175// ── LlmProvider (text-level) wrapper (ADR-010) ───────────────────────────────
176
177/// Wraps a `CandleEngine` behind the `LlmProvider` trait using a byte-level
178/// tokenizer.  A production build would swap in a HuggingFace tokenizer loaded
179/// from the model file.
180pub struct LocalLlmProvider {
181    session: std::sync::Mutex<InferenceSession<CandleEngine>>,
182    vocab: usize,
183}
184
185impl LocalLlmProvider {
186    /// Load from a consumer-supplied GGUF file.
187    pub fn from_path(
188        path: impl AsRef<std::path::Path>,
189        eos: Token,
190        permit: LoadPermit,
191    ) -> Result<Self> {
192        let engine = CandleEngine::from_path(path, eos)?;
193        let vocab = engine.vocab;
194        let session = InferenceSession::new(SessionId(1), SessionConfig::default(), engine, permit);
195        Ok(Self {
196            session: std::sync::Mutex::new(session),
197            vocab,
198        })
199    }
200
201    /// Build a toy provider for testing.
202    pub fn toy(vocab: usize, dim: usize, eos: Token, permit: LoadPermit) -> Result<Self> {
203        let engine = CandleEngine::toy(vocab, dim, eos)?;
204        let session = InferenceSession::new(SessionId(1), SessionConfig::default(), engine, permit);
205        Ok(Self {
206            session: std::sync::Mutex::new(session),
207            vocab,
208        })
209    }
210
211    fn encode(&self, text: &str) -> Vec<Token> {
212        text.bytes()
213            .map(|b| (b as Token) % self.vocab as Token)
214            .collect()
215    }
216
217    fn decode(tokens: &[Token]) -> String {
218        tokens
219            .iter()
220            .map(|&t| {
221                let b = (t & 0xFF) as u8;
222                if b.is_ascii_graphic() || b == b' ' {
223                    b as char
224                } else {
225                    '?'
226                }
227            })
228            .collect()
229    }
230
231    fn format_messages(messages: &[ChatMessage]) -> String {
232        messages
233            .iter()
234            .map(|m| {
235                let role = match m.role {
236                    ChatRole::System => "system",
237                    ChatRole::User => "user",
238                    ChatRole::Assistant => "assistant",
239                };
240                format!("{role}: {}", m.content)
241            })
242            .collect::<Vec<_>>()
243            .join("\n")
244    }
245}
246
247impl LlmProvider for LocalLlmProvider {
248    fn chat(&self, req: &ChatRequest) -> Result<ChatResponse> {
249        let prompt = Self::format_messages(&req.messages);
250        let prompt_tokens = self.encode(&prompt);
251        let prompt_len = prompt_tokens.len() as u32;
252        let max = req.max_tokens.unwrap_or(64);
253
254        let mut session = self.session.lock().unwrap();
255        session.reset();
256        let ports = Ports::permissive();
257        session.load_prompt(&ports, &prompt_tokens)?;
258        session.generate(&ports, max)?;
259
260        let output = session.output().to_vec();
261        let completion_len = output.len() as u32;
262
263        Ok(ChatResponse {
264            content: Self::decode(&output),
265            model: "local/candle".into(),
266            prompt_tokens: prompt_len,
267            completion_tokens: completion_len,
268        })
269    }
270
271    fn chat_stream(&self, req: &ChatRequest, on_token: &mut dyn FnMut(ChatToken)) -> Result<()> {
272        let resp = self.chat(req)?;
273        for ch in resp.content.chars() {
274            on_token(ChatToken {
275                text: ch.to_string(),
276                is_final: false,
277            });
278        }
279        on_token(ChatToken {
280            text: String::new(),
281            is_final: true,
282        });
283        Ok(())
284    }
285}
286
287// ── Real Qwen2 transformer engine + chat provider (ADR-002 + ADR-010) ────────
288//
289// Unlike `CandleEngine` (a single linear projection used as the engine-seam
290// proof) this runs a genuine Qwen2 transformer forward via `candle-transformers`
291// with a real HuggingFace tokenizer, so it produces coherent chat. It plugs into
292// the SAME `el_runtime::InferenceSession` decode loop as every other engine —
293// nothing in the SDK pipeline is bypassed.
294
295use candle_transformers::models::quantized_qwen2::ModelWeights as Qwen2Weights;
296use el_core::{ModelId, ModelVersion};
297use el_provenance::{ModelArtifact, SignatureVerifier};
298use tokenizers::Tokenizer;
299
300// ── Opt-in benchmark instrumentation (EL_BENCH=1) ────────────────────────────
301//
302// Zero-cost when `EL_BENCH` is unset: `enabled()` short-circuits and no timing
303// is taken. When set, `QwenChatProvider::chat` prints a per-phase breakdown and
304// per-forward attribution (model compute vs. seam quantisation vs. runtime loop)
305// to stderr. Diagnostics only — not part of the SDK's public behaviour.
306mod bench {
307    use std::cell::Cell;
308    use std::sync::OnceLock;
309    use std::time::Duration;
310
311    static ENABLED: OnceLock<bool> = OnceLock::new();
312
313    /// True iff the `EL_BENCH` environment variable is present (read once).
314    pub fn enabled() -> bool {
315        *ENABLED.get_or_init(|| std::env::var_os("EL_BENCH").is_some())
316    }
317
318    thread_local! {
319        static FWD_TOTAL: Cell<Duration> = const { Cell::new(Duration::ZERO) };
320        static FWD_MODEL: Cell<Duration> = const { Cell::new(Duration::ZERO) };
321        static FWD_CALLS: Cell<u64> = const { Cell::new(0) };
322    }
323
324    /// Accumulate one `forward_one` sample: `total` is the whole seam call,
325    /// `model` is just the candle transformer forward inside it.
326    pub fn record(total: Duration, model: Duration) {
327        FWD_TOTAL.with(|c| c.set(c.get() + total));
328        FWD_MODEL.with(|c| c.set(c.get() + model));
329        FWD_CALLS.with(|c| c.set(c.get() + 1));
330    }
331
332    /// Read and reset the forward accumulators: `(total, model, calls)`.
333    pub fn take() -> (Duration, Duration, u64) {
334        (
335            FWD_TOTAL.replace(Duration::ZERO),
336            FWD_MODEL.replace(Duration::ZERO),
337            FWD_CALLS.replace(0),
338        )
339    }
340}
341
342/// A real Qwen2 transformer `InferenceEngine`.
343///
344/// Holds candle's stateful KV cache. Within one generation it is fed
345/// incrementally (prefill, then one new token per `next_logits` call); candle
346/// exposes no public cache reset, so a fresh conversation builds a new engine.
347/// Float logits are quantised to integer milli-logits at the seam, exactly like
348/// [`CandleEngine`], so the runtime stays float-free.
349pub struct QwenEngine {
350    model: Qwen2Weights,
351    device: Device,
352    /// Absolute KV position written so far (candle's `index_pos`).
353    index_pos: usize,
354    /// How many of the runtime-`committed` tokens have already been fed.
355    fed: usize,
356    /// Milli-logits produced after the most recent forward.
357    last_logits: Vec<i32>,
358    vocab: usize,
359    eos: Token,
360}
361
362impl QwenEngine {
363    /// Load Qwen2 weights from a consumer-supplied GGUF file.
364    pub fn from_path(path: impl AsRef<std::path::Path>, eos: Token) -> Result<Self> {
365        use candle_core::quantized::gguf_file;
366        let mut file = std::fs::File::open(path.as_ref())
367            .map_err(|_| EdgeError::Engine("model file not found or not readable"))?;
368        let content = gguf_file::Content::read(&mut file)
369            .map_err(|_| EdgeError::Engine("GGUF: invalid or unrecognised file"))?;
370        let device = Device::Cpu;
371        let model = Qwen2Weights::from_gguf(content, &mut file, &device)
372            .map_err(|_| EdgeError::Engine("GGUF: failed to load Qwen2 weights"))?;
373        Ok(Self {
374            model,
375            device,
376            index_pos: 0,
377            fed: 0,
378            last_logits: Vec::new(),
379            vocab: 0,
380            eos,
381        })
382    }
383
384    /// One forward over a single token at the current position; advances the KV
385    /// cache and returns milli-logits for the next token.
386    fn forward_one(&mut self, token: Token) -> Result<Vec<i32>> {
387        let t_total = bench::enabled().then(std::time::Instant::now);
388
389        let input = Tensor::from_vec(vec![token], (1, 1), &self.device)
390            .map_err(|_| EdgeError::Engine("candle: input tensor build failed"))?;
391
392        let t_model = bench::enabled().then(std::time::Instant::now);
393        let logits = self
394            .model
395            .forward(&input, self.index_pos)
396            .map_err(|_| EdgeError::Engine("candle: Qwen2 forward failed"))?;
397        let model_dur = t_model.map(|t| t.elapsed()).unwrap_or_default();
398
399        self.index_pos += 1;
400        let row = logits
401            .squeeze(0)
402            .map_err(|_| EdgeError::Engine("candle: squeeze logits failed"))?;
403        let floats = row
404            .to_vec1::<f32>()
405            .map_err(|_| EdgeError::Engine("candle: logits to vec failed"))?;
406        let out: Vec<i32> = floats.iter().map(|x| (x * 1000.0).round() as i32).collect();
407
408        if let Some(t) = t_total {
409            bench::record(t.elapsed(), model_dur);
410        }
411        Ok(out)
412    }
413}
414
415impl InferenceEngine for QwenEngine {
416    fn prefill(&mut self, tokens: &[Token]) -> Result<u32> {
417        self.index_pos = 0;
418        self.fed = 0;
419        for &t in tokens {
420            self.last_logits = self.forward_one(t)?;
421        }
422        self.vocab = self.last_logits.len();
423        Ok(tokens.len() as u32)
424    }
425
426    fn next_logits(&mut self, committed: &[Token]) -> Vec<i32> {
427        // Feed any newly committed (generated) tokens beyond what we've seen.
428        // `committed` grows by exactly one per decode step, so this feeds the
429        // token the runtime just sampled and returns the next distribution.
430        while self.fed < committed.len() {
431            let t = committed[self.fed];
432            match self.forward_one(t) {
433                Ok(l) => self.last_logits = l,
434                Err(_) => return vec![0; self.vocab.max(1)],
435            }
436            self.fed += 1;
437        }
438        self.last_logits.clone()
439    }
440
441    fn eos_token(&self) -> Token {
442        self.eos
443    }
444}
445
446/// A real local chat backend: a Qwen2 GGUF model + its tokenizer, driven
447/// through [`el_runtime::InferenceSession`].
448///
449/// Each `chat` call renders the whole conversation to Qwen2.5 ChatML, builds a
450/// fresh [`QwenEngine`] (candle has no public KV-cache reset), then runs the
451/// SDK's standard provenance-gated session: `load_prompt` (prefill) →
452/// `generate` (grammar mask → safety steer → greedy commit). The provider holds
453/// no mutable session state, so it is `Send + Sync` without locking.
454pub struct QwenChatProvider {
455    model_path: std::path::PathBuf,
456    tokenizer: Tokenizer,
457    permit: LoadPermit,
458    eos: Token,
459    default_max_tokens: u32,
460    model_label: String,
461}
462
463impl QwenChatProvider {
464    /// Load a Qwen2 GGUF model and its `tokenizer.json` from local paths.
465    pub fn from_paths(
466        model_path: impl AsRef<std::path::Path>,
467        tokenizer_path: impl AsRef<std::path::Path>,
468    ) -> Result<Self> {
469        let model_path = model_path.as_ref().to_path_buf();
470        if !model_path.exists() {
471            return Err(EdgeError::Engine("model file not found"));
472        }
473        let tokenizer = Tokenizer::from_file(tokenizer_path.as_ref())
474            .map_err(|_| EdgeError::Engine("failed to load tokenizer.json"))?;
475
476        // Stop token: Qwen2.5 ChatML turn terminator (fallback to its known id).
477        let eos = tokenizer.token_to_id("<|im_end|>").unwrap_or(151_645);
478
479        let model_label = model_path
480            .file_stem()
481            .and_then(|s| s.to_str())
482            .map(|s| format!("local/{s}"))
483            .unwrap_or_else(|| "local/qwen2".to_string());
484
485        Ok(Self {
486            model_path,
487            tokenizer,
488            permit: local_load_permit()?,
489            eos,
490            default_max_tokens: 512,
491            model_label,
492        })
493    }
494
495    fn encode(&self, text: &str) -> Result<Vec<Token>> {
496        let enc = self
497            .tokenizer
498            .encode(text, false)
499            .map_err(|_| EdgeError::Engine("tokenizer encode failed"))?;
500        Ok(enc.get_ids().to_vec())
501    }
502
503    fn decode(&self, ids: &[Token]) -> Result<String> {
504        self.tokenizer
505            .decode(ids, true)
506            .map_err(|_| EdgeError::Engine("tokenizer decode failed"))
507    }
508}
509
510impl LlmProvider for QwenChatProvider {
511    fn chat(&self, req: &ChatRequest) -> Result<ChatResponse> {
512        let prompt = render_chatml(&req.messages);
513
514        let t_encode = bench::enabled().then(std::time::Instant::now);
515        let prompt_tokens = self.encode(&prompt)?;
516        let d_encode = t_encode.map(|t| t.elapsed()).unwrap_or_default();
517
518        // Fresh engine + session each turn (candle KV cache has no public reset);
519        // the full conversation is re-prefilled. This is the standard SDK path —
520        // provenance permit, session lifecycle, decode loop — not a shortcut.
521        let t_load = bench::enabled().then(std::time::Instant::now);
522        let engine = QwenEngine::from_path(&self.model_path, self.eos)?;
523        let d_load = t_load.map(|t| t.elapsed()).unwrap_or_default();
524
525        let mut session =
526            InferenceSession::new(SessionId(1), SessionConfig::default(), engine, self.permit);
527        let ports = Ports::permissive();
528
529        let _ = bench::take(); // clear forward accumulators before prefill
530        let t_prefill = bench::enabled().then(std::time::Instant::now);
531        session.load_prompt(&ports, &prompt_tokens)?;
532        let d_prefill = t_prefill.map(|t| t.elapsed()).unwrap_or_default();
533        let (pf_total, pf_model, pf_calls) = bench::take();
534
535        let max = req.max_tokens.unwrap_or(self.default_max_tokens);
536        let t_decode = bench::enabled().then(std::time::Instant::now);
537        session.generate(&ports, max)?;
538        let d_decode = t_decode.map(|t| t.elapsed()).unwrap_or_default();
539        let (dc_total, dc_model, dc_calls) = bench::take();
540
541        let out = session.output();
542        let completion_tokens = out.len() as u32;
543
544        let t_detok = bench::enabled().then(std::time::Instant::now);
545        let content = self.decode(out)?.trim().to_string();
546        let d_detok = t_detok.map(|t| t.elapsed()).unwrap_or_default();
547
548        if bench::enabled() {
549            report_breakdown(
550                prompt_tokens.len() as u32,
551                completion_tokens,
552                d_load,
553                d_encode,
554                d_prefill,
555                d_decode,
556                d_detok,
557                (pf_total, pf_model, pf_calls),
558                (dc_total, dc_model, dc_calls),
559            );
560        }
561
562        Ok(ChatResponse {
563            content,
564            model: self.model_label.clone(),
565            prompt_tokens: prompt_tokens.len() as u32,
566            completion_tokens,
567        })
568    }
569
570    fn chat_stream(&self, req: &ChatRequest, on_token: &mut dyn FnMut(ChatToken)) -> Result<()> {
571        // The runtime decode loop runs to completion internally (no per-token
572        // hook), so — like the toy `LocalLlmProvider` — we stream the finished
573        // reply out character by character.
574        let resp = self.chat(req)?;
575        for ch in resp.content.chars() {
576            on_token(ChatToken {
577                text: ch.to_string(),
578                is_final: false,
579            });
580        }
581        on_token(ChatToken {
582            text: String::new(),
583            is_final: true,
584        });
585        Ok(())
586    }
587}
588
589/// Print an `EL_BENCH` per-phase + per-forward breakdown for one `chat()` call.
590#[allow(clippy::too_many_arguments)]
591fn report_breakdown(
592    prompt_tokens: u32,
593    completion_tokens: u32,
594    d_load: std::time::Duration,
595    d_encode: std::time::Duration,
596    d_prefill: std::time::Duration,
597    d_decode: std::time::Duration,
598    d_detok: std::time::Duration,
599    prefill_fwd: (std::time::Duration, std::time::Duration, u64),
600    decode_fwd: (std::time::Duration, std::time::Duration, u64),
601) {
602    let ms = |d: std::time::Duration| d.as_secs_f64() * 1000.0;
603    let total = d_load + d_encode + d_prefill + d_decode + d_detok;
604    let pct = |d: std::time::Duration| {
605        if total.as_secs_f64() > 0.0 {
606            d.as_secs_f64() / total.as_secs_f64() * 100.0
607        } else {
608            0.0
609        }
610    };
611    let tps = |n: u32, d: std::time::Duration| {
612        if d.as_secs_f64() > 0.0 {
613            n as f64 / d.as_secs_f64()
614        } else {
615            0.0
616        }
617    };
618
619    let (pf_total, pf_model, pf_calls) = prefill_fwd;
620    let (dc_total, dc_model, dc_calls) = decode_fwd;
621    let dc_loop = d_decode.saturating_sub(dc_total);
622    let dc_seam = dc_total.saturating_sub(dc_model);
623    let per_tok = |d: std::time::Duration, n: u64| if n > 0 { ms(d) / n as f64 } else { 0.0 };
624
625    eprintln!("\n┌─ EL_BENCH chat() breakdown ───────────────────────────────");
626    eprintln!("│ prompt_tokens={prompt_tokens}  completion_tokens={completion_tokens}");
627    eprintln!("│ phase           wall(ms)    %total   throughput");
628    eprintln!(
629        "│ model load    {:>9.1}  {:>6.1}%   (read+dequantize GGUF)",
630        ms(d_load),
631        pct(d_load)
632    );
633    eprintln!(
634        "│ tokenize       {:>9.2}  {:>6.1}%",
635        ms(d_encode),
636        pct(d_encode)
637    );
638    eprintln!(
639        "│ prefill       {:>9.1}  {:>6.1}%   {:>7.1} tok/s",
640        ms(d_prefill),
641        pct(d_prefill),
642        tps(prompt_tokens, d_prefill)
643    );
644    eprintln!(
645        "│ decode        {:>9.1}  {:>6.1}%   {:>7.1} tok/s",
646        ms(d_decode),
647        pct(d_decode),
648        tps(completion_tokens, d_decode)
649    );
650    eprintln!(
651        "│ detokenize     {:>9.2}  {:>6.1}%",
652        ms(d_detok),
653        pct(d_detok)
654    );
655    eprintln!("│ TOTAL         {:>9.1}", ms(total));
656    eprintln!("│ ─ forward attribution (where prefill+decode time goes) ─");
657    eprintln!(
658        "│ prefill: {} fwd calls, model {:.1}ms, seam {:.1}ms, loop {:.1}ms",
659        pf_calls,
660        ms(pf_model),
661        ms(pf_total.saturating_sub(pf_model)),
662        ms(d_prefill.saturating_sub(pf_total)),
663    );
664    eprintln!(
665        "│ decode : {} fwd calls, model {:.1}ms, seam {:.1}ms, loop {:.1}ms",
666        dc_calls,
667        ms(dc_model),
668        ms(dc_seam),
669        ms(dc_loop),
670    );
671    eprintln!(
672        "│ per decoded token: {:.2}ms total = model {:.2} + seam {:.2} + loop {:.2}",
673        per_tok(d_decode, dc_calls),
674        per_tok(dc_model, dc_calls),
675        per_tok(dc_seam, dc_calls),
676        per_tok(dc_loop, dc_calls),
677    );
678    eprintln!("└───────────────────────────────────────────────────────────");
679}
680
681/// Render a conversation as Qwen2.5 ChatML and open an assistant turn.
682fn render_chatml(messages: &[ChatMessage]) -> String {
683    let mut s = String::new();
684    for m in messages {
685        let role = match m.role {
686            ChatRole::System => "system",
687            ChatRole::User => "user",
688            ChatRole::Assistant => "assistant",
689        };
690        s.push_str("<|im_start|>");
691        s.push_str(role);
692        s.push('\n');
693        s.push_str(&m.content);
694        s.push_str("<|im_end|>\n");
695    }
696    s.push_str("<|im_start|>assistant\n");
697    s
698}
699
700/// Obtain a [`LoadPermit`] through the real ADR-006 gate for a user-supplied
701/// local model. There is no detached signature to check for a file the user
702/// downloaded themselves, so a trust-the-local-file verifier is used — the
703/// point is to go through the gate API the runtime requires, not to bypass it.
704fn local_load_permit() -> Result<LoadPermit> {
705    struct LocalFileTrust;
706    impl SignatureVerifier for LocalFileTrust {
707        fn verify(&self, _bytes: &[u8], _sig: &[u8], _key: u32) -> bool {
708            true
709        }
710    }
711    let mut artifact = ModelArtifact::new(
712        ModelId(1),
713        ModelVersion::new(0, 1, 0),
714        el_core::ModelFormat::Gguf,
715    );
716    artifact.verify(&LocalFileTrust, b"local-file", b"local-file", 0);
717    artifact.ensure_loadable()
718}
719
720#[cfg(test)]
721mod tests {
722    use super::*;
723    use el_runtime::InferenceEngine;
724
725    // ── helpers ──────────────────────────────────────────────────────────────
726
727    fn ok_permit() -> LoadPermit {
728        use el_core::{ModelFormat, ModelId, ModelVersion};
729        use el_provenance::{ModelArtifact, SignatureVerifier};
730        struct OkV;
731        impl SignatureVerifier for OkV {
732            fn verify(&self, _: &[u8], _: &[u8], _: u32) -> bool {
733                true
734            }
735        }
736        let mut a = ModelArtifact::new(ModelId(1), ModelVersion::new(0, 1, 0), ModelFormat::Gguf);
737        a.verify(&OkV, b"w", b"s", 0);
738        a.ensure_loadable().unwrap()
739    }
740
741    /// Build a minimal but spec-compliant GGUF v3 file in memory.
742    ///
743    /// Layout:  no KV metadata, two F32 tensors:
744    ///   `token_embd.weight`  [vocab, dim]  at offset 0
745    ///   `output.weight`      [vocab, dim]  at offset vocab*dim*4
746    ///
747    /// GGUF stores dimensions innermost-first; candle reverses them on read.
748    fn make_minimal_gguf(vocab: usize, dim: usize) -> Vec<u8> {
749        let mut w: Vec<u8> = Vec::new();
750
751        // Header
752        w.extend_from_slice(b"GGUF");
753        w.extend_from_slice(&3u32.to_le_bytes()); // version 3
754        w.extend_from_slice(&2u64.to_le_bytes()); // n_tensors
755        w.extend_from_slice(&0u64.to_le_bytes()); // n_kv (none)
756
757        let tensor_bytes = (vocab * dim * 4) as u64;
758
759        // token_embd.weight: [vocab, dim] → GGUF dims [dim, vocab]
760        let name = b"token_embd.weight";
761        w.extend_from_slice(&(name.len() as u64).to_le_bytes());
762        w.extend_from_slice(name);
763        w.extend_from_slice(&2u32.to_le_bytes());
764        w.extend_from_slice(&(dim as u64).to_le_bytes()); // innermost
765        w.extend_from_slice(&(vocab as u64).to_le_bytes()); // outermost
766        w.extend_from_slice(&0u32.to_le_bytes()); // F32
767        w.extend_from_slice(&0u64.to_le_bytes()); // offset 0
768
769        // output.weight: [vocab, dim] → GGUF dims [dim, vocab]; loader will transpose
770        let name = b"output.weight";
771        w.extend_from_slice(&(name.len() as u64).to_le_bytes());
772        w.extend_from_slice(name);
773        w.extend_from_slice(&2u32.to_le_bytes());
774        w.extend_from_slice(&(dim as u64).to_le_bytes());
775        w.extend_from_slice(&(vocab as u64).to_le_bytes());
776        w.extend_from_slice(&0u32.to_le_bytes());
777        w.extend_from_slice(&tensor_bytes.to_le_bytes()); // offset after embed
778
779        // Pad to 32-byte alignment
780        let pad = (32usize.wrapping_sub(w.len() % 32)) % 32;
781        w.resize(w.len() + pad, 0u8);
782
783        // Tensor data (both tensors, row-major f32)
784        for i in 0..(vocab * dim * 2) {
785            w.extend_from_slice(&(i as f32 * 0.1f32).to_le_bytes());
786        }
787
788        w
789    }
790
791    // ── toy-model tests (unchanged) ──────────────────────────────────────────
792
793    #[test]
794    fn real_candle_forward_is_deterministic_and_right_shape() {
795        let mut eng = CandleEngine::toy(8, 4, 7).unwrap();
796        let a = eng.next_logits(&[2]);
797        let b = eng.next_logits(&[2]);
798        assert_eq!(a.len(), 8, "logits length == vocab");
799        assert_eq!(a, b, "fixed weights → deterministic real-tensor forward");
800        let c = eng.next_logits(&[5]);
801        assert_ne!(a, c);
802    }
803
804    #[test]
805    fn drives_the_runtime_end_to_end() {
806        use el_core::{ModelFormat, ModelId, ModelVersion, SessionConfig, SessionId, StopReason};
807        use el_provenance::{ModelArtifact, SignatureVerifier};
808
809        struct OkVerifier;
810        impl SignatureVerifier for OkVerifier {
811            fn verify(&self, _: &[u8], _: &[u8], _: u32) -> bool {
812                true
813            }
814        }
815        let mut art = ModelArtifact::new(
816            ModelId(1),
817            ModelVersion::new(0, 1, 0),
818            ModelFormat::Safetensors,
819        );
820        art.verify(&OkVerifier, b"w", b"s", 1);
821        let permit = art.ensure_loadable().unwrap();
822
823        let eng = CandleEngine::toy(16, 8, 9999).unwrap();
824        let mut session =
825            InferenceSession::new(SessionId(1), SessionConfig::default(), eng, permit);
826        let ports = Ports::permissive();
827        session.load_prompt(&ports, &[1, 2, 3]).unwrap();
828
829        let stop = session.generate(&ports, 4).unwrap();
830        assert_eq!(stop, StopReason::MaxTokens);
831        assert_eq!(session.output().len(), 4);
832    }
833
834    // ── GGUF loading tests ───────────────────────────────────────────────────
835
836    #[test]
837    fn from_bytes_rejects_invalid_magic() {
838        let r = CandleEngine::from_bytes(b"not a gguf file", 0);
839        assert!(matches!(r, Err(EdgeError::Engine(_))));
840    }
841
842    #[test]
843    fn from_bytes_loads_minimal_gguf_and_forward_has_correct_vocab() {
844        let vocab = 8;
845        let dim = 4;
846        let gguf = make_minimal_gguf(vocab, dim);
847        let mut engine = CandleEngine::from_bytes(&gguf, 7).unwrap();
848
849        let logits = engine.next_logits(&[0]);
850        assert_eq!(logits.len(), vocab, "logit vec width == vocab from GGUF");
851        assert_eq!(engine.eos_token(), 7);
852    }
853
854    #[test]
855    fn from_bytes_gguf_forward_is_deterministic() {
856        let gguf = make_minimal_gguf(8, 4);
857        let mut eng = CandleEngine::from_bytes(&gguf, 0).unwrap();
858        assert_eq!(eng.next_logits(&[3]), eng.next_logits(&[3]));
859    }
860
861    /// Same as `make_minimal_gguf` but `output.weight` has `wrong_dim` instead of `dim`,
862    /// so the embed / output dimensions are incompatible.
863    fn make_mismatched_gguf(vocab: usize, embed_dim: usize, output_dim: usize) -> Vec<u8> {
864        let mut w: Vec<u8> = Vec::new();
865        w.extend_from_slice(b"GGUF");
866        w.extend_from_slice(&3u32.to_le_bytes());
867        w.extend_from_slice(&2u64.to_le_bytes());
868        w.extend_from_slice(&0u64.to_le_bytes());
869
870        let embed_bytes = (vocab * embed_dim * 4) as u64;
871
872        let name = b"token_embd.weight";
873        w.extend_from_slice(&(name.len() as u64).to_le_bytes());
874        w.extend_from_slice(name);
875        w.extend_from_slice(&2u32.to_le_bytes());
876        w.extend_from_slice(&(embed_dim as u64).to_le_bytes());
877        w.extend_from_slice(&(vocab as u64).to_le_bytes());
878        w.extend_from_slice(&0u32.to_le_bytes());
879        w.extend_from_slice(&0u64.to_le_bytes());
880
881        let name = b"output.weight";
882        w.extend_from_slice(&(name.len() as u64).to_le_bytes());
883        w.extend_from_slice(name);
884        w.extend_from_slice(&2u32.to_le_bytes());
885        w.extend_from_slice(&(output_dim as u64).to_le_bytes()); // wrong dim
886        w.extend_from_slice(&(vocab as u64).to_le_bytes());
887        w.extend_from_slice(&0u32.to_le_bytes());
888        w.extend_from_slice(&embed_bytes.to_le_bytes());
889
890        let pad = (32usize.wrapping_sub(w.len() % 32)) % 32;
891        w.resize(w.len() + pad, 0u8);
892
893        for i in 0..(vocab * embed_dim + vocab * output_dim) {
894            w.extend_from_slice(&(i as f32 * 0.1f32).to_le_bytes());
895        }
896        w
897    }
898
899    #[test]
900    fn from_path_missing_file_returns_engine_error() {
901        let r = CandleEngine::from_path(std::path::Path::new("/nonexistent/model.gguf"), 0);
902        assert!(matches!(r, Err(EdgeError::Engine(_))));
903    }
904
905    #[test]
906    fn from_bytes_rejects_mismatched_output_dim_at_load_time() {
907        // embed dim=4, output dim=7 — incompatible; must error at load, not silently at forward.
908        let gguf = make_mismatched_gguf(8, 4, 7);
909        let r = CandleEngine::from_bytes(&gguf, 0);
910        assert!(
911            matches!(r, Err(EdgeError::Engine(_))),
912            "mismatched output weight dim must be rejected at load time"
913        );
914    }
915
916    // ── LocalLlmProvider tests (unchanged + new from_path error path) ────────
917
918    #[test]
919    fn local_provider_chat_returns_response() {
920        let p = LocalLlmProvider::toy(32, 8, 31, ok_permit()).unwrap();
921        let req = el_core::ChatRequest::new("local", vec![el_core::ChatMessage::user("hello")])
922            .with_max_tokens(4);
923        let resp = p.chat(&req).unwrap();
924        assert_eq!(resp.model, "local/candle");
925        assert_eq!(resp.completion_tokens, 4);
926        assert!(!resp.content.is_empty());
927    }
928
929    #[test]
930    fn local_provider_stream_ends_with_final_token() {
931        let p = LocalLlmProvider::toy(32, 8, 31, ok_permit()).unwrap();
932        let req = el_core::ChatRequest::new("local", vec![el_core::ChatMessage::user("hi")])
933            .with_max_tokens(3);
934        let mut tokens: Vec<el_core::ChatToken> = Vec::new();
935        p.chat_stream(&req, &mut |t| tokens.push(t)).unwrap();
936        assert!(tokens.last().unwrap().is_final);
937        assert!(tokens.len() > 1);
938    }
939
940    #[test]
941    fn local_provider_session_resets_between_calls() {
942        let p = LocalLlmProvider::toy(32, 8, 31, ok_permit()).unwrap();
943        let req = el_core::ChatRequest::new("local", vec![el_core::ChatMessage::user("a")])
944            .with_max_tokens(4);
945        let r1 = p.chat(&req).unwrap();
946        let r2 = p.chat(&req).unwrap();
947        assert_eq!(r1.content, r2.content);
948    }
949
950    #[test]
951    fn local_provider_from_path_missing_file_returns_error() {
952        let r = LocalLlmProvider::from_path(
953            std::path::Path::new("/nonexistent/model.gguf"),
954            0,
955            ok_permit(),
956        );
957        assert!(matches!(r, Err(EdgeError::Engine(_))));
958    }
959
960    // ── Qwen provider helpers ─────────────────────────────────────────────────
961
962    #[test]
963    fn render_chatml_wraps_each_turn_and_opens_assistant() {
964        let msgs = vec![
965            ChatMessage::system("be nice"),
966            ChatMessage::user("hi"),
967            ChatMessage::assistant("hello"),
968            ChatMessage::user("bye"),
969        ];
970        let got = render_chatml(&msgs);
971        let want = "<|im_start|>system\nbe nice<|im_end|>\n\
972                    <|im_start|>user\nhi<|im_end|>\n\
973                    <|im_start|>assistant\nhello<|im_end|>\n\
974                    <|im_start|>user\nbye<|im_end|>\n\
975                    <|im_start|>assistant\n";
976        assert_eq!(got, want);
977    }
978
979    #[test]
980    fn local_load_permit_passes_the_provenance_gate() {
981        // The runtime requires a LoadPermit; the local-trust path must yield one
982        // for a GGUF artifact (ADR-006 gate exercised, not bypassed).
983        let permit = local_load_permit().expect("local permit issued");
984        assert_eq!(permit.format, el_core::ModelFormat::Gguf);
985    }
986
987    #[test]
988    fn qwen_provider_from_paths_missing_model_errors() {
989        let r = QwenChatProvider::from_paths(
990            std::path::Path::new("/nonexistent/model.gguf"),
991            std::path::Path::new("/nonexistent/tokenizer.json"),
992        );
993        assert!(matches!(r, Err(EdgeError::Engine(_))));
994    }
995}