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 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 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 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}