diff --git a/.claude/rules/subproject-llama-cpp-bindings-types.md b/.claude/rules/subproject-llama-cpp-bindings-types.md index b2443d0bf..438cbad74 100644 --- a/.claude/rules/subproject-llama-cpp-bindings-types.md +++ b/.claude/rules/subproject-llama-cpp-bindings-types.md @@ -5,5 +5,5 @@ paths: # `llama-cpp-bindings-types` Context -- The purposse of `llama-cpp-bindings-types` is to provide a thin layer of types that do not need to rely on `llama.cpp` vendored library itself +- The purposse of `llama-cpp-bindings-types` is to provide a thin layer of types that do not need to rely on the `llama.cpp` library itself - `llama-cpp-bindings-types` must not depend on llama.cpp bindings themselves diff --git a/Cargo.lock b/Cargo.lock index 04d2ec5cd..01a64b40a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1163,7 +1163,7 @@ checksum = "11d3d7f243d5c5a8b9bb5d6dd2b1602c0cb0b9db1621bafc7ed66e35ff9fe092" [[package]] name = "llama-cpp-bindings" -version = "0.14.0" +version = "0.15.0" dependencies = [ "encoding_rs", "enumflags2", @@ -1177,6 +1177,7 @@ dependencies = [ "llguidance", "log", "nom 8.0.0", + "once_cell", "serde_json", "serial_test", "thiserror", @@ -1185,7 +1186,7 @@ dependencies = [ [[package]] name = "llama-cpp-bindings-build" -version = "0.14.0" +version = "0.15.0" dependencies = [ "bindgen", "cc", @@ -1197,14 +1198,14 @@ dependencies = [ [[package]] name = "llama-cpp-bindings-sys" -version = "0.14.0" +version = "0.15.0" dependencies = [ "llama-cpp-bindings-build", ] [[package]] name = "llama-cpp-bindings-tests" -version = "0.14.0" +version = "0.15.0" dependencies = [ "anyhow", "encoding_rs", @@ -1216,7 +1217,7 @@ dependencies = [ [[package]] name = "llama-cpp-bindings-types" -version = "0.14.0" +version = "0.15.0" dependencies = [ "serde", "serde_json", @@ -1225,14 +1226,14 @@ dependencies = [ [[package]] name = "llama-cpp-error-recorder" -version = "0.14.0" +version = "0.15.0" dependencies = [ "log", ] [[package]] name = "llama-cpp-ffi-status" -version = "0.14.0" +version = "0.15.0" dependencies = [ "llama-cpp-bindings-sys", "llama-cpp-wrapper-error-fixture", @@ -1241,7 +1242,7 @@ dependencies = [ [[package]] name = "llama-cpp-gbnf" -version = "0.14.0" +version = "0.15.0" dependencies = [ "llama-cpp-bindings-sys", "llama-cpp-ffi-status", @@ -1250,11 +1251,11 @@ dependencies = [ [[package]] name = "llama-cpp-log-decoder" -version = "0.14.0" +version = "0.15.0" [[package]] name = "llama-cpp-test-harness" -version = "0.14.0" +version = "0.15.0" dependencies = [ "anyhow", "hf-hub", @@ -1267,7 +1268,7 @@ dependencies = [ [[package]] name = "llama-cpp-test-harness-macros" -version = "0.14.0" +version = "0.15.0" dependencies = [ "proc-macro2", "quote", @@ -1276,14 +1277,14 @@ dependencies = [ [[package]] name = "llama-cpp-wrapper-error-fixture" -version = "0.14.0" +version = "0.15.0" dependencies = [ "llama-cpp-bindings-sys", ] [[package]] name = "llama-cpp-wrapper-sources" -version = "0.14.0" +version = "0.15.0" dependencies = [ "serde", "serde_json", diff --git a/Cargo.toml b/Cargo.toml index 9f59280e8..624722a16 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -18,7 +18,7 @@ members = [ [workspace.package] edition = "2024" -version = "0.14.0" +version = "0.15.0" license = "Apache-2.0" repository = "https://github.com/intentee/llama-cpp-bindings" @@ -33,21 +33,22 @@ find_cuda_helper = "=0.2.0" hf-hub = "=0.5.0" inventory = "=0.3.24" libtest-mimic = "=0.8.2" -llama-cpp-bindings = { path = "llama-cpp-bindings", version = "=0.14.0" } -llama-cpp-bindings-build = { path = "llama-cpp-bindings-build", version = "=0.14.0" } -llama-cpp-bindings-sys = { path = "llama-cpp-bindings-sys", version = "=0.14.0" } -llama-cpp-bindings-types = { path = "llama-cpp-bindings-types", version = "=0.14.0" } -llama-cpp-error-recorder = { path = "llama-cpp-error-recorder", version = "=0.14.0" } -llama-cpp-ffi-status = { path = "llama-cpp-ffi-status", version = "=0.14.0" } -llama-cpp-gbnf = { path = "llama-cpp-gbnf", version = "=0.14.0" } -llama-cpp-log-decoder = { path = "llama-cpp-log-decoder", version = "=0.14.0" } -llama-cpp-test-harness = { path = "llama-cpp-test-harness", version = "=0.14.0" } -llama-cpp-test-harness-macros = { path = "llama-cpp-test-harness-macros", version = "=0.14.0" } +llama-cpp-bindings = { path = "llama-cpp-bindings", version = "=0.15.0" } +llama-cpp-bindings-build = { path = "llama-cpp-bindings-build", version = "=0.15.0" } +llama-cpp-bindings-sys = { path = "llama-cpp-bindings-sys", version = "=0.15.0" } +llama-cpp-bindings-types = { path = "llama-cpp-bindings-types", version = "=0.15.0" } +llama-cpp-error-recorder = { path = "llama-cpp-error-recorder", version = "=0.15.0" } +llama-cpp-ffi-status = { path = "llama-cpp-ffi-status", version = "=0.15.0" } +llama-cpp-gbnf = { path = "llama-cpp-gbnf", version = "=0.15.0" } +llama-cpp-log-decoder = { path = "llama-cpp-log-decoder", version = "=0.15.0" } +llama-cpp-test-harness = { path = "llama-cpp-test-harness", version = "=0.15.0" } +llama-cpp-test-harness-macros = { path = "llama-cpp-test-harness-macros", version = "=0.15.0" } llama-cpp-wrapper-error-fixture = { path = "llama-cpp-wrapper-error-fixture" } -llama-cpp-wrapper-sources = { path = "llama-cpp-wrapper-sources", version = "=0.14.0" } +llama-cpp-wrapper-sources = { path = "llama-cpp-wrapper-sources", version = "=0.15.0" } llguidance = "=1.7.0" log = "=0.4.29" nom = "=8.0.0" +once_cell = "1.21" proc-macro2 = "=1.0.106" quote = "=1.0.45" serde = { version = "=1.0.228", features = ["derive"] } diff --git a/Makefile b/Makefile index f8b09df6e..886920b2a 100644 --- a/Makefile +++ b/Makefile @@ -29,7 +29,7 @@ WRAPPER_SOURCES_CRATE_FILES = \ EMIT_WRAPPER_BUILD_INPUTS = cargo run --quiet --package $(WRAPPER_SOURCES_CRATE) -- \ $(CURDIR)/llama-cpp-bindings-sys $(COMPILE_COMMANDS) $(WRAPPER_SOURCES_RESPONSE_FILE) -VENDORED_SUPPRESSIONS = \ +LLAMA_CPP_AND_GSL_SUPPRESSIONS = \ --suppress='*:*llama-cpp-bindings-sys/llama.cpp/*' \ --suppress='*:*llama-cpp-bindings-sys/GSL/*' @@ -115,7 +115,7 @@ lint.cpp.clang-tidy: $(COMPILE_COMMANDS) $(WRAPPER_SOURCES_RESPONSE_FILE) lint.cpp.cppcheck: $(COMPILE_COMMANDS) $(CPPCHECK) $(CPPCHECK) --project=$(COMPILE_COMMANDS) --enable=all --inconclusive \ --check-level=exhaustive --error-exitcode=1 \ - $(VENDORED_SUPPRESSIONS) \ + $(LLAMA_CPP_AND_GSL_SUPPRESSIONS) \ --suppress=missingIncludeSystem .PHONY: test diff --git a/llama-cpp-bindings-sys/wrapper_chat_parse.cpp b/llama-cpp-bindings-sys/wrapper_chat_parse.cpp index 9525c9ccf..b52d1534b 100644 --- a/llama-cpp-bindings-sys/wrapper_chat_parse.cpp +++ b/llama-cpp-bindings-sys/wrapper_chat_parse.cpp @@ -30,6 +30,61 @@ void dup_or_set_alloc_flag(const std::string & source, char ** out_dup, bool * o *out_dup = llama_rs_dup_string(source); *out_alloc_failed = (*out_dup == nullptr); } + +auto report_current_parse_exception( + char ** out_error, + llama_rs_parse_chat_message_status thrown_status) -> llama_rs_parse_chat_message_status { + try { + throw; + } catch (const std::bad_alloc &) { + return LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY; + } catch (const std::exception & ex) { + *out_error = llama_rs_dup_string(std::string(ex.what())); + } catch (...) { + *out_error = llama_rs_dup_string(std::string("unknown c++ exception")); + } + if (*out_error == nullptr) { + return LLAMA_RS_PARSE_CHAT_MESSAGE_ERROR_STRING_ALLOCATION_FAILED; + } + return thrown_status; +} + +auto build_tools_parser(const llama_rs_chat_parser & parser, const char * tools_json) -> common_peg_arena { + autoparser::generation_params inputs; + + if ((tools_json != nullptr) && *tools_json != '\0') { + inputs.tools = common_json::parse(tools_json); + } else { + inputs.tools = common_json::array(); + } + + return parser.parser.build_parser(inputs, std::string()); +} + +auto parse_with_tools_parser( + const common_peg_arena & chat_parser, + const char * input, + int is_partial, + llama_rs_parsed_chat_handle * out_handle, + char ** out_error) -> llama_rs_parse_chat_message_status { + try { + common_chat_parser_params parser_params; + parser_params.format = COMMON_CHAT_FORMAT_PEG_NATIVE; + + common_chat_msg parsed = + common_chat_peg_parse(chat_parser, input, is_partial != 0, parser_params); + + auto handle = std::make_unique(); + handle->message = std::move(parsed); + + *out_handle = handle.release(); + + return LLAMA_RS_PARSE_CHAT_MESSAGE_OK; + } catch (...) { + return report_current_parse_exception( + out_error, LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION); + } +} } // namespace extern "C" auto llama_rs_chat_parser_create( @@ -148,42 +203,12 @@ extern "C" auto llama_rs_parse_chat_message( } try { - autoparser::generation_params inputs; - - if ((tools_json != nullptr) && *tools_json != '\0') { - inputs.tools = common_json::parse(tools_json); - } else { - inputs.tools = common_json::array(); - } - - common_peg_arena const chat_parser = parser->parser.build_parser(inputs, std::string()); - - common_chat_parser_params parser_params; - parser_params.format = COMMON_CHAT_FORMAT_PEG_NATIVE; - - common_chat_msg parsed = - common_chat_peg_parse(chat_parser, input, is_partial != 0, parser_params); + common_peg_arena const chat_parser = build_tools_parser(*parser, tools_json); - auto handle = std::make_unique(); - handle->message = std::move(parsed); - - *out_handle = handle.release(); - - return LLAMA_RS_PARSE_CHAT_MESSAGE_OK; - } catch (const std::bad_alloc &) { - return LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY; - } catch (const std::exception & ex) { - *out_error = llama_rs_dup_string(std::string(ex.what())); - if (*out_error == nullptr) { - return LLAMA_RS_PARSE_CHAT_MESSAGE_ERROR_STRING_ALLOCATION_FAILED; - } - return LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION; + return parse_with_tools_parser(chat_parser, input, is_partial, out_handle, out_error); } catch (...) { - *out_error = llama_rs_dup_string(std::string("unknown c++ exception")); - if (*out_error == nullptr) { - return LLAMA_RS_PARSE_CHAT_MESSAGE_ERROR_STRING_ALLOCATION_FAILED; - } - return LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION; + return report_current_parse_exception( + out_error, LLAMA_RS_PARSE_CHAT_MESSAGE_TOOLS_PARSER_BUILD_THREW_CXX_EXCEPTION); } } diff --git a/llama-cpp-bindings-sys/wrapper_chat_parse.h b/llama-cpp-bindings-sys/wrapper_chat_parse.h index 8c097f0bf..afc62cf71 100644 --- a/llama-cpp-bindings-sys/wrapper_chat_parse.h +++ b/llama-cpp-bindings-sys/wrapper_chat_parse.h @@ -52,6 +52,7 @@ typedef enum llama_rs_parse_chat_message_status { LLAMA_RS_PARSE_CHAT_MESSAGE_ERROR_STRING_ALLOCATION_FAILED, LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY, LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION, + LLAMA_RS_PARSE_CHAT_MESSAGE_TOOLS_PARSER_BUILD_THREW_CXX_EXCEPTION, } llama_rs_parse_chat_message_status; llama_rs_parse_chat_message_status llama_rs_parse_chat_message( diff --git a/llama-cpp-bindings-tests/tests/chat_protocol.rs b/llama-cpp-bindings-tests/tests/chat_protocol.rs index c8899c817..09dd6d8f8 100644 --- a/llama-cpp-bindings-tests/tests/chat_protocol.rs +++ b/llama-cpp-bindings-tests/tests/chat_protocol.rs @@ -2,6 +2,8 @@ use anyhow::Result; use anyhow::bail; use llama_cpp_bindings::ChatMessageParseOutcome; use llama_cpp_bindings::ChatTemplateError; +use llama_cpp_bindings::ChatTools; +use llama_cpp_bindings::ParseChatMessageError; use llama_cpp_bindings::model::LlamaChatMessage; use llama_cpp_bindings_tests::build_user_prompt_with_media_marker::build_user_prompt_with_media_marker; use llama_cpp_test_harness::LlamaFixture; @@ -247,9 +249,11 @@ fn chat_template_with_nonexistent_name_returns_error(fixture: &LlamaFixture<'_>) n_ubatch = 64, )] fn parses_pure_content_response(fixture: &LlamaFixture<'_>) -> Result<()> { - let outcome = fixture - .model - .parse_chat_message("[]", "hello world", false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json("[]".to_owned())?, + "hello world", + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for plain content; got Unrecognized"); @@ -295,7 +299,10 @@ fn parses_pure_content_response(fixture: &LlamaFixture<'_>) -> Result<()> { )] fn parses_reasoning_section_into_reasoning_content(fixture: &LlamaFixture<'_>) -> Result<()> { let input = "step one, step two\n\nactual response"; - let outcome = fixture.model.parse_chat_message("[]", input, false)?; + let outcome = + fixture + .model + .parse_chat_message(&ChatTools::from_json("[]".to_owned())?, input, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for reasoning section; got Unrecognized"); @@ -343,7 +350,10 @@ fn parses_reasoning_section_into_reasoning_content(fixture: &LlamaFixture<'_>) - n_ubatch = 64, )] fn parses_empty_input_yields_empty_message(fixture: &LlamaFixture<'_>) -> Result<()> { - let outcome = fixture.model.parse_chat_message("[]", "", false)?; + let outcome = + fixture + .model + .parse_chat_message(&ChatTools::from_json("[]".to_owned())?, "", false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for empty input; got Unrecognized"); @@ -353,30 +363,6 @@ fn parses_empty_input_yields_empty_message(fixture: &LlamaFixture<'_>) -> Result Ok(()) } -#[llama_test( - model_source = HuggingFace("unsloth/DeepSeek-R1-Distill-Llama-8B-GGUF", "DeepSeek-R1-Distill-Llama-8B-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/GLM-4.7-Flash-GGUF", "GLM-4.7-Flash-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] #[llama_test( model_source = HuggingFace("unsloth/Qwen3.6-35B-A3B-GGUF", "Qwen3.6-35B-A3B-UD-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -385,134 +371,24 @@ fn parses_empty_input_yields_empty_message(fixture: &LlamaFixture<'_>) -> Result n_batch = 128, n_ubatch = 64, )] -fn parses_malformed_tools_json_returns_tools_json_invalid_error( - fixture: &LlamaFixture<'_>, -) -> Result<()> { - let result = fixture - .model - .parse_chat_message("not_a_json[}", "hello", false); - - assert!(matches!( - result, - Err(llama_cpp_bindings::ParseChatMessageError::ToolsJsonInvalid( - _ - )) - )); - Ok(()) -} - -#[llama_test( - model_source = HuggingFace("unsloth/DeepSeek-R1-Distill-Llama-8B-GGUF", "DeepSeek-R1-Distill-Llama-8B-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/GLM-4.7-Flash-GGUF", "GLM-4.7-Flash-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/Qwen3.6-35B-A3B-GGUF", "Qwen3.6-35B-A3B-UD-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -fn parses_non_array_tools_json_returns_tools_json_not_array_error( - fixture: &LlamaFixture<'_>, -) -> Result<()> { - let result = fixture - .model - .parse_chat_message("{\"foo\": 1}", "hello", false); - - assert!(matches!( - result, - Err(llama_cpp_bindings::ParseChatMessageError::ToolsJsonNotArray) - )); - Ok(()) -} - -#[llama_test( - model_source = HuggingFace("unsloth/DeepSeek-R1-Distill-Llama-8B-GGUF", "DeepSeek-R1-Distill-Llama-8B-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/GLM-4.7-Flash-GGUF", "GLM-4.7-Flash-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/Qwen3.6-35B-A3B-GGUF", "Qwen3.6-35B-A3B-UD-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -fn parses_with_tools_null_byte_reports_the_nul_byte_not_a_json_syntax_error( +fn parses_with_input_null_byte_reports_the_input_as_the_source( fixture: &LlamaFixture<'_>, ) -> Result<()> { - let result = fixture - .model - .parse_chat_message("[]\0extra", "hello", false); + let result = fixture.model.parse_chat_message( + &ChatTools::from_json("[]".to_owned())?, + "hello\0world", + false, + ); - let Err(llama_cpp_bindings::ParseChatMessageError::ToolsJsonContainsNulByte(nul_error)) = - result - else { - anyhow::bail!("a NUL byte in tools_json must be named as such, not reported as bad JSON"); + let Err(ParseChatMessageError::InputContainsNulByte(nul_error)) = result else { + anyhow::bail!("a NUL byte in the message must be reported against the message"); }; - assert_eq!(nul_error.nul_position(), 2); + assert_eq!(nul_error.nul_position(), 5); Ok(()) } -#[llama_test( - model_source = HuggingFace("unsloth/DeepSeek-R1-Distill-Llama-8B-GGUF", "DeepSeek-R1-Distill-Llama-8B-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -#[llama_test( - model_source = HuggingFace("unsloth/GLM-4.7-Flash-GGUF", "GLM-4.7-Flash-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] #[llama_test( model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -521,51 +397,23 @@ fn parses_with_tools_null_byte_reports_the_nul_byte_not_a_json_syntax_error( n_batch = 128, n_ubatch = 64, )] -#[llama_test( - model_source = HuggingFace("unsloth/Qwen3.6-35B-A3B-GGUF", "Qwen3.6-35B-A3B-UD-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -fn parses_with_tools_json_null_byte_reports_the_tools_as_the_source( - fixture: &LlamaFixture<'_>, -) -> Result<()> { - let result = fixture.model.parse_chat_message("[]\0", "hello", false); - - let Err(llama_cpp_bindings::ParseChatMessageError::ToolsJsonContainsNulByte(nul_error)) = - result - else { - anyhow::bail!("a NUL byte in tools_json must be reported against tools_json"); - }; - - assert_eq!(nul_error.nul_position(), 2); - - Ok(()) -} - -#[llama_test( - model_source = HuggingFace("unsloth/Qwen3.6-35B-A3B-GGUF", "Qwen3.6-35B-A3B-UD-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 512, - n_batch = 128, - n_ubatch = 64, -)] -fn parses_with_input_null_byte_reports_the_input_as_the_source( +fn parses_with_a_tool_missing_its_function_name_reports_a_tools_parser_build_failure( fixture: &LlamaFixture<'_>, ) -> Result<()> { - let result = fixture - .model - .parse_chat_message("[]", "hello\0world", false); - - let Err(llama_cpp_bindings::ParseChatMessageError::InputContainsNulByte(nul_error)) = result - else { - anyhow::bail!("a NUL byte in the message must be reported against the message"); - }; + let result = fixture.model.parse_chat_message( + &ChatTools::from_json( + r#"[{"type":"function","function":{"description":"reports the weather"}}]"#.to_owned(), + )?, + "hello", + false, + ); - assert_eq!(nul_error.nul_position(), 5); + assert_eq!( + result.unwrap_err(), + ParseChatMessageError::ToolsParserBuildFailed { + message: "[json.exception.out_of_range.403] key 'name' not found".to_owned(), + } + ); Ok(()) } diff --git a/llama-cpp-bindings-tests/tests/context_state.rs b/llama-cpp-bindings-tests/tests/context_state.rs index b9efaac74..207640c2e 100644 --- a/llama-cpp-bindings-tests/tests/context_state.rs +++ b/llama-cpp-bindings-tests/tests/context_state.rs @@ -123,6 +123,24 @@ fn context_creation_and_properties(fixture: &LlamaFixture<'_>) -> Result<()> { Ok(()) } +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 2048, + n_batch = 512, + n_ubatch = 512, + n_seq_max = 4, +)] +fn n_ctx_seq_splits_the_context_between_sequences(fixture: &LlamaFixture<'_>) -> Result<()> { + let context = fixture.build_context()?; + + assert_eq!(context.n_ctx(), 2048); + assert_eq!(context.n_ctx_seq(), 512); + + Ok(()) +} + #[llama_test( model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), n_gpu_layers = 999, diff --git a/llama-cpp-bindings-tests/tests/main.rs b/llama-cpp-bindings-tests/tests/main.rs index 6b07d941a..f147a0f90 100644 --- a/llama-cpp-bindings-tests/tests/main.rs +++ b/llama-cpp-bindings-tests/tests/main.rs @@ -7,6 +7,7 @@ mod chat_protocol; mod context_state; mod embedding_models; mod generation_control; +mod model_derived_data_caching; mod model_introspection; mod model_loading_errors; mod multimodal_audio; diff --git a/llama-cpp-bindings-tests/tests/model_derived_data_caching.rs b/llama-cpp-bindings-tests/tests/model_derived_data_caching.rs new file mode 100644 index 000000000..59bacd40a --- /dev/null +++ b/llama-cpp-bindings-tests/tests/model_derived_data_caching.rs @@ -0,0 +1,64 @@ +use std::ptr; +use std::sync::Arc; + +use anyhow::Context; +use anyhow::Result; +use llama_cpp_test_harness::LlamaFixture; +use llama_cpp_test_harness::llama_test; + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 2048, + n_batch = 512, + n_ubatch = 128 +)] +fn approximate_tok_env_is_cached_across_calls(fixture: &LlamaFixture<'_>) -> Result<()> { + let first = fixture.model.approximate_tok_env()?; + let second = fixture.model.approximate_tok_env()?; + + assert!(Arc::ptr_eq(&first, &second)); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 2048, + n_batch = 512, + n_ubatch = 128 +)] +fn streaming_markers_are_cached_across_calls(fixture: &LlamaFixture<'_>) -> Result<()> { + let first = fixture.model.streaming_markers()?; + let second = fixture.model.streaming_markers()?; + + assert!(Arc::ptr_eq(&first, &second)); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 2048, + n_batch = 512, + n_ubatch = 128 +)] +fn reasoning_markers_are_cached_across_calls(fixture: &LlamaFixture<'_>) -> Result<()> { + let first = fixture + .model + .reasoning_markers()? + .context("Qwen3.5 must expose reasoning markers")?; + let second = fixture + .model + .reasoning_markers()? + .context("Qwen3.5 must expose reasoning markers")?; + + assert!(ptr::eq(first, second)); + + Ok(()) +} diff --git a/llama-cpp-bindings-tests/tests/model_introspection.rs b/llama-cpp-bindings-tests/tests/model_introspection.rs index c04c50ed3..c19e942bb 100644 --- a/llama-cpp-bindings-tests/tests/model_introspection.rs +++ b/llama-cpp-bindings-tests/tests/model_introspection.rs @@ -1760,20 +1760,3 @@ fn debug_format_includes_struct_name_and_model_field(fixture: &LlamaFixture<'_>) Ok(()) } - -#[llama_test( - model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), - n_gpu_layers = 999, - load_mode = Mmap, - n_ctx = 2048, - n_batch = 512, - n_ubatch = 128 -)] -fn approximate_tok_env_is_cached_across_calls(fixture: &LlamaFixture<'_>) -> Result<()> { - let first = fixture.model.approximate_tok_env()?; - let second = fixture.model.approximate_tok_env()?; - - assert!(std::sync::Arc::ptr_eq(&first, &second)); - - Ok(()) -} diff --git a/llama-cpp-bindings-tests/tests/structured_chat_output.rs b/llama-cpp-bindings-tests/tests/structured_chat_output.rs index 5ffb33628..a22fe85d6 100644 --- a/llama-cpp-bindings-tests/tests/structured_chat_output.rs +++ b/llama-cpp-bindings-tests/tests/structured_chat_output.rs @@ -1,6 +1,7 @@ use anyhow::Result; use anyhow::bail; use llama_cpp_bindings::ChatMessageParseOutcome; +use llama_cpp_bindings::ChatTools; use llama_cpp_bindings::MarkerRole; use llama_cpp_bindings::ParsedChatMessage; use llama_cpp_bindings::SampledTokenSection; @@ -32,7 +33,8 @@ fn parse_partial_reasoning_response( } else { format!("{}{generated}", markers.open) }; - let parse_outcome = model.parse_chat_message("[]", &response, true)?; + let parse_outcome = + model.parse_chat_message(&ChatTools::from_json("[]".to_owned())?, &response, true)?; let ChatMessageParseOutcome::Recognized(parsed) = parse_outcome else { bail!("model chat template must recognize a partial reasoning response"); }; @@ -319,10 +321,11 @@ fn deepseek_r1_8b_duck_types_gemma_paired_quote(fixture: &LlamaFixture<'_>) -> R const GEMMA_PAIRED_QUOTE_PAYLOAD: &str = "<|tool_call>call:get_weather{location:<|\"|>Paris<|\"|>}"; - let outcome = - fixture - .model - .parse_chat_message(TOOLS_JSON, GEMMA_PAIRED_QUOTE_PAYLOAD, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + GEMMA_PAIRED_QUOTE_PAYLOAD, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -382,9 +385,11 @@ fn deepseek_r1_8b_duck_types_glm_key_value_tags(fixture: &LlamaFixture<'_>) -> R Paris\ "; - let outcome = fixture - .model - .parse_chat_message(TOOLS_JSON, GLM_KEY_VALUE_PAYLOAD, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + GLM_KEY_VALUE_PAYLOAD, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -442,10 +447,11 @@ fn deepseek_r1_8b_duck_types_mistral_bracketed_json(fixture: &LlamaFixture<'_>) const MISTRAL_BRACKETED_JSON_PAYLOAD: &str = r#"[TOOL_CALLS]get_weather[ARGS]{"location":"Paris"}"#; - let outcome = - fixture - .model - .parse_chat_message(TOOLS_JSON, MISTRAL_BRACKETED_JSON_PAYLOAD, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + MISTRAL_BRACKETED_JSON_PAYLOAD, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -508,9 +514,11 @@ Paris\n\ \n\ "; - let outcome = fixture - .model - .parse_chat_message(TOOLS_JSON, QWEN_XML_PAYLOAD, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + QWEN_XML_PAYLOAD, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -569,9 +577,11 @@ fn deepseek_r1_8b_recognizes_empty_tool_calls_when_input_is_plain_content_with_t const PLAIN_CONTENT: &str = "Sorry, I cannot help with that."; - let outcome = fixture - .model - .parse_chat_message(TOOLS_JSON, PLAIN_CONTENT, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + PLAIN_CONTENT, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -601,9 +611,11 @@ fn deepseek_r1_8b_recognizes_empty_tool_calls_when_tools_not_requested( ) -> Result<()> { const PLAIN_CONTENT: &str = "Hello there."; - let outcome = fixture - .model - .parse_chat_message("[]", PLAIN_CONTENT, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json("[]".to_owned())?, + PLAIN_CONTENT, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("plain content with empty tools array must produce Recognized; got Unrecognized"); @@ -856,10 +868,11 @@ fn gemma4_parses_tool_call_payload(fixture: &LlamaFixture<'_>) -> Result<()> { const GEMMA4_PAIRED_QUOTE_PAYLOAD: &str = "<|tool_call>call:get_weather{location:<|\"|>Paris<|\"|>}"; - let outcome = - fixture - .model - .parse_chat_message(TOOLS_JSON, GEMMA4_PAIRED_QUOTE_PAYLOAD, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + GEMMA4_PAIRED_QUOTE_PAYLOAD, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for Gemma 4 PairedQuote on a Gemma-4 model; got Unrecognized"); @@ -1125,9 +1138,11 @@ fn glm47_parses_tool_call_payload(fixture: &LlamaFixture<'_>) -> Result<()> { Paris\ "; - let outcome = fixture - .model - .parse_chat_message(TOOLS_JSON, GLM47_KEY_VALUE_PAYLOAD, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + GLM47_KEY_VALUE_PAYLOAD, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -1368,10 +1383,11 @@ fn mistral3_parses_tool_call_payload(fixture: &LlamaFixture<'_>) -> Result<()> { const MISTRAL3_BRACKETED_JSON_PAYLOAD: &str = r#"[TOOL_CALLS]get_weather[ARGS]{"location":"Paris"}"#; - let outcome = - fixture - .model - .parse_chat_message(TOOLS_JSON, MISTRAL3_BRACKETED_JSON_PAYLOAD, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + MISTRAL3_BRACKETED_JSON_PAYLOAD, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -1752,7 +1768,7 @@ get off the keyboard\n\ "; let outcome = fixture.model.parse_chat_message( - NEGOTIATE_WITH_CAT_TOOLS_JSON, + &ChatTools::from_json(NEGOTIATE_WITH_CAT_TOOLS_JSON.to_owned())?, NEGOTIATE_WITH_CAT_INPUT, false, )?; @@ -1813,9 +1829,11 @@ Paris\n\ \n\ "; - let outcome = fixture - .model - .parse_chat_message(TOOLS_JSON, QWEN_XML_PAYLOAD, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + QWEN_XML_PAYLOAD, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for Qwen XML on a Qwen-3.5 model; got Unrecognized"); @@ -1864,9 +1882,11 @@ fn qwen35_parses_partial_tool_call_returns_pending_state(fixture: &LlamaFixture< const PARTIAL_QWEN_XML_PAYLOAD: &str = "\n\n\n\ "; - let outcome = fixture - .model - .parse_chat_message(TOOLS_JSON, TWO_QWEN_XML_PAYLOADS, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + TWO_QWEN_XML_PAYLOADS, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -1965,9 +1987,11 @@ fn qwen35_recognizes_empty_tool_calls_when_input_is_plain_content_with_tools_req const PLAIN_CONTENT: &str = "Sorry, I cannot help with that."; - let outcome = fixture - .model - .parse_chat_message(TOOLS_JSON, PLAIN_CONTENT, false)?; + let outcome = fixture.model.parse_chat_message( + &ChatTools::from_json(TOOLS_JSON.to_owned())?, + PLAIN_CONTENT, + false, + )?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( diff --git a/llama-cpp-bindings/Cargo.toml b/llama-cpp-bindings/Cargo.toml index 3eaf97f0a..0d49b00c0 100644 --- a/llama-cpp-bindings/Cargo.toml +++ b/llama-cpp-bindings/Cargo.toml @@ -18,6 +18,7 @@ llama-cpp-log-decoder = { workspace = true } llguidance = { workspace = true } log = { workspace = true } nom = { workspace = true } +once_cell = { workspace = true } serde_json = { workspace = true } thiserror = { workspace = true } toktrie = { workspace = true } diff --git a/llama-cpp-bindings/src/chat_message_parse_outcome.rs b/llama-cpp-bindings/src/chat_message_parse_outcome.rs index 6a6b77c56..2cd064d1a 100644 --- a/llama-cpp-bindings/src/chat_message_parse_outcome.rs +++ b/llama-cpp-bindings/src/chat_message_parse_outcome.rs @@ -24,7 +24,6 @@ mod tests { Vec::new(), )), ChatMessageParseOutcome::Unrecognized(RawChatMessage { - tools_json: "[]".to_owned(), text: "raw input".to_owned(), is_partial: false, ffi_error_message: "parser bailed".to_owned(), @@ -42,7 +41,6 @@ mod tests { saw_recognized = true; } ChatMessageParseOutcome::Unrecognized(raw) => { - assert_eq!(raw.tools_json, "[]"); assert_eq!(raw.text, "raw input"); assert!(!raw.is_partial); assert_eq!(raw.ffi_error_message, "parser bailed"); diff --git a/llama-cpp-bindings/src/chat_tools.rs b/llama-cpp-bindings/src/chat_tools.rs new file mode 100644 index 000000000..7e5f9e5d4 --- /dev/null +++ b/llama-cpp-bindings/src/chat_tools.rs @@ -0,0 +1,77 @@ +use std::ffi::CStr; +use std::ffi::CString; + +use crate::error::ChatToolsError; + +#[derive(Debug)] +pub struct ChatTools { + json: CString, +} + +impl ChatTools { + /// # Errors + /// Returns [`ChatToolsError`] when `json` contains a NUL byte, is not valid JSON, or is not + /// a JSON array. + pub fn from_json(json: String) -> Result { + let json = CString::new(json).map_err(ChatToolsError::ContainsNulByte)?; + let json_value: serde_json::Value = + serde_json::from_slice(json.as_bytes()).map_err(ChatToolsError::InvalidJson)?; + + if !json_value.is_array() { + return Err(ChatToolsError::NotArray); + } + + Ok(Self { json }) + } + + #[must_use] + pub fn json_cstr(&self) -> &CStr { + &self.json + } +} + +#[cfg(test)] +mod tests { + use serde_json::Value; + + use super::ChatTools; + + #[test] + fn keeps_a_valid_tools_array_for_parsing() { + let tools = ChatTools::from_json("[]".to_owned()).unwrap(); + + assert_eq!(tools.json_cstr().to_bytes(), b"[]"); + } + + #[test] + fn rejects_malformed_json() { + let json_error = serde_json::from_str::("not_a_json[}").unwrap_err(); + + assert_eq!( + ChatTools::from_json("not_a_json[}".to_owned()) + .unwrap_err() + .to_string(), + format!("chat tools are not valid JSON: {json_error}") + ); + } + + #[test] + fn rejects_json_that_is_not_an_array() { + assert_eq!( + ChatTools::from_json("{\"foo\": 1}".to_owned()) + .unwrap_err() + .to_string(), + "chat tools must be a JSON array" + ); + } + + #[test] + fn reports_a_nul_byte_followed_by_extra_text_as_a_nul_byte() { + assert_eq!( + ChatTools::from_json("[]\0extra".to_owned()) + .unwrap_err() + .to_string(), + "chat tools contain an interior NUL byte at position 2" + ); + } +} diff --git a/llama-cpp-bindings/src/context.rs b/llama-cpp-bindings/src/context.rs index c8c266dd7..6701647f9 100644 --- a/llama-cpp-bindings/src/context.rs +++ b/llama-cpp-bindings/src/context.rs @@ -325,6 +325,11 @@ impl<'model> LlamaContext<'model> { unsafe { llama_cpp_bindings_sys::llama_n_ctx(self.context.as_ptr()) } } + #[must_use] + pub fn n_ctx_seq(&self) -> u32 { + unsafe { llama_cpp_bindings_sys::llama_n_ctx_seq(self.context.as_ptr()) } + } + #[expect(unsafe_code, reason = "required for FFI abort callback registration")] pub fn set_abort_flag(&mut self, flag: Arc) { let raw_ptr = Arc::as_ptr(&flag) as *mut c_void; diff --git a/llama-cpp-bindings/src/error.rs b/llama-cpp-bindings/src/error.rs index 04deed8f4..43a907f9c 100644 --- a/llama-cpp-bindings/src/error.rs +++ b/llama-cpp-bindings/src/error.rs @@ -1,6 +1,7 @@ pub mod apply_chat_template_error; pub mod bracketed_args_failure; pub mod chat_template_error; +pub mod chat_tools_error; pub mod clear_kv_cache_seq_error; pub mod copy_kv_cache_seq_error; pub mod decode_error; @@ -44,6 +45,7 @@ pub use llama_cpp_ffi_status::FfiStatusError; pub use apply_chat_template_error::ApplyChatTemplateError; pub use bracketed_args_failure::BracketedArgsFailure; pub use chat_template_error::ChatTemplateError; +pub use chat_tools_error::ChatToolsError; pub use clear_kv_cache_seq_error::ClearKvCacheSeqError; pub use copy_kv_cache_seq_error::CopyKvCacheSeqError; pub use decode_error::DecodeError; diff --git a/llama-cpp-bindings/src/error/chat_tools_error.rs b/llama-cpp-bindings/src/error/chat_tools_error.rs new file mode 100644 index 000000000..e9e5c3842 --- /dev/null +++ b/llama-cpp-bindings/src/error/chat_tools_error.rs @@ -0,0 +1,11 @@ +use std::ffi::NulError; + +#[derive(Debug, thiserror::Error)] +pub enum ChatToolsError { + #[error("chat tools are not valid JSON: {0}")] + InvalidJson(#[source] serde_json::Error), + #[error("chat tools must be a JSON array")] + NotArray, + #[error("chat tools contain an interior NUL byte at position {}", .0.nul_position())] + ContainsNulByte(#[source] NulError), +} diff --git a/llama-cpp-bindings/src/error/parse_chat_message_error.rs b/llama-cpp-bindings/src/error/parse_chat_message_error.rs index 05ff0e87c..488b50c56 100644 --- a/llama-cpp-bindings/src/error/parse_chat_message_error.rs +++ b/llama-cpp-bindings/src/error/parse_chat_message_error.rs @@ -2,7 +2,7 @@ use std::string::FromUtf8Error; use crate::error::marker_detection_error::MarkerDetectionError; -#[derive(Debug, thiserror::Error)] +#[derive(Debug, Eq, PartialEq, thiserror::Error)] pub enum ParseChatMessageError { #[error(transparent)] FfiStatus(#[from] crate::FfiStatusError), @@ -18,6 +18,8 @@ pub enum ParseChatMessageError { LlamaCppOutOfMemory, #[error("the chat parser could not be constructed: {message}")] ParserCreationFailed { message: String }, + #[error("the chat parser could not be built for the given tools: {message}")] + ToolsParserBuildFailed { message: String }, #[error("the chat parser did not recognize the message: {message}")] MessageUnrecognized { message: String }, #[error("the chat parser destructor threw: {message}")] @@ -30,12 +32,6 @@ pub enum ParseChatMessageError { ToolCallArgumentsIndexOutOfBounds { index: usize }, #[error("ffi returned non-utf8 string: {0}")] StringUtf8Error(#[from] FromUtf8Error), - #[error("tools_json is not valid JSON: {0}")] - ToolsJsonInvalid(#[source] serde_json::Error), - #[error("tools_json must be a JSON array")] - ToolsJsonNotArray, - #[error("tools_json contains an interior NUL byte")] - ToolsJsonContainsNulByte(#[source] std::ffi::NulError), #[error("the message to parse contains an interior NUL byte")] InputContainsNulByte(#[source] std::ffi::NulError), #[error("reasoning-marker detection failed: {0}")] diff --git a/llama-cpp-bindings/src/lib.rs b/llama-cpp-bindings/src/lib.rs index eea27d8fe..69b7667ad 100644 --- a/llama-cpp-bindings/src/lib.rs +++ b/llama-cpp-bindings/src/lib.rs @@ -6,6 +6,7 @@ pub mod batch_add_error; pub mod chat_message_parse_outcome; pub mod chat_template_tool_calls; +pub mod chat_tools; pub mod classified_sample; pub mod context; pub mod error; @@ -64,16 +65,18 @@ pub mod tool_call_format; pub mod tool_call_marker_pair; pub use error::{ - ApplyChatTemplateError, ChatTemplateError, ClearKvCacheSeqError, CopyKvCacheSeqError, - DecodeError, EmbeddingsError, EncodeError, EvalMultimodalChunksError, FfiContractError, - FfiStatusError, GrammarError, JsonSchemaToGrammarError, KvCacheSeqAddError, KvCacheSeqDivError, - KvCacheSeqPosMaxError, LlamaContextLoadError, LlamaCppError, LlamaLoraAdapterInitError, - LlamaLoraAdaptersError, LlamaModelLoadError, LogitsError, MarkerDetectionError, MetaValError, - ModelParamsError, NewLlamaChatMessageError, ParseChatMessageError, Result, SampleError, - SamplerAcceptError, SamplingError, StringToTokenError, TokenSamplingError, TokenToStringError, + ApplyChatTemplateError, ChatTemplateError, ChatToolsError, ClearKvCacheSeqError, + CopyKvCacheSeqError, DecodeError, EmbeddingsError, EncodeError, EvalMultimodalChunksError, + FfiContractError, FfiStatusError, GrammarError, JsonSchemaToGrammarError, KvCacheSeqAddError, + KvCacheSeqDivError, KvCacheSeqPosMaxError, LlamaContextLoadError, LlamaCppError, + LlamaLoraAdapterInitError, LlamaLoraAdaptersError, LlamaModelLoadError, LogitsError, + MarkerDetectionError, MetaValError, ModelParamsError, NewLlamaChatMessageError, + ParseChatMessageError, Result, SampleError, SamplerAcceptError, SamplingError, + StringToTokenError, TokenSamplingError, TokenToStringError, }; pub use chat_message_parse_outcome::ChatMessageParseOutcome; +pub use chat_tools::ChatTools; pub use classified_sample::ClassifiedSample; pub use eval_multimodal_chunks_params::EvalMultimodalChunksParams; pub use llama_backend_device::LlamaBackendDevice; diff --git a/llama-cpp-bindings/src/model.rs b/llama-cpp-bindings/src/model.rs index 329c76e6c..70e588aaf 100644 --- a/llama-cpp-bindings/src/model.rs +++ b/llama-cpp-bindings/src/model.rs @@ -22,8 +22,8 @@ use std::path::Path; use std::ptr; use std::ptr::NonNull; use std::sync::Arc; -use std::sync::OnceLock; +use once_cell::sync::OnceCell; use toktrie::ApproximateTokEnv; use toktrie::TokRxInfo; use toktrie::TokTrie; @@ -36,6 +36,7 @@ use llama_cpp_bindings_types::ToolCallMarkers; use crate::chat_message_parse_outcome::ChatMessageParseOutcome; use crate::chat_template_tool_calls; +use crate::chat_tools::ChatTools; use crate::llama_backend::LlamaBackend; use crate::llama_token_attrs::LlamaTokenAttrs; use crate::llama_token_attrs_from_int_error::LlamaTokenAttrsFromIntError; @@ -88,8 +89,10 @@ fn cstring_with_validated_len(text: &str) -> Result, - tok_env: OnceLock>, - chat_parser: OnceLock, + tok_env: OnceCell>, + chat_parser: OnceCell, + reasoning_markers: OnceCell>, + streaming_markers: OnceCell>, } #[derive(Debug)] @@ -229,8 +232,10 @@ unsafe fn load_model_from_file_status_to_result( })?; Ok(LlamaModel { model, - tok_env: OnceLock::new(), - chat_parser: OnceLock::new(), + tok_env: OnceCell::new(), + chat_parser: OnceCell::new(), + reasoning_markers: OnceCell::new(), + streaming_markers: OnceCell::new(), }) } llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_LLAMA_CPP_RETURNED_NULL => { @@ -285,11 +290,36 @@ unsafe fn load_model_from_file_status_to_result( } } +/// # Safety +/// +/// `out_error` must reference the pointer populated by a `llama_rs_parse_chat_message` call +/// that reported a thrown C++ exception. The error is read, freed, and the referenced pointer +/// is nulled so the later free in the caller does not double-free. +unsafe fn thrown_parse_exception_error( + out_error: *mut *mut c_char, + error_from_message: fn(String) -> ParseChatMessageError, +) -> ParseChatMessageError { + match unsafe { + read_and_free_cpp_string( + *out_error, + "llama_rs_parse_chat_message", + "reported a thrown C++ exception without an error message", + ) + } { + Ok(message) => { + unsafe { *out_error = ptr::null_mut() }; + + error_from_message(message) + } + Err(missing_message) => missing_message.into(), + } +} + /// # Safety /// /// `handle` must be the parsed-chat handle (or null) and `out_error` must reference the /// pointer populated by the preceding `llama_rs_parse_chat_message` call. In the CXX-exception -/// arm the error is read, freed, and the referenced pointer is nulled so the later free in the +/// arms the error is read, freed, and the referenced pointer is nulled so the later free in the /// caller does not double-free. unsafe fn parse_chat_message_status_to_result( status: llama_cpp_bindings_sys::llama_rs_parse_chat_message_status, @@ -315,15 +345,18 @@ unsafe fn parse_chat_message_status_to_result( Err(ParseChatMessageError::LlamaCppOutOfMemory) } llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = unsafe { - read_and_free_cpp_string( - *out_error, - "llama_rs_parse_chat_message", - "reported a thrown C++ exception without an error message", - ) - }?; - unsafe { *out_error = ptr::null_mut() }; - Err(ParseChatMessageError::MessageUnrecognized { message }) + Err(unsafe { + thrown_parse_exception_error(out_error, |message| { + ParseChatMessageError::MessageUnrecognized { message } + }) + }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_TOOLS_PARSER_BUILD_THREW_CXX_EXCEPTION => { + Err(unsafe { + thrown_parse_exception_error(out_error, |message| { + ParseChatMessageError::ToolsParserBuildFailed { message } + }) + }) } llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_PARSER_ARG => { Err(crate::FfiContractError { @@ -435,7 +468,6 @@ unsafe fn chat_parser_create_status_to_result( fn outcome_from_via_ffi_result( via_ffi_result: Result, - tools_json: &str, input: &str, is_partial: bool, ) -> Result { @@ -446,7 +478,6 @@ fn outcome_from_via_ffi_result( } Err(ParseChatMessageError::MessageUnrecognized { message }) => { Ok(ChatMessageParseOutcome::Unrecognized(RawChatMessage { - tools_json: tools_json.to_owned(), text: input.to_owned(), is_partial, ffi_error_message: message, @@ -997,15 +1028,20 @@ impl LlamaModel { pub fn sampled_token_classifier( &self, ) -> Result, MarkerDetectionError> { - let markers = self.streaming_markers()?; - - Ok(SampledTokenClassifier::new(self, markers)) + self.streaming_markers() + .map(|markers| SampledTokenClassifier::new(self, markers)) } /// # Errors /// Returns [`MarkerDetectionError`] when any underlying FFI call fails. - pub fn streaming_markers(&self) -> Result { - let reasoning_markers = invoke_detect_reasoning_markers(self.model.as_ptr())?; + pub fn streaming_markers(&self) -> Result, MarkerDetectionError> { + self.streaming_markers + .get_or_try_init(|| self.detect_streaming_markers().map(Arc::new)) + .map(Arc::clone) + } + + fn detect_streaming_markers(&self) -> Result { + let reasoning_markers = self.reasoning_markers()?; let tool_call_haystack = invoke_compute_tool_call_haystack(self.model.as_ptr())?; @@ -1024,7 +1060,7 @@ impl LlamaModel { self.resolve_tool_call_marker_strings(autoparser_open, autoparser_close)?; let mut candidates = Vec::new(); - if let Some(markers) = &reasoning_markers { + if let Some(markers) = reasoning_markers { for marker in &markers.closes { if let Some(tokens) = self.tokenize_marker(Some(marker))? { candidates.push(MarkerRoleCandidate { @@ -1035,9 +1071,7 @@ impl LlamaModel { } } - let reasoning_open = reasoning_markers - .as_ref() - .map(|markers| markers.open.as_str()); + let reasoning_open = reasoning_markers.map(|markers| markers.open.as_str()); if let Some(tokens) = self.tokenize_marker(reasoning_open)? { candidates.push(MarkerRoleCandidate { tokens, @@ -1093,8 +1127,10 @@ impl LlamaModel { /// # Errors /// Returns [`MarkerDetectionError`] when the underlying FFI call fails. - pub fn reasoning_markers(&self) -> Result, MarkerDetectionError> { - invoke_detect_reasoning_markers(self.model.as_ptr()) + pub fn reasoning_markers(&self) -> Result, MarkerDetectionError> { + self.reasoning_markers + .get_or_try_init(|| invoke_detect_reasoning_markers(self.model.as_ptr())) + .map(Option::as_ref) } /// # Errors @@ -1136,23 +1172,16 @@ impl LlamaModel { /// # Errors /// - /// Returns [`ParseChatMessageError`] when `tools_json` is not valid JSON, - /// the FFI returns a non-OK status other than `ParseException`, or - /// accessor strings are not valid UTF-8. + /// Returns [`ParseChatMessageError`] when reasoning-marker detection fails, + /// `input` contains a NUL byte, the chat parser cannot be created or built + /// for `tools`, the FFI returns a non-OK status other than a message parse + /// exception, or accessor strings are not valid UTF-8. pub fn parse_chat_message( &self, - tools_json: &str, + tools: &ChatTools, input: &str, is_partial: bool, ) -> Result { - let tools_cstring = - CString::new(tools_json).map_err(ParseChatMessageError::ToolsJsonContainsNulByte)?; - let tools_value: serde_json::Value = - serde_json::from_str(tools_json).map_err(ParseChatMessageError::ToolsJsonInvalid)?; - if !tools_value.is_array() { - return Err(ParseChatMessageError::ToolsJsonNotArray); - } - let reasoning_markers = self.reasoning_markers()?; for candidate in chat_template_tool_calls::known_marker_candidates() { @@ -1161,7 +1190,7 @@ impl LlamaModel { ToolCallFormatOutcome::Parsed(calls) => { let split = split_reasoning_prefix( input, - reasoning_markers.as_ref(), + reasoning_markers, Some(&candidate.open), is_partial, ); @@ -1175,18 +1204,13 @@ impl LlamaModel { } let via_ffi_result = self - .parse_chat_message_via_ffi(&tools_cstring, input, is_partial) + .parse_chat_message_via_ffi(tools.json_cstr(), input, is_partial) .map(|mut parsed| { - restore_partial_reasoning( - &mut parsed, - input, - reasoning_markers.as_ref(), - is_partial, - ); + restore_partial_reasoning(&mut parsed, input, reasoning_markers, is_partial); parsed }); - outcome_from_via_ffi_result(via_ffi_result, tools_json, input, is_partial) + outcome_from_via_ffi_result(via_ffi_result, input, is_partial) } fn parse_chat_message_via_ffi( @@ -1238,11 +1262,8 @@ impl LlamaModel { } fn chat_parser(&self) -> Result<&ChatParserHandle, ParseChatMessageError> { - if let Some(parser) = self.chat_parser.get() { - return Ok(parser); - } - let parser = self.create_chat_parser()?; - Ok(self.chat_parser.get_or_init(|| parser)) + self.chat_parser + .get_or_try_init(|| self.create_chat_parser()) } fn create_chat_parser(&self) -> Result { @@ -1285,11 +1306,9 @@ impl LlamaModel { /// as empty (not an error); a piece that overflows the probe buffer is /// re-read at the exact size rather than dropped. pub fn approximate_tok_env(&self) -> Result, TokenToStringError> { - if let Some(env) = self.tok_env.get() { - return Ok(Arc::clone(env)); - } - let env = build_approximate_tok_env(self)?; - Ok(Arc::clone(self.tok_env.get_or_init(|| env))) + self.tok_env + .get_or_try_init(|| build_approximate_tok_env(self)) + .map(Arc::clone) } } @@ -3094,6 +3113,51 @@ mod ffi_status_mapping_tests { ); } + #[test] + fn parse_chat_message_tools_parser_build_exception_is_a_build_failure_and_nulls_error() { + let mut out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"key 'name' not found".as_ptr()) + }; + let result = unsafe { + parse_chat_message_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_TOOLS_PARSER_BUILD_THREW_CXX_EXCEPTION, + ptr::null_mut(), + &raw mut out_error, + ) + }; + + assert_eq!( + result.unwrap_err(), + ParseChatMessageError::ToolsParserBuildFailed { + message: "key 'name' not found".to_owned(), + } + ); + assert!( + out_error.is_null(), + "the reclaimed pointer must be nulled so the caller does not free it twice" + ); + } + + #[test] + fn parse_chat_message_cxx_exception_without_an_error_message_is_a_contract_error() { + let mut out_error: *mut c_char = ptr::null_mut(); + let result = unsafe { + parse_chat_message_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_TOOLS_PARSER_BUILD_THREW_CXX_EXCEPTION, + ptr::null_mut(), + &raw mut out_error, + ) + }; + + assert_eq!( + result.unwrap_err(), + ParseChatMessageError::FfiContract(crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "reported a thrown C++ exception without an error message", + }) + ); + } + #[test] fn parse_chat_message_unknown_status_is_preserved() { let mut out_error: *mut c_char = ptr::null_mut(); @@ -4338,7 +4402,7 @@ mod ffi_status_mapping_tests { )], ); - let outcome = outcome_from_via_ffi_result(Ok(parsed), "[]", "answer", false); + let outcome = outcome_from_via_ffi_result(Ok(parsed), "answer", false); assert_eq!( outcome.unwrap(), @@ -4360,7 +4424,6 @@ mod ffi_status_mapping_tests { Err(ParseChatMessageError::MessageUnrecognized { message: "boom".to_owned(), }), - "[]", "garbled", true, ); @@ -4368,7 +4431,6 @@ mod ffi_status_mapping_tests { assert_eq!( outcome.unwrap(), ChatMessageParseOutcome::Unrecognized(RawChatMessage { - tools_json: "[]".to_owned(), text: "garbled".to_owned(), is_partial: true, ffi_error_message: "boom".to_owned(), @@ -4382,7 +4444,6 @@ mod ffi_status_mapping_tests { Err(ParseChatMessageError::ParserCreationFailed { message: "the parser could not be built".to_owned(), }), - "[]", "garbled", true, ); @@ -4397,8 +4458,7 @@ mod ffi_status_mapping_tests { #[test] fn outcome_from_via_ffi_result_other_error_propagates() { - let outcome = - outcome_from_via_ffi_result(Err(ParseChatMessageError::NoVocab), "[]", "x", false); + let outcome = outcome_from_via_ffi_result(Err(ParseChatMessageError::NoVocab), "x", false); assert_eq!( discriminant(&outcome.unwrap_err()), diff --git a/llama-cpp-bindings/src/raw_chat_message.rs b/llama-cpp-bindings/src/raw_chat_message.rs index 0108d7f54..15b2f2662 100644 --- a/llama-cpp-bindings/src/raw_chat_message.rs +++ b/llama-cpp-bindings/src/raw_chat_message.rs @@ -1,6 +1,5 @@ #[derive(Debug, Eq, PartialEq)] pub struct RawChatMessage { - pub tools_json: String, pub text: String, pub is_partial: bool, pub ffi_error_message: String, @@ -11,15 +10,13 @@ mod tests { use super::RawChatMessage; #[test] - fn carries_tools_json_text_partial_flag_and_ffi_error_message() { + fn carries_text_partial_flag_and_ffi_error_message() { let raw = RawChatMessage { - tools_json: "[]".to_owned(), text: "hello".to_owned(), is_partial: true, ffi_error_message: "parser bailed".to_owned(), }; - assert_eq!(raw.tools_json, "[]"); assert_eq!(raw.text, "hello"); assert!(raw.is_partial); assert_eq!(raw.ffi_error_message, "parser bailed"); diff --git a/llama-cpp-bindings/src/sampled_token_classifier.rs b/llama-cpp-bindings/src/sampled_token_classifier.rs index 8f73a880a..2e511ca12 100644 --- a/llama-cpp-bindings/src/sampled_token_classifier.rs +++ b/llama-cpp-bindings/src/sampled_token_classifier.rs @@ -1,4 +1,5 @@ use std::collections::VecDeque; +use std::sync::Arc; use llama_cpp_bindings_sys::llama_pos; use llama_cpp_bindings_sys::llama_seq_id; @@ -57,7 +58,7 @@ enum ProbeMode { pub struct SampledTokenClassifier<'model> { model: &'model LlamaModel, - markers: StreamingMarkers, + markers: Arc, decoder: encoding_rs::Decoder, pending: VecDeque, section: SampledTokenSection, @@ -68,7 +69,7 @@ pub struct SampledTokenClassifier<'model> { impl<'model> SampledTokenClassifier<'model> { #[must_use] - pub fn new(model: &'model LlamaModel, markers: StreamingMarkers) -> Self { + pub fn new(model: &'model LlamaModel, markers: Arc) -> Self { Self { model, markers, @@ -505,13 +506,15 @@ impl<'model> SampledTokenClassifier<'model> { } #[must_use] - pub const fn markers(&self) -> &StreamingMarkers { + pub fn markers(&self) -> &StreamingMarkers { &self.markers } } #[cfg(test)] mod tests { + use std::sync::Arc; + use super::JsonProbeState; use super::PendingMarkerStatus; use super::PendingToken; @@ -554,7 +557,7 @@ mod tests { fn synthetic_classifier(markers: StreamingMarkers) -> SampledTokenClassifier<'static> { SampledTokenClassifier { model: unsafe { &*std::ptr::NonNull::::dangling().as_ptr() }, - markers, + markers: Arc::new(markers), decoder: encoding_rs::UTF_8.new_decoder(), pending: std::collections::VecDeque::new(), section: SampledTokenSection::Pending,