1use 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
12pub 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 pub fn permit(&self) -> LoadPermit {
69 self.permit
70 }
71 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 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 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 let logits = self.engine.next_logits(&self.output);
140 let vocab = logits.len();
141
142 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 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 let token = pick(&logits, &mask, &adj);
157 self.emit(DomainEvent::TokenGenerated { sampled: false });
158
159 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 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 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
206fn 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 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 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); 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 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])), safety: Box::new(LightweightFilter::new(vec![1])), 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 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 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 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 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}