MOD = 1_000_000_007


def parity_sign(value):
    parity = 0
    while value > 0:
        parity ^= 1
        value &= value - 1
    return 1 if parity == 0 else -1


def bit_count(limit):
    bits = 0
    while (1 << bits) <= limit:
        bits += 1
    return bits


def highest_bit(value):
    bit = 0
    while (1 << (bit + 1)) <= value:
        bit += 1
    return bit


def parity_and_candidate(n):
    bits = bit_count(n)
    count = [0] * bits
    discrepancy = [0] * bits

    for value in range(1, n + 1):
        sign = parity_sign(value)
        for bit in range(bits):
            if (value >> bit) & 1:
                count[bit] += 1
                discrepancy[bit] += sign

    upper_bound = 0
    value_sum = 0
    for bit in range(bits):
        weight = 1 << bit
        m = count[bit]
        s = discrepancy[bit]
        upper_bound += weight * ((m * m) // 4)
        value_sum += weight * ((m * m - s * s) // 4)

    return value_sum, upper_bound


def max_and_value(n):
    value, upper_bound = parity_and_candidate(n)
    assert value == upper_bound
    return value


def brute_max_and(n):
    best = 0
    for mask in range(1 << n):
        current = 0
        for a in range(1, n + 1):
            for b in range(a + 1, n + 1):
                side_a = (mask >> (a - 1)) & 1
                side_b = (mask >> (b - 1)) & 1
                if side_a != side_b:
                    current += a & b
        best = max(best, current)
    return best


def max_xor_sum(n):
    squares = [value * value for value in range(n + 1)]

    edges = []
    for frm in range(1, n + 1):
        for to in range(1, n + 1):
            if frm != to:
                edges.append((squares[frm] ^ squares[to], frm, to))

    edges.sort()

    best = [0] * (n + 1)
    answer = 0

    index = 0
    total = len(edges)
    while index < total:
        weight = edges[index][0]
        pending = []
        while index < total and edges[index][0] == weight:
            _, frm, to = edges[index]
            candidate = best[frm] + weight
            pending.append((to, candidate))
            if candidate > answer:
                answer = candidate
            index += 1
        for to, candidate in pending:
            if candidate > best[to]:
                best[to] = candidate

    return answer


def count_unreachable_nim(n):
    if n <= 1:
        return 0

    top_bit = highest_bit(n - 1)
    total = 0
    for target_bit in range(top_bit + 1):
        dp = [0] * 8
        dp[7] = 1

        for bit in range(top_bit, -1, -1):
            nxt = [0] * 8
            limit_bit = ((n - 1) >> bit) & 1

            for mask in range(8):
                if dp[mask] == 0:
                    continue
                for a in (0, 1):
                    for b in (0, 1):
                        for c in (0, 1):
                            if bit > target_bit and (a ^ b ^ c) != 0:
                                continue
                            if bit == target_bit and (a == 0 or b == 0 or c == 0):
                                continue

                            chosen = (a, b, c)
                            valid = True
                            next_mask = 0
                            for idx in range(3):
                                if not ((mask >> idx) & 1):
                                    continue
                                if chosen[idx] > limit_bit:
                                    valid = False
                                    break
                                if chosen[idx] == limit_bit:
                                    next_mask |= 1 << idx

                            if valid:
                                nxt[next_mask] += dp[mask]

            dp = nxt

        total += sum(dp)

    return total


def brute_unreachable_nim(n):
    total = 0
    for a in range(n):
        for b in range(n):
            for c in range(n):
                nim_sum = a ^ b ^ c
                if nim_sum == 0:
                    continue
                bit = highest_bit(nim_sum)
                if ((a >> bit) & 1) and ((b >> bit) & 1) and ((c >> bit) & 1):
                    total += 1
    return total


def mul_mod(lhs, rhs):
    return (lhs * rhs) % MOD


def meta_value(index, first, second, third):
    if index == 0:
        return first % MOD
    if index == 1:
        return second % MOD
    if index == 2:
        return third % MOD

    a = first % MOD
    b = second % MOD
    c = third % MOD
    for _ in range(3, index + 1):
        nxt = mul_mod(mul_mod(c, b), a)
        a, b, c = b, c, nxt

    return c


def run_checkpoints():
    assert brute_max_and(10) == 50
    assert max_and_value(10) == 50
    value, upper_bound = parity_and_candidate(1000)
    assert value == upper_bound
    assert max_xor_sum(4) == 71
    assert max_xor_sum(10) == 702

    for n in range(1, 21):
        assert count_unreachable_nim(n) == brute_unreachable_nim(n)
    assert count_unreachable_nim(10) == 123


def solve(n=1000):
    max_and = max_and_value(n)
    max_xor = max_xor_sum(n)
    unreachable_nim = count_unreachable_nim(n)
    return meta_value(n, max_and, max_xor, unreachable_nim)


if __name__ == "__main__":
    run_checkpoints()

    max_and = max_and_value(1000)
    max_xor = max_xor_sum(1000)
    unreachable_nim = count_unreachable_nim(1000)
    assert meta_value(4, max_and, max_xor, unreachable_nim) == 457_587_170

    print(meta_value(1000, max_and, max_xor, unreachable_nim))
