Coverage for python/pyairflowtester/dependency_intelligence/models.py: 88%

135 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-09-19 20:43 +0530

1"""Data models for Dependency Intelligence Engine.""" 

2 

3from dataclasses import dataclass, field 

4from datetime import datetime 

5from enum import Enum 

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

7 

8 

9class NodeType(Enum): 

10 """Types of nodes in dependency graph.""" 

11 

12 DAG = "dag" 

13 TASK = "task" 

14 TASK_GROUP = "task_group" 

15 DATASET = "dataset" 

16 DBT_SOURCE = "dbt_source" 

17 DBT_MODEL = "dbt_model" 

18 DBT_TEST = "dbt_test" 

19 DBT_SNAPSHOT = "dbt_snapshot" 

20 DBT_EXPOSURE = "dbt_exposure" 

21 EXTERNAL_TABLE = "external_table" 

22 EXTERNAL_API = "external_api" 

23 DASHBOARD = "dashboard" 

24 

25 

26class NodeSeverity(Enum): 

27 """Severity levels for nodes.""" 

28 

29 CRITICAL = "critical" 

30 HIGH = "high" 

31 MEDIUM = "medium" 

32 LOW = "low" 

33 

34 

35class RelationshipType(Enum): 

36 """Types of edges/relationships.""" 

37 

38 DEPENDS_ON = "depends_on" 

39 TRIGGERS = "triggers" 

40 DATASET_CONSUMER = "dataset_consumer" 

41 DATASET_PRODUCER = "dataset_producer" 

42 TEST_OF = "test_of" 

43 EXPOSES = "exposes" 

44 CALLS = "calls" 

45 

46 

47@dataclass 

48class Node: 

49 """Represents a node in the dependency graph.""" 

50 

51 id: str 

52 name: str 

53 type: NodeType 

54 owner: str = "" 

55 severity: NodeSeverity = NodeSeverity.MEDIUM 

56 description: str = "" 

57 metadata: Dict[str, Any] = field(default_factory=dict) 

58 created_at: datetime = field(default_factory=datetime.utcnow) 

59 updated_at: datetime = field(default_factory=datetime.utcnow) 

60 

61 # Cached computed properties 

62 upstream_count: int = 0 

63 downstream_count: int = 0 

64 upstream_nodes: Set[str] = field(default_factory=set) 

65 downstream_nodes: Set[str] = field(default_factory=set) 

66 

67 def __hash__(self): 

68 return hash(self.id) 

69 

70 def __eq__(self, other): 

71 if not isinstance(other, Node): 

72 return False 

73 return self.id == other.id 

74 

75 

76@dataclass 

77class Edge: 

78 """Represents an edge/relationship between nodes.""" 

79 

80 source: str 

81 target: str 

82 relationship_type: RelationshipType = RelationshipType.DEPENDS_ON 

83 strength: float = 1.0 # 0.0-1.0, indicates importance/weight 

84 metadata: Dict[str, Any] = field(default_factory=dict) 

85 created_at: datetime = field(default_factory=datetime.utcnow) 

86 

87 def __hash__(self): 

88 return hash((self.source, self.target, self.relationship_type)) 

89 

90 def __eq__(self, other): 

91 if not isinstance(other, Edge): 

92 return False 

93 return ( 

94 self.source == other.source 

95 and self.target == other.target 

96 and self.relationship_type == other.relationship_type 

97 ) 

98 

99 

100@dataclass 

101class DependencyGraph: 

102 """Represents complete dependency graph.""" 

103 

104 nodes: Dict[str, Node] = field(default_factory=dict) 

105 edges: List[Edge] = field(default_factory=list) 

106 version: str = "1.0.0" 

107 created_at: datetime = field(default_factory=datetime.utcnow) 

108 updated_at: datetime = field(default_factory=datetime.utcnow) 

109 metadata: Dict[str, Any] = field(default_factory=dict) 

110 

111 def add_node(self, node: Node) -> None: 

112 """Add a node to the graph.""" 

113 self.nodes[node.id] = node 

114 self.updated_at = datetime.utcnow() 

115 

116 def add_edge(self, edge: Edge) -> None: 

117 """Add an edge to the graph.""" 

