#![forbid(unsafe_code)]
use el_core::{DeviceTarget, SafetyMode, Token};
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct LogitAdjustment {
penalties: Vec<(Token, i32)>,
}
impl LogitAdjustment {
pub fn none() -> Self {
Self::default()
}
pub fn with_penalties(penalties: Vec<(Token, i32)>) -> Self {
Self { penalties }
}
pub fn is_empty(&self) -> bool {
self.penalties.is_empty()
}
pub fn delta_for(&self, token: Token) -> i32 {
self.penalties
.iter()
.find(|(t, _)| *t == token)
.map(|(_, d)| *d)
.unwrap_or(0)
}
pub fn l1_norm_milli(&self) -> u32 {
self.penalties
.iter()
.fold(0u32, |acc, (_, d)| acc.saturating_add(d.unsigned_abs()))
}
pub fn penalties(&self) -> &[(Token, i32)] {
&self.penalties
}
}
pub trait SafetySteerer {
fn adjust(&self, recent_tokens: &[Token]) -> LogitAdjustment;
fn mode(&self) -> SafetyMode;
fn adjust_with_logits(&self, recent_tokens: &[Token], _base_logits: &[i32]) -> LogitAdjustment {
self.adjust(recent_tokens)
}
}
pub struct SafetyModeSelector;
impl SafetyModeSelector {
pub fn resolve(requested: SafetyMode, device: DeviceTarget) -> SafetyMode {
match (requested, device) {
(SafetyMode::SecDecoding, DeviceTarget::MidRange) => SafetyMode::Lightweight,
(m, _) => m,
}
}
}
pub struct NoSafety;
impl SafetySteerer for NoSafety {
fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
LogitAdjustment::none()
}
fn mode(&self) -> SafetyMode {
SafetyMode::Off
}
}
pub struct LightweightFilter {
banned: Vec<Token>,
}
impl LightweightFilter {
pub const HARD_BAN: i32 = -1_000_000;
pub fn new(banned: Vec<Token>) -> Self {
Self { banned }
}
}
impl SafetySteerer for LightweightFilter {
fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
LogitAdjustment::with_penalties(self.banned.iter().map(|&t| (t, Self::HARD_BAN)).collect())
}
fn mode(&self) -> SafetyMode {
SafetyMode::Lightweight
}
}
pub struct SecDecodingSteerer {
_private: (),
}
impl SecDecodingSteerer {
pub fn placeholder() -> Self {
Self { _private: () }
}
}
impl SafetySteerer for SecDecodingSteerer {
fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
LogitAdjustment::none()
}
fn mode(&self) -> SafetyMode {
SafetyMode::SecDecoding
}
}
pub trait ExpertLogits {
fn logits(&self, committed: &[Token]) -> Vec<i32>;
}
pub fn contrastive_adjustment(
base: &[i32],
expert: &[i32],
alpha_milli: i32,
top_k: usize,
) -> LogitAdjustment {
if base.is_empty() || base.len() != expert.len() {
return LogitAdjustment::none();
}
let alpha = i64::from(alpha_milli.max(0));
let mut idx: Vec<usize> = (0..base.len()).collect();
if top_k > 0 && top_k < base.len() {
idx.sort_unstable_by(|&a, &b| base[b].cmp(&base[a]).then(a.cmp(&b)));
idx.truncate(top_k);
}
let mut penalties: Vec<(Token, i32)> = Vec::new();
for i in idx {
let diff = i64::from(expert[i]) - i64::from(base[i]);
let delta = (diff.saturating_mul(alpha) / 1000)
.clamp(i64::from(i32::MIN), i64::from(i32::MAX)) as i32;
if delta != 0 {
penalties.push((i as Token, delta));
}
}
LogitAdjustment::with_penalties(penalties)
}
pub struct ContrastiveSteerer<E: ExpertLogits> {
expert: E,
banned: Vec<Token>,
alpha_milli: i32,
top_k: usize,
mode: SafetyMode,
}
impl<E: ExpertLogits> ContrastiveSteerer<E> {
pub fn new(
expert: E,
banned: Vec<Token>,
alpha_milli: i32,
top_k: usize,
mode: SafetyMode,
) -> Self {
Self {
expert,
banned,
alpha_milli,
top_k,
mode,
}
}
fn bans(&self) -> Vec<(Token, i32)> {
self.banned
.iter()
.map(|&t| (t, LightweightFilter::HARD_BAN))
.collect()
}
}
impl<E: ExpertLogits> SafetySteerer for ContrastiveSteerer<E> {
fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
LogitAdjustment::with_penalties(self.bans())
}
fn adjust_with_logits(&self, recent: &[Token], base_logits: &[i32]) -> LogitAdjustment {
let mut penalties = self.bans();
let expert = self.expert.logits(recent);
let contrast = contrastive_adjustment(base_logits, &expert, self.alpha_milli, self.top_k);
penalties.extend_from_slice(contrast.penalties());
LogitAdjustment::with_penalties(penalties)
}
fn mode(&self) -> SafetyMode {
self.mode
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
pub struct SafetyScore(u16);
impl SafetyScore {
pub const SAFE: SafetyScore = SafetyScore(0);
pub const MAX: SafetyScore = SafetyScore(1000);
pub fn from_milli(milli: u16) -> Self {
Self(milli.min(1000))
}
pub fn milli(self) -> u16 {
self.0
}
}
pub trait ChunkGuard {
fn score(&self, recent: &[Token]) -> SafetyScore;
}
#[derive(Debug, Clone, Default)]
pub struct AnchorGuard {
patterns: Vec<Vec<Token>>,
per_hit: u16,
}
impl AnchorGuard {
pub fn new(patterns: Vec<Vec<Token>>, per_hit_milli: u16) -> Self {
Self {
patterns,
per_hit: per_hit_milli,
}
}
pub fn hard(patterns: Vec<Vec<Token>>) -> Self {
Self::new(patterns, SafetyScore::MAX.milli())
}
pub fn is_empty(&self) -> bool {
self.patterns.iter().all(Vec::is_empty)
}
}
impl ChunkGuard for AnchorGuard {
fn score(&self, recent: &[Token]) -> SafetyScore {
let mut milli: u32 = 0;
for pat in &self.patterns {
if pat.is_empty() || pat.len() > recent.len() {
continue;
}
if recent.windows(pat.len()).any(|w| w == pat.as_slice()) {
milli = milli.saturating_add(u32::from(self.per_hit));
}
}
SafetyScore::from_milli(milli.min(u32::from(u16::MAX)) as u16)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RollbackPolicy {
pub guard_every: u32,
pub steer_window: u32,
pub soft_threshold: SafetyScore,
pub hard_threshold: SafetyScore,
pub max_rollbacks: u8,
pub max_checkpoints: u8,
}
impl RollbackPolicy {
pub fn for_device(device: DeviceTarget, mode: SafetyMode) -> Self {
if matches!(mode, SafetyMode::Off) {
return Self::disabled();
}
match device {
DeviceTarget::MidRange => Self {
guard_every: 16,
steer_window: 8,
soft_threshold: SafetyScore(600),
hard_threshold: SafetyScore(800),
max_rollbacks: 2,
max_checkpoints: 4,
},
DeviceTarget::HighEnd | DeviceTarget::Auto => Self {
guard_every: 4,
steer_window: 16,
soft_threshold: SafetyScore(500),
hard_threshold: SafetyScore(750),
max_rollbacks: 4,
max_checkpoints: 8,
},
}
}
pub fn disabled() -> Self {
Self {
guard_every: 0,
steer_window: 0,
soft_threshold: SafetyScore::MAX,
hard_threshold: SafetyScore::MAX,
max_rollbacks: 0,
max_checkpoints: 0,
}
}
pub fn guards(&self) -> bool {
self.guard_every > 0
}
pub fn active(&self) -> bool {
self.guard_every > 0 || self.steer_window > 0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Checkpoint {
pub output_len: u32,
pub kv_len: u32,
}
#[derive(Debug, Default)]
pub struct CheckpointManager {
ring: Vec<Checkpoint>,
cap: usize,
enabled: bool,
}
impl CheckpointManager {
pub fn new(cap: u8) -> Self {
Self {
ring: Vec::new(),
cap: cap as usize,
enabled: cap > 0,
}
}
pub fn enabled(&self) -> bool {
self.enabled
}
pub fn disable(&mut self) {
self.enabled = false;
self.ring.clear();
}
pub fn push(&mut self, checkpoint: Checkpoint) {
if !self.enabled {
return;
}
if self.ring.len() == self.cap {
self.ring.remove(0);
}
self.ring.push(checkpoint);
}
pub fn last(&self) -> Option<Checkpoint> {
self.ring.last().copied()
}
pub fn len(&self) -> usize {
self.ring.len()
}
pub fn is_empty(&self) -> bool {
self.ring.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn secdecoding_downgrades_on_midrange() {
assert_eq!(
SafetyModeSelector::resolve(SafetyMode::SecDecoding, DeviceTarget::MidRange),
SafetyMode::Lightweight
);
assert_eq!(
SafetyModeSelector::resolve(SafetyMode::SecDecoding, DeviceTarget::HighEnd),
SafetyMode::SecDecoding
);
}
#[test]
fn lightweight_bans_tokens() {
let f = LightweightFilter::new(vec![42, 99]);
let adj = f.adjust(&[]);
assert_eq!(adj.delta_for(42), LightweightFilter::HARD_BAN);
assert_eq!(adj.delta_for(7), 0);
assert!(adj.l1_norm_milli() > 0);
}
struct FixedExpert(Vec<i32>);
impl ExpertLogits for FixedExpert {
fn logits(&self, _committed: &[Token]) -> Vec<i32> {
self.0.clone()
}
}
#[test]
fn contrastive_pushes_toward_expert_and_is_noop_when_equal() {
let base = [100, 200, 300];
let expert = [600, 200, 100]; let adj = contrastive_adjustment(&base, &expert, 1000, 0); assert_eq!(adj.delta_for(0), 500); assert_eq!(adj.delta_for(1), 0); assert_eq!(adj.delta_for(2), -200);
assert_eq!(
contrastive_adjustment(&base, &expert, 500, 0).delta_for(0),
250
);
assert!(contrastive_adjustment(&base, &base, 1000, 0).is_empty());
}
#[test]
fn contrastive_top_k_restricts_to_distribution_head() {
let base = [10, 50, 40, 5]; let expert = [999, 999, 999, 999];
let adj = contrastive_adjustment(&base, &expert, 1000, 2);
assert_ne!(adj.delta_for(1), 0);
assert_ne!(adj.delta_for(2), 0);
assert_eq!(adj.delta_for(0), 0); assert_eq!(adj.delta_for(3), 0);
}
#[test]
fn contrastive_empty_on_length_mismatch_or_empty() {
assert!(contrastive_adjustment(&[1, 2, 3], &[1, 2], 1000, 0).is_empty());
assert!(contrastive_adjustment(&[], &[], 1000, 0).is_empty());
}
#[test]
fn contrastive_clamps_negative_and_saturates_extreme_alpha() {
let base = [100, 200];
let expert = [900, 0];
assert!(contrastive_adjustment(&base, &expert, -5000, 0).is_empty());
let adj = contrastive_adjustment(&base, &expert, i32::MAX, 0);
assert!(adj.delta_for(0) > 0);
assert!(adj.delta_for(1) < 0);
let _ = adj.l1_norm_milli(); }
#[test]
fn contrastive_steerer_layers_bans_and_contrast() {
let base = [100, 100, 100];
let expert = FixedExpert(vec![100, 400, 100]); let steerer = ContrastiveSteerer::new(expert, vec![0], 1000, 0, SafetyMode::SecDecoding);
let banned_only = steerer.adjust(&[]);
assert_eq!(banned_only.delta_for(0), LightweightFilter::HARD_BAN);
assert_eq!(banned_only.delta_for(1), 0);
let full = steerer.adjust_with_logits(&[], &base);
assert_eq!(full.delta_for(0), LightweightFilter::HARD_BAN); assert_eq!(full.delta_for(1), 300); assert_eq!(steerer.mode(), SafetyMode::SecDecoding);
}
#[test]
fn policy_has_tier_aware_steer_window() {
let mid = RollbackPolicy::for_device(DeviceTarget::MidRange, SafetyMode::Lightweight);
let high = RollbackPolicy::for_device(DeviceTarget::HighEnd, SafetyMode::SecDecoding);
assert!(mid.steer_window > 0 && high.steer_window > 0);
assert!(high.steer_window >= mid.steer_window);
let off = RollbackPolicy::for_device(DeviceTarget::Auto, SafetyMode::Off);
assert_eq!(off.steer_window, 0);
assert!(!off.active());
}
#[test]
fn anchor_guard_matches_token_ngrams_exactly() {
let guard = AnchorGuard::hard(vec![vec![9], vec![40, 41]]);
assert_eq!(guard.score(&[1, 2, 3]), SafetyScore::SAFE);
assert_eq!(guard.score(&[1, 9, 2]), SafetyScore::MAX);
assert_eq!(guard.score(&[7, 40, 41, 8]), SafetyScore::MAX);
assert_eq!(guard.score(&[40, 99, 41]), SafetyScore::SAFE);
}
#[test]
fn anchor_guard_accumulates_below_max_and_is_empty_safe() {
let guard = AnchorGuard::new(vec![vec![1], vec![2]], 600);
assert_eq!(guard.score(&[1, 5]).milli(), 600);
assert_eq!(guard.score(&[1, 2]), SafetyScore::MAX);
assert!(AnchorGuard::hard(vec![]).is_empty());
assert_eq!(
AnchorGuard::hard(vec![]).score(&[1, 2, 3]),
SafetyScore::SAFE
);
assert_eq!(
AnchorGuard::hard(vec![vec![]]).score(&[1]),
SafetyScore::SAFE
);
}
#[test]
fn safety_score_clamps_and_orders() {
assert_eq!(SafetyScore::from_milli(5000), SafetyScore::MAX);
assert!(SafetyScore::SAFE < SafetyScore::MAX);
assert_eq!(SafetyScore::from_milli(750).milli(), 750);
}
#[test]
fn policy_off_is_disabled() {
let p = RollbackPolicy::for_device(DeviceTarget::HighEnd, SafetyMode::Off);
assert!(!p.guards());
assert_eq!(p.max_rollbacks, 0);
}
#[test]
fn policy_is_tier_aware() {
let mid = RollbackPolicy::for_device(DeviceTarget::MidRange, SafetyMode::Lightweight);
let high = RollbackPolicy::for_device(DeviceTarget::HighEnd, SafetyMode::SecDecoding);
assert!(mid.guards() && high.guards());
assert!(high.guard_every < mid.guard_every);
assert!(high.max_rollbacks >= mid.max_rollbacks);
}
#[test]
fn checkpoint_ring_is_bounded_and_disablable() {
let mut m = CheckpointManager::new(2);
m.push(Checkpoint {
output_len: 1,
kv_len: 1,
});
m.push(Checkpoint {
output_len: 2,
kv_len: 2,
});
m.push(Checkpoint {
output_len: 3,
kv_len: 3,
}); assert_eq!(m.len(), 2);
assert_eq!(
m.last(),
Some(Checkpoint {
output_len: 3,
kv_len: 3,
})
);
m.disable();
assert!(!m.enabled() && m.is_empty());
m.push(Checkpoint {
output_len: 9,
kv_len: 9,
}); assert!(m.last().is_none());
}
}