diff options
Diffstat (limited to 'test')
54 files changed, 4473 insertions, 3669 deletions
diff --git a/test/aaa_profiling/test_pool.py b/test/aaa_profiling/test_pool.py index ece619fb5..07438f880 100644 --- a/test/aaa_profiling/test_pool.py +++ b/test/aaa_profiling/test_pool.py @@ -30,21 +30,11 @@ class QueuePoolTest(fixtures.TestBase, AssertsExecutionResults): # has the effect of initializing # class-level event listeners on Pool, # if not present already. - p1 = QueuePool( - creator=self.Connection, - pool_size=3, - max_overflow=-1, - use_threadlocal=True, - ) + p1 = QueuePool(creator=self.Connection, pool_size=3, max_overflow=-1) p1.connect() global pool - pool = QueuePool( - creator=self.Connection, - pool_size=3, - max_overflow=-1, - use_threadlocal=True, - ) + pool = QueuePool(creator=self.Connection, pool_size=3, max_overflow=-1) @profiling.function_call_count() def test_first_connect(self): @@ -60,13 +50,3 @@ class QueuePoolTest(fixtures.TestBase, AssertsExecutionResults): return conn2 go() - - def test_second_samethread_connect(self): - conn = pool.connect() - conn # strong ref - - @profiling.function_call_count() - def go(): - return pool.connect() - - go() diff --git a/test/dialect/mssql/test_query.py b/test/dialect/mssql/test_query.py index ef8c63d6b..46af4658d 100644 --- a/test/dialect/mssql/test_query.py +++ b/test/dialect/mssql/test_query.py @@ -211,34 +211,28 @@ class QueryUnicodeTest(fixtures.TestBase): @testing.requires.mssql_freetds @testing.requires.python2 + @testing.provide_metadata def test_convert_unicode(self): - meta = MetaData(testing.db) + meta = self.metadata t1 = Table( "unitest_table", meta, Column("id", Integer, primary_key=True), - Column("descr", mssql.MSText(convert_unicode=True)), + Column("descr", mssql.MSText()), ) meta.create_all() - con = testing.db.connect() - - # encode in UTF-8 (sting object) because this is the default - # dialect encoding - - con.execute( - ue( - "insert into unitest_table values ('bien u\ - umang\xc3\xa9')" - ).encode("UTF-8") - ) - try: - r = t1.select().execute().first() + with testing.db.connect() as con: + con.execute( + ue( + "insert into unitest_table values ('abc \xc3\xa9 def')" + ).encode("UTF-8") + ) + r = con.execute(t1.select()).first() assert isinstance(r[1], util.text_type), ( "%s is %s instead of unicode, working on %s" % (r[1], type(r[1]), meta.bind) ) - finally: - meta.drop_all() + eq_(r[1], util.ue("abc \xc3\xa9 def")) class QueryTest(testing.AssertsExecutionResults, fixtures.TestBase): diff --git a/test/dialect/mysql/test_types.py b/test/dialect/mysql/test_types.py index ed5bbcb2b..12f73fe24 100644 --- a/test/dialect/mysql/test_types.py +++ b/test/dialect/mysql/test_types.py @@ -1117,12 +1117,7 @@ class EnumSetTest( "t", self.metadata, Column("id", Integer, primary_key=True), - Column( - "data", - mysql.SET( - u("réveillé"), u("drôle"), u("S’il"), convert_unicode=True - ), - ), + Column("data", mysql.SET(u("réveillé"), u("drôle"), u("S’il"))), ) set_table.create() diff --git a/test/dialect/oracle/test_compiler.py b/test/dialect/oracle/test_compiler.py index a9c84ba49..596161ef2 100644 --- a/test/dialect/oracle/test_compiler.py +++ b/test/dialect/oracle/test_compiler.py @@ -211,7 +211,7 @@ class CompileTest(fixtures.TestBase, AssertsCompiledSQL): eq_(len(c._result_columns), 2) assert t.c.col1 in set(c._create_result_map()["col1"][1]) - s = select([t], for_update=True).limit(10).order_by(t.c.col2) + s = select([t]).with_for_update().limit(10).order_by(t.c.col2) self.assert_compile( s, "SELECT col1, col2 FROM (SELECT " @@ -222,7 +222,8 @@ class CompileTest(fixtures.TestBase, AssertsCompiledSQL): ) s = ( - select([t], for_update=True) + select([t]) + .with_for_update() .limit(10) .offset(20) .order_by(t.c.col2) diff --git a/test/dialect/oracle/test_dialect.py b/test/dialect/oracle/test_dialect.py index 8d0e37188..1401d40d0 100644 --- a/test/dialect/oracle/test_dialect.py +++ b/test/dialect/oracle/test_dialect.py @@ -85,13 +85,12 @@ class OutParamTest(fixtures.TestBase, AssertsExecutionResults): def test_out_params(self): result = testing.db.execute( text( - "begin foo(:x_in, :x_out, :y_out, " ":z_out); end;", - bindparams=[ - bindparam("x_in", Float), - outparam("x_out", Integer), - outparam("y_out", Float), - outparam("z_out", String), - ], + "begin foo(:x_in, :x_out, :y_out, " ":z_out); end;" + ).bindparams( + bindparam("x_in", Float), + outparam("x_out", Integer), + outparam("y_out", Float), + outparam("z_out", String), ), x_in=5, ) @@ -268,7 +267,7 @@ class ExecuteTest(fixtures.TestBase): # here, we can't use ORDER BY. eq_( - t.select(for_update=True).limit(2).execute().fetchall(), + t.select().with_for_update().limit(2).execute().fetchall(), [(1, 1), (2, 7)], ) @@ -277,7 +276,7 @@ class ExecuteTest(fixtures.TestBase): assert_raises_message( exc.DatabaseError, "ORA-02014", - t.select(for_update=True).limit(2).offset(3).execute, + t.select().with_for_update().limit(2).offset(3).execute, ) diff --git a/test/dialect/oracle/test_types.py b/test/dialect/oracle/test_types.py index dbe91ec03..5af5459ff 100644 --- a/test/dialect/oracle/test_types.py +++ b/test/dialect/oracle/test_types.py @@ -556,15 +556,12 @@ class TypesTest(fixtures.TestBase): ) row = testing.db.execute( - text( - stmt, - typemap={ - "idata": Integer(), - "ndata": Numeric(20, 2), - "ndata2": Numeric(20, 2), - "nidata": Numeric(5, 0), - "fdata": Float(), - }, + text(stmt).columns( + idata=Integer(), + ndata=Numeric(20, 2), + ndata2=Numeric(20, 2), + nidata=Numeric(5, 0), + fdata=Float(), ) ).fetchall()[0] eq_( @@ -616,15 +613,12 @@ class TypesTest(fixtures.TestBase): ) row = testing.db.execute( - text( - stmt, - typemap={ - "anon_1_idata": Integer(), - "anon_1_ndata": Numeric(20, 2), - "anon_1_ndata2": Numeric(20, 2), - "anon_1_nidata": Numeric(5, 0), - "anon_1_fdata": Float(), - }, + text(stmt).columns( + anon_1_idata=Integer(), + anon_1_ndata=Numeric(20, 2), + anon_1_ndata2=Numeric(20, 2), + anon_1_nidata=Numeric(5, 0), + anon_1_fdata=Float(), ) ).fetchall()[0] eq_( @@ -643,15 +637,12 @@ class TypesTest(fixtures.TestBase): ) row = testing.db.execute( - text( - stmt, - typemap={ - "anon_1_idata": Integer(), - "anon_1_ndata": Numeric(20, 2, asdecimal=False), - "anon_1_ndata2": Numeric(20, 2, asdecimal=False), - "anon_1_nidata": Numeric(5, 0, asdecimal=False), - "anon_1_fdata": Float(asdecimal=True), - }, + text(stmt).columns( + anon_1_idata=Integer(), + anon_1_ndata=Numeric(20, 2, asdecimal=False), + anon_1_ndata2=Numeric(20, 2, asdecimal=False), + anon_1_nidata=Numeric(5, 0, asdecimal=False), + anon_1_fdata=Float(asdecimal=True), ) ).fetchall()[0] eq_( diff --git a/test/dialect/postgresql/test_dialect.py b/test/dialect/postgresql/test_dialect.py index d4e1cddc4..cadcbdc1c 100644 --- a/test/dialect/postgresql/test_dialect.py +++ b/test/dialect/postgresql/test_dialect.py @@ -407,7 +407,7 @@ class MiscBackendTest( @testing.fails_on("+zxjdbc", "psycopg2/pg8000 specific assertion") @testing.requires.psycopg2_or_pg8000_compatibility def test_numeric_raise(self): - stmt = text("select cast('hi' as char) as hi", typemap={"hi": Numeric}) + stmt = text("select cast('hi' as char) as hi").columns(hi=Numeric) assert_raises(exc.InvalidRequestError, testing.db.execute, stmt) @testing.only_if( diff --git a/test/dialect/postgresql/test_types.py b/test/dialect/postgresql/test_types.py index f20f92251..e55754f1b 100644 --- a/test/dialect/postgresql/test_types.py +++ b/test/dialect/postgresql/test_types.py @@ -165,8 +165,9 @@ class EnumTest(fixtures.TestBase, AssertsExecutionResults): 'zxjdbc fails on ENUM: column "XXX" is of type ' "XXX but expression is of type character varying", ) + @testing.provide_metadata def test_create_table(self): - metadata = MetaData(testing.db) + metadata = self.metadata t1 = Table( "table", metadata, @@ -175,19 +176,16 @@ class EnumTest(fixtures.TestBase, AssertsExecutionResults): "value", Enum("one", "two", "three", name="onetwothreetype") ), ) - t1.create() - t1.create(checkfirst=True) # check the create - try: - t1.insert().execute(value="two") - t1.insert().execute(value="three") - t1.insert().execute(value="three") + with testing.db.connect() as conn: + t1.create(conn) + t1.create(conn, checkfirst=True) # check the create + conn.execute(t1.insert(), value="two") + conn.execute(t1.insert(), value="three") + conn.execute(t1.insert(), value="three") eq_( - t1.select().order_by(t1.c.id).execute().fetchall(), + conn.execute(t1.select().order_by(t1.c.id)).fetchall(), [(1, "two"), (2, "three"), (3, "three")], ) - finally: - metadata.drop_all() - metadata.drop_all() def test_name_required(self): metadata = MetaData(testing.db) diff --git a/test/dialect/test_sqlite.py b/test/dialect/test_sqlite.py index c480653b3..adfca5b53 100644 --- a/test/dialect/test_sqlite.py +++ b/test/dialect/test_sqlite.py @@ -116,7 +116,7 @@ class TestTypes(fixtures.TestBase, AssertsExecutionResults): ValueError, "Couldn't parse %s string." % disp, lambda: testing.db.execute( - text("select 'ASDF' as value", typemap={"value": typ}) + text("select 'ASDF' as value").columns(value=typ) ).scalar(), ) @@ -254,12 +254,12 @@ class TestTypes(fixtures.TestBase, AssertsExecutionResults): dialect = sqlite.dialect() for t in ( - String(convert_unicode=True), - sqltypes.CHAR(convert_unicode=True), + String(), + sqltypes.CHAR(), sqltypes.Unicode(), sqltypes.UnicodeText(), - String(convert_unicode=True), - sqltypes.CHAR(convert_unicode=True), + String(), + sqltypes.CHAR(), sqltypes.Unicode(), sqltypes.UnicodeText(), ): diff --git a/test/engine/test_bind.py b/test/engine/test_bind.py index 98d53ee6d..ac209c69b 100644 --- a/test/engine/test_bind.py +++ b/test/engine/test_bind.py @@ -23,10 +23,6 @@ class BindTest(fixtures.TestBase): assert not conn.closed assert conn.closed - with e.contextual_connect() as conn: - assert not conn.closed - assert conn.closed - def test_bind_close_conn(self): e = testing.db conn = e.connect() @@ -35,11 +31,6 @@ class BindTest(fixtures.TestBase): assert not conn.closed assert c2.closed - with conn.contextual_connect() as c2: - assert not c2.closed - assert not conn.closed - assert c2.closed - def test_create_drop_explicit(self): metadata = MetaData() table = Table("test_table", metadata, Column("foo", Integer)) diff --git a/test/engine/test_ddlevents.py b/test/engine/test_ddlevents.py index 2762eaa7e..c9177e6ad 100644 --- a/test/engine/test_ddlevents.py +++ b/test/engine/test_ddlevents.py @@ -1,7 +1,6 @@ import sqlalchemy as tsa from sqlalchemy import create_engine from sqlalchemy import event -from sqlalchemy import exc from sqlalchemy import Integer from sqlalchemy import MetaData from sqlalchemy import String @@ -11,7 +10,6 @@ from sqlalchemy.schema import AddConstraint from sqlalchemy.schema import CheckConstraint from sqlalchemy.schema import DDL from sqlalchemy.schema import DropConstraint -from sqlalchemy.testing import assert_raises from sqlalchemy.testing import AssertsCompiledSQL from sqlalchemy.testing import engines from sqlalchemy.testing import eq_ @@ -373,22 +371,6 @@ class DDLEventTest(fixtures.TestBase): ) eq_(metadata_canary.mock_calls, []) - def test_append_listener(self): - metadata, table, bind = self.metadata, self.table, self.bind - - def fn(*a): - return None - - table.append_ddl_listener("before-create", fn) - assert_raises( - exc.InvalidRequestError, table.append_ddl_listener, "blah", fn - ) - - metadata.append_ddl_listener("before-create", fn) - assert_raises( - exc.InvalidRequestError, metadata.append_ddl_listener, "blah", fn - ) - class DDLExecutionTest(fixtures.TestBase): def setup(self): @@ -466,66 +448,6 @@ class DDLExecutionTest(fixtures.TestBase): assert "xyzzy" in strings assert "fnord" in strings - def test_deprecated_append_ddl_listener_table(self): - metadata, users, engine = self.metadata, self.users, self.engine - canary = [] - users.append_ddl_listener( - "before-create", lambda e, t, b: canary.append("mxyzptlk") - ) - users.append_ddl_listener( - "after-create", lambda e, t, b: canary.append("klptzyxm") - ) - users.append_ddl_listener( - "before-drop", lambda e, t, b: canary.append("xyzzy") - ) - users.append_ddl_listener( - "after-drop", lambda e, t, b: canary.append("fnord") - ) - - metadata.create_all() - assert "mxyzptlk" in canary - assert "klptzyxm" in canary - assert "xyzzy" not in canary - assert "fnord" not in canary - del engine.mock[:] - canary[:] = [] - metadata.drop_all() - assert "mxyzptlk" not in canary - assert "klptzyxm" not in canary - assert "xyzzy" in canary - assert "fnord" in canary - - def test_deprecated_append_ddl_listener_metadata(self): - metadata, users, engine = self.metadata, self.users, self.engine - canary = [] - metadata.append_ddl_listener( - "before-create", - lambda e, t, b, tables=None: canary.append("mxyzptlk"), - ) - metadata.append_ddl_listener( - "after-create", - lambda e, t, b, tables=None: canary.append("klptzyxm"), - ) - metadata.append_ddl_listener( - "before-drop", lambda e, t, b, tables=None: canary.append("xyzzy") - ) - metadata.append_ddl_listener( - "after-drop", lambda e, t, b, tables=None: canary.append("fnord") - ) - - metadata.create_all() - assert "mxyzptlk" in canary - assert "klptzyxm" in canary - assert "xyzzy" not in canary - assert "fnord" not in canary - del engine.mock[:] - canary[:] = [] - metadata.drop_all() - assert "mxyzptlk" not in canary - assert "klptzyxm" not in canary - assert "xyzzy" in canary - assert "fnord" in canary - def test_metadata(self): metadata, engine = self.metadata, self.engine @@ -779,27 +701,3 @@ class DDLTest(fixtures.TestBase, AssertsCompiledSQL): ) ._should_execute(tbl, cx) ) - - @testing.uses_deprecated(r"See DDLEvents") - def test_filter_deprecated(self): - cx = self.mock_engine() - - tbl = Table("t", MetaData(), Column("id", Integer)) - target = cx.name - - assert DDL("")._should_execute_deprecated("x", tbl, cx) - assert DDL("", on=target)._should_execute_deprecated("x", tbl, cx) - assert not DDL("", on="bogus")._should_execute_deprecated("x", tbl, cx) - assert DDL("", on=lambda d, x, y, z: True)._should_execute_deprecated( - "x", tbl, cx - ) - assert DDL( - "", on=lambda d, x, y, z: z.engine.name != "bogus" - )._should_execute_deprecated("x", tbl, cx) - - def test_repr(self): - assert repr(DDL("s")) - assert repr(DDL("s", on="engine")) - assert repr(DDL("s", on=lambda x: 1)) - assert repr(DDL("s", context={"a": 1})) - assert repr(DDL("s", on="engine", context={"a": 1})) diff --git a/test/engine/test_deprecations.py b/test/engine/test_deprecations.py new file mode 100644 index 000000000..35226a097 --- /dev/null +++ b/test/engine/test_deprecations.py @@ -0,0 +1,1793 @@ +import re +import time + +import sqlalchemy as tsa +from sqlalchemy import column +from sqlalchemy import create_engine +from sqlalchemy import engine_from_config +from sqlalchemy import event +from sqlalchemy import ForeignKey +from sqlalchemy import func +from sqlalchemy import inspect +from sqlalchemy import INT +from sqlalchemy import Integer +from sqlalchemy import literal +from sqlalchemy import MetaData +from sqlalchemy import pool +from sqlalchemy import select +from sqlalchemy import Sequence +from sqlalchemy import String +from sqlalchemy import testing +from sqlalchemy import text +from sqlalchemy import TypeDecorator +from sqlalchemy import VARCHAR +from sqlalchemy.engine.base import Engine +from sqlalchemy.interfaces import ConnectionProxy +from sqlalchemy.testing import assert_raises_message +from sqlalchemy.testing import engines +from sqlalchemy.testing import eq_ +from sqlalchemy.testing import fixtures +from sqlalchemy.testing.engines import testing_engine +from sqlalchemy.testing.mock import call +from sqlalchemy.testing.mock import Mock +from sqlalchemy.testing.schema import Column +from sqlalchemy.testing.schema import Table +from sqlalchemy.testing.util import gc_collect +from sqlalchemy.testing.util import lazy_gc +from .test_parseconnect import mock_dbapi + +tlengine = None + + +class SomeException(Exception): + pass + + +def _tlengine_deprecated(): + return testing.expect_deprecated( + "The 'threadlocal' engine strategy is deprecated" + ) + + +class TableNamesOrderByTest(fixtures.TestBase): + @testing.provide_metadata + def test_order_by_foreign_key(self): + Table( + "t1", + self.metadata, + Column("id", Integer, primary_key=True), + test_needs_acid=True, + ) + Table( + "t2", + self.metadata, + Column("id", Integer, primary_key=True), + Column("t1id", Integer, ForeignKey("t1.id")), + test_needs_acid=True, + ) + Table( + "t3", + self.metadata, + Column("id", Integer, primary_key=True), + Column("t2id", Integer, ForeignKey("t2.id")), + test_needs_acid=True, + ) + self.metadata.create_all() + insp = inspect(testing.db) + with testing.expect_deprecated( + "The get_table_names.order_by parameter is deprecated " + ): + tnames = insp.get_table_names(order_by="foreign_key") + eq_(tnames, ["t1", "t2", "t3"]) + + +class CreateEngineTest(fixtures.TestBase): + def test_pool_threadlocal_from_config(self): + dbapi = mock_dbapi + + config = { + "sqlalchemy.url": "postgresql://scott:tiger@somehost/test", + "sqlalchemy.pool_threadlocal": "false", + } + + e = engine_from_config(config, module=dbapi, _initialize=False) + eq_(e.pool._use_threadlocal, False) + + config = { + "sqlalchemy.url": "postgresql://scott:tiger@somehost/test", + "sqlalchemy.pool_threadlocal": "true", + } + + with testing.expect_deprecated( + "The Pool.use_threadlocal parameter is deprecated" + ): + e = engine_from_config(config, module=dbapi, _initialize=False) + eq_(e.pool._use_threadlocal, True) + + +class RecycleTest(fixtures.TestBase): + __backend__ = True + + def test_basic(self): + with testing.expect_deprecated( + "The Pool.use_threadlocal parameter is deprecated" + ): + engine = engines.reconnecting_engine( + options={"pool_threadlocal": True} + ) + + with testing.expect_deprecated( + r"The Engine.contextual_connect\(\) method is deprecated" + ): + conn = engine.contextual_connect() + eq_(conn.execute(select([1])).scalar(), 1) + conn.close() + + # set the pool recycle down to 1. + # we aren't doing this inline with the + # engine create since cx_oracle takes way + # too long to create the 1st connection and don't + # want to build a huge delay into this test. + + engine.pool._recycle = 1 + + # kill the DB connection + engine.test_shutdown() + + # wait until past the recycle period + time.sleep(2) + + # can connect, no exception + with testing.expect_deprecated( + r"The Engine.contextual_connect\(\) method is deprecated" + ): + conn = engine.contextual_connect() + eq_(conn.execute(select([1])).scalar(), 1) + conn.close() + + +class TLTransactionTest(fixtures.TestBase): + __requires__ = ("ad_hoc_engines",) + __backend__ = True + + @classmethod + def setup_class(cls): + global users, metadata, tlengine + + with _tlengine_deprecated(): + tlengine = testing_engine(options=dict(strategy="threadlocal")) + metadata = MetaData() + users = Table( + "query_users", + metadata, + Column( + "user_id", + INT, + Sequence("query_users_id_seq", optional=True), + primary_key=True, + ), + Column("user_name", VARCHAR(20)), + test_needs_acid=True, + ) + metadata.create_all(tlengine) + + def teardown(self): + tlengine.execute(users.delete()).close() + + @classmethod + def teardown_class(cls): + tlengine.close() + metadata.drop_all(tlengine) + tlengine.dispose() + + def setup(self): + + # ensure tests start with engine closed + + tlengine.close() + + @testing.crashes( + "oracle", "TNS error of unknown origin occurs on the buildbot." + ) + def test_rollback_no_trans(self): + with _tlengine_deprecated(): + tlengine = testing_engine(options=dict(strategy="threadlocal")) + + # shouldn't fail + tlengine.rollback() + + tlengine.begin() + tlengine.rollback() + + # shouldn't fail + tlengine.rollback() + + def test_commit_no_trans(self): + with _tlengine_deprecated(): + tlengine = testing_engine(options=dict(strategy="threadlocal")) + + # shouldn't fail + tlengine.commit() + + tlengine.begin() + tlengine.rollback() + + # shouldn't fail + tlengine.commit() + + def test_prepare_no_trans(self): + with _tlengine_deprecated(): + tlengine = testing_engine(options=dict(strategy="threadlocal")) + + # shouldn't fail + tlengine.prepare() + + tlengine.begin() + tlengine.rollback() + + # shouldn't fail + tlengine.prepare() + + def test_connection_close(self): + """test that when connections are closed for real, transactions + are rolled back and disposed.""" + + c = tlengine.contextual_connect() + c.begin() + assert c.in_transaction() + c.close() + assert not c.in_transaction() + + def test_transaction_close(self): + c = tlengine.contextual_connect() + t = c.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.execute(users.insert(), user_id=2, user_name="user2") + t2 = c.begin() + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.execute(users.insert(), user_id=4, user_name="user4") + t2.close() + result = c.execute("select * from query_users") + assert len(result.fetchall()) == 4 + t.close() + external_connection = tlengine.connect() + result = external_connection.execute("select * from query_users") + try: + assert len(result.fetchall()) == 0 + finally: + c.close() + external_connection.close() + + def test_rollback(self): + """test a basic rollback""" + + tlengine.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.rollback() + external_connection = tlengine.connect() + result = external_connection.execute("select * from query_users") + try: + assert len(result.fetchall()) == 0 + finally: + external_connection.close() + + def test_commit(self): + """test a basic commit""" + + tlengine.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.commit() + external_connection = tlengine.connect() + result = external_connection.execute("select * from query_users") + try: + assert len(result.fetchall()) == 3 + finally: + external_connection.close() + + def test_with_interface(self): + trans = tlengine.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.execute(users.insert(), user_id=2, user_name="user2") + trans.commit() + + trans = tlengine.begin() + tlengine.execute(users.insert(), user_id=3, user_name="user3") + trans.__exit__(Exception, "fake", None) + trans = tlengine.begin() + tlengine.execute(users.insert(), user_id=4, user_name="user4") + trans.__exit__(None, None, None) + eq_( + tlengine.execute( + users.select().order_by(users.c.user_id) + ).fetchall(), + [(1, "user1"), (2, "user2"), (4, "user4")], + ) + + def test_commits(self): + connection = tlengine.connect() + assert ( + connection.execute("select count(*) from query_users").scalar() + == 0 + ) + connection.close() + connection = tlengine.contextual_connect() + transaction = connection.begin() + connection.execute(users.insert(), user_id=1, user_name="user1") + transaction.commit() + transaction = connection.begin() + connection.execute(users.insert(), user_id=2, user_name="user2") + connection.execute(users.insert(), user_id=3, user_name="user3") + transaction.commit() + transaction = connection.begin() + result = connection.execute("select * from query_users") + rows = result.fetchall() + assert len(rows) == 3, "expected 3 got %d" % len(rows) + transaction.commit() + connection.close() + + def test_rollback_off_conn(self): + + # test that a TLTransaction opened off a TLConnection allows + # that TLConnection to be aware of the transactional context + + conn = tlengine.contextual_connect() + trans = conn.begin() + conn.execute(users.insert(), user_id=1, user_name="user1") + conn.execute(users.insert(), user_id=2, user_name="user2") + conn.execute(users.insert(), user_id=3, user_name="user3") + trans.rollback() + external_connection = tlengine.connect() + result = external_connection.execute("select * from query_users") + try: + assert len(result.fetchall()) == 0 + finally: + conn.close() + external_connection.close() + + def test_morerollback_off_conn(self): + + # test that an existing TLConnection automatically takes place + # in a TLTransaction opened on a second TLConnection + + conn = tlengine.contextual_connect() + conn2 = tlengine.contextual_connect() + trans = conn2.begin() + conn.execute(users.insert(), user_id=1, user_name="user1") + conn.execute(users.insert(), user_id=2, user_name="user2") + conn.execute(users.insert(), user_id=3, user_name="user3") + trans.rollback() + external_connection = tlengine.connect() + result = external_connection.execute("select * from query_users") + try: + assert len(result.fetchall()) == 0 + finally: + conn.close() + conn2.close() + external_connection.close() + + def test_commit_off_connection(self): + conn = tlengine.contextual_connect() + trans = conn.begin() + conn.execute(users.insert(), user_id=1, user_name="user1") + conn.execute(users.insert(), user_id=2, user_name="user2") + conn.execute(users.insert(), user_id=3, user_name="user3") + trans.commit() + external_connection = tlengine.connect() + result = external_connection.execute("select * from query_users") + try: + assert len(result.fetchall()) == 3 + finally: + conn.close() + external_connection.close() + + def test_nesting_rollback(self): + """tests nesting of transactions, rollback at the end""" + + external_connection = tlengine.connect() + self.assert_( + external_connection.connection + is not tlengine.contextual_connect().connection + ) + tlengine.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.begin() + tlengine.execute(users.insert(), user_id=4, user_name="user4") + tlengine.execute(users.insert(), user_id=5, user_name="user5") + tlengine.commit() + tlengine.rollback() + try: + self.assert_( + external_connection.scalar("select count(*) from query_users") + == 0 + ) + finally: + external_connection.close() + + def test_nesting_commit(self): + """tests nesting of transactions, commit at the end.""" + + external_connection = tlengine.connect() + self.assert_( + external_connection.connection + is not tlengine.contextual_connect().connection + ) + tlengine.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.begin() + tlengine.execute(users.insert(), user_id=4, user_name="user4") + tlengine.execute(users.insert(), user_id=5, user_name="user5") + tlengine.commit() + tlengine.commit() + try: + self.assert_( + external_connection.scalar("select count(*) from query_users") + == 5 + ) + finally: + external_connection.close() + + def test_mixed_nesting(self): + """tests nesting of transactions off the TLEngine directly + inside of transactions off the connection from the TLEngine""" + + external_connection = tlengine.connect() + self.assert_( + external_connection.connection + is not tlengine.contextual_connect().connection + ) + conn = tlengine.contextual_connect() + trans = conn.begin() + trans2 = conn.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.begin() + tlengine.execute(users.insert(), user_id=4, user_name="user4") + tlengine.begin() + tlengine.execute(users.insert(), user_id=5, user_name="user5") + tlengine.execute(users.insert(), user_id=6, user_name="user6") + tlengine.execute(users.insert(), user_id=7, user_name="user7") + tlengine.commit() + tlengine.execute(users.insert(), user_id=8, user_name="user8") + tlengine.commit() + trans2.commit() + trans.rollback() + conn.close() + try: + self.assert_( + external_connection.scalar("select count(*) from query_users") + == 0 + ) + finally: + external_connection.close() + + def test_more_mixed_nesting(self): + """tests nesting of transactions off the connection from the + TLEngine inside of transactions off the TLEngine directly.""" + + external_connection = tlengine.connect() + self.assert_( + external_connection.connection + is not tlengine.contextual_connect().connection + ) + tlengine.begin() + connection = tlengine.contextual_connect() + connection.execute(users.insert(), user_id=1, user_name="user1") + tlengine.begin() + connection.execute(users.insert(), user_id=2, user_name="user2") + connection.execute(users.insert(), user_id=3, user_name="user3") + trans = connection.begin() + connection.execute(users.insert(), user_id=4, user_name="user4") + connection.execute(users.insert(), user_id=5, user_name="user5") + trans.commit() + tlengine.commit() + tlengine.rollback() + connection.close() + try: + self.assert_( + external_connection.scalar("select count(*) from query_users") + == 0 + ) + finally: + external_connection.close() + + @testing.requires.savepoints + def test_nested_subtransaction_rollback(self): + tlengine.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.begin_nested() + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.rollback() + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.commit() + tlengine.close() + eq_( + tlengine.execute( + select([users.c.user_id]).order_by(users.c.user_id) + ).fetchall(), + [(1,), (3,)], + ) + tlengine.close() + + @testing.requires.savepoints + @testing.crashes( + "oracle+zxjdbc", + "Errors out and causes subsequent tests to " "deadlock", + ) + def test_nested_subtransaction_commit(self): + tlengine.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.begin_nested() + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.commit() + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.commit() + tlengine.close() + eq_( + tlengine.execute( + select([users.c.user_id]).order_by(users.c.user_id) + ).fetchall(), + [(1,), (2,), (3,)], + ) + tlengine.close() + + @testing.requires.savepoints + def test_rollback_to_subtransaction(self): + tlengine.begin() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.begin_nested() + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.begin() + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.rollback() + tlengine.rollback() + tlengine.execute(users.insert(), user_id=4, user_name="user4") + tlengine.commit() + tlengine.close() + eq_( + tlengine.execute( + select([users.c.user_id]).order_by(users.c.user_id) + ).fetchall(), + [(1,), (4,)], + ) + tlengine.close() + + def test_connections(self): + """tests that contextual_connect is threadlocal""" + + c1 = tlengine.contextual_connect() + c2 = tlengine.contextual_connect() + assert c1.connection is c2.connection + c2.close() + assert not c1.closed + assert not tlengine.closed + + @testing.requires.independent_cursors + def test_result_closing(self): + """tests that contextual_connect is threadlocal""" + + r1 = tlengine.execute(select([1])) + r2 = tlengine.execute(select([1])) + row1 = r1.fetchone() + row2 = r2.fetchone() + r1.close() + assert r2.connection is r1.connection + assert not r2.connection.closed + assert not tlengine.closed + + # close again, nothing happens since resultproxy calls close() + # only once + + r1.close() + assert r2.connection is r1.connection + assert not r2.connection.closed + assert not tlengine.closed + r2.close() + assert r2.connection.closed + assert tlengine.closed + + @testing.crashes( + "oracle+cx_oracle", "intermittent failures on the buildbot" + ) + def test_dispose(self): + with _tlengine_deprecated(): + eng = testing_engine(options=dict(strategy="threadlocal")) + result = eng.execute(select([1])) + eng.dispose() + eng.execute(select([1])) + + @testing.requires.two_phase_transactions + def test_two_phase_transaction(self): + tlengine.begin_twophase() + tlengine.execute(users.insert(), user_id=1, user_name="user1") + tlengine.prepare() + tlengine.commit() + tlengine.begin_twophase() + tlengine.execute(users.insert(), user_id=2, user_name="user2") + tlengine.commit() + tlengine.begin_twophase() + tlengine.execute(users.insert(), user_id=3, user_name="user3") + tlengine.rollback() + tlengine.begin_twophase() + tlengine.execute(users.insert(), user_id=4, user_name="user4") + tlengine.prepare() + tlengine.rollback() + eq_( + tlengine.execute( + select([users.c.user_id]).order_by(users.c.user_id) + ).fetchall(), + [(1,), (2,)], + ) + + +class ConvenienceExecuteTest(fixtures.TablesTest): + __backend__ = True + + @classmethod + def define_tables(cls, metadata): + cls.table = Table( + "exec_test", + metadata, + Column("a", Integer), + Column("b", Integer), + test_needs_acid=True, + ) + + def _trans_fn(self, is_transaction=False): + def go(conn, x, value=None): + if is_transaction: + conn = conn.connection + conn.execute(self.table.insert().values(a=x, b=value)) + + return go + + def _trans_rollback_fn(self, is_transaction=False): + def go(conn, x, value=None): + if is_transaction: + conn = conn.connection + conn.execute(self.table.insert().values(a=x, b=value)) + raise SomeException("breakage") + + return go + + def _assert_no_data(self): + eq_( + testing.db.scalar( + select([func.count("*")]).select_from(self.table) + ), + 0, + ) + + def _assert_fn(self, x, value=None): + eq_(testing.db.execute(self.table.select()).fetchall(), [(x, value)]) + + def test_transaction_tlocal_engine_ctx_commit(self): + fn = self._trans_fn() + with _tlengine_deprecated(): + engine = engines.testing_engine( + options=dict(strategy="threadlocal", pool=testing.db.pool) + ) + ctx = engine.begin() + testing.run_as_contextmanager(ctx, fn, 5, value=8) + self._assert_fn(5, value=8) + + def test_transaction_tlocal_engine_ctx_rollback(self): + fn = self._trans_rollback_fn() + with _tlengine_deprecated(): + engine = engines.testing_engine( + options=dict(strategy="threadlocal", pool=testing.db.pool) + ) + ctx = engine.begin() + assert_raises_message( + Exception, + "breakage", + testing.run_as_contextmanager, + ctx, + fn, + 5, + value=8, + ) + self._assert_no_data() + + +def _proxy_execute_deprecated(): + return ( + testing.expect_deprecated("ConnectionProxy.execute is deprecated."), + testing.expect_deprecated( + "ConnectionProxy.cursor_execute is deprecated." + ), + ) + + +class ProxyConnectionTest(fixtures.TestBase): + + """These are the same tests as EngineEventsTest, except using + the deprecated ConnectionProxy interface. + + """ + + __requires__ = ("ad_hoc_engines",) + __prefer_requires__ = ("two_phase_transactions",) + + @testing.uses_deprecated(r".*Use event.listen") + @testing.fails_on("firebird", "Data type unknown") + def test_proxy(self): + + stmts = [] + cursor_stmts = [] + + class MyProxy(ConnectionProxy): + def execute( + self, conn, execute, clauseelement, *multiparams, **params + ): + stmts.append((str(clauseelement), params, multiparams)) + return execute(clauseelement, *multiparams, **params) + + def cursor_execute( + self, + execute, + cursor, + statement, + parameters, + context, + executemany, + ): + cursor_stmts.append((str(statement), parameters, None)) + return execute(cursor, statement, parameters, context) + + def assert_stmts(expected, received): + for stmt, params, posn in expected: + if not received: + assert False, "Nothing available for stmt: %s" % stmt + while received: + teststmt, testparams, testmultiparams = received.pop(0) + teststmt = ( + re.compile(r"[\n\t ]+", re.M) + .sub(" ", teststmt) + .strip() + ) + if teststmt.startswith(stmt) and ( + testparams == params or testparams == posn + ): + break + + with testing.expect_deprecated( + "ConnectionProxy.execute is deprecated.", + "ConnectionProxy.cursor_execute is deprecated.", + ): + plain_engine = engines.testing_engine( + options=dict(implicit_returning=False, proxy=MyProxy()) + ) + + with testing.expect_deprecated( + "ConnectionProxy.execute is deprecated.", + "ConnectionProxy.cursor_execute is deprecated.", + "The 'threadlocal' engine strategy is deprecated", + ): + + tl_engine = engines.testing_engine( + options=dict( + implicit_returning=False, + proxy=MyProxy(), + strategy="threadlocal", + ) + ) + + for engine in (plain_engine, tl_engine): + m = MetaData(engine) + t1 = Table( + "t1", + m, + Column("c1", Integer, primary_key=True), + Column( + "c2", + String(50), + default=func.lower("Foo"), + primary_key=True, + ), + ) + m.create_all() + try: + t1.insert().execute(c1=5, c2="some data") + t1.insert().execute(c1=6) + eq_( + engine.execute("select * from t1").fetchall(), + [(5, "some data"), (6, "foo")], + ) + finally: + m.drop_all() + engine.dispose() + compiled = [ + ("CREATE TABLE t1", {}, None), + ( + "INSERT INTO t1 (c1, c2)", + {"c2": "some data", "c1": 5}, + None, + ), + ("INSERT INTO t1 (c1, c2)", {"c1": 6}, None), + ("select * from t1", {}, None), + ("DROP TABLE t1", {}, None), + ] + + cursor = [ + ("CREATE TABLE t1", {}, ()), + ( + "INSERT INTO t1 (c1, c2)", + {"c2": "some data", "c1": 5}, + (5, "some data"), + ), + ("SELECT lower", {"lower_1": "Foo"}, ("Foo",)), + ( + "INSERT INTO t1 (c1, c2)", + {"c2": "foo", "c1": 6}, + (6, "foo"), + ), + ("select * from t1", {}, ()), + ("DROP TABLE t1", {}, ()), + ] + + assert_stmts(compiled, stmts) + assert_stmts(cursor, cursor_stmts) + + @testing.uses_deprecated(r".*Use event.listen") + def test_options(self): + canary = [] + + class TrackProxy(ConnectionProxy): + def __getattribute__(self, key): + fn = object.__getattribute__(self, key) + + def go(*arg, **kw): + canary.append(fn.__name__) + return fn(*arg, **kw) + + return go + + with testing.expect_deprecated( + *[ + "ConnectionProxy.%s is deprecated" % name + for name in [ + "execute", + "cursor_execute", + "begin", + "rollback", + "commit", + "savepoint", + "rollback_savepoint", + "release_savepoint", + "begin_twophase", + "prepare_twophase", + "rollback_twophase", + "commit_twophase", + ] + ] + ): + engine = engines.testing_engine(options={"proxy": TrackProxy()}) + conn = engine.connect() + c2 = conn.execution_options(foo="bar") + eq_(c2._execution_options, {"foo": "bar"}) + c2.execute(select([1])) + c3 = c2.execution_options(bar="bat") + eq_(c3._execution_options, {"foo": "bar", "bar": "bat"}) + eq_(canary, ["execute", "cursor_execute"]) + + @testing.uses_deprecated(r".*Use event.listen") + def test_transactional(self): + canary = [] + + class TrackProxy(ConnectionProxy): + def __getattribute__(self, key): + fn = object.__getattribute__(self, key) + + def go(*arg, **kw): + canary.append(fn.__name__) + return fn(*arg, **kw) + + return go + + with testing.expect_deprecated( + *[ + "ConnectionProxy.%s is deprecated" % name + for name in [ + "execute", + "cursor_execute", + "begin", + "rollback", + "commit", + "savepoint", + "rollback_savepoint", + "release_savepoint", + "begin_twophase", + "prepare_twophase", + "rollback_twophase", + "commit_twophase", + ] + ] + ): + engine = engines.testing_engine(options={"proxy": TrackProxy()}) + conn = engine.connect() + trans = conn.begin() + conn.execute(select([1])) + trans.rollback() + trans = conn.begin() + conn.execute(select([1])) + trans.commit() + + eq_( + canary, + [ + "begin", + "execute", + "cursor_execute", + "rollback", + "begin", + "execute", + "cursor_execute", + "commit", + ], + ) + + @testing.uses_deprecated(r".*Use event.listen") + @testing.requires.savepoints + @testing.requires.two_phase_transactions + def test_transactional_advanced(self): + canary = [] + + class TrackProxy(ConnectionProxy): + def __getattribute__(self, key): + fn = object.__getattribute__(self, key) + + def go(*arg, **kw): + canary.append(fn.__name__) + return fn(*arg, **kw) + + return go + + with testing.expect_deprecated( + *[ + "ConnectionProxy.%s is deprecated" % name + for name in [ + "execute", + "cursor_execute", + "begin", + "rollback", + "commit", + "savepoint", + "rollback_savepoint", + "release_savepoint", + "begin_twophase", + "prepare_twophase", + "rollback_twophase", + "commit_twophase", + ] + ] + ): + engine = engines.testing_engine(options={"proxy": TrackProxy()}) + conn = engine.connect() + + trans = conn.begin() + trans2 = conn.begin_nested() + conn.execute(select([1])) + trans2.rollback() + trans2 = conn.begin_nested() + conn.execute(select([1])) + trans2.commit() + trans.rollback() + + trans = conn.begin_twophase() + conn.execute(select([1])) + trans.prepare() + trans.commit() + + canary = [t for t in canary if t not in ("cursor_execute", "execute")] + eq_( + canary, + [ + "begin", + "savepoint", + "rollback_savepoint", + "savepoint", + "release_savepoint", + "rollback", + "begin_twophase", + "prepare_twophase", + "commit_twophase", + ], + ) + + +class HandleInvalidatedOnConnectTest(fixtures.TestBase): + __requires__ = ("sqlite",) + + def setUp(self): + e = create_engine("sqlite://") + + connection = Mock(get_server_version_info=Mock(return_value="5.0")) + + def connect(*args, **kwargs): + return connection + + dbapi = Mock( + sqlite_version_info=(99, 9, 9), + version_info=(99, 9, 9), + sqlite_version="99.9.9", + paramstyle="named", + connect=Mock(side_effect=connect), + ) + + sqlite3 = e.dialect.dbapi + dbapi.Error = (sqlite3.Error,) + dbapi.ProgrammingError = sqlite3.ProgrammingError + + self.dbapi = dbapi + self.ProgrammingError = sqlite3.ProgrammingError + + def test_dont_touch_non_dbapi_exception_on_contextual_connect(self): + dbapi = self.dbapi + dbapi.connect = Mock(side_effect=TypeError("I'm not a DBAPI error")) + + e = create_engine("sqlite://", module=dbapi) + e.dialect.is_disconnect = is_disconnect = Mock() + with testing.expect_deprecated( + r"The Engine.contextual_connect\(\) method is deprecated" + ): + assert_raises_message( + TypeError, "I'm not a DBAPI error", e.contextual_connect + ) + eq_(is_disconnect.call_count, 0) + + def test_invalidate_on_contextual_connect(self): + """test that is_disconnect() is called during connect. + + interpretation of connection failures are not supported by + every backend. + + """ + + dbapi = self.dbapi + dbapi.connect = Mock( + side_effect=self.ProgrammingError( + "Cannot operate on a closed database." + ) + ) + e = create_engine("sqlite://", module=dbapi) + try: + with testing.expect_deprecated( + r"The Engine.contextual_connect\(\) method is deprecated" + ): + e.contextual_connect() + assert False + except tsa.exc.DBAPIError as de: + assert de.connection_invalidated + + +class HandleErrorTest(fixtures.TestBase): + __requires__ = ("ad_hoc_engines",) + __backend__ = True + + def tearDown(self): + Engine.dispatch._clear() + Engine._has_events = False + + def test_legacy_dbapi_error(self): + engine = engines.testing_engine() + canary = Mock() + + with testing.expect_deprecated( + r"The ConnectionEvents.dbapi_error\(\) event is deprecated" + ): + event.listen(engine, "dbapi_error", canary) + + with engine.connect() as conn: + try: + conn.execute("SELECT FOO FROM I_DONT_EXIST") + assert False + except tsa.exc.DBAPIError as e: + eq_(canary.mock_calls[0][1][5], e.orig) + eq_(canary.mock_calls[0][1][2], "SELECT FOO FROM I_DONT_EXIST") + + def test_legacy_dbapi_error_no_ad_hoc_context(self): + engine = engines.testing_engine() + + listener = Mock(return_value=None) + with testing.expect_deprecated( + r"The ConnectionEvents.dbapi_error\(\) event is deprecated" + ): + event.listen(engine, "dbapi_error", listener) + + nope = SomeException("nope") + + class MyType(TypeDecorator): + impl = Integer + + def process_bind_param(self, value, dialect): + raise nope + + with engine.connect() as conn: + assert_raises_message( + tsa.exc.StatementError, + r"\(.*SomeException\) " r"nope \[SQL\: u?'SELECT 1 ", + conn.execute, + select([1]).where(column("foo") == literal("bar", MyType())), + ) + # no legacy event + eq_(listener.mock_calls, []) + + def test_legacy_dbapi_error_non_dbapi_error(self): + engine = engines.testing_engine() + + listener = Mock(return_value=None) + with testing.expect_deprecated( + r"The ConnectionEvents.dbapi_error\(\) event is deprecated" + ): + event.listen(engine, "dbapi_error", listener) + + nope = TypeError("I'm not a DBAPI error") + with engine.connect() as c: + c.connection.cursor = Mock( + return_value=Mock(execute=Mock(side_effect=nope)) + ) + + assert_raises_message( + TypeError, "I'm not a DBAPI error", c.execute, "select " + ) + # no legacy event + eq_(listener.mock_calls, []) + + +def MockDBAPI(): # noqa + def cursor(): + return Mock() + + def connect(*arg, **kw): + def close(): + conn.closed = True + + # mock seems like it might have an issue logging + # call_count correctly under threading, not sure. + # adding a side_effect for close seems to help. + conn = Mock( + cursor=Mock(side_effect=cursor), + close=Mock(side_effect=close), + closed=False, + ) + return conn + + def shutdown(value): + if value: + db.connect = Mock(side_effect=Exception("connect failed")) + else: + db.connect = Mock(side_effect=connect) + db.is_shutdown = value + + db = Mock( + connect=Mock(side_effect=connect), shutdown=shutdown, is_shutdown=False + ) + return db + + +class PoolTestBase(fixtures.TestBase): + def setup(self): + pool.clear_managers() + self._teardown_conns = [] + + def teardown(self): + for ref in self._teardown_conns: + conn = ref() + if conn: + conn.close() + + @classmethod + def teardown_class(cls): + pool.clear_managers() + + def _queuepool_fixture(self, **kw): + dbapi, pool = self._queuepool_dbapi_fixture(**kw) + return pool + + def _queuepool_dbapi_fixture(self, **kw): + dbapi = MockDBAPI() + return ( + dbapi, + pool.QueuePool(creator=lambda: dbapi.connect("foo.db"), **kw), + ) + + +class DeprecatedPoolListenerTest(PoolTestBase): + @testing.requires.predictable_gc + @testing.uses_deprecated( + r".*Use the PoolEvents", r".*'listeners' argument .* is deprecated" + ) + def test_listeners(self): + class InstrumentingListener(object): + def __init__(self): + if hasattr(self, "connect"): + self.connect = self.inst_connect + if hasattr(self, "first_connect"): + self.first_connect = self.inst_first_connect + if hasattr(self, "checkout"): + self.checkout = self.inst_checkout + if hasattr(self, "checkin"): + self.checkin = self.inst_checkin + self.clear() + + def clear(self): + self.connected = [] + self.first_connected = [] + self.checked_out = [] + self.checked_in = [] + + def assert_total(self, conn, fconn, cout, cin): + eq_(len(self.connected), conn) + eq_(len(self.first_connected), fconn) + eq_(len(self.checked_out), cout) + eq_(len(self.checked_in), cin) + + def assert_in(self, item, in_conn, in_fconn, in_cout, in_cin): + eq_((item in self.connected), in_conn) + eq_((item in self.first_connected), in_fconn) + eq_((item in self.checked_out), in_cout) + eq_((item in self.checked_in), in_cin) + + def inst_connect(self, con, record): + print("connect(%s, %s)" % (con, record)) + assert con is not None + assert record is not None + self.connected.append(con) + + def inst_first_connect(self, con, record): + print("first_connect(%s, %s)" % (con, record)) + assert con is not None + assert record is not None + self.first_connected.append(con) + + def inst_checkout(self, con, record, proxy): + print("checkout(%s, %s, %s)" % (con, record, proxy)) + assert con is not None + assert record is not None + assert proxy is not None + self.checked_out.append(con) + + def inst_checkin(self, con, record): + print("checkin(%s, %s)" % (con, record)) + # con can be None if invalidated + assert record is not None + self.checked_in.append(con) + + class ListenAll(tsa.interfaces.PoolListener, InstrumentingListener): + pass + + class ListenConnect(InstrumentingListener): + def connect(self, con, record): + pass + + class ListenFirstConnect(InstrumentingListener): + def first_connect(self, con, record): + pass + + class ListenCheckOut(InstrumentingListener): + def checkout(self, con, record, proxy, num): + pass + + class ListenCheckIn(InstrumentingListener): + def checkin(self, con, record): + pass + + def assert_listeners(p, total, conn, fconn, cout, cin): + for instance in (p, p.recreate()): + self.assert_(len(instance.dispatch.connect) == conn) + self.assert_(len(instance.dispatch.first_connect) == fconn) + self.assert_(len(instance.dispatch.checkout) == cout) + self.assert_(len(instance.dispatch.checkin) == cin) + + p = self._queuepool_fixture() + assert_listeners(p, 0, 0, 0, 0, 0) + + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["connect", "first_connect", "checkout", "checkin"] + ] + ): + p.add_listener(ListenAll()) + assert_listeners(p, 1, 1, 1, 1, 1) + + with testing.expect_deprecated( + *["PoolListener.%s is deprecated." % name for name in ["connect"]] + ): + p.add_listener(ListenConnect()) + assert_listeners(p, 2, 2, 1, 1, 1) + + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["first_connect"] + ] + ): + p.add_listener(ListenFirstConnect()) + assert_listeners(p, 3, 2, 2, 1, 1) + + with testing.expect_deprecated( + *["PoolListener.%s is deprecated." % name for name in ["checkout"]] + ): + p.add_listener(ListenCheckOut()) + assert_listeners(p, 4, 2, 2, 2, 1) + + with testing.expect_deprecated( + *["PoolListener.%s is deprecated." % name for name in ["checkin"]] + ): + p.add_listener(ListenCheckIn()) + assert_listeners(p, 5, 2, 2, 2, 2) + del p + + snoop = ListenAll() + + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["connect", "first_connect", "checkout", "checkin"] + ] + + [ + "PoolListener is deprecated in favor of the PoolEvents " + "listener interface. The Pool.listeners parameter " + "will be removed" + ] + ): + p = self._queuepool_fixture(listeners=[snoop]) + assert_listeners(p, 1, 1, 1, 1, 1) + + c = p.connect() + snoop.assert_total(1, 1, 1, 0) + cc = c.connection + snoop.assert_in(cc, True, True, True, False) + c.close() + snoop.assert_in(cc, True, True, True, True) + del c, cc + + snoop.clear() + + # this one depends on immediate gc + c = p.connect() + cc = c.connection + snoop.assert_in(cc, False, False, True, False) + snoop.assert_total(0, 0, 1, 0) + del c, cc + lazy_gc() + snoop.assert_total(0, 0, 1, 1) + + p.dispose() + snoop.clear() + + c = p.connect() + c.close() + c = p.connect() + snoop.assert_total(1, 0, 2, 1) + c.close() + snoop.assert_total(1, 0, 2, 2) + + # invalidation + p.dispose() + snoop.clear() + + c = p.connect() + snoop.assert_total(1, 0, 1, 0) + c.invalidate() + snoop.assert_total(1, 0, 1, 1) + c.close() + snoop.assert_total(1, 0, 1, 1) + del c + lazy_gc() + snoop.assert_total(1, 0, 1, 1) + c = p.connect() + snoop.assert_total(2, 0, 2, 1) + c.close() + del c + lazy_gc() + snoop.assert_total(2, 0, 2, 2) + + # detached + p.dispose() + snoop.clear() + + c = p.connect() + snoop.assert_total(1, 0, 1, 0) + c.detach() + snoop.assert_total(1, 0, 1, 0) + c.close() + del c + snoop.assert_total(1, 0, 1, 0) + c = p.connect() + snoop.assert_total(2, 0, 2, 0) + c.close() + del c + snoop.assert_total(2, 0, 2, 1) + + # recreated + p = p.recreate() + snoop.clear() + + c = p.connect() + snoop.assert_total(1, 1, 1, 0) + c.close() + snoop.assert_total(1, 1, 1, 1) + c = p.connect() + snoop.assert_total(1, 1, 2, 1) + c.close() + snoop.assert_total(1, 1, 2, 2) + + @testing.uses_deprecated(r".*Use the PoolEvents") + def test_listeners_callables(self): + def connect(dbapi_con, con_record): + counts[0] += 1 + + def checkout(dbapi_con, con_record, con_proxy): + counts[1] += 1 + + def checkin(dbapi_con, con_record): + counts[2] += 1 + + i_all = dict(connect=connect, checkout=checkout, checkin=checkin) + i_connect = dict(connect=connect) + i_checkout = dict(checkout=checkout) + i_checkin = dict(checkin=checkin) + + for cls in (pool.QueuePool, pool.StaticPool): + counts = [0, 0, 0] + + def assert_listeners(p, total, conn, cout, cin): + for instance in (p, p.recreate()): + eq_(len(instance.dispatch.connect), conn) + eq_(len(instance.dispatch.checkout), cout) + eq_(len(instance.dispatch.checkin), cin) + + p = self._queuepool_fixture() + assert_listeners(p, 0, 0, 0, 0) + + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["connect", "checkout", "checkin"] + ] + ): + p.add_listener(i_all) + assert_listeners(p, 1, 1, 1, 1) + + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["connect"] + ] + ): + p.add_listener(i_connect) + assert_listeners(p, 2, 1, 1, 1) + + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["checkout"] + ] + ): + p.add_listener(i_checkout) + assert_listeners(p, 3, 1, 1, 1) + + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["checkin"] + ] + ): + p.add_listener(i_checkin) + assert_listeners(p, 4, 1, 1, 1) + del p + + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["connect", "checkout", "checkin"] + ] + + [".*The Pool.listeners parameter will be removed"] + ): + p = self._queuepool_fixture(listeners=[i_all]) + assert_listeners(p, 1, 1, 1, 1) + + c = p.connect() + assert counts == [1, 1, 0] + c.close() + assert counts == [1, 1, 1] + + c = p.connect() + assert counts == [1, 2, 1] + with testing.expect_deprecated( + *[ + "PoolListener.%s is deprecated." % name + for name in ["checkin"] + ] + ): + p.add_listener(i_checkin) + c.close() + assert counts == [1, 2, 2] + + +class PoolTest(PoolTestBase): + def test_manager(self): + with testing.expect_deprecated( + r"The pool.manage\(\) function is deprecated," + ): + manager = pool.manage(MockDBAPI(), use_threadlocal=True) + + with testing.expect_deprecated( + r".*Pool.use_threadlocal parameter is deprecated" + ): + c1 = manager.connect("foo.db") + c2 = manager.connect("foo.db") + c3 = manager.connect("bar.db") + c4 = manager.connect("foo.db", bar="bat") + c5 = manager.connect("foo.db", bar="hoho") + c6 = manager.connect("foo.db", bar="bat") + + assert c1.cursor() is not None + assert c1 is c2 + assert c1 is not c3 + assert c4 is c6 + assert c4 is not c5 + + def test_manager_with_key(self): + + dbapi = MockDBAPI() + + with testing.expect_deprecated( + r"The pool.manage\(\) function is deprecated," + ): + manager = pool.manage(dbapi, use_threadlocal=True) + + with testing.expect_deprecated( + r".*Pool.use_threadlocal parameter is deprecated" + ): + c1 = manager.connect("foo.db", sa_pool_key="a") + c2 = manager.connect("foo.db", sa_pool_key="b") + c3 = manager.connect("bar.db", sa_pool_key="a") + + assert c1.cursor() is not None + assert c1 is not c2 + assert c1 is c3 + + eq_(dbapi.connect.mock_calls, [call("foo.db"), call("foo.db")]) + + def test_bad_args(self): + with testing.expect_deprecated( + r"The pool.manage\(\) function is deprecated," + ): + manager = pool.manage(MockDBAPI()) + manager.connect(None) + + def test_non_thread_local_manager(self): + with testing.expect_deprecated( + r"The pool.manage\(\) function is deprecated," + ): + manager = pool.manage(MockDBAPI(), use_threadlocal=False) + + connection = manager.connect("foo.db") + connection2 = manager.connect("foo.db") + + self.assert_(connection.cursor() is not None) + self.assert_(connection is not connection2) + + def test_threadlocal_del(self): + self._do_testthreadlocal(useclose=False) + + def test_threadlocal_close(self): + self._do_testthreadlocal(useclose=True) + + def _do_testthreadlocal(self, useclose=False): + dbapi = MockDBAPI() + + with testing.expect_deprecated( + r".*Pool.use_threadlocal parameter is deprecated" + ): + for p in ( + pool.QueuePool( + creator=dbapi.connect, + pool_size=3, + max_overflow=-1, + use_threadlocal=True, + ), + pool.SingletonThreadPool( + creator=dbapi.connect, use_threadlocal=True + ), + ): + c1 = p.connect() + c2 = p.connect() + self.assert_(c1 is c2) + c3 = p.unique_connection() + self.assert_(c3 is not c1) + if useclose: + c2.close() + else: + c2 = None + c2 = p.connect() + self.assert_(c1 is c2) + self.assert_(c3 is not c1) + if useclose: + c2.close() + else: + c2 = None + lazy_gc() + if useclose: + c1 = p.connect() + c2 = p.connect() + c3 = p.connect() + c3.close() + c2.close() + self.assert_(c1.connection is not None) + c1.close() + c1 = c2 = c3 = None + + # extra tests with QueuePool to ensure connections get + # __del__()ed when dereferenced + + if isinstance(p, pool.QueuePool): + lazy_gc() + self.assert_(p.checkedout() == 0) + c1 = p.connect() + c2 = p.connect() + if useclose: + c2.close() + c1.close() + else: + c2 = None + c1 = None + lazy_gc() + self.assert_(p.checkedout() == 0) + + def test_mixed_close(self): + pool._refs.clear() + with testing.expect_deprecated( + r".*Pool.use_threadlocal parameter is deprecated" + ): + p = self._queuepool_fixture( + pool_size=3, max_overflow=-1, use_threadlocal=True + ) + c1 = p.connect() + c2 = p.connect() + assert c1 is c2 + c1.close() + c2 = None + assert p.checkedout() == 1 + c1 = None + lazy_gc() + assert p.checkedout() == 0 + lazy_gc() + assert not pool._refs + + +class QueuePoolTest(PoolTestBase): + def test_threadfairy(self): + with testing.expect_deprecated( + r".*Pool.use_threadlocal parameter is deprecated" + ): + p = self._queuepool_fixture( + pool_size=3, max_overflow=-1, use_threadlocal=True + ) + c1 = p.connect() + c1.close() + c2 = p.connect() + assert c2.connection is not None + + def test_trick_the_counter(self): + """this is a "flaw" in the connection pool; since threadlocal + uses a single ConnectionFairy per thread with an open/close + counter, you can fool the counter into giving you a + ConnectionFairy with an ambiguous counter. i.e. its not true + reference counting.""" + + with testing.expect_deprecated( + r".*Pool.use_threadlocal parameter is deprecated" + ): + p = self._queuepool_fixture( + pool_size=3, max_overflow=-1, use_threadlocal=True + ) + c1 = p.connect() + c2 = p.connect() + assert c1 is c2 + c1.close() + c2 = p.connect() + c2.close() + self.assert_(p.checkedout() != 0) + c2.close() + self.assert_(p.checkedout() == 0) + + @testing.requires.predictable_gc + def test_weakref_kaboom(self): + with testing.expect_deprecated( + r".*Pool.use_threadlocal parameter is deprecated" + ): + p = self._queuepool_fixture( + pool_size=3, max_overflow=-1, use_threadlocal=True + ) + c1 = p.connect() + c2 = p.connect() + c1.close() + c2 = None + del c1 + del c2 + gc_collect() + assert p.checkedout() == 0 + c3 = p.connect() + assert c3 is not None + + +class ExplicitAutoCommitDeprecatedTest(fixtures.TestBase): + + """test the 'autocommit' flag on select() and text() objects. + + Requires PostgreSQL so that we may define a custom function which + modifies the database. """ + + __only_on__ = "postgresql" + + @classmethod + def setup_class(cls): + global metadata, foo + metadata = MetaData(testing.db) + foo = Table( + "foo", + metadata, + Column("id", Integer, primary_key=True), + Column("data", String(100)), + ) + metadata.create_all() + testing.db.execute( + "create function insert_foo(varchar) " + "returns integer as 'insert into foo(data) " + "values ($1);select 1;' language sql" + ) + + def teardown(self): + foo.delete().execute().close() + + @classmethod + def teardown_class(cls): + testing.db.execute("drop function insert_foo(varchar)") + metadata.drop_all() + + def test_explicit_compiled(self): + conn1 = testing.db.connect() + conn2 = testing.db.connect() + with testing.expect_deprecated( + "The select.autocommit parameter is deprecated" + ): + conn1.execute(select([func.insert_foo("data1")], autocommit=True)) + assert conn2.execute(select([foo.c.data])).fetchall() == [("data1",)] + with testing.expect_deprecated( + r"The SelectBase.autocommit\(\) method is deprecated," + ): + conn1.execute(select([func.insert_foo("data2")]).autocommit()) + assert conn2.execute(select([foo.c.data])).fetchall() == [ + ("data1",), + ("data2",), + ] + conn1.close() + conn2.close() + + def test_explicit_text(self): + conn1 = testing.db.connect() + conn2 = testing.db.connect() + with testing.expect_deprecated( + "The text.autocommit parameter is deprecated" + ): + conn1.execute( + text("select insert_foo('moredata')", autocommit=True) + ) + assert conn2.execute(select([foo.c.data])).fetchall() == [ + ("moredata",) + ] + conn1.close() + conn2.close() diff --git a/test/engine/test_execute.py b/test/engine/test_execute.py index 8613be5bc..061dae005 100644 --- a/test/engine/test_execute.py +++ b/test/engine/test_execute.py @@ -22,7 +22,6 @@ from sqlalchemy import util from sqlalchemy import VARCHAR from sqlalchemy.engine import default from sqlalchemy.engine.base import Engine -from sqlalchemy.interfaces import ConnectionProxy from sqlalchemy.sql import column from sqlalchemy.sql import literal from sqlalchemy.testing import assert_raises @@ -372,8 +371,7 @@ class ExecuteTest(fixtures.TestBase): def _go(conn): assert_raises_message( tsa.exc.StatementError, - r"\(test.engine.test_execute.SomeException\) " - r"nope \[SQL\: u?'SELECT 1 ", + r"\(.*.SomeException\) " r"nope \[SQL\: u?'SELECT 1 ", conn.execute, select([1]).where(column("foo") == literal("bar", MyType())), ) @@ -613,7 +611,7 @@ class ExecuteTest(fixtures.TestBase): eng = engines.testing_engine( options={"execution_options": {"foo": "bar"}} ) - with eng.contextual_connect() as conn: + with eng.connect() as conn: eq_(conn._execution_options["foo"], "bar") eq_( conn.execution_options(bat="hoho")._execution_options["foo"], @@ -628,7 +626,7 @@ class ExecuteTest(fixtures.TestBase): "hoho", ) eng.update_execution_options(foo="hoho") - conn = eng.contextual_connect() + conn = eng.connect() eq_(conn._execution_options["foo"], "hoho") @testing.requires.ad_hoc_engines @@ -787,32 +785,6 @@ class ConvenienceExecuteTest(fixtures.TablesTest): ) self._assert_no_data() - def test_transaction_tlocal_engine_ctx_commit(self): - fn = self._trans_fn() - engine = engines.testing_engine( - options=dict(strategy="threadlocal", pool=testing.db.pool) - ) - ctx = engine.begin() - testing.run_as_contextmanager(ctx, fn, 5, value=8) - self._assert_fn(5, value=8) - - def test_transaction_tlocal_engine_ctx_rollback(self): - fn = self._trans_rollback_fn() - engine = engines.testing_engine( - options=dict(strategy="threadlocal", pool=testing.db.pool) - ) - ctx = engine.begin() - assert_raises_message( - Exception, - "breakage", - testing.run_as_contextmanager, - ctx, - fn, - 5, - value=8, - ) - self._assert_no_data() - def test_transaction_connection_ctx_commit(self): fn = self._trans_fn(True) with testing.db.connect() as conn: @@ -1495,11 +1467,16 @@ class EngineEventsTest(fixtures.TestBase): ): cursor_stmts.append((str(statement), parameters, None)) + with testing.expect_deprecated( + "The 'threadlocal' engine strategy is deprecated" + ): + tl_engine = engines.testing_engine( + options=dict(implicit_returning=False, strategy="threadlocal") + ) + for engine in [ engines.testing_engine(options=dict(implicit_returning=False)), - engines.testing_engine( - options=dict(implicit_returning=False, strategy="threadlocal") - ), + tl_engine, engines.testing_engine( options=dict(implicit_returning=False) ).connect(), @@ -1999,63 +1976,6 @@ class HandleErrorTest(fixtures.TestBase): Engine.dispatch._clear() Engine._has_events = False - def test_legacy_dbapi_error(self): - engine = engines.testing_engine() - canary = Mock() - - event.listen(engine, "dbapi_error", canary) - - with engine.connect() as conn: - try: - conn.execute("SELECT FOO FROM I_DONT_EXIST") - assert False - except tsa.exc.DBAPIError as e: - eq_(canary.mock_calls[0][1][5], e.orig) - eq_(canary.mock_calls[0][1][2], "SELECT FOO FROM I_DONT_EXIST") - - def test_legacy_dbapi_error_no_ad_hoc_context(self): - engine = engines.testing_engine() - - listener = Mock(return_value=None) - event.listen(engine, "dbapi_error", listener) - - nope = SomeException("nope") - - class MyType(TypeDecorator): - impl = Integer - - def process_bind_param(self, value, dialect): - raise nope - - with engine.connect() as conn: - assert_raises_message( - tsa.exc.StatementError, - r"\(test.engine.test_execute.SomeException\) " - r"nope \[SQL\: u?'SELECT 1 ", - conn.execute, - select([1]).where(column("foo") == literal("bar", MyType())), - ) - # no legacy event - eq_(listener.mock_calls, []) - - def test_legacy_dbapi_error_non_dbapi_error(self): - engine = engines.testing_engine() - - listener = Mock(return_value=None) - event.listen(engine, "dbapi_error", listener) - - nope = TypeError("I'm not a DBAPI error") - with engine.connect() as c: - c.connection.cursor = Mock( - return_value=Mock(execute=Mock(side_effect=nope)) - ) - - assert_raises_message( - TypeError, "I'm not a DBAPI error", c.execute, "select " - ) - # no legacy event - eq_(listener.mock_calls, []) - def test_handle_error(self): engine = engines.testing_engine() canary = Mock(return_value=None) @@ -2249,8 +2169,7 @@ class HandleErrorTest(fixtures.TestBase): with engine.connect() as conn: assert_raises_message( tsa.exc.StatementError, - r"\(test.engine.test_execute.SomeException\) " - r"nope \[SQL\: u?'SELECT 1 ", + r"\(.*.SomeException\) " r"nope \[SQL\: u?'SELECT 1 ", conn.execute, select([1]).where(column("foo") == literal("bar", MyType())), ) @@ -2571,27 +2490,15 @@ class HandleInvalidatedOnConnectTest(fixtures.TestBase): except tsa.exc.DBAPIError: assert conn.invalidated - def _test_dont_touch_non_dbapi_exception_on_connect(self, connect_fn): + def test_dont_touch_non_dbapi_exception_on_connect(self): dbapi = self.dbapi dbapi.connect = Mock(side_effect=TypeError("I'm not a DBAPI error")) e = create_engine("sqlite://", module=dbapi) e.dialect.is_disconnect = is_disconnect = Mock() - assert_raises_message( - TypeError, "I'm not a DBAPI error", connect_fn, e - ) + assert_raises_message(TypeError, "I'm not a DBAPI error", e.connect) eq_(is_disconnect.call_count, 0) - def test_dont_touch_non_dbapi_exception_on_connect(self): - self._test_dont_touch_non_dbapi_exception_on_connect( - lambda engine: engine.connect() - ) - - def test_dont_touch_non_dbapi_exception_on_contextual_connect(self): - self._test_dont_touch_non_dbapi_exception_on_connect( - lambda engine: engine.contextual_connect() - ) - def test_ensure_dialect_does_is_disconnect_no_conn(self): """test that is_disconnect() doesn't choke if no connection, cursor given.""" @@ -2601,275 +2508,26 @@ class HandleInvalidatedOnConnectTest(fixtures.TestBase): dbapi.OperationalError("test"), None, None ) - def _test_invalidate_on_connect(self, connect_fn): + def test_invalidate_on_connect(self): """test that is_disconnect() is called during connect. interpretation of connection failures are not supported by every backend. """ - dbapi = self.dbapi dbapi.connect = Mock( side_effect=self.ProgrammingError( "Cannot operate on a closed database." ) ) + e = create_engine("sqlite://", module=dbapi) try: - connect_fn(create_engine("sqlite://", module=dbapi)) + e.connect() assert False except tsa.exc.DBAPIError as de: assert de.connection_invalidated - def test_invalidate_on_connect(self): - """test that is_disconnect() is called during connect. - - interpretation of connection failures are not supported by - every backend. - - """ - self._test_invalidate_on_connect(lambda engine: engine.connect()) - - def test_invalidate_on_contextual_connect(self): - """test that is_disconnect() is called during connect. - - interpretation of connection failures are not supported by - every backend. - - """ - self._test_invalidate_on_connect( - lambda engine: engine.contextual_connect() - ) - - -class ProxyConnectionTest(fixtures.TestBase): - - """These are the same tests as EngineEventsTest, except using - the deprecated ConnectionProxy interface. - - """ - - __requires__ = ("ad_hoc_engines",) - __prefer_requires__ = ("two_phase_transactions",) - - @testing.uses_deprecated(r".*Use event.listen") - @testing.fails_on("firebird", "Data type unknown") - def test_proxy(self): - - stmts = [] - cursor_stmts = [] - - class MyProxy(ConnectionProxy): - def execute( - self, conn, execute, clauseelement, *multiparams, **params - ): - stmts.append((str(clauseelement), params, multiparams)) - return execute(clauseelement, *multiparams, **params) - - def cursor_execute( - self, - execute, - cursor, - statement, - parameters, - context, - executemany, - ): - cursor_stmts.append((str(statement), parameters, None)) - return execute(cursor, statement, parameters, context) - - def assert_stmts(expected, received): - for stmt, params, posn in expected: - if not received: - assert False, "Nothing available for stmt: %s" % stmt - while received: - teststmt, testparams, testmultiparams = received.pop(0) - teststmt = ( - re.compile(r"[\n\t ]+", re.M) - .sub(" ", teststmt) - .strip() - ) - if teststmt.startswith(stmt) and ( - testparams == params or testparams == posn - ): - break - - for engine in ( - engines.testing_engine( - options=dict(implicit_returning=False, proxy=MyProxy()) - ), - engines.testing_engine( - options=dict( - implicit_returning=False, - proxy=MyProxy(), - strategy="threadlocal", - ) - ), - ): - m = MetaData(engine) - t1 = Table( - "t1", - m, - Column("c1", Integer, primary_key=True), - Column( - "c2", - String(50), - default=func.lower("Foo"), - primary_key=True, - ), - ) - m.create_all() - try: - t1.insert().execute(c1=5, c2="some data") - t1.insert().execute(c1=6) - eq_( - engine.execute("select * from t1").fetchall(), - [(5, "some data"), (6, "foo")], - ) - finally: - m.drop_all() - engine.dispose() - compiled = [ - ("CREATE TABLE t1", {}, None), - ( - "INSERT INTO t1 (c1, c2)", - {"c2": "some data", "c1": 5}, - None, - ), - ("INSERT INTO t1 (c1, c2)", {"c1": 6}, None), - ("select * from t1", {}, None), - ("DROP TABLE t1", {}, None), - ] - - cursor = [ - ("CREATE TABLE t1", {}, ()), - ( - "INSERT INTO t1 (c1, c2)", - {"c2": "some data", "c1": 5}, - (5, "some data"), - ), - ("SELECT lower", {"lower_1": "Foo"}, ("Foo",)), - ( - "INSERT INTO t1 (c1, c2)", - {"c2": "foo", "c1": 6}, - (6, "foo"), - ), - ("select * from t1", {}, ()), - ("DROP TABLE t1", {}, ()), - ] - - assert_stmts(compiled, stmts) - assert_stmts(cursor, cursor_stmts) - - @testing.uses_deprecated(r".*Use event.listen") - def test_options(self): - canary = [] - - class TrackProxy(ConnectionProxy): - def __getattribute__(self, key): - fn = object.__getattribute__(self, key) - - def go(*arg, **kw): - canary.append(fn.__name__) - return fn(*arg, **kw) - - return go - - engine = engines.testing_engine(options={"proxy": TrackProxy()}) - conn = engine.connect() - c2 = conn.execution_options(foo="bar") - eq_(c2._execution_options, {"foo": "bar"}) - c2.execute(select([1])) - c3 = c2.execution_options(bar="bat") - eq_(c3._execution_options, {"foo": "bar", "bar": "bat"}) - eq_(canary, ["execute", "cursor_execute"]) - - @testing.uses_deprecated(r".*Use event.listen") - def test_transactional(self): - canary = [] - - class TrackProxy(ConnectionProxy): - def __getattribute__(self, key): - fn = object.__getattribute__(self, key) - - def go(*arg, **kw): - canary.append(fn.__name__) - return fn(*arg, **kw) - - return go - - engine = engines.testing_engine(options={"proxy": TrackProxy()}) - conn = engine.connect() - trans = conn.begin() - conn.execute(select([1])) - trans.rollback() - trans = conn.begin() - conn.execute(select([1])) - trans.commit() - - eq_( - canary, - [ - "begin", - "execute", - "cursor_execute", - "rollback", - "begin", - "execute", - "cursor_execute", - "commit", - ], - ) - - @testing.uses_deprecated(r".*Use event.listen") - @testing.requires.savepoints - @testing.requires.two_phase_transactions - def test_transactional_advanced(self): - canary = [] - - class TrackProxy(ConnectionProxy): - def __getattribute__(self, key): - fn = object.__getattribute__(self, key) - - def go(*arg, **kw): - canary.append(fn.__name__) - return fn(*arg, **kw) - - return go - - engine = engines.testing_engine(options={"proxy": TrackProxy()}) - conn = engine.connect() - - trans = conn.begin() - trans2 = conn.begin_nested() - conn.execute(select([1])) - trans2.rollback() - trans2 = conn.begin_nested() - conn.execute(select([1])) - trans2.commit() - trans.rollback() - - trans = conn.begin_twophase() - conn.execute(select([1])) - trans.prepare() - trans.commit() - - canary = [t for t in canary if t not in ("cursor_execute", "execute")] - eq_( - canary, - [ - "begin", - "savepoint", - "rollback_savepoint", - "savepoint", - "release_savepoint", - "rollback", - "begin_twophase", - "prepare_twophase", - "commit_twophase", - ], - ) - class DialectEventTest(fixtures.TestBase): @contextmanager diff --git a/test/engine/test_parseconnect.py b/test/engine/test_parseconnect.py index 7a8918817..be90378c9 100644 --- a/test/engine/test_parseconnect.py +++ b/test/engine/test_parseconnect.py @@ -209,25 +209,6 @@ class CreateEngineTest(fixtures.TestBase): ) assert e.echo is True - def test_pool_threadlocal_from_config(self): - dbapi = mock_dbapi - - config = { - "sqlalchemy.url": "postgresql://scott:tiger@somehost/test", - "sqlalchemy.pool_threadlocal": "false", - } - - e = engine_from_config(config, module=dbapi, _initialize=False) - eq_(e.pool._use_threadlocal, False) - - config = { - "sqlalchemy.url": "postgresql://scott:tiger@somehost/test", - "sqlalchemy.pool_threadlocal": "true", - } - - e = engine_from_config(config, module=dbapi, _initialize=False) - eq_(e.pool._use_threadlocal, True) - def test_pool_reset_on_return_from_config(self): dbapi = mock_dbapi diff --git a/test/engine/test_pool.py b/test/engine/test_pool.py index 75caa233a..feff61b88 100644 --- a/test/engine/test_pool.py +++ b/test/engine/test_pool.py @@ -90,50 +90,6 @@ class PoolTestBase(fixtures.TestBase): class PoolTest(PoolTestBase): - def test_manager(self): - manager = pool.manage(MockDBAPI(), use_threadlocal=True) - - c1 = manager.connect("foo.db") - c2 = manager.connect("foo.db") - c3 = manager.connect("bar.db") - c4 = manager.connect("foo.db", bar="bat") - c5 = manager.connect("foo.db", bar="hoho") - c6 = manager.connect("foo.db", bar="bat") - - assert c1.cursor() is not None - assert c1 is c2 - assert c1 is not c3 - assert c4 is c6 - assert c4 is not c5 - - def test_manager_with_key(self): - - dbapi = MockDBAPI() - manager = pool.manage(dbapi, use_threadlocal=True) - - c1 = manager.connect("foo.db", sa_pool_key="a") - c2 = manager.connect("foo.db", sa_pool_key="b") - c3 = manager.connect("bar.db", sa_pool_key="a") - - assert c1.cursor() is not None - assert c1 is not c2 - assert c1 is c3 - - eq_(dbapi.connect.mock_calls, [call("foo.db"), call("foo.db")]) - - def test_bad_args(self): - manager = pool.manage(MockDBAPI()) - manager.connect(None) - - def test_non_thread_local_manager(self): - manager = pool.manage(MockDBAPI(), use_threadlocal=False) - - connection = manager.connect("foo.db") - connection2 = manager.connect("foo.db") - - self.assert_(connection.cursor() is not None) - self.assert_(connection is not connection2) - @testing.fails_on( "+pyodbc", "pyodbc cursor doesn't implement tuple __eq__" ) @@ -170,69 +126,6 @@ class PoolTest(PoolTestBase): p.dispose() p.recreate() - def test_threadlocal_del(self): - self._do_testthreadlocal(useclose=False) - - def test_threadlocal_close(self): - self._do_testthreadlocal(useclose=True) - - def _do_testthreadlocal(self, useclose=False): - dbapi = MockDBAPI() - for p in ( - pool.QueuePool( - creator=dbapi.connect, - pool_size=3, - max_overflow=-1, - use_threadlocal=True, - ), - pool.SingletonThreadPool( - creator=dbapi.connect, use_threadlocal=True - ), - ): - c1 = p.connect() - c2 = p.connect() - self.assert_(c1 is c2) - c3 = p.unique_connection() - self.assert_(c3 is not c1) - if useclose: - c2.close() - else: - c2 = None - c2 = p.connect() - self.assert_(c1 is c2) - self.assert_(c3 is not c1) - if useclose: - c2.close() - else: - c2 = None - lazy_gc() - if useclose: - c1 = p.connect() - c2 = p.connect() - c3 = p.connect() - c3.close() - c2.close() - self.assert_(c1.connection is not None) - c1.close() - c1 = c2 = c3 = None - - # extra tests with QueuePool to ensure connections get - # __del__()ed when dereferenced - - if isinstance(p, pool.QueuePool): - lazy_gc() - self.assert_(p.checkedout() == 0) - c1 = p.connect() - c2 = p.connect() - if useclose: - c2.close() - c1.close() - else: - c2 = None - c1 = None - lazy_gc() - self.assert_(p.checkedout() == 0) - def test_info(self): p = self._queuepool_fixture(pool_size=1, max_overflow=0) @@ -822,255 +715,6 @@ class PoolFirstConnectSyncTest(PoolTestBase): ) -class DeprecatedPoolListenerTest(PoolTestBase): - @testing.requires.predictable_gc - @testing.uses_deprecated( - r".*Use the PoolEvents", - r".*'listeners' argument .* is deprecated" - ) - def test_listeners(self): - class InstrumentingListener(object): - def __init__(self): - if hasattr(self, "connect"): - self.connect = self.inst_connect - if hasattr(self, "first_connect"): - self.first_connect = self.inst_first_connect - if hasattr(self, "checkout"): - self.checkout = self.inst_checkout - if hasattr(self, "checkin"): - self.checkin = self.inst_checkin - self.clear() - - def clear(self): - self.connected = [] - self.first_connected = [] - self.checked_out = [] - self.checked_in = [] - - def assert_total(self, conn, fconn, cout, cin): - eq_(len(self.connected), conn) - eq_(len(self.first_connected), fconn) - eq_(len(self.checked_out), cout) - eq_(len(self.checked_in), cin) - - def assert_in(self, item, in_conn, in_fconn, in_cout, in_cin): - eq_((item in self.connected), in_conn) - eq_((item in self.first_connected), in_fconn) - eq_((item in self.checked_out), in_cout) - eq_((item in self.checked_in), in_cin) - - def inst_connect(self, con, record): - print("connect(%s, %s)" % (con, record)) - assert con is not None - assert record is not None - self.connected.append(con) - - def inst_first_connect(self, con, record): - print("first_connect(%s, %s)" % (con, record)) - assert con is not None - assert record is not None - self.first_connected.append(con) - - def inst_checkout(self, con, record, proxy): - print("checkout(%s, %s, %s)" % (con, record, proxy)) - assert con is not None - assert record is not None - assert proxy is not None - self.checked_out.append(con) - - def inst_checkin(self, con, record): - print("checkin(%s, %s)" % (con, record)) - # con can be None if invalidated - assert record is not None - self.checked_in.append(con) - - class ListenAll(tsa.interfaces.PoolListener, InstrumentingListener): - pass - - class ListenConnect(InstrumentingListener): - def connect(self, con, record): - pass - - class ListenFirstConnect(InstrumentingListener): - def first_connect(self, con, record): - pass - - class ListenCheckOut(InstrumentingListener): - def checkout(self, con, record, proxy, num): - pass - - class ListenCheckIn(InstrumentingListener): - def checkin(self, con, record): - pass - - def assert_listeners(p, total, conn, fconn, cout, cin): - for instance in (p, p.recreate()): - self.assert_(len(instance.dispatch.connect) == conn) - self.assert_(len(instance.dispatch.first_connect) == fconn) - self.assert_(len(instance.dispatch.checkout) == cout) - self.assert_(len(instance.dispatch.checkin) == cin) - - p = self._queuepool_fixture() - assert_listeners(p, 0, 0, 0, 0, 0) - - p.add_listener(ListenAll()) - assert_listeners(p, 1, 1, 1, 1, 1) - - p.add_listener(ListenConnect()) - assert_listeners(p, 2, 2, 1, 1, 1) - - p.add_listener(ListenFirstConnect()) - assert_listeners(p, 3, 2, 2, 1, 1) - - p.add_listener(ListenCheckOut()) - assert_listeners(p, 4, 2, 2, 2, 1) - - p.add_listener(ListenCheckIn()) - assert_listeners(p, 5, 2, 2, 2, 2) - del p - - snoop = ListenAll() - p = self._queuepool_fixture(listeners=[snoop]) - assert_listeners(p, 1, 1, 1, 1, 1) - - c = p.connect() - snoop.assert_total(1, 1, 1, 0) - cc = c.connection - snoop.assert_in(cc, True, True, True, False) - c.close() - snoop.assert_in(cc, True, True, True, True) - del c, cc - - snoop.clear() - - # this one depends on immediate gc - c = p.connect() - cc = c.connection - snoop.assert_in(cc, False, False, True, False) - snoop.assert_total(0, 0, 1, 0) - del c, cc - lazy_gc() - snoop.assert_total(0, 0, 1, 1) - - p.dispose() - snoop.clear() - - c = p.connect() - c.close() - c = p.connect() - snoop.assert_total(1, 0, 2, 1) - c.close() - snoop.assert_total(1, 0, 2, 2) - - # invalidation - p.dispose() - snoop.clear() - - c = p.connect() - snoop.assert_total(1, 0, 1, 0) - c.invalidate() - snoop.assert_total(1, 0, 1, 1) - c.close() - snoop.assert_total(1, 0, 1, 1) - del c - lazy_gc() - snoop.assert_total(1, 0, 1, 1) - c = p.connect() - snoop.assert_total(2, 0, 2, 1) - c.close() - del c - lazy_gc() - snoop.assert_total(2, 0, 2, 2) - - # detached - p.dispose() - snoop.clear() - - c = p.connect() - snoop.assert_total(1, 0, 1, 0) - c.detach() - snoop.assert_total(1, 0, 1, 0) - c.close() - del c - snoop.assert_total(1, 0, 1, 0) - c = p.connect() - snoop.assert_total(2, 0, 2, 0) - c.close() - del c - snoop.assert_total(2, 0, 2, 1) - - # recreated - p = p.recreate() - snoop.clear() - - c = p.connect() - snoop.assert_total(1, 1, 1, 0) - c.close() - snoop.assert_total(1, 1, 1, 1) - c = p.connect() - snoop.assert_total(1, 1, 2, 1) - c.close() - snoop.assert_total(1, 1, 2, 2) - - @testing.uses_deprecated( - r".*Use the PoolEvents", - r".*'listeners' argument .* is deprecated" - ) - def test_listeners_callables(self): - def connect(dbapi_con, con_record): - counts[0] += 1 - - def checkout(dbapi_con, con_record, con_proxy): - counts[1] += 1 - - def checkin(dbapi_con, con_record): - counts[2] += 1 - - i_all = dict(connect=connect, checkout=checkout, checkin=checkin) - i_connect = dict(connect=connect) - i_checkout = dict(checkout=checkout) - i_checkin = dict(checkin=checkin) - - for cls in (pool.QueuePool, pool.StaticPool): - counts = [0, 0, 0] - - def assert_listeners(p, total, conn, cout, cin): - for instance in (p, p.recreate()): - eq_(len(instance.dispatch.connect), conn) - eq_(len(instance.dispatch.checkout), cout) - eq_(len(instance.dispatch.checkin), cin) - - p = self._queuepool_fixture() - assert_listeners(p, 0, 0, 0, 0) - - p.add_listener(i_all) - assert_listeners(p, 1, 1, 1, 1) - - p.add_listener(i_connect) - assert_listeners(p, 2, 1, 1, 1) - - p.add_listener(i_checkout) - assert_listeners(p, 3, 1, 1, 1) - - p.add_listener(i_checkin) - assert_listeners(p, 4, 1, 1, 1) - del p - - p = self._queuepool_fixture(listeners=[i_all]) - assert_listeners(p, 1, 1, 1, 1) - - c = p.connect() - assert counts == [1, 1, 0] - c.close() - assert counts == [1, 1, 1] - - c = p.connect() - assert counts == [1, 2, 1] - p.add_listener(i_checkin) - c.close() - assert counts == [1, 2, 2] - - class QueuePoolTest(PoolTestBase): def test_queuepool_del(self): self._do_testqueuepool(useclose=False) @@ -1491,30 +1135,7 @@ class QueuePoolTest(PoolTestBase): def test_max_overflow(self): self._test_overflow(40, 5) - def test_mixed_close(self): - pool._refs.clear() - p = self._queuepool_fixture( - pool_size=3, max_overflow=-1, use_threadlocal=True - ) - c1 = p.connect() - c2 = p.connect() - assert c1 is c2 - c1.close() - c2 = None - assert p.checkedout() == 1 - c1 = None - lazy_gc() - assert p.checkedout() == 0 - lazy_gc() - assert not pool._refs - - def test_overflow_no_gc_tlocal(self): - self._test_overflow_no_gc(True) - def test_overflow_no_gc(self): - self._test_overflow_no_gc(False) - - def _test_overflow_no_gc(self, threadlocal): p = self._queuepool_fixture(pool_size=2, max_overflow=2) # disable weakref collection of the @@ -1543,42 +1164,6 @@ class QueuePoolTest(PoolTestBase): set([1, 1, 1, 1, 1, 1, 1, 0, 1, 1, 1, 0]), ) - @testing.requires.predictable_gc - def test_weakref_kaboom(self): - p = self._queuepool_fixture( - pool_size=3, max_overflow=-1, use_threadlocal=True - ) - c1 = p.connect() - c2 = p.connect() - c1.close() - c2 = None - del c1 - del c2 - gc_collect() - assert p.checkedout() == 0 - c3 = p.connect() - assert c3 is not None - - def test_trick_the_counter(self): - """this is a "flaw" in the connection pool; since threadlocal - uses a single ConnectionFairy per thread with an open/close - counter, you can fool the counter into giving you a - ConnectionFairy with an ambiguous counter. i.e. its not true - reference counting.""" - - p = self._queuepool_fixture( - pool_size=3, max_overflow=-1, use_threadlocal=True - ) - c1 = p.connect() - c2 = p.connect() - assert c1 is c2 - c1.close() - c2 = p.connect() - c2.close() - self.assert_(p.checkedout() != 0) - c2.close() - self.assert_(p.checkedout() == 0) - def test_recycle(self): with patch("sqlalchemy.pool.base.time.time") as mock: mock.return_value = 10000 @@ -1957,15 +1542,6 @@ class QueuePoolTest(PoolTestBase): c2.close() eq_(c2_con.close.call_count, 0) - def test_threadfairy(self): - p = self._queuepool_fixture( - pool_size=3, max_overflow=-1, use_threadlocal=True - ) - c1 = p.connect() - c1.close() - c2 = p.connect() - assert c2.connection is not None - def test_no_double_checkin(self): p = self._queuepool_fixture(pool_size=1) diff --git a/test/engine/test_reconnect.py b/test/engine/test_reconnect.py index f6904174b..dd2ebb1c4 100644 --- a/test/engine/test_reconnect.py +++ b/test/engine/test_reconnect.py @@ -950,33 +950,30 @@ class RecycleTest(fixtures.TestBase): __backend__ = True def test_basic(self): - for threadlocal in False, True: - engine = engines.reconnecting_engine( - options={"pool_threadlocal": threadlocal} - ) + engine = engines.reconnecting_engine() - conn = engine.contextual_connect() - eq_(conn.execute(select([1])).scalar(), 1) - conn.close() + conn = engine.connect() + eq_(conn.execute(select([1])).scalar(), 1) + conn.close() - # set the pool recycle down to 1. - # we aren't doing this inline with the - # engine create since cx_oracle takes way - # too long to create the 1st connection and don't - # want to build a huge delay into this test. + # set the pool recycle down to 1. + # we aren't doing this inline with the + # engine create since cx_oracle takes way + # too long to create the 1st connection and don't + # want to build a huge delay into this test. - engine.pool._recycle = 1 + engine.pool._recycle = 1 - # kill the DB connection - engine.test_shutdown() + # kill the DB connection + engine.test_shutdown() - # wait until past the recycle period - time.sleep(2) + # wait until past the recycle period + time.sleep(2) - # can connect, no exception - conn = engine.contextual_connect() - eq_(conn.execute(select([1])).scalar(), 1) - conn.close() + # can connect, no exception + conn = engine.connect() + eq_(conn.execute(select([1])).scalar(), 1) + conn.close() class PrePingRealTest(fixtures.TestBase): diff --git a/test/engine/test_transaction.py b/test/engine/test_transaction.py index d8161a29a..81f86089b 100644 --- a/test/engine/test_transaction.py +++ b/test/engine/test_transaction.py @@ -8,7 +8,6 @@ from sqlalchemy import INT from sqlalchemy import Integer from sqlalchemy import MetaData from sqlalchemy import select -from sqlalchemy import Sequence from sqlalchemy import String from sqlalchemy import testing from sqlalchemy import text @@ -822,34 +821,6 @@ class ExplicitAutoCommitTest(fixtures.TestBase): conn1.close() conn2.close() - @testing.uses_deprecated( - r".*select.autocommit parameter is deprecated", - r".*SelectBase.autocommit\(\) .* is deprecated", - ) - def test_explicit_compiled_deprecated(self): - conn1 = testing.db.connect() - conn2 = testing.db.connect() - conn1.execute(select([func.insert_foo("data1")], autocommit=True)) - assert conn2.execute(select([foo.c.data])).fetchall() == [("data1",)] - conn1.execute(select([func.insert_foo("data2")]).autocommit()) - assert conn2.execute(select([foo.c.data])).fetchall() == [ - ("data1",), - ("data2",), - ] - conn1.close() - conn2.close() - - @testing.uses_deprecated(r"autocommit on text\(\) is deprecated") - def test_explicit_text_deprecated(self): - conn1 = testing.db.connect() - conn2 = testing.db.connect() - conn1.execute(text("select insert_foo('moredata')", autocommit=True)) - assert conn2.execute(select([foo.c.data])).fetchall() == [ - ("moredata",) - ] - conn1.close() - conn2.close() - def test_implicit_text(self): conn1 = testing.db.connect() conn2 = testing.db.connect() @@ -861,485 +832,6 @@ class ExplicitAutoCommitTest(fixtures.TestBase): conn2.close() -tlengine = None - - -class TLTransactionTest(fixtures.TestBase): - __requires__ = ("ad_hoc_engines",) - __backend__ = True - - @classmethod - def setup_class(cls): - global users, metadata, tlengine - tlengine = testing_engine(options=dict(strategy="threadlocal")) - metadata = MetaData() - users = Table( - "query_users", - metadata, - Column( - "user_id", - INT, - Sequence("query_users_id_seq", optional=True), - primary_key=True, - ), - Column("user_name", VARCHAR(20)), - test_needs_acid=True, - ) - metadata.create_all(tlengine) - - def teardown(self): - tlengine.execute(users.delete()).close() - - @classmethod - def teardown_class(cls): - tlengine.close() - metadata.drop_all(tlengine) - tlengine.dispose() - - def setup(self): - - # ensure tests start with engine closed - - tlengine.close() - - @testing.crashes( - "oracle", "TNS error of unknown origin occurs on the buildbot." - ) - def test_rollback_no_trans(self): - tlengine = testing_engine(options=dict(strategy="threadlocal")) - - # shouldn't fail - tlengine.rollback() - - tlengine.begin() - tlengine.rollback() - - # shouldn't fail - tlengine.rollback() - - def test_commit_no_trans(self): - tlengine = testing_engine(options=dict(strategy="threadlocal")) - - # shouldn't fail - tlengine.commit() - - tlengine.begin() - tlengine.rollback() - - # shouldn't fail - tlengine.commit() - - def test_prepare_no_trans(self): - tlengine = testing_engine(options=dict(strategy="threadlocal")) - - # shouldn't fail - tlengine.prepare() - - tlengine.begin() - tlengine.rollback() - - # shouldn't fail - tlengine.prepare() - - def test_connection_close(self): - """test that when connections are closed for real, transactions - are rolled back and disposed.""" - - c = tlengine.contextual_connect() - c.begin() - assert c.in_transaction() - c.close() - assert not c.in_transaction() - - def test_transaction_close(self): - c = tlengine.contextual_connect() - t = c.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.execute(users.insert(), user_id=2, user_name="user2") - t2 = c.begin() - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.execute(users.insert(), user_id=4, user_name="user4") - t2.close() - result = c.execute("select * from query_users") - assert len(result.fetchall()) == 4 - t.close() - external_connection = tlengine.connect() - result = external_connection.execute("select * from query_users") - try: - assert len(result.fetchall()) == 0 - finally: - c.close() - external_connection.close() - - def test_rollback(self): - """test a basic rollback""" - - tlengine.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.rollback() - external_connection = tlengine.connect() - result = external_connection.execute("select * from query_users") - try: - assert len(result.fetchall()) == 0 - finally: - external_connection.close() - - def test_commit(self): - """test a basic commit""" - - tlengine.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.commit() - external_connection = tlengine.connect() - result = external_connection.execute("select * from query_users") - try: - assert len(result.fetchall()) == 3 - finally: - external_connection.close() - - def test_with_interface(self): - trans = tlengine.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.execute(users.insert(), user_id=2, user_name="user2") - trans.commit() - - trans = tlengine.begin() - tlengine.execute(users.insert(), user_id=3, user_name="user3") - trans.__exit__(Exception, "fake", None) - trans = tlengine.begin() - tlengine.execute(users.insert(), user_id=4, user_name="user4") - trans.__exit__(None, None, None) - eq_( - tlengine.execute( - users.select().order_by(users.c.user_id) - ).fetchall(), - [(1, "user1"), (2, "user2"), (4, "user4")], - ) - - def test_commits(self): - connection = tlengine.connect() - assert ( - connection.execute("select count(*) from query_users").scalar() - == 0 - ) - connection.close() - connection = tlengine.contextual_connect() - transaction = connection.begin() - connection.execute(users.insert(), user_id=1, user_name="user1") - transaction.commit() - transaction = connection.begin() - connection.execute(users.insert(), user_id=2, user_name="user2") - connection.execute(users.insert(), user_id=3, user_name="user3") - transaction.commit() - transaction = connection.begin() - result = connection.execute("select * from query_users") - rows = result.fetchall() - assert len(rows) == 3, "expected 3 got %d" % len(rows) - transaction.commit() - connection.close() - - def test_rollback_off_conn(self): - - # test that a TLTransaction opened off a TLConnection allows - # that TLConnection to be aware of the transactional context - - conn = tlengine.contextual_connect() - trans = conn.begin() - conn.execute(users.insert(), user_id=1, user_name="user1") - conn.execute(users.insert(), user_id=2, user_name="user2") - conn.execute(users.insert(), user_id=3, user_name="user3") - trans.rollback() - external_connection = tlengine.connect() - result = external_connection.execute("select * from query_users") - try: - assert len(result.fetchall()) == 0 - finally: - conn.close() - external_connection.close() - - def test_morerollback_off_conn(self): - - # test that an existing TLConnection automatically takes place - # in a TLTransaction opened on a second TLConnection - - conn = tlengine.contextual_connect() - conn2 = tlengine.contextual_connect() - trans = conn2.begin() - conn.execute(users.insert(), user_id=1, user_name="user1") - conn.execute(users.insert(), user_id=2, user_name="user2") - conn.execute(users.insert(), user_id=3, user_name="user3") - trans.rollback() - external_connection = tlengine.connect() - result = external_connection.execute("select * from query_users") - try: - assert len(result.fetchall()) == 0 - finally: - conn.close() - conn2.close() - external_connection.close() - - def test_commit_off_connection(self): - conn = tlengine.contextual_connect() - trans = conn.begin() - conn.execute(users.insert(), user_id=1, user_name="user1") - conn.execute(users.insert(), user_id=2, user_name="user2") - conn.execute(users.insert(), user_id=3, user_name="user3") - trans.commit() - external_connection = tlengine.connect() - result = external_connection.execute("select * from query_users") - try: - assert len(result.fetchall()) == 3 - finally: - conn.close() - external_connection.close() - - def test_nesting_rollback(self): - """tests nesting of transactions, rollback at the end""" - - external_connection = tlengine.connect() - self.assert_( - external_connection.connection - is not tlengine.contextual_connect().connection - ) - tlengine.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.begin() - tlengine.execute(users.insert(), user_id=4, user_name="user4") - tlengine.execute(users.insert(), user_id=5, user_name="user5") - tlengine.commit() - tlengine.rollback() - try: - self.assert_( - external_connection.scalar("select count(*) from query_users") - == 0 - ) - finally: - external_connection.close() - - def test_nesting_commit(self): - """tests nesting of transactions, commit at the end.""" - - external_connection = tlengine.connect() - self.assert_( - external_connection.connection - is not tlengine.contextual_connect().connection - ) - tlengine.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.begin() - tlengine.execute(users.insert(), user_id=4, user_name="user4") - tlengine.execute(users.insert(), user_id=5, user_name="user5") - tlengine.commit() - tlengine.commit() - try: - self.assert_( - external_connection.scalar("select count(*) from query_users") - == 5 - ) - finally: - external_connection.close() - - def test_mixed_nesting(self): - """tests nesting of transactions off the TLEngine directly - inside of transactions off the connection from the TLEngine""" - - external_connection = tlengine.connect() - self.assert_( - external_connection.connection - is not tlengine.contextual_connect().connection - ) - conn = tlengine.contextual_connect() - trans = conn.begin() - trans2 = conn.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.begin() - tlengine.execute(users.insert(), user_id=4, user_name="user4") - tlengine.begin() - tlengine.execute(users.insert(), user_id=5, user_name="user5") - tlengine.execute(users.insert(), user_id=6, user_name="user6") - tlengine.execute(users.insert(), user_id=7, user_name="user7") - tlengine.commit() - tlengine.execute(users.insert(), user_id=8, user_name="user8") - tlengine.commit() - trans2.commit() - trans.rollback() - conn.close() - try: - self.assert_( - external_connection.scalar("select count(*) from query_users") - == 0 - ) - finally: - external_connection.close() - - def test_more_mixed_nesting(self): - """tests nesting of transactions off the connection from the - TLEngine inside of transactions off the TLEngine directly.""" - - external_connection = tlengine.connect() - self.assert_( - external_connection.connection - is not tlengine.contextual_connect().connection - ) - tlengine.begin() - connection = tlengine.contextual_connect() - connection.execute(users.insert(), user_id=1, user_name="user1") - tlengine.begin() - connection.execute(users.insert(), user_id=2, user_name="user2") - connection.execute(users.insert(), user_id=3, user_name="user3") - trans = connection.begin() - connection.execute(users.insert(), user_id=4, user_name="user4") - connection.execute(users.insert(), user_id=5, user_name="user5") - trans.commit() - tlengine.commit() - tlengine.rollback() - connection.close() - try: - self.assert_( - external_connection.scalar("select count(*) from query_users") - == 0 - ) - finally: - external_connection.close() - - @testing.requires.savepoints - def test_nested_subtransaction_rollback(self): - tlengine.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.begin_nested() - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.rollback() - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.commit() - tlengine.close() - eq_( - tlengine.execute( - select([users.c.user_id]).order_by(users.c.user_id) - ).fetchall(), - [(1,), (3,)], - ) - tlengine.close() - - @testing.requires.savepoints - @testing.crashes( - "oracle+zxjdbc", - "Errors out and causes subsequent tests to " "deadlock", - ) - def test_nested_subtransaction_commit(self): - tlengine.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.begin_nested() - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.commit() - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.commit() - tlengine.close() - eq_( - tlengine.execute( - select([users.c.user_id]).order_by(users.c.user_id) - ).fetchall(), - [(1,), (2,), (3,)], - ) - tlengine.close() - - @testing.requires.savepoints - def test_rollback_to_subtransaction(self): - tlengine.begin() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.begin_nested() - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.begin() - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.rollback() - tlengine.rollback() - tlengine.execute(users.insert(), user_id=4, user_name="user4") - tlengine.commit() - tlengine.close() - eq_( - tlengine.execute( - select([users.c.user_id]).order_by(users.c.user_id) - ).fetchall(), - [(1,), (4,)], - ) - tlengine.close() - - def test_connections(self): - """tests that contextual_connect is threadlocal""" - - c1 = tlengine.contextual_connect() - c2 = tlengine.contextual_connect() - assert c1.connection is c2.connection - c2.close() - assert not c1.closed - assert not tlengine.closed - - @testing.requires.independent_cursors - def test_result_closing(self): - """tests that contextual_connect is threadlocal""" - - r1 = tlengine.execute(select([1])) - r2 = tlengine.execute(select([1])) - row1 = r1.fetchone() - row2 = r2.fetchone() - r1.close() - assert r2.connection is r1.connection - assert not r2.connection.closed - assert not tlengine.closed - - # close again, nothing happens since resultproxy calls close() - # only once - - r1.close() - assert r2.connection is r1.connection - assert not r2.connection.closed - assert not tlengine.closed - r2.close() - assert r2.connection.closed - assert tlengine.closed - - @testing.crashes( - "oracle+cx_oracle", "intermittent failures on the buildbot" - ) - def test_dispose(self): - eng = testing_engine(options=dict(strategy="threadlocal")) - result = eng.execute(select([1])) - eng.dispose() - eng.execute(select([1])) - - @testing.requires.two_phase_transactions - def test_two_phase_transaction(self): - tlengine.begin_twophase() - tlengine.execute(users.insert(), user_id=1, user_name="user1") - tlengine.prepare() - tlengine.commit() - tlengine.begin_twophase() - tlengine.execute(users.insert(), user_id=2, user_name="user2") - tlengine.commit() - tlengine.begin_twophase() - tlengine.execute(users.insert(), user_id=3, user_name="user3") - tlengine.rollback() - tlengine.begin_twophase() - tlengine.execute(users.insert(), user_id=4, user_name="user4") - tlengine.prepare() - tlengine.rollback() - eq_( - tlengine.execute( - select([users.c.user_id]).order_by(users.c.user_id) - ).fetchall(), - [(1,), (2,)], - ) - - class IsolationLevelTest(fixtures.TestBase): __requires__ = ("isolation_level", "ad_hoc_engines") __backend__ = True diff --git a/test/ext/declarative/test_basic.py b/test/ext/declarative/test_basic.py index 0f0035019..4406925ff 100644 --- a/test/ext/declarative/test_basic.py +++ b/test/ext/declarative/test_basic.py @@ -1973,44 +1973,6 @@ class DeclarativeTest(DeclarativeTestBase): rt = sess.query(User).filter(User.namesyn == "someuser").one() eq_(rt, u1) - def test_comparable_using(self): - class NameComparator(sa.orm.PropComparator): - @property - def upperself(self): - cls = self.prop.parent.class_ - col = getattr(cls, "name") - return sa.func.upper(col) - - def operate(self, op, other, **kw): - return op(self.upperself, other, **kw) - - class User(Base, fixtures.ComparableEntity): - - __tablename__ = "users" - id = Column( - "id", Integer, primary_key=True, test_needs_autoincrement=True - ) - name = Column("name", String(50)) - - @decl.comparable_using(NameComparator) - @property - def uc_name(self): - return self.name is not None and self.name.upper() or None - - Base.metadata.create_all() - sess = create_session() - u1 = User(name="someuser") - eq_(u1.name, "someuser", u1.name) - eq_(u1.uc_name, "SOMEUSER", u1.uc_name) - sess.add(u1) - sess.flush() - sess.expunge_all() - rt = sess.query(User).filter(User.uc_name == "SOMEUSER").one() - eq_(rt, u1) - sess.expunge_all() - rt = sess.query(User).filter(User.uc_name.startswith("SOMEUSE")).one() - eq_(rt, u1) - def test_duplicate_classes_in_base(self): class Test(Base): __tablename__ = "a" diff --git a/test/ext/test_associationproxy.py b/test/ext/test_associationproxy.py index 75b6b8901..1bf84a77c 100644 --- a/test/ext/test_associationproxy.py +++ b/test/ext/test_associationproxy.py @@ -1325,7 +1325,7 @@ class ReconstitutionTest(fixtures.TestBase): Parent, self.parents, properties=dict(children=relationship(Child)) ) mapper(Child, self.children) - session = create_session(weak_identity_map=True) + session = create_session() def add_child(parent_name, child_name): parent = session.query(Parent).filter_by(name=parent_name).one() diff --git a/test/ext/test_horizontal_shard.py b/test/ext/test_horizontal_shard.py index 8d590c496..99c0a3c1a 100644 --- a/test/ext/test_horizontal_shard.py +++ b/test/ext/test_horizontal_shard.py @@ -23,6 +23,7 @@ from sqlalchemy.orm import relationship from sqlalchemy.orm import selectinload from sqlalchemy.orm import Session from sqlalchemy.orm import sessionmaker +from sqlalchemy.pool import SingletonThreadPool from sqlalchemy.sql import operators from sqlalchemy.testing import eq_ from sqlalchemy.testing import fixtures @@ -50,8 +51,8 @@ class ShardTest(object): def id_generator(ctx): # in reality, might want to use a separate transaction for this. - c = db1.contextual_connect() - nextid = c.execute(ids.select(for_update=True)).scalar() + c = db1.connect() + nextid = c.execute(ids.select().with_for_update()).scalar() c.execute(ids.update(values={ids.c.nextid: ids.c.nextid + 1})) return nextid @@ -411,7 +412,7 @@ class DistinctEngineShardTest(ShardTest, fixtures.TestBase): def _init_dbs(self): db1 = testing_engine( "sqlite:///shard1_%s.db" % provision.FOLLOWER_IDENT, - options=dict(pool_threadlocal=True), + options=dict(poolclass=SingletonThreadPool), ) db2 = testing_engine( "sqlite:///shard2_%s.db" % provision.FOLLOWER_IDENT @@ -551,8 +552,7 @@ class RefreshDeferExpireTest(fixtures.DeclarativeMappedTest): class LazyLoadIdentityKeyTest(fixtures.DeclarativeMappedTest): def _init_dbs(self): self.db1 = db1 = testing_engine( - "sqlite:///shard1_%s.db" % provision.FOLLOWER_IDENT, - options=dict(pool_threadlocal=True), + "sqlite:///shard1_%s.db" % provision.FOLLOWER_IDENT ) self.db2 = db2 = testing_engine( "sqlite:///shard2_%s.db" % provision.FOLLOWER_IDENT diff --git a/test/orm/inheritance/test_assorted_poly.py b/test/orm/inheritance/test_assorted_poly.py index ee869ab22..75d219b49 100644 --- a/test/orm/inheritance/test_assorted_poly.py +++ b/test/orm/inheritance/test_assorted_poly.py @@ -19,7 +19,6 @@ from sqlalchemy.orm import clear_mappers from sqlalchemy.orm import create_session from sqlalchemy.orm import join from sqlalchemy.orm import joinedload -from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import mapper from sqlalchemy.orm import polymorphic_union from sqlalchemy.orm import Query @@ -2096,7 +2095,7 @@ class Ticket2419Test(fixtures.DeclarativeMappedTest): s.commit() - q = s.query(B, B.ds.any(D.id == 1)).options(joinedload_all("es")) + q = s.query(B, B.ds.any(D.id == 1)).options(joinedload("es")) q = q.join(C, C.b_id == B.id) q = q.limit(5) eq_(q.all(), [(b, True)]) diff --git a/test/orm/inheritance/test_basic.py b/test/orm/inheritance/test_basic.py index f2336980c..ab6116256 100644 --- a/test/orm/inheritance/test_basic.py +++ b/test/orm/inheritance/test_basic.py @@ -1876,7 +1876,7 @@ class VersioningTest(fixtures.MappedTest): assert_raises( orm_exc.StaleDataError, - sess2.query(Base).with_lockmode("read").get, + sess2.query(Base).with_for_update(read=True).get, s1.id, ) diff --git a/test/orm/inheritance/test_polymorphic_rel.py b/test/orm/inheritance/test_polymorphic_rel.py index 508de986c..c16573b23 100644 --- a/test/orm/inheritance/test_polymorphic_rel.py +++ b/test/orm/inheritance/test_polymorphic_rel.py @@ -5,10 +5,9 @@ from sqlalchemy import select from sqlalchemy import testing from sqlalchemy.orm import aliased from sqlalchemy.orm import create_session +from sqlalchemy.orm import defaultload from sqlalchemy.orm import joinedload -from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import subqueryload -from sqlalchemy.orm import subqueryload_all from sqlalchemy.orm import with_polymorphic from sqlalchemy.testing import assert_raises from sqlalchemy.testing import eq_ @@ -652,12 +651,10 @@ class _PolymorphicTestBase(object): ) def test_subclass_option_pathing(self): - from sqlalchemy.orm import defer - sess = create_session() dilbert = ( sess.query(Person) - .options(defer(Engineer.machines, Machine.name)) + .options(defaultload(Engineer.machines).defer(Machine.name)) .filter(Person.name == "dilbert") .first() ) @@ -805,8 +802,8 @@ class _PolymorphicTestBase(object): eq_( sess.query(Company) .options( - joinedload_all( - Company.employees.of_type(Engineer), Engineer.machines + joinedload(Company.employees.of_type(Engineer)).joinedload( + Engineer.machines ) ) .all(), @@ -832,9 +829,9 @@ class _PolymorphicTestBase(object): eq_( sess.query(Company) .options( - subqueryload_all( - Company.employees.of_type(Engineer), Engineer.machines - ) + subqueryload( + Company.employees.of_type(Engineer) + ).subqueryload(Engineer.machines) ) .all(), expected, diff --git a/test/orm/inheritance/test_relationship.py b/test/orm/inheritance/test_relationship.py index 887453b1b..9db2a5163 100644 --- a/test/orm/inheritance/test_relationship.py +++ b/test/orm/inheritance/test_relationship.py @@ -9,12 +9,10 @@ from sqlalchemy.orm import backref from sqlalchemy.orm import contains_eager from sqlalchemy.orm import create_session from sqlalchemy.orm import joinedload -from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import mapper from sqlalchemy.orm import relationship from sqlalchemy.orm import Session from sqlalchemy.orm import subqueryload -from sqlalchemy.orm import subqueryload_all from sqlalchemy.orm import with_polymorphic from sqlalchemy.testing import AssertsCompiledSQL from sqlalchemy.testing import eq_ @@ -956,7 +954,9 @@ class EagerToSubclassTest(fixtures.MappedTest): def go(): eq_( sess.query(Parent) - .options(subqueryload_all(Parent.children, Base.related)) + .options( + subqueryload(Parent.children).subqueryload(Base.related) + ) .order_by(Parent.data) .all(), [p1, p2], @@ -973,7 +973,7 @@ class EagerToSubclassTest(fixtures.MappedTest): def go(): eq_( sess.query(pa) - .options(subqueryload_all(pa.children, Base.related)) + .options(subqueryload(pa.children).subqueryload(Base.related)) .order_by(pa.data) .all(), [p1, p2], @@ -1909,7 +1909,7 @@ class JoinedloadOverWPolyAliased( session = Session() q = session.query(cls).options( - joinedload_all(cls.links, Link.child, cls.links) + joinedload(cls.links).joinedload(Link.child).joinedload(cls.links) ) if cls is self.classes.Sub1: extra = " WHERE parent.type IN (:type_1)" @@ -1938,7 +1938,7 @@ class JoinedloadOverWPolyAliased( session = Session() q = session.query(Link).options( - joinedload_all(Link.child, parent_cls.owner) + joinedload(Link.child).joinedload(parent_cls.owner) ) if Link.child.property.mapper.class_ is self.classes.Sub1: diff --git a/test/orm/test_assorted_eager.py b/test/orm/test_assorted_eager.py index 0c2dcdf3d..8317b58f1 100644 --- a/test/orm/test_assorted_eager.py +++ b/test/orm/test_assorted_eager.py @@ -1309,9 +1309,10 @@ class EagerTest9(fixtures.MappedTest): acc = ( session.query(Account) .options( - sa.orm.joinedload_all( - "entries.transaction.entries.account" - ) + sa.orm.joinedload("entries") + .joinedload("transaction") + .joinedload("entries") + .joinedload("account") ) .order_by(Account.account_id) ).first() diff --git a/test/orm/test_attributes.py b/test/orm/test_attributes.py index 2690c7442..d99fcc77b 100644 --- a/test/orm/test_attributes.py +++ b/test/orm/test_attributes.py @@ -7,7 +7,6 @@ from sqlalchemy.orm import attributes from sqlalchemy.orm import exc as orm_exc from sqlalchemy.orm import instrumentation from sqlalchemy.orm.collections import collection -from sqlalchemy.orm.interfaces import AttributeExtension from sqlalchemy.orm.state import InstanceState from sqlalchemy.testing import assert_raises from sqlalchemy.testing import assert_raises_message @@ -574,206 +573,6 @@ class AttributesTest(fixtures.ORMTest): eq_(u.addresses[0].email_address, "lala@123.com") eq_(u.addresses[1].email_address, "foo@bar.com") - def test_extension_commit_attr(self): - """test that an extension which commits attribute history - maintains the end-result history. - - This won't work in conjunction with some unitofwork extensions. - - """ - - class Foo(fixtures.BasicEntity): - pass - - class Bar(fixtures.BasicEntity): - pass - - class ReceiveEvents(AttributeExtension): - def __init__(self, key): - self.key = key - - def append(self, state, child, initiator): - if commit: - state._commit_all(state.dict) - return child - - def remove(self, state, child, initiator): - if commit: - state._commit_all(state.dict) - return child - - def set(self, state, child, oldchild, initiator): - if commit: - state._commit_all(state.dict) - return child - - instrumentation.register_class(Foo) - instrumentation.register_class(Bar) - - b1, b2, b3, b4 = Bar(id="b1"), Bar(id="b2"), Bar(id="b3"), Bar(id="b4") - - def loadcollection(state, passive): - if passive is attributes.PASSIVE_NO_FETCH: - return attributes.PASSIVE_NO_RESULT - return [b1, b2] - - def loadscalar(state, passive): - if passive is attributes.PASSIVE_NO_FETCH: - return attributes.PASSIVE_NO_RESULT - return b2 - - attributes.register_attribute( - Foo, - "bars", - uselist=True, - useobject=True, - callable_=loadcollection, - extension=[ReceiveEvents("bars")], - ) - - attributes.register_attribute( - Foo, - "bar", - uselist=False, - useobject=True, - callable_=loadscalar, - extension=[ReceiveEvents("bar")], - ) - - attributes.register_attribute( - Foo, - "scalar", - uselist=False, - useobject=False, - extension=[ReceiveEvents("scalar")], - ) - - def create_hist(): - def hist(key, fn, *arg): - attributes.instance_state(f1)._commit_all( - attributes.instance_dict(f1) - ) - fn(*arg) - histories.append(attributes.get_history(f1, key)) - - f1 = Foo() - hist("bars", f1.bars.append, b3) - hist("bars", f1.bars.append, b4) - hist("bars", f1.bars.remove, b2) - hist("bar", setattr, f1, "bar", b3) - hist("bar", setattr, f1, "bar", None) - hist("bar", setattr, f1, "bar", b4) - hist("scalar", setattr, f1, "scalar", 5) - hist("scalar", setattr, f1, "scalar", None) - hist("scalar", setattr, f1, "scalar", 4) - - histories = [] - commit = False - create_hist() - without_commit = list(histories) - histories[:] = [] - commit = True - create_hist() - with_commit = histories - for without, with_ in zip(without_commit, with_commit): - woc = without - wic = with_ - eq_(woc, wic) - - def test_extension_lazyload_assertion(self): - class Foo(fixtures.BasicEntity): - pass - - class Bar(fixtures.BasicEntity): - pass - - class ReceiveEvents(AttributeExtension): - def append(self, state, child, initiator): - state.obj().bars - return child - - def remove(self, state, child, initiator): - state.obj().bars - return child - - def set(self, state, child, oldchild, initiator): - return child - - instrumentation.register_class(Foo) - instrumentation.register_class(Bar) - - bar1, bar2, bar3 = [Bar(id=1), Bar(id=2), Bar(id=3)] - - def func1(state, passive): - if passive is attributes.PASSIVE_NO_FETCH: - return attributes.PASSIVE_NO_RESULT - - return [bar1, bar2, bar3] - - attributes.register_attribute( - Foo, - "bars", - uselist=True, - callable_=func1, - useobject=True, - extension=[ReceiveEvents()], - ) - attributes.register_attribute( - Bar, "foos", uselist=True, useobject=True, backref="bars" - ) - - x = Foo() - assert_raises(AssertionError, Bar(id=4).foos.append, x) - - x.bars - b = Bar(id=4) - b.foos.append(x) - attributes.instance_state(x)._expire_attributes( - attributes.instance_dict(x), ["bars"] - ) - assert_raises(AssertionError, b.foos.remove, x) - - def test_scalar_listener(self): - - # listeners on ScalarAttributeImpl aren't used normally. test that - # they work for the benefit of user extensions - - class Foo(object): - - pass - - results = [] - - class ReceiveEvents(AttributeExtension): - def append(self, state, child, initiator): - assert False - - def remove(self, state, child, initiator): - results.append(("remove", state.obj(), child)) - - def set(self, state, child, oldchild, initiator): - results.append(("set", state.obj(), child, oldchild)) - return child - - instrumentation.register_class(Foo) - attributes.register_attribute( - Foo, "x", uselist=False, useobject=False, extension=ReceiveEvents() - ) - - f = Foo() - f.x = 5 - f.x = 17 - del f.x - - eq_( - results, - [ - ("set", f, 5, attributes.NEVER_SET), - ("set", f, 17, 5), - ("remove", f, 17), - ], - ) - def test_lazytrackparent(self): """test that the "hasparent" flag works properly when lazy loaders and backrefs are used diff --git a/test/orm/test_collection.py b/test/orm/test_collection.py index 4e40ac346..83f4f4451 100644 --- a/test/orm/test_collection.py +++ b/test/orm/test_collection.py @@ -1,11 +1,10 @@ from operator import and_ -import sqlalchemy as sa +from sqlalchemy import event from sqlalchemy import exc as sa_exc from sqlalchemy import ForeignKey from sqlalchemy import Integer from sqlalchemy import String -from sqlalchemy import testing from sqlalchemy import text from sqlalchemy import util from sqlalchemy.orm import attributes @@ -24,12 +23,17 @@ from sqlalchemy.testing.schema import Column from sqlalchemy.testing.schema import Table -class Canary(sa.orm.interfaces.AttributeExtension): +class Canary(object): def __init__(self): self.data = set() self.added = set() self.removed = set() + def listen(self, attr): + event.listen(attr, "append", self.append) + event.listen(attr, "remove", self.remove) + event.listen(attr, "set", self.set) + def append(self, obj, value, initiator): assert value not in self.added self.data.add(value) @@ -91,14 +95,14 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( + d = attributes.register_attribute( Foo, "attr", uselist=True, - extension=canary, typecallable=typecallable, useobject=True, ) + canary.listen(d) obj = Foo() adapter = collections.collection_adapter(obj.attr) @@ -142,14 +146,14 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( + d = attributes.register_attribute( Foo, "attr", uselist=True, - extension=canary, typecallable=typecallable, useobject=True, ) + canary.listen(d) obj = Foo() adapter = collections.collection_adapter(obj.attr) @@ -371,14 +375,14 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( + d = attributes.register_attribute( Foo, "attr", uselist=True, - extension=canary, typecallable=typecallable, useobject=True, ) + canary.listen(d) obj = Foo() direct = obj.attr @@ -578,14 +582,14 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( + d = attributes.register_attribute( Foo, "attr", uselist=True, - extension=canary, typecallable=typecallable, useobject=True, ) + canary.listen(d) obj = Foo() adapter = collections.collection_adapter(obj.attr) @@ -846,14 +850,14 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( + d = attributes.register_attribute( Foo, "attr", uselist=True, - extension=canary, typecallable=typecallable, useobject=True, ) + canary.listen(d) obj = Foo() direct = obj.attr @@ -986,14 +990,14 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( + d = attributes.register_attribute( Foo, "attr", uselist=True, - extension=canary, typecallable=typecallable, useobject=True, ) + canary.listen(d) obj = Foo() adapter = collections.collection_adapter(obj.attr) @@ -1114,14 +1118,14 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( + d = attributes.register_attribute( Foo, "attr", uselist=True, - extension=canary, typecallable=typecallable, useobject=True, ) + canary.listen(d) obj = Foo() direct = obj.attr @@ -1228,38 +1232,6 @@ class CollectionsTest(fixtures.ORMTest): self._test_dict_bulk(MyOrdered) self.assert_(getattr(MyOrdered, "_sa_instrumented") == id(MyOrdered)) - @testing.uses_deprecated(r".*Please refer to the .*bulk_replace listener") - def test_dict_subclass4(self): - # tests #2654 - class MyDict(collections.MappedCollection): - def __init__(self): - super(MyDict, self).__init__(lambda value: "k%d" % value) - - @collection.converter - def _convert(self, dictlike): - for key, value in dictlike.items(): - yield value + 5 - - class Foo(object): - pass - - canary = Canary() - - instrumentation.register_class(Foo) - attributes.register_attribute( - Foo, - "attr", - uselist=True, - extension=canary, - typecallable=MyDict, - useobject=True, - ) - - f = Foo() - f.attr = {"k1": 1, "k2": 2} - - eq_(f.attr, {"k7": 7, "k6": 6}) - def test_dict_duck(self): class DictLike(object): def __init__(self): @@ -1371,14 +1343,14 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( + d = attributes.register_attribute( Foo, "attr", uselist=True, - extension=canary, typecallable=typecallable, useobject=True, ) + canary.listen(d) obj = Foo() adapter = collections.collection_adapter(obj.attr) @@ -1532,14 +1504,10 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() instrumentation.register_class(Foo) - attributes.register_attribute( - Foo, - "attr", - uselist=True, - extension=canary, - typecallable=Custom, - useobject=True, + d = attributes.register_attribute( + Foo, "attr", uselist=True, typecallable=Custom, useobject=True ) + canary.listen(d) obj = Foo() adapter = collections.collection_adapter(obj.attr) @@ -1610,9 +1578,10 @@ class CollectionsTest(fixtures.ORMTest): canary = Canary() creator = self.entity_maker instrumentation.register_class(Foo) - attributes.register_attribute( - Foo, "attr", uselist=True, extension=canary, useobject=True + d = attributes.register_attribute( + Foo, "attr", uselist=True, useobject=True ) + canary.listen(d) obj = Foo() col1 = obj.attr @@ -2420,77 +2389,6 @@ class InstrumentationTest(fixtures.ORMTest): collections._instrument_class(Touchy) - @testing.uses_deprecated(r".*Please refer to the .*bulk_replace listener") - def test_name_setup(self): - class Base(object): - @collection.iterator - def base_iterate(self, x): - return "base_iterate" - - @collection.appender - def base_append(self, x): - return "base_append" - - @collection.converter - def base_convert(self, x): - return "base_convert" - - @collection.remover - def base_remove(self, x): - return "base_remove" - - from sqlalchemy.orm.collections import _instrument_class - - _instrument_class(Base) - - eq_(Base._sa_remover(Base(), 5), "base_remove") - eq_(Base._sa_appender(Base(), 5), "base_append") - eq_(Base._sa_iterator(Base(), 5), "base_iterate") - eq_(Base._sa_converter(Base(), 5), "base_convert") - - class Sub(Base): - @collection.converter - def base_convert(self, x): - return "sub_convert" - - @collection.remover - def sub_remove(self, x): - return "sub_remove" - - _instrument_class(Sub) - - eq_(Sub._sa_appender(Sub(), 5), "base_append") - eq_(Sub._sa_remover(Sub(), 5), "sub_remove") - eq_(Sub._sa_iterator(Sub(), 5), "base_iterate") - eq_(Sub._sa_converter(Sub(), 5), "sub_convert") - - @testing.uses_deprecated(r".*Please refer to the .*init_collection") - def test_link_event(self): - canary = [] - - class Collection(list): - @collection.linker - def _on_link(self, obj): - canary.append(obj) - - class Foo(object): - pass - - instrumentation.register_class(Foo) - attributes.register_attribute( - Foo, "attr", uselist=True, typecallable=Collection, useobject=True - ) - - f1 = Foo() - f1.attr.append(3) - - eq_(canary, [f1.attr._sa_adapter]) - adapter_1 = f1.attr._sa_adapter - - l2 = Collection() - f1.attr = l2 - eq_(canary, [adapter_1, f1.attr._sa_adapter, None]) - def test_referenced_by_owner(self): class Foo(object): pass diff --git a/test/orm/test_defaults.py b/test/orm/test_defaults.py index 40e72ad74..e870b3057 100644 --- a/test/orm/test_defaults.py +++ b/test/orm/test_defaults.py @@ -40,29 +40,27 @@ class TriggerDefaultsTest(fixtures.MappedTest): "CREATE TRIGGER dt_ins AFTER INSERT ON dt " "FOR EACH ROW BEGIN " "UPDATE dt SET col2='ins', col4='ins' " - "WHERE dt.id = NEW.id; END", - on="sqlite", - ), + "WHERE dt.id = NEW.id; END" + ).execute_if(dialect="sqlite"), sa.DDL( "CREATE TRIGGER dt_ins ON dt AFTER INSERT AS " "UPDATE dt SET col2='ins', col4='ins' " - "WHERE dt.id IN (SELECT id FROM inserted);", - on="mssql", - ), + "WHERE dt.id IN (SELECT id FROM inserted);" + ).execute_if(dialect="mssql"), sa.DDL( "CREATE TRIGGER dt_ins BEFORE INSERT " "ON dt " "FOR EACH ROW " "BEGIN " - ":NEW.col2 := 'ins'; :NEW.col4 := 'ins'; END;", - on="oracle", - ), + ":NEW.col2 := 'ins'; :NEW.col4 := 'ins'; END;" + ).execute_if(dialect="oracle"), sa.DDL( "CREATE TRIGGER dt_ins BEFORE INSERT ON dt " "FOR EACH ROW BEGIN " - "SET NEW.col2='ins'; SET NEW.col4='ins'; END", - on=lambda ddl, event, target, bind, **kw: bind.engine.name - not in ("oracle", "mssql", "sqlite"), + "SET NEW.col2='ins'; SET NEW.col4='ins'; END" + ).execute_if( + callable_=lambda ddl, target, bind, **kw: bind.engine.name + not in ("oracle", "mssql", "sqlite") ), ): event.listen(dt, "after_create", ins) @@ -74,27 +72,25 @@ class TriggerDefaultsTest(fixtures.MappedTest): "CREATE TRIGGER dt_up AFTER UPDATE ON dt " "FOR EACH ROW BEGIN " "UPDATE dt SET col3='up', col4='up' " - "WHERE dt.id = OLD.id; END", - on="sqlite", - ), + "WHERE dt.id = OLD.id; END" + ).execute_if(dialect="sqlite"), sa.DDL( "CREATE TRIGGER dt_up ON dt AFTER UPDATE AS " "UPDATE dt SET col3='up', col4='up' " - "WHERE dt.id IN (SELECT id FROM deleted);", - on="mssql", - ), + "WHERE dt.id IN (SELECT id FROM deleted);" + ).execute_if(dialect="mssql"), sa.DDL( "CREATE TRIGGER dt_up BEFORE UPDATE ON dt " "FOR EACH ROW BEGIN " - ":NEW.col3 := 'up'; :NEW.col4 := 'up'; END;", - on="oracle", - ), + ":NEW.col3 := 'up'; :NEW.col4 := 'up'; END;" + ).execute_if(dialect="oracle"), sa.DDL( "CREATE TRIGGER dt_up BEFORE UPDATE ON dt " "FOR EACH ROW BEGIN " - "SET NEW.col3='up'; SET NEW.col4='up'; END", - on=lambda ddl, event, target, bind, **kw: bind.engine.name - not in ("oracle", "mssql", "sqlite"), + "SET NEW.col3='up'; SET NEW.col4='up'; END" + ).execute_if( + callable_=lambda ddl, target, bind, **kw: bind.engine.name + not in ("oracle", "mssql", "sqlite") ), ): event.listen(dt, "after_create", up) diff --git a/test/orm/test_deferred.py b/test/orm/test_deferred.py index 3c8172850..551952cfe 100644 --- a/test/orm/test_deferred.py +++ b/test/orm/test_deferred.py @@ -966,7 +966,9 @@ class DeferredOptionsTest(AssertsCompiledSQL, _fixtures.FixtureTest): ) q = sess.query(User).options( - defer(User.orders, Order.items, Item.description) + defaultload(User.orders) + .defaultload(Order.items) + .defer(Item.description) ) self.assert_compile(q, exp) diff --git a/test/orm/test_deprecations.py b/test/orm/test_deprecations.py index 195012b99..04dff0252 100644 --- a/test/orm/test_deprecations.py +++ b/test/orm/test_deprecations.py @@ -1,80 +1,320 @@ -"""The collection of modern alternatives to deprecated & removed functionality. - -Collects specimens of old ORM code and explicitly covers the recommended -modern (i.e. not deprecated) alternative to them. The tests snippets here can -be migrated directly to the wiki, docs, etc. - -.. deprecated:: - - This test suite is interested in extremely old (pre 0.5) patterns - and in modern use illustrates trivial use cases that don't need - an additional test suite. - -""" -from sqlalchemy import ForeignKey +import sqlalchemy as sa +from sqlalchemy import event +from sqlalchemy import exc from sqlalchemy import func from sqlalchemy import Integer +from sqlalchemy import MetaData +from sqlalchemy import select from sqlalchemy import String -from sqlalchemy import text +from sqlalchemy import testing +from sqlalchemy.ext.declarative import comparable_using +from sqlalchemy.ext.declarative import declarative_base +from sqlalchemy.orm import AttributeExtension +from sqlalchemy.orm import attributes +from sqlalchemy.orm import collections +from sqlalchemy.orm import column_property +from sqlalchemy.orm import comparable_property +from sqlalchemy.orm import composite +from sqlalchemy.orm import configure_mappers from sqlalchemy.orm import create_session +from sqlalchemy.orm import defer +from sqlalchemy.orm import deferred +from sqlalchemy.orm import EXT_CONTINUE +from sqlalchemy.orm import identity +from sqlalchemy.orm import instrumentation +from sqlalchemy.orm import joinedload +from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import mapper +from sqlalchemy.orm import MapperExtension +from sqlalchemy.orm import PropComparator from sqlalchemy.orm import relationship +from sqlalchemy.orm import Session +from sqlalchemy.orm import SessionExtension from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import synonym +from sqlalchemy.orm import undefer +from sqlalchemy.orm.collections import collection +from sqlalchemy.testing import assert_raises +from sqlalchemy.testing import assert_raises_message +from sqlalchemy.testing import assertions +from sqlalchemy.testing import AssertsCompiledSQL +from sqlalchemy.testing import engines +from sqlalchemy.testing import eq_ from sqlalchemy.testing import fixtures +from sqlalchemy.testing import is_ from sqlalchemy.testing.schema import Column from sqlalchemy.testing.schema import Table +from sqlalchemy.testing.util import gc_collect +from sqlalchemy.util.compat import pypy +from . import _fixtures +from .test_options import PathTest as OptionsPathTest +from .test_transaction import _LocalFixture + + +class DeprecationWarningsTest(fixtures.DeclarativeMappedTest): + run_setup_classes = "each" + run_setup_mappers = "each" + run_define_tables = "each" + run_create_tables = None + + def test_attribute_extension(self): + class SomeExtension(AttributeExtension): + def append(self, obj, value, initiator): + pass + + def remove(self, obj, value, initiator): + pass + + def set(self, obj, value, oldvalue, initiator): + pass + + with assertions.expect_deprecated( + ".*The column_property.extension parameter will be removed in a " + "future release." + ): + + class Foo(self.DeclarativeBasic): + __tablename__ = "foo" + + id = Column(Integer, primary_key=True) + foo = column_property( + Column("q", Integer), extension=SomeExtension() + ) + + with assertions.expect_deprecated( + "AttributeExtension.append is deprecated. The " + "AttributeExtension class will be removed in a future release.", + "AttributeExtension.remove is deprecated. The " + "AttributeExtension class will be removed in a future release.", + "AttributeExtension.set is deprecated. The " + "AttributeExtension class will be removed in a future release.", + ): + configure_mappers() + + def test_attribute_extension_parameter(self): + class SomeExtension(AttributeExtension): + def append(self, obj, value, initiator): + pass + + with assertions.expect_deprecated( + ".*The relationship.extension parameter will be removed in a " + "future release." + ): + relationship("Bar", extension=SomeExtension) + + with assertions.expect_deprecated( + ".*The column_property.extension parameter will be removed in a " + "future release." + ): + column_property(Column("q", Integer), extension=SomeExtension) + + with assertions.expect_deprecated( + ".*The composite.extension parameter will be removed in a " + "future release." + ): + composite("foo", extension=SomeExtension) + + def test_session_extension(self): + class SomeExtension(SessionExtension): + def after_commit(self, session): + pass + + def after_rollback(self, session): + pass + + def before_flush(self, session, flush_context, instances): + pass + + with assertions.expect_deprecated( + ".*The Session.extension parameter will be removed", + "SessionExtension.after_commit is deprecated. " + "The SessionExtension class", + "SessionExtension.before_flush is deprecated. " + "The SessionExtension class", + "SessionExtension.after_rollback is deprecated. " + "The SessionExtension class", + ): + Session(extension=SomeExtension()) + + def test_mapper_extension(self): + class SomeExtension(MapperExtension): + def init_instance( + self, mapper, class_, oldinit, instance, args, kwargs + ): + pass + + def init_failed( + self, mapper, class_, oldinit, instance, args, kwargs + ): + pass + + with assertions.expect_deprecated( + "MapperExtension.init_instance is deprecated. " + "The MapperExtension class", + "MapperExtension.init_failed is deprecated. " + "The MapperExtension class", + ".*The mapper.extension parameter will be removed", + ): + + class Foo(self.DeclarativeBasic): + __tablename__ = "foo" + + id = Column(Integer, primary_key=True) + + __mapper_args__ = {"extension": SomeExtension()} + + def test_session_weak_identity_map(self): + with testing.expect_deprecated( + ".*Session.weak_identity_map parameter as well as the" + ): + s = Session(weak_identity_map=True) + + is_(s._identity_cls, identity.WeakInstanceDict) + + with assertions.expect_deprecated( + "The Session.weak_identity_map parameter as well as" + ): + s = Session(weak_identity_map=False) + + is_(s._identity_cls, identity.StrongInstanceDict) + + s = Session() + is_(s._identity_cls, identity.WeakInstanceDict) + + def test_session_prune(self): + s = Session() + + with assertions.expect_deprecated( + r"The Session.prune\(\) method is deprecated along with " + "Session.weak_identity_map" + ): + s.prune() + + def test_session_enable_transaction_accounting(self): + with assertions.expect_deprecated( + "the Session._enable_transaction_accounting parameter is " + "deprecated" + ): + s = Session(_enable_transaction_accounting=False) + + def test_session_is_modified(self): + class Foo(self.DeclarativeBasic): + __tablename__ = "foo" + + id = Column(Integer, primary_key=True) + + f1 = Foo() + s = Session() + with assertions.expect_deprecated( + "The Session.is_modified.passive flag is deprecated" + ): + # this flag was for a long time documented as requiring + # that it be set to True, so we've changed the default here + # so that the warning emits + s.is_modified(f1, passive=True) + + +class DeprecatedAccountingFlagsTest(_LocalFixture): + def test_rollback_no_accounting(self): + User, users = self.classes.User, self.tables.users + + with testing.expect_deprecated( + "The Session._enable_transaction_accounting parameter" + ): + sess = sessionmaker(_enable_transaction_accounting=False)() + u1 = User(name="ed") + sess.add(u1) + sess.commit() + + u1.name = "edwardo" + sess.rollback() + + testing.db.execute( + users.update(users.c.name == "ed").values(name="edward") + ) + + assert u1.name == "edwardo" + sess.expire_all() + assert u1.name == "edward" + + def test_commit_no_accounting(self): + User, users = self.classes.User, self.tables.users + with testing.expect_deprecated( + "The Session._enable_transaction_accounting parameter" + ): + sess = sessionmaker(_enable_transaction_accounting=False)() + u1 = User(name="ed") + sess.add(u1) + sess.commit() -class QueryAlternativesTest(fixtures.MappedTest): - r'''Collects modern idioms for Queries + u1.name = "edwardo" + sess.rollback() - The docstring for each test case serves as miniature documentation about - the deprecated use case, and the test body illustrates (and covers) the - intended replacement code to accomplish the same task. + testing.db.execute( + users.update(users.c.name == "ed").values(name="edward") + ) - Documenting the "old way" including the argument signature helps these - cases remain useful to readers even after the deprecated method has been - removed from the modern codebase. + assert u1.name == "edwardo" + sess.commit() - Format:: + assert testing.db.execute(select([users.c.name])).fetchall() == [ + ("edwardo",) + ] + assert u1.name == "edwardo" - def test_deprecated_thing(self): - """Query.methodname(old, arg, **signature) + sess.delete(u1) + sess.commit() - output = session.query(User).deprecatedmethod(inputs) + def test_preflush_no_accounting(self): + User, users = self.classes.User, self.tables.users - """ + with testing.expect_deprecated( + "The Session._enable_transaction_accounting parameter" + ): + sess = Session( + _enable_transaction_accounting=False, + autocommit=True, + autoflush=False, + ) + u1 = User(name="ed") + sess.add(u1) + sess.flush() - # 0.4+ - output = session.query(User).newway(inputs) - assert output is correct + sess.begin() + u1.name = "edwardo" + u2 = User(name="some other user") + sess.add(u2) - # 0.5+ - output = session.query(User).evennewerway(inputs) - assert output is correct + sess.rollback() - ''' + sess.begin() + assert testing.db.execute(select([users.c.name])).fetchall() == [ + ("ed",) + ] - run_inserts = "once" - run_deletes = None + +class TLTransactionTest(fixtures.MappedTest): + run_dispose_bind = "once" + __backend__ = True @classmethod - def define_tables(cls, metadata): - Table( - "users_table", - metadata, - Column("id", Integer, primary_key=True), - Column("name", String(64)), - ) + def setup_bind(cls): + with testing.expect_deprecated( + ".*'threadlocal' engine strategy is deprecated" + ): + return engines.testing_engine(options=dict(strategy="threadlocal")) + @classmethod + def define_tables(cls, metadata): Table( - "addresses_table", + "users", metadata, - Column("id", Integer, primary_key=True), - Column("user_id", Integer, ForeignKey("users_table.id")), - Column("email_address", String(128)), - Column("purpose", String(16)), - Column("bounces", Integer, default=0), + Column( + "id", Integer, primary_key=True, test_needs_autoincrement=True + ), + Column("name", String(20)), + test_needs_acid=True, ) @classmethod @@ -82,511 +322,1843 @@ class QueryAlternativesTest(fixtures.MappedTest): class User(cls.Basic): pass - class Address(cls.Basic): - pass - @classmethod def setup_mappers(cls): - addresses_table, User, users_table, Address = ( - cls.tables.addresses_table, - cls.classes.User, - cls.tables.users_table, - cls.classes.Address, - ) - - mapper( - User, - users_table, - properties=dict(addresses=relationship(Address, backref="user")), - ) - mapper(Address, addresses_table) + users, User = cls.tables.users, cls.classes.User - @classmethod - def fixtures(cls): - return dict( - users_table=( - ("id", "name"), - (1, "jack"), - (2, "ed"), - (3, "fred"), - (4, "chuck"), - ), - addresses_table=( - ("id", "user_id", "email_address", "purpose", "bounces"), - (1, 1, "jack@jack.home", "Personal", 0), - (2, 1, "jack@jack.bizz", "Work", 1), - (3, 2, "ed@foo.bar", "Personal", 0), - (4, 3, "fred@the.fred", "Personal", 10), - ), - ) + mapper(User, users) - ###################################################################### + @testing.exclude("mysql", "<", (5, 0, 3), "FIXME: unknown") + def test_session_nesting(self): + User = self.classes.User - def test_override_get(self): - """MapperExtension.get() + sess = create_session(bind=self.bind) + self.bind.begin() + u = User(name="ed") + sess.add(u) + sess.flush() + self.bind.commit() - x = session.query.get(5) - """ +class DeprecatedSessionFeatureTest(_fixtures.FixtureTest): + run_inserts = None - Address = self.classes.Address + def test_fast_discard_race(self): + # test issue #4068 + users, User = self.tables.users, self.classes.User - from sqlalchemy.orm.query import Query + mapper(User, users) - cache = {} + with testing.expect_deprecated(".*identity map are deprecated"): + sess = Session(weak_identity_map=False) - class MyQuery(Query): - def get(self, ident, **kwargs): - if ident in cache: - return cache[ident] - else: - x = super(MyQuery, self).get(ident) - cache[ident] = x - return x + u1 = User(name="u1") + sess.add(u1) + sess.commit() - session = sessionmaker(query_cls=MyQuery)() + u1_state = u1._sa_instance_state + sess.identity_map._dict.pop(u1_state.key) + ref = u1_state.obj + u1_state.obj = lambda: None - ad1 = session.query(Address).get(1) - assert ad1 in list(cache.values()) + u2 = sess.query(User).first() + u1_state._cleanup(ref) - def test_load(self): - """x = session.query(Address).load(1) + u3 = sess.query(User).first() - x = session.load(Address, 1) + is_(u2, u3) - """ + u2_state = u2._sa_instance_state + assert sess.identity_map.contains_state(u2._sa_instance_state) + ref = u2_state.obj + u2_state.obj = lambda: None + u2_state._cleanup(ref) + assert not sess.identity_map.contains_state(u2._sa_instance_state) - Address = self.classes.Address + def test_is_modified_passive_on(self): + User, Address = self.classes.User, self.classes.Address + users, addresses = self.tables.users, self.tables.addresses + mapper(User, users, properties={"addresses": relationship(Address)}) + mapper(Address, addresses) - session = create_session() - ad1 = session.query(Address).populate_existing().get(1) - assert bool(ad1) + s = Session() + u = User(name="fred", addresses=[Address(email_address="foo")]) + s.add(u) + s.commit() - def test_apply_max(self): - """Query.apply_max(col) + u.id - max = session.query(Address).apply_max(Address.bounces) + def go(): + assert not s.is_modified(u, passive=True) - """ + with testing.expect_deprecated( + ".*Session.is_modified.passive flag is deprecated " + ): + self.assert_sql_count(testing.db, go, 0) - Address = self.classes.Address + u.name = "newname" - session = create_session() + def go(): + assert s.is_modified(u, passive=True) - # 0.5.0 - maxes = list(session.query(Address).values(func.max(Address.bounces))) - max_ = maxes[0][0] - assert max_ == 10 + with testing.expect_deprecated( + ".*Session.is_modified.passive flag is deprecated " + ): + self.assert_sql_count(testing.db, go, 0) - max_ = session.query(func.max(Address.bounces)).one()[0] - assert max_ == 10 - def test_apply_min(self): - """Query.apply_min(col) +class StrongIdentityMapTest(_fixtures.FixtureTest): + run_inserts = None - min = session.query(Address).apply_min(Address.bounces) + def _strong_ident_fixture(self): + with testing.expect_deprecated( + ".*Session.weak_identity_map parameter as well as the" + ): + sess = create_session(weak_identity_map=False) - """ + def prune(): + with testing.expect_deprecated(".*Session.prune"): + return sess.prune() - Address = self.classes.Address + return sess, prune + def _event_fixture(self): session = create_session() - # 0.5.0 - mins = list(session.query(Address).values(func.min(Address.bounces))) - min_ = mins[0][0] - assert min_ == 0 + @event.listens_for(session, "pending_to_persistent") + @event.listens_for(session, "deleted_to_persistent") + @event.listens_for(session, "detached_to_persistent") + @event.listens_for(session, "loaded_as_persistent") + def strong_ref_object(sess, instance): + if "refs" not in sess.info: + sess.info["refs"] = refs = set() + else: + refs = sess.info["refs"] + + refs.add(instance) + + @event.listens_for(session, "persistent_to_detached") + @event.listens_for(session, "persistent_to_deleted") + @event.listens_for(session, "persistent_to_transient") + def deref_object(sess, instance): + sess.info["refs"].discard(instance) + + def prune(): + if "refs" not in session.info: + return 0 + + sess_size = len(session.identity_map) + session.info["refs"].clear() + gc_collect() + session.info["refs"] = set( + s.obj() for s in session.identity_map.all_states() + ) + return sess_size - len(session.identity_map) + + return session, prune + + def test_strong_ref_imap(self): + self._test_strong_ref(self._strong_ident_fixture) + + def test_strong_ref_events(self): + self._test_strong_ref(self._event_fixture) + + def _test_strong_ref(self, fixture): + s, prune = fixture() + + users, User = self.tables.users, self.classes.User + + mapper(User, users) + + # save user + s.add(User(name="u1")) + s.flush() + user = s.query(User).one() + user = None + print(s.identity_map) + gc_collect() + assert len(s.identity_map) == 1 + + user = s.query(User).one() + assert not s.identity_map._modified + user.name = "u2" + assert s.identity_map._modified + s.flush() + eq_(users.select().execute().fetchall(), [(user.id, "u2")]) + + def test_prune_imap(self): + self._test_prune(self._strong_ident_fixture) + + def test_prune_events(self): + self._test_prune(self._event_fixture) + + @testing.fails_if(lambda: pypy, "pypy has a real GC") + @testing.fails_on("+zxjdbc", "http://www.sqlalchemy.org/trac/ticket/1473") + def _test_prune(self, fixture): + s, prune = fixture() + + users, User = self.tables.users, self.classes.User + + mapper(User, users) + + for o in [User(name="u%s" % x) for x in range(10)]: + s.add(o) + # o is still live after this loop... + + self.assert_(len(s.identity_map) == 0) + eq_(prune(), 0) + s.flush() + gc_collect() + eq_(prune(), 9) + # o is still in local scope here, so still present + self.assert_(len(s.identity_map) == 1) + + id_ = o.id + del o + eq_(prune(), 1) + self.assert_(len(s.identity_map) == 0) + + u = s.query(User).get(id_) + eq_(prune(), 0) + self.assert_(len(s.identity_map) == 1) + u.name = "squiznart" + del u + eq_(prune(), 0) + self.assert_(len(s.identity_map) == 1) + s.flush() + eq_(prune(), 1) + self.assert_(len(s.identity_map) == 0) + + s.add(User(name="x")) + eq_(prune(), 0) + self.assert_(len(s.identity_map) == 0) + s.flush() + self.assert_(len(s.identity_map) == 1) + eq_(prune(), 1) + self.assert_(len(s.identity_map) == 0) + + u = s.query(User).get(id_) + s.delete(u) + del u + eq_(prune(), 0) + self.assert_(len(s.identity_map) == 1) + s.flush() + eq_(prune(), 0) + self.assert_(len(s.identity_map) == 0) + + +class DeprecatedMapperTest(_fixtures.FixtureTest, AssertsCompiledSQL): + __dialect__ = "default" + + def test_cancel_order_by(self): + users, User = self.tables.users, self.classes.User + + with testing.expect_deprecated( + "The Mapper.order_by parameter is deprecated, and will be " + "removed in a future release." + ): + mapper(User, users, order_by=users.c.name.desc()) + + assert ( + "order by users.name desc" + in str(create_session().query(User).statement).lower() + ) + assert ( + "order by" + not in str( + create_session().query(User).order_by(None).statement + ).lower() + ) + assert ( + "order by users.name asc" + in str( + create_session() + .query(User) + .order_by(User.name.asc()) + .statement + ).lower() + ) - min_ = session.query(func.min(Address.bounces)).one()[0] - assert min_ == 0 + eq_( + create_session().query(User).all(), + [ + User(id=7, name="jack"), + User(id=9, name="fred"), + User(id=8, name="ed"), + User(id=10, name="chuck"), + ], + ) - def test_apply_avg(self): - """Query.apply_avg(col) + eq_( + create_session().query(User).order_by(User.name).all(), + [ + User(id=10, name="chuck"), + User(id=8, name="ed"), + User(id=9, name="fred"), + User(id=7, name="jack"), + ], + ) - avg = session.query(Address).apply_avg(Address.bounces) + def test_comparable(self): + users = self.tables.users - """ + class extendedproperty(property): + attribute = 123 - Address = self.classes.Address + def method1(self): + return "method1" - session = create_session() + from sqlalchemy.orm.properties import ColumnProperty - avgs = list(session.query(Address).values(func.avg(Address.bounces))) - avg = avgs[0][0] - assert avg > 0 and avg < 10 + class UCComparator(ColumnProperty.Comparator): + __hash__ = None - avg = session.query(func.avg(Address.bounces)).one()[0] - assert avg > 0 and avg < 10 + def method1(self): + return "uccmethod1" - def test_apply_sum(self): - """Query.apply_sum(col) + def method2(self, other): + return "method2" - avg = session.query(Address).apply_avg(Address.bounces) + def __eq__(self, other): + cls = self.prop.parent.class_ + col = getattr(cls, "name") + if other is None: + return col is None + else: + return sa.func.upper(col) == sa.func.upper(other) + + def map_(with_explicit_property): + class User(object): + @extendedproperty + def uc_name(self): + if self.name is None: + return None + return self.name.upper() + + if with_explicit_property: + args = (UCComparator, User.uc_name) + else: + args = (UCComparator,) + + with assertions.expect_deprecated( + r"comparable_property\(\) is deprecated and will be " + "removed in a future release." + ): + mapper( + User, + users, + properties=dict(uc_name=sa.orm.comparable_property(*args)), + ) + return User + + for User in (map_(True), map_(False)): + sess = create_session() + sess.begin() + q = sess.query(User) + + assert hasattr(User, "name") + assert hasattr(User, "uc_name") + + eq_(User.uc_name.method1(), "method1") + eq_(User.uc_name.method2("x"), "method2") + + assert_raises_message( + AttributeError, + "Neither 'extendedproperty' object nor 'UCComparator' " + "object associated with User.uc_name has an attribute " + "'nonexistent'", + getattr, + User.uc_name, + "nonexistent", + ) - """ + # test compile + assert not isinstance(User.uc_name == "jack", bool) + u = q.filter(User.uc_name == "JACK").one() - Address = self.classes.Address + assert u.uc_name == "JACK" + assert u not in sess.dirty - session = create_session() + u.name = "some user name" + eq_(u.name, "some user name") + assert u in sess.dirty + eq_(u.uc_name, "SOME USER NAME") - avgs = list(session.query(Address).values(func.sum(Address.bounces))) - avg = avgs[0][0] - assert avg == 11 + sess.flush() + sess.expunge_all() - avg = session.query(func.sum(Address.bounces)).one()[0] - assert avg == 11 + q = sess.query(User) + u2 = q.filter(User.name == "some user name").one() + u3 = q.filter(User.uc_name == "SOME USER NAME").one() - def test_count_by(self): - r"""Query.count_by(\*args, \**params) + assert u2 is u3 - num = session.query(Address).count_by(purpose='Personal') + eq_(User.uc_name.attribute, 123) + sess.rollback() - # old-style implicit *_by join - num = session.query(User).count_by(purpose='Personal') + def test_comparable_column(self): + users, User = self.tables.users, self.classes.User - """ + class MyComparator(sa.orm.properties.ColumnProperty.Comparator): + __hash__ = None - User, Address = self.classes.User, self.classes.Address + def __eq__(self, other): + # lower case comparison + return func.lower(self.__clause_element__()) == func.lower( + other + ) - session = create_session() + def intersects(self, other): + # non-standard comparator + return self.__clause_element__().op("&=")(other) + + mapper( + User, + users, + properties={ + "name": sa.orm.column_property( + users.c.name, comparator_factory=MyComparator + ) + }, + ) - num = session.query(Address).filter_by(purpose="Personal").count() - assert num == 3, num + assert_raises_message( + AttributeError, + "Neither 'InstrumentedAttribute' object nor " + "'MyComparator' object associated with User.name has " + "an attribute 'nonexistent'", + getattr, + User.name, + "nonexistent", + ) - num = ( - session.query(User) - .join("addresses") - .filter(Address.purpose == "Personal") - ).count() - assert num == 3, num + eq_( + str( + (User.name == "ed").compile( + dialect=sa.engine.default.DefaultDialect() + ) + ), + "lower(users.name) = lower(:lower_1)", + ) + eq_( + str( + (User.name.intersects("ed")).compile( + dialect=sa.engine.default.DefaultDialect() + ) + ), + "users.name &= :name_1", + ) - def test_count_whereclause(self): - r"""Query.count(whereclause=None, params=None, \**kwargs) + def test_info(self): + users = self.tables.users + Address = self.classes.Address - num = session.query(Address).count(address_table.c.bounces > 1) + class MyComposite(object): + pass - """ + with assertions.expect_deprecated( + r"comparable_property\(\) is deprecated and will be " + "removed in a future release." + ): + for constructor, args in [(comparable_property, "foo")]: + obj = constructor(info={"x": "y"}, *args) + eq_(obj.info, {"x": "y"}) + obj.info["q"] = "p" + eq_(obj.info, {"x": "y", "q": "p"}) - Address = self.classes.Address + obj = constructor(*args) + eq_(obj.info, {}) + obj.info["q"] = "p" + eq_(obj.info, {"q": "p"}) - session = create_session() + def test_add_property(self): + users = self.tables.users - num = session.query(Address).filter(Address.bounces > 1).count() - assert num == 1, num + assert_col = [] - def test_execute(self): - r"""Query.execute(clauseelement, params=None, \*args, \**kwargs) + class User(fixtures.ComparableEntity): + def _get_name(self): + assert_col.append(("get", self._name)) + return self._name - users = session.query(User).execute(users_table.select()) + def _set_name(self, name): + assert_col.append(("set", name)) + self._name = name - """ + name = property(_get_name, _set_name) - User, users_table = self.classes.User, self.tables.users_table + def _uc_name(self): + if self._name is None: + return None + return self._name.upper() - session = create_session() + uc_name = property(_uc_name) + uc_name2 = property(_uc_name) - users = session.query(User).from_statement(users_table.select()).all() - assert len(users) == 4 + m = mapper(User, users) - def test_get_by(self): - r"""Query.get_by(\*args, \**params) + class UCComparator(PropComparator): + __hash__ = None - user = session.query(User).get_by(name='ed') + def __eq__(self, other): + cls = self.prop.parent.class_ + col = getattr(cls, "name") + if other is None: + return col is None + else: + return func.upper(col) == func.upper(other) + + m.add_property("_name", deferred(users.c.name)) + m.add_property("name", synonym("_name")) + with assertions.expect_deprecated( + r"comparable_property\(\) is deprecated and will be " + "removed in a future release." + ): + m.add_property("uc_name", comparable_property(UCComparator)) + m.add_property( + "uc_name2", comparable_property(UCComparator, User.uc_name2) + ) - # 0.3-style implicit *_by join - user = session.query(User).get_by(email_addresss='fred@the.fred') + sess = create_session(autocommit=False) + assert sess.query(User).get(7) - """ + u = sess.query(User).filter_by(name="jack").one() - User, Address = self.classes.User, self.classes.Address + def go(): + eq_(u.name, "jack") + eq_(u.uc_name, "JACK") + eq_(u.uc_name2, "JACK") + eq_(assert_col, [("get", "jack")], str(assert_col)) - session = create_session() + self.sql_count_(1, go) - user = session.query(User).filter_by(name="ed").first() - assert user.name == "ed" + def test_kwarg_accepted(self): + users, Address = self.tables.users, self.classes.Address - user = ( - session.query(User) - .join("addresses") - .filter(Address.email_address == "fred@the.fred") - ).first() - assert user.name == "fred" + class DummyComposite(object): + def __init__(self, x, y): + pass - user = ( - session.query(User) - .filter( - User.addresses.any(Address.email_address == "fred@the.fred") + class MyFactory(PropComparator): + pass + + with assertions.expect_deprecated( + r"comparable_property\(\) is deprecated and will be " + "removed in a future release." + ): + for args in ((comparable_property,),): + fn = args[0] + args = args[1:] + fn(comparator_factory=MyFactory, *args) + + def test_merge_synonym_comparable(self): + users = self.tables.users + + class User(object): + class Comparator(PropComparator): + pass + + def _getValue(self): + return self._value + + def _setValue(self, value): + setattr(self, "_value", value) + + value = property(_getValue, _setValue) + + with assertions.expect_deprecated( + r"comparable_property\(\) is deprecated and will be " + "removed in a future release." + ): + mapper( + User, + users, + properties={ + "uid": synonym("id"), + "foobar": comparable_property(User.Comparator, User.value), + }, ) - .first() + + sess = create_session() + u = User() + u.name = "ed" + sess.add(u) + sess.flush() + sess.expunge(u) + sess.merge(u) + + +class DeprecatedDeclTest(fixtures.TestBase): + @testing.provide_metadata + def test_comparable_using(self): + class NameComparator(sa.orm.PropComparator): + @property + def upperself(self): + cls = self.prop.parent.class_ + col = getattr(cls, "name") + return sa.func.upper(col) + + def operate(self, op, other, **kw): + return op(self.upperself, other, **kw) + + Base = declarative_base(metadata=self.metadata) + + with testing.expect_deprecated( + r"comparable_property\(\) is deprecated and will be " + "removed in a future release." + ): + + class User(Base, fixtures.ComparableEntity): + + __tablename__ = "users" + id = Column( + "id", + Integer, + primary_key=True, + test_needs_autoincrement=True, + ) + name = Column("name", String(50)) + + @comparable_using(NameComparator) + @property + def uc_name(self): + return self.name is not None and self.name.upper() or None + + Base.metadata.create_all() + sess = create_session() + u1 = User(name="someuser") + eq_(u1.name, "someuser", u1.name) + eq_(u1.uc_name, "SOMEUSER", u1.uc_name) + sess.add(u1) + sess.flush() + sess.expunge_all() + rt = sess.query(User).filter(User.uc_name == "SOMEUSER").one() + eq_(rt, u1) + sess.expunge_all() + rt = sess.query(User).filter(User.uc_name.startswith("SOMEUSE")).one() + eq_(rt, u1) + + +class DeprecatedMapperExtensionTest(_fixtures.FixtureTest): + + """Superseded by MapperEventsTest - test backwards + compatibility of MapperExtension.""" + + run_inserts = None + + def extension(self): + methods = [] + + class Ext(MapperExtension): + def instrument_class(self, mapper, cls): + methods.append("instrument_class") + return EXT_CONTINUE + + def init_instance( + self, mapper, class_, oldinit, instance, args, kwargs + ): + methods.append("init_instance") + return EXT_CONTINUE + + def init_failed( + self, mapper, class_, oldinit, instance, args, kwargs + ): + methods.append("init_failed") + return EXT_CONTINUE + + def reconstruct_instance(self, mapper, instance): + methods.append("reconstruct_instance") + return EXT_CONTINUE + + def before_insert(self, mapper, connection, instance): + methods.append("before_insert") + return EXT_CONTINUE + + def after_insert(self, mapper, connection, instance): + methods.append("after_insert") + return EXT_CONTINUE + + def before_update(self, mapper, connection, instance): + methods.append("before_update") + return EXT_CONTINUE + + def after_update(self, mapper, connection, instance): + methods.append("after_update") + return EXT_CONTINUE + + def before_delete(self, mapper, connection, instance): + methods.append("before_delete") + return EXT_CONTINUE + + def after_delete(self, mapper, connection, instance): + methods.append("after_delete") + return EXT_CONTINUE + + return Ext, methods + + def test_basic(self): + """test that common user-defined methods get called.""" + + User, users = self.classes.User, self.tables.users + + Ext, methods = self.extension() + + with testing.expect_deprecated( + "MapperExtension is deprecated in favor of the MapperEvents", + "MapperExtension.before_insert is deprecated", + "MapperExtension.instrument_class is deprecated", + "MapperExtension.init_instance is deprecated", + "MapperExtension.after_insert is deprecated", + "MapperExtension.reconstruct_instance is deprecated", + "MapperExtension.before_delete is deprecated", + "MapperExtension.after_delete is deprecated", + "MapperExtension.before_update is deprecated", + "MapperExtension.after_update is deprecated", + "MapperExtension.init_failed is deprecated", + ): + mapper(User, users, extension=Ext()) + sess = create_session() + u = User(name="u1") + sess.add(u) + sess.flush() + u = sess.query(User).populate_existing().get(u.id) + sess.expunge_all() + u = sess.query(User).get(u.id) + u.name = "u1 changed" + sess.flush() + sess.delete(u) + sess.flush() + eq_( + methods, + [ + "instrument_class", + "init_instance", + "before_insert", + "after_insert", + "reconstruct_instance", + "before_update", + "after_update", + "before_delete", + "after_delete", + ], + ) + + def test_inheritance(self): + users, addresses, User = ( + self.tables.users, + self.tables.addresses, + self.classes.User, + ) + + Ext, methods = self.extension() + + class AdminUser(User): + pass + + with testing.expect_deprecated( + "MapperExtension is deprecated in favor of the MapperEvents", + "MapperExtension.before_insert is deprecated", + "MapperExtension.instrument_class is deprecated", + "MapperExtension.init_instance is deprecated", + "MapperExtension.after_insert is deprecated", + "MapperExtension.reconstruct_instance is deprecated", + "MapperExtension.before_delete is deprecated", + "MapperExtension.after_delete is deprecated", + "MapperExtension.before_update is deprecated", + "MapperExtension.after_update is deprecated", + "MapperExtension.init_failed is deprecated", + ): + mapper(User, users, extension=Ext()) + mapper( + AdminUser, + addresses, + inherits=User, + properties={"address_id": addresses.c.id}, ) - assert user.name == "fred" - def test_instances_entities(self): - r"""Query.instances(cursor, \*mappers_or_columns, \**kwargs) + sess = create_session() + am = AdminUser(name="au1", email_address="au1@e1") + sess.add(am) + sess.flush() + am = sess.query(AdminUser).populate_existing().get(am.id) + sess.expunge_all() + am = sess.query(AdminUser).get(am.id) + am.name = "au1 changed" + sess.flush() + sess.delete(am) + sess.flush() + eq_( + methods, + [ + "instrument_class", + "instrument_class", + "init_instance", + "before_insert", + "after_insert", + "reconstruct_instance", + "before_update", + "after_update", + "before_delete", + "after_delete", + ], + ) - sel = users_table.join(addresses_table).select(use_labels=True) - res = session.query(User).instances(sel.execute(), Address) + def test_before_after_only_collection(self): + """before_update is called on parent for collection modifications, + after_update is called even if no columns were updated. """ - addresses_table, User, users_table, Address = ( - self.tables.addresses_table, - self.classes.User, - self.tables.users_table, - self.classes.Address, + keywords, items, item_keywords, Keyword, Item = ( + self.tables.keywords, + self.tables.items, + self.tables.item_keywords, + self.classes.Keyword, + self.classes.Item, ) - session = create_session() + Ext1, methods1 = self.extension() + Ext2, methods2 = self.extension() + + with testing.expect_deprecated( + "MapperExtension is deprecated in favor of the MapperEvents", + "MapperExtension.before_insert is deprecated", + "MapperExtension.instrument_class is deprecated", + "MapperExtension.init_instance is deprecated", + "MapperExtension.after_insert is deprecated", + "MapperExtension.reconstruct_instance is deprecated", + "MapperExtension.before_delete is deprecated", + "MapperExtension.after_delete is deprecated", + "MapperExtension.before_update is deprecated", + "MapperExtension.after_update is deprecated", + "MapperExtension.init_failed is deprecated", + ): + mapper( + Item, + items, + extension=Ext1(), + properties={ + "keywords": relationship(Keyword, secondary=item_keywords) + }, + ) + with testing.expect_deprecated( + "MapperExtension is deprecated in favor of the MapperEvents", + "MapperExtension.before_insert is deprecated", + "MapperExtension.instrument_class is deprecated", + "MapperExtension.init_instance is deprecated", + "MapperExtension.after_insert is deprecated", + "MapperExtension.reconstruct_instance is deprecated", + "MapperExtension.before_delete is deprecated", + "MapperExtension.after_delete is deprecated", + "MapperExtension.before_update is deprecated", + "MapperExtension.after_update is deprecated", + "MapperExtension.init_failed is deprecated", + ): + mapper(Keyword, keywords, extension=Ext2()) + + sess = create_session() + i1 = Item(description="i1") + k1 = Keyword(name="k1") + sess.add(i1) + sess.add(k1) + sess.flush() + eq_( + methods1, + [ + "instrument_class", + "init_instance", + "before_insert", + "after_insert", + ], + ) + eq_( + methods2, + [ + "instrument_class", + "init_instance", + "before_insert", + "after_insert", + ], + ) - sel = users_table.join(addresses_table).select(use_labels=True) - res = list(session.query(User, Address).instances(sel.execute())) + del methods1[:] + del methods2[:] + i1.keywords.append(k1) + sess.flush() + eq_(methods1, ["before_update", "after_update"]) + eq_(methods2, []) - assert len(res) == 4 - cola, colb = res[0] - assert isinstance(cola, User) and isinstance(colb, Address) + def test_inheritance_with_dupes(self): + """Inheritance with the same extension instance on both mappers.""" - def test_join_by(self): - r"""Query.join_by(\*args, \**params) + users, addresses, User = ( + self.tables.users, + self.tables.addresses, + self.classes.User, + ) - TODO - """ + Ext, methods = self.extension() - session = create_session() + class AdminUser(User): + pass - def test_join_to(self): - """Query.join_to(key) + ext = Ext() + with testing.expect_deprecated( + "MapperExtension is deprecated in favor of the MapperEvents", + "MapperExtension.before_insert is deprecated", + "MapperExtension.instrument_class is deprecated", + "MapperExtension.init_instance is deprecated", + "MapperExtension.after_insert is deprecated", + "MapperExtension.reconstruct_instance is deprecated", + "MapperExtension.before_delete is deprecated", + "MapperExtension.after_delete is deprecated", + "MapperExtension.before_update is deprecated", + "MapperExtension.after_update is deprecated", + "MapperExtension.init_failed is deprecated", + ): + mapper(User, users, extension=ext) + + with testing.expect_deprecated( + "MapperExtension is deprecated in favor of the MapperEvents" + ): + mapper( + AdminUser, + addresses, + inherits=User, + extension=ext, + properties={"address_id": addresses.c.id}, + ) - TODO - """ + sess = create_session() + am = AdminUser(name="au1", email_address="au1@e1") + sess.add(am) + sess.flush() + am = sess.query(AdminUser).populate_existing().get(am.id) + sess.expunge_all() + am = sess.query(AdminUser).get(am.id) + am.name = "au1 changed" + sess.flush() + sess.delete(am) + sess.flush() + eq_( + methods, + [ + "instrument_class", + "instrument_class", + "init_instance", + "before_insert", + "after_insert", + "reconstruct_instance", + "before_update", + "after_update", + "before_delete", + "after_delete", + ], + ) - session = create_session() + def test_unnecessary_methods_not_evented(self): + users = self.tables.users - def test_join_via(self): - """Query.join_via(keys) + class MyExtension(MapperExtension): + def before_insert(self, mapper, connection, instance): + pass + + class Foo(object): + pass + + with testing.expect_deprecated( + "MapperExtension is deprecated in favor of the MapperEvents", + "MapperExtension.before_insert is deprecated", + ): + m = mapper(Foo, users, extension=MyExtension()) + assert not m.class_manager.dispatch.load + assert not m.dispatch.before_update + assert len(m.dispatch.before_insert) == 1 + + +class DeprecatedSessionExtensionTest(_fixtures.FixtureTest): + run_inserts = None + + def test_extension(self): + User, users = self.classes.User, self.tables.users + + mapper(User, users) + log = [] + + class MyExt(SessionExtension): + def before_commit(self, session): + log.append("before_commit") + + def after_commit(self, session): + log.append("after_commit") + + def after_rollback(self, session): + log.append("after_rollback") + + def before_flush(self, session, flush_context, objects): + log.append("before_flush") + + def after_flush(self, session, flush_context): + log.append("after_flush") + + def after_flush_postexec(self, session, flush_context): + log.append("after_flush_postexec") + + def after_begin(self, session, transaction, connection): + log.append("after_begin") + + def after_attach(self, session, instance): + log.append("after_attach") + + def after_bulk_update(self, session, query, query_context, result): + log.append("after_bulk_update") + + def after_bulk_delete(self, session, query, query_context, result): + log.append("after_bulk_delete") + + with testing.expect_deprecated( + "SessionExtension is deprecated in favor of " "the SessionEvents", + "SessionExtension.before_commit is deprecated", + "SessionExtension.after_commit is deprecated", + "SessionExtension.after_begin is deprecated", + "SessionExtension.after_attach is deprecated", + "SessionExtension.before_flush is deprecated", + "SessionExtension.after_flush is deprecated", + "SessionExtension.after_flush_postexec is deprecated", + "SessionExtension.after_rollback is deprecated", + "SessionExtension.after_bulk_update is deprecated", + "SessionExtension.after_bulk_delete is deprecated", + ): + sess = create_session(extension=MyExt()) + u = User(name="u1") + sess.add(u) + sess.flush() + assert log == [ + "after_attach", + "before_flush", + "after_begin", + "after_flush", + "after_flush_postexec", + "before_commit", + "after_commit", + ] + log = [] + with testing.expect_deprecated( + "SessionExtension is deprecated in favor of " "the SessionEvents", + "SessionExtension.before_commit is deprecated", + "SessionExtension.after_commit is deprecated", + "SessionExtension.after_begin is deprecated", + "SessionExtension.after_attach is deprecated", + "SessionExtension.before_flush is deprecated", + "SessionExtension.after_flush is deprecated", + "SessionExtension.after_flush_postexec is deprecated", + "SessionExtension.after_rollback is deprecated", + "SessionExtension.after_bulk_update is deprecated", + "SessionExtension.after_bulk_delete is deprecated", + ): + sess = create_session(autocommit=False, extension=MyExt()) + u = User(name="u1") + sess.add(u) + sess.flush() + assert log == [ + "after_attach", + "before_flush", + "after_begin", + "after_flush", + "after_flush_postexec", + ] + log = [] + u.name = "ed" + sess.commit() + assert log == [ + "before_commit", + "before_flush", + "after_flush", + "after_flush_postexec", + "after_commit", + ] + log = [] + sess.commit() + assert log == ["before_commit", "after_commit"] + log = [] + sess.query(User).delete() + assert log == ["after_begin", "after_bulk_delete"] + log = [] + sess.query(User).update({"name": "foo"}) + assert log == ["after_bulk_update"] + log = [] + with testing.expect_deprecated( + "SessionExtension is deprecated in favor of " "the SessionEvents", + "SessionExtension.before_commit is deprecated", + "SessionExtension.after_commit is deprecated", + "SessionExtension.after_begin is deprecated", + "SessionExtension.after_attach is deprecated", + "SessionExtension.before_flush is deprecated", + "SessionExtension.after_flush is deprecated", + "SessionExtension.after_flush_postexec is deprecated", + "SessionExtension.after_rollback is deprecated", + "SessionExtension.after_bulk_update is deprecated", + "SessionExtension.after_bulk_delete is deprecated", + ): + sess = create_session( + autocommit=False, extension=MyExt(), bind=testing.db + ) + sess.connection() + assert log == ["after_begin"] + sess.close() + + def test_multiple_extensions(self): + User, users = self.classes.User, self.tables.users + + log = [] + + class MyExt1(SessionExtension): + def before_commit(self, session): + log.append("before_commit_one") + + class MyExt2(SessionExtension): + def before_commit(self, session): + log.append("before_commit_two") + + mapper(User, users) + with testing.expect_deprecated( + "SessionExtension is deprecated in favor of " "the SessionEvents", + "SessionExtension.before_commit is deprecated", + ): + sess = create_session(extension=[MyExt1(), MyExt2()]) + u = User(name="u1") + sess.add(u) + sess.flush() + assert log == ["before_commit_one", "before_commit_two"] + + def test_unnecessary_methods_not_evented(self): + class MyExtension(SessionExtension): + def before_commit(self, session): + pass + + with testing.expect_deprecated( + "SessionExtension is deprecated in favor of " "the SessionEvents", + "SessionExtension.before_commit is deprecated.", + ): + s = Session(extension=MyExtension()) + assert not s.dispatch.after_commit + assert len(s.dispatch.before_commit) == 1 + + +class DeprecatedAttributeExtensionTest1(fixtures.ORMTest): + def test_extension_commit_attr(self): + """test that an extension which commits attribute history + maintains the end-result history. + + This won't work in conjunction with some unitofwork extensions. - TODO """ - session = create_session() + class Foo(fixtures.BasicEntity): + pass - def test_list(self): - """Query.list() + class Bar(fixtures.BasicEntity): + pass - users = session.query(User).list() + class ReceiveEvents(AttributeExtension): + def __init__(self, key): + self.key = key + + def append(self, state, child, initiator): + if commit: + state._commit_all(state.dict) + return child + + def remove(self, state, child, initiator): + if commit: + state._commit_all(state.dict) + return child + + def set(self, state, child, oldchild, initiator): + if commit: + state._commit_all(state.dict) + return child + + instrumentation.register_class(Foo) + instrumentation.register_class(Bar) + + b1, b2, b3, b4 = Bar(id="b1"), Bar(id="b2"), Bar(id="b3"), Bar(id="b4") + + def loadcollection(state, passive): + if passive is attributes.PASSIVE_NO_FETCH: + return attributes.PASSIVE_NO_RESULT + return [b1, b2] + + def loadscalar(state, passive): + if passive is attributes.PASSIVE_NO_FETCH: + return attributes.PASSIVE_NO_RESULT + return b2 + + with testing.expect_deprecated( + "AttributeExtension.append is deprecated.", + "AttributeExtension.remove is deprecated.", + "AttributeExtension.set is deprecated.", + ): + attributes.register_attribute( + Foo, + "bars", + uselist=True, + useobject=True, + callable_=loadcollection, + extension=[ReceiveEvents("bars")], + ) - """ + with testing.expect_deprecated( + "AttributeExtension.append is deprecated.", + "AttributeExtension.remove is deprecated.", + "AttributeExtension.set is deprecated.", + ): + attributes.register_attribute( + Foo, + "bar", + uselist=False, + useobject=True, + callable_=loadscalar, + extension=[ReceiveEvents("bar")], + ) - User = self.classes.User + with testing.expect_deprecated( + "AttributeExtension.append is deprecated.", + "AttributeExtension.remove is deprecated.", + "AttributeExtension.set is deprecated.", + ): + attributes.register_attribute( + Foo, + "scalar", + uselist=False, + useobject=False, + extension=[ReceiveEvents("scalar")], + ) - session = create_session() + def create_hist(): + def hist(key, fn, *arg): + attributes.instance_state(f1)._commit_all( + attributes.instance_dict(f1) + ) + fn(*arg) + histories.append(attributes.get_history(f1, key)) + + f1 = Foo() + hist("bars", f1.bars.append, b3) + hist("bars", f1.bars.append, b4) + hist("bars", f1.bars.remove, b2) + hist("bar", setattr, f1, "bar", b3) + hist("bar", setattr, f1, "bar", None) + hist("bar", setattr, f1, "bar", b4) + hist("scalar", setattr, f1, "scalar", 5) + hist("scalar", setattr, f1, "scalar", None) + hist("scalar", setattr, f1, "scalar", 4) + + histories = [] + commit = False + create_hist() + without_commit = list(histories) + histories[:] = [] + commit = True + create_hist() + with_commit = histories + for without, with_ in zip(without_commit, with_commit): + woc = without + wic = with_ + eq_(woc, wic) + + def test_extension_lazyload_assertion(self): + class Foo(fixtures.BasicEntity): + pass - users = session.query(User).all() - assert len(users) == 4 + class Bar(fixtures.BasicEntity): + pass - def test_scalar(self): - """Query.scalar() + class ReceiveEvents(AttributeExtension): + def append(self, state, child, initiator): + state.obj().bars + return child + + def remove(self, state, child, initiator): + state.obj().bars + return child + + def set(self, state, child, oldchild, initiator): + return child + + instrumentation.register_class(Foo) + instrumentation.register_class(Bar) + + bar1, bar2, bar3 = [Bar(id=1), Bar(id=2), Bar(id=3)] + + def func1(state, passive): + if passive is attributes.PASSIVE_NO_FETCH: + return attributes.PASSIVE_NO_RESULT + + return [bar1, bar2, bar3] + + with testing.expect_deprecated( + "AttributeExtension.append is deprecated.", + "AttributeExtension.remove is deprecated.", + "AttributeExtension.set is deprecated.", + ): + attributes.register_attribute( + Foo, + "bars", + uselist=True, + callable_=func1, + useobject=True, + extension=[ReceiveEvents()], + ) + attributes.register_attribute( + Bar, "foos", uselist=True, useobject=True, backref="bars" + ) - user = session.query(User).filter(User.id==1).scalar() + x = Foo() + assert_raises(AssertionError, Bar(id=4).foos.append, x) - """ + x.bars + b = Bar(id=4) + b.foos.append(x) + attributes.instance_state(x)._expire_attributes( + attributes.instance_dict(x), ["bars"] + ) + assert_raises(AssertionError, b.foos.remove, x) - User = self.classes.User + def test_scalar_listener(self): - session = create_session() + # listeners on ScalarAttributeImpl aren't used normally. test that + # they work for the benefit of user extensions - user = session.query(User).filter(User.id == 1).first() - assert user.id == 1 + class Foo(object): - def test_select(self): - r"""Query.select(arg=None, \**kwargs) + pass - users = session.query(User).select(users_table.c.name != None) + results = [] + + class ReceiveEvents(AttributeExtension): + def append(self, state, child, initiator): + assert False + + def remove(self, state, child, initiator): + results.append(("remove", state.obj(), child)) + + def set(self, state, child, oldchild, initiator): + results.append(("set", state.obj(), child, oldchild)) + return child + + instrumentation.register_class(Foo) + with testing.expect_deprecated( + "AttributeExtension.append is deprecated.", + "AttributeExtension.remove is deprecated.", + "AttributeExtension.set is deprecated.", + ): + attributes.register_attribute( + Foo, + "x", + uselist=False, + useobject=False, + extension=ReceiveEvents(), + ) - """ + f = Foo() + f.x = 5 + f.x = 17 + del f.x + + eq_( + results, + [ + ("set", f, 5, attributes.NEVER_SET), + ("set", f, 17, 5), + ("remove", f, 17), + ], + ) - User = self.classes.User + def test_cascading_extensions(self): + t1 = Table( + "t1", + MetaData(), + Column("id", Integer, primary_key=True), + Column("type", String(40)), + Column("data", String(50)), + ) - session = create_session() + ext_msg = [] - users = session.query(User).filter(User.name != None).all() # noqa - assert len(users) == 4 + class Ex1(AttributeExtension): + def set(self, state, value, oldvalue, initiator): + ext_msg.append("Ex1 %r" % value) + return "ex1" + value - def test_select_by(self): - r"""Query.select_by(\*args, \**params) + class Ex2(AttributeExtension): + def set(self, state, value, oldvalue, initiator): + ext_msg.append("Ex2 %r" % value) + return "ex2" + value - users = session.query(User).select_by(name='fred') + class A(fixtures.BasicEntity): + pass - # 0.3 magic join on \*_by methods - users = session.query(User).select_by(email_address='fred@the.fred') + class B(A): + pass - """ + class C(B): + pass - User, Address = self.classes.User, self.classes.Address + with testing.expect_deprecated( + "AttributeExtension is deprecated in favor of the " + "AttributeEvents listener interface. " + "The column_property.extension parameter" + ): + mapper( + A, + t1, + polymorphic_on=t1.c.type, + polymorphic_identity="a", + properties={ + "data": column_property(t1.c.data, extension=Ex1()) + }, + ) + mapper(B, polymorphic_identity="b", inherits=A) + with testing.expect_deprecated( + "AttributeExtension is deprecated in favor of the " + "AttributeEvents listener interface. " + "The column_property.extension parameter" + ): + mapper( + C, + polymorphic_identity="c", + inherits=B, + properties={ + "data": column_property(t1.c.data, extension=Ex2()) + }, + ) - session = create_session() + with testing.expect_deprecated( + "AttributeExtension.set is deprecated. " + ): + configure_mappers() + + a1 = A(data="a1") + b1 = B(data="b1") + c1 = C(data="c1") + + eq_(a1.data, "ex1a1") + eq_(b1.data, "ex1b1") + eq_(c1.data, "ex2c1") + + a1.data = "a2" + b1.data = "b2" + c1.data = "c2" + eq_(a1.data, "ex1a2") + eq_(b1.data, "ex1b2") + eq_(c1.data, "ex2c2") + + eq_( + ext_msg, + [ + "Ex1 'a1'", + "Ex1 'b1'", + "Ex2 'c1'", + "Ex1 'a2'", + "Ex1 'b2'", + "Ex2 'c2'", + ], + ) - users = session.query(User).filter_by(name="fred").all() - assert len(users) == 1 - users = session.query(User).filter(User.name == "fred").all() - assert len(users) == 1 +class DeprecatedOptionAllTest(OptionsPathTest, _fixtures.FixtureTest): + run_inserts = "once" + run_deletes = None + + def _mapper_fixture_one(self): + users, User, addresses, Address, orders, Order = ( + self.tables.users, + self.classes.User, + self.tables.addresses, + self.classes.Address, + self.tables.orders, + self.classes.Order, + ) + keywords, items, item_keywords, Keyword, Item = ( + self.tables.keywords, + self.tables.items, + self.tables.item_keywords, + self.classes.Keyword, + self.classes.Item, + ) + mapper( + User, + users, + properties={ + "addresses": relationship(Address), + "orders": relationship(Order), + }, + ) + mapper(Address, addresses) + mapper( + Order, + orders, + properties={ + "items": relationship(Item, secondary=self.tables.order_items) + }, + ) + mapper( + Keyword, + keywords, + properties={ + "keywords": column_property(keywords.c.name + "some keyword") + }, + ) + mapper( + Item, + items, + properties=dict( + keywords=relationship(Keyword, secondary=item_keywords) + ), + ) - users = ( - session.query(User) - .join("addresses") - .filter_by(email_address="fred@the.fred") - ).all() - assert len(users) == 1 + def _assert_eager_with_entity_exception( + self, entity_list, options, message + ): + assert_raises_message( + sa.exc.ArgumentError, + message, + create_session().query(*entity_list).options, + *options + ) - users = ( - session.query(User) - .filter( - User.addresses.any(Address.email_address == "fred@the.fred") + def test_option_against_nonexistent_twolevel_all(self): + self._mapper_fixture_one() + Item = self.classes.Item + with testing.expect_deprecated( + r"The joinedload_all\(\) function is deprecated, and " + "will be removed in a future release. " + r"Please use method chaining with joinedload\(\)" + ): + self._assert_eager_with_entity_exception( + [Item], + (joinedload_all("keywords.foo"),), + r"Can't find property named 'foo' on the mapped entity " + r"Mapper\|Keyword\|keywords in this Query.", ) - .all() + + def test_all_path_vs_chained(self): + self._mapper_fixture_one() + User = self.classes.User + Order = self.classes.Order + Item = self.classes.Item + + with testing.expect_deprecated( + r"The joinedload_all\(\) function is deprecated, and " + "will be removed in a future release. " + r"Please use method chaining with joinedload\(\)" + ): + l1 = joinedload_all("orders.items.keywords") + + sess = Session() + q = sess.query(User) + self._assert_path_result( + l1, + q, + [ + (User, "orders"), + (User, "orders", Order, "items"), + (User, "orders", Order, "items", Item, "keywords"), + ], ) - assert len(users) == 1 - def test_selectfirst(self): - r"""Query.selectfirst(arg=None, \**kwargs) + l2 = joinedload("orders").joinedload("items").joinedload("keywords") + self._assert_path_result( + l2, + q, + [ + (User, "orders"), + (User, "orders", Order, "items"), + (User, "orders", Order, "items", Item, "keywords"), + ], + ) - bounced = session.query(Address).selectfirst( - addresses_table.c.bounces > 0) + def test_subqueryload_mapper_order_by(self): + users, User, Address, addresses = ( + self.tables.users, + self.classes.User, + self.classes.Address, + self.tables.addresses, + ) - """ + mapper(Address, addresses) + + with testing.expect_deprecated( + ".*Mapper.order_by parameter is deprecated" + ): + mapper( + User, + users, + properties={ + "addresses": relationship( + Address, lazy="subquery", order_by=addresses.c.id + ) + }, + order_by=users.c.id.desc(), + ) - Address = self.classes.Address + sess = create_session() + q = sess.query(User) - session = create_session() + result = q.limit(2).all() + eq_(result, list(reversed(self.static.user_address_result[2:4]))) - bounced = session.query(Address).filter(Address.bounces > 0).first() - assert bounced.bounces > 0 + def test_selectinload_mapper_order_by(self): + users, User, Address, addresses = ( + self.tables.users, + self.classes.User, + self.classes.Address, + self.tables.addresses, + ) - def test_selectfirst_by(self): - r"""Query.selectfirst_by(\*args, \**params) + mapper(Address, addresses) + with testing.expect_deprecated( + ".*Mapper.order_by parameter is deprecated" + ): + mapper( + User, + users, + properties={ + "addresses": relationship( + Address, lazy="selectin", order_by=addresses.c.id + ) + }, + order_by=users.c.id.desc(), + ) - onebounce = session.query(Address).selectfirst_by(bounces=1) + sess = create_session() + q = sess.query(User) - # 0.3 magic join on *_by methods - onebounce_user = session.query(User).selectfirst_by(bounces=1) + result = q.limit(2).all() + eq_(result, list(reversed(self.static.user_address_result[2:4]))) - """ + def test_join_mapper_order_by(self): + """test that mapper-level order_by is adapted to a selectable.""" - User, Address = self.classes.User, self.classes.Address + User, users = self.classes.User, self.tables.users - session = create_session() + with testing.expect_deprecated( + ".*Mapper.order_by parameter is deprecated" + ): + mapper(User, users, order_by=users.c.id) - onebounce = session.query(Address).filter_by(bounces=1).first() - assert onebounce.bounces == 1 + sel = users.select(users.c.id.in_([7, 8])) + sess = create_session() - onebounce_user = ( - session.query(User).join("addresses").filter_by(bounces=1) - ).first() - assert onebounce_user.name == "jack" + eq_( + sess.query(User).select_entity_from(sel).all(), + [User(name="jack", id=7), User(name="ed", id=8)], + ) - onebounce_user = ( - session.query(User).join("addresses").filter(Address.bounces == 1) - ).first() - assert onebounce_user.name == "jack" + def test_defer_addtl_attrs(self): + users, User, Address, addresses = ( + self.tables.users, + self.classes.User, + self.classes.Address, + self.tables.addresses, + ) - onebounce_user = ( - session.query(User) - .filter(User.addresses.any(Address.bounces == 1)) - .first() + mapper(Address, addresses) + mapper( + User, + users, + properties={ + "addresses": relationship( + Address, lazy="selectin", order_by=addresses.c.id + ) + }, ) - assert onebounce_user.name == "jack" - def test_selectone(self): - r"""Query.selectone(arg=None, \**kwargs) + sess = create_session() - ed = session.query(User).selectone(users_table.c.name == 'ed') + with testing.expect_deprecated( + r"The \*addl_attrs on orm.defer is deprecated. " + "Please use method chaining" + ): + sess.query(User).options(defer("addresses", "email_address")) - """ + with testing.expect_deprecated( + r"The \*addl_attrs on orm.undefer is deprecated. " + "Please use method chaining" + ): + sess.query(User).options(undefer("addresses", "email_address")) + +class LegacyLockModeTest(_fixtures.FixtureTest): + run_inserts = None + + @classmethod + def setup_mappers(cls): + User, users = cls.classes.User, cls.tables.users + mapper(User, users) + + def _assert_legacy(self, arg, read=False, nowait=False): User = self.classes.User + s = Session() - session = create_session() + with testing.expect_deprecated( + r"The Query.with_lockmode\(\) method is deprecated" + ): + q = s.query(User).with_lockmode(arg) + sel = q._compile_context().statement - ed = session.query(User).filter(User.name == "jack").one() + if arg is None: + assert q._for_update_arg is None + assert sel._for_update_arg is None + return - def test_selectone_by(self): - """Query.selectone_by + assert q._for_update_arg.read is read + assert q._for_update_arg.nowait is nowait - ed = session.query(User).selectone_by(name='ed') + assert sel._for_update_arg.read is read + assert sel._for_update_arg.nowait is nowait - # 0.3 magic join on *_by methods - ed = session.query(User).selectone_by(email_address='ed@foo.bar') + def test_false_legacy(self): + self._assert_legacy(None) - """ + def test_plain_legacy(self): + self._assert_legacy("update") - User, Address = self.classes.User, self.classes.Address + def test_nowait_legacy(self): + self._assert_legacy("update_nowait", nowait=True) - session = create_session() + def test_read_legacy(self): + self._assert_legacy("read", read=True) + + def test_unknown_legacy_lock_mode(self): + User = self.classes.User + sess = Session() + with testing.expect_deprecated( + r"The Query.with_lockmode\(\) method is deprecated" + ): + assert_raises_message( + exc.ArgumentError, + "Unknown with_lockmode argument: 'unknown_mode'", + sess.query(User.id).with_lockmode, + "unknown_mode", + ) - ed = session.query(User).filter_by(name="jack").one() - ed = session.query(User).filter(User.name == "jack").one() +class InstrumentationTest(fixtures.ORMTest): + def test_dict_subclass4(self): + # tests #2654 + with testing.expect_deprecated( + r"The collection.converter\(\) handler is deprecated and will " + "be removed in a future release. Please refer to the " + "AttributeEvents" + ): - ed = ( - session.query(User) - .join("addresses") - .filter(Address.email_address == "ed@foo.bar") - .one() - ) + class MyDict(collections.MappedCollection): + def __init__(self): + super(MyDict, self).__init__(lambda value: "k%d" % value) + + @collection.converter + def _convert(self, dictlike): + for key, value in dictlike.items(): + yield value + 5 - ed = ( - session.query(User) - .filter(User.addresses.any(Address.email_address == "ed@foo.bar")) - .one() + class Foo(object): + pass + + instrumentation.register_class(Foo) + d = attributes.register_attribute( + Foo, "attr", uselist=True, typecallable=MyDict, useobject=True ) - def test_select_statement(self): - r"""Query.select_statement(statement, \**params) + f = Foo() + f.attr = {"k1": 1, "k2": 2} - users = session.query(User).select_statement(users_table.select()) + eq_(f.attr, {"k7": 7, "k6": 6}) - """ + def test_name_setup(self): + with testing.expect_deprecated( + r"The collection.converter\(\) handler is deprecated and will " + "be removed in a future release. Please refer to the " + "AttributeEvents" + ): - User, users_table = self.classes.User, self.tables.users_table + class Base(object): + @collection.iterator + def base_iterate(self, x): + return "base_iterate" - session = create_session() + @collection.appender + def base_append(self, x): + return "base_append" - users = session.query(User).from_statement(users_table.select()).all() - assert len(users) == 4 + @collection.converter + def base_convert(self, x): + return "base_convert" - def test_select_text(self): - r"""Query.select_text(text, \**params) + @collection.remover + def base_remove(self, x): + return "base_remove" - users = session.query(User).select_text('SELECT * FROM users_table') + from sqlalchemy.orm.collections import _instrument_class - """ + _instrument_class(Base) - User = self.classes.User + eq_(Base._sa_remover(Base(), 5), "base_remove") + eq_(Base._sa_appender(Base(), 5), "base_append") + eq_(Base._sa_iterator(Base(), 5), "base_iterate") + eq_(Base._sa_converter(Base(), 5), "base_convert") - session = create_session() + with testing.expect_deprecated( + r"The collection.converter\(\) handler is deprecated and will " + "be removed in a future release. Please refer to the " + "AttributeEvents" + ): - users = ( - session.query(User).from_statement( - text("SELECT * FROM users_table") - ) - ).all() - assert len(users) == 4 + class Sub(Base): + @collection.converter + def base_convert(self, x): + return "sub_convert" - def test_select_whereclause(self): - r"""Query.select_whereclause(whereclause=None, params=None, \**kwargs) + @collection.remover + def sub_remove(self, x): + return "sub_remove" + _instrument_class(Sub) - users = session,query(User).select_whereclause(users.c.name=='ed') - users = session.query(User).select_whereclause("name='ed'") + eq_(Sub._sa_appender(Sub(), 5), "base_append") + eq_(Sub._sa_remover(Sub(), 5), "sub_remove") + eq_(Sub._sa_iterator(Sub(), 5), "base_iterate") + eq_(Sub._sa_converter(Sub(), 5), "sub_convert") - """ + def test_link_event(self): + canary = [] - User = self.classes.User + with testing.expect_deprecated( + r"The collection.linker\(\) handler is deprecated and will " + "be removed in a future release. Please refer to the " + "AttributeEvents" + ): - session = create_session() + class Collection(list): + @collection.linker + def _on_link(self, obj): + canary.append(obj) + + class Foo(object): + pass + + instrumentation.register_class(Foo) + attributes.register_attribute( + Foo, "attr", uselist=True, typecallable=Collection, useobject=True + ) + + f1 = Foo() + f1.attr.append(3) - users = session.query(User).filter(User.name == "ed").all() - assert len(users) == 1 and users[0].name == "ed" + eq_(canary, [f1.attr._sa_adapter]) + adapter_1 = f1.attr._sa_adapter - users = session.query(User).filter(text("name='ed'")).all() - assert len(users) == 1 and users[0].name == "ed" + l2 = Collection() + f1.attr = l2 + eq_(canary, [adapter_1, f1.attr._sa_adapter, None]) diff --git a/test/orm/test_eager_relations.py b/test/orm/test_eager_relations.py index ea8ae764d..fd272b181 100644 --- a/test/orm/test_eager_relations.py +++ b/test/orm/test_eager_relations.py @@ -21,7 +21,6 @@ from sqlalchemy.orm import create_session from sqlalchemy.orm import defaultload from sqlalchemy.orm import deferred from sqlalchemy.orm import joinedload -from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import lazyload from sqlalchemy.orm import Load from sqlalchemy.orm import load_only @@ -1437,7 +1436,7 @@ class EagerTest(_fixtures.FixtureTest, testing.AssertsCompiledSQL): self.assert_compile( sess.query(User) - .options(joinedload_all("orders.address")) + .options(joinedload("orders").joinedload("address")) .limit(10), "SELECT anon_1.users_id AS anon_1_users_id, " "anon_1.users_name AS anon_1_users_name, " @@ -1459,7 +1458,8 @@ class EagerTest(_fixtures.FixtureTest, testing.AssertsCompiledSQL): self.assert_compile( sess.query(User).options( - joinedload_all("orders.items"), joinedload("orders.address") + joinedload("orders").joinedload("items"), + joinedload("orders").joinedload("address"), ), "SELECT users.id AS users_id, users.name AS users_name, " "items_1.id AS items_1_id, " @@ -2391,7 +2391,9 @@ class EagerTest(_fixtures.FixtureTest, testing.AssertsCompiledSQL): sess.query(User) .join(User.orders) .join(Order.items) - .options(joinedload_all("orders.items.keywords")) + .options( + joinedload("orders").joinedload("items").joinedload("keywords") + ) ) # here, the eager join for keywords can catch onto @@ -2583,7 +2585,9 @@ class EagerTest(_fixtures.FixtureTest, testing.AssertsCompiledSQL): self.assert_compile( sess.query(User).options( - joinedload_all(User.orders, Order.items, innerjoin=True) + joinedload(User.orders, innerjoin=True).joinedload( + Order.items, innerjoin=True + ) ), "SELECT users.id AS users_id, users.name AS users_name, " "items_1.id AS items_1_id, " @@ -3230,7 +3234,7 @@ class SubqueryAliasingTest(fixtures.MappedTest, testing.AssertsCompiledSQL): self.assert_compile( create_session() .query(A) - .options(joinedload_all("bs")) + .options(joinedload("bs")) .order_by(A.summation) .limit(50), "SELECT anon_1.anon_2 AS anon_1_anon_2, anon_1.a_id " @@ -3253,7 +3257,7 @@ class SubqueryAliasingTest(fixtures.MappedTest, testing.AssertsCompiledSQL): self.assert_compile( create_session() .query(A) - .options(joinedload_all("bs")) + .options(joinedload("bs")) .order_by(A.summation.desc()) .limit(50), "SELECT anon_1.anon_2 AS anon_1_anon_2, anon_1.a_id " @@ -3278,7 +3282,7 @@ class SubqueryAliasingTest(fixtures.MappedTest, testing.AssertsCompiledSQL): self.assert_compile( create_session() .query(A) - .options(joinedload_all("bs")) + .options(joinedload("bs")) .order_by(A.summation) .limit(50), "SELECT anon_1.anon_2 AS anon_1_anon_2, anon_1.a_id " @@ -3307,7 +3311,7 @@ class SubqueryAliasingTest(fixtures.MappedTest, testing.AssertsCompiledSQL): self.assert_compile( create_session() .query(A) - .options(joinedload_all("bs")) + .options(joinedload("bs")) .order_by(cp) .limit(50), "SELECT anon_1.a_id AS anon_1_a_id, anon_1.anon_2 " @@ -3334,7 +3338,7 @@ class SubqueryAliasingTest(fixtures.MappedTest, testing.AssertsCompiledSQL): self.assert_compile( create_session() .query(A) - .options(joinedload_all("bs")) + .options(joinedload("bs")) .order_by(cp) .limit(50), "SELECT anon_1.a_id AS anon_1_a_id, anon_1.foo " @@ -3361,7 +3365,7 @@ class SubqueryAliasingTest(fixtures.MappedTest, testing.AssertsCompiledSQL): self.assert_compile( create_session() .query(A) - .options(joinedload_all("bs")) + .options(joinedload("bs")) .order_by(~cp) .limit(50), "SELECT anon_1.a_id AS anon_1_a_id, anon_1.anon_2 " @@ -3457,7 +3461,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): a2 = u1.addresses[0] a2.email_address = "foo" sess.query(User).options( - joinedload_all("addresses.dingaling") + joinedload("addresses").joinedload("dingaling") ).filter_by(id=8).all() assert u1.addresses[-1] is a1 for a in u1.addresses: @@ -3475,9 +3479,9 @@ class LoadOnExistingTest(_fixtures.FixtureTest): u1.orders o1 = Order() u1.orders.append(o1) - sess.query(User).options(joinedload_all("orders.items")).filter_by( - id=7 - ).all() + sess.query(User).options( + joinedload("orders").joinedload("items") + ).filter_by(id=7).all() for o in u1.orders: if o is not o1: assert "items" in o.__dict__ @@ -3494,7 +3498,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): .one() ) sess.query(User).filter_by(id=8).options( - joinedload_all("addresses.dingaling") + joinedload("addresses").joinedload("dingaling") ).first() assert "dingaling" in u1.addresses[0].__dict__ @@ -3508,7 +3512,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): .one() ) sess.query(User).filter_by(id=7).options( - joinedload_all("orders.items") + joinedload("orders").joinedload("items") ).first() assert "items" in u1.orders[0].__dict__ diff --git a/test/orm/test_events.py b/test/orm/test_events.py index bb1a935de..af5191569 100644 --- a/test/orm/test_events.py +++ b/test/orm/test_events.py @@ -6,7 +6,6 @@ from sqlalchemy import testing from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import attributes from sqlalchemy.orm import class_mapper -from sqlalchemy.orm import column_property from sqlalchemy.orm import configure_mappers from sqlalchemy.orm import create_session from sqlalchemy.orm import events @@ -2387,470 +2386,6 @@ class SessionLifecycleEventsTest(_RemoveListeners, _fixtures.FixtureTest): ) -class MapperExtensionTest(_fixtures.FixtureTest): - - """Superseded by MapperEventsTest - test backwards - compatibility of MapperExtension.""" - - run_inserts = None - - def extension(self): - methods = [] - - class Ext(sa.orm.MapperExtension): - def instrument_class(self, mapper, cls): - methods.append("instrument_class") - return sa.orm.EXT_CONTINUE - - def init_instance( - self, mapper, class_, oldinit, instance, args, kwargs - ): - methods.append("init_instance") - return sa.orm.EXT_CONTINUE - - def init_failed( - self, mapper, class_, oldinit, instance, args, kwargs - ): - methods.append("init_failed") - return sa.orm.EXT_CONTINUE - - def reconstruct_instance(self, mapper, instance): - methods.append("reconstruct_instance") - return sa.orm.EXT_CONTINUE - - def before_insert(self, mapper, connection, instance): - methods.append("before_insert") - return sa.orm.EXT_CONTINUE - - def after_insert(self, mapper, connection, instance): - methods.append("after_insert") - return sa.orm.EXT_CONTINUE - - def before_update(self, mapper, connection, instance): - methods.append("before_update") - return sa.orm.EXT_CONTINUE - - def after_update(self, mapper, connection, instance): - methods.append("after_update") - return sa.orm.EXT_CONTINUE - - def before_delete(self, mapper, connection, instance): - methods.append("before_delete") - return sa.orm.EXT_CONTINUE - - def after_delete(self, mapper, connection, instance): - methods.append("after_delete") - return sa.orm.EXT_CONTINUE - - return Ext, methods - - def test_basic(self): - """test that common user-defined methods get called.""" - - User, users = self.classes.User, self.tables.users - - Ext, methods = self.extension() - - mapper(User, users, extension=Ext()) - sess = create_session() - u = User(name="u1") - sess.add(u) - sess.flush() - u = sess.query(User).populate_existing().get(u.id) - sess.expunge_all() - u = sess.query(User).get(u.id) - u.name = "u1 changed" - sess.flush() - sess.delete(u) - sess.flush() - eq_( - methods, - [ - "instrument_class", - "init_instance", - "before_insert", - "after_insert", - "reconstruct_instance", - "before_update", - "after_update", - "before_delete", - "after_delete", - ], - ) - - def test_inheritance(self): - users, addresses, User = ( - self.tables.users, - self.tables.addresses, - self.classes.User, - ) - - Ext, methods = self.extension() - - class AdminUser(User): - pass - - mapper(User, users, extension=Ext()) - mapper( - AdminUser, - addresses, - inherits=User, - properties={"address_id": addresses.c.id}, - ) - - sess = create_session() - am = AdminUser(name="au1", email_address="au1@e1") - sess.add(am) - sess.flush() - am = sess.query(AdminUser).populate_existing().get(am.id) - sess.expunge_all() - am = sess.query(AdminUser).get(am.id) - am.name = "au1 changed" - sess.flush() - sess.delete(am) - sess.flush() - eq_( - methods, - [ - "instrument_class", - "instrument_class", - "init_instance", - "before_insert", - "after_insert", - "reconstruct_instance", - "before_update", - "after_update", - "before_delete", - "after_delete", - ], - ) - - def test_before_after_only_collection(self): - """before_update is called on parent for collection modifications, - after_update is called even if no columns were updated. - - """ - - keywords, items, item_keywords, Keyword, Item = ( - self.tables.keywords, - self.tables.items, - self.tables.item_keywords, - self.classes.Keyword, - self.classes.Item, - ) - - Ext1, methods1 = self.extension() - Ext2, methods2 = self.extension() - - mapper( - Item, - items, - extension=Ext1(), - properties={ - "keywords": relationship(Keyword, secondary=item_keywords) - }, - ) - mapper(Keyword, keywords, extension=Ext2()) - - sess = create_session() - i1 = Item(description="i1") - k1 = Keyword(name="k1") - sess.add(i1) - sess.add(k1) - sess.flush() - eq_( - methods1, - [ - "instrument_class", - "init_instance", - "before_insert", - "after_insert", - ], - ) - eq_( - methods2, - [ - "instrument_class", - "init_instance", - "before_insert", - "after_insert", - ], - ) - - del methods1[:] - del methods2[:] - i1.keywords.append(k1) - sess.flush() - eq_(methods1, ["before_update", "after_update"]) - eq_(methods2, []) - - def test_inheritance_with_dupes(self): - """Inheritance with the same extension instance on both mappers.""" - - users, addresses, User = ( - self.tables.users, - self.tables.addresses, - self.classes.User, - ) - - Ext, methods = self.extension() - - class AdminUser(User): - pass - - ext = Ext() - mapper(User, users, extension=ext) - mapper( - AdminUser, - addresses, - inherits=User, - extension=ext, - properties={"address_id": addresses.c.id}, - ) - - sess = create_session() - am = AdminUser(name="au1", email_address="au1@e1") - sess.add(am) - sess.flush() - am = sess.query(AdminUser).populate_existing().get(am.id) - sess.expunge_all() - am = sess.query(AdminUser).get(am.id) - am.name = "au1 changed" - sess.flush() - sess.delete(am) - sess.flush() - eq_( - methods, - [ - "instrument_class", - "instrument_class", - "init_instance", - "before_insert", - "after_insert", - "reconstruct_instance", - "before_update", - "after_update", - "before_delete", - "after_delete", - ], - ) - - def test_unnecessary_methods_not_evented(self): - users = self.tables.users - - class MyExtension(sa.orm.MapperExtension): - def before_insert(self, mapper, connection, instance): - pass - - class Foo(object): - pass - - m = mapper(Foo, users, extension=MyExtension()) - assert not m.class_manager.dispatch.load - assert not m.dispatch.before_update - assert len(m.dispatch.before_insert) == 1 - - -class AttributeExtensionTest(fixtures.MappedTest): - @classmethod - def define_tables(cls, metadata): - Table( - "t1", - metadata, - Column("id", Integer, primary_key=True), - Column("type", String(40)), - Column("data", String(50)), - ) - - def test_cascading_extensions(self): - t1 = self.tables.t1 - - ext_msg = [] - - class Ex1(sa.orm.AttributeExtension): - def set(self, state, value, oldvalue, initiator): - ext_msg.append("Ex1 %r" % value) - return "ex1" + value - - class Ex2(sa.orm.AttributeExtension): - def set(self, state, value, oldvalue, initiator): - ext_msg.append("Ex2 %r" % value) - return "ex2" + value - - class A(fixtures.BasicEntity): - pass - - class B(A): - pass - - class C(B): - pass - - mapper( - A, - t1, - polymorphic_on=t1.c.type, - polymorphic_identity="a", - properties={"data": column_property(t1.c.data, extension=Ex1())}, - ) - mapper(B, polymorphic_identity="b", inherits=A) - mapper( - C, - polymorphic_identity="c", - inherits=B, - properties={"data": column_property(t1.c.data, extension=Ex2())}, - ) - - a1 = A(data="a1") - b1 = B(data="b1") - c1 = C(data="c1") - - eq_(a1.data, "ex1a1") - eq_(b1.data, "ex1b1") - eq_(c1.data, "ex2c1") - - a1.data = "a2" - b1.data = "b2" - c1.data = "c2" - eq_(a1.data, "ex1a2") - eq_(b1.data, "ex1b2") - eq_(c1.data, "ex2c2") - - eq_( - ext_msg, - [ - "Ex1 'a1'", - "Ex1 'b1'", - "Ex2 'c1'", - "Ex1 'a2'", - "Ex1 'b2'", - "Ex2 'c2'", - ], - ) - - -class SessionExtensionTest(_fixtures.FixtureTest): - run_inserts = None - - def test_extension(self): - User, users = self.classes.User, self.tables.users - - mapper(User, users) - log = [] - - class MyExt(sa.orm.session.SessionExtension): - def before_commit(self, session): - log.append("before_commit") - - def after_commit(self, session): - log.append("after_commit") - - def after_rollback(self, session): - log.append("after_rollback") - - def before_flush(self, session, flush_context, objects): - log.append("before_flush") - - def after_flush(self, session, flush_context): - log.append("after_flush") - - def after_flush_postexec(self, session, flush_context): - log.append("after_flush_postexec") - - def after_begin(self, session, transaction, connection): - log.append("after_begin") - - def after_attach(self, session, instance): - log.append("after_attach") - - def after_bulk_update(self, session, query, query_context, result): - log.append("after_bulk_update") - - def after_bulk_delete(self, session, query, query_context, result): - log.append("after_bulk_delete") - - sess = create_session(extension=MyExt()) - u = User(name="u1") - sess.add(u) - sess.flush() - assert log == [ - "after_attach", - "before_flush", - "after_begin", - "after_flush", - "after_flush_postexec", - "before_commit", - "after_commit", - ] - log = [] - sess = create_session(autocommit=False, extension=MyExt()) - u = User(name="u1") - sess.add(u) - sess.flush() - assert log == [ - "after_attach", - "before_flush", - "after_begin", - "after_flush", - "after_flush_postexec", - ] - log = [] - u.name = "ed" - sess.commit() - assert log == [ - "before_commit", - "before_flush", - "after_flush", - "after_flush_postexec", - "after_commit", - ] - log = [] - sess.commit() - assert log == ["before_commit", "after_commit"] - log = [] - sess.query(User).delete() - assert log == ["after_begin", "after_bulk_delete"] - log = [] - sess.query(User).update({"name": "foo"}) - assert log == ["after_bulk_update"] - log = [] - sess = create_session( - autocommit=False, extension=MyExt(), bind=testing.db - ) - sess.connection() - assert log == ["after_begin"] - sess.close() - - def test_multiple_extensions(self): - User, users = self.classes.User, self.tables.users - - log = [] - - class MyExt1(sa.orm.session.SessionExtension): - def before_commit(self, session): - log.append("before_commit_one") - - class MyExt2(sa.orm.session.SessionExtension): - def before_commit(self, session): - log.append("before_commit_two") - - mapper(User, users) - sess = create_session(extension=[MyExt1(), MyExt2()]) - u = User(name="u1") - sess.add(u) - sess.flush() - assert log == ["before_commit_one", "before_commit_two"] - - def test_unnecessary_methods_not_evented(self): - class MyExtension(sa.orm.session.SessionExtension): - def before_commit(self, session): - pass - - s = Session(extension=MyExtension()) - assert not s.dispatch.after_commit - assert len(s.dispatch.before_commit) == 1 - - class QueryEventsTest( _RemoveListeners, _fixtures.FixtureTest, AssertsCompiledSQL ): diff --git a/test/orm/test_froms.py b/test/orm/test_froms.py index fc2fb670c..c5e1d1485 100644 --- a/test/orm/test_froms.py +++ b/test/orm/test_froms.py @@ -26,7 +26,6 @@ from sqlalchemy.orm import contains_alias from sqlalchemy.orm import contains_eager from sqlalchemy.orm import create_session from sqlalchemy.orm import joinedload -from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import mapper from sqlalchemy.orm import relation from sqlalchemy.orm import relationship @@ -2551,22 +2550,6 @@ class SelectFromTest(QueryTest, AssertsCompiledSQL): ) eq_(q.all(), [("chuck",), ("ed",), ("fred",), ("jack",)]) - @testing.uses_deprecated("Mapper.order_by") - def test_join_mapper_order_by(self): - """test that mapper-level order_by is adapted to a selectable.""" - - User, users = self.classes.User, self.tables.users - - mapper(User, users, order_by=users.c.id) - - sel = users.select(users.c.id.in_([7, 8])) - sess = create_session() - - eq_( - sess.query(User).select_entity_from(sel).all(), - [User(name="jack", id=7), User(name="ed", id=8)], - ) - def test_differentiate_self_external(self): """test some different combinations of joining a table to a subquery of itself.""" @@ -2964,7 +2947,11 @@ class SelectFromTest(QueryTest, AssertsCompiledSQL): eq_( sess.query(User) .select_entity_from(sel) - .options(joinedload_all("orders.items.keywords")) + .options( + joinedload("orders") + .joinedload("items") + .joinedload("keywords") + ) .join("orders", "items", "keywords", aliased=True) .filter(Keyword.name.in_(["red", "big", "round"])) .all(), @@ -3425,7 +3412,7 @@ class ExternalColumnsTest(QueryTest): def go(): o1 = ( sess.query(Order) - .options(joinedload_all("address.user")) + .options(joinedload("address").joinedload("user")) .get(1) ) eq_(o1.address.user.count, 1) @@ -3437,7 +3424,7 @@ class ExternalColumnsTest(QueryTest): def go(): o1 = ( sess.query(Order) - .options(joinedload_all("address.user")) + .options(joinedload("address").joinedload("user")) .first() ) eq_(o1.address.user.count, 1) diff --git a/test/orm/test_generative.py b/test/orm/test_generative.py index abc666af1..97106cafc 100644 --- a/test/orm/test_generative.py +++ b/test/orm/test_generative.py @@ -73,7 +73,6 @@ class GenerativeQueryTest(fixtures.MappedTest): assert query[10:20][5] == orig[10:20][5] - @testing.uses_deprecated("Call to deprecated function apply_max") def test_aggregate(self): foo, Foo = self.tables.foo, self.classes.Foo diff --git a/test/orm/test_lockmode.py b/test/orm/test_lockmode.py index dd928f0db..e3653c6ac 100644 --- a/test/orm/test_lockmode.py +++ b/test/orm/test_lockmode.py @@ -11,54 +11,6 @@ from sqlalchemy.testing import eq_ from test.orm import _fixtures -class LegacyLockModeTest(_fixtures.FixtureTest): - run_inserts = None - - @classmethod - def setup_mappers(cls): - User, users = cls.classes.User, cls.tables.users - mapper(User, users) - - def _assert_legacy(self, arg, read=False, nowait=False): - User = self.classes.User - s = Session() - q = s.query(User).with_lockmode(arg) - sel = q._compile_context().statement - - if arg is None: - assert q._for_update_arg is None - assert sel._for_update_arg is None - return - - assert q._for_update_arg.read is read - assert q._for_update_arg.nowait is nowait - - assert sel._for_update_arg.read is read - assert sel._for_update_arg.nowait is nowait - - def test_false_legacy(self): - self._assert_legacy(None) - - def test_plain_legacy(self): - self._assert_legacy("update") - - def test_nowait_legacy(self): - self._assert_legacy("update_nowait", nowait=True) - - def test_read_legacy(self): - self._assert_legacy("read", read=True) - - def test_unknown_legacy_lock_mode(self): - User = self.classes.User - sess = Session() - assert_raises_message( - exc.ArgumentError, - "Unknown with_lockmode argument: 'unknown_mode'", - sess.query(User.id).with_lockmode, - "unknown_mode", - ) - - class ForUpdateTest(_fixtures.FixtureTest): @classmethod def setup_mappers(cls): diff --git a/test/orm/test_mapper.py b/test/orm/test_mapper.py index 47710792e..4e42e4f9a 100644 --- a/test/orm/test_mapper.py +++ b/test/orm/test_mapper.py @@ -18,7 +18,6 @@ from sqlalchemy.orm import attributes from sqlalchemy.orm import backref from sqlalchemy.orm import class_mapper from sqlalchemy.orm import column_property -from sqlalchemy.orm import comparable_property from sqlalchemy.orm import composite from sqlalchemy.orm import configure_mappers from sqlalchemy.orm import create_session @@ -555,7 +554,6 @@ class MapperTest(_fixtures.FixtureTest, AssertsCompiledSQL): (relationship, (Address,)), (composite, (MyComposite, "id", "name")), (synonym, "foo"), - (comparable_property, "foo"), ]: obj = constructor(info={"x": "y"}, *args) eq_(obj.info, {"x": "y"}) @@ -630,35 +628,12 @@ class MapperTest(_fixtures.FixtureTest, AssertsCompiledSQL): name = property(_get_name, _set_name) - def _uc_name(self): - if self._name is None: - return None - return self._name.upper() - - uc_name = property(_uc_name) - uc_name2 = property(_uc_name) - m = mapper(User, users) mapper(Address, addresses) - class UCComparator(sa.orm.PropComparator): - __hash__ = None - - def __eq__(self, other): - cls = self.prop.parent.class_ - col = getattr(cls, "name") - if other is None: - return col is None - else: - return sa.func.upper(col) == sa.func.upper(other) - m.add_property("_name", deferred(users.c.name)) m.add_property("name", synonym("_name")) m.add_property("addresses", relationship(Address)) - m.add_property("uc_name", sa.orm.comparable_property(UCComparator)) - m.add_property( - "uc_name2", sa.orm.comparable_property(UCComparator, User.uc_name2) - ) sess = create_session(autocommit=False) assert sess.query(User).get(7) @@ -671,8 +646,6 @@ class MapperTest(_fixtures.FixtureTest, AssertsCompiledSQL): len(self.static.user_address_result[0].addresses), ) eq_(u.name, "jack") - eq_(u.uc_name, "JACK") - eq_(u.uc_name2, "JACK") eq_(assert_col, [("get", "jack")], str(assert_col)) self.sql_count_(2, go) @@ -1410,52 +1383,6 @@ class MapperTest(_fixtures.FixtureTest, AssertsCompiledSQL): eq_(result, [self.static.user_result[0]]) - @testing.uses_deprecated("Mapper.order_by") - def test_cancel_order_by(self): - users, User = self.tables.users, self.classes.User - - mapper(User, users, order_by=users.c.name.desc()) - - assert ( - "order by users.name desc" - in str(create_session().query(User).statement).lower() - ) - assert ( - "order by" - not in str( - create_session().query(User).order_by(None).statement - ).lower() - ) - assert ( - "order by users.name asc" - in str( - create_session() - .query(User) - .order_by(User.name.asc()) - .statement - ).lower() - ) - - eq_( - create_session().query(User).all(), - [ - User(id=7, name="jack"), - User(id=9, name="fred"), - User(id=8, name="ed"), - User(id=10, name="chuck"), - ], - ) - - eq_( - create_session().query(User).order_by(User.name).all(), - [ - User(id=10, name="chuck"), - User(id=8, name="ed"), - User(id=9, name="fred"), - User(id=7, name="jack"), - ], - ) - # 'Raises a "expression evaluation not supported" error at prepare time @testing.fails_on("firebird", "FIXME: unknown") def test_function(self): @@ -1809,151 +1736,6 @@ class MapperTest(_fixtures.FixtureTest, AssertsCompiledSQL): ), ) - def test_comparable(self): - users = self.tables.users - - class extendedproperty(property): - attribute = 123 - - def method1(self): - return "method1" - - from sqlalchemy.orm.properties import ColumnProperty - - class UCComparator(ColumnProperty.Comparator): - __hash__ = None - - def method1(self): - return "uccmethod1" - - def method2(self, other): - return "method2" - - def __eq__(self, other): - cls = self.prop.parent.class_ - col = getattr(cls, "name") - if other is None: - return col is None - else: - return sa.func.upper(col) == sa.func.upper(other) - - def map_(with_explicit_property): - class User(object): - @extendedproperty - def uc_name(self): - if self.name is None: - return None - return self.name.upper() - - if with_explicit_property: - args = (UCComparator, User.uc_name) - else: - args = (UCComparator,) - mapper( - User, - users, - properties=dict(uc_name=sa.orm.comparable_property(*args)), - ) - return User - - for User in (map_(True), map_(False)): - sess = create_session() - sess.begin() - q = sess.query(User) - - assert hasattr(User, "name") - assert hasattr(User, "uc_name") - - eq_(User.uc_name.method1(), "method1") - eq_(User.uc_name.method2("x"), "method2") - - assert_raises_message( - AttributeError, - "Neither 'extendedproperty' object nor 'UCComparator' " - "object associated with User.uc_name has an attribute " - "'nonexistent'", - getattr, - User.uc_name, - "nonexistent", - ) - - # test compile - assert not isinstance(User.uc_name == "jack", bool) - u = q.filter(User.uc_name == "JACK").one() - - assert u.uc_name == "JACK" - assert u not in sess.dirty - - u.name = "some user name" - eq_(u.name, "some user name") - assert u in sess.dirty - eq_(u.uc_name, "SOME USER NAME") - - sess.flush() - sess.expunge_all() - - q = sess.query(User) - u2 = q.filter(User.name == "some user name").one() - u3 = q.filter(User.uc_name == "SOME USER NAME").one() - - assert u2 is u3 - - eq_(User.uc_name.attribute, 123) - sess.rollback() - - def test_comparable_column(self): - users, User = self.tables.users, self.classes.User - - class MyComparator(sa.orm.properties.ColumnProperty.Comparator): - __hash__ = None - - def __eq__(self, other): - # lower case comparison - return func.lower(self.__clause_element__()) == func.lower( - other - ) - - def intersects(self, other): - # non-standard comparator - return self.__clause_element__().op("&=")(other) - - mapper( - User, - users, - properties={ - "name": sa.orm.column_property( - users.c.name, comparator_factory=MyComparator - ) - }, - ) - - assert_raises_message( - AttributeError, - "Neither 'InstrumentedAttribute' object nor " - "'MyComparator' object associated with User.name has " - "an attribute 'nonexistent'", - getattr, - User.name, - "nonexistent", - ) - - eq_( - str( - (User.name == "ed").compile( - dialect=sa.engine.default.DefaultDialect() - ) - ), - "lower(users.name) = lower(:lower_1)", - ) - eq_( - str( - (User.name.intersects("ed")).compile( - dialect=sa.engine.default.DefaultDialect() - ) - ), - "users.name &= :name_1", - ) - def test_reentrant_compile(self): users, Address, addresses, User = ( self.tables.users, @@ -2776,7 +2558,11 @@ class DeepOptionsTest(_fixtures.FixtureTest): result = ( sess.query(User) .order_by(User.id) - .options(sa.orm.joinedload_all("orders.items.keywords")) + .options( + sa.orm.joinedload("orders") + .joinedload("items") + .joinedload("keywords") + ) ).all() def go(): @@ -2788,7 +2574,9 @@ class DeepOptionsTest(_fixtures.FixtureTest): result = ( sess.query(User).options( - sa.orm.subqueryload_all("orders.items.keywords") + sa.orm.subqueryload("orders") + .subqueryload("items") + .subqueryload("keywords") ) ).all() @@ -2885,7 +2673,6 @@ class ComparatorFactoryTest(_fixtures.FixtureTest, AssertsCompiledSQL): (composite, DummyComposite, users.c.id, users.c.name), (relationship, Address), (backref, "address"), - (comparable_property,), (dynamic_loader, Address), ): fn = args[0] @@ -3845,9 +3632,9 @@ class RequirementsTest(fixtures.MappedTest): h1s = ( s.query(H1) .options( - sa.orm.joinedload_all("t6a.h1b"), + sa.orm.joinedload("t6a").joinedload("h1b"), sa.orm.joinedload("h2s"), - sa.orm.joinedload_all("h3s.h1s"), + sa.orm.joinedload("h3s").joinedload("h1s"), ) .all() ) diff --git a/test/orm/test_merge.py b/test/orm/test_merge.py index c3e38d1b0..995989cd9 100644 --- a/test/orm/test_merge.py +++ b/test/orm/test_merge.py @@ -12,14 +12,12 @@ from sqlalchemy import testing from sqlalchemy import Text from sqlalchemy.orm import attributes from sqlalchemy.orm import backref -from sqlalchemy.orm import comparable_property from sqlalchemy.orm import configure_mappers from sqlalchemy.orm import create_session from sqlalchemy.orm import defer from sqlalchemy.orm import deferred from sqlalchemy.orm import foreign from sqlalchemy.orm import mapper -from sqlalchemy.orm import PropComparator from sqlalchemy.orm import relationship from sqlalchemy.orm import Session from sqlalchemy.orm import sessionmaker @@ -1283,13 +1281,10 @@ class MergeTest(_fixtures.FixtureTest): except sa.exc.InvalidRequestError as e: assert "load=False option does not support" in str(e) - def test_synonym_comparable(self): + def test_synonym(self): users = self.tables.users class User(object): - class Comparator(PropComparator): - pass - def _getValue(self): return self._value @@ -1298,14 +1293,7 @@ class MergeTest(_fixtures.FixtureTest): value = property(_getValue, _setValue) - mapper( - User, - users, - properties={ - "uid": synonym("id"), - "foobar": comparable_property(User.Comparator, User.value), - }, - ) + mapper(User, users, properties={"uid": synonym("id")}) sess = create_session() u = User() diff --git a/test/orm/test_of_type.py b/test/orm/test_of_type.py index 6e980e641..61fc80cb0 100644 --- a/test/orm/test_of_type.py +++ b/test/orm/test_of_type.py @@ -8,11 +8,9 @@ from sqlalchemy.engine import default from sqlalchemy.orm import aliased from sqlalchemy.orm import contains_eager from sqlalchemy.orm import joinedload -from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import relationship from sqlalchemy.orm import Session from sqlalchemy.orm import subqueryload -from sqlalchemy.orm import subqueryload_all from sqlalchemy.orm import with_polymorphic from sqlalchemy.testing import assert_raises_message from sqlalchemy.testing import eq_ @@ -562,8 +560,8 @@ class SubclassRelationshipTest( ) s = Session(testing.db) q = s.query(ParentThing).options( - subqueryload_all( - ParentThing.container, DataContainer.jobs.of_type(SubJob) + subqueryload(ParentThing.container).subqueryload( + DataContainer.jobs.of_type(SubJob) ) ) @@ -577,8 +575,8 @@ class SubclassRelationshipTest( s = Session(testing.db) sj_alias = aliased(SubJob) q = s.query(DataContainer).options( - subqueryload_all( - DataContainer.jobs.of_type(sj_alias), sj_alias.widget + subqueryload(DataContainer.jobs.of_type(sj_alias)).subqueryload( + sj_alias.widget ) ) @@ -596,8 +594,8 @@ class SubclassRelationshipTest( ) s = Session(testing.db) q = s.query(ParentThing).options( - joinedload_all( - ParentThing.container, DataContainer.jobs.of_type(SubJob) + joinedload(ParentThing.container).joinedload( + DataContainer.jobs.of_type(SubJob) ) ) diff --git a/test/orm/test_options.py b/test/orm/test_options.py index 98b6bccbf..4d205e593 100644 --- a/test/orm/test_options.py +++ b/test/orm/test_options.py @@ -11,7 +11,6 @@ from sqlalchemy.orm import column_property from sqlalchemy.orm import create_session from sqlalchemy.orm import defaultload from sqlalchemy.orm import joinedload -from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import lazyload from sqlalchemy.orm import Load from sqlalchemy.orm import mapper @@ -978,11 +977,11 @@ class OptionsNoPropTest(_fixtures.FixtureTest): r"Mapper\|Keyword\|keywords in this Query.", ) - def test_option_against_nonexistent_twolevel_all(self): + def test_option_against_nonexistent_twolevel_chained(self): Item = self.classes.Item self._assert_eager_with_entity_exception( [Item], - (joinedload_all("keywords.foo"),), + (joinedload("keywords").joinedload("foo"),), r"Can't find property named 'foo' on the mapped entity " r"Mapper\|Keyword\|keywords in this Query.", ) @@ -996,7 +995,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Keyword, Item], - (joinedload_all("keywords"),), + (joinedload("keywords"),), r"Attribute 'keywords' of entity 'Mapper\|Keyword\|keywords' " "does not refer to a mapped entity", ) @@ -1010,7 +1009,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Keyword, Item], - (joinedload_all("keywords"),), + (joinedload("keywords"),), r"Attribute 'keywords' of entity 'Mapper\|Keyword\|keywords' " "does not refer to a mapped entity", ) @@ -1019,7 +1018,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Item = self.classes.Item self._assert_eager_with_entity_exception( [Item], - (joinedload_all("id", "keywords"),), + (joinedload("id").joinedload("keywords"),), r"Attribute 'id' of entity 'Mapper\|Item\|items' does not " r"refer to a mapped entity", ) @@ -1029,7 +1028,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Keyword, Item], - (joinedload_all("id", "keywords"),), + (joinedload("id").joinedload("keywords"),), r"Attribute 'id' of entity 'Mapper\|Keyword\|keywords' " "does not refer to a mapped entity", ) @@ -1039,7 +1038,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Keyword, Item], - (joinedload_all("description"),), + (joinedload("description"),), r"Can't find property named 'description' on the mapped " r"entity Mapper\|Keyword\|keywords in this Query.", ) @@ -1049,7 +1048,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Keyword.id, Item.id], - (joinedload_all("keywords"),), + (joinedload("keywords"),), r"Query has only expression-based entities - can't find property " "named 'keywords'.", ) @@ -1059,7 +1058,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Keyword, Item], - (joinedload_all(Keyword.id, Item.keywords),), + (joinedload(Keyword.id).joinedload(Item.keywords),), r"Attribute 'id' of entity 'Mapper\|Keyword\|keywords' " "does not refer to a mapped entity", ) @@ -1069,7 +1068,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Keyword, Item], - (joinedload_all(Keyword.keywords, Item.keywords),), + (joinedload(Keyword.keywords).joinedload(Item.keywords),), r"Attribute 'keywords' of entity 'Mapper\|Keyword\|keywords' " "does not refer to a mapped entity", ) @@ -1079,7 +1078,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Keyword.id, Item.id], - (joinedload_all(Keyword.keywords, Item.keywords),), + (joinedload(Keyword.keywords).joinedload(Item.keywords),), r"Query has only expression-based entities - " "can't find property named 'keywords'.", ) @@ -1089,7 +1088,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Keyword = self.classes.Keyword self._assert_eager_with_entity_exception( [Item], - (joinedload_all(Keyword),), + (joinedload(Keyword),), r"mapper option expects string key or list of attributes", ) @@ -1097,7 +1096,7 @@ class OptionsNoPropTest(_fixtures.FixtureTest): User = self.classes.User self._assert_eager_with_entity_exception( [User], - (joinedload_all(User.addresses, User.orders),), + (joinedload(User.addresses).joinedload(User.orders),), r"Attribute 'User.orders' does not link " "from element 'Mapper|Address|addresses'", ) @@ -1107,7 +1106,11 @@ class OptionsNoPropTest(_fixtures.FixtureTest): Order = self.classes.Order self._assert_eager_with_entity_exception( [User], - (joinedload_all(User.addresses, User.orders.of_type(Order)),), + ( + joinedload(User.addresses).joinedload( + User.orders.of_type(Order) + ), + ), r"Attribute 'User.orders' does not link " "from element 'Mapper|Address|addresses'", ) diff --git a/test/orm/test_query.py b/test/orm/test_query.py index 12f894ef3..0c8c27bb2 100644 --- a/test/orm/test_query.py +++ b/test/orm/test_query.py @@ -39,7 +39,6 @@ from sqlalchemy.orm import column_property from sqlalchemy.orm import create_session from sqlalchemy.orm import defer from sqlalchemy.orm import joinedload -from sqlalchemy.orm import joinedload_all from sqlalchemy.orm import lazyload from sqlalchemy.orm import mapper from sqlalchemy.orm import Query @@ -842,7 +841,7 @@ class GetTest(QueryTest): # eager load does s.query(User).options( - joinedload("addresses"), joinedload_all("orders.items") + joinedload("addresses"), joinedload("orders").joinedload("items") ).populate_existing().all() assert u.addresses[0].email_address == "jack@bean.com" assert u.orders[1].items[2].description == "item 5" @@ -4000,31 +3999,30 @@ class TextTest(QueryTest, AssertsCompiledSQL): None, ) - def test_fragment(self): + def test_whereclause(self): User = self.classes.User - with expect_warnings("Textual SQL expression"): - eq_( - create_session().query(User).filter("id in (8, 9)").all(), - [User(id=8), User(id=9)], - ) + eq_( + create_session().query(User).filter(text("id in (8, 9)")).all(), + [User(id=8), User(id=9)], + ) - eq_( - create_session() - .query(User) - .filter("name='fred'") - .filter("id=9") - .all(), - [User(id=9)], - ) - eq_( - create_session() - .query(User) - .filter("name='fred'") - .filter(User.id == 9) - .all(), - [User(id=9)], - ) + eq_( + create_session() + .query(User) + .filter(text("name='fred'")) + .filter(text("id=9")) + .all(), + [User(id=9)], + ) + eq_( + create_session() + .query(User) + .filter(text("name='fred'")) + .filter(User.id == 9) + .all(), + [User(id=9)], + ) def test_binds_coerce(self): User = self.classes.User diff --git a/test/orm/test_selectin_relations.py b/test/orm/test_selectin_relations.py index 0ace60725..b891835d1 100644 --- a/test/orm/test_selectin_relations.py +++ b/test/orm/test_selectin_relations.py @@ -13,7 +13,6 @@ from sqlalchemy.orm import joinedload from sqlalchemy.orm import mapper from sqlalchemy.orm import relationship from sqlalchemy.orm import selectinload -from sqlalchemy.orm import selectinload_all from sqlalchemy.orm import Session from sqlalchemy.orm import subqueryload from sqlalchemy.orm import undefer @@ -136,7 +135,7 @@ class EagerTest(_fixtures.FixtureTest, testing.AssertsCompiledSQL): self.assert_sql_count(testing.db, go, 2) q = sess.query(u).options( - selectinload_all(u.addresses, Address.dingalings) + selectinload(u.addresses).selectinload(Address.dingalings) ) def go(): @@ -1040,33 +1039,6 @@ class EagerTest(_fixtures.FixtureTest, testing.AssertsCompiledSQL): result = q.order_by(sa.desc(User.id)).limit(2).offset(2).all() eq_(list(reversed(self.static.user_all_result[0:2])), result) - @testing.uses_deprecated("Mapper.order_by") - def test_mapper_order_by(self): - users, User, Address, addresses = ( - self.tables.users, - self.classes.User, - self.classes.Address, - self.tables.addresses, - ) - - mapper(Address, addresses) - mapper( - User, - users, - properties={ - "addresses": relationship( - Address, lazy="selectin", order_by=addresses.c.id - ) - }, - order_by=users.c.id.desc(), - ) - - sess = create_session() - q = sess.query(User) - - result = q.limit(2).all() - eq_(result, list(reversed(self.static.user_address_result[2:4]))) - def test_one_to_many_scalar(self): Address, addresses, users, User = ( self.classes.Address, @@ -1320,7 +1292,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): a2 = u1.addresses[0] a2.email_address = "foo" sess.query(User).options( - selectinload_all("addresses.dingaling") + selectinload("addresses").selectinload("dingaling") ).filter_by(id=8).all() assert u1.addresses[-1] is a1 for a in u1.addresses: @@ -1338,9 +1310,9 @@ class LoadOnExistingTest(_fixtures.FixtureTest): u1.orders o1 = Order() u1.orders.append(o1) - sess.query(User).options(selectinload_all("orders.items")).filter_by( - id=7 - ).all() + sess.query(User).options( + selectinload("orders").selectinload("items") + ).filter_by(id=7).all() for o in u1.orders: if o is not o1: assert "items" in o.__dict__ @@ -1357,7 +1329,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): .one() ) sess.query(User).filter_by(id=8).options( - selectinload_all("addresses.dingaling") + selectinload("addresses").selectinload("dingaling") ).first() assert "dingaling" in u1.addresses[0].__dict__ @@ -1371,7 +1343,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): .one() ) sess.query(User).filter_by(id=7).options( - selectinload_all("orders.items") + selectinload("orders").selectinload("items") ).first() assert "items" in u1.orders[0].__dict__ @@ -2429,7 +2401,7 @@ class SelfReferentialTest(fixtures.MappedTest): sess.query(Node) .filter_by(data="n1") .order_by(Node.id) - .options(selectinload_all("children.children")) + .options(selectinload("children").selectinload("children")) .first() ) eq_( diff --git a/test/orm/test_session.py b/test/orm/test_session.py index 7221f1d12..2bc11398d 100644 --- a/test/orm/test_session.py +++ b/test/orm/test_session.py @@ -34,7 +34,6 @@ from sqlalchemy.testing.schema import Column from sqlalchemy.testing.schema import Table from sqlalchemy.testing.util import gc_collect from sqlalchemy.util import pickle -from sqlalchemy.util import pypy from sqlalchemy.util.compat import inspect_getfullargspec from test.orm import _fixtures @@ -187,8 +186,9 @@ class SessionUtilTest(_fixtures.FixtureTest): assert u2 in s2 with assertions.expect_deprecated( - r"The Session.close_all\(\) method is deprecated and will " - "be removed in a future release. "): + r"The Session.close_all\(\) method is deprecated and will " + "be removed in a future release. " + ): Session.close_all() assert u1 not in s1 @@ -628,12 +628,11 @@ class SessionStateTest(_fixtures.FixtureTest): assert u1 not in s2 assert not s2.identity_map.keys() - @testing.uses_deprecated() def test_identity_conflict(self): users, User = self.tables.users, self.classes.User mapper(User, users) - for s in (create_session(), create_session(weak_identity_map=False)): + for s in (create_session(), create_session()): users.delete().execute() u1 = User(name="ed") s.add(u1) @@ -1466,172 +1465,6 @@ class WeakIdentityMapTest(_fixtures.FixtureTest): assert not sess.identity_map.contains_state(u2._sa_instance_state) -class StrongIdentityMapTest(_fixtures.FixtureTest): - run_inserts = None - - def _strong_ident_fixture(self): - sess = create_session(weak_identity_map=False) - return sess, sess.prune - - def _event_fixture(self): - session = create_session() - - @event.listens_for(session, "pending_to_persistent") - @event.listens_for(session, "deleted_to_persistent") - @event.listens_for(session, "detached_to_persistent") - @event.listens_for(session, "loaded_as_persistent") - def strong_ref_object(sess, instance): - if "refs" not in sess.info: - sess.info["refs"] = refs = set() - else: - refs = sess.info["refs"] - - refs.add(instance) - - @event.listens_for(session, "persistent_to_detached") - @event.listens_for(session, "persistent_to_deleted") - @event.listens_for(session, "persistent_to_transient") - def deref_object(sess, instance): - sess.info["refs"].discard(instance) - - def prune(): - if "refs" not in session.info: - return 0 - - sess_size = len(session.identity_map) - session.info["refs"].clear() - gc_collect() - session.info["refs"] = set( - s.obj() for s in session.identity_map.all_states() - ) - return sess_size - len(session.identity_map) - - return session, prune - - @testing.uses_deprecated() - def test_strong_ref_imap(self): - self._test_strong_ref(self._strong_ident_fixture) - - def test_strong_ref_events(self): - self._test_strong_ref(self._event_fixture) - - def _test_strong_ref(self, fixture): - s, prune = fixture() - - users, User = self.tables.users, self.classes.User - - mapper(User, users) - - # save user - s.add(User(name="u1")) - s.flush() - user = s.query(User).one() - user = None - print(s.identity_map) - gc_collect() - assert len(s.identity_map) == 1 - - user = s.query(User).one() - assert not s.identity_map._modified - user.name = "u2" - assert s.identity_map._modified - s.flush() - eq_(users.select().execute().fetchall(), [(user.id, "u2")]) - - @testing.uses_deprecated() - def test_prune_imap(self): - self._test_prune(self._strong_ident_fixture) - - def test_prune_events(self): - self._test_prune(self._event_fixture) - - @testing.fails_if(lambda: pypy, "pypy has a real GC") - @testing.fails_on("+zxjdbc", "http://www.sqlalchemy.org/trac/ticket/1473") - def _test_prune(self, fixture): - s, prune = fixture() - - users, User = self.tables.users, self.classes.User - - mapper(User, users) - - for o in [User(name="u%s" % x) for x in range(10)]: - s.add(o) - # o is still live after this loop... - - self.assert_(len(s.identity_map) == 0) - eq_(prune(), 0) - s.flush() - gc_collect() - eq_(prune(), 9) - # o is still in local scope here, so still present - self.assert_(len(s.identity_map) == 1) - - id_ = o.id - del o - eq_(prune(), 1) - self.assert_(len(s.identity_map) == 0) - - u = s.query(User).get(id_) - eq_(prune(), 0) - self.assert_(len(s.identity_map) == 1) - u.name = "squiznart" - del u - eq_(prune(), 0) - self.assert_(len(s.identity_map) == 1) - s.flush() - eq_(prune(), 1) - self.assert_(len(s.identity_map) == 0) - - s.add(User(name="x")) - eq_(prune(), 0) - self.assert_(len(s.identity_map) == 0) - s.flush() - self.assert_(len(s.identity_map) == 1) - eq_(prune(), 1) - self.assert_(len(s.identity_map) == 0) - - u = s.query(User).get(id_) - s.delete(u) - del u - eq_(prune(), 0) - self.assert_(len(s.identity_map) == 1) - s.flush() - eq_(prune(), 0) - self.assert_(len(s.identity_map) == 0) - - @testing.uses_deprecated() - def test_fast_discard_race(self): - # test issue #4068 - users, User = self.tables.users, self.classes.User - - mapper(User, users) - - sess = Session(weak_identity_map=False) - - u1 = User(name="u1") - sess.add(u1) - sess.commit() - - u1_state = u1._sa_instance_state - sess.identity_map._dict.pop(u1_state.key) - ref = u1_state.obj - u1_state.obj = lambda: None - - u2 = sess.query(User).first() - u1_state._cleanup(ref) - - u3 = sess.query(User).first() - - is_(u2, u3) - - u2_state = u2._sa_instance_state - assert sess.identity_map.contains_state(u2._sa_instance_state) - ref = u2_state.obj - u2_state.obj = lambda: None - u2_state._cleanup(ref) - assert not sess.identity_map.contains_state(u2._sa_instance_state) - - class IsModifiedTest(_fixtures.FixtureTest): run_inserts = None @@ -1700,28 +1533,6 @@ class IsModifiedTest(_fixtures.FixtureTest): addresses_loaded = "addresses" in u.__dict__ assert mod is not addresses_loaded - def test_is_modified_passive_on(self): - User, Address = self._default_mapping_fixture() - - s = Session() - u = User(name="fred", addresses=[Address(email_address="foo")]) - s.add(u) - s.commit() - - u.id - - def go(): - assert not s.is_modified(u, passive=True) - - self.assert_sql_count(testing.db, go, 0) - - u.name = "newname" - - def go(): - assert s.is_modified(u, passive=True) - - self.assert_sql_count(testing.db, go, 0) - def test_is_modified_syn(self): User, users = self.classes.User, self.tables.users @@ -2037,49 +1848,6 @@ class SessionInterface(fixtures.TestBase): ) -class TLTransactionTest(fixtures.MappedTest): - run_dispose_bind = "once" - __backend__ = True - - @classmethod - def setup_bind(cls): - return engines.testing_engine(options=dict(strategy="threadlocal")) - - @classmethod - def define_tables(cls, metadata): - Table( - "users", - metadata, - Column( - "id", Integer, primary_key=True, test_needs_autoincrement=True - ), - Column("name", String(20)), - test_needs_acid=True, - ) - - @classmethod - def setup_classes(cls): - class User(cls.Basic): - pass - - @classmethod - def setup_mappers(cls): - users, User = cls.tables.users, cls.classes.User - - mapper(User, users) - - @testing.exclude("mysql", "<", (5, 0, 3), "FIXME: unknown") - def test_session_nesting(self): - User = self.classes.User - - sess = create_session(bind=self.bind) - self.bind.begin() - u = User(name="ed") - sess.add(u) - sess.flush() - self.bind.commit() - - class FlushWarningsTest(fixtures.MappedTest): run_setup_mappers = "each" diff --git a/test/orm/test_subquery_relations.py b/test/orm/test_subquery_relations.py index a4ee2d804..b4be6debe 100644 --- a/test/orm/test_subquery_relations.py +++ b/test/orm/test_subquery_relations.py @@ -15,7 +15,6 @@ from sqlalchemy.orm import mapper from sqlalchemy.orm import relationship from sqlalchemy.orm import Session from sqlalchemy.orm import subqueryload -from sqlalchemy.orm import subqueryload_all from sqlalchemy.orm import undefer from sqlalchemy.orm import with_polymorphic from sqlalchemy.testing import assert_raises @@ -137,7 +136,7 @@ class EagerTest(_fixtures.FixtureTest, testing.AssertsCompiledSQL): self.assert_sql_count(testing.db, go, 2) q = sess.query(u).options( - subqueryload_all(u.addresses, Address.dingalings) + subqueryload(u.addresses).subqueryload(Address.dingalings) ) def go(): @@ -1060,33 +1059,6 @@ class EagerTest(_fixtures.FixtureTest, testing.AssertsCompiledSQL): result = q.order_by(sa.desc(User.id)).limit(2).offset(2).all() eq_(list(reversed(self.static.user_all_result[0:2])), result) - @testing.uses_deprecated("Mapper.order_by") - def test_mapper_order_by(self): - users, User, Address, addresses = ( - self.tables.users, - self.classes.User, - self.classes.Address, - self.tables.addresses, - ) - - mapper(Address, addresses) - mapper( - User, - users, - properties={ - "addresses": relationship( - Address, lazy="subquery", order_by=addresses.c.id - ) - }, - order_by=users.c.id.desc(), - ) - - sess = create_session() - q = sess.query(User) - - result = q.limit(2).all() - eq_(result, list(reversed(self.static.user_address_result[2:4]))) - def test_one_to_many_scalar(self): Address, addresses, users, User = ( self.classes.Address, @@ -1340,7 +1312,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): a2 = u1.addresses[0] a2.email_address = "foo" sess.query(User).options( - subqueryload_all("addresses.dingaling") + subqueryload("addresses").subqueryload("dingaling") ).filter_by(id=8).all() assert u1.addresses[-1] is a1 for a in u1.addresses: @@ -1358,9 +1330,9 @@ class LoadOnExistingTest(_fixtures.FixtureTest): u1.orders o1 = Order() u1.orders.append(o1) - sess.query(User).options(subqueryload_all("orders.items")).filter_by( - id=7 - ).all() + sess.query(User).options( + subqueryload("orders").subqueryload("items") + ).filter_by(id=7).all() for o in u1.orders: if o is not o1: assert "items" in o.__dict__ @@ -1377,7 +1349,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): .one() ) sess.query(User).filter_by(id=8).options( - subqueryload_all("addresses.dingaling") + subqueryload("addresses").subqueryload("dingaling") ).first() assert "dingaling" in u1.addresses[0].__dict__ @@ -1391,7 +1363,7 @@ class LoadOnExistingTest(_fixtures.FixtureTest): .one() ) sess.query(User).filter_by(id=7).options( - subqueryload_all("orders.items") + subqueryload("orders").subqueryload("items") ).first() assert "items" in u1.orders[0].__dict__ @@ -2328,7 +2300,7 @@ class SelfReferentialTest(fixtures.MappedTest): sess.query(Node) .filter_by(data="n1") .order_by(Node.id) - .options(subqueryload_all("children.children")) + .options(subqueryload("children").subqueryload("children")) .first() ) eq_( diff --git a/test/orm/test_transaction.py b/test/orm/test_transaction.py index e5708808e..5c8ed4732 100644 --- a/test/orm/test_transaction.py +++ b/test/orm/test_transaction.py @@ -711,19 +711,19 @@ class SessionTransactionTest(fixtures.RemovesEvents, FixtureTest): eq_( bind.mock_calls, [ - mock.call.contextual_connect(), - mock.call.contextual_connect().execution_options( + mock.call._contextual_connect(), + mock.call._contextual_connect().execution_options( isolation_level="FOO" ), - mock.call.contextual_connect().execution_options().begin(), + mock.call._contextual_connect().execution_options().begin(), ], ) - eq_(c1, bind.contextual_connect().execution_options()) + eq_(c1, bind._contextual_connect().execution_options()) def test_execution_options_ignored_mid_transaction(self): bind = mock.Mock() conn = mock.Mock(engine=bind) - bind.contextual_connect = mock.Mock(return_value=conn) + bind._contextual_connect = mock.Mock(return_value=conn) sess = Session(bind=bind) sess.execute("select 1") with expect_warnings( @@ -1553,75 +1553,6 @@ class AccountingFlagsTest(_LocalFixture): sess.expire_all() assert u1.name == "edward" - def test_rollback_no_accounting(self): - User, users = self.classes.User, self.tables.users - - sess = sessionmaker(_enable_transaction_accounting=False)() - u1 = User(name="ed") - sess.add(u1) - sess.commit() - - u1.name = "edwardo" - sess.rollback() - - testing.db.execute( - users.update(users.c.name == "ed").values(name="edward") - ) - - assert u1.name == "edwardo" - sess.expire_all() - assert u1.name == "edward" - - def test_commit_no_accounting(self): - User, users = self.classes.User, self.tables.users - - sess = sessionmaker(_enable_transaction_accounting=False)() - u1 = User(name="ed") - sess.add(u1) - sess.commit() - - u1.name = "edwardo" - sess.rollback() - - testing.db.execute( - users.update(users.c.name == "ed").values(name="edward") - ) - - assert u1.name == "edwardo" - sess.commit() - - assert testing.db.execute(select([users.c.name])).fetchall() == [ - ("edwardo",) - ] - assert u1.name == "edwardo" - - sess.delete(u1) - sess.commit() - - def test_preflush_no_accounting(self): - User, users = self.classes.User, self.tables.users - - sess = Session( - _enable_transaction_accounting=False, - autocommit=True, - autoflush=False, - ) - u1 = User(name="ed") - sess.add(u1) - sess.flush() - - sess.begin() - u1.name = "edwardo" - u2 = User(name="some other user") - sess.add(u2) - - sess.rollback() - - sess.begin() - assert testing.db.execute(select([users.c.name])).fetchall() == [ - ("ed",) - ] - class AutoCommitTest(_LocalFixture): __backend__ = True diff --git a/test/orm/test_versioning.py b/test/orm/test_versioning.py index e54925f0b..23263d136 100644 --- a/test/orm/test_versioning.py +++ b/test/orm/test_versioning.py @@ -363,19 +363,19 @@ class VersioningTest(fixtures.MappedTest): sa.orm.exc.StaleDataError, r"Instance .* has version id '\d+' which does not " r"match database-loaded version id '\d+'", - s1.query(Foo).with_lockmode("read").get, + s1.query(Foo).with_for_update(read=True).get, f1s1.id, ) # reload it - this expires the old version first - s1.refresh(f1s1, lockmode="read") + s1.refresh(f1s1, with_for_update=dict(read=True)) # now assert version OK - s1.query(Foo).with_lockmode("read").get(f1s1.id) + s1.query(Foo).with_for_update(read=True).get(f1s1.id) # assert brand new load is OK too s1.close() - s1.query(Foo).with_lockmode("read").get(f1s1.id) + s1.query(Foo).with_for_update(read=True).get(f1s1.id) def test_versioncheck_not_versioned(self): """ensure the versioncheck logic skips if there isn't a @@ -389,7 +389,7 @@ class VersioningTest(fixtures.MappedTest): f1s1 = Foo(value="f1 value", version_id=1) s1.add(f1s1) s1.commit() - s1.query(Foo).with_lockmode("read").get(f1s1.id) + s1.query(Foo).with_for_update(read=True).get(f1s1.id) @testing.emits_warning(r".*versioning cannot be verified") @engines.close_open_connections @@ -489,7 +489,7 @@ class VersioningTest(fixtures.MappedTest): s1.commit() s2 = create_session(autocommit=False) - f1s2 = s2.query(Foo).with_lockmode("read").get(f1s1.id) + f1s2 = s2.query(Foo).with_for_update(read=True).get(f1s1.id) assert f1s2.id == f1s1.id assert f1s2.value == f1s1.value diff --git a/test/sql/test_compiler.py b/test/sql/test_compiler.py index f152d0ed9..3418eac73 100644 --- a/test/sql/test_compiler.py +++ b/test/sql/test_compiler.py @@ -1529,14 +1529,6 @@ class SelectTest(fixtures.TestBase, AssertsCompiledSQL): "FROM mytable WHERE mytable.myid = :myid_1 FOR UPDATE", ) - assert_raises_message( - exc.ArgumentError, - "Unknown for_update argument: 'unknown_mode'", - table1.select, - table1.c.myid == 7, - for_update="unknown_mode", - ) - def test_alias(self): # test the alias for a table1. column names stay the same, # table name "changes" to "foo". diff --git a/test/sql/test_defaults.py b/test/sql/test_defaults.py index 5ede46ede..4ce8cc32f 100644 --- a/test/sql/test_defaults.py +++ b/test/sql/test_defaults.py @@ -363,7 +363,7 @@ class DefaultTest(fixtures.TestBase): @testing.fails_on("firebird", "Data type unknown") def test_standalone(self): - c = testing.db.engine.contextual_connect() + c = testing.db.engine.connect() x = c.execute(t.c.col1.default) y = t.c.col2.default.execute() z = c.execute(t.c.col3.default) diff --git a/test/sql/test_deprecations.py b/test/sql/test_deprecations.py new file mode 100644 index 000000000..2e4042a1b --- /dev/null +++ b/test/sql/test_deprecations.py @@ -0,0 +1,425 @@ +#! coding: utf-8 + +from sqlalchemy import bindparam +from sqlalchemy import Column +from sqlalchemy import column +from sqlalchemy import create_engine +from sqlalchemy import exc +from sqlalchemy import ForeignKey +from sqlalchemy import Integer +from sqlalchemy import MetaData +from sqlalchemy import select +from sqlalchemy import String +from sqlalchemy import Table +from sqlalchemy import table +from sqlalchemy import testing +from sqlalchemy import text +from sqlalchemy import util +from sqlalchemy.engine import default +from sqlalchemy.schema import DDL +from sqlalchemy.sql import util as sql_util +from sqlalchemy.testing import assert_raises +from sqlalchemy.testing import assert_raises_message +from sqlalchemy.testing import AssertsCompiledSQL +from sqlalchemy.testing import engines +from sqlalchemy.testing import eq_ +from sqlalchemy.testing import fixtures +from sqlalchemy.testing import mock + + +class DeprecationWarningsTest(fixtures.TestBase): + def test_ident_preparer_force(self): + preparer = testing.db.dialect.identifier_preparer + preparer.quote("hi") + with testing.expect_deprecated( + "The IdentifierPreparer.quote.force parameter is deprecated" + ): + preparer.quote("hi", True) + + with testing.expect_deprecated( + "The IdentifierPreparer.quote.force parameter is deprecated" + ): + preparer.quote("hi", False) + + preparer.quote_schema("hi") + with testing.expect_deprecated( + "The IdentifierPreparer.quote_schema.force parameter is deprecated" + ): + preparer.quote_schema("hi", True) + + with testing.expect_deprecated( + "The IdentifierPreparer.quote_schema.force parameter is deprecated" + ): + preparer.quote_schema("hi", True) + + def test_string_convert_unicode(self): + with testing.expect_deprecated( + "The String.convert_unicode parameter is deprecated and " + "will be removed in a future release." + ): + String(convert_unicode=True) + + def test_string_convert_unicode_force(self): + with testing.expect_deprecated( + "The String.convert_unicode parameter is deprecated and " + "will be removed in a future release." + ): + String(convert_unicode="force") + + def test_engine_convert_unicode(self): + with testing.expect_deprecated( + "The create_engine.convert_unicode parameter and " + "corresponding dialect-level" + ): + create_engine("mysql://", convert_unicode=True, module=mock.Mock()) + + def test_join_condition_ignore_nonexistent_tables(self): + m = MetaData() + t1 = Table("t1", m, Column("id", Integer)) + t2 = Table( + "t2", m, Column("id", Integer), Column("t1id", ForeignKey("t1.id")) + ) + with testing.expect_deprecated( + "The join_condition.ignore_nonexistent_tables " + "parameter is deprecated" + ): + join_cond = sql_util.join_condition( + t1, t2, ignore_nonexistent_tables=True + ) + + t1t2 = t1.join(t2) + + assert t1t2.onclause.compare(join_cond) + + def test_select_autocommit(self): + with testing.expect_deprecated( + "The select.autocommit parameter is deprecated and " + "will be removed in a future release." + ): + stmt = select([column("x")], autocommit=True) + + def test_select_for_update(self): + with testing.expect_deprecated( + "The select.for_update parameter is deprecated and " + "will be removed in a future release." + ): + stmt = select([column("x")], for_update=True) + + @testing.provide_metadata + def test_table_useexisting(self): + meta = self.metadata + + Table("t", meta, Column("x", Integer)) + meta.create_all() + + with testing.expect_deprecated( + "The Table.useexisting parameter is deprecated and " + "will be removed in a future release." + ): + Table("t", meta, useexisting=True, autoload_with=testing.db) + + with testing.expect_deprecated( + "The Table.useexisting parameter is deprecated and " + "will be removed in a future release." + ): + assert_raises_message( + exc.ArgumentError, + "useexisting is synonymous with extend_existing.", + Table, + "t", + meta, + useexisting=True, + extend_existing=True, + autoload_with=testing.db, + ) + + +class DDLListenerDeprecationsTest(fixtures.TestBase): + def setup(self): + self.bind = self.engine = engines.mock_engine() + self.metadata = MetaData(self.bind) + self.table = Table("t", self.metadata, Column("id", Integer)) + self.users = Table( + "users", + self.metadata, + Column("user_id", Integer, primary_key=True), + Column("user_name", String(40)), + ) + + def test_append_listener(self): + metadata, table, bind = self.metadata, self.table, self.bind + + def fn(*a): + return None + + with testing.expect_deprecated(".* is deprecated .*"): + table.append_ddl_listener("before-create", fn) + with testing.expect_deprecated(".* is deprecated .*"): + assert_raises( + exc.InvalidRequestError, table.append_ddl_listener, "blah", fn + ) + + with testing.expect_deprecated(".* is deprecated .*"): + metadata.append_ddl_listener("before-create", fn) + with testing.expect_deprecated(".* is deprecated .*"): + assert_raises( + exc.InvalidRequestError, + metadata.append_ddl_listener, + "blah", + fn, + ) + + def test_deprecated_append_ddl_listener_table(self): + metadata, users, engine = self.metadata, self.users, self.engine + canary = [] + with testing.expect_deprecated(".* is deprecated .*"): + users.append_ddl_listener( + "before-create", lambda e, t, b: canary.append("mxyzptlk") + ) + with testing.expect_deprecated(".* is deprecated .*"): + users.append_ddl_listener( + "after-create", lambda e, t, b: canary.append("klptzyxm") + ) + with testing.expect_deprecated(".* is deprecated .*"): + users.append_ddl_listener( + "before-drop", lambda e, t, b: canary.append("xyzzy") + ) + with testing.expect_deprecated(".* is deprecated .*"): + users.append_ddl_listener( + "after-drop", lambda e, t, b: canary.append("fnord") + ) + + metadata.create_all() + assert "mxyzptlk" in canary + assert "klptzyxm" in canary + assert "xyzzy" not in canary + assert "fnord" not in canary + del engine.mock[:] + canary[:] = [] + metadata.drop_all() + assert "mxyzptlk" not in canary + assert "klptzyxm" not in canary + assert "xyzzy" in canary + assert "fnord" in canary + + def test_deprecated_append_ddl_listener_metadata(self): + metadata, users, engine = self.metadata, self.users, self.engine + canary = [] + with testing.expect_deprecated(".* is deprecated .*"): + metadata.append_ddl_listener( + "before-create", + lambda e, t, b, tables=None: canary.append("mxyzptlk"), + ) + with testing.expect_deprecated(".* is deprecated .*"): + metadata.append_ddl_listener( + "after-create", + lambda e, t, b, tables=None: canary.append("klptzyxm"), + ) + with testing.expect_deprecated(".* is deprecated .*"): + metadata.append_ddl_listener( + "before-drop", + lambda e, t, b, tables=None: canary.append("xyzzy"), + ) + with testing.expect_deprecated(".* is deprecated .*"): + metadata.append_ddl_listener( + "after-drop", + lambda e, t, b, tables=None: canary.append("fnord"), + ) + + metadata.create_all() + assert "mxyzptlk" in canary + assert "klptzyxm" in canary + assert "xyzzy" not in canary + assert "fnord" not in canary + del engine.mock[:] + canary[:] = [] + metadata.drop_all() + assert "mxyzptlk" not in canary + assert "klptzyxm" not in canary + assert "xyzzy" in canary + assert "fnord" in canary + + def test_filter_deprecated(self): + cx = self.engine + + tbl = Table("t", MetaData(), Column("id", Integer)) + target = cx.name + + assert DDL("")._should_execute_deprecated("x", tbl, cx) + with testing.expect_deprecated(".* is deprecated .*"): + assert DDL("", on=target)._should_execute_deprecated("x", tbl, cx) + with testing.expect_deprecated(".* is deprecated .*"): + assert not DDL("", on="bogus")._should_execute_deprecated( + "x", tbl, cx + ) + with testing.expect_deprecated(".* is deprecated .*"): + assert DDL( + "", on=lambda d, x, y, z: True + )._should_execute_deprecated("x", tbl, cx) + with testing.expect_deprecated(".* is deprecated .*"): + assert DDL( + "", on=lambda d, x, y, z: z.engine.name != "bogus" + )._should_execute_deprecated("x", tbl, cx) + + +class ConvertUnicodeDeprecationTest(fixtures.TestBase): + + __backend__ = True + + data = util.u( + "Alors vous imaginez ma surprise, au lever du jour, quand " + "une drôle de petite voix m’a réveillé. " + "Elle disait: « S’il vous plaît… dessine-moi un mouton! »" + ) + + def test_unicode_warnings_dialectlevel(self): + + unicodedata = self.data + + with testing.expect_deprecated( + "The create_engine.convert_unicode parameter and " + "corresponding dialect-level" + ): + dialect = default.DefaultDialect(convert_unicode=True) + dialect.supports_unicode_binds = False + + s = String() + uni = s.dialect_impl(dialect).bind_processor(dialect) + + uni(util.b("x")) + assert isinstance(uni(unicodedata), util.binary_type) + + eq_(uni(unicodedata), unicodedata.encode("utf-8")) + + def test_ignoring_unicode_error(self): + """checks String(unicode_error='ignore') is passed to + underlying codec.""" + + unicodedata = self.data + + with testing.expect_deprecated( + "The String.convert_unicode parameter is deprecated and " + "will be removed in a future release.", + "The String.unicode_errors parameter is deprecated and " + "will be removed in a future release.", + ): + type_ = String( + 248, convert_unicode="force", unicode_error="ignore" + ) + dialect = default.DefaultDialect(encoding="ascii") + proc = type_.result_processor(dialect, 10) + + utfdata = unicodedata.encode("utf8") + eq_(proc(utfdata), unicodedata.encode("ascii", "ignore").decode()) + + +class ForUpdateTest(fixtures.TestBase, AssertsCompiledSQL): + __dialect__ = "default" + + def _assert_legacy(self, leg, read=False, nowait=False): + t = table("t", column("c")) + + with testing.expect_deprecated( + "The select.for_update parameter is deprecated and " + "will be removed in a future release." + ): + s1 = select([t], for_update=leg) + + if leg is False: + assert s1._for_update_arg is None + assert s1.for_update is None + else: + eq_(s1._for_update_arg.read, read) + eq_(s1._for_update_arg.nowait, nowait) + eq_(s1.for_update, leg) + + def test_false_legacy(self): + self._assert_legacy(False) + + def test_plain_true_legacy(self): + self._assert_legacy(True) + + def test_read_legacy(self): + self._assert_legacy("read", read=True) + + def test_nowait_legacy(self): + self._assert_legacy("nowait", nowait=True) + + def test_read_nowait_legacy(self): + self._assert_legacy("read_nowait", read=True, nowait=True) + + def test_unknown_mode(self): + t = table("t", column("c")) + + with testing.expect_deprecated( + "The select.for_update parameter is deprecated and " + "will be removed in a future release." + ): + assert_raises_message( + exc.ArgumentError, + "Unknown for_update argument: 'unknown_mode'", + t.select, + t.c.c == 7, + for_update="unknown_mode", + ) + + def test_legacy_setter(self): + t = table("t", column("c")) + s = select([t]) + s.for_update = "nowait" + eq_(s._for_update_arg.nowait, True) + + +class TextTest(fixtures.TestBase, AssertsCompiledSQL): + __dialect__ = "default" + + def test_legacy_bindparam(self): + with testing.expect_deprecated( + "The text.bindparams parameter is deprecated" + ): + t = text( + "select * from foo where lala=:bar and hoho=:whee", + bindparams=[bindparam("bar", 4), bindparam("whee", 7)], + ) + + self.assert_compile( + t, + "select * from foo where lala=:bar and hoho=:whee", + checkparams={"bar": 4, "whee": 7}, + ) + + def test_legacy_typemap(self): + table1 = table( + "mytable", + column("myid", Integer), + column("name", String), + column("description", String), + ) + with testing.expect_deprecated( + "The text.typemap parameter is deprecated" + ): + t = text( + "select id, name from user", + typemap=dict(id=Integer, name=String), + ) + + stmt = select([table1.c.myid]).select_from( + table1.join(t, table1.c.myid == t.c.id) + ) + compiled = stmt.compile() + eq_( + compiled._create_result_map(), + { + "myid": ( + "myid", + (table1.c.myid, "myid", "myid"), + table1.c.myid.type, + ) + }, + ) + + def test_autocommit(self): + with testing.expect_deprecated( + "The text.autocommit parameter is deprecated" + ): + t = text("select id, name from user", autocommit=True) diff --git a/test/sql/test_generative.py b/test/sql/test_generative.py index 1e6221cdc..f030f1d9c 100644 --- a/test/sql/test_generative.py +++ b/test/sql/test_generative.py @@ -500,8 +500,8 @@ class ClauseTest(fixtures.TestBase, AssertsCompiledSQL): ) def test_text(self): - clause = text( - "select * from table where foo=:bar", bindparams=[bindparam("bar")] + clause = text("select * from table where foo=:bar").bindparams( + bindparam("bar") ) c1 = str(clause) diff --git a/test/sql/test_metadata.py b/test/sql/test_metadata.py index a6d4b2d1a..3d60fb60e 100644 --- a/test/sql/test_metadata.py +++ b/test/sql/test_metadata.py @@ -2262,18 +2262,6 @@ class UseExistingTest(fixtures.TablesTest): extend_existing=True, ) - @testing.uses_deprecated() - def test_existing_plus_useexisting_raises(self): - meta2 = self._useexisting_fixture() - assert_raises( - exc.ArgumentError, - Table, - "users", - meta2, - useexisting=True, - extend_existing=True, - ) - def test_keep_existing_no_dupe_constraints(self): meta2 = self._notexisting_fixture() users = Table( diff --git a/test/sql/test_resultset.py b/test/sql/test_resultset.py index 3bd61b1f8..5987c7746 100644 --- a/test/sql/test_resultset.py +++ b/test/sql/test_resultset.py @@ -1647,7 +1647,7 @@ class AlternateResultProxyTest(fixtures.TablesTest): "test", metadata, Column("x", Integer, primary_key=True), - Column("y", String(50, convert_unicode="force")), + Column("y", String(50)), ) @classmethod diff --git a/test/sql/test_selectable.py b/test/sql/test_selectable.py index 5456dfb4f..04c0e6102 100644 --- a/test/sql/test_selectable.py +++ b/test/sql/test_selectable.py @@ -2544,39 +2544,6 @@ class ResultMapTest(fixtures.TestBase): class ForUpdateTest(fixtures.TestBase, AssertsCompiledSQL): __dialect__ = "default" - def _assert_legacy(self, leg, read=False, nowait=False): - t = table("t", column("c")) - s1 = select([t], for_update=leg) - - if leg is False: - assert s1._for_update_arg is None - assert s1.for_update is None - else: - eq_(s1._for_update_arg.read, read) - eq_(s1._for_update_arg.nowait, nowait) - eq_(s1.for_update, leg) - - def test_false_legacy(self): - self._assert_legacy(False) - - def test_plain_true_legacy(self): - self._assert_legacy(True) - - def test_read_legacy(self): - self._assert_legacy("read", read=True) - - def test_nowait_legacy(self): - self._assert_legacy("nowait", nowait=True) - - def test_read_nowait_legacy(self): - self._assert_legacy("read_nowait", read=True, nowait=True) - - def test_legacy_setter(self): - t = table("t", column("c")) - s = select([t]) - s.for_update = "nowait" - eq_(s._for_update_arg.nowait, True) - def test_basic_clone(self): t = table("t", column("c")) s = select([t]).with_for_update(read=True, of=t.c.c) diff --git a/test/sql/test_text.py b/test/sql/test_text.py index 6b419f599..48302058d 100644 --- a/test/sql/test_text.py +++ b/test/sql/test_text.py @@ -198,18 +198,6 @@ class SelectCompositionTest(fixtures.TestBase, AssertsCompiledSQL): class BindParamTest(fixtures.TestBase, AssertsCompiledSQL): __dialect__ = "default" - def test_legacy(self): - t = text( - "select * from foo where lala=:bar and hoho=:whee", - bindparams=[bindparam("bar", 4), bindparam("whee", 7)], - ) - - self.assert_compile( - t, - "select * from foo where lala=:bar and hoho=:whee", - checkparams={"bar": 4, "whee": 7}, - ) - def test_positional(self): t = text("select * from foo where lala=:bar and hoho=:whee") t = t.bindparams(bindparam("bar", 4), bindparam("whee", 7)) diff --git a/test/sql/test_types.py b/test/sql/test_types.py index c54fe1e54..4bd182c3e 100644 --- a/test/sql/test_types.py +++ b/test/sql/test_types.py @@ -175,7 +175,7 @@ class AdaptTest(fixtures.TestBase): % (type_, expected) ) - @testing.uses_deprecated() + @testing.uses_deprecated(".*Binary.*") def test_adapt_method(self): """ensure all types have a working adapt() method, which creates a distinct copy. @@ -191,6 +191,8 @@ class AdaptTest(fixtures.TestBase): def adaptions(): for typ in self._all_types(): + # up adapt from LowerCase to UPPERCASE, + # as well as to all non-sqltypes up_adaptions = [typ] + typ.__subclasses__() yield False, typ, up_adaptions for subcl in typ.__subclasses__(): @@ -258,7 +260,6 @@ class AdaptTest(fixtures.TestBase): eq_(types.DateTime().python_type, datetime.datetime) eq_(types.String().python_type, str) eq_(types.Unicode().python_type, util.text_type) - eq_(types.String(convert_unicode=True).python_type, util.text_type) eq_(types.Enum("one", "two", "three").python_type, str) assert_raises( @@ -283,10 +284,15 @@ class AdaptTest(fixtures.TestBase): This essentially is testing the behavior of util.constructor_copy(). """ - t1 = String(length=50, convert_unicode=False) - t2 = t1.adapt(Text, convert_unicode=True) + t1 = String(length=50) + t2 = t1.adapt(Text) eq_(t2.length, 50) - eq_(t2.convert_unicode, True) + + def test_convert_unicode_text_type(self): + with testing.expect_deprecated( + "The String.convert_unicode parameter is deprecated" + ): + eq_(types.String(convert_unicode=True).python_type, util.text_type) class TypeAffinityTest(fixtures.TestBase): @@ -1245,34 +1251,6 @@ class UnicodeTest(fixtures.TestBase): ): eq_(uni(5), 5) - def test_unicode_warnings_dialectlevel(self): - - unicodedata = self.data - - dialect = default.DefaultDialect(convert_unicode=True) - dialect.supports_unicode_binds = False - - s = String() - uni = s.dialect_impl(dialect).bind_processor(dialect) - - uni(util.b("x")) - assert isinstance(uni(unicodedata), util.binary_type) - - eq_(uni(unicodedata), unicodedata.encode("utf-8")) - - def test_ignoring_unicode_error(self): - """checks String(unicode_error='ignore') is passed to - underlying codec.""" - - unicodedata = self.data - - type_ = String(248, convert_unicode="force", unicode_error="ignore") - dialect = default.DefaultDialect(encoding="ascii") - proc = type_.result_processor(dialect, 10) - - utfdata = unicodedata.encode("utf8") - eq_(proc(utfdata), unicodedata.encode("ascii", "ignore").decode()) - class EnumTest(AssertsCompiledSQL, fixtures.TablesTest): __backend__ = True @@ -1857,6 +1835,7 @@ class EnumTest(AssertsCompiledSQL, fixtures.TablesTest): # depending on backend. assert "('x'," in e.print_sql() + @testing.uses_deprecated(".*convert_unicode") def test_repr(self): e = Enum( "x", @@ -1958,13 +1937,14 @@ class BinaryTest(fixtures.TestBase, AssertsExecutionResults): binary_table.select(order_by=binary_table.c.primary_id), text( "select * from binary_table order by binary_table.primary_id", - typemap={ + bind=testing.db, + ).columns( + **{ "pickled": PickleType, "mypickle": MyPickleType, "data": LargeBinary, "data_slice": LargeBinary, - }, - bind=testing.db, + } ), ): result = stmt.execute().fetchall() |
