summaryrefslogtreecommitdiff
path: root/test
diff options
context:
space:
mode:
authorMike Bayer <mike_mp@zzzcomputing.com>2020-08-05 21:47:43 -0400
committerMike Bayer <mike_mp@zzzcomputing.com>2020-08-05 22:13:11 -0400
commitc7b489b25802f7a25ef78d0731411295c611cc1c (patch)
treef5e3b66ab8eb8bb7398c0195fa2b2f1de8ab91c4 /test
parent71a3ccbdef0d88e9231b7de9c51e4ed60b3b7181 (diff)
downloadsqlalchemy-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.py4
-rw-r--r--test/orm/inheritance/test_polymorphic_rel.py22
-rw-r--r--test/orm/test_bundle.py42
-rw-r--r--test/orm/test_cache_key.py68
-rw-r--r--test/orm/test_events.py166
-rw-r--r--test/orm/test_options.py1
-rw-r--r--test/orm/test_relationship_criteria.py867
-rw-r--r--test/sql/test_compare.py20
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)