227 lines
7.5 KiB
Rust
227 lines
7.5 KiB
Rust
//! Target classification.
|
|
//!
|
|
//! Runs the classifier registry over a target's artifacts and their ingested
|
|
//! working paths, then merges and ranks the verdicts into a [`Classification`].
|
|
//! The registry is the heuristic classifier (artifact kinds + source markers)
|
|
//! plus the tramiton firmware detector (behind a [`FirmwareDetector`] port).
|
|
|
|
mod firmware;
|
|
mod language;
|
|
|
|
pub use firmware::{
|
|
FirmwareDetection, FirmwareDetector, FirmwareTarget, MockFirmwareDetector, TramitonNative,
|
|
};
|
|
pub use language::HeuristicClassifier;
|
|
|
|
use std::collections::HashMap;
|
|
use std::path::PathBuf;
|
|
|
|
use compliance_core::error::CoreError;
|
|
use compliance_core::models::{
|
|
ArtifactKind, Classification, DetectedFact, OnboardedTarget, TargetType, TargetTypeCandidate,
|
|
};
|
|
use compliance_core::traits::{ClassificationInput, ClassifierVerdict, TargetClassifier};
|
|
|
|
use firmware::detection_to_verdict;
|
|
|
|
/// Classify a target from its artifacts and their ingested working paths, using
|
|
/// the heuristic classifier plus the tramiton firmware detector. Verdicts are
|
|
/// merged (max confidence per target type) and ranked into a [`Classification`].
|
|
pub async fn classify_target<D: FirmwareDetector>(
|
|
target: &OnboardedTarget,
|
|
working_paths: &HashMap<String, PathBuf>,
|
|
firmware_detector: &D,
|
|
) -> Result<Classification, CoreError> {
|
|
let input = ClassificationInput {
|
|
artifacts: &target.artifacts,
|
|
working_paths,
|
|
description: target.description.as_deref(),
|
|
};
|
|
|
|
let mut verdicts = Vec::new();
|
|
let mut detected_by = Vec::new();
|
|
|
|
let heuristic = HeuristicClassifier.classify(&input).await?;
|
|
if !heuristic.is_empty() {
|
|
detected_by.push("heuristic".to_string());
|
|
}
|
|
verdicts.extend(heuristic);
|
|
|
|
// Tramiton firmware detection over firmware / code working paths.
|
|
let mut tramiton_used = false;
|
|
for artifact in &target.artifacts {
|
|
if !matches!(
|
|
artifact.kind,
|
|
ArtifactKind::FirmwareImage | ArtifactKind::GitRepo | ArtifactKind::SourceArchive
|
|
) {
|
|
continue;
|
|
}
|
|
let Some(path) = working_paths.get(&artifact.id) else {
|
|
continue;
|
|
};
|
|
if let Some(detection) = firmware_detector.detect(path).await? {
|
|
verdicts.push(detection_to_verdict(&detection));
|
|
tramiton_used = true;
|
|
}
|
|
}
|
|
if tramiton_used {
|
|
detected_by.push("tramiton".to_string());
|
|
}
|
|
|
|
Ok(rank(verdicts, detected_by, target.target_type))
|
|
}
|
|
|
|
/// Merge verdicts by target type (keeping the max confidence and its rationale),
|
|
/// dedupe facts, rank by descending confidence, and assemble a [`Classification`].
|
|
/// Falls back to the declared type when no verdict is produced.
|
|
fn rank(
|
|
verdicts: Vec<ClassifierVerdict>,
|
|
detected_by: Vec<String>,
|
|
fallback: TargetType,
|
|
) -> Classification {
|
|
let mut best: HashMap<TargetType, (f32, String)> = HashMap::new();
|
|
let mut facts: Vec<DetectedFact> = Vec::new();
|
|
for verdict in verdicts {
|
|
for fact in verdict.facts {
|
|
if !facts
|
|
.iter()
|
|
.any(|e| e.key == fact.key && e.value == fact.value)
|
|
{
|
|
facts.push(fact);
|
|
}
|
|
}
|
|
let entry = best
|
|
.entry(verdict.target_type)
|
|
.or_insert((0.0, String::new()));
|
|
if verdict.confidence > entry.0 {
|
|
*entry = (verdict.confidence, verdict.rationale);
|
|
}
|
|
}
|
|
|
|
let mut candidates: Vec<TargetTypeCandidate> = best
|
|
.into_iter()
|
|
.map(
|
|
|(target_type, (confidence, rationale))| TargetTypeCandidate {
|
|
target_type,
|
|
confidence,
|
|
rationale,
|
|
},
|
|
)
|
|
.collect();
|
|
// Descending confidence; ties broken by type name for deterministic ordering.
|
|
candidates.sort_by(|a, b| {
|
|
b.confidence
|
|
.partial_cmp(&a.confidence)
|
|
.unwrap_or(std::cmp::Ordering::Equal)
|
|
.then_with(|| a.target_type.to_string().cmp(&b.target_type.to_string()))
|
|
});
|
|
|
|
let suggested = candidates
|
|
.first()
|
|
.map(|c| c.target_type)
|
|
.unwrap_or(fallback);
|
|
|
|
Classification {
|
|
suggested,
|
|
candidates,
|
|
facts,
|
|
detected_by,
|
|
detected_at: chrono::Utc::now(),
|
|
confirmed: false,
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
#[allow(clippy::expect_used, clippy::unwrap_used)]
|
|
mod tests {
|
|
use super::*;
|
|
use compliance_core::models::Artifact;
|
|
use std::fs;
|
|
use std::path::Path;
|
|
|
|
struct Scratch(PathBuf);
|
|
impl Scratch {
|
|
fn new() -> Self {
|
|
let p = std::env::temp_dir().join(format!("cs-classify-mod-{}", uuid::Uuid::new_v4()));
|
|
fs::create_dir_all(&p).expect("mkdir");
|
|
Self(p)
|
|
}
|
|
}
|
|
impl Drop for Scratch {
|
|
fn drop(&mut self) {
|
|
let _ = fs::remove_dir_all(&self.0);
|
|
}
|
|
}
|
|
|
|
fn no_firmware() -> MockFirmwareDetector {
|
|
MockFirmwareDetector { detection: None }
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn backend_repo_classifies_as_backend() {
|
|
let scratch = Scratch::new();
|
|
fs::write(scratch.0.join("go.mod"), "module x").unwrap();
|
|
|
|
let artifact = Artifact::git_repo("https://git/x", "main");
|
|
let mut wp = HashMap::new();
|
|
wp.insert(artifact.id.clone(), scratch.0.clone());
|
|
let mut target = OnboardedTarget::new("x".to_string(), TargetType::WebApp);
|
|
target.artifacts.push(artifact);
|
|
|
|
let c = classify_target(&target, &wp, &no_firmware())
|
|
.await
|
|
.expect("classify");
|
|
assert_eq!(c.suggested, TargetType::BackendService);
|
|
assert!(c.detected_by.contains(&"heuristic".to_string()));
|
|
assert!(!c.confirmed);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn firmware_detector_verdict_ranks_top() {
|
|
let scratch = Scratch::new();
|
|
fs::write(scratch.0.join("fw.bin"), b"x").unwrap();
|
|
|
|
let artifact =
|
|
Artifact::firmware_image(scratch.0.join("fw.bin").to_string_lossy().to_string());
|
|
let mut wp = HashMap::new();
|
|
wp.insert(artifact.id.clone(), scratch.0.clone());
|
|
let mut target = OnboardedTarget::new("fw".to_string(), TargetType::FirmwareBareMetal);
|
|
target.artifacts.push(artifact);
|
|
|
|
let detector = MockFirmwareDetector {
|
|
detection: Some(FirmwareDetection {
|
|
provider: "zephyr".to_string(),
|
|
confidence: "high".to_string(),
|
|
build_system: "zephyr".to_string(),
|
|
framework: Some("zephyr".to_string()),
|
|
target: FirmwareTarget {
|
|
mcu: Some("nrf52840".to_string()),
|
|
..Default::default()
|
|
},
|
|
gaps: vec![],
|
|
}),
|
|
};
|
|
|
|
let c = classify_target(&target, &wp, &detector)
|
|
.await
|
|
.expect("classify");
|
|
// tramiton's high-confidence RTOS verdict beats the weak firmware prior.
|
|
assert_eq!(c.suggested, TargetType::FirmwareRtos);
|
|
assert!(c.detected_by.contains(&"tramiton".to_string()));
|
|
assert!(c.facts.iter().any(|f| f.key == "mcu"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn no_signal_falls_back_to_declared_type() {
|
|
let scratch = Scratch::new();
|
|
let _ = Path::new(&scratch.0);
|
|
let target = OnboardedTarget::new("empty".to_string(), TargetType::DesktopApp);
|
|
let wp = HashMap::new();
|
|
let c = classify_target(&target, &wp, &no_firmware())
|
|
.await
|
|
.expect("classify");
|
|
assert_eq!(c.suggested, TargetType::DesktopApp);
|
|
assert!(c.candidates.is_empty());
|
|
}
|
|
}
|