Skip to main content

llama_cpp_bindings/
json_schema_to_grammar.rs

1use std::ffi::{CStr, CString, c_char};
2
3use crate::error::JsonSchemaToGrammarError;
4use crate::ffi_error_reader::read_and_free_cpp_error;
5
6/// # Safety
7///
8/// On `LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_OK` the function reads and frees `out` as a
9/// null-terminated C string allocated by the wrapper, so `out` must be a valid such
10/// pointer for that status. On error statuses it reads and frees `error_ptr` via
11/// [`read_and_free_cpp_error`], which tolerates a null pointer.
12unsafe fn json_schema_to_grammar_status_to_result(
13    status: llama_cpp_bindings_sys::llama_rs_json_schema_to_grammar_status,
14    out: *mut c_char,
15    error_ptr: *mut c_char,
16) -> Result<String, JsonSchemaToGrammarError> {
17    match status {
18        llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_OK => {
19            let grammar_bytes = unsafe { CStr::from_ptr(out) }.to_bytes().to_vec();
20            unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out) };
21            Ok(String::from_utf8(grammar_bytes)?)
22        }
23        llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_ERROR_STRING_ALLOCATION_FAILED => {
24            Err(JsonSchemaToGrammarError::NotEnoughMemory)
25        }
26        llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_INVALID_SCHEMA => {
27            let message = unsafe { read_and_free_cpp_error(error_ptr) };
28            Err(JsonSchemaToGrammarError::InvalidSchema { message })
29        }
30        llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_VENDORED_THREW_CXX_EXCEPTION => {
31            let message = unsafe { read_and_free_cpp_error(error_ptr) };
32            Err(JsonSchemaToGrammarError::Reported { message })
33        }
34        other => {
35            unreachable!("llama_rs_json_schema_to_grammar returned unrecognized status {other}")
36        }
37    }
38}
39
40/// # Errors
41///
42/// Returns [`JsonSchemaToGrammarError`] if the schema string contains a NUL byte,
43/// the wrapper reports any non-OK status, or the returned grammar is not valid UTF-8.
44pub fn json_schema_to_grammar(schema_json: &str) -> Result<String, JsonSchemaToGrammarError> {
45    let schema_cstr = CString::new(schema_json)?;
46    let mut out: *mut c_char = std::ptr::null_mut();
47    let mut error_ptr: *mut c_char = std::ptr::null_mut();
48
49    let status = unsafe {
50        llama_cpp_bindings_sys::llama_rs_json_schema_to_grammar(
51            schema_cstr.as_ptr(),
52            false,
53            &raw mut out,
54            &raw mut error_ptr,
55        )
56    };
57
58    unsafe { json_schema_to_grammar_status_to_result(status, out, error_ptr) }
59}
60
61#[cfg(test)]
62mod tests {
63    use std::ffi::c_char;
64
65    use super::json_schema_to_grammar;
66    use super::json_schema_to_grammar_status_to_result;
67    use crate::error::JsonSchemaToGrammarError;
68
69    unsafe extern "C" {
70        fn strdup(source: *const c_char) -> *mut c_char;
71    }
72
73    #[test]
74    fn simple_object() {
75        let schema = r#"{"type": "object", "properties": {"name": {"type": "string"}}}"#;
76        let grammar = json_schema_to_grammar(schema).expect("schema converts to grammar");
77
78        assert!(!grammar.is_empty());
79    }
80
81    #[test]
82    fn null_byte_returns_schema_contains_nul_byte_error() {
83        use std::ffi::CString;
84
85        let schema = "{\x00}";
86        let err = json_schema_to_grammar(schema).unwrap_err();
87        let representative = JsonSchemaToGrammarError::SchemaContainsNulByte(
88            CString::new(b"a\0b".to_vec()).unwrap_err(),
89        );
90
91        assert_eq!(
92            std::mem::discriminant(&err),
93            std::mem::discriminant(&representative)
94        );
95    }
96
97    #[test]
98    fn simple_string() {
99        let schema = r#"{"type": "string"}"#;
100        let grammar = json_schema_to_grammar(schema).expect("schema converts to grammar");
101
102        assert!(!grammar.is_empty());
103    }
104
105    #[test]
106    fn invalid_json_returns_reported() {
107        let schema = "not valid json at all";
108        let err = json_schema_to_grammar(schema).unwrap_err();
109        let representative = JsonSchemaToGrammarError::Reported {
110            message: String::new(),
111        };
112
113        assert_eq!(
114            std::mem::discriminant(&err),
115            std::mem::discriminant(&representative)
116        );
117    }
118
119    #[test]
120    fn unresolved_ref_returns_invalid_schema() {
121        let schema = r##"{"$ref": "#/$defs/Missing"}"##;
122        let err = json_schema_to_grammar(schema).unwrap_err();
123        let representative = JsonSchemaToGrammarError::InvalidSchema {
124            message: String::new(),
125        };
126
127        assert_eq!(
128            std::mem::discriminant(&err),
129            std::mem::discriminant(&representative)
130        );
131    }
132
133    #[test]
134    fn invalid_schema_status_returns_invalid_schema() {
135        let result = unsafe {
136            json_schema_to_grammar_status_to_result(
137                llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_INVALID_SCHEMA,
138                std::ptr::null_mut(),
139                std::ptr::null_mut(),
140            )
141        };
142
143        assert_eq!(
144            result,
145            Err(JsonSchemaToGrammarError::InvalidSchema {
146                message: "unknown error".to_owned(),
147            })
148        );
149    }
150
151    #[test]
152    fn vendored_exception_status_returns_reported() {
153        let result = unsafe {
154            json_schema_to_grammar_status_to_result(
155                llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_VENDORED_THREW_CXX_EXCEPTION,
156                std::ptr::null_mut(),
157                std::ptr::null_mut(),
158            )
159        };
160
161        assert_eq!(
162            result,
163            Err(JsonSchemaToGrammarError::Reported {
164                message: "unknown error".to_owned(),
165            })
166        );
167    }
168
169    #[test]
170    fn allocation_failed_status_returns_not_enough_memory() {
171        let result = unsafe {
172            json_schema_to_grammar_status_to_result(
173                llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_ERROR_STRING_ALLOCATION_FAILED,
174                std::ptr::null_mut(),
175                std::ptr::null_mut(),
176            )
177        };
178
179        assert_eq!(result, Err(JsonSchemaToGrammarError::NotEnoughMemory));
180    }
181
182    #[test]
183    fn ok_status_with_non_utf8_grammar_returns_grammar_not_utf8() {
184        let invalid_utf8_grammar: [u8; 2] = [0xFF, 0];
185        let out = unsafe { strdup(invalid_utf8_grammar.as_ptr().cast::<c_char>()) };
186        assert!(!out.is_null(), "strdup must allocate a copy");
187
188        let result = unsafe {
189            json_schema_to_grammar_status_to_result(
190                llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_OK,
191                out,
192                std::ptr::null_mut(),
193            )
194        };
195        let representative =
196            JsonSchemaToGrammarError::GrammarNotUtf8(String::from_utf8(vec![0xFF]).unwrap_err());
197
198        assert_eq!(
199            std::mem::discriminant(&result.unwrap_err()),
200            std::mem::discriminant(&representative),
201        );
202    }
203
204    #[test]
205    fn ok_status_with_valid_utf8_grammar_returns_grammar_string() {
206        let grammar_text: &[u8; 14] = b"root ::= \"x\"\0\0";
207        let out = unsafe { strdup(grammar_text.as_ptr().cast::<c_char>()) };
208        assert!(!out.is_null(), "strdup must allocate a copy");
209
210        let result = unsafe {
211            json_schema_to_grammar_status_to_result(
212                llama_cpp_bindings_sys::LLAMA_RS_JSON_SCHEMA_TO_GRAMMAR_OK,
213                out,
214                std::ptr::null_mut(),
215            )
216        };
217
218        assert_eq!(result, Ok("root ::= \"x\"".to_owned()));
219    }
220
221    #[test]
222    #[should_panic(expected = "llama_rs_json_schema_to_grammar returned unrecognized status")]
223    fn unrecognized_status_panics() {
224        let _result = unsafe {
225            json_schema_to_grammar_status_to_result(
226                llama_cpp_bindings_sys::llama_rs_json_schema_to_grammar_status::MAX,
227                std::ptr::null_mut(),
228                std::ptr::null_mut(),
229            )
230        };
231    }
232}