sort_governor/engine/
session.rs1use std::path::{
6 Path,
7 PathBuf,
8};
9
10use futures_util::stream;
11use serde::Serialize;
12use serde::de::DeserializeOwned;
13
14use crate::engine::cleanup::AsyncCleanupGuard;
15use crate::engine::merge::{
16 ValueStream,
17 cascade_and_stream,
18 merge_group_to_file,
19};
20use crate::engine::row::RunRow;
21use crate::engine::run::RunWriter;
22use crate::error::SorterError;
23use crate::plan::SortPlan;
24
25#[cfg(test)]
26mod tests;
27
28pub struct SortSession<K, V> {
30 plan: SortPlan,
31 dir: PathBuf,
32 dedup: bool,
33 buffer: Vec<RunRow<K, V>>,
34 buffer_bytes: usize,
35 spills: [Option<PathBuf>; u64::BITS as usize],
38 next_run: u64,
39 failed: bool,
40 hold: Option<Box<dyn Send>>,
41 cleanup: AsyncCleanupGuard,
45}
46
47impl<K, V> SortSession<K, V>
48where
49 K: Ord + Clone + Serialize + DeserializeOwned + Send + 'static,
50 V: Serialize + DeserializeOwned + Send + 'static,
51{
52 #[must_use]
56 pub fn new(plan: SortPlan, dir: PathBuf, dedup: bool) -> Self {
57 Self {
58 plan,
59 dir,
60 dedup,
61 buffer: Vec::new(),
62 buffer_bytes: 0,
63 spills: std::array::from_fn(|_| None),
64 next_run: 0,
65 failed: false,
66 hold: None,
67 cleanup: AsyncCleanupGuard::disarmed(),
68 }
69 }
70
71 #[must_use]
75 pub fn with_temp_dir(plan: SortPlan, temp_root: &Path, dedup: bool) -> Self {
76 let dir = temp_root.join(format!("sort-{}", uuid::Uuid::new_v4()));
77 Self::new(plan, dir, dedup)
78 }
79
80 #[must_use]
84 pub fn hold_resource(mut self, resource: Box<dyn Send>) -> Self {
85 self.hold = Some(resource);
86 self
87 }
88
89 #[must_use]
92 pub fn scratch_dir(&self) -> &Path {
93 &self.dir
94 }
95
96 pub async fn push_with_size(
105 &mut self,
106 key: K,
107 value: V,
108 estimated_bytes: usize,
109 ) -> Result<(), SorterError> {
110 self.ensure_usable()?;
111 let bytes = estimated_bytes.max(1);
112 if let SortPlan::External {
113 run_buffer_bytes, ..
114 } = self.plan
115 && !self.buffer.is_empty()
116 && self.buffer_bytes.saturating_add(bytes) > run_buffer_bytes
117 {
118 self.spill().await?;
119 }
120 self.buffer.push(RunRow::new(key, value));
121 self.buffer_bytes = self.buffer_bytes.saturating_add(bytes);
122 Ok(())
123 }
124
125 pub async fn push(&mut self, key: K, value: V) -> Result<(), SorterError> {
131 let estimate = std::mem::size_of::<(K, V)>().max(1);
132 self.push_with_size(key, value, estimate).await
133 }
134
135 pub async fn finish(mut self) -> Result<ValueStream<V>, SorterError> {
141 self.ensure_usable()?;
142 self.buffer.sort_by(|left, right| left.key.cmp(&right.key));
143 let fan_in = match self.plan {
144 SortPlan::InMemory => return Ok(self.in_memory_stream()),
145 SortPlan::External { .. } if self.next_run == 0 => {
146 return Ok(self.in_memory_stream());
147 }
148 SortPlan::External { max_fan_in, .. } => max_fan_in as usize,
149 };
150 if !self.buffer.is_empty() {
151 self.spill().await?;
152 }
153 let hold = self.hold.take();
154 let guard = std::mem::take(&mut self.cleanup);
155 let spills = self
158 .spills
159 .iter_mut()
160 .rev()
161 .filter_map(Option::take)
162 .collect();
163 let dir = self.dir.clone();
164 cascade_and_stream::<K, V>(spills, fan_in.max(2), dir, guard, self.dedup, hold).await
165 }
166
167 fn in_memory_stream(&mut self) -> ValueStream<V> {
172 let rows = std::mem::take(&mut self.buffer);
173 let mut values = Vec::with_capacity(rows.len());
174 let mut last: Option<K> = None;
175 for row in rows {
176 if self.dedup {
177 if last.as_ref() == Some(&row.key) {
178 continue;
179 }
180 last = Some(row.key.clone());
181 }
182 values.push(row.value);
183 }
184 let hold = self.hold.take();
185 Box::pin(stream::unfold(
186 (values.into_iter(), hold),
187 |(mut values, hold)| async move { values.next().map(|value| (Ok(value), (values, hold))) },
188 ))
189 }
190
191 async fn spill(&mut self) -> Result<(), SorterError> {
193 if self.buffer.is_empty() {
194 return Ok(());
195 }
196 self.failed = true;
199 let ordinal = self.next_run;
200 self.next_run = ordinal.checked_add(1).ok_or(SorterError::RunLimit)?;
201 self.buffer.sort_by(|left, right| left.key.cmp(&right.key));
202 async_fs_io::ensure_dir(&self.dir).await?;
203 self.cleanup.arm(self.dir.clone());
204 let path = self.dir.join(format!("run-{ordinal:020}.cbor"));
205 let mut writer = RunWriter::create(path).await?;
206 for row in &self.buffer {
207 writer.write_row(row).await?;
208 }
209 let mut run = writer.finish().await?;
210 self.buffer.clear();
211 self.buffer_bytes = 0;
212 for (level, slot) in self.spills.iter_mut().enumerate() {
213 let Some(older) = slot.take() else {
214 *slot = Some(run);
215 self.failed = false;
216 return Ok(());
217 };
218 let output = self
219 .dir
220 .join(format!("carry-{ordinal:020}-{level:02}.cbor"));
221 run = merge_group_to_file::<K, V>(&[older, run], output, false).await?;
223 }
224 unreachable!("a checked u64 spill count always fits in 64 frontier levels")
225 }
226
227 fn ensure_usable(&self) -> Result<(), SorterError> {
228 if self.failed {
229 return Err(SorterError::SessionFailed);
230 }
231 Ok(())
232 }
233}