Skip to main content

qql_core/ast/
transform.rs

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    /// Custom shard routing key for this statement, if any.
11    ///
12    /// Corresponds to QQL `SHARD '…'` on DML, lowered to request-level
13    /// `shard_key` (REST) / `ShardKeySelector` (gRPC) — never inside `Filter`.
14    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    /// Set custom shard routing (same field as QQL `SHARD '…'`).
31    ///
32    /// Prefer writing `SHARD 'tenant'` in the query when the tenant is known at
33    /// authoring time. Use this setter only when the host resolves the key after
34    /// parse (e.g. from auth context) without re-stringifying QQL.
35    ///
36    /// On `QUERY`, recurses into CTEs and nested prefetch queries so routing
37    /// matches a top-level `SHARD` clause. Empty / `None` clears the key.
38    /// Returns `false` for statement types that cannot carry routing (DDL, SHOW).
39    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
87/// Apply shard routing to a query and nested CTE / prefetch queries.
88fn 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}