Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
164 changes: 142 additions & 22 deletions django_ormql/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -185,6 +185,99 @@ class Generator(Generator):
}


# Maps Django's Field.get_internal_type() to the ORMQL sql_type vocabulary
# emitted by BaseColumn.sql_type subclasses (see columns.py). Keeping one
# vocabulary means /tables/ introspection and per-query result metadata speak
# the same language.
INTERNAL_TYPE_TO_SQL_TYPE = {
"DateField": "DATE",
"DateTimeField": "DATETIME",
"TimeField": "TIME",
"DurationField": "DURATION",
"IntegerField": "INT",
"BigIntegerField": "INT",
"SmallIntegerField": "INT",
"PositiveIntegerField": "INT",
"PositiveSmallIntegerField": "INT",
"PositiveBigIntegerField": "INT",
"AutoField": "INT",
"BigAutoField": "INT",
"SmallAutoField": "INT",
"DecimalField": "DECIMAL",
"FloatField": "FLOAT",
"BooleanField": "BOOLEAN",
"JSONField": "JSONB",
"CharField": "TEXT",
"TextField": "TEXT",
"EmailField": "TEXT",
"URLField": "TEXT",
"SlugField": "TEXT",
"UUIDField": "TEXT",
"GenericIPAddressField": "TEXT",
"FileField": "TEXT",
"FilePathField": "TEXT",
"ImageField": "TEXT",
"BinaryField": "TEXT",
}


def _describe_expression(expr, query=None):
"""Extract (sql_type, nullable, hint) from a Django expression's output_field.
Hint is dependent on type and currently only available for truncated datetimes.

Aggregates and Cast() / typed Func() have `.output_field` directly. Plain
F() references only know their field once resolved against a query; when
`query` is supplied, we resolve on a cloned query (so we don't mutate the
original with extra joins) and read `output_field` off the resolved node.

Returns ("", None, None) for anything we can't pin down — the frontend treats
that as "unknown, render opaquely" which matches the pre-metadata
behavior.
"""
out = None
try:
out = expr.output_field
except: # noqa
out = None
if out is None and query is not None:
try:
resolved = expr.resolve_expression(query=query.chain(), allow_joins=True)
out = resolved.output_field
except: # noqa
out = None
if out is None:
return "", None
try:
internal = out.get_internal_type()
except: # noqa
return "", None
sql_type = INTERNAL_TYPE_TO_SQL_TYPE.get(internal, "")
nullable = getattr(out, "null", None)
hint = None

if isinstance(expr, functions.datetime.TruncBase):
hint = {"truncated": expr.kind}

return sql_type, nullable, hint


class Result:
"""Iterable query result carrying per-query column metadata.

Iterating yields row dicts just like the previous evaluate() generator, so
existing `for row in result:` / `list(result)` callers keep working. The
`columns` attribute is a list of {"name", "type", "nullable"} dicts using
the sql_type vocabulary from INTERNAL_TYPE_TO_SQL_TYPE.
"""

def __init__(self, rows, columns):
self._rows = rows
self.columns = columns

def __iter__(self):
return iter(self._rows)


