Skip to main content

rbt/core/
select.rs

1//! Model selection (`--select`) similar to dbt node selectors (`name` / `+name` / `name+`).
2//!
3//! # Execute vs Exact
4//!
5//! * [`SelectMode::Exact`] — expand only explicit `+` modifiers (for `compile` listing).
6//! * [`SelectMode::Execute`] — always include **ancestors** of every selected node so
7//!   `ref()` dependencies exist when the subgraph is run.
8
9use super::dag::{ModelDag, ModelNode};
10use super::parser::DependencyRef;
11use anyhow::{bail, Result};
12use petgraph::Direction;
13use std::collections::HashSet;
14
15/// How selection expands relative to named seeds.
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub enum SelectMode {
18    /// Expand `+` modifiers only (compile listing).
19    Exact,
20    /// Ensure a runnable subgraph: always include ancestors of every selected node.
21    Execute,
22}
23
24/// One comma-separated token: `name`, `+name`, `name+`, `+name+`.
25#[derive(Debug, Clone, PartialEq, Eq)]
26pub struct SelectToken {
27    pub name: String,
28    pub upstream: bool,
29    pub downstream: bool,
30}
31
32impl SelectToken {
33    /// Parse a single select token.
34    pub fn parse(raw: &str) -> Result<Self> {
35        let s = raw.trim();
36        if s.is_empty() {
37            bail!("E_RBT_SELECT_EMPTY: empty --select token");
38        }
39        let upstream = s.starts_with('+');
40        let body = if upstream { &s[1..] } else { s };
41        let downstream = body.ends_with('+');
42        let name = if downstream {
43            body[..body.len().saturating_sub(1)].trim()
44        } else {
45            body.trim()
46        };
47        if name.is_empty() || name.contains('+') {
48            bail!(
49                "E_RBT_SELECT_INVALID: invalid --select token '{}'; expected name, +name, name+, or +name+",
50                raw
51            );
52        }
53        // Reject path-like or whitespace names early
54        if name.contains('/') || name.contains('\\') || name.contains(' ') {
55            bail!(
56                "E_RBT_SELECT_INVALID: model name '{}' must not contain path separators or spaces",
57                name
58            );
59        }
60        Ok(Self {
61            name: name.to_string(),
62            upstream,
63            downstream,
64        })
65    }
66}
67
68/// Parse a full `--select` string into tokens.
69pub fn parse_select_spec(spec: &str) -> Result<Vec<SelectToken>> {
70    let mut out = Vec::new();
71    for part in spec.split([',', ' ']) {
72        let part = part.trim();
73        if part.is_empty() {
74            continue;
75        }
76        out.push(SelectToken::parse(part)?);
77    }
78    if out.is_empty() {
79        bail!("E_RBT_SELECT_EMPTY: --select produced no models");
80    }
81    Ok(out)
82}
83
84/// Whether a model declares frontmatter tests / grain / unique_key worth running under `rbt test`.
85pub fn model_has_test_contract(node: &ModelNode) -> bool {
86    node.frontmatter
87        .as_ref()
88        .map(|fm| {
89            fm.tests
90                .as_ref()
91                .map(|t| !t.is_empty())
92                .unwrap_or(false)
93                || fm
94                    .unique_key
95                    .as_ref()
96                    .map(|u| !u.is_empty())
97                    .unwrap_or(false)
98                || fm.grain.as_ref().map(|g| !g.is_empty()).unwrap_or(false)
99        })
100        .unwrap_or(false)
101}
102
103impl ModelDag {
104    /// Resolve which model names are in the selection.
105    pub fn resolve_select(&self, select: Option<&str>, mode: SelectMode) -> Result<HashSet<String>> {
106        let Some(spec) = select.map(str::trim).filter(|s| !s.is_empty()) else {
107            return Ok(self.node_map.keys().cloned().collect());
108        };
109
110        let tokens = parse_select_spec(spec)?;
111        let mut keep: HashSet<String> = HashSet::new();
112
113        for token in tokens {
114            if !self.node_map.contains_key(&token.name) {
115                let available: Vec<_> = {
116                    let mut v: Vec<_> = self.node_map.keys().cloned().collect();
117                    v.sort();
118                    v
119                };
120                bail!(
121                    "E_RBT_MODEL_NOT_FOUND: model '{}' not in project (select={}). Available: {}",
122                    token.name,
123                    spec,
124                    if available.is_empty() {
125                        "(none)".to_string()
126                    } else {
127                        available.join(", ")
128                    }
129                );
130            }
131            let mut up = token.upstream;
132            let down = token.downstream;
133            // Execute mode always pulls ancestors so refs resolve at runtime.
134            if mode == SelectMode::Execute {
135                up = true;
136            }
137            keep.insert(token.name.clone());
138            if up {
139                self.collect_ancestors(&token.name, &mut keep);
140            }
141            if down {
142                self.collect_descendants(&token.name, &mut keep);
143            }
144        }
145
146        Ok(keep)
147    }
148
149    /// Return a new DAG containing only `keep` models (must include all model deps).
150    pub fn subgraph(&self, keep: &HashSet<String>) -> Result<ModelDag> {
151        if keep.is_empty() {
152            bail!("E_RBT_SELECT_EMPTY: selection resolved to zero models");
153        }
154
155        let mut out = ModelDag::new();
156        for node in self.topological_sequence()? {
157            if !keep.contains(&node.name) {
158                continue;
159            }
160            for dep in &node.dependencies {
161                if let DependencyRef::Model(dep_name) = dep {
162                    if !keep.contains(dep_name) {
163                        bail!(
164                            "E_RBT_SELECT_INCOMPLETE: model '{}' depends on '{}' which is not selected; \
165                             use SelectMode::Execute or include +upstream",
166                            node.name,
167                            dep_name
168                        );
169                    }
170                }
171            }
172            let name = node.name.clone();
173            let idx = out.graph.add_node(node);
174            out.node_map.insert(name, idx);
175        }
176
177        // Rebuild edges among kept nodes (topo order already validated).
178        let mut edges = Vec::new();
179        for &idx in out.node_map.values() {
180            let node = &out.graph[idx];
181            for dep in &node.dependencies {
182                if let DependencyRef::Model(dep_name) = dep {
183                    if let Some(&dep_idx) = out.node_map.get(dep_name) {
184                        edges.push((dep_idx, idx));
185                    }
186                }
187            }
188        }
189        for (from, to) in edges {
190            out.graph.add_edge(from, to, ());
191        }
192        Ok(out)
193    }
194
195    /// Apply `--select` and return a filtered DAG.
196    pub fn apply_select(&self, select: Option<&str>, mode: SelectMode) -> Result<ModelDag> {
197        let keep = self.resolve_select(select, mode)?;
198        self.subgraph(&keep)
199    }
200
201    /// Names of models that declare a test contract (for default `rbt test`).
202    pub fn models_with_test_contract(&self) -> Result<Vec<String>> {
203        Ok(self
204            .topological_sequence()?
205            .into_iter()
206            .filter(model_has_test_contract)
207            .map(|n| n.name)
208            .collect())
209    }
210
211    fn collect_ancestors(&self, name: &str, keep: &mut HashSet<String>) {
212        let Some(&idx) = self.node_map.get(name) else {
213            return;
214        };
215        let mut stack: Vec<_> = self
216            .graph
217            .neighbors_directed(idx, Direction::Incoming)
218            .collect();
219        while let Some(n) = stack.pop() {
220            let n_name = self.graph[n].name.clone();
221            if keep.insert(n_name) {
222                stack.extend(self.graph.neighbors_directed(n, Direction::Incoming));
223            }
224        }
225    }
226
227    fn collect_descendants(&self, name: &str, keep: &mut HashSet<String>) {
228        let Some(&idx) = self.node_map.get(name) else {
229            return;
230        };
231        let mut stack: Vec<_> = self
232            .graph
233            .neighbors_directed(idx, Direction::Outgoing)
234            .collect();
235        while let Some(n) = stack.pop() {
236            let n_name = self.graph[n].name.clone();
237            if keep.insert(n_name) {
238                stack.extend(self.graph.neighbors_directed(n, Direction::Outgoing));
239            }
240        }
241    }
242}
243
244#[cfg(test)]
245mod tests {
246    use super::*;
247    use crate::core::dag::{Materialization, OutputFormat};
248
249    fn sample_dag() -> ModelDag {
250        let mut dag = ModelDag::new();
251        dag.add_model_with_format(
252            "stg_a",
253            "SELECT 1 AS id",
254            Materialization::Table,
255            OutputFormat::Parquet,
256            None,
257            "",
258        )
259        .unwrap();
260        dag.add_model_with_format(
261            "tf_b",
262            "SELECT * FROM {{ ref('stg_a') }}",
263            Materialization::Table,
264            OutputFormat::Parquet,
265            None,
266            "",
267        )
268        .unwrap();
269        dag.add_model_with_format(
270            "fact_c",
271            "SELECT * FROM {{ ref('tf_b') }}",
272            Materialization::Table,
273            OutputFormat::Parquet,
274            None,
275            "",
276        )
277        .unwrap();
278        // parallel branch
279        dag.add_model_with_format(
280            "stg_x",
281            "SELECT 2 AS id",
282            Materialization::Table,
283            OutputFormat::Parquet,
284            None,
285            "",
286        )
287        .unwrap();
288        dag.build_graph().unwrap();
289        dag
290    }
291
292    #[test]
293    fn select_none_is_all() {
294        let dag = sample_dag();
295        let keep = dag.resolve_select(None, SelectMode::Execute).unwrap();
296        assert_eq!(keep.len(), 4);
297    }
298
299    #[test]
300    fn select_execute_includes_ancestors() {
301        let dag = sample_dag();
302        let keep = dag
303            .resolve_select(Some("fact_c"), SelectMode::Execute)
304            .unwrap();
305        assert!(keep.contains("stg_a"));
306        assert!(keep.contains("tf_b"));
307        assert!(keep.contains("fact_c"));
308        assert!(!keep.contains("stg_x"));
309    }
310
311    #[test]
312    fn select_exact_bare_name_is_only_self() {
313        let dag = sample_dag();
314        let keep = dag
315            .resolve_select(Some("fact_c"), SelectMode::Exact)
316            .unwrap();
317        assert_eq!(keep, HashSet::from(["fact_c".to_string()]));
318    }
319
320    #[test]
321    fn select_downstream_plus() {
322        let dag = sample_dag();
323        let keep = dag
324            .resolve_select(Some("stg_a+"), SelectMode::Exact)
325            .unwrap();
326        assert!(keep.contains("stg_a"));
327        assert!(keep.contains("tf_b"));
328        assert!(keep.contains("fact_c"));
329    }
330
331    #[test]
332    fn select_upstream_plus_exact() {
333        let dag = sample_dag();
334        let keep = dag
335            .resolve_select(Some("+fact_c"), SelectMode::Exact)
336            .unwrap();
337        assert!(keep.contains("stg_a") && keep.contains("tf_b") && keep.contains("fact_c"));
338    }
339
340    #[test]
341    fn select_both_plus() {
342        let dag = sample_dag();
343        let keep = dag
344            .resolve_select(Some("+tf_b+"), SelectMode::Exact)
345            .unwrap();
346        assert!(keep.contains("stg_a"));
347        assert!(keep.contains("tf_b"));
348        assert!(keep.contains("fact_c"));
349    }
350
351    #[test]
352    fn select_comma_and_space() {
353        let dag = sample_dag();
354        let keep = dag
355            .resolve_select(Some("stg_a, stg_x"), SelectMode::Exact)
356            .unwrap();
357        assert!(keep.contains("stg_a") && keep.contains("stg_x"));
358        assert!(!keep.contains("fact_c"));
359    }
360
361    #[test]
362    fn select_missing_errors() {
363        let dag = sample_dag();
364        let err = dag
365            .resolve_select(Some("nope"), SelectMode::Execute)
366            .unwrap_err()
367            .to_string();
368        assert!(err.contains("E_RBT_MODEL_NOT_FOUND"));
369        assert!(err.contains("Available:"));
370    }
371
372    #[test]
373    fn select_invalid_token() {
374        assert!(SelectToken::parse("+").is_err());
375        assert!(SelectToken::parse("a+b").is_err());
376        assert!(parse_select_spec("  ,  ").is_err());
377    }
378
379    #[test]
380    fn subgraph_execute_runnable() {
381        let dag = sample_dag();
382        let sub = dag
383            .apply_select(Some("fact_c"), SelectMode::Execute)
384            .unwrap();
385        assert_eq!(sub.node_map.len(), 3);
386        let tiers = sub.execution_tiers().unwrap();
387        assert_eq!(tiers[0][0].name, "stg_a");
388    }
389
390    #[test]
391    fn subgraph_exact_without_deps_fails() {
392        let dag = sample_dag();
393        let err = dag
394            .apply_select(Some("fact_c"), SelectMode::Exact)
395            .unwrap_err()
396            .to_string();
397        assert!(err.contains("E_RBT_SELECT_INCOMPLETE") || err.contains("depends on"));
398    }
399}