1use std::collections::{BTreeMap, BTreeSet};
9
10use anyhow::{Context, Result, bail};
11use serde::{Deserialize, Serialize};
12use serde_json::Value;
13
14use crate::{Api, Operation};
15
16#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
18pub struct MockScenario {
19 pub operation_id: String,
22 pub name: String,
24 #[serde(default)]
25 pub when: MockRequestMatch,
26 pub response: MockResponse,
27}
28
29#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
36pub struct MockRequestMatch {
37 #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
38 pub headers: BTreeMap<String, String>,
39 #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
40 pub query: BTreeMap<String, String>,
41 #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
42 pub path: BTreeMap<String, String>,
43 #[serde(default, skip_serializing_if = "Option::is_none")]
44 pub body: Option<Value>,
45}
46
47#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
49pub struct MockResponse {
50 pub status: u16,
51 #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
52 pub headers: BTreeMap<String, String>,
53 #[serde(default, skip_serializing_if = "Option::is_none")]
54 pub body: Option<Value>,
55 #[serde(default, skip_serializing_if = "Option::is_none")]
58 pub delay_ms: Option<u64>,
59}
60
61pub fn mock_scenario_matches(
69 scenario: &MockScenario,
70 headers: &BTreeMap<String, String>,
71 query: &BTreeMap<String, String>,
72 path: &BTreeMap<String, String>,
73 body: Option<&Value>,
74) -> bool {
75 scenario.when.headers.iter().all(|(name, expected)| {
76 headers
77 .iter()
78 .any(|(actual, value)| actual.eq_ignore_ascii_case(name) && value == expected)
79 }) && scenario
80 .when
81 .query
82 .iter()
83 .all(|(name, expected)| query.get(name) == Some(expected))
84 && scenario
85 .when
86 .path
87 .iter()
88 .all(|(name, expected)| path.get(name) == Some(expected))
89 && scenario
90 .when
91 .body
92 .as_ref()
93 .is_none_or(|expected| body == Some(expected))
94}
95
96const EXTENSION: &str = "x-poolster-mock";
97const MAX_DELAY_MS: u64 = 600_000;
98
99pub fn extract_mock_scenarios(api: &Api) -> Result<Vec<MockScenario>> {
106 api.operations
107 .iter()
108 .map(extract_operation_mock_scenarios)
109 .collect::<Result<Vec<_>>>()
110 .map(|sets| sets.into_iter().flatten().collect())
111}
112
113pub fn extract_operation_mock_scenarios(operation: &Operation) -> Result<Vec<MockScenario>> {
115 let Some(value) = operation.annotations.get(EXTENSION) else {
116 return Ok(Vec::new());
117 };
118 let object = value.as_object().with_context(|| {
119 format!(
120 "{EXTENSION} on operation {:?} must be an object",
121 operation.id
122 )
123 })?;
124 let scenarios = object
125 .get("scenarios")
126 .with_context(|| {
127 format!(
128 "{EXTENSION} on operation {:?} requires scenarios",
129 operation.id
130 )
131 })?
132 .as_array()
133 .with_context(|| {
134 format!(
135 "{EXTENSION}.scenarios on operation {:?} must be an array",
136 operation.id
137 )
138 })?;
139
140 let mut names = BTreeSet::new();
141 scenarios
142 .iter()
143 .enumerate()
144 .map(|(index, scenario)| {
145 parse_scenario(operation, scenario, index).and_then(|scenario| {
146 if !names.insert(scenario.name.clone()) {
147 bail!(
148 "{EXTENSION}.scenarios on operation {:?} contains duplicate name {:?}",
149 operation.id,
150 scenario.name
151 );
152 }
153 Ok(scenario)
154 })
155 })
156 .collect()
157}
158
159fn parse_scenario(operation: &Operation, value: &Value, index: usize) -> Result<MockScenario> {
160 let context = format!(
161 "{EXTENSION}.scenarios[{index}] on operation {:?}",
162 operation.id
163 );
164 let object = value
165 .as_object()
166 .with_context(|| format!("{context} must be an object"))?;
167 reject_unknown(object, &["name", "when", "response"], &context)?;
168 let name = required_string(object, "name", &context)?;
169 if name.trim().is_empty() {
170 bail!("{context}.name must not be empty");
171 }
172 let when = object
173 .get("when")
174 .map(|value| parse_match(value, &format!("{context}.when")))
175 .transpose()?
176 .unwrap_or_default();
177 let response = object
178 .get("response")
179 .with_context(|| format!("{context} requires response"))
180 .and_then(|value| parse_response(value, &format!("{context}.response")))?;
181 Ok(MockScenario {
182 operation_id: operation.id.clone(),
183 name: name.to_owned(),
184 when,
185 response,
186 })
187}
188
189fn parse_match(value: &Value, context: &str) -> Result<MockRequestMatch> {
190 let object = value
191 .as_object()
192 .with_context(|| format!("{context} must be an object"))?;
193 reject_unknown(object, &["headers", "query", "path", "body"], context)?;
194 Ok(MockRequestMatch {
195 headers: optional_string_map(object, "headers", context)?,
196 query: optional_string_map(object, "query", context)?,
197 path: optional_string_map(object, "path", context)?,
198 body: object.get("body").cloned(),
199 })
200}
201
202fn parse_response(value: &Value, context: &str) -> Result<MockResponse> {
203 let object = value
204 .as_object()
205 .with_context(|| format!("{context} must be an object"))?;
206 reject_unknown(object, &["status", "headers", "body", "delay_ms"], context)?;
207 let status = object
208 .get("status")
209 .with_context(|| format!("{context} requires status"))?
210 .as_u64()
211 .with_context(|| format!("{context}.status must be an integer"))?;
212 if !(100..=599).contains(&status) {
213 bail!("{context}.status must be between 100 and 599");
214 }
215 let delay_ms = object
216 .get("delay_ms")
217 .map(|value| {
218 value
219 .as_u64()
220 .with_context(|| format!("{context}.delay_ms must be an integer"))
221 })
222 .transpose()?;
223 if delay_ms.is_some_and(|delay| delay > MAX_DELAY_MS) {
224 bail!("{context}.delay_ms must not exceed {MAX_DELAY_MS}");
225 }
226 Ok(MockResponse {
227 status: status as u16,
228 headers: optional_string_map(object, "headers", context)?,
229 body: object.get("body").cloned(),
230 delay_ms,
231 })
232}
233
234fn required_string<'a>(
235 object: &'a serde_json::Map<String, Value>,
236 key: &str,
237 context: &str,
238) -> Result<&'a str> {
239 object
240 .get(key)
241 .with_context(|| format!("{context} requires {key}"))?
242 .as_str()
243 .with_context(|| format!("{context}.{key} must be a string"))
244}
245
246fn optional_string_map(
247 object: &serde_json::Map<String, Value>,
248 key: &str,
249 context: &str,
250) -> Result<BTreeMap<String, String>> {
251 let Some(value) = object.get(key) else {
252 return Ok(BTreeMap::new());
253 };
254 let map = value
255 .as_object()
256 .with_context(|| format!("{context}.{key} must be an object"))?;
257 map.iter()
258 .map(|(name, value)| {
259 if name.is_empty() {
260 bail!("{context}.{key} must not contain an empty name");
261 }
262 let value = value
263 .as_str()
264 .with_context(|| format!("{context}.{key}.{name} must be a string"))?;
265 if (key == "headers") && (name.contains(['\r', '\n']) || value.contains(['\r', '\n'])) {
266 bail!("{context}.{key}.{name} must not contain a line break");
267 }
268 Ok((name.clone(), value.to_owned()))
269 })
270 .collect()
271}
272
273fn reject_unknown(
274 object: &serde_json::Map<String, Value>,
275 known: &[&str],
276 context: &str,
277) -> Result<()> {
278 if let Some(key) = object.keys().find(|key| !known.contains(&key.as_str())) {
279 bail!("{context} contains unsupported field {key:?}");
280 }
281 Ok(())
282}
283
284#[cfg(test)]
285mod tests {
286 use std::collections::BTreeMap;
287
288 use serde_json::json;
289
290 use super::*;
291 use crate::{HttpMethod, Operation};
292
293 fn operation(extension: Value) -> Operation {
294 Operation {
295 id: "listContacts".into(),
296 method: HttpMethod::Get,
297 path: "/v1/contacts/{contact_id}".into(),
298 annotations: BTreeMap::from([(EXTENSION.into(), extension)]),
299 ..Operation::default()
300 }
301 }
302
303 #[test]
304 fn extracts_typed_scenarios() {
305 let scenarios = extract_operation_mock_scenarios(&operation(json!({
306 "scenarios": [{
307 "name": "rate-limited",
308 "when": {
309 "headers": {"x-test-scenario": "rate-limited"},
310 "query": {"expand": "stats"},
311 "path": {"contact_id": "contact_123"},
312 "body": {"enabled": true}
313 },
314 "response": {
315 "status": 429,
316 "headers": {"retry-after": "1"},
317 "body": {"message": "Too many requests"},
318 "delay_ms": 25
319 }
320 }]
321 })))
322 .unwrap();
323 assert_eq!(scenarios.len(), 1);
324 assert_eq!(scenarios[0].operation_id, "listContacts");
325 assert_eq!(scenarios[0].when.path["contact_id"], "contact_123");
326 assert_eq!(scenarios[0].response.status, 429);
327 assert_eq!(scenarios[0].response.delay_ms, Some(25));
328 }
329
330 #[test]
331 fn rejects_invalid_and_ambiguous_scenarios() {
332 for extension in [
333 json!({"scenarios": [{"name": "bad", "response": {"status": 99}}]}),
334 json!({"scenarios": [{"name": "bad", "when": {"header": {}}, "response": {"status": 200}}]}),
335 json!({"scenarios": [{"name": "same", "response": {"status": 200}}, {"name": "same", "response": {"status": 201}}]}),
336 json!({"scenarios": [{"name": "slow", "response": {"status": 200, "delay_ms": 600001}}]}),
337 ] {
338 assert!(extract_operation_mock_scenarios(&operation(extension)).is_err());
339 }
340 }
341
342 #[test]
343 fn ignores_operations_without_mock_extensions() {
344 let api = Api {
345 operations: vec![Operation::default()],
346 ..Api::default()
347 };
348 assert!(extract_mock_scenarios(&api).unwrap().is_empty());
349 }
350
351 #[test]
352 fn matches_all_scenario_predicates_with_case_insensitive_headers() {
353 let scenario = extract_operation_mock_scenarios(&operation(json!({
354 "scenarios": [{
355 "name": "all-predicates",
356 "when": {
357 "headers": {"x-test-scenario": "selected"},
358 "query": {"expand": "stats"},
359 "path": {"contact_id": "contact_123"},
360 "body": {"enabled": true}
361 },
362 "response": {"status": 200}
363 }]
364 })))
365 .unwrap()
366 .remove(0);
367 let headers = BTreeMap::from([("X-Test-Scenario".into(), "selected".into())]);
368 let query = BTreeMap::from([("expand".into(), "stats".into())]);
369 let path = BTreeMap::from([("contact_id".into(), "contact_123".into())]);
370 assert!(mock_scenario_matches(
371 &scenario,
372 &headers,
373 &query,
374 &path,
375 Some(&json!({"enabled": true})),
376 ));
377 assert!(!mock_scenario_matches(
378 &scenario,
379 &headers,
380 &query,
381 &path,
382 Some(&json!({"enabled": false})),
383 ));
384 }
385}