Skip to main content

eredu_runtime/
state.rs

1//! Backend-neutral mutable model-state contracts.
2//!
3//! Architectures declare state geometry through [`StateLayout`]. Runtime
4//! policies select a residency realization, while concrete backends retain
5//! their native layer-state and tensor types.
6
7use std::{marker::PhantomData, ops::Range};
8
9use eredu_core::{
10    cache::{
11        LayerCachePolicy, PromptCacheError, PromptCacheModelIdentity, PromptCacheStateSegment,
12        PromptCacheTopology, StateComponentPolicy, StateTensorRole,
13    },
14    LayerSchedule,
15};
16use eredu_nn::NeuralBackend;
17
18/// Architecture-declared placement of one contiguous mutable-state range.
19#[derive(Debug, Clone, Copy, Eq, PartialEq)]
20#[non_exhaustive]
21pub enum ArchitectureStatePlacement {
22    /// Partition the state range in lockstep with one execution group's units.
23    GroupUnits {
24        /// Canonical execution-group slot.
25        group: usize,
26    },
27    /// Attach the complete state range to the realized architecture output owner.
28    OutputOwner,
29}
30
31/// One architecture-authored rule in a mutable-state partition plan.
32#[derive(Debug, Clone, Eq, PartialEq)]
33pub struct ArchitectureStatePartitionRule {
34    layers: Range<usize>,
35    placement: ArchitectureStatePlacement,
36}
37
38impl ArchitectureStatePartitionRule {
39    /// Aligns a state range one-for-one with an execution group's unit indices.
40    pub fn group_units(group: usize, layers: Range<usize>) -> Self {
41        Self {
42            layers,
43            placement: ArchitectureStatePlacement::GroupUnits { group },
44        }
45    }
46
47    /// Attaches a complete state range to the architecture output owner.
48    pub fn output_owner(layers: Range<usize>) -> Self {
49        Self {
50            layers,
51            placement: ArchitectureStatePlacement::OutputOwner,
52        }
53    }
54
55    /// Returns the architecture-global state-layer range governed by this rule.
56    pub fn layers(&self) -> Range<usize> {
57        self.layers.clone()
58    }
59
60    /// Returns the rule's neutral placement semantics.
61    pub const fn placement(&self) -> ArchitectureStatePlacement {
62        self.placement
63    }
64}
65
66/// Complete architecture-authored mutable-state partition policy.
67///
68/// Resolution validates that the rules cover the supplied [`StateLayout`]
69/// exactly once and that unit-aligned ranges match their execution groups.
70#[derive(Debug, Clone, Eq, PartialEq)]
71pub struct ArchitectureStatePartitionPlan {
72    rules: Vec<ArchitectureStatePartitionRule>,
73}
74
75impl ArchitectureStatePartitionPlan {
76    /// Collects the architecture's state placement rules in declaration order.
77    pub fn new(rules: impl IntoIterator<Item = ArchitectureStatePartitionRule>) -> Self {
78        Self {
79            rules: rules.into_iter().collect(),
80        }
81    }
82
83    /// Returns the declared state placement rules.
84    pub fn rules(&self) -> &[ArchitectureStatePartitionRule] {
85        &self.rules
86    }
87}
88
89/// Invalid architecture-authored mutable-state partition policy.
90#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
91#[non_exhaustive]
92pub enum ArchitectureStatePartitionError {
93    /// The plan contains no placement rules.
94    #[error("architecture state partition plan must contain at least one rule")]
95    EmptyPlan,
96    /// A rule selected no state layers.
97    #[error("architecture state partition rule has empty range {start}..{end}")]
98    EmptyRange {
99        /// Inclusive invalid range start.
100        start: usize,
101        /// Exclusive invalid range end.
102        end: usize,
103    },
104    /// A rule selected state layers outside the complete layout.
105    #[error("architecture state partition range {start}..{end} exceeds the {layers}-layer layout")]
106    RangeOutOfBounds {
107        /// Inclusive invalid range start.
108        start: usize,
109        /// Exclusive invalid range end.
110        end: usize,
111        /// Complete state-layout length.
112        layers: usize,
113    },
114    /// Two rules selected the same state layer.
115    #[error(
116        "architecture state partition range starts at {start}, before prior frontier {frontier}"
117    )]
118    OverlappingRange {
119        /// Inclusive overlapping range start.
120        start: usize,
121        /// End of the preceding declared range.
122        frontier: usize,
123    },
124    /// No rule selected one or more state layers.
125    #[error("architecture state layer {layer} is not assigned by the partition plan")]
126    UnassignedLayer {
127        /// First state layer without a rule.
128        layer: usize,
129    },
130    /// A unit-aligned rule named a nonexistent execution group.
131    #[error("architecture state partition references unknown execution group {group}")]
132    UnknownGroup {
133        /// Missing canonical execution-group slot.
134        group: usize,
135    },
136    /// A unit-aligned state range and its execution group have different lengths.
137    #[error(
138        "architecture state range {start}..{end} has {} layers but group {group} has {units} units",
139        end - start
140    )]
141    GroupLengthMismatch {
142        /// Canonical execution-group slot.
143        group: usize,
144        /// Inclusive state-range start.
145        start: usize,
146        /// Exclusive state-range end.
147        end: usize,
148        /// Execution-group unit count.
149        units: usize,
150    },
151    /// This partition's selected state ranges cannot use the contiguous state representation.
152    #[error(
153        "architecture state partition selects discontiguous ranges ending at {frontier} and starting at {start}"
154    )]
155    DiscontiguousSelection {
156        /// End of the preceding selected range.
157        frontier: usize,
158        /// Start of the next selected range.
159        start: usize,
160    },
161    /// The selected state layout could not be sliced from the complete layout.
162    #[error("architecture state partition layout is invalid: {0}")]
163    InvalidLayout(String),
164}
165
166/// Stable name assigned to the implicit segment of a simple state layout.
167pub const DEFAULT_STATE_SEGMENT_ID: &str = "state";
168
169/// Validated stable identity of one contiguous mutable-state segment.
170#[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)]
171pub struct StateSegmentId(String);
172
173impl StateSegmentId {
174    /// Creates a non-empty stable segment identity.
175    pub fn new(id: impl Into<String>) -> Result<Self, StateError> {
176        let id = id.into();
177        if id.trim().is_empty() {
178            return Err(StateError::EmptySegmentId);
179        }
180        Ok(Self(id))
181    }
182
183    /// Returns the stable segment name.
184    pub fn as_str(&self) -> &str {
185        &self.0
186    }
187}
188
189impl std::fmt::Display for StateSegmentId {
190    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
191        formatter.write_str(&self.0)
192    }
193}
194
195/// Lifetime policy attached to a named mutable-state segment.
196#[derive(Debug, Clone, Copy, Eq, PartialEq, Ord, PartialOrd, Hash)]
197#[non_exhaustive]
198pub enum StateSegmentLifetime {
199    /// State survives from one model input or frame to the next.
200    Persistent,
201    /// State is reused within one frame and reset at the frame boundary.
202    FrameLocal,
203}
204
205/// One named contiguous range in an architecture's state layout.
206#[derive(Debug, Clone, Eq, PartialEq)]
207pub struct StateSegmentSpec {
208    id: StateSegmentId,
209    layers: Range<usize>,
210    lifetime: StateSegmentLifetime,
211    processed_token_offset: i32,
212}
213
214impl StateSegmentSpec {
215    /// Creates a non-empty segment range.
216    pub fn new(
217        id: impl Into<String>,
218        layers: Range<usize>,
219        lifetime: StateSegmentLifetime,
220        processed_token_offset: i32,
221    ) -> Result<Self, StateError> {
222        let id = StateSegmentId::new(id)?;
223        if layers.is_empty() {
224            return Err(StateError::EmptySegmentRange {
225                segment: id,
226                start: layers.start,
227                end: layers.end,
228            });
229        }
230        if processed_token_offset > 0 {
231            return Err(StateError::PositiveSegmentOffset {
232                segment: id,
233                offset: processed_token_offset,
234            });
235        }
236        Ok(Self {
237            id,
238            layers,
239            lifetime,
240            processed_token_offset,
241        })
242    }
243
244    /// Returns the stable segment identity.
245    pub const fn id(&self) -> &StateSegmentId {
246        &self.id
247    }
248
249    /// Returns the architecture-global state-layer range.
250    pub fn layers(&self) -> Range<usize> {
251        self.layers.clone()
252    }
253
254    /// Returns whether the segment persists or resets at a frame boundary.
255    pub const fn lifetime(&self) -> StateSegmentLifetime {
256        self.lifetime
257    }
258
259    /// Returns the segment's processed-token delta from the persisted prefix.
260    pub const fn processed_token_offset(&self) -> i32 {
261        self.processed_token_offset
262    }
263}
264
265/// Complete ordered mutable-state geometry owned by one model instance.
266#[derive(Debug, Clone, Eq, PartialEq)]
267pub struct StateLayout {
268    layers: LayerSchedule<LayerCachePolicy>,
269    components: Vec<Vec<StateComponentPolicy>>,
270    segments: Vec<StateSegmentSpec>,
271}
272
273impl StateLayout {
274    /// Creates and validates a simple ordered layout with one persistent
275    /// segment named [`DEFAULT_STATE_SEGMENT_ID`].
276    pub fn new(layers: LayerSchedule<LayerCachePolicy>) -> Result<Self, StateError> {
277        if layers.is_empty() {
278            return Err(StateError::EmptyLayout);
279        }
280        let count = layers.len();
281        Self::segmented(
282            layers,
283            [StateSegmentSpec::new(
284                DEFAULT_STATE_SEGMENT_ID,
285                0..count,
286                StateSegmentLifetime::Persistent,
287                0,
288            )?],
289        )
290    }
291
292    /// Creates an ordered layout partitioned into named contiguous segments.
293    ///
294    /// Segment declarations are sorted into layer order and must form an exact,
295    /// non-overlapping partition of every state-bearing layer. Stable segment
296    /// identity and lifetime therefore participate in layout equality and
297    /// runtime-state compatibility checks.
298    pub fn segmented(
299        layers: LayerSchedule<LayerCachePolicy>,
300        segments: impl IntoIterator<Item = StateSegmentSpec>,
301    ) -> Result<Self, StateError> {
302        if layers.is_empty() {
303            return Err(StateError::EmptyLayout);
304        }
305        for (layer, policy) in layers.iter().enumerate() {
306            policy
307                .validate()
308                .map_err(|error| StateError::InvalidLayer {
309                    layer,
310                    reason: error.to_string(),
311                })?;
312        }
313        let components = layers.iter().map(LayerCachePolicy::components).collect();
314        let mut segments = segments.into_iter().collect::<Vec<_>>();
315        if segments.is_empty() {
316            return Err(StateError::EmptySegments);
317        }
318        segments.sort_by(|left, right| {
319            left.layers
320                .start
321                .cmp(&right.layers.start)
322                .then_with(|| left.layers.end.cmp(&right.layers.end))
323                .then_with(|| left.id.cmp(&right.id))
324        });
325        let mut identities = std::collections::BTreeSet::new();
326        let mut frontier = 0usize;
327        for segment in &segments {
328            if !identities.insert(segment.id.clone()) {
329                return Err(StateError::DuplicateSegment {
330                    segment: segment.id.clone(),
331                });
332            }
333            if segment.layers.end > layers.len() {
334                return Err(StateError::SegmentOutOfBounds {
335                    segment: segment.id.clone(),
336                    start: segment.layers.start,
337                    end: segment.layers.end,
338                    layers: layers.len(),
339                });
340            }
341            if segment.layers.start < frontier {
342                return Err(StateError::OverlappingSegment {
343                    segment: segment.id.clone(),
344                    start: segment.layers.start,
345                    frontier,
346                });
347            }
348            if segment.layers.start > frontier {
349                return Err(StateError::UnassignedStateLayer { layer: frontier });
350            }
351            frontier = segment.layers.end;
352        }
353        if frontier != layers.len() {
354            return Err(StateError::UnassignedStateLayer { layer: frontier });
355        }
356        Ok(Self {
357            layers,
358            components,
359            segments,
360        })
361    }
362
363    /// Returns the number of architecture-global layers represented here.
364    pub fn len(&self) -> usize {
365        self.layers.len()
366    }
367
368    /// Returns whether this layout has no layers.
369    pub fn is_empty(&self) -> bool {
370        self.layers.is_empty()
371    }
372
373    /// Returns one layer's exact state policy.
374    pub fn layer(&self, layer: usize) -> Option<&LayerCachePolicy> {
375        self.layers.get(layer)
376    }
377
378    /// Borrows the portable ordered layer schedule.
379    pub const fn layers(&self) -> &LayerSchedule<LayerCachePolicy> {
380        &self.layers
381    }
382
383    /// Returns ordered named semantic components for one layer.
384    pub fn components(&self, layer: usize) -> Option<&[StateComponentPolicy]> {
385        self.components.get(layer).map(Vec::as_slice)
386    }
387
388    /// Returns named state segments in deterministic layer order.
389    pub fn segments(&self) -> &[StateSegmentSpec] {
390        &self.segments
391    }
392
393    /// Resolves one named state segment.
394    pub fn segment(&self, id: &StateSegmentId) -> Option<&StateSegmentSpec> {
395        self.segments.iter().find(|segment| segment.id() == id)
396    }
397
398    /// Resolves the unique named segment containing one state layer.
399    pub fn segment_for_layer(&self, layer: usize) -> Option<&StateSegmentSpec> {
400        self.segments
401            .iter()
402            .find(|segment| segment.layers.contains(&layer))
403    }
404
405    /// Expands architecture-declared segment frontiers into layer order.
406    pub fn layer_prefix_offsets(&self) -> Vec<i32> {
407        let mut offsets = Vec::with_capacity(self.len());
408        for segment in &self.segments {
409            offsets.extend(std::iter::repeat_n(
410                segment.processed_token_offset(),
411                segment.layers.len(),
412            ));
413        }
414        offsets
415    }
416
417    /// Selects a contiguous architecture-global range while preserving the
418    /// intersecting segment identities, lifetimes, and token frontiers.
419    pub fn slice(&self, layers: Range<usize>) -> Result<Self, StateError> {
420        if layers.is_empty() || layers.end > self.len() {
421            return Err(StateError::InvalidLayoutSlice {
422                start: layers.start,
423                end: layers.end,
424                layers: self.len(),
425            });
426        }
427        let policies = self
428            .layers
429            .iter()
430            .skip(layers.start)
431            .take(layers.len())
432            .cloned()
433            .collect::<Vec<_>>();
434        let mut segments = Vec::new();
435        for segment in &self.segments {
436            let start = segment.layers.start.max(layers.start);
437            let end = segment.layers.end.min(layers.end);
438            if start < end {
439                segments.push(StateSegmentSpec::new(
440                    segment.id.as_str(),
441                    start - layers.start..end - layers.start,
442                    segment.lifetime,
443                    segment.processed_token_offset,
444                )?);
445            }
446        }
447        Self::segmented(
448            LayerSchedule::new(policies.len(), policies)
449                .map_err(|error| StateError::InvalidResidency(error.to_string()))?,
450            segments,
451        )
452    }
453}
454
455/// Concrete layer state capable of exposing backend-native retained tensors.
456pub trait RuntimeLayerState<B: NeuralBackend> {
457    /// Allocation-free iterator returned for one layer.
458    type RetainedValues<'a>: Iterator<Item = &'a B::Tensor>
459    where
460        Self: 'a,
461        B::Tensor: 'a;
462
463    /// Borrows tensors that must remain alive through this layer's submission.
464    fn retained_values(&self) -> Self::RetainedValues<'_>;
465}
466
467/// Reset capability for one concrete backend-native layer state.
468///
469/// Reset drops the semantic contents of the layer state without replacing its
470/// concrete cache type or inspecting backend-native values on the host.
471pub trait ResettableRuntimeLayerState<B: NeuralBackend>: RuntimeLayerState<B> {
472    /// Clears this layer state to its initial empty value.
473    fn reset(&mut self) -> Result<(), StateError>;
474}
475
476/// Mutable access to architecture-declared fixed state components.
477///
478/// Operators address semantic roles rather than backend storage. Concrete
479/// realizations keep native tensors and may combine these slots with an
480/// append-only attention cache in the same layer state.
481pub trait RuntimeStateComponents<B: NeuralBackend>: RuntimeLayerState<B> {
482    /// Current absolute token frontier for this layer.
483    fn position(&self) -> i32;
484
485    /// Borrows the optional tensor slot for one declared fixed component.
486    fn fixed_component(
487        &mut self,
488        role: StateTensorRole,
489    ) -> Result<&mut Option<B::Tensor>, StateError>;
490
491    /// Advances a fixed-state-only layer after a successful operator call.
492    fn advance_fixed(&mut self, tokens: i32) -> Result<(), StateError>;
493}
494
495/// Mutable state realization consumed by generic resident and layerwise engines.
496pub trait RuntimeState<B: NeuralBackend> {
497    /// Concrete iterator retaining native values for one execution unit.
498    type RetainedValues<'a>: Iterator<Item = &'a B::Tensor>
499    where
500        Self: 'a,
501        B::Tensor: 'a;
502
503    /// Returns the exact layout used to create this realization.
504    fn layout(&self) -> &StateLayout;
505
506    /// Borrows tensors retained by one execution unit without cloning handles.
507    ///
508    /// The flat ordinal addresses policy storage while `address` preserves
509    /// architecture-group semantics for composite and shared state.
510    fn retained_values(
511        &self,
512        ordinal: usize,
513        address: crate::ExecutionUnitAddress,
514    ) -> Result<Self::RetainedValues<'_>, StateError>;
515}
516
517/// Additive backend mechanism for realizing an architecture-declared state layout.
518///
519/// Ordinary key/value architectures need not implement this extension: it is
520/// selected only by composition that requires a distinct concrete state
521/// representation, such as a layout combining attention and fixed components.
522pub trait ArchitectureStateFactory<B: NeuralBackend> {
523    /// Concrete state returned by this realization mechanism.
524    type State: RuntimeState<B>;
525    /// Backend-specific construction failure.
526    type Error;
527
528    /// Allocates native state for the exact selected architecture layout.
529    fn realize(&mut self, layout: &StateLayout) -> Result<Self::State, Self::Error>;
530}
531
532/// Failure while selecting and realizing architecture-authored mutable state.
533#[derive(Debug, thiserror::Error)]
534#[non_exhaustive]
535pub enum ArchitectureStateRealizationError<ArchitectureError, FactoryError> {
536    /// The architecture could not derive its authoritative state layout.
537    #[error("architecture state layout selection failed")]
538    Architecture(#[source] ArchitectureError),
539    /// The selected backend mechanism could not realize the layout.
540    #[error("backend state realization failed")]
541    Factory(#[source] FactoryError),
542    /// The backend returned state whose layout differs from the selected value.
543    #[error("backend state realization changed the selected architecture layout")]
544    LayoutMismatch,
545}
546
547/// Selects the architecture's exact state layout and realizes it through an
548/// explicitly supplied additive mechanism.
549pub fn realize_architecture_state<B, M, F>(
550    architecture: &M,
551    factory: &mut F,
552) -> Result<F::State, ArchitectureStateRealizationError<M::DefinitionError, F::Error>>
553where
554    B: NeuralBackend,
555    M: crate::ArchitectureParameters<B>,
556    F: ArchitectureStateFactory<B>,
557{
558    let layout = architecture
559        .state_layout()
560        .map_err(ArchitectureStateRealizationError::Architecture)?;
561    let state = factory
562        .realize(&layout)
563        .map_err(ArchitectureStateRealizationError::Factory)?;
564    if state.layout() != &layout {
565        return Err(ArchitectureStateRealizationError::LayoutMismatch);
566    }
567    Ok(state)
568}
569
570/// Named-segment reset supported by a concrete runtime-state realization.
571pub trait ResettableRuntimeState<B: NeuralBackend>: RuntimeState<B> {
572    /// Resets every layer in exactly one declared state segment.
573    fn reset_segment(&mut self, segment: &StateSegmentId) -> Result<(), StateError>;
574}
575
576/// Optional capability for architectures with one indexed state per layer.
577pub trait LayerRuntimeState<B: NeuralBackend>: RuntimeState<B> {
578    /// Concrete monomorphized layer state used by the architecture.
579    type LayerState: RuntimeLayerState<B>;
580
581    /// Mutably borrows one architecture-global layer state.
582    fn layer(&mut self, layer: usize) -> Result<&mut Self::LayerState, StateError>;
583}
584
585/// Fully device-resident state with one concrete value per architecture layer.
586#[derive(Debug)]
587pub struct DeviceState<B: NeuralBackend, L> {
588    layout: StateLayout,
589    layers: Vec<L>,
590    backend: PhantomData<fn() -> B>,
591}
592
593impl<B: NeuralBackend, L: Clone> Clone for DeviceState<B, L> {
594    fn clone(&self) -> Self {
595        Self {
596            layout: self.layout.clone(),
597            layers: self.layers.clone(),
598            backend: PhantomData,
599        }
600    }
601
602    fn clone_from(&mut self, source: &Self) {
603        self.layout.clone_from(&source.layout);
604        self.layers.clone_from(&source.layers);
605    }
606}
607
608impl<B: NeuralBackend, L> DeviceState<B, L> {
609    /// Realizes every layer through a backend-specific construction closure.
610    pub fn create<E>(
611        layout: StateLayout,
612        mut create: impl FnMut(usize, &LayerCachePolicy) -> Result<L, E>,
613    ) -> Result<Self, E> {
614        let layers = layout
615            .layers()
616            .iter()
617            .enumerate()
618            .map(|(layer, policy)| create(layer, policy))
619            .collect::<Result<Vec<_>, _>>()?;
620        Ok(Self {
621            layout,
622            layers,
623            backend: PhantomData,
624        })
625    }
626}
627
628impl<B, L> RuntimeState<B> for DeviceState<B, L>
629where
630    B: NeuralBackend,
631    L: RuntimeLayerState<B>,
632{
633    type RetainedValues<'a>
634        = L::RetainedValues<'a>
635    where
636        Self: 'a,
637        B::Tensor: 'a;
638
639    fn layout(&self) -> &StateLayout {
640        &self.layout
641    }
642
643    fn retained_values(
644        &self,
645        _ordinal: usize,
646        address: crate::ExecutionUnitAddress,
647    ) -> Result<Self::RetainedValues<'_>, StateError> {
648        let layer = address.index();
649        self.layers
650            .get(layer)
651            .map(|layer| layer.retained_values())
652            .ok_or(StateError::UnknownLayer {
653                layer,
654                count: self.layers.len(),
655            })
656    }
657}
658
659impl<B, L> LayerRuntimeState<B> for DeviceState<B, L>
660where
661    B: NeuralBackend,
662    L: RuntimeLayerState<B>,
663{
664    type LayerState = L;
665
666    fn layer(&mut self, layer: usize) -> Result<&mut Self::LayerState, StateError> {
667        let count = self.layers.len();
668        self.layers
669            .get_mut(layer)
670            .ok_or(StateError::UnknownLayer { layer, count })
671    }
672}
673
674impl<B, L> ResettableRuntimeState<B> for DeviceState<B, L>
675where
676    B: NeuralBackend,
677    L: ResettableRuntimeLayerState<B>,
678{
679    fn reset_segment(&mut self, segment: &StateSegmentId) -> Result<(), StateError> {
680        let range = self
681            .layout
682            .segment(segment)
683            .map(StateSegmentSpec::layers)
684            .ok_or_else(|| StateError::UnknownSegment {
685                segment: segment.clone(),
686            })?;
687        for layer in &mut self.layers[range] {
688            layer.reset()?;
689        }
690        Ok(())
691    }
692}
693
694impl<B: NeuralBackend, L> AsRef<[L]> for DeviceState<B, L> {
695    fn as_ref(&self) -> &[L] {
696        &self.layers
697    }
698}
699
700impl<B: NeuralBackend, L> AsMut<[L]> for DeviceState<B, L> {
701    fn as_mut(&mut self) -> &mut [L] {
702        &mut self.layers
703    }
704}
705
706/// Architecture and placement identity used to derive persistence identity.
707#[derive(Debug, Clone, Eq, PartialEq)]
708pub struct ModelStateIdentity {
709    /// Stable architecture family.
710    model_family: String,
711    /// Effective normalized model type.
712    effective_model_type: String,
713    /// Cache-relevant architecture fingerprint.
714    architecture_fingerprint: String,
715    /// Total architecture layer count.
716    layer_count: usize,
717    /// Inclusive first global layer owned by this runtime instance.
718    global_layer_start: usize,
719    /// Attention sink or pinned-prefix token count.
720    sink_tokens: usize,
721    /// Rank-local distributed placement.
722    topology: PromptCacheTopology,
723}
724
725impl ModelStateIdentity {
726    /// Creates a validated architecture and placement identity.
727    #[allow(clippy::too_many_arguments)]
728    pub fn new(
729        model_family: impl Into<String>,
730        effective_model_type: impl Into<String>,
731        architecture_fingerprint: impl Into<String>,
732        layer_count: usize,
733        global_layer_start: usize,
734        sink_tokens: usize,
735        topology: PromptCacheTopology,
736    ) -> Result<Self, PromptCacheError> {
737        let model_family = model_family.into();
738        let effective_model_type = effective_model_type.into();
739        let architecture_fingerprint = architecture_fingerprint.into();
740        if model_family.trim().is_empty()
741            || effective_model_type.trim().is_empty()
742            || architecture_fingerprint.trim().is_empty()
743        {
744            return Err(PromptCacheError::Malformed(
745                "model-state identity strings must be non-empty".into(),
746            ));
747        }
748        if layer_count == 0 || global_layer_start > layer_count {
749            return Err(PromptCacheError::Malformed(format!(
750                "model-state layer start {global_layer_start} is invalid for {layer_count} layers"
751            )));
752        }
753        topology.validate()?;
754        Ok(Self {
755            model_family,
756            effective_model_type,
757            architecture_fingerprint,
758            layer_count,
759            global_layer_start,
760            sink_tokens,
761            topology,
762        })
763    }
764
765    /// Total architecture layer count.
766    pub const fn layer_count(&self) -> usize {
767        self.layer_count
768    }
769
770    /// Inclusive first global layer owned by this runtime instance.
771    pub const fn global_layer_start(&self) -> usize {
772        self.global_layer_start
773    }
774
775    /// Combines architecture identity, placement, and exact state geometry.
776    pub fn prompt_cache_identity(
777        &self,
778        layout: &StateLayout,
779    ) -> Result<PromptCacheModelIdentity, PromptCacheError> {
780        let global_layer_end = self
781            .global_layer_start
782            .checked_add(layout.len())
783            .ok_or_else(|| PromptCacheError::Malformed("owned layer range overflowed".into()))?;
784        PromptCacheModelIdentity::new(
785            self.model_family.clone(),
786            self.effective_model_type.clone(),
787            self.architecture_fingerprint.clone(),
788            self.layer_count,
789            self.global_layer_start,
790            global_layer_end,
791            self.sink_tokens,
792            self.topology.clone(),
793            layout.layers().clone(),
794            layout.layer_prefix_offsets(),
795            layout
796                .segments()
797                .iter()
798                .map(|segment| {
799                    PromptCacheStateSegment::new(segment.id().as_str(), segment.layers())
800                })
801                .collect::<Result<Vec<_>, _>>()?,
802        )
803    }
804}
805
806/// Invalid architecture state geometry or runtime access.
807#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
808#[non_exhaustive]
809pub enum StateError {
810    /// A model declared no state-bearing layer slots.
811    #[error("runtime state layout must contain at least one layer")]
812    EmptyLayout,
813    /// A composite layout declared no named state segments.
814    #[error("runtime state layout must contain at least one named segment")]
815    EmptySegments,
816    /// A state segment identity was empty or whitespace-only.
817    #[error("runtime state segment identity must not be empty")]
818    EmptySegmentId,
819    /// A state segment range contained no layers.
820    #[error("runtime state segment {segment:?} has empty layer range {start}..{end}")]
821    EmptySegmentRange {
822        /// Invalid segment identity.
823        segment: StateSegmentId,
824        /// Inclusive range start.
825        start: usize,
826        /// Exclusive range end.
827        end: usize,
828    },
829    /// A state segment claimed to be ahead of the persisted token prefix.
830    #[error("runtime state segment {segment:?} has positive processed-token offset {offset}")]
831    PositiveSegmentOffset {
832        /// Invalid segment identity.
833        segment: StateSegmentId,
834        /// Invalid positive processed-token offset.
835        offset: i32,
836    },
837    /// A requested sub-layout range was empty or outside the source layout.
838    #[error("runtime state layout cannot select range {start}..{end} from {layers} layers")]
839    InvalidLayoutSlice {
840        /// Inclusive requested layer start.
841        start: usize,
842        /// Exclusive requested layer end.
843        end: usize,
844        /// Available source layer count.
845        layers: usize,
846    },
847    /// Two state segments used the same stable identity.
848    #[error("runtime state segment {segment:?} is declared more than once")]
849    DuplicateSegment {
850        /// Duplicated identity.
851        segment: StateSegmentId,
852    },
853    /// A state segment addressed layers outside the layout.
854    #[error(
855        "runtime state segment {segment:?} range {start}..{end} exceeds {layers} layout layers"
856    )]
857    SegmentOutOfBounds {
858        /// Invalid segment identity.
859        segment: StateSegmentId,
860        /// Inclusive range start.
861        start: usize,
862        /// Exclusive range end.
863        end: usize,
864        /// Total layout layer count.
865        layers: usize,
866    },
867    /// A state segment overlapped an earlier segment in layer order.
868    #[error(
869        "runtime state segment {segment:?} starts at layer {start}, before prior frontier {frontier}"
870    )]
871    OverlappingSegment {
872        /// Overlapping segment identity.
873        segment: StateSegmentId,
874        /// Inclusive range start.
875        start: usize,
876        /// End of the prior segment.
877        frontier: usize,
878    },
879    /// No named segment owned one layer in the state layout.
880    #[error("runtime state layer {layer} is not assigned to a named segment")]
881    UnassignedStateLayer {
882        /// First unassigned layer.
883        layer: usize,
884    },
885    /// A reset requested a segment absent from the realized layout.
886    #[error("runtime state layout has no segment {segment:?}")]
887    UnknownSegment {
888        /// Requested segment identity.
889        segment: StateSegmentId,
890    },
891    /// A concrete backend failed while clearing a declared state segment.
892    #[error("runtime state reset failed: {0}")]
893    ResetFailed(String),
894    /// One layer supplied an invalid portable state policy.
895    #[error("invalid runtime state policy for layer {layer}: {reason}")]
896    InvalidLayer {
897        /// Invalid global layer index.
898        layer: usize,
899        /// Validation detail.
900        reason: String,
901    },
902    /// A residency plan violates a finite-resource invariant.
903    #[error("invalid runtime state residency plan: {0}")]
904    InvalidResidency(String),
905    /// A runtime requested a layer outside the realized layout.
906    #[error("runtime state layer {layer} is outside the {count}-layer layout")]
907    UnknownLayer {
908        /// Requested layer.
909        layer: usize,
910        /// Realized layer count.
911        count: usize,
912    },
913    /// A layer state does not declare the requested fixed component.
914    #[error("runtime state layer does not declare fixed component {role:?}")]
915    UnknownComponent {
916        /// Requested semantic component.
917        role: StateTensorRole,
918    },
919    /// A fixed-state token frontier could not be advanced safely.
920    #[error("invalid fixed-state advance: {0}")]
921    InvalidAdvance(String),
922}
923
924#[cfg(test)]
925mod tests {
926    use super::*;
927    use eredu_core::{AttentionPolicy, LayerSchedule};
928
929    fn layout() -> StateLayout {
930        StateLayout::new(
931            LayerSchedule::new(
932                2,
933                vec![
934                    LayerCachePolicy::key_value(AttentionPolicy::Full, 2, 8).unwrap(),
935                    LayerCachePolicy::key_value(
936                        AttentionPolicy::from_sliding_window(Some(16)).unwrap(),
937                        2,
938                        8,
939                    )
940                    .unwrap(),
941                ],
942            )
943            .unwrap(),
944        )
945        .unwrap()
946    }
947
948    #[test]
949    fn prompt_identity_is_derived_from_layout_and_placement() {
950        let layout = layout();
951        let identity = ModelStateIdentity {
952            model_family: "fixture".into(),
953            effective_model_type: "fixture-v1".into(),
954            architecture_fingerprint: "geometry-1".into(),
955            layer_count: 4,
956            global_layer_start: 1,
957            sink_tokens: 0,
958            topology: PromptCacheTopology::default(),
959        }
960        .prompt_cache_identity(&layout)
961        .unwrap();
962        assert_eq!(identity.global_layer_start(), 1);
963        assert_eq!(identity.global_layer_end(), 3);
964        assert_eq!(identity.layer_layout(), layout.layers());
965    }
966
967    #[test]
968    fn prompt_identity_derives_offsets_from_segments() {
969        let policy = LayerCachePolicy::key_value(AttentionPolicy::Full, 1, 8).unwrap();
970        let layout = StateLayout::segmented(
971            LayerSchedule::new(2, vec![policy.clone(), policy]).unwrap(),
972            [
973                StateSegmentSpec::new("target", 0..1, StateSegmentLifetime::Persistent, 0).unwrap(),
974                StateSegmentSpec::new("prediction", 1..2, StateSegmentLifetime::Persistent, -1)
975                    .unwrap(),
976            ],
977        )
978        .unwrap();
979        let identity = ModelStateIdentity {
980            model_family: "fixture".into(),
981            effective_model_type: "fixture-v1".into(),
982            architecture_fingerprint: "geometry-1".into(),
983            layer_count: 2,
984            global_layer_start: 0,
985            sink_tokens: 0,
986            topology: PromptCacheTopology::default(),
987        }
988        .prompt_cache_identity(&layout)
989        .unwrap();
990        assert_eq!(identity.layer_prefix_offsets(), [0, -1]);
991        assert_eq!(identity.state_segments().len(), 2);
992        assert_eq!(identity.state_segments()[0].id(), "target");
993        assert_eq!(identity.state_segments()[0].layers(), 0..1);
994        assert_eq!(identity.state_segments()[1].id(), "prediction");
995        assert_eq!(identity.state_segments()[1].layers(), 1..2);
996
997        let prediction = identity.select_state_segment("prediction").unwrap();
998        assert_eq!(prediction.global_layer_start(), 1);
999        assert_eq!(prediction.global_layer_end(), 2);
1000        assert_eq!(prediction.layer_prefix_offsets(), [-1]);
1001        assert_eq!(prediction.state_segments()[0].id(), "prediction");
1002        assert_eq!(prediction.state_segments()[0].layers(), 0..1);
1003    }
1004
1005    #[test]
1006    fn state_layout_exposes_stable_semantic_component_names() {
1007        let layout = StateLayout::new(
1008            LayerSchedule::new(
1009                1,
1010                vec![
1011                    LayerCachePolicy::compressed_latent_rotary(AttentionPolicy::Full, 16, 8)
1012                        .unwrap(),
1013                ],
1014            )
1015            .unwrap(),
1016        )
1017        .unwrap();
1018        let names = layout
1019            .components(0)
1020            .unwrap()
1021            .iter()
1022            .map(|component| component.role().stable_name())
1023            .collect::<Vec<_>>();
1024        assert_eq!(
1025            names,
1026            ["attention.compressed_latent", "attention.rotary_keys"]
1027        );
1028    }
1029
1030    fn four_layer_schedule() -> LayerSchedule<LayerCachePolicy> {
1031        LayerSchedule::new(
1032            4,
1033            (0..4)
1034                .map(|_| LayerCachePolicy::key_value(AttentionPolicy::Full, 1, 8).unwrap())
1035                .collect(),
1036        )
1037        .unwrap()
1038    }
1039
1040    #[test]
1041    fn composite_state_segments_are_canonical_and_cover_every_layer() {
1042        let layout = StateLayout::segmented(
1043            four_layer_schedule(),
1044            [
1045                StateSegmentSpec::new("depth", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
1046                StateSegmentSpec::new("temporal", 0..2, StateSegmentLifetime::Persistent, 0)
1047                    .unwrap(),
1048            ],
1049        )
1050        .unwrap();
1051
1052        assert_eq!(
1053            layout
1054                .segments()
1055                .iter()
1056                .map(|segment| (segment.id().as_str(), segment.layers(), segment.lifetime()))
1057                .collect::<Vec<_>>(),
1058            [
1059                ("temporal", 0..2, StateSegmentLifetime::Persistent),
1060                ("depth", 2..4, StateSegmentLifetime::FrameLocal),
1061            ]
1062        );
1063        assert_eq!(
1064            layout.segment_for_layer(0).unwrap().id().as_str(),
1065            "temporal"
1066        );
1067        assert_eq!(layout.segment_for_layer(3).unwrap().id().as_str(), "depth");
1068        assert!(layout.segment_for_layer(4).is_none());
1069    }
1070
1071    #[test]
1072    fn state_layout_slice_preserves_and_rebases_segment_frontiers() {
1073        let layout = StateLayout::segmented(
1074            four_layer_schedule(),
1075            [
1076                StateSegmentSpec::new("target", 0..2, StateSegmentLifetime::Persistent, 0).unwrap(),
1077                StateSegmentSpec::new("prediction", 2..4, StateSegmentLifetime::Persistent, -1)
1078                    .unwrap(),
1079            ],
1080        )
1081        .unwrap();
1082
1083        let sliced = layout.slice(1..4).unwrap();
1084        assert_eq!(sliced.segments()[0].layers(), 0..1);
1085        assert_eq!(sliced.segments()[1].layers(), 1..3);
1086        assert_eq!(sliced.layer_prefix_offsets(), [0, -1, -1]);
1087    }
1088
1089    #[test]
1090    fn segment_identity_lifetime_and_offset_participate_in_layout_equality() {
1091        let layout = |depth_name, lifetime, offset| {
1092            StateLayout::segmented(
1093                four_layer_schedule(),
1094                [
1095                    StateSegmentSpec::new("temporal", 0..2, StateSegmentLifetime::Persistent, 0)
1096                        .unwrap(),
1097                    StateSegmentSpec::new(depth_name, 2..4, lifetime, offset).unwrap(),
1098                ],
1099            )
1100            .unwrap()
1101        };
1102        let canonical = layout("depth", StateSegmentLifetime::FrameLocal, 0);
1103        assert_ne!(
1104            canonical,
1105            layout("predictor", StateSegmentLifetime::FrameLocal, 0)
1106        );
1107        assert_ne!(
1108            canonical,
1109            layout("depth", StateSegmentLifetime::Persistent, 0)
1110        );
1111        assert_ne!(
1112            canonical,
1113            layout("depth", StateSegmentLifetime::FrameLocal, -1)
1114        );
1115    }
1116
1117    #[test]
1118    fn malformed_segment_partitions_fail_closed() {
1119        let duplicate = StateLayout::segmented(
1120            four_layer_schedule(),
1121            [
1122                StateSegmentSpec::new("cache", 0..2, StateSegmentLifetime::Persistent, 0).unwrap(),
1123                StateSegmentSpec::new("cache", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
1124            ],
1125        )
1126        .unwrap_err();
1127        assert!(matches!(duplicate, StateError::DuplicateSegment { .. }));
1128
1129        let overlap = StateLayout::segmented(
1130            four_layer_schedule(),
1131            [
1132                StateSegmentSpec::new("left", 0..3, StateSegmentLifetime::Persistent, 0).unwrap(),
1133                StateSegmentSpec::new("right", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
1134            ],
1135        )
1136        .unwrap_err();
1137        assert!(matches!(overlap, StateError::OverlappingSegment { .. }));
1138
1139        let gap = StateLayout::segmented(
1140            four_layer_schedule(),
1141            [
1142                StateSegmentSpec::new("left", 0..1, StateSegmentLifetime::Persistent, 0).unwrap(),
1143                StateSegmentSpec::new("right", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
1144            ],
1145        )
1146        .unwrap_err();
1147        assert_eq!(gap, StateError::UnassignedStateLayer { layer: 1 });
1148
1149        let outside = StateLayout::segmented(
1150            four_layer_schedule(),
1151            [StateSegmentSpec::new("all", 0..5, StateSegmentLifetime::Persistent, 0).unwrap()],
1152        )
1153        .unwrap_err();
1154        assert!(matches!(outside, StateError::SegmentOutOfBounds { .. }));
1155    }
1156}