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

1""" 

2DAG analysis rules. 

3""" 

4 

5import re 

6from typing import Any, Dict, List, Set, Tuple 

7 

8 

9class BaseRule: 

10 """Base rule class.""" 

11 

12 def __init__(self): 

13 self.id = "" 

14 self.name = "" 

15 self.severity = "" 

16 self.category = "" 

17 self.execution_mode = "" 

18 

19 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]: 

20 """Evaluate rule against source code.""" 

21 raise NotImplementedError 

22 

23 

24class CircularDependencyRule(BaseRule): 

25 """Detect circular dependencies in DAGs.""" 

26 

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" 

34 

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*\)") 

41 

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]] = [] 

45 

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:])) 

49 

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)) 

54 

55 for m in self._SET_DOWNSTREAM.finditer(source_code): 

56 edges.append((m.group(1), m.group(2))) 

57 

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))) 

61 

62 return edges 

63 

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()) 

71 

72 WHITE, GRAY, BLACK = 0, 1, 2 

73 color = {node: WHITE for node in graph} 

74 

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 

84 

85 return any(color[node] == WHITE and dfs(node) for node in list(graph)) 

86 

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 = [] 

90 

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 ) 

102 

103 return violations 

104 

105 

106class MissingSLARule(BaseRule): 

107 """Detect missing SLAs on production DAGs.""" 

108 

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" 

116 

117 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]: 

118 """Detect missing SLAs.""" 

119 violations = [] 

120 

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 ) 

132 

133 return violations 

134 

135 

136class ExpensiveImportsRule(BaseRule): 

137 """Detect expensive imports in DAG files.""" 

138 

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"] 

147 

148 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]: 

149 """Detect expensive imports.""" 

150 violations = [] 

151 

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 ) 

164 

165 return violations 

166 

167 

168class ParseTimeRule(BaseRule): 

169 """Analyze DAG parse time.""" 

170 

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" 

178 

179 def evaluate(self, source_code: str, file_name: str = "") -> List[Dict[str, Any]]: 

180 """Analyze parse time.""" 

181 violations = [] 

182 

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 ) 

194 

195 return violations