use std::ffi::{c_char, CString};
use std::ptr::NonNull;
use llama_cpp_sys_4 as sys;
use crate::model::LlamaModel;
pub type ChatError = crate::shim::ShimError;
use crate::shim::{check_status, last_error, read_string};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ToolChoice {
#[default]
Auto,
Required,
None,
}
impl ToolChoice {
#[allow(clippy::cast_possible_wrap)]
fn as_raw(self) -> i32 {
let raw = match self {
Self::Auto => sys::CHAT_SHIM_TOOL_CHOICE_AUTO,
Self::Required => sys::CHAT_SHIM_TOOL_CHOICE_REQUIRED,
Self::None => sys::CHAT_SHIM_TOOL_CHOICE_NONE,
};
raw as i32
}
#[allow(clippy::cast_possible_wrap)]
fn from_raw(raw: i32) -> Option<Self> {
if raw == sys::CHAT_SHIM_TOOL_CHOICE_AUTO as i32 {
Some(Self::Auto)
} else if raw == sys::CHAT_SHIM_TOOL_CHOICE_REQUIRED as i32 {
Some(Self::Required)
} else if raw == sys::CHAT_SHIM_TOOL_CHOICE_NONE as i32 {
Some(Self::None)
} else {
None
}
}
pub fn parse_oaicompat(value: &str) -> Result<Self, ChatError> {
let c_value = CString::new(value)?;
let mut raw: i32 = 0;
let status = unsafe {
sys::chat_shim_tool_choice_parse_oaicompat(c_value.as_ptr(), &raw mut raw)
};
check_status(status)?;
Self::from_raw(raw).ok_or_else(|| ChatError::Failed(format!("unknown tool_choice {raw}")))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ReasoningFormat {
#[default]
None,
Auto,
DeepSeekLegacy,
DeepSeek,
}
impl ReasoningFormat {
#[allow(clippy::cast_possible_wrap)]
fn as_raw(self) -> i32 {
let raw = match self {
Self::None => sys::CHAT_SHIM_REASONING_NONE,
Self::Auto => sys::CHAT_SHIM_REASONING_AUTO,
Self::DeepSeekLegacy => sys::CHAT_SHIM_REASONING_DEEPSEEK_LEGACY,
Self::DeepSeek => sys::CHAT_SHIM_REASONING_DEEPSEEK,
};
raw as i32
}
}
pub fn json_schema_to_grammar(schema_json: &str, force_gbnf: bool) -> Result<String, ChatError> {
let c_schema = CString::new(schema_json)?;
read_string(|buf, len, expected| unsafe {
sys::common_json_schema_to_grammar_c(c_schema.as_ptr(), force_gbnf, buf, len, expected)
})
}
#[derive(Debug)]
pub struct ChatTemplates {
raw: NonNull<sys::chat_shim_templates>,
}
unsafe impl Send for ChatTemplates {}
unsafe impl Sync for ChatTemplates {}
impl Drop for ChatTemplates {
fn drop(&mut self) {
unsafe { sys::chat_shim_templates_free(self.raw.as_ptr()) }
}
}
impl ChatTemplates {
pub fn from_model(
model: &LlamaModel,
template_override: Option<&str>,
) -> Result<Self, ChatError> {
let c_override = template_override.map(CString::new).transpose()?;
let override_ptr = c_override.as_ref().map_or(std::ptr::null(), |c| c.as_ptr());
let raw = unsafe { sys::chat_shim_templates_init(model.model.as_ptr(), override_ptr) };
NonNull::new(raw)
.map(|raw| Self { raw })
.ok_or_else(|| ChatError::Init(last_error()))
}
pub fn source(&self, variant: Option<&str>) -> Result<String, ChatError> {
let c_variant = variant.map(CString::new).transpose()?;
let variant_ptr = c_variant.as_ref().map_or(std::ptr::null(), |c| c.as_ptr());
read_string(|buf, len, expected| unsafe {
sys::chat_shim_templates_source(self.raw.as_ptr(), variant_ptr, buf, len, expected)
})
}
#[must_use]
pub fn was_explicit(&self) -> bool {
unsafe { sys::chat_shim_templates_was_explicit(self.raw.as_ptr()) }
}
#[must_use]
pub fn supports_enable_thinking(&self) -> bool {
unsafe { sys::chat_shim_templates_support_enable_thinking(self.raw.as_ptr()) }
}
pub fn caps_json(&self) -> Result<String, ChatError> {
read_string(|buf, len, expected| unsafe {
sys::chat_shim_templates_get_caps(self.raw.as_ptr(), buf, len, expected)
})
}
pub fn apply(&self, params: &ChatApplyParams) -> Result<ChatParams, ChatError> {
let messages = CString::new(params.messages_json.as_str())?;
let tools = params.tools_json.as_deref().map(CString::new).transpose()?;
let grammar = params.grammar.as_deref().map(CString::new).transpose()?;
let schema = params.json_schema.as_deref().map(CString::new).transpose()?;
let kwargs = params
.template_kwargs_json
.as_deref()
.map(CString::new)
.transpose()?;
let raw_params = sys::chat_shim_apply_params {
messages_json: messages.as_ptr(),
tools_json: opt_ptr(tools.as_ref()),
grammar: opt_ptr(grammar.as_ref()),
json_schema: opt_ptr(schema.as_ref()),
template_kwargs_json: opt_ptr(kwargs.as_ref()),
tool_choice: params.tool_choice.as_raw(),
reasoning_format: params.reasoning_format.as_raw(),
add_generation_prompt: params.add_generation_prompt,
enable_thinking: params.enable_thinking,
parallel_tool_calls: params.parallel_tool_calls,
use_jinja: params.use_jinja,
add_bos: params.add_bos,
add_eos: params.add_eos,
};
let mut result: sys::chat_shim_apply_result = unsafe { std::mem::zeroed() };
let mut needed: usize = 0;
let status = unsafe {
sys::chat_shim_templates_apply(
self.raw.as_ptr(),
&raw const raw_params,
&raw mut result,
std::ptr::null_mut(),
0,
&raw mut needed,
)
};
if status != sys::LLAMA_SHIM_BUFFER_TOO_SMALL {
check_status(status)?;
}
let mut buf = vec![0u8; needed];
let status = unsafe {
sys::chat_shim_templates_apply(
self.raw.as_ptr(),
&raw const raw_params,
&raw mut result,
buf.as_mut_ptr().cast::<c_char>(),
buf.len(),
&raw mut needed,
)
};
check_status(status)?;
ChatParams::from_packed(&buf, &result)
}
}
#[allow(clippy::struct_excessive_bools)]
#[derive(Debug, Clone)]
pub struct ChatApplyParams {
messages_json: String,
tools_json: Option<String>,
grammar: Option<String>,
json_schema: Option<String>,
template_kwargs_json: Option<String>,
tool_choice: ToolChoice,
reasoning_format: ReasoningFormat,
add_generation_prompt: bool,
enable_thinking: bool,
parallel_tool_calls: bool,
use_jinja: bool,
add_bos: bool,
add_eos: bool,
}
impl ChatApplyParams {
#[must_use]
pub fn new(messages_json: impl Into<String>) -> Self {
Self {
messages_json: messages_json.into(),
tools_json: None,
grammar: None,
json_schema: None,
template_kwargs_json: None,
tool_choice: ToolChoice::Auto,
reasoning_format: ReasoningFormat::Auto,
add_generation_prompt: true,
enable_thinking: true,
parallel_tool_calls: false,
use_jinja: true,
add_bos: false,
add_eos: false,
}
}
#[must_use]
pub fn with_tools(mut self, tools_json: impl Into<String>) -> Self {
self.tools_json = Some(tools_json.into());
self
}
#[must_use]
pub fn with_grammar(mut self, grammar: impl Into<String>) -> Self {
self.grammar = Some(grammar.into());
self
}
#[must_use]
pub fn with_json_schema(mut self, schema_json: impl Into<String>) -> Self {
self.json_schema = Some(schema_json.into());
self
}
#[must_use]
pub fn with_template_kwargs(mut self, kwargs_json: impl Into<String>) -> Self {
self.template_kwargs_json = Some(kwargs_json.into());
self
}
#[must_use]
pub fn with_tool_choice(mut self, choice: ToolChoice) -> Self {
self.tool_choice = choice;
self
}
#[must_use]
pub fn with_reasoning_format(mut self, format: ReasoningFormat) -> Self {
self.reasoning_format = format;
self
}
#[must_use]
pub fn with_add_generation_prompt(mut self, add: bool) -> Self {
self.add_generation_prompt = add;
self
}
#[must_use]
pub fn with_enable_thinking(mut self, enable: bool) -> Self {
self.enable_thinking = enable;
self
}
#[must_use]
pub fn with_parallel_tool_calls(mut self, parallel: bool) -> Self {
self.parallel_tool_calls = parallel;
self
}
#[must_use]
pub fn with_use_jinja(mut self, use_jinja: bool) -> Self {
self.use_jinja = use_jinja;
self
}
#[allow(clippy::similar_names)]
#[must_use]
pub fn with_add_bos_eos(mut self, add_bos: bool, add_eos: bool) -> Self {
self.add_bos = add_bos;
self.add_eos = add_eos;
self
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GrammarTrigger {
pub kind: String,
pub value: String,
pub token: i32,
}
#[derive(Debug, Clone)]
pub struct ChatParams {
pub prompt: String,
pub grammar: String,
pub grammar_lazy: bool,
pub grammar_triggers_json: String,
pub grammar_triggers: Vec<GrammarTrigger>,
pub preserved_tokens_json: String,
pub additional_stops_json: String,
pub supports_thinking: bool,
pub thinking_start_tag: String,
pub thinking_end_tags_json: String,
pub format: i32,
parser: String,
generation_prompt: String,
reasoning_format: ReasoningFormat,
}
impl ChatParams {
fn from_packed(
buf: &[u8],
result: &sys::chat_shim_apply_result,
) -> Result<Self, ChatError> {
let at = |off: usize| -> Result<String, ChatError> {
let rest = buf.get(off..).ok_or(ChatError::CorruptResult)?;
let end = rest
.iter()
.position(|b| *b == 0)
.ok_or(ChatError::CorruptResult)?;
String::from_utf8(rest[..end].to_vec()).map_err(ChatError::from)
};
let grammar_triggers_json = at(result.grammar_triggers_off)?;
let grammar_triggers = parse_triggers(&grammar_triggers_json);
Ok(Self {
prompt: at(result.prompt_off)?,
grammar: at(result.grammar_off)?,
grammar_lazy: result.grammar_lazy,
grammar_triggers,
grammar_triggers_json,
preserved_tokens_json: at(result.preserved_tokens_off)?,
additional_stops_json: at(result.additional_stops_off)?,
supports_thinking: result.supports_thinking,
thinking_start_tag: at(result.thinking_start_tag_off)?,
thinking_end_tags_json: at(result.thinking_end_tags_off)?,
format: result.format,
parser: at(result.parser_off)?,
generation_prompt: at(result.generation_prompt_off)?,
reasoning_format: ReasoningFormat::Auto,
})
}
pub fn format_name(&self) -> Result<String, ChatError> {
format_name(self.format)
}
#[must_use]
pub fn sampler_triggers(&self) -> (Vec<String>, Vec<crate::token::LlamaToken>) {
let mut patterns = Vec::new();
let mut tokens = Vec::new();
for trigger in &self.grammar_triggers {
match trigger.kind.as_str() {
"word" => patterns.push(regex_escape(&trigger.value)),
"pattern" => patterns.push(trigger.value.clone()),
"pattern_full" => patterns.push(anchor_pattern(&trigger.value)),
"token" => tokens.push(crate::token::LlamaToken(trigger.token)),
_ => {}
}
}
(patterns, tokens)
}
#[must_use]
pub fn generation_prompt(&self) -> &str {
&self.generation_prompt
}
#[must_use]
pub fn grammar_sampler(&self, model: &LlamaModel) -> Option<crate::sampling::LlamaSampler> {
use crate::sampling::LlamaSampler;
if self.grammar.is_empty() {
return None;
}
let (patterns, tokens) = self.sampler_triggers();
let lazy = self.grammar_lazy && !(patterns.is_empty() && tokens.is_empty());
let mut sampler = if lazy {
let refs: Vec<&str> = patterns.iter().map(String::as_str).collect();
LlamaSampler::grammar_lazy_patterns(model, &self.grammar, "root", &refs, &tokens)
} else {
LlamaSampler::grammar(model, &self.grammar, "root")
};
if !lazy {
for token in self.generation_prompt_tokens(model) {
sampler.accept(token);
}
}
Some(sampler)
}
fn generation_prompt_tokens(&self, model: &LlamaModel) -> Vec<crate::token::LlamaToken> {
if self.generation_prompt.is_empty() {
return Vec::new();
}
let Ok(tokens) = model.str_to_token(&self.generation_prompt, crate::model::AddBos::Never)
else {
return Vec::new();
};
let starts_with_space = self
.generation_prompt
.starts_with(char::is_whitespace);
let mut out = Vec::with_capacity(tokens.len());
for (i, token) in tokens.into_iter().enumerate() {
if i == 0 && !starts_with_space {
if let Ok(piece) = model.token_to_str(token, crate::model::Special::Tokenize) {
if piece.starts_with(char::is_whitespace) {
continue;
}
}
}
out.push(token);
}
out
}
pub fn parse(&self, text: &str, is_partial: bool) -> Result<String, ChatError> {
self.parse_with(text, is_partial, true, false)
}
pub fn parse_with(
&self,
text: &str,
is_partial: bool,
parse_tool_calls: bool,
reasoning_in_content: bool,
) -> Result<String, ChatError> {
let c_text = CString::new(text)?;
let c_parser = CString::new(self.parser.as_str())?;
let c_gen_prompt = CString::new(self.generation_prompt.as_str())?;
let params = sys::chat_shim_parse_params {
text: c_text.as_ptr(),
parser: c_parser.as_ptr(),
generation_prompt: c_gen_prompt.as_ptr(),
format: self.format,
reasoning_format: self.reasoning_format.as_raw(),
is_partial,
parse_tool_calls,
reasoning_in_content,
};
read_string(|buf, len, expected| unsafe {
sys::chat_shim_parse(&raw const params, buf, len, expected)
})
}
}
pub fn format_name(format: i32) -> Result<String, ChatError> {
read_string(|buf, len, expected| unsafe {
sys::chat_shim_format_name(format, buf, len, expected)
})
}
pub fn parse_messages_oaicompat(messages_json: &str) -> Result<String, ChatError> {
let c_messages = CString::new(messages_json)?;
read_string(|buf, len, expected| unsafe {
sys::chat_shim_msgs_parse_oaicompat(c_messages.as_ptr(), buf, len, expected)
})
}
pub fn parse_tools_oaicompat(tools_json: &str) -> Result<String, ChatError> {
let c_tools = CString::new(tools_json)?;
read_string(|buf, len, expected| unsafe {
sys::chat_shim_tools_parse_oaicompat(c_tools.as_ptr(), buf, len, expected)
})
}
fn regex_escape(s: &str) -> String {
const SPECIAL: &[char] = &[
'.', '^', '$', '|', '(', ')', '*', '+', '?', '[', ']', '{', '}', '\\',
];
let mut out = String::with_capacity(s.len());
for c in s.chars() {
if SPECIAL.contains(&c) {
out.push('\\');
}
out.push(c);
}
out
}
fn anchor_pattern(pattern: &str) -> String {
if pattern.is_empty() {
return "^$".to_owned();
}
let mut out = String::with_capacity(pattern.len() + 2);
if !pattern.starts_with('^') {
out.push('^');
}
out.push_str(pattern);
if !pattern.ends_with('$') {
out.push('$');
}
out
}
fn opt_ptr(s: Option<&CString>) -> *const c_char {
s.map_or(std::ptr::null(), |c| c.as_ptr())
}
fn parse_triggers(json: &str) -> Vec<GrammarTrigger> {
let mut out = Vec::new();
for chunk in json.split('{').skip(1) {
let kind = json_str_field(chunk, "type");
let value = json_str_field(chunk, "value");
let token = json_int_field(chunk, "token");
if let (Some(kind), Some(value)) = (kind, value) {
out.push(GrammarTrigger {
kind,
value,
token: token.unwrap_or(-1),
});
}
}
out
}
fn json_str_field(chunk: &str, key: &str) -> Option<String> {
let needle = format!("\"{key}\":");
let rest = &chunk[chunk.find(&needle)? + needle.len()..];
let rest = rest.trim_start();
let mut chars = rest.strip_prefix('"')?.chars();
let mut value = String::new();
while let Some(c) = chars.next() {
match c {
'"' => return Some(value),
'\\' => match chars.next()? {
'n' => value.push('\n'),
'r' => value.push('\r'),
't' => value.push('\t'),
'u' => {
let hex: String = chars.by_ref().take(4).collect();
let code = u32::from_str_radix(&hex, 16).ok()?;
value.push(char::from_u32(code)?);
}
other => value.push(other),
},
other => value.push(other),
}
}
None
}
fn json_int_field(chunk: &str, key: &str) -> Option<i32> {
let needle = format!("\"{key}\":");
let rest = &chunk[chunk.find(&needle)? + needle.len()..];
let rest = rest.trim_start();
let end = rest
.find(|c: char| !c.is_ascii_digit() && c != '-')
.unwrap_or(rest.len());
rest[..end].parse().ok()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn realistic_object_schema_converts() {
let gbnf = json_schema_to_grammar(
r#"{"type":"object","properties":{"city":{"type":"string"}},"required":["city"]}"#,
false,
)
.unwrap();
eprintln!("GBNF:\n{gbnf}");
assert!(gbnf.contains("root"));
}
#[test]
fn json_schema_to_grammar_produces_a_root_rule() {
let gbnf = json_schema_to_grammar(r#"{"type":"integer"}"#, false).unwrap();
assert!(gbnf.contains("root"), "no root rule in: {gbnf}");
}
#[test]
fn json_schema_to_grammar_constrains_object_keys() {
let gbnf = json_schema_to_grammar(
r#"{"type":"object","properties":{"city":{"type":"string"}},"required":["city"]}"#,
false,
)
.unwrap();
assert!(gbnf.contains("city"), "key not in grammar: {gbnf}");
}
#[test]
fn json_schema_to_grammar_rejects_invalid_json() {
let err = json_schema_to_grammar("{not json", false).unwrap_err();
assert!(
matches!(err, ChatError::BadJson(_)),
"expected BadJson, got {err:?}"
);
}
#[test]
fn json_schema_to_grammar_rejects_empty_input() {
assert!(json_schema_to_grammar("", false).is_err());
}
#[test]
fn json_schema_to_grammar_rejects_interior_nul() {
let err = json_schema_to_grammar("{\"type\":\"in\0teger\"}", false).unwrap_err();
assert!(matches!(err, ChatError::Nul(_)), "got {err:?}");
}
#[test]
fn tool_choice_parses_openai_values() {
assert_eq!(ToolChoice::parse_oaicompat("auto").unwrap(), ToolChoice::Auto);
assert_eq!(
ToolChoice::parse_oaicompat("required").unwrap(),
ToolChoice::Required
);
assert_eq!(ToolChoice::parse_oaicompat("none").unwrap(), ToolChoice::None);
}
#[test]
fn tool_choice_rejects_unknown_value() {
assert!(ToolChoice::parse_oaicompat("sometimes").is_err());
}
#[test]
fn tools_parse_oaicompat_extracts_name_and_parameters() {
let normalised = parse_tools_oaicompat(
r#"[{"type":"function","function":{"name":"get_weather",
"description":"Get weather","parameters":{"type":"object"}}}]"#,
)
.unwrap();
assert!(normalised.contains("get_weather"), "got {normalised}");
}
#[test]
fn tools_parse_oaicompat_rejects_garbage() {
assert!(parse_tools_oaicompat("[[[").is_err());
}
#[test]
fn messages_parse_oaicompat_roundtrips_a_simple_turn() {
let normalised =
parse_messages_oaicompat(r#"[{"role":"user","content":"hi"}]"#).unwrap();
assert!(normalised.contains("user"), "got {normalised}");
assert!(normalised.contains("hi"), "got {normalised}");
}
#[test]
fn messages_parse_oaicompat_rejects_non_array() {
assert!(parse_messages_oaicompat(r#"{"role":"user"}"#).is_err());
}
#[test]
fn trigger_parser_reads_shim_output() {
let json = r#"[{"type":"word","value":"<tool_call>","token":-1},
{"type":"token","value":"a\nb","token":42}]"#;
let triggers = parse_triggers(json);
assert_eq!(triggers.len(), 2);
assert_eq!(triggers[0].kind, "word");
assert_eq!(triggers[0].value, "<tool_call>");
assert_eq!(triggers[0].token, -1);
assert_eq!(triggers[1].kind, "token");
assert_eq!(triggers[1].value, "a\nb");
assert_eq!(triggers[1].token, 42);
}
#[test]
fn word_triggers_are_regex_escaped() {
let params = ChatParams {
grammar_triggers: vec![GrammarTrigger {
kind: "word".to_owned(),
value: "a.b[c]".to_owned(),
token: -1,
}],
..stub_params()
};
let (patterns, tokens) = params.sampler_triggers();
assert_eq!(patterns, vec![r"a\.b\[c\]".to_owned()]);
assert!(tokens.is_empty());
}
#[test]
fn pattern_triggers_pass_through_unescaped() {
let params = ChatParams {
grammar_triggers: vec![GrammarTrigger {
kind: "pattern".to_owned(),
value: "a.b".to_owned(),
token: -1,
}],
..stub_params()
};
assert_eq!(params.sampler_triggers().0, vec!["a.b".to_owned()]);
}
#[test]
fn pattern_full_triggers_are_anchored_once() {
let cases = [
("abc", "^abc$"),
("^abc", "^abc$"),
("abc$", "^abc$"),
("^abc$", "^abc$"),
("", "^$"),
];
for (input, want) in cases {
let params = ChatParams {
grammar_triggers: vec![GrammarTrigger {
kind: "pattern_full".to_owned(),
value: input.to_owned(),
token: -1,
}],
..stub_params()
};
assert_eq!(
params.sampler_triggers().0,
vec![want.to_owned()],
"anchoring {input:?}"
);
}
}
#[test]
fn token_triggers_become_tokens_not_patterns() {
let params = ChatParams {
grammar_triggers: vec![GrammarTrigger {
kind: "token".to_owned(),
value: String::new(),
token: 42,
}],
..stub_params()
};
let (patterns, tokens) = params.sampler_triggers();
assert!(patterns.is_empty());
assert_eq!(tokens, vec![crate::token::LlamaToken(42)]);
}
#[test]
fn unknown_trigger_kinds_are_dropped() {
let params = ChatParams {
grammar_triggers: vec![GrammarTrigger {
kind: "something_new".to_owned(),
value: "x".to_owned(),
token: -1,
}],
..stub_params()
};
let (patterns, tokens) = params.sampler_triggers();
assert!(patterns.is_empty());
assert!(tokens.is_empty());
}
fn stub_params() -> ChatParams {
ChatParams {
prompt: String::new(),
grammar: String::new(),
grammar_lazy: false,
grammar_triggers_json: String::new(),
grammar_triggers: Vec::new(),
preserved_tokens_json: String::new(),
additional_stops_json: String::new(),
supports_thinking: false,
thinking_start_tag: String::new(),
thinking_end_tags_json: String::new(),
format: 0,
parser: String::new(),
generation_prompt: String::new(),
reasoning_format: ReasoningFormat::Auto,
}
}
#[test]
fn trigger_parser_handles_empty_array() {
assert!(parse_triggers("[]").is_empty());
}
}