algoalgo-world
algoalgo-world/tree/splay
tree/splay

스플레이 트리

종류
쓰는 대로 자리를 바꾸는 이진 탐색 트리
균형
건드린 것을 뿌리까지 끌어올리기
평균
O(log n)
최악
O(n)
최악 높이
n - 1
노드마다 더 드는 것
없다

평균은 여러 번을 합쳐 나눈 값이다. 한 번이 트리를 통째로 훑을 수도 있지만 n 번을 합치면 n log n 을 안 넘는다

언어를 고르면 코드가 열린다

def rotate(node, to_right):
    pivot = node.left if to_right else node.right
    if to_right:
        node.left, pivot.right = pivot.right, node
    else:
        node.right, pivot.left = pivot.left, node
    return pivot


def splay(node, key):
    if node is None or node.key == key:
        return node
    if key < node.key:
        if node.left is not None:
            if key < node.left.key:
                node.left.left = splay(node.left.left, key)
                node = rotate(node, True)
            elif key > node.left.key:
                node.left.right = splay(node.left.right, key)
                if node.left.right is not None:
                    node.left = rotate(node.left, False)
        return node if node.left is None else rotate(node, True)
    if node.right is not None:
        if key > node.right.key:
            node.right.right = splay(node.right.right, key)
            node = rotate(node, False)
        elif key < node.right.key:
            node.right.left = splay(node.right.left, key)
            if node.right.left is not None:
                node.right = rotate(node.right, True)
    return node if node.right is None else rotate(node, False)


def insert(root, key):
    fresh = Node(key)
    root = splay(root, key)
    if root is None:
        return fresh
    if key < root.key:
        fresh.left, fresh.right, root.left = root.left, root, None
    else:
        fresh.left, fresh.right, root.right = root, root.right, None
    return fresh
function rotate(node, toRight) {
  const pivot = toRight ? node.left : node.right;
  if (toRight) [node.left, pivot.right] = [pivot.right, node];
  else [node.right, pivot.left] = [pivot.left, node];
  return pivot;
}

function splay(node, key) {
  if (node === null || node.key === key) return node;
  if (key < node.key) {
    if (node.left === null) return node;
    if (key < node.left.key) {
      node.left.left = splay(node.left.left, key);
      node = rotate(node, true);
    } else if (key > node.left.key) {
      node.left.right = splay(node.left.right, key);
      if (node.left.right !== null) node.left = rotate(node.left, false);
    }
    return node.left === null ? node : rotate(node, true);
  }
  if (node.right === null) return node;
  if (key > node.right.key) {
    node.right.right = splay(node.right.right, key);
    node = rotate(node, false);
  } else if (key < node.right.key) {
    node.right.left = splay(node.right.left, key);
    if (node.right.left !== null) node.right = rotate(node.right, true);
  }
  return node.right === null ? node : rotate(node, false);
}

function insert(root, key) {
  const fresh = makeNode(key);
  root = splay(root, key);
  if (root === null) return fresh;
  const less = key < root.key;
  fresh.left = less ? root.left : root;
  fresh.right = less ? root : root.right;
  if (less) root.left = null;
  else root.right = null;
  return fresh;
}
static Node* rotate(Node* n, int to_right) {
    Node* pivot = to_right ? n->left : n->right;
    if (to_right) { n->left = pivot->right; pivot->right = n; }
    else { n->right = pivot->left; pivot->left = n; }
    return pivot;
}

Node* splay(Node* n, int key) {
    if (n == NULL || n->key == key) return n;
    if (key < n->key) {
        if (n->left == NULL) return n;
        if (key < n->left->key) {
            n->left->left = splay(n->left->left, key);
            n = rotate(n, 1);
        } else if (key > n->left->key) {
            n->left->right = splay(n->left->right, key);
            if (n->left->right != NULL) n->left = rotate(n->left, 0);
        }
        return n->left == NULL ? n : rotate(n, 1);
    }
    if (n->right == NULL) return n;
    if (key > n->right->key) {
        n->right->right = splay(n->right->right, key);
        n = rotate(n, 0);
    } else if (key < n->right->key) {
        n->right->left = splay(n->right->left, key);
        if (n->right->left != NULL) n->right = rotate(n->right, 1);
    }
    return n->right == NULL ? n : rotate(n, 0);
}

