summaryrefslogtreecommitdiff
path: root/tests
diff options
context:
space:
mode:
Diffstat (limited to 'tests')
-rw-r--r--tests/aggregation/tests.py14
-rw-r--r--tests/custom_lookups/tests.py72
-rw-r--r--tests/foreign_object/models.py3
-rw-r--r--tests/queries/tests.py64
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):