Skip to content

Commit f7e82a7

Browse files
ADD: Add absolute seek support for DBN
1 parent a8fed3f commit f7e82a7

5 files changed

Lines changed: 271 additions & 5 deletions

File tree

CHANGELOG.md

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,11 @@
11
# Changelog
22

3+
## 0.60.0 - TBD
4+
5+
### Enhancements
6+
- Added `RecordDecoder::seek_to` for absolute byte-offset seeking
7+
- Added `is_compressed` helpers for dynamic DBN readers
8+
39
## 0.59.0 - 2026-06-02
410

511
### Enhancements

rust/dbn/src/decode/dbn/async.rs

Lines changed: 120 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,8 +10,8 @@ use crate::{
1010
decode::{
1111
dbn::fsm::{DbnFsm, ProcessResult},
1212
zstd::zstd_decoder,
13-
AsyncDecodeRecord, AsyncDecodeRecordRef, AsyncSkipBytes, DbnMetadata, VersionUpgradePolicy,
14-
ZSTD_FILE_BUFFER_CAPACITY,
13+
AsyncDecodeRecord, AsyncDecodeRecordRef, AsyncDynReader, AsyncSkipBytes, DbnMetadata,
14+
VersionUpgradePolicy, ZSTD_FILE_BUFFER_CAPACITY,
1515
},
1616
HasRType, Metadata, RecordRef, Result, DBN_VERSION,
1717
};
@@ -458,6 +458,73 @@ where
458458
}
459459
}
460460

461+
impl<R> RecordDecoder<R>
462+
where
463+
R: io::AsyncReadExt + io::AsyncSeekExt + Unpin,
464+
{
465+
/// Seeks to absolute byte position `pos` and clears buffered decoder state.
466+
///
467+
/// # Warning
468+
/// Callers are responsible for ensuring `pos` is a valid record boundary.
469+
///
470+
/// # Cancel safety
471+
/// This method may not be cancellation safe, depending on the cancellation safety
472+
/// of `seek()` of the inner reader `R`.
473+
///
474+
/// # Errors
475+
/// This function returns an error if it fails to seek in the inner reader.
476+
pub async fn seek_to(&mut self, pos: u64) -> crate::Result<()> {
477+
self.reader
478+
.seek(std::io::SeekFrom::Start(pos))
479+
.await
480+
.map(drop)
481+
.map_err(|err| crate::Error::io(err, format!("seeking to byte offset {pos}")))?;
482+
self.fsm.reset_for_seek();
483+
Ok(())
484+
}
485+
}
486+
487+
impl<R> RecordDecoder<AsyncDynReader<R>>
488+
where
489+
R: io::AsyncReadExt + io::AsyncBufReadExt + Unpin,
490+
{
491+
/// Seeks to absolute byte position `pos` and clears buffered decoder state.
492+
///
493+
/// # Warning
494+
/// Callers are responsible for ensuring `pos` is a valid record boundary.
495+
///
496+
/// # Cancel safety
497+
/// This method is not cancel safe.
498+
///
499+
/// # Errors
500+
/// This function returns an error if it fails to seek in the inner reader, or if
501+
/// the input is Zstandard-compressed.
502+
pub async fn seek_to(&mut self, pos: u64) -> crate::Result<()>
503+
where
504+
R: io::AsyncSeekExt,
505+
{
506+
if self.reader.is_compressed() {
507+
return Err(crate::Error::BadArgument {
508+
param_name: "self".to_owned(),
509+
desc: "absolute seek is unsupported for zstd-compressed input".to_owned(),
510+
});
511+
}
512+
self.reader
513+
.get_mut()
514+
.seek(std::io::SeekFrom::Start(pos))
515+
.await
516+
.map(drop)
517+
.map_err(|err| crate::Error::io(err, format!("seeking to byte offset {pos}")))?;
518+
self.fsm.reset_for_seek();
519+
Ok(())
520+
}
521+
522+
/// Returns whether the input is Zstandard-compressed.
523+
pub fn is_compressed(&self) -> bool {
524+
self.reader.is_compressed()
525+
}
526+
}
527+
461528
impl<R> RecordDecoder<R>
462529
where
463530
R: AsyncSkipBytes + io::AsyncReadExt + Unpin,
@@ -763,6 +830,57 @@ mod tests {
763830
assert!(decoder.decode_record::<MboMsg>().await.unwrap().is_none());
764831
}
765832

