use std::ffi::{c_char, CString};
use std::ptr::NonNull;
use llama_cpp_sys_4 as sys;
use crate::context::LlamaContext;
use crate::model::LlamaModel;
use crate::token::LlamaToken;
pub type CommonSamplerError = crate::shim::ShimError;
use crate::shim::{check_status, last_error, read_i32s, read_string, read_tokens, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(i32)]
pub enum CommonSamplerType {
Dry = 1,
TopK = 2,
TopP = 3,
MinP = 4,
TypicalP = 6,
Temperature = 7,
Xtc = 8,
Infill = 9,
Penalties = 10,
TopNSigma = 11,
AdaptiveP = 12,
}
impl CommonSamplerType {
pub fn name(self) -> Result<String> {
read_string(|buf, len, expected| unsafe {
sys::common_shim_sampler_type_to_str(self as i32, buf, len, expected)
})
}
pub fn from_names(names: &[&str]) -> Result<Vec<i32>> {
let c_names: Vec<CString> = names
.iter()
.map(|n| CString::new(*n))
.collect::<std::result::Result<_, _>>()?;
let ptrs: Vec<*const c_char> = c_names.iter().map(|c| c.as_ptr()).collect();
read_i32s(|out, cap, len| unsafe {
sys::common_shim_sampler_types_from_names(ptrs.as_ptr(), ptrs.len(), out, cap, len)
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum GrammarSource {
#[default]
None,
User,
OutputFormat,
ToolCalls,
}
impl GrammarSource {
#[allow(clippy::cast_possible_wrap)]
fn as_raw(self) -> i32 {
let raw = match self {
Self::None => sys::COMMON_SHIM_GRAMMAR_NONE,
Self::User => sys::COMMON_SHIM_GRAMMAR_USER,
Self::OutputFormat => sys::COMMON_SHIM_GRAMMAR_OUTPUT_FORMAT,
Self::ToolCalls => sys::COMMON_SHIM_GRAMMAR_TOOL_CALLS,
};
raw as i32
}
}
pub struct CommonSamplerParams {
raw: NonNull<sys::common_shim_sampler_params>,
}
impl std::fmt::Debug for CommonSamplerParams {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CommonSamplerParams")
.field("scalars", &self.scalars())
.finish_non_exhaustive()
}
}
unsafe impl Send for CommonSamplerParams {}
impl Drop for CommonSamplerParams {
fn drop(&mut self) {
unsafe { sys::common_shim_sampler_params_free(self.raw.as_ptr()) }
}
}
pub type CommonSamplerScalars = sys::common_shim_sampler_scalars;
impl Default for CommonSamplerParams {
fn default() -> Self {
Self::new()
}
}
impl CommonSamplerParams {
#[must_use]
pub fn new() -> Self {
let raw = unsafe { sys::common_shim_sampler_params_init() };
Self {
raw: NonNull::new(raw).expect("common_shim_sampler_params_init returned null"),
}
}
#[must_use]
pub fn scalars(&self) -> CommonSamplerScalars {
let mut out: CommonSamplerScalars = unsafe { std::mem::zeroed() };
unsafe { sys::common_shim_sampler_params_get_scalars(self.raw.as_ptr(), &raw mut out) };
out
}
pub fn set_scalars(&mut self, scalars: &CommonSamplerScalars) {
unsafe { sys::common_shim_sampler_params_set_scalars(self.raw.as_ptr(), scalars) }
}
pub fn set_grammar(&mut self, grammar: &str, source: GrammarSource, lazy: bool) -> Result<()> {
let c_grammar = CString::new(grammar)?;
let status = unsafe {
sys::common_shim_sampler_params_set_grammar(
self.raw.as_ptr(),
c_grammar.as_ptr(),
source.as_raw(),
lazy,
)
};
check_status(status)
}
pub fn add_grammar_trigger(&mut self, kind: &str, value: &str, token: LlamaToken) -> Result<()> {
let raw_kind = match kind {
"token" => 0,
"word" => 1,
"pattern" => 2,
"pattern_full" => 3,
_ => return Err(CommonSamplerError::InvalidArg),
};
let c_value = CString::new(value)?;
let status = unsafe {
sys::common_shim_sampler_params_add_grammar_trigger(
self.raw.as_ptr(),
raw_kind,
c_value.as_ptr(),
token.0,
)
};
check_status(status)
}
pub fn set_generation_prompt(&mut self, prompt: &str) -> Result<()> {
let c_prompt = CString::new(prompt)?;
let status = unsafe {
sys::common_shim_sampler_params_set_generation_prompt(
self.raw.as_ptr(),
c_prompt.as_ptr(),
)
};
check_status(status)
}
pub fn add_logit_bias(&mut self, token: LlamaToken, bias: f32) -> Result<()> {
let status = unsafe {
sys::common_shim_sampler_params_add_logit_bias(self.raw.as_ptr(), token.0, bias)
};
check_status(status)
}
pub fn set_samplers(&mut self, samplers: &[CommonSamplerType]) -> Result<()> {
let raw: Vec<i32> = samplers.iter().map(|s| *s as i32).collect();
let status = unsafe {
sys::common_shim_sampler_params_set_samplers(
self.raw.as_ptr(),
raw.as_ptr(),
raw.len(),
)
};
check_status(status)
}
pub fn set_dry_breakers(&mut self, breakers: &[&str]) -> Result<()> {
let c_breakers: Vec<CString> = breakers
.iter()
.map(|b| CString::new(*b))
.collect::<std::result::Result<_, _>>()?;
let ptrs: Vec<*const c_char> = c_breakers.iter().map(|c| c.as_ptr()).collect();
let status = unsafe {
sys::common_shim_sampler_params_set_dry_breakers(
self.raw.as_ptr(),
ptrs.as_ptr(),
ptrs.len(),
)
};
check_status(status)
}
pub fn set_reasoning_budget(
&mut self,
start: &[LlamaToken],
ends: &[Vec<LlamaToken>],
forced: &[LlamaToken],
message: &str,
) -> Result<()> {
let start_raw: Vec<i32> = start.iter().map(|t| t.0).collect();
let (ends_flat, end_lens) = flatten(ends);
let forced_raw: Vec<i32> = forced.iter().map(|t| t.0).collect();
let c_message = CString::new(message)?;
let status = unsafe {
sys::common_shim_sampler_params_set_reasoning_budget(
self.raw.as_ptr(),
start_raw.as_ptr(),
start_raw.len(),
ends_flat.as_ptr(),
end_lens.as_ptr(),
end_lens.len(),
forced_raw.as_ptr(),
forced_raw.len(),
c_message.as_ptr(),
)
};
check_status(status)
}
}
fn flatten(seqs: &[Vec<LlamaToken>]) -> (Vec<i32>, Vec<usize>) {
let mut data = Vec::new();
let mut lens = Vec::with_capacity(seqs.len());
for seq in seqs {
lens.push(seq.len());
data.extend(seq.iter().map(|t| t.0));
}
(data, lens)
}
pub struct CommonSampler {
raw: NonNull<sys::common_shim_sampler>,
}
impl std::fmt::Debug for CommonSampler {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CommonSampler")
.field("seed", &self.seed())
.finish_non_exhaustive()
}
}
unsafe impl Send for CommonSampler {}
impl Drop for CommonSampler {
fn drop(&mut self) {
unsafe { sys::common_shim_sampler_free(self.raw.as_ptr()) }
}
}
impl CommonSampler {
pub fn new(model: &LlamaModel, params: &mut CommonSamplerParams) -> Result<Self> {
let raw =
unsafe { sys::common_shim_sampler_init(model.model.as_ptr(), params.raw.as_ptr()) };
NonNull::new(raw)
.map(|raw| Self { raw })
.ok_or_else(|| CommonSamplerError::Failed(last_error()))
}
pub fn sample(
&mut self,
ctx: &mut LlamaContext<'_>,
idx: i32,
grammar_first: bool,
) -> Result<LlamaToken> {
let mut status = sys::LLAMA_SHIM_OK;
let token = unsafe {
sys::common_shim_sampler_sample(
self.raw.as_ptr(),
ctx.context.as_ptr(),
idx,
grammar_first,
&raw mut status,
)
};
check_status(status)?;
Ok(LlamaToken(token))
}
pub fn accept(&mut self, token: LlamaToken, is_generated: bool) {
unsafe { sys::common_shim_sampler_accept(self.raw.as_ptr(), token.0, is_generated) }
}
pub fn sample_and_accept_n(
&mut self,
ctx: &mut LlamaContext<'_>,
draft: &[LlamaToken],
grammar_first: bool,
) -> Result<Vec<LlamaToken>> {
let raw_draft: Vec<i32> = draft.iter().map(|t| t.0).collect();
read_tokens(|out, cap, len| unsafe {
sys::common_shim_sampler_sample_and_accept_n(
self.raw.as_ptr(),
ctx.context.as_ptr(),
raw_draft.as_ptr(),
raw_draft.len(),
grammar_first,
out,
cap,
len,
)
})
}
pub fn reset(&mut self) {
unsafe { sys::common_shim_sampler_reset(self.raw.as_ptr()) }
}
pub fn try_clone(&self) -> Result<Self> {
let raw = unsafe { sys::common_shim_sampler_clone(self.raw.as_ptr()) };
NonNull::new(raw)
.map(|raw| Self { raw })
.ok_or_else(|| CommonSamplerError::Failed(last_error()))
}
#[must_use]
pub fn seed(&self) -> u32 {
unsafe { sys::common_shim_sampler_get_seed(self.raw.as_ptr()) }
}
#[must_use]
pub fn last(&self) -> LlamaToken {
LlamaToken(unsafe { sys::common_shim_sampler_last(self.raw.as_ptr()) })
}
pub fn force_end_reasoning(&mut self) -> bool {
unsafe { sys::common_shim_sampler_reasoning_budget_force(self.raw.as_ptr()) }
}
pub fn describe(&self) -> Result<String> {
read_string(|buf, len, expected| unsafe {
sys::common_shim_sampler_print(self.raw.as_ptr(), buf, len, expected)
})
}
pub fn prev_str(&mut self, ctx: &mut LlamaContext<'_>, n: i32) -> Result<String> {
read_string(|buf, len, expected| unsafe {
sys::common_shim_sampler_prev_str(
self.raw.as_ptr(),
ctx.context.as_ptr(),
n,
buf,
len,
expected,
)
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReasoningBudgetState {
Idle,
Counting,
Forcing,
WaitingUtf8,
Done,
}
impl ReasoningBudgetState {
#[allow(clippy::cast_possible_wrap)]
fn from_raw(raw: i32) -> Option<Self> {
if raw == sys::COMMON_SHIM_RBUDGET_IDLE as i32 {
Some(Self::Idle)
} else if raw == sys::COMMON_SHIM_RBUDGET_COUNTING as i32 {
Some(Self::Counting)
} else if raw == sys::COMMON_SHIM_RBUDGET_FORCING as i32 {
Some(Self::Forcing)
} else if raw == sys::COMMON_SHIM_RBUDGET_WAITING_UTF8 as i32 {
Some(Self::WaitingUtf8)
} else if raw == sys::COMMON_SHIM_RBUDGET_DONE as i32 {
Some(Self::Done)
} else {
None
}
}
#[allow(clippy::cast_possible_wrap)]
fn as_raw(self) -> i32 {
let raw = match self {
Self::Idle => sys::COMMON_SHIM_RBUDGET_IDLE,
Self::Counting => sys::COMMON_SHIM_RBUDGET_COUNTING,
Self::Forcing => sys::COMMON_SHIM_RBUDGET_FORCING,
Self::WaitingUtf8 => sys::COMMON_SHIM_RBUDGET_WAITING_UTF8,
Self::Done => sys::COMMON_SHIM_RBUDGET_DONE,
};
raw as i32
}
}
#[derive(Debug)]
pub struct ReasoningBudget {
sampler: crate::sampling::LlamaSampler,
}
impl ReasoningBudget {
pub fn new(
model: &LlamaModel,
starts: &[Vec<LlamaToken>],
ends: &[Vec<LlamaToken>],
forced: &[LlamaToken],
budget: i32,
) -> Result<Self> {
Self::with_initial_state(model, starts, ends, forced, budget, ReasoningBudgetState::Idle)
}
pub fn with_initial_state(
model: &LlamaModel,
starts: &[Vec<LlamaToken>],
ends: &[Vec<LlamaToken>],
forced: &[LlamaToken],
budget: i32,
initial_state: ReasoningBudgetState,
) -> Result<Self> {
let (starts_flat, start_lens) = flatten(starts);
let (ends_flat, end_lens) = flatten(ends);
let forced_raw: Vec<i32> = forced.iter().map(|t| t.0).collect();
let raw = unsafe {
sys::common_shim_reasoning_budget_init(
model.get_vocab().vocab.as_ref(),
starts_flat.as_ptr(),
start_lens.as_ptr(),
start_lens.len(),
ends_flat.as_ptr(),
end_lens.as_ptr(),
end_lens.len(),
forced_raw.as_ptr(),
forced_raw.len(),
budget,
initial_state.as_raw(),
)
};
let ptr = NonNull::new(raw).ok_or_else(|| CommonSamplerError::Failed(last_error()))?;
Ok(Self {
sampler: unsafe { crate::sampling::LlamaSampler::from_raw_ptr(ptr) },
})
}
#[must_use]
pub fn state(&self) -> Option<ReasoningBudgetState> {
let raw =
unsafe { sys::common_shim_reasoning_budget_get_state(self.sampler.as_ptr().cast_const()) };
ReasoningBudgetState::from_raw(raw)
}
pub fn force_end(&mut self) -> bool {
unsafe { sys::common_shim_reasoning_budget_force(self.sampler.as_ptr()) }
}
#[must_use]
pub fn into_sampler(self) -> crate::sampling::LlamaSampler {
self.sampler
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn params_start_from_llama_cpp_defaults() {
let params = CommonSamplerParams::new();
let s = params.scalars();
assert!((s.temp - 0.80).abs() < 1e-6, "temp = {}", s.temp);
assert_eq!(s.top_k, 40);
assert!((s.top_p - 0.95).abs() < 1e-6, "top_p = {}", s.top_p);
assert!((s.min_p - 0.05).abs() < 1e-6, "min_p = {}", s.min_p);
assert_eq!(s.mirostat, 0);
assert_eq!(s.reasoning_budget_tokens, -1, "budget disabled by default");
}
#[test]
fn scalars_round_trip() {
let mut params = CommonSamplerParams::new();
let mut s = params.scalars();
s.temp = 0.25;
s.top_k = 7;
s.reasoning_budget_tokens = 128;
params.set_scalars(&s);
let back = params.scalars();
assert!((back.temp - 0.25).abs() < 1e-6);
assert_eq!(back.top_k, 7);
assert_eq!(back.reasoning_budget_tokens, 128);
}
#[test]
fn setting_scalars_preserves_untouched_fields() {
let mut params = CommonSamplerParams::new();
let before = params.scalars();
let mut s = before;
s.top_k = 3;
params.set_scalars(&s);
let after = params.scalars();
assert_eq!(after.top_k, 3);
assert!((after.top_p - before.top_p).abs() < 1e-6);
assert!((after.dry_base - before.dry_base).abs() < 1e-6);
assert_eq!(after.penalty_last_n, before.penalty_last_n);
assert_eq!(after.seed, before.seed);
}
#[test]
fn sampler_type_names_match_upstream() {
assert_eq!(CommonSamplerType::TopK.name().unwrap(), "top_k");
assert_eq!(CommonSamplerType::TopP.name().unwrap(), "top_p");
assert_eq!(
CommonSamplerType::Temperature.name().unwrap(),
"temperature"
);
}
#[test]
fn sampler_type_discriminants_round_trip_through_names() {
for ty in [
CommonSamplerType::Dry,
CommonSamplerType::TopK,
CommonSamplerType::TopP,
CommonSamplerType::MinP,
CommonSamplerType::TypicalP,
CommonSamplerType::Temperature,
CommonSamplerType::Xtc,
CommonSamplerType::Infill,
CommonSamplerType::Penalties,
CommonSamplerType::TopNSigma,
CommonSamplerType::AdaptiveP,
] {
let name = ty.name().unwrap();
let parsed = CommonSamplerType::from_names(&[&name]).unwrap();
assert_eq!(parsed, vec![ty as i32], "{name} did not round-trip");
}
}
#[test]
fn unknown_sampler_names_are_dropped() {
let parsed = CommonSamplerType::from_names(&["top_k", "not_a_sampler"]).unwrap();
assert_eq!(parsed, vec![CommonSamplerType::TopK as i32]);
}
#[test]
fn grammar_trigger_rejects_unknown_kind() {
let mut params = CommonSamplerParams::new();
assert!(matches!(
params.add_grammar_trigger("nonsense", "x", LlamaToken(-1)),
Err(CommonSamplerError::InvalidArg)
));
}
#[test]
fn grammar_trigger_accepts_every_known_kind() {
let mut params = CommonSamplerParams::new();
for kind in ["token", "word", "pattern", "pattern_full"] {
params
.add_grammar_trigger(kind, "<tool_call>", LlamaToken(1))
.unwrap_or_else(|e| panic!("{kind} rejected: {e}"));
}
}
#[test]
fn interior_nul_is_rejected_not_truncated() {
let mut params = CommonSamplerParams::new();
assert!(matches!(
params.set_grammar("root ::= \0 \"a\"", GrammarSource::User, false),
Err(CommonSamplerError::Nul(_))
));
assert!(matches!(
params.set_generation_prompt("a\0b"),
Err(CommonSamplerError::Nul(_))
));
}
#[test]
fn flatten_produces_matching_data_and_lengths() {
let seqs = vec![
vec![LlamaToken(1), LlamaToken(2)],
vec![],
vec![LlamaToken(3)],
];
let (data, lens) = flatten(&seqs);
assert_eq!(data, vec![1, 2, 3]);
assert_eq!(lens, vec![2, 0, 1]);
}
}