1use std::{
5 collections::{BTreeMap, HashMap},
6 slice::from_ref,
7};
8
9use postcard::{from_bytes, to_stdvec};
10use reifydb_abi::operator::capabilities::OperatorCapability;
11use reifydb_codec::encoded::{
12 row::EncodedRow,
13 shape::{RowShape, RowShapeField},
14};
15use reifydb_core::{
16 interface::{
17 catalog::flow::FlowNodeId,
18 change::{Change, Diff},
19 },
20 value::column::columns::Columns,
21};
22use reifydb_value::{
23 Result,
24 error::Error,
25 value::{Value, blob::Blob, row_number::RowNumber},
26};
27use serde::{Deserialize, Serialize};
28
29use crate::{
30 error::FlowStateError,
31 operator::{
32 Operator, OperatorCell,
33 stateful::{raw::RawStatefulOperator, single::SingleStateful, utils},
34 },
35 transaction::{FlowTransaction, slot::PersistFn},
36};
37
38#[derive(Debug, Clone, Serialize, Deserialize, Default)]
39struct TakeState {
40 by_seq: BTreeMap<u64, RowNumber>,
41 by_row: HashMap<RowNumber, (u64, usize)>,
42 candidates_by_seq: BTreeMap<u64, RowNumber>,
43 candidates_by_row: HashMap<RowNumber, (u64, usize)>,
44 next_seq: u64,
45 row_data: HashMap<RowNumber, EncodedRow>,
46}
47
48pub struct TakeOperator {
49 parent: OperatorCell,
50 node: FlowNodeId,
51 limit: usize,
52 shape: RowShape,
53}
54
55fn row_shape_from_columns(cols: &Columns) -> RowShape {
56 let fields: Vec<RowShapeField> = cols
57 .names
58 .iter()
59 .zip(cols.columns.iter())
60 .map(|(name, buf)| RowShapeField::unconstrained(name.text().to_string(), buf.get_type()))
61 .collect();
62 RowShape::new(fields)
63}
64
65fn encode_take_row(shape: &RowShape, columns: &Columns, row_idx: usize) -> EncodedRow {
66 let values: Vec<Value> = columns.columns.iter().map(|buf| buf.get_value(row_idx)).collect();
67 let mut encoded = shape.allocate();
68 shape.set_values(&mut encoded, &values);
69 encoded
70}
71
72fn decode_take_row(shape: &RowShape, row_number: RowNumber, encoded: &EncodedRow) -> Columns {
73 Columns::from_encoded_rows(shape, &[row_number], from_ref(encoded))
74}
75
76impl TakeOperator {
77 pub fn new(parent: OperatorCell, node: FlowNodeId, limit: usize) -> Self {
78 Self {
79 parent,
80 node,
81 limit,
82 shape: RowShape::operator_state(),
83 }
84 }
85
86 fn load_take_state(&self, txn: &mut FlowTransaction) -> Result<TakeState> {
87 let state_row = self.load_state(txn)?;
88
89 if state_row.is_empty() || !state_row.is_defined(0) {
90 return Ok(TakeState::default());
91 }
92
93 let blob = self.shape.get_blob(&state_row, 0);
94 if blob.is_empty() {
95 return Ok(TakeState::default());
96 }
97
98 from_bytes(blob.as_ref()).map_err(|e| {
99 Error::from(FlowStateError::Decode {
100 state: "TakeState",
101 cause: e.to_string(),
102 })
103 })
104 }
105
106 #[inline]
107 fn acquire_take_state(&self, txn: &mut FlowTransaction) -> Result<(TakeState, PersistFn)> {
108 let node_id = self.node;
109 let shape_for_persist = self.shape.clone();
110 txn.take_operator_state::<TakeState, _>(node_id, |txn| {
111 let s = self.load_take_state(txn)?;
112 let shape = shape_for_persist.clone();
113 let persist: PersistFn = Box::new(move |txn, value| {
114 let state = value.downcast::<TakeState>().expect("TakeState slot type");
115 let serialized = to_stdvec(&*state).map_err(|e| {
116 Error::from(FlowStateError::Encode {
117 state: "TakeState",
118 cause: e.to_string(),
119 })
120 })?;
121 let blob = Blob::from(serialized);
122 let key = utils::empty_key();
123 let mut row = utils::load_or_create_row(node_id, txn, &key, &shape)?;
124 shape.set_blob(&mut row, 0, &blob);
125 utils::save_row(node_id, txn, &key, row)?;
126 Ok(())
127 });
128 Ok((s, persist))
129 })
130 }
131
132 pub(crate) fn output_schema(&self) -> Option<Columns> {
133 self.parent.output_schema()
134 }
135
136 #[inline]
137 fn prune_candidates(&self, state: &mut TakeState) {
138 let cap = self.limit.saturating_mul(4);
139 while state.candidates_by_seq.len() > cap {
140 let Some((&oldest_seq, &oldest_row)) = state.candidates_by_seq.iter().next() else {
141 break;
142 };
143 state.candidates_by_seq.remove(&oldest_seq);
144 state.candidates_by_row.remove(&oldest_row);
145 state.row_data.remove(&oldest_row);
146 }
147 }
148
149 #[inline]
150 fn promote_one_candidate(&self, state: &mut TakeState, schema: &RowShape, output_diffs: &mut Vec<Diff>) {
151 let Some((&seq, &row_number)) = state.candidates_by_seq.iter().next_back() else {
152 return;
153 };
154 let count = state.candidates_by_row.get(&row_number).map(|(_, c)| *c).unwrap_or(1);
155 state.candidates_by_seq.remove(&seq);
156 state.candidates_by_row.remove(&row_number);
157 state.by_seq.insert(seq, row_number);
158 state.by_row.insert(row_number, (seq, count));
159
160 if let Some(encoded) = state.row_data.get(&row_number) {
161 let cols = decode_take_row(schema, row_number, encoded);
162 if !cols.is_empty() {
163 output_diffs.push(Diff::insert(cols));
164 }
165 }
166 }
167
168 #[inline]
169 fn admit_new_row(
170 &self,
171 state: &mut TakeState,
172 row_number: RowNumber,
173 single_row: Columns,
174 schema: &RowShape,
175 output_diffs: &mut Vec<Diff>,
176 ) {
177 if self.limit == 0 {
178 return;
179 }
180
181 let seq = state.next_seq;
182 state.next_seq += 1;
183 state.row_data.insert(row_number, encode_take_row(schema, &single_row, 0));
184 state.by_seq.insert(seq, row_number);
185 state.by_row.insert(row_number, (seq, 1));
186 output_diffs.push(Diff::insert(single_row));
187
188 if state.by_seq.len() > self.limit {
189 let oldest = state.by_seq.iter().next().map(|(s, r)| (*s, *r));
190 if let Some((oldest_seq, oldest_row)) = oldest {
191 let count = state.by_row.get(&oldest_row).map(|(_, c)| *c).unwrap_or(1);
192 state.by_seq.remove(&oldest_seq);
193 state.by_row.remove(&oldest_row);
194 state.candidates_by_seq.insert(oldest_seq, oldest_row);
195 state.candidates_by_row.insert(oldest_row, (oldest_seq, count));
196 if let Some(encoded) = state.row_data.get(&oldest_row) {
197 let cols = decode_take_row(schema, oldest_row, encoded);
198 if !cols.is_empty() {
199 output_diffs.push(Diff::remove(cols));
200 }
201 }
202 }
203 }
204
205 self.prune_candidates(state);
206 }
207
208 #[inline]
209 fn apply_insert_diff(&self, state: &mut TakeState, post: Columns, output_diffs: &mut Vec<Diff>) {
210 let schema = row_shape_from_columns(&post);
211 let row_count = post.row_count();
212 for row_idx in 0..row_count {
213 let row_number = post.row_numbers[row_idx];
214
215 if let Some(slot) = state.by_row.get_mut(&row_number) {
216 slot.1 += 1;
217 continue;
218 }
219
220 if let Some(slot) = state.candidates_by_row.get_mut(&row_number) {
221 slot.1 += 1;
222 continue;
223 }
224
225 let single = post.extract_by_indices(&[row_idx]);
226 self.admit_new_row(state, row_number, single, &schema, output_diffs);
227 }
228 }
229
230 #[inline]
231 fn apply_update_diff(&self, state: &mut TakeState, pre: Columns, post: Columns, output_diffs: &mut Vec<Diff>) {
232 let schema = row_shape_from_columns(&post);
233 let row_count = post.row_count();
234 let mut update_indices: Vec<usize> = Vec::new();
235
236 for row_idx in 0..row_count {
237 let row_number = post.row_numbers[row_idx];
238
239 if state.by_row.contains_key(&row_number) {
240 update_indices.push(row_idx);
241 state.row_data.insert(row_number, encode_take_row(&schema, &post, row_idx));
242 continue;
243 }
244
245 if state.candidates_by_row.contains_key(&row_number) {
246 state.row_data.insert(row_number, encode_take_row(&schema, &post, row_idx));
247 continue;
248 }
249
250 let single = post.extract_by_indices(&[row_idx]);
251 self.admit_new_row(state, row_number, single, &schema, output_diffs);
252 }
253
254 if !update_indices.is_empty() {
255 output_diffs.push(Diff::update(
256 pre.extract_by_indices(&update_indices),
257 post.extract_by_indices(&update_indices),
258 ));
259 }
260 }
261
262 #[inline]
263 fn apply_remove_diff(&self, state: &mut TakeState, pre: Columns, output_diffs: &mut Vec<Diff>) {
264 let schema = row_shape_from_columns(&pre);
265 let row_count = pre.row_count();
266 for row_idx in 0..row_count {
267 let row_number = pre.row_numbers[row_idx];
268
269 if let Some(slot) = state.by_row.get_mut(&row_number) {
270 if slot.1 > 1 {
271 slot.1 -= 1;
272 continue;
273 }
274 let seq = slot.0;
275 state.by_row.remove(&row_number);
276 state.by_seq.remove(&seq);
277 state.row_data.remove(&row_number);
278 output_diffs.push(Diff::remove(pre.extract_by_indices(&[row_idx])));
279
280 if state.by_seq.len() < self.limit && !state.candidates_by_seq.is_empty() {
281 self.promote_one_candidate(state, &schema, output_diffs);
282 }
283 continue;
284 }
285
286 if let Some(slot) = state.candidates_by_row.get_mut(&row_number) {
287 if slot.1 > 1 {
288 slot.1 -= 1;
289 } else {
290 let seq = slot.0;
291 state.candidates_by_row.remove(&row_number);
292 state.candidates_by_seq.remove(&seq);
293 state.row_data.remove(&row_number);
294 }
295 }
296 }
297 }
298}
299
300impl RawStatefulOperator for TakeOperator {}
301
302impl SingleStateful for TakeOperator {
303 fn layout(&self) -> RowShape {
304 self.shape.clone()
305 }
306}
307
308impl Operator for TakeOperator {
309 fn id(&self) -> FlowNodeId {
310 self.node
311 }
312
313 fn capabilities(&self) -> &[OperatorCapability] {
314 OperatorCapability::STANDARD
315 }
316
317 fn apply(&self, txn: &mut FlowTransaction, change: Change) -> Result<Change> {
318 let node_id = self.node;
319 let (mut state, persist) = self.acquire_take_state(txn)?;
320
321 let mut output_diffs = Vec::new();
322 let version = change.version;
323
324 for diff in change.diffs {
325 match diff {
326 Diff::Insert {
327 post,
328 ..
329 } => self.apply_insert_diff(&mut state, post, &mut output_diffs),
330 Diff::Update {
331 pre,
332 post,
333 ..
334 } => self.apply_update_diff(&mut state, pre, post, &mut output_diffs),
335 Diff::Remove {
336 pre,
337 ..
338 } => self.apply_remove_diff(&mut state, pre, &mut output_diffs),
339 }
340 }
341
342 txn.put_operator_state(node_id, state, persist);
343
344 Ok(Change::from_flow(self.node, version, output_diffs, change.changed_at))
345 }
346}