#![cfg(all(feature = "remote", feature = "mmap"))]
mod common;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use cera::bundle::BundleRepo;
use cera::engine::{BackendPreference, CeraEngine, EngineConfig};
use cera::session::SessionConfig;
const BUNDLE_ID: &str = "LFM2-350M-Extract-GGUF";
const QUANT: &str = "Q4_0";
fn run_with_shift_counter<R>(f: impl FnOnce(Arc<AtomicUsize>) -> R) -> R {
use tracing::subscriber::DefaultGuard;
use tracing_subscriber::layer::SubscriberExt;
let counter = Arc::new(AtomicUsize::new(0));
struct ShiftCounter(Arc<AtomicUsize>);
impl<S> tracing_subscriber::Layer<S> for ShiftCounter
where
S: tracing::Subscriber,
{
fn on_event(
&self,
event: &tracing::Event<'_>,
_ctx: tracing_subscriber::layer::Context<'_, S>,
) {
if event.metadata().target() == "cera::kv_shift" {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
}
let subscriber = tracing_subscriber::registry().with(ShiftCounter(counter.clone()));
let _guard: DefaultGuard = tracing::subscriber::set_default(subscriber);
f(counter)
}
fn run_shift_through_real_model(backend: BackendPreference) {
if std::env::var("CERA_TEST_DOWNLOAD").is_err() {
eprintln!("skipping: CERA_TEST_DOWNLOAD not set");
return;
}
const CTX: usize = 256;
let repo = BundleRepo::new(common::download::cache_dir());
let engine = CeraEngine::from_bundle_id(
BUNDLE_ID,
QUANT,
EngineConfig {
context_size: CTX,
backend,
bundle_repo: Some(repo),
..Default::default()
},
)
.expect("load engine");
let cfg = SessionConfig {
max_seq_len: Some(CTX as u32),
n_keep: 32,
seed: Some(0),
ubatch_size: 64,
..Default::default()
};
let shift_count = run_with_shift_counter(|counter| {
let mut session = engine.new_session(cfg);
let base = "The quick brown fox jumps over the lazy dog, and then walks back to inspect the fence. ";
let long_prompt = base.repeat(9);
session
.append_text(&long_prompt)
.expect("append long prompt");
let pos_after_prompt = session.position() as usize;
println!("position after prompt: {pos_after_prompt} / {CTX}");
assert!(
(CTX / 2..CTX).contains(&pos_after_prompt),
"prompt should fill most of window but stay under cap \
(pos={pos_after_prompt}); re-tune reps if vocab changed"
);
let before = counter.load(Ordering::Relaxed);
let follow = "Additionally, the narrator walks slowly along the river path. ".repeat(15);
session.append_text(&follow).expect("append forcing shift");
let pos_after_follow = session.position() as usize;
println!("position after follow: {pos_after_follow} / {CTX}");
let after = counter.load(Ordering::Relaxed);
assert!(
after > before,
"shift should have fired during follow-up append \
(before={before} after={after}, \
pos_before={pos_after_prompt}, pos_after={pos_after_follow})"
);
assert_eq!(
pos_after_follow, CTX,
"post-shift position must land at exactly max_seq_len"
);
after
});
assert!(
shift_count >= 1,
"expected ≥1 shift event, got {shift_count}"
);
}
#[test]
#[ignore = "downloads ~210 MB; set CERA_TEST_DOWNLOAD=1 and pass --ignored"]
fn shift_runs_through_real_model() {
run_shift_through_real_model(BackendPreference::Cpu);
}
#[cfg(all(feature = "metal", target_os = "macos"))]
#[test]
#[ignore = "downloads ~210 MB + needs Metal; set CERA_TEST_DOWNLOAD=1 and pass --ignored"]
fn shift_runs_through_real_model_metal() {
run_shift_through_real_model(BackendPreference::Metal);
}