Coverage for python/pyairflowtester/rules/dag.py: 98%
101 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-19 20:43 +0530
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-19 20:43 +0530
1"""
2DAG analysis rules.
3"""
5import re
6from typing import Any, Dict, List, Set, Tuple
9class BaseRule:
10 """Base rule class."""
12 def __init__(self):
13 self.id = ""
14 self.name = ""
15 self.severity = ""
16 self.category = ""
17 self.execution_mode = ""
19 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]:
20 """Evaluate rule against source code."""
21 raise NotImplementedError
24class CircularDependencyRule(BaseRule):
25 """Detect circular dependencies in DAGs."""
27 def __init__(self):
28 super().__init__()
29 self.id = "AFW001"
30 self.name = "Circular Dependency"
31 self.severity = "critical"
32 self.category = "reliability"
33 self.execution_mode = "static"
35 # Matches chains like "task_a >> task_b >> task_c" (2+ nodes)
36 _DOWNSTREAM_CHAIN = re.compile(r"\b\w+(?:\s*>>\s*\w+)+\b")
37 # Matches chains like "task_a << task_b << task_c" (2+ nodes)
38 _UPSTREAM_CHAIN = re.compile(r"\b\w+(?:\s*<<\s*\w+)+\b")
39 _SET_DOWNSTREAM = re.compile(r"(\w+)\.set_downstream\(\s*(\w+)\s*\)")
40 _SET_UPSTREAM = re.compile(r"(\w+)\.set_upstream\(\s*(\w+)\s*\)")
42 def _extract_edges(self, source_code: str) -> List[Tuple[str, str]]:
43 """Extract directed task-dependency edges (upstream -> downstream) from source."""
44 edges: List[Tuple[str, str]] = []
46 for match in self._DOWNSTREAM_CHAIN.finditer(source_code):
47 nodes = re.split(r"\s*>>\s*", match.group())
48 edges.extend(zip(nodes, nodes[1:]))
50 for match in self._UPSTREAM_CHAIN.finditer(source_code): 50 ↛ 51line 50 didn't jump to line 51 because the loop on line 50 never started
51 nodes = re.split(r"\s*<<\s*", match.group())
52 # "a << b" means b is upstream of a, i.e. edge b -> a
53 edges.extend(zip(nodes[1:], nodes))
55 for m in self._SET_DOWNSTREAM.finditer(source_code):
56 edges.append((m.group(1), m.group(2)))
58 for m in self._SET_UPSTREAM.finditer(source_code):
59 # "a.set_upstream(b)" means b -> a
60 edges.append((m.group(2), m.group(1)))
62 return edges
64 @staticmethod
65 def _has_cycle(edges: List[Tuple[str, str]]) -> bool:
66 """Detect a cycle in a directed graph via DFS with coloring."""
67 graph: Dict[str, Set[str]] = {}
68 for upstream, downstream in edges:
69 graph.setdefault(upstream, set()).add(downstream)
70 graph.setdefault(downstream, set())
72 WHITE, GRAY, BLACK = 0, 1, 2
73 color = {node: WHITE for node in graph}
75 def dfs(node: str) -> bool:
76 color[node] = GRAY
77 for neighbor in graph.get(node, ()):
78 if color.get(neighbor) == GRAY:
79 return True
80 if color.get(neighbor) == WHITE and dfs(neighbor):
81 return True
82 color[node] = BLACK
83 return False
85 return any(color[node] == WHITE and dfs(node) for node in list(graph))
87 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]:
88 """Detect circular dependencies via graph-cycle detection over parsed edges."""
89 violations = []
91 edges = self._extract_edges(source_code)
92 if edges and self._has_cycle(edges):
93 violations.append(
94 {
95 "rule_id": self.id,
96 "severity": self.severity,
97 "affected_resource": file_name,
98 "message": "Circular dependency detected in task graph",
99 "remediation": "Review task dependencies and remove cycles",
100 }
101 )
103 return violations
106class MissingSLARule(BaseRule):
107 """Detect missing SLAs on production DAGs."""
109 def __init__(self):
110 super().__init__()
111 self.id = "AFW002"
112 self.name = "Missing SLA"
113 self.severity = "high"
114 self.category = "reliability"
115 self.execution_mode = "static"
117 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]:
118 """Detect missing SLAs."""
119 violations = []
121 # Check if DAG has SLA defined
122 if "sla" not in source_code.lower() and "production" in file_name.lower():
123 violations.append(
124 {
125 "rule_id": self.id,
126 "severity": self.severity,
127 "affected_resource": file_name,
128 "message": "Production DAG missing SLA configuration",
129 "remediation": "Add 'sla' parameter to DAG definition",
130 }
131 )
133 return violations
136class ExpensiveImportsRule(BaseRule):
137 """Detect expensive imports in DAG files."""
139 def __init__(self):
140 super().__init__()
141 self.id = "AFW003"
142 self.name = "Expensive Imports"
143 self.severity = "medium"
144 self.category = "performance"
145 self.execution_mode = "static"
146 self.expensive_modules = ["tensorflow", "torch", "sklearn", "pandas", "numpy"]
148 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]:
149 """Detect expensive imports."""
150 violations = []
152 for module in self.expensive_modules:
153 pattern = rf"^import\s+{module}|^from\s+{module}\s+import"
154 if re.search(pattern, source_code, re.MULTILINE):
155 violations.append(
156 {
157 "rule_id": self.id,
158 "severity": self.severity,
159 "affected_resource": module,
160 "message": f"Expensive import detected: {module}",
161 "remediation": f"Move '{module}' import inside task or use lazy import",
162 }
163 )
165 return violations
168class ParseTimeRule(BaseRule):
169 """Analyze DAG parse time."""
171 def __init__(self):
172 super().__init__()
173 self.id = "AFW004"
174 self.name = "Parse Time Analysis"
175 self.severity = "medium"
176 self.category = "performance"
177 self.execution_mode = "static"
179 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]:
180 """Analyze parse time."""
181 violations = []
183 # Check for potentially slow patterns
184 if re.search(r"for\s+\w+\s+in\s+.*:\s+create.*DAG", source_code, re.DOTALL):
185 violations.append(
186 {
187 "rule_id": self.id,
188 "severity": self.severity,
189 "affected_resource": file_name,
190 "message": "DAG file contains loop-based DAG generation (slow parsing)",
191 "remediation": "Use task factories or DAG generation patterns",
192 }
193 )
195 return violations