Maximize Spanning Tree Stability with Upgrades

hard binary search greedy union-find spanning tree

Problem

Given n nodes and weighted edges [u, v, s, must], build a spanning tree (exactly n−1 edges, fully connected, acyclic). Edges with must = 1 have to be used and cannot change; an optional edge (must = 0) may have its strength doubled once, using at most k upgrades total. The stability of a tree is the minimum edge strength in it. Return the maximum achievable stability, or -1 if no spanning tree exists.

Inputn = 3, edges = [[0,1,2,1],[1,2,3,0]], k = 1
Output2
Edge [0,1] (strength 2) is mandatory. Edge [1,2] is optional and can be upgraded 3→6. The tree uses strengths {2, 6}; the minimum is 2.

def maxStability(n, edges, k):
    parent = list(range(n))
    def find(x):
        while parent[x] != x:
            parent[x] = parent[parent[x]]
            x = parent[x]
        return x
    musts = [e for e in edges if e[3] == 1]
    opt   = [e for e in edges if e[3] == 0]

    def feasible(x):                       # can stability >= x?
        for i in range(n): parent[i] = i
        comps = n
        for u, v, s, _ in musts:           # mandatory edges first
            if s < x: return False         # cannot upgrade -> caps min
            ru, rv = find(u), find(v)
            if ru == rv: return False      # must edge makes a cycle
            parent[ru] = rv; comps -= 1
        for u, v, s, _ in opt:             # free edges that meet x
            if s >= x:
                ru, rv = find(u), find(v)
                if ru != rv: parent[ru] = rv; comps -= 1
        up = k
        for u, v, s, _ in opt:             # spend upgrades greedily
            if up == 0: break
            if s < x and 2 * s >= x:
                ru, rv = find(u), find(v)
                if ru != rv: parent[ru] = rv; comps -= 1; up -= 1
        return comps == 1

    cand = sorted({s for *_, s, _ in [(e[0],e[1],e[2],e[3]) for e in edges]}
                  | {2 * e[2] for e in edges})
    lo, hi, ans = 0, len(cand) - 1, -1
    while lo <= hi:                         # binary search the threshold
        mid = (lo + hi) // 2
        if feasible(cand[mid]):
            ans = cand[mid]; lo = mid + 1
        else:
            hi = mid - 1
    return ans
