1#![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
25pub struct CandleEngine {
27 embed: Tensor,
28 w_out: Tensor,
29 vocab: usize,
30 eos: Token,
31}
32
33impl CandleEngine {
34 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 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 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 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 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 fn forward(&self, last: usize) -> candle_core::Result<Vec<f32>> {
147 let row = self.embed.narrow(0, last, 1)?; let logits = row.matmul(&self.w_out)?; 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
175pub struct LocalLlmProvider {
181 session: std::sync::Mutex<InferenceSession<CandleEngine>>,
182 vocab: usize,
183}
184
185impl LocalLlmProvider {
186 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 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
287use candle_transformers::models::quantized_qwen2::ModelWeights as Qwen2Weights;
296use el_core::{ModelId, ModelVersion};
297use el_provenance::{ModelArtifact, SignatureVerifier};
298use tokenizers::Tokenizer;
299
300mod 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 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 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 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
342pub struct QwenEngine {
350 model: Qwen2Weights,
351 device: Device,
352 index_pos: usize,
354 fed: usize,
356 last_logits: Vec<i32>,
358 vocab: usize,
359 eos: Token,
360}
361
362impl QwenEngine {
363 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 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 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
446pub 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 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 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 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(); 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 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#[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
681fn 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
700fn 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 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 fn make_minimal_gguf(vocab: usize, dim: usize) -> Vec<u8> {
749 let mut w: Vec<u8> = Vec::new();
750
751 w.extend_from_slice(b"GGUF");
753 w.extend_from_slice(&3u32.to_le_bytes()); w.extend_from_slice(&2u64.to_le_bytes()); w.extend_from_slice(&0u64.to_le_bytes()); let tensor_bytes = (vocab * dim * 4) as u64;
758
759 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()); w.extend_from_slice(&(vocab as u64).to_le_bytes()); w.extend_from_slice(&0u32.to_le_bytes()); w.extend_from_slice(&0u64.to_le_bytes()); 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()); let pad = (32usize.wrapping_sub(w.len() % 32)) % 32;
781 w.resize(w.len() + pad, 0u8);
782
783 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 #[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 #[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 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()); 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 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 #[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 #[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 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}