diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py index 5d818105..158303d3 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py @@ -469,8 +469,33 @@ def join(self, table: str, column1: str, equality: str, column2: str, clause: st self._joins += (join_clause,) return self - def where_column(self, column1: str, column2: str) -> "QueryBuilder": - self._wheres += (QueryExpression(column1, "=", column2, "value_equals"),) + _WHERE_COLUMN_OPERATORS = ("=", "!=", "<>", ">", ">=", "<", "<=") + + def _normalize_where_column(self, operator: str, column2: str | None): + """Resolve the where_column arity. + + Two-arg ``(col1, col2)`` defaults the operator to ``=``; three-arg + ``(col1, operator, col2)`` validates the operator. + """ + if column2 is None: + return "=", operator + if operator not in self._WHERE_COLUMN_OPERATORS: + raise ValueError( + f"Invalid where_column operator {operator!r}. " + f"Expected one of: {', '.join(self._WHERE_COLUMN_OPERATORS)}" + ) + return operator, column2 + + def where_column(self, column1: str, operator: str, column2: str | None = None) -> "QueryBuilder": + """Compare two columns (identifiers, never bound values), joined with AND.""" + operator, column2 = self._normalize_where_column(operator, column2) + self._wheres += (QueryExpression(column1, operator, column2, "value_equals"),) + return self + + def or_where_column(self, column1: str, operator: str, column2: str | None = None) -> "QueryBuilder": + """Compare two columns (identifiers, never bound values), joined with OR.""" + operator, column2 = self._normalize_where_column(operator, column2) + self._wheres += (QueryExpression(column1, operator, column2, "value_equals", keyword="or"),) return self def when(self, condition, callback) -> "QueryBuilder": diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/BaseGrammar.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/BaseGrammar.py index 44822a51..8e377d54 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/BaseGrammar.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/BaseGrammar.py @@ -634,7 +634,9 @@ def process_wheres(self, query=None, qmark=False, strip_first_where=False): keyword=keyword, ) elif value_type == "value_equals": - sql_string = self.value_equal_string().format(value1=where.column, value2=where.value, keyword=keyword) + sql_string = self.value_equal_string().format( + value1=where.column, value2=where.value, keyword=keyword, equality=where.equality + ) elif value_type == "NULL": sql_string = self.where_null_string() elif value_type == "DATE": diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/MSSQLGrammar.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/MSSQLGrammar.py index 3eec6228..7a94c0bc 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/MSSQLGrammar.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/MSSQLGrammar.py @@ -125,7 +125,7 @@ def where_in_string(self): return "WHERE IN ({values})" def value_equal_string(self): - return "{keyword} {value1} = {value2}" + return "{keyword} {value1} {equality} {value2}" def where_null_string(self): return " {keyword} {column} IS NULL" diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/MySQLGrammar.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/MySQLGrammar.py index 4ee2c167..c168cab9 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/MySQLGrammar.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/MySQLGrammar.py @@ -191,7 +191,7 @@ def where_in_string(self): return "WHERE IN ({values})" def value_equal_string(self): - return "{keyword} {value1} = {value2}" + return "{keyword} {value1} {equality} {value2}" def where_string(self): return " {keyword} {column} {equality} {value}" diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/PostgresGrammar.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/PostgresGrammar.py index 3281b829..ad0419e6 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/PostgresGrammar.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/PostgresGrammar.py @@ -185,7 +185,7 @@ def where_date_string(self): return "{keyword} DATE({column}) {equality} {value}" def value_equal_string(self): - return "{keyword} {value1} = {value2}" + return "{keyword} {value1} {equality} {value2}" def where_string(self): return " {keyword} {column} {equality} {value}" diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/SQLiteGrammar.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/SQLiteGrammar.py index 191c1591..db1df09d 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/SQLiteGrammar.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/query/grammars/SQLiteGrammar.py @@ -184,7 +184,7 @@ def where_date_string(self): return "{keyword} DATE({column}) {equality} {value}" def value_equal_string(self): - return "{keyword} {value1} = {value2}" + return "{keyword} {value1} {equality} {value2}" def where_not_null_string(self): return " {keyword} {column} IS NOT NULL" diff --git a/fastapi_startkit/tests/masoniteorm/query/grammars/test_where_column_grammar.py b/fastapi_startkit/tests/masoniteorm/query/grammars/test_where_column_grammar.py new file mode 100644 index 00000000..8f358de9 --- /dev/null +++ b/fastapi_startkit/tests/masoniteorm/query/grammars/test_where_column_grammar.py @@ -0,0 +1,95 @@ +import unittest + +from fastapi_startkit.masoniteorm.models.builder import QueryBuilder +from fastapi_startkit.masoniteorm.query.grammars.SQLiteGrammar import SQLiteGrammar +from fastapi_startkit.masoniteorm.query.grammars.MySQLGrammar import MySQLGrammar +from fastapi_startkit.masoniteorm.query.grammars.PostgresGrammar import PostgresGrammar +from fastapi_startkit.masoniteorm.query.grammars.MSSQLGrammar import MSSQLGrammar + +GRAMMARS = { + "sqlite": SQLiteGrammar, + "mysql": MySQLGrammar, + "postgres": PostgresGrammar, + "mssql": MSSQLGrammar, +} + +# Table name is quoted per dialect; the compared columns stay bare identifiers. +TABLE = {"sqlite": '"users"', "mysql": "`users`", "postgres": '"users"', "mssql": "[users]"} +ACTIVE = { + "sqlite": '"users"."active"', + "mysql": "`users`.`active`", + "postgres": '"users"."active"', + "mssql": "[users].[active]", +} + + +def qb(grammar): + q = QueryBuilder(connection=None, grammar=grammar, processor=None) + q._table = "users" + return q + + +class TestWhereColumnGrammar(unittest.TestCase): + """Grammar-level SQL parity for where_column / or_where_column across all dialects.""" + + def test_where_column_two_arg_equality(self): + for name, grammar in GRAMMARS.items(): + with self.subTest(grammar=name): + sql = qb(grammar).where_column("first_name", "last_name").to_sql() + self.assertEqual(sql, f"SELECT * FROM {TABLE[name]} WHERE first_name = last_name") + + def test_where_column_three_arg_operator(self): + for name, grammar in GRAMMARS.items(): + with self.subTest(grammar=name): + sql = qb(grammar).where_column("updated_at", ">", "created_at").to_sql() + self.assertEqual(sql, f"SELECT * FROM {TABLE[name]} WHERE updated_at > created_at") + + def test_where_column_all_operators(self): + for name, grammar in GRAMMARS.items(): + for op in ("=", "!=", "<>", ">", ">=", "<", "<="): + with self.subTest(grammar=name, operator=op): + sql = qb(grammar).where_column("a", op, "b").to_sql() + self.assertEqual(sql, f"SELECT * FROM {TABLE[name]} WHERE a {op} b") + + def test_where_column_rejects_invalid_operator(self): + for name, grammar in GRAMMARS.items(): + with self.subTest(grammar=name): + with self.assertRaises(ValueError): + qb(grammar).where_column("a", "BAD", "b") + + def test_or_where_column_two_arg_equality(self): + for name, grammar in GRAMMARS.items(): + with self.subTest(grammar=name): + sql = qb(grammar).where("active", 1).or_where_column("first_name", "last_name").to_sql() + self.assertEqual( + sql, + f"SELECT * FROM {TABLE[name]} WHERE {ACTIVE[name]} = '1' OR first_name = last_name", + ) + + def test_or_where_column_three_arg_operator(self): + for name, grammar in GRAMMARS.items(): + with self.subTest(grammar=name): + sql = qb(grammar).where("active", 1).or_where_column("updated_at", ">", "created_at").to_sql() + self.assertEqual( + sql, + f"SELECT * FROM {TABLE[name]} WHERE {ACTIVE[name]} = '1' OR updated_at > created_at", + ) + + def test_or_where_column_uses_or_and_keeps_columns_as_identifiers(self): + for name, grammar in GRAMMARS.items(): + with self.subTest(grammar=name): + q = qb(grammar).where("active", 1).or_where_column("updated_at", ">", "created_at") + sql = q.to_qmark() + # OR join, correlated identifiers preserved, and only the literal value is bound. + self.assertIn("OR updated_at > created_at", sql) + self.assertEqual(list(q.get_bindings()), [1]) + + def test_or_where_column_rejects_invalid_operator(self): + for name, grammar in GRAMMARS.items(): + with self.subTest(grammar=name): + with self.assertRaises(ValueError): + qb(grammar).or_where_column("a", "BAD", "b") + + +if __name__ == "__main__": + unittest.main() diff --git a/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_where_column.py b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_where_column.py new file mode 100644 index 00000000..d7813e01 --- /dev/null +++ b/fastapi_startkit/tests/masoniteorm/sqlite/builder/test_sqlite_where_column.py @@ -0,0 +1,67 @@ +from fastapi_startkit.masoniteorm import Model + +from ..fixtures.db import DB +from ..test_case import TestCase + + +class Pair(Model): + __table__ = "pairs" + __timestamps__ = None + id: int + left_val: int + right_val: int + name: str + nickname: str + + +class TestSqliteWhereColumn(TestCase): + async def asyncSetUp(self): + await super().asyncSetUp() + conn = DB.connection("default") + await conn.execute( + "CREATE TABLE pairs (id INTEGER PRIMARY KEY, left_val INTEGER, right_val INTEGER, name TEXT, nickname TEXT)" + ) + await conn.execute( + "INSERT INTO pairs (id, left_val, right_val, name, nickname) VALUES " + "(1, 5, 5, 'Sam', 'Sam'), " # left == right, name == nickname + "(2, 9, 3, 'Bob', 'Bobby'), " # left > right, name != nickname + "(3, 2, 8, 'Al', 'Al')" # left < right, name == nickname + ) + + async def test_where_column_two_arg_equality(self): + rows = await Pair.query().where_column("left_val", "right_val").order_by("id").get() + self.assertEqual([p.id for p in rows], [1]) + + async def test_where_column_three_arg_greater_than(self): + rows = await Pair.query().where_column("left_val", ">", "right_val").order_by("id").get() + self.assertEqual([p.id for p in rows], [2]) + + async def test_where_column_three_arg_less_than(self): + rows = await Pair.query().where_column("left_val", "<", "right_val").order_by("id").get() + self.assertEqual([p.id for p in rows], [3]) + + async def test_where_column_not_equal(self): + rows = await Pair.query().where_column("name", "!=", "nickname").order_by("id").get() + self.assertEqual([p.id for p in rows], [2]) + + async def test_or_where_column_or_joined(self): + # id == 3 OR left_val > right_val -> rows 2 (9>3) and 3 (id match) + query = Pair.query().where("id", 3).or_where_column("left_val", ">", "right_val") + self.assertIn("OR left_val > right_val", query.to_sql()) + rows = await query.order_by("id").get() + self.assertEqual([p.id for p in rows], [2, 3]) + + async def test_or_where_column_two_arg_equality(self): + # left_val == right_val OR name == nickname -> rows 1 (both) and 3 (name==nick) + rows = ( + await Pair.query() + .where_column("left_val", "right_val") + .or_where_column("name", "nickname") + .order_by("id") + .get() + ) + self.assertEqual([p.id for p in rows], [1, 3]) + + async def test_where_column_rejects_invalid_operator(self): + with self.assertRaises(ValueError): + Pair.query().where_column("left_val", "BAD", "right_val")