Assumption, outfile, bugfixes on type deduction

This commit is contained in:
2022-08-05 21:34:28 +08:00
parent 6490aab558
commit 70f29e028d
22 changed files with 192 additions and 84 deletions
+75 -30
View File
@@ -15,6 +15,7 @@ class ast_node:
self.context = parent.context if context is None else context
self.parent = parent
self.sql = ''
self.ccode = ''
if hasattr(parent, 'datasource'):
self.datasource = parent.datasource
else:
@@ -28,7 +29,9 @@ class ast_node:
self.context.emit(code)
def add(self, code):
self.sql += code + ' '
def addc(self, code):
self.ccode += code + '\n'
name = 'null'
def init(self, _):
@@ -64,7 +67,8 @@ class projection(ast_node):
self.datasource = join(self, from_clause)
if 'assumptions' in from_clause:
self.assumptions = enlist(from_clause['assumptions'])
else:
self.assumptions = []
if self.datasource is not None:
self.datasource_changed = True
self.prev_datasource = self.context.datasource
@@ -98,6 +102,7 @@ class projection(ast_node):
this_type = proj_expr.type
name = proj_expr.sql
compound = True # compound column
proj_expr.cols_mentioned = self.datasource.rec
if not proj_expr.is_special:
y = lambda x:x
name = eval('f\'' + name + '\'')
@@ -110,7 +115,6 @@ class projection(ast_node):
if self.datasource.rec is not None:
self.col_ext = self.col_ext.union(self.datasource.rec)
proj_map[i] = [this_type, proj_expr.sql, proj_expr]
if 'name' in proj: # renaming column by AS keyword
name += ' AS ' + proj['name']
if not proj_expr.is_special:
@@ -147,12 +151,25 @@ class projection(ast_node):
self.add('FROM')
finialize(self.datasource)
finialize(self.where)
finialize(self.group_node)
if self.group_node and not self.group_node.use_sp_gb:
self.add(self.group_node.sql)
if self.col_ext or self.group_node and self.group_node.use_sp_gb:
self.use_postproc = True
o = self.assumptions
if 'orderby' in node:
self.add(orderby(self, node['orderby']).sql)
o.extend(enlist(node['orderby']))
if o:
self.add(orderby(self, o).sql)
if 'outfile' in node:
self.sql = outfile(self, node['outfile'], sql = self.sql).sql
self.outfile = outfile(self, node['outfile'], sql = self.sql)
if not self.use_postproc:
self.sql += self.outfile.sql
else:
self.outfile = None
if self.parent is None:
self.emit(self.sql+';\n')
else:
@@ -175,7 +192,7 @@ class projection(ast_node):
self.context.emitc(f'auto {vname} = ColRef<{typenames[idx].cname}>({length_name}, server->getCol({idx}));')
vid2cname[idx] = vname
# Create table into context
outtable_name = 'out_' + base62uuid(6)
self.outtable_name = 'out_' + base62uuid(6)
out_typenames = [None] * len(proj_map)
for key, val in proj_map.items():
@@ -186,18 +203,23 @@ class projection(ast_node):
if callable(val[1]):
val[1] = val[1](True)
decltypestring = val[1]
if val[0] == LazyT:
decltypestring = f'value_type<decays<decltype({decltypestring})>>'
if type(val[2].udf) is udf and val[2].udf.return_pattern == udf.ReturnPattern.elemental_return:
out_typenames[key] = f'ColRef<{decltypestring}>'
else:
out_typenames[key] = decltypestring
out_typenames[key] = decltypestring
else:
out_typenames[key] = val[0].cname
if (type(val[2].udf_called) is udf and
val[2].udf_called.return_pattern == udf.ReturnPattern.elemental_return
or
self.group_node and self.group_node.use_sp_gb and
val[2].cols_mentioned.intersection(
self.datasource.all_cols.difference(self.group_node.refs))
):
out_typenames[key] = f'ColRef<{out_typenames[key]}>'
# out_typenames = [v[0].cname for v in proj_map.values()]
self.context.emitc(f'auto {outtable_name} = new TableInfo<{",".join(out_typenames)}>("{outtable_name}");')
self.context.emitc(f'auto {self.outtable_name} = new TableInfo<{",".join(out_typenames)}>("{self.outtable_name}");')
# TODO: Inject custom group by code here and flag them in proj_map
# Type of UDFs? Complex UDFs, ones with static vars?
if self.group_node is not None and self.group_node.use_sp_gb:
@@ -206,19 +228,23 @@ class projection(ast_node):
for key, val in proj_map.items():
col_name = 'col_' + base62uuid(6)
self.context.emitc(f'decltype(auto) {col_name} = {outtable_name}->get_col<{key}>();')
self.context.emitc(f'decltype(auto) {col_name} = {self.outtable_name}->get_col<{key}>();')
gb_cexprs.append((col_name, val[2]))
self.group_node.finalize(gb_cexprs, gb_vartable)
else:
for key, val in proj_map.items():
if type(val[1]) is int:
self.context.emitc(f'{outtable_name}->get_col<{key}>().initfrom({vid2cname[val[1]]});')
self.context.emitc(f'{self.outtable_name}->get_col<{key}>().initfrom({vid2cname[val[1]]});')
else:
# for funcs evaluate f_i(x, ...)
self.context.emitc(f'{outtable_name}->get_col<{key}>() = {val[1]};')
self.context.emitc(f'{self.outtable_name}->get_col<{key}>() = {val[1]};')
# print out col_is
self.context.emitc(f'print(*{outtable_name});')
self.context.emitc(f'print(*{self.outtable_name});')
if self.outfile:
self.outfile.finalize()
self.context.emitc(f'puts("done.");')
class orderby(ast_node):
name = 'order by'
def produce(self, node):
@@ -357,11 +383,11 @@ class groupby_c(ast_node):
for ce in cexprs:
ex = ce[1]
materialize_builtin = {}
if type(ex.udf) is udf:
if '_builtin_len' in ex.udf.builtin_used:
if type(ex.udf_called) is udf:
if '_builtin_len' in ex.udf_called.builtin_used:
define_len_var()
materialize_builtin['_builtin_len'] = len_var
if '_builtin_ret' in ex.udf.builtin_used:
if '_builtin_ret' in ex.udf_called.builtin_used:
define_len_var()
gscanner.add(f'{ce[0]}.emplace_back({{{len_var}}});\n')
materialize_builtin['_builtin_ret'] = f'{ce[0]}.back()'
@@ -382,6 +408,7 @@ class groupby(ast_node):
node = enlist(node)
o_list = []
self.refs = set()
self.dedicated_glist : List[Tuple[expr, Set[ColRef]]] = []
self.use_sp_gb = False
for g in node:
@@ -392,7 +419,7 @@ class groupby(ast_node):
if self.parent.col_ext:
this_sp_ref = refs.difference(self.parent.col_ext)
self.use_sp_gb = self.use_sp_gb or len(this_sp_ref) > 0
self.refs.update(refs)
self.dedicated_glist.append((g_expr, refs))
g_str = g_expr.eval(c_code = False)
if 'sort' in g and f'{g["sort"]}'.lower() == 'desc':
@@ -418,7 +445,7 @@ class join(ast_node):
name = 'join'
def init(self, _):
self.joins:list = []
self.tables = []
self.tables : List[TableInfo] = []
self.tables_dir = dict()
self.rec = None
# self.tmp_name = 'join_' + base62uuid(4)
@@ -496,7 +523,9 @@ class join(ast_node):
raise ValueError(f'Table name/alias not defined{parsedColExpr[0]}')
else:
return datasource.parse_col_names(parsedColExpr[1])
@property
def all_cols(self):
return set([c for t in self.tables for c in t.columns])
def consume(self, _):
self.sql = ', '.join(self.joins)
return super().consume(_)
@@ -581,13 +610,19 @@ class load(ast_node):
class outfile(ast_node):
name="_outfile"
def __init__(self, parent, node, context = None, *, sql = None):
self.node = node
super().__init__(parent, node, context)
self.sql = sql
if self.context.dialect == 'MonetDB':
self.produce = self.produce_monetdb
else:
self.produce = self.produce_aq
self.sql = sql if sql else ''
def init(self, _):
assert(type(self.parent) is projection)
if not self.parent.use_postproc:
if self.context.dialect == 'MonetDB':
self.produce = self.produce_monetdb
else:
self.produce = self.produce_aq
return super().init(_)
def produce_aq(self, node):
filename = node['loc']['literal'] if 'loc' in node else node['literal']
self.sql += f'INTO OUTFILE "{filename}"'
@@ -605,6 +640,15 @@ class outfile(ast_node):
d = node['term']['literal']
self.sql += f' delimiters \'{d}\', \'{e}\''
def finalize(self):
filename = self.node['loc']['literal'] if 'loc' in self.node else self.node['literal']
sep = ',' if 'term' not in self.node else self.node['term']['literal']
file_pointer = 'fp_' + base62uuid(6)
self.addc(f'FILE* {file_pointer} = fopen("{filename}", "w");')
self.addc(f'{self.parent.outtable_name}->printall("{sep}", "\\n", nullptr, {file_pointer});')
self.addc(f'fclose({file_pointer});')
self.context.ccode += self.ccode
class udf(ast_node):
name = 'udf'
first_order = name
@@ -863,6 +907,7 @@ class udf(ast_node):
else:
return udf.ReturnPattern.bulk_return
def include(objs):
import inspect
for _, cls in inspect.getmembers(objs):
+15 -8
View File
@@ -57,6 +57,8 @@ class expr(ast_node):
if type(c_code) is bool:
self.c_code = c_code
self.udf_called = None
self.cols_mentioned : Optional[set[ColRef]] = None
ast_node.__init__(self, parent, node, None)
def init(self, _):
@@ -97,20 +99,25 @@ class expr(ast_node):
if key in special_func and not self.is_special:
self.is_special = True
if key in self.context.udf_map:
self.root.udf = self.context.udf_map[key]
if key == self.root.udf.name:
self.root.udf_called = self.context.udf_map[key]
if self.is_udfexpr and key == self.root.udf.name:
self.root.is_recursive_call_inudf = True
# TODO: make udf_called a set!
p = self.parent
while type(p) is expr and not p.udf_called:
p.udf_called = self.udf_called
p = p.parent
p = self.parent
while type(p) is expr and not p.is_special:
p.is_special = True
p = p.parent
need_decltypestr = any([e.need_decltypestr for e in exp_vals])
if need_decltypestr or (self.udf and type(op) is udf):
if need_decltypestr or (self.udf_called and type(op) is udf):
decltypestr_vals = [e.udf_decltypecall for e in exp_vals]
self.udf_decltypecall = op(self.c_code, *decltypestr_vals)
if self.udf and type(op) is udf:
if self.udf_called and type(op) is udf:
self.udf_decltypecall = op.decltypecall(self.c_code, *decltypestr_vals)
elif self.is_udfexpr:
@@ -230,10 +237,10 @@ class expr(ast_node):
assert(self.is_root)
def call(decltypestr = False) -> str:
nonlocal c_code, y, materialize_builtin
if self.udf is not None:
if self.udf_called is not None:
loc = locals()
builtin_vars = self.udf.builtin_used
for b in self.udf.builtin_var.all:
builtin_vars = self.udf_called.builtin_used
for b in self.udf_called.builtin_var.all:
exec(f'loc["{b}"] = lambda: "{{{b}()}}"')
if builtin_vars:
if type(materialize_builtin) is dict:
+2 -1
View File
@@ -101,7 +101,8 @@ class Context:
self.has_dll = False
self.dialect = 'MonetDB'
self.have_hge = False
self.Info = lambda *_: None
self.Info = lambda *_: None
self.new()