今天脑抽想写 BST 结果写炸了。
求求路过聚聚帮忙看看!
#include <iostream>
#include <cstring>
#include <cstdio>
using namespace std;
const int INF = 2e9;
struct node {
node *ls, *rs;
int val, size, cnt;
node() {}
node(int val, int size, int cnt):val(val), size(size), cnt(cnt) {
ls = rs = nullptr;
}
void pushup() {
size = ((ls != nullptr) ? (ls -> size) : 0) +
((rs != nullptr) ? (rs -> size) : 0) +
cnt;
}
}*root;
void insert(node *x, int val) {
if (x -> val == val) {
x -> size ++ ; x -> cnt ++ ;
return;
}
if (x -> val < val) {
if (x -> rs == nullptr) {
x -> rs = new node(val, 1, 1);
x -> pushup();
return;
}
insert(x -> rs, val);
}
else {
if (x -> ls == nullptr) {
x -> ls = new node(val, 1, 1);
x -> pushup();
return;
}
insert(x -> ls, val);
}
x -> pushup();
}
void remove(node *x, int val) {
if (x -> val == val) {
x -> cnt -- ; x -> size -- ;
return;
}
if (x -> val < val) remove(x -> rs, val);
else remove(x -> ls, val);
x -> pushup();
}
int get_rank(node *x, int val) {
node *now = x;
int rank = 0;
while (true) {
if (now -> val == val) {
rank += now -> ls -> size;
return rank;
}
if (now -> val < val) {
rank += now -> ls -> size + now -> cnt;
if (now -> rs != nullptr) now = now -> rs;
else return rank;
}
else if (now -> val > val) {
if (now -> ls != nullptr) now = now -> ls;
else return rank;
}
}
return rank;
}
int get_val(node *x, int rank) {
node *now = x;
while (true) {
int l_size = (now -> ls == nullptr) ? 0 : now -> ls -> size;
if (rank >= l_size + 1 and rank <= l_size + now -> cnt)
return now -> val;
if (rank <= l_size) now = now -> ls;
else rank -= l_size + now -> cnt, now = now -> rs;
}
}
int get_pre(node *x, int val) {
node *now = x;
int ans = -0x3f3f3f3f;
while (true) {
if (now -> val >= val) {
if (now -> ls != nullptr) now = now -> ls;
else return ans;
}
else {
if (now -> cnt > 0) ans = max(ans, now -> val);
if (now -> rs != nullptr) now = now -> rs;
else return ans;
}
}
}
int get_next(node *x, int val) {
node *now = x;
int ans = 0x3f3f3f3f;
while (true) {
if (now -> val <= val) {
if (now -> rs != nullptr) now = now -> rs;
else return ans;
}
else {
if (now -> cnt > 0) ans = min(ans, now -> val);
if (now -> ls != nullptr) now = now -> ls;
else return ans;
}
}
}
int main() {
int n;
scanf("%d", &n);
root = new node(-INF, 1, 1);
insert(root, INF);
while (n -- ) {
int op, x;
scanf("%d%d", &op, &x);
if (op == 1) insert(root, x);
else if (op == 2) remove(root, x);
else if (op == 3) printf("%d\n", get_rank(root, x) - 1);
else if (op == 4) printf("%d\n", get_val(root, x + 1));
else if (op == 5) printf("%d\n", get_pre(root, x));
else printf("%d\n", get_next(root, x));
}
return 0;
}