@@ -40,8 +40,6 @@ pub(crate) trait SpillCodec: Sized {
4040
4141pub ( crate ) struct SpillVec < ' on_flush , T : SpillCodec > {
4242 writer : Result < WriteState < ' on_flush , T > , DatabaseError > ,
43- max_rows : usize ,
44- max_bytes : usize ,
4543}
4644
4745pub ( crate ) struct SpillReader < T : SpillCodec > {
@@ -53,6 +51,8 @@ struct WriteState<'on_flush, T: SpillCodec> {
5351 buffer_bytes : usize ,
5452 file : Option < SpillFileWriter > ,
5553 on_flush : Option < OnFlush < ' on_flush , T > > ,
54+ max_rows : usize ,
55+ max_bytes : usize ,
5656}
5757
5858type OnFlush < ' on_flush , T > = Box < dyn FnMut ( & mut Vec < T > ) -> Result < ( ) , DatabaseError > + ' on_flush > ;
@@ -64,8 +64,7 @@ enum ReadState<T: SpillCodec> {
6464 tail : std:: vec:: IntoIter < T > ,
6565 _file_guard : SpillFileGuard ,
6666 } ,
67- Failed ( DatabaseError ) ,
68- Exhausted ,
67+ Exhausted ( Option < DatabaseError > ) ,
6968}
7069
7170impl < ' on_flush , T : SpillCodec > SpillVec < ' on_flush , T > {
@@ -76,18 +75,20 @@ impl<'on_flush, T: SpillCodec> SpillVec<'on_flush, T> {
7675 buffer_bytes : 0 ,
7776 file : None ,
7877 on_flush : None ,
78+ max_rows : DEFAULT_MAX_ROWS ,
79+ max_bytes : DEFAULT_MAX_BYTES ,
7980 } ) ,
80- max_rows : DEFAULT_MAX_ROWS ,
81- max_bytes : DEFAULT_MAX_BYTES ,
8281 }
8382 }
8483
8584 #[ cfg( test) ]
8685 pub ( crate ) fn limit ( mut self , max_rows : usize , max_bytes : usize ) -> Self {
8786 assert ! ( max_rows > 0 , "spill row limit must be positive" ) ;
8887 assert ! ( max_bytes > 0 , "spill byte limit must be positive" ) ;
89- self . max_rows = max_rows;
90- self . max_bytes = max_bytes;
88+ if let Ok ( state) = & mut self . writer {
89+ state. max_rows = max_rows;
90+ state. max_bytes = max_bytes;
91+ }
9192 self
9293 }
9394
@@ -105,7 +106,7 @@ impl<'on_flush, T: SpillCodec> SpillVec<'on_flush, T> {
105106 let state = self . writer . as_mut ( ) . map_err ( |_| {
106107 DatabaseError :: InvalidValue ( "cannot append to a failed SpillVec" . to_string ( ) )
107108 } ) ?;
108- state. push ( value, self . max_rows , self . max_bytes )
109+ state. push ( value)
109110 }
110111
111112 pub ( crate ) fn is_spilled ( & self ) -> bool {
@@ -143,8 +144,10 @@ impl<T: SpillCodec> IntoIterator for SpillVec<'_, T> {
143144
144145 fn into_iter ( self ) -> Self :: IntoIter {
145146 let state = match self . writer {
146- Ok ( writer) => writer. into_read ( ) . unwrap_or_else ( ReadState :: Failed ) ,
147- Err ( error) => ReadState :: Failed ( error) ,
147+ Ok ( writer) => writer
148+ . into_read ( )
149+ . unwrap_or_else ( |error| ReadState :: Exhausted ( Some ( error) ) ) ,
150+ Err ( error) => ReadState :: Exhausted ( Some ( error) ) ,
148151 } ;
149152 SpillReader { state }
150153 }
@@ -154,14 +157,6 @@ impl<T: SpillCodec> Iterator for SpillReader<T> {
154157 type Item = Result < T , DatabaseError > ;
155158
156159 fn next ( & mut self ) -> Option < Self :: Item > {
157- if matches ! ( self . state, ReadState :: Failed ( _) ) {
158- let ReadState :: Failed ( error) = std:: mem:: replace ( & mut self . state , ReadState :: Exhausted )
159- else {
160- unreachable ! ( )
161- } ;
162- return Some ( Err ( error) ) ;
163- }
164-
165160 let result = match & mut self . state {
166161 ReadState :: Memory ( rows) => Ok ( rows. next ( ) ) ,
167162 ReadState :: Spilled { reader, tail, .. } => loop {
@@ -175,17 +170,16 @@ impl<T: SpillCodec> Iterator for SpillReader<T> {
175170 } ,
176171 }
177172 } ,
178- ReadState :: Exhausted => return None ,
179- ReadState :: Failed ( _) => unreachable ! ( ) ,
173+ ReadState :: Exhausted ( error) => return error. take ( ) . map ( Err ) ,
180174 } ;
181175 match result {
182176 Ok ( Some ( value) ) => Some ( Ok ( value) ) ,
183177 Ok ( None ) => {
184- self . state = ReadState :: Exhausted ;
178+ self . state = ReadState :: Exhausted ( None ) ;
185179 None
186180 }
187181 Err ( error) => {
188- self . state = ReadState :: Exhausted ;
182+ self . state = ReadState :: Exhausted ( None ) ;
189183 Some ( Err ( error) )
190184 }
191185 }
@@ -281,17 +275,12 @@ impl<R: Read, T: SpillCodec> Iterator for SegmentReader<'_, R, T> {
281275}
282276
283277impl < T : SpillCodec > WriteState < ' _ , T > {
284- fn push (
285- & mut self ,
286- value : T ,
287- max_rows : usize ,
288- max_bytes : usize ,
289- ) -> Result < Option < SegmentOffset > , DatabaseError > {
278+ fn push ( & mut self , value : T ) -> Result < Option < SegmentOffset > , DatabaseError > {
290279 let value_size = value. estimated_size ( ) ;
291280 self . buffer . push ( value) ;
292281 self . buffer_bytes = self . buffer_bytes . saturating_add ( value_size) ;
293282
294- if self . buffer . len ( ) >= max_rows || self . buffer_bytes >= max_bytes {
283+ if self . buffer . len ( ) >= self . max_rows || self . buffer_bytes >= self . max_bytes {
295284 self . start_spilling ( ) ?;
296285 return self . flush ( ) ;
297286 }
@@ -379,10 +368,10 @@ impl SpillFileWriter {
379368 tail : std:: vec:: IntoIter < T > ,
380369 ) -> Result < ReadState < T > , DatabaseError > {
381370 self . file . flush ( ) ?;
382- let file = File :: open ( & self . file_guard . path ) ?;
371+ self . file . seek ( SeekFrom :: Start ( 0 ) ) ?;
383372 // Flushed segments are always a prefix; the in-memory buffer is its ordered tail.
384373 Ok ( ReadState :: Spilled {
385- reader : SegmentReader :: new ( file) ,
374+ reader : SegmentReader :: new ( self . file ) ,
386375 tail,
387376 _file_guard : self . file_guard ,
388377 } )
@@ -490,8 +479,6 @@ mod tests {
490479 fn spill_vec_rejects_operations_after_failure ( ) {
491480 let mut failed = SpillVec {
492481 writer : Err ( DatabaseError :: InvalidValue ( "failed spill" . to_string ( ) ) ) ,
493- max_rows : 1 ,
494- max_bytes : 1 ,
495482 } ;
496483
497484 assert ! ( matches!(
@@ -557,6 +544,8 @@ mod tests {
557544 buffer_bytes : 0 ,
558545 file : None ,
559546 on_flush : None ,
547+ max_rows : DEFAULT_MAX_ROWS ,
548+ max_bytes : DEFAULT_MAX_BYTES ,
560549 } ;
561550 assert_eq ! ( empty. flush( ) ?, None ) ;
562551
@@ -565,6 +554,8 @@ mod tests {
565554 buffer_bytes : 0 ,
566555 file : None ,
567556 on_flush : None ,
557+ max_rows : DEFAULT_MAX_ROWS ,
558+ max_bytes : DEFAULT_MAX_BYTES ,
568559 } ;
569560 assert ! ( matches!(
570561 missing_file. flush( ) ,
0 commit comments