Skip to main content

llama_cpp_bindings/
context.rs

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    /// # Errors
208    ///
209    /// Returns [`LlamaContextLoadError`] when llama.cpp fails to allocate the context.
210    pub fn from_model(
211        model: &'model LlamaModel,
212        _backend: &LlamaBackend,
213        params: LlamaContextParams,
214    ) -> Result<Self, LlamaContextLoadError> {
215        let context_params = params.context_params;
216        let mut out_ctx: *mut llama_cpp_bindings_sys::llama_context = std::ptr::null_mut();
217        let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut();
218        let status = unsafe {
219            llama_cpp_bindings_sys::llama_rs_new_context_with_model(
220                model.model.as_ptr(),
221                context_params,
222                &raw mut out_ctx,
223                &raw mut out_error,
224            )
225        };
226        let context = new_context_with_model_status_to_result(status, out_ctx, out_error)?;
227
228        Ok(Self::new(model, context, params.embeddings()))
229    }
230
231    #[must_use]
232    pub fn n_batch(&self) -> u32 {
233        unsafe { llama_cpp_bindings_sys::llama_n_batch(self.context.as_ptr()) }
234    }
235
236    #[must_use]
237    pub fn n_ubatch(&self) -> u32 {
238        unsafe { llama_cpp_bindings_sys::llama_n_ubatch(self.context.as_ptr()) }
239    }
240
241    #[must_use]
242    pub fn n_ctx(&self) -> u32 {
243        unsafe { llama_cpp_bindings_sys::llama_n_ctx(self.context.as_ptr()) }
244    }
245
246    #[expect(unsafe_code, reason = "required for FFI abort callback registration")]
247    pub fn set_abort_flag(&mut self, flag: Arc<AtomicBool>) {
248        let raw_ptr = Arc::as_ptr(&flag) as *mut c_void;
249        self.abort_flag = Some(flag);
250
251        unsafe {
252            llama_cpp_bindings_sys::llama_set_abort_callback(
253                self.context.as_ptr(),
254                Some(abort_callback_trampoline),
255                raw_ptr,
256            );
257        }
258    }
259
260    #[expect(unsafe_code, reason = "required for FFI abort callback deregistration")]
261    pub fn clear_abort_callback(&mut self) {
262        self.abort_flag = None;
263
264        unsafe {
265            llama_cpp_bindings_sys::llama_set_abort_callback(
266                self.context.as_ptr(),
267                None,
268                std::ptr::null_mut(),
269            );
270        }
271    }
272
273    #[expect(unsafe_code, reason = "required for FFI synchronization call")]
274    pub fn synchronize(&self) {
275        unsafe { llama_cpp_bindings_sys::llama_synchronize(self.context.as_ptr()) }
276    }
277
278    #[expect(unsafe_code, reason = "required for FFI threadpool detachment")]
279    pub fn detach_threadpool(&self) {
280        unsafe { llama_cpp_bindings_sys::llama_detach_threadpool(self.context.as_ptr()) }
281    }
282
283    pub fn mark_logits_initialized(&mut self, token_index: i32) {
284        self.initialized_logits = vec![token_index];
285    }
286
287    /// # Errors
288    ///
289    /// - `DecodeError` if the decoding failed.
290    pub fn decode(&mut self, batch: &mut LlamaBatch) -> Result<(), DecodeError> {
291        let mut out_vendored_return_code: i32 = 0;
292        let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut();
293        let status = unsafe {
294            llama_cpp_bindings_sys::llama_rs_decode(
295                self.context.as_ptr(),
296                batch.llama_batch,
297                &raw mut out_vendored_return_code,
298                &raw mut out_error,
299            )
300        };
301        decode_status_to_result(status, out_vendored_return_code, out_error)?;
302
303        self.initialized_logits
304            .clone_from(&batch.initialized_logits);
305
306        Ok(())
307    }
308
309    /// # Errors
310    ///
311    /// - `EncodeError` if the encoding failed.
312    pub fn encode(&mut self, batch: &mut LlamaBatch) -> Result<(), EncodeError> {
313        let mut out_vendored_return_code: i32 = 0;
314        let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut();
315        let status = unsafe {
316            llama_cpp_bindings_sys::llama_rs_encode(
317                self.context.as_ptr(),
318                batch.llama_batch,
319                &raw mut out_vendored_return_code,
320                &raw mut out_error,
321            )
322        };
323        encode_status_to_result(status, out_vendored_return_code, out_error)?;
324
325        self.initialized_logits
326            .clone_from(&batch.initialized_logits);
327
328        Ok(())
329    }
330
331    /// # Errors
332    ///
333    /// - When the current context was constructed without enabling embeddings.
334    /// - If the current model had a pooling type of [`llama_cpp_bindings_sys::LLAMA_POOLING_TYPE_NONE`]
335    /// - If the given sequence index exceeds the max sequence id.
336    ///
337    pub fn embeddings_seq_ith(&self, sequence_index: i32) -> Result<&[f32], EmbeddingsError> {
338        if !self.embeddings_enabled {
339            return Err(EmbeddingsError::NotEnabled);
340        }
341
342        let n_embd = usize::try_from(self.model.n_embd())
343            .map_err(EmbeddingsError::InvalidEmbeddingDimension)?;
344
345        unsafe {
346            let embedding = llama_cpp_bindings_sys::llama_get_embeddings_seq(
347                self.context.as_ptr(),
348                sequence_index,
349            );
350
351            if embedding.is_null() {
352                Err(EmbeddingsError::NonePoolType)
353            } else {
354                Ok(slice::from_raw_parts(embedding, n_embd))
355            }
356        }
357    }
358
359    /// # Errors
360    ///
361    /// - When the current context was constructed without enabling embeddings.
362    /// - When the given token didn't have logits enabled when it was passed.
363    /// - If the given token index exceeds the max token id.
364    ///
365    pub fn embeddings_ith(&self, token_index: i32) -> Result<&[f32], EmbeddingsError> {
366        if !self.embeddings_enabled {
367            return Err(EmbeddingsError::NotEnabled);
368        }
369
370        let n_embd = usize::try_from(self.model.n_embd())
371            .map_err(EmbeddingsError::InvalidEmbeddingDimension)?;
372
373        unsafe {
374            let embedding = llama_cpp_bindings_sys::llama_get_embeddings_ith(
375                self.context.as_ptr(),
376                token_index,
377            );
378
379            if embedding.is_null() {
380                Err(EmbeddingsError::LogitsNotEnabled)
381            } else {
382                Ok(slice::from_raw_parts(embedding, n_embd))
383            }
384        }
385    }
386
387    /// # Errors
388    /// Returns `LogitsError` if logits are null or `n_vocab` overflows.
389    pub fn candidates(&self) -> Result<impl Iterator<Item = LlamaTokenData> + '_, LogitsError> {
390        let logits = self.get_logits()?;
391
392        Ok((0_i32..).zip(logits).map(|(token_id, logit)| {
393            let token = LlamaToken::new(token_id);
394            LlamaTokenData::new(token, *logit, 0_f32)
395        }))
396    }
397
398    /// # Errors
399    /// Returns `LogitsError` if logits are null or `n_vocab` overflows.
400    pub fn token_data_array(&self) -> Result<LlamaTokenDataArray, LogitsError> {
401        Ok(LlamaTokenDataArray::from_iter(self.candidates()?, false))
402    }
403
404    /// # Errors
405    /// Returns `LogitsError` if the logits pointer is null or `n_vocab` overflows.
406    pub fn get_logits(&self) -> Result<&[f32], LogitsError> {
407        let data = unsafe { llama_cpp_bindings_sys::llama_get_logits(self.context.as_ptr()) };
408
409        unsafe { logits_slice_from_raw_parts(data, self.model.n_vocab()) }
410    }
411
412    /// # Errors
413    /// Returns `LogitsError` if the token is not initialized or out of range.
414    pub fn candidates_ith(
415        &self,
416        token_index: i32,
417    ) -> Result<impl Iterator<Item = LlamaTokenData> + '_, LogitsError> {
418        let logits = self.get_logits_ith(token_index)?;
419
420        Ok((0_i32..).zip(logits).map(|(token_id, logit)| {
421            let token = LlamaToken::new(token_id);
422            LlamaTokenData::new(token, *logit, 0_f32)
423        }))
424    }
425
426    /// # Errors
427    /// Returns `LogitsError` if the token is not initialized or out of range.
428    pub fn token_data_array_ith(
429        &self,
430        token_index: i32,
431    ) -> Result<LlamaTokenDataArray, LogitsError> {
432        Ok(LlamaTokenDataArray::from_iter(
433            self.candidates_ith(token_index)?,
434            false,
435        ))
436    }
437
438    /// # Errors
439    /// Returns `LogitsError` if the token is not initialized, out of range, or `n_vocab` overflows.
440    pub fn get_logits_ith(&self, token_index: i32) -> Result<&[f32], LogitsError> {
441        if !self.initialized_logits.contains(&token_index) {
442            return Err(LogitsError::TokenNotInitialized(token_index));
443        }
444
445        token_index_within_context(token_index, self.n_ctx())?;
446
447        let data = unsafe {
448            llama_cpp_bindings_sys::llama_get_logits_ith(self.context.as_ptr(), token_index)
449        };
450        let len = usize::try_from(self.model.n_vocab()).map_err(LogitsError::VocabSizeOverflow)?;
451
452        Ok(unsafe { slice::from_raw_parts(data, len) })
453    }
454
455    pub fn reset_timings(&mut self) {
456        unsafe { llama_cpp_bindings_sys::llama_perf_context_reset(self.context.as_ptr()) }
457    }
458
459    pub fn timings(&mut self) -> LlamaTimings {
460        let timings = unsafe { llama_cpp_bindings_sys::llama_perf_context(self.context.as_ptr()) };
461        LlamaTimings { timings }
462    }
463
464    /// # Errors
465    ///
466    /// See [`LlamaLoraAdapterSetError`] for more information.
467    pub fn lora_adapter_set(
468        &self,
469        adapter: &mut LlamaLoraAdapter,
470        scale: f32,
471    ) -> Result<(), LlamaLoraAdapterSetError> {
472        let mut adapters = [adapter.lora_adapter.as_ptr()];
473        let mut scales = [scale];
474        let err_code = unsafe {
475            llama_cpp_bindings_sys::llama_set_adapters_lora(
476                self.context.as_ptr(),
477                adapters.as_mut_ptr(),
478                1,
479                scales.as_mut_ptr(),
480            )
481        };
482        check_lora_set_result(err_code)?;
483
484        log::debug!("Set lora adapter");
485        Ok(())
486    }
487
488    /// # Errors
489    ///
490    /// See [`LlamaLoraAdapterRemoveError`] for more information.
491    pub fn lora_adapter_remove(
492        &self,
493        _adapter: &mut LlamaLoraAdapter,
494    ) -> Result<(), LlamaLoraAdapterRemoveError> {
495        let err_code = unsafe {
496            llama_cpp_bindings_sys::llama_set_adapters_lora(
497                self.context.as_ptr(),
498                std::ptr::null_mut(),
499                0,
500                std::ptr::null_mut(),
501            )
502        };
503        check_lora_remove_result(err_code)?;
504
505        log::debug!("Remove lora adapter");
506        Ok(())
507    }
508}
509
510impl Drop for LlamaContext<'_> {
511    fn drop(&mut self) {
512        unsafe { llama_cpp_bindings_sys::llama_free(self.context.as_ptr()) }
513    }
514}
515
516#[cfg(test)]
517mod unit_tests {
518    use crate::DecodeError;
519    use crate::EncodeError;
520    use crate::LlamaContextLoadError;
521    use crate::LlamaLoraAdapterRemoveError;
522    use crate::LlamaLoraAdapterSetError;
523    use crate::LogitsError;
524
525    use super::{
526        check_lora_remove_result, check_lora_set_result, decode_status_to_result,
527        encode_status_to_result, logits_slice_from_raw_parts,
528        new_context_with_model_status_to_result, token_index_within_context,
529    };
530
531    #[test]
532    fn check_lora_set_result_ok_for_zero() {
533        assert!(check_lora_set_result(0).is_ok());
534    }
535
536    #[test]
537    fn check_lora_set_result_error_for_nonzero() {
538        let result = check_lora_set_result(-1);
539
540        assert_eq!(result, Err(LlamaLoraAdapterSetError::ErrorResult(-1)));
541    }
542
543    #[test]
544    fn check_lora_remove_result_ok_for_zero() {
545        assert!(check_lora_remove_result(0).is_ok());
546    }
547
548    #[test]
549    fn check_lora_remove_result_error_for_nonzero() {
550        let result = check_lora_remove_result(-1);
551
552        assert_eq!(result, Err(LlamaLoraAdapterRemoveError::ErrorResult(-1)));
553    }
554
555    #[test]
556    fn new_context_ok_with_null_ctx_maps_unconstructible() {
557        let result = new_context_with_model_status_to_result(
558            llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_OK,
559            std::ptr::null_mut(),
560            std::ptr::null_mut(),
561        );
562
563        assert_eq!(result, Err(LlamaContextLoadError::Unconstructible));
564    }
565
566    #[test]
567    fn new_context_vendored_returned_null_maps_unconstructible() {
568        let result = new_context_with_model_status_to_result(
569            llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_VENDORED_RETURNED_NULL,
570            std::ptr::null_mut(),
571            std::ptr::null_mut(),
572        );
573
574        assert_eq!(result, Err(LlamaContextLoadError::Unconstructible));
575    }
576
577    #[test]
578    fn new_context_allocation_failed_maps_not_enough_memory() {
579        let result = new_context_with_model_status_to_result(
580            llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_ERROR_STRING_ALLOCATION_FAILED,
581            std::ptr::null_mut(),
582            std::ptr::null_mut(),
583        );
584
585        assert_eq!(result, Err(LlamaContextLoadError::NotEnoughMemory));
586    }
587
588    #[test]
589    fn new_context_cxx_exception_maps_reported() {
590        let result = new_context_with_model_status_to_result(
591            llama_cpp_bindings_sys::LLAMA_RS_NEW_CONTEXT_WITH_MODEL_VENDORED_THREW_CXX_EXCEPTION,
592            std::ptr::null_mut(),
593            std::ptr::null_mut(),
594        );
595
596        assert_eq!(
597            result,
598            Err(LlamaContextLoadError::Reported {
599                message: "unknown error".to_owned(),
600            })
601        );
602    }
603
604    #[test]
605    #[should_panic(expected = "llama_rs_new_context_with_model returned unrecognized status")]
606    fn new_context_unrecognized_status_panics() {
607        let _result = new_context_with_model_status_to_result(
608            llama_cpp_bindings_sys::llama_rs_new_context_with_model_status::MAX,
609            std::ptr::null_mut(),
610            std::ptr::null_mut(),
611        );
612    }
613
614    #[test]
615    fn decode_nonzero_code_maps_from_code() {
616        let result = decode_status_to_result(
617            llama_cpp_bindings_sys::LLAMA_RS_DECODE_VENDORED_RETURNED_NONZERO_CODE,
618            1,
619            std::ptr::null_mut(),
620        );
621
622        assert_eq!(result, Err(DecodeError::NoKvCacheSlot));
623    }
624
625    #[test]
626    fn decode_out_of_memory_maps_decode_out_of_memory() {
627        let result = decode_status_to_result(
628            llama_cpp_bindings_sys::LLAMA_RS_DECODE_OUT_OF_MEMORY,
629            0,
630            std::ptr::null_mut(),
631        );
632
633        assert_eq!(result, Err(DecodeError::DecodeOutOfMemory));
634    }
635
636    #[test]
637    fn decode_compute_failed_maps_compute_failed() {
638        let result = decode_status_to_result(
639            llama_cpp_bindings_sys::LLAMA_RS_DECODE_COMPUTE_FAILED,
640            0,
641            std::ptr::null_mut(),
642        );
643
644        assert_eq!(result, Err(DecodeError::ComputeFailed));
645    }
646
647    #[test]
648    fn decode_allocation_failed_maps_not_enough_memory() {
649        let result = decode_status_to_result(
650            llama_cpp_bindings_sys::LLAMA_RS_DECODE_ERROR_STRING_ALLOCATION_FAILED,
651            0,
652            std::ptr::null_mut(),
653        );
654
655        assert_eq!(result, Err(DecodeError::NotEnoughMemory));
656    }
657
658    #[test]
659    fn decode_cxx_exception_maps_reported() {
660        let result = decode_status_to_result(
661            llama_cpp_bindings_sys::LLAMA_RS_DECODE_VENDORED_THREW_CXX_EXCEPTION,
662            0,
663            std::ptr::null_mut(),
664        );
665
666        assert_eq!(
667            result,
668            Err(DecodeError::Reported {
669                message: "unknown error".to_owned(),
670            })
671        );
672    }
673
674    #[test]
675    #[should_panic(expected = "llama_rs_decode reported a nonzero return code")]
676    fn decode_nonzero_code_with_zero_value_panics() {
677        let _result = decode_status_to_result(
678            llama_cpp_bindings_sys::LLAMA_RS_DECODE_VENDORED_RETURNED_NONZERO_CODE,
679            0,
680            std::ptr::null_mut(),
681        );
682    }
683
684    #[test]
685    #[should_panic(expected = "llama_rs_decode returned unrecognized status")]
686    fn decode_unrecognized_status_panics() {
687        let _result = decode_status_to_result(
688            llama_cpp_bindings_sys::llama_rs_decode_status::MAX,
689            0,
690            std::ptr::null_mut(),
691        );
692    }
693
694    #[test]
695    fn encode_model_has_no_encoder_maps_model_has_no_encoder() {
696        let result = encode_status_to_result(
697            llama_cpp_bindings_sys::LLAMA_RS_ENCODE_MODEL_HAS_NO_ENCODER,
698            0,
699            std::ptr::null_mut(),
700        );
701
702        assert_eq!(result, Err(EncodeError::ModelHasNoEncoder));
703    }
704
705    #[test]
706    fn encode_nonzero_code_maps_from_code() {
707        let result = encode_status_to_result(
708            llama_cpp_bindings_sys::LLAMA_RS_ENCODE_VENDORED_RETURNED_NONZERO_CODE,
709            1,
710            std::ptr::null_mut(),
711        );
712
713        assert_eq!(result, Err(EncodeError::NoKvCacheSlot));
714    }
715
716    #[test]
717    fn encode_out_of_memory_maps_encode_out_of_memory() {
718        let result = encode_status_to_result(
719            llama_cpp_bindings_sys::LLAMA_RS_ENCODE_OUT_OF_MEMORY,
720            0,
721            std::ptr::null_mut(),
722        );
723
724        assert_eq!(result, Err(EncodeError::EncodeOutOfMemory));
725    }
726
727    #[test]
728    fn encode_compute_failed_maps_compute_failed() {
729        let result = encode_status_to_result(
730            llama_cpp_bindings_sys::LLAMA_RS_ENCODE_COMPUTE_FAILED,
731            0,
732            std::ptr::null_mut(),
733        );
734
735        assert_eq!(result, Err(EncodeError::ComputeFailed));
736    }
737
738    #[test]
739    fn encode_allocation_failed_maps_not_enough_memory() {
740        let result = encode_status_to_result(
741            llama_cpp_bindings_sys::LLAMA_RS_ENCODE_ERROR_STRING_ALLOCATION_FAILED,
742            0,
743            std::ptr::null_mut(),
744        );
745
746        assert_eq!(result, Err(EncodeError::NotEnoughMemory));
747    }
748
749    #[test]
750    fn encode_cxx_exception_maps_reported() {
751        let result = encode_status_to_result(
752            llama_cpp_bindings_sys::LLAMA_RS_ENCODE_VENDORED_THREW_CXX_EXCEPTION,
753            0,
754            std::ptr::null_mut(),
755        );
756
757        assert_eq!(
758            result,
759            Err(EncodeError::Reported {
760                message: "unknown error".to_owned(),
761            })
762        );
763    }
764
765    #[test]
766    #[should_panic(expected = "llama_rs_encode reported a nonzero return code")]
767    fn encode_nonzero_code_with_zero_value_panics() {
768        let _result = encode_status_to_result(
769            llama_cpp_bindings_sys::LLAMA_RS_ENCODE_VENDORED_RETURNED_NONZERO_CODE,
770            0,
771            std::ptr::null_mut(),
772        );
773    }
774
775    #[test]
776    #[should_panic(expected = "llama_rs_encode returned unrecognized status")]
777    fn encode_unrecognized_status_panics() {
778        let _result = encode_status_to_result(
779            llama_cpp_bindings_sys::llama_rs_encode_status::MAX,
780            0,
781            std::ptr::null_mut(),
782        );
783    }
784
785    #[test]
786    fn token_index_beyond_context_size_maps_exceeds_context() {
787        let result = token_index_within_context(5, 4);
788
789        assert_eq!(
790            result,
791            Err(LogitsError::TokenIndexExceedsContext {
792                token_index: 5,
793                context_size: 4,
794            })
795        );
796    }
797
798    #[test]
799    fn token_index_within_context_size_is_ok() {
800        assert!(token_index_within_context(2, 4).is_ok());
801    }
802
803    #[test]
804    fn token_index_negative_skips_context_check() {
805        assert!(token_index_within_context(-1, 4).is_ok());
806    }
807
808    #[test]
809    fn logits_slice_from_null_data_maps_null_logits() {
810        let result = unsafe { logits_slice_from_raw_parts(std::ptr::null(), 4) };
811
812        assert_eq!(result, Err(LogitsError::NullLogits));
813    }
814
815    #[test]
816    fn logits_slice_from_negative_vocab_maps_vocab_size_overflow() {
817        let logit_value = 0.0_f32;
818        let result = unsafe { logits_slice_from_raw_parts(&raw const logit_value, -1) };
819
820        let conversion_error = usize::try_from(-1_i32).unwrap_err();
821
822        assert_eq!(
823            result,
824            Err(LogitsError::VocabSizeOverflow(conversion_error))
825        );
826    }
827}