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.as_ref().map(|t| !t.is_empty()).unwrap_or(false)
90                || fm
91                    .unique_key
92                    .as_ref()
93                    .map(|u| !u.is_empty())
94                    .unwrap_or(false)
95                || fm.grain.as_ref().map(|g| !g.is_empty()).unwrap_or(false)
96        })
97        .unwrap_or(false)
98}
99
100impl ModelDag {
101    /// Resolve which model names are in the selection.
102    pub fn resolve_select(
103        &self,
104        select: Option<&str>,
105        mode: SelectMode,
106    ) -> Result<HashSet<String>> {
107        let Some(spec) = select.map(str::trim).filter(|s| !s.is_empty()) else {
108            return Ok(self.node_map.keys().cloned().collect());
109        };
110
111        let tokens = parse_select_spec(spec)?;
112        let mut keep: HashSet<String> = HashSet::new();
113
114        for token in tokens {
115            if !self.node_map.contains_key(&token.name) {
116                let available: Vec<_> = {
117                    let mut v: Vec<_> = self.node_map.keys().cloned().collect();
118                    v.sort();
119                    v
120                };
121                bail!(
122                    "E_RBT_MODEL_NOT_FOUND: model '{}' not in project (select={}). Available: {}",
123                    token.name,
124                    spec,
125                    if available.is_empty() {
126                        "(none)".to_string()
127                    } else {
128                        available.join(", ")
129                    }
130                );
131            }
132            let mut up = token.upstream;
133            let down = token.downstream;
134            // Execute mode always pulls ancestors so refs resolve at runtime.
135            if mode == SelectMode::Execute {
136                up = true;
137            }
138            keep.insert(token.name.clone());
139            if up {
140                self.collect_ancestors(&token.name, &mut keep);
141            }
142            if down {
143                self.collect_descendants(&token.name, &mut keep);
144            }
145        }
146
147        Ok(keep)
148    }
149
150    /// Return a new DAG containing only `keep` models (must include all model deps).
151    pub fn subgraph(&self, keep: &HashSet<String>) -> Result<ModelDag> {
152        if keep.is_empty() {
153            bail!("E_RBT_SELECT_EMPTY: selection resolved to zero models");
154        }
155
156        let mut out = ModelDag::new();
157        for node in self.topological_sequence()? {
158            if !keep.contains(&node.name) {
159                continue;
160            }
161            for dep in &node.dependencies {
162                if let DependencyRef::Model(dep_name) = dep {
163                    if !keep.contains(dep_name) {
164                        bail!(
165                            "E_RBT_SELECT_INCOMPLETE: model '{}' depends on '{}' which is not selected; \
166                             use SelectMode::Execute or include +upstream",
167                            node.name,
168                            dep_name
169                        );
170                    }
171                }
172            }
173            let name = node.name.clone();
174            let idx = out.graph.add_node(node);
175            out.node_map.insert(name, idx);
176        }
177
178        // Rebuild edges among kept nodes (topo order already validated).
179        let mut edges = Vec::new();
180        for &idx in out.node_map.values() {
181            let node = &out.graph[idx];
182            for dep in &node.dependencies {
183                if let DependencyRef::Model(dep_name) = dep {
184                    if let Some(&dep_idx) = out.node_map.get(dep_name) {
185                        edges.push((dep_idx, idx));
186                    }
187                }
188            }
189        }
190        for (from, to) in edges {
191            out.graph.add_edge(from, to, ());
192        }
193        Ok(out)
194    }
195
196    /// Apply `--select` and return a filtered DAG.
197    pub fn apply_select(&self, select: Option<&str>, mode: SelectMode) -> Result<ModelDag> {
198        let keep = self.resolve_select(select, mode)?;
199        self.subgraph(&keep)
200    }
201
202    /// Names of models that declare a test contract (for default `rbt test`).
203    pub fn models_with_test_contract(&self) -> Result<Vec<String>> {
204        Ok(self
205            .topological_sequence()?
206            .into_iter()
207            .filter(model_has_test_contract)
208            .map(|n| n.name)
209            .collect())
210    }
211
212    fn collect_ancestors(&self, name: &str, keep: &mut HashSet<String>) {
213        let Some(&idx) = self.node_map.get(name) else {
214            return;
215        };
216        let mut stack: Vec<_> = self
217            .graph
218            .neighbors_directed(idx, Direction::Incoming)
219            .collect();
220        while let Some(n) = stack.pop() {
221            let n_name = self.graph[n].name.clone();
222            if keep.insert(n_name) {
223                stack.extend(self.graph.neighbors_directed(n, Direction::Incoming));
224            }
225        }
226    }
227
228    fn collect_descendants(&self, name: &str, keep: &mut HashSet<String>) {
229        let Some(&idx) = self.node_map.get(name) else {
230            return;
231        };
232        let mut stack: Vec<_> = self
233            .graph
234            .neighbors_directed(idx, Direction::Outgoing)
235            .collect();
236        while let Some(n) = stack.pop() {
237            let n_name = self.graph[n].name.clone();
238            if keep.insert(n_name) {
239                stack.extend(self.graph.neighbors_directed(n, Direction::Outgoing));
240            }
241        }
242    }
243}
244
245#[cfg(test)]
246mod tests {
247    use super::*;
248    use crate::core::dag::{Materialization, OutputFormat};
249
250    fn sample_dag() -> ModelDag {
251        let mut dag = ModelDag::new();
252        dag.add_model_with_format(
253            "stg_a",
254            "SELECT 1 AS id",
255            Materialization::Table,
256            OutputFormat::Parquet,
257            None,
258            "",
259        )
260        .unwrap();
261        dag.add_model_with_format(
262            "tf_b",
263            "SELECT * FROM {{ ref('stg_a') }}",
264            Materialization::Table,
265            OutputFormat::Parquet,
266            None,
267            "",
268        )
269        .unwrap();
270        dag.add_model_with_format(
271            "fact_c",
272            "SELECT * FROM {{ ref('tf_b') }}",
273            Materialization::Table,
274            OutputFormat::Parquet,
275            None,
276            "",
277        )
278        .unwrap();
279        // parallel branch
280        dag.add_model_with_format(
281            "stg_x",
282            "SELECT 2 AS id",
283            Materialization::Table,
284            OutputFormat::Parquet,
285            None,
286            "",
287        )
288        .unwrap();
289        dag.build_graph().unwrap();
290        dag
291    }
292
293    #[test]
294    fn select_none_is_all() {
295        let dag = sample_dag();
296        let keep = dag.resolve_select(None, SelectMode::Execute).unwrap();
297        assert_eq!(keep.len(), 4);
298    }
299
300    #[test]
301    fn select_execute_includes_ancestors() {
302        let dag = sample_dag();
303        let keep = dag
304            .resolve_select(Some("fact_c"), SelectMode::Execute)
305            .unwrap();
306        assert!(keep.contains("stg_a"));
307        assert!(keep.contains("tf_b"));
308        assert!(keep.contains("fact_c"));
309        assert!(!keep.contains("stg_x"));
310    }
311
312    #[test]
313    fn select_exact_bare_name_is_only_self() {
314        let dag = sample_dag();
315        let keep = dag
316            .resolve_select(Some("fact_c"), SelectMode::Exact)
317            .unwrap();
318        assert_eq!(keep, HashSet::from(["fact_c".to_string()]));
319    }
320
321    #[test]
322    fn select_downstream_plus() {
323        let dag = sample_dag();
324        let keep = dag
325            .resolve_select(Some("stg_a+"), SelectMode::Exact)
326            .unwrap();
327        assert!(keep.contains("stg_a"));
328        assert!(keep.contains("tf_b"));
329        assert!(keep.contains("fact_c"));
330    }
331
332    #[test]
333    fn select_upstream_plus_exact() {
334        let dag = sample_dag();
335        let keep = dag
336            .resolve_select(Some("+fact_c"), SelectMode::Exact)
337            .unwrap();
338        assert!(keep.contains("stg_a") && keep.contains("tf_b") && keep.contains("fact_c"));
339    }
340
341    #[test]
342    fn select_both_plus() {
343        let dag = sample_dag();
344        let keep = dag
345            .resolve_select(Some("+tf_b+"), SelectMode::Exact)
346            .unwrap();
347        assert!(keep.contains("stg_a"));
348        assert!(keep.contains("tf_b"));
349        assert!(keep.contains("fact_c"));
350    }
351
352    #[test]
353    fn select_comma_and_space() {
354        let dag = sample_dag();
355        let keep = dag
356            .resolve_select(Some("stg_a, stg_x"), SelectMode::Exact)
357            .unwrap();
358        assert!(keep.contains("stg_a") && keep.contains("stg_x"));
359        assert!(!keep.contains("fact_c"));
360    }
361
362    #[test]
363    fn select_missing_errors() {
364        let dag = sample_dag();
365        let err = dag
366            .resolve_select(Some("nope"), SelectMode::Execute)
367            .unwrap_err()
368            .to_string();
369        assert!(err.contains("E_RBT_MODEL_NOT_FOUND"));
370        assert!(err.contains("Available:"));
371    }
372
373    #[test]
374    fn select_invalid_token() {
375        assert!(SelectToken::parse("+").is_err());
376        assert!(SelectToken::parse("a+b").is_err());
377        assert!(parse_select_spec("  ,  ").is_err());
378    }
379
380    #[test]
381    fn subgraph_execute_runnable() {
382        let dag = sample_dag();
383        let sub = dag
384            .apply_select(Some("fact_c"), SelectMode::Execute)
385            .unwrap();
386        assert_eq!(sub.node_map.len(), 3);
387        let tiers = sub.execution_tiers().unwrap();
388        assert_eq!(tiers[0][0].name, "stg_a");
389    }
390
391    #[test]
392    fn subgraph_exact_without_deps_fails() {
393        let dag = sample_dag();
394        let err = dag
395            .apply_select(Some("fact_c"), SelectMode::Exact)
396            .unwrap_err()
397            .to_string();
398        assert!(err.contains("E_RBT_SELECT_INCOMPLETE") || err.contains("depends on"));
399    }
400}