118 self.edges.append(edge) 

119 

120 # Update node counts 

121 if edge.source in self.nodes: 121 ↛ 125line 121 didn't jump to line 125 because the condition on line 121 was always true

122 self.nodes[edge.source].downstream_nodes.add(edge.target) 

123 self.nodes[edge.source].downstream_count += 1 

124 

125 if edge.target in self.nodes: 125 ↛ 129line 125 didn't jump to line 129 because the condition on line 125 was always true

126 self.nodes[edge.target].upstream_nodes.add(edge.source) 

127 self.nodes[edge.target].upstream_count += 1 

128 

129 self.updated_at = datetime.utcnow() 

130 

131 def get_node(self, node_id: str) -> Optional[Node]: 

132 """Get node by ID.""" 

133 return self.nodes.get(node_id) 

134 

135 def get_edges_from(self, source_id: str) -> List[Edge]: 

136 """Get all outgoing edges from a node.""" 

137 return [e for e in self.edges if e.source == source_id] 

138 

139 def get_edges_to(self, target_id: str) -> List[Edge]: 

140 """Get all incoming edges to a node.""" 

141 return [e for e in self.edges if e.target == target_id] 

142 

143 def get_node_count_by_type(self) -> Dict[NodeType, int]: 

144 """Count nodes by type.""" 

145 counts = {} 

146 for node in self.nodes.values(): 

147 counts[node.type] = counts.get(node.type, 0) + 1 

148 return counts 

149 

150 def get_critical_nodes(self) -> List[Node]: 

151 """Get all critical severity nodes.""" 

152 return [n for n in self.nodes.values() if n.severity == NodeSeverity.CRITICAL] 

153 

154 def stats(self) -> Dict[str, Any]: 

155 """Get graph statistics.""" 

156 return { 

157 "node_count": len(self.nodes), 

158 "edge_count": len(self.edges), 

159 "node_types": self.get_node_count_by_type(), 

160 "critical_nodes": len(self.get_critical_nodes()), 

161 "version": self.version, 

162 "created_at": self.created_at.isoformat(), 

163 "updated_at": self.updated_at.isoformat(), 

164 } 

165 

166 

167# Analysis result dataclasses 

168 

169 

170@dataclass 

171class ImpactResult: 

172 """Result of impact analysis.""" 

173 

174 node_id: str 

175 impacted_nodes: List[str] 

176 impact_depth: int 

177 impact_score: float # 0.0-1.0 

178 by_severity: Dict[NodeSeverity, List[str]] = field(default_factory=dict) 

179 by_type: Dict[NodeType, List[str]] = field(default_factory=dict) 

180 metadata: Dict[str, Any] = field(default_factory=dict) 

181 

182 

183@dataclass 

184class BlastRadiusResult: 

185 """Result of blast radius analysis.""" 

186 

187 change_nodes: List[str] 

188 affected_nodes: List[str] 

189 blast_radius: int # Number of affected nodes 

190 blast_depth: int # Max depth of impact 

191 severity_distribution: Dict[NodeSeverity, int] = field(default_factory=dict) 

192 risk_level: str = "low" # low, medium, high, critical 

193 deployable: bool = True 

194 metadata: Dict[str, Any] = field(default_factory=dict) 

195 

196 

197@dataclass 

198class RiskScoreResult: 

199 """Result of risk scoring.""" 

200 

201 node_id: str 

202 risk_score: float # 0.0-10.0 

203 components: Dict[str, float] = field(default_factory=dict) 

204 factors: List[str] = field(default_factory=list) 

205 severity: NodeSeverity = NodeSeverity.MEDIUM 

206 metadata: Dict[str, Any] = field(default_factory=dict) 

207 

208 

209@dataclass 

210class DriftDetectionResult: 

211 """Result of drift detection analysis.""" 

212 

213 detected_drifts: List[Dict[str, Any]] = field(default_factory=list) 

214 drift_count: int = 0 

215 affected_nodes: List[str] = field(default_factory=list) 

216 severity: NodeSeverity = NodeSeverity.LOW 

217 details: List[str] = field(default_factory=list) 

218 metadata: Dict[str, Any] = field(default_factory=dict)