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