From af3fa1f69c077789b8a1c5078d1bb94a8d5e2240 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Sun, 2 Jun 2013 18:05:47 -0400 Subject: implement join rewriting inside of visit_select(). Currently this is global or not based on fixing nested_join_translation as True or not. --- lib/sqlalchemy/sql/compiler.py | 74 ++++++++++++++++++++++++++++++++++++------ lib/sqlalchemy/sql/util.py | 7 ++-- lib/sqlalchemy/sql/visitors.py | 7 +++- 3 files changed, 75 insertions(+), 13 deletions(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index 41ef20a7a..030d6dce9 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1077,23 +1077,64 @@ class SQLCompiler(engine.Compiled): def get_crud_hint_text(self, table, text): return None + def _transform_select_for_nested_joins(self, select): + adapters = [] + + traverse_options = {"cloned": {}} + + def visit_join(elem): + if isinstance(elem.right, sql.FromGrouping): + selectable = sql.select([elem.right.element], use_labels=True) + selectable = selectable.alias() + + while adapters: + adapt = adapters.pop(-1) + selectable = adapt.traverse(selectable) + + for c in selectable.c: + c._label = c._key_label = c.name + + elem.right = selectable + adapters.append( + sql_util.ClauseAdapter(selectable, + traverse_options=traverse_options) + ) + + select = visitors.cloned_traverse(select, + traverse_options, {"join": visit_join}) + + for adap in reversed(adapters): + select = adap.traverse(select) + return select + + def _transform_result_map_for_nested_joins(self, select, transformed_select): + d = dict(zip(transformed_select.inner_columns, select.inner_columns)) + for key, (name, objs, typ) in list(self.result_map.items()): + objs = tuple([d.get(col, col) for col in objs]) + self.result_map[key] = (name, objs, typ) + def visit_select(self, select, asfrom=False, parens=True, iswrapper=False, fromhints=None, compound_index=0, force_result_map=False, - positional_names=None, **kwargs): - entry = self.stack and self.stack[-1] or {} + positional_names=None, + nested_join_translation=False, **kwargs): + + #nested_join_translation = True + if not nested_join_translation: + transformed_select = self._transform_select_for_nested_joins(select) + text = self.visit_select( + transformed_select, asfrom=asfrom, parens=parens, + iswrapper=iswrapper, fromhints=fromhints, + compound_index=compound_index, + force_result_map=force_result_map, + positional_names=positional_names, + nested_join_translation=True, **kwargs + ) - existingfroms = entry.get('from', None) - froms = select._get_display_froms(existingfroms, asfrom=asfrom) - correlate_froms = set(sql._from_objects(*froms)) - - # TODO: might want to propagate existing froms for - # select(select(select)) where innermost select should correlate - # to outermost if existingfroms: correlate_froms = - # correlate_froms.union(existingfroms) + entry = self.stack and self.stack[-1] or {} populate_result_map = force_result_map or ( compound_index == 0 and ( @@ -1102,6 +1143,19 @@ class SQLCompiler(engine.Compiled): ) ) + if not nested_join_translation: + if populate_result_map: + self._transform_result_map_for_nested_joins( + select, transformed_select) + return text + + existingfroms = entry.get('from', None) + + froms = select._get_display_froms(existingfroms, asfrom=asfrom) + + correlate_froms = set(sql._from_objects(*froms)) + + self.stack.append({'from': correlate_froms, 'iswrapper': iswrapper}) diff --git a/lib/sqlalchemy/sql/util.py b/lib/sqlalchemy/sql/util.py index 91740dc16..ffa07d3df 100644 --- a/lib/sqlalchemy/sql/util.py +++ b/lib/sqlalchemy/sql/util.py @@ -797,8 +797,11 @@ class ClauseAdapter(visitors.ReplacingCloningVisitor): def __init__(self, selectable, equivalents=None, include=None, exclude=None, include_fn=None, exclude_fn=None, - adapt_on_names=False): + adapt_on_names=False, + traverse_options=None): self.__traverse_options__ = {'stop_on': [selectable]} + if traverse_options: + self.__traverse_options__.update(traverse_options) self.selectable = selectable if include: assert not include_fn @@ -832,7 +835,7 @@ class ClauseAdapter(visitors.ReplacingCloningVisitor): def replace(self, col): if isinstance(col, expression.FromClause) and \ self.selectable.is_derived_from(col): - return self.selectable + return self.selectable elif not isinstance(col, expression.ColumnElement): return None elif self.include_fn and not self.include_fn(col): diff --git a/lib/sqlalchemy/sql/visitors.py b/lib/sqlalchemy/sql/visitors.py index 62f46ab64..31ac686e3 100644 --- a/lib/sqlalchemy/sql/visitors.py +++ b/lib/sqlalchemy/sql/visitors.py @@ -30,6 +30,7 @@ import operator __all__ = ['VisitableType', 'Visitable', 'ClauseVisitor', 'CloningVisitor', 'ReplacingCloningVisitor', 'iterate', 'iterate_depthfirst', 'traverse_using', 'traverse', + 'traverse_depthfirst', 'cloned_traverse', 'replacement_traverse'] @@ -255,7 +256,11 @@ def cloned_traverse(obj, opts, visitors): """clone the given expression structure, allowing modifications by visitors.""" - cloned = util.column_dict() + + if "cloned" in opts: + cloned = opts['cloned'] + else: + cloned = util.column_dict() stop_on = util.column_set(opts.get('stop_on', [])) def clone(elem): -- cgit v1.2.1 From 02ae3cd54d0c47850ae1c894abae256a4717fe2d Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Sun, 2 Jun 2013 19:33:19 -0400 Subject: getting things to join without subqueries, but some glitches in the compiler step when we do query.count() are showing --- lib/sqlalchemy/sql/compiler.py | 11 ++++++++--- lib/sqlalchemy/sql/expression.py | 8 ++++---- lib/sqlalchemy/sql/util.py | 1 + 3 files changed, 13 insertions(+), 7 deletions(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index 030d6dce9..27e883c86 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1095,14 +1095,19 @@ class SQLCompiler(engine.Compiled): c._label = c._key_label = c.name elem.right = selectable - adapters.append( - sql_util.ClauseAdapter(selectable, + import pdb + pdb.set_trace() + adapter = sql_util.ClauseAdapter(selectable, traverse_options=traverse_options) - ) + adapter.__traverse_options__.pop('stop_on') + adapters.append(adapter) select = visitors.cloned_traverse(select, traverse_options, {"join": visit_join}) + if adapters: + import pdb + pdb.set_trace() for adap in reversed(adapters): select = adap.traverse(select) return select diff --git a/lib/sqlalchemy/sql/expression.py b/lib/sqlalchemy/sql/expression.py index 5820cb106..edab9e290 100644 --- a/lib/sqlalchemy/sql/expression.py +++ b/lib/sqlalchemy/sql/expression.py @@ -4212,10 +4212,10 @@ class FromGrouping(FromClause): @property def foreign_keys(self): - # this could be - # self.element.foreign_keys - # see SelectableTest.test_join_condition - return set() + return self.element.foreign_keys + + def is_derived_from(self, element): + return self.element.is_derived_from(element) @property def _hide_froms(self): diff --git a/lib/sqlalchemy/sql/util.py b/lib/sqlalchemy/sql/util.py index ffa07d3df..6a267752e 100644 --- a/lib/sqlalchemy/sql/util.py +++ b/lib/sqlalchemy/sql/util.py @@ -833,6 +833,7 @@ class ClauseAdapter(visitors.ReplacingCloningVisitor): return newcol def replace(self, col): + print "COL!", col if isinstance(col, expression.FromClause) and \ self.selectable.is_derived_from(col): return self.selectable -- cgit v1.2.1 From 35a674aab4a832e76232e7be4b16b7a635a19824 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Sun, 2 Jun 2013 19:48:30 -0400 Subject: - figured out what the from_self() thing was about, part of query.statement, but would like to improve upon query.statement needing to do this --- lib/sqlalchemy/sql/compiler.py | 8 +------- lib/sqlalchemy/sql/util.py | 1 - lib/sqlalchemy/sql/visitors.py | 4 +++- 3 files changed, 4 insertions(+), 9 deletions(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index 27e883c86..ff041d5e4 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1080,7 +1080,7 @@ class SQLCompiler(engine.Compiled): def _transform_select_for_nested_joins(self, select): adapters = [] - traverse_options = {"cloned": {}} + traverse_options = {"cloned": {}, "unconditional": True} def visit_join(elem): if isinstance(elem.right, sql.FromGrouping): @@ -1095,19 +1095,13 @@ class SQLCompiler(engine.Compiled): c._label = c._key_label = c.name elem.right = selectable - import pdb - pdb.set_trace() adapter = sql_util.ClauseAdapter(selectable, traverse_options=traverse_options) - adapter.__traverse_options__.pop('stop_on') adapters.append(adapter) select = visitors.cloned_traverse(select, traverse_options, {"join": visit_join}) - if adapters: - import pdb - pdb.set_trace() for adap in reversed(adapters): select = adap.traverse(select) return select diff --git a/lib/sqlalchemy/sql/util.py b/lib/sqlalchemy/sql/util.py index 6a267752e..ffa07d3df 100644 --- a/lib/sqlalchemy/sql/util.py +++ b/lib/sqlalchemy/sql/util.py @@ -833,7 +833,6 @@ class ClauseAdapter(visitors.ReplacingCloningVisitor): return newcol def replace(self, col): - print "COL!", col if isinstance(col, expression.FromClause) and \ self.selectable.is_derived_from(col): return self.selectable diff --git a/lib/sqlalchemy/sql/visitors.py b/lib/sqlalchemy/sql/visitors.py index 31ac686e3..c5a45ffd4 100644 --- a/lib/sqlalchemy/sql/visitors.py +++ b/lib/sqlalchemy/sql/visitors.py @@ -286,10 +286,12 @@ def replacement_traverse(obj, opts, replace): cloned = util.column_dict() stop_on = util.column_set([id(x) for x in opts.get('stop_on', [])]) + unconditional = opts.get('unconditional', False) def clone(elem, **kw): if id(elem) in stop_on or \ - 'no_replacement_traverse' in elem._annotations: + (not unconditional + and 'no_replacement_traverse' in elem._annotations): return elem else: newelem = replace(elem) -- cgit v1.2.1 From 2b0b802b4973e84871d1f266cc2e8aac6c125d17 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Sun, 2 Jun 2013 20:23:03 -0400 Subject: working through tests.... --- lib/sqlalchemy/sql/compiler.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index ff041d5e4..d5ba64938 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1079,7 +1079,10 @@ class SQLCompiler(engine.Compiled): def _transform_select_for_nested_joins(self, select): adapters = [] + stop_on = [] + # test for "unconditional" - any statement with + # no_replacement_traverse setup, i.e. query.statement, from_self(), etc. traverse_options = {"cloned": {}, "unconditional": True} def visit_join(elem): @@ -1090,6 +1093,12 @@ class SQLCompiler(engine.Compiled): while adapters: adapt = adapters.pop(-1) selectable = adapt.traverse(selectable) + #stop_on.append(selectable) + + # test: see test_subquery_relations: + # CyclicalInheritingEagerTestTwo.test_integrate + stop_on.append(elem.left) + for c in selectable.c: c._label = c._key_label = c.name @@ -1097,6 +1106,7 @@ class SQLCompiler(engine.Compiled): elem.right = selectable adapter = sql_util.ClauseAdapter(selectable, traverse_options=traverse_options) + adapter.__traverse_options__['stop_on'].extend(stop_on) adapters.append(adapter) select = visitors.cloned_traverse(select, @@ -1119,7 +1129,9 @@ class SQLCompiler(engine.Compiled): positional_names=None, nested_join_translation=False, **kwargs): - #nested_join_translation = True + + if self.dialect.supports_right_nested_joins: + nested_join_translation = True if not nested_join_translation: transformed_select = self._transform_select_for_nested_joins(select) text = self.visit_select( -- cgit v1.2.1 From 55fa83fd39a0cd572e7d6426b059235d18a91e9d Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Mon, 3 Jun 2013 21:28:53 -0400 Subject: OK this is the broken version, need to think a lot more about this --- lib/sqlalchemy/sql/compiler.py | 45 +++++++++++++++++++++++++++++++++++++++++- lib/sqlalchemy/sql/util.py | 3 ++- 2 files changed, 46 insertions(+), 2 deletions(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index d5ba64938..6c0127ba2 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1083,7 +1083,49 @@ class SQLCompiler(engine.Compiled): # test for "unconditional" - any statement with # no_replacement_traverse setup, i.e. query.statement, from_self(), etc. - traverse_options = {"cloned": {}, "unconditional": True} + #traverse_options = {"cloned": {}, "unconditional": True} + traverse_options = {"unconditional": True} + + cloned = {} + def thing(element, **kw): + if element in cloned: + return cloned[element] + + newelem = cloned[element] = element._clone() + + if newelem.__visit_name__ == 'join' and \ + isinstance(newelem.right, sql.FromGrouping): + selectable = sql.select([newelem.right.element], use_labels=True) + selectable = selectable.alias() + newelem.right = selectable + stop_on.append(selectable) + for c in selectable.c: + c._label = c._key_label = c.name + adapter = sql_util.ClauseAdapter(selectable, + traverse_options=traverse_options) + adapter.magic_flag = True + adapters.append(adapter) + else: + newelem._copy_internals(clone=thing, **kw) + + return newelem + + elem = thing(select) + while adapters: + adapt = adapters.pop(-1) + adapt.__traverse_options__['stop_on'].extend(stop_on) + elem = adapt.traverse(elem) + return elem + + + def _transform_select_for_nested_joins_orig(self, select): + adapters = [] + stop_on = [] + + # test for "unconditional" - any statement with + # no_replacement_traverse setup, i.e. query.statement, from_self(), etc. + #traverse_options = {"cloned": {}, "unconditional": True} + traverse_options = {"unconditional": True} def visit_join(elem): if isinstance(elem.right, sql.FromGrouping): @@ -1109,6 +1151,7 @@ class SQLCompiler(engine.Compiled): adapter.__traverse_options__['stop_on'].extend(stop_on) adapters.append(adapter) + select = visitors.cloned_traverse(select, traverse_options, {"join": visit_join}) diff --git a/lib/sqlalchemy/sql/util.py b/lib/sqlalchemy/sql/util.py index ffa07d3df..c80693706 100644 --- a/lib/sqlalchemy/sql/util.py +++ b/lib/sqlalchemy/sql/util.py @@ -832,8 +832,9 @@ class ClauseAdapter(visitors.ReplacingCloningVisitor): newcol = self.selectable.c.get(col.name) return newcol + magic_flag = False def replace(self, col): - if isinstance(col, expression.FromClause) and \ + if not self.magic_flag and isinstance(col, expression.FromClause) and \ self.selectable.is_derived_from(col): return self.selectable elif not isinstance(col, expression.ColumnElement): -- cgit v1.2.1 From 822786dfaea7a56b16669561b4818ca1bf3a800f Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Tue, 4 Jun 2013 13:11:03 -0400 Subject: capture the really hard one in a test (hooray) --- lib/sqlalchemy/sql/compiler.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index 6c0127ba2..3e159b112 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1118,7 +1118,7 @@ class SQLCompiler(engine.Compiled): return elem - def _transform_select_for_nested_joins_orig(self, select): + def _transform_select_for_nested_joins(self, select): adapters = [] stop_on = [] -- cgit v1.2.1 From 9998e9e0131ff83a4e38e3c17a835a0854789174 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Tue, 4 Jun 2013 14:30:29 -0400 Subject: rewriting scheme now works. --- lib/sqlalchemy/sql/compiler.py | 114 ++++++++++++++++------------------------- 1 file changed, 43 insertions(+), 71 deletions(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index 3e159b112..d245c781a 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1078,86 +1078,58 @@ class SQLCompiler(engine.Compiled): return None def _transform_select_for_nested_joins(self, select): - adapters = [] - stop_on = [] - - # test for "unconditional" - any statement with - # no_replacement_traverse setup, i.e. query.statement, from_self(), etc. - #traverse_options = {"cloned": {}, "unconditional": True} - traverse_options = {"unconditional": True} + """Rewrite any "a JOIN (b JOIN c)" expression as + "a JOIN (select * from b JOIN c) AS anon", to support + databases that can't parse a parenthesized join correctly + (i.e. sqlite the main one). + """ cloned = {} - def thing(element, **kw): - if element in cloned: - return cloned[element] - - newelem = cloned[element] = element._clone() - - if newelem.__visit_name__ == 'join' and \ - isinstance(newelem.right, sql.FromGrouping): - selectable = sql.select([newelem.right.element], use_labels=True) - selectable = selectable.alias() - newelem.right = selectable - stop_on.append(selectable) - for c in selectable.c: - c._label = c._key_label = c.name - adapter = sql_util.ClauseAdapter(selectable, - traverse_options=traverse_options) - adapter.magic_flag = True - adapters.append(adapter) - else: - newelem._copy_internals(clone=thing, **kw) - - return newelem + column_translate = [{}] - elem = thing(select) - while adapters: - adapt = adapters.pop(-1) - adapt.__traverse_options__['stop_on'].extend(stop_on) - elem = adapt.traverse(elem) - return elem + join_name = sql.Join.__visit_name__ + select_name = sql.Select.__visit_name__ + def visit(element, **kw): + if element in column_translate[-1]: + return column_translate[-1][element] - def _transform_select_for_nested_joins(self, select): - adapters = [] - stop_on = [] - - # test for "unconditional" - any statement with - # no_replacement_traverse setup, i.e. query.statement, from_self(), etc. - #traverse_options = {"cloned": {}, "unconditional": True} - traverse_options = {"unconditional": True} + elif element in cloned: + return cloned[element] - def visit_join(elem): - if isinstance(elem.right, sql.FromGrouping): - selectable = sql.select([elem.right.element], use_labels=True) - selectable = selectable.alias() + newelem = cloned[element] = element._clone() - while adapters: - adapt = adapters.pop(-1) - selectable = adapt.traverse(selectable) - #stop_on.append(selectable) + if newelem.__visit_name__ is join_name and \ + isinstance(newelem.right, sql.FromGrouping): - # test: see test_subquery_relations: - # CyclicalInheritingEagerTestTwo.test_integrate - stop_on.append(elem.left) + newelem._reset_exported() + newelem.left = visit(newelem.left, **kw) + selectable = sql.select( + [newelem.right.element], + use_labels=True).alias() for c in selectable.c: c._label = c._key_label = c.name + translate_dict = dict( + zip(newelem.right.element.c, selectable.c) + ) + translate_dict[newelem.right.element.left] = selectable + translate_dict[newelem.right.element.right] = selectable + column_translate[-1].update(translate_dict) - elem.right = selectable - adapter = sql_util.ClauseAdapter(selectable, - traverse_options=traverse_options) - adapter.__traverse_options__['stop_on'].extend(stop_on) - adapters.append(adapter) - + newelem.right = selectable + newelem.onclause = visit(newelem.onclause, **kw) + elif newelem.__visit_name__ is select_name: + column_translate.append({}) + newelem._copy_internals(clone=visit, **kw) + del column_translate[-1] + else: + newelem._copy_internals(clone=visit, **kw) - select = visitors.cloned_traverse(select, - traverse_options, {"join": visit_join}) + return newelem - for adap in reversed(adapters): - select = adap.traverse(select) - return select + return visit(select) def _transform_result_map_for_nested_joins(self, select, transformed_select): d = dict(zip(transformed_select.inner_columns, select.inner_columns)) @@ -1172,10 +1144,12 @@ class SQLCompiler(engine.Compiled): positional_names=None, nested_join_translation=False, **kwargs): + needs_nested_translation = \ + not nested_join_translation and \ + not self.stack and \ + not self.dialect.supports_right_nested_joins - if self.dialect.supports_right_nested_joins: - nested_join_translation = True - if not nested_join_translation: + if needs_nested_translation: transformed_select = self._transform_select_for_nested_joins(select) text = self.visit_select( transformed_select, asfrom=asfrom, parens=parens, @@ -1186,8 +1160,6 @@ class SQLCompiler(engine.Compiled): nested_join_translation=True, **kwargs ) - - entry = self.stack and self.stack[-1] or {} populate_result_map = force_result_map or ( @@ -1197,7 +1169,7 @@ class SQLCompiler(engine.Compiled): ) ) - if not nested_join_translation: + if needs_nested_translation: if populate_result_map: self._transform_result_map_for_nested_joins( select, transformed_select) -- cgit v1.2.1 From fa3d18a47cb6f307b18da880a5f8f4c06a6023b4 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Tue, 4 Jun 2013 14:31:56 -0400 Subject: and this comment --- lib/sqlalchemy/sql/compiler.py | 4 ++++ 1 file changed, 4 insertions(+) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index d245c781a..2024666b6 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1087,6 +1087,10 @@ class SQLCompiler(engine.Compiled): cloned = {} column_translate = [{}] + # TODO: should we be using isinstance() for this, + # as this whole system won't work for custom Join/Select + # subclasses where compilation routines + # call down to compiler.visit_join(), compiler.visit_select() join_name = sql.Join.__visit_name__ select_name = sql.Select.__visit_name__ -- cgit v1.2.1 From 51e1019f610f083ac4d8c850589cdf52cff044da Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Tue, 4 Jun 2013 16:21:25 -0400 Subject: here's the flat join thing. it just works. Changing the existing compiled SQL assertions might even be most of the tests we need (though dedicated sql tests would be needed anyway) --- lib/sqlalchemy/sql/expression.py | 19 ++++++++++++++----- 1 file changed, 14 insertions(+), 5 deletions(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/expression.py b/lib/sqlalchemy/sql/expression.py index edab9e290..633a3ddba 100644 --- a/lib/sqlalchemy/sql/expression.py +++ b/lib/sqlalchemy/sql/expression.py @@ -795,7 +795,7 @@ def intersect_all(*selects, **kwargs): return CompoundSelect(CompoundSelect.INTERSECT_ALL, *selects, **kwargs) -def alias(selectable, name=None): +def alias(selectable, name=None, flat=False): """Return an :class:`.Alias` object. An :class:`.Alias` represents any :class:`.FromClause` @@ -2634,7 +2634,7 @@ class FromClause(Selectable): return Join(self, right, onclause, True) - def alias(self, name=None): + def alias(self, name=None, flat=False): """return an alias of this :class:`.FromClause`. This is shorthand for calling:: @@ -3971,7 +3971,7 @@ class Join(FromClause): def bind(self): return self.left.bind or self.right.bind - def alias(self, name=None): + def alias(self, name=None, flat=False): """return an alias of this :class:`.Join`. Used against a :class:`.Join` object, @@ -3999,7 +3999,16 @@ class Join(FromClause): aliases. """ - return self.select(use_labels=True, correlate=False).alias(name) + if flat: + assert name is None, "Can't send name argument with flat" + left_a, right_a = self.left.alias(), self.right.alias() + adapter = sqlutil.ClauseAdapter(left_a).\ + chain(sqlutil.ClauseAdapter(right_a)) + + return left_a.join(right_a, + adapter.traverse(self.onclause), isouter=self.isouter) + else: + return self.select(use_labels=True, correlate=False).alias(name) @property def _hide_froms(self): @@ -4129,7 +4138,7 @@ class CTE(Alias): self._restates = _restates super(CTE, self).__init__(selectable, name=name) - def alias(self, name=None): + def alias(self, name=None, flat=False): return CTE( self.original, name=name, -- cgit v1.2.1 From 11578ba709ff3e5f2a2a2a9f92bf6fdc2ee6d328 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Tue, 4 Jun 2013 18:23:06 -0400 Subject: - improve overlapping selectables, apply to both query and relationship - clean up inspect() calls within query._join() - make sure join.alias(flat) propagates - fix almost all assertion tests --- lib/sqlalchemy/sql/expression.py | 3 ++- lib/sqlalchemy/sql/util.py | 23 ++++++++++++++++++----- 2 files changed, 20 insertions(+), 6 deletions(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/expression.py b/lib/sqlalchemy/sql/expression.py index 633a3ddba..e7ef3cb72 100644 --- a/lib/sqlalchemy/sql/expression.py +++ b/lib/sqlalchemy/sql/expression.py @@ -4001,7 +4001,8 @@ class Join(FromClause): """ if flat: assert name is None, "Can't send name argument with flat" - left_a, right_a = self.left.alias(), self.right.alias() + left_a, right_a = self.left.alias(flat=True), \ + self.right.alias(flat=True) adapter = sqlutil.ClauseAdapter(left_a).\ chain(sqlutil.ClauseAdapter(right_a)) diff --git a/lib/sqlalchemy/sql/util.py b/lib/sqlalchemy/sql/util.py index c80693706..6f4d27e1b 100644 --- a/lib/sqlalchemy/sql/util.py +++ b/lib/sqlalchemy/sql/util.py @@ -200,15 +200,28 @@ def clause_is_present(clause, search): """ - stack = [search] - while stack: - elem = stack.pop() + for elem in surface_selectables(search): if clause == elem: # use == here so that Annotated's compare return True - elif isinstance(elem, expression.Join): + else: + return False + +def surface_selectables(clause): + stack = [clause] + while stack: + elem = stack.pop() + yield elem + if isinstance(elem, expression.Join): stack.extend((elem.left, elem.right)) - return False +def selectables_overlap(left, right): + """Return True if left/right have some overlapping selectable""" + + return bool( + set(surface_selectables(left)).intersection( + surface_selectables(right) + ) + ) def bind_values(clause): """Return an ordered list of "bound" values in the given clause. -- cgit v1.2.1 From d8a38839483aede934b6cbeb6d0828d362767a4d Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Tue, 4 Jun 2013 19:44:57 -0400 Subject: - support for a__b_dc, i.e. two levels of nesting --- lib/sqlalchemy/sql/compiler.py | 23 +++++++++++++++++++---- 1 file changed, 19 insertions(+), 4 deletions(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index 2024666b6..116fb3971 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1109,17 +1109,32 @@ class SQLCompiler(engine.Compiled): newelem._reset_exported() newelem.left = visit(newelem.left, **kw) + right = visit(newelem.right, **kw) + selectable = sql.select( - [newelem.right.element], + [right.element], use_labels=True).alias() for c in selectable.c: c._label = c._key_label = c.name translate_dict = dict( - zip(newelem.right.element.c, selectable.c) + zip(right.element.c, selectable.c) ) - translate_dict[newelem.right.element.left] = selectable - translate_dict[newelem.right.element.right] = selectable + translate_dict[right.element.left] = selectable + translate_dict[right.element.right] = selectable + + # propagate translations that we've gained + # from nested visit(newelem.right) outwards + # to the enclosing select here. this happens + # only when we have more than one level of right + # join nesting, i.e. "a JOIN (b JOIN (c JOIN d))" + for k, v in list(column_translate[-1].items()): + if v in translate_dict: + # remarkably, no current ORM tests (May 2013) + # hit this condition, only test_join_rewriting + # does. + column_translate[-1][k] = translate_dict[v] + column_translate[-1].update(translate_dict) newelem.right = selectable -- cgit v1.2.1 From 92e599f42f4a7dde1662fe0ae428d32bb4e8cc42 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Tue, 4 Jun 2013 19:52:53 -0400 Subject: - if the select() does not have use_labels on, then we just render the joins as is, regardless of the dialect not supporting it. use_labels=True indicates a higher level of automation and also can maintain the labels between rewritten and not. use_labels=False indicates a manual use case. --- lib/sqlalchemy/sql/compiler.py | 1 + 1 file changed, 1 insertion(+) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index 116fb3971..af70d13fa 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1164,6 +1164,7 @@ class SQLCompiler(engine.Compiled): nested_join_translation=False, **kwargs): needs_nested_translation = \ + select.use_labels and \ not nested_join_translation and \ not self.stack and \ not self.dialect.supports_right_nested_joins -- cgit v1.2.1 From 69e9574fefd5fbb4673c99ad476a00b03fe22318 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Tue, 4 Jun 2013 21:36:34 -0400 Subject: - add coverage for result map rewriting - fix the result map rewriter for col mismatches, since the rewritten select at the moment typically has more columns than the original --- lib/sqlalchemy/sql/compiler.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) (limited to 'lib/sqlalchemy/sql') diff --git a/lib/sqlalchemy/sql/compiler.py b/lib/sqlalchemy/sql/compiler.py index af70d13fa..c29b45450 100644 --- a/lib/sqlalchemy/sql/compiler.py +++ b/lib/sqlalchemy/sql/compiler.py @@ -1151,7 +1151,12 @@ class SQLCompiler(engine.Compiled): return visit(select) def _transform_result_map_for_nested_joins(self, select, transformed_select): - d = dict(zip(transformed_select.inner_columns, select.inner_columns)) + inner_col = dict((c._key_label, c) for + c in transformed_select.inner_columns) + d = dict( + (inner_col[c._key_label], c) + for c in select.inner_columns + ) for key, (name, objs, typ) in list(self.result_map.items()): objs = tuple([d.get(col, col) for col in objs]) self.result_map[key] = (name, objs, typ) -- cgit v1.2.1