use std::fmt::Debug;
use std::io;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use asupersync::io::{AsyncRead, AsyncWrite, ReadBuf};
use asupersync::net::{OwnedReadHalf, OwnedWriteHalf, TcpStream};
use asupersync::tls::TlsStream;
#[cfg(feature = "cassette")]
pub use cassette_seam::{
capture_scope, CaptureScope, CassetteError, CassetteRecorder, ReplayMismatch, ReplayWriteMode,
};
#[cfg(all(test, feature = "cassette"))]
pub(crate) use cassette_seam::{replay_split, replay_split_with_audit};
type SharedTls = Arc<Mutex<TlsStream<TcpStream>>>;
pub(crate) trait WireTransport {
type Read: AsyncRead + Debug + Send + Unpin + 'static;
type Write: AsyncWrite + Debug + Send + Unpin + 'static;
}
pub(crate) type TransportHalves<T> = (<T as WireTransport>::Read, <T as WireTransport>::Write);
pub(crate) trait Connector {
type Transport: WireTransport;
fn plain_split(&self, stream: TcpStream) -> TransportHalves<Self::Transport>;
fn tls_split(&self, stream: TlsStream<TcpStream>) -> TransportHalves<Self::Transport>;
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct OracleWireTransport;
impl WireTransport for OracleWireTransport {
type Read = OracleReadHalf;
type Write = OracleWriteHalf;
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct OracleConnector;
impl Connector for OracleConnector {
type Transport = OracleWireTransport;
fn plain_split(&self, stream: TcpStream) -> TransportHalves<Self::Transport> {
plain_split(stream)
}
fn tls_split(&self, stream: TlsStream<TcpStream>) -> TransportHalves<Self::Transport> {
tls_split(stream)
}
}
#[cfg(test)]
mod tests {
use super::{Connector, OracleConnector, OracleReadHalf, OracleWriteHalf, WireTransport};
fn assert_wire_transport<T: WireTransport<Read = OracleReadHalf, Write = OracleWriteHalf>>() {}
#[test]
fn oracle_connector_uses_current_transport_halves() {
assert_wire_transport::<<OracleConnector as Connector>::Transport>();
}
}
pub(crate) enum OracleReadHalf {
Plain(OwnedReadHalf),
Tls(SharedTls),
#[cfg(feature = "cassette")]
Recording(cassette_seam::RecordingRead),
#[cfg(feature = "cassette")]
#[allow(dead_code)]
Replay(cassette_seam::ReplayRead),
}
pub(crate) enum OracleWriteHalf {
Plain(OwnedWriteHalf),
Tls(SharedTls),
#[cfg(feature = "cassette")]
Recording(cassette_seam::RecordingWrite),
#[cfg(feature = "cassette")]
#[allow(dead_code)]
Replay(cassette_seam::ReplayWrite),
}
impl std::fmt::Debug for OracleReadHalf {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Plain(_) => f.write_str("OracleReadHalf::Plain"),
Self::Tls(_) => f.write_str("OracleReadHalf::Tls"),
#[cfg(feature = "cassette")]
Self::Recording(_) => f.write_str("OracleReadHalf::Recording"),
#[cfg(feature = "cassette")]
Self::Replay(_) => f.write_str("OracleReadHalf::Replay"),
}
}
}
impl std::fmt::Debug for OracleWriteHalf {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Plain(_) => f.write_str("OracleWriteHalf::Plain"),
Self::Tls(_) => f.write_str("OracleWriteHalf::Tls"),
#[cfg(feature = "cassette")]
Self::Recording(_) => f.write_str("OracleWriteHalf::Recording"),
#[cfg(feature = "cassette")]
Self::Replay(_) => f.write_str("OracleWriteHalf::Replay"),
}
}
}
#[must_use]
pub(crate) fn plain_split(stream: TcpStream) -> (OracleReadHalf, OracleWriteHalf) {
let (read, write) = stream.into_split();
let halves = (OracleReadHalf::Plain(read), OracleWriteHalf::Plain(write));
#[cfg(feature = "cassette")]
let halves = cassette_seam::wrap_if_capturing(halves);
halves
}
#[must_use]
pub(crate) fn tls_split(stream: TlsStream<TcpStream>) -> (OracleReadHalf, OracleWriteHalf) {
let shared: SharedTls = Arc::new(Mutex::new(stream));
let halves = (
OracleReadHalf::Tls(Arc::clone(&shared)),
OracleWriteHalf::Tls(shared),
);
#[cfg(feature = "cassette")]
let halves = cassette_seam::wrap_if_capturing(halves);
halves
}
fn poisoned() -> io::Error {
io::Error::other("TLS stream mutex poisoned")
}
impl AsyncRead for OracleReadHalf {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Plain(read) => Pin::new(read).poll_read(cx, buf),
Self::Tls(shared) => {
let mut guard = shared.lock().map_err(|_| poisoned())?;
Pin::new(&mut *guard).poll_read(cx, buf)
}
#[cfg(feature = "cassette")]
Self::Recording(rec) => rec.poll_read(cx, buf),
#[cfg(feature = "cassette")]
Self::Replay(replay) => replay.poll_read(buf),
}
}
}
impl AsyncWrite for OracleWriteHalf {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
match self.get_mut() {
Self::Plain(write) => Pin::new(write).poll_write(cx, buf),
Self::Tls(shared) => {
let mut guard = shared.lock().map_err(|_| poisoned())?;
Pin::new(&mut *guard).poll_write(cx, buf)
}
#[cfg(feature = "cassette")]
Self::Recording(rec) => rec.poll_write(cx, buf),
#[cfg(feature = "cassette")]
Self::Replay(replay) => replay.poll_write(buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Plain(write) => Pin::new(write).poll_flush(cx),
Self::Tls(shared) => {
let mut guard = shared.lock().map_err(|_| poisoned())?;
Pin::new(&mut *guard).poll_flush(cx)
}
#[cfg(feature = "cassette")]
Self::Recording(rec) => rec.poll_flush(cx),
#[cfg(feature = "cassette")]
Self::Replay(_) => Poll::Ready(Ok(())),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Plain(write) => Pin::new(write).poll_shutdown(cx),
Self::Tls(shared) => {
let mut guard = shared.lock().map_err(|_| poisoned())?;
Pin::new(&mut *guard).poll_shutdown(cx)
}
#[cfg(feature = "cassette")]
Self::Recording(rec) => rec.poll_shutdown(cx),
#[cfg(feature = "cassette")]
Self::Replay(_) => Poll::Ready(Ok(())),
}
}
}
#[cfg(feature = "cassette")]
mod cassette_seam {
use std::collections::VecDeque;
use std::io;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use std::time::Instant;
use asupersync::io::{AsyncRead, AsyncWrite, ReadBuf};
use oracledb_protocol::net::cassette::{self, Direction, Frame};
#[cfg(test)]
use std::fmt;
use super::{OracleReadHalf, OracleWriteHalf};
pub use oracledb_protocol::net::cassette::CassetteError;
#[derive(Clone)]
pub struct CassetteRecorder {
inner: Arc<Mutex<RecorderState>>,
}
struct RecorderState {
start: Instant,
frames: Vec<Frame>,
}
impl CassetteRecorder {
#[must_use]
pub fn new() -> Self {
Self {
inner: Arc::new(Mutex::new(RecorderState {
start: Instant::now(),
frames: Vec::new(),
})),
}
}
fn record(&self, direction: Direction, bytes: &[u8]) {
if let Ok(mut state) = self.inner.lock() {
let micros = u64::try_from(state.start.elapsed().as_micros()).unwrap_or(u64::MAX);
state.frames.push(Frame {
direction,
micros,
bytes: bytes.to_vec(),
});
}
}
#[must_use]
pub fn frame_count(&self) -> usize {
self.inner.lock().map(|s| s.frames.len()).unwrap_or(0)
}
#[must_use]
pub fn to_cassette_bytes(&self) -> Vec<u8> {
let mut out = Vec::new();
cassette::write_header(&mut out);
if let Ok(state) = self.inner.lock() {
for frame in &state.frames {
cassette::write_frame(&mut out, frame.direction, frame.micros, &frame.bytes);
}
}
out
}
#[must_use]
pub fn into_cassette_bytes(self) -> Vec<u8> {
self.to_cassette_bytes()
}
}
impl Default for CassetteRecorder {
fn default() -> Self {
Self::new()
}
}
#[must_use]
pub(crate) fn recording_split(
read: OracleReadHalf,
write: OracleWriteHalf,
recorder: CassetteRecorder,
) -> (OracleReadHalf, OracleWriteHalf) {
(
OracleReadHalf::Recording(RecordingRead {
inner: Box::new(read),
recorder: recorder.clone(),
}),
OracleWriteHalf::Recording(RecordingWrite {
inner: Box::new(write),
recorder,
}),
)
}
thread_local! {
static ACTIVE_RECORDER: std::cell::RefCell<Option<CassetteRecorder>> =
const { std::cell::RefCell::new(None) };
}
#[must_use]
pub(super) fn wrap_if_capturing(
halves: (OracleReadHalf, OracleWriteHalf),
) -> (OracleReadHalf, OracleWriteHalf) {
match ACTIVE_RECORDER.with(|slot| slot.borrow().clone()) {
Some(recorder) => recording_split(halves.0, halves.1, recorder),
None => halves,
}
}
#[must_use = "dropping the CaptureScope immediately stops recording"]
pub struct CaptureScope {
recorder: CassetteRecorder,
previous: Option<CassetteRecorder>,
}
pub fn capture_scope() -> CaptureScope {
let recorder = CassetteRecorder::new();
let previous = ACTIVE_RECORDER.with(|slot| slot.borrow_mut().replace(recorder.clone()));
CaptureScope { recorder, previous }
}
impl CaptureScope {
#[must_use]
pub fn recorder(&self) -> &CassetteRecorder {
&self.recorder
}
#[must_use]
pub fn to_cassette_bytes(&self) -> Vec<u8> {
self.recorder.to_cassette_bytes()
}
}
impl Drop for CaptureScope {
fn drop(&mut self) {
ACTIVE_RECORDER.with(|slot| {
*slot.borrow_mut() = self.previous.take();
});
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Default)]
pub enum ReplayWriteMode {
#[default]
Ignore,
Check,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct ReplayMismatch {
pub frame_index: usize,
pub expected: Vec<u8>,
pub actual: Vec<u8>,
}
#[derive(Clone, Debug)]
pub(crate) struct ReplayAudit {
inner: Arc<Mutex<ReplayAuditState>>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
struct ReplayAuditState {
read_frames_remaining: usize,
read_bytes_remaining: usize,
write_frames_remaining: usize,
write_bytes_remaining: usize,
mismatch: Option<ReplayMismatch>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
#[cfg(test)]
pub(crate) enum ReplayAuditError {
Mismatch(ReplayMismatch),
UnreadFrames { frames: usize, bytes: usize },
UnwrittenFrames { frames: usize, bytes: usize },
Poisoned,
}
#[cfg(test)]
impl fmt::Display for ReplayAuditError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Mismatch(mismatch) => write!(
f,
"replay: write mismatch at frame {} (expected {} bytes, got {} bytes)",
mismatch.frame_index,
mismatch.expected.len(),
mismatch.actual.len()
),
Self::UnreadFrames { frames, bytes } => write!(
f,
"replay: unread server frames remain ({frames} frames, {bytes} bytes)"
),
Self::UnwrittenFrames { frames, bytes } => write!(
f,
"replay: expected client writes remain ({frames} frames, {bytes} bytes)"
),
Self::Poisoned => f.write_str("replay: audit mutex poisoned"),
}
}
}
#[cfg(test)]
impl std::error::Error for ReplayAuditError {}
impl ReplayAudit {
fn new(reads: &VecDeque<Vec<u8>>, writes: &VecDeque<Vec<u8>>) -> Self {
Self {
inner: Arc::new(Mutex::new(ReplayAuditState {
read_frames_remaining: reads.len(),
read_bytes_remaining: reads.iter().map(Vec::len).sum(),
write_frames_remaining: writes.len(),
write_bytes_remaining: writes.iter().map(Vec::len).sum(),
mismatch: None,
})),
}
}
fn note_read_bytes(&self, n: usize) {
if let Ok(mut state) = self.inner.lock() {
state.read_bytes_remaining = state.read_bytes_remaining.saturating_sub(n);
}
}
fn note_read_frame_consumed(&self) {
if let Ok(mut state) = self.inner.lock() {
state.read_frames_remaining = state.read_frames_remaining.saturating_sub(1);
}
}
fn note_write_bytes(&self, n: usize) {
if let Ok(mut state) = self.inner.lock() {
state.write_bytes_remaining = state.write_bytes_remaining.saturating_sub(n);
}
}
fn note_write_frame_consumed(&self) {
if let Ok(mut state) = self.inner.lock() {
state.write_frames_remaining = state.write_frames_remaining.saturating_sub(1);
}
}
fn note_mismatch(&self, mismatch: ReplayMismatch) {
if let Ok(mut state) = self.inner.lock() {
if state.mismatch.is_none() {
state.mismatch = Some(mismatch);
}
}
}
#[cfg(test)]
pub(crate) fn assert_finished(&self) -> Result<(), ReplayAuditError> {
let state = self
.inner
.lock()
.map_err(|_| ReplayAuditError::Poisoned)?
.clone();
if let Some(mismatch) = state.mismatch {
return Err(ReplayAuditError::Mismatch(mismatch));
}
if state.read_frames_remaining != 0 || state.read_bytes_remaining != 0 {
return Err(ReplayAuditError::UnreadFrames {
frames: state.read_frames_remaining,
bytes: state.read_bytes_remaining,
});
}
if state.write_frames_remaining != 0 || state.write_bytes_remaining != 0 {
return Err(ReplayAuditError::UnwrittenFrames {
frames: state.write_frames_remaining,
bytes: state.write_bytes_remaining,
});
}
Ok(())
}
}
#[allow(dead_code)]
pub(crate) fn replay_split(
data: &[u8],
write_mode: ReplayWriteMode,
) -> Result<(OracleReadHalf, OracleWriteHalf), CassetteError> {
replay_split_inner(data, write_mode).map(|(read, write, _audit)| (read, write))
}
#[cfg(test)]
pub(crate) fn replay_split_with_audit(
data: &[u8],
write_mode: ReplayWriteMode,
) -> Result<(OracleReadHalf, OracleWriteHalf, ReplayAudit), CassetteError> {
replay_split_inner(data, write_mode)
}
fn replay_split_inner(
data: &[u8],
write_mode: ReplayWriteMode,
) -> Result<(OracleReadHalf, OracleWriteHalf, ReplayAudit), CassetteError> {
let frames = cassette::decode_all(data)?;
let mut reads: VecDeque<Vec<u8>> = VecDeque::new();
let mut writes: VecDeque<Vec<u8>> = VecDeque::new();
for frame in frames {
match frame.direction {
Direction::ServerToClient => reads.push_back(frame.bytes),
Direction::ClientToServer => writes.push_back(frame.bytes),
}
}
let audit_writes = if matches!(write_mode, ReplayWriteMode::Check) {
writes.clone()
} else {
VecDeque::new()
};
let audit = ReplayAudit::new(&reads, &audit_writes);
let mismatch = Arc::new(Mutex::new(None));
Ok((
OracleReadHalf::Replay(ReplayRead {
pending: reads,
offset: 0,
audit: audit.clone(),
}),
OracleWriteHalf::Replay(ReplayWrite {
expected: writes,
offset: 0,
index: 0,
mode: write_mode,
mismatch,
audit: audit.clone(),
}),
audit,
))
}
pub struct RecordingRead {
inner: Box<OracleReadHalf>,
recorder: CassetteRecorder,
}
impl RecordingRead {
pub(super) fn poll_read(
&mut self,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let before = buf.filled().len();
match Pin::new(self.inner.as_mut()).poll_read(cx, buf) {
Poll::Ready(Ok(())) => {
let new = &buf.filled()[before..];
if !new.is_empty() {
self.recorder.record(Direction::ServerToClient, new);
}
Poll::Ready(Ok(()))
}
other => other,
}
}
}
pub struct RecordingWrite {
inner: Box<OracleWriteHalf>,
recorder: CassetteRecorder,
}
impl RecordingWrite {
pub(super) fn poll_write(
&mut self,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
match Pin::new(self.inner.as_mut()).poll_write(cx, buf) {
Poll::Ready(Ok(n)) => {
if n > 0 {
self.recorder.record(Direction::ClientToServer, &buf[..n]);
}
Poll::Ready(Ok(n))
}
other => other,
}
}
pub(super) fn poll_flush(&mut self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(self.inner.as_mut()).poll_flush(cx)
}
pub(super) fn poll_shutdown(&mut self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(self.inner.as_mut()).poll_shutdown(cx)
}
}
pub struct ReplayRead {
pending: VecDeque<Vec<u8>>,
offset: usize,
audit: ReplayAudit,
}
impl ReplayRead {
pub(super) fn poll_read(&mut self, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
while let Some(front) = self.pending.front() {
if self.offset >= front.len() {
self.pending.pop_front();
self.offset = 0;
self.audit.note_read_frame_consumed();
} else {
break;
}
}
let Some(front) = self.pending.front() else {
return Poll::Ready(Ok(()));
};
let available = &front[self.offset..];
let take = available.len().min(buf.remaining());
buf.put_slice(&available[..take]);
self.offset += take;
self.audit.note_read_bytes(take);
if self
.pending
.front()
.is_some_and(|front| self.offset >= front.len())
{
self.pending.pop_front();
self.offset = 0;
self.audit.note_read_frame_consumed();
}
Poll::Ready(Ok(()))
}
}
pub struct ReplayWrite {
expected: VecDeque<Vec<u8>>,
offset: usize,
index: usize,
mode: ReplayWriteMode,
mismatch: Arc<Mutex<Option<ReplayMismatch>>>,
audit: ReplayAudit,
}
impl ReplayWrite {
pub(super) fn poll_write(&mut self, buf: &[u8]) -> Poll<io::Result<usize>> {
if matches!(self.mode, ReplayWriteMode::Ignore) {
return Poll::Ready(Ok(buf.len()));
}
let mut cursor = 0usize;
while cursor < buf.len() {
while let Some(front) = self.expected.front() {
if self.offset >= front.len() {
self.expected.pop_front();
self.offset = 0;
self.index += 1;
self.audit.note_write_frame_consumed();
} else {
break;
}
}
let Some(front) = self.expected.front() else {
self.note_mismatch(self.index, &[], &buf[cursor..]);
return Poll::Ready(Err(io::Error::other(
"replay: write past recorded stream",
)));
};
let remaining = &front[self.offset..];
let chunk = &buf[cursor..];
let take = remaining.len().min(chunk.len());
if remaining[..take] != chunk[..take] {
self.note_mismatch(self.index, remaining, chunk);
return Poll::Ready(Err(io::Error::other("replay: write mismatch")));
}
cursor += take;
self.offset += take;
self.audit.note_write_bytes(take);
if self.offset >= front.len() {
self.expected.pop_front();
self.offset = 0;
self.index += 1;
self.audit.note_write_frame_consumed();
}
}
Poll::Ready(Ok(buf.len()))
}
fn note_mismatch(&self, frame_index: usize, expected: &[u8], actual: &[u8]) {
let mismatch = ReplayMismatch {
frame_index,
expected: expected.to_vec(),
actual: actual.to_vec(),
};
if let Ok(mut slot) = self.mismatch.lock() {
if slot.is_none() {
*slot = Some(mismatch.clone());
}
}
self.audit.note_mismatch(mismatch);
}
}
#[cfg(test)]
mod tests {
use super::*;
use asupersync::io::ReadBuf;
fn read_n(read: &mut OracleReadHalf, n: usize) -> Vec<u8> {
let mut out = Vec::new();
while out.len() < n {
let mut scratch = vec![0u8; n - out.len()];
let mut rb = ReadBuf::new(&mut scratch);
let OracleReadHalf::Replay(replay) = read else {
panic!("expected replay read half");
};
match replay.poll_read(&mut rb) {
Poll::Ready(Ok(())) => {
let filled = rb.filled().to_vec();
if filled.is_empty() {
break; }
out.extend_from_slice(&filled);
}
_ => break,
}
}
out
}
#[test]
fn replay_serves_server_bytes_in_order() {
let recorder = CassetteRecorder::new();
recorder.record(Direction::ClientToServer, &[0x10, 0x20]);
recorder.record(Direction::ServerToClient, &[0xAA, 0xBB, 0xCC]);
recorder.record(Direction::ServerToClient, &[0xDD]);
let bytes = recorder.into_cassette_bytes();
let (mut read, _write) =
replay_split(&bytes, ReplayWriteMode::Ignore).expect("valid cassette");
let got = read_n(&mut read, 4);
assert_eq!(got, vec![0xAA, 0xBB, 0xCC, 0xDD]);
let eof = read_n(&mut read, 1);
assert!(eof.is_empty());
}
#[test]
fn replay_splits_one_transfer_across_small_reads() {
let recorder = CassetteRecorder::new();
recorder.record(Direction::ServerToClient, &[1, 2, 3, 4, 5]);
let bytes = recorder.into_cassette_bytes();
let (mut read, _w) =
replay_split(&bytes, ReplayWriteMode::Ignore).expect("valid cassette");
assert_eq!(read_n(&mut read, 2), vec![1, 2]);
assert_eq!(read_n(&mut read, 2), vec![3, 4]);
assert_eq!(read_n(&mut read, 2), vec![5]);
}
#[test]
fn replay_write_ignore_accepts_anything() {
let recorder = CassetteRecorder::new();
recorder.record(Direction::ClientToServer, &[1, 2, 3]);
let bytes = recorder.into_cassette_bytes();
let (_r, mut write) =
replay_split(&bytes, ReplayWriteMode::Ignore).expect("valid cassette");
let OracleWriteHalf::Replay(w) = &mut write else {
panic!("expected replay write half");
};
assert!(matches!(w.poll_write(&[9, 9, 9, 9]), Poll::Ready(Ok(4))));
}
#[test]
fn replay_write_check_matches_recorded_request_stream() {
let recorder = CassetteRecorder::new();
recorder.record(Direction::ClientToServer, &[0xDE, 0xAD, 0xBE, 0xEF]);
let bytes = recorder.into_cassette_bytes();
let (_r, mut write) =
replay_split(&bytes, ReplayWriteMode::Check).expect("valid cassette");
let OracleWriteHalf::Replay(w) = &mut write else {
panic!("expected replay write half");
};
assert!(matches!(w.poll_write(&[0xDE, 0xAD]), Poll::Ready(Ok(2))));
assert!(matches!(w.poll_write(&[0xBE, 0xEF]), Poll::Ready(Ok(2))));
assert!(w.mismatch.lock().expect("lock").is_none());
}
#[test]
fn replay_write_check_flags_mismatch() {
let recorder = CassetteRecorder::new();
recorder.record(Direction::ClientToServer, &[1, 2, 3, 4]);
let bytes = recorder.into_cassette_bytes();
let (_r, mut write) =
replay_split(&bytes, ReplayWriteMode::Check).expect("valid cassette");
let OracleWriteHalf::Replay(w) = &mut write else {
panic!("expected replay write half");
};
assert!(matches!(w.poll_write(&[1, 2]), Poll::Ready(Ok(2))));
assert!(matches!(w.poll_write(&[9, 9]), Poll::Ready(Err(_))));
let mismatch = w
.mismatch
.lock()
.expect("lock")
.clone()
.expect("a mismatch");
assert_eq!(mismatch.frame_index, 0);
}
#[test]
fn replay_audit_rejects_unread_server_frames() {
let recorder = CassetteRecorder::new();
recorder.record(Direction::ServerToClient, &[0xAA, 0xBB]);
let bytes = recorder.into_cassette_bytes();
let (_read, _write, audit) =
replay_split_with_audit(&bytes, ReplayWriteMode::Check).expect("valid cassette");
let err = audit
.assert_finished()
.expect_err("strict replay must reject unread server bytes");
assert!(
err.to_string().contains("unread server frames"),
"unexpected error: {err}"
);
}
#[test]
fn replay_audit_rejects_unmatched_expected_writes() {
let recorder = CassetteRecorder::new();
recorder.record(Direction::ClientToServer, &[0x10, 0x20]);
let bytes = recorder.into_cassette_bytes();
let (_read, _write, audit) =
replay_split_with_audit(&bytes, ReplayWriteMode::Check).expect("valid cassette");
let err = audit
.assert_finished()
.expect_err("strict replay must reject unmatched client writes");
assert!(
err.to_string().contains("expected client writes"),
"unexpected error: {err}"
);
}
#[test]
fn recorder_serializes_valid_cassette() {
let recorder = CassetteRecorder::new();
recorder.record(Direction::ClientToServer, &[1]);
recorder.record(Direction::ServerToClient, &[2, 3]);
assert_eq!(recorder.frame_count(), 2);
let bytes = recorder.to_cassette_bytes();
let frames = cassette::decode_all(&bytes).expect("decodes");
assert_eq!(frames.len(), 2);
assert_eq!(frames[0].direction, Direction::ClientToServer);
assert_eq!(frames[1].bytes, vec![2, 3]);
}
#[test]
fn replay_split_rejects_garbage() {
let err = replay_split(b"not a cassette", ReplayWriteMode::Ignore)
.expect_err("garbage must fail");
assert_eq!(err, CassetteError::BadMagic);
}
fn empty_cassette() -> Vec<u8> {
CassetteRecorder::new().into_cassette_bytes()
}
#[test]
fn capture_scope_wraps_splits_then_restores_on_drop() {
let (r, w) = replay_split(&empty_cassette(), ReplayWriteMode::Ignore).expect("ok");
let (r, w) = wrap_if_capturing((r, w));
assert!(matches!(r, OracleReadHalf::Replay(_)));
assert!(matches!(w, OracleWriteHalf::Replay(_)));
{
let scope = capture_scope();
let (r, w) = replay_split(&empty_cassette(), ReplayWriteMode::Ignore).expect("ok");
let (r, w) = wrap_if_capturing((r, w));
assert!(matches!(r, OracleReadHalf::Recording(_)));
assert!(matches!(w, OracleWriteHalf::Recording(_)));
assert_eq!(scope.recorder().frame_count(), 0);
}
let (r, w) = replay_split(&empty_cassette(), ReplayWriteMode::Ignore).expect("ok");
let (r, _w) = wrap_if_capturing((r, w));
assert!(matches!(r, OracleReadHalf::Replay(_)));
}
}
}