"""Independent segmented sieve and event tally, with direct reciprocal weights."""
import json
import math
import pathlib

import mpmath as mp

LIMIT = 100_000_000
SEGMENT = 1_000_000
mp.mp.dps = 50


def primes_up_to(limit):
    small_limit = math.isqrt(limit)
    small = bytearray(b"\x01") * (small_limit + 1)
    small[:2] = b"\x00\x00"
    for p in range(2, math.isqrt(small_limit) + 1):
        if small[p]:
            small[p*p::p] = b"\x00" * (((small_limit - p*p) // p) + 1)
    divisors = [p for p in range(2, small_limit + 1) if small[p]]
    for low in range(2, limit + 1, SEGMENT):
        high = min(limit, low + SEGMENT - 1)
        flags = bytearray(b"\x01") * (high - low + 1)
        for p in divisors:
            if p*p > high:
                break
            start = max(p*p, ((low + p - 1) // p) * p)
            if start <= high:
                flags[start-low::p] = b"\x00" * (((high - start) // p) + 1)
        for offset, flag in enumerate(flags):
            if flag:
                yield low + offset


def tally(primes, limit):
    states = {
        4: {"winner": 3, "loser": 1, "prime_counts": {3: 0, 1: 0}, "durations": {"winner": 0, "tie": 0, "loser": 0}, "negative_intervals": [], "tie_intervals": []},
        3: {"winner": 2, "loser": 1, "prime_counts": {2: 0, 1: 0}, "durations": {"winner": 0, "tie": 0, "loser": 0}, "negative_intervals": [], "tie_intervals": []},
    }

    def consume(first, last):
        if last < first:
            return
        for state in states.values():
            difference = state["prime_counts"][state["winner"]] - state["prime_counts"][state["loser"]]
            side = "winner" if difference > 0 else "loser" if difference < 0 else "tie"
            state["durations"][side] += last - first + 1
            intervals = state["negative_intervals"] if side == "loser" else state["tie_intervals"] if side == "tie" else None
            if intervals is not None:
                if intervals and intervals[-1][1] + 1 == first:
                    intervals[-1][1] = last
                else:
                    intervals.append([first, last])

    previous = 1
    for prime in primes:
        consume(previous, prime - 1)
        for modulus, state in states.items():
            residue = prime % modulus
            if residue in state["prime_counts"]:
                state["prime_counts"][residue] += 1
        previous = prime
    consume(previous, limit)
    return states


def prefix_control():
    limit = 30_000
    segmented = list(primes_up_to(limit))
    trial = [n for n in range(2, limit + 1) if all(n % d for d in range(2, math.isqrt(n) + 1))]
    assert segmented == trial
    event = tally(segmented, limit)
    prime_set = set(trial)
    for modulus, state in event.items():
        counts = {state["winner"]: 0, state["loser"]: 0}
        durations = {"winner": 0, "tie": 0, "loser": 0}
        for x in range(1, limit + 1):
            if x in prime_set and x % modulus in counts:
                counts[x % modulus] += 1
            difference = counts[state["winner"]] - counts[state["loser"]]
            durations["winner" if difference > 0 else "loser" if difference < 0 else "tie"] += 1
        assert durations == state["durations"]
        assert counts == state["prime_counts"]
    return {"prefix_limit": limit, "segmented_sieve_equals_trial_division": True, "event_durations_equal_direct_integer_scan": True}


def main():
    control = prefix_control()
    states = tally(primes_up_to(LIMIT), LIMIT)
    results = {"n": LIMIT}
    weighted_audit = {}
    for modulus, state in states.items():
        intervals = state["negative_intervals"]
        longest = max(intervals, key=lambda span: span[1] - span[0], default=None)
        # Directly sum reciprocals only on the negative intervals, rather than
        # subtracting approximate harmonic numbers at each interval endpoint.
        numerator = mp.fsum(mp.mpf(1) / x for first, last in intervals for x in range(first, last + 1))
        denominator = mp.harmonic(LIMIT)
        ratio = numerator / denominator
        weighted_audit[f"mod{modulus}"] = {"numerator": str(numerator), "harmonic_denominator": str(denominator), "unrounded_share": str(ratio)}
        winner = state["winner"]
        loser = state["loser"]
        results[f"mod{modulus}"] = {
            "winner_class": winner, "loser_class": loser,
            "primes_in_winner_class": state["prime_counts"][winner],
            "primes_in_loser_class": state["prime_counts"][loser],
            "lead_at_n": state["prime_counts"][winner] - state["prime_counts"][loser],
            "integers_winner_leads": state["durations"]["winner"],
            "integers_tied": state["durations"]["tie"],
            "integers_loser_leads": state["durations"]["loser"],
            "share_loser_leads": round(state["durations"]["loser"] / LIMIT, 9),
            "loser_stretches": len(intervals),
            "first_loser_lead": intervals[0][0] if intervals else None,
            "last_loser_lead": intervals[-1][1] if intervals else None,
            "longest_loser_stretch": {"first": longest[0], "last": longest[1], "length": longest[1] - longest[0] + 1} if longest else None,
            "loser_stretch_starts": [span[0] for span in intervals[:20]],
            "last_tie": state["tie_intervals"][-1][1] if state["tie_intervals"] else None,
            "log_weighted_share_loser_leads": round(float(ratio), 9),
        }
        assert sum(state["durations"].values()) == LIMIT
    pathlib.Path("results/independent.json").write_text(json.dumps(results, indent=2) + "\n")
    pathlib.Path("results/independent-audit.json").write_text(json.dumps({**control, "weighted_audit": weighted_audit}, indent=2) + "\n")


if __name__ == "__main__":
    main()
