1use std::{
5 fs,
6 path::{Path, PathBuf},
7};
8
9use semver::Version;
10
11use crate::prelude::*;
12use crate::utils::version::FOREST_VERSION;
13
14pub(super) const FOREST_DB_DEV_MODE: &str = "FOREST_DB_DEV_MODE";
20
21fn list_versioned_databases(chain_data_path: &Path) -> anyhow::Result<Vec<Version>> {
24 let versions = fs::read_dir(chain_data_path)?
25 .filter_map(|entry| entry.ok())
26 .filter_map(|entry| {
27 let path = entry.path();
28 Version::parse(path.file_name()?.to_str()?).ok()
29 })
30 .collect();
31
32 Ok(versions)
33}
34
35pub(super) fn get_latest_versioned_database(
37 chain_data_path: &Path,
38) -> anyhow::Result<Option<Version>> {
39 let versions = list_versioned_databases(chain_data_path)?;
40 Ok(versions.iter().max().cloned())
41}
42
43pub fn choose_db(chain_data_path: &Path) -> anyhow::Result<PathBuf> {
46 let db = match DbMode::read() {
47 DbMode::Current => chain_data_path.join(FOREST_VERSION.to_string()),
48 DbMode::Latest => {
49 let versions = list_versioned_databases(chain_data_path)?;
50
51 if versions.is_empty() {
52 chain_data_path.join(FOREST_VERSION.to_string())
53 } else {
54 let latest = versions
55 .iter()
56 .max()
57 .context("Failed to find latest versioned database")?; chain_data_path.join(latest.to_string())
59 }
60 }
61 DbMode::Custom(custom) => chain_data_path.join(custom),
62 };
63
64 Ok(db)
65}
66
67#[derive(Debug, PartialEq, Eq, Clone)]
69pub enum DbMode {
70 Current,
73 Latest,
75 Custom(String),
77}
78
79impl DbMode {
80 pub fn read() -> Self {
82 match std::env::var(FOREST_DB_DEV_MODE)
83 .map(|s| s.to_lowercase())
84 .as_deref()
85 {
86 Ok("latest") => Self::Latest,
87 Ok("current") | Err(_) => Self::Current,
88 Ok(val) => Self::Custom(val.to_owned()),
89 }
90 }
91}
92
93#[cfg(test)]
94mod tests {
95 use super::*;
96 use std::env;
97
98 #[test]
99 fn test_db_mode() {
100 unsafe {
101 env::set_var(FOREST_DB_DEV_MODE, "latest");
102 assert_eq!(DbMode::read(), DbMode::Latest);
103
104 env::set_var(FOREST_DB_DEV_MODE, "current");
105 assert_eq!(DbMode::read(), DbMode::Current);
106
107 env::set_var(FOREST_DB_DEV_MODE, "cthulhu");
108 assert_eq!(DbMode::read(), DbMode::Custom("cthulhu".to_owned()));
109
110 env::remove_var(FOREST_DB_DEV_MODE);
111 assert_eq!(DbMode::read(), DbMode::Current);
112 }
113 }
114
115 #[test]
116 fn test_list_versioned_databases() {
117 use tempfile::tempdir;
118
119 let dir = tempdir().unwrap();
120 let path = dir.path();
121
122 for dir in &["0.1.0", "0.2.0", "0.3.0", "Elder God", "my0.4.0"] {
123 std::fs::create_dir(path.join(dir)).unwrap();
124 }
125
126 let versions = list_versioned_databases(path)
127 .unwrap()
128 .iter()
129 .sorted()
130 .cloned()
131 .collect_vec();
132 assert_eq!(
133 versions,
134 vec![
135 Version::parse("0.1.0").unwrap(),
136 Version::parse("0.2.0").unwrap(),
137 Version::parse("0.3.0").unwrap()
138 ]
139 );
140 }
141
142 #[test]
143 fn test_choose_db() {
144 use tempfile::tempdir;
145
146 let dir = tempdir().unwrap();
147 let path = dir.path();
148
149 for dir in &["0.1.0", "0.2.0", "0.3.0", "Elder God", "my0.4.0"] {
150 std::fs::create_dir(path.join(dir)).unwrap();
151 }
152
153 let cases = [
154 ("latest", path.join("0.3.0")),
155 ("current", path.join(FOREST_VERSION.to_string())),
156 ("cthulhu", path.join("cthulhu")),
157 ];
158
159 for (mode, expected) in &cases {
160 unsafe { env::set_var(FOREST_DB_DEV_MODE, mode) };
161 let db = choose_db(path).unwrap();
162 assert_eq!(db, *expected);
163 }
164
165 unsafe { env::remove_var(FOREST_DB_DEV_MODE) };
166 let db = choose_db(path).unwrap();
167 assert_eq!(db, path.join(FOREST_VERSION.to_string()));
168 }
169}