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 Stmt::ShowQuotas => "SHOW QUOTAS",
170 Stmt::SetQuota(_) => "SET QUOTA",
171 }
172}
173
174fn build_filter(field: &str, operator: ComparisonOp, value: Value) -> Result<FilterExpr, QqlError> {
175 if field.eq_ignore_ascii_case("id") {
176 if operator != ComparisonOp::Eq {
177 return Err(QqlError::validation(
178 "QQL-VALIDATION-ID-PREDICATE",
179 "point ID injection supports equality only",
180 None,
181 ));
182 }
183 let id = match value {
184 Value::Int(value) if value >= 0 => PointId::Number(value as u64),
185 Value::Str(value) => PointId::String(value),
186 _ => {
187 return Err(QqlError::validation(
188 "QQL-VALIDATION-POINT-ID",
189 "point IDs must be unsigned integers or strings",
190 None,
191 ));
192 }
193 };
194 Ok(FilterExpr::PointId(PointIdPredicate::Eq(id)))
195 } else {
196 Ok(FilterExpr::Compare {
197 field: field.to_string(),
198 op: operator,
199 value,
200 })
201 }
202}
203
204fn inject_query(query: &mut QueryStmt, filter: &FilterExpr) {
205 merge_filter(&mut query.filter, filter.clone());
206 for cte in &mut query.ctes {
207 inject_query(&mut cte.query, filter);
208 }
209 if let Some(prefetches) = expression_prefetch(&mut query.expression) {
210 for prefetch in prefetches {
211 merge_filter(&mut prefetch.filter, filter.clone());
212 if let PrefetchSource::Query(query) = &mut prefetch.source {
213 inject_query(query, filter);
214 }
215 }
216 }
217}
218
219fn expression_prefetch(expression: &mut QueryExpr) -> Option<&mut Vec<Prefetch>> {
220 match expression {
221 QueryExpr::Nearest { prefetch, .. }
222 | QueryExpr::Recommend { prefetch, .. }
223 | QueryExpr::Context { prefetch, .. }
224 | QueryExpr::Discover { prefetch, .. }
225 | QueryExpr::Fusion { prefetch, .. }
226 | QueryExpr::Formula { prefetch, .. }
227 | QueryExpr::RelevanceFeedback { prefetch, .. }
228 | QueryExpr::Rerank { prefetch, .. }
229 | QueryExpr::CrossRerank { prefetch, .. } => Some(prefetch),
230 QueryExpr::Points { .. }
231 | QueryExpr::OrderBy { .. }
232 | QueryExpr::SampleRandom
233 | QueryExpr::Hybrid { .. } => None,
234 }
235}
236
237fn merge_selector(selector: &mut PointSelector, filter: FilterExpr) {
238 let current = match core::mem::replace(selector, PointSelector::Ids(Vec::new())) {
239 PointSelector::Id(id) => FilterExpr::PointId(PointIdPredicate::Eq(id)),
240 PointSelector::Ids(ids) => FilterExpr::PointId(PointIdPredicate::In(ids)),
241 PointSelector::Filter(existing) => *existing,
242 };
243 *selector = PointSelector::Filter(Box::new(and(current, filter)));
244}
245
246fn merge_filter(current: &mut Option<Box<FilterExpr>>, filter: FilterExpr) {
247 *current = Some(Box::new(match current.take() {
248 Some(current) => and(*current, filter),
249 None => filter,
250 }));
251}
252
253fn and(left: FilterExpr, right: FilterExpr) -> FilterExpr {
254 match left {
255 FilterExpr::And { mut operands } => {
256 operands.push(right);
257 FilterExpr::And { operands }
258 }
259 left => FilterExpr::And {
260 operands: alloc::vec![left, right],
261 },
262 }
263}