diff --git a/Cargo.lock b/Cargo.lock index 26bd8e827..6d0e2834d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1163,7 +1163,7 @@ checksum = "11d3d7f243d5c5a8b9bb5d6dd2b1602c0cb0b9db1621bafc7ed66e35ff9fe092" [[package]] name = "llama-cpp-bindings" -version = "0.15.1" +version = "0.16.0" dependencies = [ "anyhow", "encoding_rs", @@ -1179,6 +1179,7 @@ dependencies = [ "log", "nom 8.0.0", "once_cell", + "serde", "serde_json", "serial_test", "thiserror", @@ -1187,7 +1188,7 @@ dependencies = [ [[package]] name = "llama-cpp-bindings-build" -version = "0.15.1" +version = "0.16.0" dependencies = [ "bindgen", "cc", @@ -1199,14 +1200,14 @@ dependencies = [ [[package]] name = "llama-cpp-bindings-sys" -version = "0.15.1" +version = "0.16.0" dependencies = [ "llama-cpp-bindings-build", ] [[package]] name = "llama-cpp-bindings-tests" -version = "0.15.1" +version = "0.16.0" dependencies = [ "anyhow", "encoding_rs", @@ -1218,7 +1219,7 @@ dependencies = [ [[package]] name = "llama-cpp-bindings-types" -version = "0.15.1" +version = "0.16.0" dependencies = [ "serde", "serde_json", @@ -1227,14 +1228,14 @@ dependencies = [ [[package]] name = "llama-cpp-error-recorder" -version = "0.15.1" +version = "0.16.0" dependencies = [ "log", ] [[package]] name = "llama-cpp-ffi-status" -version = "0.15.1" +version = "0.16.0" dependencies = [ "llama-cpp-bindings-sys", "llama-cpp-wrapper-error-fixture", @@ -1243,7 +1244,7 @@ dependencies = [ [[package]] name = "llama-cpp-gbnf" -version = "0.15.1" +version = "0.16.0" dependencies = [ "llama-cpp-bindings-sys", "llama-cpp-ffi-status", @@ -1252,11 +1253,11 @@ dependencies = [ [[package]] name = "llama-cpp-log-decoder" -version = "0.15.1" +version = "0.16.0" [[package]] name = "llama-cpp-test-harness" -version = "0.15.1" +version = "0.16.0" dependencies = [ "anyhow", "hf-hub", @@ -1269,7 +1270,7 @@ dependencies = [ [[package]] name = "llama-cpp-test-harness-macros" -version = "0.15.1" +version = "0.16.0" dependencies = [ "proc-macro2", "quote", @@ -1278,14 +1279,14 @@ dependencies = [ [[package]] name = "llama-cpp-wrapper-error-fixture" -version = "0.15.1" +version = "0.16.0" dependencies = [ "llama-cpp-bindings-sys", ] [[package]] name = "llama-cpp-wrapper-sources" -version = "0.15.1" +version = "0.16.0" dependencies = [ "serde", "serde_json", diff --git a/Cargo.toml b/Cargo.toml index 8145a5775..17ce3e5c6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -18,7 +18,7 @@ members = [ [workspace.package] edition = "2024" -version = "0.15.1" +version = "0.16.0" license = "Apache-2.0" repository = "https://github.com/intentee/llama-cpp-bindings" @@ -33,18 +33,18 @@ 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.15.1" } -llama-cpp-bindings-build = { path = "llama-cpp-bindings-build", version = "=0.15.1" } -llama-cpp-bindings-sys = { path = "llama-cpp-bindings-sys", version = "=0.15.1" } -llama-cpp-bindings-types = { path = "llama-cpp-bindings-types", version = "=0.15.1" } -llama-cpp-error-recorder = { path = "llama-cpp-error-recorder", version = "=0.15.1" } -llama-cpp-ffi-status = { path = "llama-cpp-ffi-status", version = "=0.15.1" } -llama-cpp-gbnf = { path = "llama-cpp-gbnf", version = "=0.15.1" } -llama-cpp-log-decoder = { path = "llama-cpp-log-decoder", version = "=0.15.1" } -llama-cpp-test-harness = { path = "llama-cpp-test-harness", version = "=0.15.1" } -llama-cpp-test-harness-macros = { path = "llama-cpp-test-harness-macros", version = "=0.15.1" } +llama-cpp-bindings = { path = "llama-cpp-bindings", version = "=0.16.0" } +llama-cpp-bindings-build = { path = "llama-cpp-bindings-build", version = "=0.16.0" } +llama-cpp-bindings-sys = { path = "llama-cpp-bindings-sys", version = "=0.16.0" } +llama-cpp-bindings-types = { path = "llama-cpp-bindings-types", version = "=0.16.0" } +llama-cpp-error-recorder = { path = "llama-cpp-error-recorder", version = "=0.16.0" } +llama-cpp-ffi-status = { path = "llama-cpp-ffi-status", version = "=0.16.0" } +llama-cpp-gbnf = { path = "llama-cpp-gbnf", version = "=0.16.0" } +llama-cpp-log-decoder = { path = "llama-cpp-log-decoder", version = "=0.16.0" } +llama-cpp-test-harness = { path = "llama-cpp-test-harness", version = "=0.16.0" } +llama-cpp-test-harness-macros = { path = "llama-cpp-test-harness-macros", version = "=0.16.0" } llama-cpp-wrapper-error-fixture = { path = "llama-cpp-wrapper-error-fixture" } -llama-cpp-wrapper-sources = { path = "llama-cpp-wrapper-sources", version = "=0.15.1" } +llama-cpp-wrapper-sources = { path = "llama-cpp-wrapper-sources", version = "=0.16.0" } llguidance = "=1.7.0" log = "=0.4.29" nom = "=8.0.0" diff --git a/llama-cpp-bindings-build/src/cmake_config.rs b/llama-cpp-bindings-build/src/cmake_config.rs index 37e834942..8058b3a9d 100644 --- a/llama-cpp-bindings-build/src/cmake_config.rs +++ b/llama-cpp-bindings-build/src/cmake_config.rs @@ -246,6 +246,10 @@ fn configure_gpu_backends(config: &mut Config, target_os: TargetOs) -> Result<() if cfg!(feature = "cuda-no-vmm") { config.define("GGML_CUDA_NO_VMM", "ON"); } + + if let Some(cuda_architectures) = optional_env("CUDAARCHS")? { + config.define("CMAKE_CUDA_ARCHITECTURES", cuda_architectures); + } } if cfg!(feature = "rocm") { diff --git a/llama-cpp-bindings-build/src/rebuild_tracking.rs b/llama-cpp-bindings-build/src/rebuild_tracking.rs index 8982e7fea..cf5523592 100644 --- a/llama-cpp-bindings-build/src/rebuild_tracking.rs +++ b/llama-cpp-bindings-build/src/rebuild_tracking.rs @@ -5,6 +5,7 @@ use llama_cpp_wrapper_sources::wrapper_sources::WRAPPER_SOURCES; pub fn register_rebuild_triggers(llama_src: &Path) { println!("cargo:rerun-if-changed=build.rs"); + println!("cargo:rerun-if-env-changed=CUDAARCHS"); for path in WRAPPER_HEADERS.iter().chain(WRAPPER_SOURCES) { println!("cargo:rerun-if-changed={path}"); diff --git a/llama-cpp-bindings-sys/wrapper.h b/llama-cpp-bindings-sys/wrapper.h index 9331a9d8c..2c735ad8e 100644 --- a/llama-cpp-bindings-sys/wrapper.h +++ b/llama-cpp-bindings-sys/wrapper.h @@ -3,6 +3,7 @@ #include "wrapper_chat_apply.h" #include "wrapper_chat_parse.h" #include "wrapper_common.h" +#include "wrapper_context.h" #include "wrapper_fit.h" #include "wrapper_gbnf.h" #include "wrapper_mtmd.h" diff --git a/llama-cpp-bindings-sys/wrapper_chat_parse.cpp b/llama-cpp-bindings-sys/wrapper_chat_parse.cpp index b52d1534b..7fa33251c 100644 --- a/llama-cpp-bindings-sys/wrapper_chat_parse.cpp +++ b/llama-cpp-bindings-sys/wrapper_chat_parse.cpp @@ -25,6 +25,10 @@ struct llama_rs_chat_parser { autoparser::autoparser parser; }; +struct llama_rs_chat_tools_parser { + common_peg_arena arena; +}; + namespace { void dup_or_set_alloc_flag(const std::string & source, char ** out_dup, bool * out_alloc_failed) { *out_dup = llama_rs_dup_string(source); @@ -49,18 +53,6 @@ auto report_current_parse_exception( 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, @@ -176,9 +168,94 @@ extern "C" auto llama_rs_chat_parser_free( } } -extern "C" auto llama_rs_parse_chat_message( +extern "C" auto llama_rs_chat_tools_parser_create( llama_rs_chat_parser_handle parser, const char * tools_json, + llama_rs_chat_tools_parser_handle * out_tools_parser, + char ** out_error) -> llama_rs_chat_tools_parser_create_status { + if (out_tools_parser != nullptr) { + *out_tools_parser = nullptr; + } + if (out_error != nullptr) { + *out_error = nullptr; + } + if (parser == nullptr) { + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_PARSER_ARG; + } + if (tools_json == nullptr) { + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_TOOLS_JSON_ARG; + } + if (out_tools_parser == nullptr) { + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_OUT_TOOLS_PARSER_ARG; + } + if (out_error == nullptr) { + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_OUT_ERROR_ARG; + } + + try { + autoparser::generation_params inputs; + + inputs.tools = common_json::parse(tools_json); + + if (!inputs.tools.is_array()) { + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_TOOLS_NOT_AN_ARRAY; + } + + auto tools_parser_handle = std::make_unique(); + tools_parser_handle->arena = parser->parser.build_parser(inputs, std::string()); + + *out_tools_parser = tools_parser_handle.release(); + + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_OK; + } catch (const std::bad_alloc &) { + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_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_CHAT_TOOLS_PARSER_CREATE_ERROR_STRING_ALLOCATION_FAILED; + } + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_THREW_CXX_EXCEPTION; + } catch (...) { + *out_error = llama_rs_dup_string(std::string("unknown c++ exception")); + if (*out_error == nullptr) { + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_ERROR_STRING_ALLOCATION_FAILED; + } + return LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_THREW_CXX_EXCEPTION; + } +} + +extern "C" auto llama_rs_chat_tools_parser_free( + llama_rs_chat_tools_parser_handle tools_parser, + char ** out_error) -> llama_rs_chat_tools_parser_free_status { + if (out_error != nullptr) { + *out_error = nullptr; + } + try { + const std::unique_ptr reclaimed(tools_parser); + return LLAMA_RS_CHAT_TOOLS_PARSER_FREE_OK; + } catch (const std::bad_alloc &) { + return LLAMA_RS_CHAT_TOOLS_PARSER_FREE_LLAMA_CPP_OUT_OF_MEMORY; + } catch (const std::exception & err) { + if (out_error != nullptr) { + *out_error = llama_rs_dup_string(err.what()); + if (*out_error == nullptr) { + return LLAMA_RS_CHAT_TOOLS_PARSER_FREE_ERROR_STRING_ALLOCATION_FAILED; + } + } + return LLAMA_RS_CHAT_TOOLS_PARSER_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION; + } catch (...) { + if (out_error != nullptr) { + *out_error = llama_rs_dup_string("unknown c++ exception"); + if (*out_error == nullptr) { + return LLAMA_RS_CHAT_TOOLS_PARSER_FREE_ERROR_STRING_ALLOCATION_FAILED; + } + } + return LLAMA_RS_CHAT_TOOLS_PARSER_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION; + } +} + +extern "C" auto llama_rs_parse_chat_message( + llama_rs_chat_tools_parser_handle tools_parser, const char * input, int is_partial, llama_rs_parsed_chat_handle * out_handle, @@ -189,8 +266,8 @@ extern "C" auto llama_rs_parse_chat_message( if (out_error != nullptr) { *out_error = nullptr; } - if (parser == nullptr) { - return LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_PARSER_ARG; + if (tools_parser == nullptr) { + return LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_TOOLS_PARSER_ARG; } if (input == nullptr) { return LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_INPUT_ARG; @@ -202,14 +279,7 @@ extern "C" auto llama_rs_parse_chat_message( return LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_ERROR_ARG; } - try { - common_peg_arena const chat_parser = build_tools_parser(*parser, tools_json); - - return parse_with_tools_parser(chat_parser, input, is_partial, out_handle, out_error); - } catch (...) { - return report_current_parse_exception( - out_error, LLAMA_RS_PARSE_CHAT_MESSAGE_TOOLS_PARSER_BUILD_THREW_CXX_EXCEPTION); - } + return parse_with_tools_parser(tools_parser->arena, input, is_partial, out_handle, out_error); } extern "C" auto llama_rs_parsed_chat_free( diff --git a/llama-cpp-bindings-sys/wrapper_chat_parse.h b/llama-cpp-bindings-sys/wrapper_chat_parse.h index afc62cf71..181d1efb7 100644 --- a/llama-cpp-bindings-sys/wrapper_chat_parse.h +++ b/llama-cpp-bindings-sys/wrapper_chat_parse.h @@ -15,6 +15,9 @@ typedef struct llama_rs_parsed_chat * llama_rs_parsed_chat_handle; struct llama_rs_chat_parser; typedef struct llama_rs_chat_parser * llama_rs_chat_parser_handle; +struct llama_rs_chat_tools_parser; +typedef struct llama_rs_chat_tools_parser * llama_rs_chat_tools_parser_handle; + typedef enum llama_rs_chat_parser_create_status { LLAMA_RS_CHAT_PARSER_CREATE_OK = 0, LLAMA_RS_CHAT_PARSER_CREATE_NULL_MODEL_ARG, @@ -43,21 +46,48 @@ llama_rs_chat_parser_free_status llama_rs_chat_parser_free( llama_rs_chat_parser_handle parser, char ** out_error); +typedef enum llama_rs_chat_tools_parser_create_status { + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_OK = 0, + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_PARSER_ARG, + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_TOOLS_JSON_ARG, + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_OUT_TOOLS_PARSER_ARG, + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_OUT_ERROR_ARG, + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_TOOLS_NOT_AN_ARRAY, + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_ERROR_STRING_ALLOCATION_FAILED, + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_OUT_OF_MEMORY, + LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_THREW_CXX_EXCEPTION, +} llama_rs_chat_tools_parser_create_status; + +llama_rs_chat_tools_parser_create_status llama_rs_chat_tools_parser_create( + llama_rs_chat_parser_handle parser, + const char * tools_json, + llama_rs_chat_tools_parser_handle * out_tools_parser, + char ** out_error); + +typedef enum llama_rs_chat_tools_parser_free_status { + LLAMA_RS_CHAT_TOOLS_PARSER_FREE_OK = 0, + LLAMA_RS_CHAT_TOOLS_PARSER_FREE_ERROR_STRING_ALLOCATION_FAILED, + LLAMA_RS_CHAT_TOOLS_PARSER_FREE_LLAMA_CPP_OUT_OF_MEMORY, + LLAMA_RS_CHAT_TOOLS_PARSER_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION, +} llama_rs_chat_tools_parser_free_status; + +llama_rs_chat_tools_parser_free_status llama_rs_chat_tools_parser_free( + llama_rs_chat_tools_parser_handle tools_parser, + char ** out_error); + typedef enum llama_rs_parse_chat_message_status { LLAMA_RS_PARSE_CHAT_MESSAGE_OK = 0, - LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_PARSER_ARG, + LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_TOOLS_PARSER_ARG, LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_INPUT_ARG, LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_HANDLE_ARG, LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_ERROR_ARG, 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( - llama_rs_chat_parser_handle parser, - const char * tools_json, + llama_rs_chat_tools_parser_handle tools_parser, const char * input, int is_partial, llama_rs_parsed_chat_handle * out_handle, diff --git a/llama-cpp-bindings-sys/wrapper_context.cpp b/llama-cpp-bindings-sys/wrapper_context.cpp new file mode 100644 index 000000000..5129711e7 --- /dev/null +++ b/llama-cpp-bindings-sys/wrapper_context.cpp @@ -0,0 +1,8 @@ +#include "wrapper_context.h" + +#include "llama.cpp/include/llama.h" +#include "llama.cpp/src/llama-context.h" + +extern "C" auto llama_rs_context_decodes_batches_in_one_micro_batch(const struct llama_context * ctx) -> bool { + return llama_get_memory(ctx) == nullptr || !ctx->get_cparams().causal_attn; +} diff --git a/llama-cpp-bindings-sys/wrapper_context.h b/llama-cpp-bindings-sys/wrapper_context.h new file mode 100644 index 000000000..943cbd916 --- /dev/null +++ b/llama-cpp-bindings-sys/wrapper_context.h @@ -0,0 +1,15 @@ +#pragma once + +#include + +#ifdef __cplusplus +extern "C" { +#endif + +struct llama_context; + +bool llama_rs_context_decodes_batches_in_one_micro_batch(const struct llama_context * ctx); + +#ifdef __cplusplus +} +#endif diff --git a/llama-cpp-bindings-sys/wrapper_gbnf.cpp b/llama-cpp-bindings-sys/wrapper_gbnf.cpp index e36a84eaf..94d52bac8 100644 --- a/llama-cpp-bindings-sys/wrapper_gbnf.cpp +++ b/llama-cpp-bindings-sys/wrapper_gbnf.cpp @@ -7,6 +7,7 @@ #include extern "C" auto llama_rs_validate_gbnf( + const struct llama_vocab * vocab, const char * grammar_str, const char * grammar_root, char ** out_error) -> llama_rs_gbnf_validation_status { @@ -23,7 +24,7 @@ extern "C" auto llama_rs_validate_gbnf( return LLAMA_RS_GBNF_VALIDATION_NULL_OUT_ERROR_ARG; } try { - llama_grammar_parser parser; + llama_grammar_parser parser(vocab); if (!parser.parse(grammar_str)) { return LLAMA_RS_GBNF_VALIDATION_SYNTAX_ERROR; @@ -38,7 +39,7 @@ extern "C" auto llama_rs_validate_gbnf( } llama_grammar * grammar = llama_grammar_init_impl( - nullptr, grammar_str, grammar_root, false, nullptr, 0, nullptr, 0); + vocab, grammar_str, grammar_root, false, nullptr, 0, nullptr, 0); if (grammar == nullptr) { return LLAMA_RS_GBNF_VALIDATION_LEFT_RECURSION; diff --git a/llama-cpp-bindings-sys/wrapper_gbnf.h b/llama-cpp-bindings-sys/wrapper_gbnf.h index 20ecefab8..337167524 100644 --- a/llama-cpp-bindings-sys/wrapper_gbnf.h +++ b/llama-cpp-bindings-sys/wrapper_gbnf.h @@ -4,6 +4,8 @@ extern "C" { #endif +struct llama_vocab; + typedef enum llama_rs_gbnf_validation_status { LLAMA_RS_GBNF_VALIDATION_OK = 0, LLAMA_RS_GBNF_VALIDATION_SYNTAX_ERROR, @@ -19,6 +21,7 @@ typedef enum llama_rs_gbnf_validation_status { } llama_rs_gbnf_validation_status; llama_rs_gbnf_validation_status llama_rs_validate_gbnf( + const struct llama_vocab * vocab, const char * grammar_str, const char * grammar_root, char ** out_error); diff --git a/llama-cpp-bindings-tests/src/classify_sample_loop.rs b/llama-cpp-bindings-tests/src/classify_sample_loop.rs index 94807197f..e6fedf598 100644 --- a/llama-cpp-bindings-tests/src/classify_sample_loop.rs +++ b/llama-cpp-bindings-tests/src/classify_sample_loop.rs @@ -1,14 +1,14 @@ use anyhow::Result; +use llama_cpp_bindings::GenerationProgress; +use llama_cpp_bindings::SampledTokenSection; use llama_cpp_bindings::context::LlamaContext; use llama_cpp_bindings::ingest_outcome::IngestOutcome; use llama_cpp_bindings::llama_batch::LlamaBatch; -use llama_cpp_bindings::model::LlamaModel; use llama_cpp_bindings::sampled_token::SampledToken; use llama_cpp_bindings::sampled_token_classifier::SampledTokenClassifier; use llama_cpp_bindings::sampling::LlamaSampler; pub struct ClassifySampleLoop<'borrow, 'model, 'tokens> { - pub model: &'model LlamaModel, pub classifier: &'borrow mut SampledTokenClassifier<'model>, pub sampler: &'borrow mut LlamaSampler, pub context: &'borrow mut LlamaContext<'model>, @@ -29,150 +29,137 @@ pub struct ClassifySampleLoopOutcome { pub eog_seen: bool, } +impl ClassifySampleLoopOutcome { + const fn record_end_of_generation(&mut self, section: SampledTokenSection) { + self.eog_seen = true; + + match section { + SampledTokenSection::Content => self.observed_content += 1, + SampledTokenSection::Reasoning => self.observed_reasoning += 1, + SampledTokenSection::ToolCall => self.observed_tool_call += 1, + SampledTokenSection::Pending => self.observed_undeterminable += 1, + } + } + + fn record_outcome(&mut self, ingest: &IngestOutcome) { + self.generated_raw.push_str(ingest.piece.raw()); + + match ingest.sampled_token { + SampledToken::Content(_) => { + self.observed_content += 1; + self.content_stream.push_str(ingest.piece.visible()); + } + SampledToken::Reasoning(_) => { + self.observed_reasoning += 1; + self.reasoning_stream.push_str(ingest.piece.visible()); + } + SampledToken::ToolCall(_) => self.observed_tool_call += 1, + SampledToken::Undeterminable(_) => self.observed_undeterminable += 1, + } + } +} + impl ClassifySampleLoop<'_, '_, '_> { /// # Errors /// Forwards [`SampledTokenClassifier::sample`] / [`LlamaContext::decode`] / - /// [`LlamaBatch::add`] errors verbatim. Stops on EOG, on + /// [`LlamaBatch::add`] errors verbatim. Stops on the end of generation, on /// `max_generated_tokens` exhaustion, or on the first error. pub fn run(self) -> Result { let mut outcome = ClassifySampleLoopOutcome::default(); + let mut ingest_outcomes = Vec::new(); let mut position = self.initial_position; let max_position = position + self.max_generated_tokens; while position < max_position { - let sampled = - self.classifier - .sample(self.sampler, self.context, self.batch.n_tokens() - 1)?; - - for ingest_outcome in &sampled.outcomes { - let is_eog = self.model.is_eog_token(&ingest_outcome.sampled_token); - if is_eog { - outcome.eog_seen = true; - } else { - outcome.generated_raw.push_str(&ingest_outcome.raw_piece); - } - record_outcome(ingest_outcome, &mut outcome, is_eog); - } + let sampled = self.classifier.sample( + self.sampler, + self.context, + self.batch.n_tokens() - 1, + &mut ingest_outcomes, + )?; + + if sampled.progress == GenerationProgress::Ended { + outcome.record_end_of_generation(self.classifier.current_section()); - let raw_as_sampled = SampledToken::Content(sampled.token); - if self.model.is_eog_token(&raw_as_sampled) { - outcome.eog_seen = true; break; } self.batch.clear(); - self.batch.add(&raw_as_sampled, position, &[0], true)?; + self.batch + .add(&SampledToken::Content(sampled.token), position, &[0], true)?; position += 1; self.context.decode(self.batch)?; } - for ingest_outcome in self.classifier.flush() { - let is_eog = self.model.is_eog_token(&ingest_outcome.sampled_token); - if is_eog { - outcome.eog_seen = true; - } else { - outcome.generated_raw.push_str(&ingest_outcome.raw_piece); - } - record_outcome(&ingest_outcome, &mut outcome, is_eog); + self.classifier.finish(&mut ingest_outcomes); + + for ingest_outcome in &ingest_outcomes { + outcome.record_outcome(ingest_outcome); } Ok(outcome) } } -fn record_outcome(ingest: &IngestOutcome, outcome: &mut ClassifySampleLoopOutcome, is_eog: bool) { - match ingest.sampled_token { - SampledToken::Content(_) => { - outcome.observed_content += 1; - if !is_eog { - outcome.content_stream.push_str(&ingest.visible_piece); - } - } - SampledToken::Reasoning(_) => { - outcome.observed_reasoning += 1; - if !is_eog { - outcome.reasoning_stream.push_str(&ingest.visible_piece); - } - } - SampledToken::ToolCall(_) => { - outcome.observed_tool_call += 1; - } - SampledToken::Undeterminable(_) => { - outcome.observed_undeterminable += 1; - } - } -} - #[cfg(test)] mod tests { + use llama_cpp_bindings::SampledTokenSection; + use llama_cpp_bindings::TokenPiece; use llama_cpp_bindings::ingest_outcome::IngestOutcome; use llama_cpp_bindings::sampled_token::SampledToken; use llama_cpp_bindings::token::LlamaToken; use super::ClassifySampleLoopOutcome; - use super::record_outcome; #[test] - fn record_outcome_tool_call_token() { - let ingest = IngestOutcome { - sampled_token: SampledToken::ToolCall(LlamaToken(42)), - visible_piece: String::new(), - raw_piece: String::new(), - }; + fn records_a_tool_call_token_without_streaming_it() { let mut outcome = ClassifySampleLoopOutcome::default(); - record_outcome(&ingest, &mut outcome, false); + outcome.record_outcome(&IngestOutcome { + sampled_token: SampledToken::ToolCall(LlamaToken(42)), + piece: TokenPiece::Visible("{".to_owned()), + }); assert_eq!(outcome.observed_tool_call, 1); - assert_eq!(outcome.observed_content, 0); - assert_eq!(outcome.observed_reasoning, 0); - assert_eq!(outcome.observed_undeterminable, 0); + assert_eq!(outcome.generated_raw, "{"); + assert!(outcome.content_stream.is_empty()); } #[test] - fn record_outcome_reasoning_token_streams_visible_piece() { - let ingest = IngestOutcome { - sampled_token: SampledToken::Reasoning(LlamaToken(7)), - visible_piece: "thinking".to_string(), - raw_piece: String::new(), - }; + fn streams_the_visible_part_of_a_reasoning_token() { let mut outcome = ClassifySampleLoopOutcome::default(); - record_outcome(&ingest, &mut outcome, false); + outcome.record_outcome(&IngestOutcome { + sampled_token: SampledToken::Reasoning(LlamaToken(7)), + piece: TokenPiece::Visible("thinking".to_owned()), + }); assert_eq!(outcome.observed_reasoning, 1); assert_eq!(outcome.reasoning_stream, "thinking"); } #[test] - fn record_outcome_reasoning_token_at_end_of_generation_is_not_streamed() { - let ingest = IngestOutcome { - sampled_token: SampledToken::Reasoning(LlamaToken(7)), - visible_piece: "thinking".to_string(), - raw_piece: String::new(), - }; + fn counts_an_undeterminable_token_without_streaming_it() { let mut outcome = ClassifySampleLoopOutcome::default(); - record_outcome(&ingest, &mut outcome, true); + outcome.record_outcome(&IngestOutcome { + sampled_token: SampledToken::Undeterminable(LlamaToken(9)), + piece: TokenPiece::Visible("ignored".to_owned()), + }); - assert_eq!(outcome.observed_reasoning, 1); + assert_eq!(outcome.observed_undeterminable, 1); + assert!(outcome.content_stream.is_empty()); assert!(outcome.reasoning_stream.is_empty()); } #[test] - fn record_outcome_undeterminable_token_counts_without_streaming() { - let ingest = IngestOutcome { - sampled_token: SampledToken::Undeterminable(LlamaToken(9)), - visible_piece: "ignored".to_string(), - raw_piece: String::new(), - }; + fn counts_the_end_of_generation_in_the_section_it_ended() { let mut outcome = ClassifySampleLoopOutcome::default(); - record_outcome(&ingest, &mut outcome, false); + outcome.record_end_of_generation(SampledTokenSection::Content); - assert_eq!(outcome.observed_undeterminable, 1); - assert!(outcome.content_stream.is_empty()); - assert!(outcome.reasoning_stream.is_empty()); + assert!(outcome.eog_seen); + assert_eq!(outcome.observed_content, 1); } } diff --git a/llama-cpp-bindings-tests/tests/chat_protocol.rs b/llama-cpp-bindings-tests/tests/chat_protocol.rs index 09dd6d8f8..6419a2c29 100644 --- a/llama-cpp-bindings-tests/tests/chat_protocol.rs +++ b/llama-cpp-bindings-tests/tests/chat_protocol.rs @@ -1,8 +1,8 @@ use anyhow::Result; use anyhow::bail; use llama_cpp_bindings::ChatMessageParseOutcome; +use llama_cpp_bindings::ChatMessageParser; 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; @@ -249,11 +249,7 @@ 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( - &ChatTools::from_json("[]".to_owned())?, - "hello world", - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, "[]")?.parse("hello world", false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for plain content; got Unrecognized"); @@ -299,10 +295,7 @@ 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(&ChatTools::from_json("[]".to_owned())?, input, false)?; + let outcome = ChatMessageParser::new(fixture.model, "[]")?.parse(input, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for reasoning section; got Unrecognized"); @@ -350,10 +343,7 @@ 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(&ChatTools::from_json("[]".to_owned())?, "", false)?; + let outcome = ChatMessageParser::new(fixture.model, "[]")?.parse("", false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for empty input; got Unrecognized"); @@ -374,11 +364,7 @@ fn parses_empty_input_yields_empty_message(fixture: &LlamaFixture<'_>) -> Result fn parses_with_input_null_byte_reports_the_input_as_the_source( fixture: &LlamaFixture<'_>, ) -> Result<()> { - let result = fixture.model.parse_chat_message( - &ChatTools::from_json("[]".to_owned())?, - "hello\0world", - false, - ); + let result = ChatMessageParser::new(fixture.model, "[]")?.parse("hello\0world", false); let Err(ParseChatMessageError::InputContainsNulByte(nul_error)) = result else { anyhow::bail!("a NUL byte in the message must be reported against the message"); @@ -397,23 +383,44 @@ fn parses_with_input_null_byte_reports_the_input_as_the_source( n_batch = 128, n_ubatch = 64, )] -fn parses_with_a_tool_missing_its_function_name_reports_a_tools_parser_build_failure( +fn chat_message_parser_for_a_tool_missing_its_function_name_fails_to_build( fixture: &LlamaFixture<'_>, ) -> Result<()> { - let result = fixture.model.parse_chat_message( - &ChatTools::from_json( - r#"[{"type":"function","function":{"description":"reports the weather"}}]"#.to_owned(), - )?, - "hello", - false, + let result = ChatMessageParser::new( + fixture.model, + r#"[{"type":"function","function":{"description":"reports the weather"}}]"#, ); assert_eq!( - result.unwrap_err(), - ParseChatMessageError::ToolsParserBuildFailed { + result.err(), + Some(ParseChatMessageError::ToolsParserBuildFailed { message: "[json.exception.out_of_range.403] key 'name' not found".to_owned(), - } + }) ); 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 = 512, + n_batch = 128, + n_ubatch = 64, +)] +fn qwen35_chat_message_parser_parses_every_message_it_is_given( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let parser = ChatMessageParser::new(fixture.model, "[]")?; + + for message in ["the first answer", "the second answer"] { + let ChatMessageParseOutcome::Recognized(parsed) = parser.parse(message, false)? else { + bail!("expected {message:?} to be recognized as content"); + }; + + assert_eq!(parsed.content, message); + } + + Ok(()) +} diff --git a/llama-cpp-bindings-tests/tests/context_state.rs b/llama-cpp-bindings-tests/tests/context_state.rs index 207640c2e..049638989 100644 --- a/llama-cpp-bindings-tests/tests/context_state.rs +++ b/llama-cpp-bindings-tests/tests/context_state.rs @@ -6,6 +6,7 @@ use anyhow::Result; use llama_cpp_bindings::DecodeError; use llama_cpp_bindings::LogitsError; use llama_cpp_bindings::context::LlamaContext; +use llama_cpp_bindings::context::params::LlamaAttentionType; use llama_cpp_bindings::error::ClearKvCacheSeqError; use llama_cpp_bindings::error::CopyKvCacheSeqError; use llama_cpp_bindings::error::KvCacheSeqAddError; @@ -13,6 +14,9 @@ use llama_cpp_bindings::error::KvCacheSeqDivError; use llama_cpp_bindings::llama_batch::LlamaBatch; use llama_cpp_bindings::model::AddBos; use llama_cpp_bindings::model::lora_adapter_scale::LoraAdapterScale; +use llama_cpp_bindings::token::LlamaToken; +use llama_cpp_bindings::token::data::LlamaTokenData; +use llama_cpp_bindings::token::data_array::LlamaTokenDataArray; use llama_cpp_bindings_tests::prime_kv_cache::prime_kv_cache; use llama_cpp_bindings_tests::prime_kv_cache_with::prime_kv_cache_with; use llama_cpp_test_harness::LlamaFixture; @@ -2639,3 +2643,76 @@ fn state_seq_get_data_ext_and_set_data_ext_round_trip(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 = 256, + n_batch = 128, + n_ubatch = 64, +)] +fn qwen35_refilled_candidate_array_matches_a_freshly_built_one( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let model = fixture.model; + let mut context = LlamaContext::from_model( + model, + fixture.backend, + (*fixture.context_params).into_llama_context_params(), + )?; + let prompt_tokens = model.str_to_token("Hello", AddBos::Never)?; + let mut batch = LlamaBatch::new(64, 1)?; + + batch.add_sequence(&prompt_tokens, 0, false)?; + context.decode(&mut batch)?; + + let last_index = batch.n_tokens() - 1; + let mut reused_candidates = LlamaTokenDataArray::new( + vec![LlamaTokenData::new(LlamaToken::new(7), 1.0, 0.5)], + true, + ); + reused_candidates.selected = Some(0); + + context.fill_token_data_array_ith(last_index, &mut reused_candidates)?; + + assert_eq!(reused_candidates, context.token_data_array_ith(last_index)?); + + 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 = 512, + n_batch = 512, + n_ubatch = 64, +)] +fn decoding_more_tokens_than_the_micro_batch_with_non_causal_attention_returns_an_error( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let mut context = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params) + .into_llama_context_params() + .with_attention_type(LlamaAttentionType::NonCausal), + )?; + let tokens = fixture + .model + .str_to_token(&"hello ".repeat(100), AddBos::Always)?; + let mut batch = LlamaBatch::new(512, 1)?; + batch.add_sequence(&tokens, 0, false)?; + let n_tokens = batch.n_tokens(); + + assert_eq!( + context.decode(&mut batch), + Err(DecodeError::BatchExceedsMicroBatch { + n_tokens, + n_ubatch: context.n_ubatch(), + }) + ); + + Ok(()) +} diff --git a/llama-cpp-bindings-tests/tests/embedding_models.rs b/llama-cpp-bindings-tests/tests/embedding_models.rs index 12a8d0742..9586a4ac1 100644 --- a/llama-cpp-bindings-tests/tests/embedding_models.rs +++ b/llama-cpp-bindings-tests/tests/embedding_models.rs @@ -4,8 +4,11 @@ use std::time::Duration; use anyhow::Context; use anyhow::Result; use anyhow::bail; +use llama_cpp_bindings::BareJsonToolCalls; use llama_cpp_bindings::ClearKvCacheSeqError; use llama_cpp_bindings::CopyKvCacheSeqError; +use llama_cpp_bindings::DecodeError; +use llama_cpp_bindings::EncodeError; use llama_cpp_bindings::KvCacheSeqAddError; use llama_cpp_bindings::KvCacheSeqDivError; use llama_cpp_bindings::KvCacheSeqPosMaxError; @@ -65,7 +68,7 @@ fn embedding_generation_produces_vectors(fixture: &LlamaFixture<'_>) -> Result<( let t_main_start = ggml_time_us(); - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let mut batch = LlamaBatch::new(n_ctx, 1)?; classifier.feed_prompt_sequence_to_batch(&mut batch, &tokens, 0, false)?; @@ -166,7 +169,7 @@ fn reranking_produces_scores(fixture: &LlamaFixture<'_>) -> Result<()> { bail!("one of the provided prompts exceeds the size of the context window"); } - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let mut batch = LlamaBatch::new(2048, i32::try_from(document_count)?)?; let t_main_start = ggml_time_us(); @@ -657,3 +660,65 @@ fn approximate_tok_env_falls_back_to_eos_when_eot_unavailable( Ok(()) } + +#[llama_test( + model_source = HuggingFace("nomic-ai/nomic-embed-text-v1.5-GGUF", "nomic-embed-text-v1.5.Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 512, + n_batch = 512, + n_ubatch = 64, + embeddings = true, +)] +fn decoding_more_tokens_than_the_micro_batch_on_an_encoder_returns_an_error( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let mut context = fixture.build_context()?; + let tokens = fixture + .model + .str_to_token(&"hello ".repeat(100), AddBos::Always)?; + let mut batch = LlamaBatch::new(512, 1)?; + batch.add_sequence(&tokens, 0, false)?; + let n_tokens = batch.n_tokens(); + + assert_eq!( + context.decode(&mut batch), + Err(DecodeError::BatchExceedsMicroBatch { + n_tokens, + n_ubatch: context.n_ubatch(), + }) + ); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("Xiaojian9992024/t5-small-GGUF", "t5-small.bf16.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 512, + n_batch = 512, + n_ubatch = 64, + embeddings = true, +)] +fn encoding_more_tokens_than_the_micro_batch_returns_an_error( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let mut context = fixture.build_context()?; + let tokens = fixture + .model + .str_to_token(&"hello ".repeat(100), AddBos::Never)?; + let mut batch = LlamaBatch::new(512, 1)?; + batch.add_sequence(&tokens, 0, false)?; + let n_tokens = batch.n_tokens(); + + assert_eq!( + context.encode(&mut batch), + Err(EncodeError::BatchExceedsMicroBatch { + n_tokens, + n_ubatch: context.n_ubatch(), + }) + ); + + Ok(()) +} diff --git a/llama-cpp-bindings-tests/tests/generation_control.rs b/llama-cpp-bindings-tests/tests/generation_control.rs index 2fe178733..5cf1e6f8a 100644 --- a/llama-cpp-bindings-tests/tests/generation_control.rs +++ b/llama-cpp-bindings-tests/tests/generation_control.rs @@ -5,6 +5,10 @@ use std::time::Duration; use anyhow::Context as _; use anyhow::Result; +use anyhow::bail; +use llama_cpp_bindings::BareJsonToolCalls; +use llama_cpp_bindings::GbnfValidationError; +use llama_cpp_bindings::GenerationProgress; use llama_cpp_bindings::GrammarError; use llama_cpp_bindings::SampledToken; use llama_cpp_bindings::SamplerAcceptError; @@ -132,15 +136,15 @@ fn grammar_sampler_constrains_output_to_yes_or_no(fixture: &LlamaFixture<'_>) -> LlamaSampler::greedy()?, ])?; - let mut classifier = model.sampled_token_classifier()?; - let turn = classifier.sample(&mut sampler, &context, batch.n_tokens() - 1)?; - let mut outcomes = turn.outcomes; - outcomes.extend(classifier.flush()); + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; + let mut outcomes = Vec::new(); + let turn = classifier.sample(&mut sampler, &context, batch.n_tokens() - 1, &mut outcomes)?; + classifier.finish(&mut outcomes); assert_eq!( outcomes.len(), 1, - "expected one finalised outcome after flush" + "expected one finalised outcome after finishing" ); let outcome = &outcomes[0]; @@ -150,7 +154,7 @@ fn grammar_sampler_constrains_output_to_yes_or_no(fixture: &LlamaFixture<'_>) -> "Grammar sampler should not allow EOS as first token" ); - let piece = &outcome.raw_piece; + let piece = outcome.piece.raw(); let first_char = piece .chars() .next() @@ -230,15 +234,15 @@ fn json_schema_grammar_sampler_constrains_output_to_json(fixture: &LlamaFixture< LlamaSampler::greedy()?, ])?; - let mut classifier = model.sampled_token_classifier()?; - let turn = classifier.sample(&mut sampler, &context, batch.n_tokens() - 1)?; - let mut outcomes = turn.outcomes; - outcomes.extend(classifier.flush()); + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; + let mut outcomes = Vec::new(); + let turn = classifier.sample(&mut sampler, &context, batch.n_tokens() - 1, &mut outcomes)?; + classifier.finish(&mut outcomes); assert_eq!( outcomes.len(), 1, - "expected one finalised outcome after flush" + "expected one finalised outcome after finishing" ); let outcome = &outcomes[0]; @@ -248,7 +252,7 @@ fn json_schema_grammar_sampler_constrains_output_to_json(fixture: &LlamaFixture< "Grammar sampler should not allow EOS as first token" ); - let piece = &outcome.raw_piece; + let piece = outcome.piece.raw(); assert!( piece.starts_with('{'), @@ -309,7 +313,7 @@ fn sample_with_grammar_produces_constrained_output_in_loop( let tokens = model.str_to_token(prompt, AddBos::Always)?; let mut batch = LlamaBatch::new(512, 1)?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; classifier.feed_prompt_sequence_to_batch(&mut batch, &tokens, 0, false)?; context.decode(&mut batch)?; @@ -323,7 +327,6 @@ fn sample_with_grammar_produces_constrained_output_in_loop( let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -408,26 +411,28 @@ fn sample_without_grammar_produces_multiple_tokens(fixture: &LlamaFixture<'_>) - let mut sampler = LlamaSampler::chain_simple([LlamaSampler::temp(0.8)?, LlamaSampler::greedy()?])?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; + let mut outcomes = Vec::new(); let mut sampled_count: u64 = 0; for (position, _) in (batch.n_tokens()..).zip(0..5) { - let turn = classifier.sample(&mut sampler, &context, -1)?; - let raw_as_sampled = SampledToken::Content(turn.token); + let turn = classifier.sample(&mut sampler, &context, -1, &mut outcomes)?; - if model.is_eog_token(&raw_as_sampled) { + if turn.progress == GenerationProgress::Ended { break; } sampled_count += 1; + let raw_as_sampled = SampledToken::Content(turn.token); + batch.clear(); batch.add(&raw_as_sampled, position, &[0], true)?; context.decode(&mut batch)?; } - let _ = classifier.flush(); + classifier.finish(&mut outcomes); assert!( sampled_count > 0, @@ -847,7 +852,7 @@ fn raw_prompt_completion_with_timing(fixture: &LlamaFixture<'_>) -> Result<()> { let prompt = "Hello my name is"; let max_generated_tokens: i32 = 64; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let tokens_list = model .str_to_token(prompt, AddBos::Always) .with_context(|| format!("failed to tokenize {prompt}"))?; @@ -881,7 +886,6 @@ fn raw_prompt_completion_with_timing(fixture: &LlamaFixture<'_>) -> Result<()> { let initial_position = batch.n_tokens(); let t_main_start = ggml_time_us(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut ctx, @@ -994,7 +998,7 @@ fn chat_inference_produces_coherent_output(fixture: &LlamaFixture<'_>) -> Result )?]; let prompt = model.apply_chat_template(&chat_template, &messages, true, true)?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let tokens = model.str_to_token(&prompt, AddBos::Always)?; let prompt_token_count = u64::try_from(tokens.len())?; @@ -1012,7 +1016,6 @@ fn chat_inference_produces_coherent_output(fixture: &LlamaFixture<'_>) -> Result let mut sampler = LlamaSampler::greedy()?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1822,7 +1825,9 @@ fn llguidance_chain_samples_a_valid_token(fixture: &LlamaFixture<'_>) -> Result< fn classifier_starts_in_pending_section_for_default_fixture( fixture: &LlamaFixture<'_>, ) -> Result<()> { - let classifier = fixture.model.sampled_token_classifier()?; + let classifier = fixture + .model + .sampled_token_classifier(BareJsonToolCalls::Detect)?; assert_eq!(classifier.current_section(), SampledTokenSection::Pending); Ok(()) @@ -1861,8 +1866,12 @@ fn classifier_starts_in_pending_section_for_default_fixture( n_ubatch = 64, )] fn classifier_construction_is_idempotent_across_calls(fixture: &LlamaFixture<'_>) -> Result<()> { - let first = fixture.model.sampled_token_classifier()?; - let second = fixture.model.sampled_token_classifier()?; + let first = fixture + .model + .sampled_token_classifier(BareJsonToolCalls::Detect)?; + let second = fixture + .model + .sampled_token_classifier(BareJsonToolCalls::Detect)?; assert_eq!(first.current_section(), second.current_section()); assert_eq!(first.usage(), second.usage()); @@ -1901,14 +1910,18 @@ fn classifier_construction_is_idempotent_across_calls(fixture: &LlamaFixture<'_> n_batch = 128, n_ubatch = 64, )] -fn ingest_flushes_an_unmatched_token_with_its_visible_and_raw_piece( +fn finishing_releases_an_unmatched_token_with_its_visible_and_raw_piece( fixture: &LlamaFixture<'_>, ) -> Result<()> { let model = fixture.model; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; + let [ordinary_token] = model.str_to_token("hello", AddBos::Never)?[..] else { + bail!("\"hello\" must be a single token"); + }; - let mut outcomes = classifier.ingest(model.token_bos())?; - outcomes.extend(classifier.flush()); + let mut outcomes = Vec::new(); + classifier.ingest(ordinary_token, &mut outcomes)?; + classifier.finish(&mut outcomes); assert_eq!(outcomes.len(), 1); let outcome = &outcomes[0]; @@ -1916,7 +1929,7 @@ fn ingest_flushes_an_unmatched_token_with_its_visible_and_raw_piece( outcome.sampled_token, SampledToken::Undeterminable(_) )); - assert_eq!(outcome.visible_piece, outcome.raw_piece); + assert_eq!(outcome.piece.visible(), outcome.piece.raw()); assert_eq!(classifier.usage().undeterminable_tokens, 1); Ok(()) } @@ -1953,13 +1966,22 @@ fn ingest_flushes_an_unmatched_token_with_its_visible_and_raw_piece( n_batch = 128, n_ubatch = 64, )] -fn ingest_accounts_for_each_unmatched_token_after_flush(fixture: &LlamaFixture<'_>) -> Result<()> { +fn ingest_counts_every_token_through_the_end_of_generation( + fixture: &LlamaFixture<'_>, +) -> Result<()> { let model = fixture.model; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; + let [ordinary_token] = model.str_to_token("hello", AddBos::Never)?[..] else { + bail!("\"hello\" must be a single token"); + }; + let mut outcomes = Vec::new(); - classifier.ingest(model.token_bos())?; - classifier.ingest(model.token_eos())?; - classifier.flush(); + classifier.ingest(ordinary_token, &mut outcomes)?; + + assert_eq!( + classifier.ingest(model.token_eos(), &mut outcomes)?, + GenerationProgress::Ended + ); assert_eq!(classifier.usage().undeterminable_tokens, 2); Ok(()) @@ -1999,7 +2021,7 @@ fn ingest_accounts_for_each_unmatched_token_after_flush(fixture: &LlamaFixture<' )] fn ingest_unmatched_prompt_tokens_does_not_record_usage(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let usage_before = *classifier.usage(); classifier.ingest_prompt_token(model.token_bos()); @@ -2044,7 +2066,7 @@ fn ingest_unmatched_prompt_tokens_does_not_record_usage(fixture: &LlamaFixture<' )] fn feed_prompt_to_batch_increments_pending_prompt_tokens(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let mut batch = LlamaBatch::new(8, 1)?; classifier.feed_prompt_to_batch(&mut batch, model.token_bos(), 0, &[0], false)?; @@ -2090,7 +2112,7 @@ fn feed_prompt_to_batch_increments_pending_prompt_tokens(fixture: &LlamaFixture< )] fn feed_prompt_sequence_to_batch_stages_all_tokens(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let mut batch = LlamaBatch::new(8, 1)?; let tokens = vec![model.token_bos(), model.token_eos(), model.token_nl()]; @@ -2138,7 +2160,7 @@ fn commit_prompt_tokens_promotes_pending_count_to_usage_and_clears( fixture: &LlamaFixture<'_>, ) -> Result<()> { let model = fixture.model; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let mut batch = LlamaBatch::new(8, 1)?; classifier.feed_prompt_to_batch(&mut batch, model.token_bos(), 0, &[0], false)?; @@ -2189,7 +2211,7 @@ fn discard_pending_prompt_tokens_clears_count_without_recording_usage( fixture: &LlamaFixture<'_>, ) -> Result<()> { let model = fixture.model; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let mut batch = LlamaBatch::new(8, 1)?; classifier.feed_prompt_to_batch(&mut batch, model.token_bos(), 0, &[0], false)?; @@ -2258,3 +2280,80 @@ fn diagnose_tool_call_synthetic_renders_applies_the_template_to_both_probes( 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 = 256, + n_batch = 128, + n_ubatch = 64, +)] +fn qwen35_grammar_with_a_token_reference_constrains_the_first_token( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let model = fixture.model; + let [think_token] = model.str_to_token("", AddBos::Never)?[..] else { + bail!(" must be a single token"); + }; + let mut context = LlamaContext::from_model( + model, + fixture.backend, + (*fixture.context_params).into_llama_context_params(), + )?; + let prompt_tokens = model.str_to_token( + "<|im_start|>user\nSay hi<|im_end|>\n<|im_start|>assistant\n", + AddBos::Never, + )?; + let mut batch = LlamaBatch::new(128, 1)?; + + batch.add_sequence(&prompt_tokens, 0, false)?; + context.decode(&mut batch)?; + + let mut sampler = LlamaSampler::chain_simple([ + LlamaSampler::grammar(model, r#"root ::= "x""#, "root")?, + LlamaSampler::greedy()?, + ])?; + + assert_eq!(sampler.sample(&context, batch.n_tokens() - 1)?, think_token); + + 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 = 256, + n_batch = 128, + n_ubatch = 64, +)] +fn qwen35_grammar_without_its_root_reports_the_missing_root( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + assert_eq!( + LlamaSampler::grammar(fixture.model, r#"answer ::= "x""#, "root").err(), + Some(GrammarError::RootNotFound) + ); + + 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 = 256, + n_batch = 128, + n_ubatch = 64, +)] +fn qwen35_grammar_with_a_syntax_error_reports_it(fixture: &LlamaFixture<'_>) -> Result<()> { + assert_eq!( + LlamaSampler::grammar(fixture.model, "root ::= (", "root").err(), + Some(GrammarError::GrammarRejected( + GbnfValidationError::SyntaxError + )) + ); + + Ok(()) +} diff --git a/llama-cpp-bindings-tests/tests/model_introspection.rs b/llama-cpp-bindings-tests/tests/model_introspection.rs index c19e942bb..0ec046ac1 100644 --- a/llama-cpp-bindings-tests/tests/model_introspection.rs +++ b/llama-cpp-bindings-tests/tests/model_introspection.rs @@ -1760,3 +1760,31 @@ fn debug_format_includes_struct_name_and_model_field(fixture: &LlamaFixture<'_>) 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 = 256, + n_batch = 128, + n_ubatch = 64, +)] +fn deepseek_r1_8b_adds_a_bos_token(fixture: &LlamaFixture<'_>) -> Result<()> { + assert!(fixture.model.adds_bos_token()); + + 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 = 256, + n_batch = 128, + n_ubatch = 64, +)] +fn qwen35_does_not_add_a_bos_token(fixture: &LlamaFixture<'_>) -> Result<()> { + assert!(!fixture.model.adds_bos_token()); + + Ok(()) +} diff --git a/llama-cpp-bindings-tests/tests/multimodal_audio.rs b/llama-cpp-bindings-tests/tests/multimodal_audio.rs index 99f6b95d7..8d2e56e1b 100644 --- a/llama-cpp-bindings-tests/tests/multimodal_audio.rs +++ b/llama-cpp-bindings-tests/tests/multimodal_audio.rs @@ -1,5 +1,6 @@ use anyhow::Context; use anyhow::Result; +use llama_cpp_bindings::BareJsonToolCalls; use llama_cpp_bindings::EvalMultimodalChunksParams; use llama_cpp_bindings::context::LlamaContext; use llama_cpp_bindings::llama_batch::LlamaBatch; @@ -64,7 +65,7 @@ fn transcribe_audio(fixture: &LlamaFixture<'_>, audio_file_name: &str) -> Result "tokenization should produce at least one chunk" ); - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let n_past = classifier .eval_multimodal_chunks( &chunks, @@ -99,7 +100,6 @@ fn transcribe_audio(fixture: &LlamaFixture<'_>, audio_file_name: &str) -> Result let mut sampler = LlamaSampler::greedy()?; let mut batch = LlamaBatch::new(512, 1)?; let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, diff --git a/llama-cpp-bindings-tests/tests/multimodal_fusion.rs b/llama-cpp-bindings-tests/tests/multimodal_fusion.rs index 0510cf7ef..6b14c7ed1 100644 --- a/llama-cpp-bindings-tests/tests/multimodal_fusion.rs +++ b/llama-cpp-bindings-tests/tests/multimodal_fusion.rs @@ -1,5 +1,6 @@ use anyhow::Context; use anyhow::Result; +use llama_cpp_bindings::BareJsonToolCalls; use llama_cpp_bindings::EvalMultimodalChunksParams; use llama_cpp_bindings::context::LlamaContext; use llama_cpp_bindings::llama_batch::LlamaBatch; @@ -109,7 +110,7 @@ fn image_and_audio_together(fixture: &LlamaFixture<'_>) -> Result<()> { .with_context(|| "unable to create llama context")?; let n_batch = i32::try_from(context.n_batch())?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let n_past = classifier .eval_multimodal_chunks( &chunks, @@ -134,7 +135,6 @@ fn image_and_audio_together(fixture: &LlamaFixture<'_>) -> Result<()> { let mut sampler = LlamaSampler::greedy()?; let mut batch = LlamaBatch::new(512, 1)?; let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, diff --git a/llama-cpp-bindings-tests/tests/multimodal_vision.rs b/llama-cpp-bindings-tests/tests/multimodal_vision.rs index 6090ef1ad..0ee90825a 100644 --- a/llama-cpp-bindings-tests/tests/multimodal_vision.rs +++ b/llama-cpp-bindings-tests/tests/multimodal_vision.rs @@ -1,13 +1,12 @@ use anyhow::Context; use anyhow::Result; +use llama_cpp_bindings::BareJsonToolCalls; use llama_cpp_bindings::EvalMultimodalChunksParams; -use llama_cpp_bindings::SampledToken; -use llama_cpp_bindings::SampledTokenClassifier; use llama_cpp_bindings::TokenUsage; use llama_cpp_bindings::context::LlamaContext; +use llama_cpp_bindings::error::EvalMultimodalChunksError; use llama_cpp_bindings::ingest_prompt_chunk::ingest_prompt_chunk; use llama_cpp_bindings::llama_batch::LlamaBatch; -use llama_cpp_bindings::model::LlamaModel; use llama_cpp_bindings::mtmd::MtmdBitmap; use llama_cpp_bindings::mtmd::MtmdContext; use llama_cpp_bindings::mtmd::MtmdContextParams; @@ -15,9 +14,9 @@ use llama_cpp_bindings::mtmd::MtmdEvalError; use llama_cpp_bindings::mtmd::MtmdInputChunkType; use llama_cpp_bindings::mtmd::MtmdInputChunks; use llama_cpp_bindings::mtmd::MtmdInputText; +use llama_cpp_bindings::mtmd::NonCausalChunkMicroBatchMismatch; use llama_cpp_bindings::mtmd::mtmd_default_marker; use llama_cpp_bindings::sampling::LlamaSampler; -use llama_cpp_bindings_sys::llama_pos; use llama_cpp_bindings_tests::build_user_prompt_with_media_marker::build_user_prompt_with_media_marker; use llama_cpp_bindings_tests::chunk_token_breakdown::ChunkTokenBreakdown; use llama_cpp_bindings_tests::classify_sample_loop::ClassifySampleLoop; @@ -660,6 +659,137 @@ fn eval_chunks_returns_batch_size_exceeds_context_limit_for_huge_batch( 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 = 512, + n_batch = 128, + n_ubatch = 64, + mmproj_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "mmproj-F16.gguf"), +)] +#[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, + mmproj_source = HuggingFace("unsloth/Qwen3.6-35B-A3B-GGUF", "mmproj-F16.gguf"), +)] +fn eval_chunks_rejects_a_zero_batch_before_evaluating(fixture: &LlamaFixture<'_>) -> Result<()> { + let mtmd_ctx = fixture + .mtmd_context + .expect("mmproj_file declared in attribute"); + let chunks = tokenize_synthetic(fixture, "Describe: <__media__>")?; + let llama_ctx = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params).into_llama_context_params(), + )?; + + assert_eq!( + chunks.eval_chunks(mtmd_ctx, &llama_ctx, 0, 0, 0, true), + Err(MtmdEvalError::NonPositiveBatchSize { requested: 0 }) + ); + assert_eq!(llama_ctx.kv_cache_seq_pos_max(0)?, -1); + + 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 = 512, + n_batch = 128, + n_ubatch = 64, + mmproj_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "mmproj-F16.gguf"), +)] +#[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, + mmproj_source = HuggingFace("unsloth/Qwen3.6-35B-A3B-GGUF", "mmproj-F16.gguf"), +)] +fn eval_single_rejects_a_zero_batch_before_evaluating(fixture: &LlamaFixture<'_>) -> Result<()> { + let mtmd_ctx = fixture + .mtmd_context + .expect("mmproj_file declared in attribute"); + let chunks = tokenize_synthetic(fixture, "Describe: <__media__>")?; + let first_chunk = chunks.get(0).context("tokenization produced no chunks")?; + let llama_ctx = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params).into_llama_context_params(), + )?; + + assert_eq!( + first_chunk.eval_single(mtmd_ctx, &llama_ctx, 0, 0, 0, true), + Err(MtmdEvalError::NonPositiveBatchSize { requested: 0 }) + ); + assert_eq!(llama_ctx.kv_cache_seq_pos_max(0)?, -1); + + 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 = 512, + n_batch = 128, + n_ubatch = 64, + mmproj_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "mmproj-F16.gguf"), +)] +#[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, + mmproj_source = HuggingFace("unsloth/Qwen3.6-35B-A3B-GGUF", "mmproj-F16.gguf"), +)] +fn classifier_rejects_a_zero_batch_before_evaluating(fixture: &LlamaFixture<'_>) -> Result<()> { + let mtmd_ctx = fixture + .mtmd_context + .expect("mmproj_file declared in attribute"); + let chunks = tokenize_synthetic(fixture, "Describe: <__media__>")?; + let llama_ctx = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params).into_llama_context_params(), + )?; + let mut classifier = fixture + .model + .sampled_token_classifier(BareJsonToolCalls::Ignore)?; + + assert!(matches!( + classifier.eval_multimodal_chunks( + &chunks, + mtmd_ctx, + &llama_ctx, + EvalMultimodalChunksParams { + start_position: 0, + seq_id: 0, + n_batch: 0, + logits_last: true, + }, + ), + Err(EvalMultimodalChunksError::EvalFailed( + MtmdEvalError::NonPositiveBatchSize { requested: 0 } + )) + )); + assert_eq!(classifier.usage().prompt_tokens, 0); + assert_eq!(llama_ctx.kv_cache_seq_pos_max(0)?, -1); + + Ok(()) +} + #[llama_test( model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -887,62 +1017,6 @@ fn tokenize_with_null_byte_in_text_returns_error(fixture: &LlamaFixture<'_>) -> Ok(()) } -struct SamplingTotals { - generated: String, - observed_content: u64, - observed_reasoning: u64, -} - -fn drive_sampling_loop( - classifier: &mut SampledTokenClassifier, - model: &LlamaModel, - ctx: &mut LlamaContext, - starting_position: llama_pos, - max_tokens: usize, -) -> Result { - let mut sampler = LlamaSampler::greedy()?; - let mut totals = SamplingTotals { - generated: String::new(), - observed_content: 0, - observed_reasoning: 0, - }; - let mut batch = LlamaBatch::new(512, 1)?; - - for (current_position, _) in (starting_position..).zip(0..max_tokens) { - let turn = classifier.sample(&mut sampler, ctx, -1)?; - for outcome in &turn.outcomes { - totals.generated.push_str(&outcome.raw_piece); - match outcome.sampled_token { - SampledToken::Content(_) => totals.observed_content += 1, - SampledToken::Reasoning(_) => totals.observed_reasoning += 1, - SampledToken::ToolCall(_) | SampledToken::Undeterminable(_) => {} - } - } - - let raw_as_sampled = SampledToken::Content(turn.token); - if model.is_eog_token(&raw_as_sampled) { - break; - } - - batch.clear(); - batch.add(&raw_as_sampled, current_position, &[0], true)?; - - ctx.decode(&mut batch) - .with_context(|| "failed to decode generated token")?; - } - - for outcome in classifier.flush() { - totals.generated.push_str(&outcome.raw_piece); - match outcome.sampled_token { - SampledToken::Content(_) => totals.observed_content += 1, - SampledToken::Reasoning(_) => totals.observed_reasoning += 1, - SampledToken::ToolCall(_) | SampledToken::Undeterminable(_) => {} - } - } - - Ok(totals) -} - #[llama_test( model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -1010,7 +1084,7 @@ fn multimodal_vision_inference_produces_output(fixture: &LlamaFixture<'_>) -> Re "vision input must produce at least one image chunk" ); - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let n_past = classifier .eval_multimodal_chunks( &chunks, @@ -1034,12 +1108,22 @@ fn multimodal_vision_inference_produces_output(fixture: &LlamaFixture<'_>) -> Re assert_eq!(usage.input_audio_tokens, expected.audio); } - let totals = drive_sampling_loop(&mut classifier, model, &mut ctx, n_past, 512)?; + let mut sampler = LlamaSampler::greedy()?; + let mut batch = LlamaBatch::new(512, 1)?; + let totals = ClassifySampleLoop { + classifier: &mut classifier, + sampler: &mut sampler, + context: &mut ctx, + batch: &mut batch, + initial_position: n_past, + max_generated_tokens: 512, + } + .run()?; - eprintln!("generated text: {}", totals.generated); + eprintln!("generated text: {}", totals.generated_raw); assert!( - !totals.generated.is_empty(), + !totals.generated_raw.is_empty(), "model should generate at least one token from image input" ); @@ -1088,7 +1172,7 @@ fn build_multimodal_chunks_and_eval_into_usage( let context_params = (*fixture.context_params).into_llama_context_params(); let context = LlamaContext::from_model(model, fixture.backend, context_params)?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; classifier.eval_multimodal_chunks( &chunks, mtmd_ctx, @@ -1233,7 +1317,7 @@ fn text_chunk_records_prompt_tokens(fixture: &LlamaFixture<'_>) -> Result<()> { let n_tokens = u64::try_from(text_chunk.n_tokens())?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; ingest_prompt_chunk(&mut classifier, &text_chunk)?; @@ -1299,7 +1383,7 @@ fn image_chunk_records_input_image_tokens_only(fixture: &LlamaFixture<'_>) -> Re anyhow::bail!("image chunk should report at least one token"); } - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; ingest_prompt_chunk(&mut classifier, &image_chunk)?; @@ -1348,7 +1432,7 @@ fn text_chunk_drives_marker_state_machine_to_reasoning(fixture: &LlamaFixture<'_ }; let chunks = mtmd_ctx.tokenize(input_text, &[])?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; for index in 0..chunks.len() { let chunk = chunks @@ -1413,7 +1497,7 @@ fn gemma4_classifier_emits_reasoning_for_multimodal_thinking_prompt( let chunks = mtmd_ctx.tokenize(input_text, &[&bitmap])?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let n_past = classifier.eval_multimodal_chunks( &chunks, mtmd_ctx, @@ -1437,7 +1521,6 @@ fn gemma4_classifier_emits_reasoning_for_multimodal_thinking_prompt( let mut batch = LlamaBatch::new(2048, 1)?; let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1517,7 +1600,7 @@ fn mistral3_classifier_emits_reasoning_for_multimodal_thinking_prompt( let chunks = mtmd_ctx.tokenize(input_text, &[&bitmap])?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let n_past = classifier.eval_multimodal_chunks( &chunks, mtmd_ctx, @@ -1533,7 +1616,6 @@ fn mistral3_classifier_emits_reasoning_for_multimodal_thinking_prompt( let mut sampler = LlamaSampler::greedy()?; let mut batch = LlamaBatch::new(2048, 1)?; let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1614,7 +1696,7 @@ fn qwen35_classifier_emits_reasoning_for_multimodal_thinking_prompt( let chunks = mtmd_ctx.tokenize(input_text, &[&bitmap])?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let n_past = classifier.eval_multimodal_chunks( &chunks, mtmd_ctx, @@ -1638,7 +1720,6 @@ fn qwen35_classifier_emits_reasoning_for_multimodal_thinking_prompt( let mut batch = LlamaBatch::new(2048, 1)?; let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1710,7 +1791,7 @@ fn qwen36_classifier_emits_reasoning_for_multimodal_thinking_prompt( let chunks = mtmd_ctx.tokenize(input_text, &[&bitmap])?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let n_past = classifier.eval_multimodal_chunks( &chunks, mtmd_ctx, @@ -1734,7 +1815,6 @@ fn qwen36_classifier_emits_reasoning_for_multimodal_thinking_prompt( let mut batch = LlamaBatch::new(2048, 1)?; let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1759,3 +1839,166 @@ fn qwen36_classifier_emits_reasoning_for_multimodal_thinking_prompt( Ok(()) } + +fn tokenize_question_about_llamas( + fixture: &LlamaFixture<'_>, + mtmd_ctx: &MtmdContext, +) -> Result { + let image_path = fixtures_dir().join("llamas.jpg"); + let image_path_str = image_path + .to_str() + .with_context(|| "image path is not valid UTF-8")?; + let bitmap = MtmdBitmap::from_file(mtmd_ctx, image_path_str)?; + + Ok(mtmd_ctx.tokenize( + MtmdInputText { + text: build_user_prompt_with_media_marker(fixture.model, PROMPT_QUESTION)?, + add_special: false, + parse_special: true, + }, + &[&bitmap], + )?) +} + +#[llama_test( + model_source = HuggingFace("unsloth/gemma-3-4b-it-GGUF", "gemma-3-4b-it-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 2048, + n_batch = 2048, + n_ubatch = 128, + mmproj_source = HuggingFace("unsloth/gemma-3-4b-it-GGUF", "mmproj-F16.gguf"), +)] +fn gemma3_eval_chunks_rejects_a_non_causal_image_larger_than_the_micro_batch_before_evaluating( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let mtmd_ctx = fixture + .mtmd_context + .expect("mmproj_file declared in attribute"); + let chunks = tokenize_question_about_llamas(fixture, mtmd_ctx)?; + let image_tokens = ChunkTokenBreakdown::from_chunks(&chunks)?.image; + let llama_ctx = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params).into_llama_context_params(), + )?; + + assert_eq!( + chunks.eval_chunks( + mtmd_ctx, + &llama_ctx, + 0, + 0, + i32::try_from(llama_ctx.n_batch())?, + true + ), + Err(MtmdEvalError::NonCausalChunkExceedsMicroBatch( + NonCausalChunkMicroBatchMismatch { + chunk_tokens: usize::try_from(image_tokens)?, + micro_batch_tokens: llama_ctx.n_ubatch(), + } + )) + ); + assert_eq!(llama_ctx.kv_cache_seq_pos_max(0)?, -1); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/gemma-3-4b-it-GGUF", "gemma-3-4b-it-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 2048, + n_batch = 2048, + n_ubatch = 128, + mmproj_source = HuggingFace("unsloth/gemma-3-4b-it-GGUF", "mmproj-F16.gguf"), +)] +fn gemma3_eval_single_rejects_a_non_causal_image_larger_than_the_micro_batch_before_evaluating( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let mtmd_ctx = fixture + .mtmd_context + .expect("mmproj_file declared in attribute"); + let chunks = tokenize_question_about_llamas(fixture, mtmd_ctx)?; + let image_chunk = (0..chunks.len()) + .filter_map(|chunk_index| chunks.get(chunk_index)) + .find(|chunk| chunk.chunk_type() == Ok(MtmdInputChunkType::Image)) + .context("the prompt must contain an image chunk")?; + let llama_ctx = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params).into_llama_context_params(), + )?; + + assert_eq!( + image_chunk.eval_single( + mtmd_ctx, + &llama_ctx, + 0, + 0, + i32::try_from(llama_ctx.n_batch())?, + true + ), + Err(MtmdEvalError::NonCausalChunkExceedsMicroBatch( + NonCausalChunkMicroBatchMismatch { + chunk_tokens: image_chunk.n_tokens(), + micro_batch_tokens: llama_ctx.n_ubatch(), + } + )) + ); + assert_eq!(llama_ctx.kv_cache_seq_pos_max(0)?, -1); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/gemma-3-4b-it-GGUF", "gemma-3-4b-it-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 2048, + n_batch = 2048, + n_ubatch = 128, + mmproj_source = HuggingFace("unsloth/gemma-3-4b-it-GGUF", "mmproj-F16.gguf"), +)] +fn gemma3_classifier_leaves_the_prompt_untouched_when_a_non_causal_image_exceeds_the_micro_batch( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let mtmd_ctx = fixture + .mtmd_context + .expect("mmproj_file declared in attribute"); + let chunks = tokenize_question_about_llamas(fixture, mtmd_ctx)?; + let llama_ctx = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params).into_llama_context_params(), + )?; + let mut classifier = fixture + .model + .sampled_token_classifier(BareJsonToolCalls::Ignore)?; + + let evaluation = classifier.eval_multimodal_chunks( + &chunks, + mtmd_ctx, + &llama_ctx, + EvalMultimodalChunksParams { + start_position: 0, + seq_id: 0, + n_batch: i32::try_from(llama_ctx.n_batch())?, + logits_last: true, + }, + ); + + assert!(matches!( + evaluation, + Err(EvalMultimodalChunksError::EvalFailed( + MtmdEvalError::NonCausalChunkExceedsMicroBatch(NonCausalChunkMicroBatchMismatch { + micro_batch_tokens, + .. + }) + )) if micro_batch_tokens == llama_ctx.n_ubatch() + )); + assert_eq!(classifier.usage().prompt_tokens, 0); + assert_eq!(llama_ctx.kv_cache_seq_pos_max(0)?, -1); + + Ok(()) +} diff --git a/llama-cpp-bindings-tests/tests/structured_chat_output.rs b/llama-cpp-bindings-tests/tests/structured_chat_output.rs index a22fe85d6..55756a8ff 100644 --- a/llama-cpp-bindings-tests/tests/structured_chat_output.rs +++ b/llama-cpp-bindings-tests/tests/structured_chat_output.rs @@ -1,7 +1,9 @@ use anyhow::Result; use anyhow::bail; +use llama_cpp_bindings::BareJsonToolCalls; use llama_cpp_bindings::ChatMessageParseOutcome; -use llama_cpp_bindings::ChatTools; +use llama_cpp_bindings::ChatMessageParser; +use llama_cpp_bindings::GenerationProgress; use llama_cpp_bindings::MarkerRole; use llama_cpp_bindings::ParsedChatMessage; use llama_cpp_bindings::SampledTokenSection; @@ -33,8 +35,7 @@ fn parse_partial_reasoning_response( } else { format!("{}{generated}", markers.open) }; - let parse_outcome = - model.parse_chat_message(&ChatTools::from_json("[]".to_owned())?, &response, true)?; + let parse_outcome = ChatMessageParser::new(model, "[]")?.parse(&response, true)?; let ChatMessageParseOutcome::Recognized(parsed) = parse_outcome else { bail!("model chat template must recognize a partial reasoning response"); }; @@ -66,7 +67,7 @@ fn deepseek_r1_8b_classifier_does_not_emit_reasoning_for_thinking_disabled_promp let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(DEEPSEEK_R1_8B_THINKING_DISABLED_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -95,7 +96,6 @@ fn deepseek_r1_8b_classifier_does_not_emit_reasoning_for_thinking_disabled_promp ])?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -177,7 +177,7 @@ fn deepseek_r1_8b_classifier_emits_reasoning_for_thinking_enabled_prompt( let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(DEEPSEEK_R1_8B_THINKING_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -205,7 +205,6 @@ fn deepseek_r1_8b_classifier_emits_reasoning_for_thinking_enabled_prompt( ])?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -321,11 +320,8 @@ 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( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - GEMMA_PAIRED_QUOTE_PAYLOAD, - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, TOOLS_JSON)? + .parse(GEMMA_PAIRED_QUOTE_PAYLOAD, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -385,11 +381,8 @@ fn deepseek_r1_8b_duck_types_glm_key_value_tags(fixture: &LlamaFixture<'_>) -> R Paris\ "; - let outcome = fixture.model.parse_chat_message( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - GLM_KEY_VALUE_PAYLOAD, - false, - )?; + let outcome = + ChatMessageParser::new(fixture.model, TOOLS_JSON)?.parse(GLM_KEY_VALUE_PAYLOAD, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -447,11 +440,8 @@ 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( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - MISTRAL_BRACKETED_JSON_PAYLOAD, - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, TOOLS_JSON)? + .parse(MISTRAL_BRACKETED_JSON_PAYLOAD, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -514,11 +504,8 @@ Paris\n\ \n\ "; - let outcome = fixture.model.parse_chat_message( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - QWEN_XML_PAYLOAD, - false, - )?; + let outcome = + ChatMessageParser::new(fixture.model, TOOLS_JSON)?.parse(QWEN_XML_PAYLOAD, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -577,11 +564,7 @@ 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( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - PLAIN_CONTENT, - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, TOOLS_JSON)?.parse(PLAIN_CONTENT, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -611,11 +594,7 @@ 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( - &ChatTools::from_json("[]".to_owned())?, - PLAIN_CONTENT, - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, "[]")?.parse(PLAIN_CONTENT, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("plain content with empty tools array must produce Recognized; got Unrecognized"); @@ -651,7 +630,7 @@ fn gemma4_classifier_does_not_emit_reasoning_for_thinking_disabled_prompt( let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(GEMMA4_THINKING_DISABLED_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -672,7 +651,6 @@ fn gemma4_classifier_does_not_emit_reasoning_for_thinking_disabled_prompt( let mut sampler = LlamaSampler::greedy()?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -752,7 +730,7 @@ fn gemma4_classifier_emits_reasoning_for_thinking_prompt(fixture: &LlamaFixture< let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(GEMMA4_THINKING_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -773,7 +751,6 @@ fn gemma4_classifier_emits_reasoning_for_thinking_prompt(fixture: &LlamaFixture< let mut sampler = LlamaSampler::greedy()?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -868,11 +845,8 @@ 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( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - GEMMA4_PAIRED_QUOTE_PAYLOAD, - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, TOOLS_JSON)? + .parse(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"); @@ -961,7 +935,7 @@ What is 2 + 2? let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(GLM47_THINKING_DISABLED_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -989,7 +963,6 @@ What is 2 + 2? ])?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1041,7 +1014,7 @@ What is 2 + 2? let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(GLM47_THINKING_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -1069,7 +1042,6 @@ What is 2 + 2? ])?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1138,11 +1110,8 @@ fn glm47_parses_tool_call_payload(fixture: &LlamaFixture<'_>) -> Result<()> { Paris\ "; - let outcome = fixture.model.parse_chat_message( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - GLM47_KEY_VALUE_PAYLOAD, - false, - )?; + let outcome = + ChatMessageParser::new(fixture.model, TOOLS_JSON)?.parse(GLM47_KEY_VALUE_PAYLOAD, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -1222,7 +1191,7 @@ fn mistral3_classifier_does_not_emit_reasoning_for_thinking_disabled_prompt( let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(MISTRAL3_THINKING_DISABLED_PROMPT, AddBos::Always)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -1243,7 +1212,6 @@ fn mistral3_classifier_does_not_emit_reasoning_for_thinking_disabled_prompt( let mut sampler = LlamaSampler::greedy()?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1299,7 +1267,7 @@ to the user.[/THINK]Here, provide a self-contained response.[/SYSTEM_PROMPT]\ let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(MISTRAL3_THINKING_PROMPT, AddBos::Always)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -1320,7 +1288,6 @@ to the user.[/THINK]Here, provide a self-contained response.[/SYSTEM_PROMPT]\ let mut sampler = LlamaSampler::greedy()?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1383,11 +1350,8 @@ 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( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - MISTRAL3_BRACKETED_JSON_PAYLOAD, - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, TOOLS_JSON)? + .parse(MISTRAL3_BRACKETED_JSON_PAYLOAD, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -1437,7 +1401,7 @@ fn qwen35_chat_inference_emits_reasoning_when_template_auto_opens( )?]; let prompt = model.apply_chat_template(&chat_template, &messages, true, true)?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let tokens = model.str_to_token(&prompt, AddBos::Always)?; let prompt_token_count = u64::try_from(tokens.len())?; @@ -1452,7 +1416,6 @@ fn qwen35_chat_inference_emits_reasoning_when_template_auto_opens( let mut sampler = LlamaSampler::greedy()?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1519,7 +1482,9 @@ fn qwen35_shared_reasoning_close_and_tool_call_open_is_one_transition( .tokens() .to_vec(); - let mut classifier = fixture.model.sampled_token_classifier()?; + let mut classifier = fixture + .model + .sampled_token_classifier(BareJsonToolCalls::Detect)?; classifier.ingest_prompt_tokens(&reasoning_open); assert_eq!(classifier.current_section(), SampledTokenSection::Reasoning); @@ -1557,7 +1522,7 @@ What is 2 + 2?<|im_end|> let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(QWEN35_THINKING_DISABLED_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -1585,7 +1550,6 @@ What is 2 + 2?<|im_end|> ])?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1637,7 +1601,7 @@ What is 2 + 2?<|im_end|> let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(QWEN35_THINKING_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -1665,7 +1629,6 @@ What is 2 + 2?<|im_end|> ])?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -1767,11 +1730,8 @@ get off the keyboard\n\ \n\ "; - let outcome = fixture.model.parse_chat_message( - &ChatTools::from_json(NEGOTIATE_WITH_CAT_TOOLS_JSON.to_owned())?, - NEGOTIATE_WITH_CAT_INPUT, - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, NEGOTIATE_WITH_CAT_TOOLS_JSON)? + .parse(NEGOTIATE_WITH_CAT_INPUT, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -1829,11 +1789,8 @@ Paris\n\ \n\ "; - let outcome = fixture.model.parse_chat_message( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - QWEN_XML_PAYLOAD, - false, - )?; + let outcome = + ChatMessageParser::new(fixture.model, TOOLS_JSON)?.parse(QWEN_XML_PAYLOAD, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!("expected Recognized for Qwen XML on a Qwen-3.5 model; got Unrecognized"); @@ -1882,11 +1839,8 @@ 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( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - TWO_QWEN_XML_PAYLOADS, - false, - )?; + let outcome = + ChatMessageParser::new(fixture.model, TOOLS_JSON)?.parse(TWO_QWEN_XML_PAYLOADS, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -1987,11 +1938,7 @@ 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( - &ChatTools::from_json(TOOLS_JSON.to_owned())?, - PLAIN_CONTENT, - false, - )?; + let outcome = ChatMessageParser::new(fixture.model, TOOLS_JSON)?.parse(PLAIN_CONTENT, false)?; let ChatMessageParseOutcome::Recognized(parsed) = outcome else { bail!( @@ -2035,7 +1982,7 @@ fn qwen36_chat_inference_emits_reasoning_when_template_auto_opens( )?]; let prompt = model.apply_chat_template(&chat_template, &messages, true, true)?; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let tokens = model.str_to_token(&prompt, AddBos::Always)?; let prompt_token_count = u64::try_from(tokens.len())?; @@ -2050,7 +1997,6 @@ fn qwen36_chat_inference_emits_reasoning_when_template_auto_opens( let mut sampler = LlamaSampler::greedy()?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -2105,7 +2051,7 @@ What is 2 + 2?<|im_end|> let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(QWEN36_THINKING_DISABLED_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -2133,7 +2079,6 @@ What is 2 + 2?<|im_end|> ])?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -2185,7 +2130,7 @@ What is 2 + 2?<|im_end|> let model = fixture.model; let backend = fixture.backend; - let mut classifier = model.sampled_token_classifier()?; + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let prompt_tokens = model.str_to_token(QWEN36_THINKING_PROMPT, AddBos::Never)?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; @@ -2213,7 +2158,6 @@ What is 2 + 2?<|im_end|> ])?; let initial_position = batch.n_tokens(); let outcome = ClassifySampleLoop { - model, classifier: &mut classifier, sampler: &mut sampler, context: &mut context, @@ -2247,3 +2191,70 @@ What is 2 + 2?<|im_end|> Ok(()) } + +fn visible_text_of_generation_ending_after( + model: &LlamaModel, + generated_text: &str, +) -> Result { + let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; + let mut outcomes = Vec::new(); + + for token in model.str_to_token(generated_text, AddBos::Never)? { + assert_eq!( + classifier.ingest(token, &mut outcomes)?, + GenerationProgress::Continues + ); + } + + assert_eq!( + classifier.ingest(model.token_eos(), &mut outcomes)?, + GenerationProgress::Ended + ); + + Ok(outcomes + .iter() + .map(|outcome| outcome.piece.visible()) + .collect()) +} + +#[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 = 256, + n_batch = 128, + n_ubatch = 64, +)] +fn qwen35_classifier_ends_generation_without_emitting_the_end_of_generation_token( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + const GENERATED_TEXT: &str = "The answer is four."; + + assert_eq!( + visible_text_of_generation_ending_after(fixture.model, GENERATED_TEXT)?, + GENERATED_TEXT + ); + + 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 = 256, + n_batch = 128, + n_ubatch = 64, +)] +fn qwen35_classifier_releases_a_held_json_prefix_when_generation_ends( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + const GENERATED_TEXT: &str = r#"{"answer": 4"#; + + assert_eq!( + visible_text_of_generation_ending_after(fixture.model, GENERATED_TEXT)?, + GENERATED_TEXT + ); + + Ok(()) +} diff --git a/llama-cpp-bindings/Cargo.toml b/llama-cpp-bindings/Cargo.toml index a00b03f50..91d4eeb11 100644 --- a/llama-cpp-bindings/Cargo.toml +++ b/llama-cpp-bindings/Cargo.toml @@ -19,6 +19,7 @@ llguidance = { workspace = true } log = { workspace = true } nom = { workspace = true } once_cell = { workspace = true } +serde = { workspace = true } serde_json = { workspace = true } thiserror = { workspace = true } toktrie = { workspace = true } diff --git a/llama-cpp-bindings/src/bare_json_tool_calls.rs b/llama-cpp-bindings/src/bare_json_tool_calls.rs new file mode 100644 index 000000000..853c839e6 --- /dev/null +++ b/llama-cpp-bindings/src/bare_json_tool_calls.rs @@ -0,0 +1,5 @@ +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum BareJsonToolCalls { + Detect, + Ignore, +} diff --git a/llama-cpp-bindings/src/chat_message_parser.rs b/llama-cpp-bindings/src/chat_message_parser.rs new file mode 100644 index 000000000..69f2f073d --- /dev/null +++ b/llama-cpp-bindings/src/chat_message_parser.rs @@ -0,0 +1,2591 @@ +use std::ffi::CStr; +use std::ffi::CString; +use std::ffi::c_char; +use std::ptr; +use std::ptr::NonNull; + +use llama_cpp_bindings_types::ParsedChatMessage; +use llama_cpp_bindings_types::ParsedToolCall; +use llama_cpp_bindings_types::ReasoningMarkers; +use llama_cpp_bindings_types::ToolCallArguments; +use llama_cpp_ffi_status::read_and_free_cpp_string; + +use crate::chat_message_parse_outcome::ChatMessageParseOutcome; +use crate::chat_template_tool_calls; +use crate::error::parse_chat_message_error::ParseChatMessageError; +use crate::model::LlamaModel; +use crate::raw_chat_message::RawChatMessage; +use crate::tool_call_format; +use crate::tool_call_format::ToolCallFormatOutcome; + +/// # Safety +/// +/// `free_error` must be the pointer populated by the preceding +/// `llama_rs_parsed_chat_free` call, or null. The destructor-threw arm reads and +/// frees it. +unsafe fn parsed_chat_free_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_parsed_chat_free_status, + free_error: *mut c_char, +) -> Result<(), ParseChatMessageError> { + match status { + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_OK => Ok(()), + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_ERROR_STRING_ALLOCATION_FAILED => { + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_LLAMA_CPP_OUT_OF_MEMORY => { + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION => { + let message = unsafe { + read_and_free_cpp_string( + free_error, + "llama_rs_parsed_chat_free", + "reported a thrown C++ exception without an error message", + ) + }?; + + Err(ParseChatMessageError::DestructorFailed { message }) + } + other => Err(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_free", + code: i64::from(other), + } + .into()), + } +} + +/// # 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 +/// 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, + handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, + out_error: *mut *mut c_char, +) -> Result { + match status { + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_OK => { + if handle.is_null() { + Err(crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "success status contained a null parsed-chat handle", + } + .into()) + } else { + collect_parsed_chat_message(handle) + } + } + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_ERROR_STRING_ALLOCATION_FAILED => { + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY => { + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION => { + Err(unsafe { + thrown_parse_exception_error(out_error, |message| { + ParseChatMessageError::MessageUnrecognized { message } + }) + }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_TOOLS_PARSER_ARG => { + Err(crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "was given a null tools parser argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_INPUT_ARG => { + Err(crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "was given a null input argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_HANDLE_ARG => { + Err(crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "was given a null out_handle argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_ERROR_ARG => { + Err(crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "was given a null out_error argument", + } + .into()) + } + other => Err(crate::FfiStatusError { + operation: "llama_rs_parse_chat_message", + code: i64::from(other), + } + .into()), + } +} + +fn outcome_from_via_ffi_result( + via_ffi_result: Result, + input: &str, + is_partial: bool, +) -> Result { + match via_ffi_result { + Ok(mut parsed) => { + synthesize_missing_tool_call_ids(&mut parsed.tool_calls); + Ok(ChatMessageParseOutcome::Recognized(parsed)) + } + Err(ParseChatMessageError::MessageUnrecognized { message }) => { + Ok(ChatMessageParseOutcome::Unrecognized(RawChatMessage { + text: input.to_owned(), + is_partial, + ffi_error_message: message, + })) + } + Err(other) => Err(other), + } +} + +fn collect_parsed_chat_message( + handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, +) -> Result { + if handle.is_null() { + return Ok(ParsedChatMessage::default()); + } + + let content = read_parsed_chat_content(handle)?; + let reasoning_content = read_parsed_chat_reasoning_content(handle)?; + let count = read_parsed_chat_tool_call_count(handle)?; + + let mut tool_calls = Vec::with_capacity(count); + for index in 0..count { + let id = read_parsed_chat_tool_call_id(handle, index)?; + let name = read_parsed_chat_tool_call_name(handle, index)?; + let arguments_json = read_parsed_chat_tool_call_arguments(handle, index)?; + + let arguments = ToolCallArguments::from_string(arguments_json); + tool_calls.push(ParsedToolCall::new(id, name, arguments)); + } + + Ok(ParsedChatMessage::new( + content, + reasoning_content, + tool_calls, + )) +} + +/// # Safety +/// +/// `out_string` and `out_error` must be the pointers populated by the preceding +/// `llama_rs_parsed_chat_content` call (or null when no value/error was produced); each is +/// read and freed in exactly one match arm. +unsafe fn parsed_chat_content_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_parsed_chat_content_status, + out_string: *mut c_char, + out_error: *mut c_char, +) -> Result { + match status { + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_OK => { + consume_accessor_string(out_string, "llama_rs_parsed_chat_content") + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_ERROR_STRING_ALLOCATION_FAILED => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_OUT_OF_MEMORY => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = unsafe { + read_and_free_cpp_string( + out_error, + "llama_rs_parsed_chat_content", + "reported a thrown C++ exception without an error message", + ) + }?; + Err(ParseChatMessageError::Reported { message }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_NULL_HANDLE_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_content", + detail: "was given a null handle argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_NULL_OUT_STRING_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_content", + detail: "was given a null out_string argument", + } + .into()) + } + other => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_content", + code: i64::from(other), + } + .into()) + } + } +} + +fn read_parsed_chat_content( + handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, +) -> Result { + let mut out_string: *mut c_char = ptr::null_mut(); + let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_parsed_chat_content( + handle, + &raw mut out_string, + &raw mut out_error, + ) + }; + unsafe { parsed_chat_content_status_to_result(status, out_string, out_error) } +} + +/// # Safety +/// +/// `out_string` and `out_error` must be the pointers populated by the preceding +/// `llama_rs_parsed_chat_reasoning_content` call (or null when no value/error was produced); +/// each is read and freed in exactly one match arm. +unsafe fn parsed_chat_reasoning_content_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_parsed_chat_reasoning_content_status, + out_string: *mut c_char, + out_error: *mut c_char, +) -> Result { + match status { + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_OK => { + consume_accessor_string(out_string, "llama_rs_parsed_chat_reasoning_content") + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_ERROR_STRING_ALLOCATION_FAILED => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_OUT_OF_MEMORY => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = + unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_reasoning_content", "reported a thrown C++ exception without an error message") }?; + Err(ParseChatMessageError::Reported { message }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_HANDLE_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_reasoning_content", + detail: "was given a null handle argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_OUT_STRING_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_reasoning_content", + detail: "was given a null out_string argument", + } + .into()) + } + other => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_reasoning_content", + code: i64::from(other), + } + .into()) + } + } +} + +fn read_parsed_chat_reasoning_content( + handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, +) -> Result { + let mut out_string: *mut c_char = ptr::null_mut(); + let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_parsed_chat_reasoning_content( + handle, + &raw mut out_string, + &raw mut out_error, + ) + }; + unsafe { parsed_chat_reasoning_content_status_to_result(status, out_string, out_error) } +} + +/// # Safety +/// +/// `out_error` must be the pointer populated by the preceding +/// `llama_rs_parsed_chat_tool_call_count` call (or null when no error was produced); it is +/// freed in exactly one match arm. +unsafe fn parsed_chat_tool_call_count_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_count_status, + out_count: usize, + out_error: *mut c_char, +) -> Result { + match status { + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_OK => Ok(out_count), + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_ERROR_STRING_ALLOCATION_FAILED => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_OUT_OF_MEMORY => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = + unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_tool_call_count", "reported a thrown C++ exception without an error message") }?; + Err(ParseChatMessageError::Reported { message }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_HANDLE_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_count", + detail: "was given a null handle argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_OUT_COUNT_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_count", + detail: "was given a null out_count argument", + } + .into()) + } + other => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_tool_call_count", + code: i64::from(other), + } + .into()) + } + } +} + +fn read_parsed_chat_tool_call_count( + handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, +) -> Result { + let mut out_count: usize = 0; + let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_count( + handle, + &raw mut out_count, + &raw mut out_error, + ) + }; + unsafe { parsed_chat_tool_call_count_status_to_result(status, out_count, out_error) } +} + +/// # Safety +/// +/// `out_string` and `out_error` must be the pointers populated by the preceding +/// `llama_rs_parsed_chat_tool_call_id` call (or null when no value/error was produced); each +/// is read and freed in exactly one match arm. +unsafe fn parsed_chat_tool_call_id_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_id_status, + index: usize, + out_string: *mut c_char, + out_error: *mut c_char, +) -> Result { + match status { + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_OK => { + consume_accessor_string(out_string, "llama_rs_parsed_chat_tool_call_id") + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_INDEX_OUT_OF_BOUNDS => { + Err(ParseChatMessageError::ToolCallIdIndexOutOfBounds { index }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_ERROR_STRING_ALLOCATION_FAILED => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_OUT_OF_MEMORY => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = + unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_tool_call_id", "reported a thrown C++ exception without an error message") }?; + Err(ParseChatMessageError::Reported { message }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_HANDLE_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_id", + detail: "was given a null handle argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_OUT_STRING_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_id", + detail: "was given a null out_string argument", + } + .into()) + } + other => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_tool_call_id", + code: i64::from(other), + } + .into()) + } + } +} + +fn read_parsed_chat_tool_call_id( + handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, + index: usize, +) -> Result { + let mut out_string: *mut c_char = ptr::null_mut(); + let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_id( + handle, + index, + &raw mut out_string, + &raw mut out_error, + ) + }; + unsafe { parsed_chat_tool_call_id_status_to_result(status, index, out_string, out_error) } +} + +/// # Safety +/// +/// `out_string` and `out_error` must be the pointers populated by the preceding +/// `llama_rs_parsed_chat_tool_call_name` call (or null when no value/error was produced); each +/// is read and freed in exactly one match arm. +unsafe fn parsed_chat_tool_call_name_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_name_status, + index: usize, + out_string: *mut c_char, + out_error: *mut c_char, +) -> Result { + match status { + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_OK => { + consume_accessor_string(out_string, "llama_rs_parsed_chat_tool_call_name") + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_INDEX_OUT_OF_BOUNDS => { + Err(ParseChatMessageError::ToolCallNameIndexOutOfBounds { index }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_ERROR_STRING_ALLOCATION_FAILED => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_OUT_OF_MEMORY => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = + unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_tool_call_name", "reported a thrown C++ exception without an error message") }?; + Err(ParseChatMessageError::Reported { message }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_HANDLE_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_name", + detail: "was given a null handle argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_OUT_STRING_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_name", + detail: "was given a null out_string argument", + } + .into()) + } + other => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_tool_call_name", + code: i64::from(other), + } + .into()) + } + } +} + +fn read_parsed_chat_tool_call_name( + handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, + index: usize, +) -> Result { + let mut out_string: *mut c_char = ptr::null_mut(); + let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_name( + handle, + index, + &raw mut out_string, + &raw mut out_error, + ) + }; + unsafe { parsed_chat_tool_call_name_status_to_result(status, index, out_string, out_error) } +} + +/// # Safety +/// +/// `out_string` and `out_error` must be the pointers populated by the preceding +/// `llama_rs_parsed_chat_tool_call_arguments` call (or null when no value/error was produced); +/// each is read and freed in exactly one match arm. +unsafe fn parsed_chat_tool_call_arguments_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_arguments_status, + index: usize, + out_string: *mut c_char, + out_error: *mut c_char, +) -> Result { + match status { + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_OK => { + consume_accessor_string(out_string, "llama_rs_parsed_chat_tool_call_arguments") + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_INDEX_OUT_OF_BOUNDS => { + Err(ParseChatMessageError::ToolCallArgumentsIndexOutOfBounds { index }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_ERROR_STRING_ALLOCATION_FAILED => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_OUT_OF_MEMORY => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = + unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_tool_call_arguments", "reported a thrown C++ exception without an error message") }?; + Err(ParseChatMessageError::Reported { message }) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_HANDLE_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_arguments", + detail: "was given a null handle argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_OUT_STRING_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_arguments", + detail: "was given a null out_string argument", + } + .into()) + } + other => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_tool_call_arguments", + code: i64::from(other), + } + .into()) + } + } +} + +fn read_parsed_chat_tool_call_arguments( + handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, + index: usize, +) -> Result { + let mut out_string: *mut c_char = ptr::null_mut(); + let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_arguments( + handle, + index, + &raw mut out_string, + &raw mut out_error, + ) + }; + unsafe { + parsed_chat_tool_call_arguments_status_to_result(status, index, out_string, out_error) + } +} + +fn consume_accessor_string( + ptr: *mut c_char, + operation: &'static str, +) -> Result { + if ptr.is_null() { + return Err(crate::FfiContractError { + operation, + detail: "success status contained a null string", + } + .into()); + } + let bytes = unsafe { CStr::from_ptr(ptr) }.to_bytes().to_vec(); + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(ptr) }; + Ok(String::from_utf8(bytes)?) +} + +struct ReasoningSplit { + reasoning: String, + content: String, +} + +fn restore_partial_reasoning( + parsed: &mut ParsedChatMessage, + input: &str, + reasoning_markers: Option<&ReasoningMarkers>, + is_partial: bool, +) { + if !is_partial { + return; + } + if reasoning_markers.is_some_and(|markers| input.contains(&markers.open)) { + let split = split_reasoning_prefix(input, reasoning_markers, None, true); + parsed.reasoning_content = split.reasoning; + parsed.content = split.content; + return; + } + if let Some(open) = reasoning_markers.map(|markers| markers.open.trim()) + && let Some(reasoning) = parsed.reasoning_content.trim_start().strip_prefix(open) + { + parsed.reasoning_content = reasoning.to_owned(); + } +} + +fn split_reasoning_prefix( + input: &str, + reasoning_markers: Option<&ReasoningMarkers>, + tool_call_open: Option<&str>, + is_partial: bool, +) -> ReasoningSplit { + let content_only = || ReasoningSplit { + reasoning: String::new(), + content: prefix_before_optional(input, tool_call_open), + }; + + let Some(reasoning_markers) = reasoning_markers else { + return content_only(); + }; + let Some(open_pos) = input.find(&reasoning_markers.open) else { + return content_only(); + }; + + let after_open = &input[open_pos + reasoning_markers.open.len()..]; + let closing_marker = reasoning_markers + .closes + .iter() + .enumerate() + .filter_map(|(marker_index, marker)| { + after_open + .find(marker) + .map(|offset| (offset, marker_index, marker)) + }) + .min_by_key(|(offset, marker_index, _)| (*offset, *marker_index)); + let Some((close_offset, _, close_marker)) = closing_marker else { + return if is_partial { + ReasoningSplit { + reasoning: prefix_before_optional(after_open, tool_call_open), + content: input[..open_pos].to_owned(), + } + } else { + content_only() + }; + }; + + let reasoning = after_open[..close_offset].to_owned(); + let after_close = &after_open[close_offset + close_marker.len()..]; + + ReasoningSplit { + reasoning, + content: prefix_before_optional(after_close, tool_call_open), + } +} + +fn prefix_before_optional(text: &str, marker: Option<&str>) -> String { + marker.map_or_else( + || text.to_owned(), + |marker| { + text.find(marker) + .map_or_else(|| text.to_owned(), |pos| text[..pos].to_owned()) + }, + ) +} + +fn synthesize_missing_tool_call_ids(tool_calls: &mut [ParsedToolCall]) { + for (index, call) in tool_calls.iter_mut().enumerate() { + if call.id.is_empty() { + call.id = format!("call_{index}"); + } + } +} + +/// # Safety +/// +/// `out_error` must reference the pointer populated by the preceding +/// `llama_rs_chat_tools_parser_create` call (or null); it is read, freed, and nulled only in +/// the CXX-exception arm. `tools_parser` must be the pointer populated by the same call. +unsafe fn chat_tools_parser_create_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_chat_tools_parser_create_status, + tools_parser: *mut llama_cpp_bindings_sys::llama_rs_chat_tools_parser, + out_error: *mut *mut c_char, +) -> Result, ParseChatMessageError> { + match status { + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_OK => NonNull::new(tools_parser) + .ok_or_else(|| { + crate::FfiContractError { + operation: "llama_rs_chat_tools_parser_create", + detail: "success status contained a null tools parser handle", + } + .into() + }), + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_TOOLS_NOT_AN_ARRAY => { + Err(ParseChatMessageError::ToolsNotAnArray) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_ERROR_STRING_ALLOCATION_FAILED => { + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_OUT_OF_MEMORY => { + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = unsafe { + read_and_free_cpp_string( + *out_error, + "llama_rs_chat_tools_parser_create", + "reported a thrown C++ exception without an error message", + ) + }?; + + unsafe { *out_error = ptr::null_mut() }; + + Err(ParseChatMessageError::ToolsParserBuildFailed { message }) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_PARSER_ARG => { + Err(crate::FfiContractError { + operation: "llama_rs_chat_tools_parser_create", + detail: "was given a null parser argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_TOOLS_JSON_ARG => { + Err(crate::FfiContractError { + operation: "llama_rs_chat_tools_parser_create", + detail: "was given a null tools_json argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_OUT_TOOLS_PARSER_ARG => { + Err(crate::FfiContractError { + operation: "llama_rs_chat_tools_parser_create", + detail: "was given a null out_tools_parser argument", + } + .into()) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_OUT_ERROR_ARG => { + Err(crate::FfiContractError { + operation: "llama_rs_chat_tools_parser_create", + detail: "was given a null out_error argument", + } + .into()) + } + other => Err(crate::FfiStatusError { + operation: "llama_rs_chat_tools_parser_create", + code: i64::from(other), + } + .into()), + } +} + +/// # Safety +/// +/// `out_error` must be the pointer populated by the preceding +/// `llama_rs_chat_tools_parser_free` call, or null. The destructor-threw arm reads and frees it. +unsafe fn chat_tools_parser_free_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_chat_tools_parser_free_status, + out_error: *mut c_char, +) -> Result<(), ParseChatMessageError> { + match status { + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_FREE_OK => Ok(()), + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_FREE_ERROR_STRING_ALLOCATION_FAILED => { + Err(ParseChatMessageError::NotEnoughMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_FREE_LLAMA_CPP_OUT_OF_MEMORY => { + Err(ParseChatMessageError::LlamaCppOutOfMemory) + } + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION => { + let message = unsafe { + read_and_free_cpp_string( + out_error, + "llama_rs_chat_tools_parser_free", + "reported a thrown C++ exception without an error message", + ) + }?; + + Err(ParseChatMessageError::DestructorFailed { message }) + } + other => Err(crate::FfiStatusError { + operation: "llama_rs_chat_tools_parser_free", + code: i64::from(other), + } + .into()), + } +} + +/// Parses generated chat messages against one set of tools, building the +/// model-specific tools parser once for every message it parses. +pub struct ChatMessageParser { + reasoning_markers: Option, + tools_parser: NonNull, +} + +unsafe impl Send for ChatMessageParser {} + +impl ChatMessageParser { + /// # Errors + /// + /// Returns [`ParseChatMessageError`] when reasoning-marker detection fails, the model has + /// no chat parser, `tools_json` contains a NUL byte or is not a JSON array, or the tools + /// parser cannot be built. + pub fn new(model: &LlamaModel, tools_json: &str) -> Result { + let reasoning_markers = model.reasoning_markers()?.cloned(); + let chat_parser = model.chat_parser_ptr()?; + let tools_json_cstring = + CString::new(tools_json).map_err(ParseChatMessageError::ToolsContainNulByte)?; + let mut out_tools_parser: *mut llama_cpp_bindings_sys::llama_rs_chat_tools_parser = + ptr::null_mut(); + let mut out_error: *mut c_char = ptr::null_mut(); + + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_chat_tools_parser_create( + chat_parser, + tools_json_cstring.as_ptr(), + &raw mut out_tools_parser, + &raw mut out_error, + ) + }; + + let tools_parser = unsafe { + chat_tools_parser_create_status_to_result(status, out_tools_parser, &raw mut out_error) + }; + + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + + Ok(Self { + reasoning_markers, + tools_parser: tools_parser?, + }) + } + + /// # Errors + /// + /// Returns [`ParseChatMessageError`] when `input` contains a NUL byte, the FFI returns a + /// non-OK status other than a message parse exception, or accessor strings are not valid + /// UTF-8. + pub fn parse( + &self, + input: &str, + is_partial: bool, + ) -> Result { + let reasoning_markers = self.reasoning_markers.as_ref(); + + for candidate in chat_template_tool_calls::known_marker_candidates() { + match tool_call_format::try_parse(input, &candidate) { + ToolCallFormatOutcome::NoMatch => {} + ToolCallFormatOutcome::Parsed(calls) => { + let split = split_reasoning_prefix( + input, + reasoning_markers, + Some(&candidate.open), + is_partial, + ); + let mut parsed = ParsedChatMessage::new(split.content, split.reasoning, calls); + synthesize_missing_tool_call_ids(&mut parsed.tool_calls); + + return Ok(ChatMessageParseOutcome::Recognized(parsed)); + } + ToolCallFormatOutcome::Failed(_shape_does_not_fit) => {} + } + } + + let via_ffi_result = self.parse_via_ffi(input, is_partial).map(|mut parsed| { + restore_partial_reasoning(&mut parsed, input, reasoning_markers, is_partial); + parsed + }); + + outcome_from_via_ffi_result(via_ffi_result, input, is_partial) + } + + fn parse_via_ffi( + &self, + input: &str, + is_partial: bool, + ) -> Result { + let input_cstring = + CString::new(input).map_err(ParseChatMessageError::InputContainsNulByte)?; + + let mut handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat = ptr::null_mut(); + let mut out_error: *mut c_char = ptr::null_mut(); + + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_parse_chat_message( + self.tools_parser.as_ptr(), + input_cstring.as_ptr(), + i32::from(is_partial), + &raw mut handle, + &raw mut out_error, + ) + }; + + let parsed = + unsafe { parse_chat_message_status_to_result(status, handle, &raw mut out_error) }; + + let mut free_error: *mut c_char = ptr::null_mut(); + let free_status = unsafe { + llama_cpp_bindings_sys::llama_rs_parsed_chat_free(handle, &raw mut free_error) + }; + let freed = unsafe { parsed_chat_free_status_to_result(free_status, free_error) }; + + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + + match parsed { + Ok(message) => freed.map(|()| message), + Err(parse_failure) => { + if let Err(destructor_failure) = freed { + log::error!("{destructor_failure}"); + } + + Err(parse_failure) + } + } + } +} + +impl Drop for ChatMessageParser { + fn drop(&mut self) { + let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { + llama_cpp_bindings_sys::llama_rs_chat_tools_parser_free( + self.tools_parser.as_ptr(), + &raw mut out_error, + ) + }; + + if let Err(destructor_failure) = + unsafe { chat_tools_parser_free_status_to_result(status, out_error) } + { + log::error!("{destructor_failure}"); + } + } +} + +#[cfg(test)] +mod tests { + use std::ffi::CStr; + use std::ffi::c_char; + use std::mem::discriminant; + use std::ptr; + + use llama_cpp_bindings_types::ParsedChatMessage; + use llama_cpp_bindings_types::ParsedToolCall; + use llama_cpp_bindings_types::ReasoningMarkers; + use llama_cpp_bindings_types::ToolCallArguments; + + use super::ReasoningSplit; + use super::chat_tools_parser_create_status_to_result; + use super::chat_tools_parser_free_status_to_result; + use super::outcome_from_via_ffi_result; + use super::parse_chat_message_status_to_result; + use super::parsed_chat_content_status_to_result; + use super::parsed_chat_free_status_to_result; + use super::parsed_chat_reasoning_content_status_to_result; + use super::parsed_chat_tool_call_arguments_status_to_result; + use super::parsed_chat_tool_call_count_status_to_result; + use super::parsed_chat_tool_call_id_status_to_result; + use super::parsed_chat_tool_call_name_status_to_result; + use super::restore_partial_reasoning; + use super::split_reasoning_prefix; + use crate::chat_message_parse_outcome::ChatMessageParseOutcome; + use crate::error::parse_chat_message_error::ParseChatMessageError; + use crate::raw_chat_message::RawChatMessage; + #[test] + fn parse_chat_message_success_with_null_handle_is_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_OK, + ptr::null_mut(), + &raw mut out_error, + ) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiContract( + crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "success status contained a null parsed-chat handle", + } + )) + )); + } + + #[test] + fn parse_chat_message_allocation_failed_is_not_enough_memory() { + 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_ERROR_STRING_ALLOCATION_FAILED, + ptr::null_mut(), + &raw mut out_error, + ) + }; + + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NotEnoughMemory) + ); + } + + #[test] + fn parse_chat_message_cxx_exception_is_message_unrecognized_and_nulls_error() { + let mut out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"the message could not be parsed".as_ptr()) + }; + let result = unsafe { + parse_chat_message_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION, + ptr::null_mut(), + &raw mut out_error, + ) + }; + + let Err(ParseChatMessageError::MessageUnrecognized { message }) = result else { + panic!("the llama.cpp exception status must surface the wrapper message"); + }; + + assert_eq!(message, "the message could not be parsed"); + 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_LLAMA_CPP_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(); + let result = unsafe { + parse_chat_message_status_to_result(255, ptr::null_mut(), &raw mut out_error) + }; + + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_parse_chat_message", + code: 255, + })) + ); + } + + #[test] + fn parsed_chat_content_success_with_null_string_is_contract_error() { + let result = unsafe { + parsed_chat_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_OK, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiContract( + crate::FfiContractError { + operation: "llama_rs_parsed_chat_content", + detail: "success status contained a null string", + } + )) + )); + } + + #[test] + fn parsed_chat_content_allocation_failed_is_not_enough_memory() { + let result = unsafe { + parsed_chat_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_ERROR_STRING_ALLOCATION_FAILED, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NotEnoughMemory) + ); + } + + #[test] + fn parsed_chat_content_cxx_exception_is_reported() { + let out_error = + unsafe { llama_cpp_bindings_sys::llama_rs_string_dup(c"content read failed".as_ptr()) }; + let result = unsafe { + parsed_chat_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_THREW_CXX_EXCEPTION, + ptr::null_mut(), + out_error, + ) + }; + + let Err(ParseChatMessageError::Reported { message }) = result else { + panic!("the llama.cpp exception status must surface the wrapper message"); + }; + + assert_eq!(message, "content read failed"); + } + + #[test] + fn parsed_chat_content_unknown_status_is_preserved() { + let result = + unsafe { parsed_chat_content_status_to_result(255, ptr::null_mut(), ptr::null_mut()) }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_content", + code: 255, + })) + )); + } + + #[test] + fn parsed_chat_reasoning_content_success_with_null_string_is_contract_error() { + let result = unsafe { + parsed_chat_reasoning_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_OK, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiContract( + crate::FfiContractError { + operation: "llama_rs_parsed_chat_reasoning_content", + detail: "success status contained a null string", + } + )) + )); + } + + #[test] + fn parsed_chat_reasoning_content_allocation_failed_is_not_enough_memory() { + let result = unsafe { + parsed_chat_reasoning_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_ERROR_STRING_ALLOCATION_FAILED, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NotEnoughMemory) + ); + } + + #[test] + fn parsed_chat_reasoning_content_cxx_exception_is_reported() { + let out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"reasoning read failed".as_ptr()) + }; + let result = unsafe { + parsed_chat_reasoning_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_THREW_CXX_EXCEPTION, + ptr::null_mut(), + out_error, + ) + }; + + let Err(ParseChatMessageError::Reported { message }) = result else { + panic!("the llama.cpp exception status must surface the wrapper message"); + }; + + assert_eq!(message, "reasoning read failed"); + } + + #[test] + fn parsed_chat_reasoning_content_unknown_status_is_preserved() { + let result = unsafe { + parsed_chat_reasoning_content_status_to_result(255, ptr::null_mut(), ptr::null_mut()) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_reasoning_content", + code: 255, + })) + )); + } + + #[test] + fn parsed_chat_tool_call_count_ok_returns_count() { + let result = unsafe { + parsed_chat_tool_call_count_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_OK, + 7, + ptr::null_mut(), + ) + }; + + assert_eq!(result.unwrap(), 7); + } + + #[test] + fn parsed_chat_tool_call_count_allocation_failed_is_not_enough_memory() { + let result = unsafe { + parsed_chat_tool_call_count_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_ERROR_STRING_ALLOCATION_FAILED, + 0, + ptr::null_mut(), + ) + }; + + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NotEnoughMemory) + ); + } + + #[test] + fn parsed_chat_tool_call_count_cxx_exception_is_reported() { + let out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"tool-call count failed".as_ptr()) + }; + let result = unsafe { + parsed_chat_tool_call_count_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_THREW_CXX_EXCEPTION, + 0, + out_error, + ) + }; + + let Err(ParseChatMessageError::Reported { message }) = result else { + panic!("the llama.cpp exception status must surface the wrapper message"); + }; + + assert_eq!(message, "tool-call count failed"); + } + + #[test] + fn parsed_chat_tool_call_count_unknown_status_is_preserved() { + let result = + unsafe { parsed_chat_tool_call_count_status_to_result(255, 0, ptr::null_mut()) }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_tool_call_count", + code: 255, + })) + )); + } + + #[test] + fn parsed_chat_tool_call_id_success_with_null_string_is_contract_error() { + let result = unsafe { + parsed_chat_tool_call_id_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_OK, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiContract( + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_id", + detail: "success status contained a null string", + } + )) + )); + } + + #[test] + fn parsed_chat_tool_call_id_out_of_bounds_carries_index() { + let result = unsafe { + parsed_chat_tool_call_id_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_INDEX_OUT_OF_BOUNDS, + 4, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + let Err(ParseChatMessageError::ToolCallIdIndexOutOfBounds { index }) = result else { + panic!("expected ToolCallIdIndexOutOfBounds, got {result:?}"); + }; + assert_eq!(index, 4); + } + + #[test] + fn parsed_chat_tool_call_id_allocation_failed_is_not_enough_memory() { + let result = unsafe { + parsed_chat_tool_call_id_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_ERROR_STRING_ALLOCATION_FAILED, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NotEnoughMemory) + ); + } + + #[test] + fn parsed_chat_tool_call_id_cxx_exception_is_reported() { + let out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"tool-call id read failed".as_ptr()) + }; + let result = unsafe { + parsed_chat_tool_call_id_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_THREW_CXX_EXCEPTION, + 0, + ptr::null_mut(), + out_error, + ) + }; + + let Err(ParseChatMessageError::Reported { message }) = result else { + panic!("the llama.cpp exception status must surface the wrapper message"); + }; + + assert_eq!(message, "tool-call id read failed"); + } + + #[test] + fn parsed_chat_tool_call_id_unknown_status_is_preserved() { + let result = unsafe { + parsed_chat_tool_call_id_status_to_result(255, 0, ptr::null_mut(), ptr::null_mut()) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_tool_call_id", + code: 255, + })) + )); + } + + #[test] + fn parsed_chat_tool_call_name_success_with_null_string_is_contract_error() { + let result = unsafe { + parsed_chat_tool_call_name_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_OK, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiContract( + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_name", + detail: "success status contained a null string", + } + )) + )); + } + + #[test] + fn parsed_chat_tool_call_name_out_of_bounds_carries_index() { + let result = unsafe { + parsed_chat_tool_call_name_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_INDEX_OUT_OF_BOUNDS, + 2, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + let Err(ParseChatMessageError::ToolCallNameIndexOutOfBounds { index }) = result else { + panic!("expected ToolCallNameIndexOutOfBounds, got {result:?}"); + }; + assert_eq!(index, 2); + } + + #[test] + fn parsed_chat_tool_call_name_allocation_failed_is_not_enough_memory() { + let result = unsafe { + parsed_chat_tool_call_name_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_ERROR_STRING_ALLOCATION_FAILED, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NotEnoughMemory) + ); + } + + #[test] + fn parsed_chat_tool_call_name_cxx_exception_is_reported() { + let out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"tool-call name read failed".as_ptr()) + }; + let result = unsafe { + parsed_chat_tool_call_name_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_THREW_CXX_EXCEPTION, + 0, + ptr::null_mut(), + out_error, + ) + }; + + let Err(ParseChatMessageError::Reported { message }) = result else { + panic!("the llama.cpp exception status must surface the wrapper message"); + }; + + assert_eq!(message, "tool-call name read failed"); + } + + #[test] + fn parsed_chat_tool_call_name_unknown_status_is_preserved() { + let result = unsafe { + parsed_chat_tool_call_name_status_to_result(255, 0, ptr::null_mut(), ptr::null_mut()) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_tool_call_name", + code: 255, + })) + )); + } + + #[test] + fn parsed_chat_tool_call_arguments_success_with_null_string_is_contract_error() { + let result = unsafe { + parsed_chat_tool_call_arguments_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_OK, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiContract( + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_arguments", + detail: "success status contained a null string", + } + )) + )); + } + + #[test] + fn parsed_chat_tool_call_arguments_out_of_bounds_carries_index() { + let result = unsafe { + parsed_chat_tool_call_arguments_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_INDEX_OUT_OF_BOUNDS, + 9, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + let Err(ParseChatMessageError::ToolCallArgumentsIndexOutOfBounds { index }) = result else { + panic!("expected ToolCallArgumentsIndexOutOfBounds, got {result:?}"); + }; + assert_eq!(index, 9); + } + + #[test] + fn parsed_chat_tool_call_arguments_allocation_failed_is_not_enough_memory() { + let result = unsafe { + parsed_chat_tool_call_arguments_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_ERROR_STRING_ALLOCATION_FAILED, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NotEnoughMemory) + ); + } + + #[test] + fn parsed_chat_tool_call_arguments_cxx_exception_is_reported() { + let out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"tool-call arguments read failed".as_ptr()) + }; + let result = unsafe { + parsed_chat_tool_call_arguments_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_THREW_CXX_EXCEPTION, + 0, + ptr::null_mut(), + out_error, + ) + }; + + let Err(ParseChatMessageError::Reported { message }) = result else { + panic!("the llama.cpp exception status must surface the wrapper message"); + }; + + assert_eq!(message, "tool-call arguments read failed"); + } + + #[test] + fn parsed_chat_tool_call_arguments_unknown_status_is_preserved() { + let result = unsafe { + parsed_chat_tool_call_arguments_status_to_result( + 255, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + + assert!(matches!( + result, + Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_parsed_chat_tool_call_arguments", + code: 255, + })) + )); + } + + #[test] + fn split_reasoning_prefix_without_markers_returns_content_up_to_tool_call_open() { + let ReasoningSplit { reasoning, content } = + split_reasoning_prefix("answerrest", None, Some(""), false); + + assert!(reasoning.is_empty()); + assert_eq!(content, "answer"); + } + + #[test] + fn split_reasoning_prefix_with_missing_open_marker_returns_content_only() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let ReasoningSplit { reasoning, content } = + split_reasoning_prefix("plain answer", Some(&markers), Some(""), false); + + assert!(reasoning.is_empty()); + assert_eq!(content, "plain answer"); + } + + #[test] + fn split_reasoning_prefix_with_missing_close_marker_returns_content_only() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let ReasoningSplit { reasoning, content } = + split_reasoning_prefix("unterminated", Some(&markers), Some(""), false); + + assert!(reasoning.is_empty()); + assert_eq!(content, "unterminated"); + } + + #[test] + fn split_reasoning_prefix_with_partial_unclosed_marker_returns_reasoning() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let ReasoningSplit { reasoning, content } = split_reasoning_prefix( + "prefixunfinishedtail", + Some(&markers), + Some(""), + true, + ); + + assert_eq!(reasoning, "unfinished"); + assert_eq!(content, "prefix"); + } + + #[test] + fn split_reasoning_prefix_without_tool_marker_preserves_all_partial_reasoning() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let ReasoningSplit { reasoning, content } = + split_reasoning_prefix("unfinished", Some(&markers), None, true); + + assert_eq!(reasoning, "unfinished"); + assert!(content.is_empty()); + } + + #[test] + fn split_reasoning_prefix_extracts_reasoning_and_trailing_content() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let ReasoningSplit { reasoning, content } = split_reasoning_prefix( + "deduceanswertail", + Some(&markers), + Some(""), + false, + ); + + assert_eq!(reasoning, "deduce"); + assert_eq!(content, "answer"); + } + + #[test] + fn restore_partial_reasoning_preserves_non_partial_parser_result() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let mut parsed = + ParsedChatMessage::new("parsed content".to_owned(), String::new(), Vec::new()); + + restore_partial_reasoning(&mut parsed, "unfinished", Some(&markers), false); + + assert_eq!(parsed.content, "parsed content"); + assert!(parsed.reasoning_content.is_empty()); + } + + #[test] + fn restore_partial_reasoning_preserves_existing_reasoning() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let mut parsed = ParsedChatMessage::new( + "parsed content".to_owned(), + "parsed reasoning".to_owned(), + Vec::new(), + ); + + restore_partial_reasoning(&mut parsed, "plain response", Some(&markers), true); + + assert_eq!(parsed.content, "parsed content"); + assert_eq!(parsed.reasoning_content, "parsed reasoning"); + } + + #[test] + fn restore_partial_reasoning_removes_open_marker_from_parser_result() { + let markers = ReasoningMarkers { + open: "\n[THINK]\n".to_owned(), + closes: vec!["[/THINK]".to_owned()], + }; + let mut parsed = ParsedChatMessage::new( + String::new(), + "[THINK]parsed reasoning".to_owned(), + Vec::new(), + ); + + restore_partial_reasoning(&mut parsed, "complete response", Some(&markers), true); + + assert!(parsed.content.is_empty()); + assert_eq!(parsed.reasoning_content, "parsed reasoning"); + } + + #[test] + fn restore_partial_reasoning_preserves_unclosed_reasoning_whitespace() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let mut parsed = + ParsedChatMessage::new(String::new(), "normalized reasoning".to_owned(), Vec::new()); + + restore_partial_reasoning(&mut parsed, "\n\nreasoning", Some(&markers), true); + + assert!(parsed.content.is_empty()); + assert_eq!(parsed.reasoning_content, "\n\nreasoning"); + } + + #[test] + fn restore_partial_reasoning_preserves_closed_reasoning_whitespace() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let mut parsed = ParsedChatMessage::new( + "answer".to_owned(), + "normalized reasoning".to_owned(), + Vec::new(), + ); + + restore_partial_reasoning( + &mut parsed, + "\n\nreasoninganswer", + Some(&markers), + true, + ); + + assert_eq!(parsed.content, "answer"); + assert_eq!(parsed.reasoning_content, "\n\nreasoning"); + } + + #[test] + fn restore_partial_reasoning_removes_open_marker_after_parser_whitespace() { + let markers = ReasoningMarkers { + open: "\n[THINK]\n".to_owned(), + closes: vec!["[/THINK]".to_owned()], + }; + let mut parsed = ParsedChatMessage::new( + String::new(), + "\n[THINK]parsed reasoning".to_owned(), + Vec::new(), + ); + + restore_partial_reasoning(&mut parsed, "complete response", Some(&markers), true); + + assert!(parsed.content.is_empty()); + assert_eq!(parsed.reasoning_content, "parsed reasoning"); + } + + #[test] + fn restore_partial_reasoning_preserves_result_without_open_marker() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let mut parsed = + ParsedChatMessage::new("parsed content".to_owned(), String::new(), Vec::new()); + + restore_partial_reasoning(&mut parsed, "unfinished", Some(&markers), true); + + assert_eq!(parsed.content, "parsed content"); + assert!(parsed.reasoning_content.is_empty()); + } + + #[test] + fn restore_partial_reasoning_recovers_unclosed_reasoning() { + let markers = ReasoningMarkers { + open: "".to_owned(), + closes: vec!["".to_owned()], + }; + let mut parsed = + ParsedChatMessage::new("unfinished".to_owned(), String::new(), Vec::new()); + + restore_partial_reasoning(&mut parsed, "unfinished", Some(&markers), true); + + assert!(parsed.content.is_empty()); + assert_eq!(parsed.reasoning_content, "unfinished"); + } + + #[test] + fn outcome_from_via_ffi_result_recognized_synthesizes_tool_call_ids() { + let parsed = ParsedChatMessage::new( + "answer".to_owned(), + String::new(), + vec![ParsedToolCall::new( + String::new(), + "tool".to_owned(), + ToolCallArguments::default(), + )], + ); + + let outcome = outcome_from_via_ffi_result(Ok(parsed), "answer", false); + + assert_eq!( + outcome.unwrap(), + ChatMessageParseOutcome::Recognized(ParsedChatMessage::new( + "answer".to_owned(), + String::new(), + vec![ParsedToolCall::new( + "call_0".to_owned(), + "tool".to_owned(), + ToolCallArguments::default(), + )], + )) + ); + } + + #[test] + fn outcome_from_via_ffi_result_message_unrecognized_is_unrecognized_with_raw_message() { + let outcome = outcome_from_via_ffi_result( + Err(ParseChatMessageError::MessageUnrecognized { + message: "boom".to_owned(), + }), + "garbled", + true, + ); + + assert_eq!( + outcome.unwrap(), + ChatMessageParseOutcome::Unrecognized(RawChatMessage { + text: "garbled".to_owned(), + is_partial: true, + ffi_error_message: "boom".to_owned(), + }) + ); + } + + #[test] + fn outcome_from_via_ffi_result_parser_creation_failure_propagates() { + let outcome = outcome_from_via_ffi_result( + Err(ParseChatMessageError::ParserCreationFailed { + message: "the parser could not be built".to_owned(), + }), + "garbled", + true, + ); + + assert_eq!( + discriminant(&outcome.unwrap_err()), + discriminant(&ParseChatMessageError::ParserCreationFailed { + message: String::new() + }) + ); + } + + #[test] + fn outcome_from_via_ffi_result_other_error_propagates() { + let outcome = outcome_from_via_ffi_result(Err(ParseChatMessageError::NoVocab), "x", false); + + assert_eq!( + discriminant(&outcome.unwrap_err()), + discriminant(&ParseChatMessageError::NoVocab) + ); + } + + #[test] + fn parsed_chat_free_ok_is_success() { + let result = unsafe { + parsed_chat_free_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_OK, + ptr::null_mut(), + ) + }; + + assert!( + result.is_ok(), + "a clean destructor must not report a failure" + ); + } + + #[test] + fn parsed_chat_free_allocation_failed_is_not_enough_memory() { + let result = unsafe { + parsed_chat_free_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_ERROR_STRING_ALLOCATION_FAILED, + ptr::null_mut(), + ) + }; + + let Err(ParseChatMessageError::NotEnoughMemory) = result else { + panic!("an error-string allocation failure must map to NotEnoughMemory"); + }; + } + + #[test] + fn parsed_chat_free_llama_cpp_out_of_memory_is_preserved() { + let result = unsafe { + parsed_chat_free_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_LLAMA_CPP_OUT_OF_MEMORY, + ptr::null_mut(), + ) + }; + + let Err(ParseChatMessageError::LlamaCppOutOfMemory) = result else { + panic!("a llama.cpp allocation failure must be reported as its own variant"); + }; + } + + #[test] + fn parsed_chat_free_destructor_threw_surfaces_the_message() { + let out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"the destructor threw".as_ptr()) + }; + let result = unsafe { + parsed_chat_free_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION, + out_error, + ) + }; + + let Err(ParseChatMessageError::DestructorFailed { message }) = result else { + panic!("a throwing destructor must surface its message"); + }; + + assert_eq!(message, "the destructor threw"); + } + + #[test] + fn parsed_chat_free_unknown_status_is_preserved() { + let result = unsafe { parsed_chat_free_status_to_result(255, ptr::null_mut()) }; + + let Err(ParseChatMessageError::FfiStatus(status_error)) = result else { + panic!("an unrecognized status must be preserved verbatim"); + }; + + assert_eq!( + status_error, + crate::FfiStatusError { + operation: "llama_rs_parsed_chat_free", + code: 255, + } + ); + } + + #[test] + fn parse_chat_message_status_to_result_maps_every_contract_status() { + let mut out_error_slot: *mut c_char = ptr::null_mut(); + let outcome_0 = unsafe { + parse_chat_message_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_TOOLS_PARSER_ARG, + ptr::null_mut(), + &raw mut out_error_slot, + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { + panic!( + "LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_TOOLS_PARSER_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_0, + crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "was given a null tools parser argument", + } + ); + let outcome_1 = unsafe { + parse_chat_message_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_INPUT_ARG, + ptr::null_mut(), + &raw mut out_error_slot, + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { + panic!("LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_INPUT_ARG must map to a contract error"); + }; + assert_eq!( + contract_1, + crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "was given a null input argument", + } + ); + let outcome_2 = unsafe { + parse_chat_message_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_HANDLE_ARG, + ptr::null_mut(), + &raw mut out_error_slot, + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_2)) = outcome_2 else { + panic!("LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_HANDLE_ARG must map to a contract error"); + }; + assert_eq!( + contract_2, + crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "was given a null out_handle argument", + } + ); + let outcome_3 = unsafe { + parse_chat_message_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_ERROR_ARG, + ptr::null_mut(), + &raw mut out_error_slot, + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_3)) = outcome_3 else { + panic!("LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_ERROR_ARG must map to a contract error"); + }; + assert_eq!( + contract_3, + crate::FfiContractError { + operation: "llama_rs_parse_chat_message", + detail: "was given a null out_error argument", + } + ); + let outcome_4 = unsafe { + parse_chat_message_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY, + ptr::null_mut(), + &raw mut out_error_slot, + ) + }; + let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_4 else { + panic!( + "LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" + ); + }; + } + + #[test] + fn parsed_chat_content_status_to_result_maps_every_contract_status() { + let outcome_0 = unsafe { + parsed_chat_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_NULL_HANDLE_ARG, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { + panic!("LLAMA_RS_PARSED_CHAT_CONTENT_NULL_HANDLE_ARG must map to a contract error"); + }; + assert_eq!( + contract_0, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_content", + detail: "was given a null handle argument", + } + ); + let outcome_1 = unsafe { + parsed_chat_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_NULL_OUT_STRING_ARG, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { + panic!("LLAMA_RS_PARSED_CHAT_CONTENT_NULL_OUT_STRING_ARG must map to a contract error"); + }; + assert_eq!( + contract_1, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_content", + detail: "was given a null out_string argument", + } + ); + let outcome_2 = unsafe { + parsed_chat_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_OUT_OF_MEMORY, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { + panic!( + "LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" + ); + }; + } + + #[test] + fn parsed_chat_reasoning_content_status_to_result_maps_every_contract_status() { + let outcome_0 = unsafe { + parsed_chat_reasoning_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_HANDLE_ARG, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { + panic!( + "LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_HANDLE_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_0, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_reasoning_content", + detail: "was given a null handle argument", + } + ); + let outcome_1 = unsafe { + parsed_chat_reasoning_content_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_OUT_STRING_ARG, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { + panic!( + "LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_OUT_STRING_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_1, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_reasoning_content", + detail: "was given a null out_string argument", + } + ); + let outcome_2 = unsafe { + parsed_chat_reasoning_content_status_to_result(llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_OUT_OF_MEMORY, ptr::null_mut(), ptr::null_mut()) + }; + let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { + panic!( + "LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" + ); + }; + } + + #[test] + fn parsed_chat_tool_call_count_status_to_result_maps_every_contract_status() { + let outcome_0 = unsafe { + parsed_chat_tool_call_count_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_HANDLE_ARG, + 0, + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_HANDLE_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_0, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_count", + detail: "was given a null handle argument", + } + ); + let outcome_1 = unsafe { + parsed_chat_tool_call_count_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_OUT_COUNT_ARG, + 0, + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_OUT_COUNT_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_1, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_count", + detail: "was given a null out_count argument", + } + ); + let outcome_2 = unsafe { + parsed_chat_tool_call_count_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_OUT_OF_MEMORY, + 0, + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" + ); + }; + } + + #[test] + fn parsed_chat_tool_call_id_status_to_result_maps_every_contract_status() { + let outcome_0 = unsafe { + parsed_chat_tool_call_id_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_HANDLE_ARG, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_HANDLE_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_0, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_id", + detail: "was given a null handle argument", + } + ); + let outcome_1 = unsafe { + parsed_chat_tool_call_id_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_OUT_STRING_ARG, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_OUT_STRING_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_1, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_id", + detail: "was given a null out_string argument", + } + ); + let outcome_2 = unsafe { + parsed_chat_tool_call_id_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_OUT_OF_MEMORY, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" + ); + }; + } + + #[test] + fn parsed_chat_tool_call_name_status_to_result_maps_every_contract_status() { + let outcome_0 = unsafe { + parsed_chat_tool_call_name_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_HANDLE_ARG, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_HANDLE_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_0, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_name", + detail: "was given a null handle argument", + } + ); + let outcome_1 = unsafe { + parsed_chat_tool_call_name_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_OUT_STRING_ARG, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_OUT_STRING_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_1, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_name", + detail: "was given a null out_string argument", + } + ); + let outcome_2 = unsafe { + parsed_chat_tool_call_name_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_OUT_OF_MEMORY, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" + ); + }; + } + + #[test] + fn parsed_chat_tool_call_arguments_status_to_result_maps_every_contract_status() { + let outcome_0 = unsafe { + parsed_chat_tool_call_arguments_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_HANDLE_ARG, + 0, + ptr::null_mut(), + ptr::null_mut(), + ) + }; + let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_HANDLE_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_0, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_arguments", + detail: "was given a null handle argument", + } + ); + let outcome_1 = unsafe { + parsed_chat_tool_call_arguments_status_to_result(llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_OUT_STRING_ARG, 0, ptr::null_mut(), ptr::null_mut()) + }; + let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_OUT_STRING_ARG must map to a contract error" + ); + }; + assert_eq!( + contract_1, + crate::FfiContractError { + operation: "llama_rs_parsed_chat_tool_call_arguments", + detail: "was given a null out_string argument", + } + ); + let outcome_2 = unsafe { + parsed_chat_tool_call_arguments_status_to_result(llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_OUT_OF_MEMORY, 0, ptr::null_mut(), ptr::null_mut()) + }; + let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { + panic!( + "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" + ); + }; + } + + fn tools_parser_creation_error( + status: llama_cpp_bindings_sys::llama_rs_chat_tools_parser_create_status, + message: Option<&CStr>, + ) -> ParseChatMessageError { + let mut out_error = message.map_or(ptr::null_mut(), |message| unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(message.as_ptr()) + }); + + let error = unsafe { + chat_tools_parser_create_status_to_result(status, ptr::null_mut(), &raw mut out_error) + } + .expect_err("the status must map to an error"); + + assert!(out_error.is_null(), "a reported message must be reclaimed"); + + error + } + + #[test] + fn tools_parser_creation_reports_a_success_without_a_parser_as_a_contract_error() { + assert_eq!( + tools_parser_creation_error( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_OK, + None + ), + ParseChatMessageError::FfiContract(crate::FfiContractError { + operation: "llama_rs_chat_tools_parser_create", + detail: "success status contained a null tools parser handle", + }) + ); + } + + #[test] + fn tools_parser_creation_reports_tools_that_are_not_an_array() { + assert_eq!( + tools_parser_creation_error( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_TOOLS_NOT_AN_ARRAY, + None + ), + ParseChatMessageError::ToolsNotAnArray + ); + } + + #[test] + fn tools_parser_creation_reports_running_out_of_memory() { + assert_eq!( + tools_parser_creation_error( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_ERROR_STRING_ALLOCATION_FAILED, + None + ), + ParseChatMessageError::NotEnoughMemory + ); + assert_eq!( + tools_parser_creation_error( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_OUT_OF_MEMORY, + None + ), + ParseChatMessageError::LlamaCppOutOfMemory + ); + } + + #[test] + fn tools_parser_creation_reports_the_build_failure_message() { + assert_eq!( + tools_parser_creation_error( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_THREW_CXX_EXCEPTION, + Some(c"tools do not fit the template") + ), + ParseChatMessageError::ToolsParserBuildFailed { + message: "tools do not fit the template".to_owned() + } + ); + } + + #[test] + fn tools_parser_creation_reports_a_build_failure_without_a_message_as_a_contract_error() { + assert_eq!( + tools_parser_creation_error( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_LLAMA_CPP_THREW_CXX_EXCEPTION, + None + ), + ParseChatMessageError::FfiContract(crate::FfiContractError { + operation: "llama_rs_chat_tools_parser_create", + detail: "reported a thrown C++ exception without an error message", + }) + ); + } + + #[test] + fn tools_parser_creation_reports_every_null_argument_as_a_contract_error() { + for (status, detail) in [ + ( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_PARSER_ARG, + "was given a null parser argument", + ), + ( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_TOOLS_JSON_ARG, + "was given a null tools_json argument", + ), + ( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_OUT_TOOLS_PARSER_ARG, + "was given a null out_tools_parser argument", + ), + ( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_CREATE_NULL_OUT_ERROR_ARG, + "was given a null out_error argument", + ), + ] { + assert_eq!( + tools_parser_creation_error(status, None), + ParseChatMessageError::FfiContract(crate::FfiContractError { + operation: "llama_rs_chat_tools_parser_create", + detail, + }) + ); + } + } + + #[test] + fn tools_parser_creation_preserves_an_unknown_status() { + assert_eq!( + tools_parser_creation_error(255, None), + ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_chat_tools_parser_create", + code: 255, + }) + ); + } + + #[test] + fn tools_parser_release_maps_every_status() { + assert_eq!( + unsafe { + chat_tools_parser_free_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_FREE_OK, + ptr::null_mut(), + ) + }, + Ok(()) + ); + assert_eq!( + unsafe { + chat_tools_parser_free_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_FREE_ERROR_STRING_ALLOCATION_FAILED, + ptr::null_mut(), + ) + }, + Err(ParseChatMessageError::NotEnoughMemory) + ); + assert_eq!( + unsafe { + chat_tools_parser_free_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_FREE_LLAMA_CPP_OUT_OF_MEMORY, + ptr::null_mut(), + ) + }, + Err(ParseChatMessageError::LlamaCppOutOfMemory) + ); + assert_eq!( + unsafe { + chat_tools_parser_free_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_TOOLS_PARSER_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION, + llama_cpp_bindings_sys::llama_rs_string_dup(c"destructor failed".as_ptr()), + ) + }, + Err(ParseChatMessageError::DestructorFailed { + message: "destructor failed".to_owned() + }) + ); + assert_eq!( + unsafe { chat_tools_parser_free_status_to_result(255, ptr::null_mut()) }, + Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_chat_tools_parser_free", + code: 255, + })) + ); + } +} diff --git a/llama-cpp-bindings/src/chat_tools.rs b/llama-cpp-bindings/src/chat_tools.rs deleted file mode 100644 index 7e5f9e5d4..000000000 --- a/llama-cpp-bindings/src/chat_tools.rs +++ /dev/null @@ -1,77 +0,0 @@ -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/classified_sample.rs b/llama-cpp-bindings/src/classified_sample.rs index 0b95c3b6f..8701ff9d6 100644 --- a/llama-cpp-bindings/src/classified_sample.rs +++ b/llama-cpp-bindings/src/classified_sample.rs @@ -1,8 +1,8 @@ -use crate::ingest_outcome::IngestOutcome; +use crate::generation_progress::GenerationProgress; use crate::token::LlamaToken; -#[derive(Clone, Debug)] +#[derive(Clone, Copy, Debug, Eq, PartialEq)] pub struct ClassifiedSample { pub token: LlamaToken, - pub outcomes: Vec, + pub progress: GenerationProgress, } diff --git a/llama-cpp-bindings/src/context.rs b/llama-cpp-bindings/src/context.rs index 6701647f9..7dfc01b2c 100644 --- a/llama-cpp-bindings/src/context.rs +++ b/llama-cpp-bindings/src/context.rs @@ -320,6 +320,10 @@ impl<'model> LlamaContext<'model> { unsafe { llama_cpp_bindings_sys::llama_n_ubatch(self.context.as_ptr()) } } + fn exceeds_micro_batch(&self, batch: &LlamaBatch) -> bool { + i64::from(batch.n_tokens()) > i64::from(self.n_ubatch()) + } + #[must_use] pub fn n_ctx(&self) -> u32 { unsafe { llama_cpp_bindings_sys::llama_n_ctx(self.context.as_ptr()) } @@ -375,6 +379,19 @@ impl<'model> LlamaContext<'model> { /// /// - `DecodeError` if the decoding failed. pub fn decode(&mut self, batch: &mut LlamaBatch) -> Result<(), DecodeError> { + let decodes_batches_in_one_micro_batch = unsafe { + llama_cpp_bindings_sys::llama_rs_context_decodes_batches_in_one_micro_batch( + self.context.as_ptr(), + ) + }; + + if decodes_batches_in_one_micro_batch && self.exceeds_micro_batch(batch) { + return Err(DecodeError::BatchExceedsMicroBatch { + n_tokens: batch.n_tokens(), + n_ubatch: self.n_ubatch(), + }); + } + let mut out_llama_cpp_return_code: i32 = 0; let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut(); let status = unsafe { @@ -397,6 +414,13 @@ impl<'model> LlamaContext<'model> { /// /// - `EncodeError` if the encoding failed. pub fn encode(&mut self, batch: &mut LlamaBatch) -> Result<(), EncodeError> { + if self.exceeds_micro_batch(batch) { + return Err(EncodeError::BatchExceedsMicroBatch { + n_tokens: batch.n_tokens(), + n_ubatch: self.n_ubatch(), + }); + } + let mut out_llama_cpp_return_code: i32 = 0; let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut(); let status = unsafe { @@ -516,10 +540,26 @@ impl<'model> LlamaContext<'model> { &self, token_index: i32, ) -> Result { - Ok(LlamaTokenDataArray::from_iter( - self.candidates_ith(token_index)?, - false, - )) + let mut token_data_array = LlamaTokenDataArray::new(Vec::new(), false); + + self.fill_token_data_array_ith(token_index, &mut token_data_array)?; + + Ok(token_data_array) + } + + /// Replaces the candidates in `token_data_array` with the logits at + /// `token_index`, reusing its allocation. + /// + /// # Errors + /// Returns `LogitsError` if the token is not initialized or out of range. + pub fn fill_token_data_array_ith( + &self, + token_index: i32, + token_data_array: &mut LlamaTokenDataArray, + ) -> Result<(), LogitsError> { + token_data_array.replace_candidates(self.candidates_ith(token_index)?); + + Ok(()) } /// # Errors diff --git a/llama-cpp-bindings/src/error.rs b/llama-cpp-bindings/src/error.rs index 43a907f9c..04deed8f4 100644 --- a/llama-cpp-bindings/src/error.rs +++ b/llama-cpp-bindings/src/error.rs @@ -1,7 +1,6 @@ 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; @@ -45,7 +44,6 @@ 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 deleted file mode 100644 index e9e5c3842..000000000 --- a/llama-cpp-bindings/src/error/chat_tools_error.rs +++ /dev/null @@ -1,11 +0,0 @@ -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/decode_error.rs b/llama-cpp-bindings/src/error/decode_error.rs index c929af077..0eb296a64 100644 --- a/llama-cpp-bindings/src/error/decode_error.rs +++ b/llama-cpp-bindings/src/error/decode_error.rs @@ -23,6 +23,8 @@ pub enum DecodeError { NotEnoughMemory, #[error("the llama.cpp library ran out of memory")] LlamaCppOutOfMemory, + #[error("decode batch of {n_tokens} tokens must fit a single micro batch of {n_ubatch} tokens")] + BatchExceedsMicroBatch { n_tokens: i32, n_ubatch: u32 }, #[error("{message}")] Reported { message: String }, } diff --git a/llama-cpp-bindings/src/error/encode_error.rs b/llama-cpp-bindings/src/error/encode_error.rs index 74a93d907..3238c040d 100644 --- a/llama-cpp-bindings/src/error/encode_error.rs +++ b/llama-cpp-bindings/src/error/encode_error.rs @@ -23,6 +23,8 @@ pub enum EncodeError { NotEnoughMemory, #[error("the llama.cpp library ran out of memory")] LlamaCppOutOfMemory, + #[error("encode batch of {n_tokens} tokens must fit a single micro batch of {n_ubatch} tokens")] + BatchExceedsMicroBatch { n_tokens: i32, n_ubatch: u32 }, #[error("{message}")] Reported { message: String }, } diff --git a/llama-cpp-bindings/src/error/grammar_error.rs b/llama-cpp-bindings/src/error/grammar_error.rs index e2e059755..e092d0053 100644 --- a/llama-cpp-bindings/src/error/grammar_error.rs +++ b/llama-cpp-bindings/src/error/grammar_error.rs @@ -19,6 +19,8 @@ pub enum GrammarError { GrammarRejected(#[source] llama_cpp_gbnf::gbnf_validation_error::GbnfValidationError), #[error("the grammar string contains an interior NUL byte")] GrammarContainsNul(#[source] NulError), + #[error("the grammar root name contains an interior NUL byte")] + RootContainsNul(#[source] NulError), #[error("a lazy-grammar trigger pattern contains an interior NUL byte")] TriggerPatternContainsNul(#[source] NulError), #[error("a DRY sequence breaker contains an interior NUL byte")] 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 488b50c56..3287d137d 100644 --- a/llama-cpp-bindings/src/error/parse_chat_message_error.rs +++ b/llama-cpp-bindings/src/error/parse_chat_message_error.rs @@ -18,6 +18,10 @@ pub enum ParseChatMessageError { LlamaCppOutOfMemory, #[error("the chat parser could not be constructed: {message}")] ParserCreationFailed { message: String }, + #[error("the chat tools contain an interior NUL byte")] + ToolsContainNulByte(#[source] std::ffi::NulError), + #[error("the chat tools are not a JSON array")] + ToolsNotAnArray, #[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}")] diff --git a/llama-cpp-bindings/src/generation_progress.rs b/llama-cpp-bindings/src/generation_progress.rs new file mode 100644 index 000000000..19bf0d951 --- /dev/null +++ b/llama-cpp-bindings/src/generation_progress.rs @@ -0,0 +1,5 @@ +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum GenerationProgress { + Continues, + Ended, +} diff --git a/llama-cpp-bindings/src/ingest_outcome.rs b/llama-cpp-bindings/src/ingest_outcome.rs index 56bc0efb9..5bf1a8d0f 100644 --- a/llama-cpp-bindings/src/ingest_outcome.rs +++ b/llama-cpp-bindings/src/ingest_outcome.rs @@ -1,8 +1,8 @@ use crate::sampled_token::SampledToken; +use crate::token_piece::TokenPiece; -#[derive(Clone, Debug)] +#[derive(Clone, Debug, Eq, PartialEq)] pub struct IngestOutcome { pub sampled_token: SampledToken, - pub visible_piece: String, - pub raw_piece: String, + pub piece: TokenPiece, } diff --git a/llama-cpp-bindings/src/json_probe_outcome.rs b/llama-cpp-bindings/src/json_probe_outcome.rs new file mode 100644 index 000000000..a5b76fe7c --- /dev/null +++ b/llama-cpp-bindings/src/json_probe_outcome.rs @@ -0,0 +1,6 @@ +#[derive(Copy, Clone, Debug, Eq, PartialEq)] +pub enum JsonProbeOutcome { + StillPossiblyValid, + CompletedValid, + Failed, +} diff --git a/llama-cpp-bindings/src/lib.rs b/llama-cpp-bindings/src/lib.rs index 69b7667ad..8f32a372a 100644 --- a/llama-cpp-bindings/src/lib.rs +++ b/llama-cpp-bindings/src/lib.rs @@ -3,15 +3,17 @@ deny(clippy::unwrap_used, clippy::expect_used, clippy::panic) )] +pub mod bare_json_tool_calls; pub mod batch_add_error; pub mod chat_message_parse_outcome; +pub mod chat_message_parser; pub mod chat_template_tool_calls; -pub mod chat_tools; pub mod classified_sample; pub mod context; pub mod error; pub mod eval_multimodal_chunks_params; pub mod extract_tool_call_markers_from_haystack; +pub mod generation_progress; pub mod ggml_time_us; pub mod gguf_context; pub mod gguf_context_error; @@ -20,6 +22,7 @@ pub mod grammar_matcher; pub mod ingest_outcome; pub mod ingest_prompt_chunk; pub mod invalid_numa_strategy; +pub mod json_probe_outcome; pub mod json_schema_to_grammar; pub mod llama_backend; pub mod llama_backend_device; @@ -61,24 +64,26 @@ pub mod streaming_markers; pub mod synthetic_tool_call_renders; pub mod timing; pub mod token; +pub mod token_piece; pub mod tool_call_format; pub mod tool_call_marker_pair; pub use error::{ - 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, + 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, }; +pub use bare_json_tool_calls::BareJsonToolCalls; pub use chat_message_parse_outcome::ChatMessageParseOutcome; -pub use chat_tools::ChatTools; +pub use chat_message_parser::ChatMessageParser; pub use classified_sample::ClassifiedSample; pub use eval_multimodal_chunks_params::EvalMultimodalChunksParams; +pub use generation_progress::GenerationProgress; pub use llama_backend_device::LlamaBackendDevice; pub use llama_backend_device_type::LlamaBackendDeviceType; pub use llama_cpp_bindings_types::{ @@ -86,6 +91,7 @@ pub use llama_cpp_bindings_types::{ ReasoningMarkers, TokenUsage, TokenUsageError, ToolCallArgsShape, ToolCallArguments, ToolCallMarkers, ToolCallValueQuote, XmlTagsShape, }; +pub use llama_cpp_gbnf::gbnf_validation_error::GbnfValidationError; pub use marker_role::MarkerRole; pub use marker_role_candidate::MarkerRoleCandidate; pub use raw_chat_message::RawChatMessage; @@ -95,6 +101,7 @@ pub use sampled_token_section::SampledTokenSection; pub use streaming_marker::StreamingMarker; pub use streaming_markers::StreamingMarkers; pub use synthetic_tool_call_renders::SyntheticToolCallRenders; +pub use token_piece::TokenPiece; pub use ggml_time_us::ggml_time_us; pub use ingest_prompt_chunk::ingest_prompt_chunk; diff --git a/llama-cpp-bindings/src/model.rs b/llama-cpp-bindings/src/model.rs index 70e588aaf..7da012175 100644 --- a/llama-cpp-bindings/src/model.rs +++ b/llama-cpp-bindings/src/model.rs @@ -28,30 +28,23 @@ use toktrie::ApproximateTokEnv; use toktrie::TokRxInfo; use toktrie::TokTrie; -use llama_cpp_bindings_types::ParsedChatMessage; -use llama_cpp_bindings_types::ParsedToolCall; use llama_cpp_bindings_types::ReasoningMarkers; -use llama_cpp_bindings_types::ToolCallArguments; use llama_cpp_bindings_types::ToolCallMarkers; -use crate::chat_message_parse_outcome::ChatMessageParseOutcome; +use crate::bare_json_tool_calls::BareJsonToolCalls; 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; use crate::marker_role::MarkerRole; use crate::marker_role_candidate::MarkerRoleCandidate; use crate::model::tokenizer_input::TokenizerInput; -use crate::raw_chat_message::RawChatMessage; use crate::resolved_tool_call_markers::ResolvedToolCallMarkers; use crate::sampled_token::SampledToken; use crate::sampled_token_classifier::SampledTokenClassifier; use crate::streaming_markers::StreamingMarkers; use crate::synthetic_tool_call_renders::SyntheticToolCallRenders; use crate::token::LlamaToken; -use crate::tool_call_format; -use crate::tool_call_format::ToolCallFormatOutcome; use crate::{ ApplyChatTemplateError, ChatTemplateError, LlamaLoraAdapterInitError, LlamaModelLoadError, MarkerDetectionError, MetaValError, ParseChatMessageError, StringToTokenError, @@ -109,42 +102,6 @@ unsafe impl Send for ChatParserHandle {} unsafe impl Sync for ChatParserHandle {} -/// # Safety -/// -/// `free_error` must be the pointer populated by the preceding -/// `llama_rs_parsed_chat_free` call, or null. The destructor-threw arm reads and -/// frees it. -unsafe fn parsed_chat_free_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_parsed_chat_free_status, - free_error: *mut c_char, -) -> Result<(), ParseChatMessageError> { - match status { - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_OK => Ok(()), - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_ERROR_STRING_ALLOCATION_FAILED => { - Err(ParseChatMessageError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_LLAMA_CPP_OUT_OF_MEMORY => { - Err(ParseChatMessageError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION => { - let message = unsafe { - read_and_free_cpp_string( - free_error, - "llama_rs_parsed_chat_free", - "reported a thrown C++ exception without an error message", - ) - }?; - - Err(ParseChatMessageError::DestructorFailed { message }) - } - other => Err(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_free", - code: i64::from(other), - } - .into()), - } -} - /// # Safety /// /// `out_error` must be the pointer populated by the preceding @@ -290,110 +247,6 @@ 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 -/// 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, - handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, - out_error: *mut *mut c_char, -) -> Result { - match status { - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_OK => { - if handle.is_null() { - Err(crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "success status contained a null parsed-chat handle", - } - .into()) - } else { - collect_parsed_chat_message(handle) - } - } - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_ERROR_STRING_ALLOCATION_FAILED => { - Err(ParseChatMessageError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY => { - Err(ParseChatMessageError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION => { - 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 { - operation: "llama_rs_parse_chat_message", - detail: "was given a null parser argument", - } - .into()) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_INPUT_ARG => { - Err(crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "was given a null input argument", - } - .into()) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_HANDLE_ARG => { - Err(crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "was given a null out_handle argument", - } - .into()) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_ERROR_ARG => { - Err(crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "was given a null out_error argument", - } - .into()) - } - other => Err(crate::FfiStatusError { - operation: "llama_rs_parse_chat_message", - code: i64::from(other), - } - .into()), - } -} - /// # Safety /// /// `out_error` must reference the pointer populated by the preceding @@ -466,27 +319,6 @@ unsafe fn chat_parser_create_status_to_result( } } -fn outcome_from_via_ffi_result( - via_ffi_result: Result, - input: &str, - is_partial: bool, -) -> Result { - match via_ffi_result { - Ok(mut parsed) => { - synthesize_missing_tool_call_ids(&mut parsed.tool_calls); - Ok(ChatMessageParseOutcome::Recognized(parsed)) - } - Err(ParseChatMessageError::MessageUnrecognized { message }) => { - Ok(ChatMessageParseOutcome::Unrecognized(RawChatMessage { - text: input.to_owned(), - is_partial, - ffi_error_message: message, - })) - } - Err(other) => Err(other), - } -} - /// # Safety /// /// `out_string` and `out_error` must be the pointers populated by the preceding @@ -624,6 +456,11 @@ impl LlamaModel { LlamaToken(token) } + #[must_use] + pub fn adds_bos_token(&self) -> bool { + unsafe { llama_cpp_bindings_sys::llama_vocab_get_add_bos(self.vocab_ptr()) } + } + #[must_use] pub fn is_eog_token(&self, token: &SampledToken) -> bool { let (SampledToken::Content(LlamaToken(id)) @@ -1027,9 +864,10 @@ impl LlamaModel { /// detection failure is surfaced to the caller instead of silently ignored. pub fn sampled_token_classifier( &self, + bare_json_tool_calls: BareJsonToolCalls, ) -> Result, MarkerDetectionError> { self.streaming_markers() - .map(|markers| SampledTokenClassifier::new(self, markers)) + .map(|markers| SampledTokenClassifier::new(self, markers, bare_json_tool_calls)) } /// # Errors @@ -1170,95 +1008,17 @@ impl LlamaModel { } } + /// Returns the model's chat parser, analysing the chat template on first use. The + /// pointer stays valid for as long as the model lives. + /// /// # Errors /// - /// 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: &ChatTools, - input: &str, - is_partial: bool, - ) -> Result { - let reasoning_markers = self.reasoning_markers()?; - - for candidate in chat_template_tool_calls::known_marker_candidates() { - match tool_call_format::try_parse(input, &candidate) { - ToolCallFormatOutcome::NoMatch => {} - ToolCallFormatOutcome::Parsed(calls) => { - let split = split_reasoning_prefix( - input, - reasoning_markers, - Some(&candidate.open), - is_partial, - ); - let mut parsed = ParsedChatMessage::new(split.content, split.reasoning, calls); - synthesize_missing_tool_call_ids(&mut parsed.tool_calls); - - return Ok(ChatMessageParseOutcome::Recognized(parsed)); - } - ToolCallFormatOutcome::Failed(_shape_does_not_fit) => {} - } - } - - let via_ffi_result = self - .parse_chat_message_via_ffi(tools.json_cstr(), input, is_partial) - .map(|mut parsed| { - restore_partial_reasoning(&mut parsed, input, reasoning_markers, is_partial); - parsed - }); - - outcome_from_via_ffi_result(via_ffi_result, input, is_partial) - } - - fn parse_chat_message_via_ffi( + /// Returns [`ParseChatMessageError`] when the model has no chat template or vocab, or + /// the parser cannot be constructed. + pub fn chat_parser_ptr( &self, - tools_cstring: &CStr, - input: &str, - is_partial: bool, - ) -> Result { - let parser = self.chat_parser()?; - - let input_cstring = - CString::new(input).map_err(ParseChatMessageError::InputContainsNulByte)?; - - let mut handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat = ptr::null_mut(); - let mut out_error: *mut c_char = ptr::null_mut(); - - let status = unsafe { - llama_cpp_bindings_sys::llama_rs_parse_chat_message( - parser.parser.as_ptr(), - tools_cstring.as_ptr(), - input_cstring.as_ptr(), - i32::from(is_partial), - &raw mut handle, - &raw mut out_error, - ) - }; - - let parsed = - unsafe { parse_chat_message_status_to_result(status, handle, &raw mut out_error) }; - - let mut free_error: *mut c_char = ptr::null_mut(); - let free_status = unsafe { - llama_cpp_bindings_sys::llama_rs_parsed_chat_free(handle, &raw mut free_error) - }; - let freed = unsafe { parsed_chat_free_status_to_result(free_status, free_error) }; - - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - - match parsed { - Ok(message) => freed.map(|()| message), - Err(parse_failure) => { - if let Err(destructor_failure) = freed { - log::error!("{destructor_failure}"); - } - - Err(parse_failure) - } - } + ) -> Result<*mut llama_cpp_bindings_sys::llama_rs_chat_parser, ParseChatMessageError> { + self.chat_parser().map(|handle| handle.parser.as_ptr()) } fn chat_parser(&self) -> Result<&ChatParserHandle, ParseChatMessageError> { @@ -1367,460 +1127,453 @@ fn build_approximate_tok_env( Ok(Arc::new(ApproximateTokEnv::new(trie))) } -fn collect_parsed_chat_message( - handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, -) -> Result { - if handle.is_null() { - return Ok(ParsedChatMessage::default()); - } - - let content = read_parsed_chat_content(handle)?; - let reasoning_content = read_parsed_chat_reasoning_content(handle)?; - let count = read_parsed_chat_tool_call_count(handle)?; - - let mut tool_calls = Vec::with_capacity(count); - for index in 0..count { - let id = read_parsed_chat_tool_call_id(handle, index)?; - let name = read_parsed_chat_tool_call_name(handle, index)?; - let arguments_json = read_parsed_chat_tool_call_arguments(handle, index)?; - - let arguments = ToolCallArguments::from_string(arguments_json); - tool_calls.push(ParsedToolCall::new(id, name, arguments)); - } - - Ok(ParsedChatMessage::new( - content, - reasoning_content, - tool_calls, - )) -} - -/// # Safety -/// -/// `out_string` and `out_error` must be the pointers populated by the preceding -/// `llama_rs_parsed_chat_content` call (or null when no value/error was produced); each is -/// read and freed in exactly one match arm. -unsafe fn parsed_chat_content_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_parsed_chat_content_status, - out_string: *mut c_char, +unsafe fn detect_reasoning_markers_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_detect_reasoning_markers_status, + out_markers: *const llama_cpp_bindings_sys::llama_rs_reasoning_markers, out_error: *mut c_char, -) -> Result { +) -> Result, MarkerDetectionError> { match status { - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_OK => { - consume_accessor_string(out_string, "llama_rs_parsed_chat_content") + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_OK => unsafe { + read_reasoning_markers(out_markers) + }, + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_NULL_MODEL_ARG => { + Err(MarkerDetectionError::WrapperRejectedArgument { + operation: "llama_rs_detect_reasoning_markers", + argument: "model", + }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_ERROR_STRING_ALLOCATION_FAILED => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::NotEnoughMemory) + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_NULL_OUT_MARKERS_ARG => { + Err(MarkerDetectionError::WrapperRejectedArgument { + operation: "llama_rs_detect_reasoning_markers", + argument: "out_markers", + }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_OUT_OF_MEMORY => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::LlamaCppOutOfMemory) + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_NULL_OUT_ERROR_ARG => { + Err(MarkerDetectionError::WrapperRejectedArgument { + operation: "llama_rs_detect_reasoning_markers", + argument: "out_error", + }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = unsafe { - read_and_free_cpp_string( - out_error, - "llama_rs_parsed_chat_content", - "reported a thrown C++ exception without an error message", - ) - }?; - Err(ParseChatMessageError::Reported { message }) + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_MODEL_HAS_NO_CHAT_TEMPLATE => Ok(None), + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_MODEL_HAS_NO_VOCAB => { + Err(MarkerDetectionError::ModelHasNoVocab { + operation: "llama_rs_detect_reasoning_markers", + }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_NULL_HANDLE_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_content", - detail: "was given a null handle argument", - } - .into()) + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_ERROR_STRING_ALLOCATION_FAILED => { + Err(MarkerDetectionError::NotEnoughMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_NULL_OUT_STRING_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_content", - detail: "was given a null out_string argument", - } - .into()) + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_LLAMA_CPP_OUT_OF_MEMORY => { + Err(MarkerDetectionError::LlamaCppOutOfMemory) } - other => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_content", - code: i64::from(other), - } - .into()) + llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = unsafe { read_and_free_cpp_string(out_error, "llama_rs_detect_reasoning_markers", "reported a thrown C++ exception without an error message") }?; + Err(MarkerDetectionError::ReasoningMarkerDetectionFailed { message }) + } + other => Err(crate::FfiStatusError { + operation: "llama_rs_detect_reasoning_markers", + code: i64::from(other), } + .into()), } } -fn read_parsed_chat_content( - handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, -) -> Result { - let mut out_string: *mut c_char = ptr::null_mut(); - let mut out_error: *mut c_char = ptr::null_mut(); - let status = unsafe { - llama_cpp_bindings_sys::llama_rs_parsed_chat_content( - handle, - &raw mut out_string, - &raw mut out_error, - ) - }; - unsafe { parsed_chat_content_status_to_result(status, out_string, out_error) } +unsafe fn read_reasoning_markers( + markers: *const llama_cpp_bindings_sys::llama_rs_reasoning_markers, +) -> Result, MarkerDetectionError> { + if markers.is_null() { + return Ok(None); + } + let open_pointer = unsafe { llama_cpp_bindings_sys::llama_rs_reasoning_markers_open(markers) }; + let open = read_optional_owned_cstr(open_pointer)?; + let close_count = + unsafe { llama_cpp_bindings_sys::llama_rs_reasoning_markers_close_count(markers) }; + let mut closes = Vec::with_capacity(close_count); + for index in 0..close_count { + let close_pointer = + unsafe { llama_cpp_bindings_sys::llama_rs_reasoning_markers_close_at(markers, index) }; + closes.push(read_optional_owned_cstr(close_pointer)?); + } + validate_reasoning_markers(open, closes).map(Some) } -/// # Safety -/// -/// `out_string` and `out_error` must be the pointers populated by the preceding -/// `llama_rs_parsed_chat_reasoning_content` call (or null when no value/error was produced); -/// each is read and freed in exactly one match arm. -unsafe fn parsed_chat_reasoning_content_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_parsed_chat_reasoning_content_status, - out_string: *mut c_char, - out_error: *mut c_char, -) -> Result { - match status { - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_OK => { - consume_accessor_string(out_string, "llama_rs_parsed_chat_reasoning_content") - } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_ERROR_STRING_ALLOCATION_FAILED => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_OUT_OF_MEMORY => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = - unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_reasoning_content", "reported a thrown C++ exception without an error message") }?; - Err(ParseChatMessageError::Reported { message }) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_HANDLE_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_reasoning_content", - detail: "was given a null handle argument", - } - .into()) +fn validate_reasoning_markers( + open: Option, + closes: Vec>, +) -> Result { + let Some(open) = open else { + return Err(crate::FfiContractError { + operation: "llama_rs_reasoning_markers_open", + detail: "non-null markers returned a null opening marker", } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_OUT_STRING_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_reasoning_content", - detail: "was given a null out_string argument", - } - .into()) - } - other => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_reasoning_content", - code: i64::from(other), - } - .into()) + .into()); + }; + if open.is_empty() || closes.is_empty() { + return Err(crate::FfiContractError { + operation: "llama_rs_detect_reasoning_markers", + detail: "detected markers must contain an opening marker and a closing marker", } + .into()); } + let closes = closes + .into_iter() + .map(|close| { + let Some(close) = close else { + return Err(crate::FfiContractError { + operation: "llama_rs_reasoning_markers_close_at", + detail: "a valid closing-marker index returned null", + } + .into()); + }; + if close.is_empty() { + return Err(crate::FfiContractError { + operation: "llama_rs_reasoning_markers_close_at", + detail: "a detected closing marker was empty", + } + .into()); + } + Ok(close) + }) + .collect::, MarkerDetectionError>>()?; + Ok(ReasoningMarkers { open, closes }) } -fn read_parsed_chat_reasoning_content( - handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, -) -> Result { - let mut out_string: *mut c_char = ptr::null_mut(); - let mut out_error: *mut c_char = ptr::null_mut(); - let status = unsafe { - llama_cpp_bindings_sys::llama_rs_parsed_chat_reasoning_content( - handle, - &raw mut out_string, - &raw mut out_error, - ) - }; - unsafe { parsed_chat_reasoning_content_status_to_result(status, out_string, out_error) } +const fn cxx_exception_owns_out_error( + parsed: &Result, +) -> bool { + matches!( + parsed, + Err(MarkerDetectionError::ReasoningMarkerDetectionFailed { .. } + | MarkerDetectionError::ToolCallHaystackComputationFailed { .. } + | MarkerDetectionError::ToolCallSyntheticRenderDiagnosisFailed { .. }) + ) } /// # Safety /// -/// `out_error` must be the pointer populated by the preceding -/// `llama_rs_parsed_chat_tool_call_count` call (or null when no error was produced); it is -/// freed in exactly one match arm. -unsafe fn parsed_chat_tool_call_count_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_count_status, - out_count: usize, - out_error: *mut c_char, -) -> Result { +/// `free_error` must be the pointer populated by the preceding +/// `llama_rs_reasoning_markers_free` call, or null. The destructor-threw arm +/// reads and frees it. +unsafe fn reasoning_markers_free_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_reasoning_markers_free_status, + free_error: *mut c_char, +) -> Result<(), MarkerDetectionError> { match status { - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_OK => Ok(out_count), - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_ERROR_STRING_ALLOCATION_FAILED => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_OUT_OF_MEMORY => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = - unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_tool_call_count", "reported a thrown C++ exception without an error message") }?; - Err(ParseChatMessageError::Reported { message }) + llama_cpp_bindings_sys::LLAMA_RS_REASONING_MARKERS_FREE_OK => Ok(()), + llama_cpp_bindings_sys::LLAMA_RS_REASONING_MARKERS_FREE_ERROR_STRING_ALLOCATION_FAILED => { + Err(MarkerDetectionError::NotEnoughMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_HANDLE_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_count", - detail: "was given a null handle argument", - } - .into()) + llama_cpp_bindings_sys::LLAMA_RS_REASONING_MARKERS_FREE_LLAMA_CPP_OUT_OF_MEMORY => { + Err(MarkerDetectionError::LlamaCppOutOfMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_OUT_COUNT_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_count", - detail: "was given a null out_count argument", - } - .into()) + llama_cpp_bindings_sys::LLAMA_RS_REASONING_MARKERS_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION => { + let message = unsafe { + read_and_free_cpp_string( + free_error, + "llama_rs_reasoning_markers_free", + "reported a thrown C++ exception without an error message", + ) + }?; + + Err(MarkerDetectionError::ReasoningMarkersFreeFailed { message }) } - other => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_tool_call_count", - code: i64::from(other), - } - .into()) + other => Err(crate::FfiStatusError { + operation: "llama_rs_reasoning_markers_free", + code: i64::from(other), } + .into()), } } -fn read_parsed_chat_tool_call_count( - handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, -) -> Result { - let mut out_count: usize = 0; +fn invoke_detect_reasoning_markers( + model: *const llama_cpp_bindings_sys::llama_model, +) -> Result, MarkerDetectionError> { + let mut out_markers: *mut llama_cpp_bindings_sys::llama_rs_reasoning_markers = ptr::null_mut(); let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { - llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_count( - handle, - &raw mut out_count, + llama_cpp_bindings_sys::llama_rs_detect_reasoning_markers( + model, + &raw mut out_markers, &raw mut out_error, ) }; - unsafe { parsed_chat_tool_call_count_status_to_result(status, out_count, out_error) } + + let parsed = + unsafe { detect_reasoning_markers_status_to_result(status, out_markers, out_error) }; + + let mut free_error: *mut c_char = ptr::null_mut(); + let free_status = unsafe { + llama_cpp_bindings_sys::llama_rs_reasoning_markers_free(out_markers, &raw mut free_error) + }; + let freed = unsafe { reasoning_markers_free_status_to_result(free_status, free_error) }; + + if !cxx_exception_owns_out_error(&parsed) { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + } + + match parsed { + Ok(markers) => freed.map(|()| markers), + Err(detection_failure) => { + if let Err(destructor_failure) = freed { + log::error!("{destructor_failure}"); + } + + Err(detection_failure) + } + } } /// # Safety /// -/// `out_string` and `out_error` must be the pointers populated by the preceding -/// `llama_rs_parsed_chat_tool_call_id` call (or null when no value/error was produced); each -/// is read and freed in exactly one match arm. -unsafe fn parsed_chat_tool_call_id_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_id_status, - index: usize, - out_string: *mut c_char, +/// `out_haystack` and `out_error` must be the pointers populated by the preceding +/// `llama_rs_compute_tool_call_haystack` call (or null). `out_haystack` is read but not freed +/// here; `out_error` is freed only in the CXX-exception arm, mirroring the conditional cleanup +/// in the caller. +unsafe fn compute_tool_call_haystack_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_compute_tool_call_haystack_status, + out_haystack: *const c_char, out_error: *mut c_char, -) -> Result { +) -> Result, MarkerDetectionError> { match status { - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_OK => { - consume_accessor_string(out_string, "llama_rs_parsed_chat_tool_call_id") + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_OK => { + read_optional_owned_cstr(out_haystack) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_INDEX_OUT_OF_BOUNDS => { - Err(ParseChatMessageError::ToolCallIdIndexOutOfBounds { index }) + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_MODEL_HAS_NO_CHAT_TEMPLATE => Ok(None), + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_MODEL_HAS_NO_VOCAB => { + Err(MarkerDetectionError::ModelHasNoVocab { + operation: "llama_rs_compute_tool_call_haystack", + }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_ERROR_STRING_ALLOCATION_FAILED => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::NotEnoughMemory) + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_ERROR_STRING_ALLOCATION_FAILED => { + Err(MarkerDetectionError::NotEnoughMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_OUT_OF_MEMORY => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::LlamaCppOutOfMemory) + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_LLAMA_CPP_OUT_OF_MEMORY => { + Err(MarkerDetectionError::LlamaCppOutOfMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = - unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_tool_call_id", "reported a thrown C++ exception without an error message") }?; - Err(ParseChatMessageError::Reported { message }) + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = unsafe { read_and_free_cpp_string(out_error, "llama_rs_compute_tool_call_haystack", "reported a thrown C++ exception without an error message") }?; + Err(MarkerDetectionError::ToolCallHaystackComputationFailed { message }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_HANDLE_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_id", - detail: "was given a null handle argument", - } - .into()) + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_NULL_MODEL_ARG => Err(crate::FfiContractError { + operation: "llama_rs_compute_tool_call_haystack", + detail: "was given a null model argument", } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_OUT_STRING_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_id", - detail: "was given a null out_string argument", - } - .into()) + .into()), + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_NULL_OUT_HAYSTACK_ARG => Err(crate::FfiContractError { + operation: "llama_rs_compute_tool_call_haystack", + detail: "was given a null out_haystack argument", } - other => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_tool_call_id", - code: i64::from(other), - } - .into()) + .into()), + llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_NULL_OUT_ERROR_ARG => Err(crate::FfiContractError { + operation: "llama_rs_compute_tool_call_haystack", + detail: "was given a null out_error argument", } + .into()), + other => Err(crate::FfiStatusError { + operation: "llama_rs_compute_tool_call_haystack", + code: i64::from(other), + } + .into()), } } -fn read_parsed_chat_tool_call_id( - handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, - index: usize, -) -> Result { - let mut out_string: *mut c_char = ptr::null_mut(); +fn invoke_compute_tool_call_haystack( + model: *const llama_cpp_bindings_sys::llama_model, +) -> Result, MarkerDetectionError> { + let mut out_haystack: *mut c_char = ptr::null_mut(); let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { - llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_id( - handle, - index, - &raw mut out_string, + llama_cpp_bindings_sys::llama_rs_compute_tool_call_haystack( + model, + &raw mut out_haystack, &raw mut out_error, ) }; - unsafe { parsed_chat_tool_call_id_status_to_result(status, index, out_string, out_error) } + + let parsed = + unsafe { compute_tool_call_haystack_status_to_result(status, out_haystack, out_error) }; + + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_haystack) }; + if !cxx_exception_owns_out_error(&parsed) { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + } + + parsed } /// # Safety /// -/// `out_string` and `out_error` must be the pointers populated by the preceding -/// `llama_rs_parsed_chat_tool_call_name` call (or null when no value/error was produced); each -/// is read and freed in exactly one match arm. -unsafe fn parsed_chat_tool_call_name_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_name_status, - index: usize, - out_string: *mut c_char, +/// `out_no_tools`, `out_with_tools`, and `out_error` must be the pointers populated by the +/// preceding `llama_rs_diagnose_tool_call_synthetic_renders` call (or null). The render +/// pointers are read but not freed here; `out_error` is freed only in the CXX-exception arm, +/// mirroring the cleanup in the caller. +unsafe fn diagnose_tool_call_synthetic_renders_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_diagnose_tool_call_synthetic_renders_status, + out_no_tools: *const c_char, + out_with_tools: *const c_char, out_error: *mut c_char, -) -> Result { +) -> Result { match status { - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_OK => { - consume_accessor_string(out_string, "llama_rs_parsed_chat_tool_call_name") + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_OK => { + collect_synthetic_tool_call_renders(out_no_tools, out_with_tools) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_INDEX_OUT_OF_BOUNDS => { - Err(ParseChatMessageError::ToolCallNameIndexOutOfBounds { index }) + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_MODEL_HAS_NO_CHAT_TEMPLATE => { + Err(MarkerDetectionError::ModelHasNoChatTemplate { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_ERROR_STRING_ALLOCATION_FAILED => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::NotEnoughMemory) + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_MODEL_HAS_NO_VOCAB => { + Err(MarkerDetectionError::ModelHasNoVocab { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_OUT_OF_MEMORY => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::LlamaCppOutOfMemory) + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_ERROR_STRING_ALLOCATION_FAILED => { + Err(MarkerDetectionError::NotEnoughMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = - unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_tool_call_name", "reported a thrown C++ exception without an error message") }?; - Err(ParseChatMessageError::Reported { message }) + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_LLAMA_CPP_OUT_OF_MEMORY => { + Err(MarkerDetectionError::LlamaCppOutOfMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_HANDLE_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_name", - detail: "was given a null handle argument", - } - .into()) + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = unsafe { read_and_free_cpp_string(out_error, "llama_rs_diagnose_tool_call_synthetic_renders", "reported a thrown C++ exception without an error message") }?; + Err(MarkerDetectionError::ToolCallSyntheticRenderDiagnosisFailed { message }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_OUT_STRING_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_name", - detail: "was given a null out_string argument", - } - .into()) + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_NULL_MODEL_ARG => Err(crate::FfiContractError { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + detail: "was given a null model argument", } - other => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_tool_call_name", - code: i64::from(other), - } - .into()) + .into()), + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_NULL_OUT_NO_TOOLS_ARG => Err(crate::FfiContractError { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + detail: "was given a null out_no_tools argument", } - } -} + .into()), + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_NULL_OUT_WITH_TOOLS_ARG => Err(crate::FfiContractError { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + detail: "was given a null out_with_tools argument", + } + .into()), + llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_NULL_OUT_ERROR_ARG => Err(crate::FfiContractError { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + detail: "was given a null out_error argument", + } + .into()), + other => Err(crate::FfiStatusError { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + code: i64::from(other), + } + .into()), + } +} -fn read_parsed_chat_tool_call_name( - handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, - index: usize, -) -> Result { - let mut out_string: *mut c_char = ptr::null_mut(); +fn invoke_diagnose_tool_call_synthetic_renders( + model: *const llama_cpp_bindings_sys::llama_model, +) -> Result { + let mut out_no_tools: *mut c_char = ptr::null_mut(); + let mut out_with_tools: *mut c_char = ptr::null_mut(); let mut out_error: *mut c_char = ptr::null_mut(); + let status = unsafe { - llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_name( - handle, - index, - &raw mut out_string, + llama_cpp_bindings_sys::llama_rs_diagnose_tool_call_synthetic_renders( + model, + &raw mut out_no_tools, + &raw mut out_with_tools, &raw mut out_error, ) }; - unsafe { parsed_chat_tool_call_name_status_to_result(status, index, out_string, out_error) } + + let parsed = unsafe { + diagnose_tool_call_synthetic_renders_status_to_result( + status, + out_no_tools, + out_with_tools, + out_error, + ) + }; + + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_no_tools) }; + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_with_tools) }; + if !cxx_exception_owns_out_error(&parsed) { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + } + + parsed +} + +fn read_optional_owned_cstr(ptr: *const c_char) -> Result, MarkerDetectionError> { + if ptr.is_null() { + return Ok(None); + } + + let bytes = unsafe { CStr::from_ptr(ptr) }.to_bytes().to_vec(); + + Ok(Some(String::from_utf8(bytes)?)) } /// # Safety /// -/// `out_string` and `out_error` must be the pointers populated by the preceding -/// `llama_rs_parsed_chat_tool_call_arguments` call (or null when no value/error was produced); -/// each is read and freed in exactly one match arm. -unsafe fn parsed_chat_tool_call_arguments_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_arguments_status, - index: usize, - out_string: *mut c_char, +/// `out_error` must be the pointer populated by the preceding `llama_rs_tokenize` call (or +/// null when no error was produced); it is read and freed only in the CXX-exception arm. +unsafe fn tokenize_status_to_result( + status: llama_cpp_bindings_sys::llama_rs_tokenize_status, + out_count: c_int, out_error: *mut c_char, -) -> Result { +) -> Result { match status { - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_OK => { - consume_accessor_string(out_string, "llama_rs_parsed_chat_tool_call_arguments") + llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_OK => Ok(out_count), + llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_ERROR_STRING_ALLOCATION_FAILED => { + Err(StringToTokenError::NotEnoughMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_INDEX_OUT_OF_BOUNDS => { - Err(ParseChatMessageError::ToolCallArgumentsIndexOutOfBounds { index }) + llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_LLAMA_CPP_OUT_OF_MEMORY => { + Err(StringToTokenError::LlamaCppOutOfMemory) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_ERROR_STRING_ALLOCATION_FAILED => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::NotEnoughMemory) + llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_LLAMA_CPP_THREW_CXX_EXCEPTION => { + let message = unsafe { + read_and_free_cpp_string( + out_error, + "llama_rs_tokenize", + "reported a thrown C++ exception without an error message", + ) + }?; + Err(StringToTokenError::Reported { message }) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_OUT_OF_MEMORY => { + llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_NULL_VOCAB_ARG => { unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(ParseChatMessageError::LlamaCppOutOfMemory) + Err(crate::FfiContractError { + operation: "llama_rs_tokenize", + detail: "was given a null vocab argument", + } + .into()) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = - unsafe { read_and_free_cpp_string(out_error, "llama_rs_parsed_chat_tool_call_arguments", "reported a thrown C++ exception without an error message") }?; - Err(ParseChatMessageError::Reported { message }) + llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_NULL_TEXT_ARG => { + unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; + Err(crate::FfiContractError { + operation: "llama_rs_tokenize", + detail: "was given a null text argument", + } + .into()) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_HANDLE_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_NULL_OUT_RETURNED_COUNT_ARG => { unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_arguments", - detail: "was given a null handle argument", + operation: "llama_rs_tokenize", + detail: "was given a null out_returned_count argument", } .into()) } - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_OUT_STRING_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; + llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_NULL_OUT_ERROR_ARG => { unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; Err(crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_arguments", - detail: "was given a null out_string argument", + operation: "llama_rs_tokenize", + detail: "was given a null out_error argument", } .into()) } other => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_string) }; unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; Err(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_tool_call_arguments", + operation: "llama_rs_tokenize", code: i64::from(other), } .into()) @@ -1828,1696 +1581,557 @@ unsafe fn parsed_chat_tool_call_arguments_status_to_result( } } -fn read_parsed_chat_tool_call_arguments( - handle: *mut llama_cpp_bindings_sys::llama_rs_parsed_chat, - index: usize, -) -> Result { - let mut out_string: *mut c_char = ptr::null_mut(); +fn invoke_rs_tokenize( + vocab: *const llama_cpp_bindings_sys::llama_vocab, + text: *const c_char, + text_len: c_int, + tokens: *mut llama_cpp_bindings_sys::llama_token, + n_tokens_max: c_int, + add_bos: bool, +) -> Result { + let mut out_count: i32 = 0; let mut out_error: *mut c_char = ptr::null_mut(); let status = unsafe { - llama_cpp_bindings_sys::llama_rs_parsed_chat_tool_call_arguments( - handle, - index, - &raw mut out_string, + llama_cpp_bindings_sys::llama_rs_tokenize( + vocab, + text, + text_len, + tokens, + n_tokens_max, + add_bos, + true, + &raw mut out_count, &raw mut out_error, ) }; - unsafe { - parsed_chat_tool_call_arguments_status_to_result(status, index, out_string, out_error) - } + unsafe { tokenize_status_to_result(status, out_count, out_error) } } -fn consume_accessor_string( - ptr: *mut c_char, - operation: &'static str, -) -> Result { - if ptr.is_null() { - return Err(crate::FfiContractError { - operation, - detail: "success status contained a null string", - } - .into()); - } - let bytes = unsafe { CStr::from_ptr(ptr) }.to_bytes().to_vec(); - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(ptr) }; - Ok(String::from_utf8(bytes)?) +fn checked_token_buffer_capacity(capacity: usize) -> Result { + Ok(c_int::try_from(capacity)?) } -struct ReasoningSplit { - reasoning: String, - content: String, +fn checked_token_count(size: i32) -> Result { + Ok(usize::try_from(size)?) } -fn restore_partial_reasoning( - parsed: &mut ParsedChatMessage, - input: &str, - reasoning_markers: Option<&ReasoningMarkers>, - is_partial: bool, -) { - if !is_partial { - return; - } - if reasoning_markers.is_some_and(|markers| input.contains(&markers.open)) { - let split = split_reasoning_prefix(input, reasoning_markers, None, true); - parsed.reasoning_content = split.reasoning; - parsed.content = split.content; - return; - } - if let Some(open) = reasoning_markers.map(|markers| markers.open.trim()) - && let Some(reasoning) = parsed.reasoning_content.trim_start().strip_prefix(open) - { - parsed.reasoning_content = reasoning.to_owned(); - } -} +fn tokenize_into_buffer( + estimated_capacity: usize, + invoke: impl Fn( + *mut llama_cpp_bindings_sys::llama_token, + c_int, + ) -> Result, +) -> Result, StringToTokenError> { + let mut buffer: Vec = Vec::with_capacity(estimated_capacity); + let buffer_capacity = checked_token_buffer_capacity(buffer.capacity())?; -fn split_reasoning_prefix( - input: &str, - reasoning_markers: Option<&ReasoningMarkers>, - tool_call_open: Option<&str>, - is_partial: bool, -) -> ReasoningSplit { - let content_only = || ReasoningSplit { - reasoning: String::new(), - content: prefix_before_optional(input, tool_call_open), - }; + let size = invoke( + buffer + .as_mut_ptr() + .cast::(), + buffer_capacity, + )?; - let Some(reasoning_markers) = reasoning_markers else { - return content_only(); - }; - let Some(open_pos) = input.find(&reasoning_markers.open) else { - return content_only(); + let size = if size.is_negative() { + buffer.reserve_exact(checked_token_count(-size)?); + invoke( + buffer + .as_mut_ptr() + .cast::(), + -size, + )? + } else { + size }; - let after_open = &input[open_pos + reasoning_markers.open.len()..]; - let closing_marker = reasoning_markers - .closes - .iter() - .enumerate() - .filter_map(|(marker_index, marker)| { - after_open - .find(marker) - .map(|offset| (offset, marker_index, marker)) - }) - .min_by_key(|(offset, marker_index, _)| (*offset, *marker_index)); - let Some((close_offset, _, close_marker)) = closing_marker else { - return if is_partial { - ReasoningSplit { - reasoning: prefix_before_optional(after_open, tool_call_open), - content: input[..open_pos].to_owned(), - } - } else { - content_only() - }; - }; + let size = checked_token_count(size)?; - let reasoning = after_open[..close_offset].to_owned(); - let after_close = &after_open[close_offset + close_marker.len()..]; + unsafe { buffer.set_len(size) } - ReasoningSplit { - reasoning, - content: prefix_before_optional(after_close, tool_call_open), - } + Ok(buffer) } -fn prefix_before_optional(text: &str, marker: Option<&str>) -> String { - marker.map_or_else( - || text.to_owned(), - |marker| { - text.find(marker) - .map_or_else(|| text.to_owned(), |pos| text[..pos].to_owned()) - }, - ) +fn collect_synthetic_tool_call_renders( + without_tools_ptr: *const c_char, + with_tools_ptr: *const c_char, +) -> Result { + let without_tools = + read_optional_owned_cstr(without_tools_ptr)?.ok_or(crate::FfiContractError { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + detail: "success status contained a null render without tools", + })?; + let with_tools = read_optional_owned_cstr(with_tools_ptr)?.ok_or(crate::FfiContractError { + operation: "llama_rs_diagnose_tool_call_synthetic_renders", + detail: "success status contained a null render with tools", + })?; + + Ok(SyntheticToolCallRenders { + without_tools, + with_tools, + }) } -fn synthesize_missing_tool_call_ids(tool_calls: &mut [ParsedToolCall]) { - for (index, call) in tool_calls.iter_mut().enumerate() { - if call.id.is_empty() { - call.id = format!("call_{index}"); - } +fn extract_meta_string( + c_function: TCFunction, + capacity: usize, +) -> Result +where + TCFunction: Fn(*mut c_char, usize) -> i32, +{ + let mut buffer = vec![0u8; capacity]; + let result = c_function(buffer.as_mut_ptr().cast::(), buffer.len()); + + if result < 0 { + return Err(MetaValError::NegativeReturn(result)); } -} -unsafe fn detect_reasoning_markers_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_detect_reasoning_markers_status, - out_markers: *const llama_cpp_bindings_sys::llama_rs_reasoning_markers, - out_error: *mut c_char, -) -> Result, MarkerDetectionError> { - match status { - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_OK => unsafe { - read_reasoning_markers(out_markers) - }, - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_NULL_MODEL_ARG => { - Err(MarkerDetectionError::WrapperRejectedArgument { - operation: "llama_rs_detect_reasoning_markers", - argument: "model", - }) - } - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_NULL_OUT_MARKERS_ARG => { - Err(MarkerDetectionError::WrapperRejectedArgument { - operation: "llama_rs_detect_reasoning_markers", - argument: "out_markers", - }) - } - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_NULL_OUT_ERROR_ARG => { - Err(MarkerDetectionError::WrapperRejectedArgument { - operation: "llama_rs_detect_reasoning_markers", - argument: "out_error", - }) - } - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_MODEL_HAS_NO_CHAT_TEMPLATE => Ok(None), - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_MODEL_HAS_NO_VOCAB => { - Err(MarkerDetectionError::ModelHasNoVocab { - operation: "llama_rs_detect_reasoning_markers", - }) - } - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_ERROR_STRING_ALLOCATION_FAILED => { - Err(MarkerDetectionError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_LLAMA_CPP_OUT_OF_MEMORY => { - Err(MarkerDetectionError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_DETECT_REASONING_MARKERS_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = unsafe { read_and_free_cpp_string(out_error, "llama_rs_detect_reasoning_markers", "reported a thrown C++ exception without an error message") }?; - Err(MarkerDetectionError::ReasoningMarkerDetectionFailed { message }) - } - other => Err(crate::FfiStatusError { - operation: "llama_rs_detect_reasoning_markers", - code: i64::from(other), - } - .into()), + let returned_len = result.cast_unsigned() as usize; + + if returned_len >= capacity { + return extract_meta_string(c_function, returned_len + 1); } -} -unsafe fn read_reasoning_markers( - markers: *const llama_cpp_bindings_sys::llama_rs_reasoning_markers, -) -> Result, MarkerDetectionError> { - if markers.is_null() { - return Ok(None); - } - let open_pointer = unsafe { llama_cpp_bindings_sys::llama_rs_reasoning_markers_open(markers) }; - let open = read_optional_owned_cstr(open_pointer)?; - let close_count = - unsafe { llama_cpp_bindings_sys::llama_rs_reasoning_markers_close_count(markers) }; - let mut closes = Vec::with_capacity(close_count); - for index in 0..close_count { - let close_pointer = - unsafe { llama_cpp_bindings_sys::llama_rs_reasoning_markers_close_at(markers, index) }; - closes.push(read_optional_owned_cstr(close_pointer)?); - } - validate_reasoning_markers(open, closes).map(Some) -} - -fn validate_reasoning_markers( - open: Option, - closes: Vec>, -) -> Result { - let Some(open) = open else { - return Err(crate::FfiContractError { - operation: "llama_rs_reasoning_markers_open", - detail: "non-null markers returned a null opening marker", - } - .into()); - }; - if open.is_empty() || closes.is_empty() { - return Err(crate::FfiContractError { - operation: "llama_rs_detect_reasoning_markers", - detail: "detected markers must contain an opening marker and a closing marker", - } - .into()); - } - let closes = closes - .into_iter() - .map(|close| { - let Some(close) = close else { - return Err(crate::FfiContractError { - operation: "llama_rs_reasoning_markers_close_at", - detail: "a valid closing-marker index returned null", - } - .into()); - }; - if close.is_empty() { - return Err(crate::FfiContractError { - operation: "llama_rs_reasoning_markers_close_at", - detail: "a detected closing marker was empty", - } - .into()); - } - Ok(close) - }) - .collect::, MarkerDetectionError>>()?; - Ok(ReasoningMarkers { open, closes }) -} - -const fn cxx_exception_owns_out_error( - parsed: &Result, -) -> bool { - matches!( - parsed, - Err(MarkerDetectionError::ReasoningMarkerDetectionFailed { .. } - | MarkerDetectionError::ToolCallHaystackComputationFailed { .. } - | MarkerDetectionError::ToolCallSyntheticRenderDiagnosisFailed { .. }) - ) -} - -/// # Safety -/// -/// `free_error` must be the pointer populated by the preceding -/// `llama_rs_reasoning_markers_free` call, or null. The destructor-threw arm -/// reads and frees it. -unsafe fn reasoning_markers_free_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_reasoning_markers_free_status, - free_error: *mut c_char, -) -> Result<(), MarkerDetectionError> { - match status { - llama_cpp_bindings_sys::LLAMA_RS_REASONING_MARKERS_FREE_OK => Ok(()), - llama_cpp_bindings_sys::LLAMA_RS_REASONING_MARKERS_FREE_ERROR_STRING_ALLOCATION_FAILED => { - Err(MarkerDetectionError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_REASONING_MARKERS_FREE_LLAMA_CPP_OUT_OF_MEMORY => { - Err(MarkerDetectionError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_REASONING_MARKERS_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION => { - let message = unsafe { - read_and_free_cpp_string( - free_error, - "llama_rs_reasoning_markers_free", - "reported a thrown C++ exception without an error message", - ) - }?; - - Err(MarkerDetectionError::ReasoningMarkersFreeFailed { message }) - } - other => Err(crate::FfiStatusError { - operation: "llama_rs_reasoning_markers_free", - code: i64::from(other), - } - .into()), - } -} - -fn invoke_detect_reasoning_markers( - model: *const llama_cpp_bindings_sys::llama_model, -) -> Result, MarkerDetectionError> { - let mut out_markers: *mut llama_cpp_bindings_sys::llama_rs_reasoning_markers = ptr::null_mut(); - let mut out_error: *mut c_char = ptr::null_mut(); - - let status = unsafe { - llama_cpp_bindings_sys::llama_rs_detect_reasoning_markers( - model, - &raw mut out_markers, - &raw mut out_error, - ) - }; - - let parsed = - unsafe { detect_reasoning_markers_status_to_result(status, out_markers, out_error) }; - - let mut free_error: *mut c_char = ptr::null_mut(); - let free_status = unsafe { - llama_cpp_bindings_sys::llama_rs_reasoning_markers_free(out_markers, &raw mut free_error) - }; - let freed = unsafe { reasoning_markers_free_status_to_result(free_status, free_error) }; - - if !cxx_exception_owns_out_error(&parsed) { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - } - - match parsed { - Ok(markers) => freed.map(|()| markers), - Err(detection_failure) => { - if let Err(destructor_failure) = freed { - log::error!("{destructor_failure}"); - } - - Err(detection_failure) - } - } -} - -/// # Safety -/// -/// `out_haystack` and `out_error` must be the pointers populated by the preceding -/// `llama_rs_compute_tool_call_haystack` call (or null). `out_haystack` is read but not freed -/// here; `out_error` is freed only in the CXX-exception arm, mirroring the conditional cleanup -/// in the caller. -unsafe fn compute_tool_call_haystack_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_compute_tool_call_haystack_status, - out_haystack: *const c_char, - out_error: *mut c_char, -) -> Result, MarkerDetectionError> { - match status { - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_OK => { - read_optional_owned_cstr(out_haystack) - } - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_MODEL_HAS_NO_CHAT_TEMPLATE => Ok(None), - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_MODEL_HAS_NO_VOCAB => { - Err(MarkerDetectionError::ModelHasNoVocab { - operation: "llama_rs_compute_tool_call_haystack", - }) - } - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_ERROR_STRING_ALLOCATION_FAILED => { - Err(MarkerDetectionError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_LLAMA_CPP_OUT_OF_MEMORY => { - Err(MarkerDetectionError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = unsafe { read_and_free_cpp_string(out_error, "llama_rs_compute_tool_call_haystack", "reported a thrown C++ exception without an error message") }?; - Err(MarkerDetectionError::ToolCallHaystackComputationFailed { message }) - } - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_NULL_MODEL_ARG => Err(crate::FfiContractError { - operation: "llama_rs_compute_tool_call_haystack", - detail: "was given a null model argument", - } - .into()), - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_NULL_OUT_HAYSTACK_ARG => Err(crate::FfiContractError { - operation: "llama_rs_compute_tool_call_haystack", - detail: "was given a null out_haystack argument", - } - .into()), - llama_cpp_bindings_sys::LLAMA_RS_COMPUTE_TOOL_CALL_HAYSTACK_NULL_OUT_ERROR_ARG => Err(crate::FfiContractError { - operation: "llama_rs_compute_tool_call_haystack", - detail: "was given a null out_error argument", - } - .into()), - other => Err(crate::FfiStatusError { - operation: "llama_rs_compute_tool_call_haystack", - code: i64::from(other), - } - .into()), - } -} - -fn invoke_compute_tool_call_haystack( - model: *const llama_cpp_bindings_sys::llama_model, -) -> Result, MarkerDetectionError> { - let mut out_haystack: *mut c_char = ptr::null_mut(); - let mut out_error: *mut c_char = ptr::null_mut(); - - let status = unsafe { - llama_cpp_bindings_sys::llama_rs_compute_tool_call_haystack( - model, - &raw mut out_haystack, - &raw mut out_error, - ) - }; - - let parsed = - unsafe { compute_tool_call_haystack_status_to_result(status, out_haystack, out_error) }; - - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_haystack) }; - if !cxx_exception_owns_out_error(&parsed) { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - } - - parsed -} - -/// # Safety -/// -/// `out_no_tools`, `out_with_tools`, and `out_error` must be the pointers populated by the -/// preceding `llama_rs_diagnose_tool_call_synthetic_renders` call (or null). The render -/// pointers are read but not freed here; `out_error` is freed only in the CXX-exception arm, -/// mirroring the cleanup in the caller. -unsafe fn diagnose_tool_call_synthetic_renders_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_diagnose_tool_call_synthetic_renders_status, - out_no_tools: *const c_char, - out_with_tools: *const c_char, - out_error: *mut c_char, -) -> Result { - match status { - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_OK => { - collect_synthetic_tool_call_renders(out_no_tools, out_with_tools) - } - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_MODEL_HAS_NO_CHAT_TEMPLATE => { - Err(MarkerDetectionError::ModelHasNoChatTemplate { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - }) - } - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_MODEL_HAS_NO_VOCAB => { - Err(MarkerDetectionError::ModelHasNoVocab { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - }) - } - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_ERROR_STRING_ALLOCATION_FAILED => { - Err(MarkerDetectionError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_LLAMA_CPP_OUT_OF_MEMORY => { - Err(MarkerDetectionError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = unsafe { read_and_free_cpp_string(out_error, "llama_rs_diagnose_tool_call_synthetic_renders", "reported a thrown C++ exception without an error message") }?; - Err(MarkerDetectionError::ToolCallSyntheticRenderDiagnosisFailed { message }) - } - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_NULL_MODEL_ARG => Err(crate::FfiContractError { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - detail: "was given a null model argument", - } - .into()), - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_NULL_OUT_NO_TOOLS_ARG => Err(crate::FfiContractError { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - detail: "was given a null out_no_tools argument", - } - .into()), - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_NULL_OUT_WITH_TOOLS_ARG => Err(crate::FfiContractError { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - detail: "was given a null out_with_tools argument", - } - .into()), - llama_cpp_bindings_sys::LLAMA_RS_DIAGNOSE_TOOL_CALL_SYNTHETIC_RENDERS_NULL_OUT_ERROR_ARG => Err(crate::FfiContractError { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - detail: "was given a null out_error argument", - } - .into()), - other => Err(crate::FfiStatusError { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - code: i64::from(other), - } - .into()), - } -} - -fn invoke_diagnose_tool_call_synthetic_renders( - model: *const llama_cpp_bindings_sys::llama_model, -) -> Result { - let mut out_no_tools: *mut c_char = ptr::null_mut(); - let mut out_with_tools: *mut c_char = ptr::null_mut(); - let mut out_error: *mut c_char = ptr::null_mut(); - - let status = unsafe { - llama_cpp_bindings_sys::llama_rs_diagnose_tool_call_synthetic_renders( - model, - &raw mut out_no_tools, - &raw mut out_with_tools, - &raw mut out_error, - ) - }; - - let parsed = unsafe { - diagnose_tool_call_synthetic_renders_status_to_result( - status, - out_no_tools, - out_with_tools, - out_error, - ) - }; - - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_no_tools) }; - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_with_tools) }; - if !cxx_exception_owns_out_error(&parsed) { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - } - - parsed -} - -fn read_optional_owned_cstr(ptr: *const c_char) -> Result, MarkerDetectionError> { - if ptr.is_null() { - return Ok(None); - } - - let bytes = unsafe { CStr::from_ptr(ptr) }.to_bytes().to_vec(); - - Ok(Some(String::from_utf8(bytes)?)) -} - -/// # Safety -/// -/// `out_error` must be the pointer populated by the preceding `llama_rs_tokenize` call (or -/// null when no error was produced); it is read and freed only in the CXX-exception arm. -unsafe fn tokenize_status_to_result( - status: llama_cpp_bindings_sys::llama_rs_tokenize_status, - out_count: c_int, - out_error: *mut c_char, -) -> Result { - match status { - llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_OK => Ok(out_count), - llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_ERROR_STRING_ALLOCATION_FAILED => { - Err(StringToTokenError::NotEnoughMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_LLAMA_CPP_OUT_OF_MEMORY => { - Err(StringToTokenError::LlamaCppOutOfMemory) - } - llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_LLAMA_CPP_THREW_CXX_EXCEPTION => { - let message = unsafe { - read_and_free_cpp_string( - out_error, - "llama_rs_tokenize", - "reported a thrown C++ exception without an error message", - ) - }?; - Err(StringToTokenError::Reported { message }) - } - llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_NULL_VOCAB_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_tokenize", - detail: "was given a null vocab argument", - } - .into()) - } - llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_NULL_TEXT_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_tokenize", - detail: "was given a null text argument", - } - .into()) - } - llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_NULL_OUT_RETURNED_COUNT_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_tokenize", - detail: "was given a null out_returned_count argument", - } - .into()) - } - llama_cpp_bindings_sys::LLAMA_RS_TOKENIZE_NULL_OUT_ERROR_ARG => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiContractError { - operation: "llama_rs_tokenize", - detail: "was given a null out_error argument", - } - .into()) - } - other => { - unsafe { llama_cpp_bindings_sys::llama_rs_string_free(out_error) }; - Err(crate::FfiStatusError { - operation: "llama_rs_tokenize", - code: i64::from(other), - } - .into()) - } - } -} - -fn invoke_rs_tokenize( - vocab: *const llama_cpp_bindings_sys::llama_vocab, - text: *const c_char, - text_len: c_int, - tokens: *mut llama_cpp_bindings_sys::llama_token, - n_tokens_max: c_int, - add_bos: bool, -) -> Result { - let mut out_count: i32 = 0; - let mut out_error: *mut c_char = ptr::null_mut(); - let status = unsafe { - llama_cpp_bindings_sys::llama_rs_tokenize( - vocab, - text, - text_len, - tokens, - n_tokens_max, - add_bos, - true, - &raw mut out_count, - &raw mut out_error, - ) - }; - unsafe { tokenize_status_to_result(status, out_count, out_error) } -} - -fn checked_token_buffer_capacity(capacity: usize) -> Result { - Ok(c_int::try_from(capacity)?) -} - -fn checked_token_count(size: i32) -> Result { - Ok(usize::try_from(size)?) -} - -fn tokenize_into_buffer( - estimated_capacity: usize, - invoke: impl Fn( - *mut llama_cpp_bindings_sys::llama_token, - c_int, - ) -> Result, -) -> Result, StringToTokenError> { - let mut buffer: Vec = Vec::with_capacity(estimated_capacity); - let buffer_capacity = checked_token_buffer_capacity(buffer.capacity())?; - - let size = invoke( - buffer - .as_mut_ptr() - .cast::(), - buffer_capacity, - )?; - - let size = if size.is_negative() { - buffer.reserve_exact(checked_token_count(-size)?); - invoke( - buffer - .as_mut_ptr() - .cast::(), - -size, - )? - } else { - size - }; - - let size = checked_token_count(size)?; - - unsafe { buffer.set_len(size) } - - Ok(buffer) -} - -fn collect_synthetic_tool_call_renders( - without_tools_ptr: *const c_char, - with_tools_ptr: *const c_char, -) -> Result { - let without_tools = - read_optional_owned_cstr(without_tools_ptr)?.ok_or(crate::FfiContractError { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - detail: "success status contained a null render without tools", - })?; - let with_tools = read_optional_owned_cstr(with_tools_ptr)?.ok_or(crate::FfiContractError { - operation: "llama_rs_diagnose_tool_call_synthetic_renders", - detail: "success status contained a null render with tools", - })?; - - Ok(SyntheticToolCallRenders { - without_tools, - with_tools, - }) -} - -fn extract_meta_string( - c_function: TCFunction, - capacity: usize, -) -> Result -where - TCFunction: Fn(*mut c_char, usize) -> i32, -{ - let mut buffer = vec![0u8; capacity]; - let result = c_function(buffer.as_mut_ptr().cast::(), buffer.len()); - - if result < 0 { - return Err(MetaValError::NegativeReturn(result)); - } - - let returned_len = result.cast_unsigned() as usize; - - if returned_len >= capacity { - return extract_meta_string(c_function, returned_len + 1); - } - - if buffer.get(returned_len) != Some(&0) { - return Err(MetaValError::NegativeReturn(-1)); - } - - buffer.truncate(returned_len); - - Ok(String::from_utf8(buffer)?) -} - -impl Drop for LlamaModel { - fn drop(&mut self) { - unsafe { llama_cpp_bindings_sys::llama_model_free(self.model.as_ptr()) } - } -} - -#[cfg(test)] -mod extract_meta_string_tests { - use super::extract_meta_string; - use crate::MetaValError; - - #[test] - fn returns_error_when_null_terminator_missing() { - let result = extract_meta_string( - |buf_ptr, buf_len| { - let buffer = - unsafe { std::slice::from_raw_parts_mut(buf_ptr.cast::(), buf_len) }; - buffer[0] = b'a'; - buffer[1] = b'b'; - buffer[2] = b'c'; - 2 - }, - 4, - ); - - assert_eq!(result.unwrap_err(), MetaValError::NegativeReturn(-1)); - } - - #[test] - fn returns_error_for_negative_return_value() { - let result = extract_meta_string(|_buf_ptr, _buf_len| -5, 4); - - assert_eq!(result.unwrap_err(), MetaValError::NegativeReturn(-5)); - } - - #[test] - fn returns_error_for_invalid_utf8_data() { - let result = extract_meta_string( - |buf_ptr, buf_len| { - let buffer = - unsafe { std::slice::from_raw_parts_mut(buf_ptr.cast::(), buf_len) }; - buffer[0] = 0xFF; - buffer[1] = 0xFE; - buffer[2] = 0; - 2 - }, - 4, - ); - - assert!(result.is_err()); - assert!(result.unwrap_err().to_string().contains("FromUtf8Error")); - } - - #[test] - fn triggers_buffer_resize_when_returned_len_exceeds_capacity() { - let initial_capacity: usize = 4; - let length_exceeding_initial_capacity = 10; - let written_length = 2; - let call_count = std::cell::Cell::new(0); - let result = extract_meta_string( - |buf_ptr, buf_len| { - let count = call_count.get(); - call_count.set(count + 1); - if count == 0 { - length_exceeding_initial_capacity - } else { - let buffer = - unsafe { std::slice::from_raw_parts_mut(buf_ptr.cast::(), buf_len) }; - buffer[0] = b'h'; - buffer[1] = b'i'; - buffer[2] = 0; - written_length - } - }, - initial_capacity, - ); - - assert_eq!(result.unwrap(), "hi"); - } - - #[test] - fn cstring_with_validated_len_null_byte_returns_error() { - let result = super::cstring_with_validated_len("null\0byte"); - - assert!(result.is_err()); - } - - #[test] - fn validate_string_length_overflow_returns_error() { - let result = super::validate_string_length_for_tokenizer(usize::MAX); - - assert!(result.is_err()); - } - - #[test] - fn checked_token_buffer_capacity_overflow_returns_error() { - assert!(super::checked_token_buffer_capacity(usize::MAX).is_err()); - } - - #[test] - fn checked_token_buffer_capacity_in_range_returns_value() { - assert_eq!(super::checked_token_buffer_capacity(8), Ok(8)); - } - - #[test] - fn checked_token_count_negative_returns_error() { - assert!(super::checked_token_count(-1).is_err()); - } - - #[test] - fn checked_token_count_non_negative_returns_value() { - assert_eq!(super::checked_token_count(5), Ok(5)); - } - - #[test] - fn tokenize_into_buffer_single_pass_sets_length() { - let buffer = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| Ok(3)).unwrap(); - - assert_eq!(buffer.len(), 3); - } - - #[test] - fn tokenize_into_buffer_grows_buffer_when_first_pass_reports_negative_size() { - let call_count = std::cell::Cell::new(0); - let buffer = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| { - let count = call_count.get(); - call_count.set(count + 1); - if count == 0 { Ok(-20) } else { Ok(15) } - }) - .unwrap(); - - assert_eq!(buffer.len(), 15); - assert_eq!(call_count.get(), 2); - } - - #[test] - fn tokenize_into_buffer_propagates_invocation_error() { - let result = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| { - Err(crate::StringToTokenError::NotEnoughMemory) - }); - - assert_eq!(result, Err(crate::StringToTokenError::NotEnoughMemory)); - } - - #[test] - fn tokenize_into_buffer_propagates_second_invocation_error() { - let call_count = std::cell::Cell::new(0); - let result = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| { - let count = call_count.get(); - call_count.set(count + 1); - if count == 0 { - Ok(-20) - } else { - Err(crate::StringToTokenError::NotEnoughMemory) - } - }); - - assert_eq!(result, Err(crate::StringToTokenError::NotEnoughMemory)); - assert_eq!(call_count.get(), 2); - } - - #[test] - fn tokenize_into_buffer_negative_final_size_returns_conversion_error() { - let call_count = std::cell::Cell::new(0); - let result = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| { - let count = call_count.get(); - call_count.set(count + 1); - if count == 0 { Ok(-20) } else { Ok(-5) } - }); - - assert_eq!( - result.unwrap_err(), - crate::StringToTokenError::CIntConversionError(usize::try_from(-5i32).unwrap_err()) - ); - } - - #[test] - fn read_optional_owned_cstr_invalid_utf8_returns_error() { - let invalid_utf8_with_terminator: [u8; 3] = [0xFF, 0xFE, 0x00]; - let result = super::read_optional_owned_cstr( - invalid_utf8_with_terminator - .as_ptr() - .cast::(), - ); - - assert_eq!( - result.unwrap_err(), - crate::MarkerDetectionError::MarkerUtf8Error( - String::from_utf8(vec![0xFF, 0xFE]).unwrap_err() - ) - ); - } - - #[test] - fn collect_synthetic_tool_call_renders_first_invalid_utf8_returns_error() { - let invalid_utf8_with_terminator: [u8; 3] = [0xFF, 0xFE, 0x00]; - let valid_with_terminator: [u8; 3] = [b'o', b'k', 0x00]; - let result = super::collect_synthetic_tool_call_renders( - invalid_utf8_with_terminator - .as_ptr() - .cast::(), - valid_with_terminator.as_ptr().cast::(), - ); - - assert_eq!( - result.unwrap_err(), - crate::MarkerDetectionError::MarkerUtf8Error( - String::from_utf8(vec![0xFF, 0xFE]).unwrap_err() - ) - ); - } - - #[test] - fn collect_synthetic_tool_call_renders_second_invalid_utf8_returns_error() { - let valid_with_terminator: [u8; 3] = [b'o', b'k', 0x00]; - let invalid_utf8_with_terminator: [u8; 3] = [0xFF, 0xFE, 0x00]; - let result = super::collect_synthetic_tool_call_renders( - valid_with_terminator.as_ptr().cast::(), - invalid_utf8_with_terminator - .as_ptr() - .cast::(), - ); - - assert_eq!( - result.unwrap_err(), - crate::MarkerDetectionError::MarkerUtf8Error( - String::from_utf8(vec![0xFF, 0xFE]).unwrap_err() - ) - ); - } -} - -#[cfg(test)] -mod ffi_status_mapping_tests { - use std::ffi::c_char; - use std::mem::discriminant; - use std::path::Path; - use std::ptr; - - use llama_cpp_bindings_types::ParsedChatMessage; - use llama_cpp_bindings_types::ParsedToolCall; - use llama_cpp_bindings_types::ReasoningMarkers; - use llama_cpp_bindings_types::ToolCallArguments; - - use super::ReasoningSplit; - use super::chat_parser_create_status_to_result; - use super::chat_parser_free_status_to_result; - use super::compute_tool_call_haystack_status_to_result; - use super::cxx_exception_owns_out_error; - use super::detect_reasoning_markers_status_to_result; - use super::diagnose_tool_call_synthetic_renders_status_to_result; - use super::initialize_lora_adapter; - use super::load_model_from_file_status_to_result; - use super::outcome_from_via_ffi_result; - use super::parse_chat_message_status_to_result; - use super::parsed_chat_content_status_to_result; - use super::parsed_chat_free_status_to_result; - use super::parsed_chat_reasoning_content_status_to_result; - use super::parsed_chat_tool_call_arguments_status_to_result; - use super::parsed_chat_tool_call_count_status_to_result; - use super::parsed_chat_tool_call_id_status_to_result; - use super::parsed_chat_tool_call_name_status_to_result; - use super::reasoning_markers_free_status_to_result; - use super::restore_partial_reasoning; - use super::split_reasoning_prefix; - use super::tokenize_status_to_result; - use super::validate_reasoning_markers; - use crate::ChatMessageParseOutcome; - use crate::LlamaLoraAdapterInitError; - use crate::LlamaModelLoadError; - use crate::MarkerDetectionError; - use crate::ParseChatMessageError; - use crate::RawChatMessage; - use crate::StringToTokenError; - - #[test] - fn cxx_exception_owns_out_error_classifies_each_failure_variant() { - assert!(cxx_exception_owns_out_error::<()>(&Err( - MarkerDetectionError::ReasoningMarkerDetectionFailed { - message: String::new() - } - ))); - assert!(cxx_exception_owns_out_error::<()>(&Err( - MarkerDetectionError::ToolCallHaystackComputationFailed { - message: String::new() - } - ))); - assert!(cxx_exception_owns_out_error::<()>(&Err( - MarkerDetectionError::ToolCallSyntheticRenderDiagnosisFailed { - message: String::new() - } - ))); - assert!(!cxx_exception_owns_out_error::<()>(&Ok(()))); - } - - #[test] - fn lora_adapter_initialization_maps_a_null_pointer_to_unloadable() { - let result = initialize_lora_adapter(ptr::null_mut); - - assert_eq!(result.unwrap_err(), LlamaLoraAdapterInitError::Unloadable); - } - - #[test] - fn lora_adapter_initialization_owns_a_valid_pointer() { - let pointer = ptr::NonNull::dangling(); - let adapter = std::mem::ManuallyDrop::new( - initialize_lora_adapter(|| pointer.as_ptr()) - .expect("a non-null adapter pointer must initialize"), - ); - - assert_eq!(adapter.as_ptr(), pointer.as_ptr()); + if buffer.get(returned_len) != Some(&0) { + return Err(MetaValError::NegativeReturn(-1)); } - #[test] - fn load_model_success_with_null_model_is_contract_error() { - let result = unsafe { - load_model_from_file_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_OK, - ptr::null_mut(), - ptr::null_mut(), - Path::new("/some/path"), - ) - }; - - assert_eq!( - result.unwrap_err(), - LlamaModelLoadError::FfiContract(crate::FfiContractError { - operation: "llama_rs_load_model_from_file", - detail: "success status contained a null model", - }) - ); - } + buffer.truncate(returned_len); - #[test] - fn load_model_from_file_llama_cpp_returned_null_for_missing_path_is_file_not_found() { - let result = unsafe { - load_model_from_file_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_LLAMA_CPP_RETURNED_NULL, - ptr::null_mut(), - ptr::null_mut(), - Path::new("/definitely/missing/model.gguf"), - ) - }; + Ok(String::from_utf8(buffer)?) +} - assert_eq!( - result.unwrap_err(), - LlamaModelLoadError::FileNotFound( - Path::new("/definitely/missing/model.gguf").to_path_buf() - ) - ); +impl Drop for LlamaModel { + fn drop(&mut self) { + unsafe { llama_cpp_bindings_sys::llama_model_free(self.model.as_ptr()) } } +} - #[test] - fn load_model_from_file_allocation_failed_is_not_enough_memory() { - let result = unsafe { - load_model_from_file_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_ERROR_STRING_ALLOCATION_FAILED, - ptr::null_mut(), - ptr::null_mut(), - Path::new("/some/path"), - ) - }; - - assert_eq!(result.unwrap_err(), LlamaModelLoadError::NotEnoughMemory); - } +#[cfg(test)] +mod extract_meta_string_tests { + use super::extract_meta_string; + use crate::MetaValError; #[test] - fn load_model_from_file_cxx_exception_without_a_message_is_a_contract_error() { - let result = unsafe { - load_model_from_file_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_LLAMA_CPP_THREW_CXX_EXCEPTION, - ptr::null_mut(), - ptr::null_mut(), - Path::new("/some/path"), - ) - }; - - assert_eq!( - result.unwrap_err(), - crate::FfiContractError { - operation: "llama_rs_load_model_from_file", - detail: "reported a thrown C++ exception without an error message", - } - .into() + fn returns_error_when_null_terminator_missing() { + let result = extract_meta_string( + |buf_ptr, buf_len| { + let buffer = + unsafe { std::slice::from_raw_parts_mut(buf_ptr.cast::(), buf_len) }; + buffer[0] = b'a'; + buffer[1] = b'b'; + buffer[2] = b'c'; + 2 + }, + 4, ); - } - - #[test] - fn load_model_from_file_unknown_status_is_preserved() { - let result = unsafe { - load_model_from_file_status_to_result( - 255, - ptr::null_mut(), - ptr::null_mut(), - Path::new("/some/path"), - ) - }; - assert_eq!( - result.unwrap_err(), - LlamaModelLoadError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_load_model_from_file", - code: 255, - }) - ); + assert_eq!(result.unwrap_err(), MetaValError::NegativeReturn(-1)); } #[test] - fn parse_chat_message_success_with_null_handle_is_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_OK, - ptr::null_mut(), - &raw mut out_error, - ) - }; + fn returns_error_for_negative_return_value() { + let result = extract_meta_string(|_buf_ptr, _buf_len| -5, 4); - assert!(matches!( - result, - Err(ParseChatMessageError::FfiContract( - crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "success status contained a null parsed-chat handle", - } - )) - )); + assert_eq!(result.unwrap_err(), MetaValError::NegativeReturn(-5)); } #[test] - fn chat_parser_create_no_chat_template_maps_to_no_chat_template() { - let mut out_error: *mut c_char = ptr::null_mut(); - let result = unsafe { - chat_parser_create_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_MODEL_HAS_NO_CHAT_TEMPLATE, - ptr::null_mut(), - &raw mut out_error, - ) - }; - - assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NoChatTemplate) + fn returns_error_for_invalid_utf8_data() { + let result = extract_meta_string( + |buf_ptr, buf_len| { + let buffer = + unsafe { std::slice::from_raw_parts_mut(buf_ptr.cast::(), buf_len) }; + buffer[0] = 0xFF; + buffer[1] = 0xFE; + buffer[2] = 0; + 2 + }, + 4, ); - } - - #[test] - fn chat_parser_create_no_vocab_maps_to_no_vocab() { - let mut out_error: *mut c_char = ptr::null_mut(); - let result = unsafe { - chat_parser_create_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_MODEL_HAS_NO_VOCAB, - ptr::null_mut(), - &raw mut out_error, - ) - }; - assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NoVocab) - ); + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("FromUtf8Error")); } #[test] - fn chat_parser_create_allocation_failed_is_not_enough_memory() { - let mut out_error: *mut c_char = ptr::null_mut(); - let result = unsafe { - chat_parser_create_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_ERROR_STRING_ALLOCATION_FAILED, - ptr::null_mut(), - &raw mut out_error, - ) - }; - - assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NotEnoughMemory) + fn triggers_buffer_resize_when_returned_len_exceeds_capacity() { + let initial_capacity: usize = 4; + let length_exceeding_initial_capacity = 10; + let written_length = 2; + let call_count = std::cell::Cell::new(0); + let result = extract_meta_string( + |buf_ptr, buf_len| { + let count = call_count.get(); + call_count.set(count + 1); + if count == 0 { + length_exceeding_initial_capacity + } else { + let buffer = + unsafe { std::slice::from_raw_parts_mut(buf_ptr.cast::(), buf_len) }; + buffer[0] = b'h'; + buffer[1] = b'i'; + buffer[2] = 0; + written_length + } + }, + initial_capacity, ); - } - - #[test] - fn chat_parser_create_cxx_exception_is_parser_creation_failed_and_nulls_error() { - let mut out_error = unsafe { - llama_cpp_bindings_sys::llama_rs_string_dup(c"the parser could not be built".as_ptr()) - }; - let result = unsafe { - chat_parser_create_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_LLAMA_CPP_THREW_CXX_EXCEPTION, - ptr::null_mut(), - &raw mut out_error, - ) - }; - - let Err(ParseChatMessageError::ParserCreationFailed { message }) = result else { - panic!("the llama.cpp exception status must surface the wrapper message"); - }; - assert_eq!(message, "the parser could not be built"); - assert!( - out_error.is_null(), - "the reclaimed pointer must be nulled so the caller does not free it twice" - ); + assert_eq!(result.unwrap(), "hi"); } #[test] - fn chat_parser_create_success_with_null_parser_is_contract_error() { - let mut out_error: *mut c_char = ptr::null_mut(); - let result = unsafe { - chat_parser_create_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_OK, - ptr::null_mut(), - &raw mut out_error, - ) - }; + fn cstring_with_validated_len_null_byte_returns_error() { + let result = super::cstring_with_validated_len("null\0byte"); - assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::FfiContract( - crate::FfiContractError { - operation: "llama_rs_chat_parser_create", - detail: "success status contained a null parser handle", - } - )) - ); + assert!(result.is_err()); } #[test] - fn chat_parser_create_unknown_status_is_preserved() { - let mut out_error: *mut c_char = ptr::null_mut(); - let result = unsafe { - chat_parser_create_status_to_result(255, ptr::null_mut(), &raw mut out_error) - }; + fn validate_string_length_overflow_returns_error() { + let result = super::validate_string_length_for_tokenizer(usize::MAX); - assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_chat_parser_create", - code: 255, - })) - ); + assert!(result.is_err()); } #[test] - fn parse_chat_message_allocation_failed_is_not_enough_memory() { - 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_ERROR_STRING_ALLOCATION_FAILED, - ptr::null_mut(), - &raw mut out_error, - ) - }; - - assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NotEnoughMemory) - ); + fn checked_token_buffer_capacity_overflow_returns_error() { + assert!(super::checked_token_buffer_capacity(usize::MAX).is_err()); } #[test] - fn parse_chat_message_cxx_exception_is_message_unrecognized_and_nulls_error() { - let mut out_error = unsafe { - llama_cpp_bindings_sys::llama_rs_string_dup(c"the message could not be parsed".as_ptr()) - }; - let result = unsafe { - parse_chat_message_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_THREW_CXX_EXCEPTION, - ptr::null_mut(), - &raw mut out_error, - ) - }; + fn checked_token_buffer_capacity_in_range_returns_value() { + assert_eq!(super::checked_token_buffer_capacity(8), Ok(8)); + } - let Err(ParseChatMessageError::MessageUnrecognized { message }) = result else { - panic!("the llama.cpp exception status must surface the wrapper message"); - }; + #[test] + fn checked_token_count_negative_returns_error() { + assert!(super::checked_token_count(-1).is_err()); + } - assert_eq!(message, "the message could not be parsed"); - assert!( - out_error.is_null(), - "the reclaimed pointer must be nulled so the caller does not free it twice" - ); + #[test] + fn checked_token_count_non_negative_returns_value() { + assert_eq!(super::checked_token_count(5), Ok(5)); } #[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, - ) - }; + fn tokenize_into_buffer_single_pass_sets_length() { + let buffer = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| Ok(3)).unwrap(); - 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" - ); + assert_eq!(buffer.len(), 3); } #[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, - ) - }; + fn tokenize_into_buffer_grows_buffer_when_first_pass_reports_negative_size() { + let call_count = std::cell::Cell::new(0); + let buffer = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| { + let count = call_count.get(); + call_count.set(count + 1); + if count == 0 { Ok(-20) } else { Ok(15) } + }) + .unwrap(); - 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", - }) - ); + assert_eq!(buffer.len(), 15); + assert_eq!(call_count.get(), 2); } #[test] - fn parse_chat_message_unknown_status_is_preserved() { - let mut out_error: *mut c_char = ptr::null_mut(); - let result = unsafe { - parse_chat_message_status_to_result(255, ptr::null_mut(), &raw mut out_error) - }; + fn tokenize_into_buffer_propagates_invocation_error() { + let result = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| { + Err(crate::StringToTokenError::NotEnoughMemory) + }); - assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_parse_chat_message", - code: 255, - })) - ); + assert_eq!(result, Err(crate::StringToTokenError::NotEnoughMemory)); } #[test] - fn parsed_chat_content_success_with_null_string_is_contract_error() { - let result = unsafe { - parsed_chat_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_OK, - ptr::null_mut(), - ptr::null_mut(), - ) - }; + fn tokenize_into_buffer_propagates_second_invocation_error() { + let call_count = std::cell::Cell::new(0); + let result = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| { + let count = call_count.get(); + call_count.set(count + 1); + if count == 0 { + Ok(-20) + } else { + Err(crate::StringToTokenError::NotEnoughMemory) + } + }); - assert!(matches!( - result, - Err(ParseChatMessageError::FfiContract( - crate::FfiContractError { - operation: "llama_rs_parsed_chat_content", - detail: "success status contained a null string", - } - )) - )); + assert_eq!(result, Err(crate::StringToTokenError::NotEnoughMemory)); + assert_eq!(call_count.get(), 2); } #[test] - fn parsed_chat_content_allocation_failed_is_not_enough_memory() { - let result = unsafe { - parsed_chat_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_ERROR_STRING_ALLOCATION_FAILED, - ptr::null_mut(), - ptr::null_mut(), - ) - }; + fn tokenize_into_buffer_negative_final_size_returns_conversion_error() { + let call_count = std::cell::Cell::new(0); + let result = super::tokenize_into_buffer(8, |_tokens, _n_tokens_max| { + let count = call_count.get(); + call_count.set(count + 1); + if count == 0 { Ok(-20) } else { Ok(-5) } + }); assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NotEnoughMemory) + result.unwrap_err(), + crate::StringToTokenError::CIntConversionError(usize::try_from(-5i32).unwrap_err()) ); } #[test] - fn parsed_chat_content_cxx_exception_is_reported() { - let out_error = - unsafe { llama_cpp_bindings_sys::llama_rs_string_dup(c"content read failed".as_ptr()) }; - let result = unsafe { - parsed_chat_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_THREW_CXX_EXCEPTION, - ptr::null_mut(), - out_error, - ) - }; - - let Err(ParseChatMessageError::Reported { message }) = result else { - panic!("the llama.cpp exception status must surface the wrapper message"); - }; + fn read_optional_owned_cstr_invalid_utf8_returns_error() { + let invalid_utf8_with_terminator: [u8; 3] = [0xFF, 0xFE, 0x00]; + let result = super::read_optional_owned_cstr( + invalid_utf8_with_terminator + .as_ptr() + .cast::(), + ); - assert_eq!(message, "content read failed"); + assert_eq!( + result.unwrap_err(), + crate::MarkerDetectionError::MarkerUtf8Error( + String::from_utf8(vec![0xFF, 0xFE]).unwrap_err() + ) + ); } #[test] - fn parsed_chat_content_unknown_status_is_preserved() { - let result = - unsafe { parsed_chat_content_status_to_result(255, ptr::null_mut(), ptr::null_mut()) }; - - assert!(matches!( - result, - Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_content", - code: 255, - })) - )); - } + fn collect_synthetic_tool_call_renders_first_invalid_utf8_returns_error() { + let invalid_utf8_with_terminator: [u8; 3] = [0xFF, 0xFE, 0x00]; + let valid_with_terminator: [u8; 3] = [b'o', b'k', 0x00]; + let result = super::collect_synthetic_tool_call_renders( + invalid_utf8_with_terminator + .as_ptr() + .cast::(), + valid_with_terminator.as_ptr().cast::(), + ); - #[test] - fn parsed_chat_reasoning_content_success_with_null_string_is_contract_error() { - let result = unsafe { - parsed_chat_reasoning_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_OK, - ptr::null_mut(), - ptr::null_mut(), + assert_eq!( + result.unwrap_err(), + crate::MarkerDetectionError::MarkerUtf8Error( + String::from_utf8(vec![0xFF, 0xFE]).unwrap_err() ) - }; - - assert!(matches!( - result, - Err(ParseChatMessageError::FfiContract( - crate::FfiContractError { - operation: "llama_rs_parsed_chat_reasoning_content", - detail: "success status contained a null string", - } - )) - )); + ); } #[test] - fn parsed_chat_reasoning_content_allocation_failed_is_not_enough_memory() { - let result = unsafe { - parsed_chat_reasoning_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_ERROR_STRING_ALLOCATION_FAILED, - ptr::null_mut(), - ptr::null_mut(), - ) - }; + fn collect_synthetic_tool_call_renders_second_invalid_utf8_returns_error() { + let valid_with_terminator: [u8; 3] = [b'o', b'k', 0x00]; + let invalid_utf8_with_terminator: [u8; 3] = [0xFF, 0xFE, 0x00]; + let result = super::collect_synthetic_tool_call_renders( + valid_with_terminator.as_ptr().cast::(), + invalid_utf8_with_terminator + .as_ptr() + .cast::(), + ); assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NotEnoughMemory) + result.unwrap_err(), + crate::MarkerDetectionError::MarkerUtf8Error( + String::from_utf8(vec![0xFF, 0xFE]).unwrap_err() + ) ); } +} + +#[cfg(test)] +mod ffi_status_mapping_tests { + use std::ffi::c_char; + use std::mem::discriminant; + use std::path::Path; + use std::ptr; + + use llama_cpp_bindings_types::ReasoningMarkers; + + use super::chat_parser_create_status_to_result; + use super::chat_parser_free_status_to_result; + use super::compute_tool_call_haystack_status_to_result; + use super::cxx_exception_owns_out_error; + use super::detect_reasoning_markers_status_to_result; + use super::diagnose_tool_call_synthetic_renders_status_to_result; + use super::initialize_lora_adapter; + use super::load_model_from_file_status_to_result; + use super::reasoning_markers_free_status_to_result; + use super::tokenize_status_to_result; + use super::validate_reasoning_markers; + use crate::LlamaLoraAdapterInitError; + use crate::LlamaModelLoadError; + use crate::MarkerDetectionError; + use crate::ParseChatMessageError; + use crate::StringToTokenError; #[test] - fn parsed_chat_reasoning_content_cxx_exception_is_reported() { - let out_error = unsafe { - llama_cpp_bindings_sys::llama_rs_string_dup(c"reasoning read failed".as_ptr()) - }; - let result = unsafe { - parsed_chat_reasoning_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_THREW_CXX_EXCEPTION, - ptr::null_mut(), - out_error, - ) - }; - - let Err(ParseChatMessageError::Reported { message }) = result else { - panic!("the llama.cpp exception status must surface the wrapper message"); - }; - - assert_eq!(message, "reasoning read failed"); + fn cxx_exception_owns_out_error_classifies_each_failure_variant() { + assert!(cxx_exception_owns_out_error::<()>(&Err( + MarkerDetectionError::ReasoningMarkerDetectionFailed { + message: String::new() + } + ))); + assert!(cxx_exception_owns_out_error::<()>(&Err( + MarkerDetectionError::ToolCallHaystackComputationFailed { + message: String::new() + } + ))); + assert!(cxx_exception_owns_out_error::<()>(&Err( + MarkerDetectionError::ToolCallSyntheticRenderDiagnosisFailed { + message: String::new() + } + ))); + assert!(!cxx_exception_owns_out_error::<()>(&Ok(()))); } #[test] - fn parsed_chat_reasoning_content_unknown_status_is_preserved() { - let result = unsafe { - parsed_chat_reasoning_content_status_to_result(255, ptr::null_mut(), ptr::null_mut()) - }; + fn lora_adapter_initialization_maps_a_null_pointer_to_unloadable() { + let result = initialize_lora_adapter(ptr::null_mut); - assert!(matches!( - result, - Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_reasoning_content", - code: 255, - })) - )); + assert_eq!(result.unwrap_err(), LlamaLoraAdapterInitError::Unloadable); } #[test] - fn parsed_chat_tool_call_count_ok_returns_count() { - let result = unsafe { - parsed_chat_tool_call_count_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_OK, - 7, - ptr::null_mut(), - ) - }; + fn lora_adapter_initialization_owns_a_valid_pointer() { + let pointer = ptr::NonNull::dangling(); + let adapter = std::mem::ManuallyDrop::new( + initialize_lora_adapter(|| pointer.as_ptr()) + .expect("a non-null adapter pointer must initialize"), + ); - assert_eq!(result.unwrap(), 7); + assert_eq!(adapter.as_ptr(), pointer.as_ptr()); } #[test] - fn parsed_chat_tool_call_count_allocation_failed_is_not_enough_memory() { + fn load_model_success_with_null_model_is_contract_error() { let result = unsafe { - parsed_chat_tool_call_count_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_ERROR_STRING_ALLOCATION_FAILED, - 0, + load_model_from_file_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_OK, ptr::null_mut(), + ptr::null_mut(), + Path::new("/some/path"), ) }; assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NotEnoughMemory) + result.unwrap_err(), + LlamaModelLoadError::FfiContract(crate::FfiContractError { + operation: "llama_rs_load_model_from_file", + detail: "success status contained a null model", + }) ); } #[test] - fn parsed_chat_tool_call_count_cxx_exception_is_reported() { - let out_error = unsafe { - llama_cpp_bindings_sys::llama_rs_string_dup(c"tool-call count failed".as_ptr()) - }; - let result = unsafe { - parsed_chat_tool_call_count_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_THREW_CXX_EXCEPTION, - 0, - out_error, - ) - }; - - let Err(ParseChatMessageError::Reported { message }) = result else { - panic!("the llama.cpp exception status must surface the wrapper message"); - }; - - assert_eq!(message, "tool-call count failed"); - } - - #[test] - fn parsed_chat_tool_call_count_unknown_status_is_preserved() { - let result = - unsafe { parsed_chat_tool_call_count_status_to_result(255, 0, ptr::null_mut()) }; - - assert!(matches!( - result, - Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_tool_call_count", - code: 255, - })) - )); - } - - #[test] - fn parsed_chat_tool_call_id_success_with_null_string_is_contract_error() { + fn load_model_from_file_llama_cpp_returned_null_for_missing_path_is_file_not_found() { let result = unsafe { - parsed_chat_tool_call_id_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_OK, - 0, + load_model_from_file_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_LLAMA_CPP_RETURNED_NULL, ptr::null_mut(), ptr::null_mut(), + Path::new("/definitely/missing/model.gguf"), ) }; - assert!(matches!( - result, - Err(ParseChatMessageError::FfiContract( - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_id", - detail: "success status contained a null string", - } - )) - )); + assert_eq!( + result.unwrap_err(), + LlamaModelLoadError::FileNotFound( + Path::new("/definitely/missing/model.gguf").to_path_buf() + ) + ); } #[test] - fn parsed_chat_tool_call_id_out_of_bounds_carries_index() { + fn load_model_from_file_allocation_failed_is_not_enough_memory() { let result = unsafe { - parsed_chat_tool_call_id_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_INDEX_OUT_OF_BOUNDS, - 4, + load_model_from_file_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_ERROR_STRING_ALLOCATION_FAILED, ptr::null_mut(), ptr::null_mut(), + Path::new("/some/path"), ) }; - let Err(ParseChatMessageError::ToolCallIdIndexOutOfBounds { index }) = result else { - panic!("expected ToolCallIdIndexOutOfBounds, got {result:?}"); - }; - assert_eq!(index, 4); + assert_eq!(result.unwrap_err(), LlamaModelLoadError::NotEnoughMemory); } #[test] - fn parsed_chat_tool_call_id_allocation_failed_is_not_enough_memory() { + fn load_model_from_file_cxx_exception_without_a_message_is_a_contract_error() { let result = unsafe { - parsed_chat_tool_call_id_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_ERROR_STRING_ALLOCATION_FAILED, - 0, + load_model_from_file_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_LLAMA_CPP_THREW_CXX_EXCEPTION, ptr::null_mut(), ptr::null_mut(), + Path::new("/some/path"), ) }; assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NotEnoughMemory) + result.unwrap_err(), + crate::FfiContractError { + operation: "llama_rs_load_model_from_file", + detail: "reported a thrown C++ exception without an error message", + } + .into() ); } #[test] - fn parsed_chat_tool_call_id_cxx_exception_is_reported() { - let out_error = unsafe { - llama_cpp_bindings_sys::llama_rs_string_dup(c"tool-call id read failed".as_ptr()) - }; + fn load_model_from_file_unknown_status_is_preserved() { let result = unsafe { - parsed_chat_tool_call_id_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_THREW_CXX_EXCEPTION, - 0, + load_model_from_file_status_to_result( + 255, ptr::null_mut(), - out_error, + ptr::null_mut(), + Path::new("/some/path"), ) }; - let Err(ParseChatMessageError::Reported { message }) = result else { - panic!("the llama.cpp exception status must surface the wrapper message"); - }; - - assert_eq!(message, "tool-call id read failed"); - } - - #[test] - fn parsed_chat_tool_call_id_unknown_status_is_preserved() { - let result = unsafe { - parsed_chat_tool_call_id_status_to_result(255, 0, ptr::null_mut(), ptr::null_mut()) - }; - - assert!(matches!( - result, - Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_tool_call_id", + assert_eq!( + result.unwrap_err(), + LlamaModelLoadError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_load_model_from_file", code: 255, - })) - )); + }) + ); } #[test] - fn parsed_chat_tool_call_name_success_with_null_string_is_contract_error() { + fn chat_parser_create_no_chat_template_maps_to_no_chat_template() { + let mut out_error: *mut c_char = ptr::null_mut(); let result = unsafe { - parsed_chat_tool_call_name_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_OK, - 0, - ptr::null_mut(), + chat_parser_create_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_MODEL_HAS_NO_CHAT_TEMPLATE, ptr::null_mut(), + &raw mut out_error, ) }; - assert!(matches!( - result, - Err(ParseChatMessageError::FfiContract( - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_name", - detail: "success status contained a null string", - } - )) - )); + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NoChatTemplate) + ); } #[test] - fn parsed_chat_tool_call_name_out_of_bounds_carries_index() { + fn chat_parser_create_no_vocab_maps_to_no_vocab() { + let mut out_error: *mut c_char = ptr::null_mut(); let result = unsafe { - parsed_chat_tool_call_name_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_INDEX_OUT_OF_BOUNDS, - 2, - ptr::null_mut(), + chat_parser_create_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_MODEL_HAS_NO_VOCAB, ptr::null_mut(), + &raw mut out_error, ) }; - let Err(ParseChatMessageError::ToolCallNameIndexOutOfBounds { index }) = result else { - panic!("expected ToolCallNameIndexOutOfBounds, got {result:?}"); - }; - assert_eq!(index, 2); + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::NoVocab) + ); } #[test] - fn parsed_chat_tool_call_name_allocation_failed_is_not_enough_memory() { + fn chat_parser_create_allocation_failed_is_not_enough_memory() { + let mut out_error: *mut c_char = ptr::null_mut(); let result = unsafe { - parsed_chat_tool_call_name_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_ERROR_STRING_ALLOCATION_FAILED, - 0, - ptr::null_mut(), + chat_parser_create_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_ERROR_STRING_ALLOCATION_FAILED, ptr::null_mut(), + &raw mut out_error, ) }; @@ -3528,136 +2142,65 @@ mod ffi_status_mapping_tests { } #[test] - fn parsed_chat_tool_call_name_cxx_exception_is_reported() { - let out_error = unsafe { - llama_cpp_bindings_sys::llama_rs_string_dup(c"tool-call name read failed".as_ptr()) + fn chat_parser_create_cxx_exception_is_parser_creation_failed_and_nulls_error() { + let mut out_error = unsafe { + llama_cpp_bindings_sys::llama_rs_string_dup(c"the parser could not be built".as_ptr()) }; let result = unsafe { - parsed_chat_tool_call_name_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_THREW_CXX_EXCEPTION, - 0, + chat_parser_create_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_LLAMA_CPP_THREW_CXX_EXCEPTION, ptr::null_mut(), - out_error, + &raw mut out_error, ) }; - let Err(ParseChatMessageError::Reported { message }) = result else { + let Err(ParseChatMessageError::ParserCreationFailed { message }) = result else { panic!("the llama.cpp exception status must surface the wrapper message"); }; - assert_eq!(message, "tool-call name read failed"); - } - - #[test] - fn parsed_chat_tool_call_name_unknown_status_is_preserved() { - let result = unsafe { - parsed_chat_tool_call_name_status_to_result(255, 0, ptr::null_mut(), ptr::null_mut()) - }; - - assert!(matches!( - result, - Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_tool_call_name", - code: 255, - })) - )); + assert_eq!(message, "the parser could not be built"); + assert!( + out_error.is_null(), + "the reclaimed pointer must be nulled so the caller does not free it twice" + ); } #[test] - fn parsed_chat_tool_call_arguments_success_with_null_string_is_contract_error() { + fn chat_parser_create_success_with_null_parser_is_contract_error() { + let mut out_error: *mut c_char = ptr::null_mut(); let result = unsafe { - parsed_chat_tool_call_arguments_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_OK, - 0, - ptr::null_mut(), + chat_parser_create_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_CHAT_PARSER_CREATE_OK, ptr::null_mut(), + &raw mut out_error, ) }; - assert!(matches!( - result, - Err(ParseChatMessageError::FfiContract( + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::FfiContract( crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_arguments", - detail: "success status contained a null string", + operation: "llama_rs_chat_parser_create", + detail: "success status contained a null parser handle", } )) - )); - } - - #[test] - fn parsed_chat_tool_call_arguments_out_of_bounds_carries_index() { - let result = unsafe { - parsed_chat_tool_call_arguments_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_INDEX_OUT_OF_BOUNDS, - 9, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - - let Err(ParseChatMessageError::ToolCallArgumentsIndexOutOfBounds { index }) = result else { - panic!("expected ToolCallArgumentsIndexOutOfBounds, got {result:?}"); - }; - assert_eq!(index, 9); - } - - #[test] - fn parsed_chat_tool_call_arguments_allocation_failed_is_not_enough_memory() { - let result = unsafe { - parsed_chat_tool_call_arguments_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_ERROR_STRING_ALLOCATION_FAILED, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - - assert_eq!( - discriminant(&result.unwrap_err()), - discriminant(&ParseChatMessageError::NotEnoughMemory) ); } #[test] - fn parsed_chat_tool_call_arguments_cxx_exception_is_reported() { - let out_error = unsafe { - llama_cpp_bindings_sys::llama_rs_string_dup(c"tool-call arguments read failed".as_ptr()) - }; - let result = unsafe { - parsed_chat_tool_call_arguments_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_THREW_CXX_EXCEPTION, - 0, - ptr::null_mut(), - out_error, - ) - }; - - let Err(ParseChatMessageError::Reported { message }) = result else { - panic!("the llama.cpp exception status must surface the wrapper message"); - }; - - assert_eq!(message, "tool-call arguments read failed"); - } - - #[test] - fn parsed_chat_tool_call_arguments_unknown_status_is_preserved() { + fn chat_parser_create_unknown_status_is_preserved() { + let mut out_error: *mut c_char = ptr::null_mut(); let result = unsafe { - parsed_chat_tool_call_arguments_status_to_result( - 255, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) + chat_parser_create_status_to_result(255, ptr::null_mut(), &raw mut out_error) }; - assert!(matches!( - result, - Err(ParseChatMessageError::FfiStatus(crate::FfiStatusError { - operation: "llama_rs_parsed_chat_tool_call_arguments", + assert_eq!( + discriminant(&result.unwrap_err()), + discriminant(&ParseChatMessageError::FfiStatus(crate::FfiStatusError { + operation: "llama_rs_chat_parser_create", code: 255, })) - )); + ); } #[test] @@ -4171,301 +2714,6 @@ mod ffi_status_mapping_tests { ); } - #[test] - fn split_reasoning_prefix_without_markers_returns_content_up_to_tool_call_open() { - let ReasoningSplit { reasoning, content } = - split_reasoning_prefix("answerrest", None, Some(""), false); - - assert!(reasoning.is_empty()); - assert_eq!(content, "answer"); - } - - #[test] - fn split_reasoning_prefix_with_missing_open_marker_returns_content_only() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let ReasoningSplit { reasoning, content } = - split_reasoning_prefix("plain answer", Some(&markers), Some(""), false); - - assert!(reasoning.is_empty()); - assert_eq!(content, "plain answer"); - } - - #[test] - fn split_reasoning_prefix_with_missing_close_marker_returns_content_only() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let ReasoningSplit { reasoning, content } = - split_reasoning_prefix("unterminated", Some(&markers), Some(""), false); - - assert!(reasoning.is_empty()); - assert_eq!(content, "unterminated"); - } - - #[test] - fn split_reasoning_prefix_with_partial_unclosed_marker_returns_reasoning() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let ReasoningSplit { reasoning, content } = split_reasoning_prefix( - "prefixunfinishedtail", - Some(&markers), - Some(""), - true, - ); - - assert_eq!(reasoning, "unfinished"); - assert_eq!(content, "prefix"); - } - - #[test] - fn split_reasoning_prefix_without_tool_marker_preserves_all_partial_reasoning() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let ReasoningSplit { reasoning, content } = - split_reasoning_prefix("unfinished", Some(&markers), None, true); - - assert_eq!(reasoning, "unfinished"); - assert!(content.is_empty()); - } - - #[test] - fn split_reasoning_prefix_extracts_reasoning_and_trailing_content() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let ReasoningSplit { reasoning, content } = split_reasoning_prefix( - "deduceanswertail", - Some(&markers), - Some(""), - false, - ); - - assert_eq!(reasoning, "deduce"); - assert_eq!(content, "answer"); - } - - #[test] - fn restore_partial_reasoning_preserves_non_partial_parser_result() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let mut parsed = - ParsedChatMessage::new("parsed content".to_owned(), String::new(), Vec::new()); - - restore_partial_reasoning(&mut parsed, "unfinished", Some(&markers), false); - - assert_eq!(parsed.content, "parsed content"); - assert!(parsed.reasoning_content.is_empty()); - } - - #[test] - fn restore_partial_reasoning_preserves_existing_reasoning() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let mut parsed = ParsedChatMessage::new( - "parsed content".to_owned(), - "parsed reasoning".to_owned(), - Vec::new(), - ); - - restore_partial_reasoning(&mut parsed, "plain response", Some(&markers), true); - - assert_eq!(parsed.content, "parsed content"); - assert_eq!(parsed.reasoning_content, "parsed reasoning"); - } - - #[test] - fn restore_partial_reasoning_removes_open_marker_from_parser_result() { - let markers = ReasoningMarkers { - open: "\n[THINK]\n".to_owned(), - closes: vec!["[/THINK]".to_owned()], - }; - let mut parsed = ParsedChatMessage::new( - String::new(), - "[THINK]parsed reasoning".to_owned(), - Vec::new(), - ); - - restore_partial_reasoning(&mut parsed, "complete response", Some(&markers), true); - - assert!(parsed.content.is_empty()); - assert_eq!(parsed.reasoning_content, "parsed reasoning"); - } - - #[test] - fn restore_partial_reasoning_preserves_unclosed_reasoning_whitespace() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let mut parsed = - ParsedChatMessage::new(String::new(), "normalized reasoning".to_owned(), Vec::new()); - - restore_partial_reasoning(&mut parsed, "\n\nreasoning", Some(&markers), true); - - assert!(parsed.content.is_empty()); - assert_eq!(parsed.reasoning_content, "\n\nreasoning"); - } - - #[test] - fn restore_partial_reasoning_preserves_closed_reasoning_whitespace() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let mut parsed = ParsedChatMessage::new( - "answer".to_owned(), - "normalized reasoning".to_owned(), - Vec::new(), - ); - - restore_partial_reasoning( - &mut parsed, - "\n\nreasoninganswer", - Some(&markers), - true, - ); - - assert_eq!(parsed.content, "answer"); - assert_eq!(parsed.reasoning_content, "\n\nreasoning"); - } - - #[test] - fn restore_partial_reasoning_removes_open_marker_after_parser_whitespace() { - let markers = ReasoningMarkers { - open: "\n[THINK]\n".to_owned(), - closes: vec!["[/THINK]".to_owned()], - }; - let mut parsed = ParsedChatMessage::new( - String::new(), - "\n[THINK]parsed reasoning".to_owned(), - Vec::new(), - ); - - restore_partial_reasoning(&mut parsed, "complete response", Some(&markers), true); - - assert!(parsed.content.is_empty()); - assert_eq!(parsed.reasoning_content, "parsed reasoning"); - } - - #[test] - fn restore_partial_reasoning_preserves_result_without_open_marker() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let mut parsed = - ParsedChatMessage::new("parsed content".to_owned(), String::new(), Vec::new()); - - restore_partial_reasoning(&mut parsed, "unfinished", Some(&markers), true); - - assert_eq!(parsed.content, "parsed content"); - assert!(parsed.reasoning_content.is_empty()); - } - - #[test] - fn restore_partial_reasoning_recovers_unclosed_reasoning() { - let markers = ReasoningMarkers { - open: "".to_owned(), - closes: vec!["".to_owned()], - }; - let mut parsed = - ParsedChatMessage::new("unfinished".to_owned(), String::new(), Vec::new()); - - restore_partial_reasoning(&mut parsed, "unfinished", Some(&markers), true); - - assert!(parsed.content.is_empty()); - assert_eq!(parsed.reasoning_content, "unfinished"); - } - - #[test] - fn outcome_from_via_ffi_result_recognized_synthesizes_tool_call_ids() { - let parsed = ParsedChatMessage::new( - "answer".to_owned(), - String::new(), - vec![ParsedToolCall::new( - String::new(), - "tool".to_owned(), - ToolCallArguments::default(), - )], - ); - - let outcome = outcome_from_via_ffi_result(Ok(parsed), "answer", false); - - assert_eq!( - outcome.unwrap(), - ChatMessageParseOutcome::Recognized(ParsedChatMessage::new( - "answer".to_owned(), - String::new(), - vec![ParsedToolCall::new( - "call_0".to_owned(), - "tool".to_owned(), - ToolCallArguments::default(), - )], - )) - ); - } - - #[test] - fn outcome_from_via_ffi_result_message_unrecognized_is_unrecognized_with_raw_message() { - let outcome = outcome_from_via_ffi_result( - Err(ParseChatMessageError::MessageUnrecognized { - message: "boom".to_owned(), - }), - "garbled", - true, - ); - - assert_eq!( - outcome.unwrap(), - ChatMessageParseOutcome::Unrecognized(RawChatMessage { - text: "garbled".to_owned(), - is_partial: true, - ffi_error_message: "boom".to_owned(), - }) - ); - } - - #[test] - fn outcome_from_via_ffi_result_parser_creation_failure_propagates() { - let outcome = outcome_from_via_ffi_result( - Err(ParseChatMessageError::ParserCreationFailed { - message: "the parser could not be built".to_owned(), - }), - "garbled", - true, - ); - - assert_eq!( - discriminant(&outcome.unwrap_err()), - discriminant(&ParseChatMessageError::ParserCreationFailed { - message: String::new() - }) - ); - } - - #[test] - fn outcome_from_via_ffi_result_other_error_propagates() { - let outcome = outcome_from_via_ffi_result(Err(ParseChatMessageError::NoVocab), "x", false); - - assert_eq!( - discriminant(&outcome.unwrap_err()), - discriminant(&ParseChatMessageError::NoVocab) - ); - } - #[test] fn chat_parser_free_ok_is_success() { let result = unsafe { @@ -4545,85 +2793,6 @@ mod ffi_status_mapping_tests { ); } - #[test] - fn parsed_chat_free_ok_is_success() { - let result = unsafe { - parsed_chat_free_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_OK, - ptr::null_mut(), - ) - }; - - assert!( - result.is_ok(), - "a clean destructor must not report a failure" - ); - } - - #[test] - fn parsed_chat_free_allocation_failed_is_not_enough_memory() { - let result = unsafe { - parsed_chat_free_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_ERROR_STRING_ALLOCATION_FAILED, - ptr::null_mut(), - ) - }; - - let Err(ParseChatMessageError::NotEnoughMemory) = result else { - panic!("an error-string allocation failure must map to NotEnoughMemory"); - }; - } - - #[test] - fn parsed_chat_free_llama_cpp_out_of_memory_is_preserved() { - let result = unsafe { - parsed_chat_free_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_LLAMA_CPP_OUT_OF_MEMORY, - ptr::null_mut(), - ) - }; - - let Err(ParseChatMessageError::LlamaCppOutOfMemory) = result else { - panic!("a llama.cpp allocation failure must be reported as its own variant"); - }; - } - - #[test] - fn parsed_chat_free_destructor_threw_surfaces_the_message() { - let out_error = unsafe { - llama_cpp_bindings_sys::llama_rs_string_dup(c"the destructor threw".as_ptr()) - }; - let result = unsafe { - parsed_chat_free_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_FREE_DESTRUCTOR_THREW_CXX_EXCEPTION, - out_error, - ) - }; - - let Err(ParseChatMessageError::DestructorFailed { message }) = result else { - panic!("a throwing destructor must surface its message"); - }; - - assert_eq!(message, "the destructor threw"); - } - - #[test] - fn parsed_chat_free_unknown_status_is_preserved() { - let result = unsafe { parsed_chat_free_status_to_result(255, ptr::null_mut()) }; - - let Err(ParseChatMessageError::FfiStatus(status_error)) = result else { - panic!("an unrecognized status must be preserved verbatim"); - }; - - assert_eq!( - status_error, - crate::FfiStatusError { - operation: "llama_rs_parsed_chat_free", - code: 255, - } - ); - } - #[test] fn reasoning_markers_free_ok_is_success() { let result = unsafe { @@ -4703,13 +2872,6 @@ mod ffi_contract_status_tests { use super::detect_reasoning_markers_status_to_result; use super::diagnose_tool_call_synthetic_renders_status_to_result; use super::load_model_from_file_status_to_result; - use super::parse_chat_message_status_to_result; - use super::parsed_chat_content_status_to_result; - use super::parsed_chat_reasoning_content_status_to_result; - use super::parsed_chat_tool_call_arguments_status_to_result; - use super::parsed_chat_tool_call_count_status_to_result; - use super::parsed_chat_tool_call_id_status_to_result; - use super::parsed_chat_tool_call_name_status_to_result; use super::tokenize_status_to_result; use crate::error::apply_chat_template_error::ApplyChatTemplateError; use crate::error::llama_model_load_error::LlamaModelLoadError; @@ -4771,108 +2933,23 @@ mod ffi_contract_status_tests { Some( crate::FfiContractError { operation: "llama_rs_load_model_from_file", - detail: "was given a null out_error argument", - } - .into() - ) - ); - let outcome_3 = unsafe { - load_model_from_file_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_LLAMA_CPP_OUT_OF_MEMORY, - ptr::null_mut(), - ptr::null_mut(), - Path::new("/missing-for-contract-test.gguf"), - ) - }; - assert_eq!( - outcome_3.err(), - Some(LlamaModelLoadError::LlamaCppOutOfMemory) - ); - } - - #[test] - fn parse_chat_message_status_to_result_maps_every_contract_status() { - let mut out_error_slot: *mut c_char = ptr::null_mut(); - let outcome_0 = unsafe { - parse_chat_message_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_PARSER_ARG, - ptr::null_mut(), - &raw mut out_error_slot, - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { - panic!("LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_PARSER_ARG must map to a contract error"); - }; - assert_eq!( - contract_0, - crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "was given a null parser argument", - } - ); - let outcome_1 = unsafe { - parse_chat_message_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_INPUT_ARG, - ptr::null_mut(), - &raw mut out_error_slot, - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { - panic!("LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_INPUT_ARG must map to a contract error"); - }; - assert_eq!( - contract_1, - crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "was given a null input argument", - } - ); - let outcome_2 = unsafe { - parse_chat_message_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_HANDLE_ARG, - ptr::null_mut(), - &raw mut out_error_slot, - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_2)) = outcome_2 else { - panic!("LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_HANDLE_ARG must map to a contract error"); - }; - assert_eq!( - contract_2, - crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "was given a null out_handle argument", - } + detail: "was given a null out_error argument", + } + .into() + ) ); let outcome_3 = unsafe { - parse_chat_message_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_ERROR_ARG, + load_model_from_file_status_to_result( + llama_cpp_bindings_sys::LLAMA_RS_LOAD_MODEL_FROM_FILE_LLAMA_CPP_OUT_OF_MEMORY, ptr::null_mut(), - &raw mut out_error_slot, + ptr::null_mut(), + Path::new("/missing-for-contract-test.gguf"), ) }; - let Err(ParseChatMessageError::FfiContract(contract_3)) = outcome_3 else { - panic!("LLAMA_RS_PARSE_CHAT_MESSAGE_NULL_OUT_ERROR_ARG must map to a contract error"); - }; assert_eq!( - contract_3, - crate::FfiContractError { - operation: "llama_rs_parse_chat_message", - detail: "was given a null out_error argument", - } + outcome_3.err(), + Some(LlamaModelLoadError::LlamaCppOutOfMemory) ); - let outcome_4 = unsafe { - parse_chat_message_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY, - ptr::null_mut(), - &raw mut out_error_slot, - ) - }; - let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_4 else { - panic!( - "LLAMA_RS_PARSE_CHAT_MESSAGE_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" - ); - }; } #[test] @@ -5043,321 +3120,6 @@ mod ffi_contract_status_tests { ); } - #[test] - fn parsed_chat_content_status_to_result_maps_every_contract_status() { - let outcome_0 = unsafe { - parsed_chat_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_NULL_HANDLE_ARG, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { - panic!("LLAMA_RS_PARSED_CHAT_CONTENT_NULL_HANDLE_ARG must map to a contract error"); - }; - assert_eq!( - contract_0, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_content", - detail: "was given a null handle argument", - } - ); - let outcome_1 = unsafe { - parsed_chat_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_NULL_OUT_STRING_ARG, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { - panic!("LLAMA_RS_PARSED_CHAT_CONTENT_NULL_OUT_STRING_ARG must map to a contract error"); - }; - assert_eq!( - contract_1, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_content", - detail: "was given a null out_string argument", - } - ); - let outcome_2 = unsafe { - parsed_chat_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_OUT_OF_MEMORY, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { - panic!( - "LLAMA_RS_PARSED_CHAT_CONTENT_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" - ); - }; - } - - #[test] - fn parsed_chat_reasoning_content_status_to_result_maps_every_contract_status() { - let outcome_0 = unsafe { - parsed_chat_reasoning_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_HANDLE_ARG, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { - panic!( - "LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_HANDLE_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_0, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_reasoning_content", - detail: "was given a null handle argument", - } - ); - let outcome_1 = unsafe { - parsed_chat_reasoning_content_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_OUT_STRING_ARG, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { - panic!( - "LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_NULL_OUT_STRING_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_1, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_reasoning_content", - detail: "was given a null out_string argument", - } - ); - let outcome_2 = unsafe { - parsed_chat_reasoning_content_status_to_result(llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_OUT_OF_MEMORY, ptr::null_mut(), ptr::null_mut()) - }; - let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { - panic!( - "LLAMA_RS_PARSED_CHAT_REASONING_CONTENT_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" - ); - }; - } - - #[test] - fn parsed_chat_tool_call_count_status_to_result_maps_every_contract_status() { - let outcome_0 = unsafe { - parsed_chat_tool_call_count_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_HANDLE_ARG, - 0, - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_HANDLE_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_0, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_count", - detail: "was given a null handle argument", - } - ); - let outcome_1 = unsafe { - parsed_chat_tool_call_count_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_OUT_COUNT_ARG, - 0, - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_NULL_OUT_COUNT_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_1, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_count", - detail: "was given a null out_count argument", - } - ); - let outcome_2 = unsafe { - parsed_chat_tool_call_count_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_OUT_OF_MEMORY, - 0, - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_COUNT_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" - ); - }; - } - - #[test] - fn parsed_chat_tool_call_id_status_to_result_maps_every_contract_status() { - let outcome_0 = unsafe { - parsed_chat_tool_call_id_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_HANDLE_ARG, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_HANDLE_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_0, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_id", - detail: "was given a null handle argument", - } - ); - let outcome_1 = unsafe { - parsed_chat_tool_call_id_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_OUT_STRING_ARG, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_NULL_OUT_STRING_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_1, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_id", - detail: "was given a null out_string argument", - } - ); - let outcome_2 = unsafe { - parsed_chat_tool_call_id_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_OUT_OF_MEMORY, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ID_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" - ); - }; - } - - #[test] - fn parsed_chat_tool_call_name_status_to_result_maps_every_contract_status() { - let outcome_0 = unsafe { - parsed_chat_tool_call_name_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_HANDLE_ARG, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_HANDLE_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_0, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_name", - detail: "was given a null handle argument", - } - ); - let outcome_1 = unsafe { - parsed_chat_tool_call_name_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_OUT_STRING_ARG, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_NULL_OUT_STRING_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_1, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_name", - detail: "was given a null out_string argument", - } - ); - let outcome_2 = unsafe { - parsed_chat_tool_call_name_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_OUT_OF_MEMORY, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_NAME_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" - ); - }; - } - - #[test] - fn parsed_chat_tool_call_arguments_status_to_result_maps_every_contract_status() { - let outcome_0 = unsafe { - parsed_chat_tool_call_arguments_status_to_result( - llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_HANDLE_ARG, - 0, - ptr::null_mut(), - ptr::null_mut(), - ) - }; - let Err(ParseChatMessageError::FfiContract(contract_0)) = outcome_0 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_HANDLE_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_0, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_arguments", - detail: "was given a null handle argument", - } - ); - let outcome_1 = unsafe { - parsed_chat_tool_call_arguments_status_to_result(llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_OUT_STRING_ARG, 0, ptr::null_mut(), ptr::null_mut()) - }; - let Err(ParseChatMessageError::FfiContract(contract_1)) = outcome_1 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_NULL_OUT_STRING_ARG must map to a contract error" - ); - }; - assert_eq!( - contract_1, - crate::FfiContractError { - operation: "llama_rs_parsed_chat_tool_call_arguments", - detail: "was given a null out_string argument", - } - ); - let outcome_2 = unsafe { - parsed_chat_tool_call_arguments_status_to_result(llama_cpp_bindings_sys::LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_OUT_OF_MEMORY, 0, ptr::null_mut(), ptr::null_mut()) - }; - let Err(ParseChatMessageError::LlamaCppOutOfMemory) = outcome_2 else { - panic!( - "LLAMA_RS_PARSED_CHAT_TOOL_CALL_ARGUMENTS_LLAMA_CPP_OUT_OF_MEMORY must map to LlamaCppOutOfMemory" - ); - }; - } - #[test] fn detect_reasoning_markers_status_to_result_maps_every_contract_status() { let outcome_0 = unsafe { diff --git a/llama-cpp-bindings/src/mtmd.rs b/llama-cpp-bindings/src/mtmd.rs index 989cf7728..36f29a2d1 100644 --- a/llama-cpp-bindings/src/mtmd.rs +++ b/llama-cpp-bindings/src/mtmd.rs @@ -1,4 +1,4 @@ -pub mod image_chunk_batch_size_mismatch; +pub mod micro_batch_tokens; pub mod mtmd_bitmap; pub mod mtmd_bitmap_error; pub mod mtmd_context; @@ -16,8 +16,10 @@ pub mod mtmd_input_chunks; pub mod mtmd_input_chunks_error; pub mod mtmd_input_text; pub mod mtmd_tokenize_error; +pub mod non_causal_chunk_micro_batch_mismatch; +pub mod positive_batch_tokens; -pub use image_chunk_batch_size_mismatch::ImageChunkBatchSizeMismatch; +pub use micro_batch_tokens::micro_batch_tokens; pub use mtmd_bitmap::MtmdBitmap; pub use mtmd_bitmap_error::MtmdBitmapError; pub use mtmd_context::MtmdContext; @@ -35,3 +37,4 @@ pub use mtmd_input_chunks::MtmdInputChunks; pub use mtmd_input_chunks_error::MtmdInputChunksError; pub use mtmd_input_text::MtmdInputText; pub use mtmd_tokenize_error::MtmdTokenizeError; +pub use non_causal_chunk_micro_batch_mismatch::NonCausalChunkMicroBatchMismatch; diff --git a/llama-cpp-bindings/src/mtmd/image_chunk_batch_size_mismatch.rs b/llama-cpp-bindings/src/mtmd/image_chunk_batch_size_mismatch.rs deleted file mode 100644 index aff6affec..000000000 --- a/llama-cpp-bindings/src/mtmd/image_chunk_batch_size_mismatch.rs +++ /dev/null @@ -1,5 +0,0 @@ -#[derive(Debug, PartialEq, Eq)] -pub struct ImageChunkBatchSizeMismatch { - pub image_tokens: usize, - pub n_batch: i32, -} diff --git a/llama-cpp-bindings/src/mtmd/micro_batch_tokens.rs b/llama-cpp-bindings/src/mtmd/micro_batch_tokens.rs new file mode 100644 index 000000000..3fb97599d --- /dev/null +++ b/llama-cpp-bindings/src/mtmd/micro_batch_tokens.rs @@ -0,0 +1,10 @@ +use std::num::NonZeroU32; + +use crate::context::LlamaContext; + +/// Returns how many tokens a single decode of `n_batch` evaluates in `llama_ctx`: the smaller of +/// `n_batch` and the context's micro batch. +#[must_use] +pub fn micro_batch_tokens(llama_ctx: &LlamaContext, n_batch: NonZeroU32) -> u32 { + n_batch.get().min(llama_ctx.n_ubatch()) +} diff --git a/llama-cpp-bindings/src/mtmd/mtmd_eval_error.rs b/llama-cpp-bindings/src/mtmd/mtmd_eval_error.rs index bc40ade6f..9027b7426 100644 --- a/llama-cpp-bindings/src/mtmd/mtmd_eval_error.rs +++ b/llama-cpp-bindings/src/mtmd/mtmd_eval_error.rs @@ -1,5 +1,5 @@ -use crate::mtmd::image_chunk_batch_size_mismatch::ImageChunkBatchSizeMismatch; use crate::mtmd::mtmd_input_chunk_type_error::MtmdInputChunkTypeError; +use crate::mtmd::non_causal_chunk_micro_batch_mismatch::NonCausalChunkMicroBatchMismatch; #[derive(thiserror::Error, Debug, PartialEq, Eq)] pub enum MtmdEvalError { @@ -9,12 +9,14 @@ pub enum MtmdEvalError { FfiContract(#[from] crate::FfiContractError), #[error("batch size {requested} exceeds context batch size {context_max}")] BatchSizeExceedsContextLimit { requested: i32, context_max: u32 }, + #[error("batch size {requested} must be positive")] + NonPositiveBatchSize { requested: i32 }, #[error( - "image chunk has {} tokens but n_batch is {}", - .0.image_tokens, - .0.n_batch, + "a chunk decoded non-causally has {} tokens but a single decode fits {}", + .0.chunk_tokens, + .0.micro_batch_tokens, )] - ImageChunkExceedsBatchSize(ImageChunkBatchSizeMismatch), + NonCausalChunkExceedsMicroBatch(NonCausalChunkMicroBatchMismatch), #[error("multimodal chunk eval failed with code: {code}")] EvalFailed { code: i32 }, #[error("the chunk type could not be classified before evaluating it: {0}")] diff --git a/llama-cpp-bindings/src/mtmd/mtmd_input_chunk.rs b/llama-cpp-bindings/src/mtmd/mtmd_input_chunk.rs index 21b5c26da..35285060b 100644 --- a/llama-cpp-bindings/src/mtmd/mtmd_input_chunk.rs +++ b/llama-cpp-bindings/src/mtmd/mtmd_input_chunk.rs @@ -7,12 +7,14 @@ use crate::context::LlamaContext; use crate::token::LlamaToken; use llama_cpp_ffi_status::read_and_free_cpp_string; -use super::image_chunk_batch_size_mismatch::ImageChunkBatchSizeMismatch; +use super::micro_batch_tokens::micro_batch_tokens; use super::mtmd_context::MtmdContext; use super::mtmd_eval_error::MtmdEvalError; use super::mtmd_input_chunk_error::MtmdInputChunkError; use super::mtmd_input_chunk_type::MtmdInputChunkType; use super::mtmd_input_chunk_type_error::MtmdInputChunkTypeError; +use super::non_causal_chunk_micro_batch_mismatch::NonCausalChunkMicroBatchMismatch; +use super::positive_batch_tokens::positive_batch_tokens; /// # Safety /// @@ -99,18 +101,18 @@ fn eval_chunk_single_status_to_result( } } -fn image_chunk_batch_size_error( - is_image_chunk: bool, - chunk_token_count: usize, - n_batch: i32, +fn non_causal_chunk_micro_batch_error( + decodes_non_causally: bool, + chunk_tokens: usize, + micro_batch_tokens: u32, ) -> Option { - if is_image_chunk - && i64::try_from(chunk_token_count).is_ok_and(|tokens| tokens > i64::from(n_batch)) + if decodes_non_causally + && u64::try_from(chunk_tokens).is_ok_and(|tokens| tokens > u64::from(micro_batch_tokens)) { - return Some(MtmdEvalError::ImageChunkExceedsBatchSize( - ImageChunkBatchSizeMismatch { - image_tokens: chunk_token_count, - n_batch, + return Some(MtmdEvalError::NonCausalChunkExceedsMicroBatch( + NonCausalChunkMicroBatchMismatch { + chunk_tokens, + micro_batch_tokens, }, )); } @@ -187,12 +189,37 @@ impl MtmdInputChunk { Ok(Self { chunk, owned: true }) } + /// Checks that this chunk can be evaluated in decodes of `micro_batch_tokens`. llama.cpp + /// splits a causal chunk across decodes, while a media chunk it decodes non-causally has to + /// fit a single decode. + /// /// # Errors /// - /// Returns [`MtmdEvalError::ImageChunkExceedsBatchSize`] when this is an - /// image chunk whose token count exceeds `n_batch`. Returns - /// [`MtmdEvalError::EvalFailure`] if the underlying encode or decode step - /// fails. + /// Returns [`MtmdEvalError::NonCausalChunkExceedsMicroBatch`] when a media chunk decoded + /// non-causally has more tokens than `micro_batch_tokens`, or + /// [`MtmdEvalError::UnknownChunkType`] when the chunk type is unknown. + pub fn fit_to_micro_batch( + &self, + mtmd_ctx: &MtmdContext, + micro_batch_tokens: u32, + ) -> Result<(), MtmdEvalError> { + let decodes_non_causally = + self.chunk_type()? != MtmdInputChunkType::Text && mtmd_ctx.decode_use_non_causal(self); + + non_causal_chunk_micro_batch_error( + decodes_non_causally, + self.n_tokens(), + micro_batch_tokens, + ) + .map_or(Ok(()), Err) + } + + /// # Errors + /// + /// Returns [`MtmdEvalError::NonPositiveBatchSize`] when `n_batch` is not positive, + /// [`MtmdEvalError::NonCausalChunkExceedsMicroBatch`] when this chunk has to fit a single + /// decode but does not, or [`MtmdEvalError::EvalFailed`] if the underlying encode or decode + /// step fails. pub fn eval_single( &self, mtmd_ctx: &MtmdContext, @@ -202,15 +229,10 @@ impl MtmdInputChunk { n_batch: i32, logits_last: bool, ) -> Result { - let chunk_token_count = self.n_tokens(); - - if let Some(error) = image_chunk_batch_size_error( - self.chunk_type()? == MtmdInputChunkType::Image, - chunk_token_count, - n_batch, - ) { - return Err(error); - } + self.fit_to_micro_batch( + mtmd_ctx, + micro_batch_tokens(llama_ctx, positive_batch_tokens(n_batch)?), + )?; let mut final_position: llama_cpp_bindings_sys::llama_pos = start_position; let mut out_llama_cpp_return_code: i32 = 0; @@ -251,10 +273,10 @@ impl Drop for MtmdInputChunk { #[cfg(test)] mod unit_tests { use super::eval_chunk_single_status_to_result; - use super::image_chunk_batch_size_error; + use super::non_causal_chunk_micro_batch_error; use super::tokens_from_raw_ptr; - use crate::mtmd::image_chunk_batch_size_mismatch::ImageChunkBatchSizeMismatch; use crate::mtmd::mtmd_eval_error::MtmdEvalError; + use crate::mtmd::non_causal_chunk_micro_batch_mismatch::NonCausalChunkMicroBatchMismatch; #[test] fn tokens_from_raw_ptr_returns_none_for_null() { @@ -345,28 +367,26 @@ mod unit_tests { } #[test] - fn image_chunk_over_batch_size_reports_mismatch() { - let error = image_chunk_batch_size_error(true, 9, 4); - + fn a_non_causal_chunk_larger_than_the_micro_batch_reports_the_mismatch() { assert_eq!( - error, - Some(MtmdEvalError::ImageChunkExceedsBatchSize( - ImageChunkBatchSizeMismatch { - image_tokens: 9, - n_batch: 4, + non_causal_chunk_micro_batch_error(true, 9, 4), + Some(MtmdEvalError::NonCausalChunkExceedsMicroBatch( + NonCausalChunkMicroBatchMismatch { + chunk_tokens: 9, + micro_batch_tokens: 4, } )) ); } #[test] - fn non_image_chunk_never_reports_mismatch() { - assert!(image_chunk_batch_size_error(false, 9, 4).is_none()); + fn a_causal_chunk_larger_than_the_micro_batch_is_split_instead() { + assert!(non_causal_chunk_micro_batch_error(false, 9, 4).is_none()); } #[test] - fn image_chunk_within_batch_size_reports_no_mismatch() { - assert!(image_chunk_batch_size_error(true, 4, 4).is_none()); + fn a_non_causal_chunk_within_the_micro_batch_fits() { + assert!(non_causal_chunk_micro_batch_error(true, 4, 4).is_none()); } } diff --git a/llama-cpp-bindings/src/mtmd/mtmd_input_chunks.rs b/llama-cpp-bindings/src/mtmd/mtmd_input_chunks.rs index f592c42ce..22b431014 100644 --- a/llama-cpp-bindings/src/mtmd/mtmd_input_chunks.rs +++ b/llama-cpp-bindings/src/mtmd/mtmd_input_chunks.rs @@ -2,10 +2,12 @@ use std::ptr::NonNull; use crate::context::LlamaContext; +use super::micro_batch_tokens::micro_batch_tokens; use super::mtmd_context::MtmdContext; use super::mtmd_eval_error::MtmdEvalError; use super::mtmd_input_chunk::MtmdInputChunk; use super::mtmd_input_chunks_error::MtmdInputChunksError; +use super::positive_batch_tokens::positive_batch_tokens; const fn check_eval_result(result: i32) -> Result<(), MtmdEvalError> { if result == 0 { @@ -20,6 +22,8 @@ pub struct MtmdInputChunks { pub chunks: NonNull, } +unsafe impl Send for MtmdInputChunks {} + impl MtmdInputChunks { /// # Errors /// @@ -58,6 +62,32 @@ impl MtmdInputChunks { }) } + /// Checks every chunk with [`MtmdInputChunk::fit_to_micro_batch`] before any of them is + /// evaluated. + /// + /// # Errors + /// + /// Returns [`MtmdEvalError::NonCausalChunkExceedsMicroBatch`] for the first chunk decoded + /// non-causally that has more tokens than `micro_batch_tokens`, + /// [`MtmdEvalError::UnknownChunkType`] when a chunk type is unknown, or + /// [`MtmdEvalError::FfiContract`] when a chunk within the chunk count is null. + pub fn fit_to_micro_batch( + &self, + mtmd_ctx: &MtmdContext, + micro_batch_tokens: u32, + ) -> Result<(), MtmdEvalError> { + for index in 0..self.len() { + self.get(index) + .ok_or(crate::FfiContractError { + operation: "mtmd_input_chunks_get", + detail: "returned a null chunk within the chunk count", + })? + .fit_to_micro_batch(mtmd_ctx, micro_batch_tokens)?; + } + + Ok(()) + } + #[must_use] pub fn total_tokens(&self) -> usize { unsafe { llama_cpp_bindings_sys::mtmd_helper_get_n_tokens(self.chunks.as_ptr()) } @@ -68,9 +98,15 @@ impl MtmdInputChunks { unsafe { llama_cpp_bindings_sys::mtmd_helper_get_n_pos(self.chunks.as_ptr()) } } + /// Checks every chunk with [`Self::fit_to_micro_batch`] before evaluating any of them, so a + /// chunk that cannot be decoded leaves the KV cache untouched. + /// /// # Errors /// - /// Returns `MtmdEvalError::EvalFailure` if any encoding or decoding operation fails. + /// Returns [`MtmdEvalError::BatchSizeExceedsContextLimit`] when `n_batch` exceeds the + /// context's batch size, [`MtmdEvalError::NonPositiveBatchSize`] when it is not positive, + /// any error of [`Self::fit_to_micro_batch`], or [`MtmdEvalError::EvalFailed`] if any + /// encoding or decoding operation fails. pub fn eval_chunks( &self, mtmd_ctx: &MtmdContext, @@ -80,15 +116,18 @@ impl MtmdInputChunks { n_batch: i32, logits_last: bool, ) -> Result { + let batch_tokens = positive_batch_tokens(n_batch)?; let context_max_batch = llama_ctx.n_batch(); - if n_batch > 0 && n_batch.cast_unsigned() > context_max_batch { + if batch_tokens.get() > context_max_batch { return Err(MtmdEvalError::BatchSizeExceedsContextLimit { requested: n_batch, context_max: context_max_batch, }); } + self.fit_to_micro_batch(mtmd_ctx, micro_batch_tokens(llama_ctx, batch_tokens))?; + let mut final_position: llama_cpp_bindings_sys::llama_pos = start_position; let result = unsafe { diff --git a/llama-cpp-bindings/src/mtmd/non_causal_chunk_micro_batch_mismatch.rs b/llama-cpp-bindings/src/mtmd/non_causal_chunk_micro_batch_mismatch.rs new file mode 100644 index 000000000..06da799a4 --- /dev/null +++ b/llama-cpp-bindings/src/mtmd/non_causal_chunk_micro_batch_mismatch.rs @@ -0,0 +1,5 @@ +#[derive(Debug, PartialEq, Eq)] +pub struct NonCausalChunkMicroBatchMismatch { + pub chunk_tokens: usize, + pub micro_batch_tokens: u32, +} diff --git a/llama-cpp-bindings/src/mtmd/positive_batch_tokens.rs b/llama-cpp-bindings/src/mtmd/positive_batch_tokens.rs new file mode 100644 index 000000000..060aa6bd1 --- /dev/null +++ b/llama-cpp-bindings/src/mtmd/positive_batch_tokens.rs @@ -0,0 +1,37 @@ +use std::num::NonZeroU32; + +use super::mtmd_eval_error::MtmdEvalError; + +/// Returns `n_batch` as a positive token count. +/// +/// # Errors +/// +/// Returns [`MtmdEvalError::NonPositiveBatchSize`] when `n_batch` is not positive. +pub fn positive_batch_tokens(n_batch: i32) -> Result { + u32::try_from(n_batch) + .ok() + .and_then(NonZeroU32::new) + .ok_or(MtmdEvalError::NonPositiveBatchSize { requested: n_batch }) +} + +#[cfg(test)] +mod tests { + use super::positive_batch_tokens; + use crate::mtmd::mtmd_eval_error::MtmdEvalError; + + #[test] + fn a_zero_batch_is_rejected() { + assert_eq!( + positive_batch_tokens(0), + Err(MtmdEvalError::NonPositiveBatchSize { requested: 0 }) + ); + } + + #[test] + fn a_negative_batch_is_rejected() { + assert_eq!( + positive_batch_tokens(-1), + Err(MtmdEvalError::NonPositiveBatchSize { requested: -1 }) + ); + } +} diff --git a/llama-cpp-bindings/src/sampled_token_classifier.rs b/llama-cpp-bindings/src/sampled_token_classifier.rs index 2e511ca12..7db1e2f2d 100644 --- a/llama-cpp-bindings/src/sampled_token_classifier.rs +++ b/llama-cpp-bindings/src/sampled_token_classifier.rs @@ -7,21 +7,27 @@ use llama_cpp_bindings_sys::llama_seq_id; use llama_cpp_bindings_types::TokenUsage; use llama_cpp_bindings_types::TokenUsageError; +use crate::bare_json_tool_calls::BareJsonToolCalls; use crate::batch_add_error::BatchAddError; use crate::context::LlamaContext; use crate::error::EvalMultimodalChunksError; use crate::error::SampleError; use crate::error::TokenToStringError; use crate::eval_multimodal_chunks_params::EvalMultimodalChunksParams; +use crate::generation_progress::GenerationProgress; +use crate::json_probe_outcome::JsonProbeOutcome; use crate::llama_batch::LlamaBatch; use crate::model::LlamaModel; use crate::mtmd::MtmdContext; use crate::mtmd::MtmdInputChunks; +use crate::mtmd::micro_batch_tokens; +use crate::mtmd::positive_batch_tokens::positive_batch_tokens; use crate::sampled_token::SampledToken; use crate::sampling::LlamaSampler; -use crate::streaming_json_probe::JsonProbeOutcome; +use crate::streaming_json_probe::StreamingJsonProbe; use crate::streaming_markers::StreamingMarkers; use crate::token::LlamaToken; +use crate::token_piece::TokenPiece; pub use crate::classified_sample::ClassifiedSample; use crate::ingest_outcome::IngestOutcome; @@ -42,23 +48,22 @@ struct PendingToken { section_before_token: SampledTokenSection, marker_status: PendingMarkerStatus, is_from_prompt: bool, - is_held_for_probe: bool, -} - -#[derive(Clone, Debug, Eq, PartialEq)] -struct JsonProbeState { - held_text: String, } #[derive(Clone, Debug, Eq, PartialEq)] enum ProbeMode { Idle, - Active(JsonProbeState), + Active { + probe: StreamingJsonProbe, + held_count: usize, + }, } pub struct SampledTokenClassifier<'model> { model: &'model LlamaModel, markers: Arc, + marker_lookback: usize, + bare_json_tool_calls: BareJsonToolCalls, decoder: encoding_rs::Decoder, pending: VecDeque, section: SampledTokenSection, @@ -69,10 +74,16 @@ pub struct SampledTokenClassifier<'model> { impl<'model> SampledTokenClassifier<'model> { #[must_use] - pub fn new(model: &'model LlamaModel, markers: Arc) -> Self { + pub fn new( + model: &'model LlamaModel, + markers: Arc, + bare_json_tool_calls: BareJsonToolCalls, + ) -> Self { Self { model, + marker_lookback: markers.max_token_len().saturating_sub(1), markers, + bare_json_tool_calls, decoder: encoding_rs::UTF_8.new_decoder(), pending: VecDeque::new(), section: SampledTokenSection::Pending, @@ -82,46 +93,91 @@ impl<'model> SampledTokenClassifier<'model> { } } + /// Classifies a generated token and appends every outcome it finalises to + /// `outcomes`. An end-of-generation token is counted in the current section + /// without being detokenised, and releases every token still held back. + /// /// # Errors /// Returns [`TokenToStringError`] when the sampled token cannot be /// detokenised. The failure is surfaced rather than substituting an empty /// piece, so classification never silently drops generated text. - pub fn ingest(&mut self, token: LlamaToken) -> Result, TokenToStringError> { + pub fn ingest( + &mut self, + token: LlamaToken, + outcomes: &mut Vec, + ) -> Result { + if self.model.is_eog_token(&SampledToken::Content(token)) { + self.record_usage_in(self.section); + self.finish(outcomes); + + return Ok(GenerationProgress::Ended); + } + + let decoded = self.decode(token)?; + if self.markers.is_empty() { self.usage.record_undeterminable_token(); - let piece = self.decode(token)?; - return Ok(vec![IngestOutcome { + outcomes.push(IngestOutcome { sampled_token: SampledToken::Undeterminable(token), - visible_piece: piece.clone(), - raw_piece: piece, - }]); + piece: TokenPiece::Visible(decoded), + }); + + return Ok(GenerationProgress::Continues); } - let decoded = self.decode(token)?; self.pending.push_back(PendingToken { token, - decoded: decoded.clone(), + decoded, section: self.section, section_before_token: self.section, marker_status: PendingMarkerStatus::Unmatched, is_from_prompt: false, - is_held_for_probe: false, }); self.try_consume_marker_at_tail(); + self.classify_pending_tail(outcomes); + self.drain_overflow(outcomes); - let mut outcomes = self.classify_pending_tail(&decoded); - - outcomes.extend(self.drain_overflow()); - Ok(outcomes) + Ok(GenerationProgress::Continues) } - fn classify_pending_tail(&mut self, decoded: &str) -> Vec { - let probe_was_active = matches!(self.probe_mode, ProbeMode::Active(_)); - if probe_was_active && self.section_disengages_probe() { - self.abandon_probe() - } else { - self.update_probe(decoded) + fn classify_pending_tail(&mut self, outcomes: &mut Vec) { + let Some(tail) = self.pending.back() else { + return; + }; + let section_disengages_probe = self.section_disengages_probe(); + let probe_engages = + matches!(self.probe_mode, ProbeMode::Idle) && self.probe_engages_on(&tail.decoded); + + let probe_outcome = match &mut self.probe_mode { + ProbeMode::Active { .. } if section_disengages_probe => { + self.abandon_probe_held_before_tail(outcomes); + + return; + } + ProbeMode::Active { probe, held_count } => { + *held_count += 1; + + probe.feed(&tail.decoded) + } + ProbeMode::Idle if probe_engages => { + let mut probe = StreamingJsonProbe::default(); + let probe_outcome = probe.feed(&tail.decoded); + + self.probe_mode = ProbeMode::Active { + probe, + held_count: 1, + }; + + probe_outcome + } + ProbeMode::Idle => return, + }; + + match probe_outcome { + JsonProbeOutcome::StillPossiblyValid => {} + JsonProbeOutcome::CompletedValid => self.commit_probe_as_tool_call(outcomes), + JsonProbeOutcome::Failed => self.abandon_probe(outcomes), } } @@ -132,6 +188,15 @@ impl<'model> SampledTokenClassifier<'model> { ) } + fn probe_engages_on(&self, piece: &str) -> bool { + self.bare_json_tool_calls == BareJsonToolCalls::Detect + && matches!( + self.section, + SampledTokenSection::Content | SampledTokenSection::Pending + ) + && piece.trim_start().starts_with('{') + } + pub fn ingest_prompt_token(&mut self, token: LlamaToken) { if self.markers.is_empty() { return; @@ -144,11 +209,10 @@ impl<'model> SampledTokenClassifier<'model> { section_before_token: self.section, marker_status: PendingMarkerStatus::Unmatched, is_from_prompt: true, - is_held_for_probe: false, }); self.try_consume_marker_at_tail(); - self.drain_overflow(); + self.drain_overflow(&mut Vec::new()); } pub fn ingest_prompt_tokens(&mut self, tokens: &[LlamaToken]) { @@ -160,16 +224,16 @@ impl<'model> SampledTokenClassifier<'model> { } } - pub fn flush(&mut self) -> Vec { + /// Releases every generated token still held back, for a generation that + /// stops without an end-of-generation token (a token limit or a stop request). + pub fn finish(&mut self, outcomes: &mut Vec) { self.probe_mode = ProbeMode::Idle; - let mut outcomes = Vec::with_capacity(self.pending.len()); + while let Some(entry) = self.pending.pop_front() { - if entry.is_from_prompt { - continue; + if !entry.is_from_prompt { + outcomes.push(self.finalize_entry(entry)); } - outcomes.push(self.finalize_entry(entry)); } - outcomes } fn decode(&mut self, token: LlamaToken) -> Result { @@ -178,8 +242,10 @@ impl<'model> SampledTokenClassifier<'model> { } fn try_consume_marker_at_tail(&mut self) { - let pending_tokens: Vec<_> = self.pending.iter().map(|entry| entry.token).collect(); - let Some(marker) = self.markers.longest_matching_suffix(&pending_tokens) else { + let Some(marker) = self + .markers + .longest_matching_suffix(&self.pending.iter().rev().map(|entry| entry.token)) + else { return; }; let span_start = self.pending.len() - marker.tokens().len(); @@ -210,173 +276,122 @@ impl<'model> SampledTokenClassifier<'model> { self.section = next_section; } - fn drain_overflow(&mut self) -> Vec { - let lookback = self.markers.max_token_len().saturating_sub(1); - let mut outcomes = Vec::new(); + const fn probe_held_count(&self) -> usize { + match self.probe_mode { + ProbeMode::Active { held_count, .. } => held_count, + ProbeMode::Idle => 0, + } + } + + fn drain_overflow(&mut self, outcomes: &mut Vec) { + let probe_held_count = self.probe_held_count(); while let Some(front) = self.pending.front() { - if front.is_held_for_probe { + let drainable = self.pending.len().saturating_sub(probe_held_count); + + if drainable == 0 { break; } - let probe_held = self - .pending - .iter() - .filter(|entry| entry.is_held_for_probe) - .count(); - let drainable = self.pending.len().saturating_sub(probe_held); - let beyond_lookback = drainable > lookback; + + let beyond_lookback = drainable > self.marker_lookback; let resolved_boundary = matches!(front.marker_status, PendingMarkerStatus::ResolvedBoundary); + if !resolved_boundary && !beyond_lookback { break; } + let Some(entry) = self.pending.pop_front() else { break; }; - if entry.is_from_prompt { - continue; - } - outcomes.push(self.finalize_entry(entry)); - } - outcomes - } - - fn update_probe(&mut self, piece: &str) -> Vec { - let probe_active = matches!(self.probe_mode, ProbeMode::Active(_)); - if !probe_active { - if !self.section_allows_probe_engagement() { - return Vec::new(); - } - if !piece.trim_start().starts_with('{') { - return Vec::new(); + if !entry.is_from_prompt { + outcomes.push(self.finalize_entry(entry)); } - if let Some(entry) = self.pending.back_mut() { - entry.is_held_for_probe = true; - } - self.probe_mode = ProbeMode::Active(JsonProbeState { - held_text: piece.to_owned(), - }); - return self.evaluate_probe(); - } - - if let Some(entry) = self.pending.back_mut() { - entry.is_held_for_probe = true; } - if let ProbeMode::Active(state) = &mut self.probe_mode { - state.held_text.push_str(piece); - } - self.evaluate_probe() } - const fn section_allows_probe_engagement(&self) -> bool { - matches!( - self.section, - SampledTokenSection::Content | SampledTokenSection::Pending - ) - } + fn take_probe_held_tokens(&mut self) -> VecDeque { + let probe_held_count = self.probe_held_count(); - fn evaluate_probe(&mut self) -> Vec { - let outcome = match &self.probe_mode { - ProbeMode::Active(state) => JsonProbeOutcome::validate_prefix(&state.held_text), - ProbeMode::Idle => return Vec::new(), - }; - match outcome { - JsonProbeOutcome::StillPossiblyValid => Vec::new(), - JsonProbeOutcome::CompletedValid => self.commit_probe_as_tool_call(), - JsonProbeOutcome::Failed => self.abandon_probe(), - } + self.probe_mode = ProbeMode::Idle; + + self.pending + .split_off(self.pending.len().saturating_sub(probe_held_count)) } - fn commit_probe_as_tool_call(&mut self) -> Vec { - if !matches!(self.probe_mode, ProbeMode::Active(_)) { - return Vec::new(); - } - self.probe_mode = ProbeMode::Idle; + fn commit_probe_as_tool_call(&mut self, outcomes: &mut Vec) { self.section = SampledTokenSection::Content; - let drained: Vec<_> = self.pending.drain(..).collect(); - let mut outcomes = Vec::new(); - for mut entry in drained { - if entry.is_held_for_probe { - entry.section = SampledTokenSection::ToolCall; - entry.is_held_for_probe = false; - if !entry.is_from_prompt { - outcomes.push(self.finalize_entry(entry)); - } - } else { - self.pending.push_back(entry); - } + for mut entry in self.take_probe_held_tokens() { + entry.section = SampledTokenSection::ToolCall; + outcomes.push(self.finalize_entry(entry)); } - outcomes } - fn abandon_probe(&mut self) -> Vec { - if !matches!(self.probe_mode, ProbeMode::Active(_)) { - return Vec::new(); + fn abandon_probe(&mut self, outcomes: &mut Vec) { + for entry in self.take_probe_held_tokens() { + outcomes.push(self.finalize_entry(entry)); } - self.probe_mode = ProbeMode::Idle; + } - let drained: Vec<_> = self.pending.drain(..).collect(); - let mut outcomes = Vec::new(); - for mut entry in drained { - if entry.is_held_for_probe { - entry.is_held_for_probe = false; - if !entry.is_from_prompt { - outcomes.push(self.finalize_entry(entry)); - } - } else { - self.pending.push_back(entry); - } - } - outcomes + fn abandon_probe_held_before_tail(&mut self, outcomes: &mut Vec) { + let Some(tail) = self.pending.pop_back() else { + return; + }; + + self.abandon_probe(outcomes); + self.pending.push_back(tail); } - fn finalize_entry(&mut self, entry: PendingToken) -> IngestOutcome { - let section = entry.section; + const fn record_usage_in(&mut self, section: SampledTokenSection) { match section { SampledTokenSection::Reasoning => self.usage.record_reasoning_token(), SampledTokenSection::Content => self.usage.record_content_token(), SampledTokenSection::ToolCall => self.usage.record_tool_call_token(), SampledTokenSection::Pending => self.usage.record_undeterminable_token(), } + } + + fn finalize_entry(&mut self, entry: PendingToken) -> IngestOutcome { + self.record_usage_in(entry.section); - let sampled_token = match section { + let sampled_token = match entry.section { SampledTokenSection::Reasoning => SampledToken::Reasoning(entry.token), SampledTokenSection::Content => SampledToken::Content(entry.token), SampledTokenSection::ToolCall => SampledToken::ToolCall(entry.token), SampledTokenSection::Pending => SampledToken::Undeterminable(entry.token), }; - let visible_piece = if matches!(entry.marker_status, PendingMarkerStatus::Unmatched) { - entry.decoded.clone() - } else { - String::new() + let piece = match entry.marker_status { + PendingMarkerStatus::Unmatched => TokenPiece::Visible(entry.decoded), + PendingMarkerStatus::ResolvedBoundary | PendingMarkerStatus::AmbiguousBoundary => { + TokenPiece::Marker(entry.decoded) + } }; IngestOutcome { sampled_token, - visible_piece, - raw_piece: entry.decoded, + piece, } } + /// Samples a token and classifies it, appending the outcomes it finalises + /// to `outcomes`. + /// /// # Errors /// Forwards [`LlamaSampler::sample`] errors verbatim. Nothing is recorded on failure. - /// - /// Returns the sampled token (for downstream `batch.add` / `is_eog_token` - /// calls) alongside the outcomes that finalised this turn — see - /// [`Self::ingest`] for buffering semantics. pub fn sample( &mut self, sampler: &mut LlamaSampler, context: &LlamaContext, idx: i32, + outcomes: &mut Vec, ) -> Result { let token = sampler.sample(context, idx)?; - let outcomes = self.ingest(token)?; + let progress = self.ingest(token, outcomes)?; - Ok(ClassifiedSample { token, outcomes }) + Ok(ClassifiedSample { token, progress }) } /// # Errors @@ -434,9 +449,13 @@ impl<'model> SampledTokenClassifier<'model> { self.pending_prompt_tokens } + /// Checks every chunk with [`MtmdInputChunks::fit_to_micro_batch`] before evaluating any of + /// them, so a chunk that cannot be decoded moves neither the KV cache nor the usage counters. + /// /// # Errors - /// Returns [`EvalMultimodalChunksError::EvalFailed`] when the underlying - /// `eval_chunks` call fails (no counters move), + /// Returns [`EvalMultimodalChunksError::EvalFailed`] when `params.n_batch` is not positive or + /// a chunk cannot be decoded (both before any chunk is evaluated), or when evaluating a + /// chunk fails, /// [`EvalMultimodalChunksError::UnknownChunkType`] when a chunk reports a /// type unknown to this binding, or /// [`EvalMultimodalChunksError::ChunkOutOfBounds`] when a valid index returns @@ -448,6 +467,11 @@ impl<'model> SampledTokenClassifier<'model> { llama_ctx: &LlamaContext, params: EvalMultimodalChunksParams, ) -> Result { + chunks.fit_to_micro_batch( + mtmd_ctx, + micro_batch_tokens(llama_ctx, positive_batch_tokens(params.n_batch)?), + )?; + let chunk_count = chunks.len(); let mut next_position = params.start_position; @@ -515,16 +539,17 @@ impl<'model> SampledTokenClassifier<'model> { mod tests { use std::sync::Arc; - use super::JsonProbeState; use super::PendingMarkerStatus; use super::PendingToken; use super::ProbeMode; use super::SampledTokenClassifier; + use crate::bare_json_tool_calls::BareJsonToolCalls; use crate::ingest_outcome::IngestOutcome; use crate::marker_role::MarkerRole; use crate::marker_role_candidate::MarkerRoleCandidate; use crate::sampled_token::SampledToken; use crate::sampled_token_section::SampledTokenSection; + use crate::streaming_json_probe::StreamingJsonProbe; use crate::streaming_markers::StreamingMarkers; use crate::token::LlamaToken; @@ -557,6 +582,8 @@ mod tests { fn synthetic_classifier(markers: StreamingMarkers) -> SampledTokenClassifier<'static> { SampledTokenClassifier { model: unsafe { &*std::ptr::NonNull::::dangling().as_ptr() }, + marker_lookback: markers.max_token_len().saturating_sub(1), + bare_json_tool_calls: BareJsonToolCalls::Detect, markers: Arc::new(markers), decoder: encoding_rs::UTF_8.new_decoder(), pending: std::collections::VecDeque::new(), @@ -575,7 +602,6 @@ mod tests { section_before_token: classifier.section, marker_status: PendingMarkerStatus::Unmatched, is_from_prompt: false, - is_held_for_probe: false, }); } @@ -587,7 +613,6 @@ mod tests { section_before_token: classifier.section, marker_status: PendingMarkerStatus::Unmatched, is_from_prompt: true, - is_held_for_probe: false, }); } @@ -598,15 +623,28 @@ mod tests { ) -> Vec { push_pending(classifier, token_id, decoded); classifier.try_consume_marker_at_tail(); - let mut outcomes = classifier.classify_pending_tail(decoded); - outcomes.extend(classifier.drain_overflow()); + let mut outcomes = Vec::new(); + classifier.classify_pending_tail(&mut outcomes); + classifier.drain_overflow(&mut outcomes); + outcomes + } + + fn drained(classifier: &mut SampledTokenClassifier<'_>) -> Vec { + let mut outcomes = Vec::new(); + classifier.drain_overflow(&mut outcomes); + outcomes + } + + fn finished(classifier: &mut SampledTokenClassifier<'_>) -> Vec { + let mut outcomes = Vec::new(); + classifier.finish(&mut outcomes); outcomes } fn outcome_pieces(outcomes: &[IngestOutcome]) -> Vec<&str> { outcomes .iter() - .map(|outcome| outcome.visible_piece.as_str()) + .map(|outcome| outcome.piece.visible()) .collect() } @@ -697,11 +735,11 @@ mod tests { push_pending(&mut classifier, 300, ""); classifier.try_consume_marker_at_tail(); - let outcomes = classifier.drain_overflow(); + let outcomes = drained(&mut classifier); assert_eq!(classifier.section, SampledTokenSection::ToolCall); assert_eq!( @@ -731,11 +769,11 @@ mod tests { push_pending(&mut classifier, 300, ""); classifier.try_consume_marker_at_tail(); - let outcomes = classifier.drain_overflow(); + let outcomes = drained(&mut classifier); assert_eq!(classifier.section, SampledTokenSection::Content); assert_eq!( @@ -760,17 +798,17 @@ mod tests { push_pending(&mut classifier, 7, "step"); classifier.try_consume_marker_at_tail(); - let mut outcomes = classifier.drain_overflow(); + let mut outcomes = drained(&mut classifier); push_pending(&mut classifier, 200, ""); classifier.try_consume_marker_at_tail(); - outcomes.extend(classifier.drain_overflow()); + outcomes.extend(drained(&mut classifier)); push_pending(&mut classifier, 9, "Hi"); classifier.try_consume_marker_at_tail(); - outcomes.extend(classifier.drain_overflow()); + outcomes.extend(drained(&mut classifier)); - outcomes.extend(classifier.flush()); + outcomes.extend(finished(&mut classifier)); assert_eq!( outcome_sections(&outcomes), @@ -797,9 +835,9 @@ mod tests { for (id, decoded) in [(7, "r"), (200, ""), (9, "OK")] { push_pending(&mut classifier, id, decoded); classifier.try_consume_marker_at_tail(); - outcomes.extend(classifier.drain_overflow()); + outcomes.extend(drained(&mut classifier)); } - outcomes.extend(classifier.flush()); + outcomes.extend(finished(&mut classifier)); assert_eq!(outcome_pieces(&outcomes), vec!["r", "", "", "", "OK"]); assert_eq!(classifier.section, SampledTokenSection::Content); @@ -818,9 +856,9 @@ mod tests { for (id, decoded) in [(7, "r"), (200, "a"), (201, "b"), (300, "x")] { push_pending(&mut classifier, id, decoded); classifier.try_consume_marker_at_tail(); - outcomes.extend(classifier.drain_overflow()); + outcomes.extend(drained(&mut classifier)); } - outcomes.extend(classifier.flush()); + outcomes.extend(finished(&mut classifier)); assert_eq!(outcome_pieces(&outcomes), vec!["r", "a", "b", "x"]); assert!(outcomes.iter().all(|outcome| { @@ -840,9 +878,9 @@ mod tests { for (id, decoded) in [(100, ""), (200, ""), (9, "Hi")] { push_pending(&mut classifier, id, decoded); classifier.try_consume_marker_at_tail(); - outcomes.extend(classifier.drain_overflow()); + outcomes.extend(drained(&mut classifier)); } - outcomes.extend(classifier.flush()); + outcomes.extend(finished(&mut classifier)); assert_eq!( outcome_sections(&outcomes), @@ -864,7 +902,7 @@ mod tests { push_pending(&mut classifier, 200, ""); classifier.try_consume_marker_at_tail(); - let outcomes = classifier.drain_overflow(); + let outcomes = drained(&mut classifier); assert_eq!( outcome_sections(&outcomes), @@ -899,7 +937,7 @@ mod tests { push_pending(&mut classifier, 400, ""); classifier.try_consume_marker_at_tail(); - let outcomes = classifier.drain_overflow(); + let outcomes = drained(&mut classifier); assert_eq!( outcome_sections(&outcomes), @@ -909,7 +947,7 @@ mod tests { } #[test] - fn flush_drains_remaining_pending_at_eog() { + fn finish_releases_every_held_generated_token() { let markers = markers_with( Some(vec![token(100)]), Some(vec![token(200), token(201), token(202)]), @@ -921,7 +959,7 @@ mod tests { push_pending(&mut classifier, 200, ""); + assert_eq!(outcomes[0].piece.visible(), ""); + assert_eq!(outcomes[0].piece.raw(), "k>"); assert_eq!(classifier.section, SampledTokenSection::Content); assert_eq!(classifier.usage().reasoning_tokens, 1); @@ -1065,7 +1102,7 @@ mod tests { for token_id in [100, 7, 200, 100, 8, 200] { push_pending_from_prompt(&mut classifier, token_id); classifier.try_consume_marker_at_tail(); - classifier.drain_overflow(); + drained(&mut classifier); } assert_eq!(classifier.section, SampledTokenSection::Content); @@ -1091,7 +1128,7 @@ mod tests { for token_id in [100, 7, 200] { push_pending_from_prompt(&mut classifier, token_id); classifier.try_consume_marker_at_tail(); - classifier.drain_overflow(); + drained(&mut classifier); } assert_eq!(classifier.section, SampledTokenSection::Content); @@ -1105,17 +1142,16 @@ mod tests { section_before_token: classifier.section, marker_status: PendingMarkerStatus::Unmatched, is_from_prompt: false, - is_held_for_probe: false, }); classifier.try_consume_marker_at_tail(); - let outcomes = classifier.drain_overflow(); + let outcomes = drained(&mut classifier); assert_eq!(outcomes.len(), 1); assert_eq!( std::mem::discriminant(&outcomes[0].sampled_token), std::mem::discriminant(&SampledToken::Content(LlamaToken::new(0))) ); - assert_eq!(outcomes[0].visible_piece, "hi"); + assert_eq!(outcomes[0].piece.visible(), "hi"); assert_eq!(classifier.usage().content_tokens, 1); assert_eq!(classifier.usage().reasoning_tokens, 0); assert_eq!(classifier.usage().undeterminable_tokens, 0); @@ -1131,9 +1167,9 @@ mod tests { for (id, decoded) in [(7, "hi"), (200, ""), (8, "ok")] { push_pending(&mut classifier, id, decoded); classifier.try_consume_marker_at_tail(); - outcomes.extend(classifier.drain_overflow()); + outcomes.extend(drained(&mut classifier)); } - outcomes.extend(classifier.flush()); + outcomes.extend(finished(&mut classifier)); assert_eq!( outcome_sections(&outcomes), @@ -1157,9 +1193,9 @@ mod tests { for (id, decoded) in [(7, "step1"), (100, ""), (8, "step2")] { push_pending(&mut classifier, id, decoded); classifier.try_consume_marker_at_tail(); - outcomes.extend(classifier.drain_overflow()); + outcomes.extend(drained(&mut classifier)); } - outcomes.extend(classifier.flush()); + outcomes.extend(finished(&mut classifier)); assert_eq!(outcome_pieces(&outcomes), vec!["step1", "", "step2"]); assert_eq!(classifier.section, SampledTokenSection::Reasoning); @@ -1235,7 +1271,7 @@ mod tests { push_pending(&mut classifier, 300, ""); classifier.try_consume_marker_at_tail(); - let outcomes = classifier.drain_overflow(); + let outcomes = drained(&mut classifier); assert_eq!( outcome_sections(&outcomes), @@ -1477,6 +1513,19 @@ mod tests { assert_eq!(classifier.probe_mode, ProbeMode::Idle); } + #[test] + fn json_probe_does_not_engage_when_bare_json_tool_calls_are_ignored() { + let markers = markers_with_tool_call_open(vec![token(900)]); + let mut classifier = synthetic_classifier(markers); + classifier.bare_json_tool_calls = BareJsonToolCalls::Ignore; + classifier.section = SampledTokenSection::Content; + + let outcomes = push_and_probe(&mut classifier, 1, "{"); + + assert_eq!(classifier.probe_mode, ProbeMode::Idle); + assert_eq!(outcome_pieces(&outcomes), vec!["{"]); + } + #[test] fn json_probe_does_not_engage_in_tool_call_section() { let markers = markers_with_tool_call_open(vec![token(900)]); @@ -1602,7 +1651,7 @@ mod tests { } #[test] - fn flush_during_active_json_probe_releases_held_tokens_as_content() { + fn finish_during_active_json_probe_releases_held_tokens_as_content() { let markers = markers_with_tool_call_open(vec![token(900)]); let mut classifier = synthetic_classifier(markers); classifier.section = SampledTokenSection::Content; @@ -1611,48 +1660,18 @@ mod tests { push_and_probe(&mut classifier, 2, r#""name""#); assert_ne!(classifier.probe_mode, ProbeMode::Idle); - let outcomes = classifier.flush(); + let outcomes = finished(&mut classifier); let sections = outcome_sections(&outcomes); assert!( sections .iter() .all(|section| *section == SampledTokenSection::Content), - "mid-probe flush must release held tokens as Content, got {sections:?}", + "finishing mid-probe must release held tokens as Content, got {sections:?}", ); assert_eq!(classifier.probe_mode, ProbeMode::Idle); } - #[test] - fn evaluate_probe_while_idle_returns_no_outcomes() { - let markers = markers_with_tool_call_open(vec![token(900)]); - let mut classifier = synthetic_classifier(markers); - - let outcomes = classifier.evaluate_probe(); - - assert!(outcomes.is_empty()); - } - - #[test] - fn commit_probe_as_tool_call_while_idle_returns_no_outcomes() { - let markers = markers_with_tool_call_open(vec![token(900)]); - let mut classifier = synthetic_classifier(markers); - - let outcomes = classifier.commit_probe_as_tool_call(); - - assert!(outcomes.is_empty()); - } - - #[test] - fn abandon_probe_while_idle_returns_no_outcomes() { - let markers = markers_with_tool_call_open(vec![token(900)]); - let mut classifier = synthetic_classifier(markers); - - let outcomes = classifier.abandon_probe(); - - assert!(outcomes.is_empty()); - } - #[test] fn commit_probe_as_tool_call_requeues_non_held_entries_and_releases_held_as_tool_call() { let markers = markers_with_tool_call_open(vec![token(900)]); @@ -1666,7 +1685,6 @@ mod tests { section_before_token: SampledTokenSection::Content, marker_status: PendingMarkerStatus::Unmatched, is_from_prompt: false, - is_held_for_probe: false, }); classifier.pending.push_back(PendingToken { token: token(2), @@ -1675,13 +1693,14 @@ mod tests { section_before_token: SampledTokenSection::Content, marker_status: PendingMarkerStatus::Unmatched, is_from_prompt: false, - is_held_for_probe: true, - }); - classifier.probe_mode = ProbeMode::Active(JsonProbeState { - held_text: "{}".to_owned(), }); + classifier.probe_mode = ProbeMode::Active { + probe: StreamingJsonProbe::default(), + held_count: 1, + }; - let outcomes = classifier.commit_probe_as_tool_call(); + let mut outcomes = Vec::new(); + classifier.commit_probe_as_tool_call(&mut outcomes); let sections = outcome_sections(&outcomes); assert_eq!(sections, vec![SampledTokenSection::ToolCall]); diff --git a/llama-cpp-bindings/src/sampling.rs b/llama-cpp-bindings/src/sampling.rs index cd9d3188d..56269c3a2 100644 --- a/llama-cpp-bindings/src/sampling.rs +++ b/llama-cpp-bindings/src/sampling.rs @@ -15,6 +15,7 @@ use crate::token::logit_bias::LlamaLogitBias; use crate::{GrammarError, SampleError, SamplerAcceptError, SamplingError}; use llama_cpp_ffi_status::read_and_free_cpp_string; use llama_cpp_gbnf::gbnf_validation_error::GbnfValidationError; +use llama_cpp_gbnf::validate_gbnf::validate_gbnf; fn check_sampler_accept_status( status: llama_cpp_bindings_sys::llama_rs_sampler_accept_status, @@ -221,6 +222,8 @@ pub struct LlamaSampler { sampler: NonNull, } +unsafe impl Send for LlamaSampler {} + fn grammar_callback_error_to_result(error: Option) -> Result<(), SampleError> { error.map_or(Ok(()), |recorded| { Err(SampleError::GrammarCallbackFailed { @@ -475,8 +478,8 @@ impl LlamaSampler { grammar_root: &str, ) -> Result { let SanitizedGrammar { - grammar: grammar_str, - root: grammar_root, + grammar: grammar_cstring, + root: root_cstring, } = Self::sanitize_grammar_strings(grammar_str, grammar_root)?; let mut sampler: *mut llama_cpp_bindings_sys::llama_sampler = std::ptr::null_mut(); let mut error_ptr: *mut c_char = std::ptr::null_mut(); @@ -484,14 +487,16 @@ impl LlamaSampler { let status = unsafe { llama_cpp_bindings_sys::llama_rs_sampler_init_grammar( model.vocab_ptr(), - grammar_str.as_ptr(), - grammar_root.as_ptr(), + grammar_cstring.as_ptr(), + root_cstring.as_ptr(), &raw mut sampler, &raw mut error_ptr, ) }; - sampler_init_grammar_status_to_result(status, sampler, error_ptr) + sampler_init_grammar_status_to_result(status, sampler, error_ptr).map_err(|init_error| { + Self::diagnose_rejected_grammar(model, grammar_str, grammar_root, init_error) + }) } /// # Errors @@ -504,8 +509,8 @@ impl LlamaSampler { trigger_tokens: &[LlamaToken], ) -> Result { let SanitizedGrammar { - grammar: grammar_str, - root: grammar_root, + grammar: grammar_cstring, + root: root_cstring, } = Self::sanitize_grammar_strings(grammar_str, grammar_root)?; let trigger_patterns = Self::sanitize_trigger_patterns(trigger_patterns)?; let mut sampler: *mut llama_cpp_bindings_sys::llama_sampler = std::ptr::null_mut(); @@ -517,8 +522,8 @@ impl LlamaSampler { let status = unsafe { llama_cpp_bindings_sys::llama_rs_sampler_init_grammar_lazy_patterns( model.vocab_ptr(), - grammar_str.as_ptr(), - grammar_root.as_ptr(), + grammar_cstring.as_ptr(), + root_cstring.as_ptr(), trigger_pattern_ptrs.as_mut_ptr(), trigger_pattern_ptrs.len(), trigger_tokens.as_ptr().cast(), @@ -528,7 +533,11 @@ impl LlamaSampler { ) }; - sampler_init_grammar_lazy_patterns_status_to_result(status, sampler, error_ptr) + sampler_init_grammar_lazy_patterns_status_to_result(status, sampler, error_ptr).map_err( + |init_error| { + Self::diagnose_rejected_grammar(model, grammar_str, grammar_root, init_error) + }, + ) } /// # Errors @@ -546,23 +555,34 @@ impl LlamaSampler { grammar_str: &str, grammar_root: &str, ) -> Result { - match llama_cpp_gbnf::validate_gbnf::validate_gbnf(grammar_str, grammar_root) { - Ok(()) => {} - Err(GbnfValidationError::RootSymbolMissing { .. }) => { - return Err(GrammarError::RootNotFound); - } - Err(GbnfValidationError::GrammarContainsNul(nul_error)) => { - return Err(GrammarError::GrammarContainsNul(nul_error)); - } - Err(rejected) => return Err(GrammarError::GrammarRejected(rejected)), - } - Ok(SanitizedGrammar { grammar: CString::new(grammar_str).map_err(GrammarError::GrammarContainsNul)?, - root: CString::new(grammar_root).map_err(GrammarError::GrammarContainsNul)?, + root: CString::new(grammar_root).map_err(GrammarError::RootContainsNul)?, }) } + fn diagnose_rejected_grammar( + model: &LlamaModel, + grammar_str: &str, + grammar_root: &str, + init_error: GrammarError, + ) -> GrammarError { + if !matches!( + init_error, + GrammarError::GrammarMalformed | GrammarError::LazyGrammarMalformed + ) { + return init_error; + } + + let vocab = unsafe { &*model.vocab_ptr() }; + + match validate_gbnf(Some(vocab), grammar_str, grammar_root) { + Ok(()) => init_error, + Err(GbnfValidationError::RootSymbolMissing { .. }) => GrammarError::RootNotFound, + Err(rejected) => GrammarError::GrammarRejected(rejected), + } + } + fn sanitize_trigger_patterns( trigger_patterns: &[String], ) -> Result, GrammarError> { @@ -780,10 +800,12 @@ mod tests { } #[test] - fn sanitize_grammar_strings_root_not_found() { + fn sanitize_grammar_strings_null_byte_in_root_only() { assert_eq!( - LlamaSampler::sanitize_grammar_strings("expr ::= \"hello\"", "root"), - Err(GrammarError::RootNotFound) + LlamaSampler::sanitize_grammar_strings("root ::= \"hello\"", "ro\0ot"), + Err(GrammarError::RootContainsNul( + CString::new("ro\0ot").expect_err("the root carries a nul byte") + )) ); } @@ -1154,15 +1176,6 @@ mod tests { }) ); } - - #[test] - fn grammar_returns_root_not_found_before_touching_model() { - let model = unsafe { &*std::ptr::NonNull::::dangling().as_ptr() }; - - let err = LlamaSampler::grammar(model, "expr ::= \"hello\"", "root").unwrap_err(); - - assert_eq!(err, GrammarError::RootNotFound); - } } #[cfg(test)] diff --git a/llama-cpp-bindings/src/streaming_json_probe.rs b/llama-cpp-bindings/src/streaming_json_probe.rs index 817881f68..816398f8c 100644 --- a/llama-cpp-bindings/src/streaming_json_probe.rs +++ b/llama-cpp-bindings/src/streaming_json_probe.rs @@ -1,122 +1,267 @@ -use serde_json::Value; -use serde_json::error::Category; +use std::mem; -const NAME_FIELD: &str = "name"; -const ARGUMENTS_FIELD: &str = "arguments"; -fn evaluate_completed_value(value: &Value) -> JsonProbeOutcome { - let Value::Object(map) = value else { - return JsonProbeOutcome::Failed; - }; +use serde::Deserialize; +use serde::de::IgnoredAny; - let Some(Value::String(name)) = map.get(NAME_FIELD) else { - return JsonProbeOutcome::Failed; - }; - if name.is_empty() { - return JsonProbeOutcome::Failed; - } +use crate::json_probe_outcome::JsonProbeOutcome; + +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct BareJsonToolCall { + name: String, + #[serde(rename = "arguments")] + _arguments: Option, +} - if let Some(arguments) = map.get(ARGUMENTS_FIELD) - && !matches!(arguments, Value::Object(_)) - { - return JsonProbeOutcome::Failed; +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum ToolCallField { + Arguments, + Name, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +enum ProbeState { + AwaitingObjectOpen, + AwaitingFirstKeyOrClose, + AwaitingKey, + InKey { + quoted_key: String, + escaped: bool, + }, + AwaitingColon(ToolCallField), + AwaitingValue(ToolCallField), + InName { + escaped: bool, + }, + InArguments { + depth: usize, + in_string: bool, + escaped: bool, + }, + AwaitingCommaOrClose, + Closed, + Failed, +} + +const fn is_json_whitespace(character: char) -> bool { + matches!(character, ' ' | '\t' | '\n' | '\r') +} + +fn field_named(quoted_key: &str) -> Option { + match serde_json::from_str::(quoted_key).ok()?.as_str() { + "arguments" => Some(ToolCallField::Arguments), + "name" => Some(ToolCallField::Name), + _ => None, } +} - for key in map.keys() { - if key != NAME_FIELD && key != ARGUMENTS_FIELD { - return JsonProbeOutcome::Failed; +fn is_named_tool_call(held_text: &str) -> bool { + serde_json::from_str::(held_text) + .is_ok_and(|tool_call| !tool_call.name.is_empty()) +} + +impl ProbeState { + fn advance(self, character: char) -> Self { + match self { + Self::AwaitingObjectOpen => match character { + '{' => Self::AwaitingFirstKeyOrClose, + _ if is_json_whitespace(character) => Self::AwaitingObjectOpen, + _ => Self::Failed, + }, + Self::AwaitingFirstKeyOrClose => match character { + '"' => Self::InKey { + quoted_key: String::from('"'), + escaped: false, + }, + '}' => Self::Closed, + _ if is_json_whitespace(character) => Self::AwaitingFirstKeyOrClose, + _ => Self::Failed, + }, + Self::AwaitingKey => match character { + '"' => Self::InKey { + quoted_key: String::from('"'), + escaped: false, + }, + _ if is_json_whitespace(character) => Self::AwaitingKey, + _ => Self::Failed, + }, + Self::InKey { + mut quoted_key, + escaped, + } => { + quoted_key.push(character); + + if escaped || character != '"' { + Self::InKey { + quoted_key, + escaped: !escaped && character == '\\', + } + } else { + field_named("ed_key).map_or(Self::Failed, Self::AwaitingColon) + } + } + Self::AwaitingColon(field) => match character { + ':' => Self::AwaitingValue(field), + _ if is_json_whitespace(character) => Self::AwaitingColon(field), + _ => Self::Failed, + }, + Self::AwaitingValue(field) => match (field, character) { + (ToolCallField::Name, '"') => Self::InName { escaped: false }, + (ToolCallField::Arguments, '{') => Self::InArguments { + depth: 1, + in_string: false, + escaped: false, + }, + _ if is_json_whitespace(character) => Self::AwaitingValue(field), + _ => Self::Failed, + }, + Self::InName { escaped } => { + if !escaped && character == '"' { + Self::AwaitingCommaOrClose + } else { + Self::InName { + escaped: !escaped && character == '\\', + } + } + } + Self::InArguments { + depth, + in_string, + escaped, + } => Self::advance_arguments(depth, in_string, escaped, character), + Self::AwaitingCommaOrClose => match character { + ',' => Self::AwaitingKey, + '}' => Self::Closed, + _ if is_json_whitespace(character) => Self::AwaitingCommaOrClose, + _ => Self::Failed, + }, + Self::Closed if is_json_whitespace(character) => Self::Closed, + Self::Closed | Self::Failed => Self::Failed, } } - JsonProbeOutcome::CompletedValid + const fn advance_arguments( + depth: usize, + in_string: bool, + escaped: bool, + character: char, + ) -> Self { + if in_string { + return Self::InArguments { + depth, + in_string: escaped || character != '"', + escaped: !escaped && character == '\\', + }; + } + + match character { + '"' => Self::InArguments { + depth, + in_string: true, + escaped: false, + }, + '{' | '[' => Self::InArguments { + depth: depth + 1, + in_string: false, + escaped: false, + }, + '}' | ']' if depth == 1 => Self::AwaitingCommaOrClose, + '}' | ']' => Self::InArguments { + depth: depth - 1, + in_string: false, + escaped: false, + }, + _ => Self::InArguments { + depth, + in_string: false, + escaped: false, + }, + } + } } -#[derive(Copy, Clone, Debug, Eq, PartialEq)] -pub enum JsonProbeOutcome { - StillPossiblyValid, - CompletedValid, - Failed, +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct StreamingJsonProbe { + held_text: String, + state: ProbeState, } -impl JsonProbeOutcome { - #[must_use] - pub fn validate_prefix(buffer: &str) -> Self { - let trimmed = buffer.trim_start(); - if trimmed.is_empty() { - return Self::StillPossiblyValid; +impl Default for StreamingJsonProbe { + fn default() -> Self { + Self { + held_text: String::new(), + state: ProbeState::AwaitingObjectOpen, } - if !trimmed.starts_with('{') { - return Self::Failed; + } +} + +impl StreamingJsonProbe { + pub fn feed(&mut self, piece: &str) -> JsonProbeOutcome { + self.held_text.push_str(piece); + + for character in piece.chars() { + let state = mem::replace(&mut self.state, ProbeState::Failed); + + self.state = state.advance(character); + + if self.state == ProbeState::Failed { + return JsonProbeOutcome::Failed; + } } - match serde_json::from_str::(trimmed) { - Ok(value) => evaluate_completed_value(&value), - Err(parse_error) => match parse_error.classify() { - Category::Eof => Self::StillPossiblyValid, - Category::Io | Category::Syntax | Category::Data => Self::Failed, - }, + match self.state { + ProbeState::Closed if is_named_tool_call(&self.held_text) => { + JsonProbeOutcome::CompletedValid + } + ProbeState::Closed | ProbeState::Failed => JsonProbeOutcome::Failed, + _ => JsonProbeOutcome::StillPossiblyValid, } } } #[cfg(test)] mod tests { - use serde_json::Value; + use super::StreamingJsonProbe; + use crate::json_probe_outcome::JsonProbeOutcome; - use super::JsonProbeOutcome; - use super::evaluate_completed_value; + fn probe(buffer: &str) -> JsonProbeOutcome { + StreamingJsonProbe::default().feed(buffer) + } #[test] fn empty_buffer_is_still_possibly_valid() { - assert_eq!( - JsonProbeOutcome::validate_prefix(""), - JsonProbeOutcome::StillPossiblyValid, - ); + assert_eq!(probe(""), JsonProbeOutcome::StillPossiblyValid,); } #[test] fn whitespace_only_buffer_is_still_possibly_valid() { - assert_eq!( - JsonProbeOutcome::validate_prefix(" \n "), - JsonProbeOutcome::StillPossiblyValid, - ); + assert_eq!(probe(" \n "), JsonProbeOutcome::StillPossiblyValid,); } #[test] fn single_open_brace_is_still_possibly_valid() { - assert_eq!( - JsonProbeOutcome::validate_prefix("{"), - JsonProbeOutcome::StillPossiblyValid, - ); + assert_eq!(probe("{"), JsonProbeOutcome::StillPossiblyValid,); } #[test] fn open_brace_with_trailing_space_is_still_possibly_valid() { - assert_eq!( - JsonProbeOutcome::validate_prefix("{ "), - JsonProbeOutcome::StillPossiblyValid, - ); + assert_eq!(probe("{ "), JsonProbeOutcome::StillPossiblyValid,); } #[test] fn open_brace_with_quote_starting_key_is_still_possibly_valid() { - assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ ""#), - JsonProbeOutcome::StillPossiblyValid, - ); + assert_eq!(probe(r#"{ ""#), JsonProbeOutcome::StillPossiblyValid,); } #[test] fn partial_name_key_is_still_possibly_valid() { - assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name""#), - JsonProbeOutcome::StillPossiblyValid, - ); + assert_eq!(probe(r#"{ "name""#), JsonProbeOutcome::StillPossiblyValid,); } #[test] fn partial_name_value_quote_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name": ""#), + probe(r#"{ "name": ""#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -124,7 +269,7 @@ mod tests { #[test] fn partial_name_value_letters_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name": "ge"#), + probe(r#"{ "name": "ge"#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -132,7 +277,7 @@ mod tests { #[test] fn complete_name_string_no_comma_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name": "get_weather""#), + probe(r#"{ "name": "get_weather""#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -140,7 +285,7 @@ mod tests { #[test] fn name_then_comma_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name": "get_weather","#), + probe(r#"{ "name": "get_weather","#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -148,7 +293,7 @@ mod tests { #[test] fn name_then_partial_arguments_key_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name": "get_weather", "argum"#), + probe(r#"{ "name": "get_weather", "argum"#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -156,7 +301,7 @@ mod tests { #[test] fn name_then_arguments_key_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name": "get_weather", "arguments""#), + probe(r#"{ "name": "get_weather", "arguments""#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -164,7 +309,7 @@ mod tests { #[test] fn name_then_arguments_open_brace_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name": "get_weather", "arguments": {"#), + probe(r#"{ "name": "get_weather", "arguments": {"#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -172,9 +317,7 @@ mod tests { #[test] fn arguments_with_partial_inner_key_value_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix( - r#"{ "name": "get_weather", "arguments": {"location":"# - ), + probe(r#"{ "name": "get_weather", "arguments": {"location":"#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -182,9 +325,7 @@ mod tests { #[test] fn arguments_with_partial_inner_string_value_is_still_possibly_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix( - r#"{ "name": "get_weather", "arguments": {"location": "Pa"# - ), + probe(r#"{ "name": "get_weather", "arguments": {"location": "Pa"#), JsonProbeOutcome::StillPossiblyValid, ); } @@ -192,7 +333,7 @@ mod tests { #[test] fn complete_simple_tool_call_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":{}}"#), + probe(r#"{"name":"f","arguments":{}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -200,7 +341,7 @@ mod tests { #[test] fn complete_tool_call_with_internal_whitespace_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name": "f", "arguments": {}}"#), + probe(r#"{"name": "f", "arguments": {}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -208,9 +349,7 @@ mod tests { #[test] fn complete_tool_call_with_string_argument_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix( - r#"{"name":"get_weather","arguments":{"location":"Paris"}}"# - ), + probe(r#"{"name":"get_weather","arguments":{"location":"Paris"}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -218,9 +357,7 @@ mod tests { #[test] fn complete_tool_call_with_multiple_arguments_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix( - r#"{"name":"book_flight","arguments":{"from":"NYC","to":"PAR","passengers":2}}"# - ), + probe(r#"{"name":"book_flight","arguments":{"from":"NYC","to":"PAR","passengers":2}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -228,7 +365,7 @@ mod tests { #[test] fn complete_tool_call_with_nested_arguments_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":{"a":{"b":[1,2,3]}}}"#), + probe(r#"{"name":"f","arguments":{"a":{"b":[1,2,3]}}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -236,7 +373,7 @@ mod tests { #[test] fn complete_tool_call_with_close_brace_inside_string_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":{"q":"a } b"}}"#), + probe(r#"{"name":"f","arguments":{"q":"a } b"}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -244,7 +381,7 @@ mod tests { #[test] fn complete_tool_call_with_escaped_quotes_in_string_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":{"q":"he said \"hi\""}}"#), + probe(r#"{"name":"f","arguments":{"q":"he said \"hi\""}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -252,7 +389,7 @@ mod tests { #[test] fn complete_tool_call_with_unicode_strings_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"日本語","arguments":{"city":"パリ"}}"#), + probe(r#"{"name":"日本語","arguments":{"city":"パリ"}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -260,7 +397,7 @@ mod tests { #[test] fn complete_tool_call_with_trailing_whitespace_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix("{\"name\":\"f\",\"arguments\":{}}\n"), + probe("{\"name\":\"f\",\"arguments\":{}}\n"), JsonProbeOutcome::CompletedValid, ); } @@ -268,7 +405,7 @@ mod tests { #[test] fn complete_tool_call_with_array_inside_arguments_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":{"items":[1,2,3]}}"#), + probe(r#"{"name":"f","arguments":{"items":[1,2,3]}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -276,47 +413,35 @@ mod tests { #[test] fn complete_tool_call_without_arguments_field_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"ping"}"#), + probe(r#"{"name":"ping"}"#), JsonProbeOutcome::CompletedValid, ); } #[test] fn top_level_array_is_failed() { - assert_eq!( - JsonProbeOutcome::validate_prefix("["), - JsonProbeOutcome::Failed - ); + assert_eq!(probe("["), JsonProbeOutcome::Failed); } #[test] fn top_level_scalar_number_is_failed() { - assert_eq!( - JsonProbeOutcome::validate_prefix("123"), - JsonProbeOutcome::Failed - ); + assert_eq!(probe("123"), JsonProbeOutcome::Failed); } #[test] fn top_level_string_is_failed() { - assert_eq!( - JsonProbeOutcome::validate_prefix(r#""hi""#), - JsonProbeOutcome::Failed - ); + assert_eq!(probe(r#""hi""#), JsonProbeOutcome::Failed); } #[test] fn complete_object_with_wrong_first_key_is_failed() { - assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"foo":"bar"}"#), - JsonProbeOutcome::Failed, - ); + assert_eq!(probe(r#"{"foo":"bar"}"#), JsonProbeOutcome::Failed,); } #[test] fn complete_object_with_non_string_name_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":123,"arguments":{}}"#), + probe(r#"{"name":123,"arguments":{}}"#), JsonProbeOutcome::Failed, ); } @@ -324,7 +449,7 @@ mod tests { #[test] fn complete_object_with_null_name_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":null,"arguments":{}}"#), + probe(r#"{"name":null,"arguments":{}}"#), JsonProbeOutcome::Failed, ); } @@ -332,7 +457,7 @@ mod tests { #[test] fn complete_object_with_arguments_as_array_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":[]}"#), + probe(r#"{"name":"f","arguments":[]}"#), JsonProbeOutcome::Failed, ); } @@ -340,7 +465,7 @@ mod tests { #[test] fn complete_object_with_arguments_as_string_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":"hi"}"#), + probe(r#"{"name":"f","arguments":"hi"}"#), JsonProbeOutcome::Failed, ); } @@ -348,7 +473,7 @@ mod tests { #[test] fn complete_object_with_third_top_level_key_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":{},"extra":1}"#), + probe(r#"{"name":"f","arguments":{},"extra":1}"#), JsonProbeOutcome::Failed, ); } @@ -356,7 +481,7 @@ mod tests { #[test] fn complete_object_with_empty_name_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"","arguments":{}}"#), + probe(r#"{"name":"","arguments":{}}"#), JsonProbeOutcome::Failed, ); } @@ -364,39 +489,30 @@ mod tests { #[test] fn complete_object_with_trailing_garbage_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":{}}garbage"#), + probe(r#"{"name":"f","arguments":{}}garbage"#), JsonProbeOutcome::Failed, ); } #[test] fn empty_object_is_failed_due_to_missing_required_name() { - assert_eq!( - JsonProbeOutcome::validate_prefix("{}"), - JsonProbeOutcome::Failed - ); + assert_eq!(probe("{}"), JsonProbeOutcome::Failed); } #[test] fn complete_object_with_arguments_only_no_name_is_failed() { - assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"arguments":{}}"#), - JsonProbeOutcome::Failed, - ); + assert_eq!(probe(r#"{"arguments":{}}"#), JsonProbeOutcome::Failed,); } #[test] fn leading_whitespace_then_open_brace_is_still_possibly_valid() { - assert_eq!( - JsonProbeOutcome::validate_prefix("\n \n{"), - JsonProbeOutcome::StillPossiblyValid, - ); + assert_eq!(probe("\n \n{"), JsonProbeOutcome::StillPossiblyValid,); } #[test] fn leading_whitespace_then_complete_tool_call_is_completed_valid() { assert_eq!( - JsonProbeOutcome::validate_prefix("\n {\"name\":\"f\",\"arguments\":{}}"), + probe("\n {\"name\":\"f\",\"arguments\":{}}"), JsonProbeOutcome::CompletedValid, ); } @@ -404,25 +520,20 @@ mod tests { #[test] fn complete_tool_call_followed_by_second_object_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix( - r#"{"name":"a","arguments":{}}{"name":"b","arguments":{}}"# - ), + probe(r#"{"name":"a","arguments":{}}{"name":"b","arguments":{}}"#), JsonProbeOutcome::Failed, ); } #[test] fn buffer_with_only_open_quote_is_still_possibly_valid() { - assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "n"#), - JsonProbeOutcome::StillPossiblyValid, - ); + assert_eq!(probe(r#"{ "n"#), JsonProbeOutcome::StillPossiblyValid,); } #[test] fn buffer_with_complete_first_field_unknown_second_key_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{ "name": "f", "foo": 1}"#), + probe(r#"{ "name": "f", "foo": 1}"#), JsonProbeOutcome::Failed, ); } @@ -430,7 +541,7 @@ mod tests { #[test] fn unicode_letter_inside_name_value_completes_validly() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"éclair","arguments":{}}"#), + probe(r#"{"name":"éclair","arguments":{}}"#), JsonProbeOutcome::CompletedValid, ); } @@ -438,23 +549,50 @@ mod tests { #[test] fn arguments_field_with_explicit_null_is_failed() { assert_eq!( - JsonProbeOutcome::validate_prefix(r#"{"name":"f","arguments":null}"#), + probe(r#"{"name":"f","arguments":null}"#), JsonProbeOutcome::Failed, ); } #[test] fn syntactically_malformed_object_is_failed() { + assert_eq!(probe("{,}"), JsonProbeOutcome::Failed,); + } + + #[test] + fn key_written_with_an_escape_sequence_is_recognized() { assert_eq!( - JsonProbeOutcome::validate_prefix("{,}"), - JsonProbeOutcome::Failed, + probe(r#"{"na\u006de":"f","arguments":{}}"#), + JsonProbeOutcome::CompletedValid, + ); + } + + #[test] + fn tool_call_fed_one_character_at_a_time_completes_on_its_last_character() { + let tool_call = r#"{"name":"f","arguments":{"q":"a } b"}}"#; + let mut streaming_probe = StreamingJsonProbe::default(); + let mut outcomes = Vec::new(); + + for character in tool_call.chars() { + outcomes.push(streaming_probe.feed(&character.to_string())); + } + + let (last_outcome, earlier_outcomes) = outcomes + .split_last() + .expect("the tool call must produce outcomes"); + + assert_eq!(*last_outcome, JsonProbeOutcome::CompletedValid); + assert!( + earlier_outcomes + .iter() + .all(|outcome| *outcome == JsonProbeOutcome::StillPossiblyValid) ); } #[test] - fn non_object_completed_value_is_failed() { + fn syntax_error_inside_arguments_fails_when_the_object_closes() { assert_eq!( - evaluate_completed_value(&Value::Bool(true)), + probe(r#"{"name":"f","arguments":{"q" 1}}"#), JsonProbeOutcome::Failed, ); } diff --git a/llama-cpp-bindings/src/streaming_markers.rs b/llama-cpp-bindings/src/streaming_markers.rs index af43ed63c..8b9ed5e13 100644 --- a/llama-cpp-bindings/src/streaming_markers.rs +++ b/llama-cpp-bindings/src/streaming_markers.rs @@ -68,10 +68,23 @@ impl StreamingMarkers { } #[must_use] - pub fn longest_matching_suffix(&self, tokens: &[LlamaToken]) -> Option<&StreamingMarker> { + pub fn longest_matching_suffix( + &self, + reversed_tokens: &TReversedTokens, + ) -> Option<&StreamingMarker> + where + TReversedTokens: Iterator + Clone, + { self.markers .iter() - .filter(|marker| tokens.ends_with(marker.tokens())) + .filter(|marker| { + marker + .tokens() + .iter() + .rev() + .copied() + .eq(reversed_tokens.clone().take(marker.tokens().len())) + }) .max_by_key(|marker| marker.tokens().len()) } @@ -173,7 +186,7 @@ mod tests { .expect("markers are valid"); let matched = markers - .longest_matching_suffix(&[token(1), token(2)]) + .longest_matching_suffix(&[token(1), token(2)].into_iter().rev()) .expect("a suffix must match"); assert_eq!(matched.tokens(), &[token(1), token(2)]); diff --git a/llama-cpp-bindings/src/token/data_array.rs b/llama-cpp-bindings/src/token/data_array.rs index e2fd0d813..8d841c502 100644 --- a/llama-cpp-bindings/src/token/data_array.rs +++ b/llama-cpp-bindings/src/token/data_array.rs @@ -79,6 +79,16 @@ impl LlamaTokenDataArray { Self::new(data.into_iter().collect(), sorted) } + pub fn replace_candidates(&mut self, candidates: TIterator) + where + TIterator: IntoIterator, + { + self.data.clear(); + self.data.extend(candidates); + self.selected = None; + self.sorted = false; + } + #[must_use] pub fn selected_token(&self) -> Option { self.data.get(self.selected?).map(LlamaTokenData::id) diff --git a/llama-cpp-bindings/src/token_piece.rs b/llama-cpp-bindings/src/token_piece.rs new file mode 100644 index 000000000..5bef2b216 --- /dev/null +++ b/llama-cpp-bindings/src/token_piece.rs @@ -0,0 +1,43 @@ +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum TokenPiece { + Marker(String), + Visible(String), +} + +impl TokenPiece { + #[must_use] + pub fn raw(&self) -> &str { + match self { + Self::Marker(text) | Self::Visible(text) => text, + } + } + + #[must_use] + pub fn visible(&self) -> &str { + match self { + Self::Marker(_) => "", + Self::Visible(text) => text, + } + } +} + +#[cfg(test)] +mod tests { + use super::TokenPiece; + + #[test] + fn a_marker_is_raw_text_without_a_visible_part() { + let piece = TokenPiece::Marker("".to_owned()); + + assert_eq!(piece.raw(), ""); + assert_eq!(piece.visible(), ""); + } + + #[test] + fn visible_text_is_both_raw_and_visible() { + let piece = TokenPiece::Visible("hello".to_owned()); + + assert_eq!(piece.raw(), "hello"); + assert_eq!(piece.visible(), "hello"); + } +} diff --git a/llama-cpp-gbnf/src/validate_gbnf.rs b/llama-cpp-gbnf/src/validate_gbnf.rs index 665f1d7f4..60f71e848 100644 --- a/llama-cpp-gbnf/src/validate_gbnf.rs +++ b/llama-cpp-gbnf/src/validate_gbnf.rs @@ -1,4 +1,5 @@ use std::ffi::{CString, c_char}; +use std::ptr; use llama_cpp_ffi_status::FfiContractError; use llama_cpp_ffi_status::FfiStatusError; @@ -12,6 +13,7 @@ use llama_cpp_bindings_sys::LLAMA_RS_GBNF_VALIDATION_SYNTAX_ERROR; use llama_cpp_bindings_sys::LLAMA_RS_GBNF_VALIDATION_THREW_CXX_EXCEPTION; use llama_cpp_bindings_sys::llama_rs_gbnf_validation_status; use llama_cpp_bindings_sys::llama_rs_validate_gbnf; +use llama_cpp_bindings_sys::llama_vocab; use crate::gbnf_validation_error::GbnfValidationError; @@ -78,13 +80,18 @@ fn validation_status_to_result( /// /// Returns [`GbnfValidationError`] when `grammar` or `root` contains an interior /// NUL byte, or when the grammar parser rejects the grammar. -pub fn validate_gbnf(grammar: &str, root: &str) -> Result<(), GbnfValidationError> { +pub fn validate_gbnf( + vocab: Option<&llama_vocab>, + grammar: &str, + root: &str, +) -> Result<(), GbnfValidationError> { let grammar_cstring = CString::new(grammar).map_err(GbnfValidationError::GrammarContainsNul)?; let root_cstring = CString::new(root).map_err(GbnfValidationError::RootContainsNul)?; - let mut out_error = std::ptr::null_mut(); + let mut out_error = ptr::null_mut(); let status = unsafe { llama_rs_validate_gbnf( + vocab.map_or(ptr::null(), ptr::from_ref), grammar_cstring.as_ptr(), root_cstring.as_ptr(), &raw mut out_error, @@ -108,13 +115,16 @@ mod tests { #[test] fn valid_grammar_is_accepted() { - assert_eq!(validate_gbnf(r#"root ::= "yes" | "no""#, "root"), Ok(())); + assert_eq!( + validate_gbnf(None, r#"root ::= "yes" | "no""#, "root"), + Ok(()) + ); } #[test] fn malformed_grammar_is_a_syntax_error() { assert_eq!( - validate_gbnf("root ::= (", "root"), + validate_gbnf(None, "root ::= (", "root"), Err(GbnfValidationError::SyntaxError) ); } @@ -122,7 +132,7 @@ mod tests { #[test] fn empty_grammar_has_no_rules() { assert_eq!( - validate_gbnf("", "root"), + validate_gbnf(None, "", "root"), Err(GbnfValidationError::EmptyRuleSet) ); } @@ -130,7 +140,7 @@ mod tests { #[test] fn grammar_without_root_reports_missing_root() { assert_eq!( - validate_gbnf(r#"expr ::= "x""#, "root"), + validate_gbnf(None, r#"expr ::= "x""#, "root"), Err(GbnfValidationError::RootSymbolMissing { root: "root".to_owned() }) @@ -140,7 +150,7 @@ mod tests { #[test] fn left_recursive_grammar_is_rejected() { assert_eq!( - validate_gbnf(r#"root ::= root "x""#, "root"), + validate_gbnf(None, r#"root ::= root "x""#, "root"), Err(GbnfValidationError::LeftRecursion) ); } @@ -150,7 +160,7 @@ mod tests { let grammar = "root ::= \"a\0b\""; assert_eq!( - validate_gbnf(grammar, "root").err(), + validate_gbnf(None, grammar, "root").err(), CString::new(grammar) .err() .map(GbnfValidationError::GrammarContainsNul) @@ -162,7 +172,7 @@ mod tests { let root = "ro\0ot"; assert_eq!( - validate_gbnf(r#"root ::= "x""#, root).err(), + validate_gbnf(None, r#"root ::= "x""#, root).err(), CString::new(root) .err() .map(GbnfValidationError::RootContainsNul) diff --git a/llama-cpp-test-harness/src/execution_phase.rs b/llama-cpp-test-harness/src/execution_phase.rs index 8be5269f8..eb29819d8 100644 --- a/llama-cpp-test-harness/src/execution_phase.rs +++ b/llama-cpp-test-harness/src/execution_phase.rs @@ -1,4 +1,5 @@ use std::sync::Arc; +use std::sync::OnceLock; use libtest_mimic::Arguments; use libtest_mimic::Conclusion; @@ -12,6 +13,8 @@ use crate::llama_test_registration::LlamaTestRegistration; use crate::load_key::LoadKey; use crate::phase_state::PhaseState; +type LazyPhaseState = Arc>>; + fn source_label(source: GgufSource) -> String { match source { GgufSource::HuggingFace { repo, file } => format!("{repo} / {file}"), @@ -42,42 +45,43 @@ impl ExecutionPhase { } pub fn run(&self, backend: &Arc, arguments: &Arguments) -> Conclusion { - let trials = match self.key.load_phase_state(backend) { - Ok(state) => self.passing_trials(&Arc::new(state)), - Err(error) => self.failing_trials(&format!("phase setup failed: {error:#}")), - }; - libtest_mimic::run(arguments, trials) + let phase_state: LazyPhaseState = Arc::new(OnceLock::new()); + + libtest_mimic::run( + arguments, + self.registrations + .iter() + .map(|registration| self.trial(registration, backend, &phase_state)) + .collect(), + ) } - fn passing_trials(&self, state: &Arc) -> Vec { - self.registrations - .iter() - .map(|registration| { - let state_for_trial = Arc::clone(state); - let registration: &'static LlamaTestRegistration = registration; - let func = registration.func; - Trial::test(registration.name, move || { - let fixture = LlamaFixture { - model: &state_for_trial.model, - backend: &state_for_trial.backend, - context_params: ®istration.context_params, - mtmd_context: state_for_trial.mtmd_context.as_ref(), - model_path: &state_for_trial.model_path, - }; - func(&fixture).map_err(|error| Failed::from(format!("{error:#}"))) + fn trial( + &self, + registration: &'static LlamaTestRegistration, + backend: &Arc, + phase_state: &LazyPhaseState, + ) -> Trial { + let key = self.key; + let backend = Arc::clone(backend); + let phase_state = Arc::clone(phase_state); + + Trial::test(registration.name, move || { + match phase_state.get_or_init(|| { + key.load_phase_state(&backend) + .map_err(|error| format!("phase setup failed: {error:#}")) + }) { + Ok(state) => (registration.func)(&LlamaFixture { + model: &state.model, + backend: &state.backend, + context_params: ®istration.context_params, + mtmd_context: state.mtmd_context.as_ref(), + model_path: &state.model_path, }) - }) - .collect() - } - - fn failing_trials(&self, error_message: &str) -> Vec { - self.registrations - .iter() - .map(|registration| { - let message = error_message.to_owned(); - Trial::test(registration.name, move || Err(Failed::from(message))) - }) - .collect() + .map_err(|error| Failed::from(format!("{error:#}"))), + Err(setup_failure) => Err(Failed::from(setup_failure)), + } + }) } } diff --git a/llama-cpp-wrapper-sources/src/wrapper_headers.rs b/llama-cpp-wrapper-sources/src/wrapper_headers.rs index 77972c1b5..96f6f5a72 100644 --- a/llama-cpp-wrapper-sources/src/wrapper_headers.rs +++ b/llama-cpp-wrapper-sources/src/wrapper_headers.rs @@ -3,6 +3,7 @@ pub const WRAPPER_HEADERS: &[&str] = &[ "wrapper_chat_apply.h", "wrapper_chat_parse.h", "wrapper_common.h", + "wrapper_context.h", "wrapper_fit.h", "wrapper_gbnf.h", "wrapper_mtmd.h", diff --git a/llama-cpp-wrapper-sources/src/wrapper_sources.rs b/llama-cpp-wrapper-sources/src/wrapper_sources.rs index 384f80c93..dcd5ffa7d 100644 --- a/llama-cpp-wrapper-sources/src/wrapper_sources.rs +++ b/llama-cpp-wrapper-sources/src/wrapper_sources.rs @@ -2,6 +2,7 @@ pub const WRAPPER_SOURCES: &[&str] = &[ "wrapper_chat_apply.cpp", "wrapper_chat_parse.cpp", "wrapper_common.cpp", + "wrapper_context.cpp", "wrapper_fit.cpp", "wrapper_gbnf.cpp", "wrapper_mtmd.cpp",