833+
#[tokio::test]
834+
async fn test_seek_to_resets_buffered_records() {
835+
let mut first = MboMsg {
836+
hd: RecordHeader::new::<MboMsg>(rtype::MBO, 1, 100, 1),
837+
..Default::default()
838+
};
839+
first.order_id = 1;
840+
let mut second = MboMsg {
841+
hd: RecordHeader::new::<MboMsg>(rtype::MBO, 1, 101, 2),
842+
..Default::default()
843+
};
844+
second.order_id = 2;
845+
let mut buffer = Vec::new();
846+
buffer.extend_from_slice(first.as_ref());
847+
buffer.extend_from_slice(second.as_ref());
848+
849+
let mut decoder = RecordDecoder::with_version(
850+
std::io::Cursor::new(buffer),
851+
DBN_VERSION,
852+
VersionUpgradePolicy::AsIs,
853+
false,
854+
)
855+
.unwrap();
856+
857+
assert_eq!(*decoder.decode::<MboMsg>().await.unwrap().unwrap(), first);
858+
decoder
859+
.seek_to(std::mem::size_of::<MboMsg>() as u64)
860+
.await
861+
.unwrap();
862+
assert_eq!(*decoder.decode::<MboMsg>().await.unwrap().unwrap(), second);
863+
decoder.seek_to(0).await.unwrap();
864+
assert_eq!(*decoder.decode::<MboMsg>().await.unwrap().unwrap(), first);
865+
}
866+
867+
#[tokio::test]
868+
async fn test_seek_to_compressed_returns_bad_argument() {
869+
let reader =
870+
AsyncDynReader::from_file(format!("{TEST_DATA_PATH}/test_data.mbo.v3.dbn.zst"))
871+
.await
872+
.unwrap();
873+
let mut decoder =
874+
RecordDecoder::with_version(reader, DBN_VERSION, VersionUpgradePolicy::AsIs, false)
875+
.unwrap();
876+
877+
assert!(matches!(
878+
decoder.seek_to(0).await.unwrap_err(),
879+
Error::BadArgument { param_name, desc }
880+
if param_name == "self" && desc.contains("zstd-compressed")
881+
));
882+
}
883+
766884
#[tokio::test]
767885
async fn test_dbn_identity_with_ts_out() {
768886
let rec1 = WithTsOut {

rust/dbn/src/decode/dbn/fsm.rs

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -288,6 +288,14 @@ impl DbnFsm {
288288
self.needs_upgrade = Self::compute_needs_upgrade(self.upgrade_policy, None);
289289
}
290290

291+
/// Resets buffered record state after seeking while preserving decoder
292+
/// configuration.
293+
pub(crate) fn reset_for_seek(&mut self) {
294+
self.state = State::Record;
295+
self.buffer.reset();
296+
self.compat_buffer.reset();
297+
}
298+
291299
/// Skips ahead `nbytes`. Returns the actual number of bytes skipped.
292300
///
293301
/// Writable space is not reclaimed until the next call to

rust/dbn/src/decode/dbn/sync.rs

Lines changed: 108 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,8 +8,8 @@ use crate::{
88
decode::{
99
dbn::fsm::{DbnFsm, ProcessResult},
1010
private::LastRecord,
11-
DbnMetadata, DecodeRecord, DecodeRecordRef, DecodeStream, SkipBytes, StreamIterDecoder,
12-
VersionUpgradePolicy,
11+
DbnMetadata, DecodeRecord, DecodeRecordRef, DecodeStream, DynReader, SkipBytes,
12+
StreamIterDecoder, VersionUpgradePolicy,
1313
},
1414
HasRType, Metadata, RecordRef, DBN_VERSION,
1515
};
@@ -368,6 +368,64 @@ where
368368
}
369369
}
370370

