bug fix on select into

This commit is contained in:
2022-09-21 17:23:50 +08:00
parent 48beab441d
commit 34a9fe105c
12 changed files with 305 additions and 111 deletions
+91 -44
View File
@@ -61,9 +61,16 @@ class projection(ast_node):
pass
def produce(self, node):
p = node['select']
self.projections = p if type(p) is list else [p]
self.add('SELECT')
self.has_postproc = False
if 'select' in node:
p = node['select']
self.distinct = False
elif 'select_distinct' in node:
p = node['select_distinct']
self.distinct = True
self.projections = p if type(p) is list else [p]
if self.parent is None:
self.context.sql_begin()
self.postproc_fname = 'dll_' + base62uuid(6)
@@ -75,8 +82,8 @@ class projection(ast_node):
if 'from' in node:
from_clause = node['from']['table_source']
self.datasource = join(self, from_clause)
if 'assumptions' in from_clause:
self.assumptions = enlist(from_clause['assumptions'])
if 'assumptions' in node['from']:
self.assumptions = enlist(node['from']['assumptions'])
if self.datasource is not None:
self.datasource_changed = True
@@ -109,61 +116,82 @@ class projection(ast_node):
proj_map : Dict[int, List[Union[Types, int, str, expr]]]= dict()
self.var_table = dict()
# self.sp_refs = set()
for i, proj in enumerate(self.projections):
i = 0
for proj in self.projections:
compound = False
self.datasource.rec = set()
name = ''
this_type = AnyT
if type(proj) is dict:
if type(proj) is dict or proj == '*':
if 'value' in proj:
e = proj['value']
proj_expr = expr(self, e)
this_type = proj_expr.type
name = proj_expr.sql
compound = True # compound column
proj_expr.cols_mentioned = self.datasource.rec
alias = ''
if 'name' in proj: # renaming column by AS keyword
alias = proj['name']
if not proj_expr.is_special:
elif proj == '*':
e = '*'
else:
print('unknown projection', proj)
proj_expr = expr(self, e)
sql_expr = expr(self, e, c_code=False)
this_type = proj_expr.type
name = proj_expr.sql
compound = True # compound column
proj_expr.cols_mentioned = self.datasource.rec
alias = ''
if 'name' in proj: # renaming column by AS keyword
alias = proj['name']
if not proj_expr.is_special:
if proj_expr.node == '*':
name = [c.get_full_name() for c in self.datasource.rec]
else:
y = lambda x:x
name = eval('f\'' + name + '\'')
count = lambda : 'count(*)'
name = enlist(sql_expr.eval(False, y, count=count))
for n in name:
offset = len(col_exprs)
if name not in self.var_table:
self.var_table[name] = offset
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 n in (proj_expr.raw_col.table.alias):
self.var_table[f'{n}.'+name] = offset
for _alias in (proj_expr.raw_col.table.alias):
self.var_table[f'{_alias}.'+n] = offset
proj_map[i] = [this_type, offset, proj_expr]
col_expr = name + ' AS ' + alias if alias else name
col_expr = n + ' AS ' + alias if alias else n
if alias:
self.var_table[alias] = offset
col_exprs.append((col_expr, proj_expr.type))
else:
self.context.headers.add('"./server/aggregations.h"')
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]
disp_name = get_legal_name(alias if alias else name)
i += 1
else:
self.context.headers.add('"./server/aggregations.h"')
self.has_postproc = True
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]
i += 1
name = enlist(name)
disp_name = [get_legal_name(alias if alias else n) for n in name]
elif type(proj) is str:
col = self.datasource.get_col(proj)
this_type = col.type
disp_name = proj
print('Unknown behavior:', proj, 'is str')
# name = col.name
self.datasource.rec = None
# TODO: Type deduction in Python
cols.append(ColRef(this_type, self.out_table, None, disp_name, i, compound=compound))
for n in disp_name:
cols.append(ColRef(this_type, self.out_table, None, n, len(cols), compound=compound))
self.out_table.add_cols(cols, new = False)
if 'groupby' in node:
self.group_node = groupby(self, node['groupby'])
if self.group_node.use_sp_gb:
self.has_postproc = True
else:
self.group_node = None
if not self.has_postproc and self.distinct:
self.add('DISTINCT')
self.col_ext = [c for c in self.col_ext if c.name not in self.var_table] # remove duplicates in self.var_table
col_ext_names = [c.name for c in self.col_ext]
self.add(', '.join([c[0] for c in col_exprs] + col_ext_names))
@@ -249,7 +277,7 @@ class projection(ast_node):
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))
self.datasource.all_cols().difference(self.group_node.refs))
) and val[2].is_compound # compound val not in key
# or
# (not self.group_node and val[2].is_compound)
@@ -282,25 +310,37 @@ class projection(ast_node):
# for funcs evaluate f_i(x, ...)
self.context.emitc(f'{self.out_table.contextname_cpp}->get_col<{key}>() = {val[1]};')
# print out col_is
self.context.emitc(f'print(*{self.out_table.contextname_cpp});')
if 'into' not in node:
self.context.emitc(f'print(*{self.out_table.contextname_cpp});')
if self.outfile:
self.outfile.finalize()
if 'into' in node:
self.context.emitc(select_into(self, node['into']).ccode)
if not self.distinct:
self.finalize()
def finalize(self):
self.context.emitc(f'puts("done.");')
if self.parent is None:
self.context.sql_end()
self.context.postproc_end(self.postproc_fname)
class select_distinct(projection):
first_order = 'select_distinct'
def consume(self, node):
super().consume(node)
if self.has_postproc:
self.context.emitc(
f'{self.out_table.table_name}->distinct();'
)
self.finalize()
class select_into(ast_node):
def init(self, node):
if type(self.parent) is projection:
if isinstance(self.parent, projection):
if self.context.has_dll:
# has postproc put back to monetdb
self.produce = self.produce_cpp
@@ -308,8 +348,8 @@ class select_into(ast_node):
self.produce = self.produce_sql
else:
raise ValueError('parent must be projection')
def produce_cpp(self, node):
assert(type(self.parent) is projection)
if not hasattr(self.parent, 'out_table'):
raise Exception('No out_table found.')
else:
@@ -508,7 +548,7 @@ class groupby(ast_node):
return False
def produce(self, node):
if type(self.parent) is not projection:
if not isinstance(self.parent, projection):
raise ValueError('groupby can only be used in projection')
node = enlist(node)
@@ -554,7 +594,7 @@ class join(ast_node):
self.tables : List[TableInfo] = []
self.tables_dir = dict()
self.rec = None
self.top_level = self.parent and type(self.parent) is projection
self.top_level = self.parent and isinstance(self.parent, projection)
self.have_sep = False
# self.tmp_name = 'join_' + base62uuid(4)
# self.datasource = TableInfo(self.tmp_name, [], self.context)
@@ -636,9 +676,16 @@ class join(ast_node):
datasource.rec = None
return ret
@property
# @property
def all_cols(self):
return set([c for t in self.tables for c in t.columns])
ret = set()
for table in self.tables:
rec = table.rec
table.rec = self.rec
ret.update(table.all_cols())
table.rec = rec
return ret
def consume(self, node):
self.sql = ''
for j in self.joins:
@@ -787,7 +834,7 @@ class outfile(ast_node):
self.sql = sql if sql else ''
def init(self, _):
assert(type(self.parent) is projection)
assert(isinstance(self.parent, projection))
if not self.parent.use_postproc:
if self.context.dialect == 'MonetDB':
self.produce = self.produce_monetdb
+17 -9
View File
@@ -1,4 +1,4 @@
from typing import Optional
from typing import Optional, Set
from reconstruct.ast import ast_node
from reconstruct.storage import ColRef, Context
from engine.types import *
@@ -199,11 +199,8 @@ class expr(ast_node):
self.udf_decltypecall = ex_vname.sql
else:
print(f'Undefined expr: {key}{val}')
if 'distinct' in val and key != count:
if self.c_code:
self.sql = 'distinct ' + self.sql
elif self.is_compound:
self.sql = '(' + self.sql + ').distinct()'
if type(node) is str:
if self.is_udfexpr:
curr_udf : udf = self.root.udf
@@ -235,8 +232,13 @@ class expr(ast_node):
# get the column from the datasource in SQL context
else:
if self.datasource is not None:
self.raw_col = self.datasource.parse_col_names(node)
self.raw_col = self.raw_col if type(self.raw_col) is ColRef else None
if (node == '*' and
not (type(self.parent) is expr
and 'count' in self.parent.node)):
self.datasource.all_cols()
else:
self.raw_col = self.datasource.parse_col_names(node)
self.raw_col = self.raw_col if type(self.raw_col) is ColRef else None
if self.raw_col is not None:
self.is_ColExpr = True
table_name = ''
@@ -259,10 +261,16 @@ class expr(ast_node):
self.is_compound = True
self.opname = self.raw_col
else:
self.sql = '\'' + node + '\''
self.sql = '\'' + node + '\'' if node != '*' else '*'
self.type = StrT
self.opname = self.sql
if self.c_code and self.datasource is not None:
if (type(self.parent) is expr and
'distinct' in self.parent.node and
not self.is_special):
# this node is executed by monetdb
# gb condition, not special
self.sql = f'distinct({self.sql})'
self.sql = f'{{y(\"{self.sql}\")}}'
elif type(node) is bool:
self.type = BoolT
+19 -3
View File
@@ -1,5 +1,5 @@
from engine.types import *
from engine.utils import base62uuid, enlist
from engine.utils import CaseInsensitiveDict, base62uuid, enlist
from typing import List, Dict, Set
class ColRef:
@@ -20,6 +20,18 @@ class ColRef:
# e.g. order by, group by, filter by expressions
self.__arr__ = (_ty, cobj, table, name, id)
def get_full_name(self):
table_name = self.table.table_name
it_alias = iter(self.table.alias)
alias = next(it_alias, table_name)
try:
while alias == table_name:
alias = next(it_alias)
except StopIteration:
alias = table_name
return f'{alias}.{self.name}'
def __getitem__(self, key):
if type(key) is str:
return getattr(self, key)
@@ -35,7 +47,7 @@ class TableInfo:
self.table_name : str = table_name
self.contextname_cpp : str = ''
self.alias : Set[str] = set([table_name])
self.columns_byname : Dict[str, ColRef] = dict() # column_name, type
self.columns_byname : Dict[str, ColRef] = CaseInsensitiveDict() # column_name, type
self.columns : List[ColRef] = []
self.cxt = cxt
# keep track of temp vars
@@ -85,7 +97,11 @@ class TableInfo:
raise ValueError(f'Table name/alias not defined{parsedColExpr[0]}')
else:
return datasource.parse_col_names(parsedColExpr[1])
def all_cols(self):
if type(self.rec) is set:
self.rec.update(self.columns)
return set(self.columns)
class Context:
def new(self):