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 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 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}