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)]
109 pub reward_training: bool,
110 #[serde(default, skip_serializing_if = "Option::is_none")]
113 pub answer_feedback_area_id: Option<String>,
114 #[serde(default, skip_serializing_if = "Option::is_none")]
116 pub pain_area_id: Option<String>,
117 #[serde(default, skip_serializing_if = "Option::is_none")]
119 pub pleasure_area_id: Option<String>,
120 #[serde(default)]
122 pub answer_latency_bursts: u32,
123 #[serde(default, skip_serializing_if = "Option::is_none")]
125 pub learn_area_id: Option<String>,
126 #[serde(default, skip_serializing_if = "Option::is_none")]
128 pub confidence_area_id: Option<String>,
129 #[serde(default)]
130 pub properties: HashMap<String, serde_json::Value>,
131}
132
133impl Classifier {
134 pub fn assembly_core_ids(&self) -> Vec<String> {
136 vec![self.kernel_memory_id.clone(), self.class_memory_id.clone()]
137 }
138
139 pub fn owned_area_ids(&self) -> Vec<String> {
141 let mut owned = self.assembly_core_ids();
142 for field in &self.fields {
143 if !field.scan_twin_id.is_empty() {
144 owned.push(field.scan_twin_id.clone());
145 }
146 }
147 if let Some(pain_area_id) = &self.pain_area_id {
148 if !pain_area_id.is_empty() {
149 owned.push(pain_area_id.clone());
150 }
151 }
152 if let Some(pleasure_area_id) = &self.pleasure_area_id {
153 if !pleasure_area_id.is_empty() {
154 owned.push(pleasure_area_id.clone());
155 }
156 }
157 owned
158 }
159
160 pub fn input_area_ids(&self) -> Vec<String> {
162 let mut inputs = Vec::new();
163 if let Some(kernel) = &self.kernel_area_id {
164 inputs.push(kernel.clone());
165 }
166 if let Some(class) = &self.class_area_id {
167 inputs.push(class.clone());
168 }
169 if let Some(mask) = &self.mask_area_id {
170 inputs.push(mask.clone());
171 }
172 for field in &self.fields {
173 inputs.push(field.field_area_id.clone());
174 }
175 inputs
176 }
177
178 pub fn owns_assembly_core(&self, area_id: &str) -> bool {
179 self.kernel_memory_id == area_id || self.class_memory_id == area_id
180 }
181
182 pub fn owns_area(&self, area_id: &str) -> bool {
184 self.owns_assembly_core(area_id) || self.field_for_twin(area_id).is_some()
185 }
186
187 pub fn field_for_twin(&self, twin_id: &str) -> Option<&ClassifierField> {
188 self.fields
189 .iter()
190 .find(|field| field.scan_twin_id == twin_id)
191 }
192
193 pub fn binding_for_field(&self, field_area_id: &str) -> Option<&ClassifierField> {
194 self.fields
195 .iter()
196 .find(|field| field.field_area_id == field_area_id)
197 }
198
199 pub fn references_input(&self, area_id: &str) -> bool {
201 self.kernel_area_id.as_deref() == Some(area_id)
202 || self.class_area_id.as_deref() == Some(area_id)
203 || self.mask_area_id.as_deref() == Some(area_id)
204 || self.binding_for_field(area_id).is_some()
205 }
206
207 pub fn apply_training_inputs(
211 &mut self,
212 mode: ClassifierTrainingMode,
213 kernel_area_id: Option<String>,
214 class_area_id: Option<String>,
215 mask_area_id: Option<String>,
216 kernel_size: Option<[u32; 3]>,
217 ) -> Result<bool, String> {
218 let previous_mode = self.training_mode;
219 let previous_size = self.kernel_size;
220 match mode {
221 ClassifierTrainingMode::Kernel => {
222 let kernel = kernel_area_id.ok_or_else(|| "kernel_area_id required".to_string())?;
223 let class = class_area_id.ok_or_else(|| "class_area_id required".to_string())?;
224 self.kernel_area_id = Some(required_area_id(kernel, "kernel_area_id")?);
225 self.class_area_id = Some(required_area_id(class, "class_area_id")?);
226 self.mask_area_id = None;
227 self.kernel_size = None;
228 }
229 ClassifierTrainingMode::Scanner => {
230 let mask = mask_area_id.ok_or_else(|| "mask_area_id required".to_string())?;
231 let size = kernel_size.ok_or_else(|| "kernel_size required".to_string())?;
232 validate_kernel_size(size)?;
233 self.mask_area_id = Some(required_area_id(mask, "mask_area_id")?);
234 self.kernel_size = Some(size);
235 self.kernel_area_id = None;
236 self.class_area_id = None;
237 }
238 }
239 self.training_mode = mode;
240 Ok(previous_mode != mode
241 || (mode == ClassifierTrainingMode::Scanner && previous_size != self.kernel_size))
242 }
243
244 pub fn apply_assembly_update(
246 &mut self,
247 name: Option<String>,
248 coordinates_3d: Option<[i32; 3]>,
249 parent_region_id: Option<String>,
250 kernel_area_id: Option<String>,
251 class_area_id: Option<String>,
252 ) -> Result<(), String> {
253 if let Some(name) = name {
254 let trimmed = name.trim();
255 if trimmed.is_empty() {
256 return Err("Classifier name cannot be blank".to_string());
257 }
258 self.name = trimmed.to_string();
259 }
260 if let Some(coordinates_3d) = coordinates_3d {
261 self.coordinates_3d = coordinates_3d;
262 }
263 if let Some(parent_region_id) = parent_region_id {
264 let trimmed = parent_region_id.trim();
265 if trimmed.is_empty() {
266 return Err("parent_region_id cannot be blank".to_string());
267 }
268 self.parent_region_id = trimmed.to_string();
269 }
270 if let Some(kernel_area_id) = kernel_area_id {
271 self.kernel_area_id = Some(required_area_id(kernel_area_id, "kernel_area_id")?);
272 }
273 if let Some(class_area_id) = class_area_id {
274 self.class_area_id = Some(required_area_id(class_area_id, "class_area_id")?);
275 }
276 Ok(())
277 }
278
279 pub fn apply_metadata_update(
281 &mut self,
282 name: Option<String>,
283 coordinates_3d: Option<[i32; 3]>,
284 ) -> Result<(), String> {
285 self.apply_assembly_update(name, coordinates_3d, None, None, None)
286 }
287
288 pub fn required_mappings(&self) -> Vec<ClassifierMapping> {
290 let mut mappings = Vec::new();
291 if let Some(kernel) = &self.kernel_area_id {
292 mappings.push(ClassifierMapping {
293 src_area_id: kernel.clone(),
294 dst_area_id: self.kernel_memory_id.clone(),
295 morphology_id: CLASSIFIER_KERNEL_MORPHOLOGY.to_string(),
296 });
297 }
298 if let Some(class) = &self.class_area_id {
299 mappings.push(ClassifierMapping {
300 src_area_id: class.clone(),
301 dst_area_id: self.class_memory_id.clone(),
302 morphology_id: CLASSIFIER_CLASS_MORPHOLOGY.to_string(),
303 });
304 }
305 mappings.push(ClassifierMapping {
306 src_area_id: self.kernel_memory_id.clone(),
307 dst_area_id: self.class_memory_id.clone(),
308 morphology_id: CLASSIFIER_ASSOCIATIVE_MORPHOLOGY.to_string(),
309 });
310 for field in &self.fields {
311 mappings.push(ClassifierMapping {
312 src_area_id: field.field_area_id.clone(),
313 dst_area_id: self.kernel_memory_id.clone(),
314 morphology_id: CLASSIFIER_SCAN_MORPHOLOGY.to_string(),
315 });
316 }
317 mappings
318 }
319
320 pub fn clear_input(&mut self, area_id: &str) {
322 if self.kernel_area_id.as_deref() == Some(area_id) {
323 self.kernel_area_id = None;
324 }
325 if self.class_area_id.as_deref() == Some(area_id) {
326 self.class_area_id = None;
327 }
328 if self.mask_area_id.as_deref() == Some(area_id) {
329 self.mask_area_id = None;
330 }
331 if self.answer_feedback_area_id.as_deref() == Some(area_id) {
332 self.answer_feedback_area_id = None;
333 }
334 if self.learn_area_id.as_deref() == Some(area_id) {
335 self.learn_area_id = None;
336 }
337 if self.confidence_area_id.as_deref() == Some(area_id) {
338 self.confidence_area_id = None;
339 }
340 self.fields.retain(|field| field.field_area_id != area_id);
341 }
342
343 pub fn detach_field(&mut self, field_area_id: &str) -> Option<String> {
345 let position = self
346 .fields
347 .iter()
348 .position(|field| field.field_area_id == field_area_id)?;
349 Some(self.fields.remove(position).scan_twin_id)
350 }
351
352 pub fn detach_twin(&mut self, twin_id: &str) -> Option<String> {
354 let position = self
355 .fields
356 .iter()
357 .position(|field| field.scan_twin_id == twin_id)?;
358 Some(self.fields.remove(position).field_area_id)
359 }
360
361 pub fn attach_field(
362 &mut self,
363 field_area_id: String,
364 scan_twin_id: String,
365 ) -> Result<(), String> {
366 let field_area_id = required_area_id(field_area_id, "field_area_id")?;
367 let scan_twin_id = required_area_id(scan_twin_id, "scan_twin_id")?;
368 if self.binding_for_field(&field_area_id).is_some() {
369 return Err(format!(
370 "field_area_id {field_area_id} is already mapped to this classifier"
371 ));
372 }
373 self.fields.push(ClassifierField {
374 field_area_id,
375 scan_twin_id,
376 });
377 Ok(())
378 }
379
380 pub fn apply_mapping_change(
382 &mut self,
383 src_area_id: &str,
384 dst_area_id: &str,
385 morphology_id: &str,
386 removed: bool,
387 ) -> bool {
388 if dst_area_id == self.kernel_memory_id && morphology_id == CLASSIFIER_KERNEL_MORPHOLOGY {
389 self.kernel_area_id = if removed {
390 None
391 } else {
392 Some(src_area_id.to_string())
393 };
394 return true;
395 }
396 if dst_area_id == self.class_memory_id && morphology_id == CLASSIFIER_CLASS_MORPHOLOGY {
397 self.class_area_id = if removed {
398 None
399 } else {
400 Some(src_area_id.to_string())
401 };
402 return true;
403 }
404 false
405 }
406}
407
408pub fn validate_kernel_size(size: [u32; 3]) -> Result<(), String> {
410 if size[0] == 0 || size[1] == 0 || size[2] == 0 {
411 return Err("kernel_size axes must be greater than zero".to_string());
412 }
413 Ok(())
414}
415
416pub fn validate_scanner_field(
418 kernel_size: [u32; 3],
419 field: [u32; 3],
420 mask: [u32; 3],
421) -> Result<(), String> {
422 validate_kernel_size(kernel_size)?;
423 if field[0] == 0 || field[1] == 0 || field[2] == 0 {
424 return Err("field dimensions must be greater than zero".to_string());
425 }
426 if mask[2] == 0 {
427 return Err("mask depth must be greater than zero".to_string());
428 }
429 if kernel_size[0] > field[0] || kernel_size[1] > field[1] {
430 return Err("kernel_size does not fit the field".to_string());
431 }
432 if kernel_size[2] != field[2] {
433 return Err("kernel depth must equal the field depth".to_string());
434 }
435 if mask[0] != field[0] || mask[1] != field[1] {
436 return Err("mask width and height must equal the field".to_string());
437 }
438 Ok(())
439}
440
441pub fn is_whole_image_class_shape(feedback: [u32; 3], class_count: u32) -> bool {
443 if class_count == 0 {
444 return false;
445 }
446 let ones = feedback.iter().filter(|axis| **axis == 1).count();
447 let volume = u64::from(feedback[0])
448 .saturating_mul(u64::from(feedback[1]))
449 .saturating_mul(u64::from(feedback[2]));
450 ones >= 2 && volume == u64::from(class_count)
451}
452
453pub fn validate_answer_feedback_shape(
458 mode: ClassifierTrainingMode,
459 feedback: [u32; 3],
460 reference: [u32; 3],
461 output_shapes: &[[u32; 3]],
462) -> Result<(), String> {
463 if feedback[0] == 0 || feedback[1] == 0 || feedback[2] == 0 {
464 return Err("answer feedback dimensions must be greater than zero".to_string());
465 }
466 match mode {
467 ClassifierTrainingMode::Kernel => {
468 if feedback != reference {
469 return Err("answer feedback dimensions must match the class area".to_string());
470 }
471 }
472 ClassifierTrainingMode::Scanner => {
473 let class_count = if output_shapes.is_empty() {
474 reference[2]
475 } else {
476 output_shapes[0][2]
477 };
478 if is_whole_image_class_shape(feedback, class_count) {
479 return Ok(());
480 }
481 let shapes = if output_shapes.is_empty() {
482 std::slice::from_ref(&reference)
483 } else {
484 output_shapes
485 };
486 if shapes.iter().any(|shape| *shape != feedback) {
487 return Err(
488 "answer feedback dimensions must match the detection output or a whole-image class"
489 .to_string(),
490 );
491 }
492 }
493 }
494 Ok(())
495}
496
497fn required_area_id(area_id: String, field: &str) -> Result<String, String> {
498 let trimmed = area_id.trim();
499 if trimmed.is_empty() {
500 return Err(format!("{field} cannot be blank"));
501 }
502 Ok(trimmed.to_string())
503}
504
505#[cfg(test)]
506mod tests {
507 use super::*;
508
509 fn sample() -> Classifier {
510 let mut classifier = Classifier {
511 classifier_id: "clf-1".to_string(),
512 name: "demo".to_string(),
513 parent_region_id: "region".to_string(),
514 coordinates_3d: [1, 2, 3],
515 training_mode: ClassifierTrainingMode::Kernel,
516 kernel_area_id: Some("kernel".to_string()),
517 class_area_id: Some("class".to_string()),
518 mask_area_id: None,
519 kernel_size: None,
520 fields: Vec::new(),
521 kernel_memory_id: "kmem".to_string(),
522 class_memory_id: "cmem".to_string(),
523 reward_training: false,
524 answer_feedback_area_id: None,
525 pain_area_id: None,
526 pleasure_area_id: None,
527 answer_latency_bursts: 0,
528 learn_area_id: None,
529 confidence_area_id: None,
530 properties: HashMap::new(),
531 };
532 classifier
533 .attach_field("field".to_string(), "twin".to_string())
534 .expect("first field");
535 classifier
536 }
537
538 #[test]
539 fn required_mappings_cover_shared_edges_and_each_field() {
540 let mut classifier = sample();
541 classifier
542 .attach_field("field-b".to_string(), "twin-b".to_string())
543 .expect("second field");
544 let mappings = classifier.required_mappings();
545 assert_eq!(mappings.len(), 5);
546 assert!(mappings.iter().any(|m| {
547 m.src_area_id == "field"
548 && m.dst_area_id == "kmem"
549 && m.morphology_id == CLASSIFIER_SCAN_MORPHOLOGY
550 }));
551 assert!(mappings.iter().any(|m| m.src_area_id == "field-b"));
552 }
553
554 #[test]
555 fn deleting_one_field_keeps_the_other_eye() {
556 let mut classifier = sample();
557 classifier
558 .attach_field("field-b".to_string(), "twin-b".to_string())
559 .expect("second field");
560 assert_eq!(classifier.detach_field("field"), Some("twin".to_string()));
561 assert!(classifier.binding_for_field("field").is_none());
562 assert_eq!(
563 classifier
564 .binding_for_field("field-b")
565 .map(|f| f.scan_twin_id.as_str()),
566 Some("twin-b")
567 );
568 assert!(classifier.owns_assembly_core("kmem"));
569 assert!(!classifier.owns_area("twin"));
570 assert!(classifier.owns_area("twin-b"));
571 }
572
573 #[test]
574 fn duplicate_field_mapping_is_rejected() {
575 let mut classifier = sample();
576 let result = classifier.attach_field("field".to_string(), "other-twin".to_string());
577 assert!(result.is_err());
578 assert_eq!(classifier.fields.len(), 1);
579 }
580
581 #[test]
582 fn metadata_update_renames_without_touching_areas() {
583 let mut classifier = sample();
584 classifier
585 .apply_metadata_update(Some(" renamed ".to_string()), Some([9, 8, 7]))
586 .expect("valid metadata");
587 assert_eq!(classifier.name, "renamed");
588 assert_eq!(classifier.coordinates_3d, [9, 8, 7]);
589 assert_eq!(classifier.kernel_memory_id, "kmem");
590 assert_eq!(classifier.fields[0].scan_twin_id, "twin");
591 assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel"));
592 }
593
594 #[test]
595 fn assembly_update_retargets_kernel_and_class_only() {
596 let mut classifier = sample();
597 classifier
598 .apply_assembly_update(
599 None,
600 None,
601 Some("other-region".to_string()),
602 Some("kernel2".to_string()),
603 Some("class2".to_string()),
604 )
605 .expect("valid assembly update");
606 assert_eq!(classifier.parent_region_id, "other-region");
607 assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel2"));
608 assert_eq!(classifier.class_area_id.as_deref(), Some("class2"));
609 assert_eq!(classifier.fields[0].field_area_id, "field");
610 }
611
612 #[test]
613 fn scanner_mode_clears_kernel_inputs_and_reports_geometry_change() {
614 let mut classifier = sample();
615 let changed = classifier
616 .apply_training_inputs(
617 ClassifierTrainingMode::Scanner,
618 None,
619 None,
620 Some("mask".to_string()),
621 Some([8, 8, 3]),
622 )
623 .expect("scanner inputs");
624 assert!(changed);
625 assert_eq!(classifier.training_mode, ClassifierTrainingMode::Scanner);
626 assert!(classifier.kernel_area_id.is_none());
627 assert!(classifier.class_area_id.is_none());
628 assert_eq!(classifier.mask_area_id.as_deref(), Some("mask"));
629 assert_eq!(classifier.kernel_size, Some([8, 8, 3]));
630 assert!(classifier.references_input("mask"));
631 let same = classifier
632 .apply_training_inputs(
633 ClassifierTrainingMode::Scanner,
634 None,
635 None,
636 Some("mask".to_string()),
637 Some([8, 8, 3]),
638 )
639 .expect("same scanner geometry");
640 assert!(!same);
641 }
642
643 #[test]
644 fn kernel_mode_clears_scanner_inputs() {
645 let mut classifier = sample();
646 classifier
647 .apply_training_inputs(
648 ClassifierTrainingMode::Scanner,
649 None,
650 None,
651 Some("mask".to_string()),
652 Some([2, 2, 1]),
653 )
654 .expect("scanner");
655 classifier
656 .apply_training_inputs(
657 ClassifierTrainingMode::Kernel,
658 Some("kernel".to_string()),
659 Some("class".to_string()),
660 None,
661 None,
662 )
663 .expect("kernel");
664 assert!(classifier.mask_area_id.is_none());
665 assert!(classifier.kernel_size.is_none());
666 assert_eq!(classifier.kernel_area_id.as_deref(), Some("kernel"));
667 }
668
669 #[test]
670 fn answer_feedback_matches_class_area_or_detection_output() {
671 assert!(validate_answer_feedback_shape(
672 ClassifierTrainingMode::Kernel,
673 [1, 1, 4],
674 [1, 1, 4],
675 &[],
676 )
677 .is_ok());
678 assert!(validate_answer_feedback_shape(
679 ClassifierTrainingMode::Kernel,
680 [8, 8, 4],
681 [1, 1, 4],
682 &[],
683 )
684 .is_err());
685 assert!(validate_answer_feedback_shape(
686 ClassifierTrainingMode::Scanner,
687 [16, 16, 4],
688 [16, 16, 4],
689 &[[16, 16, 4]],
690 )
691 .is_ok());
692 assert!(validate_answer_feedback_shape(
693 ClassifierTrainingMode::Scanner,
694 [16, 16, 4],
695 [16, 16, 4],
696 &[[16, 16, 4], [8, 8, 4]],
697 )
698 .is_err());
699 assert!(validate_answer_feedback_shape(
700 ClassifierTrainingMode::Scanner,
701 [10, 1, 1],
702 [16, 16, 10],
703 &[[16, 16, 10]],
704 )
705 .is_ok());
706 }
707
708 #[test]
709 fn scanner_field_must_match_mask_and_kernel_depth() {
710 assert!(validate_scanner_field([8, 8, 3], [256, 128, 3], [256, 128, 10]).is_ok());
711 assert!(validate_scanner_field([8, 8, 1], [256, 128, 3], [256, 128, 10]).is_err());
712 assert!(validate_scanner_field([8, 8, 3], [256, 128, 3], [200, 128, 10]).is_err());
713 }
714}