Multi-query support

This commit is contained in:
2022-08-09 22:00:46 +08:00
parent 70f29e028d
commit 5699638520
12 changed files with 184 additions and 61 deletions
+1 -1
View File
@@ -25,7 +25,7 @@ def exec(stmts, cxt = None, keep = False):
generate(s, cxt)
else:
generate(stmts_stmts, cxt)
cxt.print(cxt.sql)
cxt.print(cxt.queries)
return cxt
__all__ = ["initialize", "generate", "exec", "saved_cxt"]
+36 -10
View File
@@ -35,7 +35,10 @@ class ast_node:
name = 'null'
def init(self, _):
if self.parent is None:
self.context.sql_begin()
self.add(self.__class__.name.upper())
def produce(self, _):
pass
def spawn(self, _):
@@ -44,7 +47,7 @@ class ast_node:
def consume(self, _):
if self.parent is None:
self.emit(self.sql+';\n')
self.context.sql_end()
from reconstruct.expr import expr, fastscan
@@ -54,11 +57,17 @@ class projection(ast_node):
first_order = 'select'
def init(self, _):
# skip default init
pass
def produce(self, node):
p = node['select']
self.projections = p if type(p) is list else [p]
self.add('SELECT')
if self.parent is None:
self.context.sql_begin()
self.postproc_fname = 'dll_' + base62uuid(6)
self.context.postproc_begin(self.postproc_fname)
def spawn(self, node):
self.datasource = None # datasource is Join instead of TableInfo
@@ -245,6 +254,11 @@ class projection(ast_node):
self.outfile.finalize()
self.context.emitc(f'puts("done.");')
if self.parent is None:
self.context.sql_end()
self.context.postproc_end(self.postproc_fname)
class orderby(ast_node):
name = 'order by'
def produce(self, node):
@@ -314,15 +328,16 @@ class scan(ast_node):
class groupby_c(ast_node):
name = '_groupby'
def init(self, _):
def init(self, node : List[Tuple[expr, Set[ColRef]]]):
self.proj : projection = self.parent
return super().init(_)
self.glist : List[Tuple[expr, Set[ColRef]]] = node
return super().init(node)
def produce(self, node : List[Tuple[expr, Set[ColRef]]]):
self.context.headers.add('"./server/hasher.h"')
self.context.headers.add('unordered_map')
self.group = 'g' + base62uuid(7)
self.group_type = 'record_type' + base62uuid(7)
self.datasource = self.parent.datasource
self.datasource = self.proj.datasource
self.scanner = None
self.datasource.rec = set()
@@ -330,7 +345,7 @@ class groupby_c(ast_node):
g_contents_list = []
first_col = ''
for g in node:
for g in self.glist:
e = expr(self, g[0].node, c_code=True)
g_str = e.eval(c_code = True, y = lambda c: self.proj.pyname2cname[c])
# if v is compound expr, create tmp cols
@@ -345,7 +360,7 @@ class groupby_c(ast_node):
self.context.emitc(f'typedef record<{",".join(g_contents_decltype)}> {self.group_type};')
self.context.emitc(f'unordered_map<{self.group_type}, vector_type<uint32_t>, '
f'transTypes<{self.group_type}, hasher>> {self.group};')
self.n_grps = len(node)
self.n_grps = len(self.glist)
self.scanner = scan(self, first_col + '.size')
self.scanner.add(f'{self.group}[forward_as_tuple({g_contents}[{self.scanner.it_ver}])].emplace_back({self.scanner.it_ver});')
@@ -372,9 +387,15 @@ class groupby_c(ast_node):
if len_var is None:
len_var = 'len_'+base62uuid(7)
gscanner.add(f'auto &{len_var} = {val_var}.size;', position = 'front')
def get_key_idx (varname : str):
for i, g in enumerate(self.glist):
if varname == g[0].eval():
return i
return var_table[varname]
def get_var_names (varname : str):
var = var_table[varname]
var = get_key_idx(varname)
if type(var) is str:
return f'{var}[{val_var}]'
else:
@@ -545,6 +566,8 @@ class create_table(ast_node):
name = 'create_table'
first_order = name
def init(self, node):
if self.parent is None:
self.context.sql_begin()
self.sql = 'CREATE TABLE '
def produce(self, node):
@@ -558,10 +581,11 @@ class create_table(ast_node):
self.sql += ')'
if self.context.use_columnstore:
self.sql += ' engine=ColumnStore'
class insert(ast_node):
name = 'insert'
first_order = name
def produce(self, node):
values = node['query']['select']
tbl = node['insert']
@@ -586,6 +610,8 @@ class load(ast_node):
self.produce = self.produce_monetdb
else:
self.produce = self.produce_aq
if self.parent is None:
self.context.sql_begin()
def produce_aq(self, node):
node = node['load']
@@ -916,4 +942,4 @@ def include(objs):
import sys
include(sys.modules[__name__])
include(sys.modules[__name__])
+25 -11
View File
@@ -87,7 +87,9 @@ class Context:
self.finalized = False
self.udf = None
self.scans = []
self.procs = []
self.queries = []
def __init__(self):
self.tables_byname = dict()
self.col_byname = dict()
@@ -101,11 +103,9 @@ class Context:
self.has_dll = False
self.dialect = 'MonetDB'
self.have_hge = False
self.Error = lambda *args: print(*args)
self.Info = lambda *_: None
self.Info = lambda *_: None
self.new()
def emit(self, sql:str):
self.sql += sql + ' '
def emitc(self, c:str):
@@ -118,17 +118,31 @@ class Context:
self.emitc(str_scan)
self.scans.remove(scan)
function_head = '''
extern "C" int __DLLEXPORT__ dllmain(Context* cxt) {
using namespace std;
using namespace types;
auto server = static_cast<Server*>(cxt->alt_server);
'''
function_deco = '__AQEXPORT__(int) '
function_head = ('(Context* cxt) {\n' +
'\tusing namespace std;\n' +
'\tusing namespace types;\n' +
'\tauto server = static_cast<Server*>(cxt->alt_server);\n')
udf_head = ('#pragma once\n'
'#include \"./server/libaquery.h\"\n'
'#include \"./server/aggregations.h\"\n\n'
)
def sql_begin(self):
self.sql = ''
def sql_end(self):
self.queries.append('Q' + self.sql)
self.sql = ''
def postproc_begin(self, proc_name: str):
self.ccode = self.function_deco + proc_name + self.function_head
def postproc_end(self, proc_name: str):
self.procs.append(self.ccode + 'return 0;\n}')
self.ccode = ''
self.queries.append('P' + proc_name)
def finalize(self):
if not self.finalized:
headers = ''
@@ -137,6 +151,6 @@ class Context:
headers += '#include <' + h + '>\n'
else:
headers += '#include ' + h + '\n'
self.ccode = headers + self.function_head + self.ccode + 'return 0;\n}'
self.ccode = headers + '\n'.join(self.procs)
self.headers = set()
return self.ccode