#![allow(clippy::unwrap_used)]
use std::io::Write;
use hotcoco::types::{Annotation, Category, Dataset, Image, Rle, Segmentation};
use hotcoco::{COCO, mask};
fn gt_dataset() -> Dataset {
Dataset {
info: None,
images: vec![Image {
id: 1,
file_name: "img1.jpg".into(),
height: 20,
width: 20,
..Default::default()
}],
annotations: vec![Annotation {
id: 1,
image_id: 1,
category_id: 1,
bbox: Some([1.0, 1.0, 4.0, 4.0]),
area: Some(16.0),
..Default::default()
}],
categories: vec![Category {
id: 1,
name: "thing".into(),
..Default::default()
}],
licenses: vec![],
}
}
#[test]
fn rle_from_string_unbounded_continuation_errors() {
let s = "P".repeat(20);
assert!(mask::rle_from_string(&s, 10, 10).is_err());
}
#[test]
fn rle_from_string_count_above_u32_max_errors() {
let s = "PPPPPP8";
let res = mask::rle_from_string(s, 10, 10);
let err = res.unwrap_err().to_string();
assert!(err.contains("u32::MAX"), "unexpected error: {err}");
}
#[test]
fn rle_string_roundtrip_still_works() {
let rle = Rle {
h: 100,
w: 100,
counts: vec![10, 20, 30, 40, 50, 60, 9790],
};
let s = mask::rle_to_string(&rle);
let back = mask::rle_from_string(&s, 100, 100).unwrap();
assert_eq!(back.counts, rle.counts);
}
#[test]
fn to_bbox_zero_length_foreground_run() {
let rle = Rle {
h: 5,
w: 5,
counts: vec![5, 0, 20],
};
assert_eq!(mask::to_bbox(&rle), [0.0, 0.0, 0.0, 0.0]);
let rle2 = Rle {
h: 5,
w: 5,
counts: vec![5, 0, 1, 3, 16],
};
assert_eq!(mask::to_bbox(&rle2), [1.0, 1.0, 1.0, 3.0]);
}
#[test]
fn huge_dims_error_instead_of_overflowing() {
let h = 100_000;
let w = 100_000; assert!(mask::fr_poly(&[1.0, 1.0, 5.0, 1.0, 3.0, 4.0], h, w).is_err());
assert!(mask::fr_polys(&[vec![1.0, 1.0, 5.0, 1.0, 3.0, 4.0]], h, w).is_err());
assert!(mask::fr_bbox(&[1.0, 1.0, 2.0, 2.0], h, w).is_err());
assert!(mask::encode(&[0u8; 16], h, w).is_err());
let r = Rle {
h,
w,
counts: vec![0],
};
assert!(mask::merge(&[r.clone(), r], false).is_err());
}
#[test]
fn fr_poly_extreme_coords_no_panic() {
let poly = vec![1e9, 1e9, -1e9, 5.0, 3.0, -1e9];
let rle = mask::fr_poly(&poly, 100, 100).unwrap();
assert!(mask::area(&rle) <= 100 * 100);
let weird = vec![f64::NAN, 2.0, f64::INFINITY, 5.0, 3.0, f64::NEG_INFINITY];
let rle2 = mask::fr_poly(&weird, 100, 100).unwrap();
assert!(mask::area(&rle2) <= 100 * 100);
let tri = vec![2.0, 2.0, 7.0, 2.0, 4.0, 7.0];
let rle3 = mask::fr_poly(&tri, 10, 10).unwrap();
assert_eq!(mask::area(&rle3), 12);
}
#[test]
fn fr_bbox_non_finite_coords_no_panic() {
let rle = mask::fr_bbox(&[f64::NAN, f64::INFINITY, f64::NEG_INFINITY, 1.0], 10, 10).unwrap();
assert_eq!(mask::area(&rle), 0);
}
#[test]
fn encode_length_mismatch_errors() {
assert!(mask::encode(&[0u8; 5], 3, 4).is_err());
assert!(mask::encode(&[0u8; 12], 3, 4).is_ok());
}
#[test]
fn merge_mismatched_dims_errors() {
let a = Rle {
h: 3,
w: 4,
counts: vec![12],
};
let b = Rle {
h: 4,
w: 3,
counts: vec![12],
};
let err = mask::merge(&[a.clone(), b], false).unwrap_err().to_string();
assert!(err.contains("dimensions"), "unexpected error: {err}");
assert!(mask::merge(&[a.clone(), a], false).is_ok());
}
#[test]
fn rle_new_validates_in_release_builds_too() {
assert!(Rle::new(3, 4, vec![6, 6]).is_ok());
assert!(Rle::new(3, 4, vec![6, 7]).is_err());
}
#[test]
fn corners_to_obb_length_checked() {
use hotcoco::geometry::corners_to_obb;
assert!(corners_to_obb(&[0.0; 7]).is_err());
assert!(corners_to_obb(&[0.0; 9]).is_err());
let obb = corners_to_obb(&[0.0, 0.0, 4.0, 0.0, 4.0, 3.0, 0.0, 3.0]).unwrap();
assert!((obb[2] - 4.0).abs() < 1e-9);
assert!((obb[3] - 3.0).abs() < 1e-9);
}
#[test]
fn coco_merge_id_offset_overflow_errors() {
let mut ds1 = gt_dataset();
ds1.images[0].id = u64::MAX;
ds1.annotations[0].image_id = u64::MAX;
let ds2 = gt_dataset();
assert!(COCO::merge(&[&ds1, &ds2]).is_err());
let a = gt_dataset();
let b = gt_dataset();
let merged = COCO::merge(&[&a, &b]).unwrap();
assert_eq!(merged.images.len(), 2);
let mut ids: Vec<u64> = merged.images.iter().map(|i| i.id).collect();
ids.sort_unstable();
ids.dedup();
assert_eq!(ids.len(), 2, "image ids must stay unique after merge");
}
#[test]
fn load_res_segmentation_wins_over_keypoints() {
let gt = COCO::from_dataset(gt_dataset());
let box_rle = mask::fr_bbox(&[2.0, 2.0, 4.0, 4.0], 20, 20).unwrap();
let counts = mask::rle_to_string(&box_rle);
let det = Annotation {
image_id: 1,
category_id: 1,
segmentation: Some(Segmentation::CompressedRle {
size: [20, 20],
counts,
}),
keypoints: Some(vec![0.0, 0.0, 2.0, 10.0, 10.0, 2.0]),
score: Some(0.9),
..Default::default()
};
let res = gt.load_res_anns(vec![det]).unwrap();
let ann = &res.dataset.annotations[0];
assert_eq!(ann.bbox, Some([2.0, 2.0, 4.0, 4.0]));
assert_eq!(ann.area, Some(16.0));
}
#[test]
fn load_res_empty_keypoints_derives_nothing() {
let gt = COCO::from_dataset(gt_dataset());
let det = Annotation {
image_id: 1,
category_id: 1,
keypoints: Some(vec![]),
score: Some(0.5),
..Default::default()
};
let res = gt.load_res_anns(vec![det]).unwrap();
let ann = &res.dataset.annotations[0];
assert_eq!(
ann.bbox, None,
"no bbox may be derived from empty keypoints"
);
assert_eq!(
ann.area, None,
"no area may be derived from empty keypoints"
);
}
#[test]
fn duplicate_ann_ids_warn_and_match_pycocotools() {
let mut ds = gt_dataset();
ds.annotations.push(Annotation {
id: 1, image_id: 1,
category_id: 1,
bbox: Some([5.0, 5.0, 2.0, 2.0]),
area: Some(4.0),
..Default::default()
});
let coco = COCO::from_dataset(ds);
assert!(
coco.load_warnings()
.iter()
.any(|w| w.contains("duplicate annotation id")),
"expected a duplicate-id warning, got {:?}",
coco.load_warnings()
);
let ann = coco.get_ann(1).unwrap();
assert_eq!(ann.bbox, Some([5.0, 5.0, 2.0, 2.0]));
assert_eq!(coco.get_ann_ids_for_img(1), &[1, 1]);
}
#[test]
fn load_res_unknown_image_id_recorded_in_warnings() {
let gt = COCO::from_dataset(gt_dataset());
let det = Annotation {
image_id: 999, category_id: 1,
bbox: Some([0.0, 0.0, 1.0, 1.0]),
score: Some(0.5),
..Default::default()
};
let res = gt.load_res_anns(vec![det]).unwrap();
assert!(
res.load_warnings()
.iter()
.any(|w| w.contains("image_id 999")),
"expected an unknown-image_id warning, got {:?}",
res.load_warnings()
);
let clean = COCO::from_dataset(gt_dataset());
assert!(clean.load_warnings().is_empty());
}
#[test]
fn missing_area_excluded_from_area_range() {
let mut ds = gt_dataset();
ds.annotations.push(Annotation {
id: 2,
image_id: 1,
category_id: 1,
bbox: Some([0.0, 0.0, 3.0, 3.0]),
area: None,
..Default::default()
});
let coco = COCO::from_dataset(ds);
assert_eq!(coco.get_ann_ids(&[], &[], None, None), vec![1, 2]);
assert_eq!(coco.get_ann_ids(&[], &[], Some([0.0, 1e10]), None), vec![1]);
let filtered = coco.filter(None, None, Some([0.0, 1e10]), false);
assert_eq!(filtered.annotations.len(), 1);
assert_eq!(filtered.annotations[0].id, 1);
}
#[test]
fn extra_keys_survive_load_and_filter() {
let json = r#"{
"images": [{"id": 1, "file_name": "a.jpg", "height": 10, "width": 10,
"camera": "rig-3"}],
"annotations": [{"id": 1, "image_id": 1, "category_id": 1,
"bbox": [1.0, 1.0, 2.0, 2.0], "area": 4.0,
"confidence_source": "human", "track_id": 17}],
"categories": [{"id": 1, "name": "thing", "taxonomy_code": "T-9"}]
}"#;
let mut tmp = tempfile::NamedTempFile::new().unwrap();
tmp.write_all(json.as_bytes()).unwrap();
let coco = COCO::new(tmp.path()).unwrap();
assert_eq!(
coco.dataset.images[0].extra.get("camera"),
Some(&serde_json::Value::from("rig-3"))
);
assert_eq!(
coco.dataset.annotations[0].extra.get("track_id"),
Some(&serde_json::Value::from(17))
);
assert_eq!(
coco.dataset.categories[0].extra.get("taxonomy_code"),
Some(&serde_json::Value::from("T-9"))
);
let filtered = coco.filter(Some(&[1]), None, None, false);
assert_eq!(
filtered.annotations[0].extra.get("confidence_source"),
Some(&serde_json::Value::from("human"))
);
let merged = COCO::merge(&[&filtered]).unwrap();
assert_eq!(
merged.images[0].extra.get("camera"),
Some(&serde_json::Value::from("rig-3"))
);
}
#[test]
#[ignore = "needs data/annotations/instances_val2017.json; timing is machine-dependent"]
fn load_timing_val2017() {
let path = std::path::Path::new("../../data/annotations/instances_val2017.json");
if !path.exists() {
eprintln!("val2017 annotations not present; skipping");
return;
}
let _ = COCO::new(path).unwrap();
for i in 0..3 {
let t = std::time::Instant::now();
let coco = COCO::new(path).unwrap();
eprintln!(
"run {i}: loaded {} anns in {:?}",
coco.dataset.annotations.len(),
t.elapsed()
);
}
}