Skip to main content

el_runtime/
session.rs

1//! The `InferenceSession` aggregate root and decode-loop orchestrator.
2
3use crate::ports::{InferenceEngine, Ports};
4use el_core::{
5    DomainEvent, EdgeError, EventEnvelope, Phase, Result, SessionConfig, SessionId, StopReason,
6    Token,
7};
8use el_memory::KvRegion;
9use el_provenance::LoadPermit;
10use el_safety::LogitAdjustment;
11
12/// One live generation. Constructing it requires a [`LoadPermit`], so a model
13/// that has not passed the provenance gate (ADR-006) cannot reach the runtime —
14/// the Conformist relationship is enforced in the type system.
15pub struct InferenceSession<E: InferenceEngine> {
16    id: SessionId,
17    config: SessionConfig,
18    phase: Phase,
19    engine: E,
20    kv: KvRegion,
21    permit: LoadPermit,
22    output: Vec<Token>,
23    step: u32,
24    events: Vec<EventEnvelope>,
25}
26
27impl<E: InferenceEngine> InferenceSession<E> {
28    pub fn new(id: SessionId, config: SessionConfig, engine: E, permit: LoadPermit) -> Self {
29        let mut s = Self {
30            id,
31            config,
32            phase: Phase::Initialized,
33            engine,
34            kv: KvRegion::new(),
35            permit,
36            output: Vec::new(),
37            step: 0,
38            events: Vec::new(),
39        };
40        s.emit(DomainEvent::SessionInitialized {
41            runtime: config.format.runtime(),
42            device: config.device,
43            safety: config.safety,
44            speculation: config.speculation,
45        });
46        s.emit(DomainEvent::ModelLoaded {
47            model: permit.model,
48            version: permit.version,
49            format: permit.format,
50        });
51        s
52    }
53
54    pub fn phase(&self) -> Phase {
55        self.phase
56    }
57    pub fn output(&self) -> &[Token] {
58        &self.output
59    }
60    pub fn kv_len(&self) -> u32 {
61        self.kv.len()
62    }
63    pub fn config(&self) -> &SessionConfig {
64        &self.config
65    }
66    /// The load permit this session was constructed with — evidence the model
67    /// passed the provenance gate (ADR-006).
68    pub fn permit(&self) -> LoadPermit {
69        self.permit
70    }
71    /// Take the buffered domain events (a real build would stream these to the
72    /// Telemetry subscriber).
73    pub fn drain_events(&mut self) -> Vec<EventEnvelope> {
74        std::mem::take(&mut self.events)
75    }
76
77    fn emit(&mut self, event: DomainEvent) {
78        self.events
79            .push(EventEnvelope::new(self.id, self.step, event));
80    }
81
82    /// Compress (optional) → prefill → build KV. Valid only from `Initialized`.
83    pub fn load_prompt(&mut self, ports: &Ports, prompt: &[Token]) -> Result<()> {
84        if self.phase != Phase::Initialized {
85            return Err(EdgeError::InvalidPhase {
86                expected: "Initialized",
87                found: self.phase.as_str(),
88            });
89        }
90
91        let compressed = if self.config.compress {
92            ports.compressor.compress(prompt)
93        } else {
94            prompt.to_vec()
95        };
96        if compressed.len() < prompt.len() {
97            let ratio_milli =
98                ((compressed.len() as u64 * 1000) / (prompt.len().max(1) as u64)) as u32;
99            self.emit(DomainEvent::PromptCompressed {
100                input_tokens: prompt.len() as u32,
101                output_tokens: compressed.len() as u32,
102                ratio_milli,
103            });
104        }
105
106        self.phase = Phase::Prefilling;
107        let kv_len = self.engine.prefill(&compressed)?;
108        for _ in 0..kv_len {
109            let off = self.kv.len() as u64;
110            self.kv.push(off);
111        }
112        self.emit(DomainEvent::PrefillCompleted {
113            prompt_tokens: compressed.len() as u32,
114            kv_len,
115            prefill_tps: 0,
116        });
117        self.phase = Phase::Decoding;
118        Ok(())
119    }
120
121    /// Run the decode loop until EOS or `max_tokens`. Each step composes
122    /// collaborators in the invariant order: grammar mask → safety adjust →
123    /// sample → commit.
124    pub fn generate(&mut self, ports: &Ports, max_tokens: u32) -> Result<StopReason> {
125        if self.phase != Phase::Decoding {
126            return Err(EdgeError::InvalidPhase {
127                expected: "Decoding",
128                found: self.phase.as_str(),
129            });
130        }
131        let eos = self.engine.eos_token();
132
133        let stop = loop {
134            if self.output.len() as u32 >= max_tokens {
135                break StopReason::MaxTokens;
136            }
137
138            // 2. verify / next-token logits (1. drafting is off by default)
139            let logits = self.engine.next_logits(&self.output);
140            let vocab = logits.len();
141
142            // 3. grammar mask (BEFORE safety)
143            let mask = ports.grammar.mask(&self.output, vocab);
144            let allowed = mask.iter().filter(|b| **b).count() as u32;
145            self.emit(DomainEvent::TokenMaskApplied { allowed });
146
147            // 4. safety adjust (AFTER mask, BEFORE sampling)
148            let adj = ports.safety.adjust(&self.output);
149            if !adj.is_empty() {
150                self.emit(DomainEvent::LogitsSteered {
151                    adjustment_norm_milli: adj.l1_norm_milli(),
152                });
153            }
154
155            // 5. sample (greedy argmax over legal, steered logits)
156            let token = pick(&logits, &mask, &adj);
157            self.emit(DomainEvent::TokenGenerated { sampled: false });
158
159            // 6. commit
160            self.output.push(token);
161            self.kv.push(self.output.len() as u64);
162            self.step += 1;
163            self.emit(DomainEvent::TokenCommitted {
164                kv_len: self.kv.len(),
165            });
166
167            if token == eos {
168                break StopReason::Eos;
169            }
170        };
171
172        self.phase = Phase::Completed;
173        self.emit(DomainEvent::GenerationCompleted {
174            total_tokens: self.output.len() as u32,
175            stop,
176        });
177        Ok(stop)
178    }
179
180    /// Clear KV/output for a fresh conversation (volatile memory only).
181    pub fn reset(&mut self) {
182        self.kv = KvRegion::new();
183        self.output.clear();
184        self.step = 0;
185        self.phase = Phase::Initialized;
186        self.emit(DomainEvent::SessionReset);
187    }
188
189    /// Consult the opt-in LAN relay. Hard-fails with [`EdgeError::AirGapViolation`]
190    /// unless `hybrid_mode` is enabled AND a relay is wired (ADR-004).
191    pub fn consult_relay(&mut self, ports: &Ports, query: &[Token]) -> Result<Vec<Token>> {
192        if !self.config.hybrid_mode {
193            return Err(EdgeError::AirGapViolation);
194        }
195        match &ports.relay {
196            Some(relay) => {
197                let out = relay.consult(query);
198                self.emit(DomainEvent::HybridRelayConsulted);
199                Ok(out)
200            }
201            None => Err(EdgeError::AirGapViolation),
202        }
203    }
204}
205
206/// Greedy pick over legal, safety-steered logits. Masked-out tokens are skipped
207/// entirely; the safety delta is added to surviving logits before argmax.
208fn pick(logits: &[i32], mask: &[bool], adj: &LogitAdjustment) -> Token {
209    let mut best: Option<Token> = None;
210    let mut best_val = i32::MIN;
211    for (i, &l) in logits.iter().enumerate() {
212        if mask.get(i).copied() == Some(false) {
213            continue;
214        }
215        let v = l.saturating_add(adj.delta_for(i as Token));
216        if v > best_val {
217            best_val = v;
218            best = Some(i as Token);
219        }
220    }
221    best.unwrap_or(0)
222}
223
224#[cfg(test)]
225mod tests {
226    use super::*;
227    use crate::defaults::NullEngine;
228    use crate::ports::{GrammarMasker, Ports};
229    use el_core::{ModelFormat, ModelId, ModelVersion};
230    use el_provenance::{ModelArtifact, SignatureVerifier};
231    use el_safety::LightweightFilter;
232
233    struct OkVerifier;
234    impl SignatureVerifier for OkVerifier {
235        fn verify(&self, _b: &[u8], _s: &[u8], _k: u32) -> bool {
236            true
237        }
238    }
239
240    fn permit() -> LoadPermit {
241        let mut a = ModelArtifact::new(ModelId(1), ModelVersion::new(0, 1, 0), ModelFormat::Gguf);
242        a.verify(&OkVerifier, b"weights", b"sig", 1);
243        a.ensure_loadable().expect("verified artifact loads")
244    }
245
246    /// A deterministic engine returning fixed logits; eos out of vocab range so
247    /// it never self-terminates (used for the composition-order test).
248    struct FixedEngine {
249        logits: Vec<i32>,
250    }
251    impl InferenceEngine for FixedEngine {
252        fn prefill(&mut self, t: &[Token]) -> Result<u32> {
253            Ok(t.len() as u32)
254        }
255        fn next_logits(&mut self, _c: &[Token]) -> Vec<i32> {
256            self.logits.clone()
257        }
258        fn eos_token(&self) -> Token {
259            9999
260        }
261    }
262
263    // Grammar masker that disallows specific token ids.
264    struct DisallowMasker(Vec<Token>);
265    impl GrammarMasker for DisallowMasker {
266        fn mask(&self, _recent: &[Token], vocab: usize) -> Vec<bool> {
267            (0..vocab as Token).map(|t| !self.0.contains(&t)).collect()
268        }
269    }
270
271    #[test]
272    fn full_lifecycle_init_prefill_decode_complete_reset() {
273        let mut s = InferenceSession::new(
274            SessionId(1),
275            SessionConfig::default(),
276            NullEngine::new(3, 8),
277            permit(),
278        );
279        assert_eq!(s.phase(), Phase::Initialized);
280
281        let ports = Ports::permissive();
282        s.load_prompt(&ports, &[10, 11, 12]).unwrap();
283        assert_eq!(s.phase(), Phase::Decoding);
284
285        let stop = s.generate(&ports, 16).unwrap();
286        assert_eq!(stop, StopReason::Eos); // NullEngine emits EOS first step
287        assert_eq!(s.output(), &[3]);
288        assert_eq!(s.phase(), Phase::Completed);
289
290        s.reset();
291        assert_eq!(s.phase(), Phase::Initialized);
292        assert!(s.output().is_empty());
293    }
294
295    #[test]
296    fn decode_applies_grammar_before_safety_before_sampling() {
297        // logits favour token 0 (10), then 1 (9), then 2 (8), then 3 (7).
298        let engine = FixedEngine {
299            logits: vec![10, 9, 8, 7],
300        };
301        let mut s = InferenceSession::new(SessionId(2), SessionConfig::default(), engine, permit());
302
303        let ports = Ports {
304            compressor: Box::new(crate::defaults::IdentityCompressor),
305            grammar: Box::new(DisallowMasker(vec![0])), // grammar removes the top token
306            safety: Box::new(LightweightFilter::new(vec![1])), // safety bans the next-best
307            relay: None,
308        };
309        s.load_prompt(&ports, &[1]).unwrap();
310        let stop = s.generate(&ports, 1).unwrap();
311
312        assert_eq!(stop, StopReason::MaxTokens);
313        // Token 0 removed by grammar, token 1 banned by safety → token 2 wins.
314        // Proves order: mask → adjust → sample.
315        assert_eq!(s.output(), &[2]);
316    }
317
318    #[test]
319    fn generate_before_load_prompt_is_invalid_phase() {
320        let mut s = InferenceSession::new(
321            SessionId(3),
322            SessionConfig::default(),
323            NullEngine::new(0, 4),
324            permit(),
325        );
326        let ports = Ports::permissive();
327        let err = s.generate(&ports, 4).unwrap_err();
328        assert!(matches!(err, EdgeError::InvalidPhase { .. }));
329    }
330
331    #[test]
332    fn relay_is_blocked_unless_hybrid_mode_opted_in() {
333        struct EchoRelay;
334        impl crate::ports::HybridRelay for EchoRelay {
335            fn consult(&self, q: &[Token]) -> Vec<Token> {
336                q.to_vec()
337            }
338        }
339
340        // Air-gapped by default: even with a relay wired, consulting fails.
341        let mut s = InferenceSession::new(
342            SessionId(4),
343            SessionConfig::default(),
344            NullEngine::new(0, 4),
345            permit(),
346        );
347        let ports = Ports {
348            relay: Some(Box::new(EchoRelay)),
349            ..Ports::permissive()
350        };
351        assert_eq!(
352            s.consult_relay(&ports, &[1, 2]).unwrap_err(),
353            EdgeError::AirGapViolation
354        );
355
356        // Opt in → allowed.
357        let cfg = SessionConfig {
358            hybrid_mode: true,
359            ..SessionConfig::default()
360        };
361        let mut s2 = InferenceSession::new(SessionId(5), cfg, NullEngine::new(0, 4), permit());
362        assert_eq!(s2.consult_relay(&ports, &[1, 2]).unwrap(), vec![1, 2]);
363
364        // Opted in but no relay wired → still air-gapped.
365        let no_relay = Ports::permissive();
366        assert_eq!(
367            s2.consult_relay(&no_relay, &[1]).unwrap_err(),
368            EdgeError::AirGapViolation
369        );
370    }
371
372    #[test]
373    fn first_events_are_init_then_model_loaded() {
374        let mut s = InferenceSession::new(
375            SessionId(6),
376            SessionConfig::default(),
377            NullEngine::new(0, 4),
378            permit(),
379        );
380        let evs = s.drain_events();
381        assert!(matches!(
382            evs[0].event,
383            DomainEvent::SessionInitialized { .. }
384        ));
385        assert!(matches!(evs[1].event, DomainEvent::ModelLoaded { .. }));
386    }
387}