class Query:
def __init__(self, sql, tables, placeholders, timezone, default_limit):
self.sql = sql
Expand Down Expand Up @@ -598,7 +691,7 @@ def _resolve(e, parent_stack, depth):
elif isinstance(expression, expressions.Subquery):
if not isinstance(expression.this, expressions.Select):
raise QueryNotSupported("Only SELECT subqueries are supported")
qs, _ = self._select_to_qs(
qs, _, _ = self._select_to_qs(
expression.this, parent_table_stack=parent_table_stack + [table]
)
return db_func.AutoTypedSubquery(
Expand All @@ -607,7 +700,7 @@ def _resolve(e, parent_stack, depth):
elif isinstance(expression, expressions.Exists):
if not isinstance(expression.this, expressions.Select):
raise QueryNotSupported("Only SELECT subqueries are supported")
qs, _ = self._select_to_qs(
qs, _, _ = self._select_to_qs(
expression.this, parent_table_stack=parent_table_stack + [table]
)
return models.Exists(qs)
Expand Down Expand Up @@ -719,6 +812,7 @@ def _select_to_qs(self, root, parent_table_stack):
values_names = {}
aggregations = {}
name_to_aggregation = {}
column_types = {}
for i, e in enumerate(root.args["expressions"]):
if isinstance(e, expressions.Star):
raise QueryNotSupported("SELECT * is not supported")
Expand Down Expand Up @@ -750,6 +844,9 @@ def _select_to_qs(self, root, parent_table_stack):
parent_table_stack=parent_table_stack,
)
values_names[f"expr{i}"] = n
column_types[f"expr{i}"] = _describe_expression(
django_e, query=qs.query
)

if root.args.get("distinct"):
qs = qs.distinct()
Expand Down Expand Up @@ -858,7 +955,7 @@ def _select_to_qs(self, root, parent_table_stack):
elif offset is not None or limit is not None:
qs = qs[offset:limit]

return qs, values_names
return qs, values_names, column_types

def _flatten_unions(self, root):
if isinstance(root, expressions.Select):
Expand Down Expand Up @@ -895,7 +992,15 @@ def parse(self):
try:
queries = self._flatten_unions(ast)
results = [self._select_to_qs(query, []) for query in queries]
if len({len(values_names.keys()) for qs, values_names in results}) != 1:
if (
len(
{
len(values_names.keys())
for qs, values_names, column_types in results
}
)
!= 1
):
raise QueryError(
"All parts of UNION query must return same number of columns"
)
Expand All @@ -905,24 +1010,39 @@ def parse(self):
raise QueryError("Invalid combination of types") from e
except Exception as e:
raise QueryError("Query parsing failed") from e

return [qs for qs, values_names in results], results[0][1]
return (
[qs for qs, values_names, column_types in results],
results[0][1],
results[0][2],
)

def evaluate(self):
querysets, values_names = self.parse()
querysets, values_names, column_types = self.parse()
columns = [
{
"name": values_names[k],
"type": column_types.get(k, ("", None))[0],
"nullable": column_types.get(k, ("", None))[1],
"hint": column_types.get(k, ("", None))[2],
}
for k in values_names
]

for qs in querysets:
if isinstance(qs, dict):
yield {values_names[k]: v for k, v in qs.items()}
else:
try:
if settings.DEBUG:
print(f"Generated statement: {qs.query!s}")
for row in qs:
yield {
values_names[k]: v
for k, v in row.items()
if k in values_names
}
except (FieldError, ValueError) as e:
raise QueryError("Invalid combination of types") from e
def _iter():
for qs in querysets:
if isinstance(qs, dict):
yield {values_names[k]: v for k, v in qs.items()}
else:
try:
if settings.DEBUG:
print(f"Generated statement: {qs.query!s}")
for row in qs:
yield {
values_names[k]: v
for k, v in row.items()
if k in values_names
}
except (FieldError, ValueError) as e:
raise QueryError("Invalid combination of types") from e

return Result(_iter(), columns)
141 changes: 141 additions & 0 deletions tests/test_result_columns.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,141 @@
import pytest

from django_ormql.query import Result


@pytest.mark.django_db
def test_result_is_iterable_wrapper(engine_t1):
res = engine_t1.query("SELECT title FROM categories")
assert isinstance(res, Result)
rows = list(res)
assert rows == [{"title": "Books"}, {"title": "DVDs"}]


@pytest.mark.django_db
def test_columns_simple_fields(engine_t1):
res = engine_t1.query("SELECT title, price, publication_date FROM products")
assert res.columns == [
{"name": "title", "type": "TEXT", "nullable": False, "hint": None},
{"name": "price", "type": "DECIMAL", "nullable": False, "hint": None},
{"name": "publication_date", "type": "DATE", "nullable": False, "hint": None},
]


@pytest.mark.django_db
def test_columns_datetime_field(engine_t1):
res = engine_t1.query("SELECT created FROM orders")
assert res.columns == [
{"name": "created", "type": "DATETIME", "nullable": False, "hint": None},
]


@pytest.mark.django_db
def test_columns_nullable(engine_t1):
res = engine_t1.query("SELECT name FROM orderpositions")
assert res.columns == [
{"name": "name", "type": "TEXT", "nullable": True, "hint": None},
]


@pytest.mark.django_db
@pytest.mark.xfail(reason="Known bug, we cannot see the nullability of foreign keys")
def test_columns_nullable_foreignkey(engine_t1):
# Customer is nullable on Order, but name not on customer
res = engine_t1.query("SELECT customer FROM orders")
assert res.columns == [
{"name": "customer", "type": "TEXT", "nullable": True, "hint": None},
]
res = engine_t1.query("SELECT customer.name FROM orders")
assert res.columns == [
{"name": "customer.name", "type": "TEXT", "nullable": False, "hint": None},
]


@pytest.mark.django_db
def test_columns_boolean_field(engine_t1):
res = engine_t1.query("SELECT enabled FROM customers")
assert res.columns == [
{"name": "enabled", "type": "BOOLEAN", "nullable": False, "hint": None},
]


@pytest.mark.django_db
def test_columns_json_field(engine_t1):
res = engine_t1.query("SELECT address FROM customers")
assert res.columns == [
{"name": "address", "type": "JSONB", "nullable": False, "hint": None},
]


@pytest.mark.django_db
def test_columns_json_field_extraction(engine_t1):
res = engine_t1.query("SELECT address->city->state->code as code FROM customers")
assert res.columns == [
{"name": "code", "type": "JSONB", "nullable": False, "hint": None},
]


@pytest.mark.django_db
def test_columns_count_aggregate(engine_t1):
res = engine_t1.query("SELECT COUNT(*) AS n FROM categories")
assert res.columns == [
{"name": "n", "type": "INT", "nullable": False, "hint": None}
]


@pytest.mark.django_db
def test_columns_sum_aggregate(engine_t1):
res = engine_t1.query("SELECT SUM(price) AS total FROM products")
assert len(res.columns) == 1
assert res.columns[0]["name"] == "total"
assert res.columns[0]["type"] == "DECIMAL"


@pytest.mark.django_db
def test_columns_group_by_with_aggregate(engine_t1):
res = engine_t1.query(
"SELECT category.title, COUNT(*) AS n FROM products GROUP BY category.title"
)
assert res.columns == [
{"name": "category.title", "type": "TEXT", "nullable": False, "hint": None},
{"name": "n", "type": "INT", "nullable": False, "hint": None},
]


@pytest.mark.django_db
def test_columns_preserved_with_alias(engine_t1):
res = engine_t1.query("SELECT title AS name, price AS amount FROM products")
assert res.columns == [
{"name": "name", "type": "TEXT", "nullable": False, "hint": None},
{"name": "amount", "type": "DECIMAL", "nullable": False, "hint": None},
]


@pytest.mark.django_db
def test_columns_datetime_trunc(engine_t1):
for k in ("year", "quarter", "month", "week", "day", "hour", "minute", "second"):
res = engine_t1.query(f"SELECT DATETRUNC('{k}', created) as col FROM orders")
assert res.columns == [
{
"name": "col",
"type": "DATETIME",
"nullable": False,
"hint": {"truncated": k},
},
]


@pytest.mark.django_db
def test_columns_date_trunc(engine_t1):
for k in ("year", "quarter", "month", "week"):
res = engine_t1.query(
f"SELECT DATETRUNC('{k}', publication_date) as col FROM products"
)
assert res.columns == [
{
"name": "col",
"type": "DATE",
"nullable": False,
"hint": {"truncated": k},
},
]
1 change: 1 addition & 0 deletions tests/testapp/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,7 @@ class OrderPosition(models.Model):
quantity = models.IntegerField()
single_price = models.DecimalField(max_digits=10, decimal_places=2)
tax_rate = models.DecimalField(max_digits=10, decimal_places=2)
name = models.CharField(max_length=250, null=True)

class Meta:
ordering = ("id",)
1 change: 1 addition & 0 deletions tests/testapp/ormql.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,7 @@ class Meta:
"quantity",
"single_price",
"tax_rate",
"name",
]


Expand Down
Loading