forked from vitali87/code-graph-rag
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathgraph_loader.py
More file actions
154 lines (125 loc) · 5.21 KB
/
Copy pathgraph_loader.py
File metadata and controls
154 lines (125 loc) · 5.21 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
import json
from collections import Counter, defaultdict
from pathlib import Path
from loguru import logger
from . import constants as cs
from . import exceptions as ex
from . import logs as ls
from .decorators import ensure_loaded
from .models import GraphNode, GraphRelationship
from .types_defs import GraphData, GraphMetadata, GraphSummary, PropertyValue
class GraphLoader:
def __init__(self, file_path: str):
self.file_path = Path(file_path)
self._data: GraphData | None = None
self._nodes: list[GraphNode] | None = None
self._relationships: list[GraphRelationship] | None = None
self._nodes_by_id: dict[int, GraphNode] = {}
self._nodes_by_label: defaultdict[str, list[GraphNode]] = defaultdict(list)
self._outgoing_rels: defaultdict[int, list[GraphRelationship]] = defaultdict(
list
)
self._incoming_rels: defaultdict[int, list[GraphRelationship]] = defaultdict(
list
)
self._property_indexes: dict[str, dict[PropertyValue, list[GraphNode]]] = {}
def _ensure_loaded(self) -> None:
if self._data is None:
self.load()
def load(self) -> None:
if not self.file_path.exists():
raise FileNotFoundError(ex.GRAPH_FILE_NOT_FOUND.format(path=self.file_path))
logger.info(ls.LOADING_GRAPH.format(path=self.file_path))
with open(self.file_path, encoding=cs.ENCODING_UTF8) as f:
self._data = json.load(f)
if self._data is None:
raise RuntimeError(ex.FAILED_TO_LOAD_DATA)
self._nodes = []
for node_data in self._data[cs.KEY_NODES]:
node = GraphNode(
node_id=node_data[cs.KEY_NODE_ID],
labels=node_data[cs.KEY_LABELS],
properties=node_data[cs.KEY_PROPERTIES],
)
self._nodes.append(node)
self._nodes_by_id[node.node_id] = node
for label in node.labels:
self._nodes_by_label[label].append(node)
self._relationships = []
for rel_data in self._data[cs.KEY_RELATIONSHIPS]:
rel = GraphRelationship(
from_id=rel_data[cs.KEY_FROM_ID],
to_id=rel_data[cs.KEY_TO_ID],
type=rel_data[cs.KEY_TYPE],
properties=rel_data[cs.KEY_PROPERTIES],
)
self._relationships.append(rel)
self._outgoing_rels[rel.from_id].append(rel)
self._incoming_rels[rel.to_id].append(rel)
logger.info(
ls.LOADED_GRAPH.format(
nodes=len(self._nodes), relationships=len(self._relationships)
)
)
def _build_property_index(self, property_name: str) -> None:
if property_name in self._property_indexes:
return
index: defaultdict[PropertyValue, list[GraphNode]] = defaultdict(list)
for node in self.nodes:
value = node.properties.get(property_name)
if value is not None:
index[value].append(node)
self._property_indexes[property_name] = dict(index)
@property
@ensure_loaded
def nodes(self) -> list[GraphNode]:
assert self._nodes is not None, ex.NODES_NOT_LOADED
return self._nodes
@property
@ensure_loaded
def relationships(self) -> list[GraphRelationship]:
assert self._relationships is not None, ex.RELATIONSHIPS_NOT_LOADED
return self._relationships
@property
@ensure_loaded
def metadata(self) -> GraphMetadata:
assert self._data is not None, ex.DATA_NOT_LOADED
return self._data[cs.KEY_METADATA]
@ensure_loaded
def find_nodes_by_label(self, label: str) -> list[GraphNode]:
return self._nodes_by_label.get(label, [])
@ensure_loaded
def find_node_by_property(
self, property_name: str, value: PropertyValue
) -> list[GraphNode]:
self._build_property_index(property_name)
return self._property_indexes[property_name].get(value, [])
@ensure_loaded
def get_node_by_id(self, node_id: int) -> GraphNode | None:
return self._nodes_by_id.get(node_id)
def get_relationships_for_node(self, node_id: int) -> list[GraphRelationship]:
return self.get_outgoing_relationships(
node_id
) + self.get_incoming_relationships(node_id)
@ensure_loaded
def get_outgoing_relationships(self, node_id: int) -> list[GraphRelationship]:
return self._outgoing_rels.get(node_id, [])
@ensure_loaded
def get_incoming_relationships(self, node_id: int) -> list[GraphRelationship]:
return self._incoming_rels.get(node_id, [])
def summary(self) -> GraphSummary:
node_labels = {
label: len(nodes) for label, nodes in self._nodes_by_label.items()
}
relationship_types = dict(Counter(rel.type for rel in self.relationships))
return GraphSummary(
total_nodes=len(self.nodes),
total_relationships=len(self.relationships),
node_labels=node_labels,
relationship_types=relationship_types,
metadata=self.metadata,
)
def load_graph(file_path: str) -> GraphLoader:
loader = GraphLoader(file_path)
loader.load()
return loader