@@ -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+
461528impl < R > RecordDecoder < R >
462529where
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 {
0 commit comments