restructure

This commit is contained in:
2022-08-28 19:39:18 +08:00
parent 42c334af84
commit 780a60fa5d
22 changed files with 279 additions and 74 deletions
+59 -6
View File
@@ -601,18 +601,62 @@ class insert(ast_node):
pass
self.sql += ', '.join(list_values) + ')'
class load(ast_node):
name="load"
first_order = name
def init(self, _):
if self.context.dialect == 'MonetDB':
def init(self, node):
self.module = False
if node['load']['file_type'] == 'module':
self.produce = self.produce_module
self.module = True
elif self.context.dialect == 'MonetDB':
self.produce = self.produce_monetdb
else:
else:
self.produce = self.produce_aq
if self.parent is None:
self.context.sql_begin()
def produce_module(self, node):
# create command for exec engine -> done
# create c++ stub
# create dummy udf obj for parsing
# def decode_type(ty : str) -> str:
# ret = ''
# back = ''
# while(ty.startswith('vec')):
# ret += 'ColRef<'
# back += '>'
# ty = ty[3:]
# ret += ty
# return ret + back
node = node['load']
file = node['file']['literal']
self.context.queries.append(f'M{file}')
self.module_name = file
self.functions = {}
if 'funcs' in node:
for f in enlist(node['funcs']):
fname = f['fname']
self.context.queries.append(f'F{fname}')
ret_type = VoidT
if 'ret_type' in f:
ret_type = Types.decode(f['ret_type'])
nargs = 0
arglist = ''
if 'var' in f:
arglist = []
for v in enlist(f['var']):
arglist.append(f'{Types.decode(v["type"]).cname} {v["arg"]}')
nargs = len(arglist)
arglist = ', '.join(arglist)
# create c++ stub
cpp_stub = f'{ret_type.cname} (*{fname})({arglist});'
self.context.module_stubs += cpp_stub + '\n'
self.context.module_map[fname] = cpp_stub
#registration for parser
self.functions[fname] = user_module_function(fname, nargs, ret_type)
def produce_aq(self, node):
node = node['load']
s1 = 'LOAD DATA INFILE '
@@ -710,7 +754,11 @@ class udf(ast_node):
self.var_table = {}
self.args = []
if self.context.udf is None:
self.context.udf = Context.udf_head
self.context.udf = (
Context.udf_head
+ self.context.module_stubs
+ self.context.get_init_func()
)
self.context.headers.add('\"./udf.hpp\"')
self.vecs = set()
self.code_list = []
@@ -933,6 +981,11 @@ class udf(ast_node):
else:
return udf.ReturnPattern.bulk_return
class user_module_function(OperatorBase):
def __init__(self, name, nargs, ret_type):
super().__init__(name, nargs, lambda: ret_type, call=fn_behavior)
user_module_func[name] = self
builtin_operators[name] = self
def include(objs):
import inspect
+13 -4
View File
@@ -102,6 +102,8 @@ class Context:
self.tables = []
self.cols = []
self.datasource = None
self.module_stubs = ''
self.module_map = {}
self.udf_map = dict()
self.udf_agg_map = dict()
self.use_columnstore = False
@@ -134,21 +136,28 @@ class Context:
'#include \"./server/libaquery.h\"\n'
'#include \"./server/aggregations.h\"\n\n'
)
def get_init_func(self):
if self.module_map:
return ''
ret = 'void init(Context* cxt){\n'
for fname in self.module_map.keys():
ret += f'{fname} = (decltype({fname}))(cxt->get_module_function("{fname}"));\n'
return ret + '}\n'
def sql_begin(self):
self.sql = ''
def sql_end(self):
self.queries.append('Q' + self.sql)
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)
self.queries.append('P' + proc_name)
def finalize(self):
if not self.finalized:
headers = ''
@@ -159,4 +168,4 @@ class Context:
headers += '#include ' + h + '\n'
self.ccode = headers + '\n'.join(self.procs)
self.headers = set()
return self.ccode
return self.ccode