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.
- assets/authentication.png +0 -0
- assets/full_viz.png +0 -0
- assets/hero.png +0 -0
- assets/hero_with_properties.png +0 -0
- assets/jupyter-spin-up.png +0 -0
- assets/load_ext.png +0 -0
- assets/mock_data_result.png +0 -0
- assets/new_notebook.png +0 -0
- assets/notebook_package_load.png +0 -0
- assets/query_graph.png +0 -0
- assets/sample_jupyter.png +0 -0
- spanner_graph_notebook-1.0.0.dist-info/LICENSE +202 -0
- spanner_graph_notebook-1.0.0.dist-info/METADATA +169 -0
- spanner_graph_notebook-1.0.0.dist-info/RECORD +40 -0
- spanner_graph_notebook-1.0.0.dist-info/WHEEL +5 -0
- spanner_graph_notebook-1.0.0.dist-info/top_level.txt +2 -0
- spanner_graphs/__init__.py +19 -0
- spanner_graphs/conversion.py +299 -0
- spanner_graphs/database.py +199 -0
- spanner_graphs/graph_entities.py +405 -0
- spanner_graphs/graph_mock_data.csv +51 -0
- spanner_graphs/graph_mock_schema.json +681 -0
- spanner_graphs/magics.py +211 -0
- spanner_graphs/schema_manager.py +60 -0
- templates/assets/images/graph-bg.svg +4 -0
- templates/spanner-graph/models/edge.js +77 -0
- templates/spanner-graph/models/graph-object.js +64 -0
- templates/spanner-graph/models/node.js +77 -0
- templates/spanner-graph/models/schema.js +327 -0
- templates/spanner-graph/spanner-config.js +304 -0
- templates/spanner-graph/spanner-store.js +382 -0
- templates/spanner-graph/visualization/spanner-forcegraph.js +1380 -0
- templates/spanner-graph/visualization/spanner-sidebar.js +904 -0
- templates/template-spannergraph.html +210 -0
- tests/__init__.py +13 -0
- tests/conversion_test.py +163 -0
- tests/database_test.py +62 -0
- tests/graph_entities_test.py +124 -0
- tests/schema_manager_test.py +115 -0
- 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
|