diff options
| author | Mike Bayer <mike_mp@zzzcomputing.com> | 2020-08-05 21:47:43 -0400 |
|---|---|---|
| committer | Mike Bayer <mike_mp@zzzcomputing.com> | 2020-08-05 22:13:11 -0400 |
| commit | c7b489b25802f7a25ef78d0731411295c611cc1c (patch) | |
| tree | f5e3b66ab8eb8bb7398c0195fa2b2f1de8ab91c4 /test | |
| parent | 71a3ccbdef0d88e9231b7de9c51e4ed60b3b7181 (diff) | |
| download | sqlalchemy-c7b489b25802f7a25ef78d0731411295c611cc1c.tar.gz | |
Implement relationship AND criteria; global loader criteria
Added the ability to add arbitrary criteria to the ON clause generated
by a relationship attribute in a query, which applies to methods such
as :meth:`_query.Query.join` as well as loader options like
:func:`_orm.joinedload`. Additionally, a "global" version of the option
allows limiting criteria to be applied to particular entities in
a query globally.
Documentation is minimal at this point, new examples will
be coming in a subsequent commit.
Some adjustments to execution options in how they are represented
in the ORMExecuteState as well as well as a few ORM tests that
forgot to get merged in a preceding commit.
Fixes: #4472
Change-Id: I2b8fc57092dedf35ebd16f6343ad0f0d7d332beb
Diffstat (limited to 'test')
| -rw-r--r-- | test/ext/test_baked.py | 4 | ||||
| -rw-r--r-- | test/orm/inheritance/test_polymorphic_rel.py | 22 | ||||
| -rw-r--r-- | test/orm/test_bundle.py | 42 | ||||
| -rw-r--r-- | test/orm/test_cache_key.py | 68 | ||||
| -rw-r--r-- | test/orm/test_events.py | 166 | ||||
| -rw-r--r-- | test/orm/test_options.py | 1 | ||||
| -rw-r--r-- | test/orm/test_relationship_criteria.py | 867 | ||||
| -rw-r--r-- | test/sql/test_compare.py | 20 |
8 files changed, 1188 insertions, 2 deletions
diff --git a/test/ext/test_baked.py b/test/ext/test_baked.py index 6279dcf55..c8e83bbd7 100644 --- a/test/ext/test_baked.py +++ b/test/ext/test_baked.py @@ -1017,8 +1017,8 @@ class CustomIntegrationTest(testing.AssertsCompiledSQL, BakedTest): if ckey: break else: - if "_cache_key" in orm_context.merged_execution_options: - ckey = orm_context.merged_execution_options["_cache_key"] + if "_cache_key" in orm_context.execution_options: + ckey = orm_context.execution_options["_cache_key"] if ckey is not None: return get_value( diff --git a/test/orm/inheritance/test_polymorphic_rel.py b/test/orm/inheritance/test_polymorphic_rel.py index e33e95cc0..86e0bd360 100644 --- a/test/orm/inheritance/test_polymorphic_rel.py +++ b/test/orm/inheritance/test_polymorphic_rel.py @@ -1302,6 +1302,28 @@ class _PolymorphicTestBase(object): [e1, e3], ) + def test_join_and_thru_polymorphic_nonaliased_one(self): + sess = create_session() + eq_( + sess.query(Company) + .join(Company.employees) + .join(Person.paperwork.and_(Paperwork.description.like("%#2%"))) + .all(), + [c1], + ) + + def test_join_and_thru_polymorphic_aliased_one(self): + sess = create_session() + ea = aliased(Person) + pa = aliased(Paperwork) + eq_( + sess.query(Company) + .join(ea, Company.employees) + .join(pa, ea.paperwork.and_(pa.description.like("%#2%"))) + .all(), + [c1], + ) + def test_join_through_polymorphic_nonaliased_one(self): sess = create_session() eq_( diff --git a/test/orm/test_bundle.py b/test/orm/test_bundle.py index f4af84094..9d1d0b61b 100644 --- a/test/orm/test_bundle.py +++ b/test/orm/test_bundle.py @@ -3,6 +3,7 @@ from sqlalchemy import func from sqlalchemy import Integer from sqlalchemy import select from sqlalchemy import String +from sqlalchemy import testing from sqlalchemy.orm import aliased from sqlalchemy.orm import Bundle from sqlalchemy.orm import mapper @@ -186,6 +187,35 @@ class BundleTest(fixtures.MappedTest, AssertsCompiledSQL): ], ) + def test_multi_bundle_future(self): + Data = self.classes.Data + Other = self.classes.Other + + d1 = aliased(Data) + + b1 = Bundle("b1", d1.d1, d1.d2) + b2 = Bundle("b2", Data.d1, Other.o1) + + sess = Session(testing.db, future=True) + + stmt = ( + select(b1, b2) + .join(Data.others) + .join(d1, d1.id == Data.id) + .filter(b1.c.d1 == "d3d1") + ) + + eq_( + sess.execute(stmt).all(), + [ + (("d3d1", "d3d2"), ("d3d1", "d3o0")), + (("d3d1", "d3d2"), ("d3d1", "d3o1")), + (("d3d1", "d3d2"), ("d3d1", "d3o2")), + (("d3d1", "d3d2"), ("d3d1", "d3o3")), + (("d3d1", "d3d2"), ("d3d1", "d3o4")), + ], + ) + def test_single_entity(self): Data = self.classes.Data sess = Session() @@ -197,6 +227,18 @@ class BundleTest(fixtures.MappedTest, AssertsCompiledSQL): [("d3d1", "d3d2"), ("d4d1", "d4d2"), ("d5d1", "d5d2")], ) + def test_single_entity_future(self): + Data = self.classes.Data + sess = Session(testing.db, future=True) + + b1 = Bundle("b1", Data.d1, Data.d2, single_entity=True) + + stmt = select(b1).filter(b1.c.d1.between("d3d1", "d5d1")) + eq_( + sess.execute(stmt).scalars().all(), + [("d3d1", "d3d2"), ("d4d1", "d4d2"), ("d5d1", "d5d2")], + ) + def test_single_entity_flag_but_multi_entities(self): Data = self.classes.Data sess = Session() diff --git a/test/orm/test_cache_key.py b/test/orm/test_cache_key.py index 02b1b9fbf..45a60a5cb 100644 --- a/test/orm/test_cache_key.py +++ b/test/orm/test_cache_key.py @@ -15,6 +15,7 @@ from sqlalchemy.orm import relationship from sqlalchemy.orm import selectinload from sqlalchemy.orm import Session from sqlalchemy.orm import subqueryload +from sqlalchemy.orm import with_loader_criteria from sqlalchemy.orm import with_polymorphic from sqlalchemy.sql.base import CacheableOptions from sqlalchemy.sql.visitors import InternalTraversal @@ -65,6 +66,62 @@ class CacheKeyTest(CacheKeyFixture, _fixtures.FixtureTest): compare_values=True, ) + def test_loader_criteria(self): + User, Address = self.classes("User", "Address") + + from sqlalchemy import Column, Integer, String + + class Foo(object): + id = Column(Integer) + name = Column(String) + + self._run_cache_key_fixture( + lambda: ( + with_loader_criteria(User, User.name != "somename"), + with_loader_criteria(User, User.id != 5), + with_loader_criteria(User, lambda cls: cls.id == 10), + with_loader_criteria(Address, Address.id != 5), + with_loader_criteria(Foo, lambda cls: cls.id == 10), + ), + compare_values=True, + ) + + def test_loader_criteria_bound_param_thing(self): + from sqlalchemy import Column, Integer + + class Foo(object): + id = Column(Integer) + + def go(param): + return with_loader_criteria(Foo, lambda cls: cls.id == param) + + g1 = go(10) + g2 = go(20) + + ck1 = g1._generate_cache_key() + ck2 = g2._generate_cache_key() + + eq_(ck1.key, ck2.key) + eq_(ck1.bindparams[0].key, ck2.bindparams[0].key) + eq_(ck1.bindparams[0].value, 10) + eq_(ck2.bindparams[0].value, 20) + + def test_instrumented_attributes(self): + User, Address, Keyword, Order, Item = self.classes( + "User", "Address", "Keyword", "Order", "Item" + ) + + self._run_cache_key_fixture( + lambda: ( + User.addresses, + User.addresses.of_type(aliased(Address)), + User.orders, + User.orders.and_(Order.id != 5), + User.orders.and_(Order.description != "somename"), + ), + compare_values=True, + ) + def test_unbound_options(self): User, Address, Keyword, Order, Item = self.classes( "User", "Address", "Keyword", "Order", "Item" @@ -75,6 +132,10 @@ class CacheKeyTest(CacheKeyFixture, _fixtures.FixtureTest): joinedload(User.addresses), joinedload(User.addresses.of_type(aliased(Address))), joinedload("addresses"), + joinedload(User.orders), + joinedload(User.orders.and_(Order.id != 5)), + joinedload(User.orders.and_(Order.id == 5)), + joinedload(User.orders.and_(Order.description != "somename")), joinedload(User.orders).selectinload("items"), joinedload(User.orders).selectinload(Order.items), defer(User.id), @@ -110,6 +171,10 @@ class CacheKeyTest(CacheKeyFixture, _fixtures.FixtureTest): User.addresses.of_type(aliased(Address)) ), Load(User).joinedload(User.orders), + Load(User).joinedload(User.orders.and_(Order.id != 5)), + Load(User).joinedload( + User.orders.and_(Order.description != "somename") + ), Load(User).defer(User.id), Load(User).subqueryload("addresses"), Load(Address).defer("id"), @@ -169,6 +234,9 @@ class CacheKeyTest(CacheKeyFixture, _fixtures.FixtureTest): select(User).join(Address, User.addresses), select(User).join(a1, User.addresses), select(User).join(User.addresses.of_type(a1)), + select(User).join( + User.addresses.and_(Address.email_address == "foo") + ), select(User) .join(Address, User.addresses) .join_from(User, Order), diff --git a/test/orm/test_events.py b/test/orm/test_events.py index b68e0d2e6..df48cfe63 100644 --- a/test/orm/test_events.py +++ b/test/orm/test_events.py @@ -2,6 +2,8 @@ import sqlalchemy as sa from sqlalchemy import event from sqlalchemy import ForeignKey from sqlalchemy import Integer +from sqlalchemy import literal_column +from sqlalchemy import select from sqlalchemy import String from sqlalchemy import testing from sqlalchemy.ext.declarative import declarative_base @@ -47,6 +49,170 @@ class _RemoveListeners(object): super(_RemoveListeners, self).teardown() +class ORMExecuteTest(_RemoveListeners, _fixtures.FixtureTest): + run_setup_mappers = "once" + run_inserts = "once" + run_deletes = None + + @classmethod + def setup_mappers(cls): + cls._setup_stock_mapping() + + def _caching_session_fixture(self): + + cache = {} + + maker = sessionmaker(testing.db, future=True) + + def get_value(cache_key, cache, createfunc): + if cache_key in cache: + return cache[cache_key]() + else: + cache[cache_key] = retval = createfunc().freeze() + return retval() + + @event.listens_for(maker, "do_orm_execute", retval=True) + def do_orm_execute(orm_context): + ckey = None + for opt in orm_context.user_defined_options: + ckey = opt.get_cache_key(orm_context) + if ckey: + break + else: + if "cache_key" in orm_context.execution_options: + ckey = orm_context.execution_options["cache_key"] + + if ckey is not None: + return get_value(ckey, cache, orm_context.invoke_statement,) + + return maker() + + def test_cache_option(self): + User, Address = self.classes("User", "Address") + + with self.sql_execution_asserter(testing.db) as asserter: + + with self._caching_session_fixture() as session: + stmt = ( + select(User) + .where(User.id == 7) + .execution_options(cache_key="user7") + ) + + result = session.execute(stmt) + + eq_( + result.scalars().all(), + [User(id=7, addresses=[Address(id=1)])], + ) + + result = session.execute(stmt) + + eq_( + result.scalars().all(), + [User(id=7, addresses=[Address(id=1)])], + ) + + asserter.assert_( + CompiledSQL( + "SELECT users.id, users.name FROM users " + "WHERE users.id = :id_1", + [{"id_1": 7}], + ), + CompiledSQL( + "SELECT addresses.id AS addresses_id, addresses.user_id AS " + "addresses_user_id, " + "addresses.email_address AS addresses_email_address " + "FROM addresses WHERE :param_1 = addresses.user_id " + "ORDER BY addresses.id", + [{"param_1": 7}], + ), + ) + + def test_chained_events_one(self): + + sess = Session(testing.db, future=True) + + @event.listens_for(sess, "do_orm_execute") + def one(ctx): + ctx.update_execution_options(one=True) + + @event.listens_for(sess, "do_orm_execute") + def two(ctx): + ctx.update_execution_options(two=True) + + @event.listens_for(sess, "do_orm_execute") + def three(ctx): + ctx.update_execution_options(three=True) + + @event.listens_for(sess, "do_orm_execute") + def four(ctx): + ctx.update_execution_options(four=True) + + result = sess.execute(select(literal_column("1"))) + + eq_( + result.context.execution_options, + { + "four": True, + "future_result": True, + "one": True, + "three": True, + "two": True, + }, + ) + + def test_chained_events_two(self): + + sess = Session(testing.db, future=True) + + def added(ctx): + ctx.update_execution_options(added_evt=True) + + @event.listens_for(sess, "do_orm_execute") + def one(ctx): + ctx.update_execution_options(one=True) + + @event.listens_for(sess, "do_orm_execute", retval=True) + def two(ctx): + ctx.update_execution_options(two=True) + return ctx.invoke_statement( + statement=ctx.statement.execution_options(statement_two=True) + ) + + @event.listens_for(sess, "do_orm_execute") + def three(ctx): + ctx.update_execution_options(three=True) + + @event.listens_for(sess, "do_orm_execute") + def four(ctx): + ctx.update_execution_options(four=True) + return ctx.invoke_statement( + statement=ctx.statement.execution_options(statement_four=True) + ) + + @event.listens_for(sess, "do_orm_execute") + def five(ctx): + ctx.update_execution_options(five=True) + + result = sess.execute(select(literal_column("1")), _add_event=added) + + eq_( + result.context.execution_options, + { + "statement_two": True, + "statement_four": True, + "future_result": True, + "one": True, + "two": True, + "three": True, + "four": True, + "five": True, + "added_evt": True, + }, + ) + + class MapperEventsTest(_RemoveListeners, _fixtures.FixtureTest): run_inserts = None diff --git a/test/orm/test_options.py b/test/orm/test_options.py index 208db9d85..b5a6e3b29 100644 --- a/test/orm/test_options.py +++ b/test/orm/test_options.py @@ -1391,6 +1391,7 @@ class PickleTest(PathTest, QueryTest): "propagate_to_loaders": True, "_of_type": None, "_to_bind": to_bind, + "_extra_criteria": (), }, ) diff --git a/test/orm/test_relationship_criteria.py b/test/orm/test_relationship_criteria.py new file mode 100644 index 000000000..c4bcf0404 --- /dev/null +++ b/test/orm/test_relationship_criteria.py @@ -0,0 +1,867 @@ +import datetime +import random + +from sqlalchemy import Column +from sqlalchemy import DateTime +from sqlalchemy import event +from sqlalchemy import ForeignKey +from sqlalchemy import Integer +from sqlalchemy import orm +from sqlalchemy import select +from sqlalchemy import sql +from sqlalchemy import String +from sqlalchemy import testing +from sqlalchemy.orm import aliased +from sqlalchemy.orm import joinedload +from sqlalchemy.orm import mapper +from sqlalchemy.orm import relationship +from sqlalchemy.orm import selectinload +from sqlalchemy.orm import Session +from sqlalchemy.orm import with_loader_criteria +from sqlalchemy.testing import eq_ +from sqlalchemy.testing.assertsql import CompiledSQL +from test.orm import _fixtures + + +class _Fixtures(_fixtures.FixtureTest): + @testing.fixture + def user_address_fixture(self): + users, Address, addresses, User = ( + self.tables.users, + self.classes.Address, + self.tables.addresses, + self.classes.User, + ) + + mapper( + User, + users, + properties={ + "addresses": relationship( + mapper(Address, addresses), order_by=Address.id + ) + }, + ) + return User, Address + + @testing.fixture + def order_item_fixture(self): + Order, Item = self.classes("Order", "Item") + orders, items, order_items = self.tables( + "orders", "items", "order_items" + ) + + mapper( + Order, + orders, + properties={ + # m2m + "items": relationship( + Item, secondary=order_items, order_by=items.c.id + ), + }, + ) + mapper(Item, items) + + return Order, Item + + @testing.fixture + def mixin_fixture(self): + users = self.tables.users + + class HasFoob(object): + name = Column(String) + + class UserWFoob(HasFoob, self.Comparable): + pass + + mapper( + UserWFoob, users, + ) + return HasFoob, UserWFoob + + +class LoaderCriteriaTest(_Fixtures, testing.AssertsCompiledSQL): + """ + combinations: + + + with_loader_criteria + # for these we have mapper_criteria + + select(mapper) # select_mapper + select(mapper.col, mapper.col) # select_mapper_col + select(func.count()).select_from(mapper) # select_from_mapper + select(a).join(mapper, a.target) # select_join_mapper + select(a).options(joinedload(a.target)) # select_joinedload_mapper + + + # for these we have aliased_criteria, inclaliased_criteria + + select(aliased) # select_aliased + select(aliased.col, aliased.col) # select_aliased_col + select(func.count()).select_from(aliased) # select_from_aliased + select(a).join(aliased, a.target) # select_join_aliased + select(a).options(joinedload(a.target.of_type(aliased)) + # select_joinedload_aliased + + """ + + __dialect__ = "default" + + def test_select_mapper_mapper_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + stmt = select(User).options( + with_loader_criteria(User, User.name != "name") + ) + + self.assert_compile( + stmt, + "SELECT users.id, users.name " + "FROM users WHERE users.name != :name_1", + ) + + def test_select_from_mapper_mapper_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + stmt = ( + select(sql.func.count()) + .select_from(User) + .options(with_loader_criteria(User, User.name != "name")) + ) + + self.assert_compile( + stmt, + "SELECT count(*) AS count_1 FROM users " + "WHERE users.name != :name_1", + ) + + def test_select_mapper_columns_mapper_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + stmt = select(User.id, User.name).options( + with_loader_criteria(User, User.name != "name") + ) + + self.assert_compile( + stmt, + "SELECT users.id, users.name " + "FROM users WHERE users.name != :name_1", + ) + + def test_select_join_mapper_mapper_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + stmt = ( + select(User) + .join(User.addresses) + .options( + with_loader_criteria(Address, Address.email_address != "name") + ) + ) + + self.assert_compile( + stmt, + "SELECT users.id, users.name FROM users " + "JOIN addresses ON users.id = addresses.user_id " + "AND addresses.email_address != :email_address_1", + ) + + def test_select_joinm2m_mapper_mapper_criteria(self, order_item_fixture): + Order, Item = order_item_fixture + + stmt = ( + select(Order) + .join(Order.items) + .options( + with_loader_criteria(Item, Item.description != "description") + ) + ) + + self.assert_compile( + stmt, + "SELECT orders.id, orders.user_id, orders.address_id, " + "orders.description, orders.isopen FROM orders " + "JOIN order_items AS order_items_1 " + "ON orders.id = order_items_1.order_id " + "JOIN items ON items.id = order_items_1.item_id " + "AND items.description != :description_1", + ) + + def test_select_joinedload_mapper_mapper_criteria( + self, user_address_fixture + ): + User, Address = user_address_fixture + + stmt = select(User).options( + joinedload(User.addresses), + with_loader_criteria(Address, Address.email_address != "name"), + ) + + self.assert_compile( + stmt, + "SELECT users.id, users.name, addresses_1.id AS id_1, " + "addresses_1.user_id, addresses_1.email_address " + "FROM users LEFT OUTER JOIN addresses AS addresses_1 " + "ON users.id = addresses_1.user_id " + "AND addresses_1.email_address != :email_address_1 " + "ORDER BY addresses_1.id", + ) + + def test_select_selectinload_mapper_mapper_criteria( + self, user_address_fixture + ): + User, Address = user_address_fixture + + stmt = select(User).options( + selectinload(User.addresses), + with_loader_criteria(Address, Address.email_address != "name"), + ) + + s = Session(testing.db, future=True) + + with self.sql_execution_asserter() as asserter: + + s.execute(stmt).all() + + asserter.assert_( + CompiledSQL("SELECT users.id, users.name FROM users", [],), + CompiledSQL( + "SELECT addresses.user_id AS addresses_user_id, addresses.id " + "AS addresses_id, addresses.email_address " + "AS addresses_email_address FROM addresses " + "WHERE addresses.user_id IN ([POSTCOMPILE_primary_keys]) " + "AND addresses.email_address != :email_address_1 " + "ORDER BY addresses.id", + [{"primary_keys": [7, 8, 9, 10], "email_address_1": "name"}], + ), + ) + + def test_select_lazyload_mapper_mapper_criteria( + self, user_address_fixture + ): + User, Address = user_address_fixture + + stmt = ( + select(User) + .options( + with_loader_criteria(Address, Address.email_address != "name"), + ) + .order_by(User.id) + ) + + s = Session(testing.db, future=True) + + with self.sql_execution_asserter() as asserter: + for u in s.execute(stmt).scalars(): + u.addresses + + asserter.assert_( + CompiledSQL( + "SELECT users.id, users.name FROM users ORDER BY users.id", [], + ), + CompiledSQL( + "SELECT addresses.id AS addresses_id, " + "addresses.user_id AS addresses_user_id, " + "addresses.email_address AS addresses_email_address " + "FROM addresses WHERE :param_1 = addresses.user_id " + "AND addresses.email_address != :email_address_1 " + "ORDER BY addresses.id", + [{"param_1": 7, "email_address_1": "name"}], + ), + CompiledSQL( + "SELECT addresses.id AS addresses_id, " + "addresses.user_id AS addresses_user_id, " + "addresses.email_address AS addresses_email_address " + "FROM addresses WHERE :param_1 = addresses.user_id " + "AND addresses.email_address != :email_address_1 " + "ORDER BY addresses.id", + [{"param_1": 8, "email_address_1": "name"}], + ), + CompiledSQL( + "SELECT addresses.id AS addresses_id, " + "addresses.user_id AS addresses_user_id, " + "addresses.email_address AS addresses_email_address " + "FROM addresses WHERE :param_1 = addresses.user_id " + "AND addresses.email_address != :email_address_1 " + "ORDER BY addresses.id", + [{"param_1": 9, "email_address_1": "name"}], + ), + CompiledSQL( + "SELECT addresses.id AS addresses_id, " + "addresses.user_id AS addresses_user_id, " + "addresses.email_address AS addresses_email_address " + "FROM addresses WHERE :param_1 = addresses.user_id " + "AND addresses.email_address != :email_address_1 " + "ORDER BY addresses.id", + [{"param_1": 10, "email_address_1": "name"}], + ), + ) + + def test_select_aliased_inclaliased_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + u1 = aliased(User) + stmt = select(u1).options( + with_loader_criteria( + User, User.name != "name", include_aliases=True + ) + ) + + self.assert_compile( + stmt, + "SELECT users_1.id, users_1.name " + "FROM users AS users_1 WHERE users_1.name != :name_1", + ) + + def test_select_from_aliased_inclaliased_criteria( + self, user_address_fixture + ): + User, Address = user_address_fixture + + u1 = aliased(User) + stmt = ( + select(sql.func.count()) + .select_from(u1) + .options( + with_loader_criteria( + User, User.name != "name", include_aliases=True + ) + ) + ) + + self.assert_compile( + stmt, + "SELECT count(*) AS count_1 FROM users AS users_1 " + "WHERE users_1.name != :name_1", + ) + + def test_select_aliased_columns_inclaliased_criteria( + self, user_address_fixture + ): + User, Address = user_address_fixture + + u1 = aliased(User) + stmt = select(u1.id, u1.name).options( + with_loader_criteria( + User, User.name != "name", include_aliases=True + ) + ) + + self.assert_compile( + stmt, + "SELECT users_1.id, users_1.name " + "FROM users AS users_1 WHERE users_1.name != :name_1", + ) + + def test_select_join_aliased_inclaliased_criteria( + self, user_address_fixture + ): + User, Address = user_address_fixture + + a1 = aliased(Address) + stmt = ( + select(User) + .join(User.addresses.of_type(a1)) + .options( + with_loader_criteria( + Address, + Address.email_address != "name", + include_aliases=True, + ) + ) + ) + + self.assert_compile( + stmt, + "SELECT users.id, users.name FROM users " + "JOIN addresses AS addresses_1 ON users.id = addresses_1.user_id " + "AND addresses_1.email_address != :email_address_1", + ) + + def test_select_joinm2m_aliased_inclaliased_criteria( + self, order_item_fixture + ): + Order, Item = order_item_fixture + + i1 = aliased(Item) + + stmt = ( + select(Order) + .join(Order.items.of_type(i1)) + .options( + with_loader_criteria( + Item, + Item.description != "description", + include_aliases=True, + ) + ) + ) + + self.assert_compile( + stmt, + "SELECT orders.id, orders.user_id, orders.address_id, " + "orders.description, orders.isopen FROM orders " + "JOIN order_items AS order_items_1 " + "ON orders.id = order_items_1.order_id " + "JOIN items AS items_1 ON items_1.id = order_items_1.item_id " + "AND items_1.description != :description_1", + ) + + def test_select_aliased_aliased_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + u1 = aliased(User) + stmt = select(u1).options(with_loader_criteria(u1, u1.name != "name")) + + self.assert_compile( + stmt, + "SELECT users_1.id, users_1.name " + "FROM users AS users_1 WHERE users_1.name != :name_1", + ) + + def test_select_aliased_columns_aliased_criteria( + self, user_address_fixture + ): + User, Address = user_address_fixture + + u1 = aliased(User) + stmt = select(u1.id, u1.name).options( + with_loader_criteria(u1, u1.name != "name") + ) + + self.assert_compile( + stmt, + "SELECT users_1.id, users_1.name " + "FROM users AS users_1 WHERE users_1.name != :name_1", + ) + + def test_joinedload_global_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + s = Session(testing.db, future=True) + + stmt = select(User).options( + joinedload(User.addresses), + with_loader_criteria(Address, Address.email_address != "email"), + ) + + with self.sql_execution_asserter() as asserter: + + s.execute(stmt) + + asserter.assert_( + CompiledSQL( + "SELECT users.id, users.name, addresses_1.id AS id_1, " + "addresses_1.user_id, addresses_1.email_address FROM " + "users LEFT OUTER JOIN addresses AS addresses_1 " + "ON users.id = addresses_1.user_id " + "AND addresses_1.email_address != :email_address_1 " + "ORDER BY addresses_1.id", + [{"email_address_1": "email"}], + ), + ) + + def test_query_count_global_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + s = Session(testing.db) + + q = s.query(User).options(with_loader_criteria(User, User.id != 8)) + + with self.sql_execution_asserter() as asserter: + q.count() + + asserter.assert_( + CompiledSQL( + "SELECT count(*) AS count_1 FROM (SELECT " + "users.id AS users_id, users.name AS users_name " + "FROM users WHERE users.id != :id_1) AS anon_1", + [{"id_1": 8}], + ), + ) + + def test_query_count_after_the_fact_global_criteria( + self, user_address_fixture + ): + User, Address = user_address_fixture + + s = Session(testing.db) + + # this essentially tests that the query.from_self() which takes + # place in count() is one that can still be affected by + # the loader criteria, meaning it has to be an ORM query + + q = s.query(User) + + @event.listens_for(s, "do_orm_execute") + def add_criteria(orm_context): + orm_context.statement = orm_context.statement.options( + with_loader_criteria(User, User.id != 8) + ) + + with self.sql_execution_asserter() as asserter: + q.count() + + asserter.assert_( + CompiledSQL( + "SELECT count(*) AS count_1 FROM (SELECT " + "users.id AS users_id, users.name AS users_name " + "FROM users WHERE users.id != :id_1) AS anon_1", + [{"id_1": 8}], + ), + ) + + def test_select_count_subquery_global_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + stmt = select(User).subquery() + + stmt = ( + select(sql.func.count()) + .select_from(stmt) + .options(with_loader_criteria(User, User.id != 8)) + ) + + self.assert_compile( + stmt, + "SELECT count(*) AS count_1 FROM (SELECT users.id AS id, " + "users.name AS name FROM users WHERE users.id != :id_1) AS anon_1", + ) + + def test_query_outerjoin_global_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + s = Session(testing.db) + + q = ( + s.query(User, Address) + .outerjoin(User.addresses) + .options( + with_loader_criteria( + Address, ~Address.email_address.like("ed@%"), + ) + ) + .order_by(User.id) + ) + + self.assert_compile( + q, + "SELECT users.id AS users_id, users.name AS users_name, " + "addresses.id AS addresses_id, " + "addresses.user_id AS addresses_user_id, " + "addresses.email_address AS addresses_email_address " + "FROM users LEFT OUTER JOIN addresses " + "ON users.id = addresses.user_id AND " + "addresses.email_address NOT LIKE :email_address_1 " + "ORDER BY users.id", + ) + eq_( + q.all(), + [ + (User(id=7), Address(id=1)), + (User(id=8), None), # three addresses not here + (User(id=9), Address(id=5)), + (User(id=10), None), + ], + ) + + def test_caching_and_binds_lambda(self, mixin_fixture): + HasFoob, UserWFoob = mixin_fixture + + statement = select(UserWFoob).filter(UserWFoob.id < 10) + + def go(value): + return statement.options( + with_loader_criteria( + HasFoob, + lambda cls: cls.name == value, + include_aliases=True, + ) + ) + + s = Session(testing.db, future=True) + + for i in range(10): + name = random.choice(["ed", "fred", "jack"]) + stmt = go(name) + + eq_(s.execute(stmt).scalars().all(), [UserWFoob(name=name)]) + + +class TemporalFixtureTest(testing.fixtures.DeclarativeMappedTest): + @classmethod + def setup_classes(cls): + class HasTemporal(object): + """Mixin that identifies a class as having a timestamp column""" + + timestamp = Column( + DateTime, default=datetime.datetime.utcnow, nullable=False + ) + + cls.HasTemporal = HasTemporal + + def temporal_range(range_lower, range_upper): + return with_loader_criteria( + HasTemporal, + lambda cls: cls.timestamp.between(range_lower, range_upper), + include_aliases=True, + ) + + cls.temporal_range = staticmethod(temporal_range) + + class Parent(HasTemporal, cls.DeclarativeBasic): + __tablename__ = "parent" + id = Column(Integer, primary_key=True) + children = relationship("Child", order_by="Child.id") + + class Child(HasTemporal, cls.DeclarativeBasic): + __tablename__ = "child" + id = Column(Integer, primary_key=True) + parent_id = Column( + Integer, ForeignKey("parent.id"), nullable=False + ) + + @classmethod + def insert_data(cls, connection): + Parent, Child = cls.classes("Parent", "Child") + + sess = Session(connection) + c1, c2, c3, c4, c5 = [ + Child(timestamp=datetime.datetime(2009, 10, 15, 12, 00, 00)), + Child(timestamp=datetime.datetime(2009, 10, 17, 12, 00, 00)), + Child(timestamp=datetime.datetime(2009, 10, 20, 12, 00, 00)), + Child(timestamp=datetime.datetime(2009, 10, 12, 12, 00, 00)), + Child(timestamp=datetime.datetime(2009, 10, 17, 12, 00, 00)), + ] + + p1 = Parent( + timestamp=datetime.datetime(2009, 10, 15, 12, 00, 00), + children=[c1, c2, c3], + ) + p2 = Parent( + timestamp=datetime.datetime(2009, 10, 17, 12, 00, 00), + children=[c4, c5], + ) + + sess.add_all([p1, p2]) + sess.commit() + + @testing.combinations((True,), (False,), argnames="use_caching") + @testing.combinations( + (None,), + (orm.lazyload,), + (orm.joinedload,), + (orm.subqueryload,), + (orm.selectinload,), + argnames="loader_strategy", + ) + def test_same_relatinship_load_different_range( + self, use_caching, loader_strategy + ): + """This is the first test that exercises lazy loading, which uses + a lambda select, which then needs to transform the select to have + different bound parameters if it's not cached (or generate a working + list of parameters if it is), which then calls into a + with_loader_crieria that itself has another lambda inside of it, + which means we have to traverse and replace that lambda's expression, + but we can't evaluate it until compile time, so the inner lambda + holds onto the "transform" function so it can run it as needed. + this makes use of a new feature in visitors that exports a + "run this traversal later" function. + + All of these individual features, cloning lambdaelements, + running replacement traversals later, are very new and need a lot + of tests, most likely in test/sql/test_lambdas.py. + + the test is from the "temporal_range" example which is the whole + use case this feature is designed for and it is a whopper. + + + """ + Parent, Child = self.classes("Parent", "Child") + temporal_range = self.temporal_range + + if use_caching: + Parent.children.property.bake_queries = True + eng = testing.db + else: + Parent.children.property.bake_queries = False + eng = testing.db.execution_options(compiled_cache=None) + + sess = Session(eng, future=True) + + if loader_strategy: + loader_options = (loader_strategy(Parent.children),) + else: + loader_options = () + + p1 = sess.execute( + select(Parent).filter( + Parent.timestamp == datetime.datetime(2009, 10, 15, 12, 00, 00) + ) + ).scalar() + c1, c2 = p1.children[0:2] + c2_id = c2.id + + p2 = sess.execute( + select(Parent).filter( + Parent.timestamp == datetime.datetime(2009, 10, 17, 12, 00, 00) + ) + ).scalar() + c5 = p2.children[1] + + parents = ( + sess.execute( + select(Parent) + .execution_options(populate_existing=True) + .options( + temporal_range( + datetime.datetime(2009, 10, 16, 12, 00, 00), + datetime.datetime(2009, 10, 18, 12, 00, 00), + ), + *loader_options + ) + ) + .scalars() + .all() + ) + + assert parents[0] == p2 + assert parents[0].children == [c5] + + parents = ( + sess.execute( + select(Parent) + .execution_options(populate_existing=True) + .join(Parent.children) + .filter(Child.id == c2_id) + .options( + temporal_range( + datetime.datetime(2009, 10, 15, 11, 00, 00), + datetime.datetime(2009, 10, 18, 12, 00, 00), + ), + *loader_options + ) + ) + .scalars() + .all() + ) + + assert parents[0] == p1 + assert parents[0].children == [c1, c2] + + +class RelationshipCriteriaTest(_Fixtures, testing.AssertsCompiledSQL): + __dialect__ = "default" + + @testing.fixture + def user_address_fixture(self): + users, Address, addresses, User = ( + self.tables.users, + self.classes.Address, + self.tables.addresses, + self.classes.User, + ) + + mapper( + User, + users, + properties={ + "addresses": relationship( + mapper(Address, addresses), order_by=Address.id + ) + }, + ) + return User, Address + + def test_joinedload_local_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + s = Session(testing.db, future=True) + + stmt = select(User).options( + joinedload(User.addresses.and_(Address.email_address != "email")), + ) + + with self.sql_execution_asserter() as asserter: + + s.execute(stmt) + + asserter.assert_( + CompiledSQL( + "SELECT users.id, users.name, addresses_1.id AS id_1, " + "addresses_1.user_id, addresses_1.email_address FROM " + "users LEFT OUTER JOIN addresses AS addresses_1 " + "ON users.id = addresses_1.user_id " + "AND addresses_1.email_address != :email_address_1 " + "ORDER BY addresses_1.id", + [{"email_address_1": "email"}], + ), + ) + + def test_query_join_local_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + s = Session(testing.db) + + q = s.query(User).join( + User.addresses.and_(Address.email_address != "email") + ) + + self.assert_compile( + q, + "SELECT users.id AS users_id, users.name AS users_name " + "FROM users JOIN addresses ON users.id = addresses.user_id " + "AND addresses.email_address != :email_address_1", + ) + + def test_select_join_local_criteria(self, user_address_fixture): + User, Address = user_address_fixture + + stmt = select(User).join( + User.addresses.and_(Address.email_address != "email") + ) + + self.assert_compile( + stmt, + "SELECT users.id, users.name FROM users JOIN addresses " + "ON users.id = addresses.user_id " + "AND addresses.email_address != :email_address_1", + ) + + def test_select_joinm2m_local_criteria(self, order_item_fixture): + Order, Item = order_item_fixture + + stmt = select(Order).join( + Order.items.and_(Item.description != "description") + ) + + self.assert_compile( + stmt, + "SELECT orders.id, orders.user_id, orders.address_id, " + "orders.description, orders.isopen " + "FROM orders JOIN order_items AS order_items_1 " + "ON orders.id = order_items_1.order_id " + "JOIN items ON items.id = order_items_1.item_id " + "AND items.description != :description_1", + ) + + def test_select_joinm2m_aliased_local_criteria(self, order_item_fixture): + Order, Item = order_item_fixture + + i1 = aliased(Item) + stmt = select(Order).join( + Order.items.of_type(i1).and_(i1.description != "description") + ) + + self.assert_compile( + stmt, + "SELECT orders.id, orders.user_id, orders.address_id, " + "orders.description, orders.isopen " + "FROM orders JOIN order_items AS order_items_1 " + "ON orders.id = order_items_1.order_id " + "JOIN items AS items_1 ON items_1.id = order_items_1.item_id " + "AND items_1.description != :description_1", + ) diff --git a/test/sql/test_compare.py b/test/sql/test_compare.py index b573accbd..7aad2cab8 100644 --- a/test/sql/test_compare.py +++ b/test/sql/test_compare.py @@ -1512,3 +1512,23 @@ class CompareClausesTest(fixtures.TestBase): is_true(x_p_a.compare(x_p)) is_true(x_p.compare(x_p_a)) is_false(x_p_a.compare(x_a)) + + +class ExecutableFlagsTest(fixtures.TestBase): + @testing.combinations( + (select(column("a")),), + (table("q", column("a")).insert(),), + (table("q", column("a")).update(),), + (table("q", column("a")).delete(),), + (lambda_stmt(lambda: select(column("a"))),), + ) + def test_is_select(self, case): + if isinstance(case, LambdaElement): + resolved_case = case._resolved + else: + resolved_case = case + + if isinstance(resolved_case, Select): + is_true(case.is_select) + else: + is_false(case.is_select) |
