树剖后以开始工作时间为下标(没工作的记为m),对dfn序列顺序的每个点建立线段树,update的时候这个点开始工作时间上加一,询问时对于询问时间t,求两点间路径上开始时间小于等于t-c-1的点的个数 代码如下:
#include<bits/stdc++.h>
#define pii pair<int, int>
using namespace std;
const int N = 1000010;
int h[N], to[N], nxt[N], idx;
int top[N], dfn[N], son[N], sz[N], dep[N], fa[N], lst;
int rt[N], mp[N], t[N], cnt, n, m, root;
struct Q{int x, y, c, t; }q[N];
struct Nd{int ls, rs, v; }a[N << 5];
void build(int &p, int l, int r)
{
p = ++ cnt; if(l == r) return;
int mid = l + r >> 1;
build(a[p].ls, l, mid); build(a[p].rs, mid + 1, r);
}
void upd(int &p, int l, int r, int x, int k)
{
a[++ cnt] = a[p]; p = cnt;
if(l == r) {a[p].v += k; return; }
int mid = l + r >> 1;
if(x <= mid) upd(a[p].ls, l, mid, x, k);
else upd(a[p].rs, mid + 1, r, x, k);
a[p].v = a[a[p].ls].v + a[a[p].rs].v;
}
int ask(int u, int v, int l, int r, int k)
{
if(l == r) return a[v].v - a[u].v;
int mid = l + r >> 1; int sl = a[a[v].ls].v - a[a[u].ls].v;
if(k <= mid) return ask(a[u].ls, a[v].ls, l, mid, k);
return sl + ask(a[u].rs, a[v].rs, mid + 1, r, k);
}
void add(int u, int v)
{to[++ idx] = v; nxt[idx] = h[u]; h[u] = idx; }
void dfs1(int x, int from)
{
sz[x] = 1; dep[x] = dep[from] + 1; fa[x] = from;
for(int i = h[x]; i != -1; i = nxt[i])
{
int e = to[i]; if(e == from) continue;
dfs1(e, x); sz[x] += sz[e];
if(sz[e] > sz[son[x]]) son[x] = e;
}
}
void dfs2(int x, int t)
{
dfn[x] = ++ lst; mp[lst] = x; top[x] = t;
if(!son[x]) return;
dfs2(son[x], t);
for(int i = h[x]; i != -1; i = nxt[i])
{
int e = to[i]; if(e == fa[x] || e == son[x]) continue;
dfs2(e, e);
}
}
pii qpth(int x, int y, int c)
{
int ans = 0, s = 0;
while(top[x] != top[y])
{
if(dep[top[x]] < dep[top[y]]) swap(x, y);
ans += ask(rt[dfn[top[x]] - 1], rt[dfn[x]], 1, m, c);
s += dfn[x] - dfn[top[x]] + 1; x = fa[top[x]];
}
if(dfn[x] < dfn[y]) swap(x, y);
ans += ask(rt[dfn[y] - 1], rt[dfn[x]], 1, m, c);
s += dfn[x] - dfn[y] + 1;
return {s, ans};
}
int main()
{
memset(h, -1, sizeof(h));
scanf("%d", &n);
for(int i = 1; i <= n; i ++)
{
int x; scanf("%d", &x);
if(x == 0) root = i; else add(x, i);
}
dfs1(root, 0);
dfs2(root, root);
int sq = 0; scanf("%d", &m);
build(rt[0], 1, m);
for(int i = 1; i <= n; i ++) t[i] = m;
for(int i = 1; i <= m; i ++)
{
int k, x, y, c; scanf("%d", &k);
if(k == 1)
{scanf("%d%d%d", &x, &y, &c); q[++ sq] = {x, y, c, i}; }
else
{scanf("%d", &c); t[c] = i; }
}
for(int i = 1; i <= n; i ++)
{
int ri = mp[i]; rt[i] = rt[i - 1];
upd(rt[i], 1, m, t[ri], 1);
}
for(int i = 1; i <= sq; i ++)
{
int x = q[i].x, y = q[i].y, c = q[i].c, t = q[i].t;
auto ans = qpth(x, y, t - c - 1);
cout << ans.first << ' ' << ans.second << "\n";
}
}