use tracing::{debug, instrument};
use crate::application::services::InterpretationConfig;
use crate::domain::entities::{
ClusterContext, ClusterId, EmbeddingId, NeighborEvidence, RecordingMetadata,
SegmentId, SequenceContext,
};
use crate::Result;
#[derive(Debug, Clone)]
pub struct EvidenceBuilder {
max_neighbors: usize,
include_spectrograms: bool,
include_sequences: bool,
sequence_window: usize,
min_distance_threshold: f32,
max_distance_threshold: f32,
}
impl EvidenceBuilder {
pub fn new(config: &InterpretationConfig) -> Self {
Self {
max_neighbors: config.max_neighbors,
include_spectrograms: config.include_spectrograms,
include_sequences: config.include_sequence_context,
sequence_window: config.sequence_context_window,
min_distance_threshold: 0.0,
max_distance_threshold: 1.0,
}
}
pub fn default_builder() -> Self {
Self {
max_neighbors: 10,
include_spectrograms: true,
include_sequences: true,
sequence_window: 3,
min_distance_threshold: 0.0,
max_distance_threshold: 1.0,
}
}
pub fn with_max_neighbors(mut self, n: usize) -> Self {
self.max_neighbors = n;
self
}
pub fn with_spectrograms(mut self, include: bool) -> Self {
self.include_spectrograms = include;
self
}
pub fn with_distance_threshold(mut self, min: f32, max: f32) -> Self {
self.min_distance_threshold = min;
self.max_distance_threshold = max;
self
}
pub fn max_neighbors(&self) -> usize {
self.max_neighbors
}
pub fn spectrograms_enabled(&self) -> bool {
self.include_spectrograms
}
#[instrument(skip(self, neighbors))]
pub async fn collect_neighbor_evidence(
&self,
neighbors: &[RawNeighbor],
) -> Result<Vec<NeighborEvidence>> {
let filtered: Vec<&RawNeighbor> = neighbors
.iter()
.filter(|n| {
n.distance >= self.min_distance_threshold
&& n.distance <= self.max_distance_threshold
})
.take(self.max_neighbors)
.collect();
debug!(
"Collecting evidence from {} neighbors (filtered from {})",
filtered.len(),
neighbors.len()
);
let evidence: Vec<NeighborEvidence> = filtered
.into_iter()
.map(|n| self.build_neighbor_evidence(n))
.collect();
Ok(evidence)
}
fn build_neighbor_evidence(&self, raw: &RawNeighbor) -> NeighborEvidence {
let metadata = raw
.metadata
.clone()
.unwrap_or_else(|| RecordingMetadata::new(&raw.embedding_id.0));
let mut evidence = NeighborEvidence::new(
raw.embedding_id.clone(),
raw.distance,
metadata,
);
if let Some(cluster_id) = &raw.cluster_id {
evidence = evidence.with_cluster(cluster_id.clone());
}
if self.include_spectrograms {
if let Some(url) = &raw.spectrogram_url {
evidence = evidence.with_spectrogram(url.clone());
}
}
evidence
}
#[instrument(skip(self))]
pub async fn build_cluster_context(
&self,
cluster_id: Option<ClusterId>,
label: Option<String>,
confidence: f32,
exemplar_similarity: f32,
) -> Result<ClusterContext> {
let context = ClusterContext {
assigned_cluster: cluster_id,
cluster_label: label,
confidence,
exemplar_similarity,
};
debug!(
"Built cluster context: assigned={}, confidence={}",
context.has_cluster(),
context.confidence
);
Ok(context)
}
#[instrument(skip(self))]
pub async fn build_sequence_context(
&self,
preceding: Vec<SegmentId>,
following: Vec<SegmentId>,
motif: Option<String>,
) -> Result<Option<SequenceContext>> {
if !self.include_sequences {
return Ok(None);
}
if preceding.is_empty() && following.is_empty() {
return Ok(None);
}
let preceding = preceding
.into_iter()
.take(self.sequence_window)
.collect();
let following = following
.into_iter()
.take(self.sequence_window)
.collect();
let context = SequenceContext {
preceding_segments: preceding,
following_segments: following,
detected_motif: motif,
};
debug!(
"Built sequence context: {} preceding, {} following, motif={}",
context.preceding_segments.len(),
context.following_segments.len(),
context.detected_motif.as_deref().unwrap_or("none")
);
Ok(Some(context))
}
pub fn aggregate_evidence_scores(&self, neighbors: &[NeighborEvidence]) -> EvidenceScores {
if neighbors.is_empty() {
return EvidenceScores::default();
}
let distances: Vec<f32> = neighbors.iter().map(|n| n.distance).collect();
let avg_distance = distances.iter().sum::<f32>() / distances.len() as f32;
let min_distance = distances.iter().cloned().fold(f32::INFINITY, f32::min);
let max_distance = distances.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let similarity = (1.0 - avg_distance).max(0.0).min(1.0);
let clustered_count = neighbors
.iter()
.filter(|n| n.cluster_id.is_some())
.count();
let cluster_coherence = if clustered_count > 0 {
let mut cluster_counts = std::collections::HashMap::new();
for neighbor in neighbors {
if let Some(cid) = &neighbor.cluster_id {
*cluster_counts.entry(cid.0.clone()).or_insert(0) += 1;
}
}
let max_cluster_count = cluster_counts.values().cloned().max().unwrap_or(0);
max_cluster_count as f32 / neighbors.len() as f32
} else {
0.0
};
let taxa: Vec<&str> = neighbors
.iter()
.filter_map(|n| n.recording_metadata.taxon.as_deref())
.collect();
let taxon_coherence = if !taxa.is_empty() {
let mut taxon_counts = std::collections::HashMap::new();
for taxon in &taxa {
*taxon_counts.entry(*taxon).or_insert(0) += 1;
}
let max_taxon_count = taxon_counts.values().cloned().max().unwrap_or(0);
max_taxon_count as f32 / taxa.len() as f32
} else {
0.0
};
EvidenceScores {
neighbor_count: neighbors.len(),
avg_distance,
min_distance,
max_distance,
avg_similarity: similarity,
cluster_coherence,
taxon_coherence,
}
}
}
#[derive(Debug, Clone)]
pub struct RawNeighbor {
pub embedding_id: EmbeddingId,
pub distance: f32,
pub cluster_id: Option<ClusterId>,
pub metadata: Option<RecordingMetadata>,
pub spectrogram_url: Option<String>,
}
impl RawNeighbor {
pub fn new(embedding_id: EmbeddingId, distance: f32) -> Self {
Self {
embedding_id,
distance,
cluster_id: None,
metadata: None,
spectrogram_url: None,
}
}
pub fn with_cluster(mut self, cluster_id: ClusterId) -> Self {
self.cluster_id = Some(cluster_id);
self
}
pub fn with_metadata(mut self, metadata: RecordingMetadata) -> Self {
self.metadata = Some(metadata);
self
}
pub fn with_spectrogram(mut self, url: String) -> Self {
self.spectrogram_url = Some(url);
self
}
}
#[derive(Debug, Clone, Default)]
pub struct EvidenceScores {
pub neighbor_count: usize,
pub avg_distance: f32,
pub min_distance: f32,
pub max_distance: f32,
pub avg_similarity: f32,
pub cluster_coherence: f32,
pub taxon_coherence: f32,
}
impl EvidenceScores {
pub fn overall_strength(&self) -> f32 {
if self.neighbor_count == 0 {
return 0.0;
}
let similarity_weight = 0.4;
let cluster_weight = 0.3;
let taxon_weight = 0.3;
self.avg_similarity * similarity_weight
+ self.cluster_coherence * cluster_weight
+ self.taxon_coherence * taxon_weight
}
pub fn is_strong(&self) -> bool {
self.neighbor_count >= 3 && self.overall_strength() >= 0.6
}
pub fn is_weak(&self) -> bool {
self.neighbor_count < 2 || self.overall_strength() < 0.3
}
}
#[derive(Debug)]
pub struct EvidenceContext {
pub scores: EvidenceScores,
pub unique_taxa: Vec<String>,
pub unique_clusters: Vec<String>,
pub has_sequence: bool,
pub motif: Option<String>,
}
impl EvidenceContext {
pub fn from_evidence(
neighbors: &[NeighborEvidence],
cluster_context: &ClusterContext,
sequence_context: &Option<SequenceContext>,
) -> Self {
let builder = EvidenceBuilder::default_builder();
let scores = builder.aggregate_evidence_scores(neighbors);
let unique_taxa: Vec<String> = neighbors
.iter()
.filter_map(|n| n.recording_metadata.taxon.clone())
.collect::<std::collections::HashSet<_>>()
.into_iter()
.collect();
let mut unique_clusters = Vec::new();
if let Some(label) = &cluster_context.cluster_label {
unique_clusters.push(label.clone());
}
let has_sequence = sequence_context
.as_ref()
.map(|s| s.has_temporal_context())
.unwrap_or(false);
let motif = sequence_context
.as_ref()
.and_then(|s| s.detected_motif.clone());
Self {
scores,
unique_taxa,
unique_clusters,
has_sequence,
motif,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_evidence_builder_collect_neighbors() {
let builder = EvidenceBuilder::default_builder()
.with_max_neighbors(5)
.with_spectrograms(true);
let raw_neighbors = vec![
RawNeighbor::new(EmbeddingId::new("n1"), 0.1)
.with_metadata(RecordingMetadata::new("r1").with_taxon("Species A")),
RawNeighbor::new(EmbeddingId::new("n2"), 0.2)
.with_metadata(RecordingMetadata::new("r2").with_taxon("Species A")),
RawNeighbor::new(EmbeddingId::new("n3"), 0.3)
.with_cluster(ClusterId::new("c1")),
];
let evidence = builder.collect_neighbor_evidence(&raw_neighbors).await.unwrap();
assert_eq!(evidence.len(), 3);
assert_eq!(evidence[0].embedding_id.as_str(), "n1");
assert_eq!(evidence[0].recording_metadata.taxon, Some("Species A".to_string()));
assert!(evidence[2].cluster_id.is_some());
}
#[tokio::test]
async fn test_evidence_builder_distance_filtering() {
let builder = EvidenceBuilder::default_builder()
.with_distance_threshold(0.0, 0.5);
let raw_neighbors = vec![
RawNeighbor::new(EmbeddingId::new("close"), 0.2),
RawNeighbor::new(EmbeddingId::new("far"), 0.8),
];
let evidence = builder.collect_neighbor_evidence(&raw_neighbors).await.unwrap();
assert_eq!(evidence.len(), 1);
assert_eq!(evidence[0].embedding_id.as_str(), "close");
}
#[test]
fn test_evidence_scores_calculation() {
let builder = EvidenceBuilder::default_builder();
let neighbors = vec![
NeighborEvidence::new(
EmbeddingId::new("n1"),
0.1,
RecordingMetadata::new("r1").with_taxon("Species A"),
).with_cluster(ClusterId::new("c1")),
NeighborEvidence::new(
EmbeddingId::new("n2"),
0.2,
RecordingMetadata::new("r2").with_taxon("Species A"),
).with_cluster(ClusterId::new("c1")),
NeighborEvidence::new(
EmbeddingId::new("n3"),
0.3,
RecordingMetadata::new("r3").with_taxon("Species B"),
).with_cluster(ClusterId::new("c2")),
];
let scores = builder.aggregate_evidence_scores(&neighbors);
assert_eq!(scores.neighbor_count, 3);
assert!((scores.avg_distance - 0.2).abs() < 0.001);
assert!((scores.min_distance - 0.1).abs() < 0.001);
assert!((scores.max_distance - 0.3).abs() < 0.001);
assert!(scores.cluster_coherence > 0.0);
assert!(scores.taxon_coherence > 0.0);
}
#[test]
fn test_evidence_context_from_evidence() {
let neighbors = vec![
NeighborEvidence::new(
EmbeddingId::new("n1"),
0.1,
RecordingMetadata::new("r1").with_taxon("Species A"),
),
NeighborEvidence::new(
EmbeddingId::new("n2"),
0.2,
RecordingMetadata::new("r2").with_taxon("Species B"),
),
];
let cluster_context = ClusterContext::new(
Some(ClusterId::new("c1")),
0.9,
0.85,
).with_label("Song Type A");
let sequence_context = Some(SequenceContext::new(
vec![SegmentId::new("s1")],
vec![SegmentId::new("s3")],
).with_motif("ABAB"));
let context = EvidenceContext::from_evidence(
&neighbors,
&cluster_context,
&sequence_context,
);
assert_eq!(context.unique_taxa.len(), 2);
assert_eq!(context.unique_clusters.len(), 1);
assert!(context.has_sequence);
assert_eq!(context.motif, Some("ABAB".to_string()));
}
#[tokio::test]
async fn test_build_sequence_context() {
let builder = EvidenceBuilder::default_builder();
let context = builder
.build_sequence_context(
vec![SegmentId::new("s1"), SegmentId::new("s2")],
vec![SegmentId::new("s4")],
Some("AABB".to_string()),
)
.await
.unwrap();
assert!(context.is_some());
let ctx = context.unwrap();
assert_eq!(ctx.preceding_segments.len(), 2);
assert_eq!(ctx.following_segments.len(), 1);
assert_eq!(ctx.detected_motif, Some("AABB".to_string()));
}
}