371+
impl<R> RecordDecoder<R>
372+
where
373+
R: io::Read + io::Seek,
374+
{
375+
/// Seeks to absolute byte position `pos` and clears buffered decoder state.
376+
///
377+
/// # Warning
378+
/// Callers are responsible for ensuring `pos` is a valid record boundary.
379+
///
380+
/// # Errors
381+
/// This function returns an error if it fails to seek in the inner reader.
382+
pub fn seek_to(&mut self, pos: u64) -> crate::Result<()> {
383+
self.reader
384+
.seek(io::SeekFrom::Start(pos))
385+
.map(drop)
386+
.map_err(|err| crate::Error::io(err, format!("seeking to byte offset {pos}")))?;
387+
self.fsm.reset_for_seek();
388+
Ok(())
389+
}
390+
}
391+
392+
impl<'a, R> RecordDecoder<DynReader<'a, R>>
393+
where
394+
R: io::BufRead,
395+
{
396+
/// Seeks to absolute byte position `pos` and clears buffered decoder state.
397+
///
398+
/// # Warning
399+
/// Callers are responsible for ensuring `pos` is a valid record boundary.
400+
///
401+
/// # Errors
402+
/// This function returns an error if it fails to seek in the inner reader, or if
403+
/// the input is Zstandard-compressed.
404+
pub fn seek_to(&mut self, pos: u64) -> crate::Result<()>
405+
where
406+
R: io::Seek,
407+
{
408+
if self.reader.is_compressed() {
409+
return Err(crate::Error::BadArgument {
410+
param_name: "self".to_owned(),
411+
desc: "absolute seek is unsupported for zstd-compressed input".to_owned(),
412+
});
413+
}
414+
self.reader
415+
.get_mut()
416+
.seek(io::SeekFrom::Start(pos))
417+
.map(drop)
418+
.map_err(|err| crate::Error::io(err, format!("seeking to byte offset {pos}")))?;
419+
self.fsm.reset_for_seek();
420+
Ok(())
421+
}
422+
423+
/// Returns whether the input is Zstandard-compressed.
424+
pub fn is_compressed(&self) -> bool {
425+
self.reader.is_compressed()
426+
}
427+
}
428+
371429
impl<R> LastRecord for RecordDecoder<R>
372430
where
373431
R: io::Read,
@@ -618,6 +676,54 @@ mod tests {
618676
assert!(decoder.decode_record::<MboMsg>().unwrap().is_none());
619677
}
620678

