mod gambit;
mod json;
use cfr::{Game, GameError, GameTree, PlayerNum, RegretParams, SolveMethod, SolveParams};
use gambit::{GambitError, GambitInfoset, GambitNode};
use gambit_parser::ExtensiveFormGame;
use clap::{Parser, ValueEnum};
use serde::Serialize;
use std::borrow::Borrow;
use std::collections::HashMap;
use std::fmt::{self, Display, Formatter};
use std::fs::File;
use std::hash::Hash;
use std::io;
use std::io::{BufReader, Read};
#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)]
enum Method {
Full,
Sampled,
External,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)]
enum InputFormat {
Auto,
Gambit,
Json,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)]
enum Discount {
Vanilla,
Lcfr,
CfrPlus,
Dcfr,
DcfrPrune,
}
impl Discount {
fn into_params(self) -> RegretParams {
match self {
Discount::Vanilla => RegretParams::vanilla(),
Discount::Lcfr => RegretParams::lcfr(),
Discount::CfrPlus => RegretParams::cfr_plus(),
Discount::Dcfr => RegretParams::dcfr(),
Discount::DcfrPrune => RegretParams::dcfr_prune(),
}
}
}
#[derive(Parser, Debug)]
#[clap(author, version, about)]
struct Args {
#[clap(short, long, value_parser, default_value_t = 0.0)]
clip_threshold: f64,
#[clap(short = 'r', long, value_parser, default_value_t = 0.0)]
max_regret: f64,
#[clap(short = 't', long, value_parser, default_value_t = 1000)]
max_iters: u64,
#[clap(short, long, value_parser, default_value_t = 0)]
parallel: usize,
#[clap(short, long, value_parser, default_value_t = 0)]
seed: u64,
#[clap(short, long, value_enum, default_value_t = Method::External)]
method: Method,
#[clap(long, value_enum, default_value_t = InputFormat::Auto)]
input_format: InputFormat,
#[clap(short, long, value_enum, default_value_t = Discount::Dcfr)]
discount: Discount,
#[clap(short, long, value_parser, default_value = "-")]
input: String,
#[clap(short, long, value_parser, default_value = "-")]
output: String,
#[clap(long, value_parser)]
infoset_numbers: bool,
}
#[derive(Debug, Serialize)]
struct Strategy(HashMap<String, HashMap<String, f64>>);
impl Strategy {
fn from_named<S, A, T, N>(
named_strategies: impl IntoIterator<Item = (S, A)>,
key: impl Fn(S) -> String,
) -> Result<Self, CliError>
where
A: IntoIterator<Item = (T, N)>,
T: Display,
N: Borrow<f64>,
{
let mut map = HashMap::new();
for (info, actions) in named_strategies {
let name = key(info);
let probs = actions
.into_iter()
.filter(|(_, p)| p.borrow() > &0.0)
.map(|(a, p)| (a.to_string(), *p.borrow()))
.collect();
if map.insert(name.clone(), probs).is_some() {
return Err(CliError::DuplicateInfoset(name));
}
}
Ok(Strategy(map))
}
}
#[derive(Debug, Serialize)]
struct Output {
regret: f64,
player_one_utility: f64,
player_two_utility: f64,
player_one_regret: f64,
player_two_regret: f64,
player_one_strategy: Strategy,
player_two_strategy: Strategy,
}
fn main() {
let args = Args::parse();
let mut buff = String::new();
if args.input == "-" {
io::stdin().lock().read_to_string(&mut buff).unwrap();
} else {
BufReader::new(File::open(&args.input).unwrap())
.read_to_string(&mut buff)
.unwrap();
}
let out = match args.input_format {
InputFormat::Json => solve_json(&buff, &args),
InputFormat::Gambit => solve_gambit(&buff, &args),
InputFormat::Auto if args.input.ends_with(".json") => solve_json(&buff, &args),
InputFormat::Auto if args.input.ends_with(".efg") => solve_gambit(&buff, &args),
InputFormat::Auto => solve_auto(&buff, &args),
}
.unwrap();
if args.output == "-" {
serde_json::to_writer(io::stdout(), &out).unwrap();
} else {
serde_json::to_writer(File::create(args.output).unwrap(), &out).unwrap();
};
}
enum CliError {
ParseJson(serde_json::Error),
ParseGambit(String),
Gambit(GambitError),
Materialize(GameError),
DuplicateInfoset(String),
UnknownFormat,
}
impl Display for CliError {
fn fmt(&self, out: &mut Formatter<'_>) -> fmt::Result {
match self {
CliError::ParseJson(err) => write!(
out,
"couldn't parse json game definition: {err} : https://github.com/erikbrinkman/cfr#json-error"
),
CliError::ParseGambit(err) => write!(
out,
"couldn't parse gambit game definition: {err} : https://github.com/erikbrinkman/cfr#gambit-error"
),
CliError::Gambit(err) => Display::fmt(err, out),
CliError::Materialize(err) => write!(
out,
"couldn't extract a compact game representation due to problems with the structure ({err}) : https://github.com/erikbrinkman/cfr#game-error"
),
CliError::DuplicateInfoset(name) => write!(
out,
"two distinct infosets share the name {name:?}; pass `--infoset-numbers` to key gambit infosets by number instead : https://github.com/erikbrinkman/cfr#duplicate-infosets"
),
CliError::UnknownFormat => write!(
out,
"couldn't parse any known format; try specifying your format with `--input-format` : https://github.com/erikbrinkman/cfr#auto-error"
),
}
}
}
impl fmt::Debug for CliError {
fn fmt(&self, out: &mut Formatter<'_>) -> fmt::Result {
Display::fmt(self, out)
}
}
fn solve_json(buff: &str, args: &Args) -> Result<Output, CliError> {
let state: json::State = serde_json::from_str(buff).map_err(CliError::ParseJson)?;
solve_game(&state, 0.0, args, |info| info.to_string())
}
fn solve_gambit_game(game: GambitNode<'_, '_>, sum: f64, args: &Args) -> Result<Output, CliError> {
if args.infoset_numbers {
solve_game(game, sum, args, |info: &GambitInfoset<'_>| {
info.number().to_string()
})
} else {
solve_game(game, sum, args, |info: &GambitInfoset<'_>| info.to_string())
}
}
fn solve_gambit(buff: &str, args: &Args) -> Result<Output, CliError> {
let parsed =
ExtensiveFormGame::try_from(buff).map_err(|err| CliError::ParseGambit(err.to_string()))?;
let game = GambitNode::try_from(&parsed).map_err(CliError::Gambit)?;
let sum = game.sum();
solve_gambit_game(game, sum, args)
}
fn solve_auto(buff: &str, args: &Args) -> Result<Output, CliError> {
if let Ok(state) = serde_json::from_str::<json::State>(buff) {
solve_game(&state, 0.0, args, |info| info.to_string())
} else if let Ok(parsed) = ExtensiveFormGame::try_from(buff) {
let game = GambitNode::try_from(&parsed).map_err(CliError::Gambit)?;
let sum = game.sum();
solve_gambit_game(game, sum, args)
} else {
Err(CliError::UnknownFormat)
}
}
fn solve_game<G>(
game: G,
sum: f64,
args: &Args,
infoset_key: impl Fn(&G::Infoset) -> String,
) -> Result<Output, CliError>
where
G: Game,
G::Infoset: Eq + Hash,
G::Action: Eq + Hash + Display,
G::ChanceInfoset: Eq + Hash,
{
let game = GameTree::from_game(game).map_err(CliError::Materialize)?;
let max_iters = if args.max_iters == 0 {
u64::MAX
} else {
args.max_iters
};
let method = match args.method {
Method::Full => SolveMethod::Full,
Method::Sampled => SolveMethod::Sampled,
Method::External => SolveMethod::External,
};
let (mut strategies, _) = game
.solve(
method,
max_iters,
args.max_regret,
args.parallel,
SolveParams {
regret: args.discount.into_params(),
seed: args.seed,
..Default::default()
},
)
.unwrap();
let info = strategies.get_info();
let [one, two] = strategies.as_named();
let unpruned = [
Strategy::from_named(one, &infoset_key)?,
Strategy::from_named(two, &infoset_key)?,
];
strategies.truncate(args.clip_threshold);
let pruned_info = strategies.get_info();
let (info, [player_one_strategy, player_two_strategy]) =
if pruned_info.regret() < info.regret() {
let [one, two] = strategies.as_named();
(
pruned_info,
[
Strategy::from_named(one, &infoset_key)?,
Strategy::from_named(two, &infoset_key)?,
],
)
} else {
(info, unpruned)
};
Ok(Output {
regret: info.regret(),
player_one_utility: info.player_utility(PlayerNum::One) + sum,
player_two_utility: info.player_utility(PlayerNum::Two) + sum,
player_one_regret: info.player_regret(PlayerNum::One),
player_two_regret: info.player_regret(PlayerNum::Two),
player_one_strategy,
player_two_strategy,
})
}
#[cfg(test)]
mod tests {
use super::{Args, CliError, GameError, GambitError, solve_auto, solve_gambit, solve_json};
use clap::{CommandFactory, Parser};
const PENNIES: &str = r#"{ "player": { "player_one": true, "infoset": "p1", "actions": {
"h": { "player": { "player_one": false, "infoset": "p2", "actions": {
"h": { "terminal": 1.0 }, "t": { "terminal": -1.0 } } } },
"t": { "player": { "player_one": false, "infoset": "p2", "actions": {
"h": { "terminal": -1.0 }, "t": { "terminal": 1.0 } } } } } } }"#;
fn default_args() -> Args {
Args::parse_from(["cfr"])
}
#[test]
fn test_cli() {
Args::command().debug_assert();
}
#[test]
fn auto_detects_json() {
solve_auto(r#"{ "terminal": 0.0 }"#, &default_args()).unwrap();
}
#[test]
fn auto_detects_gambit() {
solve_auto(r#"EFG 2 R "" { "" "" } t "" 1 "" { 0 0 }"#, &default_args()).unwrap();
}
#[test]
fn auto_rejects_unknown() {
let err = solve_auto("random", &default_args()).unwrap_err();
assert!(matches!(err, CliError::UnknownFormat));
}
#[test]
fn error_messages_render() {
let json_err = serde_json::from_str::<serde_json::Value>("not json").unwrap_err();
let errors = [
CliError::ParseJson(json_err),
CliError::ParseGambit("error parsing game at: 'oops'".to_owned()),
CliError::Gambit(GambitError::NotConstantSum),
CliError::Materialize(GameError::ImperfectRecall),
CliError::DuplicateInfoset("dup".to_owned()),
CliError::UnknownFormat,
];
for err in errors {
assert!(!err.to_string().is_empty());
assert_eq!(format!("{err:?}"), err.to_string());
}
}
#[test]
fn reports_input_errors() {
let args = default_args();
assert!(matches!(
solve_json("not json", &args),
Err(CliError::ParseJson(_))
));
assert!(matches!(
solve_gambit("not gambit", &args),
Err(CliError::ParseGambit(_))
));
assert!(matches!(
solve_gambit(r#"EFG 2 R "" { "" "" "" } t "" 1 "" { 0 0 0 }"#, &args),
Err(CliError::Gambit(_))
));
let bad = r#"{ "chance": { "outcomes": { "a": { "prob": 0.0, "state": { "terminal": 0.0 } } } } }"#;
assert!(matches!(
solve_json(bad, &args),
Err(CliError::Materialize(_))
));
}
#[test]
fn solves_across_methods_and_discounts() {
for method in ["full", "sampled", "external"] {
for discount in ["vanilla", "lcfr", "cfr-plus", "dcfr", "dcfr-prune"] {
let args = Args::parse_from(["cfr", "-m", method, "-d", discount, "-t", "50"]);
let out = solve_json(PENNIES, &args).unwrap();
assert!(out.regret.is_finite());
}
}
}
const SHARED_NAME_GAME: &str = r#"EFG 2 R "" { "" "" } p "" 1 1 "" { "a" "b" } 0 p "" 2 1 "" { "x" "y" } 0 t "" 1 "" { 1 -1 } t "" 2 "" { -1 1 } p "" 2 2 "" { "x" "y" } 0 t "" 3 "" { -1 1 } t "" 4 "" { 1 -1 }"#;
#[test]
fn duplicate_infoset_names_error_by_default() {
assert!(matches!(
solve_gambit(SHARED_NAME_GAME, &default_args()),
Err(CliError::DuplicateInfoset(_))
));
}
#[test]
fn infoset_numbers_flag_resolves_collisions() {
let numbered = Args::parse_from(["cfr", "--infoset-numbers"]);
let out = solve_gambit(SHARED_NAME_GAME, &numbered).unwrap();
assert_eq!(out.player_two_strategy.0.len(), 2, "both p2 infosets present");
}
#[test]
fn reports_both_constant_sum_utilities() {
let game = r#"EFG 2 R "" { "" "" } c "" 1 "c" { "x" 1/2 "y" 1/2 } 0 t "" 1 "" { 3 1 } t "" 2 "" { 1 3 }"#;
let out = solve_gambit(game, &default_args()).unwrap();
assert!(
(out.player_one_utility - 2.0).abs() < 1e-9,
"player one: {}",
out.player_one_utility
);
assert!(
(out.player_two_utility - 2.0).abs() < 1e-9,
"player two: {}",
out.player_two_utility
);
}
#[test]
fn clip_threshold_runs_the_pruning_path() {
let args = Args::parse_from(["cfr", "-c", "0.1", "-t", "200"]);
let out = solve_json(PENNIES, &args).unwrap();
assert!(out.regret.is_finite());
}
}