1use 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
54pub(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 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 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 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 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 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
749pub 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
758pub fn stmt_has_point_params(stmt: &Stmt) -> bool {
761 run_census(stmt).has_point_params
762}