1use std::path::Path;
2use std::time::Duration;
3
4use serde::Deserialize;
5
6use super::scenario::{
7 EquivalenceTarget, NormalizationRules, OracleScenario, PlatformRequirements, ScenarioCategory,
8};
9
10const CURRENT_SCHEMA_VERSION: u32 = 1;
11
12fn default_timeout_secs() -> u64 {
13 15
14}
15
16#[derive(Debug, Clone, Deserialize)]
17pub struct ScenarioFile {
18 pub schema_version: u32,
19 pub scenarios: Vec<ScenarioDef>,
20}
21
22#[derive(Debug, Clone, Deserialize)]
23pub struct ScenarioDef {
24 pub id: String,
25 pub capability_ids: Vec<String>,
26 pub description: String,
27 pub pproxy_args: Vec<String>,
28 pub eggress_toml: String,
29 pub expected_equivalence: EquivalenceTarget,
30 pub category: ScenarioCategory,
31 #[serde(default)]
32 pub normalization: NormalizationRulesDef,
33 #[serde(default)]
34 pub platform: PlatformRequirementsDef,
35 #[serde(default = "default_timeout_secs")]
36 pub timeout_secs: u64,
37 pub client_action: ClientAction,
38 #[serde(default)]
39 pub comparison: ComparisonMode,
40 #[serde(default)]
41 pub expected_divergences: Vec<String>,
42}
43
44#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
45#[serde(rename_all = "snake_case")]
46pub enum ClientAction {
47 Socks5TcpConnect,
48 HttpConnect,
49 HttpForwardGet,
50 HttpForwardPost,
51 Socks5ConnectRefused,
52 Socks5AuthFailure,
53 Socks5TcpConnectAuth,
54 Socks4Connect,
55 Socks4aConnect,
56 UdpEchoRoundtrip,
57 None,
58}
59
60#[derive(Debug, Clone, Default, Deserialize, PartialEq, Eq)]
61#[serde(rename_all = "snake_case")]
62pub enum ComparisonMode {
63 #[default]
64 ExactPayload,
65 CoarseResult,
66 StatusCode,
67 BindAddress,
68}
69
70#[derive(Debug, Clone, Deserialize)]
71pub struct NormalizationRulesDef {
72 #[serde(default = "default_true")]
73 pub strip_log_prefixes: bool,
74 #[serde(default = "default_true")]
75 pub normalize_ports: bool,
76 #[serde(default)]
77 pub normalize_line_endings: bool,
78 #[serde(default)]
79 pub strip_versions: bool,
80}
81
82fn default_true() -> bool {
83 true
84}
85
86impl Default for NormalizationRulesDef {
87 fn default() -> Self {
88 Self {
89 strip_log_prefixes: true,
90 normalize_ports: true,
91 normalize_line_endings: false,
92 strip_versions: false,
93 }
94 }
95}
96
97#[derive(Debug, Clone, Deserialize, Default)]
98pub struct PlatformRequirementsDef {
99 #[serde(default)]
100 pub requires_root: bool,
101 #[serde(default)]
102 pub requires_ipv6: bool,
103 pub required_os: Option<String>,
104}
105
106#[derive(Debug, thiserror::Error)]
107pub enum ScenarioValidationError {
108 #[error("unsupported schema version {0} (expected {CURRENT_SCHEMA_VERSION})")]
109 UnsupportedSchemaVersion(u32),
110 #[error("scenario ID '{0}' is empty")]
111 EmptyId(String),
112 #[error("duplicate scenario ID: '{0}'")]
113 DuplicateId(String),
114 #[error("scenario '{0}' has no capability IDs")]
115 NoCapabilityIds(String),
116 #[error("scenario '{0}' has empty pproxy_args")]
117 EmptyPproxyArgs(String),
118 #[error("scenario '{0}' has empty eggress_toml")]
119 EmptyEggressToml(String),
120 #[error("scenario '{0}' has empty description")]
121 EmptyDescription(String),
122 #[error("TOML parse error: {0}")]
123 TomlParse(#[from] toml::de::Error),
124 #[error("I/O error: {0}")]
125 Io(#[from] std::io::Error),
126}
127
128pub type ScenarioValidationErrors = Vec<ScenarioValidationError>;
129
130pub fn validate_scenario_file(file: &ScenarioFile) -> ScenarioValidationErrors {
131 let mut errors = Vec::new();
132
133 if file.schema_version != CURRENT_SCHEMA_VERSION {
134 errors.push(ScenarioValidationError::UnsupportedSchemaVersion(
135 file.schema_version,
136 ));
137 }
138
139 let mut seen_ids = std::collections::HashSet::new();
140 for scenario in &file.scenarios {
141 if scenario.id.is_empty() {
142 errors.push(ScenarioValidationError::EmptyId("<unnamed>".to_string()));
143 } else if !seen_ids.insert(scenario.id.clone()) {
144 errors.push(ScenarioValidationError::DuplicateId(scenario.id.clone()));
145 }
146
147 if scenario.capability_ids.is_empty() {
148 errors.push(ScenarioValidationError::NoCapabilityIds(
149 scenario.id.clone(),
150 ));
151 }
152
153 if scenario.pproxy_args.is_empty() {
154 errors.push(ScenarioValidationError::EmptyPproxyArgs(
155 scenario.id.clone(),
156 ));
157 }
158
159 if scenario.eggress_toml.is_empty() {
160 errors.push(ScenarioValidationError::EmptyEggressToml(
161 scenario.id.clone(),
162 ));
163 }
164
165 if scenario.description.is_empty() {
166 errors.push(ScenarioValidationError::EmptyDescription(
167 scenario.id.clone(),
168 ));
169 }
170 }
171
172 errors
173}
174
175pub fn load_scenario_string(toml_str: &str) -> Result<ScenarioFile, ScenarioValidationError> {
176 let file: ScenarioFile = toml::from_str(toml_str)?;
177 let errors = validate_scenario_file(&file);
178 if errors.is_empty() {
179 Ok(file)
180 } else {
181 Err(errors.into_iter().next().unwrap())
182 }
183}
184
185pub async fn load_scenario_file(path: &Path) -> Result<ScenarioFile, ScenarioValidationError> {
186 let content = tokio::fs::read_to_string(path).await?;
187 load_scenario_string(&content)
188}
189
190pub fn scenario_def_to_oracle(def: &ScenarioDef) -> OracleScenario {
191 let leak = |s: &str| -> &'static str { Box::leak(s.to_string().into_boxed_str()) };
192
193 OracleScenario {
194 id: leak(&def.id),
195 capability_ids: def.capability_ids.iter().map(|s| leak(s)).collect(),
196 description: leak(&def.description),
197 pproxy_args: def.pproxy_args.iter().map(|s| leak(s)).collect(),
198 eggress_toml: leak(&def.eggress_toml),
199 expected_equivalence: def.expected_equivalence,
200 normalization: NormalizationRules {
201 strip_log_prefixes: def.normalization.strip_log_prefixes,
202 normalize_ports: def.normalization.normalize_ports,
203 normalize_line_endings: def.normalization.normalize_line_endings,
204 strip_versions: def.normalization.strip_versions,
205 },
206 platform: PlatformRequirements {
207 requires_root: def.platform.requires_root,
208 requires_ipv6: def.platform.requires_ipv6,
209 requires_python_package: None,
210 required_os: def.platform.required_os.as_deref().map(leak),
211 },
212 timeout: Duration::from_secs(def.timeout_secs),
213 category: def.category,
214 }
215}
216
217#[cfg(test)]
218mod tests {
219 use super::*;
220
221 #[test]
222 fn schema_validation_rejects_bad_version() {
223 let file = ScenarioFile {
224 schema_version: 999,
225 scenarios: vec![],
226 };
227 let errors = validate_scenario_file(&file);
228 assert_eq!(errors.len(), 1);
229 assert!(matches!(
230 errors[0],
231 ScenarioValidationError::UnsupportedSchemaVersion(999)
232 ));
233 }
234
235 #[test]
236 fn schema_validation_rejects_empty_id() {
237 let file = ScenarioFile {
238 schema_version: 1,
239 scenarios: vec![ScenarioDef {
240 id: String::new(),
241 capability_ids: vec!["cap1".to_string()],
242 description: "test".to_string(),
243 pproxy_args: vec!["-l".to_string()],
244 eggress_toml: "test".to_string(),
245 expected_equivalence: EquivalenceTarget::Payload,
246 category: ScenarioCategory::CliDefaults,
247 normalization: NormalizationRulesDef::default(),
248 platform: PlatformRequirementsDef::default(),
249 timeout_secs: 10,
250 client_action: ClientAction::None,
251 comparison: ComparisonMode::default(),
252 expected_divergences: vec![],
253 }],
254 };
255 let errors = validate_scenario_file(&file);
256 assert!(errors
257 .iter()
258 .any(|e| matches!(e, ScenarioValidationError::EmptyId(_))));
259 }
260
261 #[test]
262 fn schema_validation_rejects_duplicate_ids() {
263 let def = ScenarioDef {
264 id: "dup".to_string(),
265 capability_ids: vec!["cap1".to_string()],
266 description: "test".to_string(),
267 pproxy_args: vec!["-l".to_string()],
268 eggress_toml: "test".to_string(),
269 expected_equivalence: EquivalenceTarget::Payload,
270 category: ScenarioCategory::CliDefaults,
271 normalization: NormalizationRulesDef::default(),
272 platform: PlatformRequirementsDef::default(),
273 timeout_secs: 10,
274 client_action: ClientAction::None,
275 comparison: ComparisonMode::default(),
276 expected_divergences: vec![],
277 };
278 let file = ScenarioFile {
279 schema_version: 1,
280 scenarios: vec![def.clone(), def],
281 };
282 let errors = validate_scenario_file(&file);
283 assert!(errors
284 .iter()
285 .any(|e| matches!(e, ScenarioValidationError::DuplicateId(s) if s == "dup")));
286 }
287
288 #[test]
289 fn schema_validation_rejects_empty_capability_ids() {
290 let file = ScenarioFile {
291 schema_version: 1,
292 scenarios: vec![ScenarioDef {
293 id: "test".to_string(),
294 capability_ids: vec![],
295 description: "test".to_string(),
296 pproxy_args: vec!["-l".to_string()],
297 eggress_toml: "test".to_string(),
298 expected_equivalence: EquivalenceTarget::Payload,
299 category: ScenarioCategory::CliDefaults,
300 normalization: NormalizationRulesDef::default(),
301 platform: PlatformRequirementsDef::default(),
302 timeout_secs: 10,
303 client_action: ClientAction::None,
304 comparison: ComparisonMode::default(),
305 expected_divergences: vec![],
306 }],
307 };
308 let errors = validate_scenario_file(&file);
309 assert!(errors
310 .iter()
311 .any(|e| matches!(e, ScenarioValidationError::NoCapabilityIds(_))));
312 }
313
314 #[test]
315 fn schema_validation_accepts_valid_file() {
316 let file = ScenarioFile {
317 schema_version: 1,
318 scenarios: vec![ScenarioDef {
319 id: "valid".to_string(),
320 capability_ids: vec!["cap1".to_string()],
321 description: "a valid scenario".to_string(),
322 pproxy_args: vec!["-l".to_string(), "socks5://127.0.0.1:1080".to_string()],
323 eggress_toml: "version = 1".to_string(),
324 expected_equivalence: EquivalenceTarget::CoarseResult,
325 category: ScenarioCategory::HttpSocksTcp,
326 normalization: NormalizationRulesDef::default(),
327 platform: PlatformRequirementsDef::default(),
328 timeout_secs: 10,
329 client_action: ClientAction::Socks5TcpConnect,
330 comparison: ComparisonMode::default(),
331 expected_divergences: vec![],
332 }],
333 };
334 let errors = validate_scenario_file(&file);
335 assert!(errors.is_empty());
336 }
337
338 #[test]
339 fn load_minimal_scenario() {
340 let toml_str = r#"
341schema_version = 1
342
343[[scenarios]]
344id = "minimal.test"
345capability_ids = ["cap1"]
346description = "minimal scenario"
347pproxy_args = ["-l", "socks5://127.0.0.1:1080"]
348eggress_toml = "version = 1"
349expected_equivalence = "coarse_result"
350category = "cli_defaults"
351timeout_secs = 5
352client_action = "none"
353"#;
354 let file = load_scenario_string(toml_str).unwrap();
355 assert_eq!(file.schema_version, 1);
356 assert_eq!(file.scenarios.len(), 1);
357 assert_eq!(file.scenarios[0].id, "minimal.test");
358 assert_eq!(file.scenarios[0].timeout_secs, 5);
359 assert_eq!(
360 file.scenarios[0].expected_equivalence,
361 EquivalenceTarget::CoarseResult
362 );
363 assert_eq!(file.scenarios[0].category, ScenarioCategory::CliDefaults);
364 assert_eq!(file.scenarios[0].client_action, ClientAction::None);
365 }
366
367 #[test]
368 fn scenario_def_to_oracle_roundtrip() {
369 let def = ScenarioDef {
370 id: "roundtrip.test".to_string(),
371 capability_ids: vec!["cap1".to_string(), "cap2".to_string()],
372 description: "roundtrip test".to_string(),
373 pproxy_args: vec!["-l".to_string(), "socks5://127.0.0.1:1080".to_string()],
374 eggress_toml: "version = 1".to_string(),
375 expected_equivalence: EquivalenceTarget::Payload,
376 category: ScenarioCategory::Chains,
377 normalization: NormalizationRulesDef {
378 strip_log_prefixes: true,
379 normalize_ports: false,
380 normalize_line_endings: true,
381 strip_versions: true,
382 },
383 platform: PlatformRequirementsDef {
384 requires_root: true,
385 requires_ipv6: false,
386 required_os: Some("linux".to_string()),
387 },
388 timeout_secs: 30,
389 client_action: ClientAction::Socks5TcpConnect,
390 comparison: ComparisonMode::ExactPayload,
391 expected_divergences: vec!["div1".to_string()],
392 };
393
394 let oracle = scenario_def_to_oracle(&def);
395 assert_eq!(oracle.id, "roundtrip.test");
396 assert_eq!(oracle.capability_ids, vec!["cap1", "cap2"]);
397 assert_eq!(oracle.description, "roundtrip test");
398 assert_eq!(oracle.pproxy_args, vec!["-l", "socks5://127.0.0.1:1080"]);
399 assert_eq!(oracle.eggress_toml, "version = 1");
400 assert_eq!(oracle.expected_equivalence, EquivalenceTarget::Payload);
401 assert_eq!(oracle.category, ScenarioCategory::Chains);
402 assert!(oracle.normalization.strip_log_prefixes);
403 assert!(!oracle.normalization.normalize_ports);
404 assert!(oracle.normalization.normalize_line_endings);
405 assert!(oracle.normalization.strip_versions);
406 assert!(oracle.platform.requires_root);
407 assert!(!oracle.platform.requires_ipv6);
408 assert_eq!(oracle.platform.required_os, Some("linux"));
409 assert_eq!(oracle.timeout, Duration::from_secs(30));
410 }
411
412 #[test]
413 fn default_normalization_rules() {
414 let rules = NormalizationRulesDef::default();
415 assert!(rules.strip_log_prefixes);
416 assert!(rules.normalize_ports);
417 assert!(!rules.normalize_line_endings);
418 assert!(!rules.strip_versions);
419 }
420
421 #[test]
422 fn default_comparison_mode() {
423 assert_eq!(ComparisonMode::default(), ComparisonMode::ExactPayload);
424 }
425
426 #[test]
427 fn load_scenario_with_defaults() {
428 let toml_str = r#"
429schema_version = 1
430
431[[scenarios]]
432id = "defaults.test"
433capability_ids = ["cap1"]
434description = "uses all defaults"
435pproxy_args = ["-l", "socks5://127.0.0.1:1080"]
436eggress_toml = "version = 1"
437expected_equivalence = "payload"
438category = "udp"
439client_action = "udp_echo_roundtrip"
440"#;
441 let file = load_scenario_string(toml_str).unwrap();
442 let scenario = &file.scenarios[0];
443 assert_eq!(scenario.timeout_secs, 15);
444 assert!(scenario.normalization.strip_log_prefixes);
445 assert!(scenario.normalization.normalize_ports);
446 assert!(!scenario.normalization.normalize_line_endings);
447 assert!(!scenario.normalization.strip_versions);
448 assert!(!scenario.platform.requires_root);
449 assert!(!scenario.platform.requires_ipv6);
450 assert!(scenario.platform.required_os.is_none());
451 assert!(scenario.expected_divergences.is_empty());
452 assert_eq!(scenario.comparison, ComparisonMode::ExactPayload);
453 }
454
455 #[test]
456 fn load_real_scenario_files() {
457 let scenarios_dir = Path::new(env!("CARGO_MANIFEST_DIR")).join("tests/oracle/scenarios");
458 if !scenarios_dir.exists() {
459 return;
460 }
461 let entries = match std::fs::read_dir(&scenarios_dir) {
462 Ok(e) => e,
463 Err(_) => return,
464 };
465 for entry in entries.flatten() {
466 let path = entry.path();
467 if path.extension().and_then(|e| e.to_str()) == Some("toml") {
468 let content = std::fs::read_to_string(&path).unwrap_or_else(|e| {
469 panic!("failed to read {}: {e}", path.display());
470 });
471 let file = load_scenario_string(&content).unwrap_or_else(|e| {
472 panic!("failed to validate {}: {e}", path.display());
473 });
474 let errors = validate_scenario_file(&file);
475 assert!(
476 errors.is_empty(),
477 "validation errors in {}: {:?}",
478 path.display(),
479 errors
480 );
481 }
482 }
483 }
484}