1use std::ffi::c_void;
2use std::fmt::{Debug, Formatter};
3use std::num::NonZeroI32;
4use std::ptr::NonNull;
5use std::slice;
6use std::sync::Arc;
7use std::sync::atomic::AtomicBool;
8use std::sync::atomic::Ordering;
9
10use crate::context::params::LlamaContextParams;
11use crate::llama_backend::LlamaBackend;
12use crate::llama_batch::LlamaBatch;
13use crate::model::{LlamaLoraAdapter, LlamaModel};
14use crate::timing::LlamaTimings;
15use crate::token::LlamaToken;
16use crate::token::data::LlamaTokenData;
17use crate::token::data_array::LlamaTokenDataArray;
18use crate::{
19 DecodeError, EmbeddingsError, EncodeError, LlamaContextLoadError, LlamaLoraAdapterRemoveError,
20 LlamaLoraAdapterSetError, LogitsError,
21};
22
23const fn check_lora_set_result(err_code: i32) -> Result<(), LlamaLoraAdapterSetError> {
24 if err_code != 0 {
25 return Err(LlamaLoraAdapterSetError::ErrorResult(err_code));
26 }
27
28 Ok(())
29}
30
31const fn check_lora_remove_result(err_code: i32) -> Result<(), LlamaLoraAdapterRemoveError> {
32 if err_code != 0 {
33 return Err(LlamaLoraAdapterRemoveError::ErrorResult(err_code));
34 }
35
36 Ok(())
37}
38
39fn new_context_with_model_status_to_result(
40 status: llama_cpp_bindings_sys::llama_rs_new_context_with_model_status,
41 out_ctx: *mut llama_cpp_bindings_sys::llama_context,
42 out_error: *mut std::os::raw::c_char,
43) -> Result<NonNull<llama_cpp_bindings_sys::llama_context>, LlamaContextLoadError> {
44 match status {
45 llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_OK => {
46 NonNull::new(out_ctx).ok_or(LlamaContextLoadError::Unconstructible)
47 }
48 llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_VENDORED_RETURNED_NULL => {
49 Err(LlamaContextLoadError::Unconstructible)
50 }
51 llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_ERROR_STRING_ALLOCATION_FAILED => {
52 Err(LlamaContextLoadError::NotEnoughMemory)
53 }
54 llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_VENDORED_THREW_CXX_EXCEPTION => {
55 let message = unsafe { crate::ffi_error_reader::read_and_free_cpp_error(out_error) };
56 Err(LlamaContextLoadError::Reported { message })
57 }
58 other => {
59 unreachable!("llama_rs_new_context_with_model returned unrecognized status {other}")
60 }
61 }
62}
63
64fn decode_status_to_result(
65 status: llama_cpp_bindings_sys::llama_rs_decode_status,
66 out_vendored_return_code: i32,
67 out_error: *mut std::os::raw::c_char,
68) -> Result<(), DecodeError> {
69 match status {
70 llama_cpp_bindings_sys::LLAMA_RS_DECODE_OK => Ok(()),
71 llama_cpp_bindings_sys::LLAMA_RS_DECODE_VENDORED_RETURNED_NONZERO_CODE => {
72 let code = NonZeroI32::new(out_vendored_return_code).unwrap_or_else(|| {
73 unreachable!(
74 "llama_rs_decode reported a nonzero return code but the value was zero"
75 )
76 });
77 Err(DecodeError::from(code))
78 }
79 llama_cpp_bindings_sys::LLAMA_RS_DECODE_OUT_OF_MEMORY => {
80 Err(DecodeError::DecodeOutOfMemory)
81 }
82 llama_cpp_bindings_sys::LLAMA_RS_DECODE_COMPUTE_FAILED => Err(DecodeError::ComputeFailed),
83 llama_cpp_bindings_sys::LLAMA_RS_DECODE_ERROR_STRING_ALLOCATION_FAILED => {
84 Err(DecodeError::NotEnoughMemory)
85 }
86 llama_cpp_bindings_sys::LLAMA_RS_DECODE_VENDORED_THREW_CXX_EXCEPTION => {
87 let message = unsafe { crate::ffi_error_reader::read_and_free_cpp_error(out_error) };
88 Err(DecodeError::Reported { message })
89 }
90 other => unreachable!("llama_rs_decode returned unrecognized status {other}"),
91 }
92}
93
94fn encode_status_to_result(
95 status: llama_cpp_bindings_sys::llama_rs_encode_status,
96 out_vendored_return_code: i32,
97 out_error: *mut std::os::raw::c_char,
98) -> Result<(), EncodeError> {
99 match status {
100 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_OK => Ok(()),
101 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_MODEL_HAS_NO_ENCODER => {
102 Err(EncodeError::ModelHasNoEncoder)
103 }
104 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_VENDORED_RETURNED_NONZERO_CODE => {
105 let code = NonZeroI32::new(out_vendored_return_code).unwrap_or_else(|| {
106 unreachable!(
107 "llama_rs_encode reported a nonzero return code but the value was zero"
108 )
109 });
110 Err(EncodeError::from(code))
111 }
112 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_OUT_OF_MEMORY => {
113 Err(EncodeError::EncodeOutOfMemory)
114 }
115 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_COMPUTE_FAILED => Err(EncodeError::ComputeFailed),
116 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_ERROR_STRING_ALLOCATION_FAILED => {
117 Err(EncodeError::NotEnoughMemory)
118 }
119 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_VENDORED_THREW_CXX_EXCEPTION => {
120 let message = unsafe { crate::ffi_error_reader::read_and_free_cpp_error(out_error) };
121 Err(EncodeError::Reported { message })
122 }
123 other => unreachable!("llama_rs_encode returned unrecognized status {other}"),
124 }
125}
126
127fn token_index_within_context(token_index: i32, context_size: u32) -> Result<(), LogitsError> {
128 if token_index >= 0 {
129 let token_index_u32 =
130 u32::try_from(token_index).map_err(LogitsError::TokenIndexOverflow)?;
131
132 if context_size <= token_index_u32 {
133 return Err(LogitsError::TokenIndexExceedsContext {
134 token_index: token_index_u32,
135 context_size,
136 });
137 }
138 }
139
140 Ok(())
141}
142
143unsafe fn logits_slice_from_raw_parts<'logits>(
144 data: *const f32,
145 n_vocab: i32,
146) -> Result<&'logits [f32], LogitsError> {
147 if data.is_null() {
148 return Err(LogitsError::NullLogits);
149 }
150
151 let len = usize::try_from(n_vocab).map_err(LogitsError::VocabSizeOverflow)?;
152
153 Ok(unsafe { slice::from_raw_parts(data, len) })
154}
155
156pub mod kv_cache;
157pub mod kv_cache_type;
158pub mod llama_attention_type;
159pub mod llama_pooling_type;
160pub mod llama_state_seq_flags;
161pub mod load_seq_state_error;
162pub mod load_session_error;
163pub mod params;
164pub mod rope_scaling_type;
165pub mod save_seq_state_error;
166pub mod save_session_error;
167pub mod session;
168
169unsafe extern "C" fn abort_callback_trampoline(data: *mut c_void) -> bool {
170 let flag = unsafe { &*(data as *const AtomicBool) };
171
172 flag.load(Ordering::Relaxed)
173}
174
175pub struct LlamaContext<'model> {
176 pub context: NonNull<llama_cpp_bindings_sys::llama_context>,
177 pub model: &'model LlamaModel,
178 abort_flag: Option<Arc<AtomicBool>>,
179 initialized_logits: Vec<i32>,
180 embeddings_enabled: bool,
181}
182
183impl Debug for LlamaContext<'_> {
184 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
185 f.debug_struct("LlamaContext")
186 .field("context", &self.context)
187 .finish()
188 }
189}
190
191impl<'model> LlamaContext<'model> {
192 #[must_use]
193 pub const fn new(
194 llama_model: &'model LlamaModel,
195 llama_context: NonNull<llama_cpp_bindings_sys::llama_context>,
196 embeddings_enabled: bool,
197 ) -> Self {
198 Self {
199 context: llama_context,
200 model: llama_model,
201 abort_flag: None,
202 initialized_logits: Vec::new(),
203 embeddings_enabled,
204 }
205 }
206
207 #[expect(
211 clippy::needless_pass_by_value,
212 reason = "LlamaContextParams may become non-trivially copyable upstream"
213 )]
214 pub fn from_model(
215 model: &'model LlamaModel,
216 _backend: &LlamaBackend,
217 params: LlamaContextParams,
218 ) -> Result<Self, LlamaContextLoadError> {
219 let context_params = params.context_params;
220 let mut out_ctx: *mut llama_cpp_bindings_sys::llama_context = std::ptr::null_mut();
221 let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut();
222 let status = unsafe {
223 llama_cpp_bindings_sys::llama_rs_new_context_with_model(
224 model.model.as_ptr(),
225 context_params,
226 &raw mut out_ctx,
227 &raw mut out_error,
228 )
229 };
230 let context = new_context_with_model_status_to_result(status, out_ctx, out_error)?;
231
232 Ok(Self::new(model, context, params.embeddings()))
233 }
234
235 #[must_use]
236 pub fn n_batch(&self) -> u32 {
237 unsafe { llama_cpp_bindings_sys::llama_n_batch(self.context.as_ptr()) }
238 }
239
240 #[must_use]
241 pub fn n_ubatch(&self) -> u32 {
242 unsafe { llama_cpp_bindings_sys::llama_n_ubatch(self.context.as_ptr()) }
243 }
244
245 #[must_use]
246 pub fn n_ctx(&self) -> u32 {
247 unsafe { llama_cpp_bindings_sys::llama_n_ctx(self.context.as_ptr()) }
248 }
249
250 #[expect(unsafe_code, reason = "required for FFI abort callback registration")]
251 pub fn set_abort_flag(&mut self, flag: Arc<AtomicBool>) {
252 let raw_ptr = Arc::as_ptr(&flag) as *mut c_void;
253 self.abort_flag = Some(flag);
254
255 unsafe {
256 llama_cpp_bindings_sys::llama_set_abort_callback(
257 self.context.as_ptr(),
258 Some(abort_callback_trampoline),
259 raw_ptr,
260 );
261 }
262 }
263
264 #[expect(unsafe_code, reason = "required for FFI abort callback deregistration")]
265 pub fn clear_abort_callback(&mut self) {
266 self.abort_flag = None;
267
268 unsafe {
269 llama_cpp_bindings_sys::llama_set_abort_callback(
270 self.context.as_ptr(),
271 None,
272 std::ptr::null_mut(),
273 );
274 }
275 }
276
277 #[expect(unsafe_code, reason = "required for FFI synchronization call")]
278 pub fn synchronize(&self) {
279 unsafe { llama_cpp_bindings_sys::llama_synchronize(self.context.as_ptr()) }
280 }
281
282 #[expect(unsafe_code, reason = "required for FFI threadpool detachment")]
283 pub fn detach_threadpool(&self) {
284 unsafe { llama_cpp_bindings_sys::llama_detach_threadpool(self.context.as_ptr()) }
285 }
286
287 pub fn mark_logits_initialized(&mut self, token_index: i32) {
288 self.initialized_logits = vec![token_index];
289 }
290
291 pub fn decode(&mut self, batch: &mut LlamaBatch) -> Result<(), DecodeError> {
295 let mut out_vendored_return_code: i32 = 0;
296 let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut();
297 let status = unsafe {
298 llama_cpp_bindings_sys::llama_rs_decode(
299 self.context.as_ptr(),
300 batch.llama_batch,
301 &raw mut out_vendored_return_code,
302 &raw mut out_error,
303 )
304 };
305 decode_status_to_result(status, out_vendored_return_code, out_error)?;
306
307 self.initialized_logits
308 .clone_from(&batch.initialized_logits);
309
310 Ok(())
311 }
312
313 pub fn encode(&mut self, batch: &mut LlamaBatch) -> Result<(), EncodeError> {
317 let mut out_vendored_return_code: i32 = 0;
318 let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut();
319 let status = unsafe {
320 llama_cpp_bindings_sys::llama_rs_encode(
321 self.context.as_ptr(),
322 batch.llama_batch,
323 &raw mut out_vendored_return_code,
324 &raw mut out_error,
325 )
326 };
327 encode_status_to_result(status, out_vendored_return_code, out_error)?;
328
329 self.initialized_logits
330 .clone_from(&batch.initialized_logits);
331
332 Ok(())
333 }
334
335 pub fn embeddings_seq_ith(&self, sequence_index: i32) -> Result<&[f32], EmbeddingsError> {
342 if !self.embeddings_enabled {
343 return Err(EmbeddingsError::NotEnabled);
344 }
345
346 let n_embd = usize::try_from(self.model.n_embd())
347 .map_err(EmbeddingsError::InvalidEmbeddingDimension)?;
348
349 unsafe {
350 let embedding = llama_cpp_bindings_sys::llama_get_embeddings_seq(
351 self.context.as_ptr(),
352 sequence_index,
353 );
354
355 if embedding.is_null() {
356 Err(EmbeddingsError::NonePoolType)
357 } else {
358 Ok(slice::from_raw_parts(embedding, n_embd))
359 }
360 }
361 }
362
363 pub fn embeddings_ith(&self, token_index: i32) -> Result<&[f32], EmbeddingsError> {
370 if !self.embeddings_enabled {
371 return Err(EmbeddingsError::NotEnabled);
372 }
373
374 let n_embd = usize::try_from(self.model.n_embd())
375 .map_err(EmbeddingsError::InvalidEmbeddingDimension)?;
376
377 unsafe {
378 let embedding = llama_cpp_bindings_sys::llama_get_embeddings_ith(
379 self.context.as_ptr(),
380 token_index,
381 );
382
383 if embedding.is_null() {
384 Err(EmbeddingsError::LogitsNotEnabled)
385 } else {
386 Ok(slice::from_raw_parts(embedding, n_embd))
387 }
388 }
389 }
390
391 pub fn candidates(&self) -> Result<impl Iterator<Item = LlamaTokenData> + '_, LogitsError> {
394 let logits = self.get_logits()?;
395
396 Ok((0_i32..).zip(logits).map(|(token_id, logit)| {
397 let token = LlamaToken::new(token_id);
398 LlamaTokenData::new(token, *logit, 0_f32)
399 }))
400 }
401
402 pub fn token_data_array(&self) -> Result<LlamaTokenDataArray, LogitsError> {
405 Ok(LlamaTokenDataArray::from_iter(self.candidates()?, false))
406 }
407
408 pub fn get_logits(&self) -> Result<&[f32], LogitsError> {
411 let data = unsafe { llama_cpp_bindings_sys::llama_get_logits(self.context.as_ptr()) };
412
413 unsafe { logits_slice_from_raw_parts(data, self.model.n_vocab()) }
414 }
415
416 pub fn candidates_ith(
419 &self,
420 token_index: i32,
421 ) -> Result<impl Iterator<Item = LlamaTokenData> + '_, LogitsError> {
422 let logits = self.get_logits_ith(token_index)?;
423
424 Ok((0_i32..).zip(logits).map(|(token_id, logit)| {
425 let token = LlamaToken::new(token_id);
426 LlamaTokenData::new(token, *logit, 0_f32)
427 }))
428 }
429
430 pub fn token_data_array_ith(
433 &self,
434 token_index: i32,
435 ) -> Result<LlamaTokenDataArray, LogitsError> {
436 Ok(LlamaTokenDataArray::from_iter(
437 self.candidates_ith(token_index)?,
438 false,
439 ))
440 }
441
442 pub fn get_logits_ith(&self, token_index: i32) -> Result<&[f32], LogitsError> {
445 if !self.initialized_logits.contains(&token_index) {
446 return Err(LogitsError::TokenNotInitialized(token_index));
447 }
448
449 token_index_within_context(token_index, self.n_ctx())?;
450
451 let data = unsafe {
452 llama_cpp_bindings_sys::llama_get_logits_ith(self.context.as_ptr(), token_index)
453 };
454 let len = usize::try_from(self.model.n_vocab()).map_err(LogitsError::VocabSizeOverflow)?;
455
456 Ok(unsafe { slice::from_raw_parts(data, len) })
457 }
458
459 pub fn reset_timings(&mut self) {
460 unsafe { llama_cpp_bindings_sys::llama_perf_context_reset(self.context.as_ptr()) }
461 }
462
463 pub fn timings(&mut self) -> LlamaTimings {
464 let timings = unsafe { llama_cpp_bindings_sys::llama_perf_context(self.context.as_ptr()) };
465 LlamaTimings { timings }
466 }
467
468 pub fn lora_adapter_set(
472 &self,
473 adapter: &mut LlamaLoraAdapter,
474 scale: f32,
475 ) -> Result<(), LlamaLoraAdapterSetError> {
476 let mut adapters = [adapter.lora_adapter.as_ptr()];
477 let mut scales = [scale];
478 let err_code = unsafe {
479 llama_cpp_bindings_sys::llama_set_adapters_lora(
480 self.context.as_ptr(),
481 adapters.as_mut_ptr(),
482 1,
483 scales.as_mut_ptr(),
484 )
485 };
486 check_lora_set_result(err_code)?;
487
488 log::debug!("Set lora adapter");
489 Ok(())
490 }
491
492 pub fn lora_adapter_remove(
496 &self,
497 _adapter: &mut LlamaLoraAdapter,
498 ) -> Result<(), LlamaLoraAdapterRemoveError> {
499 let err_code = unsafe {
500 llama_cpp_bindings_sys::llama_set_adapters_lora(
501 self.context.as_ptr(),
502 std::ptr::null_mut(),
503 0,
504 std::ptr::null_mut(),
505 )
506 };
507 check_lora_remove_result(err_code)?;
508
509 log::debug!("Remove lora adapter");
510 Ok(())
511 }
512}
513
514impl Drop for LlamaContext<'_> {
515 fn drop(&mut self) {
516 unsafe { llama_cpp_bindings_sys::llama_free(self.context.as_ptr()) }
517 }
518}
519
520#[cfg(test)]
521mod unit_tests {
522 use crate::DecodeError;
523 use crate::EncodeError;
524 use crate::LlamaContextLoadError;
525 use crate::LlamaLoraAdapterRemoveError;
526 use crate::LlamaLoraAdapterSetError;
527 use crate::LogitsError;
528
529 use super::{
530 check_lora_remove_result, check_lora_set_result, decode_status_to_result,
531 encode_status_to_result, logits_slice_from_raw_parts,
532 new_context_with_model_status_to_result, token_index_within_context,
533 };
534
535 #[test]
536 fn check_lora_set_result_ok_for_zero() {
537 assert!(check_lora_set_result(0).is_ok());
538 }
539
540 #[test]
541 fn check_lora_set_result_error_for_nonzero() {
542 let result = check_lora_set_result(-1);
543
544 assert_eq!(result, Err(LlamaLoraAdapterSetError::ErrorResult(-1)));
545 }
546
547 #[test]
548 fn check_lora_remove_result_ok_for_zero() {
549 assert!(check_lora_remove_result(0).is_ok());
550 }
551
552 #[test]
553 fn check_lora_remove_result_error_for_nonzero() {
554 let result = check_lora_remove_result(-1);
555
556 assert_eq!(result, Err(LlamaLoraAdapterRemoveError::ErrorResult(-1)));
557 }
558
559 #[test]
560 fn new_context_ok_with_null_ctx_maps_unconstructible() {
561 let result = new_context_with_model_status_to_result(
562 llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_OK,
563 std::ptr::null_mut(),
564 std::ptr::null_mut(),
565 );
566
567 assert_eq!(result, Err(LlamaContextLoadError::Unconstructible));
568 }
569
570 #[test]
571 fn new_context_vendored_returned_null_maps_unconstructible() {
572 let result = new_context_with_model_status_to_result(
573 llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_VENDORED_RETURNED_NULL,
574 std::ptr::null_mut(),
575 std::ptr::null_mut(),
576 );
577
578 assert_eq!(result, Err(LlamaContextLoadError::Unconstructible));
579 }
580
581 #[test]
582 fn new_context_allocation_failed_maps_not_enough_memory() {
583 let result = new_context_with_model_status_to_result(
584 llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_ERROR_STRING_ALLOCATION_FAILED,
585 std::ptr::null_mut(),
586 std::ptr::null_mut(),
587 );
588
589 assert_eq!(result, Err(LlamaContextLoadError::NotEnoughMemory));
590 }
591
592 #[test]
593 fn new_context_cxx_exception_maps_reported() {
594 let result = new_context_with_model_status_to_result(
595 llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_VENDORED_THREW_CXX_EXCEPTION,
596 std::ptr::null_mut(),
597 std::ptr::null_mut(),
598 );
599
600 assert_eq!(
601 result,
602 Err(LlamaContextLoadError::Reported {
603 message: "unknown error".to_owned(),
604 })
605 );
606 }
607
608 #[test]
609 #[should_panic(expected = "llama_rs_new_context_with_model returned unrecognized status")]
610 fn new_context_unrecognized_status_panics() {
611 let _result = new_context_with_model_status_to_result(
612 llama_cpp_bindings_sys::llama_rs_new_context_with_model_status::MAX,
613 std::ptr::null_mut(),
614 std::ptr::null_mut(),
615 );
616 }
617
618 #[test]
619 fn decode_nonzero_code_maps_from_code() {
620 let result = decode_status_to_result(
621 llama_cpp_bindings_sys::LLAMA_RS_DECODE_VENDORED_RETURNED_NONZERO_CODE,
622 1,
623 std::ptr::null_mut(),
624 );
625
626 assert_eq!(result, Err(DecodeError::NoKvCacheSlot));
627 }
628
629 #[test]
630 fn decode_out_of_memory_maps_decode_out_of_memory() {
631 let result = decode_status_to_result(
632 llama_cpp_bindings_sys::LLAMA_RS_DECODE_OUT_OF_MEMORY,
633 0,
634 std::ptr::null_mut(),
635 );
636
637 assert_eq!(result, Err(DecodeError::DecodeOutOfMemory));
638 }
639
640 #[test]
641 fn decode_compute_failed_maps_compute_failed() {
642 let result = decode_status_to_result(
643 llama_cpp_bindings_sys::LLAMA_RS_DECODE_COMPUTE_FAILED,
644 0,
645 std::ptr::null_mut(),
646 );
647
648 assert_eq!(result, Err(DecodeError::ComputeFailed));
649 }
650
651 #[test]
652 fn decode_allocation_failed_maps_not_enough_memory() {
653 let result = decode_status_to_result(
654 llama_cpp_bindings_sys::LLAMA_RS_DECODE_ERROR_STRING_ALLOCATION_FAILED,
655 0,
656 std::ptr::null_mut(),
657 );
658
659 assert_eq!(result, Err(DecodeError::NotEnoughMemory));
660 }
661
662 #[test]
663 fn decode_cxx_exception_maps_reported() {
664 let result = decode_status_to_result(
665 llama_cpp_bindings_sys::LLAMA_RS_DECODE_VENDORED_THREW_CXX_EXCEPTION,
666 0,
667 std::ptr::null_mut(),
668 );
669
670 assert_eq!(
671 result,
672 Err(DecodeError::Reported {
673 message: "unknown error".to_owned(),
674 })
675 );
676 }
677
678 #[test]
679 #[should_panic(expected = "llama_rs_decode reported a nonzero return code")]
680 fn decode_nonzero_code_with_zero_value_panics() {
681 let _result = decode_status_to_result(
682 llama_cpp_bindings_sys::LLAMA_RS_DECODE_VENDORED_RETURNED_NONZERO_CODE,
683 0,
684 std::ptr::null_mut(),
685 );
686 }
687
688 #[test]
689 #[should_panic(expected = "llama_rs_decode returned unrecognized status")]
690 fn decode_unrecognized_status_panics() {
691 let _result = decode_status_to_result(
692 llama_cpp_bindings_sys::llama_rs_decode_status::MAX,
693 0,
694 std::ptr::null_mut(),
695 );
696 }
697
698 #[test]
699 fn encode_model_has_no_encoder_maps_model_has_no_encoder() {
700 let result = encode_status_to_result(
701 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_MODEL_HAS_NO_ENCODER,
702 0,
703 std::ptr::null_mut(),
704 );
705
706 assert_eq!(result, Err(EncodeError::ModelHasNoEncoder));
707 }
708
709 #[test]
710 fn encode_nonzero_code_maps_from_code() {
711 let result = encode_status_to_result(
712 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_VENDORED_RETURNED_NONZERO_CODE,
713 1,
714 std::ptr::null_mut(),
715 );
716
717 assert_eq!(result, Err(EncodeError::NoKvCacheSlot));
718 }
719
720 #[test]
721 fn encode_out_of_memory_maps_encode_out_of_memory() {
722 let result = encode_status_to_result(
723 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_OUT_OF_MEMORY,
724 0,
725 std::ptr::null_mut(),
726 );
727
728 assert_eq!(result, Err(EncodeError::EncodeOutOfMemory));
729 }
730
731 #[test]
732 fn encode_compute_failed_maps_compute_failed() {
733 let result = encode_status_to_result(
734 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_COMPUTE_FAILED,
735 0,
736 std::ptr::null_mut(),
737 );
738
739 assert_eq!(result, Err(EncodeError::ComputeFailed));
740 }
741
742 #[test]
743 fn encode_allocation_failed_maps_not_enough_memory() {
744 let result = encode_status_to_result(
745 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_ERROR_STRING_ALLOCATION_FAILED,
746 0,
747 std::ptr::null_mut(),
748 );
749
750 assert_eq!(result, Err(EncodeError::NotEnoughMemory));
751 }
752
753 #[test]
754 fn encode_cxx_exception_maps_reported() {
755 let result = encode_status_to_result(
756 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_VENDORED_THREW_CXX_EXCEPTION,
757 0,
758 std::ptr::null_mut(),
759 );
760
761 assert_eq!(
762 result,
763 Err(EncodeError::Reported {
764 message: "unknown error".to_owned(),
765 })
766 );
767 }
768
769 #[test]
770 #[should_panic(expected = "llama_rs_encode reported a nonzero return code")]
771 fn encode_nonzero_code_with_zero_value_panics() {
772 let _result = encode_status_to_result(
773 llama_cpp_bindings_sys::LLAMA_RS_ENCODE_VENDORED_RETURNED_NONZERO_CODE,
774 0,
775 std::ptr::null_mut(),
776 );
777 }
778
779 #[test]
780 #[should_panic(expected = "llama_rs_encode returned unrecognized status")]
781 fn encode_unrecognized_status_panics() {
782 let _result = encode_status_to_result(
783 llama_cpp_bindings_sys::llama_rs_encode_status::MAX,
784 0,
785 std::ptr::null_mut(),
786 );
787 }
788
789 #[test]
790 fn token_index_beyond_context_size_maps_exceeds_context() {
791 let result = token_index_within_context(5, 4);
792
793 assert_eq!(
794 result,
795 Err(LogitsError::TokenIndexExceedsContext {
796 token_index: 5,
797 context_size: 4,
798 })
799 );
800 }
801
802 #[test]
803 fn token_index_within_context_size_is_ok() {
804 assert!(token_index_within_context(2, 4).is_ok());
805 }
806
807 #[test]
808 fn token_index_negative_skips_context_check() {
809 assert!(token_index_within_context(-1, 4).is_ok());
810 }
811
812 #[test]
813 fn logits_slice_from_null_data_maps_null_logits() {
814 let result = unsafe { logits_slice_from_raw_parts(std::ptr::null(), 4) };
815
816 assert_eq!(result, Err(LogitsError::NullLogits));
817 }
818
819 #[test]
820 fn logits_slice_from_negative_vocab_maps_vocab_size_overflow() {
821 let logit_value = 0.0_f32;
822 let result = unsafe { logits_slice_from_raw_parts(&raw const logit_value, -1) };
823
824 let conversion_error = usize::try_from(-1_i32).unwrap_err();
825
826 assert_eq!(
827 result,
828 Err(LogitsError::VocabSizeOverflow(conversion_error))
829 );
830 }
831}