fixed regression on join condition awareness

This commit is contained in:
2022-10-02 00:41:10 +08:00
parent e5ba3f63d6
commit b666d6d9b2
3 changed files with 47 additions and 18 deletions
+18 -4
View File
@@ -674,8 +674,19 @@ class join(ast_node):
tablename += f' ON {_ex}'
elif keys[1].lower() == 'using':
if _ex.is_ColExpr:
self.join_conditions += (_ex.raw_col, j.get_cols(_ex.raw_col.name))
self.join_conditions.append( (_ex.raw_col, j.get_cols(_ex.raw_col.name)) )
tablename += f' USING {_ex}'
if keys[0].lower().startswith('natural'):
ltbls : List[TableInfo] = []
if isinstance(self.parent, join):
ltbls = self.parent.tables
elif isinstance(self.parent, TableInfo):
ltbls = [self.parent]
for tl in ltbls:
for cl in tl.columns:
cr = j.get_cols(cl.name)
if cr:
self.join_conditions.append( (cl, cr) )
self.joins.append((tablename, self.have_sep))
self.tables += j.tables
self.tables_dir = {**self.tables_dir, **j.tables_dir}
@@ -686,13 +697,14 @@ class join(ast_node):
else:
print(f'Error: table {node} not found.')
def get_cols(self, colExpr: str) -> ColRef:
def get_cols(self, colExpr: str) -> Optional[ColRef]:
for t in self.tables:
if colExpr in t.columns_byname:
col = t.columns_byname[colExpr]
if type(self.rec) is set:
self.rec.add(col)
return col
return None
def parse_col_names(self, colExpr:str) -> ColRef:
parsedColExpr = colExpr.split('.')
@@ -771,7 +783,7 @@ class create_table(ast_node):
def produce(self, node):
ct = node[self.name]
tbl = self.context.add_table(ct['name'], ct['columns'])
self.sql = f'CREATE TABLE {tbl.table_name}('
self.sql = f'CREATE TABLE IF NOT EXISTS {tbl.table_name}('
columns = []
for c in tbl.columns:
columns.append(f'{c.name} {c.type.sqlname}')
@@ -787,7 +799,9 @@ class drop(ast_node):
node = node['drop']
tbl_name = node['table']
if tbl_name in self.context.tables_byname:
tbl_obj = self.context.tables_byname[tbl_name]
tbl_obj : TableInfo = self.context.tables_byname[tbl_name]
for a in tbl_obj.alias:
self.context.tables_byname.pop(a, None)
# TODO: delete in postproc engine
self.context.tables_byname.pop(tbl_name)
self.context.tables.remove(tbl_obj)
+12 -2
View File
@@ -121,7 +121,7 @@ class Context:
def __init__(self):
self.tables_byname = dict()
self.col_byname = dict()
self.tables = []
self.tables : List[TableInfo] = []
self.cols = []
self.datasource = None
self.module_stubs = ''
@@ -176,6 +176,15 @@ class Context:
self.queries.insert(self.module_init_loc, 'P__builtin_init_user_module')
return ret + '}\n'
def finalize_query(self):
# clear aliases
for t in self.tables:
for a in t.alias:
if a != t.table_name:
self.tables_byname.pop(a, None)
t.alias.clear()
t.alias.add(t.table_name)
def sql_begin(self):
self.sql = ''
@@ -195,7 +204,8 @@ class Context:
self.procs.append(self.ccode + 'return 0;\n}')
self.ccode = ''
self.queries.append('P' + proc_name)
self.finalize_query()
def finalize_udf(self):
if self.udf is not None:
return (Context.udf_head