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(),
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(),
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(self.context.as_ptr(), seq_id)
187        }
188    }
189}
190
191#[cfg(test)]
192mod tests {
193    use std::ptr;
194
195    use super::kv_cache_seq_add_status_to_result;
196    use super::kv_cache_seq_div_status_to_result;
197    use crate::error::{KvCacheSeqAddError, KvCacheSeqDivError};
198
199    #[test]
200    fn add_ok_status_maps_to_ok() {
201        let result = kv_cache_seq_add_status_to_result(
202            llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_OK,
203            ptr::null_mut(),
204        );
205
206        assert!(result.is_ok());
207    }
208
209    #[test]
210    fn add_incompatible_rope_type_status_maps_to_incompatible_rope_type() {
211        assert_eq!(
212            kv_cache_seq_add_status_to_result(
213                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_INCOMPATIBLE_ROPE_TYPE,
214                ptr::null_mut(),
215            ),
216            Err(KvCacheSeqAddError::IncompatibleRopeType)
217        );
218    }
219
220    #[test]
221    fn add_null_mem_status_maps_to_memory_handle_unavailable() {
222        assert_eq!(
223            kv_cache_seq_add_status_to_result(
224                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_NULL_MEM,
225                ptr::null_mut(),
226            ),
227            Err(KvCacheSeqAddError::MemoryHandleUnavailable)
228        );
229    }
230
231    #[test]
232    fn add_allocation_failed_status_maps_to_not_enough_memory() {
233        assert_eq!(
234            kv_cache_seq_add_status_to_result(
235                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_ERROR_STRING_ALLOCATION_FAILED,
236                ptr::null_mut(),
237            ),
238            Err(KvCacheSeqAddError::NotEnoughMemory)
239        );
240    }
241
242    #[test]
243    fn add_vendored_exception_status_maps_to_reported_with_unknown_message() {
244        assert_eq!(
245            kv_cache_seq_add_status_to_result(
246                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_ADD_VENDORED_THREW_CXX_EXCEPTION,
247                ptr::null_mut(),
248            ),
249            Err(KvCacheSeqAddError::Reported {
250                message: "unknown error".to_owned(),
251            })
252        );
253    }
254
255    #[test]
256    #[should_panic(expected = "llama_rs_memory_seq_add returned unrecognized status")]
257    fn add_unrecognized_status_panics() {
258        let _ = kv_cache_seq_add_status_to_result(
259            llama_cpp_bindings_sys::llama_rs_memory_seq_add_status::MAX,
260            ptr::null_mut(),
261        );
262    }
263
264    #[test]
265    fn div_ok_status_maps_to_ok() {
266        let result = kv_cache_seq_div_status_to_result(
267            llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_OK,
268            ptr::null_mut(),
269        );
270
271        assert!(result.is_ok());
272    }
273
274    #[test]
275    fn div_incompatible_rope_type_status_maps_to_incompatible_rope_type() {
276        assert_eq!(
277            kv_cache_seq_div_status_to_result(
278                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_INCOMPATIBLE_ROPE_TYPE,
279                ptr::null_mut(),
280            ),
281            Err(KvCacheSeqDivError::IncompatibleRopeType)
282        );
283    }
284
285    #[test]
286    fn div_null_mem_status_maps_to_memory_handle_unavailable() {
287        assert_eq!(
288            kv_cache_seq_div_status_to_result(
289                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_NULL_MEM,
290                ptr::null_mut(),
291            ),
292            Err(KvCacheSeqDivError::MemoryHandleUnavailable)
293        );
294    }
295
296    #[test]
297    fn div_allocation_failed_status_maps_to_not_enough_memory() {
298        assert_eq!(
299            kv_cache_seq_div_status_to_result(
300                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_ERROR_STRING_ALLOCATION_FAILED,
301                ptr::null_mut(),
302            ),
303            Err(KvCacheSeqDivError::NotEnoughMemory)
304        );
305    }
306
307    #[test]
308    fn div_vendored_exception_status_maps_to_reported_with_unknown_message() {
309        assert_eq!(
310            kv_cache_seq_div_status_to_result(
311                llama_cpp_bindings_sys::LLAMA_RS_MEMORY_SEQ_DIV_VENDORED_THREW_CXX_EXCEPTION,
312                ptr::null_mut(),
313            ),
314            Err(KvCacheSeqDivError::Reported {
315                message: "unknown error".to_owned(),
316            })
317        );
318    }
319
320    #[test]
321    #[should_panic(expected = "llama_rs_memory_seq_div returned unrecognized status")]
322    fn div_unrecognized_status_panics() {
323        let _ = kv_cache_seq_div_status_to_result(
324            llama_cpp_bindings_sys::llama_rs_memory_seq_div_status::MAX,
325            ptr::null_mut(),
326        );
327    }
328}