use crate::error::{GraphError, GraphResult};
use std::collections::VecDeque;
#[derive(Debug, Clone)]
pub struct McfEdge {
pub from: usize,
pub to: usize,
pub capacity: i64,
pub cost: i64,
}
#[derive(Debug, Clone)]
pub struct McfResult {
pub flow: i64,
pub cost: i64,
pub flow_per_edge: Vec<i64>,
}
#[derive(Debug, Clone)]
struct ResidualArc {
to: usize,
cap: i64,
cost: i64,
rev: usize,
}
type ResidualAdj = Vec<Vec<ResidualArc>>;
type FwdIndices = Vec<(usize, usize)>;
fn build_residual(n: usize, edges: &[McfEdge]) -> GraphResult<(ResidualAdj, FwdIndices)> {
let mut adj: Vec<Vec<ResidualArc>> = vec![Vec::new(); n];
let mut fwd_indices = Vec::with_capacity(edges.len());
for edge in edges {
if edge.from >= n || edge.to >= n {
return Err(GraphError::InvalidPlan("node_out_of_range".to_owned()));
}
let fwd_node = edge.from;
let fwd_pos = adj[edge.from].len();
fwd_indices.push((fwd_node, fwd_pos));
let rev_pos = adj[edge.to].len();
adj[edge.from].push(ResidualArc {
to: edge.to,
cap: edge.capacity,
cost: edge.cost,
rev: rev_pos,
});
let fwd_back = fwd_pos;
adj[edge.to].push(ResidualArc {
to: edge.from,
cap: 0,
cost: -edge.cost,
rev: fwd_back,
});
}
Ok((adj, fwd_indices))
}
fn spfa(
adj: &[Vec<ResidualArc>],
source: usize,
sink: usize,
n: usize,
) -> Option<(Vec<usize>, Vec<usize>)> {
let mut dist = vec![i64::MAX; n];
let mut in_queue = vec![false; n];
let mut prev_node = vec![usize::MAX; n];
let mut prev_arc = vec![usize::MAX; n];
dist[source] = 0;
let mut queue: VecDeque<usize> = VecDeque::new();
queue.push_back(source);
in_queue[source] = true;
while let Some(u) = queue.pop_front() {
in_queue[u] = false;
for (arc_idx, arc) in adj[u].iter().enumerate() {
if arc.cap > 0 && dist[u] != i64::MAX {
let new_dist = dist[u].saturating_add(arc.cost);
if new_dist < dist[arc.to] {
dist[arc.to] = new_dist;
prev_node[arc.to] = u;
prev_arc[arc.to] = arc_idx;
if !in_queue[arc.to] {
in_queue[arc.to] = true;
queue.push_back(arc.to);
}
}
}
}
}
if dist[sink] == i64::MAX {
None
} else {
Some((prev_node, prev_arc))
}
}
pub fn min_cost_flow(
n_nodes: usize,
edges: &[McfEdge],
source: usize,
sink: usize,
max_flow: i64,
) -> GraphResult<McfResult> {
if n_nodes == 0 {
return Err(GraphError::InvalidPlan("n_nodes_zero".to_owned()));
}
if source == sink {
return Err(GraphError::InvalidPlan("source_equals_sink".to_owned()));
}
if source >= n_nodes || sink >= n_nodes {
return Err(GraphError::InvalidPlan("node_out_of_range".to_owned()));
}
let (mut adj, fwd_indices) = build_residual(n_nodes, edges)?;
let mut total_flow: i64 = 0;
let mut total_cost: i64 = 0;
while total_flow < max_flow {
let (prev_node, prev_arc) = match spfa(&adj, source, sink, n_nodes) {
Some(p) => p,
None => break, };
let remaining = max_flow - total_flow;
let mut bottleneck = remaining;
let mut v = sink;
while v != source {
let u = prev_node[v];
let arc = &adj[u][prev_arc[v]];
bottleneck = bottleneck.min(arc.cap);
v = u;
}
if bottleneck <= 0 {
break;
}
let mut path_cost: i64 = 0;
v = sink;
while v != source {
let u = prev_node[v];
path_cost = path_cost.saturating_add(adj[u][prev_arc[v]].cost);
v = u;
}
total_cost = total_cost.saturating_add(path_cost.saturating_mul(bottleneck));
total_flow += bottleneck;
v = sink;
while v != source {
let u = prev_node[v];
let arc_idx = prev_arc[v];
let rev_idx = adj[u][arc_idx].rev;
adj[u][arc_idx].cap -= bottleneck;
adj[v][rev_idx].cap += bottleneck;
v = u;
}
}
let flow_per_edge: Vec<i64> = fwd_indices
.iter()
.zip(edges.iter())
.map(|(&(node, pos), edge)| edge.capacity - adj[node][pos].cap)
.collect();
Ok(McfResult {
flow: total_flow,
cost: total_cost,
flow_per_edge,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn edge(from: usize, to: usize, cap: i64, cost: i64) -> McfEdge {
McfEdge {
from,
to,
capacity: cap,
cost,
}
}
#[test]
fn simple_path_flow() {
let edges = vec![edge(0, 1, 1, 1), edge(1, 2, 1, 1)];
let r = min_cost_flow(3, &edges, 0, 2, 10).expect("min_cost_flow should succeed");
assert_eq!(r.flow, 1);
assert_eq!(r.cost, 2);
}
#[test]
fn max_flow_limited() {
let edges = vec![edge(0, 1, 5, 1), edge(1, 2, 5, 1)];
let r = min_cost_flow(3, &edges, 0, 2, 3).expect("min_cost_flow should succeed");
assert_eq!(r.flow, 3);
assert_eq!(r.cost, 6);
}
#[test]
fn cost_optimal_route() {
let edges = vec![
edge(0, 1, 1, 1),
edge(1, 3, 1, 1),
edge(0, 2, 1, 5),
edge(2, 3, 1, 5),
];
let r = min_cost_flow(4, &edges, 0, 3, 1).expect("min_cost_flow should succeed");
assert_eq!(r.flow, 1);
assert_eq!(r.cost, 2); }
#[test]
fn no_augmenting_path_zero_flow() {
let edges: Vec<McfEdge> = vec![edge(0, 1, 1, 1)]; let r = min_cost_flow(4, &edges, 0, 3, 10).expect("min_cost_flow should succeed");
assert_eq!(r.flow, 0);
assert_eq!(r.cost, 0);
}
#[test]
fn negative_cost_edge() {
let edges = vec![edge(0, 1, 1, -2), edge(1, 2, 1, 3), edge(0, 2, 1, 5)];
let r = min_cost_flow(3, &edges, 0, 2, 1).expect("min_cost_flow should succeed");
assert_eq!(r.flow, 1);
assert_eq!(r.cost, 1); }
#[test]
fn parallel_edges() {
let edges = vec![edge(0, 1, 1, 1), edge(0, 1, 1, 2), edge(1, 2, 2, 1)];
let r = min_cost_flow(3, &edges, 0, 2, 2).expect("min_cost_flow should succeed");
assert_eq!(r.flow, 2);
assert_eq!(r.cost, 5);
}
#[test]
fn flow_per_edge_correct() {
let edges = vec![edge(0, 1, 2, 1), edge(1, 2, 2, 1)];
let r = min_cost_flow(3, &edges, 0, 2, 10).expect("min_cost_flow should succeed");
assert_eq!(r.flow_per_edge.len(), 2);
assert_eq!(r.flow_per_edge[0], r.flow);
assert_eq!(r.flow_per_edge[1], r.flow);
assert!(r.flow_per_edge.iter().sum::<i64>() > 0);
}
#[test]
fn source_equals_sink_error() {
let edges = vec![edge(0, 1, 1, 1)];
let err = min_cost_flow(3, &edges, 1, 1, 10);
assert!(
matches!(err, Err(GraphError::InvalidPlan(ref s)) if s == "source_equals_sink"),
"got: {err:?}"
);
}
#[test]
fn n_nodes_0_error() {
let err = min_cost_flow(0, &[], 0, 0, 0);
assert!(
matches!(err, Err(GraphError::InvalidPlan(ref s)) if s == "n_nodes_zero"),
"got: {err:?}"
);
}
#[test]
fn conservation_law() {
let edges = vec![
edge(0, 1, 2, 1),
edge(0, 2, 2, 2),
edge(1, 3, 2, 1),
edge(2, 3, 2, 2),
];
let n = 4;
let r = min_cost_flow(n, &edges, 0, 3, 4).expect("min_cost_flow should succeed");
let mut net = vec![0i64; n];
for (e, &f) in edges.iter().zip(r.flow_per_edge.iter()) {
net[e.from] -= f;
net[e.to] += f;
}
for (v, &net_v) in net.iter().enumerate().skip(1).take(n - 2) {
assert_eq!(net_v, 0, "node {v} violates conservation: net={net_v}");
}
assert_eq!(net[0], -r.flow);
assert_eq!(net[n - 1], r.flow);
}
#[test]
fn node_out_of_range_error() {
let edges = vec![edge(0, 5, 1, 1)]; let err = min_cost_flow(4, &edges, 0, 3, 10);
assert!(
matches!(err, Err(GraphError::InvalidPlan(_))),
"got: {err:?}"
);
}
#[test]
fn zero_capacity_edge_contributes_nothing() {
let edges = vec![
edge(0, 1, 0, 1), edge(0, 1, 2, 5),
edge(1, 2, 2, 1),
];
let r = min_cost_flow(3, &edges, 0, 2, 10).expect("min_cost_flow should succeed");
assert_eq!(r.flow_per_edge[0], 0);
assert!(r.flow > 0);
}
}