diff options
Diffstat (limited to 'tests')
| -rw-r--r-- | tests/aggregation/tests.py | 14 | ||||
| -rw-r--r-- | tests/custom_lookups/tests.py | 72 | ||||
| -rw-r--r-- | tests/foreign_object/models.py | 3 | ||||
| -rw-r--r-- | tests/queries/tests.py | 64 |
4 files changed, 77 insertions, 76 deletions
diff --git a/tests/aggregation/tests.py b/tests/aggregation/tests.py index 8c6529c73b..e4b821b43d 100644 --- a/tests/aggregation/tests.py +++ b/tests/aggregation/tests.py @@ -863,8 +863,8 @@ class ComplexAggregateTestCase(TestCase): def test_add_implementation(self): try: # test completely changing how the output is rendered - def lower_case_function_override(self, qn, connection): - sql, params = qn.compile(self.source_expressions[0]) + def lower_case_function_override(self, compiler, connection): + sql, params = compiler.compile(self.source_expressions[0]) substitutions = dict(function=self.function.lower(), expressions=sql) substitutions.update(self.extra) return self.template % substitutions, params @@ -877,9 +877,9 @@ class ComplexAggregateTestCase(TestCase): self.assertEqual(b1.sums, 383) # test changing the dict and delegating - def lower_case_function_super(self, qn, connection): + def lower_case_function_super(self, compiler, connection): self.extra['function'] = self.function.lower() - return super(Sum, self).as_sql(qn, connection) + return super(Sum, self).as_sql(compiler, connection) setattr(Sum, 'as_' + connection.vendor, lower_case_function_super) qs = Book.objects.annotate(sums=Sum(F('rating') + F('pages') + F('price'), @@ -889,7 +889,7 @@ class ComplexAggregateTestCase(TestCase): self.assertEqual(b1.sums, 383) # test overriding all parts of the template - def be_evil(self, qn, connection): + def be_evil(self, compiler, connection): substitutions = dict(function='MAX', expressions='2') substitutions.update(self.extra) return self.template % substitutions, () @@ -921,8 +921,8 @@ class ComplexAggregateTestCase(TestCase): class Greatest(Func): function = 'GREATEST' - def as_sqlite(self, qn, connection): - return super(Greatest, self).as_sql(qn, connection, function='MAX') + def as_sqlite(self, compiler, connection): + return super(Greatest, self).as_sql(compiler, connection, function='MAX') qs = Publisher.objects.annotate( price_or_median=Greatest(Avg('book__rating'), Avg('book__price')) diff --git a/tests/custom_lookups/tests.py b/tests/custom_lookups/tests.py index 872427e883..620db613a1 100644 --- a/tests/custom_lookups/tests.py +++ b/tests/custom_lookups/tests.py @@ -14,15 +14,15 @@ from .models import Author, MySQLUnixTimestamp class Div3Lookup(models.Lookup): lookup_name = 'div3' - def as_sql(self, qn, connection): - lhs, params = self.process_lhs(qn, connection) - rhs, rhs_params = self.process_rhs(qn, connection) + def as_sql(self, compiler, connection): + lhs, params = self.process_lhs(compiler, connection) + rhs, rhs_params = self.process_rhs(compiler, connection) params.extend(rhs_params) return '(%s) %%%% 3 = %s' % (lhs, rhs), params - def as_oracle(self, qn, connection): - lhs, params = self.process_lhs(qn, connection) - rhs, rhs_params = self.process_rhs(qn, connection) + def as_oracle(self, compiler, connection): + lhs, params = self.process_lhs(compiler, connection) + rhs, rhs_params = self.process_rhs(compiler, connection) params.extend(rhs_params) return 'mod(%s, 3) = %s' % (lhs, rhs), params @@ -30,12 +30,12 @@ class Div3Lookup(models.Lookup): class Div3Transform(models.Transform): lookup_name = 'div3' - def as_sql(self, qn, connection): - lhs, lhs_params = qn.compile(self.lhs) + def as_sql(self, compiler, connection): + lhs, lhs_params = compiler.compile(self.lhs) return '(%s) %%%% 3' % lhs, lhs_params - def as_oracle(self, qn, connection): - lhs, lhs_params = qn.compile(self.lhs) + def as_oracle(self, compiler, connection): + lhs, lhs_params = compiler.compile(self.lhs) return 'mod(%s, 3)' % lhs, lhs_params @@ -47,8 +47,8 @@ class Mult3BilateralTransform(models.Transform): bilateral = True lookup_name = 'mult3' - def as_sql(self, qn, connection): - lhs, lhs_params = qn.compile(self.lhs) + def as_sql(self, compiler, connection): + lhs, lhs_params = compiler.compile(self.lhs) return '3 * (%s)' % lhs, lhs_params @@ -56,16 +56,16 @@ class UpperBilateralTransform(models.Transform): bilateral = True lookup_name = 'upper' - def as_sql(self, qn, connection): - lhs, lhs_params = qn.compile(self.lhs) + def as_sql(self, compiler, connection): + lhs, lhs_params = compiler.compile(self.lhs) return 'UPPER(%s)' % lhs, lhs_params class YearTransform(models.Transform): lookup_name = 'year' - def as_sql(self, qn, connection): - lhs_sql, params = qn.compile(self.lhs) + def as_sql(self, compiler, connection): + lhs_sql, params = compiler.compile(self.lhs) return connection.ops.date_extract_sql('year', lhs_sql), params @property @@ -77,11 +77,11 @@ class YearTransform(models.Transform): class YearExact(models.lookups.Lookup): lookup_name = 'exact' - def as_sql(self, qn, connection): + def as_sql(self, compiler, connection): # We will need to skip the extract part, and instead go # directly with the originating field, that is self.lhs.lhs - lhs_sql, lhs_params = self.process_lhs(qn, connection, self.lhs.lhs) - rhs_sql, rhs_params = self.process_rhs(qn, connection) + lhs_sql, lhs_params = self.process_lhs(compiler, connection, self.lhs.lhs) + rhs_sql, rhs_params = self.process_rhs(compiler, connection) # Note that we must be careful so that we have params in the # same order as we have the parts in the SQL. params = lhs_params + rhs_params + lhs_params + rhs_params @@ -98,12 +98,12 @@ class YearLte(models.lookups.LessThanOrEqual): The purpose of this lookup is to efficiently compare the year of the field. """ - def as_sql(self, qn, connection): + def as_sql(self, compiler, connection): # Skip the YearTransform above us (no possibility for efficient # lookup otherwise). real_lhs = self.lhs.lhs - lhs_sql, params = self.process_lhs(qn, connection, real_lhs) - rhs_sql, rhs_params = self.process_rhs(qn, connection) + lhs_sql, params = self.process_lhs(compiler, connection, real_lhs) + rhs_sql, rhs_params = self.process_rhs(compiler, connection) params.extend(rhs_params) # Build SQL where the integer year is concatenated with last month # and day, then convert that to date. (We try to have SQL like: @@ -117,7 +117,7 @@ class SQLFunc(models.Lookup): super(SQLFunc, self).__init__(*args, **kwargs) self.name = name - def as_sql(self, qn, connection): + def as_sql(self, compiler, connection): return '%s()', [self.name] @property @@ -162,9 +162,9 @@ class InMonth(models.lookups.Lookup): """ lookup_name = 'inmonth' - def as_sql(self, qn, connection): - lhs, lhs_params = self.process_lhs(qn, connection) - rhs, rhs_params = self.process_rhs(qn, connection) + def as_sql(self, compiler, connection): + lhs, lhs_params = self.process_lhs(compiler, connection) + rhs, rhs_params = self.process_rhs(compiler, connection) # We need to be careful so that we get the params in right # places. params = lhs_params + rhs_params + lhs_params + rhs_params @@ -180,8 +180,8 @@ class DateTimeTransform(models.Transform): def output_field(self): return models.DateTimeField() - def as_sql(self, qn, connection): - lhs, params = qn.compile(self.lhs) + def as_sql(self, compiler, connection): + lhs, params = compiler.compile(self.lhs) return 'from_unixtime({})'.format(lhs), params @@ -448,9 +448,9 @@ class YearLteTests(TestCase): try: # Two ways to add a customized implementation for different backends: # First is MonkeyPatch of the class. - def as_custom_sql(self, qn, connection): - lhs_sql, lhs_params = self.process_lhs(qn, connection, self.lhs.lhs) - rhs_sql, rhs_params = self.process_rhs(qn, connection) + def as_custom_sql(self, compiler, connection): + lhs_sql, lhs_params = self.process_lhs(compiler, connection, self.lhs.lhs) + rhs_sql, rhs_params = self.process_rhs(compiler, connection) params = lhs_params + rhs_params + lhs_params + rhs_params return ("%(lhs)s >= str_to_date(concat(%(rhs)s, '-01-01'), '%%%%Y-%%%%m-%%%%d') " "AND %(lhs)s <= str_to_date(concat(%(rhs)s, '-12-31'), '%%%%Y-%%%%m-%%%%d')" % @@ -468,9 +468,9 @@ class YearLteTests(TestCase): # This method should be named "as_mysql" for MySQL, "as_postgresql" for postgres # and so on, but as we don't know which DB we are running on, we need to use # setattr. - def as_custom_sql(self, qn, connection): - lhs_sql, lhs_params = self.process_lhs(qn, connection, self.lhs.lhs) - rhs_sql, rhs_params = self.process_rhs(qn, connection) + def as_custom_sql(self, compiler, connection): + lhs_sql, lhs_params = self.process_lhs(compiler, connection, self.lhs.lhs) + rhs_sql, rhs_params = self.process_rhs(compiler, connection) params = lhs_params + rhs_params + lhs_params + rhs_params return ("%(lhs)s >= str_to_date(CONCAT(%(rhs)s, '-01-01'), '%%%%Y-%%%%m-%%%%d') " "AND %(lhs)s <= str_to_date(CONCAT(%(rhs)s, '-12-31'), '%%%%Y-%%%%m-%%%%d')" % @@ -489,8 +489,8 @@ class TrackCallsYearTransform(YearTransform): lookup_name = 'year' call_order = [] - def as_sql(self, qn, connection): - lhs_sql, params = qn.compile(self.lhs) + def as_sql(self, compiler, connection): + lhs_sql, params = compiler.compile(self.lhs) return connection.ops.date_extract_sql('year', lhs_sql), params @property diff --git a/tests/foreign_object/models.py b/tests/foreign_object/models.py index fc51118149..07d9ff4450 100644 --- a/tests/foreign_object/models.py +++ b/tests/foreign_object/models.py @@ -117,7 +117,8 @@ class ColConstraint(object): def __init__(self, alias, col, value): self.alias, self.col, self.value = alias, col, value - def as_sql(self, qn, connection): + def as_sql(self, compiler, connection): + qn = compiler.quote_name_unless_alias return '%s.%s = %%s' % (qn(self.alias), qn(self.col)), [self.value] diff --git a/tests/queries/tests.py b/tests/queries/tests.py index 5ee7c85005..7b4766519a 100644 --- a/tests/queries/tests.py +++ b/tests/queries/tests.py @@ -2817,7 +2817,7 @@ class ProxyQueryCleanupTest(TestCase): class WhereNodeTest(TestCase): class DummyNode(object): - def as_sql(self, qn, connection): + def as_sql(self, compiler, connection): return 'dummy', [] class MockCompiler(object): @@ -2828,70 +2828,70 @@ class WhereNodeTest(TestCase): return connection.ops.quote_name(name) def test_empty_full_handling_conjunction(self): - qn = WhereNodeTest.MockCompiler() + compiler = WhereNodeTest.MockCompiler() w = WhereNode(children=[EverythingNode()]) - self.assertEqual(w.as_sql(qn, connection), ('', [])) + self.assertEqual(w.as_sql(compiler, connection), ('', [])) w.negate() - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) w = WhereNode(children=[NothingNode()]) - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) w.negate() - self.assertEqual(w.as_sql(qn, connection), ('', [])) + self.assertEqual(w.as_sql(compiler, connection), ('', [])) w = WhereNode(children=[EverythingNode(), EverythingNode()]) - self.assertEqual(w.as_sql(qn, connection), ('', [])) + self.assertEqual(w.as_sql(compiler, connection), ('', [])) w.negate() - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) w = WhereNode(children=[EverythingNode(), self.DummyNode()]) - self.assertEqual(w.as_sql(qn, connection), ('dummy', [])) + self.assertEqual(w.as_sql(compiler, connection), ('dummy', [])) w = WhereNode(children=[self.DummyNode(), self.DummyNode()]) - self.assertEqual(w.as_sql(qn, connection), ('(dummy AND dummy)', [])) + self.assertEqual(w.as_sql(compiler, connection), ('(dummy AND dummy)', [])) w.negate() - self.assertEqual(w.as_sql(qn, connection), ('NOT (dummy AND dummy)', [])) + self.assertEqual(w.as_sql(compiler, connection), ('NOT (dummy AND dummy)', [])) w = WhereNode(children=[NothingNode(), self.DummyNode()]) - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) w.negate() - self.assertEqual(w.as_sql(qn, connection), ('', [])) + self.assertEqual(w.as_sql(compiler, connection), ('', [])) def test_empty_full_handling_disjunction(self): - qn = WhereNodeTest.MockCompiler() + compiler = WhereNodeTest.MockCompiler() w = WhereNode(children=[EverythingNode()], connector='OR') - self.assertEqual(w.as_sql(qn, connection), ('', [])) + self.assertEqual(w.as_sql(compiler, connection), ('', [])) w.negate() - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) w = WhereNode(children=[NothingNode()], connector='OR') - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) w.negate() - self.assertEqual(w.as_sql(qn, connection), ('', [])) + self.assertEqual(w.as_sql(compiler, connection), ('', [])) w = WhereNode(children=[EverythingNode(), EverythingNode()], connector='OR') - self.assertEqual(w.as_sql(qn, connection), ('', [])) + self.assertEqual(w.as_sql(compiler, connection), ('', [])) w.negate() - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) w = WhereNode(children=[EverythingNode(), self.DummyNode()], connector='OR') - self.assertEqual(w.as_sql(qn, connection), ('', [])) + self.assertEqual(w.as_sql(compiler, connection), ('', [])) w.negate() - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) w = WhereNode(children=[self.DummyNode(), self.DummyNode()], connector='OR') - self.assertEqual(w.as_sql(qn, connection), ('(dummy OR dummy)', [])) + self.assertEqual(w.as_sql(compiler, connection), ('(dummy OR dummy)', [])) w.negate() - self.assertEqual(w.as_sql(qn, connection), ('NOT (dummy OR dummy)', [])) + self.assertEqual(w.as_sql(compiler, connection), ('NOT (dummy OR dummy)', [])) w = WhereNode(children=[NothingNode(), self.DummyNode()], connector='OR') - self.assertEqual(w.as_sql(qn, connection), ('dummy', [])) + self.assertEqual(w.as_sql(compiler, connection), ('dummy', [])) w.negate() - self.assertEqual(w.as_sql(qn, connection), ('NOT (dummy)', [])) + self.assertEqual(w.as_sql(compiler, connection), ('NOT (dummy)', [])) def test_empty_nodes(self): - qn = WhereNodeTest.MockCompiler() + compiler = WhereNodeTest.MockCompiler() empty_w = WhereNode() w = WhereNode(children=[empty_w, empty_w]) - self.assertEqual(w.as_sql(qn, connection), (None, [])) + self.assertEqual(w.as_sql(compiler, connection), (None, [])) w.negate() - self.assertEqual(w.as_sql(qn, connection), (None, [])) + self.assertEqual(w.as_sql(compiler, connection), (None, [])) w.connector = 'OR' - self.assertEqual(w.as_sql(qn, connection), (None, [])) + self.assertEqual(w.as_sql(compiler, connection), (None, [])) w.negate() - self.assertEqual(w.as_sql(qn, connection), (None, [])) + self.assertEqual(w.as_sql(compiler, connection), (None, [])) w = WhereNode(children=[empty_w, NothingNode()], connector='OR') - self.assertRaises(EmptyResultSet, w.as_sql, qn, connection) + self.assertRaises(EmptyResultSet, w.as_sql, compiler, connection) class IteratorExceptionsTest(TestCase): |