679+
#[test]
680+
fn test_seek_to_resets_buffered_records() {
681+
let mut first = MboMsg {
682+
hd: RecordHeader::new::<MboMsg>(rtype::MBO, 1, 100, 1),
683+
..Default::default()
684+
};
685+
first.order_id = 1;
686+
let mut second = MboMsg {
687+
hd: RecordHeader::new::<MboMsg>(rtype::MBO, 1, 101, 2),
688+
..Default::default()
689+
};
690+
second.order_id = 2;
691+
let mut buffer = Vec::new();
692+
buffer.extend_from_slice(first.as_ref());
693+
buffer.extend_from_slice(second.as_ref());
694+
695+
let mut decoder = RecordDecoder::with_version(
696+
std::io::Cursor::new(buffer),
697+
DBN_VERSION,
698+
VersionUpgradePolicy::AsIs,
699+
false,
700+
)
701+
.unwrap();
702+
703+
assert_eq!(*decoder.decode::<MboMsg>().unwrap().unwrap(), first);
704+
decoder
705+
.seek_to(std::mem::size_of::<MboMsg>() as u64)
706+
.unwrap();
707+
assert_eq!(*decoder.decode::<MboMsg>().unwrap().unwrap(), second);
708+
decoder.seek_to(0).unwrap();
709+
assert_eq!(*decoder.decode::<MboMsg>().unwrap().unwrap(), first);
710+
}
711+
712+
#[test]
713+
fn test_seek_to_compressed_returns_bad_argument() {
714+
let reader =
715+
DynReader::from_file(format!("{TEST_DATA_PATH}/test_data.mbo.v3.dbn.zst")).unwrap();
716+
let mut decoder =
717+
RecordDecoder::with_version(reader, DBN_VERSION, VersionUpgradePolicy::AsIs, false)
718+
.unwrap();
719+
720+
assert!(matches!(
721+
decoder.seek_to(0).unwrap_err(),
722+
Error::BadArgument { param_name, desc }
723+
if param_name == "self" && desc.contains("zstd-compressed")
724+
));
725+
}
726+
621727
#[test]
622728
fn test_dbn_identity_with_ts_out() -> Result<()> {
623729
let rec1 = WithTsOut {

rust/dbn/src/decode/dyn_reader.rs

Lines changed: 29 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -101,6 +101,11 @@ where
101101
DynReaderImpl::Zstd(reader) => reader.get_ref(),
102102
}
103103
}
104+
105+
/// Returns whether the input is Zstandard-compressed.
106+
pub fn is_compressed(&self) -> bool {
107+
matches!(&self.0, DynReaderImpl::Zstd(_))
108+
}
104109
}
105110

106111
impl DynReader<'_, BufReader<File>> {
@@ -180,6 +185,8 @@ mod tests {
180185
DynReader::from_file(format!("{TEST_DATA_PATH}/test_data.mbo.v3.dbn")).unwrap();
181186
let mut compressed =
182187
DynReader::from_file(format!("{TEST_DATA_PATH}/test_data.mbo.v3.dbn.zst")).unwrap();
188+
assert!(!uncompressed.is_compressed());
189+
assert!(compressed.is_compressed());
183190
let mut uncompressed_res = Vec::new();
184191
uncompressed.read_to_end(&mut uncompressed_res).unwrap();
185192
let mut compressed_res = Vec::new();
@@ -300,6 +307,11 @@ mod r#async {
300307
DynReaderImpl::Zstd(reader) => reader.get_ref(),
301308
}
302309
}
310+
311+
/// Returns whether the input is Zstandard-compressed.
312+
pub fn is_compressed(&self) -> bool {
313+
matches!(&self.0, DynReaderImpl::Zstd(_))
314+
}
303315
}
304316

305317
impl DynReader<BufReader<File>> {
@@ -385,6 +397,7 @@ mod r#async {
385397
}
386398
}
387399
}
400+
388401
#[cfg(test)]
389402
mod tests {
390403
use crate::{
@@ -398,7 +411,7 @@ mod r#async {
398411
#[tokio::test]
399412
async fn test_decode_multiframe_zst() {
400413
let mut decoder = AsyncDbnRecordDecoder::with_version(
401-
DynReader::from_file(&format!(
414+
DynReader::from_file(format!(
402415
"{TEST_DATA_PATH}/multi-frame.definition.v1.dbn.frag.zst"
403416
))
404417
.await
@@ -414,5 +427,20 @@ mod r#async {
414427
}
415428
assert_eq!(count, 8);
416429
}
430+
431+
#[tokio::test]
432+
async fn test_dyn_reader_is_compressed() {
433+
let uncompressed =
434+
DynReader::from_file(format!("{TEST_DATA_PATH}/test_data.mbo.v3.dbn"))
435+
.await
436+
.unwrap();
437+
let compressed =
438+
DynReader::from_file(format!("{TEST_DATA_PATH}/test_data.mbo.v3.dbn.zst"))
439+
.await
440+
.unwrap();
441+
442+
assert!(!uncompressed.is_compressed());
443+
assert!(compressed.is_compressed());
444+
}
417445
}
418446
}

0 commit comments

Comments
 (0)