package avl import ( "reflect" "testing" ) func TestAVLBasicOperations(t *testing.T) { tree := NewAVLTree[int]() if tree.Height() != 0 { t.Fatalf("空树 Height() = %d, want 0", tree.Height()) } if _, ok := tree.Min(); ok { t.Fatal("空树 Min() 的 ok = true, want false") } if _, ok := tree.Max(); ok { t.Fatal("空树 Max() 的 ok = true, want false") } if tree.Search(1) { t.Fatal("空树 Search(1) = true, want false") } for _, v := range []int{6, 3, 2, 1, 4, 5} { tree.Insert(v) } if got, want := tree.Inorder(), []int{1, 2, 3, 4, 5, 6}; !reflect.DeepEqual(got, want) { t.Fatalf("Inorder() = %v, want %v", got, want) } if !tree.Search(1) || !tree.Search(6) || tree.Search(7) { t.Fatalf("Search 结果不正确:Search(1)=%v, Search(6)=%v, Search(7)=%v", tree.Search(1), tree.Search(6), tree.Search(7)) } if got, ok := tree.Min(); !ok || got != 1 { t.Fatalf("Min() = (%d, %v), want (1, true)", got, ok) } if got, ok := tree.Max(); !ok || got != 6 { t.Fatalf("Max() = (%d, %v), want (6, true)", got, ok) } assertAVL(t, tree.Root, nil, nil) } func TestAVLDelete(t *testing.T) { cases := []struct { name string values []int delete int want []int wantMin *int wantMax *int }{ { name: "delete only root leaf", values: []int{10}, delete: 10, want: []int{}, }, { name: "delete root with one child", values: []int{10, 5}, delete: 10, want: []int{5}, wantMin: intPtr(5), wantMax: intPtr(5), }, { name: "delete non-root leaf and update parent height", values: []int{10, 5, 15}, delete: 5, want: []int{10, 15}, wantMin: intPtr(10), wantMax: intPtr(15), }, { name: "delete non-root with one child", values: []int{10, 5, 15, 12}, delete: 15, want: []int{5, 10, 12}, wantMin: intPtr(5), wantMax: intPtr(12), }, { name: "delete root with two children successor is leaf", values: []int{10, 5, 15, 12, 18}, delete: 10, want: []int{5, 12, 15, 18}, wantMin: intPtr(5), wantMax: intPtr(18), }, { name: "delete root with two children successor has right child", values: []int{10, 5, 20, 15, 25, 17}, delete: 10, want: []int{5, 15, 17, 20, 25}, wantMin: intPtr(5), wantMax: intPtr(25), }, { name: "delete missing value leaves tree unchanged", values: []int{10, 5, 20, 15, 25}, delete: 99, want: []int{5, 10, 15, 20, 25}, wantMin: intPtr(5), wantMax: intPtr(25), }, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { tree := NewAVLTree[int]() for _, v := range tc.values { tree.Insert(v) } tree.Delete(tc.delete) if got := tree.Inorder(); !reflect.DeepEqual(got, tc.want) { t.Fatalf("Delete(%d) 后 Inorder() = %v, want %v", tc.delete, got, tc.want) } if tree.Search(tc.delete) && tc.delete != 99 { t.Fatalf("Delete(%d) 后 Search(%d) = true, want false", tc.delete, tc.delete) } if tc.wantMin == nil { if tree.Height() != 0 { t.Fatalf("删空树后 Height() = %d, want 0", tree.Height()) } if _, ok := tree.Min(); ok { t.Fatal("删空树后 Min() 的 ok = true, want false") } if _, ok := tree.Max(); ok { t.Fatal("删空树后 Max() 的 ok = true, want false") } } else { if got, ok := tree.Min(); !ok || got != *tc.wantMin { t.Fatalf("Min() = (%d, %v), want (%d, true)", got, ok, *tc.wantMin) } if got, ok := tree.Max(); !ok || got != *tc.wantMax { t.Fatalf("Max() = (%d, %v), want (%d, true)", got, ok, *tc.wantMax) } } assertAVL(t, tree.Root, nil, nil) }) } } func TestAVLDeleteRebalancesOnPathToRoot(t *testing.T) { tree := NewAVLTree[int]() for _, v := range []int{4, 2, 6, 1, 3, 5, 7} { tree.Insert(v) } for _, v := range []int{7, 6, 5} { tree.Delete(v) assertAVL(t, tree.Root, nil, nil) } if got, want := tree.Inorder(), []int{1, 2, 3, 4}; !reflect.DeepEqual(got, want) { t.Fatalf("连续删除后 Inorder() = %v, want %v", got, want) } } func assertAVL(t *testing.T, node *AVLNode[int], lower, upper *int) int { t.Helper() if node == nil { return 0 } if lower != nil && node.Val <= *lower { t.Fatalf("BST 性质被破坏:节点 %d <= 下界 %d", node.Val, *lower) } if upper != nil && node.Val >= *upper { t.Fatalf("BST 性质被破坏:节点 %d >= 上界 %d", node.Val, *upper) } leftHeight := assertAVL(t, node.Left, lower, &node.Val) rightHeight := assertAVL(t, node.Right, &node.Val, upper) if balanceFactor := leftHeight - rightHeight; balanceFactor < -1 || balanceFactor > 1 { t.Fatalf("节点 %d 的 BF = %d, want [-1, 1]", node.Val, balanceFactor) } wantHeight := max(leftHeight, rightHeight) + 1 if node.Height != wantHeight { t.Fatalf("节点 %d 的 Height = %d, want %d", node.Val, node.Height, wantHeight) } return wantHeight } func intPtr(v int) *int { return &v }