1use crate::dag::{ModelDag, ModelNode};
10use crate::parser::DependencyRef;
11use anyhow::{bail, Result};
12use petgraph::Direction;
13use std::collections::HashSet;
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub enum SelectMode {
18 Exact,
20 Execute,
22}
23
24#[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 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 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
68pub 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
84pub 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 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 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 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 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 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 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::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 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}