Skip to main content

llama_cpp_bindings/context/
kv_cache.rs

1use std::ffi::c_int;
2use std::num::{NonZeroU8, TryFromIntError};
3use std::os::raw::c_char;
4use std::ptr;
5
6use crate::context::LlamaContext;
7use crate::error::{KvCacheSeqAddError, KvCacheSeqDivError};
8use crate::ffi_error_reader::read_and_free_cpp_error;
9
10#[derive(Debug, Eq, PartialEq, thiserror::Error)]
11pub enum KvCacheConversionError {
12    #[error("Provided sequence id is too large for a i32")]
13    SeqIdTooLarge(#[source] TryFromIntError),
14    #[error("Provided start position is too large for a i32")]
15    P0TooLarge(#[source] TryFromIntError),
16    #[error("Provided end position is too large for a i32")]
17    P1TooLarge(#[source] TryFromIntError),
18}
19
20fn kv_cache_seq_add_status_to_result(
21    status: llama_cpp_bindings_sys::llama_rs_memory_seq_add_status,
22    out_error: *mut c_char,
23) -> Result<(), KvCacheSeqAddError> {
24    match status {
25        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_OK => Ok(()),
26        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_INCOMPATIBLE_ROPE_TYPE => {
27            Err(KvCacheSeqAddError::IncompatibleRopeType)
28        }
29        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_NULL_MEM => {
30            Err(KvCacheSeqAddError::MemoryHandleUnavailable)
31        }
32        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_ERROR_STRING_ALLOCATION_FAILED => {
33            Err(KvCacheSeqAddError::NotEnoughMemory)
34        }
35        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_VENDORED_THREW_CXX_EXCEPTION => {
36            let message = unsafe { read_and_free_cpp_error(out_error) };
37            Err(KvCacheSeqAddError::Reported { message })
38        }
39        other => unreachable!("llama_rs_memory_seq_add returned unrecognized status {other}"),
40    }
41}
42
43fn kv_cache_seq_div_status_to_result(
44    status: llama_cpp_bindings_sys::llama_rs_memory_seq_div_status,
45    out_error: *mut c_char,
46) -> Result<(), KvCacheSeqDivError> {
47    match status {
48        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_OK => Ok(()),
49        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_INCOMPATIBLE_ROPE_TYPE => {
50            Err(KvCacheSeqDivError::IncompatibleRopeType)
51        }
52        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_NULL_MEM => {
53            Err(KvCacheSeqDivError::MemoryHandleUnavailable)
54        }
55        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_ERROR_STRING_ALLOCATION_FAILED => {
56            Err(KvCacheSeqDivError::NotEnoughMemory)
57        }
58        llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_VENDORED_THREW_CXX_EXCEPTION => {
59            let message = unsafe { read_and_free_cpp_error(out_error) };
60            Err(KvCacheSeqDivError::Reported { message })
61        }
62        other => unreachable!("llama_rs_memory_seq_div returned unrecognized status {other}"),
63    }
64}
65
66impl LlamaContext<'_> {
67    pub fn copy_cache(&mut self, src: i32, dest: i32, size: i32) {
68        let mem = unsafe { llama_cpp_bindings_sys::llama_get_memory(self.context.as_ptr()) };
69        unsafe { llama_cpp_bindings_sys::llama_memory_seq_cp(mem, src, dest, 0, size) }
70    }
71
72    /// # Errors
73    /// If either position exceeds [`i32::MAX`].
74    pub fn copy_kv_cache_seq(
75        &mut self,
76        src: i32,
77        dest: i32,
78        p0: Option<u32>,
79        p1: Option<u32>,
80    ) -> Result<(), KvCacheConversionError> {
81        let p0 = p0
82            .map_or(Ok(-1), i32::try_from)
83            .map_err(KvCacheConversionError::P0TooLarge)?;
84        let p1 = p1
85            .map_or(Ok(-1), i32::try_from)
86            .map_err(KvCacheConversionError::P1TooLarge)?;
87        let mem = unsafe { llama_cpp_bindings_sys::llama_get_memory(self.context.as_ptr()) };
88        unsafe { llama_cpp_bindings_sys::llama_memory_seq_cp(mem, src, dest, p0, p1) };
89        Ok(())
90    }
91
92    /// # Errors
93    /// If the sequence id or either position exceeds [`i32::MAX`].
94    pub fn clear_kv_cache_seq(
95        &mut self,
96        src: Option<u32>,
97        p0: Option<u32>,
98        p1: Option<u32>,
99    ) -> Result<bool, KvCacheConversionError> {
100        let src = src
101            .map_or(Ok(-1), i32::try_from)
102            .map_err(KvCacheConversionError::SeqIdTooLarge)?;
103        let p0 = p0
104            .map_or(Ok(-1), i32::try_from)
105            .map_err(KvCacheConversionError::P0TooLarge)?;
106        let p1 = p1
107            .map_or(Ok(-1), i32::try_from)
108            .map_err(KvCacheConversionError::P1TooLarge)?;
109        let mem = unsafe { llama_cpp_bindings_sys::llama_get_memory(self.context.as_ptr()) };
110        Ok(unsafe { llama_cpp_bindings_sys::llama_memory_seq_rm(mem, src, p0, p1) })
111    }
112
113    pub fn clear_kv_cache(&mut self) {
114        let mem = unsafe { llama_cpp_bindings_sys::llama_get_memory(self.context.as_ptr()) };
115        let clear_data_buffers = true;
116        unsafe { llama_cpp_bindings_sys::llama_memory_clear(mem, clear_data_buffers) }
117    }
118
119    pub fn kv_cache_seq_keep(&mut self, seq_id: i32) {
120        let mem = unsafe { llama_cpp_bindings_sys::llama_get_memory(self.context.as_ptr()) };
121        unsafe { llama_cpp_bindings_sys::llama_memory_seq_keep(mem, seq_id) }
122    }
123
124    /// # Errors
125    /// If either position exceeds [`i32::MAX`], or the underlying memory operation reports a failure.
126    pub fn kv_cache_seq_add(
127        &mut self,
128        seq_id: i32,
129        p0: Option<u32>,
130        p1: Option<u32>,
131        delta: i32,
132    ) -> Result<(), KvCacheSeqAddError> {
133        let p0 = p0
134            .map_or(Ok(-1), i32::try_from)
135            .map_err(KvCacheSeqAddError::P0TooLarge)?;
136        let p1 = p1
137            .map_or(Ok(-1), i32::try_from)
138            .map_err(KvCacheSeqAddError::P1TooLarge)?;
139        let mut out_error: *mut c_char = ptr::null_mut();
140        let status = unsafe {
141            llama_cpp_bindings_sys::llama_rs_memory_seq_add(
142                self.context.as_ptr().cast_const(),
143                seq_id,
144                p0,
145                p1,
146                delta,
147                &raw mut out_error,
148            )
149        };
150        kv_cache_seq_add_status_to_result(status, out_error)
151    }
152
153    /// # Errors
154    /// If either position exceeds [`i32::MAX`], or the underlying memory operation reports a failure.
155    pub fn kv_cache_seq_div(
156        &mut self,
157        seq_id: i32,
158        p0: Option<u32>,
159        p1: Option<u32>,
160        d: NonZeroU8,
161    ) -> Result<(), KvCacheSeqDivError> {
162        let p0 = p0
163            .map_or(Ok(-1), i32::try_from)
164            .map_err(KvCacheSeqDivError::P0TooLarge)?;
165        let p1 = p1
166            .map_or(Ok(-1), i32::try_from)
167            .map_err(KvCacheSeqDivError::P1TooLarge)?;
168        let d = c_int::from(d.get());
169        let mut out_error: *mut c_char = ptr::null_mut();
170        let status = unsafe {
171            llama_cpp_bindings_sys::llama_rs_memory_seq_div(
172                self.context.as_ptr().cast_const(),
173                seq_id,
174                p0,
175                p1,
176                d,
177                &raw mut out_error,
178            )
179        };
180        kv_cache_seq_div_status_to_result(status, out_error)
181    }
182
183    #[must_use]
184    pub fn kv_cache_seq_pos_max(&self, seq_id: i32) -> i32 {
185        unsafe {
186            llama_cpp_bindings_sys::llama_rs_memory_seq_pos_max(
187                self.context.as_ptr().cast_const(),
188                seq_id,
189            )
190        }
191    }
192}
193
194#[cfg(test)]
195mod tests {
196    use std::ptr;
197
198    use super::kv_cache_seq_add_status_to_result;
199    use super::kv_cache_seq_div_status_to_result;
200    use crate::error::{KvCacheSeqAddError, KvCacheSeqDivError};
201
202    #[test]
203    fn add_ok_status_maps_to_ok() {
204        let result = kv_cache_seq_add_status_to_result(
205            llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_OK,
206            ptr::null_mut(),
207        );
208
209        assert!(result.is_ok());
210    }
211
212    #[test]
213    fn add_incompatible_rope_type_status_maps_to_incompatible_rope_type() {
214        assert_eq!(
215            kv_cache_seq_add_status_to_result(
216                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_INCOMPATIBLE_ROPE_TYPE,
217                ptr::null_mut(),
218            ),
219            Err(KvCacheSeqAddError::IncompatibleRopeType)
220        );
221    }
222
223    #[test]
224    fn add_null_mem_status_maps_to_memory_handle_unavailable() {
225        assert_eq!(
226            kv_cache_seq_add_status_to_result(
227                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_NULL_MEM,
228                ptr::null_mut(),
229            ),
230            Err(KvCacheSeqAddError::MemoryHandleUnavailable)
231        );
232    }
233
234    #[test]
235    fn add_allocation_failed_status_maps_to_not_enough_memory() {
236        assert_eq!(
237            kv_cache_seq_add_status_to_result(
238                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_ERROR_STRING_ALLOCATION_FAILED,
239                ptr::null_mut(),
240            ),
241            Err(KvCacheSeqAddError::NotEnoughMemory)
242        );
243    }
244
245    #[test]
246    fn add_vendored_exception_status_maps_to_reported_with_unknown_message() {
247        assert_eq!(
248            kv_cache_seq_add_status_to_result(
249                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_VENDORED_THREW_CXX_EXCEPTION,
250                ptr::null_mut(),
251            ),
252            Err(KvCacheSeqAddError::Reported {
253                message: "unknown error".to_owned(),
254            })
255        );
256    }
257
258    #[test]
259    #[should_panic(expected = "llama_rs_memory_seq_add returned unrecognized status")]
260    fn add_unrecognized_status_panics() {
261        let _ = kv_cache_seq_add_status_to_result(
262            llama_cpp_bindings_sys::llama_rs_memory_seq_add_status::MAX,
263            ptr::null_mut(),
264        );
265    }
266
267    #[test]
268    fn div_ok_status_maps_to_ok() {
269        let result = kv_cache_seq_div_status_to_result(
270            llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_OK,
271            ptr::null_mut(),
272        );
273
274        assert!(result.is_ok());
275    }
276
277    #[test]
278    fn div_incompatible_rope_type_status_maps_to_incompatible_rope_type() {
279        assert_eq!(
280            kv_cache_seq_div_status_to_result(
281                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_INCOMPATIBLE_ROPE_TYPE,
282                ptr::null_mut(),
283            ),
284            Err(KvCacheSeqDivError::IncompatibleRopeType)
285        );
286    }
287
288    #[test]
289    fn div_null_mem_status_maps_to_memory_handle_unavailable() {
290        assert_eq!(
291            kv_cache_seq_div_status_to_result(
292                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_NULL_MEM,
293                ptr::null_mut(),
294            ),
295            Err(KvCacheSeqDivError::MemoryHandleUnavailable)
296        );
297    }
298
299    #[test]
300    fn div_allocation_failed_status_maps_to_not_enough_memory() {
301        assert_eq!(
302            kv_cache_seq_div_status_to_result(
303                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_ERROR_STRING_ALLOCATION_FAILED,
304                ptr::null_mut(),
305            ),
306            Err(KvCacheSeqDivError::NotEnoughMemory)
307        );
308    }
309
310    #[test]
311    fn div_vendored_exception_status_maps_to_reported_with_unknown_message() {
312        assert_eq!(
313            kv_cache_seq_div_status_to_result(
314                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_VENDORED_THREW_CXX_EXCEPTION,
315                ptr::null_mut(),
316            ),
317            Err(KvCacheSeqDivError::Reported {
318                message: "unknown error".to_owned(),
319            })
320        );
321    }
322
323    #[test]
324    #[should_panic(expected = "llama_rs_memory_seq_div returned unrecognized status")]
325    fn div_unrecognized_status_panics() {
326        let _ = kv_cache_seq_div_status_to_result(
327            llama_cpp_bindings_sys::llama_rs_memory_seq_div_status::MAX,
328            ptr::null_mut(),
329        );
330    }
331}