use std::collections::BTreeMap;
use franken_snowflake_core::redact::redact;
use super::http::{Method, MockHttpRequest, MockHttpResponse};
const SUBMIT_PATH: &str = "/api/v2/statements";
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RecordedRequest {
pub method: Method,
pub path: String,
pub redacted_authorization: Option<String>,
pub header_count: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
enum Route {
Submit,
Statement(String),
Partition { handle: String, partition: u32 },
Cancel(String),
Unknown,
}
fn route(path: &str) -> Route {
let mut pieces = path.splitn(2, '?');
let route_path = pieces.next().unwrap_or(path);
let query = pieces.next();
if route_path == SUBMIT_PATH {
return Route::Submit;
}
if let Some(rest) = route_path.strip_prefix("/api/v2/statements/") {
if let Some(handle) = rest.strip_suffix("/cancel") {
if !handle.is_empty() && !handle.contains('/') {
return Route::Cancel(handle.to_owned());
}
} else if !rest.is_empty() && !rest.contains('/') {
if let Some(partition) = partition_query_value(query) {
return Route::Partition {
handle: rest.to_owned(),
partition,
};
}
return Route::Statement(rest.to_owned());
}
}
Route::Unknown
}
fn partition_query_value(query: Option<&str>) -> Option<u32> {
query.and_then(|query| {
query.split('&').find_map(|pair| {
let (key, value) = pair.split_once('=')?;
if key == "partition" {
value.parse::<u32>().ok()
} else {
None
}
})
})
}
#[derive(Clone, Debug)]
pub struct MockSqlApi {
statement_handle: String,
running: MockHttpResponse,
terminal: MockHttpResponse,
cancel: MockHttpResponse,
polls_before_complete: u32,
immediate: bool,
partitions: BTreeMap<u32, MockHttpResponse>,
poll_counts: BTreeMap<String, u32>,
cancelled: BTreeMap<String, bool>,
log: Vec<RecordedRequest>,
}
impl MockSqlApi {
#[must_use]
pub fn new(
statement_handle: impl Into<String>,
running: MockHttpResponse,
terminal: MockHttpResponse,
cancel: MockHttpResponse,
) -> Self {
Self {
statement_handle: statement_handle.into(),
running,
terminal,
cancel,
polls_before_complete: 1,
immediate: false,
partitions: BTreeMap::new(),
poll_counts: BTreeMap::new(),
cancelled: BTreeMap::new(),
log: Vec::new(),
}
}
#[must_use]
pub fn with_polls_before_complete(mut self, polls: u32) -> Self {
self.polls_before_complete = polls;
self
}
#[must_use]
pub fn immediate(mut self) -> Self {
self.immediate = true;
self
}
#[must_use]
pub fn with_partition(mut self, partition: u32, response: MockHttpResponse) -> Self {
self.partitions.insert(partition, response);
self
}
#[must_use]
pub fn statement_handle(&self) -> &str {
&self.statement_handle
}
pub fn respond(&mut self, request: &MockHttpRequest) -> MockHttpResponse {
self.log.push(RecordedRequest {
method: request.method.clone(),
path: redact(&request.path).into_owned(),
redacted_authorization: request
.authorization()
.map(|value| redact(value).into_owned()),
header_count: request.headers.len(),
});
match (&request.method, route(&request.path)) {
(Method::Post, Route::Submit) => self.on_submit(),
(Method::Get, Route::Statement(handle)) => self.on_poll(&handle),
(Method::Get, Route::Partition { handle, partition }) => {
self.on_partition(&handle, partition)
}
(Method::Post, Route::Cancel(handle)) => self.on_cancel(&handle),
_ => not_found(),
}
}
fn on_submit(&mut self) -> MockHttpResponse {
if self.immediate {
self.terminal.clone()
} else {
self.running.clone()
}
}
fn on_poll(&mut self, handle: &str) -> MockHttpResponse {
if handle != self.statement_handle {
return not_found();
}
if self.is_cancelled(handle) {
return cancelled_status(handle);
}
let count = self.poll_counts.entry(handle.to_owned()).or_insert(0);
*count += 1;
if *count > self.polls_before_complete {
self.terminal.clone()
} else {
self.running.clone()
}
}
fn on_partition(&self, handle: &str, partition: u32) -> MockHttpResponse {
if handle != self.statement_handle {
return not_found();
}
if self.is_cancelled(handle) {
return cancelled_status(handle);
}
self.partitions
.get(&partition)
.cloned()
.unwrap_or_else(not_found)
}
fn on_cancel(&mut self, handle: &str) -> MockHttpResponse {
if handle != self.statement_handle {
return not_found();
}
self.cancelled.insert(handle.to_owned(), true);
self.cancel.clone()
}
#[must_use]
pub fn poll_count(&self, handle: &str) -> u32 {
self.poll_counts.get(handle).copied().unwrap_or(0)
}
#[must_use]
pub fn is_cancelled(&self, handle: &str) -> bool {
self.cancelled.get(handle).copied().unwrap_or(false)
}
#[must_use]
pub fn requests(&self) -> &[RecordedRequest] {
&self.log
}
#[must_use]
pub fn observed_authorizations(&self) -> Vec<&str> {
self.log
.iter()
.filter_map(|request| request.redacted_authorization.as_deref())
.collect()
}
}
fn not_found() -> MockHttpResponse {
MockHttpResponse::json(
404,
br#"{"code":"390404","message":"Statement handle not found."}"#.to_vec(),
)
}
fn cancelled_status(handle: &str) -> MockHttpResponse {
let body = serde_json::json!({
"code": "000604",
"sqlState": "57014",
"message": "SQL execution canceled",
"statementHandle": handle,
"statementStatusUrl": format!("/api/v2/statements/{handle}"),
})
.to_string()
.into_bytes();
MockHttpResponse::json(422, body)
}
#[cfg(test)]
mod tests {
use super::super::scenarios;
use super::*;
#[test]
fn route_parses_the_three_lifecycle_paths() {
assert_eq!(route("/api/v2/statements"), Route::Submit);
assert_eq!(route("/api/v2/statements?async=true"), Route::Submit);
assert_eq!(
route("/api/v2/statements/abc-123"),
Route::Statement("abc-123".to_owned())
);
assert_eq!(
route("/api/v2/statements/abc-123?partition=1"),
Route::Partition {
handle: "abc-123".to_owned(),
partition: 1,
}
);
assert_eq!(
route("/api/v2/statements/abc-123/cancel"),
Route::Cancel("abc-123".to_owned())
);
assert_eq!(route("/api/v2/other"), Route::Unknown);
}
#[test]
fn async_lifecycle_runs_then_completes_then_cancels() -> Result<(), String> {
let mut mock = scenarios::default_async_lifecycle();
let handle = mock.statement_handle().to_owned();
let submit = mock.respond(&MockHttpRequest::post(
"/api/v2/statements?async=true",
scenarios::SUBMIT_SELECT_REQUEST.to_vec(),
));
assert_eq!(submit.status, 202);
let poll_path = format!("/api/v2/statements/{handle}");
assert_eq!(mock.respond(&MockHttpRequest::get(&poll_path)).status, 202);
assert_eq!(mock.respond(&MockHttpRequest::get(&poll_path)).status, 202);
assert_eq!(mock.respond(&MockHttpRequest::get(&poll_path)).status, 200);
assert_eq!(mock.poll_count(&handle), 3);
let partition_path = format!("/api/v2/statements/{handle}?partition=1");
let partition = mock.respond(&MockHttpRequest::get(&partition_path));
assert_eq!(partition.status, 200);
assert!(partition.has_header("Content-Encoding"));
let cancel_path = format!("/api/v2/statements/{handle}/cancel");
let cancel = mock.respond(&MockHttpRequest::post(&cancel_path, Vec::new()));
assert_eq!(cancel.status, 200);
assert!(mock.is_cancelled(&handle));
assert_eq!(
mock.respond(&MockHttpRequest::get("/api/v2/statements/nope"))
.status,
404
);
Ok(())
}
#[test]
fn authorization_is_recorded_redacted() -> Result<(), String> {
let mut mock = scenarios::default_async_lifecycle();
let request = MockHttpRequest::post("/api/v2/statements", Vec::new())
.with_bearer("eyJhbGciOiJSUzI1NiJ9.payload.signature");
mock.respond(&request);
let observed = mock.observed_authorizations();
assert_eq!(observed.len(), 1);
assert!(observed[0].contains("[REDACTED]"));
assert!(!observed[0].contains("eyJhbGciOiJSUzI1NiJ9"));
Ok(())
}
#[test]
fn request_path_is_recorded_redacted() -> Result<(), String> {
let mut mock = scenarios::default_async_lifecycle();
mock.respond(&MockHttpRequest::get(
"/api/v2/statements/abc-123?token=sfpat_SECRET123",
));
let recorded = mock
.requests()
.first()
.ok_or_else(|| "request should be recorded".to_string())?;
assert!(recorded.path.contains("[REDACTED]"));
assert!(!recorded.path.contains("sfpat_SECRET123"));
Ok(())
}
#[test]
fn cancelled_handle_no_longer_completes_or_serves_partitions() -> Result<(), String> {
let mut mock = scenarios::default_async_lifecycle();
let handle = mock.statement_handle().to_owned();
let cancel_path = format!("/api/v2/statements/{handle}/cancel");
assert_eq!(
mock.respond(&MockHttpRequest::post(cancel_path, Vec::new()))
.status,
200
);
let poll_path = format!("/api/v2/statements/{handle}");
let poll = mock.respond(&MockHttpRequest::get(&poll_path));
assert_eq!(poll.status, 422);
let poll_body = std::str::from_utf8(&poll.body).map_err(|error| error.to_string())?;
assert!(poll_body.contains("\"code\":\"000604\""));
assert!(poll_body.contains("\"sqlState\":\"57014\""));
assert!(poll_body.contains("SQL execution canceled"));
let partition_path = format!("/api/v2/statements/{handle}?partition=1");
let partition = mock.respond(&MockHttpRequest::get(&partition_path));
assert_eq!(partition.status, 422);
Ok(())
}
}