spanner-graph-notebook 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (40) hide show
  1. assets/authentication.png +0 -0
  2. assets/full_viz.png +0 -0
  3. assets/hero.png +0 -0
  4. assets/hero_with_properties.png +0 -0
  5. assets/jupyter-spin-up.png +0 -0
  6. assets/load_ext.png +0 -0
  7. assets/mock_data_result.png +0 -0
  8. assets/new_notebook.png +0 -0
  9. assets/notebook_package_load.png +0 -0
  10. assets/query_graph.png +0 -0
  11. assets/sample_jupyter.png +0 -0
  12. spanner_graph_notebook-1.0.0.dist-info/LICENSE +202 -0
  13. spanner_graph_notebook-1.0.0.dist-info/METADATA +169 -0
  14. spanner_graph_notebook-1.0.0.dist-info/RECORD +40 -0
  15. spanner_graph_notebook-1.0.0.dist-info/WHEEL +5 -0
  16. spanner_graph_notebook-1.0.0.dist-info/top_level.txt +2 -0
  17. spanner_graphs/__init__.py +19 -0
  18. spanner_graphs/conversion.py +299 -0
  19. spanner_graphs/database.py +199 -0
  20. spanner_graphs/graph_entities.py +405 -0
  21. spanner_graphs/graph_mock_data.csv +51 -0
  22. spanner_graphs/graph_mock_schema.json +681 -0
  23. spanner_graphs/magics.py +211 -0
  24. spanner_graphs/schema_manager.py +60 -0
  25. templates/assets/images/graph-bg.svg +4 -0
  26. templates/spanner-graph/models/edge.js +77 -0
  27. templates/spanner-graph/models/graph-object.js +64 -0
  28. templates/spanner-graph/models/node.js +77 -0
  29. templates/spanner-graph/models/schema.js +327 -0
  30. templates/spanner-graph/spanner-config.js +304 -0
  31. templates/spanner-graph/spanner-store.js +382 -0
  32. templates/spanner-graph/visualization/spanner-forcegraph.js +1380 -0
  33. templates/spanner-graph/visualization/spanner-sidebar.js +904 -0
  34. templates/template-spannergraph.html +210 -0
  35. tests/__init__.py +13 -0
  36. tests/conversion_test.py +163 -0
  37. tests/database_test.py +62 -0
  38. tests/graph_entities_test.py +124 -0
  39. tests/schema_manager_test.py +115 -0
  40. tests/test_notebook.json +23 -0
