1use super::dag::{ModelDag, ModelNode};
10use super::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.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 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 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 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 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 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 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 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}