use std::fmt;
use std::io;
use std::path::{Path, PathBuf};
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use turnbase::{Game, PlayerId, Prng};
use turnbase_protocol::{PROTOCOL_VERSION, Request, Response};
use crate::{LocalSession, Session};
#[derive(Serialize, Deserialize)]
#[serde(bound(
serialize = "G: Serialize, G::State: Serialize",
deserialize = "G: Deserialize<'de>, G::State: Deserialize<'de>"
))]
struct SaveFile<G: Game> {
protocol_version: u32,
version: u64,
chance: Prng,
game: G,
state: G::State,
}
pub struct FileSession;
impl FileSession {
pub fn create<G>(
game: G,
path: Option<PathBuf>,
state: G::State,
seed: u64,
) -> Result<PathBuf, Error>
where
G: Game + Serialize,
G::State: Serialize,
{
let path = path.unwrap_or_else(generate_temp_path);
if path.exists() {
return Err(Error::AlreadyExists(path));
}
let (game, state, version, chance) = LocalSession::new(game, state, seed).into_parts();
write_save(
&path,
&SaveFile {
protocol_version: PROTOCOL_VERSION,
version,
chance,
game,
state,
},
)?;
Ok(path)
}
pub fn handle<G>(
path: &Path,
player: PlayerId,
request: Request<G::Action>,
) -> Result<Response<G::View>, Error>
where
G: Game + Serialize + DeserializeOwned,
G::State: Serialize + DeserializeOwned,
{
if !path.exists() {
return Err(Error::NotFound(path.to_path_buf()));
}
let save: SaveFile<G> = read_save(path)?;
let mut session = LocalSession::resume(save.game, save.state, save.version, save.chance);
let response = session.submit(player, request);
if matches!(response, Response::Ack) {
let (game, state, version, chance) = session.into_parts();
write_save(
path,
&SaveFile {
protocol_version: PROTOCOL_VERSION,
version,
chance,
game,
state,
},
)?;
}
Ok(response)
}
}
#[derive(Debug)]
pub enum Error {
AlreadyExists(PathBuf),
NotFound(PathBuf),
ProtocolMismatch {
found: u32,
expected: u32,
},
Io(io::Error),
Serde(serde_json::Error),
}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::AlreadyExists(path) => write!(f, "session already exists at {}", path.display()),
Self::NotFound(path) => write!(f, "no session found at {}", path.display()),
Self::ProtocolMismatch { found, expected } => write!(
f,
"save file protocol version {found} is incompatible with {expected}"
),
Self::Io(err) => write!(f, "session I/O error: {err}"),
Self::Serde(err) => write!(f, "session data error: {err}"),
}
}
}
impl std::error::Error for Error {}
fn generate_temp_path() -> PathBuf {
std::env::temp_dir().join(format!(
"turnbase-{}-{}.json",
std::process::id(),
unique_nanos()
))
}
fn unique_nanos() -> u128 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos()
}
fn read_save<G>(path: &Path) -> Result<SaveFile<G>, Error>
where
G: Game + DeserializeOwned,
G::State: DeserializeOwned,
{
let bytes = std::fs::read(path).map_err(Error::Io)?;
let save: SaveFile<G> = serde_json::from_slice(&bytes).map_err(Error::Serde)?;
if save.protocol_version != PROTOCOL_VERSION {
return Err(Error::ProtocolMismatch {
found: save.protocol_version,
expected: PROTOCOL_VERSION,
});
}
Ok(save)
}
fn write_save<G>(path: &Path, save: &SaveFile<G>) -> Result<(), Error>
where
G: Game + Serialize,
G::State: Serialize,
{
let bytes = serde_json::to_vec_pretty(save).map_err(Error::Serde)?;
let tmp = path.with_extension(format!("tmp-{}-{}", std::process::id(), unique_nanos()));
std::fs::write(&tmp, &bytes).map_err(Error::Io)?;
std::fs::rename(&tmp, path).map_err(|err| {
let _ = std::fs::remove_file(&tmp);
Error::Io(err)
})
}
#[cfg(test)]
mod tests {
use super::{Error, FileSession};
use serde::{Deserialize, Serialize};
use turnbase::{ActivePlayers, Game, PlayerId};
use turnbase_protocol::{Request, Response};
const P0: PlayerId = PlayerId::new(0);
const P1: PlayerId = PlayerId::new(1);
#[derive(Serialize, Deserialize)]
struct CountToThree;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
struct Bump;
impl Game for CountToThree {
type State = u32;
type Action = Bump;
type View = u32;
fn new_initial_state(&self, _seed: u64) -> Self::State {
0
}
fn num_players(&self) -> usize {
2
}
fn active_players(&self, state: &Self::State) -> ActivePlayers {
if self.is_terminal(state) {
ActivePlayers::none()
} else {
ActivePlayers::one(PlayerId::new(state % 2))
}
}
fn legal_actions(&self, state: &Self::State, _player: PlayerId) -> Vec<Self::Action> {
if self.is_terminal(state) {
Vec::new()
} else {
vec![Bump]
}
}
fn apply(&self, state: &mut Self::State, _player: PlayerId, _action: Self::Action) {
*state += 1;
}
fn is_terminal(&self, state: &Self::State) -> bool {
*state >= 3
}
fn reward(&self, state: &Self::State, player: PlayerId) -> f64 {
let winner = (state + 1) % 2;
if player.index() == winner { 1.0 } else { -1.0 }
}
fn view(&self, state: &Self::State, _viewer: Option<PlayerId>) -> Self::View {
*state
}
}
#[derive(Serialize, Deserialize)]
struct Reveal;
impl Game for Reveal {
type State = Option<u8>;
type Action = u8;
type View = Option<u8>;
fn new_initial_state(&self, _seed: u64) -> Self::State {
None
}
fn num_players(&self) -> usize {
0
}
fn active_players(&self, state: &Self::State) -> ActivePlayers {
if state.is_some() {
ActivePlayers::none()
} else {
ActivePlayers::one(PlayerId::CHANCE)
}
}
fn legal_actions(&self, state: &Self::State, player: PlayerId) -> Vec<Self::Action> {
if player.is_chance() && state.is_none() {
vec![0, 1, 2]
} else {
Vec::new()
}
}
fn apply(&self, state: &mut Self::State, _player: PlayerId, action: Self::Action) {
*state = Some(action);
}
fn is_terminal(&self, state: &Self::State) -> bool {
state.is_some()
}
fn reward(&self, _state: &Self::State, _player: PlayerId) -> f64 {
0.0
}
fn view(&self, state: &Self::State, _viewer: Option<PlayerId>) -> Self::View {
*state
}
}
fn bump() -> Request<Bump> {
Request::Act(Bump)
}
fn state_of(response: &Response<u32>) -> (u64, u32) {
let Response::State { version, view } = response else {
panic!("expected a State response, got {response:?}");
};
(*version, *view)
}
#[test]
fn create_refuses_to_overwrite_an_existing_file() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("game.json");
std::fs::write(&path, b"anything").unwrap();
let result = FileSession::create(CountToThree, Some(path), 0, 0);
assert!(matches!(result, Err(Error::AlreadyExists(_))));
}
#[test]
fn handle_round_trips_state_and_bumps_version_only_on_apply() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("game.json");
let path = FileSession::create(CountToThree, Some(path), 0, 0).unwrap();
let response = FileSession::handle::<CountToThree>(&path, P0, Request::Query).unwrap();
assert_eq!(state_of(&response), (0, 0));
assert!(matches!(
FileSession::handle::<CountToThree>(&path, P0, bump()).unwrap(),
Response::Ack
));
let response = FileSession::handle::<CountToThree>(&path, P1, Request::Query).unwrap();
assert_eq!(state_of(&response), (1, 1));
}
#[test]
fn a_rejected_action_does_not_persist() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("game.json");
let path = FileSession::create(CountToThree, Some(path), 0, 0).unwrap();
assert!(matches!(
FileSession::handle::<CountToThree>(&path, P1, bump()).unwrap(),
Response::Error(_)
));
let response = FileSession::handle::<CountToThree>(&path, P0, Request::Query).unwrap();
assert_eq!(state_of(&response), (0, 0), "version stayed at 0");
}
#[test]
fn handle_errors_on_a_missing_file() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("missing.json");
let result = FileSession::handle::<CountToThree>(&path, P0, Request::Query);
assert!(matches!(result, Err(Error::NotFound(_))));
}
#[test]
fn a_full_game_reaches_a_terminal_state_through_the_file() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("game.json");
let path = FileSession::create(CountToThree, Some(path), 0, 0).unwrap();
for round in 0..3 {
let seat = PlayerId::new(round % 2);
FileSession::handle::<CountToThree>(&path, seat, bump()).unwrap();
}
let response = FileSession::handle::<CountToThree>(&path, P0, Request::Query).unwrap();
assert_eq!(
state_of(&response),
(3, 3),
"three bumps, version 3, total 3"
);
}
#[test]
fn create_resolves_an_opening_chance_node_and_persists_it() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("chance.json");
let path = FileSession::create(Reveal, Some(path), None, 5).unwrap();
let response = FileSession::handle::<Reveal>(&path, P0, Request::Query).unwrap();
let Response::State { version, view } = response else {
panic!("expected a state response");
};
assert_eq!(version, 1);
assert!(matches!(view, Some(0..=2)));
}
}