llama_cpp_bindings/
gguf_context.rs1use 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 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 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 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 #[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 #[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 #[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 #[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 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}