-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathserializer.py
More file actions
138 lines (124 loc) · 6.42 KB
/
Copy pathserializer.py
File metadata and controls
138 lines (124 loc) · 6.42 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
import argparse
import logging
import os
import re
import shutil
import subprocess
import sys
from version import _get_lbug_version
base_dir = os.path.dirname(os.path.realpath(__file__))
def serialize(lbug_exec_path, dataset_name, dataset_path, serialized_graph_path, benchmark_copy_log_dir,
single_thread: bool = False):
bin_version = _get_lbug_version()
if os.path.exists(os.path.join(serialized_graph_path, 'version.txt')):
with open(os.path.join(serialized_graph_path, 'version.txt'), encoding="utf-8") as f:
dataset_version = f.readline().strip()
if dataset_version == bin_version:
logging.info(
'Dataset %s has version of %s, which matches the database version, skip serializing', dataset_name,
bin_version)
return
else:
logging.info(
'Dataset %s has version of %s, which does not match the database version %s, serializing dataset...',
dataset_name, dataset_version, bin_version)
else:
logging.info('Dataset %s does not exist or does not have a version file, serializing dataset...', dataset_name)
shutil.rmtree(serialized_graph_path, ignore_errors=True)
os.mkdir(serialized_graph_path)
serialize_queries = []
if single_thread:
serialize_queries.append("CALL THREADS=1")
if os.path.exists(os.path.join(dataset_path, 'schema.cypher')):
with open(os.path.join(dataset_path, 'schema.cypher'), 'r') as f:
serialize_queries += f.readlines()
with open(os.path.join(dataset_path, 'copy.cypher'), 'r') as f:
copy_lines = f.readlines()
# Fix relative paths in copy.cypher
for line in copy_lines:
# Replace quoted paths with absolute paths
def replace_path(match):
path = match.group(1)
if not os.path.isabs(path):
return '"' + os.path.join(dataset_path, path) + '"'
return match.group(0)
fixed_line = re.sub(r'"([^"]*)"', replace_path, line)
serialize_queries.append(fixed_line.strip())
else:
with open(os.path.join(base_dir, 'serialize.cypher'), 'r') as f:
serialize_queries += f.readlines()
serialize_queries = [q.strip().replace('{}', dataset_path)
for q in serialize_queries]
serialize_queries = [q for q in serialize_queries if q]
table_types = {}
for s in serialize_queries:
logging.info('Executing query: %s', s)
create_match = re.match(r'create\s+(.+?)\s+table\s+(.+?)\s*\(', s, re.IGNORECASE)
copy_match = re.match(r'copy\s+(.+?)\s+from', s, re.IGNORECASE)
# Run lbug shell one query at a time. This ensures a new process is
# created for each query to avoid memory leaks.
stdout = sys.stdout if create_match or not benchmark_copy_log_dir else subprocess.PIPE
db_path = os.path.join(serialized_graph_path, 'db.lbdb')
process = subprocess.Popen([lbug_exec_path, db_path],
stdin=subprocess.PIPE, stdout=stdout, encoding="utf-8")
process.stdin.write(s)
process.stdin.close()
if create_match:
table_types[create_match.group(2)] = create_match.group(1).lower()
elif copy_match:
filename = table_types[copy_match.group(1)] + '-' + copy_match.group(1).replace('_', '-') + '_log.txt'
if benchmark_copy_log_dir:
os.makedirs(benchmark_copy_log_dir, exist_ok=True)
with open(os.path.join(benchmark_copy_log_dir, filename), 'a', encoding="utf-8") as f:
for line in process.stdout.readlines():
print(line, end="", flush=True)
print(line, file=f)
elif s == "CALL THREADS=1":
pass
else:
raise RuntimeError(f"Unrecognized query {s}")
process.wait()
if process.returncode != 0:
raise RuntimeError(f'Error {process.returncode} executing query: {s}')
with open(os.path.join(serialized_graph_path, 'version.txt'), 'w', encoding="utf-8") as f:
print(bin_version, file=f)
if __name__ == '__main__':
logging.basicConfig(level=logging.INFO)
parser = argparse.ArgumentParser(description='Serializes dataset to a lbug database')
parser.add_argument("dataset_name", help="Name of the dataset for display purposes")
parser.add_argument("dataset_path", help="Input path of the dataset to serialize")
parser.add_argument("serialized_graph_path",
help="Output path of the database. Will be created if it does not exist already")
parser.add_argument("benchmark_copy_log_dir", help="Optional directory to store copy logs", nargs="?")
parser.add_argument("--single-thread",
help="If true, copy single threaded, which makes the results more reproducible",
action="store_true")
parser.add_argument("--lbug-shell-mode",
help="debug, release or relwithdebinfo",
default="release")
default_mode = "release"
if sys.platform == "win32":
default_lbug_exec_path = os.path.join(
base_dir, '..', 'build', default_mode, 'tools', 'shell', 'lbug_shell')
else:
default_lbug_exec_path = os.path.join(
base_dir, '..', 'build', default_mode, 'tools', 'shell', 'lbug')
parser.add_argument("--lbug-shell",
help="Path of the lbug shell executable. Defaults to the path as built in the default release build directory",
default=default_lbug_exec_path)
args = parser.parse_args()
if args.lbug_shell == default_lbug_exec_path:
mode = args.lbug_shell_mode
if sys.platform == "win32":
args.lbug_shell = os.path.join(base_dir, '..', 'build', mode, 'tools', 'shell', 'lbug_shell')
else:
args.lbug_shell = os.path.join(base_dir, '..', 'build', mode, 'tools', 'shell', 'lbug')
try:
serialize(args.lbug_shell, args.dataset_name, args.dataset_path, args.serialized_graph_path,
args.benchmark_copy_log_dir, args.single_thread)
except Exception as e:
logging.error(f'Error serializing dataset {args.dataset_name}')
raise e
finally:
shutil.rmtree(os.path.join(base_dir, 'history.txt'),
ignore_errors=True)