@@ -0,0 +1,299 @@
1
+ # Copyright 2024 Google LLC
2
+
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+
7
+ # https://www.apache.org/licenses/LICENSE-2.0
8
+
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ """
16
+ This module contains implementation to convert columns from
17
+ a database column into usable data for building a graph
18
+ """
19
+
20
+ from __future__ import annotations
21
+ from typing import Any, List, Dict, Set, Tuple
22
+ import json
23
+ from enum import Enum, auto
24
+
25
+ from google.cloud.spanner_v1.types import TypeCode, StructType
26
+ import numpy as np
27
+ import networkx as nx
28
+
29
+ from spanner_graphs.graph_entities import Node, Edge
30
+ from spanner_graphs.schema_manager import SchemaManager
31
+
32
+
33
+ # Sizing rule for the nodes
34
+ class SizeMode(Enum):
35
+ STATIC = auto()
36
+ CARDINALITY = auto()
37
+ PROPERTY = auto()
38
+
39
+ def __str__(self):
40
+ return self.name.lower()
41
+
42
+ @classmethod
43
+ def from_string(cls, s: str):
44
+ try:
45
+ return cls[s.upper()]
46
+ except KeyError:
47
+ raise ValueError(f"'{s}' is not a valid {cls.__name__}")
48
+
49
+
50
+ # Map spanner TypeCode to numpy types
51
+ _TYPE_CODE_MAP = {
52
+ TypeCode.JSON: np.object_,
53
+ TypeCode.ARRAY: np.object_,
54
+ }
55
+
56
+ NODE_DEGREE = "value"
57
+
58
+
59
+ def _column_to_native_numpy(column: List[Any],
60
+ datatype: TypeCode,
61
+ array_type: TypeCode = None) -> np.ndarray:
62
+ """Convert Spanner column (list) to a numpy array, with appropriate dtype.
63
+
64
+ Args:
65
+ column: List of Spanner column values
66
+ datatype: TypeCode representing the type of all values in `column`
67
+ array_type: If `datatype == Typecode.ARRAY`, the type of the values
68
+ store inside each array.
69
+
70
+ Returns:
71
+ np.ndarray of the values in `column`
72
+
73
+ Raises:
74
+ ValueError: If any of the fields contain unsupported types.
75
+ """
76
+ if datatype == TypeCode.JSON:
77
+ flattenned_data = []
78
+
79
+ for x in column:
80
+ if isinstance(x, list):
81
+ flattenned_data.extend(x)
82
+ elif isinstance(x, dict):
83
+ array_value = getattr(x, '_array_value', None)
84
+ if array_value is not None:
85
+ flattenned_data.extend(array_value)
86
+ else:
87
+ flattenned_data.append(x)
88
+ else:
89
+ flattenned_data.append(json.loads(x))
90
+
91
+ return np.array(
92
+ flattenned_data,
93
+ dtype=_TYPE_CODE_MAP[datatype],
94
+ )
95
+
96
+ if datatype == TypeCode.ARRAY:
97
+ if array_type == TypeCode.JSON:
98
+ return np.array(
99
+ [
100
+ np.array(
101
+ [
102
+ y if isinstance(y, dict) else json.loads(y)
103
+ for y in x
104
+ ],
105
+ dtype=_TYPE_CODE_MAP[array_type],
106
+ ) for x in column
107
+ ],
108
+ dtype=_TYPE_CODE_MAP[datatype],
109
+ )
110
+
111
+ raise ValueError(
112
+ f"Only JSON and array of JSON are allowed, got type: {datatype.name}")
113
+
114
+
115
+ def columns_to_native_numpy(
116
+ data: Dict[str, List[Any]], fields: List[StructType.Field]
117
+ ) -> Tuple[Dict[str, np.ndarray], List[str]]:
118
+ """Cast values in all columns to corresponding typed Numpy arrays.
119
+
120
+ Args:
121
+ data: dict whose values are all lists of the same length
122
+ fields: One entry per column `data`, column name should
123
+ match the key in `data` and column type should match
124
+ the desired type of values in the corresponding list in `data`
125
+
126
+ Returns:
127
+ Tuple[Dict[str, np.ndarray], List[str]]: A tuple containing a dict with
128
+ the same shape and keys as `data`, but with columns converted to
129
+ `np.array`s with a `dtype` based on the type in `fields`.
130
+ It also returns the list of ignored_columns that do not match the
131
+ supported json or array of json types.
132
+ """
133
+ output = {}
134
+ ignored_columns = []
135
+
136
+ for field in fields:
137
+ try:
138
+ column_name = field.name
139
+ if field.type_.code == TypeCode.JSON or (
140
+ field.type_.code == TypeCode.ARRAY
141
+ and field.type_.array_element_type == TypeCode.JSON):
142
+ output[column_name] = _column_to_native_numpy(
143
+ data[column_name],
144
+ field.type_.code,
145
+ (field.type_.array_element_type
146
+ if field.type_.array_element_type else None),
147
+ )
148
+ else:
149
+ # Log the column name that does not match our expectation
150
+ ignored_columns.append(field.name)
151
+ except ValueError:
152
+ # We simply advance to the next, we return any valid
153
+ # that we need for our purpose.
154
+ # When no data matches our expectation, `len(output) == 0`
155
+ pass
156
+
157
+ return output, ignored_columns
158
+
159
+
160
+ def prepare_data_for_graphing(
161
+ incoming: Dict[str, np.ndarray],
162
+ schema_json: dict = None,
163
+ apply_sort=False,
164
+ size_mode: SizeMode = SizeMode.CARDINALITY,
165
+ size_property: str = None,
166
+ node_display_props: dict[str, str] = None,
167
+ edge_display_props: dict[str, str] = None) -> nx.DiGraph:
168
+ """
169
+ Prepare the output from column conversions which may contain duplicates
170
+ by using a `set` and also ensuring both JSON and Array of JSON values
171
+ are flattened, returning a single list of json data.
172
+
173
+ Args:
174
+ incoming: A dictionay of str to np.ndarray values, where each np.ndarray
175
+ values may be a nested array.
176
+ schema_json: The property graph metadata JSON schema.
177
+ apply_sort: A boolean that tells if the data should be sorted before
178
+ the output graph is generated for deterministic output.
179
+ size_mode: A SizeMode enum indicating the rule for sizing the nodes
180
+ node_display_props: A dictionary that maps from label type
181
+ to property to display for nodes.
182
+ edge_display_props: A dictionary that maps from label type
183
+ to property to display for edges.
184
+ Returns:
185
+ A networkx graph with all nodes and edges add from the
186
+ deduplicated input
187
+ """
188
+
189
+ schema_manager = SchemaManager(schema_json)
190
+ unique_items: Set[str] = set()
191
+
192
+ for column in incoming.values():
193
+ for item in column:
194
+ if isinstance(item, np.ndarray): # Flatten array of JSON
195
+ for sub_item in item:
196
+ unique_items.add(json.dumps(sub_item))
197
+ else:
198
+ unique_items.add(json.dumps(item))
199
+
200
+ unique_json_list = [json.loads(item) for item in unique_items]
201
+ if apply_sort:
202
+ unique_json_list = sorted(
203
+ unique_json_list,
204
+ key=lambda x: (
205
+ x["kind"],
206
+ x.get("identifier", x.get("source_node_identifier", "")),
207
+ ),
208
+ )
209
+
210
+ g = nx.MultiDiGraph()
211
+
212
+ node_mapping = {}
213
+ edge_counter = 1
214
+
215
+ for item in unique_json_list:
216
+ if "kind" not in item:
217
+ continue
218
+ if item["kind"] != "node":
219
+ continue
220
+
221
+ if not Node.is_valid_node_json(item):
222
+ continue
223
+
224
+ node = Node.from_json(item)
225
+ node.key_property_names = schema_manager.get_key_property_names(node)
226
+ if node.identifier not in node_mapping:
227
+ node_mapping[node.identifier] = len(node_mapping) + 1
228
+ node.decide_label_string(node_display_props)
229
+ node.add_to_graph(g, node_mapping)
230
+
231
+ if size_mode == SizeMode.PROPERTY:
232
+ if size_property is None:
233
+ raise ValueError(
234
+ "size_property must be specified when using SizeMode.PROPERTY"
235
+ )
236
+
237
+ value = node.properties.get(size_property)
238
+ if value is not None:
239
+ try:
240
+ numeric_value = float(value)
241
+ g.nodes[node_mapping[
242
+ node.identifier]][NODE_DEGREE] = numeric_value
243
+ except ValueError:
244
+ print(
245
+ f"Warning: Property '{size_property}' for node {node.identifier} is not numeric. Using default size."
246
+ )
247
+ else:
248
+ print(
249
+ f"Warning: Property '{size_property}' not found for node {node.identifier}. Using default size."
250
+ )
251
+
252
+ # Second pass to find source and destination nodes
253
+ # from edge data. They may not be added to the graph in the
254
+ # first pass node calculation above.
255
+ for item in unique_json_list:
256
+ if "kind" not in item:
257
+ continue
258
+ if item["kind"] != "edge":
259
+ continue
260
+
261
+ if not Edge.is_valid_edge_json(item):
262
+ continue
263
+ edge = Edge.from_json(item)
264
+ if edge.source not in node_mapping:
265
+ src_id = len(node_mapping) + 1
266
+ node_mapping[edge.source] = src_id
267
+ src_node = Node(edge.source, [], {"id": src_id})
268
+ src_node.add_to_graph(g, node_mapping)
269
+ if edge.destination not in node_mapping:
270
+ dst_id = len(node_mapping) + 1
271
+ node_mapping[edge.destination] = dst_id
272
+ dst_node = Node(edge.destination, [], {"id": dst_id})
273
+ dst_node.add_to_graph(g, node_mapping)
274
+
275
+ for item in unique_json_list:
276
+ if "kind" not in item:
277
+ continue
278
+ if item["kind"] != "edge":
279
+ continue
280
+
281
+ if not Edge.is_valid_edge_json(item):
282
+ continue
283
+
284
+ edge = Edge.from_json(item)
285
+ edge.decide_label_string(edge_display_props)
286
+ numerical_id = (len(node_mapping) + 1) + edge_counter
287
+ edge_counter += 1
288
+ edge.add_to_graph(g, node_mapping, numerical_id)
289
+
290
+ if size_mode == SizeMode.CARDINALITY:
291
+ # Calculate the in-degree and out-degree using NetworkX functions
292
+ in_degrees = dict(g.in_degree())
293
+ out_degress = dict(g.out_degree())
294
+
295
+ for node_id in g.nodes():
296
+ node_size = in_degrees[node_id] + out_degress[node_id]
297
+ g.nodes[node_id][NODE_DEGREE] = node_size
298
+
299
+ return g
@@ -0,0 +1,199 @@
1
+ # Copyright 2024 Google LLC
2
+
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+
7
+ # https://www.apache.org/licenses/LICENSE-2.0
8
+
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ """
16
+ This module contains implementation for talking to spanner database
17
+ via snapshot queries.
18
+ """
19
+
20
+ from __future__ import annotations
21
+ from typing import Any, Dict, List, Tuple
22
+ import json
23
+ import os
24
+ import csv
25
+
26
+ from google.cloud import spanner
27
+ from google.cloud.spanner_v1.types import StructType, TypeCode, Type
28
+
29
+
30
+ class SpannerDatabase:
31
+ """The spanner class holding the database connection"""
32
+ def __init__(self, project_id: str, instance_id: str,
33
+ database_id: str) -> None:
34
+ self.client = spanner.Client(project=project_id)
35
+ self.instance = self.client.instance(instance_id)
36
+ self.database = self.instance.database(database_id)
37
+
38
+ # In order to ensure that the database connection was properly
39
+ # created and that customers won't be confused by connectivity
40
+ # errors happening mysteriously, let's firstly run a dummy
41
+ # query on the database.
42
+ sql = 'SELECT table_name FROM information_schema.tables WHERE table_schema = ""'
43
+ _ = self.execute_query(sql, None, is_test_query=True)
44
+
45
+ def __repr__(self) -> str:
46
+ return (f"<SpannerDatabase["
47
+ f"project:{self.client.project_name},"
48
+ f"instance{self.instance.name},"
49
+ f"db:{self.database.name}]")
50
+
51
+ def _extract_graph_name(self, query: str) -> str:
52
+ words = query.strip().split()
53
+ if len(words) < 3:
54
+ raise ValueError("invalid query: must contain at least (GRAPH, graph_name and query)")
55
+
56
+ if words[0].upper() != "GRAPH":
57
+ raise ValueError("invalid query: GRAPH must be the first word")
58
+
59
+ return words[1]
60
+
61
+ def _get_schema_for_graph(self, graph_query: str):
62
+ try:
63
+ graph_name = self._extract_graph_name(graph_query)
64
+ except ValueError as e:
65
+ print(f"Error extracting graph name: {str(e)}")
66
+ return None
67
+
68
+ with self.database.snapshot() as snapshot:
69
+ schema_query = """
70
+ SELECT property_graph_name, property_graph_metadata_json
71
+ FROM information_schema.property_graphs
72
+ WHERE property_graph_name = @graph_name
73
+ """
74
+ params = {"graph_name": graph_name}
75
+ param_type = {"graph_name": spanner.param_types.STRING}
76
+
77
+ result = snapshot.execute_sql(schema_query, params=params, param_types=param_type)
78
+ schema_rows = list(result)
79
+
80
+ if schema_rows:
81
+ return schema_rows[0][1]
82
+ else:
83
+ print(f"No schema found for graph: {graph_name}")
84
+ return None
85
+
86
+ def execute_query(
87
+ self,
88
+ query: str,
89
+ limit: int = None,
90
+ is_test_query: bool = False,
91
+ ):
92
+ """
93
+ This method executes the provided `query`
94
+
95
+ Args:
96
+ query: The SQL query to execute against the database
97
+ limit: An optional limit for the number of rows to return
98
+
99
+ Returns:
100
+ A tuple containing:
101
+ - Dict[str, List[Any]]: A dict where each key is a field name
102
+ returned in the query and the list contains all items of the same
103
+ type found for the given field.
104
+ - A list of StructType.Fields representing the fields in the result set.
105
+ - A list of rows as returned by the query execution.
106
+ """
107
+ self.schema_json = None
108
+ if not is_test_query:
109
+ self.schema_json = self._get_schema_for_graph(query)
110
+
111
+ with self.database.snapshot() as snapshot:
112
+ params = None
113
+ if limit and limit > 0:
114
+ params = dict(limit=limit)
115
+
116
+ results = snapshot.execute_sql(query, params=params)
117
+ rows = list(results)
118
+ fields: List[StructType.Field] = results.fields
119
+
120
+ data = {field.name: [] for field in fields}
121
+
122
+ if len(fields) == 0:
123
+ return data, fields, rows
124
+
125
+ for row in rows:
126
+ for field, value in zip(fields, row):
127
+ data[field.name].append(value)
128
+
129
+ return data, fields, rows, self.schema_json
130
+
131
+
132
+ class MockSpannerResult:
133
+
134
+ def __init__(self, file_path: str):
135
+ self.file_path = file_path
136
+ self.fields: List[StructType] = []
137
+ self._rows: List[List[Any]] = []
138
+ self._load_data()
139
+
140
+ def _load_data(self):
141
+ with open(self.file_path, "r", encoding="utf-8") as csvfile:
142
+ csv_reader = csv.reader(csvfile)
143
+ headers = next(csv_reader)
144
+ self.fields = [
145
+ StructType.Field(name=header, type_=Type(code=TypeCode.JSON))
146
+ for header in headers
147
+ ]
148
+
149
+ for row in csv_reader:
150
+ parsed_row = []
151
+ for value in row:
152
+ try:
153
+ js = bytes(value, "utf-8").decode("unicode_escape")
154
+ parsed_row.append(json.loads(js))
155
+ except json.JSONDecodeError:
156
+ pass
157
+ self._rows.append(parsed_row)
158
+
159
+ def __iter__(self):
160
+ return iter(self._rows)
161
+
162
+
163
+ class MockSpannerDatabase:
164
+ """Mock database class"""
165
+
166
+ def __init__(self):
167
+ dirname = os.path.dirname(__file__)
168
+ self.graph_csv_path = os.path.join(
169
+ dirname, "graph_mock_data.csv")
170
+ self.schema_json_path = os.path.join(
171
+ dirname, "graph_mock_schema.json")
172
+ self.schema_json: dict = {}
173
+
174
+ def execute_query(
175
+ self,
176
+ _: str,
177
+ limit: int = None
178
+ ) -> Tuple[Dict[str, List[Any]], List[StructType.Field], List, str]:
179
+ """Mock execution of query"""
180
+
181
+ # Before the actual query we fetch the schema as well
182
+ with open(self.schema_json_path, "r", encoding="utf-8") as js:
183
+ self.schema_json = json.load(js)
184
+
185
+ results = MockSpannerResult(self.graph_csv_path)
186
+ fields: List[StructType.Field] = results.fields
187
+ rows = list(results)
188
+ data = {field.name: [] for field in fields}
189
+
190
+ if len(fields) == 0:
191
+ return data, fields, rows
192
+
193
+ for i, row in enumerate(results):
194
+ if limit is not None and i >= limit:
195
+ break
196
+ for field, value in zip(fields, row):
197
+ data[field.name].append(value)
198
+
199
+ return data, fields, rows, self.schema_json