1use crate::ports::{InferenceEngine, Ports};
4use el_core::{
5 DegradeReason, DomainEvent, EdgeError, EventEnvelope, Phase, Result, SessionConfig, SessionId,
6 StopReason, Token,
7};
8use el_memory::KvRegion;
9use el_provenance::LoadPermit;
10use el_safety::{
11 Checkpoint, CheckpointManager, LogitAdjustment, RollbackPolicy, SafetyModeSelector,
12};
13
14const MIN_CHECKPOINT_BUDGET_BYTES: u64 = 64 * 1024 * 1024;
18
19pub struct InferenceSession<E: InferenceEngine> {
23 id: SessionId,
24 config: SessionConfig,
25 phase: Phase,
26 engine: E,
27 kv: KvRegion,
28 permit: LoadPermit,
29 prompt: Vec<Token>,
32 output: Vec<Token>,
33 step: u32,
34 events: Vec<EventEnvelope>,
35}
36
37enum GuardVerdict {
39 Pass,
41 RolledBack,
43 FailClosed,
45}
46
47struct GuardState {
50 checkpoints: CheckpointManager,
51 rollback_count: u8,
52 banned: Vec<Token>,
53 start_out: u32,
56 start_kv: u32,
57}
58
59impl<E: InferenceEngine> InferenceSession<E> {
60 pub fn new(id: SessionId, config: SessionConfig, engine: E, permit: LoadPermit) -> Self {
61 let mut s = Self {
62 id,
63 config,
64 phase: Phase::Initialized,
65 engine,
66 kv: KvRegion::new(),
67 permit,
68 prompt: Vec::new(),
69 output: Vec::new(),
70 step: 0,
71 events: Vec::new(),
72 };
73 s.emit(DomainEvent::SessionInitialized {
74 runtime: config.format.runtime(),
75 device: config.device,
76 safety: config.safety,
77 speculation: config.speculation,
78 });
79 s.emit(DomainEvent::ModelLoaded {
80 model: permit.model,
81 version: permit.version,
82 format: permit.format,
83 });
84 s
85 }
86
87 pub fn phase(&self) -> Phase {
88 self.phase
89 }
90 pub fn output(&self) -> &[Token] {
91 &self.output
92 }
93 pub fn kv_len(&self) -> u32 {
94 self.kv.len()
95 }
96 pub fn config(&self) -> &SessionConfig {
97 &self.config
98 }
99 pub fn permit(&self) -> LoadPermit {
102 self.permit
103 }
104 pub fn drain_events(&mut self) -> Vec<EventEnvelope> {
107 std::mem::take(&mut self.events)
108 }
109
110 fn emit(&mut self, event: DomainEvent) {
111 self.events
112 .push(EventEnvelope::new(self.id, self.step, event));
113 }
114
115 pub fn load_prompt(&mut self, ports: &Ports, prompt: &[Token]) -> Result<()> {
117 if self.phase != Phase::Initialized {
118 return Err(EdgeError::InvalidPhase {
119 expected: "Initialized",
120 found: self.phase.as_str(),
121 });
122 }
123
124 self.prompt = prompt.to_vec();
127
128 let compressed = if self.config.compress {
129 ports.compressor.compress(prompt)
130 } else {
131 prompt.to_vec()
132 };
133 if compressed.len() < prompt.len() {
134 let ratio_milli =
135 ((compressed.len() as u64 * 1000) / (prompt.len().max(1) as u64)) as u32;
136 self.emit(DomainEvent::PromptCompressed {
137 input_tokens: prompt.len() as u32,
138 output_tokens: compressed.len() as u32,
139 ratio_milli,
140 });
141 }
142
143 self.phase = Phase::Prefilling;
144 let kv_len = self.engine.prefill(&compressed)?;
145 for _ in 0..kv_len {
146 let off = self.kv.len() as u64;
147 self.kv.push(off);
148 }
149 self.emit(DomainEvent::PrefillCompleted {
150 prompt_tokens: compressed.len() as u32,
151 kv_len,
152 prefill_tps: 0,
153 });
154 self.phase = Phase::Decoding;
155 Ok(())
156 }
157
158 pub fn generate(&mut self, ports: &Ports, max_tokens: u32) -> Result<StopReason> {
161 let effective = SafetyModeSelector::resolve(self.config.safety, self.config.device);
166 self.emit(DomainEvent::SafetyModeSelected { mode: effective });
167 let policy = RollbackPolicy::for_device(self.config.device, effective);
168 self.generate_with_policy(ports, max_tokens, policy)
169 }
170
171 pub fn generate_with_policy(
191 &mut self,
192 ports: &Ports,
193 max_tokens: u32,
194 policy: RollbackPolicy,
195 ) -> Result<StopReason> {
196 if self.phase != Phase::Decoding {
197 return Err(EdgeError::InvalidPhase {
198 expected: "Decoding",
199 found: self.phase.as_str(),
200 });
201 }
202
203 if policy.active() {
209 if let Some(ingress) = ports.ingress.as_deref() {
210 let score = ingress.score(&self.prompt);
211 if score >= policy.hard_threshold {
212 self.emit(DomainEvent::SafetyViolationDetected {
213 score_milli: score.milli(),
214 threshold_milli: policy.hard_threshold.milli(),
215 });
216 self.phase = Phase::Completed;
217 self.emit(DomainEvent::GenerationCompleted {
218 total_tokens: self.output.len() as u32,
219 stop: StopReason::Stopped,
220 });
221 return Ok(StopReason::Stopped);
222 }
223 }
224 }
225
226 let eos = self.engine.eos_token();
227 let guarding = policy.guards() && ports.guard.is_some();
228
229 let checkpoints = if guarding {
232 if self.config.memory_budget_bytes < MIN_CHECKPOINT_BUDGET_BYTES {
233 self.emit(DomainEvent::SafetyDisabled {
234 reason: DegradeReason::MemoryPressure,
235 });
236 CheckpointManager::new(0)
237 } else {
238 CheckpointManager::new(policy.max_checkpoints)
239 }
240 } else {
241 CheckpointManager::new(0)
242 };
243 let mut state = GuardState {
244 checkpoints,
245 rollback_count: 0,
246 banned: Vec::new(),
247 start_out: self.output.len() as u32,
252 start_kv: self.kv.len(),
253 };
254 if state.checkpoints.enabled() {
256 state.checkpoints.push(Checkpoint {
257 output_len: state.start_out,
258 kv_len: state.start_kv,
259 });
260 }
261
262 let stop = loop {
263 let mut terminating: Option<StopReason> = None;
269
270 if self.output.len() as u32 >= max_tokens {
271 terminating = Some(StopReason::MaxTokens);
272 } else {
273 let logits = self.engine.next_logits(&self.output);
275 let vocab = logits.len();
276
277 let mut mask = ports.grammar.mask(&self.output, vocab);
280 for &t in &state.banned {
281 if let Some(slot) = mask.get_mut(t as usize) {
282 *slot = false;
283 }
284 }
285 let allowed = mask.iter().filter(|b| **b).count() as u32;
286 self.emit(DomainEvent::TokenMaskApplied { allowed });
287
288 let adj = if (self.output.len() as u32) < policy.steer_window {
295 if mask.iter().any(|&legal| !legal) {
300 let legal_logits: Vec<i32> = logits
301 .iter()
302 .zip(mask.iter())
303 .map(|(&l, &legal)| if legal { l } else { i32::MIN })
304 .collect();
305 ports.safety.adjust_with_logits(&self.output, &legal_logits)
306 } else {
307 ports.safety.adjust_with_logits(&self.output, &logits)
308 }
309 } else {
310 ports.safety.adjust(&self.output)
311 };
312 if !adj.is_empty() {
313 self.emit(DomainEvent::LogitsSteered {
314 adjustment_norm_milli: adj.l1_norm_milli(),
315 });
316 }
317
318 let token = match pick(&logits, &mask, &adj) {
322 Some(t) => t,
323 None => {
324 self.emit(DomainEvent::GrammarViolationBlocked);
325 break StopReason::Stopped;
326 }
327 };
328 self.emit(DomainEvent::TokenGenerated { sampled: false });
329
330 self.output.push(token);
332 self.kv.push(self.output.len() as u64);
333 self.step += 1;
334 self.emit(DomainEvent::TokenCommitted {
335 kv_len: self.kv.len(),
336 });
337
338 if token == eos {
339 terminating = Some(StopReason::Eos);
340 }
341 }
342
343 if guarding {
349 let guard = ports
350 .guard
351 .as_deref()
352 .expect("guarding implies a guard is wired");
353 let at_boundary = (self.output.len() as u32).is_multiple_of(policy.guard_every);
354 if terminating.is_some() || at_boundary {
355 match self.guard_chunk(guard, &policy, &mut state) {
356 GuardVerdict::FailClosed => break StopReason::Stopped,
358 GuardVerdict::RolledBack => continue,
361 GuardVerdict::Pass => {}
362 }
363 }
364 }
365
366 if let Some(reason) = terminating {
367 break reason;
368 }
369 };
370
371 self.phase = Phase::Completed;
372 self.emit(DomainEvent::GenerationCompleted {
373 total_tokens: self.output.len() as u32,
374 stop,
375 });
376 Ok(stop)
377 }
378
379 fn guard_chunk(
393 &mut self,
394 guard: &dyn el_safety::ChunkGuard,
395 policy: &RollbackPolicy,
396 state: &mut GuardState,
397 ) -> GuardVerdict {
398 let score = guard.score(&self.output);
399 if score >= policy.hard_threshold {
400 self.emit(DomainEvent::SafetyViolationDetected {
401 score_milli: score.milli(),
402 threshold_milli: policy.hard_threshold.milli(),
403 });
404 match state.checkpoints.last() {
405 Some(cp) if state.rollback_count < policy.max_rollbacks => {
406 if self.engine.rollback(cp.output_len).is_err() {
410 self.output.truncate(cp.output_len as usize);
411 self.kv.truncate(cp.kv_len);
412 return GuardVerdict::FailClosed;
413 }
414 if let Some(&bad) = self.output.get(cp.output_len as usize) {
416 state.banned.push(bad);
417 }
418 self.output.truncate(cp.output_len as usize);
419 self.kv.truncate(cp.kv_len);
420 state.rollback_count += 1;
421 self.emit(DomainEvent::ClaimBacktracked {
422 claim_index: cp.output_len,
423 });
424 GuardVerdict::RolledBack
425 }
426 _ => {
427 let (safe_out, safe_kv) = state
428 .checkpoints
429 .last()
430 .map_or((state.start_out, state.start_kv), |c| {
431 (c.output_len, c.kv_len)
432 });
433 self.output.truncate(safe_out as usize);
434 self.kv.truncate(safe_kv);
435 GuardVerdict::FailClosed
436 }
437 }
438 } else if score < policy.soft_threshold {
439 if state.checkpoints.enabled() {
441 state.checkpoints.push(Checkpoint {
442 output_len: self.output.len() as u32,
443 kv_len: self.kv.len(),
444 });
445 }
446 state.banned.clear();
447 GuardVerdict::Pass
448 } else {
449 GuardVerdict::Pass
451 }
452 }
453
454 pub fn reset(&mut self) {
456 self.kv = KvRegion::new();
457 self.prompt.clear();
458 self.output.clear();
459 self.step = 0;
460 self.phase = Phase::Initialized;
461 self.emit(DomainEvent::SessionReset);
462 }
463
464 pub fn consult_relay(&mut self, ports: &Ports, query: &[Token]) -> Result<Vec<Token>> {
467 if !self.config.hybrid_mode {
468 return Err(EdgeError::AirGapViolation);
469 }
470 match &ports.relay {
471 Some(relay) => {
472 let out = relay.consult(query);
473 self.emit(DomainEvent::HybridRelayConsulted);
474 Ok(out)
475 }
476 None => Err(EdgeError::AirGapViolation),
477 }
478 }
479}
480
481fn pick(logits: &[i32], mask: &[bool], adj: &LogitAdjustment) -> Option<Token> {
486 let mut best: Option<Token> = None;
487 let mut best_val = i32::MIN;
488 for (i, &l) in logits.iter().enumerate() {
489 if mask.get(i).copied() == Some(false) {
490 continue;
491 }
492 let v = l.saturating_add(adj.delta_for(i as Token));
493 if v > best_val {
494 best_val = v;
495 best = Some(i as Token);
496 }
497 }
498 best
499}
500
501#[cfg(test)]
502mod tests {
503 use super::*;
504 use crate::defaults::NullEngine;
505 use crate::ports::{GrammarMasker, Ports};
506 use el_core::{ModelFormat, ModelId, ModelVersion};
507 use el_provenance::{ModelArtifact, SignatureVerifier};
508 use el_safety::LightweightFilter;
509
510 struct OkVerifier;
511 impl SignatureVerifier for OkVerifier {
512 fn verify(&self, _b: &[u8], _s: &[u8], _k: u32) -> bool {
513 true
514 }
515 }
516
517 fn permit() -> LoadPermit {
518 let mut a = ModelArtifact::new(ModelId(1), ModelVersion::new(0, 1, 0), ModelFormat::Gguf);
519 a.verify(&OkVerifier, b"weights", b"sig", 1);
520 a.ensure_loadable().expect("verified artifact loads")
521 }
522
523 struct FixedEngine {
526 logits: Vec<i32>,
527 }
528 impl InferenceEngine for FixedEngine {
529 fn prefill(&mut self, t: &[Token]) -> Result<u32> {
530 Ok(t.len() as u32)
531 }
532 fn next_logits(&mut self, _c: &[Token]) -> Vec<i32> {
533 self.logits.clone()
534 }
535 fn eos_token(&self) -> Token {
536 9999
537 }
538 fn rollback(&mut self, _keep: u32) -> Result<()> {
539 Ok(()) }
541 }
542
543 struct DisallowMasker(Vec<Token>);
545 impl GrammarMasker for DisallowMasker {
546 fn mask(&self, _recent: &[Token], vocab: usize) -> Vec<bool> {
547 (0..vocab as Token).map(|t| !self.0.contains(&t)).collect()
548 }
549 }
550
551 #[test]
552 fn full_lifecycle_init_prefill_decode_complete_reset() {
553 let mut s = InferenceSession::new(
554 SessionId(1),
555 SessionConfig::default(),
556 NullEngine::new(3, 8),
557 permit(),
558 );
559 assert_eq!(s.phase(), Phase::Initialized);
560
561 let ports = Ports::permissive();
562 s.load_prompt(&ports, &[10, 11, 12]).unwrap();
563 assert_eq!(s.phase(), Phase::Decoding);
564
565 let stop = s.generate(&ports, 16).unwrap();
566 assert_eq!(stop, StopReason::Eos); assert_eq!(s.output(), &[3]);
568 assert_eq!(s.phase(), Phase::Completed);
569
570 s.reset();
571 assert_eq!(s.phase(), Phase::Initialized);
572 assert!(s.output().is_empty());
573 }
574
575 #[test]
576 fn decode_applies_grammar_before_safety_before_sampling() {
577 let engine = FixedEngine {
579 logits: vec![10, 9, 8, 7],
580 };
581 let mut s = InferenceSession::new(SessionId(2), SessionConfig::default(), engine, permit());
582
583 let ports = Ports {
584 compressor: Box::new(crate::defaults::IdentityCompressor),
585 grammar: Box::new(DisallowMasker(vec![0])), safety: Box::new(LightweightFilter::new(vec![1])), guard: None,
588 ingress: None,
589 relay: None,
590 };
591 s.load_prompt(&ports, &[1]).unwrap();
592 let stop = s.generate(&ports, 1).unwrap();
593
594 assert_eq!(stop, StopReason::MaxTokens);
595 assert_eq!(s.output(), &[2]);
598 }
599
600 #[test]
601 fn generate_before_load_prompt_is_invalid_phase() {
602 let mut s = InferenceSession::new(
603 SessionId(3),
604 SessionConfig::default(),
605 NullEngine::new(0, 4),
606 permit(),
607 );
608 let ports = Ports::permissive();
609 let err = s.generate(&ports, 4).unwrap_err();
610 assert!(matches!(err, EdgeError::InvalidPhase { .. }));
611 }
612
613 #[test]
614 fn relay_is_blocked_unless_hybrid_mode_opted_in() {
615 struct EchoRelay;
616 impl crate::ports::HybridRelay for EchoRelay {
617 fn consult(&self, q: &[Token]) -> Vec<Token> {
618 q.to_vec()
619 }
620 }
621
622 let mut s = InferenceSession::new(
624 SessionId(4),
625 SessionConfig::default(),
626 NullEngine::new(0, 4),
627 permit(),
628 );
629 let ports = Ports {
630 relay: Some(Box::new(EchoRelay)),
631 ..Ports::permissive()
632 };
633 assert_eq!(
634 s.consult_relay(&ports, &[1, 2]).unwrap_err(),
635 EdgeError::AirGapViolation
636 );
637
638 let cfg = SessionConfig {
640 hybrid_mode: true,
641 ..SessionConfig::default()
642 };
643 let mut s2 = InferenceSession::new(SessionId(5), cfg, NullEngine::new(0, 4), permit());
644 assert_eq!(s2.consult_relay(&ports, &[1, 2]).unwrap(), vec![1, 2]);
645
646 let no_relay = Ports::permissive();
648 assert_eq!(
649 s2.consult_relay(&no_relay, &[1]).unwrap_err(),
650 EdgeError::AirGapViolation
651 );
652 }
653
654 #[test]
655 fn first_events_are_init_then_model_loaded() {
656 let mut s = InferenceSession::new(
657 SessionId(6),
658 SessionConfig::default(),
659 NullEngine::new(0, 4),
660 permit(),
661 );
662 let evs = s.drain_events();
663 assert!(matches!(
664 evs[0].event,
665 DomainEvent::SessionInitialized { .. }
666 ));
667 assert!(matches!(evs[1].event, DomainEvent::ModelLoaded { .. }));
668 }
669
670 use el_safety::{ChunkGuard, SafetyScore};
673
674 struct BanToken(Token);
676 impl ChunkGuard for BanToken {
677 fn score(&self, recent: &[Token]) -> SafetyScore {
678 if recent.contains(&self.0) {
679 SafetyScore::MAX
680 } else {
681 SafetyScore::SAFE
682 }
683 }
684 }
685
686 struct AlwaysHot;
688 impl ChunkGuard for AlwaysHot {
689 fn score(&self, _recent: &[Token]) -> SafetyScore {
690 SafetyScore::MAX
691 }
692 }
693
694 fn tiny_policy(max_rollbacks: u8) -> RollbackPolicy {
695 RollbackPolicy {
696 guard_every: 1,
697 steer_window: 0,
698 soft_threshold: SafetyScore::from_milli(500),
699 hard_threshold: SafetyScore::from_milli(800),
700 max_rollbacks,
701 max_checkpoints: 8,
702 }
703 }
704
705 struct DenyAllMasker;
708 impl GrammarMasker for DenyAllMasker {
709 fn mask(&self, _recent: &[Token], vocab: usize) -> Vec<bool> {
710 vec![false; vocab]
711 }
712 }
713
714 fn descending_engine() -> FixedEngine {
717 FixedEngine {
718 logits: vec![5, 4, 3, 2],
719 }
720 }
721
722 struct UnsafeThenEos {
726 eos: Token,
727 vocab: usize,
728 }
729 impl InferenceEngine for UnsafeThenEos {
730 fn prefill(&mut self, t: &[Token]) -> Result<u32> {
731 Ok(t.len() as u32)
732 }
733 fn next_logits(&mut self, ctx: &[Token]) -> Vec<i32> {
734 let mut v = vec![0i32; self.vocab];
735 if ctx.is_empty() {
736 v[0] = 10; } else {
738 v[self.eos as usize] = 10; }
740 v
741 }
742 fn eos_token(&self) -> Token {
743 self.eos
744 }
745 fn rollback(&mut self, _keep: u32) -> Result<()> {
746 Ok(()) }
748 }
749
750 fn coarse_policy(max_rollbacks: u8) -> RollbackPolicy {
753 RollbackPolicy {
754 guard_every: 16,
755 steer_window: 0,
756 soft_threshold: SafetyScore::from_milli(500),
757 hard_threshold: SafetyScore::from_milli(800),
758 max_rollbacks,
759 max_checkpoints: 8,
760 }
761 }
762
763 #[test]
764 fn eos_terminated_short_completion_is_scored_not_bypassed() {
765 let mut s = InferenceSession::new(
768 SessionId(26),
769 SessionConfig::default(),
770 UnsafeThenEos { eos: 5, vocab: 8 },
771 permit(),
772 );
773 let ports = guarded_ports(Box::new(BanToken(0)));
774 s.load_prompt(&ports, &[]).unwrap();
775
776 let stop = s.generate_with_policy(&ports, 8, coarse_policy(0)).unwrap();
778
779 assert_eq!(stop, StopReason::Stopped);
780 assert!(s.output().is_empty()); let evs = s.drain_events();
782 assert!(evs
783 .iter()
784 .any(|e| matches!(e.event, DomainEvent::SafetyViolationDetected { .. })));
785 }
786
787 #[test]
788 fn max_tokens_partial_chunk_is_scored_not_bypassed() {
789 let mut s = InferenceSession::new(
792 SessionId(27),
793 SessionConfig::default(),
794 descending_engine(), permit(),
796 );
797 let ports = guarded_ports(Box::new(BanToken(0)));
798 s.load_prompt(&ports, &[]).unwrap();
799
800 let stop = s.generate_with_policy(&ports, 2, coarse_policy(0)).unwrap();
802
803 assert_eq!(stop, StopReason::Stopped);
804 assert!(s.output().is_empty());
805 let evs = s.drain_events();
806 assert!(evs
807 .iter()
808 .any(|e| matches!(e.event, DomainEvent::SafetyViolationDetected { .. })));
809 }
810
811 #[test]
812 fn eos_unsafe_tail_rolls_back_and_recovers() {
813 let mut s = InferenceSession::new(
817 SessionId(28),
818 SessionConfig::default(),
819 UnsafeThenEos { eos: 5, vocab: 8 },
820 permit(),
821 );
822 let ports = guarded_ports(Box::new(BanToken(0)));
823 s.load_prompt(&ports, &[]).unwrap();
824
825 let stop = s.generate_with_policy(&ports, 8, coarse_policy(1)).unwrap();
826
827 assert_eq!(stop, StopReason::Eos);
828 assert!(!s.output().contains(&0)); assert_eq!(s.output().last(), Some(&5)); let evs = s.drain_events();
831 assert!(evs
832 .iter()
833 .any(|e| matches!(e.event, DomainEvent::ClaimBacktracked { .. })));
834 }
835
836 struct StatefulEngine {
843 fed: usize,
844 logits: Vec<i32>,
845 eos: Token,
846 rollbacks: std::rc::Rc<std::cell::Cell<u32>>,
847 last_keep: std::rc::Rc<std::cell::Cell<u32>>,
848 desynced: std::rc::Rc<std::cell::Cell<bool>>,
849 }
850 impl InferenceEngine for StatefulEngine {
851 fn prefill(&mut self, t: &[Token]) -> Result<u32> {
852 self.fed = 0;
853 Ok(t.len() as u32)
854 }
855 fn next_logits(&mut self, committed: &[Token]) -> Vec<i32> {
856 if self.fed > committed.len() {
859 self.desynced.set(true);
860 }
861 while self.fed < committed.len() {
862 self.fed += 1;
863 }
864 self.logits.clone()
865 }
866 fn eos_token(&self) -> Token {
867 self.eos
868 }
869 fn rollback(&mut self, keep_committed: u32) -> Result<()> {
870 self.fed = keep_committed as usize; self.rollbacks.set(self.rollbacks.get() + 1);
872 self.last_keep.set(keep_committed);
873 Ok(())
874 }
875 }
876
877 #[test]
878 fn rollback_restores_engine_state_not_just_session_metadata() {
879 let rollbacks = std::rc::Rc::new(std::cell::Cell::new(0u32));
883 let last_keep = std::rc::Rc::new(std::cell::Cell::new(u32::MAX));
884 let desynced = std::rc::Rc::new(std::cell::Cell::new(false));
885 let engine = StatefulEngine {
886 fed: 0,
887 logits: vec![5, 4, 3, 2], eos: 9999,
889 rollbacks: rollbacks.clone(),
890 last_keep: last_keep.clone(),
891 desynced: desynced.clone(),
892 };
893 let mut s =
894 InferenceSession::new(SessionId(29), SessionConfig::default(), engine, permit());
895 let ports = guarded_ports(Box::new(BanToken(0)));
896 s.load_prompt(&ports, &[7, 8]).unwrap(); let stop = s.generate_with_policy(&ports, 3, tiny_policy(3)).unwrap();
899
900 assert_eq!(stop, StopReason::MaxTokens);
901 assert_eq!(s.output(), &[1, 1, 1]); assert!(
904 rollbacks.get() >= 1,
905 "session must propagate the backtrack to the engine"
906 );
907 assert!(last_keep.get() < 3);
909 assert!(
910 !desynced.get(),
911 "engine cache must track the session rollback"
912 );
913 }
914
915 fn guarded_ports(guard: Box<dyn ChunkGuard>) -> Ports {
916 Ports {
917 compressor: Box::new(crate::defaults::IdentityCompressor),
918 grammar: Box::new(crate::defaults::AllowAllMasker),
919 safety: Box::new(el_safety::NoSafety),
920 guard: Some(guard),
921 ingress: None,
922 relay: None,
923 }
924 }
925
926 #[test]
927 fn hard_breach_rolls_back_kv_and_recovers() {
928 let mut s = InferenceSession::new(
929 SessionId(20),
930 SessionConfig::default(),
931 descending_engine(),
932 permit(),
933 );
934 let ports = guarded_ports(Box::new(BanToken(0)));
935 s.load_prompt(&ports, &[]).unwrap();
936
937 let stop = s.generate_with_policy(&ports, 3, tiny_policy(3)).unwrap();
938
939 assert_eq!(stop, StopReason::MaxTokens);
940 assert_eq!(s.output(), &[1, 1, 1]);
943 assert!(!s.output().contains(&0));
944 assert_eq!(s.kv_len(), s.output().len() as u32);
946
947 let evs = s.drain_events();
948 assert!(evs
949 .iter()
950 .any(|e| matches!(e.event, DomainEvent::ClaimBacktracked { .. })));
951 assert!(evs
952 .iter()
953 .any(|e| matches!(e.event, DomainEvent::SafetyViolationDetected { .. })));
954 }
955
956 #[test]
957 fn fail_closed_refusal_when_no_rollback_budget() {
958 let mut s = InferenceSession::new(
959 SessionId(21),
960 SessionConfig::default(),
961 descending_engine(),
962 permit(),
963 );
964 let ports = guarded_ports(Box::new(BanToken(0)));
965 s.load_prompt(&ports, &[]).unwrap();
966
967 let stop = s.generate_with_policy(&ports, 5, tiny_policy(0)).unwrap();
969
970 assert_eq!(stop, StopReason::Stopped);
971 assert!(s.output().is_empty());
972 let evs = s.drain_events();
973 assert!(evs
974 .iter()
975 .any(|e| matches!(e.event, DomainEvent::SafetyViolationDetected { .. })));
976 assert!(!evs
977 .iter()
978 .any(|e| matches!(e.event, DomainEvent::ClaimBacktracked { .. })));
979 }
980
981 #[test]
982 fn rollbacks_are_bounded_then_fail_closed() {
983 let mut s = InferenceSession::new(
984 SessionId(22),
985 SessionConfig::default(),
986 descending_engine(),
987 permit(),
988 );
989 let ports = guarded_ports(Box::new(AlwaysHot));
990 s.load_prompt(&ports, &[]).unwrap();
991
992 let stop = s.generate_with_policy(&ports, 8, tiny_policy(2)).unwrap();
993
994 assert_eq!(stop, StopReason::Stopped);
995 let evs = s.drain_events();
996 let rollbacks = evs
997 .iter()
998 .filter(|e| matches!(e.event, DomainEvent::ClaimBacktracked { .. }))
999 .count();
1000 assert_eq!(rollbacks, 2);
1002 }
1003
1004 #[test]
1005 fn memory_pressure_disables_checkpoints_and_fails_closed() {
1006 let cfg = SessionConfig {
1007 memory_budget_bytes: 1024, ..SessionConfig::default()
1009 };
1010 let mut s = InferenceSession::new(SessionId(23), cfg, descending_engine(), permit());
1011 let ports = guarded_ports(Box::new(BanToken(0)));
1012 s.load_prompt(&ports, &[]).unwrap();
1013
1014 let stop = s.generate_with_policy(&ports, 5, tiny_policy(3)).unwrap();
1015
1016 assert_eq!(stop, StopReason::Stopped); let evs = s.drain_events();
1018 assert!(evs.iter().any(|e| matches!(
1019 e.event,
1020 DomainEvent::SafetyDisabled {
1021 reason: DegradeReason::MemoryPressure
1022 }
1023 )));
1024 }
1025
1026 #[test]
1027 fn no_legal_token_fails_closed() {
1028 let mut s = InferenceSession::new(
1031 SessionId(24),
1032 SessionConfig::default(),
1033 descending_engine(),
1034 permit(),
1035 );
1036 let ports = Ports {
1037 compressor: Box::new(crate::defaults::IdentityCompressor),
1038 grammar: Box::new(DenyAllMasker),
1039 safety: Box::new(el_safety::NoSafety),
1040 guard: None,
1041 ingress: None,
1042 relay: None,
1043 };
1044 s.load_prompt(&ports, &[]).unwrap();
1045
1046 let stop = s.generate(&ports, 4).unwrap();
1047
1048 assert_eq!(stop, StopReason::Stopped);
1049 assert!(s.output().is_empty()); let evs = s.drain_events();
1051 assert!(evs
1052 .iter()
1053 .any(|e| matches!(e.event, DomainEvent::GrammarViolationBlocked)));
1054 }
1055
1056 #[test]
1057 fn fail_closed_preserves_prefill_kv() {
1058 let mut s = InferenceSession::new(
1062 SessionId(25),
1063 SessionConfig::default(),
1064 descending_engine(),
1065 permit(),
1066 );
1067 let ports = guarded_ports(Box::new(AlwaysHot));
1068 s.load_prompt(&ports, &[7, 8, 9]).unwrap(); assert_eq!(s.kv_len(), 3);
1070
1071 let stop = s.generate_with_policy(&ports, 8, tiny_policy(1)).unwrap();
1072
1073 assert_eq!(stop, StopReason::Stopped);
1074 assert!(s.output().is_empty()); assert_eq!(s.kv_len(), 3); }
1077
1078 use el_core::SafetyMode;
1081 use el_safety::SafetySteerer;
1082
1083 struct RecordingSteerer {
1086 log: std::rc::Rc<std::cell::RefCell<Vec<(bool, usize)>>>,
1087 }
1088 impl SafetySteerer for RecordingSteerer {
1089 fn adjust(&self, recent: &[Token]) -> LogitAdjustment {
1090 self.log.borrow_mut().push((false, recent.len()));
1091 LogitAdjustment::none()
1092 }
1093 fn adjust_with_logits(&self, recent: &[Token], _base: &[i32]) -> LogitAdjustment {
1094 self.log.borrow_mut().push((true, recent.len()));
1095 LogitAdjustment::none()
1096 }
1097 fn mode(&self) -> SafetyMode {
1098 SafetyMode::SecDecoding
1099 }
1100 }
1101
1102 #[test]
1103 fn soft_steer_applies_only_inside_the_early_token_window() {
1104 let log = std::rc::Rc::new(std::cell::RefCell::new(Vec::new()));
1107 let mut s = InferenceSession::new(
1108 SessionId(30),
1109 SessionConfig::default(),
1110 descending_engine(), permit(),
1112 );
1113 let ports = Ports {
1114 compressor: Box::new(crate::defaults::IdentityCompressor),
1115 grammar: Box::new(crate::defaults::AllowAllMasker),
1116 safety: Box::new(RecordingSteerer { log: log.clone() }),
1117 guard: None,
1118 ingress: None,
1119 relay: None,
1120 };
1121 s.load_prompt(&ports, &[]).unwrap();
1122 let policy = RollbackPolicy {
1123 guard_every: 0,
1124 steer_window: 2,
1125 soft_threshold: SafetyScore::MAX,
1126 hard_threshold: SafetyScore::MAX,
1127 max_rollbacks: 0,
1128 max_checkpoints: 0,
1129 };
1130 s.generate_with_policy(&ports, 4, policy).unwrap();
1131
1132 let calls = log.borrow();
1133 assert_eq!(calls.len(), 4);
1134 for &(with_logits, len) in calls.iter() {
1135 assert_eq!(
1136 with_logits,
1137 len < 2,
1138 "window gate wrong at output len {len}"
1139 );
1140 }
1141 }
1142
1143 #[test]
1144 fn ingress_triage_fails_closed_before_generation() {
1145 let mut s = InferenceSession::new(
1147 SessionId(31),
1148 SessionConfig::default(),
1149 descending_engine(),
1150 permit(),
1151 );
1152 let ports = Ports {
1153 compressor: Box::new(crate::defaults::IdentityCompressor),
1154 grammar: Box::new(crate::defaults::AllowAllMasker),
1155 safety: Box::new(el_safety::NoSafety),
1156 guard: None,
1157 ingress: Some(Box::new(AlwaysHot)), relay: None,
1159 };
1160 s.load_prompt(&ports, &[1, 2, 3]).unwrap();
1161
1162 let stop = s.generate_with_policy(&ports, 8, coarse_policy(0)).unwrap();
1163
1164 assert_eq!(stop, StopReason::Stopped);
1165 assert!(s.output().is_empty());
1166 let evs = s.drain_events();
1167 assert!(evs
1168 .iter()
1169 .any(|e| matches!(e.event, DomainEvent::SafetyViolationDetected { .. })));
1170 assert!(!evs
1172 .iter()
1173 .any(|e| matches!(e.event, DomainEvent::TokenCommitted { .. })));
1174 }
1175
1176 #[test]
1177 fn generate_applies_safety_mode_selector_and_records_effective_mode() {
1178 let cfg = SessionConfig {
1181 device: el_core::DeviceTarget::MidRange,
1182 safety: SafetyMode::SecDecoding,
1183 ..SessionConfig::default()
1184 };
1185 let mut s = InferenceSession::new(SessionId(32), cfg, NullEngine::new(0, 4), permit());
1186 let ports = Ports::permissive();
1187 s.load_prompt(&ports, &[1]).unwrap();
1188 s.generate(&ports, 4).unwrap();
1189
1190 let evs = s.drain_events();
1191 assert!(evs.iter().any(|e| matches!(
1192 e.event,
1193 DomainEvent::SafetyModeSelected {
1194 mode: SafetyMode::Lightweight
1195 }
1196 )));
1197 }
1198}