use std::fmt;
use franken_snowflake_core::redact::redact;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::mock::http::{MockHttpRequest, MockHttpResponse, ResponseClass};
use crate::mock::{scenarios, server::MockSqlApi};
#[derive(Debug)]
pub enum ReplayError {
Json(serde_json::Error),
}
impl fmt::Display for ReplayError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Json(error) => write!(f, "replay json error: {error}"),
}
}
}
impl std::error::Error for ReplayError {}
impl From<serde_json::Error> for ReplayError {
fn from(error: serde_json::Error) -> Self {
Self::Json(error)
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ProtocolPacket {
pub name: String,
pub request_method: String,
pub request_path: String,
pub status: u16,
pub response_class: String,
pub headers: Vec<(String, String)>,
pub body_hex: String,
pub wire_hex: String,
}
impl ProtocolPacket {
#[must_use]
pub fn from_exchange(
name: impl Into<String>,
request: &MockHttpRequest,
response: &MockHttpResponse,
) -> Self {
let response = recordable_response(response);
Self {
name: name.into(),
request_method: request.method.as_str().to_owned(),
request_path: redact(&request.path).into_owned(),
status: response.status,
response_class: response_class_label(&response).to_owned(),
headers: response.headers.clone(),
body_hex: hex(&response.body),
wire_hex: hex(&response.to_wire()),
}
}
pub fn body_json(&self) -> Result<Value, ReplayError> {
let bytes = unhex(&self.body_hex);
Ok(serde_json::from_slice(&bytes)?)
}
}
fn recordable_response(response: &MockHttpResponse) -> MockHttpResponse {
let mut response = response.clone();
for (_, value) in &mut response.headers {
*value = redact(value).into_owned();
}
response.body = recordable_body(&response.body);
update_content_length(&mut response);
response
}
fn recordable_body(body: &[u8]) -> Vec<u8> {
let Ok(text) = std::str::from_utf8(body) else {
return body.to_vec();
};
redact(text).as_bytes().to_vec()
}
fn update_content_length(response: &mut MockHttpResponse) {
let len = response.body.len().to_string();
for (name, value) in &mut response.headers {
if name.eq_ignore_ascii_case("Content-Length") {
*value = len;
return;
}
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ReplayStep {
pub name: String,
pub packet: ProtocolPacket,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ReplaySummary {
pub schema_version: u32,
pub scenario: String,
pub statement_handle: String,
pub poll_count: u32,
pub cancelled: bool,
pub steps: Vec<ReplayStep>,
}
#[derive(Clone, Debug)]
pub struct ReplayHarness {
scenario: String,
mock: MockSqlApi,
steps: Vec<ReplayStep>,
}
impl ReplayHarness {
#[must_use]
pub fn new(scenario: impl Into<String>, mock: MockSqlApi) -> Self {
Self {
scenario: scenario.into(),
mock,
steps: Vec::new(),
}
}
#[must_use]
pub fn statement_handle(&self) -> &str {
self.mock.statement_handle()
}
pub fn send(&mut self, name: impl Into<String>, request: MockHttpRequest) -> MockHttpResponse {
let name = name.into();
let response = self.mock.respond(&request);
let packet = ProtocolPacket::from_exchange(name.clone(), &request, &response);
self.steps.push(ReplayStep { name, packet });
response
}
pub fn submit_default_select(&mut self) -> MockHttpResponse {
self.send(
"submit",
MockHttpRequest::post(
"/api/v2/statements?async=true",
scenarios::SUBMIT_SELECT_REQUEST.to_vec(),
),
)
}
pub fn poll(&mut self, name: impl Into<String>) -> MockHttpResponse {
let path = format!("/api/v2/statements/{}", self.statement_handle());
self.send(name, MockHttpRequest::get(path))
}
pub fn fetch_partition(&mut self, partition: u32) -> MockHttpResponse {
let path = format!(
"/api/v2/statements/{}?partition={partition}",
self.statement_handle()
);
self.send(format!("partition-{partition}"), MockHttpRequest::get(path))
}
pub fn cancel(&mut self) -> MockHttpResponse {
let path = format!("/api/v2/statements/{}/cancel", self.statement_handle());
self.send("cancel", MockHttpRequest::post(path, Vec::new()))
}
#[must_use]
pub fn finish(self) -> ReplaySummary {
let handle = self.mock.statement_handle().to_owned();
ReplaySummary {
schema_version: 1,
scenario: self.scenario,
poll_count: self.mock.poll_count(&handle),
cancelled: self.mock.is_cancelled(&handle),
statement_handle: handle,
steps: self.steps,
}
}
}
#[must_use]
pub fn default_protocol_replay() -> ReplaySummary {
let mut replay = ReplayHarness::new(
"default-async-select-with-partition-and-cancel",
scenarios::default_async_lifecycle(),
);
replay.submit_default_select();
replay.poll("poll-1-running");
replay.poll("poll-2-running");
replay.poll("poll-3-complete");
replay.fetch_partition(1);
replay.cancel();
replay.finish()
}
fn response_class_label(response: &MockHttpResponse) -> &'static str {
match response.class() {
ResponseClass::Completed => "completed",
ResponseClass::Running => "running",
ResponseClass::StatementTimeout => "statement_timeout",
ResponseClass::StatementFailed => "statement_failed",
ResponseClass::RateLimited => "rate_limited",
ResponseClass::Other(_) => "other",
}
}
fn hex(bytes: &[u8]) -> String {
const TABLE: &[u8; 16] = b"0123456789abcdef";
let mut out = String::with_capacity(bytes.len() * 2);
for byte in bytes {
out.push(TABLE[(byte >> 4) as usize] as char);
out.push(TABLE[(byte & 0x0f) as usize] as char);
}
out
}
fn unhex(hex_text: &str) -> Vec<u8> {
let mut out = Vec::with_capacity(hex_text.len() / 2);
let (pairs, _remainder) = hex_text.as_bytes().as_chunks::<2>();
for pair in pairs {
if let (Some(high), Some(low)) = (hex_value(pair[0]), hex_value(pair[1])) {
out.push((high << 4) | low);
}
}
out
}
fn hex_value(byte: u8) -> Option<u8> {
match byte {
b'0'..=b'9' => Some(byte - b'0'),
b'a'..=b'f' => Some(byte - b'a' + 10),
b'A'..=b'F' => Some(byte - b'A' + 10),
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::harness::golden::{
GoldenConfig, assert_no_cr, check_golden_file, to_canonical_json,
};
#[test]
fn default_replay_captures_lifecycle_partition_and_cancel() {
let replay = default_protocol_replay();
assert_eq!(replay.steps.len(), 6);
assert_eq!(replay.poll_count, 3);
assert!(replay.cancelled);
assert_eq!(replay.steps[0].packet.status, 202);
assert_eq!(replay.steps[3].packet.status, 200);
assert_eq!(replay.steps[4].packet.response_class, "completed");
assert!(
replay.steps[4]
.packet
.headers
.iter()
.any(|(name, value)| name == "Content-Encoding" && value == "gzip")
);
}
#[test]
fn default_protocol_packets_match_committed_golden() -> Result<(), Box<dyn std::error::Error>> {
let replay = default_protocol_replay();
let value = serde_json::to_value(replay)?;
let cfg = GoldenConfig::strict();
let canonical = to_canonical_json(&value, &cfg);
assert_no_cr(&canonical)?;
let path = std::path::Path::new(concat!(
env!("CARGO_MANIFEST_DIR"),
"/fixtures/golden/default_protocol_replay.golden.json"
));
check_golden_file(path, &value, &cfg)?;
Ok(())
}
#[test]
fn replay_packets_redact_recorded_paths_headers_and_text_bodies() -> Result<(), String> {
let request =
MockHttpRequest::get("/api/v2/statements?requestId=sfpat_replay_path_secret_123");
let response =
MockHttpResponse::json(200, br#"{"token":"sfpat_replay_body_secret_123"}"#.to_vec())
.with_header("X-Debug-Token", "ghp_replay_header_secret_123")
.with_header("Content-Length", "40");
let packet = ProtocolPacket::from_exchange("secret-probe", &request, &response);
let body = String::from_utf8(unhex(&packet.body_hex))
.map_err(|e| format!("body utf8 failed: {e}"))?;
let wire = String::from_utf8(unhex(&packet.wire_hex))
.map_err(|e| format!("wire utf8 failed: {e}"))?;
assert!(!packet.request_path.contains("sfpat_replay_path_secret_123"));
assert!(packet.request_path.contains("[REDACTED]"));
assert!(!body.contains("sfpat_replay_body_secret_123"));
assert!(body.contains("[REDACTED]"));
assert!(!wire.contains("sfpat_replay_body_secret_123"));
assert!(!wire.contains("ghp_replay_header_secret_123"));
assert!(wire.contains("[REDACTED]"));
assert!(
packet
.headers
.iter()
.all(|(_, value)| !value.contains("ghp_replay_header_secret_123"))
);
Ok(())
}
}