use crate::agent::provider::LlmProvider;
use crate::agent::types::{Block, Msg, Role, Stop, ToolSpec, Turn};
use async_trait::async_trait;
use candle_core::quantized::{gguf_file, GgmlDType, QTensor};
use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder;
use candle_transformers::generation::LogitsProcessor;
use candle_transformers::models::qwen3::{Config, ModelForCausalLM};
use candle_transformers::models::quantized_qwen3::ModelWeights as QModelWeights;
use candle_transformers::utils::apply_repeat_penalty;
use serde_json::json;
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex};
use tokenizers::Tokenizer;
pub const GGUF_NAME: &str = "steeldb-qwen3-1.7b-q4k.gguf";
#[derive(Clone, Copy, Debug)]
pub struct GenConfig {
pub max_tokens: usize,
pub temperature: f64,
pub top_p: f64,
pub repeat_penalty: f32,
pub repeat_last_n: usize,
pub seed: u64,
}
impl Default for GenConfig {
fn default() -> Self {
GenConfig { max_tokens: 1024, temperature: 0.2, top_p: 0.9, repeat_penalty: 1.15, repeat_last_n: 128, seed: 0 }
}
}
fn best_device() -> Device {
#[cfg(target_os = "macos")]
if std::env::var("STEELDB_NATIVE_METAL").as_deref() == Ok("1") {
if let Ok(d) = Device::new_metal(0) {
return d;
}
}
Device::Cpu
}
enum Backend {
Full(ModelForCausalLM),
Quant(QModelWeights),
}
impl Backend {
fn forward(&mut self, input: &Tensor, offset: usize) -> candle_core::Result<Tensor> {
match self {
Backend::Full(m) => m.forward(input, offset),
Backend::Quant(m) => m.forward(input, offset),
}
}
fn clear_kv_cache(&mut self) {
match self {
Backend::Full(m) => m.clear_kv_cache(),
Backend::Quant(m) => m.clear_kv_cache(),
}
}
}
pub struct NativeLlm {
model: Backend,
tokenizer: Tokenizer,
device: Device,
eos: Vec<u32>,
gen: GenConfig,
}
impl NativeLlm {
pub fn load(base_dir: &Path, lora_dir: Option<&Path>) -> Result<NativeLlm, String> {
let tokenizer = Tokenizer::from_file(base_dir.join("tokenizer.json")).map_err(|e| format!("tokenizer: {e}"))?;
if std::env::var("STEELDB_NATIVE_FP16").as_deref() == Ok("1") {
return Self::load_fp16(base_dir, lora_dir, tokenizer);
}
let gguf = crate::paths::model_roots()
.into_iter()
.map(|r| r.join(GGUF_NAME))
.find(|p| p.exists())
.unwrap_or_else(|| default_gguf_path());
if !gguf.exists() {
eprintln!("native: building quantized model (one-time) → {}", gguf.display());
build_gguf(base_dir, lora_dir, &gguf).map_err(|e| format!("build quantized GGUF: {e}"))?;
}
Self::load_gguf(&gguf, tokenizer)
}
fn load_fp16(base_dir: &Path, lora_dir: Option<&Path>, tokenizer: Tokenizer) -> Result<NativeLlm, String> {
let device = best_device();
let cfg = read_config(base_dir).map_err(|e| e.to_string())?;
let mut tensors = load_base_tensors(base_dir, &device).map_err(|e| e.to_string())?;
if let Some(dir) = lora_dir {
merge_lora(&mut tensors, dir, &device).map_err(|e| format!("merge LoRA: {e}"))?;
}
let dtype = if device.is_cpu() { DType::F32 } else { DType::BF16 };
let vb = VarBuilder::from_tensors(tensors, dtype, &device);
let model = ModelForCausalLM::new(&cfg, vb).map_err(|e| format!("build model: {e}"))?;
Ok(NativeLlm { model: Backend::Full(model), tokenizer, device, eos: vec![151645, 151643], gen: GenConfig::default() })
}
fn load_gguf(gguf: &Path, tokenizer: Tokenizer) -> Result<NativeLlm, String> {
let device = best_device();
let mut f = std::fs::File::open(gguf).map_err(|e| format!("open {}: {e}", gguf.display()))?;
let content = gguf_file::Content::read(&mut f).map_err(|e| format!("read gguf: {e}"))?;
let model = QModelWeights::from_gguf(content, &mut f, &device).map_err(|e| format!("load gguf model: {e}"))?;
Ok(NativeLlm { model: Backend::Quant(model), tokenizer, device, eos: vec![151645, 151643], gen: GenConfig::default() })
}
pub fn generate(&mut self, prompt: &str) -> Result<String, String> {
self.model.clear_kv_cache();
let enc = self.tokenizer.encode(prompt, false).map_err(|e| format!("encode: {e}"))?;
let mut tokens: Vec<u32> = enc.get_ids().to_vec();
let mut logits_proc = if self.gen.temperature <= 0.0 {
LogitsProcessor::new(self.gen.seed, None, None)
} else {
LogitsProcessor::new(self.gen.seed, Some(self.gen.temperature), Some(self.gen.top_p))
};
let mut out_tokens: Vec<u32> = Vec::new();
let mut offset = 0usize;
for step in 0..self.gen.max_tokens {
let ctx = if step == 0 { &tokens[..] } else { &tokens[tokens.len() - 1..] };
let input = Tensor::new(ctx, &self.device).map_err(|e| e.to_string())?.unsqueeze(0).map_err(|e| e.to_string())?;
let logits = self.model.forward(&input, offset).map_err(|e| format!("forward: {e}"))?;
let logits = logits.flatten_all().map_err(|e| e.to_string())?.to_dtype(DType::F32).map_err(|e| e.to_string())?;
let logits = if self.gen.repeat_penalty != 1.0 {
let start = tokens.len().saturating_sub(self.gen.repeat_last_n);
apply_repeat_penalty(&logits, self.gen.repeat_penalty, &tokens[start..]).map_err(|e| e.to_string())?
} else {
logits
};
let next = logits_proc.sample(&logits).map_err(|e| format!("sample: {e}"))?;
offset += ctx.len();
if self.eos.contains(&next) {
break;
}
tokens.push(next);
out_tokens.push(next);
}
let text = self.tokenizer.decode(&out_tokens, true).map_err(|e| format!("decode: {e}"))?;
Ok(strip_think(&text))
}
}
fn strip_think(s: &str) -> String {
let t = s.trim_start();
if let Some(rest) = t.strip_prefix("<think>") {
if let Some(end) = rest.find("</think>") {
return rest[end + "</think>".len()..].trim_start().to_string();
}
}
s.trim().to_string()
}
fn default_gguf_path() -> PathBuf {
if let Ok(home) = std::env::var("HOME") {
let d = PathBuf::from(home).join(".steeldb/models");
let _ = std::fs::create_dir_all(&d);
return d.join(GGUF_NAME);
}
PathBuf::from(GGUF_NAME)
}
fn read_config(base_dir: &Path) -> candle_core::Result<Config> {
let bytes = std::fs::read(base_dir.join("config.json")).map_err(candle_core::Error::wrap)?;
serde_json::from_slice(&bytes).map_err(candle_core::Error::wrap)
}
fn load_base_tensors(base_dir: &Path, device: &Device) -> candle_core::Result<HashMap<String, Tensor>> {
let mut shards: Vec<PathBuf> = std::fs::read_dir(base_dir)
.map_err(candle_core::Error::wrap)?
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| p.extension().map(|x| x == "safetensors").unwrap_or(false))
.collect();
shards.sort();
if shards.is_empty() {
candle_core::bail!("no *.safetensors in {}", base_dir.display());
}
let mut tensors = HashMap::new();
for shard in &shards {
tensors.extend(candle_core::safetensors::load(shard, device)?);
}
Ok(tensors)
}
fn merge_lora(tensors: &mut HashMap<String, Tensor>, lora_dir: &Path, device: &Device) -> candle_core::Result<()> {
let cfg: serde_json::Value =
serde_json::from_slice(&std::fs::read(lora_dir.join("adapter_config.json")).map_err(candle_core::Error::wrap)?)
.map_err(candle_core::Error::wrap)?;
let r = cfg.get("r").and_then(|v| v.as_f64()).unwrap_or(16.0);
let alpha = cfg.get("lora_alpha").and_then(|v| v.as_f64()).unwrap_or(r);
let scale = alpha / r;
let lora = candle_core::safetensors::load(lora_dir.join("adapter_model.safetensors"), device)?;
let mut targets: Vec<String> =
lora.keys().filter_map(|k| k.strip_suffix(".lora_A.weight").map(|s| s.to_string())).collect();
targets.sort();
let mut merged = 0usize;
for t in &targets {
let base_key = format!("{}.weight", t.strip_prefix("base_model.model.").unwrap_or(t));
let (Some(a), Some(b)) = (lora.get(&format!("{t}.lora_A.weight")), lora.get(&format!("{t}.lora_B.weight"))) else {
continue;
};
let Some(w) = tensors.get(&base_key) else { continue };
let delta = (b.to_dtype(DType::F32)?.matmul(&a.to_dtype(DType::F32)?)? * scale)?;
let w2 = (w.to_dtype(DType::F32)? + delta)?;
tensors.insert(base_key, w2);
merged += 1;
}
if merged == 0 {
candle_core::bail!("LoRA merged 0 tensors — key mismatch with base model");
}
Ok(())
}
pub fn build_gguf(base_dir: &Path, lora_dir: Option<&Path>, out: &Path) -> candle_core::Result<()> {
let device = Device::Cpu; let cfg = read_config(base_dir)?;
let mut t = load_base_tensors(base_dir, &device)?;
if let Some(dir) = lora_dir {
merge_lora(&mut t, dir, &device)?;
}
let f32 = |t: &HashMap<String, Tensor>, k: &str| -> candle_core::Result<Tensor> {
t.get(k).ok_or_else(|| candle_core::Error::Msg(format!("missing tensor {k}")))?.to_dtype(DType::F32)
};
let mut q: Vec<(String, QTensor)> = Vec::new();
let mut push = |name: String, src: &Tensor, dtype: GgmlDType| -> candle_core::Result<()> {
q.push((name, QTensor::quantize(src, dtype)?));
Ok(())
};
push("token_embd.weight".into(), &f32(&t, "model.embed_tokens.weight")?, GgmlDType::Q4K)?;
push("output_norm.weight".into(), &f32(&t, "model.norm.weight")?, GgmlDType::F16)?;
for i in 0..cfg.num_hidden_layers {
let p = format!("model.layers.{i}");
let g = format!("blk.{i}");
push(format!("{g}.attn_q.weight"), &f32(&t, &format!("{p}.self_attn.q_proj.weight"))?, GgmlDType::Q4K)?;
push(format!("{g}.attn_k.weight"), &f32(&t, &format!("{p}.self_attn.k_proj.weight"))?, GgmlDType::Q4K)?;
push(format!("{g}.attn_v.weight"), &f32(&t, &format!("{p}.self_attn.v_proj.weight"))?, GgmlDType::Q4K)?;
push(format!("{g}.attn_output.weight"), &f32(&t, &format!("{p}.self_attn.o_proj.weight"))?, GgmlDType::Q4K)?;
push(format!("{g}.attn_q_norm.weight"), &f32(&t, &format!("{p}.self_attn.q_norm.weight"))?, GgmlDType::F16)?;
push(format!("{g}.attn_k_norm.weight"), &f32(&t, &format!("{p}.self_attn.k_norm.weight"))?, GgmlDType::F16)?;
push(format!("{g}.ffn_gate.weight"), &f32(&t, &format!("{p}.mlp.gate_proj.weight"))?, GgmlDType::Q4K)?;
push(format!("{g}.ffn_up.weight"), &f32(&t, &format!("{p}.mlp.up_proj.weight"))?, GgmlDType::Q4K)?;
push(format!("{g}.ffn_down.weight"), &f32(&t, &format!("{p}.mlp.down_proj.weight"))?, GgmlDType::Q4K)?;
push(format!("{g}.attn_norm.weight"), &f32(&t, &format!("{p}.input_layernorm.weight"))?, GgmlDType::F16)?;
push(format!("{g}.ffn_norm.weight"), &f32(&t, &format!("{p}.post_attention_layernorm.weight"))?, GgmlDType::F16)?;
}
let md: Vec<(&str, gguf_file::Value)> = vec![
("general.architecture", gguf_file::Value::String("qwen3".into())),
("qwen3.block_count", gguf_file::Value::U32(cfg.num_hidden_layers as u32)),
("qwen3.context_length", gguf_file::Value::U32(cfg.max_position_embeddings as u32)),
("qwen3.embedding_length", gguf_file::Value::U32(cfg.hidden_size as u32)),
("qwen3.attention.head_count", gguf_file::Value::U32(cfg.num_attention_heads as u32)),
("qwen3.attention.head_count_kv", gguf_file::Value::U32(cfg.num_key_value_heads as u32)),
("qwen3.attention.key_length", gguf_file::Value::U32(cfg.head_dim as u32)),
("qwen3.attention.layer_norm_rms_epsilon", gguf_file::Value::F32(cfg.rms_norm_eps as f32)),
("qwen3.rope.freq_base", gguf_file::Value::F32(cfg.rope_theta as f32)),
];
let md_ref: Vec<(&str, &gguf_file::Value)> = md.iter().map(|(k, v)| (*k, v)).collect();
let tensors_ref: Vec<(&str, &QTensor)> = q.iter().map(|(n, t)| (n.as_str(), t)).collect();
let tmp = out.with_extension("gguf.tmp");
{
let mut w = std::fs::File::create(&tmp).map_err(candle_core::Error::wrap)?;
gguf_file::write(&mut w, &md_ref, &tensors_ref)?;
}
std::fs::rename(&tmp, out).map_err(candle_core::Error::wrap)?;
Ok(())
}
fn to_chatml(system: &str, msgs: &[Msg], tools: &[ToolSpec]) -> String {
let mut p = String::from("<|im_start|>system\n");
p.push_str(system);
if !tools.is_empty() {
p.push_str(
"\n\n# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>\n",
);
for t in tools {
let sig = json!({"type": "function", "function": {"name": t.name, "description": t.description, "parameters": t.schema}});
p.push_str(&sig.to_string());
p.push('\n');
}
p.push_str(
"</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call>",
);
}
p.push_str("<|im_end|>\n");
for m in msgs {
match m.role {
Role::User => {
let mut text = String::new();
for b in &m.blocks {
match b {
Block::Text(t) => {
if !text.is_empty() {
text.push('\n');
}
text.push_str(t);
}
Block::ToolResult { content, .. } => {
p.push_str("<|im_start|>user\n<tool_response>\n");
p.push_str(content);
p.push_str("\n</tool_response><|im_end|>\n");
}
_ => {}
}
}
if !text.is_empty() {
p.push_str("<|im_start|>user\n");
p.push_str(&text);
p.push_str("<|im_end|>\n");
}
}
Role::Assistant => {
p.push_str("<|im_start|>assistant\n");
for b in &m.blocks {
match b {
Block::Text(t) => p.push_str(t),
Block::ToolUse { name, input, .. } => {
let call = json!({"name": name, "arguments": input});
p.push_str("\n<tool_call>\n");
p.push_str(&call.to_string());
p.push_str("\n</tool_call>");
}
_ => {}
}
}
p.push_str("<|im_end|>\n");
}
}
}
p.push_str("<|im_start|>assistant\n");
p
}
pub struct NativeProvider {
llm: Arc<Mutex<NativeLlm>>,
}
impl NativeProvider {
pub fn new(llm: NativeLlm) -> NativeProvider {
NativeProvider { llm: Arc::new(Mutex::new(llm)) }
}
}
#[async_trait]
impl LlmProvider for NativeProvider {
fn name(&self) -> &str {
"native"
}
fn synthesizes(&self) -> bool {
false
}
async fn chat(&self, system: &str, msgs: &[Msg], tools: &[ToolSpec]) -> Result<Turn, String> {
let prompt = to_chatml(system, msgs, tools);
let llm = self.llm.clone();
let text = tokio::task::spawn_blocking(move || {
let mut guard = llm.lock().map_err(|_| "native model mutex poisoned".to_string())?;
guard.generate(&prompt)
})
.await
.map_err(|e| format!("native inference task: {e}"))??;
Ok(Turn { text, tool_uses: Vec::new(), stop: Stop::EndTurn })
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn chatml_renders_qwen3_tool_template() {
let tools = vec![ToolSpec {
name: "ikl_query".into(),
description: "run a token query".into(),
schema: json!({"type": "object", "properties": {"q": {"type": "string"}}}),
}];
let msgs = vec![
Msg::user_text("how many electric vehicles?"),
Msg { role: Role::Assistant, blocks: vec![Block::ToolUse { id: "1".into(), name: "ikl_query".into(), input: json!({"q": "powertrain/electric"}) }] },
Msg { role: Role::User, blocks: vec![Block::ToolResult { id: "1".into(), content: "count=5".into(), is_error: false }] },
];
let p = to_chatml("You are SteelDB.", &msgs, &tools);
assert!(p.starts_with("<|im_start|>system\nYou are SteelDB."));
assert!(p.contains("<tools>") && p.contains("\"name\":\"ikl_query\""));
assert!(p.contains("<tool_call>") && p.contains("\"name\":\"ikl_query\""));
assert!(p.contains("<tool_response>\ncount=5\n</tool_response>"));
assert!(p.trim_end().ends_with("<|im_start|>assistant"));
}
#[test]
fn strip_think_removes_leading_block() {
assert_eq!(strip_think("<think>\n\n</think>\n\nThe answer is 38m."), "The answer is 38m.");
assert_eq!(strip_think("<think>reasoning here</think>done"), "done");
assert_eq!(strip_think("no think block"), "no think block");
}
#[test]
fn chatml_without_tools_has_no_tools_block() {
let p = to_chatml("sys", &[Msg::user_text("hi")], &[]);
assert!(!p.contains("<tools>"));
assert!(p.contains("<|im_start|>user\nhi<|im_end|>"));
}
}