llama_cpp_bindings/context/
kv_cache.rs1use 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 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 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 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 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}