#![forbid(unsafe_code)]
use candle_core::{Device, Tensor};
use el_core::{
ChatMessage, ChatRequest, ChatResponse, ChatRole, ChatToken, DomainEvent, EdgeError,
LlmProvider, Result, SafetyMode, SessionConfig, SessionId, StopReason, Token,
};
use el_provenance::LoadPermit;
use el_runtime::{
AnchorGuard, ContrastiveSteerer, ExpertLogits, InferenceEngine, InferenceSession,
LightweightFilter, NoSafety, Ports, SafetyModeSelector, SafetySteerer,
};
pub struct CandleEngine {
embed: Tensor,
w_out: Tensor,
vocab: usize,
eos: Token,
}
impl CandleEngine {
pub fn toy(vocab: usize, dim: usize, eos: Token) -> Result<Self> {
let device = Device::Cpu;
let embed_data: Vec<f32> = (0..vocab * dim)
.map(|k| {
let (i, j) = (k / dim, k % dim);
(((i + j) % 7) as f32) * 0.1
})
.collect();
let wout_data: Vec<f32> = (0..dim * vocab)
.map(|k| {
let (a, b) = (k / vocab, k % vocab);
((((a * 31 + b * 17) % 13) as f32) * 0.1) - 0.6
})
.collect();
let embed = Tensor::from_vec(embed_data, (vocab, dim), &device)
.map_err(|_| EdgeError::Engine("candle: embed tensor build failed"))?;
let w_out = Tensor::from_vec(wout_data, (dim, vocab), &device)
.map_err(|_| EdgeError::Engine("candle: w_out tensor build failed"))?;
Ok(Self {
embed,
w_out,
vocab,
eos,
})
}
pub fn from_path(path: impl AsRef<std::path::Path>, eos: Token) -> Result<Self> {
let file = std::fs::File::open(path.as_ref())
.map_err(|_| EdgeError::Engine("model file not found or not readable"))?;
Self::load_gguf(&mut std::io::BufReader::new(file), eos)
}
pub fn from_bytes(data: &[u8], eos: Token) -> Result<Self> {
Self::load_gguf(&mut std::io::Cursor::new(data), eos)
}
fn load_gguf<R: std::io::Read + std::io::Seek>(reader: &mut R, eos: Token) -> Result<Self> {
use candle_core::quantized::gguf_file;
let content = gguf_file::Content::read(reader)
.map_err(|_| EdgeError::Engine("GGUF: invalid or unrecognised file"))?;
let device = Device::Cpu;
let embed = content
.tensor(reader, "token_embd.weight", &device)
.map_err(|_| EdgeError::Engine("GGUF: missing 'token_embd.weight'"))?
.dequantize(&device)
.map_err(|_| EdgeError::Engine("GGUF: cannot dequantize embed tensor"))?;
let (vocab, dim) = match embed.shape().dims() {
[v, d] => (*v, *d),
_ => return Err(EdgeError::Engine("GGUF: 'token_embd.weight' must be 2-D")),
};
let raw_w_q = match content.tensor(reader, "output.weight", &device) {
Ok(t) => t,
Err(_) => content
.tensor(reader, "lm_head.weight", &device)
.map_err(|_| {
EdgeError::Engine("GGUF: missing 'output.weight' / 'lm_head.weight'")
})?,
};
let raw_w = raw_w_q
.dequantize(&device)
.map_err(|_| EdgeError::Engine("GGUF: cannot dequantize output weight"))?;
let w_out = match raw_w.shape().dims() {
[v, _d] if *v == vocab => raw_w
.t()
.map_err(|_| EdgeError::Engine("GGUF: failed to transpose output weight"))?,
_ => raw_w,
};
match w_out.shape().dims() {
[d, v] if *d == dim && *v == vocab => {}
_ => return Err(EdgeError::Engine(
"GGUF: output weight shape incompatible with embed dim — expected [dim, vocab] after transpose",
)),
}
Ok(Self {
embed,
w_out,
vocab,
eos,
})
}
fn forward(&self, last: usize) -> candle_core::Result<Vec<f32>> {
let row = self.embed.narrow(0, last, 1)?; let logits = row.matmul(&self.w_out)?; Ok(logits.to_vec2::<f32>()?.remove(0))
}
}
impl InferenceEngine for CandleEngine {
fn prefill(&mut self, tokens: &[Token]) -> Result<u32> {
Ok(tokens.len() as u32)
}
fn next_logits(&mut self, committed: &[Token]) -> Vec<i32> {
let last = committed
.last()
.copied()
.unwrap_or(0)
.min(self.vocab as u32 - 1) as usize;
match self.forward(last) {
Ok(logits) => logits.iter().map(|x| (x * 1000.0).round() as i32).collect(),
Err(_) => vec![0; self.vocab],
}
}
fn eos_token(&self) -> Token {
self.eos
}
fn rollback(&mut self, _keep_committed: u32) -> Result<()> {
Ok(())
}
fn reset_cache(&mut self) -> Result<()> {
Ok(())
}
}
pub struct LocalLlmProvider {
session: std::sync::Mutex<InferenceSession<CandleEngine>>,
vocab: usize,
}
impl LocalLlmProvider {
pub fn from_path(
path: impl AsRef<std::path::Path>,
eos: Token,
permit: LoadPermit,
) -> Result<Self> {
let engine = CandleEngine::from_path(path, eos)?;
let vocab = engine.vocab;
let session = InferenceSession::new(SessionId(1), SessionConfig::default(), engine, permit);
Ok(Self {
session: std::sync::Mutex::new(session),
vocab,
})
}
pub fn toy(vocab: usize, dim: usize, eos: Token, permit: LoadPermit) -> Result<Self> {
let engine = CandleEngine::toy(vocab, dim, eos)?;
let session = InferenceSession::new(SessionId(1), SessionConfig::default(), engine, permit);
Ok(Self {
session: std::sync::Mutex::new(session),
vocab,
})
}
fn encode(&self, text: &str) -> Vec<Token> {
text.bytes()
.map(|b| (b as Token) % self.vocab as Token)
.collect()
}
fn decode(tokens: &[Token]) -> String {
tokens
.iter()
.map(|&t| {
let b = (t & 0xFF) as u8;
if b.is_ascii_graphic() || b == b' ' {
b as char
} else {
'?'
}
})
.collect()
}
fn format_messages(messages: &[ChatMessage]) -> String {
messages
.iter()
.map(|m| {
let role = match m.role {
ChatRole::System => "system",
ChatRole::User => "user",
ChatRole::Assistant => "assistant",
};
format!("{role}: {}", m.content)
})
.collect::<Vec<_>>()
.join("\n")
}
}
impl LlmProvider for LocalLlmProvider {
fn chat(&self, req: &ChatRequest) -> Result<ChatResponse> {
let prompt = Self::format_messages(&req.messages);
let prompt_tokens = self.encode(&prompt);
let prompt_len = prompt_tokens.len() as u32;
let max = req.max_tokens.unwrap_or(64);
let mut session = self.session.lock().unwrap();
session.reset()?;
let _ = session.drain_events(); let ports = Ports::permissive();
session.load_prompt(&ports, &prompt_tokens)?;
session.generate(&ports, max)?;
let output = session.output().to_vec();
let completion_len = output.len() as u32;
Ok(ChatResponse {
content: Self::decode(&output),
model: "local/candle".into(),
prompt_tokens: prompt_len,
completion_tokens: completion_len,
})
}
fn chat_stream(&self, req: &ChatRequest, on_token: &mut dyn FnMut(ChatToken)) -> Result<()> {
let resp = self.chat(req)?;
for ch in resp.content.chars() {
on_token(ChatToken {
text: ch.to_string(),
is_final: false,
});
}
on_token(ChatToken {
text: String::new(),
is_final: true,
});
Ok(())
}
}
use candle_transformers::models::quantized_qwen2::ModelWeights as Qwen2Weights;
use el_core::{ModelId, ModelVersion};
use el_provenance::{ModelArtifact, SignatureVerifier};
use tokenizers::Tokenizer;
mod bench {
use std::cell::Cell;
use std::sync::OnceLock;
use std::time::Duration;
static ENABLED: OnceLock<bool> = OnceLock::new();
pub fn enabled() -> bool {
*ENABLED.get_or_init(|| std::env::var_os("EL_BENCH").is_some())
}
thread_local! {
static FWD_TOTAL: Cell<Duration> = const { Cell::new(Duration::ZERO) };
static FWD_MODEL: Cell<Duration> = const { Cell::new(Duration::ZERO) };
static FWD_CALLS: Cell<u64> = const { Cell::new(0) };
}
pub fn record(total: Duration, model: Duration) {
FWD_TOTAL.with(|c| c.set(c.get() + total));
FWD_MODEL.with(|c| c.set(c.get() + model));
FWD_CALLS.with(|c| c.set(c.get() + 1));
}
pub fn take() -> (Duration, Duration, u64) {
(
FWD_TOTAL.replace(Duration::ZERO),
FWD_MODEL.replace(Duration::ZERO),
FWD_CALLS.replace(0),
)
}
}
pub struct QwenEngine {
model: Qwen2Weights,
device: Device,
index_pos: usize,
fed: usize,
prompt: Vec<Token>,
last_logits: Vec<i32>,
vocab: usize,
eos: Token,
cache_dirty: bool,
}
impl QwenEngine {
pub fn from_path(path: impl AsRef<std::path::Path>, eos: Token) -> Result<Self> {
use candle_core::quantized::gguf_file;
let mut file = std::fs::File::open(path.as_ref())
.map_err(|_| EdgeError::Engine("model file not found or not readable"))?;
let content = gguf_file::Content::read(&mut file)
.map_err(|_| EdgeError::Engine("GGUF: invalid or unrecognised file"))?;
let device = Device::Cpu;
let model = Qwen2Weights::from_gguf(content, &mut file, &device)
.map_err(|_| EdgeError::Engine("GGUF: failed to load Qwen2 weights"))?;
Ok(Self {
model,
device,
index_pos: 0,
fed: 0,
prompt: Vec::new(),
last_logits: Vec::new(),
vocab: 0,
eos,
cache_dirty: false,
})
}
fn forward_one(&mut self, token: Token) -> Result<Vec<i32>> {
self.cache_dirty = true;
let t_total = bench::enabled().then(std::time::Instant::now);
let input = Tensor::from_vec(vec![token], (1, 1), &self.device)
.map_err(|_| EdgeError::Engine("candle: input tensor build failed"))?;
let t_model = bench::enabled().then(std::time::Instant::now);
let logits = self
.model
.forward(&input, self.index_pos)
.map_err(|_| EdgeError::Engine("candle: Qwen2 forward failed"))?;
let model_dur = t_model.map(|t| t.elapsed()).unwrap_or_default();
self.index_pos += 1;
let row = logits
.squeeze(0)
.map_err(|_| EdgeError::Engine("candle: squeeze logits failed"))?;
let floats = row
.to_vec1::<f32>()
.map_err(|_| EdgeError::Engine("candle: logits to vec failed"))?;
let out: Vec<i32> = floats.iter().map(|x| (x * 1000.0).round() as i32).collect();
if let Some(t) = t_total {
bench::record(t.elapsed(), model_dur);
}
Ok(out)
}
}
impl InferenceEngine for QwenEngine {
fn prefill(&mut self, tokens: &[Token]) -> Result<u32> {
self.index_pos = 0;
self.fed = 0;
self.prompt = tokens.to_vec(); for &t in tokens {
self.last_logits = self.forward_one(t)?;
}
self.vocab = self.last_logits.len();
Ok(tokens.len() as u32)
}
fn next_logits(&mut self, committed: &[Token]) -> Vec<i32> {
while self.fed < committed.len() {
let t = committed[self.fed];
match self.forward_one(t) {
Ok(l) => self.last_logits = l,
Err(_) => return vec![0; self.vocab.max(1)],
}
self.fed += 1;
}
self.last_logits.clone()
}
fn eos_token(&self) -> Token {
self.eos
}
fn rollback(&mut self, _keep_committed: u32) -> Result<()> {
self.index_pos = 0;
self.fed = 0;
for i in 0..self.prompt.len() {
let t = self.prompt[i];
self.last_logits = self.forward_one(t)?;
}
Ok(())
}
fn reset_cache(&mut self) -> Result<()> {
if self.cache_dirty {
self.index_pos = 0;
self.forward_one(0)?;
self.cache_dirty = false;
}
self.index_pos = 0;
self.fed = 0;
self.prompt = Vec::new();
self.last_logits = Vec::new();
Ok(())
}
}
const DEFAULT_UNSAFE_WORDS: &[&str] = &[
"bomb",
"explosive",
"detonator",
"methamphetamine",
"ricin",
"anthrax",
"sarin",
"nerve agent",
];
const SAFETY_REFUSAL: &str = "I can't help with that request.";
const CONTRASTIVE_TOP_K: usize = 64;
fn word_to_patterns(tokenizer: &Tokenizer, word: &str) -> Vec<Vec<Token>> {
let mut out: Vec<Vec<Token>> = Vec::new();
for variant in [format!(" {word}"), word.to_string()] {
if let Ok(enc) = tokenizer.encode(variant, false) {
let seq = enc.get_ids().to_vec();
if !seq.is_empty() && !out.contains(&seq) {
out.push(seq);
}
}
}
out
}
#[derive(Debug, Clone)]
struct SafetyConfig {
mode: SafetyMode,
banned: Vec<Token>,
patterns: Vec<Vec<Token>>,
extra_guard_patterns: Vec<Vec<Token>>,
}
impl SafetyConfig {
fn lightweight(tokenizer: &Tokenizer) -> Self {
let mut banned = Vec::new();
let mut patterns = Vec::new();
for &word in DEFAULT_UNSAFE_WORDS {
for seq in word_to_patterns(tokenizer, word) {
if seq.len() == 1 && !banned.contains(&seq[0]) {
banned.push(seq[0]);
}
if !patterns.contains(&seq) {
patterns.push(seq);
}
}
}
Self {
mode: SafetyMode::Lightweight,
banned,
patterns,
extra_guard_patterns: Vec::new(),
}
}
fn ports(&self) -> Ports {
let mut ports = Ports::permissive();
if matches!(self.mode, SafetyMode::Off) {
return ports;
}
let steerer: Box<dyn SafetySteerer> = if self.banned.is_empty() {
Box::new(NoSafety)
} else {
Box::new(LightweightFilter::new(self.banned.clone()))
};
ports.safety = steerer;
let guard_patterns: Vec<Vec<Token>> = self
.patterns
.iter()
.chain(self.extra_guard_patterns.iter())
.cloned()
.collect();
if !guard_patterns.is_empty() {
ports.guard = Some(Box::new(AnchorGuard::hard(guard_patterns)));
}
if !self.patterns.is_empty() {
ports.ingress = Some(Box::new(AnchorGuard::hard(self.patterns.clone())));
}
ports
}
}
pub struct QwenExpert {
engine: std::cell::RefCell<QwenEngine>,
fed: std::cell::Cell<usize>,
_permit: LoadPermit,
}
impl QwenExpert {
pub fn from_path_primed(
path: impl AsRef<std::path::Path>,
eos: Token,
prompt: &[Token],
permit: LoadPermit,
) -> Result<Self> {
let mut engine = QwenEngine::from_path(path, eos)?;
engine.prefill(prompt)?;
Ok(Self {
engine: std::cell::RefCell::new(engine),
fed: std::cell::Cell::new(0),
_permit: permit,
})
}
}
impl ExpertLogits for QwenExpert {
fn logits(&self, committed: &[Token]) -> Vec<i32> {
let mut engine = self.engine.borrow_mut();
if committed.len() < self.fed.get() {
if engine.rollback(0).is_err() {
return Vec::new();
}
self.fed.set(0);
}
let out = engine.next_logits(committed);
self.fed.set(committed.len());
out
}
}
enum ChatSession {
Loaded(QwenEngine),
Active(InferenceSession<QwenEngine>),
Swapping,
}
pub struct QwenChatProvider {
tokenizer: Tokenizer,
permit: LoadPermit,
eos: Token,
default_max_tokens: u32,
model_label: String,
safety: SafetyConfig,
expert_model: Option<std::path::PathBuf>,
steer_alpha_milli: i32,
session: std::sync::Mutex<ChatSession>,
}
impl QwenChatProvider {
pub fn from_paths(
model_path: impl AsRef<std::path::Path>,
tokenizer_path: impl AsRef<std::path::Path>,
) -> Result<Self> {
let model_path = model_path.as_ref().to_path_buf();
if !model_path.exists() {
return Err(EdgeError::Engine("model file not found"));
}
let tokenizer = Tokenizer::from_file(tokenizer_path.as_ref())
.map_err(|_| EdgeError::Engine("failed to load tokenizer.json"))?;
let eos = tokenizer.token_to_id("<|im_end|>").unwrap_or(151_645);
let model_label = model_path
.file_stem()
.and_then(|s| s.to_str())
.map(|s| format!("local/{s}"))
.unwrap_or_else(|| "local/qwen2".to_string());
let safety = SafetyConfig::lightweight(&tokenizer);
let permit = local_load_permit(&model_path)?;
let engine = QwenEngine::from_path(&model_path, eos)?;
Ok(Self {
tokenizer,
permit,
eos,
default_max_tokens: 512,
model_label,
safety,
expert_model: None,
steer_alpha_milli: 1000,
session: std::sync::Mutex::new(ChatSession::Loaded(engine)),
})
}
pub fn with_safety(mut self, mode: SafetyMode) -> Self {
self.safety.mode = mode;
self
}
pub fn with_extra_guard_words<I, S>(mut self, words: I) -> Self
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
for word in words {
for seq in word_to_patterns(&self.tokenizer, word.as_ref()) {
if !self.safety.extra_guard_patterns.contains(&seq) {
self.safety.extra_guard_patterns.push(seq);
}
}
}
self
}
pub fn with_expert_model(mut self, path: impl AsRef<std::path::Path>) -> Self {
self.expert_model = Some(path.as_ref().to_path_buf());
self
}
pub fn with_steer_alpha(mut self, alpha_milli: i32) -> Self {
self.steer_alpha_milli = alpha_milli;
self
}
pub fn end_session(&self) -> Result<()> {
let mut cell = self
.session
.lock()
.map_err(|_| EdgeError::Engine("chat session mutex poisoned"))?;
if let ChatSession::Active(session) = &mut *cell {
session.close()?;
}
Ok(())
}
fn encode(&self, text: &str) -> Result<Vec<Token>> {
let enc = self
.tokenizer
.encode(text, false)
.map_err(|_| EdgeError::Engine("tokenizer encode failed"))?;
Ok(enc.get_ids().to_vec())
}
fn decode(&self, ids: &[Token]) -> Result<String> {
self.tokenizer
.decode(ids, true)
.map_err(|_| EdgeError::Engine("tokenizer decode failed"))
}
}
impl LlmProvider for QwenChatProvider {
fn chat(&self, req: &ChatRequest) -> Result<ChatResponse> {
let prompt = render_chatml(&req.messages);
let t_encode = bench::enabled().then(std::time::Instant::now);
let prompt_tokens = self.encode(&prompt)?;
let d_encode = t_encode.map(|t| t.elapsed()).unwrap_or_default();
let requested = requested_session_safety(self.safety.mode, self.expert_model.is_some());
let cfg = SessionConfig {
safety: requested,
..SessionConfig::default()
};
let effective = SafetyModeSelector::resolve(requested, cfg.device);
let mut cell = self
.session
.lock()
.map_err(|_| EdgeError::Engine("chat session mutex poisoned"))?;
let t_load = bench::enabled().then(std::time::Instant::now);
if matches!(&*cell, ChatSession::Loaded(_)) {
let engine = match std::mem::replace(&mut *cell, ChatSession::Swapping) {
ChatSession::Loaded(e) => e,
_ => unreachable!("guarded by the matches! above"),
};
*cell = ChatSession::Active(InferenceSession::new(
SessionId(1),
cfg,
engine,
self.permit,
));
}
let d_load = t_load.map(|t| t.elapsed()).unwrap_or_default();
let session = match &mut *cell {
ChatSession::Active(s) => s,
_ => return Err(EdgeError::Engine("chat session not initialized")),
};
session.reset()?;
let _ = session.drain_events();
let mut ports = self.safety.ports();
if matches!(effective, SafetyMode::SecDecoding) {
if let Some(expert_path) = &self.expert_model {
let expert = QwenExpert::from_path_primed(
expert_path,
self.eos,
&prompt_tokens,
local_load_permit(expert_path)?,
)?;
ports.safety = Box::new(ContrastiveSteerer::new(
expert,
self.safety.banned.clone(),
self.steer_alpha_milli,
CONTRASTIVE_TOP_K,
effective,
));
}
}
let _ = bench::take(); let t_prefill = bench::enabled().then(std::time::Instant::now);
session.load_prompt(&ports, &prompt_tokens)?;
let d_prefill = t_prefill.map(|t| t.elapsed()).unwrap_or_default();
let (pf_total, pf_model, pf_calls) = bench::take();
let max = req.max_tokens.unwrap_or(self.default_max_tokens);
let t_decode = bench::enabled().then(std::time::Instant::now);
let stop = session.generate(&ports, max)?;
let d_decode = t_decode.map(|t| t.elapsed()).unwrap_or_default();
let (dc_total, dc_model, dc_calls) = bench::take();
let out = session.output().to_vec();
let completion_tokens = out.len() as u32;
let t_detok = bench::enabled().then(std::time::Instant::now);
let decoded = self.decode(&out)?.trim().to_string();
let d_detok = t_detok.map(|t| t.elapsed()).unwrap_or_default();
let events = session.drain_events();
let safety_active = !matches!(self.safety.mode, SafetyMode::Off);
let content = if safety_active {
let violations = events
.iter()
.filter(|e| matches!(e.event, DomainEvent::SafetyViolationDetected { .. }))
.count();
let rollbacks = events
.iter()
.filter(|e| matches!(e.event, DomainEvent::ClaimBacktracked { .. }))
.count();
let refused = stop == StopReason::Stopped && violations > 0;
if violations > 0 || rollbacks > 0 {
eprintln!(
"[safety] {violations} violation(s), {rollbacks} rollback(s){}",
if refused {
" → refused (fail-closed)"
} else {
" → recovered"
}
);
}
if refused {
SAFETY_REFUSAL.to_string()
} else {
decoded
}
} else {
decoded
};
if bench::enabled() {
report_breakdown(
prompt_tokens.len() as u32,
completion_tokens,
d_load,
d_encode,
d_prefill,
d_decode,
d_detok,
(pf_total, pf_model, pf_calls),
(dc_total, dc_model, dc_calls),
);
}
Ok(ChatResponse {
content,
model: self.model_label.clone(),
prompt_tokens: prompt_tokens.len() as u32,
completion_tokens,
})
}
fn chat_stream(&self, req: &ChatRequest, on_token: &mut dyn FnMut(ChatToken)) -> Result<()> {
let resp = self.chat(req)?;
for ch in resp.content.chars() {
on_token(ChatToken {
text: ch.to_string(),
is_final: false,
});
}
on_token(ChatToken {
text: String::new(),
is_final: true,
});
Ok(())
}
}
#[allow(clippy::too_many_arguments)]
fn report_breakdown(
prompt_tokens: u32,
completion_tokens: u32,
d_load: std::time::Duration,
d_encode: std::time::Duration,
d_prefill: std::time::Duration,
d_decode: std::time::Duration,
d_detok: std::time::Duration,
prefill_fwd: (std::time::Duration, std::time::Duration, u64),
decode_fwd: (std::time::Duration, std::time::Duration, u64),
) {
let ms = |d: std::time::Duration| d.as_secs_f64() * 1000.0;
let total = d_load + d_encode + d_prefill + d_decode + d_detok;
let pct = |d: std::time::Duration| {
if total.as_secs_f64() > 0.0 {
d.as_secs_f64() / total.as_secs_f64() * 100.0
} else {
0.0
}
};
let tps = |n: u32, d: std::time::Duration| {
if d.as_secs_f64() > 0.0 {
n as f64 / d.as_secs_f64()
} else {
0.0
}
};
let (pf_total, pf_model, pf_calls) = prefill_fwd;
let (dc_total, dc_model, dc_calls) = decode_fwd;
let dc_loop = d_decode.saturating_sub(dc_total);
let dc_seam = dc_total.saturating_sub(dc_model);
let per_tok = |d: std::time::Duration, n: u64| if n > 0 { ms(d) / n as f64 } else { 0.0 };
eprintln!("\n┌─ EL_BENCH chat() breakdown ───────────────────────────────");
eprintln!("│ prompt_tokens={prompt_tokens} completion_tokens={completion_tokens}");
eprintln!("│ phase wall(ms) %total throughput");
eprintln!(
"│ session setup {:>9.1} {:>6.1}% (weights loaded once at startup — ADR-018)",
ms(d_load),
pct(d_load)
);
eprintln!(
"│ tokenize {:>9.2} {:>6.1}%",
ms(d_encode),
pct(d_encode)
);
eprintln!(
"│ prefill {:>9.1} {:>6.1}% {:>7.1} tok/s",
ms(d_prefill),
pct(d_prefill),
tps(prompt_tokens, d_prefill)
);
eprintln!(
"│ decode {:>9.1} {:>6.1}% {:>7.1} tok/s",
ms(d_decode),
pct(d_decode),
tps(completion_tokens, d_decode)
);
eprintln!(
"│ detokenize {:>9.2} {:>6.1}%",
ms(d_detok),
pct(d_detok)
);
eprintln!("│ TOTAL {:>9.1}", ms(total));
eprintln!("│ ─ forward attribution (where prefill+decode time goes) ─");
eprintln!(
"│ prefill: {} fwd calls, model {:.1}ms, seam {:.1}ms, loop {:.1}ms",
pf_calls,
ms(pf_model),
ms(pf_total.saturating_sub(pf_model)),
ms(d_prefill.saturating_sub(pf_total)),
);
eprintln!(
"│ decode : {} fwd calls, model {:.1}ms, seam {:.1}ms, loop {:.1}ms",
dc_calls,
ms(dc_model),
ms(dc_seam),
ms(dc_loop),
);
eprintln!(
"│ per decoded token: {:.2}ms total = model {:.2} + seam {:.2} + loop {:.2}",
per_tok(d_decode, dc_calls),
per_tok(dc_model, dc_calls),
per_tok(dc_seam, dc_calls),
per_tok(dc_loop, dc_calls),
);
eprintln!("└───────────────────────────────────────────────────────────");
}
fn render_chatml(messages: &[ChatMessage]) -> String {
let mut s = String::new();
for m in messages {
let role = match m.role {
ChatRole::System => "system",
ChatRole::User => "user",
ChatRole::Assistant => "assistant",
};
s.push_str("<|im_start|>");
s.push_str(role);
s.push('\n');
s.push_str(&m.content);
s.push_str("<|im_end|>\n");
}
s.push_str("<|im_start|>assistant\n");
s
}
fn requested_session_safety(configured: SafetyMode, has_expert: bool) -> SafetyMode {
match (configured, has_expert) {
(SafetyMode::Off, _) => SafetyMode::Off,
(_, true) => SafetyMode::SecDecoding,
(SafetyMode::SecDecoding | SafetyMode::Csd, false) => SafetyMode::Lightweight,
(mode, false) => mode,
}
}
fn local_load_permit(path: &std::path::Path) -> Result<LoadPermit> {
struct LocalFileTrust;
impl SignatureVerifier for LocalFileTrust {
fn verify(&self, _bytes: &[u8], _sig: &[u8], _key: u32) -> bool {
true
}
}
let _ = path;
let mut artifact = ModelArtifact::new(
ModelId(1),
ModelVersion::new(0, 1, 0),
el_core::ModelFormat::Gguf,
);
artifact.verify(&LocalFileTrust, b"local-trust", b"", 0);
artifact.ensure_loadable()
}
#[cfg(test)]
mod tests {
use super::*;
use el_runtime::InferenceEngine;
fn ok_permit() -> LoadPermit {
use el_core::{ModelFormat, ModelId, ModelVersion};
use el_provenance::{ModelArtifact, SignatureVerifier};
struct OkV;
impl SignatureVerifier for OkV {
fn verify(&self, _: &[u8], _: &[u8], _: u32) -> bool {
true
}
}
let mut a = ModelArtifact::new(ModelId(1), ModelVersion::new(0, 1, 0), ModelFormat::Gguf);
a.verify(&OkV, b"w", b"s", 0);
a.ensure_loadable().unwrap()
}
fn make_minimal_gguf(vocab: usize, dim: usize) -> Vec<u8> {
let mut w: Vec<u8> = Vec::new();
w.extend_from_slice(b"GGUF");
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;
let name = b"token_embd.weight";
w.extend_from_slice(&(name.len() as u64).to_le_bytes());
w.extend_from_slice(name);
w.extend_from_slice(&2u32.to_le_bytes());
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";
w.extend_from_slice(&(name.len() as u64).to_le_bytes());
w.extend_from_slice(name);
w.extend_from_slice(&2u32.to_le_bytes());
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(&tensor_bytes.to_le_bytes());
let pad = (32usize.wrapping_sub(w.len() % 32)) % 32;
w.resize(w.len() + pad, 0u8);
for i in 0..(vocab * dim * 2) {
w.extend_from_slice(&(i as f32 * 0.1f32).to_le_bytes());
}
w
}
#[test]
fn real_candle_forward_is_deterministic_and_right_shape() {
let mut eng = CandleEngine::toy(8, 4, 7).unwrap();
let a = eng.next_logits(&[2]);
let b = eng.next_logits(&[2]);
assert_eq!(a.len(), 8, "logits length == vocab");
assert_eq!(a, b, "fixed weights → deterministic real-tensor forward");
let c = eng.next_logits(&[5]);
assert_ne!(a, c);
}
#[test]
fn drives_the_runtime_end_to_end() {
use el_core::{ModelFormat, ModelId, ModelVersion, SessionConfig, SessionId, StopReason};
use el_provenance::{ModelArtifact, SignatureVerifier};
struct OkVerifier;
impl SignatureVerifier for OkVerifier {
fn verify(&self, _: &[u8], _: &[u8], _: u32) -> bool {
true
}
}
let mut art = ModelArtifact::new(
ModelId(1),
ModelVersion::new(0, 1, 0),
ModelFormat::Safetensors,
);
art.verify(&OkVerifier, b"w", b"s", 1);
let permit = art.ensure_loadable().unwrap();
let eng = CandleEngine::toy(16, 8, 9999).unwrap();
let mut session =
InferenceSession::new(SessionId(1), SessionConfig::default(), eng, permit);
let ports = Ports::permissive();
session.load_prompt(&ports, &[1, 2, 3]).unwrap();
let stop = session.generate(&ports, 4).unwrap();
assert_eq!(stop, StopReason::MaxTokens);
assert_eq!(session.output().len(), 4);
}
#[test]
fn from_bytes_rejects_invalid_magic() {
let r = CandleEngine::from_bytes(b"not a gguf file", 0);
assert!(matches!(r, Err(EdgeError::Engine(_))));
}
#[test]
fn from_bytes_loads_minimal_gguf_and_forward_has_correct_vocab() {
let vocab = 8;
let dim = 4;
let gguf = make_minimal_gguf(vocab, dim);
let mut engine = CandleEngine::from_bytes(&gguf, 7).unwrap();
let logits = engine.next_logits(&[0]);
assert_eq!(logits.len(), vocab, "logit vec width == vocab from GGUF");
assert_eq!(engine.eos_token(), 7);
}
#[test]
fn from_bytes_gguf_forward_is_deterministic() {
let gguf = make_minimal_gguf(8, 4);
let mut eng = CandleEngine::from_bytes(&gguf, 0).unwrap();
assert_eq!(eng.next_logits(&[3]), eng.next_logits(&[3]));
}
fn make_mismatched_gguf(vocab: usize, embed_dim: usize, output_dim: usize) -> Vec<u8> {
let mut w: Vec<u8> = Vec::new();
w.extend_from_slice(b"GGUF");
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 embed_bytes = (vocab * embed_dim * 4) as u64;
let name = b"token_embd.weight";
w.extend_from_slice(&(name.len() as u64).to_le_bytes());
w.extend_from_slice(name);
w.extend_from_slice(&2u32.to_le_bytes());
w.extend_from_slice(&(embed_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";
w.extend_from_slice(&(name.len() as u64).to_le_bytes());
w.extend_from_slice(name);
w.extend_from_slice(&2u32.to_le_bytes());
w.extend_from_slice(&(output_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(&embed_bytes.to_le_bytes());
let pad = (32usize.wrapping_sub(w.len() % 32)) % 32;
w.resize(w.len() + pad, 0u8);
for i in 0..(vocab * embed_dim + vocab * output_dim) {
w.extend_from_slice(&(i as f32 * 0.1f32).to_le_bytes());
}
w
}
#[test]
fn from_path_missing_file_returns_engine_error() {
let r = CandleEngine::from_path(std::path::Path::new("/nonexistent/model.gguf"), 0);
assert!(matches!(r, Err(EdgeError::Engine(_))));
}
#[test]
fn from_bytes_rejects_mismatched_output_dim_at_load_time() {
let gguf = make_mismatched_gguf(8, 4, 7);
let r = CandleEngine::from_bytes(&gguf, 0);
assert!(
matches!(r, Err(EdgeError::Engine(_))),
"mismatched output weight dim must be rejected at load time"
);
}
#[test]
fn local_provider_chat_returns_response() {
let p = LocalLlmProvider::toy(32, 8, 31, ok_permit()).unwrap();
let req = el_core::ChatRequest::new("local", vec![el_core::ChatMessage::user("hello")])
.with_max_tokens(4);
let resp = p.chat(&req).unwrap();
assert_eq!(resp.model, "local/candle");
assert_eq!(resp.completion_tokens, 4);
assert!(!resp.content.is_empty());
}
#[test]
fn local_provider_stream_ends_with_final_token() {
let p = LocalLlmProvider::toy(32, 8, 31, ok_permit()).unwrap();
let req = el_core::ChatRequest::new("local", vec![el_core::ChatMessage::user("hi")])
.with_max_tokens(3);
let mut tokens: Vec<el_core::ChatToken> = Vec::new();
p.chat_stream(&req, &mut |t| tokens.push(t)).unwrap();
assert!(tokens.last().unwrap().is_final);
assert!(tokens.len() > 1);
}
#[test]
fn local_provider_session_resets_between_calls() {
let p = LocalLlmProvider::toy(32, 8, 31, ok_permit()).unwrap();
let req = el_core::ChatRequest::new("local", vec![el_core::ChatMessage::user("a")])
.with_max_tokens(4);
let r1 = p.chat(&req).unwrap();
let r2 = p.chat(&req).unwrap();
assert_eq!(r1.content, r2.content);
}
#[test]
fn local_provider_from_path_missing_file_returns_error() {
let r = LocalLlmProvider::from_path(
std::path::Path::new("/nonexistent/model.gguf"),
0,
ok_permit(),
);
assert!(matches!(r, Err(EdgeError::Engine(_))));
}
#[test]
fn render_chatml_wraps_each_turn_and_opens_assistant() {
let msgs = vec![
ChatMessage::system("be nice"),
ChatMessage::user("hi"),
ChatMessage::assistant("hello"),
ChatMessage::user("bye"),
];
let got = render_chatml(&msgs);
let want = "<|im_start|>system\nbe nice<|im_end|>\n\
<|im_start|>user\nhi<|im_end|>\n\
<|im_start|>assistant\nhello<|im_end|>\n\
<|im_start|>user\nbye<|im_end|>\n\
<|im_start|>assistant\n";
assert_eq!(got, want);
}
#[test]
fn local_load_permit_passes_the_provenance_gate() {
let permit = local_load_permit(std::path::Path::new("models/qwen.gguf"))
.expect("local permit issued");
assert_eq!(permit.format, el_core::ModelFormat::Gguf);
}
#[test]
fn requested_safety_matches_the_backed_steerer_surface() {
assert_eq!(
requested_session_safety(SafetyMode::Off, true),
SafetyMode::Off,
"Off must stay off even if an expert path is configured"
);
assert_eq!(
requested_session_safety(SafetyMode::Lightweight, true),
SafetyMode::SecDecoding,
"an expert promotes the concrete model-backed path"
);
assert_eq!(
requested_session_safety(SafetyMode::SecDecoding, false),
SafetyMode::Lightweight,
"unbacked SecDecoding must not be reported as active"
);
assert_eq!(
requested_session_safety(SafetyMode::Csd, false),
SafetyMode::Lightweight,
"unbacked Csd must not be reported as active"
);
}
#[test]
fn qwen_provider_from_paths_missing_model_errors() {
let r = QwenChatProvider::from_paths(
std::path::Path::new("/nonexistent/model.gguf"),
std::path::Path::new("/nonexistent/tokenizer.json"),
);
assert!(matches!(r, Err(EdgeError::Engine(_))));
}
#[test]
fn safety_off_wires_no_guard_or_steering() {
let cfg = SafetyConfig {
mode: SafetyMode::Off,
banned: vec![1],
patterns: vec![vec![2]],
extra_guard_patterns: vec![],
};
let ports = cfg.ports();
assert!(ports.guard.is_none(), "Off must not wire the chunk guard");
assert!(ports.ingress.is_none(), "Off must not wire ingress triage");
assert_eq!(
ports.safety.mode(),
SafetyMode::Off,
"Off must keep the no-op steerer"
);
}
#[test]
fn lightweight_wires_guard_and_hard_ban_steerer() {
let cfg = SafetyConfig {
mode: SafetyMode::Lightweight,
banned: vec![1],
patterns: vec![vec![2, 3]],
extra_guard_patterns: vec![],
};
let ports = cfg.ports();
assert!(
ports.guard.is_some(),
"Lightweight must wire the chunk guard"
);
assert!(
ports.ingress.is_some(),
"Lightweight must wire prompt ingress triage (ADR-013)"
);
assert_eq!(
ports.safety.mode(),
SafetyMode::Lightweight,
"a non-empty ban list selects the LightweightFilter steerer"
);
}
#[test]
fn lightweight_without_patterns_has_no_guard_or_ingress() {
let cfg = SafetyConfig {
mode: SafetyMode::Lightweight,
banned: vec![7],
patterns: vec![],
extra_guard_patterns: vec![],
};
let ports = cfg.ports();
assert!(ports.guard.is_none());
assert!(ports.ingress.is_none());
}
#[test]
fn extra_guard_words_drive_guard_but_not_ingress() {
let cfg = SafetyConfig {
mode: SafetyMode::Lightweight,
banned: vec![],
patterns: vec![], extra_guard_patterns: vec![vec![42]], };
let ports = cfg.ports();
assert!(
ports.guard.is_some(),
"extra guard words must drive the output guard"
);
assert!(
ports.ingress.is_none(),
"extra guard words must NOT drive ingress (trajectory demo, not refusal)"
);
}
#[test]
fn qwen_expert_missing_file_errors_and_is_permit_gated() {
let r = QwenExpert::from_path_primed(
std::path::Path::new("/nonexistent/expert.gguf"),
0,
&[1, 2],
ok_permit(),
);
assert!(matches!(r, Err(EdgeError::Engine(_))));
}
}