1use crate::drift::detector::DriftLevel;
9
10#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16#[non_exhaustive]
17pub enum DriftAction {
18 #[default]
22 NotifyOnly,
23 ReduceConfidence,
27 ResetModel,
31 ResetPreprocessor,
35 ReplaceWithBaseline,
38 IncreaseAdaptationRate,
42}
43
44impl DriftAction {
45 pub const fn as_str(&self) -> &'static str {
51 match self {
52 DriftAction::NotifyOnly => "notify_only",
53 DriftAction::ReduceConfidence => "reduce_confidence",
54 DriftAction::ResetModel => "reset_model",
55 DriftAction::ResetPreprocessor => "reset_preprocessor",
56 DriftAction::ReplaceWithBaseline => "replace_with_baseline",
57 DriftAction::IncreaseAdaptationRate => "increase_adaptation_rate",
58 }
59 }
60
61 pub const fn is_destructive(self) -> bool {
66 matches!(
67 self,
68 DriftAction::ResetModel
69 | DriftAction::ResetPreprocessor
70 | DriftAction::ReplaceWithBaseline
71 | DriftAction::IncreaseAdaptationRate
72 )
73 }
74}
75
76#[derive(Debug, Clone)]
81#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
82pub struct DriftEvent {
83 pub sample_index: u64,
85 pub level: DriftLevel,
87 pub action: DriftAction,
89 pub detector_value: f64,
92}
93
94impl DriftEvent {
95 pub const fn new(
97 sample_index: u64,
98 level: DriftLevel,
99 action: DriftAction,
100 detector_value: f64,
101 ) -> Self {
102 Self {
103 sample_index,
104 level,
105 action,
106 detector_value,
107 }
108 }
109}
110
111#[cfg(test)]
112mod tests {
113 use super::*;
114
115 #[test]
116 fn action_as_str() {
117 assert_eq!(DriftAction::NotifyOnly.as_str(), "notify_only");
118 assert_eq!(DriftAction::ReduceConfidence.as_str(), "reduce_confidence");
119 assert_eq!(DriftAction::ResetModel.as_str(), "reset_model");
120 assert_eq!(
121 DriftAction::ResetPreprocessor.as_str(),
122 "reset_preprocessor"
123 );
124 assert_eq!(
125 DriftAction::ReplaceWithBaseline.as_str(),
126 "replace_with_baseline"
127 );
128 assert_eq!(
129 DriftAction::IncreaseAdaptationRate.as_str(),
130 "increase_adaptation_rate"
131 );
132 }
133
134 #[test]
135 fn action_is_destructive() {
136 assert!(!DriftAction::NotifyOnly.is_destructive());
137 assert!(!DriftAction::ReduceConfidence.is_destructive());
138 assert!(DriftAction::ResetModel.is_destructive());
139 assert!(DriftAction::ResetPreprocessor.is_destructive());
140 assert!(DriftAction::ReplaceWithBaseline.is_destructive());
141 assert!(DriftAction::IncreaseAdaptationRate.is_destructive());
142 }
143
144 #[test]
145 fn action_default_is_notify_only() {
146 assert_eq!(DriftAction::default(), DriftAction::NotifyOnly);
147 }
148
149 #[test]
150 fn drift_event_construction() {
151 let event = DriftEvent::new(42, DriftLevel::Drift, DriftAction::ResetModel, 1.23);
152 assert_eq!(event.sample_index, 42);
153 assert_eq!(event.level, DriftLevel::Drift);
154 assert_eq!(event.action, DriftAction::ResetModel);
155 assert!((event.detector_value - 1.23).abs() < 1e-12);
156 }
157
158 #[cfg(feature = "serde")]
159 #[test]
160 fn action_serde_roundtrip() {
161 for action in [
162 DriftAction::NotifyOnly,
163 DriftAction::ReduceConfidence,
164 DriftAction::ResetModel,
165 DriftAction::ResetPreprocessor,
166 DriftAction::ReplaceWithBaseline,
167 DriftAction::IncreaseAdaptationRate,
168 ] {
169 let json = serde_json::to_string(&action).unwrap();
170 let restored: DriftAction = serde_json::from_str(&json).unwrap();
171 assert_eq!(restored, action);
172 }
173 }
174
175 #[cfg(feature = "serde")]
176 #[test]
177 fn drift_event_serde_roundtrip() {
178 let event = DriftEvent::new(10, DriftLevel::Warning, DriftAction::ReduceConfidence, 1.5);
179 let json = serde_json::to_string(&event).unwrap();
180 let restored: DriftEvent = serde_json::from_str(&json).unwrap();
181 assert_eq!(restored.sample_index, 10);
182 assert_eq!(restored.level, DriftLevel::Warning);
183 assert_eq!(restored.action, DriftAction::ReduceConfidence);
184 assert!((restored.detector_value - 1.5).abs() < 1e-12);
185 }
186}