use ddo::core::abstraction::dp::Problem;
use ddo::core::common::{Decision, Domain, Variable, VarSet};
use std::cmp::{max, min};
use std::fs::File;
use std::io::{BufRead, BufReader, Lines, Read};
use crate::graph::Graph;
#[derive(Debug, Clone, Hash, Eq, PartialEq)]
pub struct McpState {
pub benef : Vec<i32>,
pub initial: bool
}
const S : i32 = 1;
const T : i32 =-1;
const ONLY_S : [i32; 1] = [S];
const BOTH_ST: [i32; 2] = [S, T];
#[derive(Debug, Clone)]
pub struct Mcp {
pub graph : Graph
}
impl Mcp {
pub fn new(g: Graph) -> Self { Mcp {graph: g} }
}
impl Problem<McpState> for Mcp {
fn nb_vars(&self) -> usize {
self.graph.nb_vertices
}
fn initial_state(&self) -> McpState {
McpState {initial: true, benef: vec![0; self.nb_vars()]}
}
fn initial_value(&self) -> i32 {
self.graph.sum_of_negative_edges()
}
fn domain_of<'a>(&self, state: &'a McpState, _var: Variable) -> Domain<'a> {
if state.initial { Domain::Slice(&ONLY_S) } else { Domain::Slice(&BOTH_ST) }
}
fn transition(&self, state: &McpState, vars: &VarSet, d: Decision) -> McpState {
let mut benefits = vec![0; self.nb_vars()];
for v in vars.iter() { benefits[v.id()] = state.benef[v.id()] + d.value * self.graph[(d.variable, v)];
}
McpState {initial: false, benef: benefits}
}
fn transition_cost(&self, state: &McpState, vars: &VarSet, d: Decision) -> i32 {
match d.value {
S => if state.initial { 0 } else { self.branch_on_s(state, vars, d) },
T => if state.initial { 0 } else { self.branch_on_t(state, vars, d) },
_ => unreachable!()
}
}
}
impl Mcp {
fn branch_on_s(&self, state: &McpState, vars: &VarSet, d: Decision) -> i32 {
let res = max(0, -state.benef[d.variable.id()]);
let mut sum = 0;
for v in vars.iter() {
let skl = state.benef[v.id()];
let wkl = self.graph[(d.variable, v)];
if skl * wkl <= 0 { sum += min(skl.abs(), wkl.abs()); }
}
res + sum
}
fn branch_on_t(&self, state: &McpState, vars: &VarSet, d: Decision) -> i32 {
let res = max(0, state.benef[d.variable.id()]);
let mut sum = 0;
for v in vars.iter() {
let skl = state.benef[v.id()];
let wkl = self.graph[(d.variable, v)];
if skl * wkl >= 0 { sum += min(skl.abs(), wkl.abs()); }
}
res + sum
}
}
impl From<Graph> for Mcp {
fn from(g: Graph) -> Self {
Mcp::new(g)
}
}
impl From<File> for Mcp {
fn from(f: File) -> Self {
Mcp::new(f.into())
}
}
impl <S: Read> From<BufReader<S>> for Mcp {
fn from(buf: BufReader<S>) -> Mcp {
Mcp::new(buf.into())
}
}
impl <B: BufRead> From<Lines<B>> for Mcp {
fn from(lines: Lines<B>) -> Mcp {
Mcp::new(lines.into())
}
}