Skip to main content

reifydb_sub_flow/operator/
take.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2026 ReifyDB
3
4use 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}