Skip to content

Commit 506bd78

Browse files
committed
refactor: simplify spill state handling
1 parent ffd6583 commit 506bd78

1 file changed

Lines changed: 25 additions & 34 deletions

File tree

src/execution/spill/mod.rs

Lines changed: 25 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -40,8 +40,6 @@ pub(crate) trait SpillCodec: Sized {
4040

4141
pub(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

4745
pub(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

5858
type 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

7170
impl<'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

283277
impl<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

Comments
 (0)