Skip to main content

qql_core/params/
ast.rs

1//! AST-level parameter binding across statement nodes.
2
3use super::filter::{bind_filter, bind_point_selector};
4use super::formula::bind_formula;
5use super::input::{bind_context_pair, bind_feedback_item, bind_query_input};
6pub use super::value::{bind_point_id, bind_shard_key, bind_value, resolve_param_u64};
7
8use crate::ast::Value;
9use crate::ast::statement::{
10    PageSpec, PointEntry, PointVectors, Prefetch, PrefetchSource, QueryExpr, QueryStmt, ShardKey,
11    Stmt, UpsertPoint, VectorValue,
12};
13use crate::error::{QqlError, Span};
14use alloc::format;
15use alloc::vec::Vec;
16
17/// Bind parameters into a `PageSpec` in-place.
18pub fn bind_page_spec<F>(
19    page: &mut PageSpec,
20    lookup: &F,
21    positional: &[Value],
22) -> Result<(), QqlError>
23where
24    F: Fn(&str) -> Option<Value>,
25{
26    if let Some(param) = page.limit_param.take() {
27        let span = page.limit_span.take();
28        page.limit = Some(resolve_param_u64(
29            &param, span, lookup, positional, "LIMIT", true,
30        )?);
31    }
32    if let Some(param) = page.offset_param.take() {
33        let span = page.offset_span.take();
34        page.offset = Some(resolve_param_u64(
35            &param, span, lookup, positional, "OFFSET", false,
36        )?);
37    }
38    Ok(())
39}
40
41fn bind_prefetch<F>(
42    prefetch: &mut Prefetch,
43    lookup: &F,
44    positional: &[Value],
45) -> Result<(), QqlError>
46where
47    F: Fn(&str) -> Option<Value>,
48{
49    match &mut prefetch.source {
50        PrefetchSource::Query(sub) => bind_query_stmt(sub, lookup, positional)?,
51        PrefetchSource::Cte(_) => {}
52    }
53    if let Some(f) = &mut prefetch.filter {
54        bind_filter(f, lookup, positional)?;
55    }
56    if let Some(spec) = &mut prefetch.lookup {
57        bind_shard_key(&mut spec.shard_key, lookup, positional)?;
58    }
59    Ok(())
60}
61
62/// Recursively bind parameters into a `QueryExpr` in-place.
63pub fn bind_query_expr<F>(
64    expr: &mut QueryExpr,
65    lookup: &F,
66    positional: &[Value],
67) -> Result<(), QqlError>
68where
69    F: Fn(&str) -> Option<Value>,
70{
71    match expr {
72        QueryExpr::Points { ids } => {
73            for id in ids {
74                bind_point_id(id, lookup, positional)?;
75            }
76        }
77        QueryExpr::Nearest {
78            input, prefetch, ..
79        } => {
80            bind_query_input(input, lookup, positional)?;
81            for p in prefetch {
82                bind_prefetch(p, lookup, positional)?;
83            }
84        }
85        QueryExpr::Recommend {
86            positive,
87            negative,
88            prefetch,
89            ..
90        } => {
91            for pos in positive {
92                bind_query_input(pos, lookup, positional)?;
93            }
94            for neg in negative {
95                bind_query_input(neg, lookup, positional)?;
96            }
97            for p in prefetch {
98                bind_prefetch(p, lookup, positional)?;
99            }
100        }
101        QueryExpr::Context {
102            pairs, prefetch, ..
103        } => {
104            for pair in pairs {
105                bind_context_pair(pair, lookup, positional)?;
106            }
107            for p in prefetch {
108                bind_prefetch(p, lookup, positional)?;
109            }
110        }
111        QueryExpr::Discover {
112            target,
113            context,
114            prefetch,
115            ..
116        } => {
117            bind_query_input(target, lookup, positional)?;
118            for pair in context {
119                bind_context_pair(pair, lookup, positional)?;
120            }
121            for p in prefetch {
122                bind_prefetch(p, lookup, positional)?;
123            }
124        }
125        QueryExpr::OrderBy { start_from, .. } => {
126            if let Some(value) = start_from {
127                bind_value(value, lookup, positional)?;
128            }
129        }
130        QueryExpr::SampleRandom => {}
131        QueryExpr::Fusion { prefetch, .. } => {
132            for p in prefetch {
133                bind_prefetch(p, lookup, positional)?;
134            }
135        }
136        QueryExpr::Formula {
137            expression,
138            defaults,
139            prefetch,
140        } => {
141            bind_formula(expression, lookup, positional, &bind_filter)?;
142            for (_k, v) in defaults {
143                bind_value(v, lookup, positional)?;
144            }
145            for p in prefetch {
146                bind_prefetch(p, lookup, positional)?;
147            }
148        }
149        QueryExpr::RelevanceFeedback {
150            target,
151            feedback,
152            prefetch,
153            ..
154        } => {
155            bind_query_input(target, lookup, positional)?;
156            for item in feedback {
157                bind_feedback_item(item, lookup, positional)?;
158            }
159            for p in prefetch {
160                bind_prefetch(p, lookup, positional)?;
161            }
162        }
163        QueryExpr::Hybrid {
164            text, text_param, ..
165        } => {
166            if let Some(param) = text_param.take() {
167                let val = if let Some(param_name) = param.strip_prefix(':') {
168                    super::value::resolve_param(param_name, None, lookup)?
169                } else if let Some(idx_str) = param.strip_prefix('?') {
170                    let idx = idx_str.parse::<usize>().map_err(|_| {
171                        QqlError::validation(
172                            "QQL-BIND-INVALID-PARAMS",
173                            format!("invalid positional parameter index '?{idx_str}'"),
174                            None,
175                        )
176                    })?;
177                    super::value::resolve_positional(idx, None, positional)?
178                } else {
179                    super::value::resolve_param(&param, None, lookup)?
180                };
181                if let Value::Str(s) = val {
182                    *text = s;
183                } else {
184                    return Err(QqlError::validation(
185                        "QQL-BIND-TYPE-MISMATCH",
186                        format!("parameter '{param}' for HYBRID query must be a string"),
187                        None,
188                    ));
189                }
190            }
191        }
192        QueryExpr::Rerank {
193            input, prefetch, ..
194        } => {
195            bind_query_input(input, lookup, positional)?;
196            for p in prefetch {
197                bind_prefetch(p, lookup, positional)?;
198            }
199        }
200        QueryExpr::CrossRerank {
201            query,
202            query_param,
203            prefetch,
204            ..
205        } => {
206            if let Some(param) = query_param.take() {
207                let val = if let Some(param_name) = param.strip_prefix(':') {
208                    super::value::resolve_param(param_name, None, lookup)?
209                } else if let Some(idx_str) = param.strip_prefix('?') {
210                    let idx = idx_str.parse::<usize>().map_err(|_| {
211                        QqlError::validation(
212                            "QQL-BIND-INVALID-PARAMS",
213                            format!("invalid positional parameter index '?{idx_str}'"),
214                            None,
215                        )
216                    })?;
217                    super::value::resolve_positional(idx, None, positional)?
218                } else {
219                    super::value::resolve_param(&param, None, lookup)?
220                };
221                if let Value::Str(s) = val {
222                    *query = s;
223                } else {
224                    return Err(QqlError::validation(
225                        "QQL-BIND-TYPE-MISMATCH",
226                        format!("parameter '{param}' for CROSS RERANK query must be a string"),
227                        None,
228                    ));
229                }
230            }
231            for p in prefetch {
232                bind_prefetch(p, lookup, positional)?;
233            }
234        }
235    }
236    Ok(())
237}
238
239/// Recursively bind parameters into a `QueryStmt` in-place.
240pub fn bind_query_stmt<F>(
241    query: &mut QueryStmt,
242    lookup: &F,
243    positional: &[Value],
244) -> Result<(), QqlError>
245where
246    F: Fn(&str) -> Option<Value>,
247{
248    for cte in &mut query.ctes {
249        bind_query_stmt(&mut cte.query, lookup, positional)?;
250    }
251    bind_query_expr(&mut query.expression, lookup, positional)?;
252    if let Some(filter) = &mut query.filter {
253        bind_filter(filter, lookup, positional)?;
254    }
255    bind_page_spec(&mut query.page, lookup, positional)?;
256    bind_shard_key(&mut query.shard_key, lookup, positional)?;
257    Ok(())
258}
259
260/// Bind parameters into a required (non-optional) `ShardKey` in-place.
261///
262/// DDL keys (`CREATE`/`DROP SHARD KEY`, `WITH PARAMS shard_keys`) are required
263/// positions, so they bind through a one-slot `Option` and back — reusing
264/// [`bind_shard_key`] instead of duplicating its `Param` resolution.
265fn bind_required_shard_key<F>(
266    key: &mut ShardKey,
267    lookup: &F,
268    positional: &[Value],
269) -> Result<(), QqlError>
270where
271    F: Fn(&str) -> Option<Value>,
272{
273    let mut slot = Some(key.clone());
274    bind_shard_key(&mut slot, lookup, positional)?;
275    if let Some(bound) = slot {
276        *key = bound;
277    }
278    Ok(())
279}
280
281/// Bind parameters into an optional `WITH PARAMS shard_keys` list in-place.
282fn bind_shard_key_list<F>(
283    keys: Option<&mut Vec<ShardKey>>,
284    lookup: &F,
285    positional: &[Value],
286) -> Result<(), QqlError>
287where
288    F: Fn(&str) -> Option<Value>,
289{
290    if let Some(keys) = keys {
291        for key in keys {
292            bind_required_shard_key(key, lookup, positional)?;
293        }
294    }
295    Ok(())
296}
297
298/// Bind parameters into a parsed AST `Stmt` in-place.
299pub fn bind_stmt<F>(stmt: &mut Stmt, lookup: F, positional: &[Value]) -> Result<(), QqlError>
300where
301    F: Fn(&str) -> Option<Value>,
302{
303    match stmt {
304        Stmt::Query(query) => bind_query_stmt(query, &lookup, positional),
305        Stmt::Scroll(scroll) => {
306            if let Some(filter) = &mut scroll.filter {
307                bind_filter(filter, &lookup, positional)?;
308            }
309            if let Some(after) = &mut scroll.after {
310                bind_point_id(after, &lookup, positional)?;
311            }
312            if let Some(order) = scroll.order_by.as_mut()
313                && let Some(value) = order.start_from.as_mut()
314            {
315                bind_value(value, &lookup, positional)?;
316            }
317            if let Some(param) = scroll.limit_param.take() {
318                let span = scroll.limit_span.take();
319                scroll.limit =
320                    resolve_param_u64(&param, span, &lookup, positional, "SCROLL LIMIT", true)?;
321            }
322            bind_shard_key(&mut scroll.shard_key, &lookup, positional)?;
323            Ok(())
324        }
325        Stmt::Upsert(upsert) => {
326            // Whole-point placeholders splice in place: one entry may become
327            // several points (`:rows` bound to a list of point dicts), so the
328            // vec is rebuilt rather than updated in place.
329            let mut bound = Vec::with_capacity(upsert.points.len());
330            for entry in core::mem::take(&mut upsert.points) {
331                match entry {
332                    PointEntry::Inline(mut point) => {
333                        bind_point_id(&mut point.id, &lookup, positional)?;
334                        if let Some(vectors) = &mut point.vectors {
335                            bind_point_vectors(vectors, &lookup, positional)?;
336                        }
337                        for (_k, v) in &mut point.payload {
338                            bind_value(v, &lookup, positional)?;
339                        }
340                        bound.push(PointEntry::Inline(point));
341                    }
342                    PointEntry::Param(name, span) => {
343                        let val = lookup(&name).ok_or_else(|| {
344                            QqlError::validation(
345                                "QQL-BIND-UNBOUND-PARAM",
346                                format!("unbound named parameter ':{name}'"),
347                                span.as_deref().copied(),
348                            )
349                        })?;
350                        bind_point_entry_value(
351                            &mut bound,
352                            val,
353                            span.as_deref().copied(),
354                            &lookup,
355                            positional,
356                        )?;
357                    }
358                    PointEntry::PositionalParam(idx, span) => {
359                        let val = positional.get(idx).cloned().ok_or_else(|| {
360                            QqlError::validation(
361                                "QQL-BIND-MISSING-POSITIONAL",
362                                format!("missing positional parameter ?{}", idx + 1),
363                                span.as_deref().copied(),
364                            )
365                        })?;
366                        bind_point_entry_value(
367                            &mut bound,
368                            val,
369                            span.as_deref().copied(),
370                            &lookup,
371                            positional,
372                        )?;
373                    }
374                }
375            }
376            upsert.points = bound;
377            if let Some(filter) = &mut upsert.update_filter {
378                super::filter::bind_filter(filter, &lookup, positional)?;
379            }
380            bind_shard_key(&mut upsert.shard_key, &lookup, positional)?;
381            Ok(())
382        }
383        Stmt::Delete(del) => {
384            bind_point_selector(&mut del.selector, &lookup, positional)?;
385            bind_shard_key(&mut del.shard_key, &lookup, positional)?;
386            Ok(())
387        }
388        Stmt::ClearPayload(cp) => {
389            bind_point_selector(&mut cp.selector, &lookup, positional)?;
390            bind_shard_key(&mut cp.shard_key, &lookup, positional)?;
391            Ok(())
392        }
393        Stmt::DeletePayload(dp) => {
394            bind_point_selector(&mut dp.selector, &lookup, positional)?;
395            bind_shard_key(&mut dp.shard_key, &lookup, positional)?;
396            Ok(())
397        }
398        Stmt::DeleteVector(dv) => {
399            bind_point_selector(&mut dv.selector, &lookup, positional)?;
400            bind_shard_key(&mut dv.shard_key, &lookup, positional)?;
401            Ok(())
402        }
403        Stmt::UpdateVector(uv) => {
404            for point in &mut uv.points {
405                bind_point_id(&mut point.id, &lookup, positional)?;
406                bind_point_vectors(&mut point.vectors, &lookup, positional)?;
407            }
408            bind_shard_key(&mut uv.shard_key, &lookup, positional)?;
409            Ok(())
410        }
411        Stmt::UpdatePayload(up) => {
412            bind_point_selector(&mut up.selector, &lookup, positional)?;
413            for (_k, v) in &mut up.payload {
414                bind_value(v, &lookup, positional)?;
415            }
416            bind_shard_key(&mut up.shard_key, &lookup, positional)?;
417            Ok(())
418        }
419        Stmt::Count(count) => {
420            if let Some(filter) = &mut count.filter {
421                bind_filter(filter, &lookup, positional)?;
422            }
423            bind_shard_key(&mut count.shard_key, &lookup, positional)?;
424            Ok(())
425        }
426        Stmt::Facet(facet) => {
427            if let Some(filter) = &mut facet.filter {
428                bind_filter(filter, &lookup, positional)?;
429            }
430            if let Some(param) = facet.limit_param.take() {
431                let span = facet.limit_span.take();
432                facet.limit = Some(resolve_param_u64(
433                    &param,
434                    span,
435                    &lookup,
436                    positional,
437                    "FACET LIMIT",
438                    true,
439                )?);
440            }
441            bind_shard_key(&mut facet.shard_key, &lookup, positional)?;
442            Ok(())
443        }
444        Stmt::Batch(batch) => {
445            // Erase to `dyn` so nested batches reuse one instantiation
446            // instead of growing `&&&&F` forever.
447            let lookup = &lookup as &dyn Fn(&str) -> Option<Value>;
448            for member in &mut batch.statements {
449                bind_stmt(member, lookup, positional)?;
450            }
451            Ok(())
452        }
453        Stmt::CreateShardKey(create) => {
454            bind_required_shard_key(&mut create.shard_key, &lookup, positional)
455        }
456        Stmt::DropShardKey(drop) => {
457            bind_required_shard_key(&mut drop.shard_key, &lookup, positional)
458        }
459        Stmt::CreateCollection(create) => {
460            let keys = create
461                .config
462                .as_mut()
463                .and_then(|config| config.params.as_mut())
464                .and_then(|params| params.shard_keys.as_mut());
465            bind_shard_key_list(keys, &lookup, positional)
466        }
467        Stmt::AlterCollection(alter) => {
468            let keys = alter
469                .config
470                .as_mut()
471                .and_then(|config| config.params.as_mut())
472                .and_then(|params| params.shard_keys.as_mut());
473            bind_shard_key_list(keys, &lookup, positional)
474        }
475        other => Err(QqlError::validation(
476            "QQL-BIND-UNSUPPORTED-STATEMENT",
477            format!(
478                "cannot bind parameters into statement type: {}",
479                other.stmt_kind()
480            ),
481            None,
482        )),
483    }
484}
485
486/// Bind parameters into a `VectorValue` in-place.
487pub fn bind_vector_value<F>(
488    vec: &mut VectorValue,
489    lookup: &F,
490    positional: &[Value],
491) -> Result<(), QqlError>
492where
493    F: Fn(&str) -> Option<Value>,
494{
495    match vec {
496        VectorValue::Param(name, span) => {
497            let val = lookup(name).ok_or_else(|| {
498                QqlError::validation(
499                    "QQL-BIND-UNBOUND-PARAM",
500                    format!("unbound named parameter ':{name}'"),
501                    span.as_deref().copied(),
502                )
503            })?;
504            *vec = crate::parser::helpers::vector_from_value(val, span.as_deref().copied())?;
505        }
506        VectorValue::PositionalParam(idx, span) => {
507            let val = positional.get(*idx).cloned().ok_or_else(|| {
508                QqlError::validation(
509                    "QQL-BIND-MISSING-POSITIONAL",
510                    format!("missing positional parameter ?{}", *idx + 1),
511                    span.as_deref().copied(),
512                )
513            })?;
514            *vec = crate::parser::helpers::vector_from_value(val, span.as_deref().copied())?;
515        }
516        VectorValue::Document { options, .. } | VectorValue::Image { options, .. } => {
517            for (_, value) in options {
518                bind_value(value, lookup, positional)?;
519            }
520        }
521        VectorValue::Object {
522            object, options, ..
523        } => {
524            bind_value(object, lookup, positional)?;
525            for (_, value) in options {
526                bind_value(value, lookup, positional)?;
527            }
528        }
529        _ => {}
530    }
531    Ok(())
532}
533
534fn point_param_shape_error(span: Option<Span>) -> QqlError {
535    QqlError::validation(
536        "QQL-BIND-TYPE-MISMATCH",
537        "point parameter must be an object ({id: …, …}) or a list of point objects",
538        span,
539    )
540}
541
542/// Build one concrete point from a bound point-dict value, mirroring
543/// `VALUES {…}` row parsing: `id` is required (case-insensitive), `vector`
544/// routes through vector lowering, every other key becomes payload.
545fn upsert_point_from_items<F>(
546    mut row: Vec<(String, Value)>,
547    span: Option<Span>,
548    lookup: &F,
549    positional: &[Value],
550) -> Result<UpsertPoint, QqlError>
551where
552    F: Fn(&str) -> Option<Value>,
553{
554    let id_index = row
555        .iter()
556        .position(|(key, _)| key.eq_ignore_ascii_case("id"))
557        .ok_or_else(|| {
558            QqlError::validation(
559                "QQL-VALIDATION-UPSERT-ID",
560                "each UPSERT row requires an id",
561                span,
562            )
563        })?;
564    let (_, id) = row.remove(id_index);
565    let mut id = crate::parser::helpers::point_id_from_value(id, span.unwrap_or(Span::new(0, 0)))?;
566    let mut vectors = if let Some(index) = row
567        .iter()
568        .position(|(key, _)| key.eq_ignore_ascii_case("vector"))
569    {
570        let (_, value) = row.remove(index);
571        Some(crate::parser::helpers::point_vectors_from_value(
572            value, span,
573        )?)
574    } else {
575        None
576    };
577    let mut payload = row;
578    // Nested placeholders inside the dict compose like inline rows.
579    bind_point_id(&mut id, lookup, positional)?;
580    if let Some(vectors) = &mut vectors {
581        bind_point_vectors(vectors, lookup, positional)?;
582    }
583    for (_k, v) in &mut payload {
584        bind_value(v, lookup, positional)?;
585    }
586    Ok(UpsertPoint {
587        id,
588        vectors,
589        payload,
590    })
591}
592
593/// Splice one bound whole-point value into the output vec: a dict becomes a
594/// single point, a list of dicts becomes several points.
595fn bind_point_entry_value<F>(
596    bound: &mut Vec<PointEntry>,
597    value: Value,
598    span: Option<Span>,
599    lookup: &F,
600    positional: &[Value],
601) -> Result<(), QqlError>
602where
603    F: Fn(&str) -> Option<Value>,
604{
605    match value {
606        Value::Dict(row) => {
607            bound.push(PointEntry::Inline(upsert_point_from_items(
608                row, span, lookup, positional,
609            )?));
610        }
611        Value::List(items) => {
612            for item in items {
613                match item {
614                    Value::Dict(row) => {
615                        bound.push(PointEntry::Inline(upsert_point_from_items(
616                            row, span, lookup, positional,
617                        )?));
618                    }
619                    _ => return Err(point_param_shape_error(span)),
620                }
621            }
622        }
623        _ => return Err(point_param_shape_error(span)),
624    }
625    Ok(())
626}
627
628/// Bind parameters into `PointVectors` in-place.
629pub fn bind_point_vectors<F>(
630    pv: &mut PointVectors,
631    lookup: &F,
632    positional: &[Value],
633) -> Result<(), QqlError>
634where
635    F: Fn(&str) -> Option<Value>,
636{
637    match pv {
638        PointVectors::Param(name, span) => {
639            let val = lookup(name).ok_or_else(|| {
640                QqlError::validation(
641                    "QQL-BIND-UNBOUND-PARAM",
642                    format!("unbound named parameter ':{name}'"),
643                    span.as_deref().copied(),
644                )
645            })?;
646            *pv = crate::parser::helpers::point_vectors_from_value(val, span.as_deref().copied())?;
647        }
648        PointVectors::PositionalParam(idx, span) => {
649            let val = positional.get(*idx).cloned().ok_or_else(|| {
650                QqlError::validation(
651                    "QQL-BIND-MISSING-POSITIONAL",
652                    format!("missing positional parameter ?{}", *idx + 1),
653                    span.as_deref().copied(),
654                )
655            })?;
656            *pv = crate::parser::helpers::point_vectors_from_value(val, span.as_deref().copied())?;
657        }
658        PointVectors::Unnamed(v) => {
659            bind_vector_value(v, lookup, positional)?;
660        }
661        PointVectors::Named(list) => {
662            for (_, v) in list {
663                bind_vector_value(v, lookup, positional)?;
664            }
665        }
666    }
667    Ok(())
668}