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_AQ_MESSAGE_ID_LENGTH,
TNS_PACKET_TYPE_ACCEPT, TNS_PACKET_TYPE_CONNECT, TNS_PACKET_TYPE_RESEND,
};
use oracledb_protocol::wire::{encode_packet, PacketLengthWidth, ProtocolLimits, TtcWriter};
use oracledb_protocol::ProtocolError;
use crate::transport::{self, scan_for_secret_fields, 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 SYNTHETIC_19C_PROTOCOL_VERSION: u16 = 318;
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)),
}
}
}
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:?}");
}
fn synthetic_19c_accept_payload() -> Vec<u8> {
let mut writer = TtcWriter::new();
writer.write_u16be(SYNTHETIC_19C_PROTOCOL_VERSION);
writer.write_u16be(0); writer.write_raw(&[0; 10]);
writer.write_u8(0); writer.write_raw(&[0; 9]);
writer.write_u32be(u32::from(ADVERTISED_SDU));
writer.write_raw(&[0; 5]);
writer.write_u32be(0); writer.into_bytes()
}
fn synthetic_19c_caps_cassette() -> Result<Vec<u8>> {
let connect_data = capture_connect_descriptor("NINETEEN_C_PROFILE");
let connect_payload = build_connect_packet_payload(&connect_data, ADVERTISED_SDU)?;
let connect_packet = encode_packet(
TNS_PACKET_TYPE_CONNECT,
0,
None,
&connect_payload,
PacketLengthWidth::Legacy16,
)?;
let accept_packet = encode_packet(
TNS_PACKET_TYPE_ACCEPT,
0,
None,
&synthetic_19c_accept_payload(),
PacketLengthWidth::Legacy16,
)?;
let mut cassette = Vec::new();
cassette::write_header(&mut cassette);
cassette::write_frame(&mut cassette, Direction::ClientToServer, 0, &connect_packet);
cassette::write_frame(&mut cassette, Direction::ServerToClient, 1, &accept_packet);
Ok(cassette)
}
#[test]
fn replay_synthetic_19c_caps_cassette_offline() -> Result<()> {
let cassette_bytes = synthetic_19c_caps_cassette()?;
assert!(
scan_for_secret_fields(&cassette_bytes).is_empty(),
"synthetic handshake must remain secret-free"
);
let (read, write, audit) =
transport::replay_split_with_audit(&cassette_bytes, ReplayWriteMode::Check)
.map_err(|err| Error::Runtime(format!("invalid 19c profile cassette: {err}")))?;
let mut core = ConnectionCore::<DriverTransport>::from_halves(read, write, "19c-profile");
core.set_protocol_limits(ProtocolLimits::DEFAULT)?;
let connect_data = capture_connect_descriptor("NINETEEN_C_PROFILE");
let runtime = build_io_runtime()?;
runtime.block_on(async move {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("missing ambient Cx in 19c replay runtime".into()))?;
let accept = drive_connect_handshake(&mut core, &cx, &connect_data).await?;
let info = parse_accept_payload(&accept.payload)?;
assert_eq!(info.protocol_version, SYNTHETIC_19C_PROTOCOL_VERSION);
assert!(!info.supports_fast_auth);
assert!(!info.supports_end_of_response);
Ok::<_, Error>(())
})?;
audit
.assert_finished()
.map_err(|err| Error::Runtime(format!("synthetic 19c replay: {err}")))?;
Ok(())
}
#[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,
db_unique_name: None,
capabilities,
ttc_seq_num,
sdu,
protocol_version: 0,
supports_fast_auth: false,
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()),
capture_guard: None,
}
}
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;
const AQ_QUEUE_NAME: &str = "RUST_CASS_RAWQ";
const AQ_SCENARIO_DESC: &str = "aq: raw enqueue + dequeue-by-msgid (single-consumer)";
const AQ_EXPECTED_VALUE: &str = "rust-oracledb cassette aq raw payload";
const DPL_TABLE_NAME: &str = "RUST_CASS_DPL";
const DPL_SCHEMA: &str = "PYTHONTEST";
const DPL_COLUMN: &str = "v";
const DPL_READBACK_SQL: &str = "select v from RUST_CASS_DPL where rownum = 1";
const DPL_SCENARIO_DESC: &str = "dpl: direct-path load one number row + read-back";
const DPL_EXPECTED_VALUE: &str = "4242";
#[derive(Clone, Debug, PartialEq, Eq)]
enum ReplayMode {
Check,
DecodedAssert { id_lengths: Vec<usize> },
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum PostAuthScenario {
ExecuteSelect,
LobBlob,
AqRawRoundTrip,
DplLoadReadback,
}
impl PostAuthScenario {
fn suffix(self) -> &'static str {
match self {
PostAuthScenario::ExecuteSelect => "postauth",
PostAuthScenario::LobBlob => "lob",
PostAuthScenario::AqRawRoundTrip => "aq",
PostAuthScenario::DplLoadReadback => "dpl",
}
}
fn sql(self) -> &'static str {
match self {
PostAuthScenario::ExecuteSelect => POSTAUTH_SQL,
PostAuthScenario::LobBlob => LOB_SCENARIO_DESC,
PostAuthScenario::AqRawRoundTrip => AQ_SCENARIO_DESC,
PostAuthScenario::DplLoadReadback => DPL_SCENARIO_DESC,
}
}
fn expected_value(self) -> &'static str {
match self {
PostAuthScenario::ExecuteSelect => POSTAUTH_EXPECTED_VALUE,
PostAuthScenario::LobBlob => LOB_EXPECTED_VALUE,
PostAuthScenario::AqRawRoundTrip => AQ_EXPECTED_VALUE,
PostAuthScenario::DplLoadReadback => DPL_EXPECTED_VALUE,
}
}
fn tag(self) -> &'static str {
match self {
PostAuthScenario::ExecuteSelect => "execute_select",
PostAuthScenario::LobBlob => "temp_lob_create_write_read",
PostAuthScenario::AqRawRoundTrip => "aq_raw_enq_deq_by_msgid",
PostAuthScenario::DplLoadReadback => "dpl_load_readback",
}
}
fn client_writes(self) -> usize {
match self {
PostAuthScenario::ExecuteSelect => 1,
PostAuthScenario::LobBlob => 3,
PostAuthScenario::AqRawRoundTrip => 2,
PostAuthScenario::DplLoadReadback => 4,
}
}
fn replay_mode(self) -> ReplayMode {
match self {
PostAuthScenario::ExecuteSelect | PostAuthScenario::LobBlob => ReplayMode::Check,
PostAuthScenario::AqRawRoundTrip => ReplayMode::DecodedAssert {
id_lengths: vec![TNS_AQ_MESSAGE_ID_LENGTH],
},
PostAuthScenario::DplLoadReadback => ReplayMode::DecodedAssert {
id_lengths: vec![1, 2, 3],
},
}
}
fn applies_to_lane(self, lane_id: &str) -> bool {
match self {
PostAuthScenario::AqRawRoundTrip | PostAuthScenario::DplLoadReadback => {
lane_id == "free23"
}
_ => true,
}
}
}
#[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))
}
PostAuthScenario::AqRawRoundTrip => {
use oracledb_protocol::thin::aq::{
AqDeqOptions, AqDeqPayload, AqEnqOptions, AqMsgProps, AqPayloadKind,
AqPayloadValue, AqQueueDesc,
};
let queue = AqQueueDesc::new(AQ_QUEUE_NAME.to_string(), AqPayloadKind::Raw, None);
let props = AqMsgProps {
payload: Some(AqPayloadValue::Raw(AQ_EXPECTED_VALUE.as_bytes().to_vec())),
..AqMsgProps::default()
};
let enq_options = AqEnqOptions {
visibility: 1,
..AqEnqOptions::default()
};
let msgid = conn
.aq_enq_one(cx, &queue, &props, &enq_options)
.await?
.ok_or_else(|| Error::Runtime("AQ enqueue returned no message id".into()))?;
let deq_options = AqDeqOptions {
visibility: 1,
msgid: Some(msgid),
..AqDeqOptions::default()
};
let result = conn.aq_deq_one(cx, &queue, &deq_options).await?;
let text = match result.message.and_then(|m| m.payload) {
Some(AqDeqPayload::Raw(bytes)) => Some(
String::from_utf8(bytes)
.map_err(|e| Error::Runtime(format!("AQ RAW payload not UTF-8: {e}")))?,
),
Some(_) => {
return Err(Error::Runtime("AQ dequeue returned non-RAW payload".into()))
}
None => None,
};
Ok(text)
}
PostAuthScenario::DplLoadReadback => {
use oracledb_protocol::dpl::DirectPathColumnValue;
use oracledb_protocol::thin::{ExecuteOptions, QueryValue};
let columns = [DPL_COLUMN.to_string()];
let rows = [vec![DirectPathColumnValue::Number(
DPL_EXPECTED_VALUE.to_string(),
)]];
conn.direct_path_load(cx, DPL_SCHEMA, DPL_TABLE_NAME, &columns, &rows, 1000)
.await?;
let exec = conn
.execute_raw(
cx,
DPL_READBACK_SQL,
2,
&[],
ExecuteOptions::default(),
None,
)
.await?;
Ok(exec
.cell(0, 0)
.and_then(QueryValue::as_number_text)
.map(|c| c.to_string()))
}
}
}
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,
})
}
fn client_write_stream(cassette_bytes: &[u8]) -> Result<Vec<u8>> {
let frames = cassette::decode_all(cassette_bytes)
.map_err(|e| Error::Runtime(format!("cassette decode: {e}")))?;
let mut out = Vec::new();
for frame in frames {
if frame.direction == Direction::ClientToServer {
out.extend_from_slice(&frame.bytes);
}
}
Ok(out)
}
fn assert_masked_writes_match(
produced: &[u8],
recorded: &[u8],
id_lengths: &[usize],
) -> Result<()> {
if produced.len() != recorded.len() {
return Err(Error::Runtime(format!(
"decoded-assert: re-issued client writes are {} bytes, recording is {} — a \
server-assigned id substitution never changes length, so this is a real \
request regression",
produced.len(),
recorded.len()
)));
}
let mut i = 0usize;
while i < produced.len() {
if produced[i] == recorded[i] {
i += 1;
continue;
}
let run_start = i;
while i < produced.len() && produced[i] != recorded[i] {
i += 1;
}
let run_len = i - run_start;
if !id_lengths.contains(&run_len) {
return Err(Error::Runtime(format!(
"decoded-assert: re-issued client writes differ from the recording in a \
{run_len}-byte run at offset {run_start}, whose length is not a known \
server-assigned id length ({id_lengths:?}) — a real request regression, not \
id noise"
)));
}
}
Ok(())
}
#[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,
mode: &ReplayMode,
) -> Result<Option<String>> {
let write_mode = match mode {
ReplayMode::Check => ReplayWriteMode::Check,
ReplayMode::DecodedAssert { .. } => ReplayWriteMode::Ignore,
};
let capture = matches!(mode, ReplayMode::DecodedAssert { .. }).then(transport::capture_scope);
let (read, write, audit) = transport::replay_split_with_audit(sliced, write_mode)
.map_err(|e| Error::Runtime(format!("replay split: {e}")))?;
let (read, write) = if capture.is_some() {
transport::wrap_if_capturing((read, write))
} else {
(read, write)
};
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
})?;
if let ReplayMode::DecodedAssert { id_lengths } = mode {
let scope = capture.expect("DecodedAssert installs a capture scope");
let produced = client_write_stream(&scope.to_cassette_bytes())?;
let recorded = client_write_stream(sliced)?;
assert_masked_writes_match(&produced, &recorded, id_lengths)?;
}
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",
"replay_mode = \"{}\"\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(""),
match scenario.replay_mode() {
ReplayMode::Check => "check",
ReplayMode::DecodedAssert { .. } => "decoded_assert",
},
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);
}
#[test]
#[ignore = "records the live AQ post-auth cassette; needs free23 + a provisioned queue"]
fn record_postauth_aq_cassettes() {
record_postauth_scenario(PostAuthScenario::AqRawRoundTrip);
}
#[test]
#[ignore = "records the live DPL post-auth cassette; needs free23 + a provisioned table"]
fn record_postauth_dpl_cassettes() {
record_postauth_scenario(PostAuthScenario::DplLoadReadback);
}
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() {
if !scenario.applies_to_lane(lane.id) {
continue;
}
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,
&scenario.replay_mode(),
) {
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()
);
}
struct PostAuthFixture {
cassette_bytes: Vec<u8>,
caps: oracledb_protocol::thin::ClientCapabilities,
ttc_seq_num: u8,
supports_end_of_response: bool,
supports_oob: bool,
sdu: usize,
expected_value: String,
}
fn read_postauth_fixture(lane_id: &str, scenario: PostAuthScenario) -> Result<PostAuthFixture> {
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")?;
Ok(PostAuthFixture {
cassette_bytes,
caps: oracledb_protocol::thin::ClientCapabilities {
ttc_field_version,
max_string_size,
charset_id,
},
ttc_seq_num,
supports_end_of_response: eor,
supports_oob: oob,
sdu,
expected_value,
})
}
fn replay_postauth_lane_offline(lane_id: &str, scenario: PostAuthScenario) -> Result<()> {
let fx = read_postauth_fixture(lane_id, scenario)?;
let value = replay_postauth(
&fx.cassette_bytes,
fx.caps,
fx.ttc_seq_num,
fx.supports_end_of_response,
fx.supports_oob,
fx.sdu,
scenario,
&scenario.replay_mode(),
)?;
if value.as_deref() != Some(fx.expected_value.as_str()) {
return Err(Error::Runtime(format!(
"{lane_id}: replay value {value:?} != {:?}",
fx.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);
}
#[test]
fn replay_postauth_aq_cassettes_offline() {
replay_postauth_scenario_offline(PostAuthScenario::AqRawRoundTrip);
}
#[test]
fn replay_postauth_dpl_cassettes_offline() {
replay_postauth_scenario_offline(PostAuthScenario::DplLoadReadback);
}
fn replay_postauth_scenario_offline(scenario: PostAuthScenario) {
let mut failures = Vec::new();
for lane in postauth_lanes() {
if !scenario.applies_to_lane(lane.id) {
continue;
}
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 ECHOED_ID_LEN: usize = TNS_AQ_MESSAGE_ID_LENGTH;
fn contains_subslice(haystack: &[u8], needle: &[u8]) -> bool {
needle.len() <= haystack.len() && haystack.windows(needle.len()).any(|w| w == needle)
}
fn find_echoed_id_region(frames: &[cassette::Frame]) -> Option<(usize, usize)> {
for (i, frame) in frames.iter().enumerate() {
if frame.direction != Direction::ClientToServer || frame.bytes.len() < ECHOED_ID_LEN {
continue;
}
let earlier_server: Vec<&[u8]> = frames[..i]
.iter()
.filter(|f| f.direction == Direction::ServerToClient)
.map(|f| f.bytes.as_slice())
.collect();
for offset in 0..=(frame.bytes.len() - ECHOED_ID_LEN) {
let window = &frame.bytes[offset..offset + ECHOED_ID_LEN];
if window.iter().all(|&b| b == window[0]) {
continue; }
if earlier_server.iter().any(|s| contains_subslice(s, window)) {
return Some((i, offset));
}
}
}
None
}
fn mutate_echoed_id(cassette_bytes: &[u8]) -> Result<(Vec<u8>, usize)> {
let mut frames = cassette::decode_all(cassette_bytes)
.map_err(|e| Error::Runtime(format!("cassette decode: {e}")))?;
let (frame_index, offset) = find_echoed_id_region(&frames).ok_or_else(|| {
Error::Runtime("cassette has no server-assigned id echoed in a request frame".into())
})?;
for byte in &mut frames[frame_index].bytes[offset..offset + ECHOED_ID_LEN] {
*byte ^= 0xFF;
}
Ok((reencode_frames(&frames), frame_index))
}
fn id_echo_proof_lane() -> Option<(PostAuthScenario, &'static str)> {
for lane in postauth_lanes() {
let aq = PostAuthScenario::AqRawRoundTrip;
if aq.applies_to_lane(lane.id) && postauth_cassette_path(lane.id, aq).exists() {
return Some((aq, lane.id));
}
}
for lane in postauth_lanes() {
let lob = PostAuthScenario::LobBlob;
if postauth_cassette_path(lane.id, lob).exists() {
return Some((lob, lane.id));
}
}
None
}
#[test]
fn decoded_assert_survives_server_id_divergence_that_check_rejects() {
let Some((scenario, lane_id)) = id_echo_proof_lane() else {
eprintln!("skip: no committed id-echo cassette (AQ or LOB) available");
return;
};
let fx = read_postauth_fixture(lane_id, scenario).expect("id-echo fixture loads");
let replay = |bytes: &[u8], mode: &ReplayMode| {
replay_postauth(
bytes,
fx.caps,
fx.ttc_seq_num,
fx.supports_end_of_response,
fx.supports_oob,
fx.sdu,
scenario,
mode,
)
};
let pristine = replay(&fx.cassette_bytes, &ReplayMode::Check)
.expect("pristine cassette replays byte-exact under Check");
assert_eq!(
pristine.as_deref(),
Some(fx.expected_value.as_str()),
"{lane_id}/{}: pristine decoded value",
scenario.suffix()
);
let (mutated, _frame) =
mutate_echoed_id(&fx.cassette_bytes).expect("cassette echoes a server id to mutate");
assert_ne!(
mutated, fx.cassette_bytes,
"mutation must actually change the cassette"
);
let checked = replay(&mutated, &ReplayMode::Check);
assert!(
checked.is_err(),
"{lane_id}/{}: byte-exact Check must reject a request whose server-id bytes diverge \
from the recording, but it passed: {checked:?}",
scenario.suffix()
);
let decoded = replay(
&mutated,
&ReplayMode::DecodedAssert {
id_lengths: vec![ECHOED_ID_LEN],
},
)
.expect("DecodedAssert masks the server-id run and replays green");
assert_eq!(
decoded.as_deref(),
Some(fx.expected_value.as_str()),
"{lane_id}/{}: decoded-assert decoded value",
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));
if PostAuthScenario::AqRawRoundTrip.applies_to_lane(lane.id) {
expected.insert(format!("{}-aq.tns-cassette", lane.id));
}
if PostAuthScenario::DplLoadReadback.applies_to_lane(lane.id) {
expected.insert(format!("{}-dpl.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:?}");
}