restructure
This commit is contained in:
+59
-6
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user