#!/usr/bin/env python3
"""Checked binary-classification metrics and calibration walkthrough."""

FN_COST = 5.0
FP_COST = 1.0
REVIEW_COST = 0.2


def metrics(tp, fp, fn, tn):
    total = tp + fp + fn + tn
    precision = tp / (tp + fp)
    recall = tp / (tp + fn)
    specificity = tn / (tn + fp)
    f1 = 2 * precision * recall / (precision + recall)
    reviews = tp + fp
    cost = FN_COST * fn + FP_COST * fp + REVIEW_COST * reviews
    return {
        "accuracy": (tp + tn) / total,
        "precision": precision,
        "recall": recall,
        "specificity": specificity,
        "f1": f1,
        "reviews": reviews,
        "cost": cost,
    }


def calibration_bins(labels, probabilities):
    bins = ((0.0, 0.5), (0.5, 1.0))
    output = []
    for lower, upper in bins:
        members = [
            (label, probability)
            for label, probability in zip(labels, probabilities)
            if lower <= probability < upper or (upper == 1.0 and probability == 1.0)
        ]
        mean_probability = sum(probability for _, probability in members) / len(members)
        observed_rate = sum(label for label, _ in members) / len(members)
        output.append((lower, upper, mean_probability, observed_rate, len(members)))
    return output


def main():
    matrices = {
        "A": (72, 108, 28, 792),
        "B": (60, 40, 40, 860),
    }
    for name, matrix in matrices.items():
        result = metrics(*matrix)
        print(
            f"Schwelle {name}: accuracy={result['accuracy']:.3f}, "
            f"precision={result['precision']:.3f}, recall={result['recall']:.3f}, "
            f"specificity={result['specificity']:.3f}, f1={result['f1']:.3f}, "
            f"reviews={result['reviews']}, cost={result['cost']:.1f}"
        )

    print("Always-negative accuracy: 0.900")
    labels = [0, 0, 1, 1]
    probabilities = [0.1, 0.3, 0.7, 0.9]
    brier = sum((probability - label) ** 2 for label, probability in zip(labels, probabilities)) / len(labels)
    print(f"Brier Loss: {brier:.3f}")
    for lower, upper, mean_probability, observed_rate, count in calibration_bins(labels, probabilities):
        print(
            f"Bin {lower:.1f}-{upper:.1f}: mean_p={mean_probability:.3f}, "
            f"observed={observed_rate:.3f}, n={count}"
        )


if __name__ == "__main__":
    main()
