Skip to main content

qql_core/params/
validate_collect.rs

1//! Single-pass parameter census for AST statements.
2//!
3//! One traversal produces every parameter answer: the named/positional census
4//! (`collect_statement_params`), the whole-point flag (`stmt_has_point_params`),
5//! the full unbound check (`validate_no_unbound_params`), and the template
6//! check (`validate_no_unbound_scalar_params` with its vector-param flag).
7//! The three historical walkers covered one `Stmt` shape and drifted apart
8//! (missing prefetch arms, DDL gaps, `QueryInput::Vector` handling). This
9//! module visits each node once and folds all answers from that visit.
10
11use crate::ast::Value;
12use crate::ast::filter::{FilterExpr, PointIdPredicate};
13use crate::ast::formula::FormulaExpr;
14use crate::ast::statement::{
15    CollectionConfig, PointId, PointSelector, PointVectors, Prefetch, PrefetchSource, QueryExpr,
16    QueryInput, QueryStmt, ShardKey, Stmt, VectorValue,
17};
18use crate::error::{QqlError, Span};
19
20fn unbound_named_err(name: &str, span: Option<Span>) -> QqlError {
21    QqlError::validation(
22        "QQL-BIND-MISSING-PARAM",
23        alloc::format!("missing value for named parameter ':{}'", name),
24        span,
25    )
26}
27
28fn unbound_positional_err(idx: usize, span: Option<Span>) -> QqlError {
29    QqlError::validation(
30        "QQL-BIND-MISSING-PARAM",
31        alloc::format!("missing value for positional parameter '?{}'", idx),
32        span,
33    )
34}
35
36fn unbound_param_str_err(param: &str, span: Option<Span>) -> QqlError {
37    if let Some(name) = param.strip_prefix(':') {
38        unbound_named_err(name, span)
39    } else if let Some(idx_str) = param.strip_prefix('?') {
40        if let Ok(idx) = idx_str.parse::<usize>() {
41            unbound_positional_err(idx, span)
42        } else {
43            QqlError::validation(
44                "QQL-BIND-INVALID-PARAMS",
45                alloc::format!("invalid positional parameter index '?{}'", idx_str),
46                span,
47            )
48        }
49    } else {
50        unbound_named_err(param, span)
51    }
52}
53
54/// Every parameter answer for one statement, gathered in a single walk.
55pub(crate) struct Census {
56    pub(crate) named: alloc::collections::BTreeSet<alloc::string::String>,
57    pub(crate) max_pos: usize,
58    pub(crate) has_vec_params: bool,
59    pub(crate) has_point_params: bool,
60    pub(crate) full_err: Option<QqlError>,
61    pub(crate) scalar_err: Option<QqlError>,
62}
63
64impl Census {
65    fn new() -> Self {
66        Self {
67            named: alloc::collections::BTreeSet::new(),
68            max_pos: 0,
69            has_vec_params: false,
70            has_point_params: false,
71            full_err: None,
72            scalar_err: None,
73        }
74    }
75
76    fn record_named(&mut self, name: &str) {
77        self.named.insert(alloc::string::String::from(name));
78    }
79
80    fn record_pos(&mut self, idx: usize) {
81        self.max_pos = self.max_pos.max(idx + 1);
82    }
83
84    fn note_full(&mut self, err: QqlError) {
85        if self.full_err.is_none() {
86            self.full_err = Some(err);
87        }
88    }
89
90    fn note_scalar(&mut self, err: QqlError) {
91        if self.scalar_err.is_none() {
92            self.scalar_err = Some(err);
93        }
94    }
95
96    fn note_both(&mut self, full: QqlError, scalar: QqlError) {
97        self.note_full(full);
98        self.note_scalar(scalar);
99    }
100
101    /// Scalar placeholder (`Value`, point ID, `OPTIONS` entry, page limit).
102    /// Both validators reject it; the census records it.
103    fn scalar_named(&mut self, name: &str, span: Option<Span>) {
104        self.record_named(name);
105        let full = unbound_named_err(name, span);
106        let scalar = unbound_named_err(name, span);
107        self.note_both(full, scalar);
108    }
109
110    fn scalar_pos(&mut self, idx: usize, span: Option<Span>) {
111        self.record_pos(idx);
112        let full = unbound_positional_err(idx, span);
113        let scalar = unbound_positional_err(idx, span);
114        self.note_both(full, scalar);
115    }
116
117    /// Vector-tolerant placeholder (`QueryInput` / `VectorValue` /
118    /// `PointVectors` whole-slot params). Full validation rejects it; the
119    /// template check records the vector flag instead.
120    fn vec_named(&mut self, name: &str, span: Option<Span>) {
121        self.record_named(name);
122        self.has_vec_params = true;
123        self.note_full(unbound_named_err(name, span));
124    }
125
126    fn vec_pos(&mut self, idx: usize, span: Option<Span>) {
127        self.record_pos(idx);
128        self.has_vec_params = true;
129        self.note_full(unbound_positional_err(idx, span));
130    }
131
132    /// Routing-shard placeholder. Full validation rejects it; templates bind
133    /// it later like vector params, so the scalar check skips it.
134    fn shard_named(&mut self, name: &str, span: Option<Span>) {
135        self.record_named(name);
136        self.note_full(unbound_named_err(name, span));
137    }
138
139    fn shard_pos(&mut self, idx: usize, span: Option<Span>) {
140        self.record_pos(idx);
141        self.note_full(unbound_positional_err(idx, span));
142    }
143
144    /// `":name"` / `"?N"` string placeholder (text params, page limits,
145    /// hybrid/cross-rerank query text). Always scalar in both validators.
146    fn param_str(&mut self, param: &str, span: Option<Span>) {
147        if let Some(name) = param.strip_prefix(':') {
148            self.record_named(name);
149        } else if let Some(idx_str) = param.strip_prefix('?') {
150            if let Ok(idx) = idx_str.parse::<usize>() {
151                self.record_pos(idx);
152            } else {
153                self.max_pos = self.max_pos.max(1);
154            }
155        } else {
156            self.record_named(param);
157        }
158        let full = unbound_param_str_err(param, span);
159        let scalar = unbound_param_str_err(param, span);
160        self.note_both(full, scalar);
161    }
162
163    /// Formula `Variable` name. Only `:` / `?` prefixed names are
164    /// placeholders; bare names (`$score`, `rank`) are variables or
165    /// `DEFAULTS` keys and contribute nothing.
166    fn formula_var(&mut self, name: &str) {
167        if let Some(param_name) = name.strip_prefix(':') {
168            self.record_named(param_name);
169            let full = unbound_named_err(param_name, None);
170            let scalar = unbound_named_err(param_name, None);
171            self.note_both(full, scalar);
172        } else if let Some(idx_str) = name.strip_prefix('?') {
173            if let Ok(idx) = idx_str.parse::<usize>() {
174                self.record_pos(idx);
175                let full = unbound_positional_err(idx, None);
176                let scalar = unbound_positional_err(idx, None);
177                self.note_both(full, scalar);
178            } else {
179                self.max_pos = self.max_pos.max(1);
180                let full = unbound_param_str_err(name, None);
181                let scalar = unbound_param_str_err(name, None);
182                self.note_both(full, scalar);
183            }
184        }
185    }
186
187    fn value(&mut self, val: &Value) {
188        match val {
189            Value::Param(name, span) => {
190                self.scalar_named(name, span.as_deref().copied());
191            }
192            Value::PositionalParam(idx, span) => {
193                self.scalar_pos(*idx, span.as_deref().copied());
194            }
195            Value::List(items) => {
196                for item in items {
197                    self.value(item);
198                }
199            }
200            Value::Dict(entries) => {
201                for (_, v) in entries {
202                    self.value(v);
203                }
204            }
205            _ => {}
206        }
207    }
208
209    fn options(&mut self, options: &[(alloc::string::String, Value)]) {
210        for (_, v) in options {
211            self.value(v);
212        }
213    }
214
215    fn point_id(&mut self, id: &PointId) {
216        match id {
217            PointId::Param(name, span) => {
218                self.scalar_named(name, span.as_deref().copied());
219            }
220            PointId::PositionalParam(idx, span) => {
221                self.scalar_pos(*idx, span.as_deref().copied());
222            }
223            _ => {}
224        }
225    }
226
227    fn shard_opt(&mut self, key: &Option<ShardKey>) {
228        match key {
229            Some(ShardKey::Param(name, span)) => {
230                self.shard_named(name, span.as_deref().copied());
231            }
232            Some(ShardKey::PositionalParam(idx, span)) => {
233                self.shard_pos(*idx, span.as_deref().copied());
234            }
235            _ => {}
236        }
237    }
238
239    fn shard_list(&mut self, keys: Option<&alloc::vec::Vec<ShardKey>>) {
240        if let Some(keys) = keys {
241            for key in keys {
242                match key {
243                    ShardKey::Param(name, span) => {
244                        let span = span.as_deref().copied();
245                        self.record_named(name);
246                        let full = unbound_named_err(name, span);
247                        let scalar = unbound_named_err(name, span);
248                        self.note_both(full, scalar);
249                    }
250                    ShardKey::PositionalParam(idx, span) => {
251                        let span = span.as_deref().copied();
252                        self.record_pos(*idx);
253                        let full = unbound_positional_err(*idx, span);
254                        let scalar = unbound_positional_err(*idx, span);
255                        self.note_both(full, scalar);
256                    }
257                    _ => {}
258                }
259            }
260        }
261    }
262
263    fn ddl_options(&mut self, config: &CollectionConfig) {
264        for options in [&config.wal, &config.strict_mode, &config.metadata]
265            .into_iter()
266            .flatten()
267        {
268            for (_, v) in options {
269                self.value(v);
270            }
271        }
272    }
273
274    fn vector_value(&mut self, vec: &VectorValue) {
275        match vec {
276            VectorValue::Param(name, span) => {
277                self.vec_named(name, span.as_deref().copied());
278            }
279            VectorValue::PositionalParam(idx, span) => {
280                self.vec_pos(*idx, span.as_deref().copied());
281            }
282            VectorValue::Document { options, .. } | VectorValue::Image { options, .. } => {
283                self.options(options);
284            }
285            VectorValue::Object {
286                object, options, ..
287            } => {
288                self.value(object);
289                self.options(options);
290            }
291            _ => {}
292        }
293    }
294
295    fn point_vectors(&mut self, pv: &PointVectors) {
296        match pv {
297            PointVectors::Param(name, span) => {
298                self.vec_named(name, span.as_deref().copied());
299            }
300            PointVectors::PositionalParam(idx, span) => {
301                self.vec_pos(*idx, span.as_deref().copied());
302            }
303            PointVectors::Unnamed(v) => self.vector_value(v),
304            PointVectors::Named(list) => {
305                for (_, v) in list {
306                    self.vector_value(v);
307                }
308            }
309        }
310    }
311
312    fn query_input(&mut self, input: &QueryInput) {
313        match input {
314            QueryInput::Param(name, span) => {
315                self.vec_named(name, span.as_deref().copied());
316            }
317            QueryInput::PositionalParam(idx, span) => {
318                self.vec_pos(*idx, span.as_deref().copied());
319            }
320            QueryInput::Vector(vec) => self.vector_value(vec),
321            QueryInput::Point(point) => self.point_id(point),
322            QueryInput::Text {
323                text_param: Some(param),
324                options,
325                ..
326            } => {
327                self.param_str(param, None);
328                self.options(options);
329            }
330            QueryInput::Text { options, .. } => self.options(options),
331            QueryInput::Image { options, .. } => self.options(options),
332            QueryInput::Object {
333                object, options, ..
334            } => {
335                self.value(object);
336                self.options(options);
337            }
338        }
339    }
340
341    fn filter(&mut self, filter: &FilterExpr) {
342        match filter {
343            FilterExpr::PointId(pred) => match pred {
344                PointIdPredicate::Eq(id) => self.point_id(id),
345                PointIdPredicate::In(ids) => {
346                    for id in ids {
347                        self.point_id(id);
348                    }
349                }
350            },
351            FilterExpr::Compare { value, .. } => self.value(value),
352            FilterExpr::Between { low, high, .. } => {
353                self.value(low);
354                self.value(high);
355            }
356            FilterExpr::In { values, .. }
357            | FilterExpr::MatchAny { values, .. }
358            | FilterExpr::MatchExcept { values, .. } => {
359                for v in values {
360                    self.value(v);
361                }
362            }
363            FilterExpr::And { operands }
364            | FilterExpr::Or { operands }
365            | FilterExpr::MinShould { operands, .. } => {
366                for op in operands {
367                    self.filter(op);
368                }
369            }
370            FilterExpr::Not { operand } => self.filter(operand),
371            FilterExpr::Nested { filter, .. } => self.filter(filter),
372            _ => {}
373        }
374    }
375
376    fn formula(&mut self, expr: &FormulaExpr) {
377        match expr {
378            FormulaExpr::Variable { name } => self.formula_var(name),
379            FormulaExpr::Sum { left, right }
380            | FormulaExpr::Sub { left, right }
381            | FormulaExpr::Mul { left, right }
382            | FormulaExpr::Div { left, right, .. }
383            | FormulaExpr::Pow {
384                base: left,
385                exponent: right,
386            } => {
387                self.formula(left);
388                self.formula(right);
389            }
390            FormulaExpr::Neg { operand }
391            | FormulaExpr::Abs { x: operand }
392            | FormulaExpr::Sqrt { x: operand, .. }
393            | FormulaExpr::Log { x: operand, .. }
394            | FormulaExpr::Ln { x: operand, .. }
395            | FormulaExpr::Exp { x: operand }
396            | FormulaExpr::Acosh { x: operand, .. } => self.formula(operand),
397            FormulaExpr::Max { args } | FormulaExpr::Min { args } => {
398                for arg in args {
399                    self.formula(arg);
400                }
401            }
402            FormulaExpr::Decay { x, target, .. } => {
403                self.formula(x);
404                if let Some(t) = target {
405                    self.formula(t);
406                }
407            }
408            FormulaExpr::Case { cond, then_, else_ } => {
409                self.filter(cond);
410                self.formula(then_);
411                self.formula(else_);
412            }
413            FormulaExpr::MatchCondition { values, .. } => {
414                for v in values {
415                    self.value(v);
416                }
417            }
418            _ => {}
419        }
420    }
421
422    fn prefetch(&mut self, prefetch: &Prefetch) {
423        if let PrefetchSource::Query(q) = &prefetch.source {
424            self.query_stmt(q);
425        }
426        if let Some(f) = &prefetch.filter {
427            self.filter(f);
428        }
429        if let Some(spec) = &prefetch.lookup {
430            self.shard_opt(&spec.shard_key);
431        }
432    }
433
434    fn query_expr(&mut self, expr: &QueryExpr) {
435        match expr {
436            QueryExpr::Points { ids } => {
437                for id in ids {
438                    self.point_id(id);
439                }
440            }
441            QueryExpr::Nearest {
442                input, prefetch, ..
443            } => {
444                self.query_input(input);
445                for p in prefetch {
446                    self.prefetch(p);
447                }
448            }
449            QueryExpr::Recommend {
450                positive,
451                negative,
452                prefetch,
453                ..
454            } => {
455                for item in positive.iter().chain(negative.iter()) {
456                    self.query_input(item);
457                }
458                for p in prefetch {
459                    self.prefetch(p);
460                }
461            }
462            QueryExpr::Context {
463                pairs, prefetch, ..
464            } => {
465                for pair in pairs {
466                    self.query_input(&pair.positive);
467                    self.query_input(&pair.negative);
468                }
469                for p in prefetch {
470                    self.prefetch(p);
471                }
472            }
473            QueryExpr::Discover {
474                target,
475                context,
476                prefetch,
477                ..
478            } => {
479                self.query_input(target);
480                for pair in context {
481                    self.query_input(&pair.positive);
482                    self.query_input(&pair.negative);
483                }
484                for p in prefetch {
485                    self.prefetch(p);
486                }
487            }
488            QueryExpr::OrderBy { start_from, .. } => {
489                if let Some(v) = start_from {
490                    self.value(v);
491                }
492            }
493            QueryExpr::SampleRandom => {}
494            QueryExpr::Fusion { prefetch, .. } => {
495                for p in prefetch {
496                    self.prefetch(p);
497                }
498            }
499            QueryExpr::Formula {
500                expression,
501                defaults,
502                prefetch,
503            } => {
504                self.formula(expression);
505                for (_, v) in defaults {
506                    self.value(v);
507                }
508                for p in prefetch {
509                    self.prefetch(p);
510                }
511            }
512            QueryExpr::RelevanceFeedback {
513                target,
514                feedback,
515                prefetch,
516                ..
517            } => {
518                self.query_input(target);
519                for item in feedback {
520                    self.query_input(&item.example);
521                }
522                for p in prefetch {
523                    self.prefetch(p);
524                }
525            }
526            QueryExpr::Hybrid { text_param, .. } => {
527                if let Some(param) = text_param {
528                    self.param_str(param, None);
529                }
530            }
531            QueryExpr::Rerank {
532                input, prefetch, ..
533            } => {
534                self.query_input(input);
535                for p in prefetch {
536                    self.prefetch(p);
537                }
538            }
539            QueryExpr::CrossRerank {
540                query_param,
541                prefetch,
542                ..
543            } => {
544                if let Some(param) = query_param {
545                    self.param_str(param, None);
546                }
547                for p in prefetch {
548                    self.prefetch(p);
549                }
550            }
551        }
552    }
553
554    fn query_stmt(&mut self, query: &QueryStmt) {
555        for cte in &query.ctes {
556            self.query_stmt(&cte.query);
557        }
558        self.query_expr(&query.expression);
559        if let Some(filter) = &query.filter {
560            self.filter(filter);
561        }
562        self.shard_opt(&query.shard_key);
563        if let Some(param) = &query.page.limit_param {
564            let span = query.page.limit_span;
565            if let Some(name) = param.strip_prefix(':') {
566                self.record_named(name);
567            } else if let Some(idx_str) = param.strip_prefix('?') {
568                if let Ok(idx) = idx_str.parse::<usize>() {
569                    self.record_pos(idx);
570                } else {
571                    self.max_pos = self.max_pos.max(1);
572                }
573            } else {
574                self.record_named(param);
575            }
576            let full = unbound_param_str_err(param, span);
577            let scalar = unbound_param_str_err(param, span);
578            self.note_both(full, scalar);
579        }
580        if let Some(param) = &query.page.offset_param {
581            let span = query.page.offset_span;
582            if let Some(name) = param.strip_prefix(':') {
583                self.record_named(name);
584            } else if let Some(idx_str) = param.strip_prefix('?') {
585                if let Ok(idx) = idx_str.parse::<usize>() {
586                    self.record_pos(idx);
587                } else {
588                    self.max_pos = self.max_pos.max(1);
589                }
590            } else {
591                self.record_named(param);
592            }
593            let full = unbound_param_str_err(param, span);
594            let scalar = unbound_param_str_err(param, span);
595            self.note_both(full, scalar);
596        }
597    }
598
599    fn point_selector(&mut self, sel: &PointSelector) {
600        match sel {
601            PointSelector::Id(id) => self.point_id(id),
602            PointSelector::Ids(ids) => {
603                for id in ids {
604                    self.point_id(id);
605                }
606            }
607            PointSelector::Filter(f) => self.filter(f),
608        }
609    }
610
611    fn page_limit_param(&mut self, param: &str, span: Option<Span>) {
612        self.param_str(param, span);
613    }
614
615    fn stmt(&mut self, stmt: &Stmt) {
616        match stmt {
617            Stmt::Query(query) => self.query_stmt(query),
618            Stmt::Scroll(scroll) => {
619                if let Some(filter) = &scroll.filter {
620                    self.filter(filter);
621                }
622                if let Some(after) = &scroll.after {
623                    self.point_id(after);
624                }
625                if let Some(order) = &scroll.order_by
626                    && let Some(v) = &order.start_from
627                {
628                    self.value(v);
629                }
630                if let Some(param) = &scroll.limit_param {
631                    self.page_limit_param(param, scroll.limit_span);
632                }
633                self.shard_opt(&scroll.shard_key);
634            }
635            Stmt::Upsert(upsert) => {
636                for point in &upsert.points {
637                    match point {
638                        crate::ast::PointEntry::Inline(inline) => {
639                            self.point_id(&inline.id);
640                            if let Some(vectors) = &inline.vectors {
641                                self.point_vectors(vectors);
642                            }
643                            for (_, v) in &inline.payload {
644                                self.value(v);
645                            }
646                        }
647                        crate::ast::PointEntry::Param(name, span) => {
648                            self.record_named(name);
649                            self.has_point_params = true;
650                            self.note_full(unbound_named_err(name, span.as_deref().copied()));
651                        }
652                        crate::ast::PointEntry::PositionalParam(idx, span) => {
653                            self.record_pos(*idx);
654                            self.has_point_params = true;
655                            self.note_full(unbound_positional_err(*idx, span.as_deref().copied()));
656                        }
657                    }
658                }
659                if let Some(filter) = &upsert.update_filter {
660                    self.filter(filter);
661                }
662                self.shard_opt(&upsert.shard_key);
663            }
664            Stmt::Delete(del) => {
665                self.point_selector(&del.selector);
666                self.shard_opt(&del.shard_key);
667            }
668            Stmt::ClearPayload(cp) => {
669                self.point_selector(&cp.selector);
670                self.shard_opt(&cp.shard_key);
671            }
672            Stmt::DeletePayload(dp) => {
673                self.point_selector(&dp.selector);
674                self.shard_opt(&dp.shard_key);
675            }
676            Stmt::DeleteVector(dv) => {
677                self.point_selector(&dv.selector);
678                self.shard_opt(&dv.shard_key);
679            }
680            Stmt::UpdateVector(uv) => {
681                for point in &uv.points {
682                    self.point_id(&point.id);
683                    self.point_vectors(&point.vectors);
684                }
685                self.shard_opt(&uv.shard_key);
686            }
687            Stmt::UpdatePayload(up) => {
688                self.point_selector(&up.selector);
689                for (_, v) in &up.payload {
690                    self.value(v);
691                }
692                self.shard_opt(&up.shard_key);
693            }
694            Stmt::Count(count) => {
695                if let Some(filter) = &count.filter {
696                    self.filter(filter);
697                }
698                self.shard_opt(&count.shard_key);
699            }
700            Stmt::Facet(facet) => {
701                if let Some(filter) = &facet.filter {
702                    self.filter(filter);
703                }
704                if let Some(param) = &facet.limit_param {
705                    self.page_limit_param(param, facet.limit_span);
706                }
707                self.shard_opt(&facet.shard_key);
708            }
709            Stmt::CreateShardKey(sk) => self.shard_opt(&Some(sk.shard_key.clone())),
710            Stmt::DropShardKey(sk) => self.shard_opt(&Some(sk.shard_key.clone())),
711            Stmt::CreateCollection(cc) => {
712                self.shard_list(
713                    cc.config
714                        .as_ref()
715                        .and_then(|config| config.params.as_ref())
716                        .and_then(|params| params.shard_keys.as_ref()),
717                );
718                if let Some(config) = &cc.config {
719                    self.ddl_options(config);
720                }
721            }
722            Stmt::AlterCollection(ac) => {
723                self.shard_list(
724                    ac.config
725                        .as_ref()
726                        .and_then(|config| config.params.as_ref())
727                        .and_then(|params| params.shard_keys.as_ref()),
728                );
729                if let Some(config) = &ac.config {
730                    self.ddl_options(config);
731                }
732            }
733            Stmt::Batch(batch) => {
734                for member in &batch.statements {
735                    self.stmt(member);
736                }
737            }
738            _ => {}
739        }
740    }
741}
742
743pub(crate) fn run_census(stmt: &Stmt) -> Census {
744    let mut census = Census::new();
745    census.stmt(stmt);
746    census
747}
748
749/// Collect all named parameter names and the maximum positional parameter
750/// index present anywhere in a statement AST.
751pub fn collect_statement_params(
752    stmt: &Stmt,
753) -> (alloc::collections::BTreeSet<alloc::string::String>, usize) {
754    let census = run_census(stmt);
755    (census.named, census.max_pos)
756}
757
758/// Whether an upsert template carries whole-point placeholders (`VALUES :p`
759/// / `VALUES ?`) needing dict splice at execution time.
760pub fn stmt_has_point_params(stmt: &Stmt) -> bool {
761    run_census(stmt).has_point_params
762}