Skip to main content

el_safety/
lib.rs

1//! `el-safety` — on-device, tiered, decoder-time safety (ADR-005).
2//!
3//! The [`SafetyMode`] tier is budget-gated by device profile via
4//! [`SafetyModeSelector`]. The `Lightweight` anchor/blacklist filter is fully
5//! implemented here. `SecDecoding` (two ~1B models) and `Csd` (claim
6//! backtracking) require model assets and are scaffolded as follow-ups
7//! ([`SecDecodingSteerer`]). **No safety path touches the network.**
8
9#![forbid(unsafe_code)]
10
11use el_core::{DeviceTarget, SafetyMode, Token};
12
13/// A vector subtracted from target logits to steer away from unsafe output.
14/// Sparse and integer (milli-logits) for deterministic, allocation-light steps.
15#[derive(Debug, Clone, Default, PartialEq, Eq)]
16pub struct LogitAdjustment {
17    penalties: Vec<(Token, i32)>,
18}
19
20impl LogitAdjustment {
21    pub fn none() -> Self {
22        Self::default()
23    }
24
25    pub fn with_penalties(penalties: Vec<(Token, i32)>) -> Self {
26        Self { penalties }
27    }
28
29    pub fn is_empty(&self) -> bool {
30        self.penalties.is_empty()
31    }
32
33    /// The milli-logit delta to add for `token` (0 if unaffected).
34    pub fn delta_for(&self, token: Token) -> i32 {
35        self.penalties
36            .iter()
37            .find(|(t, _)| *t == token)
38            .map(|(_, d)| *d)
39            .unwrap_or(0)
40    }
41
42    /// L1 norm in milli-units — what `LogitsSteered.adjustment_norm_milli`
43    /// reports to telemetry.
44    pub fn l1_norm_milli(&self) -> u32 {
45        self.penalties.iter().map(|(_, d)| d.unsigned_abs()).sum()
46    }
47}
48
49/// Per-step safety intervention. The runtime applies this **after** the grammar
50/// mask and **before** sampling.
51pub trait SafetySteerer {
52    fn adjust(&self, recent_tokens: &[Token]) -> LogitAdjustment;
53    fn mode(&self) -> SafetyMode;
54}
55
56/// Chooses the affordable mode for the device (ADR-005).
57pub struct SafetyModeSelector;
58
59impl SafetyModeSelector {
60    /// `SecDecoding` (two ~1B models) is rejected on `MidRange` and downgraded
61    /// to `Lightweight`; everything else passes through.
62    pub fn resolve(requested: SafetyMode, device: DeviceTarget) -> SafetyMode {
63        match (requested, device) {
64            (SafetyMode::SecDecoding, DeviceTarget::MidRange) => SafetyMode::Lightweight,
65            (m, _) => m,
66        }
67    }
68}
69
70/// `SafetyMode::Off` — a no-op steerer.
71pub struct NoSafety;
72
73impl SafetySteerer for NoSafety {
74    fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
75        LogitAdjustment::none()
76    }
77    fn mode(&self) -> SafetyMode {
78        SafetyMode::Off
79    }
80}
81
82/// `SafetyMode::Lightweight` — a training-free blacklist filter (real). Banned
83/// tokens receive a very large negative logit so they cannot be sampled.
84pub struct LightweightFilter {
85    banned: Vec<Token>,
86}
87
88impl LightweightFilter {
89    pub const HARD_BAN: i32 = -1_000_000;
90
91    pub fn new(banned: Vec<Token>) -> Self {
92        Self { banned }
93    }
94}
95
96impl SafetySteerer for LightweightFilter {
97    fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
98        LogitAdjustment::with_penalties(self.banned.iter().map(|&t| (t, Self::HARD_BAN)).collect())
99    }
100    fn mode(&self) -> SafetyMode {
101        SafetyMode::Lightweight
102    }
103}
104
105/// `SafetyMode::SecDecoding` — base-vs-safety-model logit steering.
106///
107/// FOLLOW-UP (ADR-005): requires two ~1B models run on Candle. Until model
108/// assets are wired, this returns no adjustment and reports its intended mode,
109/// so callers can select it without it silently mis-steering.
110pub struct SecDecodingSteerer {
111    _private: (),
112}
113
114impl SecDecodingSteerer {
115    pub fn placeholder() -> Self {
116        Self { _private: () }
117    }
118}
119
120impl SafetySteerer for SecDecodingSteerer {
121    fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
122        // TODO(adr-005): run base + safety models on Candle, derive adjustment
123        // from their divergence.
124        LogitAdjustment::none()
125    }
126    fn mode(&self) -> SafetyMode {
127        SafetyMode::SecDecoding
128    }
129}
130
131#[cfg(test)]
132mod tests {
133    use super::*;
134
135    #[test]
136    fn secdecoding_downgrades_on_midrange() {
137        assert_eq!(
138            SafetyModeSelector::resolve(SafetyMode::SecDecoding, DeviceTarget::MidRange),
139            SafetyMode::Lightweight
140        );
141        assert_eq!(
142            SafetyModeSelector::resolve(SafetyMode::SecDecoding, DeviceTarget::HighEnd),
143            SafetyMode::SecDecoding
144        );
145    }
146
147    #[test]
148    fn lightweight_bans_tokens() {
149        let f = LightweightFilter::new(vec![42, 99]);
150        let adj = f.adjust(&[]);
151        assert_eq!(adj.delta_for(42), LightweightFilter::HARD_BAN);
152        assert_eq!(adj.delta_for(7), 0);
153        assert!(adj.l1_norm_milli() > 0);
154    }
155}