1use serde::{Deserialize, Serialize};
18use std::collections::HashMap;
19
20pub const CLASSIFIER_KERNEL_MORPHOLOGY: &str = "episodic_memory";
22pub const CLASSIFIER_CLASS_MORPHOLOGY: &str = "episodic_memory";
24pub const CLASSIFIER_ASSOCIATIVE_MORPHOLOGY: &str = "associative_memory";
26pub const CLASSIFIER_SCAN_MORPHOLOGY: &str = "episodic_scan";
28
29#[derive(Debug, Clone, PartialEq, Eq)]
31pub struct ClassifierMapping {
32 pub src_area_id: String,
33 pub dst_area_id: String,
34 pub morphology_id: String,
35}
36
37#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
39pub struct ClassifierField {
40 pub field_area_id: String,
41 pub scan_twin_id: String,
42}
43
44#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Default)]
46#[serde(rename_all = "snake_case")]
47pub enum ClassifierTrainingMode {
48 #[default]
50 Kernel,
51 Scanner,
53}
54
55impl ClassifierTrainingMode {
56 pub fn as_str(self) -> &'static str {
58 match self {
59 Self::Kernel => "kernel",
60 Self::Scanner => "scanner",
61 }
62 }
63
64 pub fn parse(value: &str) -> Result<Self, String> {
66 match value {
67 "kernel" => Ok(Self::Kernel),
68 "scanner" => Ok(Self::Scanner),
69 other => Err(format!(
70 "training_mode must be \"kernel\" or \"scanner\", got \"{other}\""
71 )),
72 }
73 }
74}
75
76#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
78pub struct Classifier {
79 pub classifier_id: String,
81 pub name: String,
82 pub parent_region_id: String,
84 pub coordinates_3d: [i32; 3],
85 #[serde(default)]
88 pub training_mode: ClassifierTrainingMode,
89 #[serde(default, skip_serializing_if = "Option::is_none")]
91 pub kernel_area_id: Option<String>,
92 #[serde(default, skip_serializing_if = "Option::is_none")]
93 pub class_area_id: Option<String>,
94 #[serde(default, skip_serializing_if = "Option::is_none")]
96 pub mask_area_id: Option<String>,
97 #[serde(default, skip_serializing_if = "Option::is_none")]
99 pub kernel_size: Option<[u32; 3]>,
100 #[serde(default)]
102 pub fields: Vec<ClassifierField>,
103 pub kernel_memory_id: String,
105 pub class_memory_id: String,
106 #[serde(default)]
107 pub properties: HashMap<String, serde_json::Value>,
108}
109
110impl Classifier {
111 pub fn assembly_core_ids(&self) -> Vec<String> {
113 vec![self.kernel_memory_id.clone(), self.class_memory_id.clone()]
114 }
115
116 pub fn owned_area_ids(&self) -> Vec<String> {
118 let mut owned = self.assembly_core_ids();
119 for field in &self.fields {
120 if !field.scan_twin_id.is_empty() {
121 owned.push(field.scan_twin_id.clone());
122 }
123 }
124 owned
125 }
126
127 pub fn input_area_ids(&self) -> Vec<String> {
129 let mut inputs = Vec::new();
130 if let Some(kernel) = &self.kernel_area_id {
131 inputs.push(kernel.clone());
132 }
133 if let Some(class) = &self.class_area_id {
134 inputs.push(class.clone());
135 }
136 if let Some(mask) = &self.mask_area_id {
137 inputs.push(mask.clone());
138 }
139 for field in &self.fields {
140 inputs.push(field.field_area_id.clone());
141 }
142 inputs
143 }
144
145 pub fn owns_assembly_core(&self, area_id: &str) -> bool {
146 self.kernel_memory_id == area_id || self.class_memory_id == area_id
147 }
148
149 pub fn owns_area(&self, area_id: &str) -> bool {
151 self.owns_assembly_core(area_id) || self.field_for_twin(area_id).is_some()
152 }
153
154 pub fn field_for_twin(&self, twin_id: &str) -> Option<&ClassifierField> {
155 self.fields
156 .iter()
157 .find(|field| field.scan_twin_id == twin_id)
158 }
159
160 pub fn binding_for_field(&self, field_area_id: &str) -> Option<&ClassifierField> {
161 self.fields
162 .iter()
163 .find(|field| field.field_area_id == field_area_id)
164 }
165
166 pub fn references_input(&self, area_id: &str) -> bool {
168 self.kernel_area_id.as_deref() == Some(area_id)
169 || self.class_area_id.as_deref() == Some(area_id)
170 || self.mask_area_id.as_deref() == Some(area_id)
171 || self.binding_for_field(area_id).is_some()
172 }
173
174 pub fn apply_training_inputs(
178 &mut self,
179 mode: ClassifierTrainingMode,
180 kernel_area_id: Option<String>,
181 class_area_id: Option<String>,
182 mask_area_id: Option<String>,
183 kernel_size: Option<[u32; 3]>,
184 ) -> Result<bool, String> {
185 let previous_mode = self.training_mode;
186 let previous_size = self.kernel_size;
187 match mode {
188 ClassifierTrainingMode::Kernel => {
189 let kernel = kernel_area_id.ok_or_else(|| "kernel_area_id required".to_string())?;
190 let class = class_area_id.ok_or_else(|| "class_area_id required".to_string())?;
191 self.kernel_area_id = Some(required_area_id(kernel, "kernel_area_id")?);
192 self.class_area_id = Some(required_area_id(class, "class_area_id")?);
193 self.mask_area_id = None;
194 self.kernel_size = None;
195 }
196 ClassifierTrainingMode::Scanner => {
197 let mask = mask_area_id.ok_or_else(|| "mask_area_id required".to_string())?;
198 let size = kernel_size.ok_or_else(|| "kernel_size required".to_string())?;
199 validate_kernel_size(size)?;
200 self.mask_area_id = Some(required_area_id(mask, "mask_area_id")?);
201 self.kernel_size = Some(size);
202 self.kernel_area_id = None;
203 self.class_area_id = None;
204 }
205 }
206 self.training_mode = mode;
207 Ok(previous_mode != mode
208 || (mode == ClassifierTrainingMode::Scanner && previous_size != self.kernel_size))
209 }
210
211 pub fn apply_assembly_update(
213 &mut self,
214 name: Option<String>,
215 coordinates_3d: Option<[i32; 3]>,
216 parent_region_id: Option<String>,
217 kernel_area_id: Option<String>,
218 class_area_id: Option<String>,
219 ) -> Result<(), String> {
220 if let Some(name) = name {
221 let trimmed = name.trim();
222 if trimmed.is_empty() {
223 return Err("Classifier name cannot be blank".to_string());
224 }
225 self.name = trimmed.to_string();
226 }
227 if let Some(coordinates_3d) = coordinates_3d {
228 self.coordinates_3d = coordinates_3d;
229 }
230 if let Some(parent_region_id) = parent_region_id {
231 let trimmed = parent_region_id.trim();
232 if trimmed.is_empty() {
233 return Err("parent_region_id cannot be blank".to_string());
234 }
235 self.parent_region_id = trimmed.to_string();
236 }
237 if let Some(kernel_area_id) = kernel_area_id {
238 self.kernel_area_id = Some(required_area_id(kernel_area_id, "kernel_area_id")?);
239 }
240 if let Some(class_area_id) = class_area_id {
241 self.class_area_id = Some(required_area_id(class_area_id, "class_area_id")?);
242 }
243 Ok(())
244 }
245
246 pub fn apply_metadata_update(
248 &mut self,
249 name: Option<String>,
250 coordinates_3d: Option<[i32; 3]>,
251 ) -> Result<(), String> {
252 self.apply_assembly_update(name, coordinates_3d, None, None, None)
253 }
254
255 pub fn required_mappings(&self) -> Vec<ClassifierMapping> {
257 let mut mappings = Vec::new();
258 if let Some(kernel) = &self.kernel_area_id {
259 mappings.push(ClassifierMapping {
260 src_area_id: kernel.clone(),
261 dst_area_id: self.kernel_memory_id.clone(),
262 morphology_id: CLASSIFIER_KERNEL_MORPHOLOGY.to_string(),
263 });
264 }
265 if let Some(class) = &self.class_area_id {
266 mappings.push(ClassifierMapping {
267 src_area_id: class.clone(),
268 dst_area_id: self.class_memory_id.clone(),
269 morphology_id: CLASSIFIER_CLASS_MORPHOLOGY.to_string(),
270 });
271 }
272 mappings.push(ClassifierMapping {
273 src_area_id: self.kernel_memory_id.clone(),
274 dst_area_id: self.class_memory_id.clone(),
275 morphology_id: CLASSIFIER_ASSOCIATIVE_MORPHOLOGY.to_string(),
276 });
277 for field in &self.fields {
278 mappings.push(ClassifierMapping {
279 src_area_id: field.field_area_id.clone(),
280 dst_area_id: self.kernel_memory_id.clone(),
281 morphology_id: CLASSIFIER_SCAN_MORPHOLOGY.to_string(),
282 });
283 }
284 mappings
285 }
286
287 pub fn clear_input(&mut self, area_id: &str) {
289 if self.kernel_area_id.as_deref() == Some(area_id) {
290 self.kernel_area_id = None;
291 }
292 if self.class_area_id.as_deref() == Some(area_id) {
293 self.class_area_id = None;
294 }
295 if self.mask_area_id.as_deref() == Some(area_id) {
296 self.mask_area_id = None;
297 }
298 self.fields.retain(|field| field.field_area_id != area_id);
299 }
300
301 pub fn detach_field(&mut self, field_area_id: &str) -> Option<String> {
303 let position = self
304 .fields
305 .iter()
306 .position(|field| field.field_area_id == field_area_id)?;
307 Some(self.fields.remove(position).scan_twin_id)
308 }
309
310 pub fn detach_twin(&mut self, twin_id: &str) -> Option<String> {
312 let position = self
313 .fields
314 .iter()
315 .position(|field| field.scan_twin_id == twin_id)?;
316 Some(self.fields.remove(position).field_area_id)
317 }
318
319 pub fn attach_field(
320 &mut self,
321 field_area_id: String,
322 scan_twin_id: String,
323 ) -> Result<(), String> {
324 let field_area_id = required_area_id(field_area_id, "field_area_id")?;
325 let scan_twin_id = required_area_id(scan_twin_id, "scan_twin_id")?;
326 if self.binding_for_field(&field_area_id).is_some() {
327 return Err(format!(
328 "field_area_id {field_area_id} is already mapped to this classifier"
329 ));
330 }
331 self.fields.push(ClassifierField {
332 field_area_id,
333 scan_twin_id,
334 });
335 Ok(())
336 }
337
338 pub fn apply_mapping_change(
340 &mut self,
341 src_area_id: &str,
342 dst_area_id: &str,
343 morphology_id: &str,
344 removed: bool,
345 ) -> bool {
346 if dst_area_id == self.kernel_memory_id && morphology_id == CLASSIFIER_KERNEL_MORPHOLOGY {
347 self.kernel_area_id = if removed {
348 None
349 } else {
350 Some(src_area_id.to_string())
351 };
352 return true;
353 }
354 if dst_area_id == self.class_memory_id && morphology_id == CLASSIFIER_CLASS_MORPHOLOGY {
355 self.class_area_id = if removed {
356 None
357 } else {
358 Some(src_area_id.to_string())
359 };
360 return true;
361 }
362 false
363 }
364}
365
366pub fn validate_kernel_size(size: [u32; 3]) -> Result<(), String> {
368 if size[0] == 0 || size[1] == 0 || size[2] == 0 {
369 return Err("kernel_size axes must be greater than zero".to_string());
370 }
371 Ok(())
372}
373
374pub fn validate_scanner_field(
376 kernel_size: [u32; 3],
377 field: [u32; 3],
378 mask: [u32; 3],
379) -> Result<(), String> {
380 validate_kernel_size(kernel_size)?;
381 if field[0] == 0 || field[1] == 0 || field[2] == 0 {
382 return Err("field dimensions must be greater than zero".to_string());
383 }
384 if mask[2] == 0 {
385 return Err("mask depth must be greater than zero".to_string());
386 }
387 if kernel_size[0] > field[0] || kernel_size[1] > field[1] {
388 return Err("kernel_size does not fit the field".to_string());
389 }
390 if kernel_size[2] != field[2] {
391 return Err("kernel depth must equal the field depth".to_string());
392 }
393 if mask[0] != field[0] || mask[1] != field[1] {
394 return Err("mask width and height must equal the field".to_string());
395 }
396 Ok(())
397}
398
399fn required_area_id(area_id: String, field: &str) -> Result<String, String> {
400 let trimmed = area_id.trim();
401 if trimmed.is_empty() {
402 return Err(format!("{field} cannot be blank"));
403 }
404 Ok(trimmed.to_string())
405}
406
407#[cfg(test)]
408mod tests {
409 use super::*;
410
411 fn sample() -> Classifier {
412 let mut classifier = Classifier {
413 classifier_id: "clf-1".to_string(),
414 name: "demo".to_string(),
415 parent_region_id: "region".to_string(),
416 coordinates_3d: [1, 2, 3],
417 training_mode: ClassifierTrainingMode::Kernel,
418 kernel_area_id: Some("kernel".to_string()),
419 class_area_id: Some("class".to_string()),
420 mask_area_id: None,
421 kernel_size: None,
422 fields: Vec::new(),
423 kernel_memory_id: "kmem".to_string(),
424 class_memory_id: "cmem".to_string(),
425 properties: HashMap::new(),
426 };
427 classifier
428 .attach_field("field".to_string(), "twin".to_string())
429 .expect("first field");
430 classifier
431 }
432
433 #[test]
434 fn required_mappings_cover_shared_edges_and_each_field() {
435 let mut classifier = sample();
436 classifier
437 .attach_field("field-b".to_string(), "twin-b".to_string())
438 .expect("second field");
439 let mappings = classifier.required_mappings();
440 assert_eq!(mappings.len(), 5);
441 assert!(mappings.iter().any(|m| {
442 m.src_area_id == "field"
443 && m.dst_area_id == "kmem"
444 && m.morphology_id == CLASSIFIER_SCAN_MORPHOLOGY
445 }));
446 assert!(mappings.iter().any(|m| m.src_area_id == "field-b"));
447 }
448
449 #[test]
450 fn deleting_one_field_keeps_the_other_eye() {
451 let mut classifier = sample();
452 classifier
453 .attach_field("field-b".to_string(), "twin-b".to_string())
454 .expect("second field");
455 assert_eq!(classifier.detach_field("field"), Some("twin".to_string()));
456 assert!(classifier.binding_for_field("field").is_none());
457 assert_eq!(
458 classifier
459 .binding_for_field("field-b")
460 .map(|f| f.scan_twin_id.as_str()),
461 Some("twin-b")
462 );
463 assert!(classifier.owns_assembly_core("kmem"));
464 assert!(!classifier.owns_area("twin"));
465 assert!(classifier.owns_area("twin-b"));
466 }
467
468 #[test]
469 fn duplicate_field_mapping_is_rejected() {
470 let mut classifier = sample();
471 let result = classifier.attach_field("field".to_string(), "other-twin".to_string());
472 assert!(result.is_err());
473 assert_eq!(classifier.fields.len(), 1);
474 }
475
476 #[test]
477 fn metadata_update_renames_without_touching_areas() {
478 let mut classifier = sample();
479 classifier
480 .apply_metadata_update(Some(" renamed ".to_string()), Some([9, 8, 7]))
481 .expect("valid metadata");
482 assert_eq!(classifier.name, "renamed");
483 assert_eq!(classifier.coordinates_3d, [9, 8, 7]);
484 assert_eq!(classifier.kernel_memory_id, "kmem");
485 assert_eq!(classifier.fields[0].scan_twin_id, "twin");
486 assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel"));
487 }
488
489 #[test]
490 fn assembly_update_retargets_kernel_and_class_only() {
491 let mut classifier = sample();
492 classifier
493 .apply_assembly_update(
494 None,
495 None,
496 Some("other-region".to_string()),
497 Some("kernel2".to_string()),
498 Some("class2".to_string()),
499 )
500 .expect("valid assembly update");
501 assert_eq!(classifier.parent_region_id, "other-region");
502 assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel2"));
503 assert_eq!(classifier.class_area_id.as_deref(), Some("class2"));
504 assert_eq!(classifier.fields[0].field_area_id, "field");
505 }
506
507 #[test]
508 fn scanner_mode_clears_kernel_inputs_and_reports_geometry_change() {
509 let mut classifier = sample();
510 let changed = classifier
511 .apply_training_inputs(
512 ClassifierTrainingMode::Scanner,
513 None,
514 None,
515 Some("mask".to_string()),
516 Some([8, 8, 3]),
517 )
518 .expect("scanner inputs");
519 assert!(changed);
520 assert_eq!(classifier.training_mode, ClassifierTrainingMode::Scanner);
521 assert!(classifier.kernel_area_id.is_none());
522 assert!(classifier.class_area_id.is_none());
523 assert_eq!(classifier.mask_area_id.as_deref(), Some("mask"));
524 assert_eq!(classifier.kernel_size, Some([8, 8, 3]));
525 assert!(classifier.references_input("mask"));
526 let same = classifier
527 .apply_training_inputs(
528 ClassifierTrainingMode::Scanner,
529 None,
530 None,
531 Some("mask".to_string()),
532 Some([8, 8, 3]),
533 )
534 .expect("same scanner geometry");
535 assert!(!same);
536 }
537
538 #[test]
539 fn kernel_mode_clears_scanner_inputs() {
540 let mut classifier = sample();
541 classifier
542 .apply_training_inputs(
543 ClassifierTrainingMode::Scanner,
544 None,
545 None,
546 Some("mask".to_string()),
547 Some([2, 2, 1]),
548 )
549 .expect("scanner");
550 classifier
551 .apply_training_inputs(
552 ClassifierTrainingMode::Kernel,
553 Some("kernel".to_string()),
554 Some("class".to_string()),
555 None,
556 None,
557 )
558 .expect("kernel");
559 assert!(classifier.mask_area_id.is_none());
560 assert!(classifier.kernel_size.is_none());
561 assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel"));
562 }
563
564 #[test]
565 fn scanner_field_must_match_mask_and_kernel_depth() {
566 assert!(validate_scanner_field([8, 8, 3], [256, 128, 3], [256, 128, 10]).is_ok());
567 assert!(validate_scanner_field([8, 8, 1], [256, 128, 3], [256, 128, 10]).is_err());
568 assert!(validate_scanner_field([8, 8, 3], [256, 128, 3], [200, 128, 10]).is_err());
569 }
570}