use serde::{Deserialize, Serialize};
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum FinishReason {
Eos,
StopSequence,
GrammarComplete,
MaxTokens,
Cancelled,
}
#[derive(Debug, Clone, Default)]
pub struct GenerationCancellationToken {
cancelled: Arc<AtomicBool>,
}
impl GenerationCancellationToken {
pub fn new() -> Self {
Self::default()
}
pub fn cancel(&self) {
self.cancelled.store(true, Ordering::Release);
}
pub fn is_cancelled(&self) -> bool {
self.cancelled.load(Ordering::Acquire)
}
}
#[derive(Debug, Clone, Copy, Default, Eq, PartialEq)]
pub struct TokenTerminalSignals {
pub stop_sequence: bool,
pub grammar_complete: bool,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub struct TokenCommit {
pub token_id: u32,
pub position: usize,
pub finish_reason: Option<FinishReason>,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct GenerationSequence {
max_tokens: usize,
eos_token_ids: Vec<u32>,
tokens: Vec<u32>,
finish_reason: Option<FinishReason>,
}
impl GenerationSequence {
pub fn new(max_tokens: usize, eos_token_ids: impl IntoIterator<Item = u32>) -> Self {
let mut eos_token_ids = eos_token_ids.into_iter().collect::<Vec<_>>();
eos_token_ids.sort_unstable();
eos_token_ids.dedup();
Self {
max_tokens,
eos_token_ids,
tokens: Vec::with_capacity(max_tokens),
finish_reason: (max_tokens == 0).then_some(FinishReason::MaxTokens),
}
}
pub fn commit(
&mut self,
token_id: u32,
signals: TokenTerminalSignals,
) -> Result<TokenCommit, GenerationError> {
if self.finish_reason.is_some() {
return Err(GenerationError::AlreadyFinished);
}
let position = self.tokens.len();
self.tokens.push(token_id);
let finish_reason = signals
.stop_sequence
.then_some(FinishReason::StopSequence)
.or_else(|| {
signals
.grammar_complete
.then_some(FinishReason::GrammarComplete)
})
.or_else(|| {
self.eos_token_ids
.binary_search(&token_id)
.is_ok()
.then_some(FinishReason::Eos)
})
.or_else(|| (self.tokens.len() == self.max_tokens).then_some(FinishReason::MaxTokens));
self.finish_reason = finish_reason;
Ok(TokenCommit {
token_id,
position,
finish_reason,
})
}
pub fn cancel(&mut self) -> bool {
if self.tokens.is_empty() && self.finish_reason == Some(FinishReason::MaxTokens) {
self.finish_reason = Some(FinishReason::Cancelled);
true
} else if self.finish_reason.is_some() {
false
} else {
self.finish_reason = Some(FinishReason::Cancelled);
true
}
}
pub fn observe_cancellation(&mut self, cancellation: &GenerationCancellationToken) -> bool {
cancellation.is_cancelled() && self.cancel()
}
pub fn tokens(&self) -> &[u32] {
&self.tokens
}
pub fn into_tokens(self) -> Vec<u32> {
self.tokens
}
pub fn remaining(&self) -> usize {
self.max_tokens.saturating_sub(self.tokens.len())
}
pub const fn max_tokens(&self) -> usize {
self.max_tokens
}
pub const fn finish_reason(&self) -> Option<FinishReason> {
self.finish_reason
}
pub const fn is_finished(&self) -> bool {
self.finish_reason.is_some()
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub enum SpeculativeTail {
Replacement,
Bonus,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct SpeculativeRound {
proposal_count: usize,
accepted: usize,
committed_tokens: Vec<u32>,
tail: Option<SpeculativeTail>,
terminal: bool,
}
impl SpeculativeRound {
pub fn new(proposal_count: usize) -> Result<Self, GenerationError> {
if proposal_count == 0 {
return Err(GenerationError::EmptyProposalBlock);
}
Ok(Self {
proposal_count,
accepted: 0,
committed_tokens: Vec::with_capacity(proposal_count + 1),
tail: None,
terminal: false,
})
}
pub fn accept(&mut self, token: u32, terminal: bool) -> Result<(), GenerationError> {
if self.tail.is_some() || self.accepted == self.proposal_count || self.terminal {
return Err(GenerationError::InvalidSpeculativeTransition);
}
self.accepted += 1;
self.committed_tokens.push(token);
self.terminal = terminal;
Ok(())
}
pub fn reject_with(&mut self, token: u32, terminal: bool) -> Result<(), GenerationError> {
if self.tail.is_some() || self.accepted == self.proposal_count || self.terminal {
return Err(GenerationError::InvalidSpeculativeTransition);
}
self.tail = Some(SpeculativeTail::Replacement);
self.committed_tokens.push(token);
self.terminal = terminal;
Ok(())
}
pub fn bonus(&mut self, token: u32, terminal: bool) -> Result<(), GenerationError> {
if self.tail.is_some() || self.accepted != self.proposal_count || self.terminal {
return Err(GenerationError::InvalidSpeculativeTransition);
}
self.tail = Some(SpeculativeTail::Bonus);
self.committed_tokens.push(token);
self.terminal = terminal;
Ok(())
}
pub const fn is_full_acceptance(&self) -> bool {
self.accepted == self.proposal_count
}
pub fn commit_plan(&self) -> Result<SpeculativeCommitPlan<'_>, GenerationError> {
if !self.terminal && self.tail.is_none() {
return Err(GenerationError::IncompleteSpeculativeRound);
}
Ok(SpeculativeCommitPlan {
accepted_proposals: self.accepted,
committed_tokens: &self.committed_tokens,
verified_inputs: if self.tail.is_some() {
1 + self.accepted
} else {
self.accepted
},
full_acceptance: self.accepted == self.proposal_count,
tail: self.tail,
terminal: self.terminal,
})
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub struct SpeculativeCommitPlan<'a> {
pub accepted_proposals: usize,
pub committed_tokens: &'a [u32],
pub verified_inputs: usize,
pub full_acceptance: bool,
pub tail: Option<SpeculativeTail>,
pub terminal: bool,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub enum OptimisticReuseDecision {
DiscardTerminal,
DiscardMismatch,
MatchedConsumed,
MatchedRetained,
}
pub fn resolve_optimistic_reuse(
assumed_prefix: &[u32],
canonical_prefix: &[u32],
optimistic_tokens: &[u32],
bonus: u32,
terminal: bool,
) -> Result<OptimisticReuseDecision, GenerationError> {
if assumed_prefix != canonical_prefix {
return Err(GenerationError::OptimisticPrefixDiverged);
}
let first = optimistic_tokens
.first()
.ok_or(GenerationError::EmptyOptimisticBranch)?;
if terminal {
return Ok(OptimisticReuseDecision::DiscardTerminal);
}
if *first != bonus {
return Ok(OptimisticReuseDecision::DiscardMismatch);
}
Ok(if optimistic_tokens.len() == 1 {
OptimisticReuseDecision::MatchedConsumed
} else {
OptimisticReuseDecision::MatchedRetained
})
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SpeculativeConfig {
pub max_tokens: usize,
pub max_draft_tokens: usize,
pub temperature: f32,
pub eos_token_ids: Vec<u32>,
}
impl Default for SpeculativeConfig {
fn default() -> Self {
Self {
max_tokens: 256,
max_draft_tokens: 4,
temperature: 0.0,
eos_token_ids: Vec::new(),
}
}
}
impl SpeculativeConfig {
pub fn validate(&self) -> Result<(), GenerationError> {
if self.max_draft_tokens == 0 {
return Err(GenerationError::ZeroDraftTokens);
}
if !self.temperature.is_finite() || self.temperature < 0.0 {
return Err(GenerationError::InvalidTemperature(self.temperature));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct SpeculativeSchedulerOptions {
pub max_in_flight_verifications: usize,
pub max_optimistic_branches: usize,
pub lookahead_blocks: usize,
pub adaptive_lookahead: bool,
pub adaptive_lookahead_min_blocks: usize,
}
impl Default for SpeculativeSchedulerOptions {
fn default() -> Self {
Self {
max_in_flight_verifications: 1,
max_optimistic_branches: 1,
lookahead_blocks: 1,
adaptive_lookahead: true,
adaptive_lookahead_min_blocks: 4,
}
}
}
impl SpeculativeSchedulerOptions {
pub fn with_lookahead(mut self, enabled: bool) -> Self {
self.lookahead_blocks = usize::from(enabled);
if enabled {
self.max_optimistic_branches = self.max_optimistic_branches.max(1);
}
self
}
pub fn validate(self) -> Result<Self, GenerationError> {
if self.max_in_flight_verifications == 0 {
return Err(GenerationError::ZeroInFlightVerifications);
}
if self.lookahead_blocks > 1 {
return Err(GenerationError::TooManyLookaheadBlocks);
}
if self.lookahead_blocks > 0 && self.max_optimistic_branches == 0 {
return Err(GenerationError::LookaheadWithoutBranchCapacity);
}
if self.lookahead_blocks > 0
&& self.adaptive_lookahead
&& self.adaptive_lookahead_min_blocks == 0
{
return Err(GenerationError::ZeroAdaptiveLookaheadWindow);
}
Ok(self)
}
}
#[derive(Debug, Clone, Copy, Eq, Hash, PartialEq, Serialize, Deserialize)]
pub struct SpeculativeRequestId(usize);
impl SpeculativeRequestId {
pub const fn new(index: usize) -> Self {
Self(index)
}
pub const fn index(self) -> usize {
self.0
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum SpeculativeRequestStatus {
Prefill,
ReadyToDraft,
ReadyToSubmitVerification,
TargetVerificationInFlight,
OptimisticDraftRunning,
OptimisticDraftReady,
VerificationResolution,
Completed,
Cancelled,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub enum SpeculativeCancellationDisposition {
AlreadyTerminal,
CancelNow,
Deferred,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct SpeculativeRequestLifecycle {
status: SpeculativeRequestStatus,
cancellation_pending: bool,
}
impl Default for SpeculativeRequestLifecycle {
fn default() -> Self {
Self::new()
}
}
impl SpeculativeRequestLifecycle {
pub const fn new() -> Self {
Self {
status: SpeculativeRequestStatus::Prefill,
cancellation_pending: false,
}
}
pub const fn completed() -> Self {
Self {
status: SpeculativeRequestStatus::Completed,
cancellation_pending: false,
}
}
pub const fn cancelled() -> Self {
Self {
status: SpeculativeRequestStatus::Cancelled,
cancellation_pending: false,
}
}
pub const fn status(&self) -> SpeculativeRequestStatus {
self.status
}
pub const fn cancellation_pending(&self) -> bool {
self.cancellation_pending
}
pub const fn is_terminal(&self) -> bool {
matches!(
self.status,
SpeculativeRequestStatus::Completed | SpeculativeRequestStatus::Cancelled
)
}
pub fn request_cancellation(
&mut self,
submission_retained: bool,
) -> Result<SpeculativeCancellationDisposition, GenerationError> {
if self.is_terminal() {
return Ok(SpeculativeCancellationDisposition::AlreadyTerminal);
}
if submission_retained {
self.cancellation_pending = true;
Ok(SpeculativeCancellationDisposition::Deferred)
} else {
self.transition(SpeculativeRequestStatus::Cancelled)?;
Ok(SpeculativeCancellationDisposition::CancelNow)
}
}
pub fn transition(&mut self, next: SpeculativeRequestStatus) -> Result<(), GenerationError> {
let allowed = matches!(
(self.status, next),
(
SpeculativeRequestStatus::Prefill,
SpeculativeRequestStatus::ReadyToDraft
) | (
SpeculativeRequestStatus::Prefill,
SpeculativeRequestStatus::Completed
) | (
SpeculativeRequestStatus::Prefill,
SpeculativeRequestStatus::Cancelled
) | (
SpeculativeRequestStatus::ReadyToDraft,
SpeculativeRequestStatus::ReadyToSubmitVerification
) | (
SpeculativeRequestStatus::ReadyToDraft,
SpeculativeRequestStatus::Completed
) | (
SpeculativeRequestStatus::ReadyToDraft,
SpeculativeRequestStatus::Cancelled
) | (
SpeculativeRequestStatus::ReadyToSubmitVerification,
SpeculativeRequestStatus::TargetVerificationInFlight
) | (
SpeculativeRequestStatus::ReadyToSubmitVerification,
SpeculativeRequestStatus::Cancelled
) | (
SpeculativeRequestStatus::TargetVerificationInFlight,
SpeculativeRequestStatus::OptimisticDraftRunning
) | (
SpeculativeRequestStatus::TargetVerificationInFlight,
SpeculativeRequestStatus::VerificationResolution
) | (
SpeculativeRequestStatus::OptimisticDraftRunning,
SpeculativeRequestStatus::OptimisticDraftReady
) | (
SpeculativeRequestStatus::OptimisticDraftReady,
SpeculativeRequestStatus::VerificationResolution
) | (
SpeculativeRequestStatus::VerificationResolution,
SpeculativeRequestStatus::ReadyToDraft
) | (
SpeculativeRequestStatus::VerificationResolution,
SpeculativeRequestStatus::Completed
) | (
SpeculativeRequestStatus::VerificationResolution,
SpeculativeRequestStatus::Cancelled
)
);
if !allowed {
return Err(GenerationError::InvalidSpeculativeStatusTransition {
from: self.status,
to: next,
});
}
self.status = next;
if self.is_terminal() {
self.cancellation_pending = false;
}
Ok(())
}
}
#[derive(Debug, Clone, Default, PartialEq, Deserialize, Serialize)]
pub struct CheckpointGenerationConfig {
#[serde(default)]
pub do_sample: Option<bool>,
#[serde(default)]
pub temperature: Option<f32>,
#[serde(default)]
pub top_k: Option<i32>,
#[serde(default)]
pub top_p: Option<f32>,
#[serde(default)]
pub min_p: Option<f32>,
#[serde(default)]
pub repetition_penalty: Option<f32>,
#[serde(default)]
pub repeat_last_n: Option<i32>,
#[serde(default)]
pub frequency_penalty: Option<f32>,
#[serde(default)]
pub presence_penalty: Option<f32>,
#[serde(default)]
pub max_new_tokens: Option<usize>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Serialize, Deserialize)]
pub struct GenerationConfigOverrides {
pub do_sample: Option<bool>,
pub temperature: Option<f32>,
pub top_k: Option<i32>,
pub top_p: Option<f32>,
pub min_p: Option<f32>,
pub repetition_penalty: Option<f32>,
pub repeat_last_n: Option<i32>,
pub frequency_penalty: Option<f32>,
pub presence_penalty: Option<f32>,
pub max_new_tokens: Option<usize>,
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct ResolvedGenerationConfig {
pub do_sample: bool,
pub temperature: f32,
pub top_k: i32,
pub top_p: f32,
pub min_p: f32,
pub repetition_penalty: f32,
pub repeat_last_n: i32,
pub frequency_penalty: f32,
pub presence_penalty: f32,
pub max_new_tokens: Option<usize>,
}
pub fn resolve_generation_config(
checkpoint: Option<&CheckpointGenerationConfig>,
overrides: GenerationConfigOverrides,
) -> Result<ResolvedGenerationConfig, GenerationError> {
let checkpoint_present = checkpoint.is_some();
let checkpoint = checkpoint.cloned().unwrap_or_default();
let (do_sample, temperature) = if let Some(do_sample) = overrides.do_sample {
if do_sample {
(
true,
overrides
.temperature
.or(checkpoint.temperature)
.unwrap_or(1.0),
)
} else {
(false, 0.0)
}
} else if let Some(temperature) = overrides.temperature {
(temperature > 0.0, temperature)
} else if checkpoint.do_sample.unwrap_or(false) {
(true, checkpoint.temperature.unwrap_or(1.0))
} else {
(false, 0.0)
};
let resolved = ResolvedGenerationConfig {
do_sample,
temperature,
top_k: overrides
.top_k
.or(checkpoint.top_k)
.unwrap_or(if checkpoint_present { 50 } else { 40 }),
top_p: overrides
.top_p
.or(checkpoint.top_p)
.unwrap_or(if checkpoint_present { 1.0 } else { 0.95 }),
min_p: overrides
.min_p
.or(checkpoint.min_p)
.unwrap_or(if checkpoint_present { 0.0 } else { 0.05 }),
repetition_penalty: overrides
.repetition_penalty
.or(checkpoint.repetition_penalty)
.unwrap_or(1.0),
repeat_last_n: overrides
.repeat_last_n
.or(checkpoint.repeat_last_n)
.unwrap_or(64),
frequency_penalty: overrides
.frequency_penalty
.or(checkpoint.frequency_penalty)
.unwrap_or(0.0),
presence_penalty: overrides
.presence_penalty
.or(checkpoint.presence_penalty)
.unwrap_or(0.0),
max_new_tokens: overrides.max_new_tokens.or(checkpoint.max_new_tokens),
};
if !resolved.temperature.is_finite() || resolved.temperature < 0.0 {
return Err(GenerationError::InvalidTemperature(resolved.temperature));
}
if resolved.do_sample && resolved.temperature == 0.0 {
return Err(GenerationError::StochasticZeroTemperature);
}
if resolved.top_k < 0 {
return Err(GenerationError::InvalidTopK(resolved.top_k));
}
if !resolved.top_p.is_finite() || !(0.0..=1.0).contains(&resolved.top_p) {
return Err(GenerationError::InvalidTopP(resolved.top_p));
}
if !resolved.min_p.is_finite() || !(0.0..=1.0).contains(&resolved.min_p) {
return Err(GenerationError::InvalidMinP(resolved.min_p));
}
if !resolved.repetition_penalty.is_finite() || resolved.repetition_penalty <= 0.0 {
return Err(GenerationError::InvalidRepetitionPenalty(
resolved.repetition_penalty,
));
}
if !resolved.frequency_penalty.is_finite() {
return Err(GenerationError::InvalidFrequencyPenalty(
resolved.frequency_penalty,
));
}
if !resolved.presence_penalty.is_finite() {
return Err(GenerationError::InvalidPresencePenalty(
resolved.presence_penalty,
));
}
if resolved.max_new_tokens == Some(0) {
return Err(GenerationError::ZeroTokenBudget);
}
Ok(resolved)
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub enum SemanticEvent {
ReasoningDelta(String),
TextDelta(String),
ToolCallStart {
index: usize,
id: String,
name: String,
},
ToolArgumentsDelta {
index: usize,
json_fragment: String,
},
ToolCallEnd,
Finished {
reason: FinishReason,
},
}
#[derive(Debug, Clone, PartialEq, thiserror::Error)]
pub enum GenerationError {
#[error("generation has already finished")]
AlreadyFinished,
#[error("speculative verification requires at least one proposal")]
EmptyProposalBlock,
#[error("invalid speculative verification state transition")]
InvalidSpeculativeTransition,
#[error("speculative verification round is not resolved")]
IncompleteSpeculativeRound,
#[error("optimistic proposal prefix diverged from the canonical committed prefix")]
OptimisticPrefixDiverged,
#[error("optimistic proposal branch is empty")]
EmptyOptimisticBranch,
#[error("speculative verification already retains an optimistic branch")]
OptimisticBranchAlreadyPresent,
#[error("unknown speculative request id {index}")]
UnknownSpeculativeRequest {
index: usize,
},
#[error(
"promoted speculative block has {proposed} proposals but canonical capacity is {capacity}"
)]
ProposalCapacityExceeded {
proposed: usize,
capacity: usize,
},
#[error("cannot finish a speculative scheduler with active requests")]
ActiveSpeculativeRequests,
#[error("completed speculative request {index} has no finish reason")]
MissingSpeculativeFinishReason {
index: usize,
},
#[error("speculative max_draft_tokens must be positive")]
ZeroDraftTokens,
#[error("speculative backend does not permit any draft tokens")]
NoBackendDraftCapacity,
#[error("temperature must be finite and non-negative, got {0}")]
InvalidTemperature(f32),
#[error("do_sample=true requires a temperature greater than zero")]
StochasticZeroTemperature,
#[error("Mirostat V2 tau must be finite and positive, got {0}")]
InvalidMirostatTau(f32),
#[error("Mirostat V2 eta must be finite and positive, got {0}")]
InvalidMirostatEta(f32),
#[error("top_k must be non-negative, got {0}")]
InvalidTopK(i32),
#[error("top_p must be between zero and one, got {0}")]
InvalidTopP(f32),
#[error("min_p must be between zero and one, got {0}")]
InvalidMinP(f32),
#[error("repetition_penalty must be finite and positive, got {0}")]
InvalidRepetitionPenalty(f32),
#[error("frequency_penalty must be finite, got {0}")]
InvalidFrequencyPenalty(f32),
#[error("presence_penalty must be finite, got {0}")]
InvalidPresencePenalty(f32),
#[error("max_new_tokens must be positive when supplied")]
ZeroTokenBudget,
#[error("speculative max_in_flight_verifications must be positive")]
ZeroInFlightVerifications,
#[error("speculative scheduler currently supports at most one lookahead block")]
TooManyLookaheadBlocks,
#[error("speculative lookahead requires at least one optimistic branch slot")]
LookaheadWithoutBranchCapacity,
#[error("speculative adaptive_lookahead_min_blocks must be positive")]
ZeroAdaptiveLookaheadWindow,
#[error("speculative scheduler reached a non-terminal state with no eligible operation")]
StalledSpeculativeSchedule,
#[error("invalid speculative request status transition from {from:?} to {to:?}")]
InvalidSpeculativeStatusTransition {
from: SpeculativeRequestStatus,
to: SpeculativeRequestStatus,
},
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn terminal_precedence_and_cancellation_are_canonical() {
let mut sequence = GenerationSequence::new(2, [7]);
let first = sequence.commit(1, TokenTerminalSignals::default()).unwrap();
assert_eq!(first.finish_reason, None);
let second = sequence
.commit(
7,
TokenTerminalSignals {
stop_sequence: true,
grammar_complete: true,
},
)
.unwrap();
assert_eq!(second.finish_reason, Some(FinishReason::StopSequence));
assert!(!sequence.cancel());
let token = GenerationCancellationToken::new();
let mut active = GenerationSequence::new(3, []);
token.cancel();
assert!(active.observe_cancellation(&token));
assert_eq!(active.finish_reason(), Some(FinishReason::Cancelled));
}
#[test]
fn speculative_commit_plan_preserves_trailing_token_cache_semantics() {
let mut rejected = SpeculativeRound::new(3).unwrap();
rejected.accept(10, false).unwrap();
rejected.reject_with(99, false).unwrap();
let plan = rejected.commit_plan().unwrap();
assert_eq!(plan.accepted_proposals, 1);
assert_eq!(plan.committed_tokens, &[10, 99]);
assert_eq!(plan.verified_inputs, 2);
assert_eq!(plan.tail, Some(SpeculativeTail::Replacement));
let mut terminal_accept = SpeculativeRound::new(2).unwrap();
terminal_accept.accept(10, false).unwrap();
terminal_accept.accept(11, true).unwrap();
let plan = terminal_accept.commit_plan().unwrap();
assert!(plan.full_acceptance);
assert_eq!(plan.verified_inputs, 2);
assert_eq!(plan.tail, None);
let mut bonus = SpeculativeRound::new(2).unwrap();
bonus.accept(10, false).unwrap();
bonus.accept(11, false).unwrap();
bonus.bonus(12, false).unwrap();
assert_eq!(bonus.commit_plan().unwrap().verified_inputs, 3);
let mut incomplete = SpeculativeRound::new(2).unwrap();
incomplete.accept(10, false).unwrap();
assert!(matches!(
incomplete.commit_plan(),
Err(GenerationError::IncompleteSpeculativeRound)
));
}
#[test]
fn optimistic_reuse_is_pure_and_fail_closed() {
assert_eq!(
resolve_optimistic_reuse(&[1], &[1], &[2, 3], 2, false).unwrap(),
OptimisticReuseDecision::MatchedRetained
);
assert_eq!(
resolve_optimistic_reuse(&[1], &[1], &[2], 2, false).unwrap(),
OptimisticReuseDecision::MatchedConsumed
);
assert_eq!(
resolve_optimistic_reuse(&[1], &[1], &[2], 9, false).unwrap(),
OptimisticReuseDecision::DiscardMismatch
);
assert!(matches!(
resolve_optimistic_reuse(&[1], &[9], &[2], 2, false),
Err(GenerationError::OptimisticPrefixDiverged)
));
}
#[test]
fn sampler_and_scheduler_configuration_validate_without_a_backend() {
let checkpoint = CheckpointGenerationConfig {
do_sample: Some(true),
temperature: Some(0.8),
top_k: Some(64),
repetition_penalty: Some(1.1),
..CheckpointGenerationConfig::default()
};
let resolved =
resolve_generation_config(Some(&checkpoint), GenerationConfigOverrides::default())
.unwrap();
assert!(resolved.do_sample);
assert_eq!(resolved.top_k, 64);
assert_eq!(resolved.repetition_penalty, 1.1);
assert_eq!(resolved.repeat_last_n, 64);
assert!(matches!(
resolve_generation_config(
None,
GenerationConfigOverrides {
frequency_penalty: Some(f32::NAN),
..GenerationConfigOverrides::default()
}
),
Err(GenerationError::InvalidFrequencyPenalty(value)) if value.is_nan()
));
assert!(SpeculativeConfig::default().validate().is_ok());
assert!(SpeculativeSchedulerOptions::default().validate().is_ok());
assert!(matches!(
SpeculativeSchedulerOptions {
max_in_flight_verifications: 0,
..SpeculativeSchedulerOptions::default()
}
.validate(),
Err(GenerationError::ZeroInFlightVerifications)
));
let config_json = serde_json::to_string(&resolved).unwrap();
assert_eq!(
serde_json::from_str::<ResolvedGenerationConfig>(&config_json).unwrap(),
resolved
);
let options = SpeculativeSchedulerOptions::default();
let options_json = serde_json::to_string(&options).unwrap();
assert_eq!(
serde_json::from_str::<SpeculativeSchedulerOptions>(&options_json).unwrap(),
options
);
}
#[test]
fn semantic_events_round_trip_without_a_backend() {
let event = SemanticEvent::ToolCallStart {
index: 2,
id: "call_2".into(),
name: "lookup".into(),
};
let json = serde_json::to_string(&event).unwrap();
assert_eq!(serde_json::from_str::<SemanticEvent>(&json).unwrap(), event);
let mut zero_budget = GenerationSequence::new(0, []);
assert_eq!(zero_budget.finish_reason(), Some(FinishReason::MaxTokens));
assert!(zero_budget.cancel());
assert_eq!(zero_budget.finish_reason(), Some(FinishReason::Cancelled));
}
#[test]
fn speculative_request_lifecycle_defers_cancellation_exactly() {
let mut lifecycle = SpeculativeRequestLifecycle::new();
lifecycle
.transition(SpeculativeRequestStatus::ReadyToDraft)
.unwrap();
lifecycle
.transition(SpeculativeRequestStatus::ReadyToSubmitVerification)
.unwrap();
lifecycle
.transition(SpeculativeRequestStatus::TargetVerificationInFlight)
.unwrap();
assert_eq!(
lifecycle.request_cancellation(true).unwrap(),
SpeculativeCancellationDisposition::Deferred
);
assert!(lifecycle.cancellation_pending());
lifecycle
.transition(SpeculativeRequestStatus::VerificationResolution)
.unwrap();
lifecycle
.transition(SpeculativeRequestStatus::Cancelled)
.unwrap();
assert!(lifecycle.is_terminal());
assert!(!lifecycle.cancellation_pending());
assert!(matches!(
lifecycle.transition(SpeculativeRequestStatus::ReadyToDraft),
Err(GenerationError::InvalidSpeculativeStatusTransition { .. })
));
}
}