Node* insert(Node* root, int key) {
    Node* fresh = calloc(1, sizeof(Node));
    fresh->key = key;
    root = splay(root, key);
    if (root == NULL) return fresh;
    int less = key < root->key;
    Node** inner = less ? &root->left : &root->right;
    fresh->left = less ? *inner : root;
    fresh->right = less ? root : *inner;
    *inner = NULL;
    return fresh;
}
Node* rotate(Node* n, bool to_right) {
    Node* pivot = to_right ? n->left : n->right;
    if (to_right) { n->left = pivot->right; pivot->right = n; }
    else { n->right = pivot->left; pivot->left = n; }
    return pivot;
}

Node* splay(Node* n, int key) {
    if (n == nullptr || n->key == key) return n;
    if (key < n->key) {
        if (n->left == nullptr) return n;
        if (key < n->left->key) {
            n->left->left = splay(n->left->left, key);
            n = rotate(n, true);
        } else if (key > n->left->key) {
            n->left->right = splay(n->left->right, key);
            if (n->left->right != nullptr) n->left = rotate(n->left, false);
        }
        return n->left == nullptr ? n : rotate(n, true);
    }
    if (n->right == nullptr) return n;
    if (key > n->right->key) {
        n->right->right = splay(n->right->right, key);
        n = rotate(n, false);
    } else if (key < n->right->key) {
        n->right->left = splay(n->right->left, key);
        if (n->right->left != nullptr) n->right = rotate(n->right, true);
    }
    return n->right == nullptr ? n : rotate(n, false);
}

Node* insert(Node* root, int key) {
    Node* fresh = new Node{key};
    root = splay(root, key);
    if (root == nullptr) return fresh;
    bool less = key < root->key;
    Node*& inner = less ? root->left : root->right;
    fresh->left = less ? inner : root;
    fresh->right = less ? root : inner;
    inner = nullptr;
    return fresh;
}
static Node Rotate(Node n, bool toRight) {
    Node pivot = toRight ? n.Left : n.Right;
    if (toRight) { n.Left = pivot.Right; pivot.Right = n; }
    else { n.Right = pivot.Left; pivot.Left = n; }
    return pivot;
}

static Node Splay(Node n, int key) {
    if (n == null || n.Key == key) return n;
    if (key < n.Key) {
        if (n.Left == null) return n;
        if (key < n.Left.Key) {
            n.Left.Left = Splay(n.Left.Left, key);
            n = Rotate(n, true);
        } else if (key > n.Left.Key) {
            n.Left.Right = Splay(n.Left.Right, key);
            if (n.Left.Right != null) n.Left = Rotate(n.Left, false);
        }
        return n.Left == null ? n : Rotate(n, true);
    }
    if (n.Right == null) return n;
    if (key > n.Right.Key) {
        n.Right.Right = Splay(n.Right.Right, key);
        n = Rotate(n, false);
    } else if (key < n.Right.Key) {
        n.Right.Left = Splay(n.Right.Left, key);
        if (n.Right.Left != null) n.Right = Rotate(n.Right, true);
    }
    return n.Right == null ? n : Rotate(n, false);
}

static Node Insert(Node root, int key) {
    Node fresh = new Node(key);
    root = Splay(root, key);
    if (root == null) return fresh;
    bool less = key < root.Key;
    fresh.Left = less ? root.Left : root;
    fresh.Right = less ? root : root.Right;
    if (less) root.Left = null;
    else root.Right = null;
    return fresh;
}
static Node rotate(Node n, boolean toRight) {
    Node pivot = toRight ? n.left : n.right;
    if (toRight) { n.left = pivot.right; pivot.right = n; }
    else { n.right = pivot.left; pivot.left = n; }
    return pivot;
}

static Node splay(Node n, int key) {
    if (n == null || n.key == key) return n;
    if (key < n.key) {
        if (n.left == null) return n;
        if (key < n.left.key) {
            n.left.left = splay(n.left.left, key);
            n = rotate(n, true);
        } else if (key > n.left.key) {
            n.left.right = splay(n.left.right, key);
            if (n.left.right != null) n.left = rotate(n.left, false);
        }
        return n.left == null ? n : rotate(n, true);
    }
    if (n.right == null) return n;
    if (key > n.right.key) {
        n.right.right = splay(n.right.right, key);
        n = rotate(n, false);
    } else if (key < n.right.key) {
        n.right.left = splay(n.right.left, key);
        if (n.right.left != null) n.right = rotate(n.right, true);
    }
    return n.right == null ? n : rotate(n, false);
}

static Node insert(Node root, int key) {
    Node fresh = new Node(key);
    root = splay(root, key);
    if (root == null) return fresh;
    boolean less = key < root.key;
    fresh.left = less ? root.left : root;
    fresh.right = less ? root : root.right;
    if (less) root.left = null;
    else root.right = null;
    return fresh;
}
돌려 보고 코드도 본다