1use super::{
2 ComparisonOp, FilterExpr, PointId, PointIdPredicate, PointSelector, Prefetch, PrefetchSource,
3 QueryExpr, QueryStmt, Stmt, Value,
4};
5use crate::error::QqlError;
6use alloc::boxed::Box;
7use alloc::string::{String, ToString};
8
9impl Stmt {
10 pub fn shard_key(&self) -> Option<&str> {
15 match self {
16 Self::Query(query) => query.shard_key.as_deref(),
17 Self::Scroll(scroll) => scroll.shard_key.as_deref(),
18 Self::Count(count) => count.shard_key.as_deref(),
19 Self::Upsert(upsert) => upsert.shard_key.as_deref(),
20 Self::Delete(delete) => delete.shard_key.as_deref(),
21 Self::ClearPayload(clear) => clear.shard_key.as_deref(),
22 Self::DeletePayload(delete) => delete.shard_key.as_deref(),
23 Self::DeleteVector(delete) => delete.shard_key.as_deref(),
24 Self::UpdateVector(update) => update.shard_key.as_deref(),
25 Self::UpdatePayload(update) => update.shard_key.as_deref(),
26 _ => None,
27 }
28 }
29
30 pub fn set_shard_key(&mut self, shard_key: Option<String>) -> bool {
40 let key = shard_key.filter(|k| !k.is_empty());
41 match self {
42 Self::Query(query) => {
43 apply_query_shard(query, key.as_deref());
44 true
45 }
46 Self::Scroll(scroll) => {
47 scroll.shard_key = key;
48 true
49 }
50 Self::Count(count) => {
51 count.shard_key = key;
52 true
53 }
54 Self::Upsert(upsert) => {
55 upsert.shard_key = key;
56 true
57 }
58 Self::Delete(delete) => {
59 delete.shard_key = key;
60 true
61 }
62 Self::ClearPayload(clear) => {
63 clear.shard_key = key;
64 true
65 }
66 Self::DeletePayload(delete) => {
67 delete.shard_key = key;
68 true
69 }
70 Self::DeleteVector(delete) => {
71 delete.shard_key = key;
72 true
73 }
74 Self::UpdateVector(update) => {
75 update.shard_key = key;
76 true
77 }
78 Self::UpdatePayload(update) => {
79 update.shard_key = key;
80 true
81 }
82 _ => false,
83 }
84 }
85}
86
87fn apply_query_shard(query: &mut QueryStmt, key: Option<&str>) {
89 query.shard_key = key.map(str::to_string);
90 for cte in &mut query.ctes {
91 apply_query_shard(&mut cte.query, key);
92 }
93 if let Some(prefetches) = expression_prefetch(&mut query.expression) {
94 for prefetch in prefetches {
95 if let PrefetchSource::Query(nested) = &mut prefetch.source {
96 apply_query_shard(nested, key);
97 }
98 }
99 }
100}
101
102pub fn inject_filter(
103 statement: &mut Stmt,
104 field: &str,
105 operator: ComparisonOp,
106 value: Value,
107) -> Result<(), QqlError> {
108 let filter = build_filter(field, operator, value.clone())?;
109 match statement {
110 Stmt::Query(query) => inject_query(query, &filter),
111 Stmt::Scroll(scroll) => merge_filter(&mut scroll.filter, filter),
112 Stmt::Delete(delete) => merge_selector(&mut delete.selector, filter),
113 Stmt::Count(count) => merge_filter(&mut count.filter, filter),
114 Stmt::ClearPayload(clear) => merge_selector(&mut clear.selector, filter),
115 Stmt::DeletePayload(del) => merge_selector(&mut del.selector, filter),
116 Stmt::DeleteVector(del_vec) => merge_selector(&mut del_vec.selector, filter),
117 Stmt::UpdatePayload(update) => merge_selector(&mut update.selector, filter),
118 Stmt::Upsert(upsert)
119 if operator == ComparisonOp::Eq && !field.eq_ignore_ascii_case("id") =>
120 {
121 for point in &mut upsert.points {
122 if let Some((_, current)) = point
123 .payload
124 .iter_mut()
125 .find(|(key, _)| key.eq_ignore_ascii_case(field))
126 {
127 *current = value.clone();
128 } else {
129 point.payload.push((field.to_string(), value.clone()));
130 }
131 }
132 }
133 other => {
134 return Err(QqlError::validation(
135 "QQL-VALIDATION-FILTER-INJECT",
136 format!(
137 "inject_filter does not apply to this statement type ({})",
138 stmt_kind(other)
139 ),
140 None,
141 ));
142 }
143 }
144 Ok(())
145}
146
147fn stmt_kind(statement: &Stmt) -> &'static str {
148 match statement {
149 Stmt::Query(_) => "QUERY",
150 Stmt::Scroll(_) => "SCROLL",
151 Stmt::Count(_) => "COUNT",
152 Stmt::Upsert(_) => "UPSERT",
153 Stmt::Delete(_) => "DELETE",
154 Stmt::ClearPayload(_) => "CLEAR PAYLOAD",
155 Stmt::DeletePayload(_) => "DELETE PAYLOAD",
156 Stmt::DeleteVector(_) => "DELETE VECTOR",
157 Stmt::UpdateVector(_) => "UPDATE VECTOR",
158 Stmt::UpdatePayload(_) => "UPDATE PAYLOAD",
159 Stmt::CreateCollection(_) => "CREATE COLLECTION",
160 Stmt::AlterCollection(_) => "ALTER COLLECTION",
161 Stmt::DropCollection(_) => "DROP COLLECTION",
162 Stmt::CreateIndex(_) => "CREATE INDEX",
163 Stmt::DropIndex(_) => "DROP INDEX",
164 Stmt::CreateShardKey(_) => "CREATE SHARD KEY",
165 Stmt::DropShardKey(_) => "DROP SHARD KEY",
166 Stmt::ShowCollections => "SHOW COLLECTIONS",
167 Stmt::ShowCollection(_) => "SHOW COLLECTION",
168 Stmt::ShowShardKeys(_) => "SHOW SHARD KEYS",
169 }
170}
171
172fn build_filter(field: &str, operator: ComparisonOp, value: Value) -> Result<FilterExpr, QqlError> {
173 if field.eq_ignore_ascii_case("id") {
174 if operator != ComparisonOp::Eq {
175 return Err(QqlError::validation(
176 "QQL-VALIDATION-ID-PREDICATE",
177 "point ID injection supports equality only",
178 None,
179 ));
180 }
181 let id = match value {
182 Value::Int(value) if value >= 0 => PointId::Number(value as u64),
183 Value::Str(value) => PointId::String(value),
184 _ => {
185 return Err(QqlError::validation(
186 "QQL-VALIDATION-POINT-ID",
187 "point IDs must be unsigned integers or strings",
188 None,
189 ));
190 }
191 };
192 Ok(FilterExpr::PointId(PointIdPredicate::Eq(id)))
193 } else {
194 Ok(FilterExpr::Compare {
195 field: field.to_string(),
196 op: operator,
197 value,
198 })
199 }
200}
201
202fn inject_query(query: &mut QueryStmt, filter: &FilterExpr) {
203 merge_filter(&mut query.filter, filter.clone());
204 for cte in &mut query.ctes {
205 inject_query(&mut cte.query, filter);
206 }
207 if let Some(prefetches) = expression_prefetch(&mut query.expression) {
208 for prefetch in prefetches {
209 merge_filter(&mut prefetch.filter, filter.clone());
210 if let PrefetchSource::Query(query) = &mut prefetch.source {
211 inject_query(query, filter);
212 }
213 }
214 }
215}
216
217fn expression_prefetch(expression: &mut QueryExpr) -> Option<&mut Vec<Prefetch>> {
218 match expression {
219 QueryExpr::Nearest { prefetch, .. }
220 | QueryExpr::Recommend { prefetch, .. }
221 | QueryExpr::Context { prefetch, .. }
222 | QueryExpr::Discover { prefetch, .. }
223 | QueryExpr::Fusion { prefetch, .. }
224 | QueryExpr::Formula { prefetch, .. }
225 | QueryExpr::RelevanceFeedback { prefetch, .. }
226 | QueryExpr::Rerank { prefetch, .. }
227 | QueryExpr::CrossRerank { prefetch, .. } => Some(prefetch),
228 QueryExpr::Points { .. }
229 | QueryExpr::OrderBy { .. }
230 | QueryExpr::SampleRandom
231 | QueryExpr::Hybrid { .. } => None,
232 }
233}
234
235fn merge_selector(selector: &mut PointSelector, filter: FilterExpr) {
236 let current =
237 match core::mem::replace(selector, PointSelector::Filter(Box::new(filter.clone()))) {
238 PointSelector::Id(id) => FilterExpr::PointId(PointIdPredicate::Eq(id)),
239 PointSelector::Ids(ids) => FilterExpr::PointId(PointIdPredicate::In(ids)),
240 PointSelector::Filter(filter) => *filter,
241 };
242 *selector = PointSelector::Filter(Box::new(and(current, filter)));
243}
244
245fn merge_filter(current: &mut Option<Box<FilterExpr>>, filter: FilterExpr) {
246 *current = Some(Box::new(match current.take() {
247 Some(current) => and(*current, filter),
248 None => filter,
249 }));
250}
251
252fn and(left: FilterExpr, right: FilterExpr) -> FilterExpr {
253 match left {
254 FilterExpr::And { mut operands } => {
255 operands.push(right);
256 FilterExpr::And { operands }
257 }
258 left => FilterExpr::And {
259 operands: alloc::vec![left, right],
260 },
261 }
262}