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            text_len: text_cstring.as_bytes().len(),
182            add_special: text.add_special,
183            parse_special: text.parse_special,
184        };
185
186        let bitmap_ptrs: Vec<*const llama_cpp_bindings_sys::mtmd_bitmap> = bitmaps
187            .iter()
188            .map(|bitmap| bitmap.bitmap.as_ptr().cast_const())
189            .collect();
190
191        let mut out_undocumented_return_code: i32 = 0;
192        let mut out_error: *mut c_char = std::ptr::null_mut();
193
194        let status = unsafe {
195            llama_cpp_bindings_sys::llama_rs_mtmd_tokenize(
196                self.context.as_ptr(),
197                chunks.chunks.as_ptr(),
198                &raw const input_text,
199                bitmap_ptrs.as_ptr().cast_mut(),
200                bitmaps.len(),
201                &raw mut out_undocumented_return_code,
202                &raw mut out_error,
203            )
204        };
205
206        map_tokenize_status(status, out_undocumented_return_code, out_error)?;
207        Ok(chunks)
208    }
209
210    /// # Errors
211    ///
212    /// Returns an [`MtmdEncodeError`] variant matching the wrapper's status code.
213    pub fn encode_chunk(&self, chunk: &MtmdInputChunk) -> Result<(), MtmdEncodeError> {
214        let mut out_vendored_return_code: i32 = 0;
215        let mut out_error: *mut c_char = std::ptr::null_mut();
216
217        let status = unsafe {
218            llama_cpp_bindings_sys::llama_rs_mtmd_encode_chunk(
219                self.context.as_ptr(),
220                chunk.chunk.as_ptr(),
221                &raw mut out_vendored_return_code,
222                &raw mut out_error,
223            )
224        };
225
226        map_encode_chunk_status(status, out_vendored_return_code, out_error)
227    }
228}
229
230impl Drop for MtmdContext {
231    fn drop(&mut self) {
232        unsafe { llama_cpp_bindings_sys::mtmd_free(self.context.as_ptr()) }
233    }
234}
235
236#[cfg(test)]
237mod unit_tests {
238    use super::map_encode_chunk_status;
239    use super::map_init_from_file_status;
240    use super::map_tokenize_status;
241    use crate::mtmd::mtmd_encode_error::MtmdEncodeError;
242    use crate::mtmd::mtmd_init_error::MtmdInitError;
243    use crate::mtmd::mtmd_tokenize_error::MtmdTokenizeError;
244
245    #[test]
246    fn tokenize_status_maps_bitmap_count_mismatch() {
247        let result = map_tokenize_status(
248            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_REPORTED_BITMAP_COUNT_DOES_NOT_MATCH_MARKER_COUNT,
249            0,
250            std::ptr::null_mut(),
251        );
252
253        assert_eq!(
254            result,
255            Err(MtmdTokenizeError::BitmapCountDoesNotMatchMarkerCount)
256        );
257    }
258
259    #[test]
260    fn tokenize_status_maps_media_preprocessing_failed() {
261        let result = map_tokenize_status(
262            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_REPORTED_IMAGE_PREPROCESSING_ERROR,
263            0,
264            std::ptr::null_mut(),
265        );
266
267        assert_eq!(result, Err(MtmdTokenizeError::MediaPreprocessingFailed));
268    }
269
270    #[test]
271    fn tokenize_status_maps_unknown_status_with_value() {
272        let result = map_tokenize_status(
273            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_RETURNED_UNDOCUMENTED_NONZERO_CODE,
274            42,
275            std::ptr::null_mut(),
276        );
277
278        assert_eq!(result, Err(MtmdTokenizeError::UnknownStatus { code: 42 }));
279    }
280
281    #[test]
282    fn tokenize_status_maps_ok_to_unit() {
283        let result = map_tokenize_status(
284            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_OK,
285            0,
286            std::ptr::null_mut(),
287        );
288
289        assert_eq!(result, Ok(()));
290    }
291
292    #[test]
293    fn encode_chunk_status_maps_ok_to_unit() {
294        let result = map_encode_chunk_status(
295            llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_OK,
296            0,
297            std::ptr::null_mut(),
298        );
299
300        assert_eq!(result, Ok(()));
301    }
302
303    #[test]
304    fn encode_chunk_status_maps_encoding_failed_with_code() {
305        let result = map_encode_chunk_status(
306            llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_VENDORED_RETURNED_NONZERO_CODE,
307            5,
308            std::ptr::null_mut(),
309        );
310
311        assert_eq!(result, Err(MtmdEncodeError::EncodingFailed { code: 5 }));
312    }
313
314    #[test]
315    fn tokenize_status_maps_string_allocation_failed_to_not_enough_memory() {
316        let result = map_tokenize_status(
317            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_ERROR_STRING_ALLOCATION_FAILED,
318            0,
319            std::ptr::null_mut(),
320        );
321
322        assert_eq!(result, Err(MtmdTokenizeError::NotEnoughMemory));
323    }
324
325    #[test]
326    fn tokenize_status_maps_cxx_exception_to_reported() {
327        let result = map_tokenize_status(
328            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_VENDORED_THREW_CXX_EXCEPTION,
329            0,
330            std::ptr::null_mut(),
331        );
332
333        assert_eq!(
334            result,
335            Err(MtmdTokenizeError::Reported {
336                message: "unknown error".to_string()
337            })
338        );
339    }
340
341    #[test]
342    #[should_panic(expected = "NULL_BITMAPS_ARG")]
343    fn tokenize_status_null_bitmaps_arg_panics() {
344        let _result = map_tokenize_status(
345            llama_cpp_bindings_sys::LLAMA_RS_MTMD_TOKENIZE_NULL_BITMAPS_ARG_WHEN_NUM_BITMAPS_NONZERO,
346            0,
347            std::ptr::null_mut(),
348        );
349    }
350
351    #[test]
352    #[should_panic(expected = "llama_rs_mtmd_tokenize returned unrecognized status")]
353    fn tokenize_status_unrecognized_panics() {
354        let _result = map_tokenize_status(
355            llama_cpp_bindings_sys::llama_rs_mtmd_tokenize_status::MAX,
356            0,
357            std::ptr::null_mut(),
358        );
359    }
360
361    #[test]
362    fn encode_chunk_status_maps_string_allocation_failed_to_not_enough_memory() {
363        let result = map_encode_chunk_status(
364            llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_ERROR_STRING_ALLOCATION_FAILED,
365            0,
366            std::ptr::null_mut(),
367        );
368
369        assert_eq!(result, Err(MtmdEncodeError::NotEnoughMemory));
370    }
371
372    #[test]
373    fn encode_chunk_status_maps_cxx_exception_to_reported() {
374        let result = map_encode_chunk_status(
375            llama_cpp_bindings_sys::LLAMA_RS_MTMD_ENCODE_CHUNK_VENDORED_THREW_CXX_EXCEPTION,
376            0,
377            std::ptr::null_mut(),
378        );
379
380        assert_eq!(
381            result,
382            Err(MtmdEncodeError::Reported {
383                message: "unknown error".to_string()
384            })
385        );
386    }
387
388    #[test]
389    #[should_panic(expected = "llama_rs_mtmd_encode_chunk returned unrecognized status")]
390    fn encode_chunk_status_unrecognized_panics() {
391        let _result = map_encode_chunk_status(
392            llama_cpp_bindings_sys::llama_rs_mtmd_encode_chunk_status::MAX,
393            0,
394            std::ptr::null_mut(),
395        );
396    }
397
398    #[test]
399    fn init_from_file_status_ok_with_null_ctx_maps_unloadable() {
400        let result = map_init_from_file_status(
401            llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_OK,
402            std::ptr::null_mut(),
403            std::ptr::null_mut(),
404            "mmproj.gguf",
405        );
406
407        assert_eq!(
408            result.unwrap_err(),
409            MtmdInitError::Unloadable {
410                path: std::path::PathBuf::from("mmproj.gguf")
411            }
412        );
413    }
414
415    #[test]
416    fn init_from_file_status_maps_string_allocation_failed_to_not_enough_memory() {
417        let result = map_init_from_file_status(
418            llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_ERROR_STRING_ALLOCATION_FAILED,
419            std::ptr::null_mut(),
420            std::ptr::null_mut(),
421            "mmproj.gguf",
422        );
423
424        assert_eq!(result.unwrap_err(), MtmdInitError::NotEnoughMemory);
425    }
426
427    #[test]
428    fn init_from_file_status_maps_cxx_exception_to_reported() {
429        let result = map_init_from_file_status(
430            llama_cpp_bindings_sys::LLAMA_RS_MTMD_INIT_FROM_FILE_VENDORED_THREW_CXX_EXCEPTION,
431            std::ptr::null_mut(),
432            std::ptr::null_mut(),
433            "mmproj.gguf",
434        );
435
436        assert_eq!(
437            result.unwrap_err(),
438            MtmdInitError::Reported {
439                message: "unknown error".to_string()
440            }
441        );
442    }
443
444    #[test]
445    #[should_panic(expected = "llama_rs_mtmd_init_from_file returned unrecognized status")]
446    fn init_from_file_status_unrecognized_panics() {
447        let _result = map_init_from_file_status(
448            llama_cpp_bindings_sys::llama_rs_mtmd_init_from_file_status::MAX,
449            std::ptr::null_mut(),
450            std::ptr::null_mut(),
451            "mmproj.gguf",
452        );
453    }
454}