use crate::domain::metrics::{constants, MetricsError};
use nutype::nutype;
#[allow(unused_imports)] use serde::{Deserialize, Serialize};
#[nutype(
validate(finite, greater_or_equal = 0.0, less_or_equal = 1.0),
derive(Debug, Clone, Copy, PartialEq, PartialOrd, Serialize, Deserialize)
)]
pub struct FScore(f64);
impl Eq for FScore {}
impl FScore {
pub fn perfect() -> Self {
Self::try_new(1.0).unwrap()
}
pub fn zero() -> Self {
Self::try_new(0.0).unwrap()
}
pub fn from_precision_recall(
precision: Precision,
recall: Recall,
) -> Result<Self, MetricsError> {
let p = precision.into_inner();
let r = recall.into_inner();
if p + r == 0.0 {
return Ok(Self::zero());
}
let f_score = constants::calculation::F1_MULTIPLIER * (p * r) / (p + r);
Self::try_new(f_score).map_err(|_| MetricsError::InvalidValue(f_score))
}
pub fn from_precision_recall_beta(
precision: Precision,
recall: Recall,
beta: Beta,
) -> Result<Self, MetricsError> {
let p = precision.into_inner();
let r = recall.into_inner();
if p + r == 0.0 {
return Ok(Self::zero());
}
let f_score = Self::calculate_f_beta_formula(p, r, beta.into_inner());
Self::try_new(f_score).map_err(|_| MetricsError::InvalidValue(f_score))
}
fn calculate_f_beta_formula(precision: f64, recall: f64, beta: f64) -> f64 {
let beta_squared = beta * beta;
let numerator = (1.0 + beta_squared) * (precision * recall);
let denominator = (beta_squared * precision) + recall;
numerator / denominator
}
}
#[nutype(
validate(finite, greater_or_equal = 0.0, less_or_equal = 1.0),
derive(Debug, Clone, Copy, PartialEq, PartialOrd, Serialize, Deserialize)
)]
pub struct Precision(f64);
impl Eq for Precision {}
impl Precision {
pub fn perfect() -> Self {
Self::try_new(1.0).unwrap()
}
pub fn zero() -> Self {
Self::try_new(0.0).unwrap()
}
}
#[nutype(
validate(finite, greater_or_equal = 0.0, less_or_equal = 1.0),
derive(Debug, Clone, Copy, PartialEq, PartialOrd, Serialize, Deserialize)
)]
pub struct Recall(f64);
impl Eq for Recall {}
impl Recall {
pub fn perfect() -> Self {
Self::try_new(1.0).unwrap()
}
pub fn zero() -> Self {
Self::try_new(0.0).unwrap()
}
}
#[nutype(
validate(finite, greater = 0.0, less_or_equal = 10.0),
derive(Debug, Clone, Copy, PartialEq, PartialOrd, Serialize, Deserialize)
)]
pub struct Beta(f64);
impl Beta {
pub fn f1() -> Self {
Self::try_new(1.0).unwrap()
}
pub fn f2() -> Self {
Self::try_new(2.0).unwrap()
}
pub fn f05() -> Self {
Self::try_new(0.5).unwrap()
}
}
#[nutype(
validate(finite, greater = 0.0, less = 1.0),
derive(Debug, Clone, Copy, PartialEq, PartialOrd, Serialize, Deserialize)
)]
pub struct ConfidenceLevel(f64);
impl ConfidenceLevel {
pub fn ninety_five_percent() -> Self {
Self::try_new(0.95).unwrap()
}
pub fn ninety_nine_percent() -> Self {
Self::try_new(0.99).unwrap()
}
pub fn ninety_percent() -> Self {
Self::try_new(0.90).unwrap()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_f_score_validation() {
assert!(FScore::try_new(0.0).is_ok());
assert!(FScore::try_new(0.5).is_ok());
assert!(FScore::try_new(1.0).is_ok());
assert!(FScore::try_new(-0.1).is_err());
assert!(FScore::try_new(1.1).is_err());
assert!(FScore::try_new(f64::NAN).is_err());
assert!(FScore::try_new(f64::INFINITY).is_err());
}
#[test]
fn test_precision_recall_validation() {
assert!(Precision::try_new(0.8).is_ok());
assert!(Recall::try_new(0.6).is_ok());
assert!(Precision::try_new(-0.1).is_err());
assert!(Recall::try_new(1.1).is_err());
}
#[test]
fn test_f_score_calculation() {
let precision = Precision::try_new(0.8).unwrap();
let recall = Recall::try_new(0.6).unwrap();
let f_score = FScore::from_precision_recall(precision, recall).unwrap();
assert!((f_score.into_inner() - 0.6857142857142857).abs() < 1e-10);
}
#[test]
fn test_f_beta_calculation() {
let precision = Precision::try_new(0.8).unwrap();
let recall = Recall::try_new(0.6).unwrap();
let beta = Beta::try_new(2.0).unwrap();
let f_score = FScore::from_precision_recall_beta(precision, recall, beta).unwrap();
assert!((f_score.into_inner() - 0.631_578_947_368_421).abs() < 1e-10);
}
#[test]
fn test_edge_cases() {
let zero_precision = Precision::zero();
let zero_recall = Recall::zero();
let f_score = FScore::from_precision_recall(zero_precision, zero_recall).unwrap();
assert_eq!(f_score.into_inner(), 0.0);
let perfect_precision = Precision::perfect();
let perfect_recall = Recall::perfect();
let f_score = FScore::from_precision_recall(perfect_precision, perfect_recall).unwrap();
assert_eq!(f_score.into_inner(), 1.0);
}
#[test]
fn test_beta_validation() {
assert!(Beta::try_new(0.1).is_ok());
assert!(Beta::try_new(1.0).is_ok());
assert!(Beta::try_new(2.0).is_ok());
assert!(Beta::try_new(10.0).is_ok());
assert!(Beta::try_new(0.0).is_err());
assert!(Beta::try_new(-1.0).is_err());
assert!(Beta::try_new(11.0).is_err());
}
#[test]
fn test_confidence_level_validation() {
assert!(ConfidenceLevel::try_new(0.95).is_ok());
assert!(ConfidenceLevel::try_new(0.99).is_ok());
assert!(ConfidenceLevel::try_new(0.01).is_ok());
assert!(ConfidenceLevel::try_new(0.0).is_err());
assert!(ConfidenceLevel::try_new(1.0).is_err());
assert!(ConfidenceLevel::try_new(1.1).is_err());
}
#[test]
fn test_convenience_constructors() {
assert_eq!(FScore::perfect().into_inner(), 1.0);
assert_eq!(FScore::zero().into_inner(), 0.0);
assert_eq!(Precision::perfect().into_inner(), 1.0);
assert_eq!(Recall::zero().into_inner(), 0.0);
assert_eq!(Beta::f1().into_inner(), 1.0);
assert_eq!(Beta::f2().into_inner(), 2.0);
assert_eq!(ConfidenceLevel::ninety_five_percent().into_inner(), 0.95);
}
}