laminar_core/lookup/
align.rs1use std::sync::Arc;
12
13use arrow::row::{RowConverter, SortField};
14use arrow_array::{ArrayRef, RecordBatch};
15use rustc_hash::FxHashMap;
16
17use crate::lookup::source::LookupError;
18
19pub struct KeyAligner {
21 converter: RowConverter,
22 pk_columns: Vec<String>,
23}
24
25impl KeyAligner {
26 pub fn new(
34 pk_sort_fields: Vec<SortField>,
35 pk_columns: Vec<String>,
36 ) -> Result<Self, LookupError> {
37 if pk_columns.is_empty() {
38 return Err(LookupError::Internal(
39 "primary_key_columns must not be empty".into(),
40 ));
41 }
42 let converter = RowConverter::new(pk_sort_fields)
43 .map_err(|e| LookupError::Internal(format!("row converter: {e}")))?;
44 Ok(Self {
45 converter,
46 pk_columns,
47 })
48 }
49
50 #[must_use]
52 pub fn pk_columns(&self) -> &[String] {
53 &self.pk_columns
54 }
55
56 pub fn decode_keys(&self, keys: &[&[u8]]) -> Result<Vec<ArrayRef>, LookupError> {
64 let parser = self.converter.parser();
65 let parsed = keys.iter().map(|k| parser.parse(k));
66 self.converter
67 .convert_rows(parsed)
68 .map_err(|e| LookupError::Internal(format!("decode keys: {e}")))
69 }
70
71 pub fn align(
81 &self,
82 keys: &[&[u8]],
83 fetched: &[RecordBatch],
84 ) -> Result<Vec<Option<RecordBatch>>, LookupError> {
85 let mut index: FxHashMap<Vec<u8>, (usize, usize)> = FxHashMap::default();
86 for (batch_idx, batch) in fetched.iter().enumerate() {
87 if batch.num_rows() == 0 {
88 continue;
89 }
90 let pk_cols = self
91 .pk_columns
92 .iter()
93 .map(|name| {
94 let idx = batch.schema().index_of(name).map_err(|_| {
95 LookupError::Internal(format!("pk column not found in result: {name}"))
96 })?;
97 Ok(Arc::clone(batch.column(idx)))
98 })
99 .collect::<Result<Vec<ArrayRef>, LookupError>>()?;
100 let rows = self
101 .converter
102 .convert_columns(&pk_cols)
103 .map_err(|e| LookupError::Internal(format!("encode result keys: {e}")))?;
104 for row in 0..batch.num_rows() {
105 if index
106 .insert(rows.row(row).as_ref().to_vec(), (batch_idx, row))
107 .is_some()
108 {
109 return Err(LookupError::Internal(
110 "lookup source returned multiple rows for one key".into(),
111 ));
112 }
113 }
114 }
115 Ok(keys
116 .iter()
117 .map(|key| index.get(*key).map(|&(bi, row)| fetched[bi].slice(row, 1)))
118 .collect())
119 }
120}
121
122#[cfg(test)]
123mod tests {
124 use super::*;
125 use arrow_array::Int64Array;
126 use arrow_schema::{DataType, Field, Schema};
127
128 fn aligner() -> KeyAligner {
129 KeyAligner::new(vec![SortField::new(DataType::Int64)], vec!["id".into()]).unwrap()
130 }
131
132 fn batch(ids: &[i64]) -> RecordBatch {
133 let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int64, false)]));
134 RecordBatch::try_new(schema, vec![Arc::new(Int64Array::from(ids.to_vec()))]).unwrap()
135 }
136
137 fn encode(ids: &[i64]) -> Vec<Vec<u8>> {
138 let conv = RowConverter::new(vec![SortField::new(DataType::Int64)]).unwrap();
139 let rows = conv
140 .convert_columns(&[Arc::new(Int64Array::from(ids.to_vec()))])
141 .unwrap();
142 (0..ids.len())
143 .map(|i| rows.row(i).as_ref().to_vec())
144 .collect()
145 }
146
147 #[test]
148 fn aligns_out_of_order_with_misses_and_dups() {
149 let aligner = aligner();
150 let fetched = vec![batch(&[2, 5])];
153 let keys = encode(&[5, 2, 99, 2]);
154 let key_refs: Vec<&[u8]> = keys.iter().map(Vec::as_slice).collect();
155
156 let out = aligner.align(&key_refs, &fetched).unwrap();
157 let id = |b: &Option<RecordBatch>| {
158 b.as_ref().map(|b| {
159 b.column(0)
160 .as_any()
161 .downcast_ref::<Int64Array>()
162 .unwrap()
163 .value(0)
164 })
165 };
166 assert_eq!(id(&out[0]), Some(5));
167 assert_eq!(id(&out[1]), Some(2));
168 assert_eq!(id(&out[2]), None); assert_eq!(id(&out[3]), Some(2)); }
171
172 #[test]
173 fn decode_round_trips_to_pk_columns() {
174 let aligner = aligner();
175 let keys = encode(&[7, 8]);
176 let key_refs: Vec<&[u8]> = keys.iter().map(Vec::as_slice).collect();
177 let cols = aligner.decode_keys(&key_refs).unwrap();
178 let ids = cols[0].as_any().downcast_ref::<Int64Array>().unwrap();
179 assert_eq!(ids.values(), &[7, 8]);
180 }
181
182 #[test]
183 fn rejects_duplicate_fetched_keys_instead_of_choosing_a_row() {
184 let aligner = aligner();
185 let keys = encode(&[2]);
186 let error = aligner
187 .align(&[keys[0].as_slice()], &[batch(&[2, 2])])
188 .unwrap_err();
189 assert!(error.to_string().contains("multiple rows"), "{error}");
190 }
191}