Files
leetcode/avl/avl.go
T
2026-07-20 19:57:25 +08:00

244 lines
5.3 KiB
Go

package avl
import "cmp"
// avl.go
// 泛型 AVL 自平衡二叉搜索树。
//
// 约定:空子树高度为 0,叶子高度为 1;不插入重复值。
// 每次插入或删除后,都要更新 Height,再按 LL/LR/RR/RL 四种情况重平衡。
// AVLNode AVL 树节点。Height 是以该节点为根的子树高度。
type AVLNode[T cmp.Ordered] struct {
Val T
Left *AVLNode[T]
Right *AVLNode[T]
Height int
}
// AVLTree AVL 自平衡二叉搜索树。
type AVLTree[T cmp.Ordered] struct {
Root *AVLNode[T]
}
// NewAVLTree 创建一棵空 AVL 树。
func NewAVLTree[T cmp.Ordered]() *AVLTree[T] {
return &AVLTree[T]{}
}
// Insert 插入 v;重复值不插入。
func (t *AVLTree[T]) Insert(v T) {
var insert func(node *AVLNode[T]) *AVLNode[T]
insert = func(node *AVLNode[T]) *AVLNode[T] {
if node == nil {
return &AVLNode[T]{Val: v, Height: 1}
}
if v == node.Val {
return node
} else if v < node.Val {
node.Left = insert(node.Left)
} else {
node.Right = insert(node.Right)
}
node.Height = t.calculateHeight(node)
return t.balancing(node)
}
t.Root = insert(t.Root)
}
func (t *AVLTree[T]) balancing(node *AVLNode[T]) *AVLNode[T] {
bf := t.calculateBF(node)
if bf < -1 {
// 右重 R
_bf := t.calculateBF(node.Right)
if _bf <= 0 {
// RR:左旋自身
return t.rotateLeft(node)
}
if _bf > 0 {
// RL: 右旋右孩子再左旋自身
node.Right = t.rotateRight(node.Right)
return t.rotateLeft(node)
}
}
if bf > 1 {
// 左重 L
_bf := t.calculateBF(node.Left)
if _bf >= 0 {
// LL: 右旋自身
return t.rotateRight(node)
}
if _bf < 0 {
// LR: 左旋左孩子再右旋自身
node.Left = t.rotateLeft(node.Left)
return t.rotateRight(node)
}
}
return node
}
func (t *AVLTree[T]) calculateHeight(node *AVLNode[T]) int {
if node == nil {
return 0
}
l, r := 0, 0
if node.Left != nil {
l = node.Left.Height
}
if node.Right != nil {
r = node.Right.Height
}
return max(l, r) + 1
}
func (t *AVLTree[T]) calculateBF(node *AVLNode[T]) int {
if node == nil {
return 0
}
leftHeight, rightHeight := 0, 0
if node.Left != nil {
leftHeight = node.Left.Height
}
if node.Right != nil {
rightHeight = node.Right.Height
}
bf := leftHeight - rightHeight
return bf
}
// rotateLeft 对 node 左旋,返回旋转后的新子树根。
//
// node right
// / \ / \
// A right → node C
// / \ / \
// B C A B
//
// B 的值介于 node 和 right 之间,因此旋转后要接到 node.Right。
func (t *AVLTree[T]) rotateLeft(node *AVLNode[T]) *AVLNode[T] {
rNode := node.Right
node.Right = rNode.Left
rNode.Left = node
node.Height = t.calculateHeight(node)
rNode.Height = t.calculateHeight(rNode)
return rNode
}
// rotateRight 对 node 右旋,返回旋转后的新子树根。
//
// node left
// / \ / \
// left C → A node
// / \ / \
// A B B C
//
// B 的值介于 left 和 node 之间,因此旋转后要接到 node.Left。
func (t *AVLTree[T]) rotateRight(node *AVLNode[T]) *AVLNode[T] {
lNode := node.Left
node.Left = lNode.Right
lNode.Right = node
node.Height = t.calculateHeight(node)
lNode.Height = t.calculateHeight(lNode)
return lNode
}
func (t *AVLTree[T]) Inorder() []T {
res := []T{}
var helper func(node *AVLNode[T])
helper = func(node *AVLNode[T]) {
if node == nil {
return
}
helper(node.Left)
res = append(res, node.Val)
helper(node.Right)
}
helper(t.Root)
return res
}
// Search 查找 v 是否存在。只读操作,不需要更新高度或旋转。
func (t *AVLTree[T]) Search(v T) bool {
node := t.Root
for node != nil {
if node.Val == v {
return true
}
if v < node.Val {
node = node.Left
} else {
node = node.Right
}
}
return false
}
// Height 返回整棵树的高度;空树为 0。
func (t *AVLTree[T]) Height() int {
if t.Root == nil {
return 0
}
return t.Root.Height
}
// Min 返回最小值及其是否存在。
func (t *AVLTree[T]) Min() (T, bool) {
node := t.Root
if node == nil {
var zero T
return zero, false
}
for node.Left != nil {
node = node.Left
}
return node.Val, true
}
// Max 返回最大值及其是否存在。
func (t *AVLTree[T]) Max() (T, bool) {
node := t.Root
if node == nil {
var zero T
return zero, false
}
for node.Right != nil {
node = node.Right
}
return node.Val, true
}
// Delete 删除 v;若 v 不存在,树保持不变。
func (t *AVLTree[T]) Delete(v T) {
var deleteNode func(node *AVLNode[T], target T) *AVLNode[T]
deleteNode = func(node *AVLNode[T], target T) *AVLNode[T] {
if node == nil {
return nil
}
if target < node.Val {
node.Left = deleteNode(node.Left, target)
} else if target > node.Val {
node.Right = deleteNode(node.Right, target)
} else {
if node.Left == nil && node.Right != nil {
return node.Right
}
if node.Left != nil && node.Right == nil {
return node.Left
}
if node.Left != nil && node.Right != nil {
n := node.Right
for n.Left != nil {
n = n.Left
}
node.Val = n.Val
node.Right = deleteNode(node.Right, n.Val)
} else {
return nil
}
}
node.Height = t.calculateHeight(node)
return t.balancing(node)
}
t.Root = deleteNode(t.Root, v)
}