diff options
Diffstat (limited to 'django/db/models/sql/compiler.py')
| -rw-r--r-- | django/db/models/sql/compiler.py | 51 |
1 files changed, 47 insertions, 4 deletions
diff --git a/django/db/models/sql/compiler.py b/django/db/models/sql/compiler.py index 11ff51f60f..14d44d3eef 100644 --- a/django/db/models/sql/compiler.py +++ b/django/db/models/sql/compiler.py @@ -702,7 +702,7 @@ class SQLCompiler: """ result = [] params = [] - for alias in self.query.alias_map: + for alias in tuple(self.query.alias_map): if not self.query.alias_refcount[alias]: continue try: @@ -737,7 +737,7 @@ class SQLCompiler: f.field.related_query_name() for f in opts.related_objects if f.field.unique ) - return chain(direct_choices, reverse_choices) + return chain(direct_choices, reverse_choices, self.query._filtered_relations) related_klass_infos = [] if not restricted and cur_depth > self.query.max_depth: @@ -788,7 +788,8 @@ class SQLCompiler: klass_info = { 'model': f.remote_field.model, 'field': f, - 'reverse': False, + 'local_setter': f.set_cached_value, + 'remote_setter': f.remote_field.set_cached_value if f.unique else lambda x, y: None, 'from_parent': False, } related_klass_infos.append(klass_info) @@ -825,7 +826,8 @@ class SQLCompiler: klass_info = { 'model': model, 'field': f, - 'reverse': True, + 'local_setter': f.remote_field.set_cached_value, + 'remote_setter': f.set_cached_value, 'from_parent': from_parent, } related_klass_infos.append(klass_info) @@ -842,6 +844,47 @@ class SQLCompiler: next, restricted) get_related_klass_infos(klass_info, next_klass_infos) fields_not_found = set(requested).difference(fields_found) + for name in list(requested): + # Filtered relations work only on the topmost level. + if cur_depth > 1: + break + if name in self.query._filtered_relations: + fields_found.add(name) + f, _, join_opts, joins, _ = self.query.setup_joins([name], opts, root_alias) + model = join_opts.model + alias = joins[-1] + from_parent = issubclass(model, opts.model) and model is not opts.model + + def local_setter(obj, from_obj): + f.remote_field.set_cached_value(from_obj, obj) + + def remote_setter(obj, from_obj): + setattr(from_obj, name, obj) + klass_info = { + 'model': model, + 'field': f, + 'local_setter': local_setter, + 'remote_setter': remote_setter, + 'from_parent': from_parent, + } + related_klass_infos.append(klass_info) + select_fields = [] + columns = self.get_default_columns( + start_alias=alias, opts=model._meta, + from_parent=opts.model, + ) + for col in columns: + select_fields.append(len(select)) + select.append((col, None)) + klass_info['select_fields'] = select_fields + next_requested = requested.get(name, {}) + next_klass_infos = self.get_related_selections( + select, opts=model._meta, root_alias=alias, + cur_depth=cur_depth + 1, requested=next_requested, + restricted=restricted, + ) + get_related_klass_infos(klass_info, next_klass_infos) + fields_not_found = set(requested).difference(fields_found) if fields_not_found: invalid_fields = ("'%s'" % s for s in fields_not_found) raise FieldError( |
