use plotters::coord::Shift;
use plotters::prelude::*;
use plotters_statistical::{
BoxPlot, BoxPlotSeries, PrecisionRecallCurve, RegularizationPath, ResidualPlot, RocCurve,
StatsError, ViolinPlot,
};
fn render<F>(f: F) -> String
where
F: FnOnce(&DrawingArea<SVGBackend, Shift>) -> Result<(), Box<dyn std::error::Error>>,
{
let mut buf = String::new();
{
let root = SVGBackend::with_string(&mut buf, (400, 300)).into_drawing_area();
root.fill(&WHITE).unwrap();
f(&root).unwrap();
root.present().unwrap();
}
buf
}
fn has_shape(svg: &str) -> bool {
[
"<rect",
"<circle",
"<polygon",
"<polyline",
"<path",
"<line",
]
.iter()
.any(|tag| svg.contains(tag))
}
fn assert_well_formed(svg: &str) {
assert!(svg.contains("<svg"), "missing <svg> root");
assert!(svg.contains("</svg>"), "missing </svg> close");
assert!(has_shape(svg), "no drawn content — silent empty render?");
}
#[test]
fn renders_box_plot() {
let svg = render(|root| {
let mut chart = ChartBuilder::on(root).build_cartesian_2d(0.5f64..2.5f64, 0f64..12f64)?;
chart.draw_series(BoxPlotSeries::from_samples(vec![
(1.0f64, vec![1.0, 2.0, 3.0, 4.0, 10.0]),
(2.0f64, vec![2.0, 3.0, 4.0, 5.0]),
])?)?;
Ok(())
});
assert_well_formed(&svg);
}
#[test]
fn renders_violin_plot() {
let svg = render(|root| {
let mut chart = ChartBuilder::on(root).build_cartesian_2d(0.5f64..1.5f64, 0f64..12f64)?;
let data: Vec<f64> = (0..80)
.map(|i| 6.0 + (i as f64 * 0.1).sin() * 2.0)
.collect();
chart.draw_series(std::iter::once(
ViolinPlot::vertical(1.0f64, &data)?.show_box(true),
))?;
Ok(())
});
assert_well_formed(&svg);
}
#[test]
fn renders_roc_curve() {
let svg = render(|root| {
let mut chart = ChartBuilder::on(root).build_cartesian_2d(0f64..1f64, 0f64..1f64)?;
let scores = [0.1, 0.4, 0.35, 0.8, 0.6, 0.2];
let labels = [false, false, true, true, true, false];
chart.draw_series(std::iter::once(
RocCurve::from_scores(&scores, &labels)?
.with_baseline()
.shade_area(true),
))?;
Ok(())
});
assert_well_formed(&svg);
}
#[test]
fn renders_precision_recall_curve() {
let svg = render(|root| {
let mut chart = ChartBuilder::on(root).build_cartesian_2d(0f64..1f64, 0f64..1.05f64)?;
let scores = [0.1, 0.4, 0.35, 0.8, 0.6, 0.2];
let labels = [false, false, true, true, true, false];
chart.draw_series(std::iter::once(
PrecisionRecallCurve::from_scores(&scores, &labels)?.with_baseline(),
))?;
Ok(())
});
assert_well_formed(&svg);
}
#[test]
fn renders_regularization_path() {
let svg = render(|root| {
let mut chart = ChartBuilder::on(root).build_cartesian_2d(0f64..1f64, -1f64..3f64)?;
let strengths = vec![0.1, 0.3, 0.5, 0.7, 0.9];
let coefs = vec![
vec![3.0, 2.0],
vec![2.0, 1.0],
vec![1.0, 0.0],
vec![0.0, -0.5],
vec![0.0, -1.0],
];
chart.draw_series(RegularizationPath::new(&strengths, &coefs)?)?;
Ok(())
});
assert_well_formed(&svg);
}
#[test]
fn renders_residual_plot() {
let svg = render(|root| {
let mut chart = ChartBuilder::on(root).build_cartesian_2d(0f64..10f64, -5f64..5f64)?;
let fitted: Vec<f64> = (0..50).map(|i| i as f64 / 5.0).collect();
let residuals: Vec<f64> = fitted.iter().map(|x| 0.3 * (x - 5.0)).collect();
chart.draw_series(std::iter::once(
ResidualPlot::from_residuals(&fitted, &residuals)?.trend(true),
))?;
Ok(())
});
assert_well_formed(&svg);
}
#[test]
fn renders_single_point_box() {
let svg = render(|root| {
let mut chart = ChartBuilder::on(root).build_cartesian_2d(0.5f64..1.5f64, 0f64..10f64)?;
chart.draw_series(std::iter::once(BoxPlot::vertical(1.0f64, &[5.0])?))?;
Ok(())
});
assert_well_formed(&svg);
}
#[test]
fn renders_all_identical_box() {
let svg = render(|root| {
let mut chart = ChartBuilder::on(root).build_cartesian_2d(0.5f64..1.5f64, 0f64..10f64)?;
chart.draw_series(std::iter::once(BoxPlot::vertical(
1.0f64,
&[3.0, 3.0, 3.0],
)?))?;
Ok(())
});
assert_well_formed(&svg);
}
#[test]
fn empty_group_errors_before_render() {
let err = BoxPlotSeries::from_samples(vec![(1.0f64, Vec::<f64>::new())]).unwrap_err();
assert_eq!(err, StatsError::EmptyInput);
}
#[test]
fn single_class_roc_errors_before_render() {
let err = RocCurve::from_scores(&[0.1, 0.9], &[true, true]).unwrap_err();
assert_eq!(err, StatsError::NoNegativeLabels);
}
#[test]
fn degenerate_violin_errors_before_render() {
let err = ViolinPlot::vertical(1.0f64, &[5.0]).unwrap_err();
assert_eq!(err, StatsError::InvalidBandwidth);
}