Skip to main content

llama_cpp_bindings/mtmd/
mtmd_context.rs

1use std::ffi::CString;
2use std::ffi::c_char;
3use std::ptr::NonNull;
4
5use crate::ffi_error_reader::read_and_free_cpp_error;
6use crate::model::LlamaModel;
7
8use super::mtmd_bitmap::MtmdBitmap;
9use super::mtmd_context_params::MtmdContextParams;
10use super::mtmd_encode_error::MtmdEncodeError;
11use super::mtmd_init_error::MtmdInitError;
12use super::mtmd_input_chunk::MtmdInputChunk;
13use super::mtmd_input_chunks::MtmdInputChunks;
14use super::mtmd_input_text::MtmdInputText;
15use super::mtmd_tokenize_error::MtmdTokenizeError;
16
17fn map_tokenize_status(
18    status: llama_cpp_bindings_sys::llama_rs_mtmd_tokenize_status,
19    undocumented_return_code: i32,
20    out_error: *mut c_char,
21) -> Result<(), MtmdTokenizeError> {
22    match status {
23        llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_OK => Ok(()),
24        llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_REPORTED_BITMAP_COUNT_DOES_NOT_MATCH_MARKER_COUNT => {
25            Err(MtmdTokenizeError::BitmapCountDoesNotMatchMarkerCount)
26        }
27        llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_REPORTED_IMAGE_PREPROCESSING_ERROR => {
28            Err(MtmdTokenizeError::MediaPreprocessingFailed)
29        }
30        llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_RETURNED_UNDOCUMENTED_NONZERO_CODE => {
31            Err(MtmdTokenizeError::UnknownStatus {
32                code: undocumented_return_code,
33            })
34        }
35        llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_ERROR_STRING_ALLOCATION_FAILED => {
36            Err(MtmdTokenizeError::NotEnoughMemory)
37        }
38        llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_THREW_CXX_EXCEPTION => {
39            let message = unsafe { read_and_free_cpp_error(out_error) };
40            Err(MtmdTokenizeError::Reported { message })
41        }
42        llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_NULL_BITMAPS_ARG_WHEN_NUM_BITMAPS_NONZERO => unreachable!("llama_rs_mtmd_tokenize NULL_BITMAPS_ARG: Rust always passes a non-null bitmaps pointer when count > 0"),
43        other => unreachable!("llama_rs_mtmd_tokenize returned unrecognized status: {other}"),
44    }
45}
46
47fn map_encode_chunk_status(
48    status: llama_cpp_bindings_sys::llama_rs_mtmd_encode_chunk_status,
49    vendored_return_code: i32,
50    out_error: *mut c_char,
51) -> Result<(), MtmdEncodeError> {
52    match status {
53        llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_OK => Ok(()),
54        llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_VENDORED_RETURNED_NONZERO_CODE => {
55            Err(MtmdEncodeError::EncodingFailed {
56                code: vendored_return_code,
57            })
58        }
59        llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_ERROR_STRING_ALLOCATION_FAILED => {
60            Err(MtmdEncodeError::NotEnoughMemory)
61        }
62        llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_VENDORED_THREW_CXX_EXCEPTION => {
63            let message = unsafe { read_and_free_cpp_error(out_error) };
64            Err(MtmdEncodeError::Reported { message })
65        }
66        other => unreachable!("llama_rs_mtmd_encode_chunk returned unrecognized status: {other}"),
67    }
68}
69
70fn map_init_from_file_status(
71    status: llama_cpp_bindings_sys::llama_rs_mtmd_init_from_file_status,
72    out_ctx: *mut llama_cpp_bindings_sys::mtmd_context,
73    out_error: *mut c_char,
74    mmproj_path: &str,
75) -> Result<MtmdContext, MtmdInitError> {
76    match status {
77        llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_OK => {
78            let context = NonNull::new(out_ctx).ok_or_else(|| MtmdInitError::Unloadable {
79                path: std::path::PathBuf::from(mmproj_path),
80            })?;
81            Ok(MtmdContext { context })
82        }
83        llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_VENDORED_RETURNED_NULL => {
84            Err(MtmdInitError::Unloadable {
85                path: std::path::PathBuf::from(mmproj_path),
86            })
87        }
88        llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_ERROR_STRING_ALLOCATION_FAILED => {
89            Err(MtmdInitError::NotEnoughMemory)
90        }
91        llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_VENDORED_THREW_CXX_EXCEPTION => {
92            let message = unsafe { read_and_free_cpp_error(out_error) };
93            Err(MtmdInitError::Reported { message })
94        }
95        other => {
96            unreachable!("llama_rs_mtmd_init_from_file returned unrecognized status: {other}")
97        }
98    }
99}
100
101#[derive(Debug)]
102pub struct MtmdContext {
103    pub context: NonNull<llama_cpp_bindings_sys::mtmd_context>,
104}
105
106unsafe impl Send for MtmdContext {}
107unsafe impl Sync for MtmdContext {}
108
109impl MtmdContext {
110    /// # Errors
111    ///
112    /// Returns an [`MtmdInitError`] variant matching the wrapper's status code.
113    pub fn init_from_file(
114        mmproj_path: &str,
115        text_model: &LlamaModel,
116        params: &MtmdContextParams,
117    ) -> Result<Self, MtmdInitError> {
118        let path_cstr = CString::new(mmproj_path)?;
119        let ctx_params = llama_cpp_bindings_sys::mtmd_context_params::from(params);
120
121        let mut out_ctx: *mut llama_cpp_bindings_sys::mtmd_context = std::ptr::null_mut();
122        let mut out_error: *mut c_char = std::ptr::null_mut();
123
124        let status = unsafe {
125            llama_cpp_bindings_sys::llama_rs_mtmd_init_from_file(
126                path_cstr.as_ptr(),
127                text_model.model.as_ptr(),
128                ctx_params,
129                &raw mut out_ctx,
130                &raw mut out_error,
131            )
132        };
133
134        map_init_from_file_status(status, out_ctx, out_error, mmproj_path)
135    }
136
137    #[must_use]
138    pub fn decode_use_non_causal(&self, chunk: &MtmdInputChunk) -> bool {
139        unsafe {
140            llama_cpp_bindings_sys::mtmd_decode_use_non_causal(
141                self.context.as_ptr(),
142                chunk.chunk.as_ptr(),
143            )
144        }
145    }
146
147    #[must_use]
148    pub fn decode_use_mrope(&self) -> bool {
149        unsafe { llama_cpp_bindings_sys::mtmd_decode_use_mrope(self.context.as_ptr()) }
150    }
151
152    #[must_use]
153    pub fn support_vision(&self) -> bool {
154        unsafe { llama_cpp_bindings_sys::mtmd_support_vision(self.context.as_ptr()) }
155    }
156
157    #[must_use]
158    pub fn support_audio(&self) -> bool {
159        unsafe { llama_cpp_bindings_sys::mtmd_support_audio(self.context.as_ptr()) }
160    }
161
162    #[must_use]
163    pub fn get_audio_sample_rate(&self) -> Option<u32> {
164        let rate =
165            unsafe { llama_cpp_bindings_sys::mtmd_get_audio_sample_rate(self.context.as_ptr()) };
166        (rate > 0).then_some(rate.unsigned_abs())
167    }
168
169    /// # Errors
170    ///
171    /// Returns an [`MtmdTokenizeError`] variant matching the wrapper's status code.
172    pub fn tokenize(
173        &self,
174        text: MtmdInputText,
175        bitmaps: &[&MtmdBitmap],
176    ) -> Result<MtmdInputChunks, MtmdTokenizeError> {
177        let chunks = MtmdInputChunks::new()?;
178        let text_cstring = CString::new(text.text)?;
179        let input_text = llama_cpp_bindings_sys::mtmd_input_text {
180            text: text_cstring.as_ptr(),
181            add_special: text.add_special,
182            parse_special: text.parse_special,
183        };
184
185        let bitmap_ptrs: Vec<*const llama_cpp_bindings_sys::mtmd_bitmap> = bitmaps
186            .iter()
187            .map(|bitmap| bitmap.bitmap.as_ptr().cast_const())
188            .collect();
189
190        let mut out_undocumented_return_code: i32 = 0;
191        let mut out_error: *mut c_char = std::ptr::null_mut();
192
193        let status = unsafe {
194            llama_cpp_bindings_sys::llama_rs_mtmd_tokenize(
195                self.context.as_ptr(),
196                chunks.chunks.as_ptr(),
197                &raw const input_text,
198                bitmap_ptrs.as_ptr().cast_mut(),
199                bitmaps.len(),
200                &raw mut out_undocumented_return_code,
201                &raw mut out_error,
202            )
203        };
204
205        map_tokenize_status(status, out_undocumented_return_code, out_error)?;
206        Ok(chunks)
207    }
208
209    /// # Errors
210    ///
211    /// Returns an [`MtmdEncodeError`] variant matching the wrapper's status code.
212    pub fn encode_chunk(&self, chunk: &MtmdInputChunk) -> Result<(), MtmdEncodeError> {
213        let mut out_vendored_return_code: i32 = 0;
214        let mut out_error: *mut c_char = std::ptr::null_mut();
215
216        let status = unsafe {
217            llama_cpp_bindings_sys::llama_rs_mtmd_encode_chunk(
218                self.context.as_ptr(),
219                chunk.chunk.as_ptr(),
220                &raw mut out_vendored_return_code,
221                &raw mut out_error,
222            )
223        };
224
225        map_encode_chunk_status(status, out_vendored_return_code, out_error)
226    }
227}
228
229impl Drop for MtmdContext {
230    fn drop(&mut self) {
231        unsafe { llama_cpp_bindings_sys::mtmd_free(self.context.as_ptr()) }
232    }
233}
234
235#[cfg(test)]
236mod unit_tests {
237    use super::map_encode_chunk_status;
238    use super::map_init_from_file_status;
239    use super::map_tokenize_status;
240    use crate::mtmd::mtmd_encode_error::MtmdEncodeError;
241    use crate::mtmd::mtmd_init_error::MtmdInitError;
242    use crate::mtmd::mtmd_tokenize_error::MtmdTokenizeError;
243
244    #[test]
245    fn tokenize_status_maps_bitmap_count_mismatch() {
246        let result = map_tokenize_status(
247            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_REPORTED_BITMAP_COUNT_DOES_NOT_MATCH_MARKER_COUNT,
248            0,
249            std::ptr::null_mut(),
250        );
251
252        assert_eq!(
253            result,
254            Err(MtmdTokenizeError::BitmapCountDoesNotMatchMarkerCount)
255        );
256    }
257
258    #[test]
259    fn tokenize_status_maps_media_preprocessing_failed() {
260        let result = map_tokenize_status(
261            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_REPORTED_IMAGE_PREPROCESSING_ERROR,
262            0,
263            std::ptr::null_mut(),
264        );
265
266        assert_eq!(result, Err(MtmdTokenizeError::MediaPreprocessingFailed));
267    }
268
269    #[test]
270    fn tokenize_status_maps_unknown_status_with_value() {
271        let result = map_tokenize_status(
272            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_RETURNED_UNDOCUMENTED_NONZERO_CODE,
273            42,
274            std::ptr::null_mut(),
275        );
276
277        assert_eq!(result, Err(MtmdTokenizeError::UnknownStatus { code: 42 }));
278    }
279
280    #[test]
281    fn tokenize_status_maps_ok_to_unit() {
282        let result = map_tokenize_status(
283            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_OK,
284            0,
285            std::ptr::null_mut(),
286        );
287
288        assert_eq!(result, Ok(()));
289    }
290
291    #[test]
292    fn encode_chunk_status_maps_ok_to_unit() {
293        let result = map_encode_chunk_status(
294            llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_OK,
295            0,
296            std::ptr::null_mut(),
297        );
298
299        assert_eq!(result, Ok(()));
300    }
301
302    #[test]
303    fn encode_chunk_status_maps_encoding_failed_with_code() {
304        let result = map_encode_chunk_status(
305            llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_VENDORED_RETURNED_NONZERO_CODE,
306            5,
307            std::ptr::null_mut(),
308        );
309
310        assert_eq!(result, Err(MtmdEncodeError::EncodingFailed { code: 5 }));
311    }
312
313    #[test]
314    fn tokenize_status_maps_string_allocation_failed_to_not_enough_memory() {
315        let result = map_tokenize_status(
316            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_ERROR_STRING_ALLOCATION_FAILED,
317            0,
318            std::ptr::null_mut(),
319        );
320
321        assert_eq!(result, Err(MtmdTokenizeError::NotEnoughMemory));
322    }
323
324    #[test]
325    fn tokenize_status_maps_cxx_exception_to_reported() {
326        let result = map_tokenize_status(
327            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_THREW_CXX_EXCEPTION,
328            0,
329            std::ptr::null_mut(),
330        );
331
332        assert_eq!(
333            result,
334            Err(MtmdTokenizeError::Reported {
335                message: "unknown error".to_string()
336            })
337        );
338    }
339
340    #[test]
341    #[should_panic(expected = "NULL_BITMAPS_ARG")]
342    fn tokenize_status_null_bitmaps_arg_panics() {
343        let _result = map_tokenize_status(
344            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_NULL_BITMAPS_ARG_WHEN_NUM_BITMAPS_NONZERO,
345            0,
346            std::ptr::null_mut(),
347        );
348    }
349
350    #[test]
351    #[should_panic(expected = "llama_rs_mtmd_tokenize returned unrecognized status")]
352    fn tokenize_status_unrecognized_panics() {
353        let _result = map_tokenize_status(
354            llama_cpp_bindings_sys::llama_rs_mtmd_tokenize_status::MAX,
355            0,
356            std::ptr::null_mut(),
357        );
358    }
359
360    #[test]
361    fn encode_chunk_status_maps_string_allocation_failed_to_not_enough_memory() {
362        let result = map_encode_chunk_status(
363            llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_ERROR_STRING_ALLOCATION_FAILED,
364            0,
365            std::ptr::null_mut(),
366        );
367
368        assert_eq!(result, Err(MtmdEncodeError::NotEnoughMemory));
369    }
370
371    #[test]
372    fn encode_chunk_status_maps_cxx_exception_to_reported() {
373        let result = map_encode_chunk_status(
374            llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_VENDORED_THREW_CXX_EXCEPTION,
375            0,
376            std::ptr::null_mut(),
377        );
378
379        assert_eq!(
380            result,
381            Err(MtmdEncodeError::Reported {
382                message: "unknown error".to_string()
383            })
384        );
385    }
386
387    #[test]
388    #[should_panic(expected = "llama_rs_mtmd_encode_chunk returned unrecognized status")]
389    fn encode_chunk_status_unrecognized_panics() {
390        let _result = map_encode_chunk_status(
391            llama_cpp_bindings_sys::llama_rs_mtmd_encode_chunk_status::MAX,
392            0,
393            std::ptr::null_mut(),
394        );
395    }
396
397    #[test]
398    fn init_from_file_status_ok_with_null_ctx_maps_unloadable() {
399        let result = map_init_from_file_status(
400            llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_OK,
401            std::ptr::null_mut(),
402            std::ptr::null_mut(),
403            "mmproj.gguf",
404        );
405
406        assert_eq!(
407            result.unwrap_err(),
408            MtmdInitError::Unloadable {
409                path: std::path::PathBuf::from("mmproj.gguf")
410            }
411        );
412    }
413
414    #[test]
415    fn init_from_file_status_maps_string_allocation_failed_to_not_enough_memory() {
416        let result = map_init_from_file_status(
417            llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_ERROR_STRING_ALLOCATION_FAILED,
418            std::ptr::null_mut(),
419            std::ptr::null_mut(),
420            "mmproj.gguf",
421        );
422
423        assert_eq!(result.unwrap_err(), MtmdInitError::NotEnoughMemory);
424    }
425
426    #[test]
427    fn init_from_file_status_maps_cxx_exception_to_reported() {
428        let result = map_init_from_file_status(
429            llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_VENDORED_THREW_CXX_EXCEPTION,
430            std::ptr::null_mut(),
431            std::ptr::null_mut(),
432            "mmproj.gguf",
433        );
434
435        assert_eq!(
436            result.unwrap_err(),
437            MtmdInitError::Reported {
438                message: "unknown error".to_string()
439            }
440        );
441    }
442
443    #[test]
444    #[should_panic(expected = "llama_rs_mtmd_init_from_file returned unrecognized status")]
445    fn init_from_file_status_unrecognized_panics() {
446        let _result = map_init_from_file_status(
447            llama_cpp_bindings_sys::llama_rs_mtmd_init_from_file_status::MAX,
448            std::ptr::null_mut(),
449            std::ptr::null_mut(),
450            "mmproj.gguf",
451        );
452    }
453}