from avl_tree import AVLTree


def assert_valid(tree: AVLTree) -> None:
    assert tree.check_avl()
    assert tree.inorder_keys() == sorted(tree.inorder_keys())


def test_insert_into_empty_tree() -> None:
    tree = AVLTree()
    node = tree.insert(10, "ten")

    assert tree.root is node
    assert tree.root.key == 10
    assert tree.root.value == "ten"
    assert tree.size() == 1
    assert_valid(tree)


def test_left_rotation() -> None:
    tree = AVLTree()
    for key in [10, 20, 30]:
        tree.insert(key)

    assert tree.root is not None
    assert tree.root.key == 20
    assert tree.inorder_keys() == [10, 20, 30]
    assert_valid(tree)


def test_right_rotation() -> None:
    tree = AVLTree()
    for key in [30, 20, 10]:
        tree.insert(key)

    assert tree.root is not None
    assert tree.root.key == 20
    assert tree.inorder_keys() == [10, 20, 30]
    assert_valid(tree)


def test_left_right_rotation() -> None:
    tree = AVLTree()
    for key in [30, 10, 20]:
        tree.insert(key)

    assert tree.root is not None
    assert tree.root.key == 20
    assert tree.inorder_keys() == [10, 20, 30]
    assert_valid(tree)


def test_right_left_rotation() -> None:
    tree = AVLTree()
    for key in [10, 30, 20]:
        tree.insert(key)

    assert tree.root is not None
    assert tree.root.key == 20
    assert tree.inorder_keys() == [10, 20, 30]
    assert_valid(tree)


def test_longer_sequence_has_sorted_inorder() -> None:
    keys = [10, 20, 30, 40, 50, 25, 5, 35, 45, 60, 1, 7]
    tree = AVLTree()
    for key in keys:
        tree.insert(key)
        assert_valid(tree)

    assert tree.inorder_keys() == sorted(keys)
    assert tree.size() == len(keys)


def test_successor_for_smallest_inner_and_largest_key() -> None:
    tree = AVLTree()
    for key in [20, 10, 30, 5, 15, 25, 35, 13, 17]:
        tree.insert(key)

    assert tree.successor(tree.search(5)).key == 10
    assert tree.successor(tree.search(15)).key == 17
    assert tree.successor(tree.search(17)).key == 20
    assert tree.successor(tree.search(35)) is None
    assert_valid(tree)


def test_repeated_insert_updates_value_without_increasing_size() -> None:
    tree = AVLTree()
    tree.insert(10, "old")
    tree.insert(5, "five")
    tree.insert(15, "fifteen")

    old_size = tree.size()
    node = tree.insert(10, "new")

    assert tree.size() == old_size
    assert node.key == 10
    assert node.value == "new"
    assert tree.search(10).value == "new"
    assert tree.inorder_keys() == [5, 10, 15]
    assert_valid(tree)


if __name__ == "__main__":
    test_insert_into_empty_tree()
    test_left_rotation()
    test_right_rotation()
    test_left_right_rotation()
    test_right_left_rotation()
    test_longer_sequence_has_sorted_inorder()
    test_successor_for_smallest_inner_and_largest_key()
    test_repeated_insert_updates_value_without_increasing_size()
    print("all tests passed")
