1use crate::{matrix::FdMatrix, FdarError};
34
35#[derive(Debug, Clone, PartialEq)]
41#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
42pub struct FdComponent {
43 pub data: FdMatrix,
46 pub argvals: Vec<f64>,
49}
50
51#[derive(Debug, Clone, PartialEq)]
84#[non_exhaustive]
85#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
86pub struct MultiFunData {
87 components: Vec<FdComponent>,
88}
89
90impl MultiFunData {
91 pub fn new(components: Vec<FdComponent>) -> Result<Self, FdarError> {
119 if components.is_empty() {
120 return Err(FdarError::InvalidParameter {
121 parameter: "components",
122 message: "MultiFunData requires at least one component".to_string(),
123 });
124 }
125
126 let n_obs = components[0].data.nrows();
127
128 if components[0].argvals.len() != components[0].data.ncols() {
130 return Err(FdarError::InvalidDimension {
131 parameter: "components[0].argvals",
132 expected: format!("{}", components[0].data.ncols()),
133 actual: format!("{}", components[0].argvals.len()),
134 });
135 }
136
137 for (k, comp) in components.iter().enumerate().skip(1) {
138 if comp.data.nrows() != n_obs {
140 return Err(FdarError::InvalidDimension {
141 parameter: "components[k].data.nrows",
142 expected: format!("{n_obs} (same as component 0)"),
143 actual: format!("{} (component {k})", comp.data.nrows()),
144 });
145 }
146 if comp.argvals.len() != comp.data.ncols() {
148 return Err(FdarError::InvalidDimension {
149 parameter: "components[k].argvals",
150 expected: format!("{} (data.ncols for component {k})", comp.data.ncols()),
151 actual: format!("{}", comp.argvals.len()),
152 });
153 }
154 }
155
156 Ok(Self { components })
157 }
158
159 #[inline]
165 pub fn n_obs(&self) -> usize {
166 self.components[0].data.nrows()
167 }
168
169 #[inline]
171 pub fn n_components(&self) -> usize {
172 self.components.len()
173 }
174
175 pub fn component(&self, k: usize) -> Result<&FdComponent, FdarError> {
181 if k >= self.components.len() {
182 return Err(FdarError::InvalidParameter {
183 parameter: "k",
184 message: format!(
185 "component index {k} out of range (n_components = {})",
186 self.components.len()
187 ),
188 });
189 }
190 Ok(&self.components[k])
191 }
192
193 pub fn argvals(&self, k: usize) -> Result<&[f64], FdarError> {
199 if k >= self.components.len() {
200 return Err(FdarError::InvalidParameter {
201 parameter: "k",
202 message: format!(
203 "argvals index {k} out of range (n_components = {})",
204 self.components.len()
205 ),
206 });
207 }
208 Ok(&self.components[k].argvals)
209 }
210}
211
212#[cfg(test)]
213mod tests {
214 use super::*;
215 use crate::matrix::FdMatrix;
216
217 fn make_component(nrows: usize, ncols: usize) -> FdComponent {
218 FdComponent {
219 data: FdMatrix::zeros(nrows, ncols),
220 argvals: (0..ncols).map(|i| i as f64).collect(),
221 }
222 }
223
224 fn make_component_argvals(nrows: usize, argvals: Vec<f64>) -> FdComponent {
225 let ncols = argvals.len();
226 FdComponent {
227 data: FdMatrix::zeros(nrows, ncols),
228 argvals,
229 }
230 }
231
232 #[test]
235 fn test_two_component_different_grids_ok() {
236 let comp1 = make_component(5, 10);
238 let comp2 = make_component(5, 4);
239 let mfd = MultiFunData::new(vec![comp1, comp2]).unwrap();
240 assert_eq!(mfd.n_obs(), 5);
241 assert_eq!(mfd.n_components(), 2);
242 }
243
244 #[test]
245 fn test_single_component_ok() {
246 let comp = make_component(3, 6);
247 let mfd = MultiFunData::new(vec![comp]).unwrap();
248 assert_eq!(mfd.n_obs(), 3);
249 assert_eq!(mfd.n_components(), 1);
250 }
251
252 #[test]
253 fn test_three_components_same_nrows_ok() {
254 let comp1 = make_component(7, 5);
255 let comp2 = make_component(7, 10);
256 let comp3 = make_component(7, 3);
257 let mfd = MultiFunData::new(vec![comp1, comp2, comp3]).unwrap();
258 assert_eq!(mfd.n_obs(), 7);
259 assert_eq!(mfd.n_components(), 3);
260 }
261
262 #[test]
263 fn test_empty_components_err() {
264 let result = MultiFunData::new(vec![]);
265 assert!(matches!(result, Err(FdarError::InvalidParameter { .. })));
266 }
267
268 #[test]
269 fn test_mismatched_nrows_err() {
270 let comp1 = make_component(5, 10);
271 let comp2 = make_component(4, 10); let result = MultiFunData::new(vec![comp1, comp2]);
273 assert!(matches!(result, Err(FdarError::InvalidDimension { .. })));
274 }
275
276 #[test]
277 fn test_argvals_len_mismatch_first_component_err() {
278 let comp = FdComponent {
280 data: FdMatrix::zeros(5, 10),
281 argvals: vec![0.0, 1.0, 2.0], };
283 let result = MultiFunData::new(vec![comp]);
284 assert!(matches!(result, Err(FdarError::InvalidDimension { .. })));
285 }
286
287 #[test]
288 fn test_argvals_len_mismatch_later_component_err() {
289 let comp1 = make_component(5, 10);
291 let comp2 = FdComponent {
292 data: FdMatrix::zeros(5, 4),
293 argvals: vec![0.0, 1.0], };
295 let result = MultiFunData::new(vec![comp1, comp2]);
296 assert!(matches!(result, Err(FdarError::InvalidDimension { .. })));
297 }
298
299 #[test]
302 fn test_component_accessor_valid() {
303 let argvals1: Vec<f64> = (0..10).map(|i| i as f64 / 9.0).collect();
304 let argvals2: Vec<f64> = vec![0.0, 1.0, 2.0, 3.0];
305 let comp1 = make_component_argvals(5, argvals1.clone());
306 let comp2 = make_component_argvals(5, argvals2.clone());
307 let mfd = MultiFunData::new(vec![comp1, comp2]).unwrap();
308
309 let c0 = mfd.component(0).unwrap();
310 assert_eq!(c0.argvals, argvals1);
311 assert_eq!(c0.data.nrows(), 5);
312 assert_eq!(c0.data.ncols(), 10);
313
314 let c1 = mfd.component(1).unwrap();
315 assert_eq!(c1.argvals, argvals2);
316 assert_eq!(c1.data.ncols(), 4);
317 }
318
319 #[test]
320 fn test_component_accessor_out_of_range_err() {
321 let mfd = MultiFunData::new(vec![make_component(3, 5)]).unwrap();
322 let result = mfd.component(1);
323 assert!(matches!(result, Err(FdarError::InvalidParameter { .. })));
324 }
325
326 #[test]
327 fn test_argvals_accessor_valid() {
328 let argvals: Vec<f64> = vec![0.0, 0.5, 1.0];
329 let comp = make_component_argvals(4, argvals.clone());
330 let mfd = MultiFunData::new(vec![comp]).unwrap();
331 assert_eq!(mfd.argvals(0).unwrap(), argvals.as_slice());
332 }
333
334 #[test]
335 fn test_argvals_accessor_out_of_range_err() {
336 let mfd = MultiFunData::new(vec![make_component(3, 5)]).unwrap();
337 let result = mfd.argvals(5);
338 assert!(matches!(result, Err(FdarError::InvalidParameter { .. })));
339 }
340
341 #[test]
342 fn test_component_accessor_preserves_argvals_per_component() {
343 let argvals1: Vec<f64> = vec![0.0, 1.0, 2.0, 3.0, 4.0];
345 let argvals2: Vec<f64> = vec![10.0, 20.0];
346 let comp1 = make_component_argvals(6, argvals1.clone());
347 let comp2 = make_component_argvals(6, argvals2.clone());
348 let mfd = MultiFunData::new(vec![comp1, comp2]).unwrap();
349
350 assert_eq!(mfd.argvals(0).unwrap(), argvals1.as_slice());
351 assert_eq!(mfd.argvals(1).unwrap(), argvals2.as_slice());
352 }
353
354 #[test]
355 fn test_no_panic_on_out_of_range_component() {
356 let mfd = MultiFunData::new(vec![make_component(2, 3)]).unwrap();
357 assert!(mfd.component(100).is_err());
359 assert!(mfd.argvals(100).is_err());
360 assert!(mfd.component(usize::MAX).is_err());
361 }
362
363 #[test]
366 fn test_debug_clone_partialeq() {
367 let comp = make_component(2, 3);
368 let mfd = MultiFunData::new(vec![comp]).unwrap();
369 let mfd2 = mfd.clone();
370 assert_eq!(mfd, mfd2);
371 let s = format!("{:?}", mfd);
372 assert!(s.contains("MultiFunData"));
373 }
374
375 #[test]
376 fn test_fdcomponent_debug_clone_partialeq() {
377 let comp = make_component(2, 4);
378 let comp2 = comp.clone();
379 assert_eq!(comp, comp2);
380 let s = format!("{:?}", comp);
381 assert!(s.contains("FdComponent"));
382 }
383}