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 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 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 pub fn accept(&mut self, token: LlamaToken) -> Result<(), SamplerAcceptError> {
211 self.try_accept(token)
212 }
213
214 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 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 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 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 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 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 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 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 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 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}