diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c0734dc..6a3604e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -32,3 +32,10 @@ jobs: - name: Test run: swift test + + - name: Test token healing without a model + run: | + clang++ -std=c++17 -I Sources/CotabbyInferenceEngine \ + Sources/CotabbyInferenceEngine/TokenHealing.cpp Tests/TokenHealingTests.cpp \ + -o /tmp/cotabby-token-healing-tests + /tmp/cotabby-token-healing-tests diff --git a/README.md b/README.md index 1b2aca3..673e8de 100644 --- a/README.md +++ b/README.md @@ -22,7 +22,7 @@ Cotabby LlamaRuntimeCore Cotabby serializes generation and prefill through its runtime lock, so the middleware does not reserve unused secondary sequence capacity or run a batching worker. Prompt decode, feedback decode, KV trim, and sequence destruction use one native context mutex. Cancellation is the intentional -cross-thread operation and uses a one-way atomic flag. +cross-thread operation and uses an atomic flag, rearmed only after successful cache restoration. The public sequence ID changes whenever a sequence is recreated even though llama's internal slot is always zero. That prevents a late cancellation from accidentally targeting a replacement @@ -85,25 +85,28 @@ guard engine.decodePrompt(sequenceID, &tokens, Int32(tokens.count), 0) == .ok el } // The caller owns the generation budget. +var completionBytes: [UInt8] = [] for _ in 0 ..< 8 { let result = engine.sampleNext(sequenceID) if result.is_eos || result.was_cancelled { break } if let piece = result.piece, result.piece_length > 0 { - let text = String( - bytes: UnsafeBufferPointer( + completionBytes += Array( + UnsafeBufferPointer( start: UnsafeRawPointer(piece).assumingMemoryBound(to: UInt8.self), count: Int(result.piece_length) - ), - encoding: .utf8 - ) ?? "" - print(text, terminator: "") + ) + ) } } +if let text = String(bytes: completionBytes, encoding: .utf8) { + print(text) +} ~~~ `SampleResult.piece` is borrowed sequence storage. Copy it before another sampling call or sequence -destruction. +destruction. Its bytes can end inside a UTF-8 scalar; streaming clients must accumulate bytes before +converting to text instead of discarding undecodable individual pieces. ## Generation Semantics @@ -119,16 +122,55 @@ Cotabby controls the maximum token count in Swift. The engine controls token sel sampling selected visible text; - `logprob`: the selected token's raw-model log-probability when enabled. +## Caret Token Healing + +`tokenPiece(token)` returns printable bytes for planning a completion. When the final prompt token +exactly matches the typed suffix, the client can remove that token and pass its bytes to +`setCompletionPrefix(sequence, bytes, length)` before `decodePrompt`. The sampler then admits only +tokens compatible with that prefix, including byte-fallback tokens that cover part of it. This +allows `sched` to participate in the token for `schedule` without altering the writer's text. + +`TokenHealingVocabulary` is built once per loaded model; the per-sequence `TokenPrefix` owns only +the unconsumed replay bytes. The caller strips exactly those replayed bytes, preserves incomplete +UTF-8 fragments, and budgets replay separately from visible continuation. Cotabby caps replay at +16 bytes/tokens. While constrained, the legacy whitespace mask is bypassed and `argmax_is_eog` +is false. No EOS/control tokens can complete an unfinished replay. Clear the prefix for requests +that do not heal; sampler settings and cache lifetime do not imply the next request's prefix. + ## KV Reuse and Cancellation -`trimKV` removes a suffix from fixed llama sequence slot zero and invalidates any saved seed or -pending feedback token. Cotabby independently validates request continuity, UTF-8 prefix, token -prefix, and sampling compatibility before calling it. +`trimKV` restores a prefix in slot zero and invalidates its pending seed/feedback token. Ordinary +attention uses direct suffix removal. Recurrent, hybrid, and sliding-window models use one +`PARTIAL_ONLY | ON_DEVICE` checkpoint near the prompt tail, retaining full attention KV in place. +The checkpoint's tensor memory is capped at 128 MiB and restoration replays at most eight prompt +tokens without sampling. Device storage belongs to the llama context's slot-zero checkpoint and +is released when that context unloads. Checkpoint metadata belongs to `SequenceState`. + +A request that edits before the saved checkpoint, a very short prompt without a nonempty checkpoint, +or a model exceeding the cap returns a cache miss; callers rebuild that request. A miss must not +permanently disable reuse for the model. Saving a newer near-caret checkpoint intentionally gives +up deeper backspace history. Gemma's sliding-window cache supports suffix removal, but restoring +its saved window also protects rows evicted during prediction. No model-family-name heuristics +are used: llama model metadata selects the partial-state path. + +`decodePrompt` resets the sampler and accepts only committed prompt tokens. Discarded predictions +therefore never pollute penalties or RNG state; `llama_sampler_sample` already accepts its sample, +so the wrapper must not accept it a second time. A sampled seed is not in KV until feedback decode. +Every successful decode advances tracked positions even when the next result is EOS/cancelled. +A failed native decode invalidates cache reuse until the sequence is destroyed; an unchanged +tracked position must never make the equal-position trim shortcut certify uncertain memory. + +`getCacheDiagnostics(sequence)` exposes actual token position, checkpoint tensor bytes/position, +restoration replay count, and whether partial checkpoints are required. It contains no text. +Cotabby independently validates field continuity, byte/token prefix, and sampling compatibility. `cancelSequence` is thread-safe and nonblocking. Prompt decode checks cancellation between chunks; sample generation checks before work and after feedback decode. An active llama decode is not -preempted mid-call. Cotabby destroys a natively cancelled sequence because the flag is intentionally -one-way. +preempted mid-call. Successful `trimKV` clears the cancellation flag only after memory is valid; +a failed restoration leaves the sequence cancelled and requires destruction. Cotabby closes its +operation-specific cancellation target before restoring, so a late task cancellation cannot poison +the next request that reuses the same sequence. Destruction and cancellation serialize ownership +through the sequence mutex. ## Testing @@ -144,10 +186,31 @@ Run the full native path with a local GGUF: COTABBY_TEST_MODEL_PATH=/absolute/path/model.gguf swift test ~~~ +If the default Xcode-backed SwiftPM runner fails code signing because of local Finder metadata, +`swift test --build-system native` runs the same tests with SwiftPM's native runner. + +Run the deterministic vocabulary-prefix tests without downloading a model or linking llama: + +~~~bash +clang++ -std=c++17 -I Sources/CotabbyInferenceEngine \ + Sources/CotabbyInferenceEngine/TokenHealing.cpp Tests/TokenHealingTests.cpp \ + -o /tmp/cotabby-token-healing-tests +/tmp/cotabby-token-healing-tests +~~~ + The model-backed suite covers single-sequence admission/replacement, prompt decode, sampling, KV trim, cancellation, mid-word continuation, optional log-probability, scaffolding-token masking, and argmax-EOG behavior. CI does not currently provide a GGUF, so these tests skip there unless the -environment variable is configured. +environment variable is configured. New coverage compares cold/restored token output, cancellation +rearming, exact replay of unfinished words and trailing whitespace, and checkpoint/replay bounds. +An oversized-prompt regression also verifies failed-decode invalidation and fresh-sequence recovery. + +`testWarmPromptDecodeReportsLatency` prints medians for 32-, 214-, and 838-token prompts (the exact +counts vary by tokenizer), excluding the first pass and asserting no wall-clock threshold. On one +local Qwen3.5-0.8B-Base Q6_K run, cold/warm prompt processing measured about 40/26, 89/25, and 284/25 +milliseconds with a 20.2 MB checkpoint. This is native prompt work, not keystroke-to-visible-word +latency, and is not a Gemma or cross-hardware speed claim. Host-memory checkpoints were slower in +the same experiment; keeping tensor copies on device was necessary for the measured improvement. ## Requirements diff --git a/Sources/CotabbyInferenceEngine/CotabbyInferenceEngine.cpp b/Sources/CotabbyInferenceEngine/CotabbyInferenceEngine.cpp index 5fc867f..b9b6aa6 100644 --- a/Sources/CotabbyInferenceEngine/CotabbyInferenceEngine.cpp +++ b/Sources/CotabbyInferenceEngine/CotabbyInferenceEngine.cpp @@ -1,4 +1,5 @@ #include "CotabbyInferenceEngine.h" +#include "TokenHealing.h" #include #include @@ -102,6 +103,23 @@ struct SequenceState { int kv_position_count = 0; std::atomic cancelled{false}; std::string last_piece; + // A failed llama_decode may have partially changed native memory without advancing our + // committed-token count. Never let trimKV's equal-position fast path certify that state. + bool cache_valid = true; + + // Exact committed tokens, excluding a sampled token that has not reached llama_decode. + // The sampler can then be reset to the *writer's* prompt when a prediction is discarded. + std::vector decoded_tokens; + + // Hybrid/recurrent models cannot erase arbitrary suffixes, and a sliding window may already + // have evicted the old prompt's attention rows. Keep one partial-state checkpoint near the + // caret; full attention KV stays in llama's device buffers. This is per sequence, never disk. + std::vector prompt_checkpoint; + size_t checkpoint_memory_bytes = 0; + int checkpoint_position = 0; + int last_restore_replayed_tokens = 0; + cotabby::TokenPrefix completion_prefix; + bool single_line = false; llama_token seed_token = 0; bool has_seed_token = false; @@ -147,6 +165,12 @@ struct CotabbyInferenceEngine::Impl { int batch_size = 0; int thread_count = 0; int gpu_layer_count = 0; + bool needs_prompt_checkpoint = false; + + // Eight tokens cover ordinary retokenization/backspace near the caret while keeping replay + // bounded. One saved state avoids n_rs_seq's multiplication of every recurrent state tensor. + static constexpr int CHECKPOINT_TAIL_TOKENS = 8; + static constexpr size_t MAX_CHECKPOINT_BYTES = 128 * 1024 * 1024; // Token masks built once per model load (see buildTokenMasks). EOG tokens are deliberately // excluded so the stop check still fires; they are never emitted as text. `starts_new_word` @@ -154,6 +178,7 @@ struct CotabbyInferenceEngine::Impl { std::vector nonprintable_bias; std::vector linebreak_bias; std::vector starts_new_word; + cotabby::TokenHealingVocabulary healing_vocabulary; // One product sequence with a monotonically changing external identity. The mutex protects // create/destroy and lookup; callers still must not destroy the sequence while another method @@ -166,6 +191,104 @@ struct CotabbyInferenceEngine::Impl { // prefill, trim, and reset, while cancellation only touches the sequence's atomic flag. std::mutex decode_mutex; + std::string tokenPiece(llama_token token) const { + if (!vocab || token < 0 || token >= llama_vocab_n_tokens(vocab) || + llama_vocab_is_control(vocab, token) || llama_vocab_is_eog(vocab, token) || + (llama_vocab_get_attr(vocab, token) & + (LLAMA_TOKEN_ATTR_UNKNOWN | LLAMA_TOKEN_ATTR_UNUSED)) != 0) return {}; + std::string result(64, '\0'); + int size = llama_token_to_piece(vocab, token, result.data(), static_cast(result.size()), 0, false); + if (size < 0) { + result.resize(-size); + size = llama_token_to_piece(vocab, token, result.data(), static_cast(result.size()), 0, false); + } + if (size <= 0) return {}; + result.resize(size); + return result; + } + + // Hard-constrain only the short replay prefix. A compatible token may finish the typed + // prefix and include new letters, or may cover only its first bytes (byte fallback). + bool maskCompletionPrefix(const SequenceState& seq, int logits_row) const { + if (seq.completion_prefix.empty()) return true; + const auto allowed = healing_vocabulary.matchingTokens(seq.completion_prefix.remaining(), seq.single_line); + if (allowed.empty()) return false; + float* logits = llama_get_logits_ith(shared_ctx, logits_row); + if (!logits) return false; + std::vector values; + values.reserve(allowed.size()); + for (const auto token : allowed) values.push_back(logits[token]); + std::fill_n(logits, llama_vocab_n_tokens(vocab), -INFINITY); + for (size_t index = 0; index < allowed.size(); ++index) { + logits[allowed[index]] = values[index]; + } + return true; + } + + // Called under decode_mutex after a complete batch. Saving only partial memory is sufficient: + // ordinary attention retains all earlier KV, while recurrent/SWA memory must be restored. + void savePromptCheckpoint(SequenceState& seq) { + seq.prompt_checkpoint.clear(); + seq.checkpoint_memory_bytes = 0; + seq.checkpoint_position = 0; + if (!needs_prompt_checkpoint || seq.kv_position_count <= 0) return; + // Query the host size only to bound memory; ON_DEVICE retains tensor data in llama's + // single slot-zero checkpoint buffer instead of transferring ~20 MiB through the CPU on + // every keystroke. Its small serialized blob contains metadata, not ownership of tensors. + const size_t memory_bytes = llama_state_seq_get_size_ext( + shared_ctx, SEQUENCE_ID, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + if (memory_bytes == 0 || memory_bytes > MAX_CHECKPOINT_BYTES) return; + constexpr auto flags = LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY | LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + const size_t size = llama_state_seq_get_size_ext(shared_ctx, SEQUENCE_ID, flags); + if (size == 0) return; + seq.prompt_checkpoint.resize(size); + if (llama_state_seq_get_data_ext(shared_ctx, seq.prompt_checkpoint.data(), size, + SEQUENCE_ID, flags) != size) { + seq.prompt_checkpoint.clear(); + return; + } + seq.checkpoint_memory_bytes = memory_bytes; + seq.checkpoint_position = seq.kv_position_count; + } + + // Restoring must precede seq_rm: the hybrid implementation asks its recurrent cache first, + // and that cache cannot erase a suffix until its earlier state has been reinstalled. + bool restorePromptCheckpoint(SequenceState& seq, int keep_positions) { + if (seq.prompt_checkpoint.empty() || keep_positions < seq.checkpoint_position || + keep_positions > static_cast(seq.decoded_tokens.size()) || + keep_positions - seq.checkpoint_position > CHECKPOINT_TAIL_TOKENS) return false; + constexpr auto flags = LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY | LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + if (llama_state_seq_set_data_ext(shared_ctx, seq.prompt_checkpoint.data(), + seq.prompt_checkpoint.size(), SEQUENCE_ID, flags) != seq.prompt_checkpoint.size()) { + return false; + } + if (!llama_memory_seq_rm(llama_get_memory(shared_ctx), SEQUENCE_ID, + seq.checkpoint_position, -1)) return false; + + // At most CHECKPOINT_TAIL_TOKENS prompt tokens are replayed in the normal path. Do not + // sample during restoration: it must not advance RNG, penalties, or the visible output. + const int count = keep_positions - seq.checkpoint_position; + if (count > 0) { + llama_batch batch = llama_batch_init(count, 0, 1); + batch.n_tokens = count; + for (int i = 0; i < count; ++i) { + batch.token[i] = seq.decoded_tokens[seq.checkpoint_position + i]; + batch.pos[i] = seq.checkpoint_position + i; + batch.n_seq_id[i] = 1; + batch.seq_id[i][0] = SEQUENCE_ID; + batch.logits[i] = 0; + } + const int status = llama_decode(shared_ctx, batch); + llama_batch_free(batch); + if (status != 0) { + seq.cache_valid = false; + return false; + } + } + seq.last_restore_replayed_tokens = count; + return true; + } + SequenceState* findSequence(int32_t id) { std::lock_guard lock(sequence_mutex); return sequence && sequence->external_id == id ? sequence.get() : nullptr; @@ -246,6 +369,7 @@ struct CotabbyInferenceEngine::Impl { nonprintable_bias.clear(); linebreak_bias.clear(); starts_new_word.clear(); + healing_vocabulary.clear(); if (!vocab) return; const int32_t n = llama_vocab_n_tokens(vocab); @@ -256,6 +380,7 @@ struct CotabbyInferenceEngine::Impl { const llama_token bos_token = llama_vocab_bos(vocab); char piece[64]; + std::vector healing_entries; for (llama_token t = 0; t < n; ++t) { const bool is_eog = llama_vocab_is_eog(vocab, t); @@ -282,6 +407,11 @@ struct CotabbyInferenceEngine::Impl { } } + auto plain_piece = tokenPiece(t); + if (!plain_piece.empty() && !isScaffoldingMarkerPiece(plain_piece.data(), static_cast(plain_piece.size()))) { + healing_entries.push_back({t, std::move(plain_piece)}); + } + const int written = llama_token_to_piece(vocab, t, piece, sizeof(piece), 0, false); if (written <= 0) { continue; @@ -299,6 +429,7 @@ struct CotabbyInferenceEngine::Impl { } } } + healing_vocabulary.reset(std::move(healing_entries)); } // Masks every "starts a new word" token (decoded text begins with whitespace) in the logits @@ -424,6 +555,8 @@ EngineStatus CotabbyInferenceEngine::loadModel(const char* path, int gpu_layers, impl_->context_window_tokens = context_window_tokens; impl_->batch_size = batch_size; impl_->gpu_layer_count = gpu_layers; + impl_->needs_prompt_checkpoint = llama_model_is_recurrent(impl_->model) || + llama_model_is_hybrid(impl_->model) || llama_model_n_swa(impl_->model) > 0; // Performance cores only — see resolveDecodeThreadCount. hardware_concurrency() counted // every logical core including efficiency cores, which both slows barriered matmuls and // burns extra package power for nothing when layers are Metal-offloaded anyway. @@ -471,6 +604,7 @@ void CotabbyInferenceEngine::unloadModel() { } impl_->vocab = nullptr; impl_->model_path.clear(); + impl_->needs_prompt_checkpoint = false; if (impl_->backend_initialized) { llama_backend_free(); @@ -495,6 +629,7 @@ int32_t CotabbyInferenceEngine::createSequence(SamplingConfig config) { auto state = std::make_unique(); state->external_id = id; state->sampler = sampler; + state->single_line = config.single_line; impl_->sequence = std::move(state); return id; } @@ -574,6 +709,9 @@ EngineStatus CotabbyInferenceEngine::decodePrompt(int32_t sequence_id, SequenceState* seq = impl_->findSequence(sequence_id); if (!seq) return EngineStatus::error; + if (!seq->cache_valid) return EngineStatus::error; + if (start_position < 0 || start_position != seq->kv_position_count || + start_position != static_cast(seq->decoded_tokens.size())) return EngineStatus::error; if (seq->cancelled.load(std::memory_order_acquire)) { return EngineStatus::cancelled; @@ -587,6 +725,12 @@ EngineStatus CotabbyInferenceEngine::decodePrompt(int32_t sequence_id, int cursor = 0; int end = token_count; int total_end_position = start_position + token_count; + const int checkpoint_position = std::max(start_position, + total_end_position - Impl::CHECKPOINT_TAIL_TOKENS); + + if (impl_->needs_prompt_checkpoint && checkpoint_position == start_position) { + impl_->savePromptCheckpoint(*seq); + } while (cursor < end) { if (seq->cancelled.load(std::memory_order_acquire)) { @@ -595,6 +739,10 @@ EngineStatus CotabbyInferenceEngine::decodePrompt(int32_t sequence_id, } int chunk_end = std::min(cursor + batch_cap, end); + // Split at the checkpoint boundary once, before decoding the final prompt tail. + if (impl_->needs_prompt_checkpoint && start_position + cursor < checkpoint_position) { + chunk_end = std::min(chunk_end, checkpoint_position - start_position); + } int chunk_size = chunk_end - cursor; batch.n_tokens = static_cast(chunk_size); @@ -612,20 +760,37 @@ EngineStatus CotabbyInferenceEngine::decodePrompt(int32_t sequence_id, } if (llama_decode(impl_->shared_ctx, batch) != 0) { + seq->cache_valid = false; llama_batch_free(batch); return EngineStatus::error; } cursor = chunk_end; + seq->decoded_tokens.insert(seq->decoded_tokens.end(), tokens + cursor - chunk_size, + tokens + cursor); + seq->kv_position_count = start_position + cursor; + if (impl_->needs_prompt_checkpoint && seq->kv_position_count == checkpoint_position) { + impl_->savePromptCheckpoint(*seq); + } } llama_batch_free(batch); seq->kv_position_count = total_end_position; + // Reuse must produce the same sampler state as a fresh request. Discarded model text must + // never enter repetition history or advance this request's random stream. Only actual prompt + // tokens are accepted here; llama_sampler_sample accepts each generated token itself. + llama_sampler_reset(seq->sampler); + for (const llama_token token : seq->decoded_tokens) { + llama_sampler_accept(seq->sampler, token); + } + // First-token word-continuation constraint: when the caret is mid-word, mask new-word-start // tokens for this seed only so the completion continues the current word instead of starting // a new one. The flag clears after this single token. - if (seq->force_word_continuation) { + const bool healing = !seq->completion_prefix.empty(); + if (healing && !impl_->maskCompletionPrefix(*seq, -1)) return EngineStatus::error; + if (!healing && seq->force_word_continuation) { impl_->maskNewWordStarts(-1); seq->force_word_continuation = false; } @@ -633,10 +798,10 @@ EngineStatus CotabbyInferenceEngine::decodePrompt(int32_t sequence_id, // Seed sample: take one token from the prompt's final logits row. The seed will be returned by // the next sampleNext call as-is and feedback-decoded by the call after that. llama_token seed = llama_sampler_sample(seq->sampler, impl_->shared_ctx, -1); - llama_sampler_accept(seq->sampler, seed); + if (healing && !seq->completion_prefix.consume(impl_->tokenPiece(seed))) return EngineStatus::error; seq->seed_token = seed; seq->seed_logprob = seq->compute_logprob ? impl_->computeLogprob(-1, seed) : 0.0f; - seq->seed_argmax_is_eog = impl_->argmaxIsEOG(-1); + seq->seed_argmax_is_eog = !healing && impl_->argmaxIsEOG(-1); seq->has_seed_token = true; seq->has_pending_input = false; @@ -668,6 +833,10 @@ SampleResult CotabbyInferenceEngine::sampleNext(int32_t sequence_id) { result.is_eos = true; return result; } + if (!seq->cache_valid) { + result.is_eos = true; + return result; + } if (seq->cancelled.load(std::memory_order_acquire)) { result.was_cancelled = true; @@ -732,20 +901,37 @@ SampleResult CotabbyInferenceEngine::sampleNext(int32_t sequence_id) { batch.logits[0] = 1; const int status = llama_decode(impl_->shared_ctx, batch); + if (status == 0) { + // Decoding advances memory even if the *next* sample is EOS or cancellation arrives + // immediately afterward. Track that committed position before either early return. + seq->decoded_tokens.push_back(seq->pending_input_token); + seq->kv_position_count++; + } if (status != 0) { + seq->cache_valid = false; result.is_eos = true; } else if (seq->cancelled.load(std::memory_order_acquire)) { result.was_cancelled = true; } else { + const bool healing = !seq->completion_prefix.empty(); + if (healing && !impl_->maskCompletionPrefix(*seq, 0)) { + llama_batch_free(batch); + result.is_eos = true; + return result; + } const llama_token next = llama_sampler_sample(seq->sampler, impl_->shared_ctx, 0); - result.argmax_is_eog = impl_->argmaxIsEOG(0); + if (healing && !seq->completion_prefix.consume(impl_->tokenPiece(next))) { + llama_batch_free(batch); + result.is_eos = true; + return result; + } + result.argmax_is_eog = !healing && impl_->argmaxIsEOG(0); result.token = next; if (next == llama_vocab_eos(impl_->vocab) || llama_vocab_is_eog(impl_->vocab, next)) { result.is_eos = true; } else { - llama_sampler_accept(seq->sampler, next); seq->last_piece.resize(64); while (true) { const int written = llama_token_to_piece( @@ -775,7 +961,6 @@ SampleResult CotabbyInferenceEngine::sampleNext(int32_t sequence_id) { // Feedback decode advanced KV by one position; record the just-sampled // token as input for the next call. - seq->kv_position_count++; seq->pending_input_token = result.token; seq->has_pending_input = true; return result; @@ -789,6 +974,10 @@ bool CotabbyInferenceEngine::trimKV(int32_t sequence_id, int keep_positions) { if (!impl_->shared_ctx) return false; SequenceState* seq = impl_->findSequence(sequence_id); if (!seq) return false; + // Only destruction can recover an uncertain native decode. This check must precede the + // keep_positions == kv_position_count shortcut: an unchanged counter is not proof of valid KV. + if (!seq->cache_valid) return false; + if (keep_positions < 0 || keep_positions > seq->kv_position_count) return false; llama_memory_t memory = llama_get_memory(impl_->shared_ctx); if (!memory) return false; @@ -796,20 +985,34 @@ bool CotabbyInferenceEngine::trimKV(int32_t sequence_id, int keep_positions) { // Serialize with prompt and feedback decode; never remove KV while llama is mutating it. std::lock_guard lock(impl_->decode_mutex); - bool ok = llama_memory_seq_rm( - memory, - Impl::SEQUENCE_ID, - static_cast(keep_positions), - -1 - ); + bool ok; + seq->last_restore_replayed_tokens = 0; + if (keep_positions == 0) { + ok = llama_memory_seq_rm(memory, Impl::SEQUENCE_ID, 0, -1); + seq->prompt_checkpoint.clear(); + seq->checkpoint_memory_bytes = 0; + seq->checkpoint_position = 0; + } else if (keep_positions == seq->kv_position_count) { + // A sampled seed is not decoded until the following sampleNext. Prefill therefore + // already has exactly prompt KV and must not require a recurrent rollback. + ok = true; + } else if (impl_->needs_prompt_checkpoint) { + ok = impl_->restorePromptCheckpoint(*seq, keep_positions); + } else { + ok = llama_memory_seq_rm(memory, Impl::SEQUENCE_ID, keep_positions, -1); + } if (ok) { seq->kv_position_count = keep_positions; + seq->decoded_tokens.resize(keep_positions); // Any seed/pending input is now stale (it would feedback-decode into // a trimmed-away position). Caller must call decodePrompt to re-seed // before the next sampleNext. seq->has_seed_token = false; seq->has_pending_input = false; + // A cancellation applies to the abandoned operation. Rearm only after memory has been + // restored successfully; a failed restore requires the caller to destroy the sequence. + seq->cancelled.store(false, std::memory_order_release); } return ok; } @@ -822,6 +1025,23 @@ void CotabbyInferenceEngine::setForceWordContinuation(int32_t sequence_id, bool } } +std::vector CotabbyInferenceEngine::tokenPiece(int32_t token) const { + if (!impl_) return {}; + const auto bytes = impl_->tokenPiece(token); + return {bytes.begin(), bytes.end()}; +} + +void CotabbyInferenceEngine::setCompletionPrefix(int32_t sequence_id, const uint8_t* bytes, int length) { + if (!impl_) return; + auto* seq = impl_->findSequence(sequence_id); + if (!seq) return; + if (!bytes || length <= 0) { + seq->completion_prefix.clear(); + } else { + seq->completion_prefix.reset(std::string(reinterpret_cast(bytes), length)); + } +} + void CotabbyInferenceEngine::setComputeLogprob(int32_t sequence_id, bool enabled) { if (!impl_) return; SequenceState* seq = impl_->findSequence(sequence_id); @@ -835,9 +1055,12 @@ void CotabbyInferenceEngine::setComputeLogprob(int32_t sequence_id, bool enabled // --------------------------------------------------------------------------- void CotabbyInferenceEngine::cancelSequence(int32_t sequence_id) { - SequenceState* seq = impl_->findSequence(sequence_id); - if (seq) { - seq->cancelled.store(true, std::memory_order_release); + if (!impl_) return; + // Keep ownership locked through the store. Looking up a raw pointer and dropping the lock + // first races sequence destruction from the generation thread. + std::lock_guard lock(impl_->sequence_mutex); + if (impl_->sequence && impl_->sequence->external_id == sequence_id) { + impl_->sequence->cancelled.store(true, std::memory_order_release); } } @@ -860,3 +1083,16 @@ int CotabbyInferenceEngine::getThreadCount() const { int CotabbyInferenceEngine::getGPULayerCount() const { return impl_->gpu_layer_count; } + +CacheDiagnostics CotabbyInferenceEngine::getCacheDiagnostics(int32_t sequence_id) const { + CacheDiagnostics result; + if (!impl_) return result; + const auto* seq = impl_->findSequence(sequence_id); + if (!seq) return result; + result.decoded_token_count = seq->kv_position_count; + result.checkpoint_position = seq->checkpoint_position; + result.checkpoint_bytes = seq->checkpoint_memory_bytes; + result.last_restore_replayed_tokens = seq->last_restore_replayed_tokens; + result.uses_partial_checkpoint = impl_->needs_prompt_checkpoint; + return result; +} diff --git a/Sources/CotabbyInferenceEngine/TokenHealing.cpp b/Sources/CotabbyInferenceEngine/TokenHealing.cpp new file mode 100644 index 0000000..20db5a6 --- /dev/null +++ b/Sources/CotabbyInferenceEngine/TokenHealing.cpp @@ -0,0 +1,88 @@ +#include "TokenHealing.h" + +#include +#include + +namespace cotabby { + +void TokenHealingVocabulary::reset(std::vector entries) { + entries.erase(std::remove_if(entries.begin(), entries.end(), [](const Entry& entry) { + return entry.piece.empty(); + }), entries.end()); + for (auto& entry : entries) { + entry.contains_line_break = entry.piece.find_first_of("\r\n") != std::string::npos; + } + std::sort(entries.begin(), entries.end(), [](const Entry& lhs, const Entry& rhs) { + return lhs.piece < rhs.piece; + }); + entries_ = std::move(entries); +} + +void TokenHealingVocabulary::clear() { + entries_ = {}; +} + +std::vector TokenHealingVocabulary::matchingTokens(std::string_view prefix, bool single_line) const { + std::vector matches; + if (prefix.empty()) return matches; + + const auto lowerBound = [&](std::string_view value) { + return std::lower_bound(entries_.begin(), entries_.end(), value, + [](const Entry& entry, std::string_view key) { return entry.piece < key; }); + }; + + // Tokens that finish the replay and may also introduce new text form one contiguous range. + // Binary search avoids rescanning hundreds of thousands of vocabulary strings per keystroke. + for (auto it = lowerBound(prefix); it != entries_.end(); ++it) { + if (std::string_view(it->piece).substr(0, prefix.size()) != prefix) break; + if (!single_line || !it->contains_line_break) matches.push_back(it->token); + } + + // Prefer a token that accounts for ALL known bytes. A short token such as " i" has high + // probability because it can start many words that contradict the already-typed " int". + // Letting it compete with " intelligence", then forcing "nt" afterward, ranks incompatible + // paths before their likelihood is known. Covering tokens can be compared at the same prefix + // boundary instead. Exact replay remains allowed, so a completed word can end normally. + // Ignore line-break tokens before choosing this branch: a covering token that the sampler + // later forbids must not prevent a valid shorter spelling from being replayed. + if (!matches.empty()) return matches; + + // Multi-token/byte-fallback text may have no covering token. Only then allow shorter pieces; + // the next step applies this same rule to the remaining bytes. Include duplicate spellings. + for (size_t length = 1; length < prefix.size(); ++length) { + const auto fragment = prefix.substr(0, length); + for (auto it = lowerBound(fragment); it != entries_.end() && it->piece == fragment; ++it) { + if (!single_line || !it->contains_line_break) matches.push_back(it->token); + } + } + return matches; +} + +void TokenPrefix::reset(std::string bytes) { + bytes_ = std::move(bytes); + consumed_ = 0; +} + +void TokenPrefix::clear() { + bytes_.clear(); + consumed_ = 0; +} + +bool TokenPrefix::empty() const { + return consumed_ == bytes_.size(); +} + +std::string_view TokenPrefix::remaining() const { + return std::string_view(bytes_).substr(consumed_); +} + +bool TokenPrefix::consume(std::string_view piece) { + if (empty()) return true; + if (piece.empty()) return false; + const auto count = std::min(remaining().size(), piece.size()); + if (remaining().substr(0, count) != piece.substr(0, count)) return false; + consumed_ += count; + return true; +} + +} // namespace cotabby diff --git a/Sources/CotabbyInferenceEngine/TokenHealing.h b/Sources/CotabbyInferenceEngine/TokenHealing.h new file mode 100644 index 0000000..15712fb --- /dev/null +++ b/Sources/CotabbyInferenceEngine/TokenHealing.h @@ -0,0 +1,55 @@ +#pragma once + +#include +#include +#include +#include + +namespace cotabby { + +/// Read-only index of decoded vocabulary bytes, owned once by the loaded engine. +/// +/// A caret can split a token: fixing the existing tokenization of "sched" prevents the model +/// from choosing a token for "schedule". The caller backs up a token and asks this index which +/// replacements can reproduce the exact bytes already typed. This is byte matching, deliberately +/// independent of UTF-8 character boundaries, so byte-fallback tokenizers work as well. +class TokenHealingVocabulary { +public: + struct Entry { + int32_t token; + std::string piece; + // Derived once at reset, so single-line filtering adds no per-step string scanning. + bool contains_line_break = false; + }; + + void reset(std::vector entries); + void clear(); + + /// Returns tokens covering the complete remaining prefix, including exact replay. Falls back + /// to shorter matching pieces only when no covering token exists (e.g. byte-fallback text). + /// Empty/control tokens must not be indexed: they would make no progress through the prefix. + std::vector matchingTokens(std::string_view prefix, bool single_line = false) const; + +private: + std::vector entries_; +}; + +/// Per-request state for replaying the prompt's backed-up bytes before exposing new text. +/// The engine owns this beside its sampler; resetting a generation also resets this value. +class TokenPrefix { +public: + void reset(std::string bytes); + void clear(); + bool empty() const; + std::string_view remaining() const; + + /// Returns false on a mismatching sample without advancing. Callers fail closed rather than + /// exposing an altered spelling of text that is already in the user's document. + bool consume(std::string_view piece); + +private: + std::string bytes_; + size_t consumed_ = 0; +}; + +} // namespace cotabby diff --git a/Sources/CotabbyInferenceEngine/include/CotabbyInferenceEngine.h b/Sources/CotabbyInferenceEngine/include/CotabbyInferenceEngine.h index 383bcaf..d15151e 100644 --- a/Sources/CotabbyInferenceEngine/include/CotabbyInferenceEngine.h +++ b/Sources/CotabbyInferenceEngine/include/CotabbyInferenceEngine.h @@ -40,6 +40,15 @@ enum class EngineStatus : int { not_loaded = 3, }; +// A snapshot for logs/tests; no prompt or generated content crosses this diagnostics boundary. +struct CacheDiagnostics { + int decoded_token_count = 0; + int checkpoint_position = 0; + uint64_t checkpoint_bytes = 0; + int last_restore_replayed_tokens = 0; + bool uses_partial_checkpoint = false; +}; + class CotabbyInferenceEngine { public: CotabbyInferenceEngine(); @@ -64,6 +73,11 @@ class CotabbyInferenceEngine { // Tokenization (thread-safe, read-only on vocab) std::vector tokenize(const char* text, int text_length) const; + // Plain decoded bytes for one printable token; special/unknown tokens return empty. + std::vector tokenPiece(int32_t token) const; + // The next continuation must reproduce these already-typed bytes before adding new text. + // Set before decodePrompt. Matching may span several byte-fallback tokens. + void setCompletionPrefix(int32_t sequence_id, const uint8_t* bytes, int length); // Prompt decoding EngineStatus decodePrompt(int32_t sequence_id, const int32_t* tokens, int token_count, @@ -94,6 +108,7 @@ class CotabbyInferenceEngine { int getBatchSize() const; int getThreadCount() const; int getGPULayerCount() const; + CacheDiagnostics getCacheDiagnostics(int32_t sequence_id) const; private: struct Impl; diff --git a/Tests/CotabbyInferenceTests/LlamaMiddlewareTests.swift b/Tests/CotabbyInferenceTests/LlamaMiddlewareTests.swift index b7741ba..6b48fe5 100644 --- a/Tests/CotabbyInferenceTests/LlamaMiddlewareTests.swift +++ b/Tests/CotabbyInferenceTests/LlamaMiddlewareTests.swift @@ -92,9 +92,8 @@ final class LlamaMiddlewareTests: XCTestCase { } XCTAssertFalse(generated.isEmpty, "Expected at least one generated token") - // Hybrid/recurrent and SWA model caches can reject partial KV removal. Cotabby treats - // that as a cache-reuse miss and rebuilds the sequence, so lifecycle coverage must not - // require a model-specific optimization to succeed. + // This short prompt may not have a nonempty partial-state checkpoint. Cotabby treats + // that as a miss and rebuilds, so lifecycle coverage does not require reuse here. _ = engine.trimKV(sequence, Int32(tokens.count)) engine.destroySequence(sequence) @@ -248,9 +247,208 @@ final class LlamaMiddlewareTests: XCTestCase { XCTAssertGreaterThan(steps, 0) engine.destroySequence(sequence) } + /// Reuse is correct only when it agrees with a fresh request: abandoned predictions must + /// neither change the model state nor enter the repetition penalty's token history. + func testRestoredPromptMatchesColdGenerationAndBoundsReplay() throws { + var engine = CotabbyInferenceEngine() + let modelPath = try Self.modelPath() + XCTAssertEqual(engine.loadModel(modelPath, -1, 1024, 256), .ok) + defer { engine.unloadModel() } + let prompt = + "Hi Alex, thanks for sending the project update. I will review the schedule and send you my" + var tokens = Array(engine.tokenize(prompt, Int32(prompt.utf8.count))) + let config = Self.samplingConfig(temperature: 0, repetitionPenalty: 1.1) + let sequence = engine.createSequence(config) + XCTAssertEqual(engine.decodePrompt(sequence, &tokens, Int32(tokens.count), 0), .ok) + let cold = Self.sampleTokens(engine: &engine, sequence: sequence, count: 16) + XCTAssertFalse(cold.isEmpty) + XCTAssertTrue(engine.trimKV(sequence, Int32(tokens.count))) + let restored = engine.getCacheDiagnostics(sequence) + XCTAssertEqual(restored.decoded_token_count, Int32(tokens.count)) + XCTAssertLessThanOrEqual(restored.last_restore_replayed_tokens, 8) + if restored.uses_partial_checkpoint { + XCTAssertGreaterThan(restored.checkpoint_bytes, 0) + XCTAssertLessThanOrEqual(restored.checkpoint_bytes, 128 * 1024 * 1024) + } + + XCTAssertTrue(engine.trimKV(sequence, Int32(tokens.count - 1))) + var last = [tokens.last!] + XCTAssertEqual(engine.decodePrompt(sequence, &last, 1, Int32(tokens.count - 1)), .ok) + XCTAssertEqual(Self.sampleTokens(engine: &engine, sequence: sequence, count: 16), cold) + engine.destroySequence(sequence) + } + + /// Greedy equivalence cannot detect a stale random stream. A fixed nonzero seed and positive + /// temperature make this request exercise the distribution sampler across repeated restores. + func testRestoredSeededSamplingMatchesColdGeneration() throws { + var engine = CotabbyInferenceEngine() + let modelPath = try Self.modelPath() + XCTAssertEqual(engine.loadModel(modelPath, -1, 1024, 256), .ok) + defer { engine.unloadModel() } + let prompt = "Hi Alex, thanks for sending the project update. I will review the schedule and send you my" + var tokens = Array(engine.tokenize(prompt, Int32(prompt.utf8.count))) + let config = Self.samplingConfig(temperature: 0.7, repetitionPenalty: 1.1, seed: 42) + let sequence = engine.createSequence(config) + XCTAssertEqual(engine.decodePrompt(sequence, &tokens, Int32(tokens.count), 0), .ok) + let cold = Self.sampleTokens(engine: &engine, sequence: sequence, count: 24) + XCTAssertGreaterThan(cold.count, 1, "The test must advance the random stream beyond its seed sample") + + for _ in 0 ..< 2 { + XCTAssertTrue(engine.trimKV(sequence, Int32(tokens.count - 1))) + var last = [tokens.last!] + XCTAssertEqual(engine.decodePrompt(sequence, &last, 1, Int32(tokens.count - 1)), .ok) + XCTAssertEqual( + Self.sampleTokens(engine: &engine, sequence: sequence, count: 24), + cold, + "Restoration must reset both random state and prompt-only repetition history" + ) + } + engine.destroySequence(sequence) + } + + func testCancellationCanBeRearmedOnlyBySuccessfulRestoration() throws { + var engine = CotabbyInferenceEngine() + let modelPath = try Self.modelPath() + XCTAssertEqual(engine.loadModel(modelPath, -1, 1024, 256), .ok) + defer { engine.unloadModel() } + let prompt = + "Thank you for your thoughtful comments on the document. I have updated the draft to include" + var tokens = Array(engine.tokenize(prompt, Int32(prompt.utf8.count))) + let sequence = engine.createSequence(Self.samplingConfig(temperature: 0)) + XCTAssertEqual(engine.decodePrompt(sequence, &tokens, Int32(tokens.count), 0), .ok) + _ = Self.sampleTokens(engine: &engine, sequence: sequence, count: 4) + engine.cancelSequence(sequence) + XCTAssertTrue(engine.sampleNext(sequence).was_cancelled) + XCTAssertFalse(engine.trimKV(sequence, Int32(tokens.count + 100))) + XCTAssertTrue(engine.sampleNext(sequence).was_cancelled) + XCTAssertTrue(engine.trimKV(sequence, Int32(tokens.count - 1))) + var last = [tokens.last!] + XCTAssertEqual(engine.decodePrompt(sequence, &last, 1, Int32(tokens.count - 1)), .ok) + XCTAssertFalse(engine.sampleNext(sequence).was_cancelled) + engine.destroySequence(sequence) + } + + func testFailedPromptDecodeCannotBeReusedAtTrackedPosition() throws { + var engine = CotabbyInferenceEngine() + let modelPath = try Self.modelPath() + XCTAssertEqual(engine.loadModel(modelPath, -1, 64, 32), .ok) + defer { engine.unloadModel() } + let sequence = engine.createSequence(Self.samplingConfig(temperature: 0)) + // Exceed a deliberately small attention context without an artificial failure hook. + // Earlier chunks commit successfully before the first chunk with no free KV cells fails. + let oversized = String(repeating: "one two three four five six seven eight ", count: 128) + var tokens = Array(engine.tokenize(oversized, Int32(oversized.utf8.count))) + let status = engine.decodePrompt(sequence, &tokens, Int32(tokens.count), 0) + if status == .ok { + engine.destroySequence(sequence) + throw XCTSkip("This model does not exhaust attention KV for an oversized prompt") + } + XCTAssertEqual(status, .error) + let committed = engine.getCacheDiagnostics(sequence).decoded_token_count + XCTAssertGreaterThan(committed, 0, "The failure should happen after an earlier successful batch") + XCTAssertFalse( + engine.trimKV(sequence, committed), + "A failed decode must not pass the equal-position cache shortcut" + ) + engine.destroySequence(sequence) + + let replacement = engine.createSequence(Self.samplingConfig(temperature: 0)) + let shortPrompt = "The quick brown fox" + var shortTokens = Array(engine.tokenize(shortPrompt, Int32(shortPrompt.utf8.count))) + XCTAssertEqual(engine.decodePrompt(replacement, &shortTokens, Int32(shortTokens.count), 0), .ok) + engine.destroySequence(replacement) + } + + func testHealingReproducesTypedBytesIncludingTrailingWhitespace() throws { + var engine = CotabbyInferenceEngine() + let modelPath = try Self.modelPath() + XCTAssertEqual(engine.loadModel(modelPath, -1, 1024, 256), .ok) + defer { engine.unloadModel() } + for prompt in ["Please send me the sched", "Please send me the ", "The café serves "] { + var tokens = Array(engine.tokenize(prompt, Int32(prompt.utf8.count))) + XCTAssertGreaterThan(tokens.count, 1) + let prefix = Array(engine.tokenPiece(tokens.removeLast())) + XCTAssertFalse(prefix.isEmpty) + XCTAssertTrue(Array(prompt.utf8).suffix(prefix.count).elementsEqual(prefix)) + let sequence = engine.createSequence(Self.samplingConfig(temperature: 0)) + prefix.withUnsafeBufferPointer { + engine.setCompletionPrefix(sequence, $0.baseAddress, Int32($0.count)) + } + XCTAssertEqual(engine.decodePrompt(sequence, &tokens, Int32(tokens.count), 0), .ok) + var generated: [UInt8] = [] + for _ in 0.. prefix.count { break } + } + XCTAssertTrue(generated.starts(with: prefix), "Healing must reproduce the exact typed prefix") + XCTAssertGreaterThan(generated.count, prefix.count) + engine.destroySequence(sequence) + } + } + + /// This reports native prompt work separately from typing-to-display latency in the app. + /// No wall-clock threshold is asserted: CI hardware and Metal warmup vary substantially. + func testWarmPromptDecodeReportsLatency() throws { + var engine = CotabbyInferenceEngine() + let modelPath = try Self.modelPath() + XCTAssertEqual(engine.loadModel(modelPath, -1, 1024, 256), .ok) + defer { engine.unloadModel() } + for paragraphCount in [2, 16, 64] { + let prompt = + String( + repeating: + "We are reviewing the project schedule and collecting feedback from the team. ", + count: paragraphCount) + + "Please send the updated schedule by" + var tokens = Array(engine.tokenize(prompt, Int32(prompt.utf8.count))) + var coldMilliseconds: [Double] = [] + var warmMilliseconds: [Double] = [] + var checkpointBytes: UInt64 = 0 + for iteration in 0..<4 { + let sequence = engine.createSequence(Self.samplingConfig(temperature: 0)) + let coldStart = ContinuousClock.now + XCTAssertEqual(engine.decodePrompt(sequence, &tokens, Int32(tokens.count), 0), .ok) + let coldDuration = coldStart.duration(to: .now) + let coldToken = engine.sampleNext(sequence).token + let warmStart = ContinuousClock.now + XCTAssertTrue(engine.trimKV(sequence, Int32(tokens.count - 1))) + var last = [tokens.last!] + XCTAssertEqual(engine.decodePrompt(sequence, &last, 1, Int32(tokens.count - 1)), .ok) + let warmDuration = warmStart.duration(to: .now) + XCTAssertEqual(engine.sampleNext(sequence).token, coldToken) + checkpointBytes = engine.getCacheDiagnostics(sequence).checkpoint_bytes + // Exclude the first pass so graph/kernel initialization does not dominate the report. + if iteration > 0 { + coldMilliseconds.append(Self.milliseconds(coldDuration)) + warmMilliseconds.append(Self.milliseconds(warmDuration)) + } + engine.destroySequence(sequence) + } + print( + "CACHE_BENCHMARK prompt_tokens=\(tokens.count) cold_median_ms=\(coldMilliseconds.sorted()[1]) warm_median_ms=\(warmMilliseconds.sorted()[1]) checkpoint_bytes=\(checkpointBytes)" + ) + } + } } private extension LlamaMiddlewareTests { + static func milliseconds(_ duration: Duration) -> Double { + Double(duration.components.seconds) * 1_000 + + Double(duration.components.attoseconds) / 1_000_000_000_000_000 + } + static func sampleTokens( + engine: inout CotabbyInferenceEngine, sequence: Int32, count: Int + ) -> [Int32] { + var tokens: [Int32] = [] + for _ in 0.. String { guard let path = ProcessInfo.processInfo.environment["COTABBY_TEST_MODEL_PATH"], FileManager.default.fileExists(atPath: path) else { diff --git a/Tests/TokenHealingTests.cpp b/Tests/TokenHealingTests.cpp new file mode 100644 index 0000000..6fb140b --- /dev/null +++ b/Tests/TokenHealingTests.cpp @@ -0,0 +1,75 @@ +#include "../Sources/CotabbyInferenceEngine/TokenHealing.h" + +#include +#include +#include + +using cotabby::TokenHealingVocabulary; +using cotabby::TokenPrefix; + +static std::vector sorted(std::vector tokens) { + std::sort(tokens.begin(), tokens.end()); + return tokens; +} + +int main() { + TokenHealingVocabulary vocabulary; + vocabulary.reset({ + {1, "sched"}, {2, "schedule"}, {3, "scheduling"}, {4, "sch"}, + {5, "send"}, {6, " schedule"}, {7, " "}, {8, ""}, {9, "sched"}, + {10, "\n"}, {11, "\n\n"}, {12, std::string("\xC3", 1)}, {13, "éclair"} + }); + + // Unfinished words can choose a whole-word token instead of being trapped in the original + // tokenization. A completed word still permits replaying itself and predicting a space next. + assert(sorted(vocabulary.matchingTokens("sched")) == std::vector({1, 2, 3, 9})); + assert(sorted(vocabulary.matchingTokens("schedule")) == std::vector({2})); + assert(sorted(vocabulary.matchingTokens(" ")) == std::vector({6, 7})); + assert(sorted(vocabulary.matchingTokens("\n")) == std::vector({10, 11})); + assert(sorted(vocabulary.matchingTokens("é")) == std::vector({13})); + // A fragment must not lose its context just because a high-probability shorter token can + // start it. Neither " i" nor " in" accounts for the known final 't' of " int". + TokenHealingVocabulary technical; + technical.reset({{20, " i"}, {21, " in"}, {22, " int"}, {23, " intelligence"}, {24, " internal"}}); + assert(sorted(technical.matchingTokens(" int")) == std::vector({22, 23, 24})); + // A multi-token spelling that has no complete covering token still progresses byte by byte. + assert(sorted(vocabulary.matchingTokens("schedx")) == std::vector({1, 4, 9})); + assert(sorted(vocabulary.matchingTokens("éx")) == std::vector({12})); + TokenHealingVocabulary line_breaks; + line_breaks.reset({{30, "foo\n"}, {31, "f"}, {32, "oo"}, {33, "foo\r\n"}}); + assert(sorted(line_breaks.matchingTokens("foo")) == std::vector({30, 33})); + // Single-line sampling masks the only covering tokens. Shorter safe pieces must remain + // available, or all logits become -infinity before the sampler can complete the replay. + assert(sorted(line_breaks.matchingTokens("foo", true)) == std::vector({31})); + assert(sorted(line_breaks.matchingTokens("oo", true)) == std::vector({32})); + assert(vocabulary.matchingTokens("xyz").empty()); + assert(vocabulary.matchingTokens("").empty()); + + TokenPrefix prefix; + prefix.reset("sched"); + assert(!prefix.consume("send")); + assert(prefix.remaining() == "sched"); + assert(!prefix.consume("")); + assert(prefix.consume("sch")); + assert(prefix.remaining() == "ed"); + assert(prefix.consume("edule")); + assert(prefix.empty()); + assert(prefix.consume(" tomorrow")); + + // Prefix state can end inside a UTF-8 code point. Neither comparison nor consumption may + // decode an isolated byte into a replacement character or silently drop it. + prefix.reset("é"); + assert(prefix.consume(std::string("\xC3", 1))); + assert(prefix.consume(std::string("\xA9", 1))); + assert(prefix.empty()); + prefix.reset(" "); + assert(prefix.consume(" schedule")); + assert(prefix.empty()); + prefix.reset("old"); + prefix.clear(); + assert(prefix.empty()); + vocabulary.clear(); + assert(vocabulary.matchingTokens("sched").empty()); + + std::cout << "Token healing tests passed\n"; +}