use std::path::Path;
use std::sync::Arc;
use std::sync::mpsc::Receiver;
use anyhow::{Context, Result};
use kopitiam_runtime::{GenerationConfig, QwenModel, generate, tokenizer_from_gguf};
use kopitiam_tokenizer::BpeTokenizer;
use super::chat_template::render_chatml;
use super::generation::{resolve_eos_token_id, resolve_max_new_tokens};
use crate::{CompletionRequest, CompletionResponse, ModelAdapter, StreamChunk};
const IM_END: &str = "<|im_end|>";
const ENDOFTEXT: &str = "<|endoftext|>";
const GGUF_EOS_METADATA_KEY: &str = "tokenizer.ggml.eos_token_id";
pub struct LocalAdapter {
model: Arc<QwenModel>,
tokenizer: Arc<BpeTokenizer>,
model_name: String,
eos_token_id: Option<u32>,
}
impl LocalAdapter {
pub fn load(path: impl AsRef<Path>) -> Result<Self> {
let path = path.as_ref();
let loaded = kopitiam_loader::load_model(path)
.with_context(|| format!("loading model file {}", path.display()))?;
let model = QwenModel::from_loaded_model(&loaded)
.with_context(|| format!("building a Qwen model from {}", path.display()))?;
let tokenizer = tokenizer_from_gguf(&loaded).with_context(|| {
format!("building a tokenizer from {}'s embedded GGUF vocabulary", path.display())
})?;
let model_name = loaded
.metadata()
.name
.clone()
.or_else(|| loaded.metadata().architecture.clone())
.unwrap_or_else(|| "local-gguf-model".to_string());
let gguf_eos_metadata = loaded.metadata().raw.get_u32(GGUF_EOS_METADATA_KEY);
let eos_token_id = resolve_eos_token_id(
tokenizer.special_token_id(IM_END),
gguf_eos_metadata,
tokenizer.special_token_id(ENDOFTEXT),
);
Ok(Self { model: Arc::new(model), tokenizer: Arc::new(tokenizer), model_name, eos_token_id })
}
}
impl ModelAdapter for LocalAdapter {
fn name(&self) -> &str {
"local-qwen"
}
fn complete(&self, request: &CompletionRequest) -> Result<CompletionResponse> {
let prompt = render_chatml(&request.messages);
let config = GenerationConfig {
max_new_tokens: resolve_max_new_tokens(request.max_tokens),
eos_token_id: self.eos_token_id,
};
let content = generate(&*self.model, &*self.tokenizer, &prompt, &config, |_id, _text| {})
.context("local Qwen generation failed")?;
Ok(CompletionResponse { content, model: self.model_name.clone() })
}
fn stream(&self, request: &CompletionRequest) -> Receiver<StreamChunk> {
let (tx, rx) = std::sync::mpsc::channel();
let prompt = render_chatml(&request.messages);
let config = GenerationConfig {
max_new_tokens: resolve_max_new_tokens(request.max_tokens),
eos_token_id: self.eos_token_id,
};
let model = Arc::clone(&self.model);
let tokenizer = Arc::clone(&self.tokenizer);
std::thread::spawn(move || {
let result = generate(&*model, &*tokenizer, &prompt, &config, |_id, text| {
let _ = tx.send(StreamChunk::Token(text.to_string()));
});
match result {
Ok(_) => {
let _ = tx.send(StreamChunk::Done);
}
Err(error) => {
let _ = tx.send(StreamChunk::Error(format!("{error:#}")));
}
}
});
rx
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Message;
use crate::local::test_support::synthetic_gguf::{build_local_adapter_fixture, write_temp_gguf};
#[test]
fn load_on_a_nonexistent_path_returns_err_not_a_panic() {
let result = LocalAdapter::load("/does/not/exist/kopitiam-nonexistent.gguf");
assert!(result.is_err());
}
#[test]
fn load_on_a_non_gguf_file_returns_err_not_a_panic() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("not-a-model.gguf");
std::fs::write(&path, b"this is plainly not a GGUF or SafeTensors file").unwrap();
let result = LocalAdapter::load(&path);
assert!(result.is_err());
}
#[test]
fn end_to_end_against_a_synthetic_gguf_with_no_real_weights() {
let bytes = build_local_adapter_fixture();
let path = write_temp_gguf(&bytes, "local-adapter-e2e");
let adapter = LocalAdapter::load(&path).expect("synthetic GGUF must load");
assert_eq!(adapter.name(), "local-qwen");
let request = CompletionRequest::new([
Message::system("you are a test fixture"),
Message::user("hello"),
])
.with_max_tokens(5);
let response = adapter.complete(&request).expect("generation against synthetic weights must not error");
assert_eq!(response.model, "kopitiam-test-qwen");
}
#[test]
fn max_tokens_bounds_the_completion_length() {
let bytes = build_local_adapter_fixture();
let path = write_temp_gguf(&bytes, "local-adapter-max-tokens");
let adapter = LocalAdapter::load(&path).unwrap();
let short = adapter
.complete(&CompletionRequest::new([Message::user("hi")]).with_max_tokens(1))
.unwrap();
let long = adapter
.complete(&CompletionRequest::new([Message::user("hi")]).with_max_tokens(20))
.unwrap();
assert!(
long.content.len() >= short.content.len(),
"a larger max_tokens budget must never produce a shorter completion \
(short={:?}, long={:?})",
short.content,
long.content
);
}
#[test]
fn stream_is_wellformed_against_a_synthetic_gguf() {
let bytes = build_local_adapter_fixture();
let path = write_temp_gguf(&bytes, "local-adapter-stream");
let adapter = LocalAdapter::load(&path).unwrap();
let request = CompletionRequest::new([
Message::system("you are a test fixture"),
Message::user("hello"),
])
.with_max_tokens(8);
let chunks: Vec<StreamChunk> = adapter.stream(&request).iter().collect();
assert_eq!(chunks.last(), Some(&StreamChunk::Done));
assert_eq!(chunks.iter().filter(|c| **c == StreamChunk::Done).count(), 1);
assert!(!chunks.iter().any(|c| matches!(c, StreamChunk::Error(_))));
let token_count = chunks
.iter()
.filter(|c| matches!(c, StreamChunk::Token(_)))
.count();
assert_eq!(token_count, chunks.len() - 1, "every chunk but the terminal Done must be a Token");
assert!(token_count <= 8, "generated more tokens than the max_tokens budget");
}
#[test]
#[ignore = "no real Qwen GGUF present on this machine; point KOPITIAM_QWEN_GGUF at one to run this"]
fn a_real_local_model_answers_a_chatml_prompt() {
let path = std::env::var("KOPITIAM_QWEN_GGUF").expect("set KOPITIAM_QWEN_GGUF to a real Qwen .gguf file");
let adapter = LocalAdapter::load(&path).expect("a real Qwen GGUF must load");
let request = CompletionRequest::new([
Message::system("You are a helpful assistant."),
Message::user("Say hello in one short sentence."),
])
.with_max_tokens(64);
let response = adapter.complete(&request).expect("a real Qwen model must generate a completion");
assert!(!response.content.trim().is_empty(), "a real model should not answer with empty text");
assert_ne!(response.model, adapter.name(), "CompletionResponse::model must name the model, not the adapter");
println!("real model {:?} answered: {:?}", response.model, response.content);
}
}