function maxStability(n, edges, k) {
  const parent = Array.from({ length: n }, (_, i) => i);
  const find = x => { while (parent[x] !== x) { parent[x] = parent[parent[x]]; x = parent[x]; } return x; };
  const musts = edges.filter(e => e[3] === 1);
  const opt   = edges.filter(e => e[3] === 0);

  const feasible = x => {                  // can stability >= x?
    for (let i = 0; i < n; i++) parent[i] = i;
    let comps = n;
    for (const [u, v, s] of musts) {       // mandatory edges first
      if (s < x) return false;             // cannot upgrade -> caps min
      const ru = find(u), rv = find(v);
      if (ru === rv) return false;         // must edge makes a cycle
      parent[ru] = rv; comps--;
    }
    for (const [u, v, s] of opt) {         // free edges that meet x
      if (s >= x) { const ru = find(u), rv = find(v); if (ru !== rv) { parent[ru] = rv; comps--; } }
    }
    let up = k;
    for (const [u, v, s] of opt) {         // spend upgrades greedily
      if (up === 0) break;
      if (s < x && 2 * s >= x) { const ru = find(u), rv = find(v); if (ru !== rv) { parent[ru] = rv; comps--; up--; } }
    }
    return comps === 1;
  };

  const cand = [...new Set(edges.flatMap(e => [e[2], 2 * e[2]]))].sort((a, b) => a - b);
  let lo = 0, hi = cand.length - 1, ans = -1;
  while (lo <= hi) {                       // binary search the threshold
    const mid = (lo + hi) >> 1;
    if (feasible(cand[mid])) { ans = cand[mid]; lo = mid + 1; }
    else hi = mid - 1;
  }
  return ans;
}
int[] parent;
int find(int x) {
    while (parent[x] != x) { parent[x] = parent[parent[x]]; x = parent[x]; }
    return x;
}
int maxStability(int n, int[][] edges, int k) {
    parent = new int[n];
    List<int[]> musts = new ArrayList<>(), opt = new ArrayList<>();
    TreeSet<Integer> cset = new TreeSet<>();
    for (int[] e : edges) {
        (e[3] == 1 ? musts : opt).add(e);
        cset.add(e[2]); cset.add(2 * e[2]);
    }
    int[] cand = cset.stream().mapToInt(Integer::intValue).toArray();
    int lo = 0, hi = cand.length - 1, ans = -1;
    while (lo <= hi) {                       // binary search the threshold
        int mid = (lo + hi) >>> 1;
        if (feasible(n, cand[mid], musts, opt, k)) { ans = cand[mid]; lo = mid + 1; }
        else hi = mid - 1;
    }
    return ans;
}
boolean feasible(int n, int x, List<int[]> musts, List<int[]> opt, int k) {
    for (int i = 0; i < n; i++) parent[i] = i;
    int comps = n;
    for (int[] e : musts) {                 // mandatory edges first
        if (e[2] < x) return false;         // cannot upgrade -> caps min
        int ru = find(e[0]), rv = find(e[1]);
        if (ru == rv) return false;         // must edge makes a cycle
        parent[ru] = rv; comps--;
    }
    for (int[] e : opt) if (e[2] >= x) {    // free edges that meet x
        int ru = find(e[0]), rv = find(e[1]);
        if (ru != rv) { parent[ru] = rv; comps--; }
    }
    int up = k;
    for (int[] e : opt) {                    // spend upgrades greedily
        if (up == 0) break;
        if (e[2] < x && 2 * e[2] >= x) {
            int ru = find(e[0]), rv = find(e[1]);
            if (ru != rv) { parent[ru] = rv; comps--; up--; }
        }
    }
    return comps == 1;
}
vector<int> parent;
int find(int x) {
    while (parent[x] != x) { parent[x] = parent[parent[x]]; x = parent[x]; }
    return x;
}
bool feasible(int n, int x, vector<vector<int>>& musts, vector<vector<int>>& opt, int k) {
    for (int i = 0; i < n; i++) parent[i] = i;
    int comps = n;
    for (auto& e : musts) {                  // mandatory edges first
        if (e[2] < x) return false;         // cannot upgrade -> caps min
        int ru = find(e[0]), rv = find(e[1]);
        if (ru == rv) return false;         // must edge makes a cycle
        parent[ru] = rv; comps--;
    }
    for (auto& e : opt) if (e[2] >= x) {    // free edges that meet x
        int ru = find(e[0]), rv = find(e[1]);
        if (ru != rv) { parent[ru] = rv; comps--; }
    }
    int up = k;
    for (auto& e : opt) {                    // spend upgrades greedily
        if (up == 0) break;
        if (e[2] < x && 2 * e[2] >= x) {
            int ru = find(e[0]), rv = find(e[1]);
            if (ru != rv) { parent[ru] = rv; comps--; up--; }
        }
    }
    return comps == 1;
}
int maxStability(int n, vector<vector<int>>& edges, int k) {
    parent.resize(n);
    vector<vector<int>> musts, opt; set<int> cset;
    for (auto& e : edges) {
        (e[3] == 1 ? musts : opt).push_back(e);
        cset.insert(e[2]); cset.insert(2 * e[2]);
    }
    vector<int> cand(cset.begin(), cset.end());
    int lo = 0, hi = cand.size() - 1, ans = -1;
    while (lo <= hi) {                       // binary search the threshold
        int mid = (lo + hi) / 2;
        if (feasible(n, cand[mid], musts, opt, k)) { ans = cand[mid]; lo = mid + 1; }
        else hi = mid - 1;
    }
    return ans;
}
Time: O(E log E · α(n)) Space: O(n + E)