rt,用的倍增+kruscal
#include <bits/stdc++.h>
#define ll long long
using namespace std;
const int N = 1e5 + 10, M = 3e5 + 10, INF = 0x3f3f3f3f;
struct node
{
int x, y, z;
bool tag;
bool operator<(const node &e) const
{
return z < e.z;
}
} e[M];
vector<int> G[N], E[N];
int p[N];
int dep[N];
int f[N][21], d1[N][21], d2[N][21];
int n, m;
inline int read()
{
int x = 0, y = 1; char c = getchar();
while (c < '0' || c > '9') {if (c == '-') y = -1; c = getchar();}
while (c >= '0' && c <= '9') x = x * 10 + c - '0', c = getchar();
return x * y;
}
inline void add(int a, int b, int c)
{
G[a].push_back(b);
E[a].push_back(c);
}
inline int find(int x)
{
if (p[x] != x) p[x] = find(p[x]);
return p[x];
}
inline ll kruscal()
{
sort(e + 1, e + 1 + m);
for (int i = 1; i <= n; i++)
p[i] = i;
ll res = 0;
int cnt = 0;
for (int i = 1; i <= m; i++)
{
int px = find(e[i].x), py = find(e[i].y);
if (px == py) continue;
res += e[i].z;
e[i].tag = true;
if (++cnt == n - 1) break;
}
return res;
}
inline void build()
{
for (int i = 1; i <= m; i++)
if (e[i].tag)
add(e[i].x, e[i].y, e[i].z), add(e[i].y, e[i].x, e[i].z);
}
inline void bfs()
{
memset(dep, 0x3f, sizeof dep);
dep[1] = 1, dep[0] = 0;;
queue<int> q;
q.push(1);
while (q.size())
{
int u = q.front();
q.pop();
for (int i = 0; i < G[u].size(); i++)
{
int v = G[u][i], w = E[u][i];
if (dep[v] <= dep[u] + 1) continue;
q.push(v);
dep[v] = dep[u] + 1, f[v][0] = u, d1[v][0] = w, d2[v][0] = -INF;
for (int j = 1; j <= 16; j++)
{
int fa = f[v][j - 1];
f[v][j] = f[fa][j - 1];
int dist[4] = {d1[v][j - 1], d1[fa][j - 1], d2[v][j - 1], d2[fa][j - 1]};
d1[v][j] = d2[v][j] = -INF;
for (int k = 0; k < 4; k++)
{
if (dist[k] > d1[v][j]) d2[v][j] = d1[v][j], d1[v][j] = dist[k];
else if (dist[k] != d1[v][j] && dist[k] > d2[v][j]) d2[v][j] = dist[k];
}
}
}
}
}
inline int lca(int x, int y, int z)
{
int dist[N << 1];
int length = 0;
if (dep[x] < dep[y]) swap(x, y);
for (int i = log2(n); i >= 0; i--)
if (dep[f[x][i]] >= dep[y])
{
dist[++length] = d1[x][i];
dist[++length] = d2[x][i];
x = f[x][i];
}
if (x != y)
{
for (int i = log2(n); i >= 0; i--)
if (f[x][i] != f[y][i])
{
dist[++length] = d1[x][i];
dist[++length] = d2[x][i];
dist[++length] = d1[y][i];
dist[++length] = d2[y][i];
x = f[x][i];
y = f[y][i];
}
dist[++length] = d1[x][0];
dist[++length] = d1[y][0];
}
int t1 = -INF, t2 = -INF;
for (int i = 1; i <= length; i++)
{
if (dist[i] > t1) t2 = t1, t1 = dist[i];
else if (dist[i] != t1 && dist[i] > t2) t2 = dist[i];
}
if (z > t1) return z - t1;
if (z > t2) return z - t2;
return INF;
}
int main()
{
n = read(), m = read();
for (int i = 1; i <= m; i++)
{
int a = read(), b = read(), c = read();
e[i] = {a, b, c};
}
ll sum = kruscal();
build();
bfs();
ll res = (1ll << 62);
for (int i = 1; i <= m; i++)
{
if (e[i].tag) continue;
res = min(res, sum + lca(e[i].x, e[i].y, e[i].z));
}
printf("%lld\n", res);
return 0;
}