use std::collections::BTreeMap;
use std::fs;
use std::path::PathBuf;
use asupersync::net::TcpStream;
use asupersync::Cx;
use sha2::{Digest, Sha256};
use oracledb_protocol::net::cassette::{self, Direction};
use oracledb_protocol::thin::{
build_connect_packet_payload, parse_accept_payload, TNS_PACKET_TYPE_ACCEPT,
TNS_PACKET_TYPE_CONNECT, TNS_PACKET_TYPE_RESEND,
};
use oracledb_protocol::wire::{encode_packet, PacketLengthWidth, ProtocolLimits};
use oracledb_protocol::ProtocolError;
use crate::transport::{self, ReplayWriteMode};
use crate::{
build_io_runtime, ConnectionCore, DriverTransport, Error, IncomingPacket, Result,
MAX_CONNECT_RESEND_ROUNDS,
};
const ADVERTISED_SDU: u16 = 8192;
const ORACLE_11G_PROTOCOL_VERSION: u16 = 314;
const MANIFEST_SCHEMA_VERSION: &str = "1";
const CASSETTE_FORMAT_VERSION: &str = "1";
const SOURCE_COMMIT: &str = include_str!("../../../docs/baseline/source_commit.txt");
struct Lane {
id: &'static str,
default_connect: &'static str,
outcome: Outcome,
}
enum Outcome {
Refusal { version: u16 },
Accept {
supports_fast_auth: bool,
supports_end_of_response: bool,
},
}
enum ExpectedOutcome {
Refusal {
version: u16,
},
Accept {
supports_fast_auth: bool,
supports_end_of_response: bool,
},
}
fn lanes() -> Vec<Lane> {
vec![
Lane {
id: "xe11",
default_connect: "localhost:1511/XE",
outcome: Outcome::Refusal {
version: ORACLE_11G_PROTOCOL_VERSION,
},
},
Lane {
id: "xe18",
default_connect: "localhost:1518/XEPDB1",
outcome: Outcome::Accept {
supports_fast_auth: false,
supports_end_of_response: false,
},
},
Lane {
id: "xe21",
default_connect: "localhost:1520/XEPDB1",
outcome: Outcome::Accept {
supports_fast_auth: false,
supports_end_of_response: false,
},
},
Lane {
id: "free23",
default_connect: "localhost:1522/FREEPDB1",
outcome: Outcome::Accept {
supports_fast_auth: true,
supports_end_of_response: true,
},
},
]
}
fn fixtures_dir() -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fixtures")
.join("cassettes")
}
fn cassette_path(lane_id: &str) -> PathBuf {
fixtures_dir().join(format!("{lane_id}-connect.tns-cassette"))
}
fn manifest_path(lane_id: &str) -> PathBuf {
fixtures_dir().join(format!("{lane_id}-connect.tns-cassette.manifest"))
}
fn split_connect(connect: &str) -> Result<(String, u16, String)> {
let (addr, service) = connect
.rsplit_once('/')
.ok_or_else(|| Error::Runtime(format!("connect string {connect:?} has no /service")))?;
let (host, port) = addr
.rsplit_once(':')
.ok_or_else(|| Error::Runtime(format!("address {addr:?} has no :port")))?;
let port: u16 = port
.parse()
.map_err(|_| Error::Runtime(format!("bad port in {addr:?}")))?;
Ok((host.to_string(), port, service.to_string()))
}
fn capture_connect_descriptor(service: &str) -> String {
format!(
"(DESCRIPTION=(ADDRESS=(PROTOCOL=tcp)(HOST=cassette-capture)(PORT=0))\
(CONNECT_DATA=(SERVICE_NAME={service})(CID=(PROGRAM=rust-oracledb-cassette)\
(HOST=cassette-capture)(USER=cassette))))"
)
}
async fn drive_connect_handshake(
core: &mut ConnectionCore<DriverTransport>,
cx: &Cx,
connect_data: &str,
) -> Result<IncomingPacket> {
let mut resend_rounds = 0u8;
loop {
let payload = build_connect_packet_payload(connect_data, ADVERTISED_SDU)?;
let packet = encode_packet(
TNS_PACKET_TYPE_CONNECT,
0,
None,
&payload,
PacketLengthWidth::Legacy16,
)?;
core.write_all(cx, &packet).await?;
let reply = core.read_packet(PacketLengthWidth::Legacy16).await?;
match reply.packet_type {
TNS_PACKET_TYPE_ACCEPT => return Ok(reply),
TNS_PACKET_TYPE_RESEND => {
resend_rounds += 1;
if resend_rounds > MAX_CONNECT_RESEND_ROUNDS {
return Err(Error::ConnectResendLoop(resend_rounds));
}
continue;
}
other => return Err(Error::UnexpectedPacket(other)),
}
}
}
const SECRET_FIELD_NAMES: &[&str] = &[
"AUTH_PASSWORD",
"AUTH_SESSKEY",
"AUTH_VFR_DATA",
"AUTH_PBKDF2_CSK_SALT",
"AUTH_PBKDF2_SPEEDY_KEY",
"AUTH_TOKEN",
"SESSION_TOKEN",
"SESSION_KEY",
"ACCESS_TOKEN",
"REFRESH_TOKEN",
"PRIVATE_KEY",
];
fn scan_for_secret_fields(bytes: &[u8]) -> Vec<&'static str> {
let haystack = String::from_utf8_lossy(bytes).to_ascii_uppercase();
SECRET_FIELD_NAMES
.iter()
.copied()
.filter(|field| haystack.contains(field))
.collect()
}
fn sha256_hex(bytes: &[u8]) -> String {
let digest = Sha256::digest(bytes);
let mut out = String::with_capacity(digest.len() * 2);
for byte in digest {
use std::fmt::Write as _;
write!(&mut out, "{byte:02x}").expect("writing to String cannot fail");
}
out
}
fn write_frame_hashes(cassette_bytes: &[u8]) -> Result<Vec<String>> {
let frames = cassette::decode_all(cassette_bytes)
.map_err(|err| Error::Runtime(format!("cassette decode: {err}")))?;
Ok(frames
.iter()
.filter(|frame| frame.direction == Direction::ClientToServer)
.map(|frame| sha256_hex(&frame.bytes))
.collect())
}
fn build_manifest(
lane: &Lane,
service: &str,
cassette_bytes: &[u8],
accept_payload: &[u8],
) -> Result<String> {
let (outcome, version, fast_auth, eor) = describe_accept(accept_payload);
let write_hashes = write_frame_hashes(cassette_bytes)?;
Ok(format!(
concat!(
"schema_version = {}\n",
"format_version = {}\n",
"commit = \"{}\"\n",
"profile = \"connect-negotiation\"\n",
"lane = \"{}\"\n",
"service = \"{}\"\n",
"scenario = \"connect_accept\"\n",
"outcome = \"{}\"\n",
"protocol_version = {}\n",
"supports_fast_auth = {}\n",
"supports_end_of_response = {}\n",
"sanitized = true\n",
"checksum_sha256 = \"{}\"\n",
"expected_writes = {}\n",
"expected_write_sha256 = \"{}\"\n",
),
MANIFEST_SCHEMA_VERSION,
CASSETTE_FORMAT_VERSION,
SOURCE_COMMIT.trim(),
lane.id,
service,
outcome,
version,
fast_auth,
eor,
sha256_hex(cassette_bytes),
write_hashes.len(),
write_hashes.join(","),
))
}
fn describe_accept(payload: &[u8]) -> (&'static str, u16, bool, bool) {
match parse_accept_payload(payload) {
Ok(info) => (
"accept",
info.protocol_version,
info.supports_fast_auth,
info.supports_end_of_response,
),
Err(ProtocolError::UnsupportedVersion { version, .. }) => {
("refusal", version, false, false)
}
Err(_) => ("unknown", 0, false, false),
}
}
fn parse_manifest(text: &str) -> BTreeMap<String, String> {
text.lines()
.filter_map(|line| {
let line = line.trim();
if line.is_empty() || line.starts_with('#') {
return None;
}
let (key, value) = line.split_once('=')?;
Some((
key.trim().to_string(),
value.trim().trim_matches('"').to_string(),
))
})
.collect()
}
fn record_lane(lane: &Lane) -> Result<()> {
let connect = std::env::var(format!("ORACLEDB_CASSETTE_{}", lane.id.to_uppercase()))
.unwrap_or_else(|_| lane.default_connect.to_string());
let (host, port, service) = split_connect(&connect)?;
let connect_data = capture_connect_descriptor(&service);
let runtime = build_io_runtime()?;
let cassette_bytes = runtime.block_on(async move {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("missing ambient Cx in capture runtime".into()))?;
let stream =
TcpStream::connect_timeout((host.clone(), port), std::time::Duration::from_secs(15))
.await
.map_err(|err| Error::Runtime(format!("dial {host}:{port}: {err}")))?;
stream.set_nodelay(true).ok();
let scope = transport::capture_scope();
let (read, write) = transport::plain_split(stream);
let mut core = ConnectionCore::<DriverTransport>::from_halves(read, write, "cassette");
core.set_protocol_limits(ProtocolLimits::DEFAULT)?;
let accept = drive_connect_handshake(&mut core, &cx, &connect_data).await?;
Ok::<_, Error>((scope.to_cassette_bytes(), accept.payload))
})?;
let (cassette_bytes, accept_payload) = cassette_bytes;
let leaks = scan_for_secret_fields(&cassette_bytes);
if !leaks.is_empty() {
return Err(Error::Runtime(format!(
"REFUSING to write {}: secret field(s) present: {leaks:?}",
lane.id
)));
}
let manifest = build_manifest(lane, &service, &cassette_bytes, &accept_payload)?;
let out_dir = std::env::var("ORACLEDB_CASSETTE_RECORD")
.map(PathBuf::from)
.unwrap_or_else(|_| fixtures_dir());
fs::create_dir_all(&out_dir).map_err(|e| Error::Runtime(e.to_string()))?;
let cass = out_dir.join(format!("{}-connect.tns-cassette", lane.id));
let man = out_dir.join(format!("{}-connect.tns-cassette.manifest", lane.id));
fs::write(&cass, &cassette_bytes).map_err(|e| Error::Runtime(e.to_string()))?;
fs::write(&man, manifest).map_err(|e| Error::Runtime(e.to_string()))?;
eprintln!(
"recorded {} ({} bytes) -> {}",
lane.id,
cassette_bytes.len(),
cass.display()
);
Ok(())
}
#[test]
#[ignore = "records the live version-connect cassettes; needs the Docker lanes"]
fn record_version_connect_cassettes() {
let mut failures = Vec::new();
for lane in lanes() {
match record_lane(&lane) {
Ok(()) => {}
Err(err) => failures.push(format!("{}: {err}", lane.id)),
}
}
assert!(failures.is_empty(), "capture failures: {failures:?}");
}
fn replay_lane(lane: &Lane) -> Result<()> {
let cassette_bytes =
fs::read(cassette_path(lane.id)).map_err(|e| Error::Runtime(e.to_string()))?;
let manifest_text =
fs::read_to_string(manifest_path(lane.id)).map_err(|e| Error::Runtime(e.to_string()))?;
let manifest = parse_manifest(&manifest_text);
let expected_checksum = manifest
.get("checksum_sha256")
.ok_or_else(|| Error::Runtime("manifest missing checksum_sha256".into()))?;
if &sha256_hex(&cassette_bytes) != expected_checksum {
return Err(Error::Runtime(format!(
"{}: cassette checksum != manifest",
lane.id
)));
}
let leaks = scan_for_secret_fields(&cassette_bytes);
if !leaks.is_empty() {
return Err(Error::Runtime(format!(
"{}: secret leak {leaks:?}",
lane.id
)));
}
let service = manifest
.get("service")
.cloned()
.ok_or_else(|| Error::Runtime("manifest missing service".into()))?;
let connect_data = capture_connect_descriptor(&service);
let (read, write, audit) =
transport::replay_split_with_audit(&cassette_bytes, ReplayWriteMode::Check)
.map_err(|err| Error::Runtime(format!("invalid replay cassette: {err}")))?;
let mut core = ConnectionCore::<DriverTransport>::from_halves(read, write, "replay");
core.set_protocol_limits(ProtocolLimits::DEFAULT)?;
let lane_id = lane.id.to_string();
let outcome = match &lane.outcome {
Outcome::Refusal { version } => ExpectedOutcome::Refusal { version: *version },
Outcome::Accept {
supports_fast_auth,
supports_end_of_response,
} => ExpectedOutcome::Accept {
supports_fast_auth: *supports_fast_auth,
supports_end_of_response: *supports_end_of_response,
},
};
let expected_version = manifest.get("protocol_version").cloned();
let runtime = build_io_runtime()?;
runtime.block_on(async move {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("missing ambient Cx in replay runtime".into()))?;
let accept = drive_connect_handshake(&mut core, &cx, &connect_data).await?;
let payload = accept.payload.as_slice();
match outcome {
ExpectedOutcome::Refusal { version } => match parse_accept_payload(payload) {
Err(ProtocolError::UnsupportedVersion {
version: got,
minimum,
}) => {
assert_eq!(got, version, "{lane_id}: refused version");
assert_eq!(
minimum,
oracledb_protocol::TNS_VERSION_MIN_ACCEPTED,
"{lane_id}: floor"
);
}
other => {
return Err(Error::Runtime(format!(
"{lane_id}: expected UnsupportedVersion, got {other:?}"
)));
}
},
ExpectedOutcome::Accept {
supports_fast_auth,
supports_end_of_response,
} => {
let info = parse_accept_payload(payload)
.map_err(|e| Error::Runtime(format!("{lane_id}: parse ACCEPT: {e}")))?;
assert_eq!(
info.supports_fast_auth, supports_fast_auth,
"{lane_id}: fast_auth"
);
assert_eq!(
info.supports_end_of_response, supports_end_of_response,
"{lane_id}: end_of_response"
);
if let Some(expected) = expected_version {
assert_eq!(
info.protocol_version.to_string(),
expected,
"{lane_id}: protocol_version"
);
}
}
}
Ok::<_, Error>(())
})?;
audit
.assert_finished()
.map_err(|err| Error::Runtime(format!("{}: {err}", lane.id)))?;
Ok(())
}
#[test]
fn replay_version_connect_cassettes_offline() {
let mut failures = Vec::new();
for lane in lanes() {
if !cassette_path(lane.id).exists() {
eprintln!("skip {}: no committed cassette", lane.id);
continue;
}
if let Err(err) = replay_lane(&lane) {
failures.push(err.to_string());
}
}
assert!(failures.is_empty(), "replay failures: {failures:?}");
}
#[cfg(test)]
#[allow(clippy::too_many_arguments)]
fn loopback_for_replay(
core: ConnectionCore<DriverTransport>,
capabilities: oracledb_protocol::thin::ClientCapabilities,
ttc_seq_num: u8,
supports_end_of_response: bool,
supports_oob: bool,
sdu: usize,
) -> crate::Connection {
use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet, VecDeque};
crate::Connection {
descriptor: crate::EasyConnect::parse("127.0.0.1:1521/FREEPDB1")
.expect("loopback descriptor parses"),
identity: oracledb_protocol::ClientIdentity::new(
"cassette-replay",
"cassette-replay",
"cassette",
"unknown",
"rust-oracledb-cassette",
)
.expect("loopback identity is valid"),
core,
protocol_limits: ProtocolLimits::DEFAULT,
session_id: 0,
serial_num: 0,
server_version: None,
server_version_tuple: None,
capabilities,
ttc_seq_num,
sdu,
supports_end_of_response,
supports_oob,
cursor_columns: BTreeMap::new(),
fetch_metadata_by_sql: HashMap::new(),
fetch_metadata_order: VecDeque::new(),
dead: false,
user: "cassette-replay".into(),
combo_key: Vec::new(),
statement_cache: Vec::new(),
statement_cache_size: crate::STATEMENT_CACHE_SIZE,
in_use_cursors: HashSet::new(),
lob_prefetch_cursors: BTreeSet::new(),
copied_cursors: HashSet::new(),
cursors_to_close: Vec::new(),
sessionless_data: None,
notification_buffer: Vec::new(),
notification_header_consumed: false,
transaction_context: None,
txn_in_progress: false,
shape_cache: std::sync::Arc::new(crate::StatementShapeCache::new()),
}
}
fn reencode_frames(frames: &[cassette::Frame]) -> Vec<u8> {
let mut out = Vec::new();
cassette::write_header(&mut out);
for frame in frames {
cassette::write_frame(&mut out, frame.direction, 0, &frame.bytes);
}
out
}
fn slice_scenario_frames(post: &[cassette::Frame], client_writes: usize) -> &[cassette::Frame] {
let mut seen = 0usize;
for (idx, frame) in post.iter().enumerate() {
if frame.direction == cassette::Direction::ClientToServer {
seen += 1;
if seen > client_writes {
return &post[..idx];
}
}
}
post
}
const POSTAUTH_SQL: &str = "select cast(7 + 5 as number(6)) as v from dual";
const POSTAUTH_EXPECTED_VALUE: &str = "12";
const LOB_SCENARIO_DESC: &str = "lob: create_temp_lob + write_lob + read_lob (blob)";
const LOB_EXPECTED_VALUE: &str = "rust-oracledb cassette blob payload";
const LOB_READ_AMOUNT: u64 = 4000;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum PostAuthScenario {
ExecuteSelect,
LobBlob,
}
impl PostAuthScenario {
fn suffix(self) -> &'static str {
match self {
PostAuthScenario::ExecuteSelect => "postauth",
PostAuthScenario::LobBlob => "lob",
}
}
fn sql(self) -> &'static str {
match self {
PostAuthScenario::ExecuteSelect => POSTAUTH_SQL,
PostAuthScenario::LobBlob => LOB_SCENARIO_DESC,
}
}
fn expected_value(self) -> &'static str {
match self {
PostAuthScenario::ExecuteSelect => POSTAUTH_EXPECTED_VALUE,
PostAuthScenario::LobBlob => LOB_EXPECTED_VALUE,
}
}
fn tag(self) -> &'static str {
match self {
PostAuthScenario::ExecuteSelect => "execute_select",
PostAuthScenario::LobBlob => "temp_lob_create_write_read",
}
}
fn client_writes(self) -> usize {
match self {
PostAuthScenario::ExecuteSelect => 1,
PostAuthScenario::LobBlob => 3,
}
}
}
#[cfg(test)]
async fn drive_scenario(
conn: &mut crate::Connection,
cx: &Cx,
scenario: PostAuthScenario,
) -> Result<Option<String>> {
match scenario {
PostAuthScenario::ExecuteSelect => {
use oracledb_protocol::thin::{ExecuteOptions, QueryValue};
let exec = conn
.execute_raw(cx, scenario.sql(), 2, &[], ExecuteOptions::default(), None)
.await?;
Ok(exec
.cell(0, 0)
.and_then(QueryValue::as_number_text)
.map(|c| c.to_string()))
}
PostAuthScenario::LobBlob => {
use oracledb_protocol::thin::{CS_FORM_IMPLICIT, ORA_TYPE_NUM_BLOB};
let temp = conn
.create_temp_lob(cx, ORA_TYPE_NUM_BLOB, CS_FORM_IMPLICIT)
.await?;
let locator = temp.locator;
conn.write_lob(cx, &locator, 1, LOB_EXPECTED_VALUE.as_bytes())
.await?;
let read = conn.read_lob(cx, &locator, 1, LOB_READ_AMOUNT).await?;
let bytes = read.data.unwrap_or_default();
let text = String::from_utf8(bytes)
.map_err(|e| Error::Runtime(format!("BLOB read not UTF-8: {e}")))?;
Ok(Some(text))
}
}
}
struct PostAuthLane {
id: &'static str,
default_connect: &'static str,
default_user: &'static str,
default_password: &'static str,
}
fn postauth_lanes() -> Vec<PostAuthLane> {
vec![
PostAuthLane {
id: "xe18",
default_connect: "localhost:1518/XEPDB1",
default_user: "testuser",
default_password: "testpw",
},
PostAuthLane {
id: "xe21",
default_connect: "localhost:1520/XEPDB1",
default_user: "testuser",
default_password: "testpw",
},
PostAuthLane {
id: "free23",
default_connect: "localhost:1522/FREEPDB1",
default_user: "pythontest",
default_password: "pythontest",
},
]
}
fn postauth_cassette_path(lane_id: &str, scenario: PostAuthScenario) -> PathBuf {
fixtures_dir().join(format!("{lane_id}-{}.tns-cassette", scenario.suffix()))
}
fn postauth_manifest_path(lane_id: &str, scenario: PostAuthScenario) -> PathBuf {
fixtures_dir().join(format!(
"{lane_id}-{}.tns-cassette.manifest",
scenario.suffix()
))
}
struct PostAuthCapture {
service: String,
sliced: Vec<u8>,
capabilities: oracledb_protocol::thin::ClientCapabilities,
ttc_seq_num: u8,
supports_end_of_response: bool,
supports_oob: bool,
sdu: usize,
value: Option<String>,
}
fn capture_postauth(
connect: &str,
user: &str,
password: &str,
scenario: PostAuthScenario,
) -> Result<PostAuthCapture> {
let (_, _, service) = split_connect(connect)?;
let runtime = build_io_runtime()?;
let (full_bytes, prefix_frames, seq0, caps, eor, oob, sdu, value) =
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("missing ambient Cx in capture runtime".into()))?;
let scope = transport::capture_scope();
let identity = oracledb_protocol::ClientIdentity::new(
"cassette-postauth",
"cassette-capture",
"cassette",
"unknown",
"rust-oracledb-cassette",
)
.map_err(|e| Error::Runtime(e.to_string()))?;
let options = crate::ConnectOptions::new(
connect.to_string(),
user.to_string(),
password.to_string(),
identity,
);
let mut conn = crate::Connection::connect(&cx, options).await?;
let prefix_frames = scope.recorder().frame_count();
let seq0 = conn.ttc_seq_num;
let caps = conn.capabilities;
let eor = conn.supports_end_of_response;
let oob = conn.supports_oob;
let sdu = conn.sdu;
let value = drive_scenario(&mut conn, &cx, scenario).await?;
let full = scope.to_cassette_bytes();
conn.close(&cx).await.ok();
Ok::<_, Error>((full, prefix_frames, seq0, caps, eor, oob, sdu, value))
})?;
let all_frames = cassette::decode_all(&full_bytes)
.map_err(|e| Error::Runtime(format!("decode capture: {e}")))?;
if prefix_frames >= all_frames.len() {
return Err(Error::Runtime(format!(
"no post-auth frames after the connect+auth prefix ({prefix_frames} of {})",
all_frames.len()
)));
}
let post = &all_frames[prefix_frames..];
let sliced = reencode_frames(slice_scenario_frames(post, scenario.client_writes()));
let leaks = scan_for_secret_fields(&sliced);
if !leaks.is_empty() {
return Err(Error::Runtime(format!(
"post-auth slice leaked secrets: {leaks:?}"
)));
}
Ok(PostAuthCapture {
service,
sliced,
capabilities: caps,
ttc_seq_num: seq0,
supports_end_of_response: eor,
supports_oob: oob,
sdu,
value,
})
}
#[allow(clippy::too_many_arguments)]
fn replay_postauth(
sliced: &[u8],
capabilities: oracledb_protocol::thin::ClientCapabilities,
ttc_seq_num: u8,
supports_end_of_response: bool,
supports_oob: bool,
sdu: usize,
scenario: PostAuthScenario,
) -> Result<Option<String>> {
let (read, write, audit) = transport::replay_split_with_audit(sliced, ReplayWriteMode::Check)
.map_err(|e| Error::Runtime(format!("replay split: {e}")))?;
let core = ConnectionCore::<DriverTransport>::from_halves(read, write, "postauth_replay");
let mut conn = loopback_for_replay(
core,
capabilities,
ttc_seq_num,
supports_end_of_response,
supports_oob,
sdu,
);
let runtime = build_io_runtime()?;
let value = runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("missing ambient Cx in replay runtime".into()))?;
drive_scenario(&mut conn, &cx, scenario).await
})?;
audit
.assert_finished()
.map_err(|e| Error::Runtime(format!("post-auth replay audit: {e}")))?;
Ok(value)
}
fn build_postauth_manifest(
lane_id: &str,
cap: &PostAuthCapture,
scenario: PostAuthScenario,
) -> Result<String> {
let write_hashes = write_frame_hashes(&cap.sliced)?;
Ok(format!(
concat!(
"schema_version = {}\n",
"format_version = {}\n",
"commit = \"{}\"\n",
"profile = \"post-auth-query\"\n",
"lane = \"{}\"\n",
"service = \"{}\"\n",
"scenario = \"{}\"\n",
"sql = \"{}\"\n",
"ttc_field_version = {}\n",
"charset_id = {}\n",
"max_string_size = {}\n",
"ttc_seq_num = {}\n",
"supports_end_of_response = {}\n",
"supports_oob = {}\n",
"sdu = {}\n",
"expected_value = \"{}\"\n",
"sanitized = true\n",
"checksum_sha256 = \"{}\"\n",
"expected_writes = {}\n",
"expected_write_sha256 = \"{}\"\n",
),
MANIFEST_SCHEMA_VERSION,
CASSETTE_FORMAT_VERSION,
SOURCE_COMMIT.trim(),
lane_id,
cap.service,
scenario.tag(),
scenario.sql(),
cap.capabilities.ttc_field_version,
cap.capabilities.charset_id,
cap.capabilities.max_string_size,
cap.ttc_seq_num,
cap.supports_end_of_response,
cap.supports_oob,
cap.sdu,
cap.value.as_deref().unwrap_or(""),
sha256_hex(&cap.sliced),
write_hashes.len(),
write_hashes.join(","),
))
}
#[test]
#[ignore = "records the live post-auth query cassettes; needs the Docker lanes"]
fn record_postauth_query_cassettes() {
record_postauth_scenario(PostAuthScenario::ExecuteSelect);
}
#[test]
#[ignore = "records the live LOB post-auth cassettes; needs the Docker lanes"]
fn record_postauth_lob_cassettes() {
record_postauth_scenario(PostAuthScenario::LobBlob);
}
fn record_postauth_scenario(scenario: PostAuthScenario) {
let out_dir = std::env::var("ORACLEDB_CASSETTE_RECORD")
.map(PathBuf::from)
.unwrap_or_else(|_| fixtures_dir());
fs::create_dir_all(&out_dir).expect("create fixtures dir");
let mut failures = Vec::new();
for lane in postauth_lanes() {
let up = lane.id.to_uppercase();
let connect = std::env::var(format!("ORACLEDB_CASSETTE_{up}"))
.unwrap_or_else(|_| lane.default_connect.to_string());
let user = std::env::var(format!("ORACLEDB_CASSETTE_{up}_USER"))
.unwrap_or_else(|_| lane.default_user.to_string());
let password = std::env::var(format!("ORACLEDB_CASSETTE_{up}_PASSWORD"))
.unwrap_or_else(|_| lane.default_password.to_string());
let cap = match capture_postauth(&connect, &user, &password, scenario) {
Ok(cap) => cap,
Err(err) => {
failures.push(format!("{}: {err}", lane.id));
continue;
}
};
match replay_postauth(
&cap.sliced,
cap.capabilities,
cap.ttc_seq_num,
cap.supports_end_of_response,
cap.supports_oob,
cap.sdu,
scenario,
) {
Ok(v) if v.as_deref() == Some(scenario.expected_value()) => {}
Ok(v) => {
failures.push(format!(
"{}: replay value {v:?} != {:?}",
lane.id,
scenario.expected_value()
));
continue;
}
Err(err) => {
failures.push(format!("{}: pre-commit replay {err}", lane.id));
continue;
}
}
let manifest = match build_postauth_manifest(lane.id, &cap, scenario) {
Ok(m) => m,
Err(err) => {
failures.push(format!("{}: manifest {err}", lane.id));
continue;
}
};
let cass = out_dir.join(format!("{}-{}.tns-cassette", lane.id, scenario.suffix()));
let man = out_dir.join(format!(
"{}-{}.tns-cassette.manifest",
lane.id,
scenario.suffix()
));
if let Err(err) = fs::write(&cass, &cap.sliced) {
failures.push(format!("{}: write cassette {err}", lane.id));
continue;
}
if let Err(err) = fs::write(&man, manifest) {
failures.push(format!("{}: write manifest {err}", lane.id));
continue;
}
eprintln!(
"recorded {} {} ({} bytes) -> {}",
lane.id,
scenario.suffix(),
cap.sliced.len(),
cass.display()
);
}
assert!(
failures.is_empty(),
"{} capture failures: {failures:?}",
scenario.suffix()
);
}
fn replay_postauth_lane_offline(lane_id: &str, scenario: PostAuthScenario) -> Result<()> {
let cassette_bytes = fs::read(postauth_cassette_path(lane_id, scenario))
.map_err(|e| Error::Runtime(e.to_string()))?;
let manifest_text = fs::read_to_string(postauth_manifest_path(lane_id, scenario))
.map_err(|e| Error::Runtime(e.to_string()))?;
let manifest = parse_manifest(&manifest_text);
let expected_checksum = manifest
.get("checksum_sha256")
.ok_or_else(|| Error::Runtime("manifest missing checksum_sha256".into()))?;
if &sha256_hex(&cassette_bytes) != expected_checksum {
return Err(Error::Runtime(format!(
"{lane_id}: cassette checksum != manifest"
)));
}
let leaks = scan_for_secret_fields(&cassette_bytes);
if !leaks.is_empty() {
return Err(Error::Runtime(format!("{lane_id}: secret leak {leaks:?}")));
}
let need = |key: &str| -> Result<String> {
manifest
.get(key)
.cloned()
.ok_or_else(|| Error::Runtime(format!("{lane_id}: manifest missing {key}")))
};
let parse_u8 = |key: &str| -> Result<u8> {
need(key)?
.parse()
.map_err(|_| Error::Runtime(format!("{lane_id}: bad {key}")))
};
let ttc_field_version = parse_u8("ttc_field_version")?;
let charset_id: u16 = need("charset_id")?
.parse()
.map_err(|_| Error::Runtime(format!("{lane_id}: bad charset_id")))?;
let max_string_size: u32 = need("max_string_size")?
.parse()
.map_err(|_| Error::Runtime(format!("{lane_id}: bad max_string_size")))?;
let ttc_seq_num = parse_u8("ttc_seq_num")?;
let eor: bool = need("supports_end_of_response")?
.parse()
.map_err(|_| Error::Runtime(format!("{lane_id}: bad supports_end_of_response")))?;
let oob: bool = need("supports_oob")?
.parse()
.map_err(|_| Error::Runtime(format!("{lane_id}: bad supports_oob")))?;
let sdu: usize = need("sdu")?
.parse()
.map_err(|_| Error::Runtime(format!("{lane_id}: bad sdu")))?;
let expected_value = need("expected_value")?;
let caps = oracledb_protocol::thin::ClientCapabilities {
ttc_field_version,
max_string_size,
charset_id,
};
let value = replay_postauth(&cassette_bytes, caps, ttc_seq_num, eor, oob, sdu, scenario)?;
if value.as_deref() != Some(expected_value.as_str()) {
return Err(Error::Runtime(format!(
"{lane_id}: replay value {value:?} != {expected_value:?}"
)));
}
Ok(())
}
#[test]
fn replay_postauth_query_cassettes_offline() {
replay_postauth_scenario_offline(PostAuthScenario::ExecuteSelect);
}
#[test]
fn replay_postauth_lob_cassettes_offline() {
replay_postauth_scenario_offline(PostAuthScenario::LobBlob);
}
fn replay_postauth_scenario_offline(scenario: PostAuthScenario) {
let mut failures = Vec::new();
for lane in postauth_lanes() {
if !postauth_cassette_path(lane.id, scenario).exists() {
eprintln!(
"skip {}: no committed {} cassette",
lane.id,
scenario.suffix()
);
continue;
}
if let Err(err) = replay_postauth_lane_offline(lane.id, scenario) {
failures.push(err.to_string());
}
}
assert!(
failures.is_empty(),
"{} replay failures: {failures:?}",
scenario.suffix()
);
}
const SYNTHETIC_FIXTURE: &str = "select_7_plus_5.tns-cassette";
fn expected_cassette_files() -> std::collections::BTreeSet<String> {
let mut expected = std::collections::BTreeSet::new();
for lane in lanes() {
expected.insert(format!("{}-connect.tns-cassette", lane.id));
}
for lane in postauth_lanes() {
expected.insert(format!("{}-postauth.tns-cassette", lane.id));
expected.insert(format!("{}-lob.tns-cassette", lane.id));
}
expected.insert(SYNTHETIC_FIXTURE.to_string());
expected
}
#[test]
fn every_committed_cassette_is_covered_by_a_replay_lane() {
let dir = fixtures_dir();
let expected = expected_cassette_files();
let mut orphans = Vec::new();
let mut empty = Vec::new();
for entry in fs::read_dir(&dir).expect("fixtures/cassettes dir is readable") {
let entry = entry.expect("dir entry is readable");
let name = entry.file_name().to_string_lossy().into_owned();
if !name.ends_with(".tns-cassette") {
continue;
}
if !expected.contains(&name) {
orphans.push(name.clone());
}
let len = fs::metadata(entry.path())
.expect("cassette metadata is readable")
.len();
if len == 0 {
empty.push(name);
}
}
assert!(
orphans.is_empty(),
"orphan cassette(s) with no replay lane: {orphans:?}; \
wire them onto a lane in version_cassettes.rs (or the synthetic fixture set)"
);
assert!(empty.is_empty(), "empty committed cassette(s): {empty:?}");
}