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
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-19 20:43 +0530
1"""Data models for Dependency Intelligence Engine."""
3from dataclasses import dataclass, field
4from datetime import datetime
5from enum import Enum
6from typing import Any, Dict, List, Optional, Set
9class NodeType(Enum):
10 """Types of nodes in dependency graph."""
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"
26class NodeSeverity(Enum):
27 """Severity levels for nodes."""
29 CRITICAL = "critical"
30 HIGH = "high"
31 MEDIUM = "medium"
32 LOW = "low"
35class RelationshipType(Enum):
36 """Types of edges/relationships."""
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"
47@dataclass
48class Node:
49 """Represents a node in the dependency graph."""
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)
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)
67 def __hash__(self):
68 return hash(self.id)
70 def __eq__(self, other):
71 if not isinstance(other, Node):
72 return False
73 return self.id == other.id
76@dataclass
77class Edge:
78 """Represents an edge/relationship between nodes."""
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)
87 def __hash__(self):
88 return hash((self.source, self.target, self.relationship_type))
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 )
100@dataclass
101class DependencyGraph:
102 """Represents complete dependency graph."""
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)
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()
116 def add_edge(self, edge: Edge) -> None:
117 """Add an edge to the graph."""
118 self.edges.append(edge)
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
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
129 self.updated_at = datetime.utcnow()
131 def get_node(self, node_id: str) -> Optional[Node]:
132 """Get node by ID."""
133 return self.nodes.get(node_id)
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]
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]
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
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]
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 }
167# Analysis result dataclasses
170@dataclass
171class ImpactResult:
172 """Result of impact analysis."""
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)
183@dataclass
184class BlastRadiusResult:
185 """Result of blast radius analysis."""
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)
197@dataclass
198class RiskScoreResult:
199 """Result of risk scoring."""
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)
209@dataclass
210class DriftDetectionResult:
211 """Result of drift detection analysis."""
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)