Skip to main content

llama_cpp_bindings/
sampling.rs

1use std::borrow::Borrow;
2use std::ffi::{CString, c_char};
3use std::fmt::{Debug, Formatter};
4
5use llama_cpp_error_recorder::ErrorScope;
6use llama_cpp_error_recorder::RecordedError;
7
8use crate::context::LlamaContext;
9use crate::ffi_error_reader::read_and_free_cpp_error;
10use crate::model::LlamaModel;
11use crate::token::LlamaToken;
12use crate::token::data_array::LlamaTokenDataArray;
13use crate::token::logit_bias::LlamaLogitBias;
14use crate::{GrammarError, SampleError, SamplerAcceptError, SamplingError};
15
16fn check_sampler_accept_status(
17    status: llama_cpp_bindings_sys::llama_rs_sampler_accept_status,
18    error_ptr: *mut c_char,
19) -> Result<(), SamplerAcceptError> {
20    match status {
21        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_ACCEPT_OK => Ok(()),
22        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_ACCEPT_ERROR_STRING_ALLOCATION_FAILED => {
23            Err(SamplerAcceptError::NotEnoughMemory)
24        }
25        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_ACCEPT_VENDORED_THREW_CXX_EXCEPTION => {
26            let message = unsafe { read_and_free_cpp_error(error_ptr) };
27            Err(SamplerAcceptError::GrammarStateCorrupted { message })
28        }
29        other => unreachable!("llama_rs_sampler_accept returned unrecognized status {other}"),
30    }
31}
32
33fn sampler_sample_status_to_result(
34    status: llama_cpp_bindings_sys::llama_rs_sampler_sample_status,
35    token: i32,
36    error_ptr: *mut c_char,
37) -> Result<LlamaToken, SampleError> {
38    match status {
39        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_SAMPLE_OK => Ok(LlamaToken(token)),
40        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_SAMPLE_ERROR_STRING_ALLOCATION_FAILED => {
41            Err(SampleError::NotEnoughMemory)
42        }
43        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_SAMPLE_VENDORED_THREW_CXX_EXCEPTION => {
44            let message = unsafe { read_and_free_cpp_error(error_ptr) };
45            Err(SampleError::Reported { message })
46        }
47        other => unreachable!("llama_rs_sampler_sample returned unrecognized status {other}"),
48    }
49}
50
51fn sampler_init_grammar_status_to_result(
52    status: llama_cpp_bindings_sys::llama_rs_sampler_init_grammar_status,
53    sampler: *mut llama_cpp_bindings_sys::llama_sampler,
54    error_ptr: *mut c_char,
55) -> Result<LlamaSampler, GrammarError> {
56    match status {
57        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_OK => Ok(LlamaSampler { sampler }),
58        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_VENDORED_RETURNED_NULL => {
59            Err(GrammarError::GrammarMalformed)
60        }
61        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_ERROR_STRING_ALLOCATION_FAILED => {
62            Err(GrammarError::NotEnoughMemory)
63        }
64        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_VENDORED_THREW_CXX_EXCEPTION => {
65            let message = unsafe { read_and_free_cpp_error(error_ptr) };
66            Err(GrammarError::Reported { message })
67        }
68        other => {
69            unreachable!("llama_rs_sampler_init_grammar returned unrecognized status {other}")
70        }
71    }
72}
73
74fn sampler_init_grammar_lazy_status_to_result(
75    status: llama_cpp_bindings_sys::llama_rs_sampler_init_grammar_lazy_status,
76    sampler: *mut llama_cpp_bindings_sys::llama_sampler,
77    error_ptr: *mut c_char,
78) -> Result<LlamaSampler, GrammarError> {
79    match status {
80        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_OK => {
81            Ok(LlamaSampler { sampler })
82        }
83        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_VENDORED_RETURNED_NULL => {
84            Err(GrammarError::LazyGrammarMalformed)
85        }
86        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_ERROR_STRING_ALLOCATION_FAILED => {
87            Err(GrammarError::NotEnoughMemory)
88        }
89        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_VENDORED_THREW_CXX_EXCEPTION => {
90            let message = unsafe { read_and_free_cpp_error(error_ptr) };
91            Err(GrammarError::Reported { message })
92        }
93        other => {
94            unreachable!("llama_rs_sampler_init_grammar_lazy returned unrecognized status {other}")
95        }
96    }
97}
98
99fn sampler_init_grammar_lazy_patterns_status_to_result(
100    status: llama_cpp_bindings_sys::llama_rs_sampler_init_grammar_lazy_patterns_status,
101    sampler: *mut llama_cpp_bindings_sys::llama_sampler,
102    error_ptr: *mut c_char,
103) -> Result<LlamaSampler, GrammarError> {
104    match status {
105        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_OK => {
106            Ok(LlamaSampler { sampler })
107        }
108        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_VENDORED_RETURNED_NULL => {
109            Err(GrammarError::LazyPatternsGrammarMalformed)
110        }
111        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_ERROR_STRING_ALLOCATION_FAILED => {
112            Err(GrammarError::NotEnoughMemory)
113        }
114        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_INVALID_TRIGGER_PATTERN => {
115            let message = unsafe { read_and_free_cpp_error(error_ptr) };
116            Err(GrammarError::InvalidTriggerPattern { message })
117        }
118        llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_VENDORED_THREW_CXX_EXCEPTION => {
119            let message = unsafe { read_and_free_cpp_error(error_ptr) };
120            Err(GrammarError::Reported { message })
121        }
122        other => unreachable!(
123            "llama_rs_sampler_init_grammar_lazy_patterns returned unrecognized status {other}"
124        ),
125    }
126}
127
128fn n_ctx_train_overflow_to_grammar_error(convert_error: std::num::TryFromIntError) -> GrammarError {
129    GrammarError::IntegerOverflow(format!(
130        "n_ctx_train does not fit into u32: {convert_error}"
131    ))
132}
133
134fn checked_u32_as_i32(value: u32) -> Result<i32, GrammarError> {
135    i32::try_from(value).map_err(|convert_error| {
136        GrammarError::IntegerOverflow(format!("value exceeds i32::MAX: {convert_error}"))
137    })
138}
139
140fn checked_usize_as_i32_sampling(value: usize) -> Result<i32, SamplingError> {
141    i32::try_from(value).map_err(|convert_error| {
142        SamplingError::IntegerOverflow(format!("value exceeds i32::MAX: {convert_error}"))
143    })
144}
145
146pub struct LlamaSampler {
147    pub sampler: *mut llama_cpp_bindings_sys::llama_sampler,
148}
149
150fn grammar_callback_error_to_result(error: Option<RecordedError>) -> Result<(), SampleError> {
151    error.map_or(Ok(()), |recorded| {
152        Err(SampleError::GrammarCallbackFailed {
153            message: recorded.into_message(),
154        })
155    })
156}
157
158fn grammar_callback_error_to_accept_result(
159    error: Option<RecordedError>,
160) -> Result<(), SamplerAcceptError> {
161    error.map_or(Ok(()), |recorded| {
162        Err(SamplerAcceptError::GrammarCallbackFailed {
163            message: recorded.into_message(),
164        })
165    })
166}
167
168impl Debug for LlamaSampler {
169    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
170        f.debug_struct("LlamaSamplerChain").finish()
171    }
172}
173
174impl LlamaSampler {
175    /// # Errors
176    ///
177    /// Returns [`SampleError`] if the C++ sampler throws an exception, the index is invalid, or the
178    /// grammar sampler callback recorded a failure during sampling.
179    pub fn sample(&mut self, ctx: &LlamaContext, idx: i32) -> Result<LlamaToken, SampleError> {
180        let mut token: i32 = -1;
181        let mut error_ptr: *mut c_char = std::ptr::null_mut();
182
183        let scope = ErrorScope::enter();
184        let status = unsafe {
185            llama_cpp_bindings_sys::llama_rs_sampler_sample(
186                self.sampler,
187                ctx.context.as_ptr(),
188                idx,
189                &raw mut token,
190                &raw mut error_ptr,
191            )
192        };
193        grammar_callback_error_to_result(scope.take())?;
194
195        sampler_sample_status_to_result(status, token, error_ptr)
196    }
197
198    /// # Errors
199    ///
200    /// Returns [`SampleError`] if the grammar sampler callback recorded a failure during application.
201    pub fn apply(&self, data_array: &mut LlamaTokenDataArray) -> Result<(), SampleError> {
202        let scope = ErrorScope::enter();
203        data_array.apply_sampler(self)?;
204
205        grammar_callback_error_to_result(scope.take())
206    }
207
208    /// # Errors
209    /// Returns [`SamplerAcceptError`] if the underlying sampler rejects the token.
210    pub fn accept(&mut self, token: LlamaToken) -> Result<(), SamplerAcceptError> {
211        self.try_accept(token)
212    }
213
214    /// # Errors
215    /// Returns [`SamplerAcceptError`] if the underlying sampler rejects any token.
216    pub fn accept_many(
217        &mut self,
218        tokens: impl IntoIterator<Item = impl Borrow<LlamaToken>>,
219    ) -> Result<(), SamplerAcceptError> {
220        for token in tokens {
221            self.try_accept(*token.borrow())?;
222        }
223
224        Ok(())
225    }
226
227    /// # Errors
228    /// Returns [`SamplerAcceptError`] if the underlying sampler rejects any token.
229    pub fn with_tokens(
230        mut self,
231        tokens: impl IntoIterator<Item = impl Borrow<LlamaToken>>,
232    ) -> Result<Self, SamplerAcceptError> {
233        self.accept_many(tokens)?;
234
235        Ok(self)
236    }
237
238    /// # Errors
239    /// Returns an error if the underlying sampler rejects the token.
240    pub fn try_accept(&mut self, token: LlamaToken) -> Result<(), SamplerAcceptError> {
241        let mut error_ptr: *mut c_char = std::ptr::null_mut();
242
243        let scope = ErrorScope::enter();
244        let status = unsafe {
245            llama_cpp_bindings_sys::llama_rs_sampler_accept(
246                self.sampler,
247                token.0,
248                &raw mut error_ptr,
249            )
250        };
251        grammar_callback_error_to_accept_result(scope.take())?;
252
253        check_sampler_accept_status(status, error_ptr)
254    }
255
256    /// # Errors
257    ///
258    /// Returns [`SampleError`] if the grammar sampler callback recorded a failure during reset.
259    pub fn reset(&mut self) -> Result<(), SampleError> {
260        let scope = ErrorScope::enter();
261        unsafe {
262            llama_cpp_bindings_sys::llama_sampler_reset(self.sampler);
263        }
264
265        grammar_callback_error_to_result(scope.take())
266    }
267
268    #[must_use]
269    pub fn get_seed(&self) -> u32 {
270        unsafe { llama_cpp_bindings_sys::llama_sampler_get_seed(self.sampler) }
271    }
272
273    #[must_use]
274    pub fn chain(samplers: impl IntoIterator<Item = Self>, no_perf: bool) -> Self {
275        unsafe {
276            let chain = llama_cpp_bindings_sys::llama_sampler_chain_init(
277                llama_cpp_bindings_sys::llama_sampler_chain_params { no_perf },
278            );
279
280            for sampler in samplers {
281                llama_cpp_bindings_sys::llama_sampler_chain_add(chain, sampler.sampler);
282                std::mem::forget(sampler);
283            }
284
285            Self { sampler: chain }
286        }
287    }
288
289    #[must_use]
290    pub fn chain_simple(samplers: impl IntoIterator<Item = Self>) -> Self {
291        Self::chain(samplers, false)
292    }
293
294    #[must_use]
295    pub fn temp(t: f32) -> Self {
296        let sampler = unsafe { llama_cpp_bindings_sys::llama_sampler_init_temp(t) };
297        Self { sampler }
298    }
299
300    #[must_use]
301    pub fn temp_ext(t: f32, delta: f32, exponent: f32) -> Self {
302        let sampler =
303            unsafe { llama_cpp_bindings_sys::llama_sampler_init_temp_ext(t, delta, exponent) };
304        Self { sampler }
305    }
306
307    #[must_use]
308    pub fn top_k(k: i32) -> Self {
309        let sampler = unsafe { llama_cpp_bindings_sys::llama_sampler_init_top_k(k) };
310        Self { sampler }
311    }
312
313    #[must_use]
314    pub fn top_n_sigma(n: f32) -> Self {
315        let sampler = unsafe { llama_cpp_bindings_sys::llama_sampler_init_top_n_sigma(n) };
316        Self { sampler }
317    }
318
319    #[must_use]
320    pub fn typical(p: f32, min_keep: usize) -> Self {
321        let sampler = unsafe { llama_cpp_bindings_sys::llama_sampler_init_typical(p, min_keep) };
322        Self { sampler }
323    }
324
325    #[must_use]
326    pub fn top_p(p: f32, min_keep: usize) -> Self {
327        let sampler = unsafe { llama_cpp_bindings_sys::llama_sampler_init_top_p(p, min_keep) };
328        Self { sampler }
329    }
330
331    #[must_use]
332    pub fn min_p(p: f32, min_keep: usize) -> Self {
333        let sampler = unsafe { llama_cpp_bindings_sys::llama_sampler_init_min_p(p, min_keep) };
334        Self { sampler }
335    }
336
337    #[must_use]
338    pub fn xtc(p: f32, t: f32, min_keep: usize, seed: u32) -> Self {
339        let sampler =
340            unsafe { llama_cpp_bindings_sys::llama_sampler_init_xtc(p, t, min_keep, seed) };
341        Self { sampler }
342    }
343
344    /// # Errors
345    /// Returns an error if the grammar is invalid or the sampler cannot be initialized.
346    pub fn grammar(
347        model: &LlamaModel,
348        grammar_str: &str,
349        grammar_root: &str,
350    ) -> Result<Self, GrammarError> {
351        let (grammar_str, grammar_root) =
352            Self::sanitize_grammar_strings(grammar_str, grammar_root)?;
353        let mut sampler: *mut llama_cpp_bindings_sys::llama_sampler = std::ptr::null_mut();
354        let mut error_ptr: *mut c_char = std::ptr::null_mut();
355
356        let status = unsafe {
357            llama_cpp_bindings_sys::llama_rs_sampler_init_grammar(
358                model.vocab_ptr(),
359                grammar_str.as_ptr(),
360                grammar_root.as_ptr(),
361                &raw mut sampler,
362                &raw mut error_ptr,
363            )
364        };
365
366        sampler_init_grammar_status_to_result(status, sampler, error_ptr)
367    }
368
369    /// # Errors
370    /// Returns an error if the grammar or trigger words are invalid.
371    pub fn grammar_lazy(
372        model: &LlamaModel,
373        grammar_str: &str,
374        grammar_root: &str,
375        trigger_words: impl IntoIterator<Item = impl AsRef<[u8]>>,
376        trigger_tokens: &[LlamaToken],
377    ) -> Result<Self, GrammarError> {
378        let (grammar_str, grammar_root) =
379            Self::sanitize_grammar_strings(grammar_str, grammar_root)?;
380        let trigger_words = Self::sanitize_trigger_words(trigger_words)?;
381        let mut sampler: *mut llama_cpp_bindings_sys::llama_sampler = std::ptr::null_mut();
382        let mut error_ptr: *mut c_char = std::ptr::null_mut();
383
384        let mut trigger_word_ptrs: Vec<*const c_char> =
385            trigger_words.iter().map(|cs| cs.as_ptr()).collect();
386
387        let status = unsafe {
388            llama_cpp_bindings_sys::llama_rs_sampler_init_grammar_lazy(
389                model.vocab_ptr(),
390                grammar_str.as_ptr(),
391                grammar_root.as_ptr(),
392                trigger_word_ptrs.as_mut_ptr(),
393                trigger_word_ptrs.len(),
394                trigger_tokens.as_ptr().cast(),
395                trigger_tokens.len(),
396                &raw mut sampler,
397                &raw mut error_ptr,
398            )
399        };
400
401        sampler_init_grammar_lazy_status_to_result(status, sampler, error_ptr)
402    }
403
404    /// # Errors
405    /// Returns an error if the grammar or trigger patterns are invalid.
406    pub fn grammar_lazy_patterns(
407        model: &LlamaModel,
408        grammar_str: &str,
409        grammar_root: &str,
410        trigger_patterns: &[String],
411        trigger_tokens: &[LlamaToken],
412    ) -> Result<Self, GrammarError> {
413        let (grammar_str, grammar_root) =
414            Self::sanitize_grammar_strings(grammar_str, grammar_root)?;
415        let trigger_patterns = Self::sanitize_trigger_patterns(trigger_patterns)?;
416        let mut sampler: *mut llama_cpp_bindings_sys::llama_sampler = std::ptr::null_mut();
417        let mut error_ptr: *mut c_char = std::ptr::null_mut();
418
419        let mut trigger_pattern_ptrs: Vec<*const c_char> =
420            trigger_patterns.iter().map(|cs| cs.as_ptr()).collect();
421
422        let status = unsafe {
423            llama_cpp_bindings_sys::llama_rs_sampler_init_grammar_lazy_patterns(
424                model.vocab_ptr(),
425                grammar_str.as_ptr(),
426                grammar_root.as_ptr(),
427                trigger_pattern_ptrs.as_mut_ptr(),
428                trigger_pattern_ptrs.len(),
429                trigger_tokens.as_ptr().cast(),
430                trigger_tokens.len(),
431                &raw mut sampler,
432                &raw mut error_ptr,
433            )
434        };
435
436        sampler_init_grammar_lazy_patterns_status_to_result(status, sampler, error_ptr)
437    }
438
439    /// # Errors
440    ///
441    /// Returns [`GrammarError`] if the grammar is invalid or the sampler cannot be initialized.
442    pub fn llguidance(
443        model: &LlamaModel,
444        grammar_kind: &str,
445        grammar_data: &str,
446    ) -> Result<Self, GrammarError> {
447        crate::llguidance_sampler::create_llg_sampler(model, grammar_kind, grammar_data)
448    }
449
450    fn sanitize_grammar_strings(
451        grammar_str: &str,
452        grammar_root: &str,
453    ) -> Result<(CString, CString), GrammarError> {
454        if !grammar_str.contains(grammar_root) {
455            return Err(GrammarError::RootNotFound);
456        }
457
458        let grammar = CString::new(grammar_str).map_err(GrammarError::GrammarNullBytes)?;
459        let root = CString::new(grammar_root).map_err(GrammarError::GrammarNullBytes)?;
460
461        Ok((grammar, root))
462    }
463
464    fn sanitize_trigger_words(
465        trigger_words: impl IntoIterator<Item = impl AsRef<[u8]>>,
466    ) -> Result<Vec<CString>, GrammarError> {
467        trigger_words
468            .into_iter()
469            .map(|word| CString::new(word.as_ref()).map_err(GrammarError::TriggerWordNullBytes))
470            .collect()
471    }
472
473    fn sanitize_trigger_patterns(
474        trigger_patterns: &[String],
475    ) -> Result<Vec<CString>, GrammarError> {
476        trigger_patterns
477            .iter()
478            .map(|pattern| CString::new(pattern.as_str()).map_err(GrammarError::GrammarNullBytes))
479            .collect()
480    }
481
482    /// # Errors
483    /// Returns an error if any string in `seq_breakers` contains null bytes.
484    pub fn dry(
485        model: &LlamaModel,
486        multiplier: f32,
487        base: f32,
488        allowed_length: i32,
489        penalty_last_n: i32,
490        seq_breakers: impl IntoIterator<Item = impl AsRef<[u8]>>,
491    ) -> Result<Self, GrammarError> {
492        let seq_breakers: Vec<CString> = seq_breakers
493            .into_iter()
494            .map(|seq_breaker| CString::new(seq_breaker.as_ref()))
495            .collect::<Result<Vec<_>, _>>()?;
496        let mut seq_breaker_pointers: Vec<*const c_char> = seq_breakers
497            .iter()
498            .map(|seq_breaker| seq_breaker.as_ptr())
499            .collect();
500
501        let n_ctx_train_value = model
502            .n_ctx_train()
503            .map_err(n_ctx_train_overflow_to_grammar_error)?;
504        let n_ctx_train = checked_u32_as_i32(n_ctx_train_value)?;
505        let sampler = unsafe {
506            llama_cpp_bindings_sys::llama_sampler_init_dry(
507                model.vocab_ptr(),
508                n_ctx_train,
509                multiplier,
510                base,
511                allowed_length,
512                penalty_last_n,
513                seq_breaker_pointers.as_mut_ptr(),
514                seq_breaker_pointers.len(),
515            )
516        };
517
518        Ok(Self { sampler })
519    }
520
521    #[must_use]
522    pub fn penalties(
523        penalty_last_n: i32,
524        penalty_repeat: f32,
525        penalty_freq: f32,
526        penalty_present: f32,
527    ) -> Self {
528        let sampler = unsafe {
529            llama_cpp_bindings_sys::llama_sampler_init_penalties(
530                penalty_last_n,
531                penalty_repeat,
532                penalty_freq,
533                penalty_present,
534            )
535        };
536        Self { sampler }
537    }
538
539    #[must_use]
540    pub fn mirostat(n_vocab: i32, seed: u32, tau: f32, eta: f32, m: i32) -> Self {
541        let sampler = unsafe {
542            llama_cpp_bindings_sys::llama_sampler_init_mirostat(n_vocab, seed, tau, eta, m)
543        };
544        Self { sampler }
545    }
546
547    #[must_use]
548    pub fn mirostat_v2(seed: u32, tau: f32, eta: f32) -> Self {
549        let sampler =
550            unsafe { llama_cpp_bindings_sys::llama_sampler_init_mirostat_v2(seed, tau, eta) };
551        Self { sampler }
552    }
553
554    #[must_use]
555    pub fn dist(seed: u32) -> Self {
556        let sampler = unsafe { llama_cpp_bindings_sys::llama_sampler_init_dist(seed) };
557        Self { sampler }
558    }
559
560    #[must_use]
561    pub fn greedy() -> Self {
562        let sampler = unsafe { llama_cpp_bindings_sys::llama_sampler_init_greedy() };
563        Self { sampler }
564    }
565
566    /// # Errors
567    /// Returns [`SamplingError::IntegerOverflow`] if `biases.len()` exceeds `i32::MAX`.
568    ///
569    pub fn logit_bias(n_vocab: i32, biases: &[LlamaLogitBias]) -> Result<Self, SamplingError> {
570        let bias_count = checked_usize_as_i32_sampling(biases.len())?;
571        let data = biases
572            .as_ptr()
573            .cast::<llama_cpp_bindings_sys::llama_logit_bias>();
574
575        let sampler = unsafe {
576            llama_cpp_bindings_sys::llama_sampler_init_logit_bias(n_vocab, bias_count, data)
577        };
578
579        Ok(Self { sampler })
580    }
581}
582
583impl Drop for LlamaSampler {
584    fn drop(&mut self) {
585        unsafe {
586            llama_cpp_bindings_sys::llama_sampler_free(self.sampler);
587        }
588    }
589}
590
591#[cfg(test)]
592mod tests {
593    use std::ffi::CString;
594    use std::mem::Discriminant;
595
596    use llama_cpp_error_recorder::RecordedError;
597
598    use super::LlamaSampler;
599    use super::grammar_callback_error_to_accept_result;
600    use super::grammar_callback_error_to_result;
601    use crate::GrammarError;
602    use crate::SampleError;
603    use crate::SamplerAcceptError;
604
605    #[test]
606    fn grammar_callback_error_to_result_maps_recorded_error() {
607        let result =
608            grammar_callback_error_to_result(Some(RecordedError::new("mask failed".to_string())));
609
610        assert_eq!(
611            result.unwrap_err(),
612            SampleError::GrammarCallbackFailed {
613                message: "mask failed".to_string()
614            }
615        );
616    }
617
618    #[test]
619    fn grammar_callback_error_to_result_maps_absence_to_ok() {
620        assert!(grammar_callback_error_to_result(None).is_ok());
621    }
622
623    #[test]
624    fn grammar_callback_error_to_accept_result_maps_recorded_error() {
625        let result = grammar_callback_error_to_accept_result(Some(RecordedError::new(
626            "consume failed".to_string(),
627        )));
628
629        assert_eq!(
630            result,
631            Err(SamplerAcceptError::GrammarCallbackFailed {
632                message: "consume failed".to_string()
633            })
634        );
635    }
636
637    #[test]
638    fn grammar_callback_error_to_accept_result_maps_absence_to_ok() {
639        assert!(grammar_callback_error_to_accept_result(None).is_ok());
640    }
641
642    fn nul_error() -> std::ffi::NulError {
643        CString::new(b"a\0b".to_vec()).unwrap_err()
644    }
645
646    fn root_not_found_disc() -> Discriminant<GrammarError> {
647        std::mem::discriminant(&GrammarError::RootNotFound)
648    }
649
650    fn grammar_null_bytes_disc() -> Discriminant<GrammarError> {
651        std::mem::discriminant(&GrammarError::GrammarNullBytes(nul_error()))
652    }
653
654    fn trigger_word_null_bytes_disc() -> Discriminant<GrammarError> {
655        std::mem::discriminant(&GrammarError::TriggerWordNullBytes(nul_error()))
656    }
657
658    #[test]
659    fn sanitize_grammar_strings_valid() {
660        let result = LlamaSampler::sanitize_grammar_strings("root ::= \"hello\"", "root");
661
662        assert!(result.is_ok());
663    }
664
665    #[test]
666    fn sanitize_grammar_strings_root_not_found() {
667        let err = LlamaSampler::sanitize_grammar_strings("expr ::= \"hello\"", "root").unwrap_err();
668
669        assert_eq!(std::mem::discriminant(&err), root_not_found_disc());
670    }
671
672    #[test]
673    fn sanitize_grammar_strings_null_byte_in_grammar() {
674        let err = LlamaSampler::sanitize_grammar_strings("root ::= \"\0\"", "root").unwrap_err();
675
676        assert_eq!(std::mem::discriminant(&err), grammar_null_bytes_disc());
677    }
678
679    #[test]
680    fn sanitize_grammar_strings_null_byte_in_root() {
681        let err =
682            LlamaSampler::sanitize_grammar_strings("ro\0ot ::= \"hello\"", "ro\0ot").unwrap_err();
683
684        assert_eq!(std::mem::discriminant(&err), grammar_null_bytes_disc());
685    }
686
687    #[test]
688    fn sanitize_trigger_words_valid() {
689        let words: Vec<&[u8]> = vec![b"hello", b"world"];
690        let result = LlamaSampler::sanitize_trigger_words(words);
691
692        assert!(result.is_ok());
693        assert_eq!(result.expect("valid trigger words").len(), 2);
694    }
695
696    #[test]
697    fn sanitize_trigger_words_empty_list() {
698        let words: Vec<&[u8]> = vec![];
699        let result = LlamaSampler::sanitize_trigger_words(words);
700
701        assert!(result.is_ok());
702        assert!(result.expect("valid trigger words").is_empty());
703    }
704
705    #[test]
706    fn sanitize_trigger_words_null_byte() {
707        let words: Vec<&[u8]> = vec![b"hel\0lo"];
708        let err = LlamaSampler::sanitize_trigger_words(words).unwrap_err();
709
710        assert_eq!(std::mem::discriminant(&err), trigger_word_null_bytes_disc());
711    }
712
713    #[test]
714    fn sanitize_trigger_patterns_valid() {
715        let patterns = vec!["^hello$".to_string(), "world.*".to_string()];
716        let result = LlamaSampler::sanitize_trigger_patterns(&patterns);
717
718        assert!(result.is_ok());
719        assert_eq!(result.expect("valid trigger patterns").len(), 2);
720    }
721
722    #[test]
723    fn sanitize_trigger_patterns_empty_list() {
724        let patterns: Vec<String> = vec![];
725        let result = LlamaSampler::sanitize_trigger_patterns(&patterns);
726
727        assert!(result.is_ok());
728        assert!(result.expect("valid trigger patterns").is_empty());
729    }
730
731    #[test]
732    fn sanitize_trigger_patterns_null_byte() {
733        let patterns = vec!["hel\0lo".to_string()];
734        let err = LlamaSampler::sanitize_trigger_patterns(&patterns).unwrap_err();
735
736        assert_eq!(std::mem::discriminant(&err), grammar_null_bytes_disc());
737    }
738
739    #[test]
740    fn apply_modifies_data_array() {
741        use crate::token::LlamaToken;
742        use crate::token::data::LlamaTokenData;
743        use crate::token::data_array::LlamaTokenDataArray;
744
745        let sampler = LlamaSampler::greedy();
746        let mut data_array = LlamaTokenDataArray::new(
747            vec![
748                LlamaTokenData::new(LlamaToken::new(0), 1.0, 0.0),
749                LlamaTokenData::new(LlamaToken::new(1), 5.0, 0.0),
750            ],
751            false,
752        );
753
754        assert!(sampler.apply(&mut data_array).is_ok());
755
756        assert_eq!(data_array.selected_token(), Some(LlamaToken::new(1)));
757    }
758
759    #[test]
760    fn apply_with_null_sampler_surfaces_sampler_apply_error() {
761        use crate::error::SampleError;
762        use crate::error::SamplerApplyError;
763        use crate::token::LlamaToken;
764        use crate::token::data::LlamaTokenData;
765        use crate::token::data_array::LlamaTokenDataArray;
766
767        let null_sampler = LlamaSampler {
768            sampler: std::ptr::null_mut(),
769        };
770        let mut data_array = LlamaTokenDataArray::new(
771            vec![LlamaTokenData::new(LlamaToken::new(0), 1.0, 0.0)],
772            false,
773        );
774
775        assert_eq!(
776            null_sampler.apply(&mut data_array),
777            Err(SampleError::SamplerApply(SamplerApplyError::NullSampler)),
778        );
779    }
780
781    #[test]
782    fn accept_succeeds() {
783        let mut sampler = LlamaSampler::chain_simple([
784            LlamaSampler::penalties(64, 1.1, 0.0, 0.0),
785            LlamaSampler::greedy(),
786        ]);
787
788        sampler
789            .accept(crate::token::LlamaToken::new(1))
790            .expect("test: accept should succeed");
791    }
792
793    #[test]
794    fn try_accept_succeeds_on_penalties_sampler() {
795        let mut sampler = LlamaSampler::chain_simple([
796            LlamaSampler::penalties(64, 1.1, 0.0, 0.0),
797            LlamaSampler::greedy(),
798        ]);
799
800        let result = sampler.try_accept(crate::token::LlamaToken::new(42));
801
802        assert!(result.is_ok());
803    }
804
805    #[test]
806    fn accept_many_multiple_tokens() {
807        use crate::token::LlamaToken;
808
809        let mut sampler = LlamaSampler::chain_simple([
810            LlamaSampler::penalties(64, 1.1, 0.0, 0.0),
811            LlamaSampler::greedy(),
812        ]);
813
814        sampler
815            .accept_many([LlamaToken::new(1), LlamaToken::new(2), LlamaToken::new(3)])
816            .expect("test: accept_many should succeed");
817    }
818
819    #[test]
820    fn with_tokens_builder_pattern() {
821        use crate::token::LlamaToken;
822
823        let _sampler = LlamaSampler::chain_simple([
824            LlamaSampler::penalties(64, 1.1, 0.0, 0.0),
825            LlamaSampler::greedy(),
826        ])
827        .with_tokens([LlamaToken::new(10), LlamaToken::new(20)])
828        .expect("test: with_tokens should succeed");
829    }
830
831    #[test]
832    fn all_sampler_constructors() {
833        use crate::token::LlamaToken;
834        use crate::token::logit_bias::LlamaLogitBias;
835
836        let _temp = LlamaSampler::temp(0.8);
837        let _temp_ext = LlamaSampler::temp_ext(0.8, 0.1, 1.0);
838        let _top_k = LlamaSampler::top_k(40);
839        let _top_n_sigma = LlamaSampler::top_n_sigma(2.0);
840        let _top_p = LlamaSampler::top_p(0.9, 1);
841        let _min_p = LlamaSampler::min_p(0.05, 1);
842        let _typical = LlamaSampler::typical(0.9, 1);
843        let _xtc = LlamaSampler::xtc(0.1, 0.5, 1, 42);
844        let _dist = LlamaSampler::dist(42);
845        let _mirostat = LlamaSampler::mirostat(32000, 42, 5.0, 0.1, 100);
846        let _mirostat_v2 = LlamaSampler::mirostat_v2(42, 5.0, 0.1);
847        let biases = vec![LlamaLogitBias::new(LlamaToken::new(0), -100.0)];
848        let _logit_bias = LlamaSampler::logit_bias(32000, &biases);
849        let _chain = LlamaSampler::chain([LlamaSampler::greedy()], true);
850    }
851
852    #[test]
853    fn reset_and_get_seed() {
854        let mut sampler = LlamaSampler::dist(42);
855        assert!(sampler.reset().is_ok());
856        let _seed = sampler.get_seed();
857    }
858
859    #[test]
860    fn debug_formatting() {
861        let sampler = LlamaSampler::greedy();
862        let debug_output = format!("{sampler:?}");
863        assert!(debug_output.contains("LlamaSampler"));
864    }
865
866    #[test]
867    fn checked_u32_as_i32_overflow() {
868        let result = super::checked_u32_as_i32(u32::MAX);
869        assert!(result.is_err());
870    }
871
872    #[test]
873    fn checked_usize_as_i32_sampling_overflow() {
874        let result = super::checked_usize_as_i32_sampling(usize::MAX);
875        assert!(result.is_err());
876    }
877
878    #[test]
879    fn check_sampler_accept_status_ok() {
880        let result = super::check_sampler_accept_status(
881            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_ACCEPT_OK,
882            std::ptr::null_mut(),
883        );
884
885        assert!(result.is_ok());
886    }
887
888    #[test]
889    fn check_sampler_accept_status_exception_maps_to_typed_variant() {
890        let err = super::check_sampler_accept_status(
891            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_ACCEPT_VENDORED_THREW_CXX_EXCEPTION,
892            std::ptr::null_mut(),
893        )
894        .unwrap_err();
895        let grammar_state_corrupted_disc =
896            std::mem::discriminant(&SamplerAcceptError::GrammarStateCorrupted {
897                message: String::new(),
898            });
899
900        assert_eq!(std::mem::discriminant(&err), grammar_state_corrupted_disc);
901    }
902
903    #[test]
904    fn check_sampler_accept_status_allocation_failure_maps_to_not_enough_memory() {
905        let result = super::check_sampler_accept_status(
906            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_ACCEPT_ERROR_STRING_ALLOCATION_FAILED,
907            std::ptr::null_mut(),
908        );
909
910        assert_eq!(result, Err(SamplerAcceptError::NotEnoughMemory));
911    }
912
913    #[test]
914    #[should_panic(expected = "llama_rs_sampler_accept returned unrecognized status")]
915    fn check_sampler_accept_status_unrecognized_panics() {
916        let _result = super::check_sampler_accept_status(
917            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_ACCEPT_NULL_SAMPLER_ARG,
918            std::ptr::null_mut(),
919        );
920    }
921
922    #[test]
923    fn sampler_sample_status_allocation_failure_maps_to_not_enough_memory() {
924        let result = super::sampler_sample_status_to_result(
925            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_SAMPLE_ERROR_STRING_ALLOCATION_FAILED,
926            -1,
927            std::ptr::null_mut(),
928        );
929
930        assert_eq!(result.unwrap_err(), SampleError::NotEnoughMemory);
931    }
932
933    #[test]
934    fn sampler_sample_status_exception_maps_to_reported() {
935        let result = super::sampler_sample_status_to_result(
936            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_SAMPLE_VENDORED_THREW_CXX_EXCEPTION,
937            -1,
938            std::ptr::null_mut(),
939        );
940
941        assert_eq!(
942            result.unwrap_err(),
943            SampleError::Reported {
944                message: "unknown error".to_string()
945            }
946        );
947    }
948
949    #[test]
950    #[should_panic(expected = "llama_rs_sampler_sample returned unrecognized status")]
951    fn sampler_sample_status_unrecognized_panics() {
952        let _result = super::sampler_sample_status_to_result(
953            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_SAMPLE_NULL_CTX_ARG,
954            -1,
955            std::ptr::null_mut(),
956        );
957    }
958
959    #[test]
960    fn sampler_init_grammar_status_null_maps_to_grammar_malformed() {
961        let result = super::sampler_init_grammar_status_to_result(
962            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_VENDORED_RETURNED_NULL,
963            std::ptr::null_mut(),
964            std::ptr::null_mut(),
965        );
966
967        assert_eq!(result.unwrap_err(), GrammarError::GrammarMalformed);
968    }
969
970    #[test]
971    fn sampler_init_grammar_status_allocation_failure_maps_to_not_enough_memory() {
972        let result = super::sampler_init_grammar_status_to_result(
973            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_ERROR_STRING_ALLOCATION_FAILED,
974            std::ptr::null_mut(),
975            std::ptr::null_mut(),
976        );
977
978        assert_eq!(result.unwrap_err(), GrammarError::NotEnoughMemory);
979    }
980
981    #[test]
982    fn sampler_init_grammar_status_exception_maps_to_reported() {
983        let result = super::sampler_init_grammar_status_to_result(
984            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_VENDORED_THREW_CXX_EXCEPTION,
985            std::ptr::null_mut(),
986            std::ptr::null_mut(),
987        );
988
989        assert_eq!(
990            result.unwrap_err(),
991            GrammarError::Reported {
992                message: "unknown error".to_string()
993            }
994        );
995    }
996
997    #[test]
998    #[should_panic(expected = "llama_rs_sampler_init_grammar returned unrecognized status")]
999    fn sampler_init_grammar_status_unrecognized_panics() {
1000        let _result = super::sampler_init_grammar_status_to_result(
1001            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_NULL_OUT_SAMPLER_ARG,
1002            std::ptr::null_mut(),
1003            std::ptr::null_mut(),
1004        );
1005    }
1006
1007    #[test]
1008    fn sampler_init_grammar_lazy_status_null_maps_to_lazy_grammar_malformed() {
1009        let result = super::sampler_init_grammar_lazy_status_to_result(
1010            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_VENDORED_RETURNED_NULL,
1011            std::ptr::null_mut(),
1012            std::ptr::null_mut(),
1013        );
1014
1015        assert_eq!(result.unwrap_err(), GrammarError::LazyGrammarMalformed);
1016    }
1017
1018    #[test]
1019    fn sampler_init_grammar_lazy_status_allocation_failure_maps_to_not_enough_memory() {
1020        let result = super::sampler_init_grammar_lazy_status_to_result(
1021            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_ERROR_STRING_ALLOCATION_FAILED,
1022            std::ptr::null_mut(),
1023            std::ptr::null_mut(),
1024        );
1025
1026        assert_eq!(result.unwrap_err(), GrammarError::NotEnoughMemory);
1027    }
1028
1029    #[test]
1030    fn sampler_init_grammar_lazy_status_exception_maps_to_reported() {
1031        let result = super::sampler_init_grammar_lazy_status_to_result(
1032            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_VENDORED_THREW_CXX_EXCEPTION,
1033            std::ptr::null_mut(),
1034            std::ptr::null_mut(),
1035        );
1036
1037        assert_eq!(
1038            result.unwrap_err(),
1039            GrammarError::Reported {
1040                message: "unknown error".to_string()
1041            }
1042        );
1043    }
1044
1045    #[test]
1046    #[should_panic(expected = "llama_rs_sampler_init_grammar_lazy returned unrecognized status")]
1047    fn sampler_init_grammar_lazy_status_unrecognized_panics() {
1048        let _result = super::sampler_init_grammar_lazy_status_to_result(
1049            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_NULL_OUT_SAMPLER_ARG,
1050            std::ptr::null_mut(),
1051            std::ptr::null_mut(),
1052        );
1053    }
1054
1055    #[test]
1056    fn sampler_init_grammar_lazy_patterns_status_null_maps_to_lazy_patterns_grammar_malformed() {
1057        let result = super::sampler_init_grammar_lazy_patterns_status_to_result(
1058            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_VENDORED_RETURNED_NULL,
1059            std::ptr::null_mut(),
1060            std::ptr::null_mut(),
1061        );
1062
1063        assert_eq!(
1064            result.unwrap_err(),
1065            GrammarError::LazyPatternsGrammarMalformed
1066        );
1067    }
1068
1069    #[test]
1070    fn sampler_init_grammar_lazy_patterns_status_allocation_failure_maps_to_not_enough_memory() {
1071        let result = super::sampler_init_grammar_lazy_patterns_status_to_result(
1072            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_ERROR_STRING_ALLOCATION_FAILED,
1073            std::ptr::null_mut(),
1074            std::ptr::null_mut(),
1075        );
1076
1077        assert_eq!(result.unwrap_err(), GrammarError::NotEnoughMemory);
1078    }
1079
1080    #[test]
1081    fn sampler_init_grammar_lazy_patterns_status_exception_maps_to_reported() {
1082        let result = super::sampler_init_grammar_lazy_patterns_status_to_result(
1083            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_VENDORED_THREW_CXX_EXCEPTION,
1084            std::ptr::null_mut(),
1085            std::ptr::null_mut(),
1086        );
1087
1088        assert_eq!(
1089            result.unwrap_err(),
1090            GrammarError::Reported {
1091                message: "unknown error".to_string()
1092            }
1093        );
1094    }
1095
1096    #[test]
1097    #[should_panic(
1098        expected = "llama_rs_sampler_init_grammar_lazy_patterns returned unrecognized status"
1099    )]
1100    fn sampler_init_grammar_lazy_patterns_status_unrecognized_panics() {
1101        let _result = super::sampler_init_grammar_lazy_patterns_status_to_result(
1102            llama_cpp_bindings_sys::LLAMA_RS_SAMPLER_INIT_GRAMMAR_LAZY_PATTERNS_NULL_OUT_SAMPLER_ARG,
1103            std::ptr::null_mut(),
1104            std::ptr::null_mut(),
1105        );
1106    }
1107
1108    #[test]
1109    fn n_ctx_train_overflow_maps_to_integer_overflow() {
1110        let convert_error = u32::try_from(-1_i64).expect_err("-1 cannot convert to u32");
1111        let grammar_error = super::n_ctx_train_overflow_to_grammar_error(convert_error);
1112
1113        assert_eq!(
1114            std::mem::discriminant(&grammar_error),
1115            std::mem::discriminant(&GrammarError::IntegerOverflow(String::new())),
1116        );
1117    }
1118
1119    #[test]
1120    fn grammar_returns_root_not_found_before_touching_model() {
1121        let model = unsafe { &*std::ptr::NonNull::<crate::model::LlamaModel>::dangling().as_ptr() };
1122
1123        let err = LlamaSampler::grammar(model, "expr ::= \"hello\"", "root").unwrap_err();
1124
1125        assert_eq!(err, GrammarError::RootNotFound);
1126    }
1127}