Skip to main content

proof_engine/editor/
animation_state_machine.rs

1
2//! Animation state machine editor — states, transitions, blend trees, parameters.
3
4use glam::{Vec2, Vec3};
5use std::collections::HashMap;
6
7// ---------------------------------------------------------------------------
8// Parameters
9// ---------------------------------------------------------------------------
10
11#[derive(Debug, Clone)]
12pub enum AnimParam {
13    Float(f32),
14    Int(i32),
15    Bool(bool),
16    Trigger(bool),
17}
18
19impl AnimParam {
20    pub fn as_float(&self) -> f32 {
21        match self { AnimParam::Float(v) => *v, AnimParam::Int(v) => *v as f32, AnimParam::Bool(v) => if *v { 1.0 } else { 0.0 }, AnimParam::Trigger(v) => if *v { 1.0 } else { 0.0 } }
22    }
23    pub fn as_bool(&self) -> bool {
24        match self { AnimParam::Bool(v) | AnimParam::Trigger(v) => *v, AnimParam::Float(v) => *v != 0.0, AnimParam::Int(v) => *v != 0 }
25    }
26    pub fn type_label(&self) -> &'static str {
27        match self { AnimParam::Float(_) => "Float", AnimParam::Int(_) => "Int", AnimParam::Bool(_) => "Bool", AnimParam::Trigger(_) => "Trigger" }
28    }
29}
30
31// ---------------------------------------------------------------------------
32// Transition conditions
33// ---------------------------------------------------------------------------
34
35#[derive(Debug, Clone, Copy, PartialEq)]
36pub enum ConditionOp { Greater, Less, Equal, NotEqual, True, False }
37
38impl ConditionOp {
39    pub fn label(self) -> &'static str {
40        match self { ConditionOp::Greater => ">", ConditionOp::Less => "<", ConditionOp::Equal => "==", ConditionOp::NotEqual => "!=", ConditionOp::True => "is true", ConditionOp::False => "is false" }
41    }
42    pub fn evaluate(&self, lhs: f32, rhs: f32) -> bool {
43        match self { ConditionOp::Greater => lhs > rhs, ConditionOp::Less => lhs < rhs, ConditionOp::Equal => (lhs - rhs).abs() < 1e-4, ConditionOp::NotEqual => (lhs - rhs).abs() >= 1e-4, ConditionOp::True => lhs != 0.0, ConditionOp::False => lhs == 0.0 }
44    }
45}
46
47#[derive(Debug, Clone)]
48pub struct TransitionCondition {
49    pub parameter: String,
50    pub op: ConditionOp,
51    pub threshold: f32,
52}
53
54impl TransitionCondition {
55    pub fn new(param: impl Into<String>, op: ConditionOp, threshold: f32) -> Self {
56        Self { parameter: param.into(), op, threshold }
57    }
58
59    pub fn evaluate(&self, params: &HashMap<String, AnimParam>) -> bool {
60        let val = params.get(&self.parameter).map(|p| p.as_float()).unwrap_or(0.0);
61        self.op.evaluate(val, self.threshold)
62    }
63}
64
65// ---------------------------------------------------------------------------
66// Transition
67// ---------------------------------------------------------------------------
68
69#[derive(Debug, Clone, Copy, PartialEq)]
70pub enum InterruptionSource { None, Source, Destination, SourceThenDestination, DestinationThenSource }
71
72#[derive(Debug, Clone)]
73pub struct StateTransition {
74    pub id: u32,
75    pub from_state: u32,
76    pub to_state: u32,
77    pub conditions: Vec<TransitionCondition>,
78    pub duration: f32,
79    pub offset: f32,
80    pub has_exit_time: bool,
81    pub exit_time: f32,
82    pub can_transition_to_self: bool,
83    pub interruption_source: InterruptionSource,
84    pub ordered_interruption: bool,
85    pub blend_curve: BlendCurveType,
86    pub mute: bool,
87    pub solo: bool,
88}
89
90#[derive(Debug, Clone, Copy, PartialEq)]
91pub enum BlendCurveType { Linear, Fixed, EaseIn, EaseOut, EaseInOut, Custom }
92
93impl StateTransition {
94    pub fn new(id: u32, from: u32, to: u32) -> Self {
95        Self {
96            id, from_state: from, to_state: to,
97            conditions: Vec::new(),
98            duration: 0.25, offset: 0.0,
99            has_exit_time: false, exit_time: 0.75,
100            can_transition_to_self: false,
101            interruption_source: InterruptionSource::None,
102            ordered_interruption: true,
103            blend_curve: BlendCurveType::Linear,
104            mute: false, solo: false,
105        }
106    }
107
108    pub fn can_trigger(&self, params: &HashMap<String, AnimParam>, normalized_time: f32) -> bool {
109        if self.mute { return false; }
110        if self.has_exit_time && normalized_time < self.exit_time { return false; }
111        if self.conditions.is_empty() { return self.has_exit_time && normalized_time >= self.exit_time; }
112        self.conditions.iter().all(|c| c.evaluate(params))
113    }
114
115    pub fn blend_weight(&self, t: f32) -> f32 {
116        let t = t.clamp(0.0, 1.0);
117        match self.blend_curve {
118            BlendCurveType::Linear => t,
119            BlendCurveType::Fixed => if t > 0.5 { 1.0 } else { 0.0 },
120            BlendCurveType::EaseIn => t * t,
121            BlendCurveType::EaseOut => 1.0 - (1.0 - t) * (1.0 - t),
122            BlendCurveType::EaseInOut => t * t * (3.0 - 2.0 * t),
123            BlendCurveType::Custom => t,
124        }
125    }
126}
127
128// ---------------------------------------------------------------------------
129// Blend tree
130// ---------------------------------------------------------------------------
131
132#[derive(Debug, Clone, Copy, PartialEq)]
133pub enum BlendTreeType { Simple1D, SimpleDirectional2D, FreeformDirectional2D, FreeformCartesian2D, Direct }
134
135#[derive(Debug, Clone)]
136pub struct BlendTreeChild {
137    pub clip_name: String,
138    pub threshold: f32,
139    pub position_2d: Vec2,
140    pub direct_blend_param: String,
141    pub speed: f32,
142    pub mirror: bool,
143    pub time_scale: f32,
144}
145
146impl BlendTreeChild {
147    pub fn new_1d(clip: impl Into<String>, threshold: f32) -> Self {
148        Self {
149            clip_name: clip.into(), threshold,
150            position_2d: Vec2::ZERO, direct_blend_param: String::new(),
151            speed: 1.0, mirror: false, time_scale: 1.0,
152        }
153    }
154    pub fn new_2d(clip: impl Into<String>, pos: Vec2) -> Self {
155        Self {
156            clip_name: clip.into(), threshold: 0.0,
157            position_2d: pos, direct_blend_param: String::new(),
158            speed: 1.0, mirror: false, time_scale: 1.0,
159        }
160    }
161}
162
163#[derive(Debug, Clone)]
164pub struct BlendTree {
165    pub name: String,
166    pub blend_type: BlendTreeType,
167    pub blend_param_x: String,
168    pub blend_param_y: String,
169    pub children: Vec<BlendTreeChild>,
170    pub use_auto_thresholds: bool,
171    pub compute_threshold_automatically: bool,
172}
173
174impl BlendTree {
175    pub fn new_1d(name: impl Into<String>, param: impl Into<String>) -> Self {
176        Self {
177            name: name.into(),
178            blend_type: BlendTreeType::Simple1D,
179            blend_param_x: param.into(),
180            blend_param_y: String::new(),
181            children: Vec::new(),
182            use_auto_thresholds: true,
183            compute_threshold_automatically: true,
184        }
185    }
186
187    pub fn new_2d(name: impl Into<String>, param_x: impl Into<String>, param_y: impl Into<String>) -> Self {
188        Self {
189            name: name.into(),
190            blend_type: BlendTreeType::SimpleDirectional2D,
191            blend_param_x: param_x.into(),
192            blend_param_y: param_y.into(),
193            children: Vec::new(),
194            use_auto_thresholds: false,
195            compute_threshold_automatically: false,
196        }
197    }
198
199    pub fn add_child(&mut self, child: BlendTreeChild) {
200        self.children.push(child);
201    }
202
203    /// Compute weights for 1D blend.
204    pub fn compute_weights_1d(&self, value: f32) -> Vec<f32> {
205        if self.children.is_empty() { return Vec::new(); }
206        let n = self.children.len();
207        let mut weights = vec![0.0_f32; n];
208        let thresholds: Vec<f32> = self.children.iter().map(|c| c.threshold).collect();
209        let clamped = value.clamp(thresholds[0], thresholds[n - 1]);
210        let i = thresholds.partition_point(|&t| t <= clamped);
211        if i == 0 { weights[0] = 1.0; }
212        else if i >= n { weights[n-1] = 1.0; }
213        else {
214            let t0 = thresholds[i-1];
215            let t1 = thresholds[i];
216            let u = (clamped - t0) / (t1 - t0).max(1e-6);
217            weights[i-1] = 1.0 - u;
218            weights[i] = u;
219        }
220        weights
221    }
222
223    /// Compute weights for 2D directional blend.
224    pub fn compute_weights_2d(&self, x: f32, y: f32) -> Vec<f32> {
225        if self.children.is_empty() { return Vec::new(); }
226        let n = self.children.len();
227        let query = Vec2::new(x, y);
228        // Simple nearest-neighbor with distance weighting
229        let inv_dists: Vec<f32> = self.children.iter()
230            .map(|c| 1.0 / (c.position_2d.distance(query) + 0.001))
231            .collect();
232        let total: f32 = inv_dists.iter().sum();
233        if total < 1e-6 { return vec![1.0 / n as f32; n]; }
234        inv_dists.iter().map(|&d| d / total).collect()
235    }
236}
237
238// ---------------------------------------------------------------------------
239// State
240// ---------------------------------------------------------------------------
241
242#[derive(Debug, Clone, Copy, PartialEq)]
243pub enum StateKind {
244    Normal,
245    BlendTree,
246    Empty,
247    Any,
248    Entry,
249    Exit,
250}
251
252#[derive(Debug, Clone)]
253pub struct AnimState {
254    pub id: u32,
255    pub name: String,
256    pub kind: StateKind,
257    pub position: Vec2,
258    pub clip_name: Option<String>,
259    pub blend_tree: Option<BlendTree>,
260    pub speed: f32,
261    pub speed_multiplier_param: Option<String>,
262    pub mirror: bool,
263    pub mirror_param: Option<String>,
264    pub foot_ik: bool,
265    pub write_defaults: bool,
266    pub loop_time: bool,
267    pub cyclic_offset: f32,
268    pub tag: String,
269    pub motion_time_param: Option<String>,
270    pub color: [u8; 3],
271}
272
273impl AnimState {
274    pub fn new(id: u32, name: impl Into<String>) -> Self {
275        Self {
276            id, name: name.into(), kind: StateKind::Normal,
277            position: Vec2::ZERO,
278            clip_name: None, blend_tree: None,
279            speed: 1.0, speed_multiplier_param: None,
280            mirror: false, mirror_param: None,
281            foot_ik: false, write_defaults: true,
282            loop_time: true, cyclic_offset: 0.0,
283            tag: String::new(), motion_time_param: None,
284            color: [150, 150, 150],
285        }
286    }
287
288    pub fn with_clip(mut self, clip: impl Into<String>) -> Self {
289        self.clip_name = Some(clip.into());
290        self
291    }
292
293    pub fn with_blend_tree(mut self, tree: BlendTree) -> Self {
294        self.blend_tree = Some(tree);
295        self.kind = StateKind::BlendTree;
296        self
297    }
298
299    pub fn effective_speed(&self, params: &HashMap<String, AnimParam>) -> f32 {
300        let mult = self.speed_multiplier_param.as_ref()
301            .and_then(|p| params.get(p))
302            .map(|p| p.as_float())
303            .unwrap_or(1.0);
304        self.speed * mult
305    }
306}
307
308// ---------------------------------------------------------------------------
309// Layer
310// ---------------------------------------------------------------------------
311
312#[derive(Debug, Clone, Copy, PartialEq)]
313pub enum AvatarMask { None, Body, UpperBody, LowerBody, LeftHand, RightHand, LeftFoot, RightFoot, Custom }
314#[derive(Debug, Clone, Copy, PartialEq)]
315pub enum LayerBlendType { Override, Additive }
316
317#[derive(Debug, Clone)]
318pub struct AnimLayer {
319    pub name: String,
320    pub weight: f32,
321    pub blend_type: LayerBlendType,
322    pub avatar_mask: AvatarMask,
323    pub sync: bool,
324    pub sync_layer: Option<usize>,
325    pub ik_pass: bool,
326    pub states: Vec<AnimState>,
327    pub transitions: Vec<StateTransition>,
328    pub default_state: u32,
329    pub current_state: u32,
330    pub next_state: Option<u32>,
331    pub transition_progress: f32,
332    pub current_time: f32,
333}
334
335impl AnimLayer {
336    pub fn new(name: impl Into<String>) -> Self {
337        Self {
338            name: name.into(),
339            weight: 1.0,
340            blend_type: LayerBlendType::Override,
341            avatar_mask: AvatarMask::None,
342            sync: false,
343            sync_layer: None,
344            ik_pass: false,
345            states: Vec::new(),
346            transitions: Vec::new(),
347            default_state: 0,
348            current_state: 0,
349            next_state: None,
350            transition_progress: 0.0,
351            current_time: 0.0,
352        }
353    }
354
355    pub fn add_state(&mut self, state: AnimState) -> u32 {
356        let id = state.id;
357        if self.states.is_empty() { self.default_state = id; self.current_state = id; }
358        self.states.push(state);
359        id
360    }
361
362    pub fn add_transition(&mut self, t: StateTransition) { self.transitions.push(t); }
363
364    pub fn current_state(&self) -> Option<&AnimState> {
365        self.states.iter().find(|s| s.id == self.current_state)
366    }
367
368    pub fn update(&mut self, dt: f32, params: &HashMap<String, AnimParam>) {
369        let speed = self.current_state()
370            .map(|s| s.effective_speed(params))
371            .unwrap_or(1.0);
372        self.current_time += dt * speed;
373
374        // Check transitions
375        if self.next_state.is_none() {
376            let outgoing: Vec<StateTransition> = self.transitions.iter()
377                .filter(|t| t.from_state == self.current_state)
378                .cloned()
379                .collect();
380            for t in &outgoing {
381                if t.can_trigger(params, self.current_time) {
382                    self.next_state = Some(t.to_state);
383                    self.transition_progress = 0.0;
384                    // Consume triggers
385                    break;
386                }
387            }
388        }
389
390        if let Some(next) = self.next_state {
391            let dur = self.transitions.iter()
392                .find(|t| t.from_state == self.current_state && t.to_state == next)
393                .map(|t| t.duration)
394                .unwrap_or(0.25);
395            self.transition_progress += dt / dur.max(0.001);
396            if self.transition_progress >= 1.0 {
397                self.current_state = next;
398                self.next_state = None;
399                self.transition_progress = 0.0;
400                self.current_time = 0.0;
401            }
402        }
403    }
404}
405
406// ---------------------------------------------------------------------------
407// Animator controller
408// ---------------------------------------------------------------------------
409
410#[derive(Debug, Clone)]
411pub struct AnimatorController {
412    pub name: String,
413    pub layers: Vec<AnimLayer>,
414    pub parameters: HashMap<String, AnimParam>,
415    pub default_values: HashMap<String, AnimParam>,
416}
417
418impl AnimatorController {
419    pub fn new(name: impl Into<String>) -> Self {
420        Self {
421            name: name.into(),
422            layers: Vec::new(),
423            parameters: HashMap::new(),
424            default_values: HashMap::new(),
425        }
426    }
427
428    pub fn add_layer(&mut self, layer: AnimLayer) {
429        self.layers.push(layer);
430    }
431
432    pub fn add_param(&mut self, name: impl Into<String>, value: AnimParam) {
433        let name = name.into();
434        self.default_values.insert(name.clone(), value.clone());
435        self.parameters.insert(name, value);
436    }
437
438    pub fn set_float(&mut self, name: &str, value: f32) {
439        if let Some(p) = self.parameters.get_mut(name) {
440            *p = AnimParam::Float(value);
441        }
442    }
443    pub fn set_int(&mut self, name: &str, value: i32) {
444        if let Some(p) = self.parameters.get_mut(name) {
445            *p = AnimParam::Int(value);
446        }
447    }
448    pub fn set_bool(&mut self, name: &str, value: bool) {
449        if let Some(p) = self.parameters.get_mut(name) {
450            *p = AnimParam::Bool(value);
451        }
452    }
453    pub fn set_trigger(&mut self, name: &str) {
454        if let Some(p) = self.parameters.get_mut(name) {
455            *p = AnimParam::Trigger(true);
456        }
457    }
458    pub fn reset_trigger(&mut self, name: &str) {
459        if let Some(p) = self.parameters.get_mut(name) {
460            if matches!(p, AnimParam::Trigger(_)) {
461                *p = AnimParam::Trigger(false);
462            }
463        }
464    }
465
466    pub fn update(&mut self, dt: f32) {
467        for layer in &mut self.layers {
468            layer.update(dt, &self.parameters);
469        }
470        // Reset triggers after evaluation
471        for p in self.parameters.values_mut() {
472            if matches!(p, AnimParam::Trigger(true)) {
473                *p = AnimParam::Trigger(false);
474            }
475        }
476    }
477
478    pub fn reset_to_defaults(&mut self) {
479        for (k, v) in &self.default_values {
480            self.parameters.insert(k.clone(), v.clone());
481        }
482        for layer in &mut self.layers {
483            layer.current_state = layer.default_state;
484            layer.next_state = None;
485            layer.current_time = 0.0;
486        }
487    }
488
489    /// Build a sample biped controller.
490    pub fn build_biped_controller() -> Self {
491        let mut ctrl = AnimatorController::new("BipedController");
492        ctrl.add_param("Speed", AnimParam::Float(0.0));
493        ctrl.add_param("Direction", AnimParam::Float(0.0));
494        ctrl.add_param("IsGrounded", AnimParam::Bool(true));
495        ctrl.add_param("Jump", AnimParam::Trigger(false));
496        ctrl.add_param("IsAiming", AnimParam::Bool(false));
497        ctrl.add_param("AttackIndex", AnimParam::Int(0));
498
499        let mut base_layer = AnimLayer::new("Base Layer");
500        let mut id = 1u32;
501
502        // Locomotion blend tree
503        let mut loco_tree = BlendTree::new_2d("Locomotion", "Speed", "Direction");
504        loco_tree.add_child(BlendTreeChild::new_2d("Idle", Vec2::ZERO));
505        loco_tree.add_child(BlendTreeChild::new_2d("Walk_Forward", Vec2::new(0.0, 0.5)));
506        loco_tree.add_child(BlendTreeChild::new_2d("Run_Forward", Vec2::new(0.0, 1.0)));
507        loco_tree.add_child(BlendTreeChild::new_2d("Walk_Left", Vec2::new(-0.5, 0.5)));
508        loco_tree.add_child(BlendTreeChild::new_2d("Walk_Right", Vec2::new(0.5, 0.5)));
509        loco_tree.add_child(BlendTreeChild::new_2d("Run_Left", Vec2::new(-1.0, 1.0)));
510        loco_tree.add_child(BlendTreeChild::new_2d("Run_Right", Vec2::new(1.0, 1.0)));
511
512        let loco = AnimState::new(id, "Locomotion")
513            .with_blend_tree(loco_tree);
514        id += 1;
515        let loco_id = base_layer.add_state(loco);
516
517        let jump = AnimState::new(id, "Jump").with_clip("Jump");
518        id += 1;
519        let jump_id = base_layer.add_state(jump);
520
521        let fall = AnimState::new(id, "Fall").with_clip("Fall");
522        id += 1;
523        let fall_id = base_layer.add_state(fall);
524
525        let land = AnimState::new(id, "Land").with_clip("Land");
526        id += 1;
527        let land_id = base_layer.add_state(land);
528
529        // Transitions
530        let mut t1 = StateTransition::new(id, loco_id, jump_id);
531        t1.conditions.push(TransitionCondition::new("Jump", ConditionOp::True, 1.0));
532        t1.duration = 0.1;
533        base_layer.add_transition(t1);
534        id += 1;
535
536        let mut t2 = StateTransition::new(id, jump_id, fall_id);
537        t2.has_exit_time = true;
538        t2.exit_time = 0.5;
539        t2.duration = 0.1;
540        base_layer.add_transition(t2);
541        id += 1;
542
543        let mut t3 = StateTransition::new(id, fall_id, land_id);
544        t3.conditions.push(TransitionCondition::new("IsGrounded", ConditionOp::True, 1.0));
545        t3.duration = 0.1;
546        base_layer.add_transition(t3);
547        id += 1;
548
549        let mut t4 = StateTransition::new(id, land_id, loco_id);
550        t4.has_exit_time = true;
551        t4.exit_time = 0.9;
552        t4.duration = 0.15;
553        base_layer.add_transition(t4);
554
555        ctrl.add_layer(base_layer);
556        ctrl
557    }
558}
559
560// ---------------------------------------------------------------------------
561// State machine editor
562// ---------------------------------------------------------------------------
563
564#[derive(Debug, Clone, Copy, PartialEq)]
565pub enum StateMachineEditorTool { Select, AddState, AddTransition, Pan }
566
567#[derive(Debug, Clone)]
568pub struct StateMachineEditor {
569    pub controller: AnimatorController,
570    pub active_layer: usize,
571    pub selected_state: Option<u32>,
572    pub selected_transition: Option<u32>,
573    pub tool: StateMachineEditorTool,
574    pub zoom: f32,
575    pub pan: Vec2,
576    pub pending_transition_from: Option<u32>,
577    pub show_parameters: bool,
578    pub show_grid: bool,
579    pub preview_dt: f32,
580    pub preview_playing: bool,
581    pub state_name_input: String,
582    pub search_query: String,
583}
584
585impl StateMachineEditor {
586    pub fn new() -> Self {
587        let controller = AnimatorController::build_biped_controller();
588        Self {
589            controller,
590            active_layer: 0,
591            selected_state: None,
592            selected_transition: None,
593            tool: StateMachineEditorTool::Select,
594            zoom: 1.0,
595            pan: Vec2::ZERO,
596            pending_transition_from: None,
597            show_parameters: true,
598            show_grid: true,
599            preview_dt: 0.016,
600            preview_playing: false,
601            state_name_input: String::new(),
602            search_query: String::new(),
603        }
604    }
605
606    pub fn active_layer(&self) -> Option<&AnimLayer> {
607        self.controller.layers.get(self.active_layer)
608    }
609
610    pub fn update(&mut self, dt: f32) {
611        if self.preview_playing {
612            self.controller.update(dt);
613        }
614    }
615
616    pub fn active_state_name(&self) -> Option<&str> {
617        self.active_layer()
618            .and_then(|l| l.current_state())
619            .map(|s| s.name.as_str())
620    }
621
622    pub fn parameter_count(&self) -> usize {
623        self.controller.parameters.len()
624    }
625}
626
627// ---------------------------------------------------------------------------
628// Tests
629// ---------------------------------------------------------------------------
630#[cfg(test)]
631mod tests {
632    use super::*;
633
634    #[test]
635    fn test_blend_tree_1d() {
636        let mut tree = BlendTree::new_1d("Speed", "speed");
637        tree.add_child(BlendTreeChild::new_1d("Idle", 0.0));
638        tree.add_child(BlendTreeChild::new_1d("Walk", 0.5));
639        tree.add_child(BlendTreeChild::new_1d("Run", 1.0));
640        let w = tree.compute_weights_1d(0.25);
641        assert_eq!(w.len(), 3);
642        assert!((w[0] + w[1] + w[2] - 1.0).abs() < 0.001);
643    }
644
645    #[test]
646    fn test_transition_condition() {
647        let mut params = HashMap::new();
648        params.insert("Speed".to_string(), AnimParam::Float(5.0));
649        let cond = TransitionCondition::new("Speed", ConditionOp::Greater, 3.0);
650        assert!(cond.evaluate(&params));
651        let cond2 = TransitionCondition::new("Speed", ConditionOp::Less, 3.0);
652        assert!(!cond2.evaluate(&params));
653    }
654
655    #[test]
656    fn test_controller() {
657        let mut ctrl = AnimatorController::build_biped_controller();
658        assert!(!ctrl.layers.is_empty());
659        ctrl.set_float("Speed", 1.0);
660        ctrl.update(0.016);
661    }
662
663    #[test]
664    fn test_editor() {
665        let mut ed = StateMachineEditor::new();
666        assert!(ed.active_layer().is_some());
667        ed.update(0.016);
668    }
669}