Skip to main content

shap_rs/
metadata.rs

1use crate::{Result, ShapError};
2use serde::{Deserialize, Serialize};
3#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
4pub enum FeatureKind {
5    #[default]
6    Continuous,
7    Categorical,
8    Ordinal,
9    Boolean,
10    TextToken,
11    ImageValue,
12}
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
14pub enum OutputKind {
15    #[default]
16    Regression,
17    Probability,
18    LogOdds,
19    ClassScore,
20    Embedding,
21}
22
23#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
24#[serde(try_from = "FeatureMetadataPayload")]
25pub struct FeatureMetadata {
26    pub names: Vec<String>,
27    pub display_names: Option<Vec<String>>,
28    pub kinds: Option<Vec<FeatureKind>>,
29    pub units: Option<Vec<Option<String>>>,
30}
31#[derive(Deserialize)]
32struct FeatureMetadataPayload {
33    names: Vec<String>,
34    display_names: Option<Vec<String>>,
35    kinds: Option<Vec<FeatureKind>>,
36    units: Option<Vec<Option<String>>>,
37}
38impl TryFrom<FeatureMetadataPayload> for FeatureMetadata {
39    type Error = ShapError;
40    fn try_from(payload: FeatureMetadataPayload) -> Result<Self> {
41        let metadata = Self {
42            names: payload.names,
43            display_names: payload.display_names,
44            kinds: payload.kinds,
45            units: payload.units,
46        };
47        metadata.validate()?;
48        Ok(metadata)
49    }
50}
51impl FeatureMetadata {
52    pub fn new(names: Vec<String>) -> Result<Self> {
53        if names.iter().any(|n| n.is_empty()) {
54            return Err(ShapError::InvalidConfiguration(
55                "feature names cannot be empty".into(),
56            ));
57        }
58        let mut unique = names.clone();
59        unique.sort();
60        unique.dedup();
61        if unique.len() != names.len() {
62            return Err(ShapError::InvalidConfiguration(
63                "feature names must be unique".into(),
64            ));
65        }
66        let metadata = Self {
67            names,
68            display_names: None,
69            kinds: None,
70            units: None,
71        };
72        metadata.validate()?;
73        Ok(metadata)
74    }
75    pub fn with_kinds(mut self, kinds: Vec<FeatureKind>) -> Result<Self> {
76        if kinds.len() != self.names.len() {
77            return Err(ShapError::DimensionMismatch {
78                expected: format!("{} feature kinds", self.names.len()),
79                found: format!("{}", kinds.len()),
80            });
81        }
82        self.kinds = Some(kinds);
83        Ok(self)
84    }
85    pub fn with_units(mut self, units: Vec<Option<String>>) -> Result<Self> {
86        if units.len() != self.names.len() {
87            return Err(ShapError::DimensionMismatch {
88                expected: format!("{} feature units", self.names.len()),
89                found: format!("{}", units.len()),
90            });
91        }
92        self.units = Some(units);
93        Ok(self)
94    }
95    pub fn with_display_names(mut self, names: Vec<String>) -> Result<Self> {
96        if names.len() != self.names.len() {
97            return Err(ShapError::DimensionMismatch {
98                expected: format!("{} display names", self.names.len()),
99                found: format!("{}", names.len()),
100            });
101        }
102        self.display_names = Some(names);
103        Ok(self)
104    }
105
106    /// Validates names and the lengths of every optional metadata field.
107    pub fn validate(&self) -> Result<()> {
108        if self.names.is_empty() || self.names.iter().any(|name| name.is_empty()) {
109            return Err(ShapError::InvalidConfiguration(
110                "feature names cannot be empty".into(),
111            ));
112        }
113        let mut unique = self.names.clone();
114        unique.sort();
115        unique.dedup();
116        if unique.len() != self.names.len() {
117            return Err(ShapError::InvalidConfiguration(
118                "feature names must be unique".into(),
119            ));
120        }
121        if self
122            .display_names
123            .as_ref()
124            .is_some_and(|values| values.len() != self.names.len())
125            || self
126                .kinds
127                .as_ref()
128                .is_some_and(|values| values.len() != self.names.len())
129            || self
130                .units
131                .as_ref()
132                .is_some_and(|values| values.len() != self.names.len())
133        {
134            return Err(ShapError::DimensionMismatch {
135                expected: format!(
136                    "{} entries in all feature metadata fields",
137                    self.names.len()
138                ),
139                found: "inconsistent feature metadata".into(),
140            });
141        }
142        Ok(())
143    }
144}
145
146#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
147#[serde(try_from = "OutputMetadataPayload")]
148pub struct OutputMetadata {
149    pub names: Vec<String>,
150    pub kinds: Option<Vec<OutputKind>>,
151}
152#[derive(Deserialize)]
153struct OutputMetadataPayload {
154    names: Vec<String>,
155    kinds: Option<Vec<OutputKind>>,
156}
157impl TryFrom<OutputMetadataPayload> for OutputMetadata {
158    type Error = ShapError;
159    fn try_from(payload: OutputMetadataPayload) -> Result<Self> {
160        let metadata = Self {
161            names: payload.names,
162            kinds: payload.kinds,
163        };
164        metadata.validate()?;
165        Ok(metadata)
166    }
167}
168impl OutputMetadata {
169    pub fn new(names: Vec<String>) -> Result<Self> {
170        if names.iter().any(|n| n.is_empty()) {
171            return Err(ShapError::InvalidConfiguration(
172                "output names cannot be empty".into(),
173            ));
174        }
175        let metadata = Self { names, kinds: None };
176        metadata.validate()?;
177        Ok(metadata)
178    }
179    pub fn with_kinds(mut self, kinds: Vec<OutputKind>) -> Result<Self> {
180        if kinds.len() != self.names.len() {
181            return Err(ShapError::DimensionMismatch {
182                expected: format!("{} output kinds", self.names.len()),
183                found: format!("{}", kinds.len()),
184            });
185        }
186        self.kinds = Some(kinds);
187        Ok(self)
188    }
189
190    /// Validates output names and optional output kinds.
191    pub fn validate(&self) -> Result<()> {
192        if self.names.is_empty() || self.names.iter().any(|name| name.is_empty()) {
193            return Err(ShapError::InvalidConfiguration(
194                "output names cannot be empty".into(),
195            ));
196        }
197        if self
198            .kinds
199            .as_ref()
200            .is_some_and(|values| values.len() != self.names.len())
201        {
202            return Err(ShapError::DimensionMismatch {
203                expected: format!("{} output kinds", self.names.len()),
204                found: "inconsistent output metadata".into(),
205            });
206        }
207        Ok(())
208    }
209}
210
211#[cfg(test)]
212mod tests {
213    use super::*;
214
215    #[test]
216    fn validates_directly_constructed_feature_metadata() {
217        let metadata = FeatureMetadata {
218            names: vec!["a".into(), "b".into()],
219            display_names: Some(vec!["A".into()]),
220            kinds: None,
221            units: None,
222        };
223        assert!(matches!(
224            metadata.validate(),
225            Err(ShapError::DimensionMismatch { .. })
226        ));
227    }
228
229    #[test]
230    fn validates_directly_constructed_output_metadata() {
231        let metadata = OutputMetadata {
232            names: vec!["prediction".into()],
233            kinds: Some(vec![OutputKind::Regression, OutputKind::Probability]),
234        };
235        assert!(matches!(
236            metadata.validate(),
237            Err(ShapError::DimensionMismatch { .. })
238        ));
239    }
240}