import bisect

MOD = 1_003_443_221


def load_csv(path):
    with open(path) as fin:
        data = fin.readline()
    return [int(token) for token in data.split(",") if token.strip() != ""]


def build_intervals(values):
    positions = {}
    for i, value in enumerate(values):
        positions.setdefault(value, []).append(i)

    intervals = []
    for occ in positions.values():
        assert len(occ) == 2
        intervals.append((occ[0], occ[1]))  # (left, right)

    intervals.sort(key=lambda iv: iv[0])
    for i in range(1, len(intervals)):
        assert intervals[i - 1][0] < intervals[i][0]
    return intervals


def connectivity_number(values, mod):
    intervals = build_intervals(values)
    n = len(intervals)

    left = [iv[0] for iv in intervals]
    right = [iv[1] for iv in intervals]

    # next[i] = first interval whose left endpoint lies strictly after right[i]
    nxt = [bisect.bisect_right(left, right[i]) for i in range(n)]

    by_right = sorted(range(n), key=lambda i: right[i])

    inside = [0] * n
    ways = [1] * (n + 1)
    delta = [0] * (n + 1)
    active = [0] * n

    for i in by_right:
        inside[i] = ways[i + 1]

        inc = inside[i] * ways[nxt[i]] % mod
        delta[i] = inc
        ways[i] = (ways[i] + inc) % mod

        for p in range(i - 1, -1, -1):
            d = delta[p + 1]
            if active[p] and nxt[p] <= i:
                d = (d + inside[p] * delta[nxt[p]]) % mod
            delta[p] = d
            ways[p] = (ways[p] + d) % mod

        active[i] = 1

    return ways[0]


def crosses(a, b):
    if b[0] < a[0]:
        a, b = b, a
    return a[0] < b[0] and b[0] < a[1] and a[1] < b[1]


def brute_connectivity(values):
    intervals = build_intervals(values)
    n = len(intervals)
    assert n <= 20

    total = 0
    for mask in range(1 << n):
        ok = True
        for i in range(n):
            if not (mask >> i) & 1:
                continue
            for j in range(i + 1, n):
                if (mask >> j) & 1 and crosses(intervals[i], intervals[j]):
                    ok = False
                    break
            if not ok:
                break
        if ok:
            total += 1
    return total


def run_checkpoints():
    cases = [
        [0, 1, 0, 1],
        [0, 0, 1, 2, 2, 1],
        [0, 1, 2, 1, 0, 2],
        [0, 1, 2, 2, 1, 0],
        [0, 1, 2, 3, 1, 4, 0, 5, 4, 2, 6, 7, 3, 8, 6, 5, 9, 8, 9, 7],
    ]
    expected = [3, 8, 5, 8, 86]
    for case, want in zip(cases, expected):
        brute = brute_connectivity(case)
        assert brute == want
        assert connectivity_number(case, MOD) == brute % MOD


if __name__ == "__main__":
    run_checkpoints()

    values = load_csv("resources/documents/1001_input.txt")
    assert len(values) == 40_000

    print(connectivity_number(values, MOD))
