use std::io;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use thiserror::Error;
use crate::time::Instant;
use crate::kv_cache::{InferenceState, KvCompression};
use crate::model::Model;
use crate::model::audio_encoder::AudioEncoderWeights;
use crate::sampler::{Sampler, SamplerConfig};
use crate::tokenizer::BpeTokenizer;
#[derive(Debug, Clone)]
pub struct SessionConfig {
pub max_seq_len: Option<u32>,
pub kv_compression: KvCompression,
pub n_keep: u32,
pub seed: Option<u64>,
pub ubatch_size: u32,
}
impl Default for SessionConfig {
fn default() -> Self {
Self {
max_seq_len: None,
kv_compression: KvCompression::None,
n_keep: 0,
seed: None,
ubatch_size: 512,
}
}
}
#[derive(Debug, Clone)]
pub struct GenerateOpts {
pub max_tokens: u32,
pub temperature: f32,
pub top_p: f32,
pub top_k: u32,
pub repetition_penalty: f32,
pub stop_tokens: Vec<u32>,
pub flush_every_tokens: u32,
pub flush_every_ms: u32,
}
impl Default for GenerateOpts {
fn default() -> Self {
Self {
max_tokens: 256,
temperature: 0.7,
top_p: 0.9,
top_k: 40,
repetition_penalty: 1.0,
stop_tokens: Vec::new(),
flush_every_tokens: 16,
flush_every_ms: 50,
}
}
}
#[derive(Debug, Clone)]
pub struct GenerateSummary {
pub tokens_generated: u32,
pub prompt_eval_tokens: u32,
pub prompt_eval_ms: u32,
pub decode_ms: u32,
pub finish_reason: FinishReason,
}
#[derive(Debug, Clone)]
pub enum FinishReason {
MaxTokens,
Stop,
Cancelled,
ContextFull,
Error(String),
}
pub trait ModalitySink {
fn on_text_tokens(&mut self, _tokens: &[u32]) {}
fn on_audio_frames(&mut self, _pcm: &[f32], _sample_rate: u32) {}
fn on_done(&mut self, reason: FinishReason);
}
#[derive(Debug, Clone, Copy, Default)]
pub struct ModalityCapabilities {
pub text_in: bool,
pub text_out: bool,
pub image_in: bool,
pub audio_in: bool,
pub audio_out: bool,
}
impl ModalityCapabilities {
pub fn text_only() -> Self {
Self {
text_in: true,
text_out: true,
..Default::default()
}
}
pub fn text_and_audio() -> Self {
Self {
text_in: true,
text_out: true,
audio_in: true,
audio_out: true,
..Default::default()
}
}
pub fn text_and_image_in() -> Self {
Self {
text_in: true,
text_out: true,
image_in: true,
..Default::default()
}
}
pub fn from_inference_type(it: &crate::manifest::InferenceType) -> Self {
use crate::manifest::InferenceType::*;
match it {
LlamaCppTextToText => Self::text_only(),
LlamaCppImageToText => Self::text_and_image_in(),
LlamaCppLfm2AudioV1 => Self::text_and_audio(),
Unknown(_) => Self::text_only(),
}
}
}
#[derive(Error, Debug)]
pub enum CeraError {
#[error("modality not supported by this model")]
UnsupportedModality,
#[error("inference_type `{0}` is not supported in this version of cera")]
UnsupportedInferenceType(String),
#[error("session is busy with another operation")]
Busy,
#[error("cancelled")]
Cancelled,
#[error("context window ({max_seq_len}) exceeded by {by} tokens")]
ContextOverflow { max_seq_len: u32, by: u32 },
#[error("empty input")]
EmptyInput,
#[error("backend: {0}")]
Backend(String),
#[error("io: {0}")]
Io(#[from] io::Error),
}
pub fn can_shift(
supports_kv_shift: bool,
n_keep: usize,
is_compressed: bool,
current_pos: usize,
shift_needed: usize,
) -> bool {
supports_kv_shift && n_keep > 0 && !is_compressed && current_pos >= n_keep + shift_needed
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ChatTemplateSegment {
Text { start: usize, end: usize },
Image,
}
pub(crate) fn splice_image_markers(
tokens: &[u32],
image_marker_id: u32,
) -> Vec<ChatTemplateSegment> {
let mut segments = Vec::new();
let mut start = 0;
for (i, &t) in tokens.iter().enumerate() {
if t == image_marker_id {
if i > start {
segments.push(ChatTemplateSegment::Text { start, end: i });
}
segments.push(ChatTemplateSegment::Image);
start = i + 1;
}
}
if start < tokens.len() {
segments.push(ChatTemplateSegment::Text {
start,
end: tokens.len(),
});
}
segments
}
pub struct Session {
model: Arc<dyn Model>,
tokenizer: Arc<BpeTokenizer>,
state: InferenceState,
sampler: Sampler,
current_pos: usize,
position_atomic: Arc<AtomicU32>,
cancel: Arc<AtomicBool>,
last_logits: Option<Vec<f32>>,
max_seq_len: usize,
capabilities: ModalityCapabilities,
config: SessionConfig,
audio_encoder: Option<Arc<AudioEncoderWeights>>,
vision_encoder: Option<Arc<crate::model::vision_encoder::VisionEncoderWeights>>,
}
impl Session {
pub fn new(
model: Arc<dyn Model>,
tokenizer: Arc<BpeTokenizer>,
capabilities: ModalityCapabilities,
config: SessionConfig,
) -> Self {
let model_cfg = model.config();
let max_seq_len = config
.max_seq_len
.map(|v| v as usize)
.unwrap_or(model_cfg.max_seq_len)
.min(model_cfg.max_seq_len);
if config.n_keep > 0 && (config.n_keep as usize) >= max_seq_len {
tracing::warn!(
target: "cera::session",
n_keep = config.n_keep,
max_seq_len,
"n_keep >= max_seq_len; context shift will never fire \
because there's no room left for shifted cells. Lower \
n_keep to enable shifting."
);
}
if config.n_keep > 0 && !matches!(config.kv_compression, KvCompression::None) {
tracing::warn!(
target: "cera::session",
n_keep = config.n_keep,
"n_keep configured alongside TurboQuant KV compression; \
shift not yet supported for compressed caches, so \
overflow will still return ContextOverflow. Disable \
compression to enable n_keep."
);
}
if config.n_keep > 0 && !model.supports_kv_shift() {
tracing::warn!(
target: "cera::session",
n_keep = config.n_keep,
architecture = model_cfg.architecture.as_str(),
"n_keep configured but this backend doesn't support KV shift; \
overflow will still return ContextOverflow. CPU backend \
(BackendPreference::Cpu) supports shift today; Metal / GPU \
paths land in a follow-up."
);
}
let state = InferenceState::from_config_with_compression(model_cfg, &config.kv_compression);
let sampler_cfg = SamplerConfig {
seed: config.seed,
..SamplerConfig::default()
};
let sampler = Sampler::new(sampler_cfg);
Self {
model,
tokenizer,
state,
sampler,
current_pos: 0,
position_atomic: Arc::new(AtomicU32::new(0)),
cancel: Arc::new(AtomicBool::new(false)),
last_logits: None,
max_seq_len,
capabilities,
config,
audio_encoder: None,
vision_encoder: None,
}
}
pub fn attach_audio_encoder(&mut self, encoder: Arc<AudioEncoderWeights>) {
self.audio_encoder = Some(encoder);
}
pub fn attach_vision_encoder(
&mut self,
encoder: Arc<crate::model::vision_encoder::VisionEncoderWeights>,
) {
self.vision_encoder = Some(encoder);
}
pub fn capabilities(&self) -> ModalityCapabilities {
self.capabilities
}
pub fn tokenizer(&self) -> &BpeTokenizer {
self.tokenizer.as_ref()
}
pub fn model(&self) -> &dyn Model {
self.model.as_ref()
}
pub fn position(&self) -> u32 {
self.position_atomic.load(Ordering::Relaxed)
}
pub fn position_handle(&self) -> Arc<AtomicU32> {
Arc::clone(&self.position_atomic)
}
pub fn cancel_handle(&self) -> Arc<AtomicBool> {
Arc::clone(&self.cancel)
}
pub fn cancel(&self) {
self.cancel.store(true, Ordering::Relaxed);
}
pub fn reset(&mut self) {
let model_cfg = self.model.config();
self.state =
InferenceState::from_config_with_compression(model_cfg, &self.config.kv_compression);
self.current_pos = 0;
self.position_atomic.store(0, Ordering::Relaxed);
self.last_logits = None;
self.cancel.store(false, Ordering::Relaxed);
let sampler_cfg = SamplerConfig {
seed: self.config.seed,
..SamplerConfig::default()
};
self.sampler = Sampler::new(sampler_cfg);
}
pub fn append_text(&mut self, text: &str) -> Result<(), CeraError> {
if text.is_empty() {
return Err(CeraError::EmptyInput);
}
let tokens = self.tokenizer.encode(text);
self.append_tokens(&tokens)
}
pub fn append_audio(&mut self, samples: &[f32], sample_rate: u32) -> Result<(), CeraError> {
if !self.capabilities.audio_in {
return Err(CeraError::UnsupportedModality);
}
let Some(encoder) = self.audio_encoder.as_ref() else {
return Err(CeraError::Backend(
"Session::append_audio: no audio encoder attached. Call \
attach_audio_encoder() with weights loaded via \
AudioEncoderWeights::from_gguf(...) on the bundle's \
multimodal_projector GGUF before calling append_audio."
.to_string(),
));
};
let llm_hidden = self.model.config().hidden_size;
let enc_hidden = encoder.config.llm_hidden_size;
if enc_hidden != llm_hidden {
return Err(CeraError::Backend(format!(
"Session::append_audio: attached encoder's llm_hidden_size ({enc_hidden}) \
does not match the LLM's hidden_size ({llm_hidden}). \
The encoder must be the multimodal_projector trained for the \
currently-loaded LLM."
)));
}
if sample_rate != crate::model::audio_encoder::SAMPLE_RATE {
return Err(CeraError::Backend(format!(
"Session::append_audio: sample_rate {} != {} required by \
encoder (resampling is out of scope; resample externally \
before passing samples in)",
sample_rate,
crate::model::audio_encoder::SAMPLE_RATE,
)));
}
if samples.is_empty() {
return Err(CeraError::EmptyInput);
}
let (embeddings, n_frames) =
crate::model::audio_encoder::encode_audio_pcm(samples, encoder.as_ref());
if n_frames == 0 {
return Err(CeraError::EmptyInput);
}
self.append_embeddings(&embeddings, n_frames)
}
pub fn append_tokens(&mut self, tokens: &[u32]) -> Result<(), CeraError> {
if tokens.is_empty() {
return Err(CeraError::EmptyInput);
}
let new_end = self
.current_pos
.checked_add(tokens.len())
.ok_or(CeraError::Backend("position overflow".into()))?;
if new_end > self.max_seq_len {
let n_keep = self.config.n_keep as usize;
let shift_needed = new_end - self.max_seq_len;
if !can_shift(
self.model.supports_kv_shift(),
n_keep,
self.state.is_compressed(),
self.current_pos,
shift_needed,
) {
return Err(CeraError::ContextOverflow {
max_seq_len: self.max_seq_len as u32,
by: (new_end - self.max_seq_len) as u32,
});
}
self.model.shift_kv(&mut self.state, n_keep, shift_needed);
let before = self.current_pos;
self.current_pos -= shift_needed;
self.position_atomic
.store(self.current_pos as u32, Ordering::Relaxed);
self.last_logits = None;
tracing::info!(
target: "cera::kv_shift",
n_keep = n_keep,
shift = shift_needed,
seq_len_before = before,
seq_len_after = self.current_pos,
"kv context shift"
);
}
let (consumed, logits) = self.model.forward_prefill_chunked(
tokens,
self.current_pos,
&mut self.state,
self.config.ubatch_size as usize,
&self.cancel,
);
self.current_pos += consumed;
self.position_atomic
.store(self.current_pos as u32, Ordering::Relaxed);
if consumed < tokens.len() {
self.last_logits = None;
Err(CeraError::Cancelled)
} else {
self.last_logits = logits;
Ok(())
}
}
pub fn append_embeddings(
&mut self,
embeddings: &[f32],
n_tokens: usize,
) -> Result<(), CeraError> {
if n_tokens == 0 {
return Err(CeraError::EmptyInput);
}
if !self.model.supports_embedding_input() {
return Err(CeraError::UnsupportedModality);
}
let hidden_size = self.model.config().hidden_size;
let expected_len = n_tokens.checked_mul(hidden_size).ok_or_else(|| {
CeraError::Backend("append_embeddings: n_tokens * hidden_size overflow".into())
})?;
if embeddings.len() != expected_len {
return Err(CeraError::Backend(format!(
"append_embeddings: embeddings.len() ({}) != n_tokens ({n_tokens}) * hidden_size ({hidden_size}) = {expected_len}",
embeddings.len()
)));
}
let new_end = self
.current_pos
.checked_add(n_tokens)
.ok_or(CeraError::Backend("position overflow".into()))?;
if new_end > self.max_seq_len {
let n_keep = self.config.n_keep as usize;
let shift_needed = new_end - self.max_seq_len;
if !can_shift(
self.model.supports_kv_shift(),
n_keep,
self.state.is_compressed(),
self.current_pos,
shift_needed,
) {
return Err(CeraError::ContextOverflow {
max_seq_len: self.max_seq_len as u32,
by: (new_end - self.max_seq_len) as u32,
});
}
self.model.shift_kv(&mut self.state, n_keep, shift_needed);
let before = self.current_pos;
self.current_pos -= shift_needed;
self.position_atomic
.store(self.current_pos as u32, Ordering::Relaxed);
self.last_logits = None;
tracing::info!(
target: "cera::kv_shift",
n_keep = n_keep,
shift = shift_needed,
seq_len_before = before,
seq_len_after = self.current_pos,
"kv context shift (append_embeddings)"
);
}
let ubatch = self.config.ubatch_size as usize;
let chunk_size = if ubatch == 0 { n_tokens } else { ubatch };
let mut last_logits: Option<Vec<f32>> = None;
let mut ti = 0usize;
while ti < n_tokens {
let end = (ti + chunk_size).min(n_tokens);
let chunk = &embeddings[ti * hidden_size..end * hidden_size];
let logits = self.model.forward_prefill_from_embeddings(
chunk,
end - ti,
self.current_pos,
&mut self.state,
);
self.current_pos += end - ti;
self.position_atomic
.store(self.current_pos as u32, Ordering::Relaxed);
last_logits = Some(logits);
ti = end;
if self.cancel.load(Ordering::Relaxed) && ti < n_tokens {
self.last_logits = None;
return Err(CeraError::Cancelled);
}
}
self.last_logits = last_logits;
Ok(())
}
pub fn clear_cancel(&self) {
self.cancel.store(false, Ordering::Relaxed);
}
#[cfg(feature = "vl-preprocess")]
pub fn append_image(&mut self, bytes: &[u8]) -> Result<(), CeraError> {
if bytes.is_empty() {
return Err(CeraError::EmptyInput);
}
if !self.capabilities.image_in {
return Err(CeraError::UnsupportedModality);
}
let Some(encoder) = self.vision_encoder.as_ref() else {
return Err(CeraError::Backend(
"Session::append_image: no vision encoder attached. \
Construct via CeraEngine on a VL bundle so the \
encoder is auto-attached, or call \
attach_vision_encoder(...) in test setup."
.into(),
));
};
let llm_hidden = self.model.config().hidden_size;
let proj_dim = encoder.config.projection_dim;
if proj_dim != llm_hidden {
return Err(CeraError::Backend(format!(
"Session::append_image: vision encoder's projection_dim \
({proj_dim}) does not match LLM hidden_size ({llm_hidden}). \
The mmproj must pair with the LLM it was trained against."
)));
}
let pre = crate::model::vision_preprocessor::preprocess_image(bytes, &encoder.config)?;
let img_tokens = encoder
.encode_image(&pre.pixels, pre.grid_w, pre.grid_h)
.map_err(|e| CeraError::Backend(format!("encode_image: {e:#}")))?;
if img_tokens.len() % proj_dim != 0 {
return Err(CeraError::Backend(format!(
"Session::append_image: encode_image returned {} f32s, \
not a multiple of projection_dim ({proj_dim}) — encoder \
produced a malformed image-token tensor",
img_tokens.len(),
)));
}
let n_tokens = img_tokens.len() / proj_dim;
if n_tokens == 0 {
return Err(CeraError::Backend(
"Session::append_image: encoder produced zero image tokens \
(preprocess + encode succeeded but yielded an empty tensor)"
.into(),
));
}
self.append_embeddings(&img_tokens, n_tokens)
}
#[cfg(not(feature = "vl-preprocess"))]
pub fn append_image(&mut self, _bytes: &[u8]) -> Result<(), CeraError> {
Err(CeraError::UnsupportedModality)
}
#[cfg(feature = "vl-preprocess")]
pub fn append_chat_with_images(
&mut self,
messages: &[crate::tokenizer::ChatMessageMultimodal],
images: &[&[u8]],
add_generation_prompt: bool,
) -> Result<(), CeraError> {
if !self.capabilities.image_in {
return Err(CeraError::UnsupportedModality);
}
let rendered =
crate::tokenizer::apply_chat_template(&self.tokenizer, messages, add_generation_prompt)
.map_err(|e| CeraError::Backend(format!("chat template render: {e:#}")))?;
let probed = self.tokenizer.encode("<image>");
if probed.len() != 1 {
return Err(CeraError::Backend(format!(
"tokenizer doesn't have `<image>` as a single token \
(got {} tokens) — model isn't VL-shaped",
probed.len(),
)));
}
let image_marker_id = probed[0];
let img_start = self
.tokenizer
.special_token_id("<|image_start|>")
.ok_or_else(|| {
CeraError::Backend(
"tokenizer missing `<|image_start|>` special token — \
model isn't VL-shaped"
.into(),
)
})?;
let img_end = self
.tokenizer
.special_token_id("<|image_end|>")
.ok_or_else(|| {
CeraError::Backend(
"tokenizer missing `<|image_end|>` special token — \
model isn't VL-shaped"
.into(),
)
})?;
let tokens = self.tokenizer.encode(&rendered);
let segments = splice_image_markers(&tokens, image_marker_id);
let marker_count = segments
.iter()
.filter(|s| matches!(s, ChatTemplateSegment::Image))
.count();
if marker_count != images.len() {
return Err(CeraError::Backend(format!(
"rendered chat template has {marker_count} `<image>` \
markers but caller supplied {} images",
images.len(),
)));
}
let mut img_idx = 0;
for seg in &segments {
match *seg {
ChatTemplateSegment::Text { start, end } => {
self.append_tokens(&tokens[start..end])?;
}
ChatTemplateSegment::Image => {
self.append_tokens(&[img_start])?;
self.append_image(images[img_idx])?;
self.append_tokens(&[img_end])?;
img_idx += 1;
}
}
}
Ok(())
}
#[cfg(not(feature = "vl-preprocess"))]
pub fn append_chat_with_images(
&mut self,
_messages: &[crate::tokenizer::ChatMessageMultimodal],
_images: &[&[u8]],
_add_generation_prompt: bool,
) -> Result<(), CeraError> {
Err(CeraError::UnsupportedModality)
}
pub fn generate<S: ModalitySink + ?Sized>(
&mut self,
opts: &GenerateOpts,
sink: &mut S,
) -> Result<GenerateSummary, CeraError> {
self.cancel.store(false, Ordering::Relaxed);
let prompt_eval_tokens = self.current_pos as u32;
let prompt_eval_ms: u32 = 0;
let decode_start = Instant::now();
let mut finish = FinishReason::MaxTokens;
let mut generated: u32 = 0;
let mut pos = self.current_pos;
let mut pending: Vec<u32> = Vec::with_capacity(opts.flush_every_tokens.max(1) as usize);
let mut last_flush = Instant::now();
let flush_n = opts.flush_every_tokens.max(1) as usize;
let flush_ms = opts.flush_every_ms;
if opts.max_tokens == 0 {
sink.on_done(FinishReason::MaxTokens);
let decode_ms = decode_start.elapsed().as_millis() as u32;
return Ok(GenerateSummary {
tokens_generated: 0,
prompt_eval_tokens,
prompt_eval_ms,
decode_ms,
finish_reason: FinishReason::MaxTokens,
});
}
if self.current_pos >= self.max_seq_len {
sink.on_done(FinishReason::ContextFull);
let decode_ms = decode_start.elapsed().as_millis() as u32;
return Ok(GenerateSummary {
tokens_generated: 0,
prompt_eval_tokens,
prompt_eval_ms,
decode_ms,
finish_reason: FinishReason::ContextFull,
});
}
let mut logits = self.last_logits.take().ok_or(CeraError::EmptyInput)?;
let greedy = opts.temperature <= 0.0 || opts.top_k == 1;
let mut sample_scratch: Vec<f32> = if greedy {
Vec::new()
} else {
self.sync_sampler_from_opts(opts);
Vec::with_capacity(logits.len())
};
let mut greedy_next: u32 = if greedy {
crate::sampler::cpu_argmax(&logits)
} else {
0
};
loop {
if self.cancel.load(Ordering::Relaxed) {
finish = FinishReason::Cancelled;
break;
}
if generated >= opts.max_tokens {
break;
}
if pos >= self.max_seq_len {
finish = FinishReason::ContextFull;
break;
}
let token = if greedy {
greedy_next
} else {
sample_scratch.clear();
sample_scratch.extend_from_slice(&logits);
self.sampler.sample(&mut sample_scratch)
};
if self.tokenizer.eos_token() == Some(token) || opts.stop_tokens.contains(&token) {
finish = FinishReason::Stop;
break;
}
pending.push(token);
generated += 1;
let should_flush_n = pending.len() >= flush_n;
let should_flush_t =
flush_ms > 0 && last_flush.elapsed().as_millis() >= flush_ms as u128;
if should_flush_n || should_flush_t {
sink.on_text_tokens(&pending);
pending.clear();
last_flush = Instant::now();
}
if greedy {
greedy_next = self.model.forward_greedy(&[token], pos, &mut self.state);
} else {
logits = self.model.forward(&[token], pos, &mut self.state);
}
pos += 1;
self.position_atomic.store(pos as u32, Ordering::Relaxed);
if self.cancel.load(Ordering::Relaxed) {
finish = FinishReason::Cancelled;
break;
}
}
if !pending.is_empty() {
sink.on_text_tokens(&pending);
}
self.current_pos = pos;
self.position_atomic
.store(self.current_pos as u32, Ordering::Relaxed);
if greedy && generated > 0 {
self.last_logits = None;
} else {
self.last_logits = Some(logits);
}
sink.on_done(finish.clone());
let decode_ms = decode_start.elapsed().as_millis() as u32;
Ok(GenerateSummary {
tokens_generated: generated,
prompt_eval_tokens,
prompt_eval_ms,
decode_ms,
finish_reason: finish,
})
}
fn sync_sampler_from_opts(&mut self, opts: &GenerateOpts) {
let cfg = SamplerConfig {
temperature: opts.temperature,
top_k: opts.top_k as usize,
top_p: opts.top_p,
seed: self.config.seed,
};
self.sampler.set_config(cfg);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::manifest::InferenceType;
#[test]
fn capabilities_from_inference_type_covers_every_variant() {
let text = ModalityCapabilities::from_inference_type(&InferenceType::LlamaCppTextToText);
assert!(text.text_in && text.text_out);
assert!(!text.audio_in && !text.audio_out && !text.image_in);
let audio = ModalityCapabilities::from_inference_type(&InferenceType::LlamaCppLfm2AudioV1);
assert!(audio.text_in && audio.text_out);
assert!(audio.audio_in && audio.audio_out);
assert!(!audio.image_in);
let vl = ModalityCapabilities::from_inference_type(&InferenceType::LlamaCppImageToText);
assert!(vl.text_in && vl.text_out && vl.image_in);
assert!(!vl.audio_in && !vl.audio_out);
let unknown = ModalityCapabilities::from_inference_type(&InferenceType::Unknown(
"llama.cpp/x".into(),
));
assert!(unknown.text_in && unknown.text_out);
assert!(!unknown.audio_in && !unknown.audio_out && !unknown.image_in);
}
#[derive(Default)]
struct RecordingSink {
tokens: Vec<u32>,
done: Option<FinishReason>,
flushes: u32,
}
impl ModalitySink for RecordingSink {
fn on_text_tokens(&mut self, tokens: &[u32]) {
self.tokens.extend_from_slice(tokens);
self.flushes += 1;
}
fn on_done(&mut self, reason: FinishReason) {
self.done = Some(reason);
}
}
#[test]
fn session_config_default_is_sane() {
let c = SessionConfig::default();
assert_eq!(c.n_keep, 0);
assert_eq!(c.ubatch_size, 512);
assert!(matches!(c.kv_compression, KvCompression::None));
}
#[test]
fn generate_opts_default_batching() {
let o = GenerateOpts::default();
assert_eq!(o.flush_every_tokens, 16);
assert_eq!(o.flush_every_ms, 50);
}
#[test]
fn capabilities_text_only_shape() {
let c = ModalityCapabilities::text_only();
assert!(c.text_in && c.text_out);
assert!(!c.image_in && !c.audio_in && !c.audio_out);
}
#[test]
fn recording_sink_collects_tokens() {
let mut s = RecordingSink::default();
s.on_text_tokens(&[1, 2, 3]);
s.on_text_tokens(&[4]);
s.on_done(FinishReason::MaxTokens);
assert_eq!(s.tokens, vec![1, 2, 3, 4]);
assert_eq!(s.flushes, 2);
assert!(matches!(s.done, Some(FinishReason::MaxTokens)));
}
#[test]
fn splice_image_markers_no_markers_one_text_run() {
let segs = splice_image_markers(&[1, 2, 3, 4], 99);
assert_eq!(segs, vec![ChatTemplateSegment::Text { start: 0, end: 4 }]);
}
#[test]
fn splice_image_markers_mid_stream() {
let segs = splice_image_markers(&[1, 2, 99, 3, 4], 99);
assert_eq!(
segs,
vec![
ChatTemplateSegment::Text { start: 0, end: 2 },
ChatTemplateSegment::Image,
ChatTemplateSegment::Text { start: 3, end: 5 },
]
);
}
#[test]
fn splice_image_markers_at_start() {
let segs = splice_image_markers(&[99, 1, 2], 99);
assert_eq!(
segs,
vec![
ChatTemplateSegment::Image,
ChatTemplateSegment::Text { start: 1, end: 3 },
]
);
}
#[test]
fn splice_image_markers_at_end() {
let segs = splice_image_markers(&[1, 2, 99], 99);
assert_eq!(
segs,
vec![
ChatTemplateSegment::Text { start: 0, end: 2 },
ChatTemplateSegment::Image,
]
);
}
#[test]
fn splice_image_markers_adjacent_markers() {
let segs = splice_image_markers(&[1, 99, 99, 2], 99);
assert_eq!(
segs,
vec![
ChatTemplateSegment::Text { start: 0, end: 1 },
ChatTemplateSegment::Image,
ChatTemplateSegment::Image,
ChatTemplateSegment::Text { start: 3, end: 4 },
]
);
}
#[test]
fn splice_image_markers_two_separated() {
let segs = splice_image_markers(&[1, 99, 2, 99, 3], 99);
let images = segs
.iter()
.filter(|s| matches!(s, ChatTemplateSegment::Image))
.count();
assert_eq!(images, 2);
assert_eq!(
segs,
vec![
ChatTemplateSegment::Text { start: 0, end: 1 },
ChatTemplateSegment::Image,
ChatTemplateSegment::Text { start: 2, end: 3 },
ChatTemplateSegment::Image,
ChatTemplateSegment::Text { start: 4, end: 5 },
]
);
}
#[test]
fn splice_image_markers_all_markers() {
let segs = splice_image_markers(&[99, 99, 99], 99);
assert_eq!(
segs,
vec![
ChatTemplateSegment::Image,
ChatTemplateSegment::Image,
ChatTemplateSegment::Image,
]
);
}
#[test]
fn splice_image_markers_empty() {
let segs = splice_image_markers(&[], 99);
assert!(segs.is_empty());
}
}