import java.io.BufferedReader;
import java.io.FileReader;
import java.io.IOException;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.HashMap;
import java.util.List;
import java.util.Map;

public class Euler1001 {
    static final long MOD = 1_003_443_221L;

    static int[] loadCsv(String path) throws IOException {
        String data;
        try (BufferedReader reader = new BufferedReader(new FileReader(path))) {
            data = reader.readLine();
        }
        List<Integer> values = new ArrayList<>();
        for (String token : data.split(",")) {
            if (!token.trim().isEmpty()) {
                values.add(Integer.parseInt(token.trim()));
            }
        }
        int[] result = new int[values.size()];
        for (int i = 0; i < result.length; ++i) {
            result[i] = values.get(i);
        }
        return result;
    }

    // intervals[i] = { left, right }, sorted by left
    static int[][] buildIntervals(int[] values) {
        Map<Integer, int[]> positions = new HashMap<>();
        Map<Integer, Integer> seen = new HashMap<>();
        for (int i = 0; i < values.length; ++i) {
            int v = values[i];
            int c = seen.getOrDefault(v, 0);
            if (c == 0) {
                positions.put(v, new int[] { i, -1 });
            } else {
                assert c == 1;
                positions.get(v)[1] = i;
            }
            seen.put(v, c + 1);
        }

        List<int[]> intervals = new ArrayList<>();
        for (int[] pair : positions.values()) {
            assert pair[1] != -1;
            intervals.add(pair);
        }
        intervals.sort((a, b) -> Integer.compare(a[0], b[0]));
        for (int i = 1; i < intervals.size(); ++i) {
            assert intervals.get(i - 1)[0] < intervals.get(i)[0];
        }
        return intervals.toArray(new int[0][]);
    }

    // first index whose left endpoint is strictly greater than key (upper_bound)
    static int upperBound(int[] sortedLeft, int key) {
        int lo = 0;
        int hi = sortedLeft.length;
        while (lo < hi) {
            int mid = (lo + hi) >>> 1;
            if (sortedLeft[mid] <= key) {
                lo = mid + 1;
            } else {
                hi = mid;
            }
        }
        return lo;
    }

    static long mulMod(long a, long b, long mod) {
        return a * b % mod;
    }

    static long addMod(long a, long b, long mod) {
        long s = a + b;
        return s >= mod ? s - mod : s;
    }

    static long connectivityNumber(int[] values, long mod) {
        int[][] intervals = buildIntervals(values);
        int n = intervals.length;

        int[] left = new int[n];
        int[] right = new int[n];
        for (int i = 0; i < n; ++i) {
            left[i] = intervals[i][0];
            right[i] = intervals[i][1];
        }

        int[] next = new int[n];
        for (int i = 0; i < n; ++i) {
            next[i] = upperBound(left, right[i]);
        }

        Integer[] byRight = new Integer[n];
        for (int i = 0; i < n; ++i) {
            byRight[i] = i;
        }
        Arrays.sort(byRight, (a, b) -> Integer.compare(right[a], right[b]));

        long[] inside = new long[n];
        long[] ways = new long[n + 1];
        Arrays.fill(ways, 1L);
        long[] delta = new long[n + 1];
        boolean[] active = new boolean[n];

        for (int i : byRight) {
            inside[i] = ways[i + 1];

            long inc = mulMod(inside[i], ways[next[i]], mod);
            delta[i] = inc;
            ways[i] = addMod(ways[i], inc, mod);

            for (int p = i - 1; p >= 0; --p) {
                long d = delta[p + 1];
                if (active[p] && next[p] <= i) {
                    d = addMod(d, mulMod(inside[p], delta[next[p]], mod), mod);
                }
                delta[p] = d;
                ways[p] = addMod(ways[p], d, mod);
            }

            active[i] = true;
        }

        return ways[0];
    }

    static boolean crosses(int[] a, int[] b) {
        int al = a[0], ar = a[1], bl = b[0], br = b[1];
        if (bl < al) {
            int tl = al, tr = ar;
            al = bl; ar = br; bl = tl; br = tr;
        }
        return al < bl && bl < ar && ar < br;
    }

    static long bruteConnectivity(int[] values) {
        int[][] intervals = buildIntervals(values);
        int n = intervals.length;
        assert n <= 20;

        long total = 0;
        for (long mask = 0; mask < (1L << n); ++mask) {
            boolean ok = true;
            for (int i = 0; i < n && ok; ++i) {
                if (((mask >> i) & 1L) == 0L) {
                    continue;
                }
                for (int j = i + 1; j < n; ++j) {
                    if (((mask >> j) & 1L) != 0L && crosses(intervals[i], intervals[j])) {
                        ok = false;
                        break;
                    }
                }
            }
            if (ok) {
                ++total;
            }
        }
        return total;
    }

    static void runCheckpoints() {
        int[][] 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 },
        };
        long[] expected = { 3, 8, 5, 8, 86 };
        for (int i = 0; i < cases.length; ++i) {
            long brute = bruteConnectivity(cases[i]);
            assert brute == expected[i];
            assert connectivityNumber(cases[i], MOD) == brute % MOD;
        }
    }

    public static void main(String[] args) throws IOException {
        runCheckpoints();

        int[] values = loadCsv("resources/documents/1001_input.txt");
        assert values.length == 40_000;

        System.out.println(connectivityNumber(values, MOD));
    }
}
