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    #[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    /// # Errors
292    ///
293    /// - `DecodeError` if the decoding failed.
294    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    /// # Errors
314    ///
315    /// - `EncodeError` if the encoding failed.
316    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    /// # Errors
336    ///
337    /// - When the current context was constructed without enabling embeddings.
338    /// - If the current model had a pooling type of [`llama_cpp_bindings_sys::LLAMA_POOLING_TYPE_NONE`]
339    /// - If the given sequence index exceeds the max sequence id.
340    ///
341    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    /// # Errors
364    ///
365    /// - When the current context was constructed without enabling embeddings.
366    /// - When the given token didn't have logits enabled when it was passed.
367    /// - If the given token index exceeds the max token id.
368    ///
369    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    /// # Errors
392    /// Returns `LogitsError` if logits are null or `n_vocab` overflows.
393    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    /// # Errors
403    /// Returns `LogitsError` if logits are null or `n_vocab` overflows.
404    pub fn token_data_array(&self) -> Result<LlamaTokenDataArray, LogitsError> {
405        Ok(LlamaTokenDataArray::from_iter(self.candidates()?, false))
406    }
407
408    /// # Errors
409    /// Returns `LogitsError` if the logits pointer is null or `n_vocab` overflows.
410    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    /// # Errors
417    /// Returns `LogitsError` if the token is not initialized or out of range.
418    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    /// # Errors
431    /// Returns `LogitsError` if the token is not initialized or out of range.
432    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    /// # Errors
443    /// Returns `LogitsError` if the token is not initialized, out of range, or `n_vocab` overflows.
444    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    /// # Errors
469    ///
470    /// See [`LlamaLoraAdapterSetError`] for more information.
471    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    /// # Errors
493    ///
494    /// See [`LlamaLoraAdapterRemoveError`] for more information.
495    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}