#!/usr/bin/env python3
"""Exact synthetic correlated-vote example; no market data or LLM calls.

Each hypothetical voter is correct with probability p. With probability rho,
all voters share one Bernoulli(p) outcome. Otherwise their outcomes are
independent Bernoulli(p) draws. The resulting pairwise correlation is rho.
Enumerating the distribution avoids simulation noise and external dependencies.
"""

from fractions import Fraction
from itertools import product


def distribution(n: int, p: Fraction, rho: Fraction):
    if n < 3 or n % 2 == 0:
        raise ValueError("n must be odd and at least 3")
    if not 0 < p < 1 or not 0 <= rho <= 1:
        raise ValueError("require 0 < p < 1 and 0 <= rho <= 1")
    for votes in product((0, 1), repeat=n):
        correct = sum(votes)
        probability = (1 - rho) * p**correct * (1 - p) ** (n - correct)
        if correct == n:
            probability += rho * p
        elif correct == 0:
            probability += rho * (1 - p)
        yield votes, probability


def summarize(n: int, p: Fraction, rho: Fraction):
    rows = tuple(distribution(n, p, rho))
    assert sum(weight for _, weight in rows) == 1
    means = [sum(v[i] * w for v, w in rows) for i in range(n)]
    assert all(mean == p for mean in means)
    variance = p * (1 - p)
    for i in range(n):
        for j in range(i + 1, n):
            covariance = sum((v[i] - p) * (v[j] - p) * w for v, w in rows)
            assert covariance / variance == rho
    mean_variance = sum((Fraction(sum(v), n) - p) ** 2 * w for v, w in rows)
    effective_n = Fraction(n, 1) / (1 + (n - 1) * rho)
    assert mean_variance == variance / effective_n
    majority = sum(w for v, w in rows if sum(v) > n // 2)
    return effective_n, majority


def main():
    n, p = 5, Fraction(3, 5)
    expected_majorities = (
        Fraction(2133, 3125),
        Fraction(20556, 31250),
        Fraction(9891, 15625),
        Fraction(3, 5),
    )
    print("SYNTHETIC EXACT DISTRIBUTION: no LLM calls, no trading returns")
    print(f"n={n}; individual_accuracy={float(p):.6f}; states={2**n}")
    print("rho  effective_n  majority_accuracy")
    for rho, expected in zip(
        (Fraction(0), Fraction(3, 10), Fraction(3, 5), Fraction(1)),
        expected_majorities,
    ):
        effective_n, majority = summarize(n, p, rho)
        assert majority == expected
        print(f"{float(rho):.1f}  {float(effective_n):.6f}     {float(majority):.6f}")
    assert Fraction(10) / (1 + 9 * Fraction(3, 5)) == Fraction(25, 16)
    print("PASS: probability mass, marginals, pairwise correlations, variance identity, majority values")


if __name__ == "__main__":
    main()
