Skip to main content

llama_cpp_bindings/
gguf_context.rs

1use std::ffi::{CStr, CString};
2use std::path::Path;
3use std::ptr::NonNull;
4
5use crate::gguf_context_error::GgufContextError;
6use crate::gguf_type::GgufType;
7
8#[derive(Debug)]
9pub struct GgufContext {
10    context: NonNull<llama_cpp_bindings_sys::gguf_context>,
11}
12
13impl GgufContext {
14    /// # Errors
15    ///
16    /// Returns [`GgufContextError::InitFailed`] if the file cannot be opened or parsed.
17    /// Returns [`GgufContextError::PathToStrError`] if the path is not valid UTF-8.
18    /// Returns [`GgufContextError::NulError`] if the path contains a null byte.
19    pub fn from_file(path: impl AsRef<Path>) -> Result<Self, GgufContextError> {
20        let path_ref = path.as_ref();
21        let path_str = path_ref
22            .to_str()
23            .ok_or_else(|| GgufContextError::PathToStrError(path_ref.to_path_buf()))?;
24        let c_path = CString::new(path_str)?;
25
26        let init_params = llama_cpp_bindings_sys::gguf_init_params {
27            no_alloc: true,
28            ctx: std::ptr::null_mut(),
29        };
30
31        let raw =
32            unsafe { llama_cpp_bindings_sys::gguf_init_from_file(c_path.as_ptr(), init_params) };
33        let context = NonNull::new(raw)
34            .ok_or_else(|| GgufContextError::InitFailed(path_ref.to_path_buf()))?;
35
36        Ok(Self { context })
37    }
38
39    #[must_use]
40    pub fn n_kv(&self) -> i64 {
41        unsafe { llama_cpp_bindings_sys::gguf_get_n_kv(self.context.as_ptr()) }
42    }
43
44    /// # Errors
45    ///
46    /// Returns [`GgufContextError::KeyNotFound`] if the key does not exist.
47    /// Returns [`GgufContextError::NulError`] if the key contains a null byte.
48    pub fn find_key(&self, key: &str) -> Result<i64, GgufContextError> {
49        let c_key = CString::new(key)?;
50        let index =
51            unsafe { llama_cpp_bindings_sys::gguf_find_key(self.context.as_ptr(), c_key.as_ptr()) };
52
53        if index < 0 {
54            return Err(GgufContextError::KeyNotFound {
55                key: key.to_string(),
56            });
57        }
58
59        Ok(index)
60    }
61
62    /// # Safety considerations
63    ///
64    /// The caller must ensure `key_id` is in range `[0, n_kv())`.
65    ///
66    /// # Errors
67    ///
68    /// Returns [`GgufContextError::Utf8Error`] if the key name is not valid UTF-8.
69    pub fn key_at(&self, key_id: i64) -> Result<&str, GgufContextError> {
70        let c_str = unsafe {
71            CStr::from_ptr(llama_cpp_bindings_sys::gguf_get_key(
72                self.context.as_ptr(),
73                key_id,
74            ))
75        };
76
77        Ok(c_str.to_str()?)
78    }
79
80    /// # Safety considerations
81    ///
82    /// The caller must ensure `key_id` is in range `[0, n_kv())`.
83    #[must_use]
84    pub fn kv_type(&self, key_id: i64) -> Option<GgufType> {
85        let raw =
86            unsafe { llama_cpp_bindings_sys::gguf_get_kv_type(self.context.as_ptr(), key_id) };
87
88        GgufType::from_raw(raw)
89    }
90
91    /// # Safety considerations
92    ///
93    /// The caller must ensure the key at `key_id` has type [`GgufType::Uint32`].
94    #[must_use]
95    pub fn val_u32(&self, key_id: i64) -> u32 {
96        unsafe { llama_cpp_bindings_sys::gguf_get_val_u32(self.context.as_ptr(), key_id) }
97    }
98
99    /// # Safety considerations
100    ///
101    /// The caller must ensure the key at `key_id` has type [`GgufType::Int32`].
102    #[must_use]
103    pub fn val_i32(&self, key_id: i64) -> i32 {
104        unsafe { llama_cpp_bindings_sys::gguf_get_val_i32(self.context.as_ptr(), key_id) }
105    }
106
107    /// # Safety considerations
108    ///
109    /// The caller must ensure the key at `key_id` has type [`GgufType::Uint64`].
110    #[must_use]
111    pub fn val_u64(&self, key_id: i64) -> u64 {
112        unsafe { llama_cpp_bindings_sys::gguf_get_val_u64(self.context.as_ptr(), key_id) }
113    }
114
115    /// # Safety considerations
116    ///
117    /// The caller must ensure the key at `key_id` has type [`GgufType::String`].
118    ///
119    /// # Errors
120    ///
121    /// Returns [`GgufContextError::Utf8Error`] if the string value is not valid UTF-8.
122    pub fn val_str(&self, key_id: i64) -> Result<&str, GgufContextError> {
123        let c_str = unsafe {
124            CStr::from_ptr(llama_cpp_bindings_sys::gguf_get_val_str(
125                self.context.as_ptr(),
126                key_id,
127            ))
128        };
129
130        Ok(c_str.to_str()?)
131    }
132
133    #[must_use]
134    pub fn n_tensors(&self) -> i64 {
135        unsafe { llama_cpp_bindings_sys::gguf_get_n_tensors(self.context.as_ptr()) }
136    }
137}
138
139impl Drop for GgufContext {
140    fn drop(&mut self) {
141        unsafe { llama_cpp_bindings_sys::gguf_free(self.context.as_ptr()) }
142    }
143}
144
145#[cfg(test)]
146mod tests {
147    use std::ffi::CString;
148    use std::mem::Discriminant;
149    use std::path::PathBuf;
150
151    use super::GgufContext;
152    use crate::gguf_context_error::GgufContextError;
153    use crate::gguf_type::GgufType;
154
155    fn fixture_path() -> PathBuf {
156        PathBuf::from(env!("CARGO_MANIFEST_DIR"))
157            .join("fixtures")
158            .join("ggml-vocab-bert-bge.gguf")
159    }
160
161    fn init_failed_disc() -> Discriminant<GgufContextError> {
162        std::mem::discriminant(&GgufContextError::InitFailed(PathBuf::new()))
163    }
164
165    fn key_not_found_disc() -> Discriminant<GgufContextError> {
166        std::mem::discriminant(&GgufContextError::KeyNotFound { key: String::new() })
167    }
168
169    fn nul_error_disc() -> Discriminant<GgufContextError> {
170        let nul_err = CString::new(b"a\0b".to_vec()).unwrap_err();
171        std::mem::discriminant(&GgufContextError::NulError(nul_err))
172    }
173
174    #[cfg(unix)]
175    fn path_to_str_error_disc() -> Discriminant<GgufContextError> {
176        std::mem::discriminant(&GgufContextError::PathToStrError(PathBuf::new()))
177    }
178
179    fn utf8_error_disc() -> Discriminant<GgufContextError> {
180        let invalid_utf8_bytes: Vec<u8> = vec![0xFF];
181        let utf8_err = std::str::from_utf8(&invalid_utf8_bytes).unwrap_err();
182        std::mem::discriminant(&GgufContextError::Utf8Error(utf8_err))
183    }
184
185    #[test]
186    fn from_file_opens_valid_gguf() {
187        let context = GgufContext::from_file(fixture_path());
188
189        assert!(context.is_ok());
190    }
191
192    #[test]
193    fn from_file_nonexistent_returns_init_failed() {
194        let err = GgufContext::from_file("/nonexistent/file.gguf").unwrap_err();
195
196        assert_eq!(std::mem::discriminant(&err), init_failed_disc());
197    }
198
199    #[test]
200    fn n_kv_returns_positive_count() {
201        let context = GgufContext::from_file(fixture_path()).unwrap();
202
203        assert!(context.n_kv() > 0);
204    }
205
206    #[test]
207    fn n_tensors_returns_count() {
208        let context = GgufContext::from_file(fixture_path()).unwrap();
209
210        assert!(context.n_tensors() >= 0);
211    }
212
213    #[test]
214    fn find_key_returns_valid_index_for_known_key() {
215        let context = GgufContext::from_file(fixture_path()).unwrap();
216        let index = context.find_key("general.architecture");
217
218        assert!(index.is_ok());
219        assert!(index.unwrap() >= 0);
220    }
221
222    #[test]
223    fn find_key_returns_error_for_missing_key() {
224        let context = GgufContext::from_file(fixture_path()).unwrap();
225        let err = context.find_key("nonexistent.key").unwrap_err();
226
227        assert_eq!(std::mem::discriminant(&err), key_not_found_disc());
228    }
229
230    #[test]
231    fn key_at_returns_expected_name() {
232        let context = GgufContext::from_file(fixture_path()).unwrap();
233        let index = context.find_key("general.architecture").unwrap();
234        let key_name = context.key_at(index).unwrap();
235
236        assert_eq!(key_name, "general.architecture");
237    }
238
239    #[test]
240    fn kv_type_returns_expected_type_for_string_key() {
241        let context = GgufContext::from_file(fixture_path()).unwrap();
242        let index = context.find_key("general.architecture").unwrap();
243        let value_type = context.kv_type(index);
244
245        assert_eq!(value_type, Some(GgufType::String));
246    }
247
248    #[test]
249    fn val_str_returns_architecture_value() {
250        let context = GgufContext::from_file(fixture_path()).unwrap();
251        let index = context.find_key("general.architecture").unwrap();
252        let value = context.val_str(index).unwrap();
253
254        assert!(!value.is_empty());
255    }
256
257    #[cfg(unix)]
258    #[test]
259    fn from_file_non_utf8_path_returns_error() {
260        use std::ffi::OsStr;
261        use std::os::unix::ffi::OsStrExt;
262
263        let non_utf8_path = std::path::Path::new(OsStr::from_bytes(b"/tmp/\xff\xfe.gguf"));
264        let err = GgufContext::from_file(non_utf8_path).unwrap_err();
265
266        assert_eq!(std::mem::discriminant(&err), path_to_str_error_disc());
267    }
268
269    #[test]
270    fn from_file_with_null_byte_in_path_returns_error() {
271        let err = GgufContext::from_file("/tmp/foo\0bar.gguf").unwrap_err();
272
273        assert_eq!(std::mem::discriminant(&err), nul_error_disc());
274    }
275
276    #[test]
277    fn find_key_with_null_byte_in_key_returns_error() {
278        let context = GgufContext::from_file(fixture_path()).unwrap();
279        let err = context.find_key("foo\0bar").unwrap_err();
280
281        assert_eq!(std::mem::discriminant(&err), nul_error_disc());
282    }
283
284    #[test]
285    fn val_u32_returns_value_for_uint32_key() {
286        let context = GgufContext::from_file(fixture_path()).unwrap();
287
288        let key_id = (0..context.n_kv())
289            .find(|&id| context.kv_type(id) == Some(GgufType::Uint32))
290            .expect("fixture must contain at least one uint32 key");
291
292        let _ = context.val_u32(key_id);
293    }
294
295    struct SyntheticGgufFile {
296        path: PathBuf,
297    }
298
299    impl SyntheticGgufFile {
300        fn from_bytes(test_name: &str, bytes: &[u8]) -> Self {
301            use std::io::Write as _;
302
303            let path = std::env::temp_dir().join(format!(
304                "llama_cpp_bindings_synthetic_{}_{}.gguf",
305                std::process::id(),
306                test_name,
307            ));
308
309            let mut file = std::fs::File::create(&path).unwrap();
310            file.write_all(bytes).unwrap();
311
312            Self { path }
313        }
314
315        fn new(test_name: &str) -> Self {
316            let mut bytes: Vec<u8> = Vec::new();
317            bytes.extend_from_slice(b"GGUF");
318            bytes.extend_from_slice(&3u32.to_le_bytes());
319            bytes.extend_from_slice(&0u64.to_le_bytes());
320            bytes.extend_from_slice(&3u64.to_le_bytes());
321
322            let arch_key = b"general.architecture";
323            bytes.extend_from_slice(&(arch_key.len() as u64).to_le_bytes());
324            bytes.extend_from_slice(arch_key);
325            bytes.extend_from_slice(&8u32.to_le_bytes());
326            let arch_val = b"synthetic";
327            bytes.extend_from_slice(&(arch_val.len() as u64).to_le_bytes());
328            bytes.extend_from_slice(arch_val);
329
330            let i32_key = b"synthetic.i32_value";
331            bytes.extend_from_slice(&(i32_key.len() as u64).to_le_bytes());
332            bytes.extend_from_slice(i32_key);
333            bytes.extend_from_slice(&5u32.to_le_bytes());
334            bytes.extend_from_slice(&(-12345i32).to_le_bytes());
335
336            let u64_key = b"synthetic.u64_value";
337            bytes.extend_from_slice(&(u64_key.len() as u64).to_le_bytes());
338            bytes.extend_from_slice(u64_key);
339            bytes.extend_from_slice(&10u32.to_le_bytes());
340            bytes.extend_from_slice(&987_654_321u64.to_le_bytes());
341
342            Self::from_bytes(test_name, &bytes)
343        }
344    }
345
346    impl Drop for SyntheticGgufFile {
347        fn drop(&mut self) {
348            std::fs::remove_file(&self.path).ok();
349        }
350    }
351
352    #[test]
353    fn val_i32_and_val_u64_round_trip_through_synthetic_fixture() {
354        let fixture = SyntheticGgufFile::new("val_i32_and_val_u64_round_trip");
355
356        let context = GgufContext::from_file(&fixture.path).unwrap();
357
358        let i32_index = context.find_key("synthetic.i32_value").unwrap();
359        assert_eq!(context.kv_type(i32_index), Some(GgufType::Int32));
360        assert_eq!(context.val_i32(i32_index), -12345);
361
362        let u64_index = context.find_key("synthetic.u64_value").unwrap();
363        assert_eq!(context.kv_type(u64_index), Some(GgufType::Uint64));
364        assert_eq!(context.val_u64(u64_index), 987_654_321);
365    }
366
367    #[test]
368    fn val_str_returns_utf8_error_for_non_utf8_value() {
369        let mut bytes: Vec<u8> = Vec::new();
370        bytes.extend_from_slice(b"GGUF");
371        bytes.extend_from_slice(&3u32.to_le_bytes());
372        bytes.extend_from_slice(&0u64.to_le_bytes());
373        bytes.extend_from_slice(&1u64.to_le_bytes());
374
375        let value_key = b"synthetic.str_value";
376        bytes.extend_from_slice(&(value_key.len() as u64).to_le_bytes());
377        bytes.extend_from_slice(value_key);
378        bytes.extend_from_slice(&8u32.to_le_bytes());
379        let non_utf8_value: [u8; 2] = [0xFF, 0xFE];
380        bytes.extend_from_slice(&(non_utf8_value.len() as u64).to_le_bytes());
381        bytes.extend_from_slice(&non_utf8_value);
382
383        let fixture =
384            SyntheticGgufFile::from_bytes("val_str_returns_utf8_error_for_non_utf8_value", &bytes);
385        let context = GgufContext::from_file(&fixture.path).unwrap();
386
387        let value_index = context.find_key("synthetic.str_value").unwrap();
388        let err = context.val_str(value_index).unwrap_err();
389
390        assert_eq!(std::mem::discriminant(&err), utf8_error_disc());
391    }
392
393    #[test]
394    fn key_at_returns_utf8_error_for_non_utf8_key() {
395        let mut bytes: Vec<u8> = Vec::new();
396        bytes.extend_from_slice(b"GGUF");
397        bytes.extend_from_slice(&3u32.to_le_bytes());
398        bytes.extend_from_slice(&0u64.to_le_bytes());
399        bytes.extend_from_slice(&1u64.to_le_bytes());
400
401        let non_utf8_key: [u8; 2] = [0xFF, 0xFE];
402        bytes.extend_from_slice(&(non_utf8_key.len() as u64).to_le_bytes());
403        bytes.extend_from_slice(&non_utf8_key);
404        bytes.extend_from_slice(&5u32.to_le_bytes());
405        bytes.extend_from_slice(&42i32.to_le_bytes());
406
407        let fixture =
408            SyntheticGgufFile::from_bytes("key_at_returns_utf8_error_for_non_utf8_key", &bytes);
409        let context = GgufContext::from_file(&fixture.path).unwrap();
410
411        let err = context.key_at(0).unwrap_err();
412
413        assert_eq!(std::mem::discriminant(&err), utf8_error_disc());
414    }
415}