use serde::Deserialize;
use snafu::Snafu;
#[derive(Debug, Snafu)]
pub enum ActiveQueryError {
#[snafu(display(
"No active query '{query_id}' found. It may have already finished, it was submitted under a different principal, or it is running on another runtime instance."
))]
NotFound { query_id: String },
#[snafu(display(
"Query id '{query_id}' is not a valid UUID. Use the query_id from active_queries()."
))]
InvalidQueryId { query_id: String },
#[snafu(display(
"The configured API key does not allow cancelling queries. Use a key with write access."
))]
WriteAccessRequired,
#[snafu(display("Request failed (HTTP {status_code}): {response_body}"))]
RequestFailed {
status_code: u16,
response_body: String,
},
#[snafu(display("Request failed: {message}"))]
HttpError { message: String },
#[snafu(display("Failed to parse response: {message}"))]
ParseError { message: String },
}
#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
pub struct ActiveQuery {
pub query_id: String,
pub protocol: String,
pub sql_preview: String,
pub started_at_ms: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
pub struct ActiveQueryList {
pub queries: Vec<ActiveQuery>,
pub total_count: usize,
}
pub(crate) fn is_uuid(query_id: &str) -> bool {
let bytes = query_id.as_bytes();
if bytes.len() != 36 {
return false;
}
bytes.iter().enumerate().all(|(index, byte)| {
if matches!(index, 8 | 13 | 18 | 23) {
*byte == b'-'
} else {
byte.is_ascii_hexdigit()
}
})
}
#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
pub struct CancelActiveQueryResponse {
pub query_id: String,
pub status: String,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deserializes_an_active_query_list() {
let body = r#"{
"queries": [
{
"query_id": "0198f0a1-9c3d-7c4e-8a11-2b3c4d5e6f70",
"protocol": "flight",
"sql_preview": "SELECT * FROM taxi_trips",
"started_at_ms": 1750000000000
}
],
"total_count": 1
}"#;
let list: ActiveQueryList = serde_json::from_str(body).expect("deserialize list");
assert_eq!(list.total_count, 1);
assert_eq!(list.queries.len(), 1);
assert_eq!(list.queries[0].protocol, "flight");
assert_eq!(list.queries[0].sql_preview, "SELECT * FROM taxi_trips");
assert_eq!(list.queries[0].started_at_ms, 1_750_000_000_000);
}
#[test]
fn deserializes_an_empty_active_query_list() {
let list: ActiveQueryList =
serde_json::from_str(r#"{"queries":[],"total_count":0}"#).expect("deserialize list");
assert_eq!(list.total_count, 0);
assert!(list.queries.is_empty());
}
#[test]
fn deserializes_a_cancel_response() {
let response: CancelActiveQueryResponse = serde_json::from_str(
r#"{"query_id":"0198f0a1-9c3d-7c4e-8a11-2b3c4d5e6f70","status":"cancelled"}"#,
)
.expect("deserialize cancel response");
assert_eq!(response.status, "cancelled");
}
#[test]
fn not_found_error_explains_both_causes() {
let message = ActiveQueryError::NotFound {
query_id: "abc".to_string(),
}
.to_string();
assert!(message.contains("abc"));
assert!(message.contains("already finished"));
}
#[test]
fn is_uuid_accepts_the_ids_the_runtime_hands_out() {
assert!(is_uuid("0198f0a1-9c3d-7c4e-8a11-2b3c4d5e6f70"));
assert!(is_uuid("0198F0A1-9C3D-7C4E-8A11-2B3C4D5E6F70"));
}
#[test]
fn is_uuid_rejects_anything_that_could_reroute_a_request() {
for id in [
"",
".",
"..",
"../queries/escape",
"not-a-uuid",
"0198f0a19c3d-7c4e-8a11-2b3c4d5e6f70-",
"0198f0a1-9c3d-7c4e-8a11-2b3c4d5e6f7g",
"0198f0a1-9c3d-7c4e-8a11-2b3c4d5e6f70/cancel",
] {
assert!(!is_uuid(id), "{id:?} should not be accepted as a UUID");
}
}
#[test]
fn invalid_query_id_error_points_at_active_queries() {
let message = ActiveQueryError::InvalidQueryId {
query_id: "not-a-uuid".to_string(),
}
.to_string();
assert!(message.contains("not-a-uuid"));
assert!(message.contains("active_queries()"));
}
}