use std::ffi::CString;
use std::marker::PhantomData;
use std::os::raw::c_char;
use std::ptr::NonNull;
use ik_llama_cpp_sys as sys;
use crate::context::LlamaContext;
use crate::model::LlamaModel;
use crate::sampling::LlamaTokenDataArray;
use crate::token::LlamaToken;
#[derive(Debug, thiserror::Error)]
pub enum GrammarInitError {
#[error("grammar string contained an interior NUL byte")]
Nul(#[from] std::ffi::NulError),
#[error("failed to parse GBNF grammar")]
Parse,
}
#[derive(Debug)]
pub struct LlamaGrammar<'model> {
grammar: NonNull<sys::llama_grammar>,
_model: PhantomData<&'model LlamaModel>,
}
impl<'model> LlamaGrammar<'model> {
pub fn new(
model: &'model LlamaModel,
grammar_str: &str,
root: &str,
) -> Result<Self, GrammarInitError> {
let c_grammar = CString::new(grammar_str)?;
let c_root = CString::new(root)?;
let vocab = unsafe { sys::llama_model_get_vocab(model.model.as_ptr()) };
let raw =
unsafe { sys::llama_sampler_init_grammar(vocab, c_grammar.as_ptr(), c_root.as_ptr()) };
NonNull::new(raw)
.map(|grammar| Self {
grammar,
_model: PhantomData,
})
.ok_or(GrammarInitError::Parse)
}
pub fn apply(&self, ctx: &mut LlamaContext, arr: &mut LlamaTokenDataArray) {
let mut c = arr.as_c();
unsafe {
sys::llama_grammar_apply(self.grammar.as_ptr(), ctx.as_ptr(), &mut c);
}
}
pub fn accept_token(&mut self, ctx: &mut LlamaContext, token: LlamaToken) {
unsafe {
sys::llama_grammar_accept_token(self.grammar.as_ptr(), ctx.as_ptr(), token.0);
}
}
#[must_use]
pub fn try_clone(&self) -> Option<Self> {
let raw = unsafe { sys::llama_grammar_copy(self.grammar.as_ptr()) };
NonNull::new(raw).map(|grammar| Self {
grammar,
_model: PhantomData,
})
}
}
impl Drop for LlamaGrammar<'_> {
fn drop(&mut self) {
unsafe { sys::llama_grammar_free(self.grammar.as_ptr()) };
}
}
#[derive(Debug, Clone)]
pub struct DryParams {
pub multiplier: f32,
pub base: f32,
pub allowed_length: i32,
pub penalty_last_n: i32,
pub seq_breakers: Vec<String>,
}
impl Default for DryParams {
fn default() -> Self {
Self {
multiplier: 0.0,
base: 1.75,
allowed_length: 2,
penalty_last_n: -1,
seq_breakers: ["\n", ":", "\"", "*"]
.iter()
.map(|s| (*s).to_string())
.collect(),
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum DryInitError {
#[error("a DRY sequence breaker contained an interior NUL byte")]
Nul(#[from] std::ffi::NulError),
#[error("failed to initialize the DRY sampler")]
Init,
}
#[derive(Debug)]
pub struct LlamaDrySampler {
dry: NonNull<sys::llama_sampler_dry>,
}
impl LlamaDrySampler {
pub fn new(model: &LlamaModel, params: &DryParams) -> Result<Self, DryInitError> {
let c_breakers: Vec<CString> = params
.seq_breakers
.iter()
.map(|s| CString::new(s.as_str()))
.collect::<Result<_, _>>()?;
let mut ptrs: Vec<*const c_char> = c_breakers.iter().map(|c| c.as_ptr()).collect();
let vocab = unsafe { sys::llama_model_get_vocab(model.model.as_ptr()) };
let raw = unsafe {
sys::llama_sampler_init_dry(
vocab,
params.multiplier,
params.base,
params.allowed_length,
params.penalty_last_n,
ptrs.as_mut_ptr(),
ptrs.len(),
)
};
NonNull::new(raw)
.map(|dry| Self { dry })
.ok_or(DryInitError::Init)
}
pub fn apply(&mut self, ctx: &mut LlamaContext, arr: &mut LlamaTokenDataArray) {
let mut c = arr.as_c();
unsafe { sys::llama_sample_dry(ctx.as_ptr(), self.dry.as_ptr(), &mut c) };
}
pub fn accept(&mut self, token: LlamaToken) {
unsafe { sys::llama_sampler_dry_accept(self.dry.as_ptr(), token.0) };
}
pub fn reset(&mut self) {
unsafe { sys::llama_sampler_dry_reset(self.dry.as_ptr()) };
}
#[must_use]
pub fn try_clone(&self) -> Option<Self> {
let raw = unsafe { sys::llama_sampler_dry_clone(self.dry.as_ptr()) };
NonNull::new(raw).map(|dry| Self { dry })
}
}
impl Drop for LlamaDrySampler {
fn drop(&mut self) {
unsafe { sys::llama_sampler_dry_free(self.dry.as_ptr()) };
}
}
#[cfg(feature = "common")]
#[derive(Debug, thiserror::Error)]
pub enum JsonSchemaError {
#[error("JSON schema string contained an interior NUL byte")]
Nul(#[from] std::ffi::NulError),
#[error("JSON schema to grammar conversion failed (status {0})")]
Convert(i32),
#[error("converted grammar was not valid UTF-8")]
Utf8,
}
#[cfg(feature = "common")]
pub fn json_schema_to_grammar(schema_json: &str) -> Result<String, JsonSchemaError> {
let schema = CString::new(schema_json)?;
let mut out: *mut c_char = std::ptr::null_mut();
let status =
unsafe { sys::ik_llama_rs_json_schema_to_grammar(schema.as_ptr(), false, &mut out) };
if status as i32 != 0 || out.is_null() {
if !out.is_null() {
unsafe { sys::ik_llama_rs_string_free(out) };
}
return Err(JsonSchemaError::Convert(status as i32));
}
let bytes = unsafe { std::ffi::CStr::from_ptr(out) }.to_bytes().to_vec();
unsafe { sys::ik_llama_rs_string_free(out) };
String::from_utf8(bytes).map_err(|_| JsonSchemaError::Utf8)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dry_params_default_is_disabled() {
let p = DryParams::default();
assert_eq!(p.multiplier, 0.0, "DRY is disabled by default");
assert_eq!(p.seq_breakers.len(), 4);
assert!(p.seq_breakers.iter().any(|b| b == "\n"));
}
}