fixed wildcard compound cols, ratios, etc.
This commit is contained in:
+32
-14
@@ -133,7 +133,7 @@ class projection(ast_node):
|
||||
sql_expr = expr(self, e, c_code=False)
|
||||
this_type = proj_expr.type
|
||||
name = proj_expr.sql
|
||||
compound = True # compound column
|
||||
compound = [proj_expr.is_compound > 1] # compound column
|
||||
proj_expr.cols_mentioned = self.datasource.rec
|
||||
alias = ''
|
||||
if 'name' in proj: # renaming column by AS keyword
|
||||
@@ -142,23 +142,28 @@ class projection(ast_node):
|
||||
if not proj_expr.is_special:
|
||||
if proj_expr.node == '*':
|
||||
name = [c.get_full_name() for c in self.datasource.rec]
|
||||
this_type = [c.type for c in self.datasource.rec]
|
||||
compound = [c.compound for c in self.datasource.rec]
|
||||
proj_expr = [expr(self, c.name) for c in self.datasource.rec]
|
||||
else:
|
||||
y = lambda x:x
|
||||
count = lambda : 'count(*)'
|
||||
name = enlist(sql_expr.eval(False, y, count=count))
|
||||
for n in name:
|
||||
this_type = enlist(this_type)
|
||||
proj_expr = enlist(proj_expr)
|
||||
for t, n, pexpr in zip(this_type, name, proj_expr):
|
||||
offset = len(col_exprs)
|
||||
if n not in self.var_table:
|
||||
self.var_table[n] = offset
|
||||
if proj_expr.is_ColExpr and type(proj_expr.raw_col) is ColRef:
|
||||
for _alias in (proj_expr.raw_col.table.alias):
|
||||
if pexpr.is_ColExpr and type(pexpr.raw_col) is ColRef:
|
||||
for _alias in (pexpr.raw_col.table.alias):
|
||||
self.var_table[f'{_alias}.'+n] = offset
|
||||
proj_map[i] = [this_type, offset, proj_expr]
|
||||
proj_map[i] = [t, offset, pexpr]
|
||||
col_expr = n + ' AS ' + alias if alias else n
|
||||
if alias:
|
||||
self.var_table[alias] = offset
|
||||
|
||||
col_exprs.append((col_expr, proj_expr.type))
|
||||
col_exprs.append((col_expr, t))
|
||||
i += 1
|
||||
else:
|
||||
self.context.headers.add('"./server/aggregations.h"')
|
||||
@@ -169,7 +174,8 @@ class projection(ast_node):
|
||||
i += 1
|
||||
name = enlist(name)
|
||||
disp_name = [get_legal_name(alias if alias else n) for n in name]
|
||||
|
||||
this_type = enlist(this_type)
|
||||
|
||||
elif type(proj) is str:
|
||||
col = self.datasource.get_col(proj)
|
||||
this_type = col.type
|
||||
@@ -178,8 +184,8 @@ class projection(ast_node):
|
||||
# name = col.name
|
||||
self.datasource.rec = None
|
||||
# TODO: Type deduction in Python
|
||||
for n in disp_name:
|
||||
cols.append(ColRef(this_type, self.out_table, None, n, len(cols), compound=compound))
|
||||
for t, n, c in zip(this_type, disp_name, compound):
|
||||
cols.append(ColRef(t, self.out_table, None, n, len(cols), compound=c))
|
||||
|
||||
self.out_table.add_cols(cols, new = False)
|
||||
|
||||
@@ -213,7 +219,7 @@ class projection(ast_node):
|
||||
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
|
||||
self.has_postproc = True
|
||||
|
||||
o = self.assumptions
|
||||
if 'orderby' in node:
|
||||
@@ -223,7 +229,7 @@ class projection(ast_node):
|
||||
|
||||
if 'outfile' in node:
|
||||
self.outfile = outfile(self, node['outfile'], sql = self.sql)
|
||||
if not self.use_postproc:
|
||||
if not self.has_postproc:
|
||||
self.sql += self.outfile.sql
|
||||
else:
|
||||
self.outfile = None
|
||||
@@ -279,11 +285,12 @@ class projection(ast_node):
|
||||
val[2].cols_mentioned.intersection(
|
||||
self.datasource.all_cols().difference(self.group_node.refs))
|
||||
) and val[2].is_compound # compound val not in key
|
||||
# or
|
||||
or
|
||||
val[2].is_compound > 1
|
||||
# (not self.group_node and val[2].is_compound)
|
||||
):
|
||||
out_typenames[key] = f'ColRef<{out_typenames[key]}>'
|
||||
|
||||
self.out_table.columns[key].compound = True
|
||||
outtable_col_nameslist = ', '.join([f'"{c.name}"' for c in self.out_table.columns])
|
||||
self.outtable_col_names = 'names_' + base62uuid(4)
|
||||
self.context.emitc(f'const char* {self.outtable_col_names}[] = {{{outtable_col_nameslist}}};')
|
||||
@@ -729,6 +736,17 @@ class create_table(ast_node):
|
||||
name = 'create_table'
|
||||
first_order = name
|
||||
def init(self, node):
|
||||
node = node[self.name]
|
||||
if 'query' in node:
|
||||
if 'name' not in node:
|
||||
raise ValueError("Table name not specified")
|
||||
projection_node = node['query']
|
||||
projection_node['into'] = node['name']
|
||||
projection(None, projection_node, self.context)
|
||||
self.produce = lambda *_: None
|
||||
self.spawn = lambda *_: None
|
||||
self.consume = lambda *_: None
|
||||
return
|
||||
if self.parent is None:
|
||||
self.context.sql_begin()
|
||||
self.sql = 'CREATE TABLE '
|
||||
@@ -851,7 +869,7 @@ class outfile(ast_node):
|
||||
|
||||
def init(self, _):
|
||||
assert(isinstance(self.parent, projection))
|
||||
if not self.parent.use_postproc:
|
||||
if not self.parent.has_postproc:
|
||||
if self.context.dialect == 'MonetDB':
|
||||
self.produce = self.produce_monetdb
|
||||
else:
|
||||
|
||||
+5
-4
@@ -80,7 +80,7 @@ class expr(ast_node):
|
||||
self.udf_map = parent.context.udf_map
|
||||
self.func_maps = {**builtin_func, **self.udf_map, **user_module_func}
|
||||
self.operators = {**builtin_operators, **self.udf_map, **user_module_func}
|
||||
self.ext_aggfuncs = ['sum', 'avg', 'count', 'min', 'max', 'last']
|
||||
self.ext_aggfuncs = ['sum', 'avg', 'count', 'min', 'max', 'last', 'first']
|
||||
|
||||
def produce(self, node):
|
||||
from engine.utils import enlist
|
||||
@@ -114,9 +114,9 @@ class expr(ast_node):
|
||||
|
||||
str_vals = [e.sql for e in exp_vals]
|
||||
type_vals = [e.type for e in exp_vals]
|
||||
is_compound = any([e.is_compound for e in exp_vals])
|
||||
is_compound = max([e.is_compound for e in exp_vals])
|
||||
if key in self.ext_aggfuncs:
|
||||
self.is_compound = False
|
||||
self.is_compound = max(0, is_compound - 1)
|
||||
else:
|
||||
self.is_compound = is_compound
|
||||
try:
|
||||
@@ -134,7 +134,7 @@ class expr(ast_node):
|
||||
self.sql = op(self.c_code, *str_vals)
|
||||
|
||||
special_func = [*self.context.udf_map.keys(), *self.context.module_map.keys(),
|
||||
"maxs", "mins", "avgs", "sums", "deltas", "last"]
|
||||
"maxs", "mins", "avgs", "sums", "deltas", "last", "first", "ratios"]
|
||||
if self.context.special_gb:
|
||||
special_func = [*special_func, *self.ext_aggfuncs]
|
||||
|
||||
@@ -259,6 +259,7 @@ class expr(ast_node):
|
||||
self.sql = table_name + self.raw_col.name
|
||||
self.type = self.raw_col.type
|
||||
self.is_compound = True
|
||||
self.is_compound += self.raw_col.compound
|
||||
self.opname = self.raw_col
|
||||
else:
|
||||
self.sql = '\'' + node + '\'' if node != '*' else '*'
|
||||
|
||||
@@ -16,7 +16,7 @@ class expr_base(ast_node, metaclass = abc.ABCMeta):
|
||||
self.udf_map = self.context.udf_map
|
||||
self.func_maps = {**builtin_func, **self.udf_map, **user_module_func}
|
||||
self.operators = {**builtin_operators, **self.udf_map, **user_module_func}
|
||||
self.narrow_funcs = ['sum', 'avg', 'count', 'min', 'max', 'last']
|
||||
self.narrow_funcs = ['sum', 'avg', 'count', 'min', 'max', 'last', 'first']
|
||||
|
||||
def get_variable(self):
|
||||
pass
|
||||
@@ -56,7 +56,7 @@ class expr_base(ast_node, metaclass = abc.ABCMeta):
|
||||
raise ValueError(f'Parse Error: more than 1 entry in {node}.')
|
||||
key, val = next(iter(node.items()))
|
||||
if key in self.operators:
|
||||
self.child_exprs = [__class__(self, v) for v in val]
|
||||
self.child_exprs = [self.__class__(self, v) for v in val]
|
||||
self.process_child_nodes()
|
||||
else:
|
||||
self.process_non_operator(key, val)
|
||||
|
||||
Reference in New Issue
Block a user