summaryrefslogtreecommitdiff
path: root/test
diff options
context:
space:
mode:
authorMike Bayer <mike_mp@zzzcomputing.com>2007-07-27 04:08:53 +0000
committerMike Bayer <mike_mp@zzzcomputing.com>2007-07-27 04:08:53 +0000
commited4fc64bb0ac61c27bc4af32962fb129e74a36bf (patch)
treec1cf2fb7b1cafced82a8898e23d2a0bf5ced8526 /test
parent3a8e235af64e36b3b711df1f069d32359fe6c967 (diff)
downloadsqlalchemy-ed4fc64bb0ac61c27bc4af32962fb129e74a36bf.tar.gz
merging 0.4 branch to trunk. see CHANGES for details. 0.3 moves to maintenance branch in branches/rel_0_3.
Diffstat (limited to 'test')
-rw-r--r--test/base/alltests.py1
-rw-r--r--test/base/dependency.py7
-rw-r--r--test/base/utils.py67
-rw-r--r--test/dialect/alltests.py1
-rw-r--r--test/dialect/mysql.py58
-rw-r--r--test/dialect/oracle.py32
-rw-r--r--test/dialect/postgres.py128
-rw-r--r--test/engine/alltests.py2
-rw-r--r--test/engine/autoconnect_engine.py90
-rw-r--r--test/engine/bind.py49
-rw-r--r--test/engine/execute.py17
-rw-r--r--test/engine/metadata.py18
-rw-r--r--test/engine/parseconnect.py12
-rw-r--r--test/engine/pool.py44
-rw-r--r--test/engine/proxy_engine.py204
-rw-r--r--test/engine/reconnect.py12
-rw-r--r--test/engine/reflection.py130
-rw-r--r--test/engine/transaction.py294
-rw-r--r--test/ext/activemapper.py61
-rw-r--r--test/ext/alltests.py1
-rw-r--r--test/ext/assignmapper.py58
-rw-r--r--test/ext/associationproxy.py126
-rw-r--r--test/ext/legacy_objectstore.py113
-rw-r--r--test/ext/orderinglist.py49
-rw-r--r--test/ext/selectresults.py239
-rw-r--r--test/ext/wsgi_test.py122
-rw-r--r--test/orm/alltests.py23
-rw-r--r--test/orm/association.py7
-rw-r--r--test/orm/assorted_eager.py (renamed from test/orm/eagertest3.py)309
-rw-r--r--test/orm/attributes.py104
-rw-r--r--test/orm/cascade.py26
-rw-r--r--test/orm/collection.py1140
-rw-r--r--test/orm/compile.py5
-rw-r--r--test/orm/cycles.py41
-rw-r--r--test/orm/eager_relations.py133
-rw-r--r--test/orm/eagertest1.py69
-rw-r--r--test/orm/eagertest2.py239
-rw-r--r--test/orm/entity.py10
-rw-r--r--test/orm/fixtures.py4
-rw-r--r--test/orm/generative.py183
-rw-r--r--test/orm/inheritance/__init__.py0
-rw-r--r--test/orm/inheritance/abc_inheritance.py (renamed from test/orm/abc_inheritance.py)44
-rw-r--r--test/orm/inheritance/alltests.py28
-rw-r--r--test/orm/inheritance/basic.py (renamed from test/orm/inheritance.py)416
-rw-r--r--test/orm/inheritance/concrete.py (renamed from test/orm/inheritance4.py)9
-rw-r--r--test/orm/inheritance/magazine.py (renamed from test/orm/inheritance3.py)81
-rw-r--r--test/orm/inheritance/manytomany.py255
-rw-r--r--test/orm/inheritance/poly_linked_list.py (renamed from test/orm/poly_linked_list.py)35
-rw-r--r--test/orm/inheritance/polymorph.py (renamed from test/orm/polymorph.py)246
-rw-r--r--test/orm/inheritance/polymorph2.py (renamed from test/orm/inheritance5.py)142
-rw-r--r--test/orm/inheritance/productspec.py (renamed from test/orm/inheritance2.py)7
-rw-r--r--test/orm/inheritance/single.py (renamed from test/orm/single.py)7
-rw-r--r--test/orm/lazy_relations.py4
-rw-r--r--test/orm/lazytest1.py5
-rw-r--r--test/orm/manytomany.py18
-rw-r--r--test/orm/mapper.py538
-rw-r--r--test/orm/memusage.py19
-rw-r--r--test/orm/merge.py9
-rw-r--r--test/orm/onetoone.py4
-rw-r--r--test/orm/query.py564
-rw-r--r--test/orm/relationships.py141
-rw-r--r--test/orm/session.py262
-rw-r--r--test/orm/sessioncontext.py10
-rw-r--r--test/orm/sharding/__init__.py0
-rw-r--r--test/orm/sharding/alltests.py18
-rw-r--r--test/orm/sharding/shard.py154
-rw-r--r--test/orm/unitofwork.py115
-rw-r--r--test/perf/cascade_speed.py2
-rw-r--r--test/perf/masscreate.py5
-rw-r--r--test/perf/masscreate2.py6
-rw-r--r--test/perf/masseagerload.py102
-rw-r--r--test/perf/massload.py17
-rw-r--r--test/perf/massload2.py1
-rw-r--r--test/perf/masssave.py14
-rw-r--r--test/perf/ormsession.py225
-rw-r--r--test/perf/poolload.py5
-rw-r--r--test/perf/threaded_compile.py3
-rw-r--r--test/perf/wsgi.py54
-rw-r--r--test/rundocs.py242
-rw-r--r--test/sql/alltests.py4
-rw-r--r--test/sql/case_statement.py16
-rw-r--r--test/sql/constraints.py11
-rw-r--r--test/sql/defaults.py129
-rw-r--r--test/sql/generative.py275
-rw-r--r--test/sql/labels.py18
-rw-r--r--test/sql/query.py246
-rw-r--r--test/sql/quote.py5
-rw-r--r--test/sql/rowcount.py6
-rw-r--r--test/sql/select.py188
-rwxr-xr-xtest/sql/selectable.py32
-rw-r--r--test/sql/testtypes.py131
-rw-r--r--test/sql/unicode.py56
-rw-r--r--test/testbase.py474
-rw-r--r--test/testlib/__init__.py11
-rw-r--r--test/testlib/config.py255
-rw-r--r--test/testlib/coverage.py (renamed from test/coverage.py)271
-rw-r--r--test/testlib/profiling.py74
-rw-r--r--test/testlib/schema.py28
-rw-r--r--test/testlib/tables.py (renamed from test/tables.py)27
-rw-r--r--test/testlib/testing.py363
-rw-r--r--test/zblog/mappers.py1
-rw-r--r--test/zblog/tables.py7
-rw-r--r--test/zblog/tests.py28
-rw-r--r--test/zblog/user.py10
104 files changed, 6774 insertions, 3927 deletions
diff --git a/test/base/alltests.py b/test/base/alltests.py
index 70ff83ab8..44fa9b2ec 100644
--- a/test/base/alltests.py
+++ b/test/base/alltests.py
@@ -5,6 +5,7 @@ def suite():
modules_to_test = (
# core utilities
'base.dependency',
+ 'base.utils',
)
alltests = unittest.TestSuite()
for name in modules_to_test:
diff --git a/test/base/dependency.py b/test/base/dependency.py
index c5e54fc9f..ddadd1b31 100644
--- a/test/base/dependency.py
+++ b/test/base/dependency.py
@@ -1,7 +1,8 @@
-from testbase import PersistTest
+import testbase
import sqlalchemy.topological as topological
-import unittest, sys, os
from sqlalchemy import util
+from testlib import *
+
# TODO: need assertion conditions in this suite
@@ -190,4 +191,4 @@ class DependencySortTest(PersistTest):
if __name__ == "__main__":
- unittest.main()
+ testbase.main()
diff --git a/test/base/utils.py b/test/base/utils.py
new file mode 100644
index 000000000..97f3db06f
--- /dev/null
+++ b/test/base/utils.py
@@ -0,0 +1,67 @@
+import testbase
+from sqlalchemy import util, column, sql, exceptions
+from testlib import *
+
+
+class OrderedDictTest(PersistTest):
+ def test_odict(self):
+ o = util.OrderedDict()
+ o['a'] = 1
+ o['b'] = 2
+ o['snack'] = 'attack'
+ o['c'] = 3
+
+ self.assert_(o.keys() == ['a', 'b', 'snack', 'c'])
+ self.assert_(o.values() == [1, 2, 'attack', 3])
+
+ o.pop('snack')
+
+ self.assert_(o.keys() == ['a', 'b', 'c'])
+ self.assert_(o.values() == [1, 2, 3])
+
+ o2 = util.OrderedDict(d=4)
+ o2['e'] = 5
+
+ self.assert_(o2.keys() == ['d', 'e'])
+ self.assert_(o2.values() == [4, 5])
+
+ o.update(o2)
+ self.assert_(o.keys() == ['a', 'b', 'c', 'd', 'e'])
+ self.assert_(o.values() == [1, 2, 3, 4, 5])
+
+ o.setdefault('c', 'zzz')
+ o.setdefault('f', 6)
+ self.assert_(o.keys() == ['a', 'b', 'c', 'd', 'e', 'f'])
+ self.assert_(o.values() == [1, 2, 3, 4, 5, 6])
+
+class ColumnCollectionTest(PersistTest):
+ def test_in(self):
+ cc = sql.ColumnCollection()
+ cc.add(column('col1'))
+ cc.add(column('col2'))
+ cc.add(column('col3'))
+ assert 'col1' in cc
+ assert 'col2' in cc
+
+ try:
+ cc['col1'] in cc
+ assert False
+ except exceptions.ArgumentError, e:
+ assert str(e) == "__contains__ requires a string argument"
+
+ def test_compare(self):
+ cc1 = sql.ColumnCollection()
+ cc2 = sql.ColumnCollection()
+ cc3 = sql.ColumnCollection()
+ c1 = column('col1')
+ c2 = c1.label('col2')
+ c3 = column('col3')
+ cc1.add(c1)
+ cc2.add(c2)
+ cc3.add(c3)
+ assert (cc1==cc2).compare(c1 == c2)
+ assert not (cc1==cc3).compare(c2 == c3)
+
+
+if __name__ == "__main__":
+ testbase.main()
diff --git a/test/dialect/alltests.py b/test/dialect/alltests.py
index f4b39dd6f..890073625 100644
--- a/test/dialect/alltests.py
+++ b/test/dialect/alltests.py
@@ -5,6 +5,7 @@ def suite():
modules_to_test = (
'dialect.mysql',
'dialect.postgres',
+ 'dialect.oracle',
)
alltests = unittest.TestSuite()
for name in modules_to_test:
diff --git a/test/dialect/mysql.py b/test/dialect/mysql.py
index d9227383f..dbba78893 100644
--- a/test/dialect/mysql.py
+++ b/test/dialect/mysql.py
@@ -1,15 +1,13 @@
-from testbase import PersistTest, AssertMixin
import testbase
from sqlalchemy import *
from sqlalchemy.databases import mysql
-import sys, StringIO
+from testlib import *
-db = testbase.db
class TypesTest(AssertMixin):
"Test MySQL column types"
- @testbase.supported('mysql')
+ @testing.supported('mysql')
def test_numeric(self):
"Exercise type specification and options for numeric types."
@@ -104,13 +102,13 @@ class TypesTest(AssertMixin):
'SMALLINT(4) UNSIGNED ZEROFILL'),
]
- table_args = ['test_mysql_numeric', db]
+ table_args = ['test_mysql_numeric', MetaData(testbase.db)]
for index, spec in enumerate(columns):
type_, args, kw, res = spec
table_args.append(Column('c%s' % index, type_(*args, **kw)))
numeric_table = Table(*table_args)
- gen = db.dialect.schemagenerator(db, None, None)
+ gen = testbase.db.dialect.schemagenerator(testbase.db, None, None)
for col in numeric_table.c:
index = int(col.name[1:])
@@ -124,7 +122,7 @@ class TypesTest(AssertMixin):
raise
numeric_table.drop()
- @testbase.supported('mysql')
+ @testing.supported('mysql')
def test_charset(self):
"""Exercise CHARACTER SET and COLLATE-related options on string-type
columns."""
@@ -188,13 +186,13 @@ class TypesTest(AssertMixin):
'''ENUM('foo','bar') UNICODE''')
]
- table_args = ['test_mysql_charset', db]
+ table_args = ['test_mysql_charset', MetaData(testbase.db)]
for index, spec in enumerate(columns):
type_, args, kw, res = spec
table_args.append(Column('c%s' % index, type_(*args, **kw)))
charset_table = Table(*table_args)
- gen = db.dialect.schemagenerator(db, None, None)
+ gen = testbase.db.dialect.schemagenerator(testbase.db, None, None)
for col in charset_table.c:
index = int(col.name[1:])
@@ -208,11 +206,12 @@ class TypesTest(AssertMixin):
raise
charset_table.drop()
- @testbase.supported('mysql')
+ @testing.supported('mysql')
def test_enum(self):
"Exercise the ENUM type"
-
- enum_table = Table('mysql_enum', db,
+
+ db = testbase.db
+ enum_table = Table('mysql_enum', MetaData(testbase.db),
Column('e1', mysql.MSEnum('"a"', "'b'")),
Column('e2', mysql.MSEnum('"a"', "'b'"), nullable=False),
Column('e3', mysql.MSEnum('"a"', "'b'", strict=True)),
@@ -242,38 +241,17 @@ class TypesTest(AssertMixin):
enum_table.insert().execute(e1='a', e2='a', e3='a', e4='a')
enum_table.insert().execute(e1='b', e2='b', e3='b', e4='b')
- # Insert out of range enums, push stderr aside to avoid expected
- # warnings cluttering test output
- con = db.connect()
- if not hasattr(con.connection, 'show_warnings'):
- con.execute(insert(enum_table, {'e1':'c', 'e2':'c',
- 'e3':'a', 'e4':'a'}))
- else:
- try:
- aside = sys.stderr
- sys.stderr = StringIO.StringIO()
-
- self.assert_(not con.connection.show_warnings())
-
- con.execute(insert(enum_table, {'e1':'c', 'e2':'c',
- 'e3':'a', 'e4':'a'}))
-
- self.assert_(con.connection.show_warnings())
- finally:
- sys.stderr = aside
-
res = enum_table.select().execute().fetchall()
expected = [(None, 'a', None, 'a'),
('a', 'a', 'a', 'a'),
- ('b', 'b', 'b', 'b'),
- ('', '', 'a', 'a')]
+ ('b', 'b', 'b', 'b')]
# This is known to fail with MySQLDB 1.2.2 beta versions
# which return these as sets.Set(['a']), sets.Set(['b'])
# (even on Pythons with __builtin__.set)
- if db.dialect.dbapi.version_info < (1, 2, 2, 'beta', 3) and \
- db.dialect.dbapi.version_info >= (1, 2, 2):
+ if testbase.db.dialect.dbapi.version_info < (1, 2, 2, 'beta', 3) and \
+ testbase.db.dialect.dbapi.version_info >= (1, 2, 2):
# these mysqldb seem to always uses 'sets', even on later pythons
import sets
def convert(value):
@@ -292,10 +270,10 @@ class TypesTest(AssertMixin):
self.assertEqual(res, expected)
enum_table.drop()
- @testbase.supported('mysql')
+ @testing.supported('mysql')
def test_type_reflection(self):
# FIXME: older versions need their own test
- if db.dialect.get_version_info(db) < (5, 0):
+ if testbase.db.dialect.get_version_info(testbase.db) < (5, 0):
return
# (ask_for, roundtripped_as_if_different)
@@ -325,12 +303,12 @@ class TypesTest(AssertMixin):
columns = [Column('c%i' % (i + 1), t[0]) for i, t in enumerate(specs)]
- m = MetaData(db)
+ m = MetaData(testbase.db)
t_table = Table('mysql_types', m, *columns)
m.drop_all()
m.create_all()
- m2 = MetaData(db)
+ m2 = MetaData(testbase.db)
rt = Table('mysql_types', m2, autoload=True)
#print
diff --git a/test/dialect/oracle.py b/test/dialect/oracle.py
new file mode 100644
index 000000000..14de8960b
--- /dev/null
+++ b/test/dialect/oracle.py
@@ -0,0 +1,32 @@
+import testbase
+from sqlalchemy import *
+from sqlalchemy.databases import mysql
+
+from testlib import *
+
+
+class OutParamTest(AssertMixin):
+ @testing.supported('oracle')
+ def setUpAll(self):
+ testbase.db.execute("""
+create or replace procedure foo(x_in IN number, x_out OUT number, y_out OUT number) IS
+ retval number;
+ begin
+ retval := 6;
+ x_out := 10;
+ y_out := x_in * 15;
+ end;
+ """)
+
+ @testing.supported('oracle')
+ def test_out_params(self):
+ result = testbase.db.execute(text("begin foo(:x, :y, :z); end;", bindparams=[bindparam('x', Numeric), outparam('y', Numeric), outparam('z', Numeric)]), x=5)
+ assert result.out_parameters == {'y':10, 'z':75}, result.out_parameters
+ print result.out_parameters
+
+ @testing.supported('oracle')
+ def tearDownAll(self):
+ testbase.db.execute("DROP PROCEDURE foo")
+
+if __name__ == '__main__':
+ testbase.main()
diff --git a/test/dialect/postgres.py b/test/dialect/postgres.py
index 0507b7c5b..f80ddcadd 100644
--- a/test/dialect/postgres.py
+++ b/test/dialect/postgres.py
@@ -1,68 +1,68 @@
-from testbase import AssertMixin
import testbase
+import datetime
from sqlalchemy import *
from sqlalchemy.databases import postgres
-import datetime
+from testlib import *
-db = testbase.db
class DomainReflectionTest(AssertMixin):
"Test PostgreSQL domains"
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def setUpAll(self):
- self.con = db.connect()
- self.con.execute('CREATE DOMAIN testdomain INTEGER NOT NULL DEFAULT 42')
- self.con.execute('CREATE DOMAIN alt_schema.testdomain INTEGER DEFAULT 0')
- self.con.execute('CREATE TABLE testtable (question integer, answer testdomain)')
- self.con.execute('CREATE TABLE alt_schema.testtable(question integer, answer alt_schema.testdomain, anything integer)')
- self.con.execute('CREATE TABLE crosschema (question integer, answer alt_schema.testdomain)')
+ con = testbase.db.connect()
+ con.execute('CREATE DOMAIN testdomain INTEGER NOT NULL DEFAULT 42')
+ con.execute('CREATE DOMAIN alt_schema.testdomain INTEGER DEFAULT 0')
+ con.execute('CREATE TABLE testtable (question integer, answer testdomain)')
+ con.execute('CREATE TABLE alt_schema.testtable(question integer, answer alt_schema.testdomain, anything integer)')
+ con.execute('CREATE TABLE crosschema (question integer, answer alt_schema.testdomain)')
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def tearDownAll(self):
- self.con.execute('DROP TABLE testtable')
- self.con.execute('DROP TABLE alt_schema.testtable')
- self.con.execute('DROP TABLE crosschema')
- self.con.execute('DROP DOMAIN testdomain')
- self.con.execute('DROP DOMAIN alt_schema.testdomain')
+ con = testbase.db.connect()
+ con.execute('DROP TABLE testtable')
+ con.execute('DROP TABLE alt_schema.testtable')
+ con.execute('DROP TABLE crosschema')
+ con.execute('DROP DOMAIN testdomain')
+ con.execute('DROP DOMAIN alt_schema.testdomain')
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_table_is_reflected(self):
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
table = Table('testtable', metadata, autoload=True)
self.assertEquals(set(table.columns.keys()), set(['question', 'answer']), "Columns of reflected table didn't equal expected columns")
self.assertEquals(table.c.answer.type.__class__, postgres.PGInteger)
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_domain_is_reflected(self):
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
table = Table('testtable', metadata, autoload=True)
self.assertEquals(str(table.columns.answer.default.arg), '42', "Reflected default value didn't equal expected value")
self.assertFalse(table.columns.answer.nullable, "Expected reflected column to not be nullable.")
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_table_is_reflected_alt_schema(self):
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
table = Table('testtable', metadata, autoload=True, schema='alt_schema')
self.assertEquals(set(table.columns.keys()), set(['question', 'answer', 'anything']), "Columns of reflected table didn't equal expected columns")
self.assertEquals(table.c.anything.type.__class__, postgres.PGInteger)
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_schema_domain_is_reflected(self):
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
table = Table('testtable', metadata, autoload=True, schema='alt_schema')
self.assertEquals(str(table.columns.answer.default.arg), '0', "Reflected default value didn't equal expected value")
self.assertTrue(table.columns.answer.nullable, "Expected reflected column to be nullable.")
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_crosschema_domain_is_reflected(self):
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
table = Table('crosschema', metadata, autoload=True)
self.assertEquals(str(table.columns.answer.default.arg), '0', "Reflected default value didn't equal expected value")
self.assertTrue(table.columns.answer.nullable, "Expected reflected column to be nullable.")
class MiscTest(AssertMixin):
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_date_reflection(self):
m1 = MetaData(testbase.db)
t1 = Table('pgdate', m1,
@@ -78,7 +78,7 @@ class MiscTest(AssertMixin):
finally:
m1.drop_all()
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_pg_weirdchar_reflection(self):
meta1 = MetaData(testbase.db)
subject = Table("subject", meta1,
@@ -99,18 +99,18 @@ class MiscTest(AssertMixin):
finally:
meta1.drop_all()
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_checksfor_sequence(self):
meta1 = MetaData(testbase.db)
t = Table('mytable', meta1,
Column('col1', Integer, Sequence('fooseq')))
try:
testbase.db.execute("CREATE SEQUENCE fooseq")
- t.create()
+ t.create(checkfirst=True)
finally:
- t.drop()
+ t.drop(checkfirst=True)
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_schema_reflection(self):
"""note: this test requires that the 'alt_schema' schema be separate and accessible by the test user"""
@@ -141,7 +141,7 @@ class MiscTest(AssertMixin):
finally:
meta1.drop_all()
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_schema_reflection_2(self):
meta1 = MetaData(testbase.db)
subject = Table("subject", meta1,
@@ -162,7 +162,7 @@ class MiscTest(AssertMixin):
finally:
meta1.drop_all()
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_schema_reflection_3(self):
meta1 = MetaData(testbase.db)
subject = Table("subject", meta1,
@@ -185,7 +185,7 @@ class MiscTest(AssertMixin):
finally:
meta1.drop_all()
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_preexecute_passivedefault(self):
"""test that when we get a primary key column back
from reflecting a table which has a default value on it, we pre-execute
@@ -216,7 +216,7 @@ class TimezoneTest(AssertMixin):
if postgres returns it. python then will not let you compare a datetime with a tzinfo to a datetime
that doesnt have one. this test illustrates two ways to have datetime types with and without timezone
info. """
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def setUpAll(self):
global tztable, notztable, metadata
metadata = MetaData(testbase.db)
@@ -233,11 +233,11 @@ class TimezoneTest(AssertMixin):
Column("name", String(20)),
)
metadata.create_all()
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def tearDownAll(self):
metadata.drop_all()
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_with_timezone(self):
# get a date with a tzinfo
somedate = testbase.db.connect().scalar(func.current_timestamp().select())
@@ -246,7 +246,7 @@ class TimezoneTest(AssertMixin):
x = c.last_updated_params()
print x['date'] == somedate
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_without_timezone(self):
# get a date without a tzinfo
somedate = datetime.datetime(2005, 10,20, 11, 52, 00)
@@ -255,6 +255,56 @@ class TimezoneTest(AssertMixin):
x = c.last_updated_params()
print x['date'] == somedate
+class ArrayTest(AssertMixin):
+ @testing.supported('postgres')
+ def setUpAll(self):
+ global metadata, arrtable
+ metadata = MetaData(testbase.db)
+
+ arrtable = Table('arrtable', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('intarr', postgres.PGArray(Integer)),
+ Column('strarr', postgres.PGArray(String), nullable=False)
+ )
+ metadata.create_all()
+ @testing.supported('postgres')
+ def tearDownAll(self):
+ metadata.drop_all()
+
+ @testing.supported('postgres')
+ def test_reflect_array_column(self):
+ metadata2 = MetaData(testbase.db)
+ tbl = Table('arrtable', metadata2, autoload=True)
+ self.assertTrue(isinstance(tbl.c.intarr.type, postgres.PGArray))
+ self.assertTrue(isinstance(tbl.c.strarr.type, postgres.PGArray))
+ self.assertTrue(isinstance(tbl.c.intarr.type.item_type, Integer))
+ self.assertTrue(isinstance(tbl.c.strarr.type.item_type, String))
+
+ @testing.supported('postgres')
+ def test_insert_array(self):
+ arrtable.insert().execute(intarr=[1,2,3], strarr=['abc', 'def'])
+ results = arrtable.select().execute().fetchall()
+ self.assertEquals(len(results), 1)
+ self.assertEquals(results[0]['intarr'], [1,2,3])
+ self.assertEquals(results[0]['strarr'], ['abc','def'])
+ arrtable.delete().execute()
+
+ @testing.supported('postgres')
+ def test_array_where(self):
+ arrtable.insert().execute(intarr=[1,2,3], strarr=['abc', 'def'])
+ arrtable.insert().execute(intarr=[4,5,6], strarr='ABC')
+ results = arrtable.select().where(arrtable.c.intarr == [1,2,3]).execute().fetchall()
+ self.assertEquals(len(results), 1)
+ self.assertEquals(results[0]['intarr'], [1,2,3])
+ arrtable.delete().execute()
+ @testing.supported('postgres')
+ def test_array_concat(self):
+ arrtable.insert().execute(intarr=[1,2,3], strarr=['abc', 'def'])
+ results = select([arrtable.c.intarr + [4,5,6]]).execute().fetchall()
+ self.assertEquals(len(results), 1)
+ self.assertEquals(results[0][0], [1,2,3,4,5,6])
+ arrtable.delete().execute()
+
if __name__ == "__main__":
testbase.main()
diff --git a/test/engine/alltests.py b/test/engine/alltests.py
index ec8a47390..a34a82ed7 100644
--- a/test/engine/alltests.py
+++ b/test/engine/alltests.py
@@ -10,12 +10,12 @@ def suite():
'engine.bind',
'engine.reconnect',
'engine.execute',
+ 'engine.metadata',
'engine.transaction',
# schema/tables
'engine.reflection',
- 'engine.proxy_engine'
)
alltests = unittest.TestSuite()
for name in modules_to_test:
diff --git a/test/engine/autoconnect_engine.py b/test/engine/autoconnect_engine.py
deleted file mode 100644
index 69c2c33f5..000000000
--- a/test/engine/autoconnect_engine.py
+++ /dev/null
@@ -1,90 +0,0 @@
-from testbase import PersistTest
-import testbase
-from sqlalchemy import *
-from sqlalchemy.ext.proxy import AutoConnectEngine
-
-import os
-
-#
-# Define an engine, table and mapper at the module level, to show that the
-# table and mapper can be used with different real engines in multiple threads
-#
-
-
-module_engine = AutoConnectEngine( testbase.db_uri )
-users = Table('users', module_engine,
- Column('user_id', Integer, primary_key=True),
- Column('user_name', String(16)),
- Column('password', String(20))
- )
-
-class User(object):
- pass
-
-
-class AutoConnectEngineTest1(PersistTest):
-
- def setUp(self):
- clear_mappers()
- objectstore.clear()
-
- def test_engine_connect(self):
- users.create()
- assign_mapper(User, users)
- try:
- trans = objectstore.begin()
-
- user = User()
- user.user_name='fred'
- user.password='*'
- trans.commit()
-
- # select
- sqluser = User.select_by(user_name='fred')[0]
- assert sqluser.user_name == 'fred'
-
- # modify
- sqluser.user_name = 'fred jones'
-
- # commit - saves everything that changed
- objectstore.commit()
-
- allusers = [ user.user_name for user in User.select() ]
- assert allusers == [ 'fred jones' ]
- finally:
- users.drop()
-
-
-
-
-if __name__ == "__main__":
- testbase.main()
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
diff --git a/test/engine/bind.py b/test/engine/bind.py
index b9e53e6b1..6a0c78f57 100644
--- a/test/engine/bind.py
+++ b/test/engine/bind.py
@@ -2,12 +2,10 @@
including the deprecated versions of these arguments"""
import testbase
-import unittest, sys, datetime
-import tables
-db = testbase.db
from sqlalchemy import *
+from testlib import *
-class BindTest(testbase.PersistTest):
+class BindTest(PersistTest):
def test_create_drop_explicit(self):
metadata = MetaData()
table = Table('test_table', metadata,
@@ -17,7 +15,6 @@ class BindTest(testbase.PersistTest):
testbase.db.connect()
):
for args in [
- ([], {'connectable':bind}),
([], {'bind':bind}),
([bind], {})
]:
@@ -57,7 +54,7 @@ class BindTest(testbase.PersistTest):
table = Table('test_table', metadata,
Column('foo', Integer))
metadata.bind = bind
- assert metadata.bind is metadata.engine is table.bind is table.engine is bind
+ assert metadata.bind is table.bind is bind
metadata.create_all()
assert table.exists()
metadata.drop_all()
@@ -70,7 +67,7 @@ class BindTest(testbase.PersistTest):
Column('foo', Integer))
metadata.connect(bind)
- assert metadata.bind is metadata.engine is table.bind is table.engine is bind
+ assert metadata.bind is table.bind is bind
metadata.create_all()
assert table.exists()
metadata.drop_all()
@@ -88,15 +85,12 @@ class BindTest(testbase.PersistTest):
try:
for args in (
([bind], {}),
- ([], {'engine_or_url':bind}),
([], {'bind':bind}),
- ([], {'engine':bind})
):
metadata = MetaData(*args[0], **args[1])
table = Table('test_table', metadata,
- Column('foo', Integer))
-
- assert metadata.bind is metadata.engine is table.bind is table.engine is bind
+ Column('foo', Integer))
+ assert metadata.bind is table.bind is bind
metadata.create_all()
assert table.exists()
metadata.drop_all()
@@ -111,7 +105,8 @@ class BindTest(testbase.PersistTest):
metadata = MetaData()
table = Table('test_table', metadata,
Column('foo', Integer),
- mysql_engine='InnoDB')
+ test_needs_acid=True,
+ )
conn = testbase.db.connect()
metadata.create_all(bind=conn)
try:
@@ -124,7 +119,7 @@ class BindTest(testbase.PersistTest):
table.insert().execute(foo=7)
trans.rollback()
metadata.bind = None
- assert testbase.db.execute("select count(1) from test_table").scalar() == 0
+ assert conn.execute("select count(1) from test_table").scalar() == 0
finally:
metadata.drop_all(bind=conn)
@@ -147,10 +142,7 @@ class BindTest(testbase.PersistTest):
):
try:
e = elem(bind=bind)
- assert e.bind is e.engine is bind
- e.execute()
- e = elem(engine=bind)
- assert e.bind is e.engine is bind
+ assert e.bind is bind
e.execute()
finally:
if isinstance(bind, engine.Connection):
@@ -158,16 +150,19 @@ class BindTest(testbase.PersistTest):
try:
e = elem()
- assert e.bind is e.engine is None
+ assert e.bind is None
e.execute()
assert False
except exceptions.InvalidRequestError, e:
assert str(e) == "This Compiled object is not bound to any Engine or Connection."
-
+
finally:
+ if isinstance(bind, engine.Connection):
+ bind.close()
metadata.drop_all(bind=testbase.db)
def test_session(self):
+ from sqlalchemy.orm import create_session, mapper
metadata = MetaData()
table = Table('test_table', metadata,
Column('foo', Integer, primary_key=True),
@@ -177,11 +172,13 @@ class BindTest(testbase.PersistTest):
mapper(Foo, table)
metadata.create_all(bind=testbase.db)
try:
- for bind in (testbase.db, testbase.db.connect()):
+ for bind in (testbase.db,
+ testbase.db.connect()
+ ):
try:
- for args in ({'bind':bind}, {'bind_to':bind}):
+ for args in ({'bind':bind},):
sess = create_session(**args)
- assert sess.bind is sess.bind_to is bind
+ assert sess.bind is bind
f = Foo()
sess.save(f)
sess.flush()
@@ -189,6 +186,9 @@ class BindTest(testbase.PersistTest):
finally:
if isinstance(bind, engine.Connection):
bind.close()
+
+ if isinstance(bind, engine.Connection):
+ bind.close()
sess = create_session()
f = Foo()
@@ -198,8 +198,9 @@ class BindTest(testbase.PersistTest):
assert False
except exceptions.InvalidRequestError, e:
assert str(e).startswith("Could not locate any Engine or Connection bound to mapper")
-
finally:
+ if isinstance(bind, engine.Connection):
+ bind.close()
metadata.drop_all(bind=testbase.db)
diff --git a/test/engine/execute.py b/test/engine/execute.py
index 283006cfa..3d3b43f9b 100644
--- a/test/engine/execute.py
+++ b/test/engine/execute.py
@@ -1,19 +1,14 @@
-
import testbase
-import unittest, sys, datetime
-import tables
-db = testbase.db
from sqlalchemy import *
+from testlib import *
-
-class ExecuteTest(testbase.PersistTest):
+class ExecuteTest(PersistTest):
def setUpAll(self):
global users, metadata
metadata = MetaData(testbase.db)
users = Table('users', metadata,
Column('user_id', INT, primary_key = True),
Column('user_name', VARCHAR(20)),
- mysql_engine='InnoDB'
)
metadata.create_all()
@@ -22,7 +17,7 @@ class ExecuteTest(testbase.PersistTest):
def tearDownAll(self):
metadata.drop_all()
- @testbase.supported('sqlite')
+ @testing.supported('sqlite')
def test_raw_qmark(self):
for conn in (testbase.db, testbase.db.connect()):
conn.execute("insert into users (user_id, user_name) values (?, ?)", (1,"jack"))
@@ -34,7 +29,7 @@ class ExecuteTest(testbase.PersistTest):
assert res.fetchall() == [(1, "jack"), (2, "fred"), (3, "ed"), (4, "horse"), (5, "barney"), (6, "donkey"), (7, 'sally')]
conn.execute("delete from users")
- @testbase.supported('mysql', 'postgres')
+ @testing.supported('mysql', 'postgres')
def test_raw_sprintf(self):
for conn in (testbase.db, testbase.db.connect()):
conn.execute("insert into users (user_id, user_name) values (%s, %s)", [1,"jack"])
@@ -47,7 +42,7 @@ class ExecuteTest(testbase.PersistTest):
# pyformat is supported for mysql, but skipping because a few driver
# versions have a bug that bombs out on this test. (1.2.2b3, 1.2.2c1, 1.2.2)
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_raw_python(self):
for conn in (testbase.db, testbase.db.connect()):
conn.execute("insert into users (user_id, user_name) values (%(id)s, %(name)s)", {'id':1, 'name':'jack'})
@@ -57,7 +52,7 @@ class ExecuteTest(testbase.PersistTest):
assert res.fetchall() == [(1, "jack"), (2, "ed"), (3, "horse"), (4, 'sally')]
conn.execute("delete from users")
- @testbase.supported('sqlite')
+ @testing.supported('sqlite')
def test_raw_named(self):
for conn in (testbase.db, testbase.db.connect()):
conn.execute("insert into users (user_id, user_name) values (:id, :name)", {'id':1, 'name':'jack'})
diff --git a/test/engine/metadata.py b/test/engine/metadata.py
new file mode 100644
index 000000000..973007fab
--- /dev/null
+++ b/test/engine/metadata.py
@@ -0,0 +1,18 @@
+import testbase
+from sqlalchemy import *
+from testlib import *
+
+class MetaDataTest(PersistTest):
+ def test_metadata_connect(self):
+ metadata = MetaData()
+ t1 = Table('table1', metadata, Column('col1', Integer, primary_key=True),
+ Column('col2', String(20)))
+ metadata.bind = testbase.db
+ metadata.create_all()
+ try:
+ assert t1.count().scalar() == 0
+ finally:
+ metadata.drop_all()
+
+if __name__ == '__main__':
+ testbase.main()
diff --git a/test/engine/parseconnect.py b/test/engine/parseconnect.py
index 967a20ed5..3e186275d 100644
--- a/test/engine/parseconnect.py
+++ b/test/engine/parseconnect.py
@@ -1,8 +1,7 @@
-from testbase import PersistTest
import testbase
-import sqlalchemy.engine.url as url
from sqlalchemy import *
-import unittest
+import sqlalchemy.engine.url as url
+from testlib import *
class ParseConnectTest(PersistTest):
@@ -65,7 +64,7 @@ class CreateEngineTest(PersistTest):
def testrecycle(self):
dbapi = MockDBAPI(foober=12, lala=18, hoho={'this':'dict'}, fooz='somevalue')
e = create_engine('postgres://', pool_recycle=472, module=dbapi)
- assert e.connection_provider._pool._recycle == 472
+ assert e.pool._recycle == 472
def testbadargs(self):
# good arg, use MockDBAPI to prevent oracle import errors
@@ -116,7 +115,6 @@ class CreateEngineTest(PersistTest):
except TypeError:
assert True
- e = create_engine('sqlite://', echo=True)
e = create_engine('mysql://', module=MockDBAPI(), connect_args={'use_unicode':True}, convert_unicode=True)
e = create_engine('sqlite://', connect_args={'use_unicode':True}, convert_unicode=True)
@@ -139,8 +137,8 @@ class CreateEngineTest(PersistTest):
def testpoolargs(self):
"""test that connection pool args make it thru"""
e = create_engine('postgres://', creator=None, pool_recycle=-1, echo_pool=None, auto_close_cursors=False, disallow_open_cursors=True, module=MockDBAPI())
- assert e.connection_provider._pool.auto_close_cursors is False
- assert e.connection_provider._pool.disallow_open_cursors is True
+ assert e.pool.auto_close_cursors is False
+ assert e.pool.disallow_open_cursors is True
# these args work for QueuePool
e = create_engine('postgres://', max_overflow=8, pool_timeout=60, poolclass=pool.QueuePool, module=MockDBAPI())
diff --git a/test/engine/pool.py b/test/engine/pool.py
index 85e9d59fd..364afa9d7 100644
--- a/test/engine/pool.py
+++ b/test/engine/pool.py
@@ -1,10 +1,9 @@
import testbase
-from testbase import PersistTest
-import unittest, sys, os, time
-import threading, thread
-
+import threading, thread, time
import sqlalchemy.pool as pool
import sqlalchemy.exceptions as exceptions
+from testlib import *
+
mcid = 1
class MockDBAPI(object):
@@ -45,7 +44,7 @@ class PoolTest(PersistTest):
connection2 = manager.connect('foo.db')
connection3 = manager.connect('bar.db')
- self.echo( "connection " + repr(connection))
+ print "connection " + repr(connection)
self.assert_(connection.cursor() is not None)
self.assert_(connection is connection2)
self.assert_(connection2 is not connection3)
@@ -64,7 +63,7 @@ class PoolTest(PersistTest):
connection = manager.connect('foo.db')
connection2 = manager.connect('foo.db')
- self.echo( "connection " + repr(connection))
+ print "connection " + repr(connection)
self.assert_(connection.cursor() is not None)
self.assert_(connection is not connection2)
@@ -80,7 +79,7 @@ class PoolTest(PersistTest):
def status(pool):
tup = (pool.size(), pool.checkedin(), pool.overflow(), pool.checkedout())
- self.echo( "Pool size: %d Connections in pool: %d Current Overflow: %d Current Checked out connections: %d" % tup)
+ print "Pool size: %d Connections in pool: %d Current Overflow: %d Current Checked out connections: %d" % tup
return tup
c1 = p.connect()
@@ -160,7 +159,7 @@ class PoolTest(PersistTest):
print timeouts
assert len(timeouts) > 0
for t in timeouts:
- assert abs(t - 3) < 1
+ assert abs(t - 3) < 1, "Not all timeouts were 3 seconds: " + repr(timeouts)
def _test_overflow(self, thread_count, max_overflow):
def creator():
@@ -352,6 +351,35 @@ class PoolTest(PersistTest):
c2 = None
c1 = None
self.assert_(p.checkedout() == 0)
+
+ def test_properties(self):
+ dbapi = MockDBAPI()
+ p = pool.QueuePool(creator=lambda: dbapi.connect('foo.db'),
+ pool_size=1, max_overflow=0)
+
+ c = p.connect()
+ self.assert_(not c.properties)
+ self.assert_(c.properties is c._connection_record.properties)
+
+ c.properties['foo'] = 'bar'
+ c.close()
+ del c
+
+ c = p.connect()
+ self.assert_('foo' in c.properties)
+
+ c.invalidate()
+ c = p.connect()
+ self.assert_('foo' not in c.properties)
+
+ c.properties['foo2'] = 'bar2'
+ c.detach()
+ self.assert_('foo2' in c.properties)
+
+ c2 = p.connect()
+ self.assert_(c.connection is not c2.connection)
+ self.assert_(not c2.properties)
+ self.assert_('foo2' in c.properties)
def tearDown(self):
pool.clear_managers()
diff --git a/test/engine/proxy_engine.py b/test/engine/proxy_engine.py
deleted file mode 100644
index 26b738e41..000000000
--- a/test/engine/proxy_engine.py
+++ /dev/null
@@ -1,204 +0,0 @@
-from testbase import PersistTest
-import testbase
-import os
-
-from sqlalchemy import *
-from sqlalchemy.ext.proxy import ProxyEngine
-
-
-#
-# Define an engine, table and mapper at the module level, to show that the
-# table and mapper can be used with different real engines in multiple threads
-#
-
-
-class ProxyTestBase(PersistTest):
- def setUpAll(self):
-
- global users, User, module_engine, module_metadata
-
- module_engine = ProxyEngine(echo=testbase.echo)
- module_metadata = MetaData()
-
- users = Table('users', module_metadata,
- Column('user_id', Integer, primary_key=True),
- Column('user_name', String(16)),
- Column('password', String(20))
- )
-
- class User(object):
- pass
-
- User.mapper = mapper(User, users)
- def tearDownAll(self):
- clear_mappers()
-
-class ConstructTest(ProxyTestBase):
- """tests that we can build SQL constructs without engine-specific parameters, particulary
- oid_column, being needed, as the proxy engine is usually not connected yet."""
-
- def test_join(self):
- engine = ProxyEngine()
- t = Table('table1', engine,
- Column('col1', Integer, primary_key=True))
- t2 = Table('table2', engine,
- Column('col2', Integer, ForeignKey('table1.col1')))
- j = join(t, t2)
-
-
-class ProxyEngineTest1(ProxyTestBase):
-
- def test_engine_connect(self):
- # connect to a real engine
- module_engine.connect(testbase.db_uri)
- module_metadata.create_all(module_engine)
-
- session = create_session(bind_to=module_engine)
- try:
-
- user = User()
- user.user_name='fred'
- user.password='*'
-
- session.save(user)
- session.flush()
-
- query = session.query(User)
-
- # select
- sqluser = query.select_by(user_name='fred')[0]
- assert sqluser.user_name == 'fred'
-
- # modify
- sqluser.user_name = 'fred jones'
-
- # flush - saves everything that changed
- session.flush()
-
- allusers = [ user.user_name for user in query.select() ]
- assert allusers == ['fred jones']
-
- finally:
- module_metadata.drop_all(module_engine)
-
-
-class ThreadProxyTest(ProxyTestBase):
-
- def tearDownAll(self):
- try:
- os.remove('threadtesta.db')
- except OSError:
- pass
- try:
- os.remove('threadtestb.db')
- except OSError:
- pass
-
- @testbase.supported('sqlite')
- def test_multi_thread(self):
-
- from threading import Thread
- from Queue import Queue
-
- # start 2 threads with different connection params
- # and perform simultaneous operations, showing that the
- # 2 threads don't share a connection
- qa = Queue()
- qb = Queue()
- def run(db_uri, uname, queue):
- def test():
-
- try:
- module_engine.connect(db_uri)
- module_metadata.create_all(module_engine)
- try:
- session = create_session(bind_to=module_engine)
-
- query = session.query(User)
-
- all = list(query.select())
- assert all == []
-
- u = User()
- u.user_name = uname
- u.password = 'whatever'
-
- session.save(u)
- session.flush()
-
- names = [u.user_name for u in query.select()]
- assert names == [uname]
- finally:
- module_metadata.drop_all(module_engine)
- module_engine.get_engine().dispose()
- except Exception, e:
- import traceback
- traceback.print_exc()
- queue.put(e)
- else:
- queue.put(False)
- return test
-
- a = Thread(target=run('sqlite:///threadtesta.db', 'jim', qa))
- b = Thread(target=run('sqlite:///threadtestb.db', 'joe', qb))
-
- a.start()
- b.start()
-
- # block and wait for the threads to push their results
- res = qa.get()
- if res != False:
- raise res
-
- res = qb.get()
- if res != False:
- raise res
-
-
-class ProxyEngineTest2(ProxyTestBase):
-
- def test_table_singleton_a(self):
- """set up for table singleton check
- """
- #
- # For this 'test', create a proxy engine instance, connect it
- # to a real engine, and make it do some work
- #
- engine = ProxyEngine()
- cats = Table('cats', engine,
- Column('cat_id', Integer, primary_key=True),
- Column('cat_name', String))
-
- engine.connect(testbase.db_uri)
-
- cats.create(engine)
- cats.drop(engine)
-
- ProxyEngineTest2.cats_table_a = cats
- assert isinstance(cats, Table)
-
- def test_table_singleton_b(self):
- """check that a table on a 2nd proxy engine instance gets 2nd table
- instance
- """
- #
- # Now create a new proxy engine instance and attach the same
- # table as the first test. This should result in 2 table instances,
- # since different proxy engine instances can't attach to the
- # same table instance
- #
- engine = ProxyEngine()
- cats = Table('cats', engine,
- Column('cat_id', Integer, primary_key=True),
- Column('cat_name', String))
- assert id(cats) != id(ProxyEngineTest2.cats_table_a)
-
- # the real test -- if we're still using the old engine reference,
- # this will fail because the old reference's local storage will
- # not have the default attributes
- engine.connect(testbase.db_uri)
- cats.create(engine)
- cats.drop(engine)
-
-if __name__ == "__main__":
- testbase.main()
diff --git a/test/engine/reconnect.py b/test/engine/reconnect.py
index defc878ab..7c213695f 100644
--- a/test/engine/reconnect.py
+++ b/test/engine/reconnect.py
@@ -1,6 +1,8 @@
import testbase
+import sys, weakref
from sqlalchemy import create_engine, exceptions
-import gc, weakref, sys
+from testlib import *
+
class MockDisconnect(Exception):
pass
@@ -37,7 +39,7 @@ class MockCursor(object):
def close(self):
pass
-class ReconnectTest(testbase.PersistTest):
+class ReconnectTest(PersistTest):
def test_reconnect(self):
"""test that an 'is_disconnect' condition will invalidate the connection, and additionally
dispose the previous connection pool and recreate."""
@@ -50,7 +52,7 @@ class ReconnectTest(testbase.PersistTest):
# monkeypatch disconnect checker
db.dialect.is_disconnect = lambda e: isinstance(e, MockDisconnect)
- pid = id(db.connection_provider._pool)
+ pid = id(db.pool)
# make a connection
conn = db.connect()
@@ -81,7 +83,7 @@ class ReconnectTest(testbase.PersistTest):
# close shouldnt break
conn.close()
- assert id(db.connection_provider._pool) != pid
+ assert id(db.pool) != pid
# ensure all connections closed (pool was recycled)
assert len(dbapi.connections) == 0
@@ -92,4 +94,4 @@ class ReconnectTest(testbase.PersistTest):
assert len(dbapi.connections) == 1
if __name__ == '__main__':
- testbase.main() \ No newline at end of file
+ testbase.main()
diff --git a/test/engine/reflection.py b/test/engine/reflection.py
index 74ae75e2e..00c1276ee 100644
--- a/test/engine/reflection.py
+++ b/test/engine/reflection.py
@@ -1,13 +1,12 @@
-from testbase import PersistTest
import testbase
-import pickle
-import sqlalchemy.ansisql as ansisql
+import pickle, StringIO
from sqlalchemy import *
+import sqlalchemy.ansisql as ansisql
from sqlalchemy.exceptions import NoSuchTableError
import sqlalchemy.databases.mysql as mysql
+from testlib import *
-import unittest, re, StringIO
class ReflectionTest(PersistTest):
def testbasic(self):
@@ -15,6 +14,10 @@ class ReflectionTest(PersistTest):
use_string_defaults = use_function_defaults or testbase.db.engine.__module__.endswith('sqlite')
+ if (testbase.db.engine.name == 'mysql' and
+ testbase.db.dialect.get_version_info(testbase.db) < (4, 1, 1)):
+ return
+
if use_function_defaults:
defval = func.current_date()
deftype = Date
@@ -54,14 +57,14 @@ class ReflectionTest(PersistTest):
Column('test_passivedefault4', deftype3, PassiveDefault(defval3)),
Column('test9', Binary(100)),
Column('test_numeric', Numeric(None, None)),
- mysql_engine='InnoDB'
+ test_needs_fk=True,
)
-
+
addresses = Table('engine_email_addresses', meta,
Column('address_id', Integer, primary_key = True),
Column('remote_user_id', Integer, ForeignKey(users.c.user_id)),
Column('email_address', String(20)),
- mysql_engine='InnoDB'
+ test_needs_fk=True,
)
meta.drop_all()
@@ -106,6 +109,29 @@ class ReflectionTest(PersistTest):
addresses.drop()
users.drop()
+ def test_autoload_partial(self):
+ meta = MetaData(testbase.db)
+ foo = Table('foo', meta,
+ Column('a', String(30)),
+ Column('b', String(30)),
+ Column('c', String(30)),
+ Column('d', String(30)),
+ Column('e', String(30)),
+ Column('f', String(30)),
+ )
+ meta.create_all()
+ try:
+ meta2 = MetaData(testbase.db)
+ foo2 = Table('foo', meta2, autoload=True, include_columns=['b', 'f', 'e'])
+ # test that cols come back in original order
+ assert [c.name for c in foo2.c] == ['b', 'e', 'f']
+ for c in ('b', 'f', 'e'):
+ assert c in foo2.c
+ for c in ('a', 'c', 'd'):
+ assert c not in foo2.c
+ finally:
+ meta.drop_all()
+
def testoverridecolumns(self):
"""test that you can override columns which contain foreign keys to other reflected tables"""
meta = MetaData(testbase.db)
@@ -203,7 +229,7 @@ class ReflectionTest(PersistTest):
finally:
meta.drop_all()
- @testbase.supported('mysql')
+ @testing.supported('mysql')
def testmysqltypes(self):
meta1 = MetaData(testbase.db)
table = Table(
@@ -250,7 +276,7 @@ class ReflectionTest(PersistTest):
PRIMARY KEY(id)
)""")
try:
- metadata = MetaData(engine=testbase.db)
+ metadata = MetaData(bind=testbase.db)
book = Table('book', metadata, autoload=True)
assert book.c.id in book.primary_key
assert book.c.series not in book.primary_key
@@ -271,7 +297,7 @@ class ReflectionTest(PersistTest):
PRIMARY KEY(id, isbn)
)""")
try:
- metadata = MetaData(engine=testbase.db)
+ metadata = MetaData(bind=testbase.db)
book = Table('book', metadata, autoload=True)
assert book.c.id in book.primary_key
assert book.c.isbn in book.primary_key
@@ -280,7 +306,7 @@ class ReflectionTest(PersistTest):
finally:
testbase.db.execute("drop table book")
- @testbase.supported('sqlite')
+ @testing.supported('sqlite')
def test_goofy_sqlite(self):
"""test autoload of table where quotes were used with all the colnames. quirky in sqlite."""
testbase.db.execute("""CREATE TABLE "django_content_type" (
@@ -309,7 +335,12 @@ class ReflectionTest(PersistTest):
def test_composite_fk(self):
"""test reflection of composite foreign keys"""
+
+ if (testbase.db.engine.name == 'mysql' and
+ testbase.db.dialect.get_version_info(testbase.db) < (4, 1, 1)):
+ return
meta = MetaData(testbase.db)
+
table = Table(
'multi', meta,
Column('multi_id', Integer, primary_key=True),
@@ -317,7 +348,7 @@ class ReflectionTest(PersistTest):
Column('multi_hoho', Integer, primary_key=True),
Column('name', String(50), nullable=False),
Column('val', String(100)),
- mysql_engine='InnoDB'
+ test_needs_fk=True,
)
table2 = Table('multi2', meta,
Column('id', Integer, primary_key=True),
@@ -326,7 +357,7 @@ class ReflectionTest(PersistTest):
Column('lala', Integer),
Column('data', String(50)),
ForeignKeyConstraint(['foo', 'bar', 'lala'], ['multi.multi_id', 'multi.multi_rev', 'multi.multi_hoho']),
- mysql_engine='InnoDB'
+ test_needs_fk=True,
)
assert table.c.multi_hoho
meta.create_all()
@@ -345,7 +376,6 @@ class ReflectionTest(PersistTest):
finally:
meta.drop_all()
-
def test_to_metadata(self):
meta = MetaData()
@@ -372,17 +402,17 @@ class ReflectionTest(PersistTest):
def test_pickle():
meta.connect(testbase.db)
meta2 = pickle.loads(pickle.dumps(meta))
- assert meta2.engine is None
+ assert meta2.bind is None
return (meta2.tables['mytable'], meta2.tables['othertable'])
def test_pickle_via_reflect():
# this is the most common use case, pickling the results of a
# database reflection
- meta2 = MetaData(engine=testbase.db)
+ meta2 = MetaData(bind=testbase.db)
t1 = Table('mytable', meta2, autoload=True)
t2 = Table('othertable', meta2, autoload=True)
meta3 = pickle.loads(pickle.dumps(meta2))
- assert meta3.engine is None
+ assert meta3.bind is None
assert meta3.tables['mytable'] is not t1
return (meta3.tables['mytable'], meta3.tables['othertable'])
@@ -392,6 +422,8 @@ class ReflectionTest(PersistTest):
table_c, table2_c = test()
assert table is not table_c
assert table_c.c.myid.primary_key
+ assert isinstance(table_c.c.myid.type, Integer)
+ assert isinstance(table_c.c.name.type, String)
assert not table_c.c.name.nullable
assert table_c.c.description.nullable
assert table.primary_key is not table_c.primary_key
@@ -418,14 +450,10 @@ class ReflectionTest(PersistTest):
finally:
meta.drop_all(testbase.db)
- # mysql throws its own exception for no such table, resulting in
- # a sqlalchemy.SQLError instead of sqlalchemy.NoSuchTableError.
- # this could probably be fixed at some point.
- @testbase.unsupported('mysql')
def test_nonexistent(self):
self.assertRaises(NoSuchTableError, Table,
'fake_table',
- testbase.db, autoload=True)
+ MetaData(testbase.db), autoload=True)
def testoverride(self):
meta = MetaData(testbase.db)
@@ -452,7 +480,7 @@ class ReflectionTest(PersistTest):
finally:
table.drop()
- @testbase.supported('mssql')
+ @testing.supported('mssql')
def testidentity(self):
meta = MetaData(testbase.db)
table = Table(
@@ -505,20 +533,6 @@ class ReflectionTest(PersistTest):
finally:
meta.drop_all()
-
- meta = MetaData(testbase.db)
- table = Table(
- 'select', meta,
- Column('col1', Integer, primary_key=True)
- )
- table.create()
-
- meta2 = MetaData(testbase.db)
- try:
- table2 = Table('select', meta2, autoload=True)
- finally:
- table.drop()
-
class CreateDropTest(PersistTest):
def setUpAll(self):
global metadata, users
@@ -558,33 +572,33 @@ class CreateDropTest(PersistTest):
def testcheckfirst(self):
try:
assert not users.exists(testbase.db)
- users.create(connectable=testbase.db)
+ users.create(bind=testbase.db)
assert users.exists(testbase.db)
- users.create(connectable=testbase.db, checkfirst=True)
- users.drop(connectable=testbase.db)
- users.drop(connectable=testbase.db, checkfirst=True)
- assert not users.exists(connectable=testbase.db)
- users.create(connectable=testbase.db, checkfirst=True)
- users.drop(connectable=testbase.db)
+ users.create(bind=testbase.db, checkfirst=True)
+ users.drop(bind=testbase.db)
+ users.drop(bind=testbase.db, checkfirst=True)
+ assert not users.exists(bind=testbase.db)
+ users.create(bind=testbase.db, checkfirst=True)
+ users.drop(bind=testbase.db)
finally:
- metadata.drop_all(connectable=testbase.db)
+ metadata.drop_all(bind=testbase.db)
def test_createdrop(self):
- metadata.create_all(connectable=testbase.db)
+ metadata.create_all(bind=testbase.db)
self.assertEqual( testbase.db.has_table('items'), True )
self.assertEqual( testbase.db.has_table('email_addresses'), True )
- metadata.create_all(connectable=testbase.db)
+ metadata.create_all(bind=testbase.db)
self.assertEqual( testbase.db.has_table('items'), True )
- metadata.drop_all(connectable=testbase.db)
+ metadata.drop_all(bind=testbase.db)
self.assertEqual( testbase.db.has_table('items'), False )
self.assertEqual( testbase.db.has_table('email_addresses'), False )
- metadata.drop_all(connectable=testbase.db)
+ metadata.drop_all(bind=testbase.db)
self.assertEqual( testbase.db.has_table('items'), False )
class SchemaTest(PersistTest):
# this test should really be in the sql tests somewhere, not engine
- @testbase.unsupported('sqlite')
+ @testing.unsupported('sqlite')
def testiteration(self):
metadata = MetaData()
table1 = Table('table1', metadata,
@@ -607,14 +621,17 @@ class SchemaTest(PersistTest):
print buf
assert buf.index("CREATE TABLE someschema.table1") > -1
assert buf.index("CREATE TABLE someschema.table2") > -1
-
- @testbase.unsupported('sqlite', 'postgres')
- def test_create_with_defaultschema(self):
+
+ @testing.supported('mysql','postgres')
+ def testcreate(self):
engine = testbase.db
schema = engine.dialect.get_default_schema_name(engine)
+ #engine.echo = True
- # test reflection of tables with an explcit schemaname
- # matching the default
+ if testbase.db.name == 'mysql':
+ schema = testbase.db.url.database
+ else:
+ schema = 'public'
metadata = MetaData(testbase.db)
table1 = Table('table1', metadata,
Column('col1', Integer, primary_key=True),
@@ -628,10 +645,7 @@ class SchemaTest(PersistTest):
metadata.clear()
table1 = Table('table1', metadata, autoload=True, schema=schema)
table2 = Table('table2', metadata, autoload=True, schema=schema)
- assert table1.schema == table2.schema == schema
- assert len(metadata.tables) == 2
metadata.drop_all()
-
if __name__ == "__main__":
testbase.main()
diff --git a/test/engine/transaction.py b/test/engine/transaction.py
index c89bf4b14..593a069a9 100644
--- a/test/engine/transaction.py
+++ b/test/engine/transaction.py
@@ -1,19 +1,19 @@
-
import testbase
-import unittest, sys, datetime
-import tables
-db = testbase.db
+import sys, time, threading
+
from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
-class TransactionTest(testbase.PersistTest):
+class TransactionTest(PersistTest):
def setUpAll(self):
global users, metadata
metadata = MetaData()
users = Table('query_users', metadata,
Column('user_id', INT, primary_key = True),
Column('user_name', VARCHAR(20)),
- mysql_engine='InnoDB'
+ test_needs_acid=True,
)
users.create(testbase.db)
@@ -114,8 +114,154 @@ class TransactionTest(testbase.PersistTest):
result = connection.execute("select * from query_users")
assert len(result.fetchall()) == 0
connection.close()
+
+ @testing.unsupported('sqlite')
+ def testnestedsubtransactionrollback(self):
+ connection = testbase.db.connect()
+ transaction = connection.begin()
+ connection.execute(users.insert(), user_id=1, user_name='user1')
+ trans2 = connection.begin_nested()
+ connection.execute(users.insert(), user_id=2, user_name='user2')
+ trans2.rollback()
+ connection.execute(users.insert(), user_id=3, user_name='user3')
+ transaction.commit()
+
+ self.assertEquals(
+ connection.execute(select([users.c.user_id]).order_by(users.c.user_id)).fetchall(),
+ [(1,),(3,)]
+ )
+ connection.close()
+
+ @testing.unsupported('sqlite')
+ def testnestedsubtransactioncommit(self):
+ connection = testbase.db.connect()
+ transaction = connection.begin()
+ connection.execute(users.insert(), user_id=1, user_name='user1')
+ trans2 = connection.begin_nested()
+ connection.execute(users.insert(), user_id=2, user_name='user2')
+ trans2.commit()
+ connection.execute(users.insert(), user_id=3, user_name='user3')
+ transaction.commit()
+
+ self.assertEquals(
+ connection.execute(select([users.c.user_id]).order_by(users.c.user_id)).fetchall(),
+ [(1,),(2,),(3,)]
+ )
+ connection.close()
+
+ @testing.unsupported('sqlite')
+ def testrollbacktosubtransaction(self):
+ connection = testbase.db.connect()
+ transaction = connection.begin()
+ connection.execute(users.insert(), user_id=1, user_name='user1')
+ trans2 = connection.begin_nested()
+ connection.execute(users.insert(), user_id=2, user_name='user2')
+ trans3 = connection.begin()
+ connection.execute(users.insert(), user_id=3, user_name='user3')
+ trans3.rollback()
+ connection.execute(users.insert(), user_id=4, user_name='user4')
+ transaction.commit()
+
+ self.assertEquals(
+ connection.execute(select([users.c.user_id]).order_by(users.c.user_id)).fetchall(),
+ [(1,),(4,)]
+ )
+ connection.close()
+
+ @testing.supported('postgres', 'mysql')
+ def testtwophasetransaction(self):
+ connection = testbase.db.connect()
+
+ transaction = connection.begin_twophase()
+ connection.execute(users.insert(), user_id=1, user_name='user1')
+ transaction.prepare()
+ transaction.commit()
+
+ transaction = connection.begin_twophase()
+ connection.execute(users.insert(), user_id=2, user_name='user2')
+ transaction.commit()
+
+ transaction = connection.begin_twophase()
+ connection.execute(users.insert(), user_id=3, user_name='user3')
+ transaction.rollback()
+
+ transaction = connection.begin_twophase()
+ connection.execute(users.insert(), user_id=4, user_name='user4')
+ transaction.prepare()
+ transaction.rollback()
+
+ self.assertEquals(
+ connection.execute(select([users.c.user_id]).order_by(users.c.user_id)).fetchall(),
+ [(1,),(2,)]
+ )
+ connection.close()
+
+ @testing.supported('postgres', 'mysql')
+ def testmixedtransaction(self):
+ connection = testbase.db.connect()
+
+ transaction = connection.begin_twophase()
+ connection.execute(users.insert(), user_id=1, user_name='user1')
+
+ transaction2 = connection.begin()
+ connection.execute(users.insert(), user_id=2, user_name='user2')
+
+ transaction3 = connection.begin_nested()
+ connection.execute(users.insert(), user_id=3, user_name='user3')
+
+ transaction4 = connection.begin()
+ connection.execute(users.insert(), user_id=4, user_name='user4')
+ transaction4.commit()
+
+ transaction3.rollback()
+
+ connection.execute(users.insert(), user_id=5, user_name='user5')
+
+ transaction2.commit()
+
+ transaction.prepare()
+
+ transaction.commit()
+
+ self.assertEquals(
+ connection.execute(select([users.c.user_id]).order_by(users.c.user_id)).fetchall(),
+ [(1,),(2,),(5,)]
+ )
+ connection.close()
-class AutoRollbackTest(testbase.PersistTest):
+ @testing.supported('postgres')
+ def testtwophaserecover(self):
+ # MySQL recovery doesn't currently seem to work correctly
+ # Prepared transactions disappear when connections are closed and even
+ # when they aren't it doesn't seem possible to use the recovery id.
+ connection = testbase.db.connect()
+
+ transaction = connection.begin_twophase()
+ connection.execute(users.insert(), user_id=1, user_name='user1')
+ transaction.prepare()
+
+ connection.close()
+ connection2 = testbase.db.connect()
+
+ self.assertEquals(
+ connection2.execute(select([users.c.user_id]).order_by(users.c.user_id)).fetchall(),
+ []
+ )
+
+ recoverables = connection2.recover_twophase()
+ self.assertTrue(
+ transaction.xid in recoverables
+ )
+
+ connection2.commit_prepared(transaction.xid, recover=True)
+
+ self.assertEquals(
+ connection2.execute(select([users.c.user_id]).order_by(users.c.user_id)).fetchall(),
+ [(1,)]
+ )
+ connection2.close()
+
+class AutoRollbackTest(PersistTest):
def setUpAll(self):
global metadata
metadata = MetaData()
@@ -123,7 +269,7 @@ class AutoRollbackTest(testbase.PersistTest):
def tearDownAll(self):
metadata.drop_all(testbase.db)
- @testbase.unsupported('sqlite')
+ @testing.unsupported('sqlite')
def testrollback_deadlock(self):
"""test that returning connections to the pool clears any object locks."""
conn1 = testbase.db.connect()
@@ -131,6 +277,7 @@ class AutoRollbackTest(testbase.PersistTest):
users = Table('deadlock_users', metadata,
Column('user_id', INT, primary_key = True),
Column('user_name', VARCHAR(20)),
+ test_needs_acid=True,
)
users.create(conn1)
conn1.execute("select * from deadlock_users")
@@ -141,15 +288,15 @@ class AutoRollbackTest(testbase.PersistTest):
users.drop(conn2)
conn2.close()
-class TLTransactionTest(testbase.PersistTest):
+class TLTransactionTest(PersistTest):
def setUpAll(self):
global users, metadata, tlengine
- tlengine = create_engine(testbase.db_uri, strategy='threadlocal')
+ tlengine = create_engine(testbase.db.url, strategy='threadlocal')
metadata = MetaData()
users = Table('query_users', metadata,
Column('user_id', INT, primary_key = True),
Column('user_name', VARCHAR(20)),
- mysql_engine='InnoDB'
+ test_needs_acid=True,
)
users.create(tlengine)
def tearDown(self):
@@ -254,7 +401,7 @@ class TLTransactionTest(testbase.PersistTest):
finally:
external_connection.close()
- @testbase.unsupported('sqlite')
+ @testing.unsupported('sqlite')
def testnesting(self):
"""tests nesting of tranacstions"""
external_connection = tlengine.connect()
@@ -330,7 +477,7 @@ class TLTransactionTest(testbase.PersistTest):
try:
mapper(User, users)
- sess = create_session(bind_to=tlengine)
+ sess = create_session(bind=tlengine)
tlengine.begin()
u = User()
sess.save(u)
@@ -347,6 +494,127 @@ class TLTransactionTest(testbase.PersistTest):
assert c1.connection is c2.connection
c2.close()
assert c1.connection.connection is not None
+
+class ForUpdateTest(PersistTest):
+ def setUpAll(self):
+ global counters, metadata
+ metadata = MetaData()
+ counters = Table('forupdate_counters', metadata,
+ Column('counter_id', INT, primary_key = True),
+ Column('counter_value', INT),
+ test_needs_acid=True,
+ )
+ counters.create(testbase.db)
+ def tearDown(self):
+ testbase.db.connect().execute(counters.delete())
+ def tearDownAll(self):
+ counters.drop(testbase.db)
+
+ def increment(self, count, errors, update_style=True, delay=0.005):
+ con = testbase.db.connect()
+ sel = counters.select(for_update=update_style,
+ whereclause=counters.c.counter_id==1)
+
+ for i in xrange(count):
+ trans = con.begin()
+ try:
+ existing = con.execute(sel).fetchone()
+ incr = existing['counter_value'] + 1
+
+ time.sleep(delay)
+ con.execute(counters.update(counters.c.counter_id==1,
+ values={'counter_value':incr}))
+ time.sleep(delay)
+
+ readback = con.execute(sel).fetchone()
+ if (readback['counter_value'] != incr):
+ raise AssertionError("Got %s post-update, expected %s" %
+ (readback['counter_value'], incr))
+ trans.commit()
+ except Exception, e:
+ trans.rollback()
+ errors.append(e)
+ break
+
+ con.close()
+
+ @testing.supported('mysql', 'oracle', 'postgres')
+ def testqueued_update(self):
+ """Test SELECT FOR UPDATE with concurrent modifications.
+
+ Runs concurrent modifications on a single row in the users table,
+ with each mutator trying to increment a value stored in user_name.
+ """
+
+ db = testbase.db
+ db.execute(counters.insert(), counter_id=1, counter_value=0)
+
+ iterations, thread_count = 10, 5
+ threads, errors = [], []
+ for i in xrange(thread_count):
+ thread = threading.Thread(target=self.increment,
+ args=(iterations,),
+ kwargs={'errors': errors,
+ 'update_style': True})
+ thread.start()
+ threads.append(thread)
+ for thread in threads:
+ thread.join()
+
+ for e in errors:
+ sys.stderr.write("Failure: %s\n" % e)
+
+ self.assert_(len(errors) == 0)
+
+ sel = counters.select(whereclause=counters.c.counter_id==1)
+ final = db.execute(sel).fetchone()
+ self.assert_(final['counter_value'] == iterations * thread_count)
+
+ def overlap(self, ids, errors, update_style):
+ sel = counters.select(for_update=update_style,
+ whereclause=counters.c.counter_id.in_(*ids))
+ con = testbase.db.connect()
+ trans = con.begin()
+ try:
+ rows = con.execute(sel).fetchall()
+ time.sleep(0.25)
+ trans.commit()
+ except Exception, e:
+ trans.rollback()
+ errors.append(e)
+
+ def _threaded_overlap(self, thread_count, groups, update_style=True, pool=5):
+ db = testbase.db
+ for cid in range(pool - 1):
+ db.execute(counters.insert(), counter_id=cid + 1, counter_value=0)
+
+ errors, threads = [], []
+ for i in xrange(thread_count):
+ thread = threading.Thread(target=self.overlap,
+ args=(groups.pop(0), errors, update_style))
+ thread.start()
+ threads.append(thread)
+ for thread in threads:
+ thread.join()
+
+ return errors
+
+ @testing.supported('mysql', 'oracle', 'postgres')
+ def testqueued_select(self):
+ """Simple SELECT FOR UPDATE conflict test"""
+
+ errors = self._threaded_overlap(2, [(1,2,3),(3,4,5)])
+ for e in errors:
+ sys.stderr.write("Failure: %s\n" % e)
+ self.assert_(len(errors) == 0)
+
+ @testing.supported('oracle', 'postgres')
+ def testnowait_select(self):
+ """Simple SELECT FOR UPDATE NOWAIT conflict test"""
+
+ errors = self._threaded_overlap(2, [(1,2,3),(3,4,5)],
+ update_style='nowait')
+ self.assert_(len(errors) != 0)
if __name__ == "__main__":
testbase.main()
diff --git a/test/ext/activemapper.py b/test/ext/activemapper.py
index ebb832fdc..e28c72cd7 100644
--- a/test/ext/activemapper.py
+++ b/test/ext/activemapper.py
@@ -1,16 +1,18 @@
import testbase
+from datetime import datetime
+
from sqlalchemy.ext.activemapper import ActiveMapper, column, one_to_many, one_to_one, many_to_many, objectstore
-from sqlalchemy import and_, or_, clear_mappers, backref, create_session, exceptions
+from sqlalchemy import and_, or_, exceptions
from sqlalchemy import ForeignKey, String, Integer, DateTime, Table, Column
-from datetime import datetime
-import sqlalchemy
-
+from sqlalchemy.orm import clear_mappers, backref, create_session, class_mapper
import sqlalchemy.ext.activemapper as activemapper
+import sqlalchemy
+from testlib import *
-class testcase(testbase.PersistTest):
+class testcase(PersistTest):
def setUpAll(self):
- sqlalchemy.clear_mappers()
+ clear_mappers()
objectstore.clear()
global Person, Preferences, Address
@@ -133,7 +135,7 @@ class testcase(testbase.PersistTest):
objectstore.flush()
objectstore.clear()
- results = Person.select()
+ results = Person.query.select()
self.assertEquals(len(results), 1)
@@ -142,30 +144,30 @@ class testcase(testbase.PersistTest):
self.assertEquals(len(person.addresses), 2)
self.assertEquals(person.addresses[0].postal_code, '30338')
- @testbase.unsupported('mysql')
+ @testing.unsupported('mysql')
def test_update(self):
p1 = self.create_person_one()
objectstore.flush()
objectstore.clear()
- person = Person.select()[0]
+ person = Person.query.select()[0]
person.gender = 'F'
objectstore.flush()
objectstore.clear()
self.assertEquals(person.row_version, 2)
- person = Person.select()[0]
+ person = Person.query.select()[0]
person.gender = 'M'
objectstore.flush()
objectstore.clear()
self.assertEquals(person.row_version, 3)
#TODO: check that a concurrent modification raises exception
- p1 = Person.select()[0]
+ p1 = Person.query.select()[0]
s1 = objectstore.session
s2 = create_session()
objectstore.context.current = s2
- p2 = Person.select()[0]
+ p2 = Person.query.select()[0]
p1.first_name = "jack"
p2.first_name = "ed"
objectstore.flush()
@@ -185,14 +187,14 @@ class testcase(testbase.PersistTest):
objectstore.flush()
objectstore.clear()
- results = Person.select()
+ results = Person.query.select()
self.assertEquals(len(results), 1)
results[0].delete()
objectstore.flush()
objectstore.clear()
- results = Person.select()
+ results = Person.query.select()
self.assertEquals(len(results), 0)
@@ -204,7 +206,7 @@ class testcase(testbase.PersistTest):
objectstore.clear()
# select and make sure we get back two results
- people = Person.select()
+ people = Person.query.select()
self.assertEquals(len(people), 2)
# make sure that our backwards relationships work
@@ -212,7 +214,7 @@ class testcase(testbase.PersistTest):
self.assertEquals(people[1].addresses[0].person.id, p2.id)
# try a more complex select
- results = Person.select(
+ results = Person.query.select(
or_(
and_(
Address.c.person_id == Person.c.id,
@@ -253,17 +255,16 @@ class testcase(testbase.PersistTest):
objectstore.flush()
objectstore.clear()
- results = Person.select(
- Address.c.postal_code.like('30075') &
- Person.join_to('addresses')
+ results = Person.query.join('addresses').select(
+ Address.c.postal_code.like('30075')
)
self.assertEquals(len(results), 1)
- self.assertEquals(Person.count(), 2)
+ self.assertEquals(Person.query.count(), 2)
-class testmanytomany(testbase.PersistTest):
+class testmanytomany(PersistTest):
def setUpAll(self):
- sqlalchemy.clear_mappers()
+ clear_mappers()
objectstore.clear()
global secondarytable, foo, baz
secondarytable = Table("secondarytable",
@@ -299,8 +300,8 @@ class testmanytomany(testbase.PersistTest):
objectstore.flush()
objectstore.clear()
- foo1 = foo.get_by(name='foo1')
- baz1 = baz.get_by(name='baz1')
+ foo1 = foo.query.get_by(name='foo1')
+ baz1 = baz.query.get_by(name='baz1')
# Just checking ...
assert (foo1.name == 'foo1')
@@ -313,14 +314,12 @@ class testmanytomany(testbase.PersistTest):
# Optimistically based on activemapper one_to_many test, try to append
# baz1 to foo1.bazrel - (AttributeError: 'foo' object has no attribute 'bazrel')
- print sqlalchemy.class_mapper(foo).props
- print sqlalchemy.class_mapper(baz).props
foo1.bazrel.append(baz1)
assert (foo1.bazrel == [baz1])
-class testselfreferential(testbase.PersistTest):
+class testselfreferential(PersistTest):
def setUpAll(self):
- sqlalchemy.clear_mappers()
+ clear_mappers()
objectstore.clear()
global TreeNode
class TreeNode(activemapper.ActiveMapper):
@@ -343,15 +342,15 @@ class testselfreferential(testbase.PersistTest):
objectstore.flush()
objectstore.clear()
- t = TreeNode.get_by(name='node1')
+ t = TreeNode.query.get_by(name='node1')
assert (t.name == 'node1')
assert (t.children[0].name == 'node2')
assert (t.children[1].name == 'node3')
assert (t.children[1].parent is t)
objectstore.clear()
- t = TreeNode.get_by(name='node3')
- assert (t.parent is TreeNode.get_by(name='node1'))
+ t = TreeNode.query.get_by(name='node3')
+ assert (t.parent is TreeNode.query.get_by(name='node1'))
if __name__ == '__main__':
testbase.main()
diff --git a/test/ext/alltests.py b/test/ext/alltests.py
index 713601c3b..589f0f68f 100644
--- a/test/ext/alltests.py
+++ b/test/ext/alltests.py
@@ -3,7 +3,6 @@ import unittest, doctest
def suite():
unittest_modules = ['ext.activemapper',
- 'ext.selectresults',
'ext.assignmapper',
'ext.orderinglist',
'ext.associationproxy']
diff --git a/test/ext/assignmapper.py b/test/ext/assignmapper.py
index 650994987..31b3dd576 100644
--- a/test/ext/assignmapper.py
+++ b/test/ext/assignmapper.py
@@ -1,12 +1,13 @@
-from testbase import PersistTest, AssertMixin
import testbase
from sqlalchemy import *
-
+from sqlalchemy.orm import create_session, clear_mappers, relation, class_mapper
from sqlalchemy.ext.assignmapper import assign_mapper
from sqlalchemy.ext.sessioncontext import SessionContext
+from testlib import *
+
-class OverrideAttributesTest(PersistTest):
+class AssignMapperTest(PersistTest):
def setUpAll(self):
global metadata, table, table2
metadata = MetaData(testbase.db)
@@ -18,25 +19,18 @@ class OverrideAttributesTest(PersistTest):
Column('someid', None, ForeignKey('sometable.id'))
)
metadata.create_all()
- def tearDownAll(self):
- metadata.drop_all()
- def tearDown(self):
- clear_mappers()
+
def setUp(self):
- pass
- def test_override_attributes(self):
+ global SomeObject, SomeOtherObject, ctx
class SomeObject(object):pass
class SomeOtherObject(object):pass
ctx = SessionContext(create_session)
assign_mapper(ctx, SomeObject, table, properties={
- # this is the current workaround for class attribute name/collection collision: specify collection_class
- # explicitly. when we do away with class attributes specifying collection classes, this wont be
- # needed anymore.
- 'options':relation(SomeOtherObject, collection_class=list)
+ 'options':relation(SomeOtherObject)
})
assign_mapper(ctx, SomeOtherObject, table2)
- class_mapper(SomeObject)
+
s = SomeObject()
s.id = 1
s.data = 'hello'
@@ -44,8 +38,42 @@ class OverrideAttributesTest(PersistTest):
s.options.append(sso)
ctx.current.flush()
ctx.current.clear()
+
+ def tearDownAll(self):
+ metadata.drop_all()
+ def tearDown(self):
+ for table in metadata.table_iterator(reverse=True):
+ table.delete().execute()
+ clear_mappers()
+
+ def test_override_attributes(self):
+
+ sso = SomeOtherObject.query().first()
- assert SomeObject.get_by(id=s.id).options[0].id == sso.id
+ assert SomeObject.query.filter_by(id=1).one().options[0].id == sso.id
+
+ s2 = SomeObject(someid=12)
+ s3 = SomeOtherObject(someid=123, bogus=345)
+
+ class ValidatedOtherObject(object):pass
+ assign_mapper(ctx, ValidatedOtherObject, table2, validate=True)
+
+ v1 = ValidatedOtherObject(someid=12)
+ try:
+ v2 = ValidatedOtherObject(someid=12, bogus=345)
+ assert False
+ except exceptions.ArgumentError:
+ pass
+
+ def test_dont_clobber_methods(self):
+ class MyClass(object):
+ def expunge(self):
+ return "an expunge !"
+
+ assign_mapper(ctx, MyClass, table2)
+
+ assert MyClass().expunge() == "an expunge !"
+
if __name__ == '__main__':
testbase.main()
diff --git a/test/ext/associationproxy.py b/test/ext/associationproxy.py
index 3b18581bb..f602871c2 100644
--- a/test/ext/associationproxy.py
+++ b/test/ext/associationproxy.py
@@ -1,17 +1,19 @@
-from testbase import PersistTest
-import sqlalchemy.util as util
-import unittest
import testbase
+
from sqlalchemy import *
+from sqlalchemy.orm import *
+from sqlalchemy.orm.collections import collection
from sqlalchemy.ext.associationproxy import *
+from testlib import *
-db = testbase.db
class DictCollection(dict):
+ @collection.appender
def append(self, obj):
self[obj.foo] = obj
- def __iter__(self):
- return self.itervalues()
+ @collection.remover
+ def remove(self, obj):
+ del self[obj.foo]
class SetCollection(set):
pass
@@ -22,18 +24,20 @@ class ListCollection(list):
class ObjectCollection(object):
def __init__(self):
self.values = list()
+ @collection.appender
def append(self, obj):
self.values.append(obj)
+ @collection.remover
+ def remove(self, obj):
+ self.values.remove(obj)
def __iter__(self):
return iter(self.values)
- def clear(self):
- self.values.clear()
class _CollectionOperations(PersistTest):
def setUp(self):
collection_class = self.collection_class
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
parents_table = Table('Parent', metadata,
Column('id', Integer, primary_key=True),
@@ -131,9 +135,57 @@ class _CollectionOperations(PersistTest):
self.assert_(len(p1._children) == 3)
self.assert_(len(p1.children) == 3)
+ popped = p1.children.pop()
+ self.assert_(len(p1.children) == 2)
+ self.assert_(popped not in p1.children)
+ p1 = self.roundtrip(p1)
+ self.assert_(len(p1.children) == 2)
+ self.assert_(popped not in p1.children)
+
+ p1.children[1] = 'changed-in-place'
+ self.assert_(p1.children[1] == 'changed-in-place')
+ inplace_id = p1._children[1].id
+ p1 = self.roundtrip(p1)
+ self.assert_(p1.children[1] == 'changed-in-place')
+ assert p1._children[1].id == inplace_id
+
+ p1.children.append('changed-in-place')
+ self.assert_(p1.children.count('changed-in-place') == 2)
+
+ p1.children.remove('changed-in-place')
+ self.assert_(p1.children.count('changed-in-place') == 1)
+
+ p1 = self.roundtrip(p1)
+ self.assert_(p1.children.count('changed-in-place') == 1)
+
p1._children = []
self.assert_(len(p1.children) == 0)
+ after = ['a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j']
+ p1.children = ['a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j']
+ self.assert_(len(p1.children) == 10)
+ self.assert_([c.name for c in p1._children] == after)
+
+ p1.children[2:6] = ['x'] * 4
+ after = ['a', 'b', 'x', 'x', 'x', 'x', 'g', 'h', 'i', 'j']
+ self.assert_(p1.children == after)
+ self.assert_([c.name for c in p1._children] == after)
+
+ p1.children[2:6] = ['y']
+ after = ['a', 'b', 'y', 'g', 'h', 'i', 'j']
+ self.assert_(p1.children == after)
+ self.assert_([c.name for c in p1._children] == after)
+
+ p1.children[2:3] = ['z'] * 4
+ after = ['a', 'b', 'z', 'z', 'z', 'z', 'g', 'h', 'i', 'j']
+ self.assert_(p1.children == after)
+ self.assert_([c.name for c in p1._children] == after)
+
+ p1.children[2::2] = ['O'] * 4
+ after = ['a', 'b', 'O', 'z', 'O', 'z', 'O', 'h', 'O', 'j']
+ self.assert_(p1.children == after)
+ self.assert_([c.name for c in p1._children] == after)
+
class DefaultTest(_CollectionOperations):
def __init__(self, *args, **kw):
super(DefaultTest, self).__init__(*args, **kw)
@@ -218,12 +270,27 @@ class CustomDictTest(DictTest):
self.assert_(len(p1._children) == 3)
self.assert_(len(p1.children) == 3)
- p1.children['d'] = 'new d'
- assert p1.children['d'] == 'new d'
+ p1.children['e'] = 'changed-in-place'
+ self.assert_(p1.children['e'] == 'changed-in-place')
+ inplace_id = p1._children['e'].id
+ p1 = self.roundtrip(p1)
+ self.assert_(p1.children['e'] == 'changed-in-place')
+ self.assert_(p1._children['e'].id == inplace_id)
p1._children = {}
self.assert_(len(p1.children) == 0)
+ try:
+ p1._children = []
+ self.assert_(False)
+ except exceptions.ArgumentError:
+ self.assert_(True)
+
+ try:
+ p1._children = None
+ self.assert_(False)
+ except exceptions.ArgumentError:
+ self.assert_(True)
class SetTest(_CollectionOperations):
def __init__(self, *args, **kw):
@@ -239,7 +306,7 @@ class SetTest(_CollectionOperations):
self.assert_(not p1.children)
ch1 = Child('regular')
- p1._children.append(ch1)
+ p1._children.add(ch1)
self.assert_(ch1 in p1._children)
self.assert_(len(p1._children) == 1)
@@ -256,7 +323,8 @@ class SetTest(_CollectionOperations):
self.assert_(len(p1.children) == 2)
self.assert_(len(p1._children) == 2)
- self.assert_(set([o.name for o in p1._children]) == set(['regular', 'proxied']))
+ self.assert_(set([o.name for o in p1._children]) ==
+ set(['regular', 'proxied']))
ch2 = None
for o in p1._children:
@@ -322,9 +390,22 @@ class SetTest(_CollectionOperations):
p1 = self.roundtrip(p1)
self.assert_(p1.children == set(['c']))
- p1._children = []
+ p1._children = set()
self.assert_(len(p1.children) == 0)
+ try:
+ p1._children = []
+ self.assert_(False)
+ except exceptions.ArgumentError:
+ self.assert_(True)
+
+ try:
+ p1._children = None
+ self.assert_(False)
+ except exceptions.ArgumentError:
+ self.assert_(True)
+
+
def test_set_comparisons(self):
Parent, Child = self.Parent, self.Child
@@ -393,14 +474,7 @@ class SetTest(_CollectionOperations):
print 'want', repr(control)
print 'got', repr(p.children)
raise
-
- # workaround for bug #548
- def test_set_pop(self):
- Parent, Child = self.Parent, self.Child
- p = Parent('p1')
- p.children.add('a')
- p.children.pop()
- self.assert_(True)
+
class CustomSetTest(SetTest):
def __init__(self, *args, **kw):
@@ -434,7 +508,7 @@ class CustomObjectTest(_CollectionOperations):
class ScalarTest(PersistTest):
def test_scalar_proxy(self):
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
parents_table = Table('Parent', metadata,
Column('id', Integer, primary_key=True),
@@ -550,7 +624,7 @@ class ScalarTest(PersistTest):
class LazyLoadTest(PersistTest):
def setUp(self):
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
parents_table = Table('Parent', metadata,
Column('id', Integer, primary_key=True),
@@ -606,7 +680,7 @@ class LazyLoadTest(PersistTest):
# Is there a better way to ensure that the association_proxy
# didn't convert a lazy load to an eager load? This does work though.
self.assert_('_children' not in p.__dict__)
- self.assert_(len(p._children.data) == 3)
+ self.assert_(len(p._children) == 3)
self.assert_('_children' in p.__dict__)
def test_eager_list(self):
@@ -622,7 +696,7 @@ class LazyLoadTest(PersistTest):
p = self.roundtrip(p)
self.assert_('_children' in p.__dict__)
- self.assert_(len(p._children.data) == 3)
+ self.assert_(len(p._children) == 3)
def test_lazy_scalar(self):
Parent, Child = self.Parent, self.Child
diff --git a/test/ext/legacy_objectstore.py b/test/ext/legacy_objectstore.py
deleted file mode 100644
index 3aa99a1ae..000000000
--- a/test/ext/legacy_objectstore.py
+++ /dev/null
@@ -1,113 +0,0 @@
-from testbase import PersistTest, AssertMixin
-import unittest, sys, os
-from sqlalchemy import *
-import StringIO
-import testbase
-
-from tables import *
-import tables
-
-install_mods('legacy_session')
-
-
-class LegacySessionTest(AssertMixin):
- def setUpAll(self):
- db.echo = False
- users.create()
- db.echo = testbase.echo
- def tearDownAll(self):
- db.echo = False
- users.drop()
- db.echo = testbase.echo
- def setUp(self):
- objectstore.get_session().clear()
- clear_mappers()
- tables.user_data()
- #db.echo = "debug"
- def tearDown(self):
- tables.delete_user_data()
-
- def test_nested_begin_commit(self):
- """tests that nesting objectstore transactions with multiple commits
- affects only the outermost transaction"""
- class User(object):pass
- m = mapper(User, users)
- def name_of(id):
- return users.select(users.c.user_id == id).execute().fetchone().user_name
- name1 = "Oliver Twist"
- name2 = 'Mr. Bumble'
- self.assert_(name_of(7) != name1, msg="user_name should not be %s" % name1)
- self.assert_(name_of(8) != name2, msg="user_name should not be %s" % name2)
- s = objectstore.get_session()
- trans = s.begin()
- trans2 = s.begin()
- m.get(7).user_name = name1
- trans3 = s.begin()
- m.get(8).user_name = name2
- trans3.commit()
- s.commit() # should do nothing
- self.assert_(name_of(7) != name1, msg="user_name should not be %s" % name1)
- self.assert_(name_of(8) != name2, msg="user_name should not be %s" % name2)
- trans2.commit()
- s.commit() # should do nothing
- self.assert_(name_of(7) != name1, msg="user_name should not be %s" % name1)
- self.assert_(name_of(8) != name2, msg="user_name should not be %s" % name2)
- trans.commit()
- self.assert_(name_of(7) == name1, msg="user_name should be %s" % name1)
- self.assert_(name_of(8) == name2, msg="user_name should be %s" % name2)
-
- def test_nested_rollback(self):
- """tests that nesting objectstore transactions with a rollback inside
- affects only the outermost transaction"""
- class User(object):pass
- m = mapper(User, users)
- def name_of(id):
- return users.select(users.c.user_id == id).execute().fetchone().user_name
- name1 = "Oliver Twist"
- name2 = 'Mr. Bumble'
- self.assert_(name_of(7) != name1, msg="user_name should not be %s" % name1)
- self.assert_(name_of(8) != name2, msg="user_name should not be %s" % name2)
- s = objectstore.get_session()
- trans = s.begin()
- trans2 = s.begin()
- m.get(7).user_name = name1
- trans3 = s.begin()
- m.get(8).user_name = name2
- trans3.rollback()
- self.assert_(name_of(7) != name1, msg="user_name should not be %s" % name1)
- self.assert_(name_of(8) != name2, msg="user_name should not be %s" % name2)
- trans2.commit()
- self.assert_(name_of(7) != name1, msg="user_name should not be %s" % name1)
- self.assert_(name_of(8) != name2, msg="user_name should not be %s" % name2)
- trans.commit()
- self.assert_(name_of(7) != name1, msg="user_name should not be %s" % name1)
- self.assert_(name_of(8) != name2, msg="user_name should not be %s" % name2)
-
- def test_true_nested(self):
- """tests creating a new Session inside a database transaction, in
- conjunction with an engine-level nested transaction, which uses
- a second connection in order to achieve a nested transaction that commits, inside
- of another engine session that rolls back."""
-# testbase.db.echo='debug'
- class User(object):
- pass
- testbase.db.begin()
- try:
- m = mapper(User, users)
- name1 = "Oliver Twist"
- name2 = 'Mr. Bumble'
- m.get(7).user_name = name1
- s = objectstore.Session(nest_on=testbase.db)
- m.using(s).get(8).user_name = name2
- s.commit()
- objectstore.commit()
- testbase.db.rollback()
- except:
- testbase.db.rollback()
- raise
- objectstore.clear()
- self.assert_(m.get(8).user_name == name2)
- self.assert_(m.get(7).user_name != name1)
-
-if __name__ == "__main__":
- testbase.main()
diff --git a/test/ext/orderinglist.py b/test/ext/orderinglist.py
index 6dcf057d4..d16e20da7 100644
--- a/test/ext/orderinglist.py
+++ b/test/ext/orderinglist.py
@@ -1,11 +1,10 @@
-from testbase import PersistTest
-import sqlalchemy.util as util
-import unittest, sys, os
import testbase
+
from sqlalchemy import *
+from sqlalchemy.orm import *
from sqlalchemy.ext.orderinglist import *
+from testlib import *
-db = testbase.db
metadata = None
# order in whole steps
@@ -52,7 +51,7 @@ class OrderingListTest(PersistTest):
global metadata, slides_table, bullets_table, Slide, Bullet
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
slides_table = Table('test_Slides', metadata,
Column('id', Integer, primary_key=True),
Column('name', String))
@@ -297,43 +296,7 @@ class OrderingListTest(PersistTest):
self.assert_(srt.bullets[i].position == i)
self.assert_(srt.bullets[i].text == text)
- def test_replace1(self):
- self._setup(ordering_list('position'))
-
- s1 = Slide('Slide #1')
- s1.bullets = [ Bullet('1'), Bullet('2'), Bullet('3') ]
-
- self.assert_(len(s1.bullets) == 3)
- self.assert_(s1.bullets[2].position == 2)
-
- session = create_session()
- session.save(s1)
- session.flush()
-
- new_bullet = Bullet('new 2')
- self.assert_(new_bullet.position is None)
-
- # naive replacement, no database deletion should occur
- # with current InstrumentedList __setitem__ semantics
- s1.bullets[1] = new_bullet
-
- self.assert_(new_bullet.position == 1)
- self.assert_(len(s1.bullets) == 3)
-
- id = s1.id
-
- session.flush()
- session.clear()
-
- srt = session.query(Slide).get(id)
-
- self.assert_(srt.bullets)
- self.assert_(len(srt.bullets) == 4)
-
- self.assert_(srt.bullets[1].text == '2')
- self.assert_(srt.bullets[2].text == 'new 2')
-
- def test_replace2(self):
+ def test_replace(self):
self._setup(ordering_list('position'))
s1 = Slide('Slide #1')
@@ -350,7 +313,7 @@ class OrderingListTest(PersistTest):
self.assert_(new_bullet.position is None)
# mark existing bullet as db-deleted before replacement.
- session.delete(s1.bullets[1])
+ #session.delete(s1.bullets[1])
s1.bullets[1] = new_bullet
self.assert_(new_bullet.position == 1)
diff --git a/test/ext/selectresults.py b/test/ext/selectresults.py
deleted file mode 100644
index 1ec724c3d..000000000
--- a/test/ext/selectresults.py
+++ /dev/null
@@ -1,239 +0,0 @@
-from testbase import PersistTest, AssertMixin
-import testbase
-import tables
-
-from sqlalchemy import *
-
-from sqlalchemy.ext.selectresults import SelectResultsExt, SelectResults
-
-class Foo(object):
- pass
-
-class SelectResultsTest(PersistTest):
- def setUpAll(self):
- self.install_threadlocal()
- global foo, metadata
- metadata = MetaData(testbase.db)
- foo = Table('foo', metadata,
- Column('id', Integer, Sequence('foo_id_seq'), primary_key=True),
- Column('bar', Integer),
- Column('range', Integer))
-
- assign_mapper(Foo, foo, extension=SelectResultsExt())
- metadata.create_all()
- for i in range(100):
- Foo(bar=i, range=i%10)
- objectstore.flush()
-
- def setUp(self):
- self.query = Query(Foo)
- self.orig = self.query.select_whereclause()
- self.res = self.query.select()
-
- def tearDownAll(self):
- metadata.drop_all()
- self.uninstall_threadlocal()
- clear_mappers()
-
- def test_selectby(self):
- res = self.query.select_by(range=5)
- assert res.order_by([Foo.c.bar])[0].bar == 5
- assert res.order_by([desc(Foo.c.bar)])[0].bar == 95
-
- @testbase.unsupported('mssql')
- def test_slice(self):
- assert self.res[1] == self.orig[1]
- assert list(self.res[10:20]) == self.orig[10:20]
- assert list(self.res[10:]) == self.orig[10:]
- assert list(self.res[:10]) == self.orig[:10]
- assert list(self.res[:10]) == self.orig[:10]
- assert list(self.res[10:40:3]) == self.orig[10:40:3]
- assert list(self.res[-5:]) == self.orig[-5:]
- assert self.res[10:20][5] == self.orig[10:20][5]
-
- @testbase.supported('mssql')
- def test_slice_mssql(self):
- assert list(self.res[:10]) == self.orig[:10]
- assert list(self.res[:10]) == self.orig[:10]
-
- def test_aggregate(self):
- assert self.res.count() == 100
- assert self.res.filter(foo.c.bar<30).min(foo.c.bar) == 0
- assert self.res.filter(foo.c.bar<30).max(foo.c.bar) == 29
-
- @testbase.unsupported('mysql')
- def test_aggregate_1(self):
- # this one fails in mysql as the result comes back as a string
- assert self.res.filter(foo.c.bar<30).sum(foo.c.bar) == 435
-
- @testbase.unsupported('postgres', 'mysql', 'firebird', 'mssql')
- def test_aggregate_2(self):
- assert self.res.filter(foo.c.bar<30).avg(foo.c.bar) == 14.5
-
- @testbase.supported('postgres', 'mysql', 'firebird', 'mssql')
- def test_aggregate_2_int(self):
- assert int(self.res.filter(foo.c.bar<30).avg(foo.c.bar)) == 14
-
- def test_filter(self):
- assert self.res.count() == 100
- assert self.res.filter(Foo.c.bar < 30).count() == 30
- res2 = self.res.filter(Foo.c.bar < 30).filter(Foo.c.bar > 10)
- assert res2.count() == 19
-
- def test_options(self):
- class ext1(MapperExtension):
- def populate_instance(self, mapper, selectcontext, row, instance, identitykey, isnew):
- instance.TEST = "hello world"
- return EXT_PASS
- objectstore.clear()
- assert self.res.options(extension(ext1()))[0].TEST == "hello world"
-
- def test_order_by(self):
- assert self.res.order_by([Foo.c.bar])[0].bar == 0
- assert self.res.order_by([desc(Foo.c.bar)])[0].bar == 99
-
- def test_offset(self):
- assert list(self.res.order_by([Foo.c.bar]).offset(10))[0].bar == 10
-
- def test_offset(self):
- assert len(list(self.res.limit(10))) == 10
-
-class Obj1(object):
- pass
-class Obj2(object):
- pass
-
-class SelectResultsTest2(PersistTest):
- def setUpAll(self):
- self.install_threadlocal()
- global metadata, table1, table2
- metadata = MetaData(testbase.db)
- table1 = Table('Table1', metadata,
- Column('id', Integer, primary_key=True),
- )
- table2 = Table('Table2', metadata,
- Column('t1id', Integer, ForeignKey("Table1.id"), primary_key=True),
- Column('num', Integer, primary_key=True),
- )
- assign_mapper(Obj1, table1, extension=SelectResultsExt())
- assign_mapper(Obj2, table2, extension=SelectResultsExt())
- metadata.create_all()
- table1.insert().execute({'id':1},{'id':2},{'id':3},{'id':4})
- table2.insert().execute({'num':1,'t1id':1},{'num':2,'t1id':1},{'num':3,'t1id':1},\
-{'num':4,'t1id':2},{'num':5,'t1id':2},{'num':6,'t1id':3})
-
- def setUp(self):
- self.query = Query(Obj1)
- #self.orig = self.query.select_whereclause()
- #self.res = self.query.select()
-
- def tearDownAll(self):
- metadata.drop_all()
- self.uninstall_threadlocal()
- clear_mappers()
-
- def test_distinctcount(self):
- res = self.query.select()
- assert res.count() == 4
- res = self.query.select(and_(table1.c.id==table2.c.t1id,table2.c.t1id==1))
- assert res.count() == 3
- res = self.query.select(and_(table1.c.id==table2.c.t1id,table2.c.t1id==1), distinct=True)
- self.assertEqual(res.count(), 1)
-
-class RelationsTest(AssertMixin):
- def setUpAll(self):
- tables.create()
- tables.data()
- def tearDownAll(self):
- tables.drop()
- def tearDown(self):
- clear_mappers()
- def test_jointo(self):
- """test the join_to and outerjoin_to functions on SelectResults"""
- mapper(tables.User, tables.users, properties={
- 'orders':relation(mapper(tables.Order, tables.orders, properties={
- 'items':relation(mapper(tables.Item, tables.orderitems))
- }))
- })
- session = create_session()
- query = SelectResults(session.query(tables.User))
- x = query.join_to('orders').join_to('items').select(tables.Item.c.item_id==2)
- print x.compile()
- self.assert_result(list(x), tables.User, tables.user_result[2])
- def test_outerjointo(self):
- """test the join_to and outerjoin_to functions on SelectResults"""
- mapper(tables.User, tables.users, properties={
- 'orders':relation(mapper(tables.Order, tables.orders, properties={
- 'items':relation(mapper(tables.Item, tables.orderitems))
- }))
- })
- session = create_session()
- query = SelectResults(session.query(tables.User))
- x = query.outerjoin_to('orders').outerjoin_to('items').select(or_(tables.Order.c.order_id==None,tables.Item.c.item_id==2))
- print x.compile()
- self.assert_result(list(x), tables.User, *tables.user_result[1:3])
- def test_outerjointo_count(self):
- """test the join_to and outerjoin_to functions on SelectResults"""
- mapper(tables.User, tables.users, properties={
- 'orders':relation(mapper(tables.Order, tables.orders, properties={
- 'items':relation(mapper(tables.Item, tables.orderitems))
- }))
- })
- session = create_session()
- query = SelectResults(session.query(tables.User))
- x = query.outerjoin_to('orders').outerjoin_to('items').select(or_(tables.Order.c.order_id==None,tables.Item.c.item_id==2)).count()
- assert x==2
- def test_from(self):
- mapper(tables.User, tables.users, properties={
- 'orders':relation(mapper(tables.Order, tables.orders, properties={
- 'items':relation(mapper(tables.Item, tables.orderitems))
- }))
- })
- session = create_session()
- query = SelectResults(session.query(tables.User))
- x = query.select_from([tables.users.outerjoin(tables.orders).outerjoin(tables.orderitems)]).\
- filter(or_(tables.Order.c.order_id==None,tables.Item.c.item_id==2))
- print x.compile()
- self.assert_result(list(x), tables.User, *tables.user_result[1:3])
-
-
-class CaseSensitiveTest(PersistTest):
- def setUpAll(self):
- self.install_threadlocal()
- global metadata, table1, table2
- metadata = MetaData(testbase.db)
- table1 = Table('Table1', metadata,
- Column('ID', Integer, primary_key=True),
- )
- table2 = Table('Table2', metadata,
- Column('T1ID', Integer, ForeignKey("Table1.ID"), primary_key=True),
- Column('NUM', Integer, primary_key=True),
- )
- assign_mapper(Obj1, table1, extension=SelectResultsExt())
- assign_mapper(Obj2, table2, extension=SelectResultsExt())
- metadata.create_all()
- table1.insert().execute({'ID':1},{'ID':2},{'ID':3},{'ID':4})
- table2.insert().execute({'NUM':1,'T1ID':1},{'NUM':2,'T1ID':1},{'NUM':3,'T1ID':1},\
-{'NUM':4,'T1ID':2},{'NUM':5,'T1ID':2},{'NUM':6,'T1ID':3})
-
- def setUp(self):
- self.query = Query(Obj1)
- #self.orig = self.query.select_whereclause()
- #self.res = self.query.select()
-
- def tearDownAll(self):
- metadata.drop_all()
- self.uninstall_threadlocal()
- clear_mappers()
-
- def test_distinctcount(self):
- res = self.query.select()
- assert res.count() == 4
- res = self.query.select(and_(table1.c.ID==table2.c.T1ID,table2.c.T1ID==1))
- assert res.count() == 3
- res = self.query.select(and_(table1.c.ID==table2.c.T1ID,table2.c.T1ID==1), distinct=True)
- self.assertEqual(res.count(), 1)
-
-
-if __name__ == "__main__":
- testbase.main()
diff --git a/test/ext/wsgi_test.py b/test/ext/wsgi_test.py
deleted file mode 100644
index 1330f88b6..000000000
--- a/test/ext/wsgi_test.py
+++ /dev/null
@@ -1,122 +0,0 @@
-"""Interactive wsgi test
-
-Small WSGI application that uses a table and mapper defined at the module
-level, with per-application uris enabled by the ProxyEngine.
-
-Requires the wsgiutils package from:
-
-http://www.owlfish.com/software/wsgiutils/
-
-Run the script with python wsgi_test.py, then visit http://localhost:8080/a
-and http://localhost:8080/b with a browser. You should see two distinct lists
-of colors.
-"""
-
-from sqlalchemy import *
-from sqlalchemy.ext.proxy import ProxyEngine
-from wsgiutils import wsgiServer
-
-engine = ProxyEngine()
-
-colors = Table('colors', engine,
- Column('id', Integer, primary_key=True),
- Column('name', String(32)),
- Column('hex', String(6)))
-
-class Color(object):
- pass
-
-assign_mapper(Color, colors)
-
-data = { 'a': (('fff','white'), ('aaa','gray'), ('000','black'),
- ('f00', 'red'), ('0f0', 'green')),
- 'b': (('00f','blue'), ('ff0', 'yellow'), ('0ff','purple')) }
-
-db_uri = { 'a': 'sqlite://filename=wsgi_db_a.db',
- 'b': 'sqlite://filename=wsgi_db_b.db' }
-
-def app(dataset):
- print '... connecting to database %s: %s' % (dataset, db_uri[dataset])
- engine.connect(db_uri[dataset], echo=True, echo_pool=True)
- colors.create()
-
- print '... populating data into %s' % db_uri[dataset]
- for hex, name in data[dataset]:
- c = Color()
- c.hex = hex
- c.name = name
- objectstore.commit()
- objectstore.clear()
-
- def call(environ, start_response):
- engine.connect(db_uri[dataset], echo=True, echo_pool=True)
-
- # NOTE: must clear objectstore on each request, or you'll see
- # objects from another thread here
- objectstore.clear()
- objectstore.begin()
-
- c = Color.select()
-
- start_response('200 OK', [('content-type','text/html')])
- yield '<html><head><title>Test dataset %s</title></head>' % dataset
- yield '<body>'
- yield '<p>uri: %s</p>' % db_uri[dataset]
- yield '<p>engine: <xmp>%s</xmp></p>' % engine.engine
- yield '<p>Colors!</p>'
- for color in c:
- yield '<div style="background: #%s">%s</div>' % (color.hex,
- color.name)
- yield '</body></html>'
- return call
-
-def cleanup():
- for uri in db_uri.values():
- print "Cleaning db %s" % uri
- engine.connect(uri)
- colors.drop()
-
-def run_server(apps, host='localhost', port=8080):
- print "Serving test app at http://%s:%s/" % (host, port)
- print "Visit http://%(host)s:%(port)s/a and " \
- "http://%(host)s:%(port)s/b to test apps" % {'host': host,
- 'port': port}
-
- server = wsgiServer.WSGIServer((host, port), apps, serveFiles=False)
- try:
- server.serve_forever()
- except:
- cleanup()
- raise
-
-if __name__ == '__main__':
- run_server({'/a':app('a'), '/b':app('b')})
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
diff --git a/test/orm/alltests.py b/test/orm/alltests.py
index 35650c136..4f8f4b6b7 100644
--- a/test/orm/alltests.py
+++ b/test/orm/alltests.py
@@ -1,18 +1,20 @@
import testbase
import unittest
+import inheritance.alltests as inheritance
+import sharding.alltests as sharding
+
def suite():
modules_to_test = (
- 'orm.attributes',
- 'orm.mapper',
+ 'orm.attributes',
'orm.query',
'orm.lazy_relations',
'orm.eager_relations',
+ 'orm.mapper',
+ 'orm.collection',
'orm.generative',
'orm.lazytest1',
- 'orm.eagertest1',
- 'orm.eagertest2',
- 'orm.eagertest3',
+ 'orm.assorted_eager',
'orm.sessioncontext',
'orm.unitofwork',
@@ -24,20 +26,11 @@ def suite():
'orm.memusage',
'orm.cycles',
- 'orm.poly_linked_list',
'orm.entity',
'orm.compile',
'orm.manytomany',
'orm.onetoone',
- 'orm.inheritance',
- 'orm.inheritance2',
- 'orm.inheritance3',
- 'orm.inheritance4',
- 'orm.inheritance5',
- 'orm.abc_inheritance',
- 'orm.single',
- 'orm.polymorph'
)
alltests = unittest.TestSuite()
for name in modules_to_test:
@@ -45,6 +38,8 @@ def suite():
for token in name.split('.')[1:]:
mod = getattr(mod, token)
alltests.addTest(unittest.findTestCases(mod, suiteClass=None))
+ alltests.addTest(inheritance.suite())
+ alltests.addTest(sharding.suite())
return alltests
diff --git a/test/orm/association.py b/test/orm/association.py
index 416cfabbb..a2b899418 100644
--- a/test/orm/association.py
+++ b/test/orm/association.py
@@ -1,9 +1,10 @@
import testbase
from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
-
-class AssociationTest(testbase.PersistTest):
+class AssociationTest(PersistTest):
def setUpAll(self):
global items, item_keywords, keywords, metadata, Item, Keyword, KeywordAssociation
metadata = MetaData(testbase.db)
@@ -138,7 +139,7 @@ class AssociationTest(testbase.PersistTest):
sess.flush()
self.assert_(item_keywords.count().scalar() == 0)
-class AssociationTest2(testbase.PersistTest):
+class AssociationTest2(PersistTest):
def setUpAll(self):
global table_originals, table_people, table_isauthor, metadata, Originals, People, IsAuthor
metadata = MetaData(testbase.db)
diff --git a/test/orm/eagertest3.py b/test/orm/assorted_eager.py
index 8e7735812..652186b8e 100644
--- a/test/orm/eagertest3.py
+++ b/test/orm/assorted_eager.py
@@ -1,8 +1,11 @@
-from testbase import PersistTest, AssertMixin
+"""eager loading unittests derived from mailing list-reported problems and trac tickets."""
+
import testbase
+import random, datetime
from sqlalchemy import *
-from sqlalchemy.ext.selectresults import SelectResults
-import random
+from sqlalchemy.orm import *
+from sqlalchemy.ext.sessioncontext import SessionContext
+from testlib import *
class EagerTest(AssertMixin):
def setUpAll(self):
@@ -119,12 +122,12 @@ class EagerTest(AssertMixin):
assert result == [u'1 Some Category', u'3 Some Category']
def test_dslish(self):
- """test the same as witheagerload except building the query via SelectResults"""
+ """test the same as witheagerload except using generative"""
s = create_session()
- q=SelectResults(s.query(Test).options(eagerload('category')))
- l=q.select (
+ q=s.query(Test).options(eagerload('category'))
+ l=q.filter (
and_(tests.c.owner_id==1,or_(options.c.someoption==None,options.c.someoption==False))
- ).outerjoin_to('owner_option')
+ ).outerjoin('owner_option')
result = ["%d %s" % ( t.id,t.category.name ) for t in l]
print result
@@ -170,6 +173,7 @@ class EagerTest2(AssertMixin):
def tearDown(self):
for t in metadata.table_iterator(reverse=True):
t.delete().execute()
+
def testeagerterminate(self):
"""test that eager query generation does not include the same mapper's table twice.
@@ -189,7 +193,7 @@ class EagerTest2(AssertMixin):
'right': relation(Right, lazy=False, backref=backref('middle', lazy=False)),
}
)
- session = create_session(bind_to=testbase.db)
+ session = create_session(bind=testbase.db)
p = Middle('test1')
p.left.append(Left('tag1'))
p.right.append(Right('tag2'))
@@ -199,7 +203,7 @@ class EagerTest2(AssertMixin):
obj = session.query(Left).get_by(tag='tag1')
print obj.middle.right[0]
-class EagerTest3(testbase.ORMTest):
+class EagerTest3(ORMTest):
"""test eager loading combined with nested SELECT statements, functions, and aggregates"""
def define_tables(self, metadata):
global datas, foo, stats
@@ -267,7 +271,7 @@ class EagerTest3(testbase.ORMTest):
# algorithms and there are repeated 'somedata' values in the list)
assert verify_result == arb_result
-class EagerTest4(testbase.ORMTest):
+class EagerTest4(ORMTest):
def define_tables(self, metadata):
global departments, employees
departments = Table('departments', metadata,
@@ -315,17 +319,11 @@ class EagerTest4(testbase.ORMTest):
sess.flush()
q = sess.query(Department)
- filters = [q.join_to('employees'),
- Employee.c.name.startswith('J')]
+ q = q.join('employees').filter(Employee.c.name.startswith('J')).distinct().order_by([desc(Department.c.name)])
+ assert q.count() == 2
+ assert q[0] is d2
- d = SelectResults(q)
- d = d.join_to('employees').filter(Employee.c.name.startswith('J'))
- d = d.distinct()
- d = d.order_by([desc(Department.c.name)])
- assert d.count() == 2
- assert d[0] is d2
-
-class EagerTest5(testbase.ORMTest):
+class EagerTest5(ORMTest):
"""test the construction of AliasedClauses for the same eager load property but different
parent mappers, due to inheritance"""
def define_tables(self, metadata):
@@ -416,10 +414,270 @@ class EagerTest5(testbase.ORMTest):
# eager load had to succeed
assert len([c for c in d2.comments]) == 1
-class EagerTest6(testbase.ORMTest):
+class EagerTest6(ORMTest):
+ def define_tables(self, metadata):
+ global designType, design, part, inheritedPart
+ designType = Table('design_types', metadata,
+ Column('design_type_id', Integer, primary_key=True),
+ )
+
+ design =Table('design', metadata,
+ Column('design_id', Integer, primary_key=True),
+ Column('design_type_id', Integer, ForeignKey('design_types.design_type_id')))
+
+ part = Table('parts', metadata,
+ Column('part_id', Integer, primary_key=True),
+ Column('design_id', Integer, ForeignKey('design.design_id')),
+ Column('design_type_id', Integer, ForeignKey('design_types.design_type_id')))
+
+ inheritedPart = Table('inherited_part', metadata,
+ Column('ip_id', Integer, primary_key=True),
+ Column('part_id', Integer, ForeignKey('parts.part_id')),
+ Column('design_id', Integer, ForeignKey('design.design_id')),
+ )
+
+ def testone(self):
+ class Part(object):pass
+ class Design(object):pass
+ class DesignType(object):pass
+ class InheritedPart(object):pass
+
+ mapper(Part, part)
+
+ mapper(InheritedPart, inheritedPart, properties=dict(
+ part=relation(Part, lazy=False)
+ ))
+
+ mapper(Design, design, properties=dict(
+ parts=relation(Part, private=True, backref="design"),
+ inheritedParts=relation(InheritedPart, private=True, backref="design"),
+ ))
+
+ mapper(DesignType, designType, properties=dict(
+ # designs=relation(Design, private=True, backref="type"),
+ ))
+
+ class_mapper(Design).add_property("type", relation(DesignType, lazy=False, backref="designs"))
+ class_mapper(Part).add_property("design", relation(Design, lazy=False, backref="parts"))
+ #Part.mapper.add_property("designType", relation(DesignType))
+
+ d = Design()
+ sess = create_session()
+ sess.save(d)
+ sess.flush()
+ sess.clear()
+ x = sess.query(Design).get(1)
+ x.inheritedParts
+
+class EagerTest7(ORMTest):
+ def define_tables(self, metadata):
+ global companies_table, addresses_table, invoice_table, phones_table, items_table, ctx
+ global Company, Address, Phone, Item,Invoice
+
+ ctx = SessionContext(create_session)
+
+ companies_table = Table('companies', metadata,
+ Column('company_id', Integer, Sequence('company_id_seq', optional=True), primary_key = True),
+ Column('company_name', String(40)),
+
+ )
+
+ addresses_table = Table('addresses', metadata,
+ Column('address_id', Integer, Sequence('address_id_seq', optional=True), primary_key = True),
+ Column('company_id', Integer, ForeignKey("companies.company_id")),
+ Column('address', String(40)),
+ )
+
+ phones_table = Table('phone_numbers', metadata,
+ Column('phone_id', Integer, Sequence('phone_id_seq', optional=True), primary_key = True),
+ Column('address_id', Integer, ForeignKey('addresses.address_id')),
+ Column('type', String(20)),
+ Column('number', String(10)),
+ )
+
+ invoice_table = Table('invoices', metadata,
+ Column('invoice_id', Integer, Sequence('invoice_id_seq', optional=True), primary_key = True),
+ Column('company_id', Integer, ForeignKey("companies.company_id")),
+ Column('date', DateTime),
+ )
+
+ items_table = Table('items', metadata,
+ Column('item_id', Integer, Sequence('item_id_seq', optional=True), primary_key = True),
+ Column('invoice_id', Integer, ForeignKey('invoices.invoice_id')),
+ Column('code', String(20)),
+ Column('qty', Integer),
+ )
+
+ class Company(object):
+ def __init__(self):
+ self.company_id = None
+ def __repr__(self):
+ return "Company:" + repr(getattr(self, 'company_id', None)) + " " + repr(getattr(self, 'company_name', None)) + " " + str([repr(addr) for addr in self.addresses])
+
+ class Address(object):
+ def __repr__(self):
+ return "Address: " + repr(getattr(self, 'address_id', None)) + " " + repr(getattr(self, 'company_id', None)) + " " + repr(self.address) + str([repr(ph) for ph in getattr(self, 'phones', [])])
+
+ class Phone(object):
+ def __repr__(self):
+ return "Phone: " + repr(getattr(self, 'phone_id', None)) + " " + repr(getattr(self, 'address_id', None)) + " " + repr(self.type) + " " + repr(self.number)
+
+ class Invoice(object):
+ def __init__(self):
+ self.invoice_id = None
+ def __repr__(self):
+ return "Invoice:" + repr(getattr(self, 'invoice_id', None)) + " " + repr(getattr(self, 'date', None)) + " " + repr(self.company) + " " + str([repr(item) for item in self.items])
+
+ class Item(object):
+ def __repr__(self):
+ return "Item: " + repr(getattr(self, 'item_id', None)) + " " + repr(getattr(self, 'invoice_id', None)) + " " + repr(self.code) + " " + repr(self.qty)
+
+ def testone(self):
+ """tests eager load of a many-to-one attached to a one-to-many. this testcase illustrated
+ the bug, which is that when the single Company is loaded, no further processing of the rows
+ occurred in order to load the Company's second Address object."""
+
+ mapper(Address, addresses_table, properties={
+ }, extension=ctx.mapper_extension)
+ mapper(Company, companies_table, properties={
+ 'addresses' : relation(Address, lazy=False),
+ }, extension=ctx.mapper_extension)
+ mapper(Invoice, invoice_table, properties={
+ 'company': relation(Company, lazy=False, )
+ }, extension=ctx.mapper_extension)
+
+ c1 = Company()
+ c1.company_name = 'company 1'
+ a1 = Address()
+ a1.address = 'a1 address'
+ c1.addresses.append(a1)
+ a2 = Address()
+ a2.address = 'a2 address'
+ c1.addresses.append(a2)
+ i1 = Invoice()
+ i1.date = datetime.datetime.now()
+ i1.company = c1
+
+ ctx.current.flush()
+
+ company_id = c1.company_id
+ invoice_id = i1.invoice_id
+
+ ctx.current.clear()
+
+ c = ctx.current.query(Company).get(company_id)
+
+ ctx.current.clear()
+
+ i = ctx.current.query(Invoice).get(invoice_id)
+
+ print repr(c)
+ print repr(i.company)
+ self.assert_(repr(c) == repr(i.company))
+
+ def testtwo(self):
+ """this is the original testcase that includes various complicating factors"""
+
+ mapper(Phone, phones_table, extension=ctx.mapper_extension)
+
+ mapper(Address, addresses_table, properties={
+ 'phones': relation(Phone, lazy=False, backref='address')
+ }, extension=ctx.mapper_extension)
+
+ mapper(Company, companies_table, properties={
+ 'addresses' : relation(Address, lazy=False, backref='company'),
+ }, extension=ctx.mapper_extension)
+
+ mapper(Item, items_table, extension=ctx.mapper_extension)
+
+ mapper(Invoice, invoice_table, properties={
+ 'items': relation(Item, lazy=False, backref='invoice'),
+ 'company': relation(Company, lazy=False, backref='invoices')
+ }, extension=ctx.mapper_extension)
+
+ ctx.current.clear()
+ c1 = Company()
+ c1.company_name = 'company 1'
+
+ a1 = Address()
+ a1.address = 'a1 address'
+
+ p1 = Phone()
+ p1.type = 'home'
+ p1.number = '1111'
+
+ a1.phones.append(p1)
+
+ p2 = Phone()
+ p2.type = 'work'
+ p2.number = '22222'
+ a1.phones.append(p2)
+
+ c1.addresses.append(a1)
+
+ a2 = Address()
+ a2.address = 'a2 address'
+
+ p3 = Phone()
+ p3.type = 'home'
+ p3.number = '3333'
+ a2.phones.append(p3)
+
+ p4 = Phone()
+ p4.type = 'work'
+ p4.number = '44444'
+ a2.phones.append(p4)
+
+ c1.addresses.append(a2)
+
+ ctx.current.flush()
+
+ company_id = c1.company_id
+
+ ctx.current.clear()
+
+ a = ctx.current.query(Company).get(company_id)
+ print repr(a)
+
+ # set up an invoice
+ i1 = Invoice()
+ i1.date = datetime.datetime.now()
+ i1.company = c1
+
+ item1 = Item()
+ item1.code = 'aaaa'
+ item1.qty = 1
+ item1.invoice = i1
+
+ item2 = Item()
+ item2.code = 'bbbb'
+ item2.qty = 2
+ item2.invoice = i1
+
+ item3 = Item()
+ item3.code = 'cccc'
+ item3.qty = 3
+ item3.invoice = i1
+
+ ctx.current.flush()
+
+ invoice_id = i1.invoice_id
+
+ ctx.current.clear()
+
+ c = ctx.current.query(Company).get(company_id)
+ print repr(c)
+
+ ctx.current.clear()
+
+ i = ctx.current.query(Invoice).get(invoice_id)
+
+ assert repr(i.company) == repr(c), repr(i.company) + " does not match " + repr(c)
+
+class EagerTest8(ORMTest):
def define_tables(self, metadata):
global project_t, task_t, task_status_t, task_type_t, message_t, message_type_t
-
+
project_t = Table('prj', metadata,
Column('id', Integer, primary_key=True),
Column('created', DateTime , ),
@@ -460,12 +718,12 @@ class EagerTest6(testbase.ORMTest):
testbase.db.execute("INSERT INTO task_status (id) values(1);")
testbase.db.execute("INSERT INTO task_type(id) values(1);")
testbase.db.execute("INSERT INTO task (title, task_type_id, status_id, prj_id) values('task 1',1,1,1);")
-
+
def test_nested_joins(self):
# this is testing some subtle column resolution stuff,
# concerning corresponding_column() being extremely accurate
# as well as how mapper sets up its column properties
-
+
class Task(object):pass
class Task_Type(object):pass
class Message(object):pass
@@ -510,6 +768,7 @@ class EagerTest6(testbase.ORMTest):
for t in session.query(cls.mapper).limit(10).offset(0).list():
print t.id, t.title, t.props_cnt
-
+
+
if __name__ == "__main__":
testbase.main()
diff --git a/test/orm/attributes.py b/test/orm/attributes.py
index 7e0a22aff..9b5f738bf 100644
--- a/test/orm/attributes.py
+++ b/test/orm/attributes.py
@@ -1,10 +1,9 @@
-from testbase import PersistTest
-import sqlalchemy.util as util
+import testbase
+import pickle
import sqlalchemy.orm.attributes as attributes
+from sqlalchemy.orm.collections import collection
from sqlalchemy import exceptions
-import unittest, sys, os
-import pickle
-import testbase
+from testlib import *
class MyTest(object):pass
class MyTest2(object):pass
@@ -50,14 +49,53 @@ class AttributesTest(PersistTest):
# shouldnt be pickling callables at the class level
def somecallable(*args):
return None
- manager.register_attribute(MyTest, 'mt2', uselist = True, trackparent=True, callable_=somecallable)
- x = MyTest()
- x.mt2.append(MyTest2())
-
- x.user_id=7
- s = pickle.dumps(x)
- x2 = pickle.loads(s)
- assert s == pickle.dumps(x2)
+ attr_name = 'mt2'
+ manager.register_attribute(MyTest, attr_name, uselist = True, trackparent=True, callable_=somecallable)
+
+ o = MyTest()
+ o.mt2.append(MyTest2())
+ o.user_id=7
+ o.mt2[0].a = 'abcde'
+ pk_o = pickle.dumps(o)
+
+ o2 = pickle.loads(pk_o)
+
+ # so... pickle is creating a new 'mt2' string after a roundtrip here,
+ # so we'll brute-force set it to be id-equal to the original string
+ o_mt2_str = [ k for k in o.__dict__ if k == 'mt2'][0]
+ o2_mt2_str = [ k for k in o2.__dict__ if k == 'mt2'][0]
+ self.assert_(o_mt2_str == o2_mt2_str)
+ self.assert_(o_mt2_str is not o2_mt2_str)
+ # change the id of o2.__dict__['mt2']
+ former = o2.__dict__['mt2']
+ del o2.__dict__['mt2']
+ o2.__dict__[o_mt2_str] = former
+
+ pk_o2 = pickle.dumps(o2)
+
+ self.assert_(pk_o == pk_o2)
+
+ # the above is kind of distrurbing, so let's do it again a little
+ # differently. the string-id in serialization thing is just an
+ # artifact of pickling that comes up in the first round-trip.
+ # a -> b differs in pickle memoization of 'mt2', but b -> c will
+ # serialize identically.
+
+ o3 = pickle.loads(pk_o2)
+ pk_o3 = pickle.dumps(o3)
+ o4 = pickle.loads(pk_o3)
+ pk_o4 = pickle.dumps(o4)
+
+ self.assert_(pk_o3 == pk_o4)
+
+ # and lastly make sure we still have our data after all that.
+ # identical serialzation is great, *if* it's complete :)
+ self.assert_(o4.user_id == 7)
+ self.assert_(o4.user_name is None)
+ self.assert_(o4.email_address is None)
+ self.assert_(len(o4.mt2) == 1)
+ self.assert_(o4.mt2[0].a == 'abcde')
+ self.assert_(o4.mt2[0].b is None)
def testlist(self):
class User(object):pass
@@ -110,13 +148,12 @@ class AttributesTest(PersistTest):
s = Student()
c = Course()
s.courses.append(c)
- print c.students
- print [s]
self.assert_(c.students == [s])
s.courses.remove(c)
self.assert_(c.students == [])
(s1, s2, s3) = (Student(), Student(), Student())
+
c.students = [s1, s2, s3]
self.assert_(s2.courses == [c])
self.assert_(s1.courses == [c])
@@ -126,9 +163,7 @@ class AttributesTest(PersistTest):
print c
print c.students
s1.courses.remove(c)
- self.assert_(c.students == [s2,s3])
-
-
+ self.assert_(c.students == [s2,s3])
class Post(object):pass
class Blog(object):pass
@@ -334,44 +369,47 @@ class AttributesTest(PersistTest):
manager = attributes.AttributeManager()
class Foo(object):pass
manager.register_attribute(Foo, "collection", uselist=True, typecallable=set)
- assert isinstance(Foo().collection.data, set)
+ assert isinstance(Foo().collection, set)
- manager.register_attribute(Foo, "collection", uselist=True, typecallable=dict)
try:
- Foo().collection
+ manager.register_attribute(Foo, "collection", uselist=True, typecallable=dict)
assert False
except exceptions.ArgumentError, e:
- assert str(e) == "Dictionary collection class 'dict' must implement an append() method"
-
+ assert str(e) == "Type InstrumentedDict must elect an appender method to be a collection class"
+
class MyDict(dict):
+ @collection.appender
def append(self, item):
self[item.foo] = item
+ @collection.remover
+ def remove(self, item):
+ del self[item.foo]
manager.register_attribute(Foo, "collection", uselist=True, typecallable=MyDict)
- assert isinstance(Foo().collection.data, MyDict)
+ assert isinstance(Foo().collection, MyDict)
class MyColl(object):pass
- manager.register_attribute(Foo, "collection", uselist=True, typecallable=MyColl)
try:
- Foo().collection
+ manager.register_attribute(Foo, "collection", uselist=True, typecallable=MyColl)
assert False
except exceptions.ArgumentError, e:
- assert str(e) == "Collection class 'MyColl' is not of type 'list', 'set', or 'dict' and has no append() or add() method"
+ assert str(e) == "Type MyColl must elect an appender method to be a collection class"
class MyColl(object):
+ @collection.iterator
def __iter__(self):
return iter([])
+ @collection.appender
def append(self, item):
pass
+ @collection.remover
+ def remove(self, item):
+ pass
manager.register_attribute(Foo, "collection", uselist=True, typecallable=MyColl)
try:
Foo().collection
- assert False
+ assert True
except exceptions.ArgumentError, e:
- assert str(e) == "Collection class 'MyColl' is not of type 'list', 'set', or 'dict' and has no clear() method"
-
- def foo(self):pass
- MyColl.clear = foo
- assert isinstance(Foo().collection.data, MyColl)
+ assert False
if __name__ == "__main__":
testbase.main()
diff --git a/test/orm/cascade.py b/test/orm/cascade.py
index 16c4db40f..b832c427e 100644
--- a/test/orm/cascade.py
+++ b/test/orm/cascade.py
@@ -1,10 +1,12 @@
-import testbase, tables
-import unittest, sys, datetime
+import testbase
-from sqlalchemy.ext.sessioncontext import SessionContext
from sqlalchemy import *
+from sqlalchemy.orm import *
+from sqlalchemy.ext.sessioncontext import SessionContext
+from testlib import *
+import testlib.tables as tables
-class O2MCascadeTest(testbase.AssertMixin):
+class O2MCascadeTest(AssertMixin):
def tearDown(self):
tables.delete()
@@ -112,7 +114,7 @@ class O2MCascadeTest(testbase.AssertMixin):
sess = create_session()
l = sess.query(tables.User).select()
for u in l:
- self.echo( repr(u.orders))
+ print repr(u.orders)
self.assert_result(l, data[0], *data[1:])
ids = (l[0].user_id, l[2].user_id)
@@ -172,7 +174,7 @@ class O2MCascadeTest(testbase.AssertMixin):
self.assert_(tables.orderitems.count(tables.orders.c.user_id.in_(*ids) &(tables.orderitems.c.order_id==tables.orders.c.order_id)).scalar() == 0)
-class M2OCascadeTest(testbase.AssertMixin):
+class M2OCascadeTest(AssertMixin):
def tearDown(self):
ctx.current.clear()
for t in metadata.table_iterator(reverse=True):
@@ -260,7 +262,7 @@ class M2OCascadeTest(testbase.AssertMixin):
-class M2MCascadeTest(testbase.AssertMixin):
+class M2MCascadeTest(AssertMixin):
def setUpAll(self):
global metadata, a, b, atob
metadata = MetaData(testbase.db)
@@ -335,7 +337,7 @@ class M2MCascadeTest(testbase.AssertMixin):
assert b.count().scalar() == 0
assert a.count().scalar() == 0
-class UnsavedOrphansTest(testbase.ORMTest):
+class UnsavedOrphansTest(ORMTest):
"""tests regarding pending entities that are orphans"""
def define_tables(self, metadata):
@@ -395,7 +397,7 @@ class UnsavedOrphansTest(testbase.ORMTest):
assert a.address_id is None, "Error: address should not be persistent"
-class UnsavedOrphansTest2(testbase.ORMTest):
+class UnsavedOrphansTest2(ORMTest):
"""same test as UnsavedOrphans only three levels deep"""
def define_tables(self, meta):
@@ -455,7 +457,7 @@ class UnsavedOrphansTest2(testbase.ORMTest):
assert item.id is None
assert attr.id is None
-class DoubleParentOrphanTest(testbase.AssertMixin):
+class DoubleParentOrphanTest(AssertMixin):
"""test orphan detection for an entity with two parent relations"""
def setUpAll(self):
@@ -521,7 +523,7 @@ class DoubleParentOrphanTest(testbase.AssertMixin):
assert False
except exceptions.FlushError, e:
assert True
-
-
+
+
if __name__ == "__main__":
testbase.main()
diff --git a/test/orm/collection.py b/test/orm/collection.py
new file mode 100644
index 000000000..1f4f64928
--- /dev/null
+++ b/test/orm/collection.py
@@ -0,0 +1,1140 @@
+import testbase
+from sqlalchemy import *
+import sqlalchemy.exceptions as exceptions
+from sqlalchemy.orm import create_session, mapper, relation, \
+ interfaces, attributes
+import sqlalchemy.orm.collections as collections
+from sqlalchemy.orm.collections import collection
+from sqlalchemy import util
+from operator import and_
+from testlib import *
+
+class Canary(interfaces.AttributeExtension):
+ def __init__(self):
+ self.data = set()
+ self.added = set()
+ self.removed = set()
+ def append(self, obj, value, initiator):
+ assert value not in self.added
+ self.data.add(value)
+ self.added.add(value)
+ def remove(self, obj, value, initiator):
+ assert value not in self.removed
+ self.data.remove(value)
+ self.removed.add(value)
+ def set(self, obj, value, oldvalue, initiator):
+ if oldvalue is not None:
+ self.remove(obj, oldvalue, None)
+ self.append(obj, value, None)
+
+class Entity(object):
+ def __init__(self, a=None, b=None, c=None):
+ self.a = a
+ self.b = b
+ self.c = c
+ def __repr__(self):
+ return str((id(self), self.a, self.b, self.c))
+
+manager = attributes.AttributeManager()
+
+_id = 1
+def entity_maker():
+ global _id
+ _id += 1
+ return Entity(_id)
+def dictable_entity(a=None, b=None, c=None):
+ global _id
+ _id += 1
+ return Entity(a or str(_id), b or 'value %s' % _id, c)
+
+
+class CollectionsTest(PersistTest):
+ def _test_adapter(self, typecallable, creator=entity_maker,
+ to_set=None):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ manager.register_attribute(Foo, 'attr', True, extension=canary,
+ typecallable=typecallable)
+
+ obj = Foo()
+ adapter = collections.collection_adapter(obj.attr)
+ direct = obj.attr
+ if to_set is None:
+ to_set = lambda col: set(col)
+
+ def assert_eq():
+ self.assert_(to_set(direct) == canary.data)
+ self.assert_(set(adapter) == canary.data)
+ assert_ne = lambda: self.assert_(to_set(direct) != canary.data)
+
+ e1, e2 = creator(), creator()
+
+ adapter.append_with_event(e1)
+ assert_eq()
+
+ adapter.append_without_event(e2)
+ assert_ne()
+ canary.data.add(e2)
+ assert_eq()
+
+ adapter.remove_without_event(e2)
+ assert_ne()
+ canary.data.remove(e2)
+ assert_eq()
+
+ adapter.remove_with_event(e1)
+ assert_eq()
+
+ def _test_list(self, typecallable, creator=entity_maker):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ manager.register_attribute(Foo, 'attr', True, extension=canary,
+ typecallable=typecallable)
+
+ obj = Foo()
+ adapter = collections.collection_adapter(obj.attr)
+ direct = obj.attr
+ control = list()
+
+ def assert_eq():
+ self.assert_(set(direct) == canary.data)
+ self.assert_(set(adapter) == canary.data)
+ self.assert_(direct == control)
+
+ # assume append() is available for list tests
+ e = creator()
+ direct.append(e)
+ control.append(e)
+ assert_eq()
+
+ if hasattr(direct, 'pop'):
+ direct.pop()
+ control.pop()
+ assert_eq()
+
+ if hasattr(direct, '__setitem__'):
+ e = creator()
+ direct.append(e)
+ control.append(e)
+
+ e = creator()
+ direct[0] = e
+ control[0] = e
+ assert_eq()
+
+ if reduce(and_, [hasattr(direct, a) for a in
+ ('__delitem', 'insert', '__len__')], True):
+ values = [creator(), creator(), creator(), creator()]
+ direct[slice(0,1)] = values
+ control[slice(0,1)] = values
+ assert_eq()
+
+ values = [creator(), creator()]
+ direct[slice(0,-1,2)] = values
+ control[slice(0,-1,2)] = values
+ assert_eq()
+
+ values = [creator()]
+ direct[slice(0,-1)] = values
+ control[slice(0,-1)] = values
+ assert_eq()
+
+ if hasattr(direct, '__delitem__'):
+ e = creator()
+ direct.append(e)
+ control.append(e)
+ del direct[-1]
+ del control[-1]
+ assert_eq()
+
+ if hasattr(direct, '__getslice__'):
+ for e in [creator(), creator(), creator(), creator()]:
+ direct.append(e)
+ control.append(e)
+
+ del direct[:-3]
+ del control[:-3]
+ assert_eq()
+
+ del direct[0:1]
+ del control[0:1]
+ assert_eq()
+
+ del direct[::2]
+ del control[::2]
+ assert_eq()
+
+ if hasattr(direct, 'remove'):
+ e = creator()
+ direct.append(e)
+ control.append(e)
+
+ direct.remove(e)
+ control.remove(e)
+ assert_eq()
+
+ if hasattr(direct, '__setslice__'):
+ values = [creator(), creator()]
+ direct[0:1] = values
+ control[0:1] = values
+ assert_eq()
+
+ values = [creator()]
+ direct[0:] = values
+ control[0:] = values
+ assert_eq()
+
+ if hasattr(direct, '__delslice__'):
+ for i in range(1, 4):
+ e = creator()
+ direct.append(e)
+ control.append(e)
+
+ del direct[-1:]
+ del control[-1:]
+ assert_eq()
+
+ del direct[1:2]
+ del control[1:2]
+ assert_eq()
+
+ del direct[:]
+ del control[:]
+ assert_eq()
+
+ if hasattr(direct, 'extend'):
+ values = [creator(), creator(), creator()]
+
+ direct.extend(values)
+ control.extend(values)
+ assert_eq()
+
+ def _test_list_bulk(self, typecallable, creator=entity_maker):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ manager.register_attribute(Foo, 'attr', True, extension=canary,
+ typecallable=typecallable)
+
+ obj = Foo()
+ direct = obj.attr
+
+ e1 = creator()
+ obj.attr.append(e1)
+
+ like_me = typecallable()
+ e2 = creator()
+ like_me.append(e2)
+
+ self.assert_(obj.attr is direct)
+ obj.attr = like_me
+ self.assert_(obj.attr is not direct)
+ self.assert_(obj.attr is not like_me)
+ self.assert_(set(obj.attr) == set([e2]))
+ self.assert_(e1 in canary.removed)
+ self.assert_(e2 in canary.added)
+
+ e3 = creator()
+ real_list = [e3]
+ obj.attr = real_list
+ self.assert_(obj.attr is not real_list)
+ self.assert_(set(obj.attr) == set([e3]))
+ self.assert_(e2 in canary.removed)
+ self.assert_(e3 in canary.added)
+
+ e4 = creator()
+ try:
+ obj.attr = set([e4])
+ self.assert_(False)
+ except exceptions.ArgumentError:
+ self.assert_(e4 not in canary.data)
+ self.assert_(e3 in canary.data)
+
+ def test_list(self):
+ self._test_adapter(list)
+ self._test_list(list)
+ self._test_list_bulk(list)
+
+ def test_list_subclass(self):
+ class MyList(list):
+ pass
+ self._test_adapter(MyList)
+ self._test_list(MyList)
+ self._test_list_bulk(MyList)
+ self.assert_(getattr(MyList, '_sa_instrumented') == id(MyList))
+
+ def test_list_duck(self):
+ class ListLike(object):
+ def __init__(self):
+ self.data = list()
+ def append(self, item):
+ self.data.append(item)
+ def remove(self, item):
+ self.data.remove(item)
+ def insert(self, index, item):
+ self.data.insert(index, item)
+ def pop(self, index=-1):
+ return self.data.pop(index)
+ def extend(self):
+ assert False
+ def __iter__(self):
+ return iter(self.data)
+ def __eq__(self, other):
+ return self.data == other
+ def __repr__(self):
+ return 'ListLike(%s)' % repr(self.data)
+
+ self._test_adapter(ListLike)
+ self._test_list(ListLike)
+ self._test_list_bulk(ListLike)
+ self.assert_(getattr(ListLike, '_sa_instrumented') == id(ListLike))
+
+ def test_list_emulates(self):
+ class ListIsh(object):
+ __emulates__ = list
+ def __init__(self):
+ self.data = list()
+ def append(self, item):
+ self.data.append(item)
+ def remove(self, item):
+ self.data.remove(item)
+ def insert(self, index, item):
+ self.data.insert(index, item)
+ def pop(self, index=-1):
+ return self.data.pop(index)
+ def extend(self):
+ assert False
+ def __iter__(self):
+ return iter(self.data)
+ def __eq__(self, other):
+ return self.data == other
+ def __repr__(self):
+ return 'ListIsh(%s)' % repr(self.data)
+
+ self._test_adapter(ListIsh)
+ self._test_list(ListIsh)
+ self._test_list_bulk(ListIsh)
+ self.assert_(getattr(ListIsh, '_sa_instrumented') == id(ListIsh))
+
+ def _test_set(self, typecallable, creator=entity_maker):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ manager.register_attribute(Foo, 'attr', True, extension=canary,
+ typecallable=typecallable)
+
+ obj = Foo()
+ adapter = collections.collection_adapter(obj.attr)
+ direct = obj.attr
+ control = set()
+
+ def assert_eq():
+ self.assert_(set(direct) == canary.data)
+ self.assert_(set(adapter) == canary.data)
+ self.assert_(direct == control)
+
+ def addall(*values):
+ for item in values:
+ direct.add(item)
+ control.add(item)
+ assert_eq()
+ def zap():
+ for item in list(direct):
+ direct.remove(item)
+ control.clear()
+
+ # assume add() is available for list tests
+ addall(creator())
+
+ if hasattr(direct, 'pop'):
+ direct.pop()
+ control.pop()
+ assert_eq()
+
+ if hasattr(direct, 'remove'):
+ e = creator()
+ addall(e)
+
+ direct.remove(e)
+ control.remove(e)
+ assert_eq()
+
+ e = creator()
+ try:
+ direct.remove(e)
+ except KeyError:
+ assert_eq()
+ self.assert_(e not in canary.removed)
+ else:
+ self.assert_(False)
+
+ if hasattr(direct, 'discard'):
+ e = creator()
+ addall(e)
+
+ direct.discard(e)
+ control.discard(e)
+ assert_eq()
+
+ e = creator()
+ direct.discard(e)
+ self.assert_(e not in canary.removed)
+ assert_eq()
+
+ if hasattr(direct, 'update'):
+ e = creator()
+ addall(e)
+
+ values = set([e, creator(), creator()])
+
+ direct.update(values)
+ control.update(values)
+ assert_eq()
+
+ if hasattr(direct, 'clear'):
+ addall(creator(), creator())
+ direct.clear()
+ control.clear()
+ assert_eq()
+
+ if hasattr(direct, 'difference_update'):
+ zap()
+ addall(creator(), creator())
+ values = set([creator()])
+
+ direct.difference_update(values)
+ control.difference_update(values)
+ assert_eq()
+ values.update(set([e, creator()]))
+ direct.difference_update(values)
+ control.difference_update(values)
+ assert_eq()
+
+ if hasattr(direct, 'intersection_update'):
+ zap()
+ e = creator()
+ addall(e, creator(), creator())
+ values = set(control)
+
+ direct.intersection_update(values)
+ control.intersection_update(values)
+ assert_eq()
+
+ values.update(set([e, creator()]))
+ direct.intersection_update(values)
+ control.intersection_update(values)
+ assert_eq()
+
+ if hasattr(direct, 'symmetric_difference_update'):
+ zap()
+ e = creator()
+ addall(e, creator(), creator())
+
+ values = set([e, creator()])
+ direct.symmetric_difference_update(values)
+ control.symmetric_difference_update(values)
+ assert_eq()
+
+ e = creator()
+ addall(e)
+ values = set([e])
+ direct.symmetric_difference_update(values)
+ control.symmetric_difference_update(values)
+ assert_eq()
+
+ values = set()
+ direct.symmetric_difference_update(values)
+ control.symmetric_difference_update(values)
+ assert_eq()
+
+ def _test_set_bulk(self, typecallable, creator=entity_maker):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ manager.register_attribute(Foo, 'attr', True, extension=canary,
+ typecallable=typecallable)
+
+ obj = Foo()
+ direct = obj.attr
+
+ e1 = creator()
+ obj.attr.add(e1)
+
+ like_me = typecallable()
+ e2 = creator()
+ like_me.add(e2)
+
+ self.assert_(obj.attr is direct)
+ obj.attr = like_me
+ self.assert_(obj.attr is not direct)
+ self.assert_(obj.attr is not like_me)
+ self.assert_(obj.attr == set([e2]))
+ self.assert_(e1 in canary.removed)
+ self.assert_(e2 in canary.added)
+
+ e3 = creator()
+ real_set = set([e3])
+ obj.attr = real_set
+ self.assert_(obj.attr is not real_set)
+ self.assert_(obj.attr == set([e3]))
+ self.assert_(e2 in canary.removed)
+ self.assert_(e3 in canary.added)
+
+ e4 = creator()
+ try:
+ obj.attr = [e4]
+ self.assert_(False)
+ except exceptions.ArgumentError:
+ self.assert_(e4 not in canary.data)
+ self.assert_(e3 in canary.data)
+
+ def test_set(self):
+ self._test_adapter(set)
+ self._test_set(set)
+ self._test_set_bulk(set)
+
+ def test_set_subclass(self):
+ class MySet(set):
+ pass
+ self._test_adapter(MySet)
+ self._test_set(MySet)
+ self._test_set_bulk(MySet)
+ self.assert_(getattr(MySet, '_sa_instrumented') == id(MySet))
+
+ def test_set_duck(self):
+ class SetLike(object):
+ def __init__(self):
+ self.data = set()
+ def add(self, item):
+ self.data.add(item)
+ def remove(self, item):
+ self.data.remove(item)
+ def discard(self, item):
+ self.data.discard(item)
+ def pop(self):
+ return self.data.pop()
+ def update(self, other):
+ self.data.update(other)
+ def __iter__(self):
+ return iter(self.data)
+ def __eq__(self, other):
+ return self.data == other
+
+ self._test_adapter(SetLike)
+ self._test_set(SetLike)
+ self._test_set_bulk(SetLike)
+ self.assert_(getattr(SetLike, '_sa_instrumented') == id(SetLike))
+
+ def test_set_emulates(self):
+ class SetIsh(object):
+ __emulates__ = set
+ def __init__(self):
+ self.data = set()
+ def add(self, item):
+ self.data.add(item)
+ def remove(self, item):
+ self.data.remove(item)
+ def discard(self, item):
+ self.data.discard(item)
+ def pop(self):
+ return self.data.pop()
+ def update(self, other):
+ self.data.update(other)
+ def __iter__(self):
+ return iter(self.data)
+ def __eq__(self, other):
+ return self.data == other
+
+ self._test_adapter(SetIsh)
+ self._test_set(SetIsh)
+ self._test_set_bulk(SetIsh)
+ self.assert_(getattr(SetIsh, '_sa_instrumented') == id(SetIsh))
+
+ def _test_dict(self, typecallable, creator=dictable_entity):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ manager.register_attribute(Foo, 'attr', True, extension=canary,
+ typecallable=typecallable)
+
+ obj = Foo()
+ adapter = collections.collection_adapter(obj.attr)
+ direct = obj.attr
+ control = dict()
+
+ def assert_eq():
+ self.assert_(set(direct.values()) == canary.data)
+ self.assert_(set(adapter) == canary.data)
+ self.assert_(direct == control)
+
+ def addall(*values):
+ for item in values:
+ direct.set(item)
+ control[item.a] = item
+ assert_eq()
+ def zap():
+ for item in list(adapter):
+ direct.remove(item)
+ control.clear()
+
+ # assume an 'set' method is available for tests
+ addall(creator())
+
+ if hasattr(direct, '__setitem__'):
+ e = creator()
+ direct[e.a] = e
+ control[e.a] = e
+ assert_eq()
+
+ e = creator(e.a, e.b)
+ direct[e.a] = e
+ control[e.a] = e
+ assert_eq()
+
+ if hasattr(direct, '__delitem__'):
+ e = creator()
+ addall(e)
+
+ del direct[e.a]
+ del control[e.a]
+ assert_eq()
+
+ e = creator()
+ try:
+ del direct[e.a]
+ except KeyError:
+ self.assert_(e not in canary.removed)
+
+ if hasattr(direct, 'clear'):
+ addall(creator(), creator(), creator())
+
+ direct.clear()
+ control.clear()
+ assert_eq()
+
+ direct.clear()
+ control.clear()
+ assert_eq()
+
+ if hasattr(direct, 'pop'):
+ e = creator()
+ addall(e)
+
+ direct.pop(e.a)
+ control.pop(e.a)
+ assert_eq()
+
+ e = creator()
+ try:
+ direct.pop(e.a)
+ except KeyError:
+ self.assert_(e not in canary.removed)
+
+ if hasattr(direct, 'popitem'):
+ zap()
+ e = creator()
+ addall(e)
+
+ direct.popitem()
+ control.popitem()
+ assert_eq()
+
+ if hasattr(direct, 'setdefault'):
+ e = creator()
+
+ val_a = direct.setdefault(e.a, e)
+ val_b = control.setdefault(e.a, e)
+ assert_eq()
+ self.assert_(val_a is val_b)
+
+ val_a = direct.setdefault(e.a, e)
+ val_b = control.setdefault(e.a, e)
+ assert_eq()
+ self.assert_(val_a is val_b)
+
+ if hasattr(direct, 'update'):
+ e = creator()
+ d = dict([(ee.a, ee) for ee in [e, creator(), creator()]])
+ addall(e, creator())
+
+ direct.update(d)
+ control.update(d)
+ assert_eq()
+
+ kw = dict([(ee.a, ee) for ee in [e, creator()]])
+ direct.update(**kw)
+ control.update(**kw)
+ assert_eq()
+
+ def _test_dict_bulk(self, typecallable, creator=dictable_entity):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ manager.register_attribute(Foo, 'attr', True, extension=canary,
+ typecallable=typecallable)
+
+ obj = Foo()
+ direct = obj.attr
+
+ e1 = creator()
+ collections.collection_adapter(direct).append_with_event(e1)
+
+ like_me = typecallable()
+ e2 = creator()
+ like_me.set(e2)
+
+ self.assert_(obj.attr is direct)
+ obj.attr = like_me
+ self.assert_(obj.attr is not direct)
+ self.assert_(obj.attr is not like_me)
+ self.assert_(set(collections.collection_adapter(obj.attr)) == set([e2]))
+ self.assert_(e1 in canary.removed)
+ self.assert_(e2 in canary.added)
+
+ e3 = creator()
+ real_dict = dict(keyignored1=e3)
+ obj.attr = real_dict
+ self.assert_(obj.attr is not real_dict)
+ self.assert_('keyignored1' not in obj.attr)
+ self.assert_(set(collections.collection_adapter(obj.attr)) == set([e3]))
+ self.assert_(e2 in canary.removed)
+ self.assert_(e3 in canary.added)
+
+ e4 = creator()
+ try:
+ obj.attr = [e4]
+ self.assert_(False)
+ except exceptions.ArgumentError:
+ self.assert_(e4 not in canary.data)
+ self.assert_(e3 in canary.data)
+
+ def test_dict(self):
+ try:
+ self._test_adapter(dict, dictable_entity,
+ to_set=lambda c: set(c.values()))
+ self.assert_(False)
+ except exceptions.ArgumentError, e:
+ self.assert_(e.args[0] == 'Type InstrumentedDict must elect an appender method to be a collection class')
+
+ try:
+ self._test_dict(dict)
+ self.assert_(False)
+ except exceptions.ArgumentError, e:
+ self.assert_(e.args[0] == 'Type InstrumentedDict must elect an appender method to be a collection class')
+
+ def test_dict_subclass(self):
+ class MyDict(dict):
+ @collection.appender
+ @collection.internally_instrumented
+ def set(self, item, _sa_initiator=None):
+ self.__setitem__(item.a, item, _sa_initiator=_sa_initiator)
+ @collection.remover
+ @collection.internally_instrumented
+ def _remove(self, item, _sa_initiator=None):
+ self.__delitem__(item.a, _sa_initiator=_sa_initiator)
+
+ self._test_adapter(MyDict, dictable_entity,
+ to_set=lambda c: set(c.values()))
+ self._test_dict(MyDict)
+ self._test_dict_bulk(MyDict)
+ self.assert_(getattr(MyDict, '_sa_instrumented') == id(MyDict))
+
+ def test_dict_subclass2(self):
+ class MyEasyDict(collections.MappedCollection):
+ def __init__(self):
+ super(MyEasyDict, self).__init__(lambda e: e.a)
+
+ self._test_adapter(MyEasyDict, dictable_entity,
+ to_set=lambda c: set(c.values()))
+ self._test_dict(MyEasyDict)
+ self._test_dict_bulk(MyEasyDict)
+ self.assert_(getattr(MyEasyDict, '_sa_instrumented') == id(MyEasyDict))
+
+ def test_dict_subclass3(self):
+ class MyOrdered(util.OrderedDict, collections.MappedCollection):
+ def __init__(self):
+ collections.MappedCollection.__init__(self, lambda e: e.a)
+ util.OrderedDict.__init__(self)
+
+ self._test_adapter(MyOrdered, dictable_entity,
+ to_set=lambda c: set(c.values()))
+ self._test_dict(MyOrdered)
+ self._test_dict_bulk(MyOrdered)
+ self.assert_(getattr(MyOrdered, '_sa_instrumented') == id(MyOrdered))
+
+ def test_dict_duck(self):
+ class DictLike(object):
+ def __init__(self):
+ self.data = dict()
+
+ @collection.appender
+ @collection.replaces(1)
+ def set(self, item):
+ current = self.data.get(item.a, None)
+ self.data[item.a] = item
+ return current
+ @collection.remover
+ def _remove(self, item):
+ del self.data[item.a]
+ def __setitem__(self, key, value):
+ self.data[key] = value
+ def __getitem__(self, key):
+ return self.data[key]
+ def __delitem__(self, key):
+ del self.data[key]
+ def values(self):
+ return self.data.values()
+ def __contains__(self, key):
+ return key in self.data
+ @collection.iterator
+ def itervalues(self):
+ return self.data.itervalues()
+ def __eq__(self, other):
+ return self.data == other
+ def __repr__(self):
+ return 'DictLike(%s)' % repr(self.data)
+
+ self._test_adapter(DictLike, dictable_entity,
+ to_set=lambda c: set(c.itervalues()))
+ self._test_dict(DictLike)
+ self._test_dict_bulk(DictLike)
+ self.assert_(getattr(DictLike, '_sa_instrumented') == id(DictLike))
+
+ def test_dict_emulates(self):
+ class DictIsh(object):
+ __emulates__ = dict
+ def __init__(self):
+ self.data = dict()
+
+ @collection.appender
+ @collection.replaces(1)
+ def set(self, item):
+ current = self.data.get(item.a, None)
+ self.data[item.a] = item
+ return current
+ @collection.remover
+ def _remove(self, item):
+ del self.data[item.a]
+ def __setitem__(self, key, value):
+ self.data[key] = value
+ def __getitem__(self, key):
+ return self.data[key]
+ def __delitem__(self, key):
+ del self.data[key]
+ def values(self):
+ return self.data.values()
+ def __contains__(self, key):
+ return key in self.data
+ @collection.iterator
+ def itervalues(self):
+ return self.data.itervalues()
+ def __eq__(self, other):
+ return self.data == other
+ def __repr__(self):
+ return 'DictIsh(%s)' % repr(self.data)
+
+ self._test_adapter(DictIsh, dictable_entity,
+ to_set=lambda c: set(c.itervalues()))
+ self._test_dict(DictIsh)
+ self._test_dict_bulk(DictIsh)
+ self.assert_(getattr(DictIsh, '_sa_instrumented') == id(DictIsh))
+
+ def _test_object(self, typecallable, creator=entity_maker):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ manager.register_attribute(Foo, 'attr', True, extension=canary,
+ typecallable=typecallable)
+
+ obj = Foo()
+ adapter = collections.collection_adapter(obj.attr)
+ direct = obj.attr
+ control = set()
+
+ def assert_eq():
+ self.assert_(set(direct) == canary.data)
+ self.assert_(set(adapter) == canary.data)
+ self.assert_(direct == control)
+
+ # There is no API for object collections. We'll make one up
+ # for the purposes of the test.
+ e = creator()
+ direct.push(e)
+ control.add(e)
+ assert_eq()
+
+ direct.zark(e)
+ control.remove(e)
+ assert_eq()
+
+ e = creator()
+ direct.maybe_zark(e)
+ control.discard(e)
+ assert_eq()
+
+ e = creator()
+ direct.push(e)
+ control.add(e)
+ assert_eq()
+
+ e = creator()
+ direct.maybe_zark(e)
+ control.discard(e)
+ assert_eq()
+
+ def test_object_duck(self):
+ class MyCollection(object):
+ def __init__(self):
+ self.data = set()
+ @collection.appender
+ def push(self, item):
+ self.data.add(item)
+ @collection.remover
+ def zark(self, item):
+ self.data.remove(item)
+ @collection.removes_return()
+ def maybe_zark(self, item):
+ if item in self.data:
+ self.data.remove(item)
+ return item
+ @collection.iterator
+ def __iter__(self):
+ return iter(self.data)
+ def __eq__(self, other):
+ return self.data == other
+
+ self._test_adapter(MyCollection)
+ self._test_object(MyCollection)
+ self.assert_(getattr(MyCollection, '_sa_instrumented') ==
+ id(MyCollection))
+
+ def test_object_emulates(self):
+ class MyCollection2(object):
+ __emulates__ = None
+ def __init__(self):
+ self.data = set()
+ # looks like a list
+ def append(self, item):
+ assert False
+ @collection.appender
+ def push(self, item):
+ self.data.add(item)
+ @collection.remover
+ def zark(self, item):
+ self.data.remove(item)
+ @collection.removes_return()
+ def maybe_zark(self, item):
+ if item in self.data:
+ self.data.remove(item)
+ return item
+ @collection.iterator
+ def __iter__(self):
+ return iter(self.data)
+ def __eq__(self, other):
+ return self.data == other
+
+ self._test_adapter(MyCollection2)
+ self._test_object(MyCollection2)
+ self.assert_(getattr(MyCollection2, '_sa_instrumented') ==
+ id(MyCollection2))
+
+ def test_lifecycle(self):
+ class Foo(object):
+ pass
+
+ canary = Canary()
+ creator = entity_maker
+ manager.register_attribute(Foo, 'attr', True, extension=canary)
+
+ obj = Foo()
+ col1 = obj.attr
+
+ e1 = creator()
+ obj.attr.append(e1)
+
+ e2 = creator()
+ bulk1 = [e2]
+ # empty & sever col1 from obj
+ obj.attr = bulk1
+ self.assert_(len(col1) == 0)
+ self.assert_(len(canary.data) == 1)
+ self.assert_(obj.attr is not col1)
+ self.assert_(obj.attr is not bulk1)
+ self.assert_(obj.attr == bulk1)
+
+ e3 = creator()
+ col1.append(e3)
+ self.assert_(e3 not in canary.data)
+ self.assert_(collections.collection_adapter(col1) is None)
+
+ obj.attr[0] = e3
+ self.assert_(e3 in canary.data)
+
+class DictHelpersTest(ORMTest):
+ def define_tables(self, metadata):
+ global parents, children, Parent, Child
+
+ parents = Table('parents', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('label', String))
+ children = Table('children', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('parent_id', Integer, ForeignKey('parents.id'),
+ nullable=False),
+ Column('a', String),
+ Column('b', String),
+ Column('c', String))
+
+ class Parent(object):
+ def __init__(self, label=None):
+ self.label = label
+ class Child(object):
+ def __init__(self, a=None, b=None, c=None):
+ self.a = a
+ self.b = b
+ self.c = c
+
+ def _test_scalar_mapped(self, collection_class):
+ mapper(Child, children)
+ mapper(Parent, parents, properties={
+ 'children': relation(Child, collection_class=collection_class,
+ cascade="all, delete-orphan")
+ })
+
+ p = Parent()
+ p.children['foo'] = Child('foo', 'value')
+ p.children['bar'] = Child('bar', 'value')
+ session = create_session()
+ session.save(p)
+ session.flush()
+ pid = p.id
+ session.clear()
+
+ p = session.query(Parent).get(pid)
+
+ self.assert_(set(p.children.keys()) == set(['foo', 'bar']))
+ cid = p.children['foo'].id
+
+ collections.collection_adapter(p.children).append_with_event(
+ Child('foo', 'newvalue'))
+
+ session.save(p)
+ session.flush()
+ session.clear()
+
+ p = session.query(Parent).get(pid)
+
+ self.assert_(set(p.children.keys()) == set(['foo', 'bar']))
+ self.assert_(p.children['foo'].id != cid)
+
+ self.assert_(len(list(collections.collection_adapter(p.children))) == 2)
+ session.flush()
+ session.clear()
+
+ p = session.query(Parent).get(pid)
+ self.assert_(len(list(collections.collection_adapter(p.children))) == 2)
+
+ collections.collection_adapter(p.children).remove_with_event(
+ p.children['foo'])
+
+ self.assert_(len(list(collections.collection_adapter(p.children))) == 1)
+ session.flush()
+ session.clear()
+
+ p = session.query(Parent).get(pid)
+ self.assert_(len(list(collections.collection_adapter(p.children))) == 1)
+
+ del p.children['bar']
+ self.assert_(len(list(collections.collection_adapter(p.children))) == 0)
+ session.flush()
+ session.clear()
+
+ p = session.query(Parent).get(pid)
+ self.assert_(len(list(collections.collection_adapter(p.children))) == 0)
+
+
+ def _test_composite_mapped(self, collection_class):
+ mapper(Child, children)
+ mapper(Parent, parents, properties={
+ 'children': relation(Child, collection_class=collection_class,
+ cascade="all, delete-orphan")
+ })
+
+ p = Parent()
+ p.children[('foo', '1')] = Child('foo', '1', 'value 1')
+ p.children[('foo', '2')] = Child('foo', '2', 'value 2')
+
+ session = create_session()
+ session.save(p)
+ session.flush()
+ pid = p.id
+ session.clear()
+
+ p = session.query(Parent).get(pid)
+
+ self.assert_(set(p.children.keys()) == set([('foo', '1'), ('foo', '2')]))
+ cid = p.children[('foo', '1')].id
+
+ collections.collection_adapter(p.children).append_with_event(
+ Child('foo', '1', 'newvalue'))
+
+ session.save(p)
+ session.flush()
+ session.clear()
+
+ p = session.query(Parent).get(pid)
+
+ self.assert_(set(p.children.keys()) == set([('foo', '1'), ('foo', '2')]))
+ self.assert_(p.children[('foo', '1')].id != cid)
+
+ self.assert_(len(list(collections.collection_adapter(p.children))) == 2)
+
+ def test_mapped_collection(self):
+ collection_class = collections.mapped_collection(lambda c: c.a)
+ self._test_scalar_mapped(collection_class)
+
+ def test_mapped_collection2(self):
+ collection_class = collections.mapped_collection(lambda c: (c.a, c.b))
+ self._test_composite_mapped(collection_class)
+
+ def test_attr_mapped_collection(self):
+ collection_class = collections.attribute_mapped_collection('a')
+ self._test_scalar_mapped(collection_class)
+
+ def test_column_mapped_collection(self):
+ collection_class = collections.column_mapped_collection(children.c.a)
+ self._test_scalar_mapped(collection_class)
+
+ def test_column_mapped_collection2(self):
+ collection_class = collections.column_mapped_collection((children.c.a,
+ children.c.b))
+ self._test_composite_mapped(collection_class)
+
+ def test_mixin(self):
+ class Ordered(util.OrderedDict, collections.MappedCollection):
+ def __init__(self):
+ collections.MappedCollection.__init__(self, lambda v: v.a)
+ util.OrderedDict.__init__(self)
+ collection_class = Ordered
+ self._test_scalar_mapped(collection_class)
+
+ def test_mixin2(self):
+ class Ordered2(util.OrderedDict, collections.MappedCollection):
+ def __init__(self, keyfunc):
+ collections.MappedCollection.__init__(self, keyfunc)
+ util.OrderedDict.__init__(self)
+ collection_class = lambda: Ordered2(lambda v: (v.a, v.b))
+ self._test_composite_mapped(collection_class)
+
+if __name__ == "__main__":
+ testbase.main()
diff --git a/test/orm/compile.py b/test/orm/compile.py
index 61107ce8e..23f04db85 100644
--- a/test/orm/compile.py
+++ b/test/orm/compile.py
@@ -1,7 +1,10 @@
import testbase
from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
-class CompileTest(testbase.AssertMixin):
+
+class CompileTest(AssertMixin):
"""test various mapper compilation scenarios"""
def tearDownAll(self):
clear_mappers()
diff --git a/test/orm/cycles.py b/test/orm/cycles.py
index c53e9e846..ce3065f77 100644
--- a/test/orm/cycles.py
+++ b/test/orm/cycles.py
@@ -1,11 +1,8 @@
-from testbase import PersistTest, AssertMixin, ORMTest
-import unittest, sys, os
-from sqlalchemy import *
-import StringIO
import testbase
-
-from tables import *
-import tables
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
+from testlib.tables import *
"""test cyclical mapper relationships. Many of the assertions are provided
via running with postgres, which is strict about foreign keys.
@@ -107,22 +104,6 @@ class SelfReferentialTest(AssertMixin):
sess.delete(a)
sess.flush()
- def testeagerassertion(self):
- """test that an eager self-referential relationship raises an error."""
- class C1(Tester):
- pass
- class C2(Tester):
- pass
-
- m1 = mapper(C1, t1, properties = {
- 'c1s' : relation(C1, lazy=False),
- })
-
- try:
- m1.compile()
- assert False
- except exceptions.ArgumentError:
- assert True
class SelfReferentialNoPKTest(AssertMixin):
"""test self-referential relationship that joins on a column other than the primary key column"""
@@ -541,8 +522,6 @@ class OneToManyManyToOneTest(AssertMixin):
)
)
- print str(Person.mapper.props['balls'].primaryjoin)
-
b = Ball('some data')
p = Person('some data')
p.balls.append(b)
@@ -554,7 +533,7 @@ class OneToManyManyToOneTest(AssertMixin):
sess.save(b)
sess.save(p)
- self.assert_sql(db, lambda: sess.flush(), [
+ self.assert_sql(testbase.db, lambda: sess.flush(), [
(
"INSERT INTO person (favorite_ball_id, data) VALUES (:favorite_ball_id, :data)",
{'favorite_ball_id': None, 'data':'some data'}
@@ -608,7 +587,7 @@ class OneToManyManyToOneTest(AssertMixin):
)
])
sess.delete(p)
- self.assert_sql(db, lambda: sess.flush(), [
+ self.assert_sql(testbase.db, lambda: sess.flush(), [
# heres the post update (which is a pre-update with deletes)
(
"UPDATE person SET favorite_ball_id=:favorite_ball_id WHERE person.id = :person_id",
@@ -645,8 +624,6 @@ class OneToManyManyToOneTest(AssertMixin):
)
)
- print str(Person.mapper.props['balls'].primaryjoin)
-
b = Ball('some data')
p = Person('some data')
p.balls.append(b)
@@ -660,7 +637,7 @@ class OneToManyManyToOneTest(AssertMixin):
sess = create_session()
[sess.save(x) for x in [b,p,b2,b3,b4]]
- self.assert_sql(db, lambda: sess.flush(), [
+ self.assert_sql(testbase.db, lambda: sess.flush(), [
(
"INSERT INTO ball (person_id, data) VALUES (:person_id, :data)",
{'person_id':None, 'data':'some data'}
@@ -739,7 +716,7 @@ class OneToManyManyToOneTest(AssertMixin):
])
sess.delete(p)
- self.assert_sql(db, lambda: sess.flush(), [
+ self.assert_sql(testbase.db, lambda: sess.flush(), [
(
"UPDATE ball SET person_id=:person_id WHERE ball.id = :ball_id",
lambda ctx:{'person_id': None, 'ball_id': b.id}
@@ -851,7 +828,7 @@ class SelfReferentialPostUpdateTest(AssertMixin):
remove_child(root, cats)
# pre-trigger lazy loader on 'cats' to make the test easier
cats.children
- self.assert_sql(db, lambda: session.flush(), [
+ self.assert_sql(testbase.db, lambda: session.flush(), [
(
"UPDATE node SET prev_sibling_id=:prev_sibling_id WHERE node.id = :node_id",
lambda ctx:{'prev_sibling_id':about.id, 'node_id':stories.id}
diff --git a/test/orm/eager_relations.py b/test/orm/eager_relations.py
index 37b5ecdf7..a109be56f 100644
--- a/test/orm/eager_relations.py
+++ b/test/orm/eager_relations.py
@@ -1,9 +1,9 @@
"""basic tests of eager loaded attributes"""
+import testbase
from sqlalchemy import *
from sqlalchemy.orm import *
-import testbase
-
+from testlib import *
from fixtures import *
from query import QueryTest
@@ -20,7 +20,7 @@ class EagerTest(QueryTest):
sess = create_session()
q = sess.query(User)
- assert [User(id=7, addresses=[Address(id=1, email_address='jack@bean.com')])] == q.filter(users.c.id == 7).all()
+ assert [User(id=7, addresses=[Address(id=1, email_address='jack@bean.com')])] == q.filter(User.id==7).all()
assert fixtures.user_address_result == q.all()
def test_no_orphan(self):
@@ -66,7 +66,7 @@ class EagerTest(QueryTest):
))
q = create_session().query(User)
- l = q.filter(users.c.id==addresses.c.user_id).order_by(addresses.c.email_address).all()
+ l = q.filter(User.id==Address.user_id).order_by(Address.email_address).all()
assert [
User(id=8, addresses=[
@@ -148,7 +148,7 @@ class EagerTest(QueryTest):
assert fixtures.user_address_result == sess.query(User).all()
def test_double(self):
- """tests lazy loading with two relations simulatneously, from the same table, using aliases. """
+ """tests eager loading with two relations simulatneously, from the same table, using aliases. """
openorders = alias(orders, 'openorders')
closedorders = alias(orders, 'closedorders')
@@ -185,6 +185,46 @@ class EagerTest(QueryTest):
] == q.all()
self.assert_sql_count(testbase.db, go, 1)
+
+ def test_double_same_mappers(self):
+ """tests eager loading with two relations simulatneously, from the same table, using aliases. """
+
+ mapper(Address, addresses)
+ mapper(Order, orders, properties={
+ 'items':relation(Item, secondary=order_items, lazy=False, order_by=items.c.id),
+ })
+ mapper(Item, items)
+ mapper(User, users, properties = dict(
+ addresses = relation(Address, lazy=False),
+ open_orders = relation(Order, primaryjoin = and_(orders.c.isopen == 1, users.c.id==orders.c.user_id), lazy=False),
+ closed_orders = relation(Order, primaryjoin = and_(orders.c.isopen == 0, users.c.id==orders.c.user_id), lazy=False)
+ ))
+ q = create_session().query(User)
+
+ def go():
+ assert [
+ User(
+ id=7,
+ addresses=[Address(id=1)],
+ open_orders = [Order(id=3, items=[Item(id=3), Item(id=4), Item(id=5)])],
+ closed_orders = [Order(id=1, items=[Item(id=1), Item(id=2), Item(id=3)]), Order(id=5, items=[Item(id=5)])]
+ ),
+ User(
+ id=8,
+ addresses=[Address(id=2), Address(id=3), Address(id=4)],
+ open_orders = [],
+ closed_orders = []
+ ),
+ User(
+ id=9,
+ addresses=[Address(id=5)],
+ open_orders = [Order(id=4, items=[Item(id=1), Item(id=5)])],
+ closed_orders = [Order(id=2, items=[Item(id=1), Item(id=2), Item(id=3)])]
+ ),
+ User(id=10)
+
+ ] == q.all()
+ self.assert_sql_count(testbase.db, go, 1)
def test_limit(self):
"""test limit operations combined with lazy-load relationships."""
@@ -236,7 +276,7 @@ class EagerTest(QueryTest):
sess = create_session()
q = sess.query(Item)
l = q.filter((Item.c.description=='item 2') | (Item.c.description=='item 5') | (Item.c.description=='item 3')).\
- order_by(Item.c.id).limit(2).all()
+ order_by(Item.id).limit(2).all()
assert fixtures.item_keyword_result[1:3] == l
@@ -259,7 +299,7 @@ class EagerTest(QueryTest):
q = sess.query(User)
if testbase.db.engine.name != 'mssql':
- l = q.join('orders').order_by(desc(orders.c.user_id)).limit(2).offset(1)
+ l = q.join('orders').order_by(desc(Order.user_id)).limit(2).offset(1)
assert [
User(id=9,
orders=[Order(id=2), Order(id=4)],
@@ -271,7 +311,7 @@ class EagerTest(QueryTest):
)
] == l.all()
- l = q.join('addresses').order_by(desc(addresses.c.email_address)).limit(1).offset(0)
+ l = q.join('addresses').order_by(desc(Address.email_address)).limit(1).offset(0)
assert [
User(id=7,
orders=[Order(id=1), Order(id=3), Order(id=5)],
@@ -375,6 +415,7 @@ class EagerTest(QueryTest):
'user':relation(User, lazy=False)
})
mapper(User, users)
+ mapper(Item, items)
q = create_session().query(Order)
assert [
@@ -382,7 +423,7 @@ class EagerTest(QueryTest):
Order(id=4, user=User(id=9))
] == q.all()
- q = q.select_from(s.join(order_items).join(items)).filter(~items.c.id.in_(1, 2, 5))
+ q = q.select_from(s.join(order_items).join(items)).filter(~Item.id.in_(1, 2, 5))
assert [
Order(id=3, user=User(id=7)),
] == q.all()
@@ -394,8 +435,80 @@ class EagerTest(QueryTest):
addresses = relation(mapper(Address, addresses), lazy=False)
))
q = create_session().query(User)
- l = q.filter(addresses.c.email_address == 'ed@lala.com').filter(addresses.c.user_id==users.c.id)
+ l = q.filter(addresses.c.email_address == 'ed@lala.com').filter(Address.user_id==User.id)
assert fixtures.user_address_result[1:2] == l.all()
+class SelfReferentialEagerTest(ORMTest):
+ def define_tables(self, metadata):
+ global nodes
+ nodes = Table('nodes', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('parent_id', Integer, ForeignKey('nodes.id')),
+ Column('data', String(30)))
+
+ def test_basic(self):
+ class Node(Base):
+ def append(self, node):
+ self.children.append(node)
+
+ mapper(Node, nodes, properties={
+ 'children':relation(Node, lazy=False, join_depth=3)
+ })
+ sess = create_session()
+ n1 = Node(data='n1')
+ n1.append(Node(data='n11'))
+ n1.append(Node(data='n12'))
+ n1.append(Node(data='n13'))
+ n1.children[1].append(Node(data='n121'))
+ n1.children[1].append(Node(data='n122'))
+ n1.children[1].append(Node(data='n123'))
+ sess.save(n1)
+ sess.flush()
+ sess.clear()
+ def go():
+ d = sess.query(Node).filter_by(data='n1').first()
+ assert Node(data='n1', children=[
+ Node(data='n11'),
+ Node(data='n12', children=[
+ Node(data='n121'),
+ Node(data='n122'),
+ Node(data='n123')
+ ]),
+ Node(data='n13')
+ ]) == d
+ self.assert_sql_count(testbase.db, go, 1)
+
+ def test_no_depth(self):
+ class Node(Base):
+ def append(self, node):
+ self.children.append(node)
+
+ mapper(Node, nodes, properties={
+ 'children':relation(Node, lazy=False)
+ })
+ sess = create_session()
+ n1 = Node(data='n1')
+ n1.append(Node(data='n11'))
+ n1.append(Node(data='n12'))
+ n1.append(Node(data='n13'))
+ n1.children[1].append(Node(data='n121'))
+ n1.children[1].append(Node(data='n122'))
+ n1.children[1].append(Node(data='n123'))
+ sess.save(n1)
+ sess.flush()
+ sess.clear()
+ def go():
+ d = sess.query(Node).filter_by(data='n1').first()
+ assert Node(data='n1', children=[
+ Node(data='n11'),
+ Node(data='n12', children=[
+ Node(data='n121'),
+ Node(data='n122'),
+ Node(data='n123')
+ ]),
+ Node(data='n13')
+ ]) == d
+ self.assert_sql_count(testbase.db, go, 3)
+
if __name__ == '__main__':
testbase.main()
diff --git a/test/orm/eagertest1.py b/test/orm/eagertest1.py
deleted file mode 100644
index 9765379f4..000000000
--- a/test/orm/eagertest1.py
+++ /dev/null
@@ -1,69 +0,0 @@
-from testbase import PersistTest, AssertMixin
-import testbase
-import unittest, sys, os
-from sqlalchemy import *
-import datetime
-
-class EagerTest(AssertMixin):
- def setUpAll(self):
- global designType, design, part, inheritedPart
- designType = Table('design_types', testbase.metadata,
- Column('design_type_id', Integer, primary_key=True),
- )
-
- design =Table('design', testbase.metadata,
- Column('design_id', Integer, primary_key=True),
- Column('design_type_id', Integer, ForeignKey('design_types.design_type_id')))
-
- part = Table('parts', testbase.metadata,
- Column('part_id', Integer, primary_key=True),
- Column('design_id', Integer, ForeignKey('design.design_id')),
- Column('design_type_id', Integer, ForeignKey('design_types.design_type_id')))
-
- inheritedPart = Table('inherited_part', testbase.metadata,
- Column('ip_id', Integer, primary_key=True),
- Column('part_id', Integer, ForeignKey('parts.part_id')),
- Column('design_id', Integer, ForeignKey('design.design_id')),
- )
-
- testbase.metadata.create_all()
- def tearDownAll(self):
- testbase.metadata.drop_all()
- testbase.metadata.clear()
- def testone(self):
- class Part(object):pass
- class Design(object):pass
- class DesignType(object):pass
- class InheritedPart(object):pass
-
- mapper(Part, part)
-
- mapper(InheritedPart, inheritedPart, properties=dict(
- part=relation(Part, lazy=False)
- ))
-
- mapper(Design, design, properties=dict(
- parts=relation(Part, private=True, backref="design"),
- inheritedParts=relation(InheritedPart, private=True, backref="design"),
- ))
-
- mapper(DesignType, designType, properties=dict(
- # designs=relation(Design, private=True, backref="type"),
- ))
-
- class_mapper(Design).add_property("type", relation(DesignType, lazy=False, backref="designs"))
- class_mapper(Part).add_property("design", relation(Design, lazy=False, backref="parts"))
- #Part.mapper.add_property("designType", relation(DesignType))
-
- d = Design()
- sess = create_session()
- sess.save(d)
- sess.flush()
- sess.clear()
- x = sess.query(Design).get(1)
- x.inheritedParts
-
-if __name__ == "__main__":
- testbase.main()
-
-
diff --git a/test/orm/eagertest2.py b/test/orm/eagertest2.py
deleted file mode 100644
index 04de56f01..000000000
--- a/test/orm/eagertest2.py
+++ /dev/null
@@ -1,239 +0,0 @@
-from testbase import PersistTest, AssertMixin
-import testbase
-import unittest, sys, os
-from sqlalchemy import *
-import datetime
-from sqlalchemy.ext.sessioncontext import SessionContext
-
-class EagerTest(AssertMixin):
- def setUpAll(self):
- global companies_table, addresses_table, invoice_table, phones_table, items_table, ctx, metadata
-
- metadata = MetaData(testbase.db)
- ctx = SessionContext(create_session)
-
- companies_table = Table('companies', metadata,
- Column('company_id', Integer, Sequence('company_id_seq', optional=True), primary_key = True),
- Column('company_name', String(40)),
-
- )
-
- addresses_table = Table('addresses', metadata,
- Column('address_id', Integer, Sequence('address_id_seq', optional=True), primary_key = True),
- Column('company_id', Integer, ForeignKey("companies.company_id")),
- Column('address', String(40)),
- )
-
- phones_table = Table('phone_numbers', metadata,
- Column('phone_id', Integer, Sequence('phone_id_seq', optional=True), primary_key = True),
- Column('address_id', Integer, ForeignKey('addresses.address_id')),
- Column('type', String(20)),
- Column('number', String(10)),
- )
-
- invoice_table = Table('invoices', metadata,
- Column('invoice_id', Integer, Sequence('invoice_id_seq', optional=True), primary_key = True),
- Column('company_id', Integer, ForeignKey("companies.company_id")),
- Column('date', DateTime),
- )
-
- items_table = Table('items', metadata,
- Column('item_id', Integer, Sequence('item_id_seq', optional=True), primary_key = True),
- Column('invoice_id', Integer, ForeignKey('invoices.invoice_id')),
- Column('code', String(20)),
- Column('qty', Integer),
- )
-
- metadata.create_all()
-
- def tearDownAll(self):
- metadata.drop_all()
-
- def tearDown(self):
- clear_mappers()
- for t in metadata.table_iterator(reverse=True):
- t.delete().execute()
-
- def testone(self):
- """tests eager load of a many-to-one attached to a one-to-many. this testcase illustrated
- the bug, which is that when the single Company is loaded, no further processing of the rows
- occurred in order to load the Company's second Address object."""
- class Company(object):
- def __init__(self):
- self.company_id = None
- def __repr__(self):
- return "Company:" + repr(getattr(self, 'company_id', None)) + " " + repr(getattr(self, 'company_name', None)) + " " + str([repr(addr) for addr in self.addresses])
-
- class Address(object):
- def __repr__(self):
- return "Address: " + repr(getattr(self, 'address_id', None)) + " " + repr(getattr(self, 'company_id', None)) + " " + repr(self.address)
-
- class Invoice(object):
- def __init__(self):
- self.invoice_id = None
- def __repr__(self):
- return "Invoice:" + repr(getattr(self, 'invoice_id', None)) + " " + repr(getattr(self, 'date', None)) + " " + repr(self.company)
-
- mapper(Address, addresses_table, properties={
- }, extension=ctx.mapper_extension)
- mapper(Company, companies_table, properties={
- 'addresses' : relation(Address, lazy=False),
- }, extension=ctx.mapper_extension)
- mapper(Invoice, invoice_table, properties={
- 'company': relation(Company, lazy=False, )
- }, extension=ctx.mapper_extension)
-
- c1 = Company()
- c1.company_name = 'company 1'
- a1 = Address()
- a1.address = 'a1 address'
- c1.addresses.append(a1)
- a2 = Address()
- a2.address = 'a2 address'
- c1.addresses.append(a2)
- i1 = Invoice()
- i1.date = datetime.datetime.now()
- i1.company = c1
-
- ctx.current.flush()
-
- company_id = c1.company_id
- invoice_id = i1.invoice_id
-
- ctx.current.clear()
-
- c = ctx.current.query(Company).get(company_id)
-
- ctx.current.clear()
-
- i = ctx.current.query(Invoice).get(invoice_id)
-
- self.echo(repr(c))
- self.echo(repr(i.company))
- self.assert_(repr(c) == repr(i.company))
-
- def testtwo(self):
- """this is the original testcase that includes various complicating factors"""
- class Company(object):
- def __init__(self):
- self.company_id = None
- def __repr__(self):
- return "Company:" + repr(getattr(self, 'company_id', None)) + " " + repr(getattr(self, 'company_name', None)) + " " + str([repr(addr) for addr in self.addresses])
-
- class Address(object):
- def __repr__(self):
- return "Address: " + repr(getattr(self, 'address_id', None)) + " " + repr(getattr(self, 'company_id', None)) + " " + repr(self.address) + str([repr(ph) for ph in self.phones])
-
- class Phone(object):
- def __repr__(self):
- return "Phone: " + repr(getattr(self, 'phone_id', None)) + " " + repr(getattr(self, 'address_id', None)) + " " + repr(self.type) + " " + repr(self.number)
-
- class Invoice(object):
- def __init__(self):
- self.invoice_id = None
- def __repr__(self):
- return "Invoice:" + repr(getattr(self, 'invoice_id', None)) + " " + repr(getattr(self, 'date', None)) + " " + repr(self.company) + " " + str([repr(item) for item in self.items])
-
- class Item(object):
- def __repr__(self):
- return "Item: " + repr(getattr(self, 'item_id', None)) + " " + repr(getattr(self, 'invoice_id', None)) + " " + repr(self.code) + " " + repr(self.qty)
-
- mapper(Phone, phones_table, extension=ctx.mapper_extension)
-
- mapper(Address, addresses_table, properties={
- 'phones': relation(Phone, lazy=False, backref='address')
- }, extension=ctx.mapper_extension)
-
- mapper(Company, companies_table, properties={
- 'addresses' : relation(Address, lazy=False, backref='company'),
- }, extension=ctx.mapper_extension)
-
- mapper(Item, items_table, extension=ctx.mapper_extension)
-
- mapper(Invoice, invoice_table, properties={
- 'items': relation(Item, lazy=False, backref='invoice'),
- 'company': relation(Company, lazy=False, backref='invoices')
- }, extension=ctx.mapper_extension)
-
- ctx.current.clear()
- c1 = Company()
- c1.company_name = 'company 1'
-
- a1 = Address()
- a1.address = 'a1 address'
-
- p1 = Phone()
- p1.type = 'home'
- p1.number = '1111'
-
- a1.phones.append(p1)
-
- p2 = Phone()
- p2.type = 'work'
- p2.number = '22222'
- a1.phones.append(p2)
-
- c1.addresses.append(a1)
-
- a2 = Address()
- a2.address = 'a2 address'
-
- p3 = Phone()
- p3.type = 'home'
- p3.number = '3333'
- a2.phones.append(p3)
-
- p4 = Phone()
- p4.type = 'work'
- p4.number = '44444'
- a2.phones.append(p4)
-
- c1.addresses.append(a2)
-
- ctx.current.flush()
-
- company_id = c1.company_id
-
- ctx.current.clear()
-
- a = ctx.current.query(Company).get(company_id)
- self.echo(repr(a))
-
- # set up an invoice
- i1 = Invoice()
- i1.date = datetime.datetime.now()
- i1.company = c1
-
- item1 = Item()
- item1.code = 'aaaa'
- item1.qty = 1
- item1.invoice = i1
-
- item2 = Item()
- item2.code = 'bbbb'
- item2.qty = 2
- item2.invoice = i1
-
- item3 = Item()
- item3.code = 'cccc'
- item3.qty = 3
- item3.invoice = i1
-
- ctx.current.flush()
-
- invoice_id = i1.invoice_id
-
- ctx.current.clear()
-
- c = ctx.current.query(Company).get(company_id)
- self.echo(repr(c))
-
- ctx.current.clear()
-
- i = ctx.current.query(Invoice).get(invoice_id)
- self.echo(repr(i))
-
- self.assert_(repr(i.company) == repr(c))
-
-if __name__ == "__main__":
- testbase.main()
diff --git a/test/orm/entity.py b/test/orm/entity.py
index 86486cafc..da76e8df0 100644
--- a/test/orm/entity.py
+++ b/test/orm/entity.py
@@ -1,11 +1,9 @@
-from testbase import PersistTest, AssertMixin
-import unittest
-from sqlalchemy import *
import testbase
+from sqlalchemy import *
+from sqlalchemy.orm import *
from sqlalchemy.ext.sessioncontext import SessionContext
-
-from tables import *
-import tables
+from testlib import *
+from testlib.tables import *
class EntityTest(AssertMixin):
"""tests mappers that are constructed based on "entity names", which allows the same class
diff --git a/test/orm/fixtures.py b/test/orm/fixtures.py
index aed74c170..4a7d41459 100644
--- a/test/orm/fixtures.py
+++ b/test/orm/fixtures.py
@@ -1,4 +1,6 @@
+import testbase
from sqlalchemy import *
+from testlib import *
_recursion_stack = util.Set()
class Base(object):
@@ -35,7 +37,7 @@ class Base(object):
continue
else:
if value is not None:
- if value != getattr(other, attr):
+ if value != getattr(other, attr, None):
return False
else:
return True
diff --git a/test/orm/generative.py b/test/orm/generative.py
index 75280deed..4a90c13cb 100644
--- a/test/orm/generative.py
+++ b/test/orm/generative.py
@@ -1,16 +1,19 @@
-from testbase import PersistTest, AssertMixin, ORMTest
import testbase
-import tables
-
from sqlalchemy import *
+from sqlalchemy.orm import *
from sqlalchemy import exceptions
+from testlib import *
+import testlib.tables as tables
+
+# TODO: these are more tests that should be updated to be part of test/orm/query.py
class Foo(object):
- pass
+ def __init__(self, **kwargs):
+ for k in kwargs:
+ setattr(self, k, kwargs[k])
class GenerativeQueryTest(PersistTest):
def setUpAll(self):
- self.install_threadlocal()
global foo, metadata
metadata = MetaData(testbase.db)
foo = Table('foo', metadata,
@@ -18,89 +21,103 @@ class GenerativeQueryTest(PersistTest):
Column('bar', Integer),
Column('range', Integer))
- assign_mapper(Foo, foo)
+ mapper(Foo, foo)
metadata.create_all()
+
+ sess = create_session(bind=testbase.db)
for i in range(100):
- Foo(bar=i, range=i%10)
- objectstore.flush()
+ sess.save(Foo(bar=i, range=i%10))
+ sess.flush()
- def setUp(self):
- self.query = Foo.query()
- self.orig = self.query.select_whereclause()
- self.res = self.query
-
def tearDownAll(self):
metadata.drop_all()
- self.uninstall_threadlocal()
clear_mappers()
def test_selectby(self):
- res = self.query.filter_by(range=5)
+ res = create_session(bind=testbase.db).query(Foo).filter_by(range=5)
assert res.order_by([Foo.c.bar])[0].bar == 5
assert res.order_by([desc(Foo.c.bar)])[0].bar == 95
- @testbase.unsupported('mssql')
+ @testing.unsupported('mssql')
def test_slice(self):
- assert self.query[1] == self.orig[1]
- assert list(self.query[10:20]) == self.orig[10:20]
- assert list(self.query[10:]) == self.orig[10:]
- assert list(self.query[:10]) == self.orig[:10]
- assert list(self.query[:10]) == self.orig[:10]
- assert list(self.query[10:40:3]) == self.orig[10:40:3]
- assert list(self.query[-5:]) == self.orig[-5:]
- assert self.query[10:20][5] == self.orig[10:20][5]
+ sess = create_session(bind=testbase.db)
+ query = sess.query(Foo)
+ orig = query.all()
+ assert query[1] == orig[1]
+ assert list(query[10:20]) == orig[10:20]
+ assert list(query[10:]) == orig[10:]
+ assert list(query[:10]) == orig[:10]
+ assert list(query[:10]) == orig[:10]
+ assert list(query[10:40:3]) == orig[10:40:3]
+ assert list(query[-5:]) == orig[-5:]
+ assert query[10:20][5] == orig[10:20][5]
- @testbase.supported('mssql')
+ @testing.supported('mssql')
def test_slice_mssql(self):
- assert list(self.query[:10]) == self.orig[:10]
- assert list(self.query[:10]) == self.orig[:10]
+ sess = create_session(bind=testbase.db)
+ query = sess.query(Foo)
+ orig = query.all()
+ assert list(query[:10]) == orig[:10]
+ assert list(query[:10]) == orig[:10]
def test_aggregate(self):
- assert self.query.count() == 100
- assert self.query.filter(foo.c.bar<30).min(foo.c.bar) == 0
- assert self.query.filter(foo.c.bar<30).max(foo.c.bar) == 29
- assert self.query.filter(foo.c.bar<30).apply_max(foo.c.bar).scalar() == 29
+ sess = create_session(bind=testbase.db)
+ query = sess.query(Foo)
+ assert query.count() == 100
+ assert query.filter(foo.c.bar<30).min(foo.c.bar) == 0
+ assert query.filter(foo.c.bar<30).max(foo.c.bar) == 29
+ assert query.filter(foo.c.bar<30).apply_max(foo.c.bar).first() == 29
+ assert query.filter(foo.c.bar<30).apply_max(foo.c.bar).one() == 29
- @testbase.unsupported('mysql')
+ @testing.unsupported('mysql')
def test_aggregate_1(self):
# this one fails in mysql as the result comes back as a string
- assert self.query.filter(foo.c.bar<30).sum(foo.c.bar) == 435
+ query = create_session(bind=testbase.db).query(Foo)
+ assert query.filter(foo.c.bar<30).sum(foo.c.bar) == 435
- @testbase.unsupported('postgres', 'mysql', 'firebird', 'mssql')
+ @testing.unsupported('postgres', 'mysql', 'firebird', 'mssql')
def test_aggregate_2(self):
- assert self.res.filter(foo.c.bar<30).avg(foo.c.bar) == 14.5
+ query = create_session(bind=testbase.db).query(Foo)
+ assert query.filter(foo.c.bar<30).avg(foo.c.bar) == 14.5
- @testbase.supported('postgres', 'mysql', 'firebird', 'mssql')
+ @testing.supported('postgres', 'mysql', 'firebird', 'mssql')
def test_aggregate_2_int(self):
- assert int(self.res.filter(foo.c.bar<30).avg(foo.c.bar)) == 14
+ query = create_session(bind=testbase.db).query(Foo)
+ assert int(query.filter(foo.c.bar<30).avg(foo.c.bar)) == 14
- @testbase.unsupported('postgres', 'mysql', 'firebird', 'mssql')
+ @testing.unsupported('postgres', 'mysql', 'firebird', 'mssql')
def test_aggregate_3(self):
- assert self.res.filter(foo.c.bar<30).apply_avg(foo.c.bar).scalar() == 14.5
+ query = create_session(bind=testbase.db).query(Foo)
+ assert query.filter(foo.c.bar<30).apply_avg(foo.c.bar).first() == 14.5
+ assert query.filter(foo.c.bar<30).apply_avg(foo.c.bar).one() == 14.5
def test_filter(self):
- assert self.query.count() == 100
- assert self.query.filter(Foo.c.bar < 30).count() == 30
- res2 = self.query.filter(Foo.c.bar < 30).filter(Foo.c.bar > 10)
+ query = create_session(bind=testbase.db).query(Foo)
+ assert query.count() == 100
+ assert query.filter(Foo.c.bar < 30).count() == 30
+ res2 = query.filter(Foo.c.bar < 30).filter(Foo.c.bar > 10)
assert res2.count() == 19
def test_options(self):
+ query = create_session(bind=testbase.db).query(Foo)
class ext1(MapperExtension):
- def populate_instance(self, mapper, selectcontext, row, instance, identitykey, isnew):
+ def populate_instance(self, mapper, selectcontext, row, instance, **flags):
instance.TEST = "hello world"
return EXT_PASS
- objectstore.clear()
- assert self.res.options(extension(ext1()))[0].TEST == "hello world"
+ assert query.options(extension(ext1()))[0].TEST == "hello world"
def test_order_by(self):
- assert self.res.order_by([Foo.c.bar])[0].bar == 0
- assert self.res.order_by([desc(Foo.c.bar)])[0].bar == 99
+ query = create_session(bind=testbase.db).query(Foo)
+ assert query.order_by([Foo.c.bar])[0].bar == 0
+ assert query.order_by([desc(Foo.c.bar)])[0].bar == 99
def test_offset(self):
- assert list(self.res.order_by([Foo.c.bar]).offset(10))[0].bar == 10
+ query = create_session(bind=testbase.db).query(Foo)
+ assert list(query.order_by([Foo.c.bar]).offset(10))[0].bar == 10
def test_offset(self):
- assert len(list(self.res.limit(10))) == 10
+ query = create_session(bind=testbase.db).query(Foo)
+ assert len(list(query.limit(10))) == 10
class Obj1(object):
pass
@@ -109,9 +126,8 @@ class Obj2(object):
class GenerativeTest2(PersistTest):
def setUpAll(self):
- self.install_threadlocal()
global metadata, table1, table2
- metadata = MetaData(testbase.db)
+ metadata = MetaData()
table1 = Table('Table1', metadata,
Column('id', Integer, primary_key=True),
)
@@ -119,29 +135,23 @@ class GenerativeTest2(PersistTest):
Column('t1id', Integer, ForeignKey("Table1.id"), primary_key=True),
Column('num', Integer, primary_key=True),
)
- assign_mapper(Obj1, table1)
- assign_mapper(Obj2, table2)
- metadata.create_all()
- table1.insert().execute({'id':1},{'id':2},{'id':3},{'id':4})
- table2.insert().execute({'num':1,'t1id':1},{'num':2,'t1id':1},{'num':3,'t1id':1},\
+ mapper(Obj1, table1)
+ mapper(Obj2, table2)
+ metadata.create_all(bind=testbase.db)
+ testbase.db.execute(table1.insert(), {'id':1},{'id':2},{'id':3},{'id':4})
+ testbase.db.execute(table2.insert(), {'num':1,'t1id':1},{'num':2,'t1id':1},{'num':3,'t1id':1},\
{'num':4,'t1id':2},{'num':5,'t1id':2},{'num':6,'t1id':3})
- def setUp(self):
- self.query = Query(Obj1)
- #self.orig = self.query.select_whereclause()
- #self.res = self.query.select()
-
def tearDownAll(self):
- metadata.drop_all()
- self.uninstall_threadlocal()
+ metadata.drop_all(bind=testbase.db)
clear_mappers()
def test_distinctcount(self):
- res = self.query
- assert res.count() == 4
- res = self.query.filter(and_(table1.c.id==table2.c.t1id,table2.c.t1id==1))
+ query = create_session(bind=testbase.db).query(Obj1)
+ assert query.count() == 4
+ res = query.filter(and_(table1.c.id==table2.c.t1id,table2.c.t1id==1))
assert res.count() == 3
- res = self.query.filter(and_(table1.c.id==table2.c.t1id,table2.c.t1id==1)).distinct()
+ res = query.filter(and_(table1.c.id==table2.c.t1id,table2.c.t1id==1)).distinct()
self.assertEqual(res.count(), 1)
class RelationsTest(AssertMixin):
@@ -159,9 +169,9 @@ class RelationsTest(AssertMixin):
'items':relation(mapper(tables.Item, tables.orderitems))
}))
})
- session = create_session()
+ session = create_session(bind=testbase.db)
query = session.query(tables.User)
- x = query.join('orders').join('items').filter(tables.Item.c.item_id==2)
+ x = query.join(['orders', 'items']).filter(tables.Item.c.item_id==2)
print x.compile()
self.assert_result(list(x), tables.User, tables.user_result[2])
def test_outerjointo(self):
@@ -171,9 +181,9 @@ class RelationsTest(AssertMixin):
'items':relation(mapper(tables.Item, tables.orderitems))
}))
})
- session = create_session()
+ session = create_session(bind=testbase.db)
query = session.query(tables.User)
- x = query.outerjoin('orders').outerjoin('items').filter(or_(tables.Order.c.order_id==None,tables.Item.c.item_id==2))
+ x = query.outerjoin(['orders', 'items']).filter(or_(tables.Order.c.order_id==None,tables.Item.c.item_id==2))
print x.compile()
self.assert_result(list(x), tables.User, *tables.user_result[1:3])
def test_outerjointo_count(self):
@@ -183,9 +193,9 @@ class RelationsTest(AssertMixin):
'items':relation(mapper(tables.Item, tables.orderitems))
}))
})
- session = create_session()
+ session = create_session(bind=testbase.db)
query = session.query(tables.User)
- x = query.outerjoin('orders').outerjoin('items').filter(or_(tables.Order.c.order_id==None,tables.Item.c.item_id==2)).count()
+ x = query.outerjoin(['orders', 'items']).filter(or_(tables.Order.c.order_id==None,tables.Item.c.item_id==2)).count()
assert x==2
def test_from(self):
mapper(tables.User, tables.users, properties={
@@ -193,7 +203,7 @@ class RelationsTest(AssertMixin):
'items':relation(mapper(tables.Item, tables.orderitems))
}))
})
- session = create_session()
+ session = create_session(bind=testbase.db)
query = session.query(tables.User)
x = query.select_from([tables.users.outerjoin(tables.orders).outerjoin(tables.orderitems)]).\
filter(or_(tables.Order.c.order_id==None,tables.Item.c.item_id==2))
@@ -203,7 +213,6 @@ class RelationsTest(AssertMixin):
class CaseSensitiveTest(PersistTest):
def setUpAll(self):
- self.install_threadlocal()
global metadata, table1, table2
metadata = MetaData(testbase.db)
table1 = Table('Table1', metadata,
@@ -213,29 +222,23 @@ class CaseSensitiveTest(PersistTest):
Column('T1ID', Integer, ForeignKey("Table1.ID"), primary_key=True),
Column('NUM', Integer, primary_key=True),
)
- assign_mapper(Obj1, table1)
- assign_mapper(Obj2, table2)
+ mapper(Obj1, table1)
+ mapper(Obj2, table2)
metadata.create_all()
table1.insert().execute({'ID':1},{'ID':2},{'ID':3},{'ID':4})
table2.insert().execute({'NUM':1,'T1ID':1},{'NUM':2,'T1ID':1},{'NUM':3,'T1ID':1},\
{'NUM':4,'T1ID':2},{'NUM':5,'T1ID':2},{'NUM':6,'T1ID':3})
- def setUp(self):
- self.query = Query(Obj1)
- #self.orig = self.query.select_whereclause()
- #self.res = self.query.select()
-
def tearDownAll(self):
metadata.drop_all()
- self.uninstall_threadlocal()
clear_mappers()
def test_distinctcount(self):
- res = self.query
- assert res.count() == 4
- res = self.query.filter(and_(table1.c.ID==table2.c.T1ID,table2.c.T1ID==1))
+ q = create_session(bind=testbase.db).query(Obj1)
+ assert q.count() == 4
+ res = q.filter(and_(table1.c.ID==table2.c.T1ID,table2.c.T1ID==1))
assert res.count() == 3
- res = self.query.filter(and_(table1.c.ID==table2.c.T1ID,table2.c.T1ID==1)).distinct()
+ res = q.filter(and_(table1.c.ID==table2.c.T1ID,table2.c.T1ID==1)).distinct()
self.assertEqual(res.count(), 1)
class SelfRefTest(ORMTest):
@@ -248,18 +251,18 @@ class SelfRefTest(ORMTest):
def test_noautojoin(self):
class T(object):pass
mapper(T, t1, properties={'children':relation(T)})
- sess = create_session()
+ sess = create_session(bind=testbase.db)
try:
sess.query(T).join('children').select_by(id=7)
assert False
except exceptions.InvalidRequestError, e:
- assert str(e) == "Self-referential query on 'T.children (T)' property must be constructed manually using an Alias object for the related table.", str(e)
+ assert str(e) == "Self-referential query on 'T.children (T)' property requires create_aliases=True argument.", str(e)
try:
sess.query(T).join(['children']).select_by(id=7)
assert False
except exceptions.InvalidRequestError, e:
- assert str(e) == "Self-referential query on 'T.children (T)' property must be constructed manually using an Alias object for the related table.", str(e)
+ assert str(e) == "Self-referential query on 'T.children (T)' property requires create_aliases=True argument.", str(e)
diff --git a/test/orm/inheritance/__init__.py b/test/orm/inheritance/__init__.py
new file mode 100644
index 000000000..e69de29bb
--- /dev/null
+++ b/test/orm/inheritance/__init__.py
diff --git a/test/orm/abc_inheritance.py b/test/orm/inheritance/abc_inheritance.py
index 7689bd543..3b35b3713 100644
--- a/test/orm/abc_inheritance.py
+++ b/test/orm/inheritance/abc_inheritance.py
@@ -1,12 +1,14 @@
+import testbase
from sqlalchemy import *
+from sqlalchemy.orm import *
from sqlalchemy.orm.sync import ONETOMANY, MANYTOONE
-import testbase
+from testlib import *
def produce_test(parent, child, direction):
"""produce a testcase for A->B->C inheritance with a self-referential
relationship between two of the classes, using either one-to-many or
many-to-one."""
- class ABCTest(testbase.ORMTest):
+ class ABCTest(ORMTest):
def define_tables(self, meta):
global ta, tb, tc
ta = ["a", meta]
@@ -53,6 +55,8 @@ def produce_test(parent, child, direction):
parent_table = {"a":ta, "b":tb, "c": tc}[parent]
child_table = {"a":ta, "b":tb, "c": tc}[child]
+ remote_side = None
+
if direction == MANYTOONE:
foreign_keys = [parent_table.c.child_id]
elif direction == ONETOMANY:
@@ -65,6 +69,8 @@ def produce_test(parent, child, direction):
relationjoin = parent_table.c.id==child_table.c.parent_id
elif direction == MANYTOONE:
relationjoin = parent_table.c.child_id==child_table.c.id
+ if parent is child:
+ remote_side = [child_table.c.id]
abcjoin = polymorphic_union(
{"a":ta.select(tb.c.id==None, from_obj=[ta.outerjoin(tb, onclause=atob)]),
@@ -79,19 +85,15 @@ def produce_test(parent, child, direction):
"c":tc.join(tb, onclause=btoc).join(ta, onclause=atob)
},"type", "bcjoin"
)
-
- class A(object):pass
+ class A(object):
+ def __init__(self, name):
+ self.a_data = name
class B(A):pass
class C(B):pass
- mapper(A, ta, polymorphic_on=abcjoin.c.type, select_table=abcjoin, polymorphic_identity="a", )
- mapper(B, tb, polymorphic_on=bcjoin.c.type, select_table=bcjoin, polymorphic_identity="b", inherits=A, inherit_condition=atob,)
- mapper(C, tc, polymorphic_identity="c", inherits=B, inherit_condition=btoc, )
-
- #print "KEYS:"
- #print [c.key for c in class_mapper(A).primary_key]
- #print [c.key for c in class_mapper(B).primary_key]
- #print [c.key for c in class_mapper(C).primary_key]
+ mapper(A, ta, polymorphic_on=abcjoin.c.type, select_table=abcjoin, polymorphic_identity="a")
+ mapper(B, tb, polymorphic_on=bcjoin.c.type, select_table=bcjoin, polymorphic_identity="b", inherits=A, inherit_condition=atob)
+ mapper(C, tc, polymorphic_identity="c", inherits=B, inherit_condition=btoc)
parent_mapper = class_mapper({ta:A, tb:B, tc:C}[parent_table])
child_mapper = class_mapper({ta:A, tb:B, tc:C}[child_table])
@@ -99,24 +101,24 @@ def produce_test(parent, child, direction):
parent_class = parent_mapper.class_
child_class = child_mapper.class_
- parent_mapper.add_property("collection", relation(child_mapper, primaryjoin=relationjoin, foreign_keys=foreign_keys, uselist=True))
+ parent_mapper.add_property("collection", relation(child_mapper, primaryjoin=relationjoin, foreign_keys=foreign_keys, remote_side=remote_side, uselist=True))
sess = create_session()
- parent_obj = parent_class()
- child_obj = child_class()
- somea = A()
- someb = B()
- somec = C()
+ parent_obj = parent_class('parent1')
+ child_obj = child_class('child1')
+ somea = A('somea')
+ someb = B('someb')
+ somec = C('somec')
print "APPENDING", parent.__class__.__name__ , "TO", child.__class__.__name__
sess.save(parent_obj)
parent_obj.collection.append(child_obj)
if direction == ONETOMANY:
- child2 = child_class()
+ child2 = child_class('child2')
parent_obj.collection.append(child2)
sess.save(child2)
elif direction == MANYTOONE:
- parent2 = parent_class()
+ parent2 = parent_class('parent2')
parent2.collection.append(child_obj)
sess.save(parent2)
sess.save(somea)
@@ -155,8 +157,6 @@ def produce_test(parent, child, direction):
# test all combinations of polymorphic a/b/c related to another of a/b/c
for parent in ["a", "b", "c"]:
for child in ["a", "b", "c"]:
- if parent == child:
- continue
for direction in [ONETOMANY, MANYTOONE]:
testclass = produce_test(parent, child, direction)
exec("%s = testclass" % testclass.__name__)
diff --git a/test/orm/inheritance/alltests.py b/test/orm/inheritance/alltests.py
new file mode 100644
index 000000000..1ab10c060
--- /dev/null
+++ b/test/orm/inheritance/alltests.py
@@ -0,0 +1,28 @@
+import testbase
+import unittest
+
+def suite():
+ modules_to_test = (
+ 'orm.inheritance.basic',
+ 'orm.inheritance.manytomany',
+ 'orm.inheritance.single',
+ 'orm.inheritance.concrete',
+ 'orm.inheritance.polymorph',
+ 'orm.inheritance.polymorph2',
+ 'orm.inheritance.poly_linked_list',
+ 'orm.inheritance.abc_inheritance',
+ 'orm.inheritance.productspec',
+ 'orm.inheritance.magazine',
+
+ )
+ alltests = unittest.TestSuite()
+ for name in modules_to_test:
+ mod = __import__(name)
+ for token in name.split('.')[1:]:
+ mod = getattr(mod, token)
+ alltests.addTest(unittest.findTestCases(mod, suiteClass=None))
+ return alltests
+
+
+if __name__ == '__main__':
+ testbase.main(suite())
diff --git a/test/orm/inheritance.py b/test/orm/inheritance/basic.py
index 2281a0597..be623e1b8 100644
--- a/test/orm/inheritance.py
+++ b/test/orm/inheritance/basic.py
@@ -1,159 +1,13 @@
import testbase
from sqlalchemy import *
-import string
-import sys
+from sqlalchemy.orm import *
+from testlib import *
-class Principal( object ):
- def __init__(self, **kwargs):
- for key, value in kwargs.iteritems():
- setattr(self, key, value)
-class User( Principal ):
- pass
-
-class Group( Principal ):
- pass
-
-class InheritTest(testbase.ORMTest):
- """deals with inheritance and many-to-many relationships"""
- def define_tables(self, metadata):
- global principals
- global users
- global groups
- global user_group_map
-
- principals = Table(
- 'principals',
- metadata,
- Column('principal_id', Integer, Sequence('principal_id_seq', optional=False), primary_key=True),
- Column('name', String(50), nullable=False),
- )
-
- users = Table(
- 'prin_users',
- metadata,
- Column('principal_id', Integer, ForeignKey('principals.principal_id'), primary_key=True),
- Column('password', String(50), nullable=False),
- Column('email', String(50), nullable=False),
- Column('login_id', String(50), nullable=False),
-
- )
-
- groups = Table(
- 'prin_groups',
- metadata,
- Column( 'principal_id', Integer, ForeignKey('principals.principal_id'), primary_key=True),
-
- )
-
- user_group_map = Table(
- 'prin_user_group_map',
- metadata,
- Column('user_id', Integer, ForeignKey( "prin_users.principal_id"), primary_key=True ),
- Column('group_id', Integer, ForeignKey( "prin_groups.principal_id"), primary_key=True ),
- #Column('user_id', Integer, ForeignKey( "prin_users.principal_id"), ),
- #Column('group_id', Integer, ForeignKey( "prin_groups.principal_id"), ),
-
- )
-
- def testbasic(self):
- mapper( Principal, principals )
- mapper(
- User,
- users,
- inherits=Principal
- )
-
- mapper(
- Group,
- groups,
- inherits=Principal,
- properties=dict( users = relation(User, secondary=user_group_map, lazy=True, backref="groups") )
- )
-
- g = Group(name="group1")
- g.users.append(User(name="user1", password="pw", email="foo@bar.com", login_id="lg1"))
- sess = create_session()
- sess.save(g)
- sess.flush()
- # TODO: put an assertion
-
-class InheritTest2(testbase.ORMTest):
- """deals with inheritance and many-to-many relationships"""
- def define_tables(self, metadata):
- global foo, bar, foo_bar
- foo = Table('foo', metadata,
- Column('id', Integer, Sequence('foo_id_seq'), primary_key=True),
- Column('data', String(20)),
- )
-
- bar = Table('bar', metadata,
- Column('bid', Integer, ForeignKey('foo.id'), primary_key=True),
- #Column('fid', Integer, ForeignKey('foo.id'), )
- )
-
- foo_bar = Table('foo_bar', metadata,
- Column('foo_id', Integer, ForeignKey('foo.id')),
- Column('bar_id', Integer, ForeignKey('bar.bid')))
-
- def testget(self):
- class Foo(object):
- def __init__(self, data=None):
- self.data = data
- class Bar(Foo):pass
-
- mapper(Foo, foo)
- mapper(Bar, bar, inherits=Foo)
-
- b = Bar('somedata')
- sess = create_session()
- sess.save(b)
- sess.flush()
- sess.clear()
-
- # test that "bar.bid" does not need to be referenced in a get
- # (ticket 185)
- assert sess.query(Bar).get(b.id).id == b.id
-
- def testbasic(self):
- class Foo(object):
- def __init__(self, data=None):
- self.data = data
-
- mapper(Foo, foo)
- class Bar(Foo):
- pass
-
- mapper(Bar, bar, inherits=Foo, properties={
- 'foos': relation(Foo, secondary=foo_bar, lazy=False)
- })
-
- sess = create_session()
- b = Bar('barfoo')
- sess.save(b)
- sess.flush()
-
- f1 = Foo('subfoo1')
- f2 = Foo('subfoo2')
- b.foos.append(f1)
- b.foos.append(f2)
-
- sess.flush()
- sess.clear()
-
- l = sess.query(Bar).select()
- print l[0]
- print l[0].foos
- self.assert_result(l, Bar,
-# {'id':1, 'data':'barfoo', 'bid':1, 'foos':(Foo, [{'id':2,'data':'subfoo1'}, {'id':3,'data':'subfoo2'}])},
- {'id':b.id, 'data':'barfoo', 'foos':(Foo, [{'id':f1.id,'data':'subfoo1'}, {'id':f2.id,'data':'subfoo2'}])},
- )
-
-class InheritTest3(testbase.ORMTest):
- """deals with inheritance and many-to-many relationships"""
+class O2MTest(ORMTest):
+ """deals with inheritance and one-to-many relationships"""
def define_tables(self, metadata):
- global foo, bar, blub, bar_foo, blub_bar, blub_foo
-
+ global foo, bar, blub
# the 'data' columns are to appease SQLite which cant handle a blank INSERT
foo = Table('foo', metadata,
Column('id', Integer, Sequence('foo_seq'), primary_key=True),
@@ -165,20 +19,9 @@ class InheritTest3(testbase.ORMTest):
blub = Table('blub', metadata,
Column('id', Integer, ForeignKey('bar.id'), primary_key=True),
+ Column('foo_id', Integer, ForeignKey('foo.id'), nullable=False),
Column('data', String(20)))
- bar_foo = Table('bar_foo', metadata,
- Column('bar_id', Integer, ForeignKey('bar.id')),
- Column('foo_id', Integer, ForeignKey('foo.id')))
-
- blub_bar = Table('bar_blub', metadata,
- Column('blub_id', Integer, ForeignKey('blub.id')),
- Column('bar_id', Integer, ForeignKey('bar.id')))
-
- blub_foo = Table('blub_foo', metadata,
- Column('blub_id', Integer, ForeignKey('blub.id')),
- Column('foo_id', Integer, ForeignKey('foo.id')))
-
def testbasic(self):
class Foo(object):
def __init__(self, data=None):
@@ -190,71 +33,41 @@ class InheritTest3(testbase.ORMTest):
class Bar(Foo):
def __repr__(self):
return "Bar id %d, data %s" % (self.id, self.data)
-
- mapper(Bar, bar, inherits=Foo, properties={
- 'foos' :relation(Foo, secondary=bar_foo, lazy=True)
- })
-
- sess = create_session()
- b = Bar('bar #1', _sa_session=sess)
- b.foos.append(Foo("foo #1"))
- b.foos.append(Foo("foo #2"))
- sess.flush()
- compare = repr(b) + repr(b.foos)
- sess.clear()
- l = sess.query(Bar).select()
- self.echo(repr(l[0]) + repr(l[0].foos))
- self.assert_(repr(l[0]) + repr(l[0].foos) == compare)
-
- def testadvanced(self):
- class Foo(object):
- def __init__(self, data=None):
- self.data = data
- def __repr__(self):
- return "Foo id %d, data %s" % (self.id, self.data)
- mapper(Foo, foo)
- class Bar(Foo):
- def __repr__(self):
- return "Bar id %d, data %s" % (self.id, self.data)
mapper(Bar, bar, inherits=Foo)
-
+
class Blub(Bar):
def __repr__(self):
- return "Blub id %d, data %s, bars %s, foos %s" % (self.id, self.data, repr([b for b in self.bars]), repr([f for f in self.foos]))
-
+ return "Blub id %d, data %s" % (self.id, self.data)
+
mapper(Blub, blub, inherits=Bar, properties={
- 'bars':relation(Bar, secondary=blub_bar, lazy=False),
- 'foos':relation(Foo, secondary=blub_foo, lazy=False),
+ 'parent_foo':relation(Foo)
})
sess = create_session()
- f1 = Foo("foo #1", _sa_session=sess)
- b1 = Bar("bar #1", _sa_session=sess)
- b2 = Bar("bar #2", _sa_session=sess)
- bl1 = Blub("blub #1", _sa_session=sess)
- bl1.foos.append(f1)
- bl1.bars.append(b2)
+ b1 = Blub("blub #1")
+ b2 = Blub("blub #2")
+ f = Foo("foo #1")
+ sess.save(b1)
+ sess.save(b2)
+ sess.save(f)
+ b1.parent_foo = f
+ b2.parent_foo = f
sess.flush()
- compare = repr(bl1)
- blubid = bl1.id
+ compare = repr(b1) + repr(b2) + repr(b1.parent_foo) + repr(b2.parent_foo)
sess.clear()
-
l = sess.query(Blub).select()
- self.echo(l)
- self.assert_(repr(l[0]) == compare)
- sess.clear()
- x = sess.query(Blub).get_by(id=blubid)
- self.echo(x)
- self.assert_(repr(x) == compare)
-
-class InheritTest4(testbase.ORMTest):
- """deals with inheritance and one-to-many relationships"""
+ result = repr(l[0]) + repr(l[1]) + repr(l[0].parent_foo) + repr(l[1].parent_foo)
+ print result
+ self.assert_(compare == result)
+ self.assert_(l[0].parent_foo.data == 'foo #1' and l[1].parent_foo.data == 'foo #1')
+
+class GetTest(ORMTest):
def define_tables(self, metadata):
global foo, bar, blub
- # the 'data' columns are to appease SQLite which cant handle a blank INSERT
foo = Table('foo', metadata,
Column('id', Integer, Sequence('foo_seq'), primary_key=True),
+ Column('type', String(30)),
Column('data', String(20)))
bar = Table('bar', metadata,
@@ -262,50 +75,80 @@ class InheritTest4(testbase.ORMTest):
Column('data', String(20)))
blub = Table('blub', metadata,
- Column('id', Integer, ForeignKey('bar.id'), primary_key=True),
- Column('foo_id', Integer, ForeignKey('foo.id'), nullable=False),
+ Column('id', Integer, primary_key=True),
+ Column('foo_id', Integer, ForeignKey('foo.id')),
+ Column('bar_id', Integer, ForeignKey('bar.id')),
Column('data', String(20)))
+
+ def create_test(polymorphic):
+ def test_get(self):
+ class Foo(object):
+ pass
- def testbasic(self):
- class Foo(object):
- def __init__(self, data=None):
- self.data = data
- def __repr__(self):
- return "Foo id %d, data %s" % (self.id, self.data)
- mapper(Foo, foo)
-
- class Bar(Foo):
- def __repr__(self):
- return "Bar id %d, data %s" % (self.id, self.data)
-
- mapper(Bar, bar, inherits=Foo)
+ class Bar(Foo):
+ pass
- class Blub(Bar):
- def __repr__(self):
- return "Blub id %d, data %s" % (self.id, self.data)
-
- mapper(Blub, blub, inherits=Bar, properties={
- 'parent_foo':relation(Foo)
- })
+ class Blub(Bar):
+ pass
+
+ if polymorphic:
+ mapper(Foo, foo, polymorphic_on=foo.c.type, polymorphic_identity='foo')
+ mapper(Bar, bar, inherits=Foo, polymorphic_identity='bar')
+ mapper(Blub, blub, inherits=Bar, polymorphic_identity='blub')
+ else:
+ mapper(Foo, foo)
+ mapper(Bar, bar, inherits=Foo)
+ mapper(Blub, blub, inherits=Bar)
+
+ sess = create_session()
+ f = Foo()
+ b = Bar()
+ bl = Blub()
+ sess.save(f)
+ sess.save(b)
+ sess.save(bl)
+ sess.flush()
+
+ if polymorphic:
+ def go():
+ assert sess.query(Foo).get(f.id) == f
+ assert sess.query(Foo).get(b.id) == b
+ assert sess.query(Foo).get(bl.id) == bl
+ assert sess.query(Bar).get(b.id) == b
+ assert sess.query(Bar).get(bl.id) == bl
+ assert sess.query(Blub).get(bl.id) == bl
+
+ self.assert_sql_count(testbase.db, go, 0)
+ else:
+ # this is testing the 'wrong' behavior of using get()
+ # polymorphically with mappers that are not configured to be
+ # polymorphic. the important part being that get() always
+ # returns an instance of the query's type.
+ def go():
+ assert sess.query(Foo).get(f.id) == f
+
+ bb = sess.query(Foo).get(b.id)
+ assert isinstance(b, Foo) and bb.id==b.id
+
+ bll = sess.query(Foo).get(bl.id)
+ assert isinstance(bll, Foo) and bll.id==bl.id
+
+ assert sess.query(Bar).get(b.id) == b
+
+ bll = sess.query(Bar).get(bl.id)
+ assert isinstance(bll, Bar) and bll.id == bl.id
+
+ assert sess.query(Blub).get(bl.id) == bl
+
+ self.assert_sql_count(testbase.db, go, 3)
+
+ return test_get
+
+ test_get_polymorphic = create_test(True)
+ test_get_nonpolymorphic = create_test(False)
- sess = create_session()
- b1 = Blub("blub #1", _sa_session=sess)
- b2 = Blub("blub #2", _sa_session=sess)
- f = Foo("foo #1", _sa_session=sess)
- b1.parent_foo = f
- b2.parent_foo = f
- sess.flush()
- compare = repr(b1) + repr(b2) + repr(b1.parent_foo) + repr(b2.parent_foo)
- sess.clear()
- l = sess.query(Blub).select()
- result = repr(l[0]) + repr(l[1]) + repr(l[0].parent_foo) + repr(l[1].parent_foo)
- self.echo(result)
- self.assert_(compare == result)
- self.assert_(l[0].parent_foo.data == 'foo #1' and l[1].parent_foo.data == 'foo #1')
-class InheritTest5(testbase.ORMTest):
- """testing that construction of inheriting mappers works regardless of when extra properties
- are added to the superclass mapper"""
+class ConstructionTest(ORMTest):
def define_tables(self, metadata):
global content_type, content, product
content_type = Table('content_type', metadata,
@@ -313,7 +156,8 @@ class InheritTest5(testbase.ORMTest):
)
content = Table('content', metadata,
Column('id', Integer, primary_key=True),
- Column('content_type_id', Integer, ForeignKey('content_type.id'))
+ Column('content_type_id', Integer, ForeignKey('content_type.id')),
+ Column('type', String(30))
)
product = Table('product', metadata,
Column('id', Integer, ForeignKey('content.id'), primary_key=True)
@@ -327,11 +171,15 @@ class InheritTest5(testbase.ORMTest):
content_types = mapper(ContentType, content_type)
contents = mapper(Content, content, properties={
'content_type':relation(content_types)
- })
- #contents.add_property('content_type', relation(content_types)) #adding this makes the inheritance stop working
- # shouldnt throw exception
- products = mapper(Product, product, inherits=contents)
- # TODO: assertion ??
+ }, polymorphic_identity='contents')
+
+ products = mapper(Product, product, inherits=contents, polymorphic_identity='products')
+
+ try:
+ compile_mappers()
+ assert False
+ except exceptions.ArgumentError, e:
+ assert str(e) == "Mapper 'Mapper|Content|content' specifies a polymorphic_identity of 'contents', but no mapper in it's hierarchy specifies the 'polymorphic_on' column argument"
def testbackref(self):
"""tests adding a property to the superclass mapper"""
@@ -339,8 +187,8 @@ class InheritTest5(testbase.ORMTest):
class Content(object): pass
class Product(Content): pass
- contents = mapper(Content, content)
- products = mapper(Product, product, inherits=contents)
+ contents = mapper(Content, content, polymorphic_on=content.c.type, polymorphic_identity='content')
+ products = mapper(Product, product, inherits=contents, polymorphic_identity='product')
content_types = mapper(ContentType, content_type, properties={
'content':relation(contents, backref='contenttype')
})
@@ -348,7 +196,7 @@ class InheritTest5(testbase.ORMTest):
p.contenttype = ContentType()
# TODO: assertion ??
-class InheritTest6(testbase.ORMTest):
+class EagerLazyTest(ORMTest):
"""tests eager load/lazy load of child items off inheritance mappers, tests that
LazyLoader constructs the right query condition."""
def define_tables(self, metadata):
@@ -370,7 +218,6 @@ class InheritTest6(testbase.ORMTest):
foos = mapper(Foo, foo)
bars = mapper(Bar, bar, inherits=foos)
bars.add_property('lazy', relation(foos, bar_foo, lazy=True))
- print bars.props['lazy'].primaryjoin, bars.props['lazy'].secondaryjoin
bars.add_property('eager', relation(foos, bar_foo, lazy=False))
foo.insert().execute(data='foo1')
@@ -391,7 +238,7 @@ class InheritTest6(testbase.ORMTest):
self.assert_(len(q.selectfirst().eager) == 1)
-class InheritTest7(testbase.ORMTest):
+class FlushTest(ORMTest):
"""test dependency sorting among inheriting mappers"""
def define_tables(self, metadata):
global users, roles, user_roles, admins
@@ -412,22 +259,20 @@ class InheritTest7(testbase.ORMTest):
)
admins = Table('admin', metadata,
- Column('id', Integer, primary_key=True),
+ Column('admin_id', Integer, primary_key=True),
Column('user_id', Integer, ForeignKey('users.id'))
)
def testone(self):
class User(object):pass
- class Role(object):
- def __init__(self, description):
- self.description = description
+ class Role(object):pass
class Admin(User):pass
role_mapper = mapper(Role, roles)
user_mapper = mapper(User, users, properties = {
'roles' : relation(Role, secondary=user_roles, lazy=False, private=False)
}
)
- admin_mapper = mapper(Admin, admins, inherits=user_mapper, properties={'aid':admins.c.id})
+ admin_mapper = mapper(Admin, admins, inherits=user_mapper)
sess = create_session()
adminrole = Role('admin')
sess.save(adminrole)
@@ -435,7 +280,7 @@ class InheritTest7(testbase.ORMTest):
# create an Admin, and append a Role. the dependency processors
# corresponding to the "roles" attribute for the Admin mapper and the User mapper
- # have to insure that two dependency processors dont fire off and insert the
+ # have to ensure that two dependency processors dont fire off and insert the
# many to many row twice.
a = Admin()
a.roles.append(adminrole)
@@ -463,7 +308,7 @@ class InheritTest7(testbase.ORMTest):
}
)
- admin_mapper = mapper(Admin, admins, inherits=user_mapper, properties={'aid':admins.c.id})
+ admin_mapper = mapper(Admin, admins, inherits=user_mapper)
# create roles
adminrole = Role('admin')
@@ -482,14 +327,14 @@ class InheritTest7(testbase.ORMTest):
sess.flush()
assert user_roles.count().scalar() == 1
-class InheritTest8(testbase.ORMTest):
+class DistinctPKTest(ORMTest):
"""test the construction of mapper.primary_key when an inheriting relationship
joins on a column other than primary key column."""
keep_data = True
-
+
def define_tables(self, metadata):
global person_table, employee_table, Person, Employee
-
+
person_table = Table("persons", metadata,
Column("id", Integer, primary_key=True),
Column("name", String(80)),
@@ -509,7 +354,7 @@ class InheritTest8(testbase.ORMTest):
import warnings
warnings.filterwarnings("error", r".*On mapper.*distinct primary key")
-
+
def insert_data(self):
person_insert = person_table.insert()
person_insert.execute(id=1, name='alice')
@@ -518,22 +363,17 @@ class InheritTest8(testbase.ORMTest):
employee_insert = employee_table.insert()
employee_insert.execute(id=2, salary=250, person_id=1) # alice
employee_insert.execute(id=3, salary=200, person_id=2) # bob
-
+
def test_implicit(self):
person_mapper = mapper(Person, person_table)
mapper(Employee, employee_table, inherits=person_mapper)
- try:
- print class_mapper(Employee).primary_key
- assert list(class_mapper(Employee).primary_key) == [person_table.c.id, employee_table.c.id]
- assert False
- except RuntimeWarning, e:
- assert str(e) == "On mapper Mapper|Employee|employees, primary key column 'employees.id' is being combined with distinct primary key column 'persons.id' in attribute 'id'. Use explicit properties to give each column its own mapped attribute name."
+ assert list(class_mapper(Employee).primary_key) == [person_table.c.id]
def test_explicit_props(self):
person_mapper = mapper(Person, person_table)
mapper(Employee, employee_table, inherits=person_mapper, properties={'pid':person_table.c.id, 'eid':employee_table.c.id})
self._do_test(True)
-
+
def test_explicit_composite_pk(self):
person_mapper = mapper(Person, person_table)
mapper(Employee, employee_table, inherits=person_mapper, primary_key=[person_table.c.id, employee_table.c.id])
@@ -547,17 +387,12 @@ class InheritTest8(testbase.ORMTest):
person_mapper = mapper(Person, person_table)
mapper(Employee, employee_table, inherits=person_mapper, primary_key=[person_table.c.id])
self._do_test(False)
-
+
def _do_test(self, composite):
session = create_session()
query = session.query(Employee)
if composite:
- try:
- query.get(1)
- assert False
- except exceptions.InvalidRequestError, e:
- assert str(e) == "Could not find enough values to formulate primary key for query.get(); primary key columns are 'persons.id', 'employees.id'"
alice1 = query.get([1,2])
bob = query.get([2,3])
alice2 = query.get([1,2])
@@ -565,11 +400,10 @@ class InheritTest8(testbase.ORMTest):
alice1 = query.get(1)
bob = query.get(2)
alice2 = query.get(1)
-
+
assert alice1.name == alice2.name == 'alice'
assert bob.name == 'bob'
-
-
+
if __name__ == "__main__":
testbase.main()
diff --git a/test/orm/inheritance4.py b/test/orm/inheritance/concrete.py
index 9f4e275ae..d95a96da5 100644
--- a/test/orm/inheritance4.py
+++ b/test/orm/inheritance/concrete.py
@@ -1,7 +1,9 @@
-from sqlalchemy import *
import testbase
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
-class ConcreteTest1(testbase.ORMTest):
+class ConcreteTest1(ORMTest):
def define_tables(self, metadata):
global managers_table, engineers_table
managers_table = Table('managers', metadata,
@@ -52,6 +54,7 @@ class ConcreteTest1(testbase.ORMTest):
session.flush()
session.clear()
+ print set([repr(x) for x in session.query(Employee).select()])
assert set([repr(x) for x in session.query(Employee).select()]) == set(["Engineer Kurt knows how to hack", "Manager Tom knows how to manage things"])
assert set([repr(x) for x in session.query(Manager).select()]) == set(["Manager Tom knows how to manage things"])
assert set([repr(x) for x in session.query(Engineer).select()]) == set(["Engineer Kurt knows how to hack"])
@@ -63,4 +66,4 @@ class ConcreteTest1(testbase.ORMTest):
if __name__ == '__main__':
- testbase.main() \ No newline at end of file
+ testbase.main()
diff --git a/test/orm/inheritance3.py b/test/orm/inheritance/magazine.py
index a9c88ef60..a0bf24148 100644
--- a/test/orm/inheritance3.py
+++ b/test/orm/inheritance/magazine.py
@@ -1,5 +1,8 @@
import testbase
from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
+
class BaseObject(object):
def __init__(self, *args, **kwargs):
@@ -20,7 +23,7 @@ class Location(BaseObject):
def _set_name(self, name):
session = create_session()
- s = session.query(LocationName).selectfirst(location_name_table.c.name==name)
+ s = session.query(LocationName).filter(LocationName.name==name).first()
session.clear()
if s is not None:
self._name = s
@@ -64,8 +67,8 @@ class MagazinePage(Page):
class ClassifiedPage(MagazinePage):
pass
-class InheritTest(testbase.ORMTest):
- """tests a large polymorphic relationship"""
+
+class MagazineTest(ORMTest):
def define_tables(self, metadata):
global publication_table, issue_table, location_table, location_name_table, magazine_table, \
page_table, magazine_page_table, classified_page_table, page_size_table
@@ -116,6 +119,8 @@ class InheritTest(testbase.ORMTest):
Column('name', String(45), default=''),
)
+def generate_round_trip_test(use_unions=False, use_joins=False):
+ def test_roundtrip(self):
publication_mapper = mapper(Publication, publication_table)
issue_mapper = mapper(Issue, issue_table, properties = {
@@ -133,33 +138,50 @@ class InheritTest(testbase.ORMTest):
page_size_mapper = mapper(PageSize, page_size_table)
- page_join = polymorphic_union(
- {
- 'm': page_table.join(magazine_page_table),
- 'c': page_table.join(magazine_page_table).join(classified_page_table),
- 'p': page_table.select(page_table.c.type=='p'),
- }, None, 'page_join')
-
- magazine_join = polymorphic_union(
- {
- 'm': page_table.join(magazine_page_table),
- 'c': page_table.join(magazine_page_table).join(classified_page_table),
- }, None, 'page_join')
-
magazine_mapper = mapper(Magazine, magazine_table, properties = {
'location': relation(Location, backref=backref('magazine', uselist=False)),
'size': relation(PageSize),
})
- page_mapper = mapper(Page, page_table, select_table=page_join, polymorphic_on=page_join.c.type, polymorphic_identity='p')
+ if use_unions:
+ page_join = polymorphic_union(
+ {
+ 'm': page_table.join(magazine_page_table),
+ 'c': page_table.join(magazine_page_table).join(classified_page_table),
+ 'p': page_table.select(page_table.c.type=='p'),
+ }, None, 'page_join')
+ page_mapper = mapper(Page, page_table, select_table=page_join, polymorphic_on=page_join.c.type, polymorphic_identity='p')
+ elif use_joins:
+ page_join = page_table.outerjoin(magazine_page_table).outerjoin(classified_page_table)
+ page_mapper = mapper(Page, page_table, select_table=page_join, polymorphic_on=page_table.c.type, polymorphic_identity='p')
+ else:
+ page_mapper = mapper(Page, page_table, polymorphic_on=page_table.c.type, polymorphic_identity='p')
+
+ if use_unions:
+ magazine_join = polymorphic_union(
+ {
+ 'm': page_table.join(magazine_page_table),
+ 'c': page_table.join(magazine_page_table).join(classified_page_table),
+ }, None, 'page_join')
+ magazine_page_mapper = mapper(MagazinePage, magazine_page_table, select_table=magazine_join, inherits=page_mapper, polymorphic_identity='m', properties={
+ 'magazine': relation(Magazine, backref=backref('pages', order_by=magazine_join.c.page_no))
+ })
+ elif use_joins:
+ magazine_join = page_table.join(magazine_page_table).outerjoin(classified_page_table)
+ magazine_page_mapper = mapper(MagazinePage, magazine_page_table, select_table=magazine_join, inherits=page_mapper, polymorphic_identity='m', properties={
+ 'magazine': relation(Magazine, backref=backref('pages', order_by=page_table.c.page_no))
+ })
+ else:
+ magazine_page_mapper = mapper(MagazinePage, magazine_page_table, inherits=page_mapper, polymorphic_identity='m', properties={
+ 'magazine': relation(Magazine, backref=backref('pages', order_by=page_table.c.page_no))
+ })
+
+ classified_page_mapper = mapper(ClassifiedPage, classified_page_table, inherits=magazine_page_mapper, polymorphic_identity='c', primary_key=[page_table.c.id])
+ #compile_mappers()
+ #print [str(s) for s in classified_page_mapper.primary_key]
+ #print classified_page_mapper.columntoproperty[page_table.c.id]
- magazine_page_mapper = mapper(MagazinePage, magazine_page_table, select_table=magazine_join, inherits=page_mapper, polymorphic_identity='m', properties={
- 'magazine': relation(Magazine, backref=backref('pages', order_by=magazine_join.c.page_no))
- })
-
- classified_page_mapper = mapper(ClassifiedPage, classified_page_table, inherits=magazine_page_mapper, polymorphic_identity='c')
- def testone(self):
session = create_session()
pub = Publication(name='Test')
@@ -174,18 +196,25 @@ class InheritTest(testbase.ORMTest):
page2 = MagazinePage(magazine=magazine,page_no=2)
page3 = ClassifiedPage(magazine=magazine,page_no=3)
session.save(pub)
-
+
session.flush()
print [x for x in session]
session.clear()
session.flush()
session.clear()
- p = session.query(Publication).selectone_by(name='Test')
+ p = session.query(Publication).filter(Publication.name=="Test").one()
print p.issues[0].locations[0].magazine.pages
print [page, page2, page3]
- assert repr(p.issues[0].locations[0].magazine.pages) == repr([page, page2, page3])
+ assert repr(p.issues[0].locations[0].magazine.pages) == repr([page, page2, page3]), repr(p.issues[0].locations[0].magazine.pages)
+
+ test_roundtrip.__name__ = "test_%s" % (not use_union and (use_joins and "joins" or "select") or "unions")
+ setattr(MagazineTest, test_roundtrip.__name__, test_roundtrip)
+
+for (use_union, use_join) in [(True, False), (False, True), (False, False)]:
+ generate_round_trip_test(use_union, use_join)
+
if __name__ == '__main__':
testbase.main()
diff --git a/test/orm/inheritance/manytomany.py b/test/orm/inheritance/manytomany.py
new file mode 100644
index 000000000..df00f39d0
--- /dev/null
+++ b/test/orm/inheritance/manytomany.py
@@ -0,0 +1,255 @@
+import testbase
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
+
+
+class InheritTest(ORMTest):
+ """deals with inheritance and many-to-many relationships"""
+ def define_tables(self, metadata):
+ global principals
+ global users
+ global groups
+ global user_group_map
+
+ principals = Table(
+ 'principals',
+ metadata,
+ Column('principal_id', Integer, Sequence('principal_id_seq', optional=False), primary_key=True),
+ Column('name', String(50), nullable=False),
+ )
+
+ users = Table(
+ 'prin_users',
+ metadata,
+ Column('principal_id', Integer, ForeignKey('principals.principal_id'), primary_key=True),
+ Column('password', String(50), nullable=False),
+ Column('email', String(50), nullable=False),
+ Column('login_id', String(50), nullable=False),
+
+ )
+
+ groups = Table(
+ 'prin_groups',
+ metadata,
+ Column( 'principal_id', Integer, ForeignKey('principals.principal_id'), primary_key=True),
+
+ )
+
+ user_group_map = Table(
+ 'prin_user_group_map',
+ metadata,
+ Column('user_id', Integer, ForeignKey( "prin_users.principal_id"), primary_key=True ),
+ Column('group_id', Integer, ForeignKey( "prin_groups.principal_id"), primary_key=True ),
+ )
+
+ def testbasic(self):
+ class Principal(object):
+ def __init__(self, **kwargs):
+ for key, value in kwargs.iteritems():
+ setattr(self, key, value)
+
+ class User(Principal):
+ pass
+
+ class Group(Principal):
+ pass
+
+ mapper(Principal, principals)
+ mapper(
+ User,
+ users,
+ inherits=Principal
+ )
+
+ mapper(
+ Group,
+ groups,
+ inherits=Principal,
+ properties=dict( users = relation(User, secondary=user_group_map, lazy=True, backref="groups") )
+ )
+
+ g = Group(name="group1")
+ g.users.append(User(name="user1", password="pw", email="foo@bar.com", login_id="lg1"))
+ sess = create_session()
+ sess.save(g)
+ sess.flush()
+ # TODO: put an assertion
+
+class InheritTest2(ORMTest):
+ """deals with inheritance and many-to-many relationships"""
+ def define_tables(self, metadata):
+ global foo, bar, foo_bar
+ foo = Table('foo', metadata,
+ Column('id', Integer, Sequence('foo_id_seq'), primary_key=True),
+ Column('data', String(20)),
+ )
+
+ bar = Table('bar', metadata,
+ Column('bid', Integer, ForeignKey('foo.id'), primary_key=True),
+ #Column('fid', Integer, ForeignKey('foo.id'), )
+ )
+
+ foo_bar = Table('foo_bar', metadata,
+ Column('foo_id', Integer, ForeignKey('foo.id')),
+ Column('bar_id', Integer, ForeignKey('bar.bid')))
+
+ def testget(self):
+ class Foo(object):pass
+ def __init__(self, data=None):
+ self.data = data
+ class Bar(Foo):pass
+
+ mapper(Foo, foo)
+ mapper(Bar, bar, inherits=Foo)
+ print foo.join(bar).primary_key
+ print class_mapper(Bar).primary_key
+ b = Bar('somedata')
+ sess = create_session()
+ sess.save(b)
+ sess.flush()
+ sess.clear()
+
+ # test that "bar.bid" does not need to be referenced in a get
+ # (ticket 185)
+ assert sess.query(Bar).get(b.id).id == b.id
+
+ def testbasic(self):
+ class Foo(object):
+ def __init__(self, data=None):
+ self.data = data
+
+ mapper(Foo, foo)
+ class Bar(Foo):
+ pass
+
+ mapper(Bar, bar, inherits=Foo, properties={
+ 'foos': relation(Foo, secondary=foo_bar, lazy=False)
+ })
+
+ sess = create_session()
+ b = Bar('barfoo')
+ sess.save(b)
+ sess.flush()
+
+ f1 = Foo('subfoo1')
+ f2 = Foo('subfoo2')
+ b.foos.append(f1)
+ b.foos.append(f2)
+
+ sess.flush()
+ sess.clear()
+
+ l = sess.query(Bar).select()
+ print l[0]
+ print l[0].foos
+ self.assert_result(l, Bar,
+# {'id':1, 'data':'barfoo', 'bid':1, 'foos':(Foo, [{'id':2,'data':'subfoo1'}, {'id':3,'data':'subfoo2'}])},
+ {'id':b.id, 'data':'barfoo', 'foos':(Foo, [{'id':f1.id,'data':'subfoo1'}, {'id':f2.id,'data':'subfoo2'}])},
+ )
+
+class InheritTest3(ORMTest):
+ """deals with inheritance and many-to-many relationships"""
+ def define_tables(self, metadata):
+ global foo, bar, blub, bar_foo, blub_bar, blub_foo
+
+ # the 'data' columns are to appease SQLite which cant handle a blank INSERT
+ foo = Table('foo', metadata,
+ Column('id', Integer, Sequence('foo_seq'), primary_key=True),
+ Column('data', String(20)))
+
+ bar = Table('bar', metadata,
+ Column('id', Integer, ForeignKey('foo.id'), primary_key=True),
+ Column('data', String(20)))
+
+ blub = Table('blub', metadata,
+ Column('id', Integer, ForeignKey('bar.id'), primary_key=True),
+ Column('data', String(20)))
+
+ bar_foo = Table('bar_foo', metadata,
+ Column('bar_id', Integer, ForeignKey('bar.id')),
+ Column('foo_id', Integer, ForeignKey('foo.id')))
+
+ blub_bar = Table('bar_blub', metadata,
+ Column('blub_id', Integer, ForeignKey('blub.id')),
+ Column('bar_id', Integer, ForeignKey('bar.id')))
+
+ blub_foo = Table('blub_foo', metadata,
+ Column('blub_id', Integer, ForeignKey('blub.id')),
+ Column('foo_id', Integer, ForeignKey('foo.id')))
+
+ def testbasic(self):
+ class Foo(object):
+ def __init__(self, data=None):
+ self.data = data
+ def __repr__(self):
+ return "Foo id %d, data %s" % (self.id, self.data)
+ mapper(Foo, foo)
+
+ class Bar(Foo):
+ def __repr__(self):
+ return "Bar id %d, data %s" % (self.id, self.data)
+
+ mapper(Bar, bar, inherits=Foo, properties={
+ 'foos' :relation(Foo, secondary=bar_foo, lazy=True)
+ })
+
+ sess = create_session()
+ b = Bar('bar #1')
+ sess.save(b)
+ b.foos.append(Foo("foo #1"))
+ b.foos.append(Foo("foo #2"))
+ sess.flush()
+ compare = repr(b) + repr(b.foos)
+ sess.clear()
+ l = sess.query(Bar).select()
+ print repr(l[0]) + repr(l[0].foos)
+ self.assert_(repr(l[0]) + repr(l[0].foos) == compare)
+
+ def testadvanced(self):
+ class Foo(object):
+ def __init__(self, data=None):
+ self.data = data
+ def __repr__(self):
+ return "Foo id %d, data %s" % (self.id, self.data)
+ mapper(Foo, foo)
+
+ class Bar(Foo):
+ def __repr__(self):
+ return "Bar id %d, data %s" % (self.id, self.data)
+ mapper(Bar, bar, inherits=Foo)
+
+ class Blub(Bar):
+ def __repr__(self):
+ return "Blub id %d, data %s, bars %s, foos %s" % (self.id, self.data, repr([b for b in self.bars]), repr([f for f in self.foos]))
+
+ mapper(Blub, blub, inherits=Bar, properties={
+ 'bars':relation(Bar, secondary=blub_bar, lazy=False),
+ 'foos':relation(Foo, secondary=blub_foo, lazy=False),
+ })
+
+ sess = create_session()
+ f1 = Foo("foo #1")
+ b1 = Bar("bar #1")
+ b2 = Bar("bar #2")
+ bl1 = Blub("blub #1")
+ for o in (f1, b1, b2, bl1):
+ sess.save(o)
+ bl1.foos.append(f1)
+ bl1.bars.append(b2)
+ sess.flush()
+ compare = repr(bl1)
+ blubid = bl1.id
+ sess.clear()
+
+ l = sess.query(Blub).select()
+ print l
+ self.assert_(repr(l[0]) == compare)
+ sess.clear()
+ x = sess.query(Blub).get_by(id=blubid)
+ print x
+ self.assert_(repr(x) == compare)
+
+
+if __name__ == "__main__":
+ testbase.main()
diff --git a/test/orm/poly_linked_list.py b/test/orm/inheritance/poly_linked_list.py
index 30cda4bb6..7297002f5 100644
--- a/test/orm/poly_linked_list.py
+++ b/test/orm/inheritance/poly_linked_list.py
@@ -1,7 +1,10 @@
import testbase
from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
-class PolymorphicCircularTest(testbase.ORMTest):
+
+class PolymorphicCircularTest(ORMTest):
keep_mappers = True
def define_tables(self, metadata):
global Table1, Table1B, Table2, Table3, Data
@@ -26,14 +29,15 @@ class PolymorphicCircularTest(testbase.ORMTest):
Column('data', String(30))
)
- join = polymorphic_union(
- {
- 'table3' : table1.join(table3),
- 'table2' : table1.join(table2),
- 'table1' : table1.select(table1.c.type.in_('table1', 'table1b')),
- }, None, 'pjoin')
-
- # still with us so far ?
+ #join = polymorphic_union(
+ # {
+ # 'table3' : table1.join(table3),
+ # 'table2' : table1.join(table2),
+ # 'table1' : table1.select(table1.c.type.in_('table1', 'table1b')),
+ # }, None, 'pjoin')
+
+ join = table1.outerjoin(table2).outerjoin(table3).alias('pjoin')
+ #join = None
class Table1(object):
def __init__(self, name, data=None):
@@ -59,10 +63,10 @@ class PolymorphicCircularTest(testbase.ORMTest):
return "%s(%d, %s)" % (self.__class__.__name__, self.id, repr(str(self.data)))
try:
- # this is how the mapping used to work. insure that this raises an error now
+ # this is how the mapping used to work. ensure that this raises an error now
table1_mapper = mapper(Table1, table1,
select_table=join,
- polymorphic_on=join.c.type,
+ polymorphic_on=table1.c.type,
polymorphic_identity='table1',
properties={
'next': relation(Table1,
@@ -83,8 +87,8 @@ class PolymorphicCircularTest(testbase.ORMTest):
# exception now. since eager loading would never work for that relation anyway, its better that the user
# gets an exception instead of it silently not eager loading.
table1_mapper = mapper(Table1, table1,
- select_table=join,
- polymorphic_on=join.c.type,
+ #select_table=join,
+ polymorphic_on=table1.c.type,
polymorphic_identity='table1',
properties={
'next': relation(Table1,
@@ -101,7 +105,10 @@ class PolymorphicCircularTest(testbase.ORMTest):
polymorphic_identity='table2')
table3_mapper = mapper(Table3, table3, inherits=table1_mapper, polymorphic_identity='table3')
-
+
+ table1_mapper.compile()
+ assert table1_mapper.primary_key == [table1.c.id], table1_mapper.primary_key
+
def testone(self):
self.do_testlist([Table1, Table2, Table1, Table2])
diff --git a/test/orm/polymorph.py b/test/orm/inheritance/polymorph.py
index 9d886cf3f..3eb2e032f 100644
--- a/test/orm/polymorph.py
+++ b/test/orm/inheritance/polymorph.py
@@ -1,8 +1,11 @@
+"""tests basic polymorphic mapper loading/saving, minimal relations"""
+
import testbase
-from sqlalchemy import *
import sets
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
-# tests basic polymorphic mapper loading/saving, minimal relations
class Person(object):
def __init__(self, **kwargs):
@@ -21,6 +24,10 @@ class Engineer(Person):
class Manager(Person):
def __repr__(self):
return "Manager %s, status %s, manager_name %s" % (self.get_name(), self.status, self.manager_name)
+class Boss(Manager):
+ def __repr__(self):
+ return "Boss %s, status %s, manager_name %s golf swing %s" % (self.get_name(), self.status, self.manager_name, self.golf_swing)
+
class Company(object):
def __init__(self, **kwargs):
for key, value in kwargs.iteritems():
@@ -28,9 +35,9 @@ class Company(object):
def __repr__(self):
return "Company %s" % self.name
-class PolymorphTest(testbase.ORMTest):
+class PolymorphTest(ORMTest):
def define_tables(self, metadata):
- global companies, people, engineers, managers
+ global companies, people, engineers, managers, boss
# a table to store companies
companies = Table('companies', metadata,
@@ -58,6 +65,11 @@ class PolymorphTest(testbase.ORMTest):
Column('manager_name', String(50))
)
+ boss = Table('boss', metadata,
+ Column('boss_id', Integer, ForeignKey('managers.person_id'), primary_key=True),
+ Column('golf_swing', String(30)),
+ )
+
metadata.create_all()
class CompileTest(PolymorphTest):
@@ -100,29 +112,6 @@ class CompileTest(PolymorphTest):
#person_mapper.compile()
class_mapper(Manager).compile()
- def testcompile3(self):
- """test that a mapper referencing an inheriting mapper in a self-referential relationship does
- not allow an eager load to be set up."""
- person_join = polymorphic_union( {
- 'engineer':people.join(engineers),
- 'manager':people.join(managers),
- 'person':people.select(people.c.type=='person'),
- }, None, 'pjoin')
-
- person_mapper = mapper(Person, people, select_table=person_join, polymorphic_on=person_join.c.type,
- polymorphic_identity='person',
- properties = dict(managers = relation(Manager, lazy=False))
- )
-
- mapper(Engineer, engineers, inherits=person_mapper, polymorphic_identity='engineer')
- mapper(Manager, managers, inherits=person_mapper, polymorphic_identity='manager')
-
- try:
- class_mapper(Manager).compile()
- assert False
- except exceptions.ArgumentError:
- assert True
-
class InsertOrderTest(PolymorphTest):
def test_insert_order(self):
"""test that classes of multiple types mix up mapper inserts
@@ -191,8 +180,11 @@ class RelationToSubclassTest(PolymorphTest):
sess.query(Company).get_by(company_id=c.company_id)
assert sets.Set([e.get_name() for e in c.managers]) == sets.Set(['pointy haired boss'])
assert c.managers[0].company is c
-
-def generate_round_trip_test(include_base=False, lazy_relation=True, redefine_colprop=False, use_literal_join=False):
+
+class RoundTripTest(PolymorphTest):
+ pass
+
+def generate_round_trip_test(include_base=False, lazy_relation=True, redefine_colprop=False, use_literal_join=False, polymorphic_fetch=None, use_outer_joins=False):
"""generates a round trip test.
include_base - whether or not to include the base 'person' type in the union.
@@ -200,117 +192,145 @@ def generate_round_trip_test(include_base=False, lazy_relation=True, redefine_co
redefine_colprop - if we redefine the 'name' column to be 'people_name' on the base Person class
use_literal_join - primary join condition is explicitly specified
"""
- class RoundTripTest(PolymorphTest):
- def test_roundtrip(self):
- # create a union that represents both types of joins.
- if include_base:
+ def test_roundtrip(self):
+ # create a union that represents both types of joins.
+ if not polymorphic_fetch == 'union':
+ person_join = None
+ manager_join = None
+ elif include_base:
+ if use_outer_joins:
+ person_join = people.outerjoin(engineers).outerjoin(managers).outerjoin(boss)
+ manager_join = people.join(managers).outerjoin(boss)
+ else:
person_join = polymorphic_union(
{
'engineer':people.join(engineers),
'manager':people.join(managers),
'person':people.select(people.c.type=='person'),
}, None, 'pjoin')
+
+ manager_join = people.join(managers).outerjoin(boss)
+ else:
+ if use_outer_joins:
+ person_join = people.outerjoin(engineers).outerjoin(managers).outerjoin(boss)
+ manager_join = people.join(managers).outerjoin(boss)
else:
person_join = polymorphic_union(
{
'engineer':people.join(engineers),
'manager':people.join(managers),
}, None, 'pjoin')
+ manager_join = people.join(managers).outerjoin(boss)
- if redefine_colprop:
- person_mapper = mapper(Person, people, select_table=person_join, polymorphic_on=person_join.c.type, polymorphic_identity='person', properties= {'person_name':people.c.name})
- else:
- person_mapper = mapper(Person, people, select_table=person_join, polymorphic_on=person_join.c.type, polymorphic_identity='person')
-
- mapper(Engineer, engineers, inherits=person_mapper, polymorphic_identity='engineer')
- mapper(Manager, managers, inherits=person_mapper, polymorphic_identity='manager')
-
- if use_literal_join:
- mapper(Company, companies, properties={
- 'employees': relation(Person, lazy=lazy_relation, primaryjoin=people.c.company_id==companies.c.company_id, private=True,
- backref="company"
- )
- })
- else:
- mapper(Company, companies, properties={
- 'employees': relation(Person, lazy=lazy_relation, private=True,
- backref="company"
- )
- })
-
- if redefine_colprop:
- person_attribute_name = 'person_name'
- else:
- person_attribute_name = 'name'
+ if redefine_colprop:
+ person_mapper = mapper(Person, people, select_table=person_join, polymorphic_fetch=polymorphic_fetch, polymorphic_on=people.c.type, polymorphic_identity='person', properties= {'person_name':people.c.name})
+ else:
+ person_mapper = mapper(Person, people, select_table=person_join, polymorphic_fetch=polymorphic_fetch, polymorphic_on=people.c.type, polymorphic_identity='person')
- session = create_session()
- c = Company(name='company1')
- c.employees.append(Manager(status='AAB', manager_name='manager1', **{person_attribute_name:'pointy haired boss'}))
- c.employees.append(Engineer(status='BBA', engineer_name='engineer1', primary_language='java', **{person_attribute_name:'dilbert'}))
- if include_base:
- c.employees.append(Person(status='HHH', **{person_attribute_name:'joesmith'}))
- c.employees.append(Engineer(status='CGG', engineer_name='engineer2', primary_language='python', **{person_attribute_name:'wally'}))
- c.employees.append(Manager(status='ABA', manager_name='manager2', **{person_attribute_name:'jsmith'}))
- session.save(c)
- print session.new
- session.flush()
- session.clear()
- id = c.company_id
- c = session.query(Company).get(id)
- for e in c.employees:
- print e, e._instance_key, e.company
- if include_base:
- assert sets.Set([e.get_name() for e in c.employees]) == sets.Set(['pointy haired boss', 'dilbert', 'joesmith', 'wally', 'jsmith'])
- else:
- assert sets.Set([e.get_name() for e in c.employees]) == sets.Set(['pointy haired boss', 'dilbert', 'wally', 'jsmith'])
- print "\n"
+ mapper(Engineer, engineers, inherits=person_mapper, polymorphic_identity='engineer')
+ mapper(Manager, managers, inherits=person_mapper, select_table=manager_join, polymorphic_identity='manager')
+ mapper(Boss, boss, inherits=Manager, polymorphic_identity='boss')
- # test selecting from the query, using the base mapped table (people) as the selection criterion.
- # in the case of the polymorphic Person query, the "people" selectable should be adapted to be "person_join"
- dilbert = session.query(Person).selectfirst(people.c.name=='dilbert')
- dilbert2 = session.query(Engineer).selectfirst(people.c.name=='dilbert')
- assert dilbert is dilbert2
-
- # test selecting from the query, joining against an alias of the base "people" table. test that
- # the "palias" alias does *not* get sucked up into the "person_join" conversion.
- palias = people.alias("palias")
- session.query(Person).selectfirst((palias.c.name=='dilbert') & (palias.c.person_id==people.c.person_id))
- dilbert2 = session.query(Engineer).selectfirst((palias.c.name=='dilbert') & (palias.c.person_id==people.c.person_id))
- assert dilbert is dilbert2
-
- session.query(Person).selectfirst((engineers.c.engineer_name=="engineer1") & (engineers.c.person_id==people.c.person_id))
- dilbert2 = session.query(Engineer).selectfirst(engineers.c.engineer_name=="engineer1")
- assert dilbert is dilbert2
+ if use_literal_join:
+ mapper(Company, companies, properties={
+ 'employees': relation(Person, lazy=lazy_relation, primaryjoin=people.c.company_id==companies.c.company_id, private=True,
+ backref="company"
+ )
+ })
+ else:
+ mapper(Company, companies, properties={
+ 'employees': relation(Person, lazy=lazy_relation, private=True,
+ backref="company"
+ )
+ })
+ if redefine_colprop:
+ person_attribute_name = 'person_name'
+ else:
+ person_attribute_name = 'name'
+
+ session = create_session()
+ c = Company(name='company1')
+ c.employees.append(Manager(status='AAB', manager_name='manager1', **{person_attribute_name:'pointy haired boss'}))
+ c.employees.append(Engineer(status='BBA', engineer_name='engineer1', primary_language='java', **{person_attribute_name:'dilbert'}))
+ if include_base:
+ c.employees.append(Person(status='HHH', **{person_attribute_name:'joesmith'}))
+ c.employees.append(Engineer(status='CGG', engineer_name='engineer2', primary_language='python', **{person_attribute_name:'wally'}))
+ c.employees.append(Manager(status='ABA', manager_name='manager2', **{person_attribute_name:'jsmith'}))
+ session.save(c)
+ print session.new
+ session.flush()
+ session.clear()
+ id = c.company_id
+ c = session.query(Company).get(id)
+ for e in c.employees:
+ print e, e._instance_key, e.company
+ if include_base:
+ assert sets.Set([(e.get_name(), getattr(e, 'status', None)) for e in c.employees]) == sets.Set([('pointy haired boss', 'AAB'), ('dilbert', 'BBA'), ('joesmith', None), ('wally', 'CGG'), ('jsmith', 'ABA')])
+ else:
+ assert sets.Set([(e.get_name(), e.status) for e in c.employees]) == sets.Set([('pointy haired boss', 'AAB'), ('dilbert', 'BBA'), ('wally', 'CGG'), ('jsmith', 'ABA')])
+ print "\n"
+
+ # test selecting from the query, using the base mapped table (people) as the selection criterion.
+ # in the case of the polymorphic Person query, the "people" selectable should be adapted to be "person_join"
+ dilbert = session.query(Person).filter(getattr(Person, person_attribute_name)=='dilbert').first()
+ dilbert2 = session.query(Engineer).filter(getattr(Person, person_attribute_name)=='dilbert').first()
+ assert dilbert is dilbert2
+
+ # test selecting from the query, joining against an alias of the base "people" table. test that
+ # the "palias" alias does *not* get sucked up into the "person_join" conversion.
+ palias = people.alias("palias")
+ session.query(Person).filter((palias.c.name=='dilbert') & (palias.c.person_id==Person.person_id)).first()
+ dilbert2 = session.query(Engineer).filter((palias.c.name=='dilbert') & (palias.c.person_id==Person.person_id)).first()
+ assert dilbert is dilbert2
+
+ session.query(Person).filter((Engineer.engineer_name=="engineer1") & (Engineer.person_id==people.c.person_id)).first()
- dilbert.engineer_name = 'hes dibert!'
-
- session.flush()
- session.clear()
+ dilbert2 = session.query(Engineer).filter(Engineer.engineer_name=="engineer1")[0]
+ assert dilbert is dilbert2
+
+ dilbert.engineer_name = 'hes dibert!'
- c = session.query(Company).get(id)
- for e in c.employees:
- print e, e._instance_key
+ session.flush()
+ session.clear()
- session.delete(c)
- session.flush()
+ # save/load some managers/bosses
+ b = Boss(status='BBB', manager_name='boss', golf_swing='fore', **{person_attribute_name:'daboss'})
+ session.save(b)
+ session.flush()
+ session.clear()
+ c = session.query(Manager).all()
+ assert sets.Set([repr(x) for x in c]) == sets.Set(["Manager pointy haired boss, status AAB, manager_name manager1", "Manager jsmith, status ABA, manager_name manager2", "Boss daboss, status BBB, manager_name boss golf swing fore"]), repr([repr(x) for x in c])
+
+ c = session.query(Company).get(id)
+ for e in c.employees:
+ print e, e._instance_key
- RoundTripTest.__name__ = "Test%s%s%s%s" % (
- (lazy_relation and "Lazy" or "Eager"),
- (include_base and "Inclbase" or ""),
- (redefine_colprop and "Redefcol" or ""),
- (use_literal_join and "Litjoin" or "")
+ session.delete(c)
+ session.flush()
+
+
+ test_roundtrip.__name__ = "test_%s%s%s%s%s" % (
+ (lazy_relation and "lazy" or "eager"),
+ (include_base and "_inclbase" or ""),
+ (redefine_colprop and "_redefcol" or ""),
+ (polymorphic_fetch != 'union' and '_' + polymorphic_fetch or (use_literal_join and "_litjoin" or "")),
+ (use_outer_joins and '_outerjoins' or '')
)
- return RoundTripTest
+ setattr(RoundTripTest, test_roundtrip.__name__, test_roundtrip)
for include_base in [True, False]:
for lazy_relation in [True, False]:
for redefine_colprop in [True, False]:
for use_literal_join in [True, False]:
- testclass = generate_round_trip_test(include_base, lazy_relation, redefine_colprop, use_literal_join)
- exec("%s = testclass" % testclass.__name__)
-
+ for polymorphic_fetch in ['union', 'select', 'deferred']:
+ if polymorphic_fetch == 'union':
+ for use_outer_joins in [True, False]:
+ generate_round_trip_test(include_base, lazy_relation, redefine_colprop, use_literal_join, polymorphic_fetch, use_outer_joins)
+ else:
+ generate_round_trip_test(include_base, lazy_relation, redefine_colprop, use_literal_join, polymorphic_fetch, False)
+
if __name__ == "__main__":
testbase.main()
diff --git a/test/orm/inheritance5.py b/test/orm/inheritance/polymorph2.py
index cf7224fa4..a2f9c4a5f 100644
--- a/test/orm/inheritance5.py
+++ b/test/orm/inheritance/polymorph2.py
@@ -1,5 +1,8 @@
-from sqlalchemy import *
import testbase
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
+
class AttrSettable(object):
def __init__(self, **kwargs):
@@ -8,7 +11,7 @@ class AttrSettable(object):
return self.__class__.__name__ + "(%s)" % (hex(id(self)))
-class RelationTest1(testbase.ORMTest):
+class RelationTest1(ORMTest):
"""test self-referential relationships on polymorphic mappers"""
def define_tables(self, metadata):
global people, managers
@@ -88,7 +91,7 @@ class RelationTest1(testbase.ORMTest):
print p, m, m.employee
assert m.employee is p
-class RelationTest2(testbase.ORMTest):
+class RelationTest2(ORMTest):
"""test self-referential relationships on polymorphic mappers"""
def define_tables(self, metadata):
global people, managers, data
@@ -116,6 +119,10 @@ class RelationTest2(testbase.ORMTest):
self.do_test("join1", True)
def testrelationonsubclass_j2_data(self):
self.do_test("join2", True)
+ def testrelationonsubclass_j3_nodata(self):
+ self.do_test("join3", False)
+ def testrelationonsubclass_j3_data(self):
+ self.do_test("join3", True)
def do_test(self, jointype="join1", usedata=False):
class Person(AttrSettable):
@@ -128,19 +135,24 @@ class RelationTest2(testbase.ORMTest):
'person':people.select(people.c.type=='person'),
'manager':join(people, managers, people.c.person_id==managers.c.person_id)
}, None)
+ polymorphic_on=poly_union.c.type
elif jointype == "join2":
poly_union = polymorphic_union({
'person':people.select(people.c.type=='person'),
'manager':managers.join(people, people.c.person_id==managers.c.person_id)
}, None)
-
+ polymorphic_on=poly_union.c.type
+ elif jointype == "join3":
+ poly_union = None
+ polymorphic_on = people.c.type
+
if usedata:
class Data(object):
def __init__(self, data):
self.data = data
mapper(Data, data)
- mapper(Person, people, select_table=poly_union, polymorphic_identity='person', polymorphic_on=poly_union.c.type)
+ mapper(Person, people, select_table=poly_union, polymorphic_identity='person', polymorphic_on=polymorphic_on)
if usedata:
mapper(Manager, managers, inherits=Person, inherit_condition=people.c.person_id==managers.c.person_id, polymorphic_identity='manager',
@@ -174,7 +186,7 @@ class RelationTest2(testbase.ORMTest):
if usedata:
assert m.data.data == 'ms data'
-class RelationTest3(testbase.ORMTest):
+class RelationTest3(ORMTest):
"""test self-referential relationships on polymorphic mappers"""
def define_tables(self, metadata):
global people, managers, data
@@ -194,16 +206,8 @@ class RelationTest3(testbase.ORMTest):
Column('data', String(30))
)
- def testrelationonbaseclass_j1_nodata(self):
- self.do_test("join1", False)
- def testrelationonbaseclass_j2_nodata(self):
- self.do_test("join2", False)
- def testrelationonbaseclass_j1_data(self):
- self.do_test("join1", True)
- def testrelationonbaseclass_j2_data(self):
- self.do_test("join2", True)
-
- def do_test(self, jointype="join1", usedata=False):
+def generate_test(jointype="join1", usedata=False):
+ def do_test(self):
class Person(AttrSettable):
pass
class Manager(Person):
@@ -224,10 +228,14 @@ class RelationTest3(testbase.ORMTest):
'manager':join(people, managers, people.c.person_id==managers.c.person_id),
'person':people.select(people.c.type=='person')
}, None)
-
+ elif jointype == 'join3':
+ poly_union = people.outerjoin(managers)
+ elif jointype == "join4":
+ poly_union=None
+
if usedata:
mapper(Data, data)
-
+
mapper(Manager, managers, inherits=Person, inherit_condition=people.c.person_id==managers.c.person_id, polymorphic_identity='manager')
if usedata:
mapper(Person, people, select_table=poly_union, polymorphic_identity='person', polymorphic_on=people.c.type,
@@ -258,7 +266,7 @@ class RelationTest3(testbase.ORMTest):
sess.save(m)
sess.save(p)
sess.flush()
-
+
sess.clear()
p = sess.query(Person).get(p.person_id)
p2 = sess.query(Person).get(p2.person_id)
@@ -271,9 +279,17 @@ class RelationTest3(testbase.ORMTest):
if usedata:
assert p.data.data == 'ps data'
assert m.data.data == 'ms data'
+
+ do_test.__name__ = 'test_relationonbaseclass_%s_%s' % (jointype, data and "nodata" or "data")
+ return do_test
+for jointype in ["join1", "join2", "join3", "join4"]:
+ for data in (True, False):
+ func = generate_test(jointype, data)
+ setattr(RelationTest3, func.__name__, func)
+
-class RelationTest4(testbase.ORMTest):
+class RelationTest4(ORMTest):
def define_tables(self, metadata):
global people, engineers, managers, cars
people = Table('people', metadata,
@@ -329,6 +345,9 @@ class RelationTest4(testbase.ORMTest):
manager_mapper = mapper(Manager, managers, inherits=person_mapper, polymorphic_identity='manager')
car_mapper = mapper(Car, cars, properties= {'employee':relation(person_mapper)})
+ print class_mapper(Person).primary_key
+ print person_mapper.get_select_mapper().primary_key
+
# so the primaryjoin is "people.c.person_id==cars.c.owner". the "lazy" clause will be
# "people.c.person_id=?". the employee_join is two selects union'ed together, one of which
# will contain employee.c.person_id the other contains manager.c.person_id. people.c.person_id is not explicitly in
@@ -350,8 +369,8 @@ class RelationTest4(testbase.ORMTest):
session.flush()
- engineer4 = session.query(Engineer).selectfirst_by(name="E4")
- manager3 = session.query(Manager).selectfirst_by(name="M3")
+ engineer4 = session.query(Engineer).filter(Engineer.name=="E4").first()
+ manager3 = session.query(Manager).filter(Manager.name=="M3").first()
car1 = Car(employee=engineer4)
session.save(car1)
@@ -361,27 +380,32 @@ class RelationTest4(testbase.ORMTest):
session.clear()
+ print "----------------------------"
car1 = session.query(Car).get(car1.car_id)
+ print "----------------------------"
usingGet = session.query(person_mapper).get(car1.owner)
+ print "----------------------------"
usingProperty = car1.employee
+ print "----------------------------"
# All print should output the same person (engineer E4)
assert str(engineer4) == "Engineer E4, status X"
+ print str(usingGet)
assert str(usingGet) == "Engineer E4, status X"
assert str(usingProperty) == "Engineer E4, status X"
session.clear()
-
+ print "-----------------------------------------------------------------"
# and now for the lightning round, eager !
car1 = session.query(Car).options(eagerload('employee')).get(car1.car_id)
assert str(car1.employee) == "Engineer E4, status X"
session.clear()
s = session.query(Car)
- c = s.join("employee").select(employee_join.c.name=="E4")[0]
+ c = s.join("employee").filter(Person.name=="E4")[0]
assert c.car_id==car1.car_id
-class RelationTest5(testbase.ORMTest):
+class RelationTest5(ORMTest):
def define_tables(self, metadata):
global people, engineers, managers, cars
people = Table('people', metadata,
@@ -441,7 +465,7 @@ class RelationTest5(testbase.ORMTest):
assert carlist[0].manager is None
assert carlist[1].manager.person_id == car2.manager.person_id
-class RelationTest6(testbase.ORMTest):
+class RelationTest6(ORMTest):
"""test self-referential relationships on a single joined-table inheritance mapper"""
def define_tables(self, metadata):
global people, managers, data
@@ -484,7 +508,7 @@ class RelationTest6(testbase.ORMTest):
m2 = sess.query(Manager).get(m2.person_id)
assert m.colleague is m2
-class RelationTest7(testbase.ORMTest):
+class RelationTest7(ORMTest):
def define_tables(self, metadata):
global people, engineers, managers, cars, offroad_cars
cars = Table('cars', metadata,
@@ -583,7 +607,7 @@ class RelationTest7(testbase.ORMTest):
for p in r:
assert p.car_id == p.car.car_id
-class GenerativeTest(testbase.AssertMixin):
+class GenerativeTest(AssertMixin):
def setUpAll(self):
# cars---owned by--- people (abstract) --- has a --- status
# | ^ ^ |
@@ -698,7 +722,7 @@ class GenerativeTest(testbase.AssertMixin):
# test these twice because theres caching involved, as well previous issues that modified the polymorphic union
for x in range(0, 2):
- r = session.query(Person).filter_by(people.c.name.like('%2')).join('status').filter_by(name="active")
+ r = session.query(Person).filter(people.c.name.like('%2')).join('status').filter_by(name="active")
assert str(list(r)) == "[Manager M2, category YYYYYYYYY, status Status active, Engineer E2, field X, status Status active]"
r = session.query(Engineer).join('status').filter(people.c.name.in_('E2', 'E3', 'E4', 'M4', 'M2', 'M1') & (status.c.name=="active"))
assert str(list(r)) == "[Engineer E2, field X, status Status active, Engineer E3, field X, status Status active]"
@@ -709,22 +733,22 @@ class GenerativeTest(testbase.AssertMixin):
r = session.query(Person).filter(exists([Car.c.owner], Car.c.owner==employee_join.c.person_id))
assert str(list(r)) == "[Engineer E4, field X, status Status dead]"
-class MultiLevelTest(testbase.ORMTest):
+class MultiLevelTest(ORMTest):
def define_tables(self, metadata):
global table_Employee, table_Engineer, table_Manager
table_Employee = Table( 'Employee', metadata,
- Column( 'name', type= String(100), ),
- Column( 'id', primary_key= True, type= Integer, ),
- Column( 'atype', type= String(100), ),
+ Column( 'name', type_= String(100), ),
+ Column( 'id', primary_key= True, type_= Integer, ),
+ Column( 'atype', type_= String(100), ),
)
table_Engineer = Table( 'Engineer', metadata,
- Column( 'machine', type= String(100), ),
+ Column( 'machine', type_= String(100), ),
Column( 'id', Integer, ForeignKey( 'Employee.id', ), primary_key= True, ),
)
table_Manager = Table( 'Manager', metadata,
- Column( 'duties', type= String(100), ),
+ Column( 'duties', type_= String(100), ),
Column( 'id', Integer, ForeignKey( 'Engineer.id', ), primary_key= True, ),
)
def test_threelevels(self):
@@ -786,7 +810,7 @@ class MultiLevelTest(testbase.ORMTest):
assert set(session.query( Engineer).select()) == set([b,c])
assert session.query( Manager).select() == [c]
-class ManyToManyPolyTest(testbase.ORMTest):
+class ManyToManyPolyTest(ORMTest):
def define_tables(self, metadata):
global base_item_table, item_table, base_item_collection_table, collection_table
base_item_table = Table(
@@ -836,7 +860,7 @@ class ManyToManyPolyTest(testbase.ORMTest):
class_mapper(BaseItem)
-class CustomPKTest(testbase.ORMTest):
+class CustomPKTest(ORMTest):
def define_tables(self, metadata):
global t1, t2
t1 = Table('t1', metadata,
@@ -847,7 +871,7 @@ class CustomPKTest(testbase.ORMTest):
t2 = Table('t2', metadata,
Column('t2id', Integer, ForeignKey('t1.id'), primary_key=True),
Column('t2data', String(30)))
-
+
def test_custompk(self):
"""test that the primary_key attribute is propigated to the polymorphic mapper"""
@@ -885,6 +909,48 @@ class CustomPKTest(testbase.ORMTest):
ot1 = sess.query(T1).get(ot1.id)
ot1.data = 'hi'
sess.flush()
+
+ def test_pk_collapses(self):
+ """test that a composite primary key attribute formed by a join is "collapsed" into its
+ minimal columns"""
+
+ class T1(object):pass
+ class T2(T1):pass
+
+ # create a polymorphic union with the select against the base table first.
+ # with the join being second, the alias of the union will
+ # pick up two "primary key" columns. technically the alias should have a
+ # 2-col pk in any case but the leading select has a NULL for the "t2id" column
+ d = util.OrderedDict()
+ d['t1'] = t1.select(t1.c.type=='t1')
+ d['t2'] = t1.join(t2)
+ pjoin = polymorphic_union(d, None, 'pjoin')
+
+ #print pjoin.original.primary_key
+ #print pjoin.primary_key
+ assert len(pjoin.primary_key) == 2
+
+ mapper(T1, t1, polymorphic_on=t1.c.type, polymorphic_identity='t1', select_table=pjoin)
+ mapper(T2, t2, inherits=T1, polymorphic_identity='t2')
+ assert len(class_mapper(T1).primary_key) == 1
+ assert len(class_mapper(T1).get_select_mapper().compile().primary_key) == 1
+
+ print [str(c) for c in class_mapper(T1).primary_key]
+ ot1 = T1()
+ ot2 = T2()
+ sess = create_session()
+ sess.save(ot1)
+ sess.save(ot2)
+ sess.flush()
+ sess.clear()
+
+ # query using get(), using only one value. this requires the select_table mapper
+ # has the same single-col primary key.
+ assert sess.query(T1).get(ot1.id).id == ot1.id
+
+ ot1 = sess.query(T1).get(ot1.id)
+ ot1.data = 'hi'
+ sess.flush()
if __name__ == "__main__":
testbase.main()
diff --git a/test/orm/inheritance2.py b/test/orm/inheritance/productspec.py
index 906526456..2459cd36e 100644
--- a/test/orm/inheritance2.py
+++ b/test/orm/inheritance/productspec.py
@@ -1,8 +1,11 @@
import testbase
-from sqlalchemy import *
from datetime import datetime
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
+
-class InheritTest(testbase.ORMTest):
+class InheritTest(ORMTest):
"""tests some various inheritance round trips involving a particular set of polymorphic inheritance relationships"""
def define_tables(self, metadata):
global products_table, specification_table, documents_table
diff --git a/test/orm/single.py b/test/orm/inheritance/single.py
index 31a90da21..68fe821af 100644
--- a/test/orm/single.py
+++ b/test/orm/inheritance/single.py
@@ -1,7 +1,10 @@
-from sqlalchemy import *
import testbase
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
+
-class SingleInheritanceTest(testbase.AssertMixin):
+class SingleInheritanceTest(AssertMixin):
def setUpAll(self):
metadata = MetaData(testbase.db)
global employees_table
diff --git a/test/orm/lazy_relations.py b/test/orm/lazy_relations.py
index e9d77e09c..6684c6288 100644
--- a/test/orm/lazy_relations.py
+++ b/test/orm/lazy_relations.py
@@ -1,9 +1,9 @@
"""basic tests of lazy loaded attributes"""
+import testbase
from sqlalchemy import *
from sqlalchemy.orm import *
-import testbase
-
+from testlib import *
from fixtures import *
from query import QueryTest
diff --git a/test/orm/lazytest1.py b/test/orm/lazytest1.py
index 2cabac3a2..b5296120b 100644
--- a/test/orm/lazytest1.py
+++ b/test/orm/lazytest1.py
@@ -1,8 +1,7 @@
-from testbase import PersistTest, AssertMixin
import testbase
-import unittest, sys, os
from sqlalchemy import *
-import datetime
+from sqlalchemy.orm import *
+from testlib import *
class LazyTest(AssertMixin):
def setUpAll(self):
diff --git a/test/orm/manytomany.py b/test/orm/manytomany.py
index f6e9197a2..8b310f86c 100644
--- a/test/orm/manytomany.py
+++ b/test/orm/manytomany.py
@@ -1,6 +1,8 @@
import testbase
from sqlalchemy import *
-import string
+from sqlalchemy.orm import *
+from testlib import *
+
class Place(object):
'''represents a place'''
@@ -25,7 +27,7 @@ class Transition(object):
def __repr__(self):
return object.__repr__(self)+ " " + repr(self.inputs) + " " + repr(self.outputs)
-class M2MTest(testbase.ORMTest):
+class M2MTest(ORMTest):
def define_tables(self, metadata):
global place
place = Table('place', metadata,
@@ -110,7 +112,7 @@ class M2MTest(testbase.ORMTest):
for p in l:
pp = p.places
- self.echo("Place " + str(p) +" places " + repr(pp))
+ print "Place " + str(p) +" places " + repr(pp)
[sess.delete(p) for p in p1,p2,p3,p4,p5,p6,p7]
sess.flush()
@@ -176,7 +178,7 @@ class M2MTest(testbase.ORMTest):
self.assert_result([t1], Transition, {'outputs': (Place, [{'name':'place3'}, {'name':'place1'}])})
self.assert_result([p2], Place, {'inputs': (Transition, [{'name':'transition1'},{'name':'transition2'}])})
-class M2MTest2(testbase.ORMTest):
+class M2MTest2(ORMTest):
def define_tables(self, metadata):
global studentTbl
studentTbl = Table('student', metadata, Column('name', String(20), primary_key=True))
@@ -243,7 +245,7 @@ class M2MTest2(testbase.ORMTest):
sess.flush()
assert enrolTbl.count().scalar() == 0
-class M2MTest3(testbase.ORMTest):
+class M2MTest3(ORMTest):
def define_tables(self, metadata):
global c, c2a1, c2a2, b, a
c = Table('c', metadata,
@@ -277,15 +279,15 @@ class M2MTest3(testbase.ORMTest):
class A(object):pass
class B(object):pass
- assign_mapper(B, b)
+ mapper(B, b)
- assign_mapper(A, a,
+ mapper(A, a,
properties = {
'tbs' : relation(B, primaryjoin=and_(b.c.a1==a.c.a1, b.c.b2 == True), lazy=False),
}
)
- assign_mapper(C, c,
+ mapper(C, c,
properties = {
'a1s' : relation(A, secondary=c2a1, lazy=False),
'a2s' : relation(A, secondary=c2a2, lazy=False)
diff --git a/test/orm/mapper.py b/test/orm/mapper.py
index c0297e514..b72a10516 100644
--- a/test/orm/mapper.py
+++ b/test/orm/mapper.py
@@ -1,13 +1,14 @@
-from testbase import PersistTest, AssertMixin
+"""tests general mapper operations with an emphasis on selecting/loading"""
+
import testbase
-import unittest, sys, os
from sqlalchemy import *
+from sqlalchemy.orm import *
import sqlalchemy.exceptions as exceptions
-from sqlalchemy.ext.sessioncontext import SessionContext
-from tables import *
-import tables
+from sqlalchemy.ext.sessioncontext import SessionContext, SessionContextExt
+from testlib import *
+from testlib.tables import *
+import testlib.tables as tables
-"""tests general mapper operations with an emphasis on selecting/loading"""
class MapperSuperTest(AssertMixin):
def setUpAll(self):
@@ -21,36 +22,6 @@ class MapperSuperTest(AssertMixin):
pass
class MapperTest(MapperSuperTest):
- # TODO: MapperTest has grown much larger than it originally was and needs
- # to be broken up among various functions, including querying, session operations,
- # mapper configurational issues
- def testget(self):
- s = create_session()
- mapper(User, users)
- self.assert_(s.get(User, 19) is None)
- u = s.get(User, 7)
- u2 = s.get(User, 7)
- self.assert_(u is u2)
- s.clear()
- u2 = s.get(User, 7)
- self.assert_(u is not u2)
-
- def testunicodeget(self):
- """test that Query.get properly sets up the type for the bind parameter. using unicode would normally fail
- on postgres, mysql and oracle unless it is converted to an encoded string"""
- metadata = MetaData(db)
- table = Table('foo', metadata,
- Column('id', Unicode(10), primary_key=True),
- Column('data', Unicode(40)))
- try:
- table.create()
- class LocalFoo(object):pass
- mapper(LocalFoo, table)
- crit = 'petit voix m\xe2\x80\x99a '.decode('utf-8')
- print repr(crit)
- create_session().query(LocalFoo).get(crit)
- finally:
- table.drop()
def testpropconflict(self):
"""test that a backref created against an existing mapper with a property name
@@ -76,25 +47,20 @@ class MapperTest(MapperSuperTest):
assert str(e) == "Invalid cascade option 'fake'"
def testcolumnprefix(self):
- mapper(User, users, column_prefix='_', properties={
- 'user_name':synonym('_user_name')
- })
+ mapper(User, users, column_prefix='_')
s = create_session()
u = s.get(User, 7)
assert u._user_name=='jack'
assert u._user_id ==7
assert not hasattr(u, 'user_name')
- u2 = s.query(User).filter_by(user_name='jack').one()
- assert u is u2
def testrefresh(self):
- mapper(User, users, properties={'addresses':relation(mapper(Address, addresses))})
+ mapper(User, users, properties={'addresses':relation(mapper(Address, addresses), backref='user')})
s = create_session()
u = s.get(User, 7)
u.user_name = 'foo'
a = Address()
- import sqlalchemy.orm.session
- assert sqlalchemy.orm.session.object_session(a) is None
+ assert object_session(a) is None
u.addresses.append(a)
self.assert_(a in u.addresses)
@@ -120,6 +86,7 @@ class MapperTest(MapperSuperTest):
# get the attribute, it refreshes
self.assert_(u.user_name == 'jack')
self.assert_(a not in u.addresses)
+
def testexpirecascade(self):
mapper(User, users, properties={'addresses':relation(mapper(Address, addresses), cascade="all, refresh-expire")})
@@ -171,25 +138,28 @@ class MapperTest(MapperSuperTest):
def __init__(self):
raise ex
mapper(Foo, users)
-
+
try:
Foo()
assert False
except Exception, e:
assert e is ex
+ clear_mappers()
+ mapper(Foo, users, extension=SessionContextExt(SessionContext()))
def bad_expunge(foo):
raise Exception("this exception should be stated as a warning")
import warnings
warnings.filterwarnings("always", r".*this exception should be stated as a warning")
+
sess.expunge = bad_expunge
try:
Foo(_sa_session=sess)
assert False
except Exception, e:
assert e is ex
-
+
def testrefresh_lazy(self):
"""test that when a lazy loader is set as a trigger on an object's attribute (at the attribute level, not the class level), a refresh() operation doesnt fire the lazy loader or create any problems"""
s = create_session()
@@ -198,7 +168,7 @@ class MapperTest(MapperSuperTest):
u = q2.selectfirst(users.c.user_id==8)
def go():
s.refresh(u)
- self.assert_sql_count(db, go, 1)
+ self.assert_sql_count(testbase.db, go, 1)
def testexpire(self):
"""test the expire function"""
@@ -233,12 +203,19 @@ class MapperTest(MapperSuperTest):
def testrefresh2(self):
"""test a hang condition that was occuring on expire/refresh"""
+
s = create_session()
- mapper(Address, addresses)
-
- mapper(User, users, properties = dict(addresses=relation(Address,private=True,lazy=False)) )
+ m1 = mapper(Address, addresses)
+ m2 = mapper(User, users, properties = dict(addresses=relation(Address,private=True,lazy=False)) )
+ assert m1._Mapper__is_compiled is False
+ assert m2._Mapper__is_compiled is False
+
+# compile_mappers()
+ print "NEW USER"
u=User()
+ print "NEW USER DONE"
+ assert m2._Mapper__is_compiled is True
u.user_name='Justin'
a = Address()
a.address_id=17 # to work around the hardcoded IDs in this test suite....
@@ -259,15 +236,8 @@ class MapperTest(MapperSuperTest):
m = mapper(User, users, properties = {
'addresses' : relation(mapper(Address, addresses))
}).compile()
- self.assert_(User.addresses.property is m.props['addresses'])
+ self.assert_(User.addresses.property is m.get_property('addresses'))
- def testquery(self):
- """test a basic Query.select() operation."""
- mapper(User, users)
- l = create_session().query(User).select()
- self.assert_result(l, User, *user_result)
- l = create_session().query(User).select(users.c.user_name.endswith('ed'))
- self.assert_result(l, User, *user_result[1:3])
def testrecursiveselectby(self):
"""test that no endless loop occurs when traversing for select_by"""
@@ -302,151 +272,6 @@ class MapperTest(MapperSuperTest):
l = q.select()
self.assert_result(l, User, *result)
- def testwithparent(self):
- """test the with_parent()) method and one-to-many relationships"""
-
- m = mapper(User, users, properties={
- 'user_name_syn':synonym('user_name'),
- 'orders':relation(mapper(Order, orders, properties={
- 'items':relation(mapper(Item, orderitems)),
- 'items_syn':synonym('items')
- })),
- 'orders_syn':synonym('orders')
- })
-
- sess = create_session()
- q = sess.query(m)
- u1 = q.get_by(user_name='jack')
-
- # test auto-lookup of property
- o = sess.query(Order).with_parent(u1).list()
- self.assert_result(o, Order, *user_all_result[0]['orders'][1])
-
- # test with explicit property
- o = sess.query(Order).with_parent(u1, property='orders').list()
- self.assert_result(o, Order, *user_all_result[0]['orders'][1])
-
- # test static method
- o = Query.query_from_parent(u1, property='orders', session=sess).list()
- self.assert_result(o, Order, *user_all_result[0]['orders'][1])
-
- # test generative criterion
- o = sess.query(Order).with_parent(u1).select_by(orders.c.order_id>2)
- self.assert_result(o, Order, *user_all_result[0]['orders'][1][1:])
-
- try:
- q = sess.query(Item).with_parent(u1)
- assert False
- except exceptions.InvalidRequestError, e:
- assert str(e) == "Could not locate a property which relates instances of class 'Item' to instances of class 'User'"
-
-
- for nameprop, orderprop in (
- ('user_name', 'orders'),
- ('user_name_syn', 'orders'),
- ('user_name', 'orders_syn'),
- ('user_name_syn', 'orders_syn'),
- ):
- sess = create_session()
- q = sess.query(User)
-
- u1 = q.filter_by(**{nameprop:'jack'}).one()
-
- o = sess.query(Order).with_parent(u1, property=orderprop).list()
- self.assert_result(o, Order, *user_all_result[0]['orders'][1])
-
- def testwithparentm2m(self):
- """test the with_parent() method and many-to-many relationships"""
-
- m = mapper(Item, orderitems, properties = {
- 'keywords' : relation(mapper(Keyword, keywords), itemkeywords)
- })
- sess = create_session()
- i1 = sess.query(Item).get_by(item_id=2)
- k = sess.query(Keyword).with_parent(i1)
- self.assert_result(k, Keyword, *item_keyword_result[1]['keywords'][1])
-
-
- def test_join(self):
- """test functions derived from Query's _join_to function."""
-
- m = mapper(User, users, properties={
- 'orders':relation(mapper(Order, orders, properties={
- 'items':relation(mapper(Item, orderitems)),
- 'items_syn':synonym('items')
- })),
-
- 'orders_syn':synonym('orders'),
- })
-
- sess = create_session()
- q = sess.query(m)
-
- for j in (
- ['orders', 'items'],
- ['orders', 'items_syn'],
- ['orders_syn', 'items'],
- ['orders_syn', 'items_syn'],
- ):
- for q in (
- q.filter(orderitems.c.item_name=='item 4').join(j),
- q.filter(orderitems.c.item_name=='item 4').join(j[-1]),
- q.filter(orderitems.c.item_name=='item 4').filter(q.join_via(j)),
- q.filter(orderitems.c.item_name=='item 4').filter(q.join_to(j[-1])),
- ):
- l = q.all()
- self.assert_result(l, User, user_result[0])
-
- l = q.select_by(item_name='item 4')
- self.assert_result(l, User, user_result[0])
-
- l = q.filter(orderitems.c.item_name=='item 4').join('item_name').list()
- self.assert_result(l, User, user_result[0])
-
- l = q.filter(orderitems.c.item_name=='item 4').join('items').list()
- self.assert_result(l, User, user_result[0])
-
- # test comparing to an object instance
- item = sess.query(Item).get_by(item_name='item 4')
-
- l = sess.query(Order).select_by(items=item)
- self.assert_result(l, Order, user_all_result[0]['orders'][1][1])
-
- l = q.select_by(items=item)
- self.assert_result(l, User, user_result[0])
-
- # TODO: this works differently from:
- #q = sess.query(User).join(['orders', 'items']).select_by(order_id=3)
- # because select_by() doesnt respect query._joinpoint, whereas filter_by does
- q = sess.query(User).join(['orders', 'items']).filter_by(order_id=3).list()
- self.assert_result(l, User, user_result[0])
-
- try:
- # this should raise AttributeError
- l = q.select_by(items=5)
- assert False
- except AttributeError:
- assert True
-
- def testautojoinm2m(self):
- """test functions derived from Query's _join_to function."""
-
- m = mapper(Order, orders, properties = {
- 'items' : relation(mapper(Item, orderitems, properties = {
- 'keywords' : relation(mapper(Keyword, keywords), itemkeywords)
- }))
- })
-
- sess = create_session()
- q = sess.query(m)
-
- l = q.filter(keywords.c.name=='square').join(['items', 'keywords']).list()
- self.assert_result(l, Order, order_result[1])
-
- # test comparing to an object instance
- item = sess.query(Item).selectfirst()
- l = sess.query(Item).select_by(keywords=item.keywords[0])
- assert item == l[0]
def testcustomjoin(self):
"""test that the from_obj parameter to query.select() can be used
@@ -475,7 +300,7 @@ class MapperTest(MapperSuperTest):
# l = create_session().query(User).select(order_by=None)
- @testbase.unsupported('firebird')
+ @testing.unsupported('firebird')
def testfunction(self):
"""test mapping to a SELECT statement that has functions in it."""
s = select([users, (users.c.user_id * 2).label('concat'), func.count(addresses.c.address_id).label('count')],
@@ -488,52 +313,7 @@ class MapperTest(MapperSuperTest):
assert l[0].concat == l[0].user_id * 2 == 14
assert l[1].concat == l[1].user_id * 2 == 16
- def testexternalcolumns(self):
- """test creating mappings that reference external columns or functions"""
-
- f = (users.c.user_id *2).label('concat')
- try:
- mapper(User, users, properties={
- 'concat': f,
- })
- class_mapper(User)
- except exceptions.ArgumentError, e:
- assert str(e) == "Column '%s' is not represented in mapper's table. Use the `column_property()` function to force this column to be mapped as a read-only attribute." % str(f)
- clear_mappers()
-
- mapper(User, users, properties={
- 'concat': column_property(f),
- 'count': column_property(select([func.count(addresses.c.address_id)], users.c.user_id==addresses.c.user_id, scalar=True).label('count'))
- })
-
- sess = create_session()
- l = sess.query(User).select()
- for u in l:
- print "User", u.user_id, u.user_name, u.concat, u.count
- assert l[0].concat == l[0].user_id * 2 == 14
- assert l[1].concat == l[1].user_id * 2 == 16
-
- ### eager loads, not really working across all DBs, no column aliasing in place so
- # results still wont be good for larger situations
- clear_mappers()
- mapper(Address, addresses, properties={
- 'user':relation(User, lazy=False)
- })
-
- mapper(User, users, properties={
- 'concat': column_property(f),
- })
-
- for x in range(0, 2):
- sess.clear()
- l = sess.query(Address).select()
- for a in l:
- print "User", a.user.user_id, a.user.user_name, a.user.concat
- assert l[0].user.concat == l[0].user.user_id * 2 == 14
- assert l[1].user.concat == l[1].user.user_id * 2 == 16
-
-
- @testbase.unsupported('firebird')
+ @testing.unsupported('firebird')
def testcount(self):
"""test the count function on Query.
@@ -610,12 +390,12 @@ class MapperTest(MapperSuperTest):
def go():
u = sess.query(User).options(eagerload('adlist')).get_by(user_name='jack')
self.assert_result(u.adlist, Address, *(user_address_result[0]['addresses'][1]))
- self.assert_sql_count(db, go, 1)
+ self.assert_sql_count(testbase.db, go, 1)
def testextensionoptions(self):
sess = create_session()
class ext1(MapperExtension):
- def populate_instance(self, mapper, selectcontext, row, instance, identitykey, isnew):
+ def populate_instance(self, mapper, selectcontext, row, instance, **flags):
"""test options at the Mapper._instance level"""
instance.TEST = "hello world"
return EXT_PASS
@@ -626,7 +406,7 @@ class MapperTest(MapperSuperTest):
def select_by(self, *args, **kwargs):
"""test options at the Query level"""
return "HI"
- def populate_instance(self, mapper, selectcontext, row, instance, identitykey, isnew):
+ def populate_instance(self, mapper, selectcontext, row, instance, **flags):
"""test options at the Mapper._instance level"""
instance.TEST_2 = "also hello world"
return EXT_PASS
@@ -649,7 +429,7 @@ class MapperTest(MapperSuperTest):
def go():
self.assert_result(l, User, *user_address_result)
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
def testeageroptionswithlimit(self):
sess = create_session()
@@ -661,7 +441,7 @@ class MapperTest(MapperSuperTest):
def go():
assert u.user_id == 8
assert len(u.addresses) == 3
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
sess.clear()
@@ -670,7 +450,7 @@ class MapperTest(MapperSuperTest):
u = sess.query(User).get_by(user_id=8)
assert u.user_id == 8
assert len(u.addresses) == 3
- assert "tbl_row_count" not in self.capture_sql(db, go)
+ assert "tbl_row_count" not in self.capture_sql(testbase.db, go)
def testlazyoptionswithlimit(self):
sess = create_session()
@@ -682,7 +462,7 @@ class MapperTest(MapperSuperTest):
def go():
assert u.user_id == 8
assert len(u.addresses) == 3
- self.assert_sql_count(db, go, 1)
+ self.assert_sql_count(testbase.db, go, 1)
def testeagerdegrade(self):
"""tests that an eager relation automatically degrades to a lazy relation if eager columns are not available"""
@@ -695,7 +475,7 @@ class MapperTest(MapperSuperTest):
def go():
l = sess.query(usermapper).select()
self.assert_result(l, User, *user_address_result)
- self.assert_sql_count(db, go, 1)
+ self.assert_sql_count(testbase.db, go, 1)
sess.clear()
@@ -706,7 +486,7 @@ class MapperTest(MapperSuperTest):
r = users.select().execute()
l = usermapper.instances(r, sess)
self.assert_result(l, User, *user_address_result)
- self.assert_sql_count(db, go, 4)
+ self.assert_sql_count(testbase.db, go, 4)
clear_mappers()
@@ -733,7 +513,7 @@ class MapperTest(MapperSuperTest):
def go():
l = sess.query(usermapper).select()
self.assert_result(l, User, *user_all_result)
- self.assert_sql_count(db, go, 1)
+ self.assert_sql_count(testbase.db, go, 1)
sess.clear()
@@ -743,7 +523,7 @@ class MapperTest(MapperSuperTest):
r = users.select().execute()
l = usermapper.instances(r, sess)
self.assert_result(l, User, *user_all_result)
- self.assert_sql_count(db, go, 7)
+ self.assert_sql_count(testbase.db, go, 7)
def testlazyoptions(self):
@@ -755,7 +535,7 @@ class MapperTest(MapperSuperTest):
l = sess.query(User).options(lazyload('addresses')).select()
def go():
self.assert_result(l, User, *user_address_result)
- self.assert_sql_count(db, go, 3)
+ self.assert_sql_count(testbase.db, go, 3)
def testlatecompile(self):
"""tests mappers compiling late in the game"""
@@ -769,7 +549,7 @@ class MapperTest(MapperSuperTest):
u = sess.query(User).select()
def go():
print u[0].orders[1].items[0].keywords[1]
- self.assert_sql_count(db, go, 3)
+ self.assert_sql_count(testbase.db, go, 3)
def testdeepoptions(self):
mapper(User, users,
@@ -787,18 +567,18 @@ class MapperTest(MapperSuperTest):
u = sess.query(User).select()
def go():
print u[0].orders[1].items[0].keywords[1]
- self.assert_sql_count(db, go, 3)
+ self.assert_sql_count(testbase.db, go, 3)
sess.clear()
print "-------MARK----------"
- # eagerload orders, orders.items, orders.items.keywords
- q2 = sess.query(User).options(eagerload('orders'), eagerload('orders.items'), eagerload('orders.items.keywords'))
+ # eagerload orders.items.keywords; eagerload_all() implies eager load of orders, orders.items
+ q2 = sess.query(User).options(eagerload_all('orders.items.keywords'))
u = q2.select()
def go():
print u[0].orders[1].items[0].keywords[1]
print "-------MARK2----------"
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
sess.clear()
@@ -808,7 +588,7 @@ class MapperTest(MapperSuperTest):
def go():
print u[0].orders[1].items[0].keywords[1]
print "-------MARK3----------"
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
print "-------MARK4----------"
sess.clear()
@@ -818,7 +598,7 @@ class MapperTest(MapperSuperTest):
print "-------MARK5----------"
q3 = sess.query(User).options(eagerload('orders.items.keywords'))
u = q3.select()
- self.assert_sql_count(db, go, 2)
+ self.assert_sql_count(testbase.db, go, 2)
class DeferredTest(MapperSuperTest):
@@ -839,8 +619,8 @@ class DeferredTest(MapperSuperTest):
o2 = l[2]
print o2.description
- orderby = str(orders.default_order_by()[0].compile(engine=db))
- self.assert_sql(db, go, [
+ orderby = str(orders.default_order_by()[0].compile(bind=testbase.db))
+ self.assert_sql(testbase.db, go, [
("SELECT orders.order_id AS orders_order_id, orders.user_id AS orders_user_id, orders.isopen AS orders_isopen FROM orders ORDER BY %s" % orderby, {}),
("SELECT orders.description AS orders_description FROM orders WHERE orders.order_id = :orders_order_id", {'orders_order_id':3})
])
@@ -893,7 +673,8 @@ class DeferredTest(MapperSuperTest):
'description':deferred(orders.c.description, group='primary'),
'opened':deferred(orders.c.isopen, group='primary')
})
- q = create_session().query(m)
+ sess = create_session()
+ q = sess.query(m)
def go():
l = q.select()
o2 = l[2]
@@ -901,12 +682,43 @@ class DeferredTest(MapperSuperTest):
assert o2.opened == 1
assert o2.userident == 7
assert o2.description == 'order 3'
- orderby = str(orders.default_order_by()[0].compile(db))
- self.assert_sql(db, go, [
+ orderby = str(orders.default_order_by()[0].compile(testbase.db))
+ self.assert_sql(testbase.db, go, [
("SELECT orders.order_id AS orders_order_id FROM orders ORDER BY %s" % orderby, {}),
("SELECT orders.user_id AS orders_user_id, orders.description AS orders_description, orders.isopen AS orders_isopen FROM orders WHERE orders.order_id = :orders_order_id", {'orders_order_id':3})
])
+ o2 = q.select()[2]
+# assert o2.opened == 1
+ assert o2.description == 'order 3'
+ assert o2 not in sess.dirty
+ o2.description = 'order 3'
+ def go():
+ sess.flush()
+ self.assert_sql_count(testbase.db, go, 0)
+
+ def testcommitsstate(self):
+ """test that when deferred elements are loaded via a group, they get the proper CommittedState
+ and dont result in changes being committed"""
+
+ m = mapper(Order, orders, properties = {
+ 'userident':deferred(orders.c.user_id, group='primary'),
+ 'description':deferred(orders.c.description, group='primary'),
+ 'opened':deferred(orders.c.isopen, group='primary')
+ })
+ sess = create_session()
+ q = sess.query(m)
+ o2 = q.select()[2]
+ # this will load the group of attributes
+ assert o2.description == 'order 3'
+ assert o2 not in sess.dirty
+ # this will mark it as 'dirty', but nothing actually changed
+ o2.description = 'order 3'
+ def go():
+ # therefore the flush() shouldnt actually issue any SQL
+ sess.flush()
+ self.assert_sql_count(testbase.db, go, 0)
+
def testoptions(self):
"""tests using options on a mapper to create deferred and undeferred columns"""
m = mapper(Order, orders)
@@ -917,8 +729,8 @@ class DeferredTest(MapperSuperTest):
l = q2.select()
print l[2].user_id
- orderby = str(orders.default_order_by()[0].compile(db))
- self.assert_sql(db, go, [
+ orderby = str(orders.default_order_by()[0].compile(testbase.db))
+ self.assert_sql(testbase.db, go, [
("SELECT orders.order_id AS orders_order_id, orders.description AS orders_description, orders.isopen AS orders_isopen FROM orders ORDER BY %s" % orderby, {}),
("SELECT orders.user_id AS orders_user_id FROM orders WHERE orders.order_id = :orders_order_id", {'orders_order_id':3})
])
@@ -927,10 +739,31 @@ class DeferredTest(MapperSuperTest):
def go():
l = q3.select()
print l[3].user_id
- self.assert_sql(db, go, [
+ self.assert_sql(testbase.db, go, [
("SELECT orders.order_id AS orders_order_id, orders.user_id AS orders_user_id, orders.description AS orders_description, orders.isopen AS orders_isopen FROM orders ORDER BY %s" % orderby, {}),
])
+ def testundefergroup(self):
+ """tests undefer_group()"""
+ m = mapper(Order, orders, properties = {
+ 'userident':deferred(orders.c.user_id, group='primary'),
+ 'description':deferred(orders.c.description, group='primary'),
+ 'opened':deferred(orders.c.isopen, group='primary')
+ })
+ sess = create_session()
+ q = sess.query(m)
+ def go():
+ l = q.options(undefer_group('primary')).select()
+ o2 = l[2]
+ print o2.opened, o2.description, o2.userident
+ assert o2.opened == 1
+ assert o2.userident == 7
+ assert o2.description == 'order 3'
+ orderby = str(orders.default_order_by()[0].compile(testbase.db))
+ self.assert_sql(testbase.db, go, [
+ ("SELECT orders.user_id AS orders_user_id, orders.description AS orders_description, orders.isopen AS orders_isopen, orders.order_id AS orders_order_id FROM orders ORDER BY %s" % orderby, {}),
+ ])
+
def testdeepoptions(self):
m = mapper(User, users, properties={
@@ -946,7 +779,7 @@ class DeferredTest(MapperSuperTest):
item = l[0].orders[1].items[1]
def go():
print item.item_name
- self.assert_sql_count(db, go, 1)
+ self.assert_sql_count(testbase.db, go, 1)
self.assert_(item.item_name == 'item 4')
sess.clear()
q2 = q.options(undefer('orders.items.item_name'))
@@ -954,10 +787,138 @@ class DeferredTest(MapperSuperTest):
item = l[0].orders[1].items[1]
def go():
print item.item_name
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
self.assert_(item.item_name == 'item 4')
-
+class CompositeTypesTest(ORMTest):
+ def define_tables(self, metadata):
+ global graphs, edges
+ graphs = Table('graphs', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('version_id', Integer, primary_key=True),
+ Column('name', String(30)))
+
+ edges = Table('edges', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('graph_id', Integer, nullable=False),
+ Column('graph_version_id', Integer, nullable=False),
+ Column('x1', Integer),
+ Column('y1', Integer),
+ Column('x2', Integer),
+ Column('y2', Integer),
+ ForeignKeyConstraint(['graph_id', 'graph_version_id'], ['graphs.id', 'graphs.version_id'])
+ )
+
+ def test_basic(self):
+ class Point(object):
+ def __init__(self, x, y):
+ self.x = x
+ self.y = y
+ def __colset__(self):
+ return [self.x, self.y]
+ def __eq__(self, other):
+ return other.x == self.x and other.y == self.y
+ def __ne__(self, other):
+ return not self.__eq__(other)
+
+ class Graph(object):
+ pass
+ class Edge(object):
+ def __init__(self, start, end):
+ self.start = start
+ self.end = end
+
+ mapper(Graph, graphs, properties={
+ 'edges':relation(Edge)
+ })
+ mapper(Edge, edges, properties={
+ 'start':composite(Point, edges.c.x1, edges.c.y1),
+ 'end':composite(Point, edges.c.x2, edges.c.y2)
+ })
+
+ sess = create_session()
+ g = Graph()
+ g.id = 1
+ g.version_id=1
+ g.edges.append(Edge(Point(3, 4), Point(5, 6)))
+ g.edges.append(Edge(Point(14, 5), Point(2, 7)))
+ sess.save(g)
+ sess.flush()
+
+ sess.clear()
+ g2 = sess.query(Graph).get([g.id, g.version_id])
+ for e1, e2 in zip(g.edges, g2.edges):
+ assert e1.start == e2.start
+ assert e1.end == e2.end
+
+ g2.edges[1].end = Point(18, 4)
+ sess.flush()
+ sess.clear()
+ e = sess.query(Edge).get(g2.edges[1].id)
+ assert e.end == Point(18, 4)
+
+ e.end.x = 19
+ e.end.y = 5
+ sess.flush()
+ sess.clear()
+ assert sess.query(Edge).get(g2.edges[1].id).end == Point(19, 5)
+
+ g.edges[1].end = Point(19, 5)
+
+ sess.clear()
+ def go():
+ g2 = sess.query(Graph).options(eagerload('edges')).get([g.id, g.version_id])
+ for e1, e2 in zip(g.edges, g2.edges):
+ assert e1.start == e2.start
+ assert e1.end == e2.end
+ self.assert_sql_count(testbase.db, go, 1)
+
+ # test comparison of CompositeProperties to their object instances
+ g = sess.query(Graph).get([1, 1])
+ assert sess.query(Edge).filter(Edge.start==Point(3, 4)).one() is g.edges[0]
+
+ assert sess.query(Edge).filter(Edge.start!=Point(3, 4)).first() is g.edges[1]
+
+ assert sess.query(Edge).filter(Edge.start==None).all() == []
+
+
+ def test_pk(self):
+ """test using a composite type as a primary key"""
+
+ class Version(object):
+ def __init__(self, id, version):
+ self.id = id
+ self.version = version
+ def __colset__(self):
+ return [self.id, self.version]
+ def __eq__(self, other):
+ return other.id == self.id and other.version == self.version
+ def __ne__(self, other):
+ return not self.__eq__(other)
+
+ class Graph(object):
+ def __init__(self, version):
+ self.version = version
+
+ mapper(Graph, graphs, properties={
+ 'version':composite(Version, graphs.c.id, graphs.c.version_id)
+ })
+
+ sess = create_session()
+ g = Graph(Version(1, 1))
+ sess.save(g)
+ sess.flush()
+
+ sess.clear()
+ g2 = sess.query(Graph).get([1, 1])
+ assert g.version == g2.version
+ sess.clear()
+
+ g2 = sess.query(Graph).get(Version(1, 1))
+ assert g.version == g2.version
+
+
+
class NoLoadTest(MapperSuperTest):
def testbasic(self):
"""tests a basic one-to-many lazy load"""
@@ -975,7 +936,6 @@ class NoLoadTest(MapperSuperTest):
self.assert_result(l[0], User,
{'user_id' : 7, 'addresses' : (Address, [])},
)
-
def testoptions(self):
m = mapper(User, users, properties = dict(
addresses = relation(mapper(Address, addresses), lazy=None)
@@ -992,8 +952,20 @@ class NoLoadTest(MapperSuperTest):
{'user_id' : 7, 'addresses' : (Address, [{'address_id' : 1}])},
)
-
-
+class MapperExtensionTest(MapperSuperTest):
+ def testcreateinstance(self):
+ class Ext(MapperExtension):
+ def create_instance(self, *args, **kwargs):
+ return User()
+ m = mapper(Address, addresses)
+ m = mapper(User, users, extension=Ext(), properties = dict(
+ addresses = relation(Address, lazy=True),
+ ))
+
+ q = create_session().query(m)
+ l = q.select();
+ self.assert_result(l, User, *user_address_result)
+
if __name__ == "__main__":
testbase.main()
diff --git a/test/orm/memusage.py b/test/orm/memusage.py
index 4e961a6d7..26da7c010 100644
--- a/test/orm/memusage.py
+++ b/test/orm/memusage.py
@@ -1,21 +1,18 @@
-from sqlalchemy import *
-from sqlalchemy.orm import mapperlib, session, unitofwork, attributes
-Mapper = mapperlib.Mapper
-import gc
import testbase
-import tables
+import gc
+from sqlalchemy import MetaData, Integer, String, ForeignKey
+from sqlalchemy.orm import mapper, relation, clear_mappers, create_session
+from sqlalchemy.orm.mapper import Mapper
+from testlib import *
class A(object):pass
class B(object):pass
-class MapperCleanoutTest(testbase.AssertMixin):
+class MapperCleanoutTest(AssertMixin):
"""test that clear_mappers() removes everything related to the class.
does not include classes that use the assignmapper extension."""
- def setUp(self):
- global engine
- engine = testbase.db
-
+
def test_mapper_cleanup(self):
for x in range(0, 5):
self.do_test()
@@ -33,7 +30,7 @@ class MapperCleanoutTest(testbase.AssertMixin):
assert True
def do_test(self):
- metadata = MetaData(engine)
+ metadata = MetaData(testbase.db)
table1 = Table("mytable", metadata,
Column('col1', Integer, primary_key=True),
diff --git a/test/orm/merge.py b/test/orm/merge.py
index cca01f2a5..3dd0a95a4 100644
--- a/test/orm/merge.py
+++ b/test/orm/merge.py
@@ -1,8 +1,9 @@
-from testbase import PersistTest, AssertMixin
import testbase
from sqlalchemy import *
-from tables import *
-import tables
+from sqlalchemy.orm import *
+from testlib import *
+from testlib.tables import *
+import testlib.tables as tables
class MergeTest(AssertMixin):
"""tests session.merge() functionality"""
@@ -164,5 +165,3 @@ class MergeTest(AssertMixin):
if __name__ == "__main__":
testbase.main()
-
- \ No newline at end of file
diff --git a/test/orm/onetoone.py b/test/orm/onetoone.py
index 6ac7c514d..e41fa1d20 100644
--- a/test/orm/onetoone.py
+++ b/test/orm/onetoone.py
@@ -1,6 +1,8 @@
import testbase
from sqlalchemy import *
+from sqlalchemy.orm import *
from sqlalchemy.ext.sessioncontext import SessionContext
+from testlib import *
class Jack(object):
def __repr__(self):
@@ -22,7 +24,7 @@ class Port(object):
self.name=name
self.description = description
-class O2OTest(testbase.AssertMixin):
+class O2OTest(AssertMixin):
def setUpAll(self):
global jack, port, metadata, ctx
metadata = MetaData(testbase.db)
diff --git a/test/orm/query.py b/test/orm/query.py
index 872d1772e..3783e1fa0 100644
--- a/test/orm/query.py
+++ b/test/orm/query.py
@@ -1,43 +1,12 @@
import testbase
+import operator
from sqlalchemy import *
+from sqlalchemy import ansisql
from sqlalchemy.orm import *
+from testlib import *
from fixtures import *
-class Base(object):
- def __init__(self, **kwargs):
- for k in kwargs:
- setattr(self, k, kwargs[k])
-
- def __ne__(self, other):
- return not self.__eq__(other)
-
- def __eq__(self, other):
- """'passively' compare this object to another.
-
- only look at attributes that are present on the source object.
-
- """
- # use __dict__ to avoid instrumented properties
- for attr in self.__dict__.keys():
- if attr[0] == '_':
- continue
- value = getattr(self, attr)
- if hasattr(value, '__iter__') and not isinstance(value, basestring):
- if len(value) == 0:
- continue
- for (us, them) in zip(value, getattr(other, attr)):
- if us != them:
- return False
- else:
- continue
- else:
- if value is not None:
- if value != getattr(other, attr):
- return False
- else:
- return True
-
-class QueryTest(testbase.ORMTest):
+class QueryTest(ORMTest):
keep_mappers = True
keep_data = True
@@ -53,16 +22,16 @@ class QueryTest(testbase.ORMTest):
def define_tables(self, meta):
# a slight dirty trick here.
meta.tables = metadata.tables
- metadata.connect(meta.engine)
+ metadata.connect(meta.bind)
def setup_mappers(self):
mapper(User, users, properties={
- 'addresses':relation(Address),
+ 'addresses':relation(Address, backref='user'),
'orders':relation(Order, backref='user'), # o2m, m2o
})
mapper(Address, addresses)
mapper(Order, orders, properties={
- 'items':relation(Item, secondary=order_items), #m2m
+ 'items':relation(Item, secondary=order_items, order_by=items.c.id), #m2m
'address':relation(Address), # m2o
})
mapper(Item, items, properties={
@@ -70,7 +39,6 @@ class QueryTest(testbase.ORMTest):
})
mapper(Keyword, keywords)
-
class GetTest(QueryTest):
def test_get(self):
s = create_session()
@@ -82,6 +50,33 @@ class GetTest(QueryTest):
u2 = s.query(User).get(7)
assert u is not u2
+ def test_load(self):
+ s = create_session()
+
+ try:
+ assert s.query(User).load(19) is None
+ assert False
+ except exceptions.InvalidRequestError:
+ assert True
+
+ u = s.query(User).load(7)
+ u2 = s.query(User).load(7)
+ assert u is u2
+ s.clear()
+ u2 = s.query(User).load(7)
+ assert u is not u2
+
+ u2.name = 'some name'
+ a = Address(name='some other name')
+ u2.addresses.append(a)
+ assert u2 in s.dirty
+ assert a in u2.addresses
+
+ s.query(User).load(7)
+ assert u2 not in s.dirty
+ assert u2.name =='jack'
+ assert a not in u2.addresses
+
def test_unicode(self):
"""test that Query.get properly sets up the type for the bind parameter. using unicode would normally fail
on postgres, mysql and oracle unless it is converted to an encoded string"""
@@ -90,16 +85,116 @@ class GetTest(QueryTest):
Column('id', Unicode(40), primary_key=True),
Column('data', Unicode(40)))
table.create()
- ustring = 'petit voix m\xe2\x80\x99a'.decode('utf-8')
+ ustring = 'petit voix m\xe2\x80\x99a '.decode('utf-8')
table.insert().execute(id=ustring, data=ustring)
class LocalFoo(Base):pass
mapper(LocalFoo, table)
assert create_session().query(LocalFoo).get(ustring) == LocalFoo(id=ustring, data=ustring)
+ def test_populate_existing(self):
+ s = create_session()
+
+ userlist = s.query(User).all()
+
+ u = userlist[0]
+ u.name = 'foo'
+ a = Address(name='ed')
+ u.addresses.append(a)
+
+ self.assert_(a in u.addresses)
+
+ s.query(User).populate_existing().all()
+
+ self.assert_(u not in s.dirty)
+
+ self.assert_(u.name == 'jack')
+
+ self.assert_(a not in u.addresses)
+
+ u.addresses[0].email_address = 'lala'
+ u.orders[1].items[2].description = 'item 12'
+ # test that lazy load doesnt change child items
+ s.query(User).populate_existing().all()
+ assert u.addresses[0].email_address == 'lala'
+ assert u.orders[1].items[2].description == 'item 12'
+
+ # eager load does
+ s.query(User).options(eagerload('addresses'), eagerload_all('orders.items')).populate_existing().all()
+ assert u.addresses[0].email_address == 'jack@bean.com'
+ assert u.orders[1].items[2].description == 'item 5'
+
+class OperatorTest(QueryTest):
+ """test sql.Comparator implementation for MapperProperties"""
+
+ def _test(self, clause, expected):
+ c = str(clause.compile(dialect=ansisql.ANSIDialect()))
+ assert c == expected, "%s != %s" % (c, expected)
+
+ def test_arithmetic(self):
+ create_session().query(User)
+ for (py_op, sql_op) in ((operator.add, '+'), (operator.mul, '*'),
+ (operator.sub, '-'), (operator.div, '/'),
+ ):
+ for (lhs, rhs, res) in (
+ (5, User.id, ':users_id %s users.id'),
+ (5, literal(6), ':literal %s :literal_1'),
+ (User.id, 5, 'users.id %s :users_id'),
+ (User.id, literal('b'), 'users.id %s :literal'),
+ (User.id, User.id, 'users.id %s users.id'),
+ (literal(5), 'b', ':literal %s :literal_1'),
+ (literal(5), User.id, ':literal %s users.id'),
+ (literal(5), literal(6), ':literal %s :literal_1'),
+ ):
+ self._test(py_op(lhs, rhs), res % sql_op)
+
+ def test_comparison(self):
+ create_session().query(User)
+ for (py_op, fwd_op, rev_op) in ((operator.lt, '<', '>'),
+ (operator.gt, '>', '<'),
+ (operator.eq, '=', '='),
+ (operator.ne, '!=', '!='),
+ (operator.le, '<=', '>='),
+ (operator.ge, '>=', '<=')):
+ for (lhs, rhs, l_sql, r_sql) in (
+ ('a', User.id, ':users_id', 'users.id'),
+ ('a', literal('b'), ':literal_1', ':literal'), # note swap!
+ (User.id, 'b', 'users.id', ':users_id'),
+ (User.id, literal('b'), 'users.id', ':literal'),
+ (User.id, User.id, 'users.id', 'users.id'),
+ (literal('a'), 'b', ':literal', ':literal_1'),
+ (literal('a'), User.id, ':literal', 'users.id'),
+ (literal('a'), literal('b'), ':literal', ':literal_1'),
+ ):
+
+ # the compiled clause should match either (e.g.):
+ # 'a' < 'b' -or- 'b' > 'a'.
+ compiled = str(py_op(lhs, rhs).compile(dialect=ansisql.ANSIDialect()))
+ fwd_sql = "%s %s %s" % (l_sql, fwd_op, r_sql)
+ rev_sql = "%s %s %s" % (r_sql, rev_op, l_sql)
+
+ self.assert_(compiled == fwd_sql or compiled == rev_sql,
+ "\n'" + compiled + "'\n does not match\n'" +
+ fwd_sql + "'\n or\n'" + rev_sql + "'")
+
+ def test_in(self):
+ self._test(User.id.in_('a', 'b'), "users.id IN (:users_id, :users_id_1)")
+
+ def test_clauses(self):
+ for (expr, compare) in (
+ (func.max(User.id), "max(users.id)"),
+ (desc(User.id), "users.id DESC"),
+ (between(5, User.id, Address.id), ":literal BETWEEN users.id AND addresses.id"),
+ # this one would require adding compile() to InstrumentedScalarAttribute. do we want this ?
+ #(User.id, "users.id")
+ ):
+ c = expr.compile(dialect=ansisql.ANSIDialect())
+ assert str(c) == compare, "%s != %s" % (str(c), compare)
+
+
class CompileTest(QueryTest):
def test_deferred(self):
session = create_session()
- s = session.query(User).filter(and_(addresses.c.email_address == bindparam('emailad'), addresses.c.user_id==users.c.id)).compile()
+ s = session.query(User).filter(and_(addresses.c.email_address == bindparam('emailad'), Address.user_id==User.id)).compile()
l = session.query(User).instances(s.execute(emailad = 'jack@bean.com'))
assert [User(id=7)] == l
@@ -108,7 +203,23 @@ class SliceTest(QueryTest):
def test_first(self):
assert User(id=7) == create_session().query(User).first()
- assert create_session().query(User).filter(users.c.id==27).first() is None
+ assert create_session().query(User).filter(User.id==27).first() is None
+
+ # more slice tests are available in test/orm/generative.py
+
+class TextTest(QueryTest):
+ def test_fulltext(self):
+ assert [User(id=7), User(id=8), User(id=9),User(id=10)] == create_session().query(User).from_statement("select * from users").all()
+
+ def test_fragment(self):
+ assert [User(id=8), User(id=9)] == create_session().query(User).filter("id in (8, 9)").all()
+
+ assert [User(id=9)] == create_session().query(User).filter("name='fred'").filter("id=9").all()
+
+ assert [User(id=9)] == create_session().query(User).filter("name='fred'").filter(User.id==9).all()
+
+ def test_binds(self):
+ assert [User(id=8), User(id=9)] == create_session().query(User).filter("id in (:id1, :id2)").params(id1=8, id2=9).all()
class FilterTest(QueryTest):
def test_basic(self):
@@ -122,8 +233,75 @@ class FilterTest(QueryTest):
assert User(id=8) == create_session().query(User)[1]
def test_onefilter(self):
- assert [User(id=8), User(id=9)] == create_session().query(User).filter(users.c.name.endswith('ed')).all()
+ assert [User(id=8), User(id=9)] == create_session().query(User).filter(User.name.endswith('ed')).all()
+
+ def test_contains(self):
+ """test comparing a collection to an object instance."""
+
+ sess = create_session()
+ address = sess.query(Address).get(3)
+ assert [User(id=8)] == sess.query(User).filter(User.addresses.contains(address)).all()
+ try:
+ sess.query(User).filter(User.addresses == address)
+ assert False
+ except exceptions.InvalidRequestError:
+ assert True
+
+ assert [User(id=10)] == sess.query(User).filter(User.addresses==None).all()
+
+ try:
+ assert [User(id=7), User(id=9), User(id=10)] == sess.query(User).filter(User.addresses!=address).all()
+ assert False
+ except exceptions.InvalidRequestError:
+ assert True
+
+ #assert [User(id=7), User(id=9), User(id=10)] == sess.query(User).filter(User.addresses!=address).all()
+
+ def test_any(self):
+ sess = create_session()
+
+ assert [User(id=8), User(id=9)] == sess.query(User).filter(User.addresses.any(Address.email_address.like('%ed%'))).all()
+
+ assert [User(id=8)] == sess.query(User).filter(User.addresses.any(Address.email_address.like('%ed%'), id=4)).all()
+
+ assert [User(id=9)] == sess.query(User).filter(User.addresses.any(email_address='fred@fred.com')).all()
+
+ def test_has(self):
+ sess = create_session()
+ assert [Address(id=5)] == sess.query(Address).filter(Address.user.has(name='fred')).all()
+
+ assert [Address(id=2), Address(id=3), Address(id=4), Address(id=5)] == sess.query(Address).filter(Address.user.has(User.name.like('%ed%'))).all()
+
+ assert [Address(id=2), Address(id=3), Address(id=4)] == sess.query(Address).filter(Address.user.has(User.name.like('%ed%'), id=8)).all()
+
+ def test_contains_m2m(self):
+ sess = create_session()
+ item = sess.query(Item).get(3)
+ assert [Order(id=1), Order(id=2), Order(id=3)] == sess.query(Order).filter(Order.items.contains(item)).all()
+
+ assert [Order(id=4), Order(id=5)] == sess.query(Order).filter(~Order.items.contains(item)).all()
+
+ def test_comparison(self):
+ """test scalar comparison to an object instance"""
+
+ sess = create_session()
+ user = sess.query(User).get(8)
+ assert [Address(id=2), Address(id=3), Address(id=4)] == sess.query(Address).filter(Address.user==user).all()
+
+ assert [Address(id=1), Address(id=5)] == sess.query(Address).filter(Address.user!=user).all()
+
+class AggregateTest(QueryTest):
+ def test_sum(self):
+ sess = create_session()
+ orders = sess.query(Order).filter(Order.id.in_(2, 3, 4))
+ assert orders.sum(Order.user_id * Order.address_id) == 79
+
+ def test_apply(self):
+ sess = create_session()
+ assert sess.query(Order).apply_sum(Order.user_id * Order.address_id).filter(Order.id.in_(2, 3, 4)).one() == 79
+
+
class CountTest(QueryTest):
def test_basic(self):
assert 4 == create_session().query(User).count()
@@ -139,7 +317,7 @@ class TextTest(QueryTest):
assert [User(id=9)] == create_session().query(User).filter("name='fred'").filter("id=9").all()
- assert [User(id=9)] == create_session().query(User).filter("name='fred'").filter(users.c.id==9).all()
+ assert [User(id=9)] == create_session().query(User).filter("name='fred'").filter(User.id==9).all()
def test_binds(self):
assert [User(id=8), User(id=9)] == create_session().query(User).filter("id in (:id1, :id2)").params(id1=8, id2=9).all()
@@ -188,14 +366,25 @@ class ParentTest(QueryTest):
class JoinTest(QueryTest):
+
def test_overlapping_paths(self):
- # load a user who has an order that contains item id 3 and address id 1 (order 3, owned by jack)
- result = create_session().query(User).join(['orders', 'items']).filter_by(id=3).reset_joinpoint().join(['orders','address']).filter_by(id=1).all()
- assert [User(id=7, name='jack')] == result
+ for aliased in (True,False):
+ # load a user who has an order that contains item id 3 and address id 1 (order 3, owned by jack)
+ result = create_session().query(User).join(['orders', 'items'], aliased=aliased).filter_by(id=3).join(['orders','address'], aliased=aliased).filter_by(id=1).all()
+ assert [User(id=7, name='jack')] == result
def test_overlapping_paths_outerjoin(self):
- result = create_session().query(User).outerjoin(['orders', 'items']).filter_by(id=3).reset_joinpoint().outerjoin(['orders','address']).filter_by(id=1).all()
+ result = create_session().query(User).outerjoin(['orders', 'items']).filter_by(id=3).outerjoin(['orders','address']).filter_by(id=1).all()
assert [User(id=7, name='jack')] == result
+
+ def test_reset_joinpoint(self):
+ for aliased in (True, False):
+ # load a user who has an order that contains item id 3 and address id 1 (order 3, owned by jack)
+ result = create_session().query(User).join(['orders', 'items'], aliased=aliased).filter_by(id=3).reset_joinpoint().join(['orders','address'], aliased=aliased).filter_by(id=1).all()
+ assert [User(id=7, name='jack')] == result
+
+ result = create_session().query(User).outerjoin(['orders', 'items'], aliased=aliased).filter_by(id=3).reset_joinpoint().outerjoin(['orders','address'], aliased=aliased).filter_by(id=1).all()
+ assert [User(id=7, name='jack')] == result
def test_overlap_with_aliases(self):
oalias = orders.alias('oalias')
@@ -206,7 +395,64 @@ class JoinTest(QueryTest):
result = create_session().query(User).select_from(users.join(oalias)).filter(oalias.c.description.in_("order 1", "order 2", "order 3")).join(['orders', 'items']).filter_by(id=4).all()
assert [User(id=7, name='jack')] == result
-class MultiplePathTest(testbase.ORMTest):
+ def test_aliased(self):
+ """test automatic generation of aliased joins."""
+
+ sess = create_session()
+
+ # test a basic aliasized path
+ q = sess.query(User).join('addresses', aliased=True).filter_by(email_address='jack@bean.com')
+ assert [User(id=7)] == q.all()
+
+ q = sess.query(User).join('addresses', aliased=True).filter(Address.email_address=='jack@bean.com')
+ assert [User(id=7)] == q.all()
+
+ # test two aliasized paths, one to 'orders' and the other to 'orders','items'.
+ # one row is returned because user 7 has order 3 and also has order 1 which has item 1
+ # this tests a o2m join and a m2m join.
+ q = sess.query(User).join('orders', aliased=True).filter(Order.description=="order 3").join(['orders', 'items'], aliased=True).filter(Item.description=="item 1")
+ assert q.count() == 1
+ assert [User(id=7)] == q.all()
+
+ # test the control version - same joins but not aliased. rows are not returned because order 3 does not have item 1
+ # addtionally by placing this test after the previous one, test that the "aliasing" step does not corrupt the
+ # join clauses that are cached by the relationship.
+ q = sess.query(User).join('orders').filter(Order.description=="order 3").join(['orders', 'items']).filter(Order.description=="item 1")
+ assert [] == q.all()
+ assert q.count() == 0
+
+ q = sess.query(User).join('orders', aliased=True).filter(Order.items.any(Item.description=='item 4'))
+ assert [User(id=7)] == q.all()
+
+ def test_aliased_add_entity(self):
+ """test the usage of aliased joins with add_entity()"""
+ sess = create_session()
+ q = sess.query(User).join('orders', aliased=True, id='order1').filter(Order.description=="order 3").join(['orders', 'items'], aliased=True, id='item1').filter(Item.description=="item 1")
+
+ try:
+ q.add_entity(Order, id='fakeid').compile()
+ assert False
+ except exceptions.InvalidRequestError, e:
+ assert str(e) == "Query has no alias identified by 'fakeid'"
+
+ try:
+ q.add_entity(Order, id='fakeid').instances(None)
+ assert False
+ except exceptions.InvalidRequestError, e:
+ assert str(e) == "Query has no alias identified by 'fakeid'"
+
+ q = q.add_entity(Order, id='order1').add_entity(Item, id='item1')
+ assert q.count() == 1
+ assert [(User(id=7), Order(description='order 3'), Item(description='item 1'))] == q.all()
+
+ q = sess.query(User).add_entity(Order).join('orders', aliased=True).filter(Order.description=="order 3").join('orders', aliased=True).filter(Order.description=='order 4')
+ try:
+ q.compile()
+ assert False
+ except exceptions.InvalidRequestError, e:
+ assert str(e) == "Ambiguous join for entity 'Mapper|Order|orders'; specify id=<someid> to query.join()/query.add_entity()"
+
+class MultiplePathTest(ORMTest):
def define_tables(self, metadata):
global t1, t2, t1t2_1, t1t2_2
t1 = Table('t1', metadata,
@@ -217,7 +463,7 @@ class MultiplePathTest(testbase.ORMTest):
Column('id', Integer, primary_key=True),
Column('data', String(30))
)
-
+
t1t2_1 = Table('t1t2_1', metadata,
Column('t1id', Integer, ForeignKey('t1.id')),
Column('t2id', Integer, ForeignKey('t2.id'))
@@ -227,23 +473,28 @@ class MultiplePathTest(testbase.ORMTest):
Column('t1id', Integer, ForeignKey('t1.id')),
Column('t2id', Integer, ForeignKey('t2.id'))
)
-
+
def test_basic(self):
class T1(object):pass
class T2(object):pass
-
+
mapper(T1, t1, properties={
't2s_1':relation(T2, secondary=t1t2_1),
't2s_2':relation(T2, secondary=t1t2_2),
})
mapper(T2, t2)
-
+
try:
- create_session().query(T1).join('t2s_1').filter_by(t2.c.id==5).reset_joinpoint().join('t2s_2')
+ create_session().query(T1).join('t2s_1').filter(t2.c.id==5).reset_joinpoint().join('t2s_2')
assert False
except exceptions.InvalidRequestError, e:
- assert str(e) == "Can't join to property 't2s_2'; a path to this table along a different secondary table already exists. Use explicit `Alias` objects."
+ assert str(e) == "Can't join to property 't2s_2'; a path to this table along a different secondary table already exists. Use the `alias=True` argument to `join()`."
+
+ create_session().query(T1).join('t2s_1', aliased=True).filter(t2.c.id==5).reset_joinpoint().join('t2s_2').all()
+ create_session().query(T1).join('t2s_1').filter(t2.c.id==5).reset_joinpoint().join('t2s_2', aliased=True).all()
+
+
class SynonymTest(QueryTest):
keep_mappers = True
keep_data = True
@@ -372,22 +623,48 @@ class InstancesTest(QueryTest):
l = q.instances(selectquery.execute(), Address)
assert l == expected
+ for aliased in (False, True):
+ q = sess.query(User)
+ q = q.add_entity(Address).outerjoin('addresses', aliased=aliased)
+ l = q.all()
+ assert l == expected
+
+ q = sess.query(User).add_entity(Address)
+ l = q.join('addresses', aliased=aliased).filter_by(email_address='ed@bettyboop.com').all()
+ assert l == [(user8, address3)]
+
+ q = sess.query(User, Address).join('addresses', aliased=aliased).filter_by(email_address='ed@bettyboop.com')
+ assert q.all() == [(user8, address3)]
+
+ q = sess.query(User, Address).join('addresses', aliased=aliased).options(eagerload('addresses')).filter_by(email_address='ed@bettyboop.com')
+ assert q.all() == [(user8, address3)]
+
+ def test_aliased_multi_mappers(self):
+ sess = create_session()
+
+ (user7, user8, user9, user10) = sess.query(User).all()
+ (address1, address2, address3, address4, address5) = sess.query(Address).all()
+
+ # note the result is a cartesian product
+ expected = [(user7, address1),
+ (user8, address2),
+ (user8, address3),
+ (user8, address4),
+ (user9, address5),
+ (user10, None)]
+
q = sess.query(User)
- q = q.add_entity(Address).outerjoin('addresses')
+ adalias = addresses.alias('adalias')
+ q = q.add_entity(Address, alias=adalias).select_from(users.outerjoin(adalias))
l = q.all()
assert l == expected
- q = sess.query(User).add_entity(Address)
- l = q.join('addresses').filter_by(email_address='ed@bettyboop.com').all()
+ q = sess.query(User).add_entity(Address, alias=adalias)
+ l = q.select_from(users.outerjoin(adalias)).filter(adalias.c.email_address=='ed@bettyboop.com').all()
assert l == [(user8, address3)]
- q = sess.query(User, Address).join('addresses').filter_by(email_address='ed@bettyboop.com')
- assert q.all() == [(user8, address3)]
-
- q = sess.query(User, Address).join('addresses').options(eagerload('addresses')).filter_by(email_address='ed@bettyboop.com')
- assert q.all() == [(user8, address3)]
-
def test_multi_columns(self):
+ """test aliased/nonalised joins with the usage of add_column()"""
sess = create_session()
(user7, user8, user9, user10) = sess.query(User).all()
expected = [(user7, 1),
@@ -395,18 +672,18 @@ class InstancesTest(QueryTest):
(user9, 1),
(user10, 0)
]
-
- q = sess.query(User)
- q = q.group_by([c for c in users.c]).order_by(User.c.id).outerjoin('addresses').add_column(func.count(addresses.c.id).label('count'))
- l = q.all()
- assert l == expected
+
+ for aliased in (False, True):
+ q = sess.query(User)
+ q = q.group_by([c for c in users.c]).order_by(User.id).outerjoin('addresses', aliased=aliased).add_column(func.count(Address.id).label('count'))
+ l = q.all()
+ assert l == expected
- s = select([users, func.count(addresses.c.id).label('count')], from_obj=[users.outerjoin(addresses)], group_by=[c for c in users.c], order_by=users.c.id)
+ s = select([users, func.count(addresses.c.id).label('count')]).select_from(users.outerjoin(addresses)).group_by(*[c for c in users.c]).order_by(User.id)
q = sess.query(User)
l = q.add_column("count").from_statement(s).all()
assert l == expected
- @testbase.unsupported('mysql') # only because of "+" operator requiring "concat" in mysql (fix #475)
def test_two_columns(self):
sess = create_session()
(user7, user8, user9, user10) = sess.query(User).all()
@@ -416,17 +693,162 @@ class InstancesTest(QueryTest):
(user9, 1, "Name:fred"),
(user10, 0, "Name:chuck")]
+ # test with a straight statement
s = select([users, func.count(addresses.c.id).label('count'), ("Name:" + users.c.name).label('concat')], from_obj=[users.outerjoin(addresses)], group_by=[c for c in users.c], order_by=[users.c.id])
q = create_session().query(User)
l = q.add_column("count").add_column("concat").from_statement(s).all()
assert l == expected
+ # test with select_from()
q = create_session().query(User).add_column(func.count(addresses.c.id))\
.add_column(("Name:" + users.c.name)).select_from(users.outerjoin(addresses))\
.group_by([c for c in users.c]).order_by(users.c.id)
assert q.all() == expected
+ # test with outerjoin() both aliased and non
+ for aliased in (False, True):
+ q = create_session().query(User).add_column(func.count(addresses.c.id))\
+ .add_column(("Name:" + users.c.name)).outerjoin('addresses', aliased=aliased)\
+ .group_by([c for c in users.c]).order_by(users.c.id)
+
+ assert q.all() == expected
+
+class CustomJoinTest(QueryTest):
+ keep_mappers = False
+
+ def setup_mappers(self):
+ pass
+
+ def test_double_same_mappers(self):
+ """test aliasing of joins with a custom join condition"""
+ mapper(Address, addresses)
+ mapper(Order, orders, properties={
+ 'items':relation(Item, secondary=order_items, lazy=True, order_by=items.c.id),
+ })
+ mapper(Item, items)
+ mapper(User, users, properties = dict(
+ addresses = relation(Address, lazy=True),
+ open_orders = relation(Order, primaryjoin = and_(orders.c.isopen == 1, users.c.id==orders.c.user_id), lazy=True),
+ closed_orders = relation(Order, primaryjoin = and_(orders.c.isopen == 0, users.c.id==orders.c.user_id), lazy=True)
+ ))
+ q = create_session().query(User)
+
+ assert [User(id=7)] == q.join(['open_orders', 'items'], aliased=True).filter(Item.id==4).join(['closed_orders', 'items'], aliased=True).filter(Item.id==3).all()
+
+class SelfReferentialJoinTest(ORMTest):
+ def define_tables(self, metadata):
+ global nodes
+ nodes = Table('nodes', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('parent_id', Integer, ForeignKey('nodes.id')),
+ Column('data', String(30)))
+
+ def test_join(self):
+ class Node(Base):
+ def append(self, node):
+ self.children.append(node)
+
+ mapper(Node, nodes, properties={
+ 'children':relation(Node, lazy=True, join_depth=3,
+ backref=backref('parent', remote_side=[nodes.c.id])
+ )
+ })
+ sess = create_session()
+ n1 = Node(data='n1')
+ n1.append(Node(data='n11'))
+ n1.append(Node(data='n12'))
+ n1.append(Node(data='n13'))
+ n1.children[1].append(Node(data='n121'))
+ n1.children[1].append(Node(data='n122'))
+ n1.children[1].append(Node(data='n123'))
+ sess.save(n1)
+ sess.flush()
+ sess.clear()
+
+ # TODO: the aliasing of the join in query._join_to has to limit the aliasing
+ # among local_side / remote_side (add local_side as an attribute on PropertyLoader)
+ # also implement this idea in EagerLoader
+ node = sess.query(Node).join('children', aliased=True).filter_by(data='n122').first()
+ assert node.data=='n12'
+
+ node = sess.query(Node).join(['children', 'children'], aliased=True).filter_by(data='n122').first()
+ assert node.data=='n1'
+
+ node = sess.query(Node).filter_by(data='n122').join('parent', aliased=True).filter_by(data='n12').\
+ join('parent', aliased=True, from_joinpoint=True).filter_by(data='n1').first()
+ assert node.data == 'n122'
+
+class ExternalColumnsTest(QueryTest):
+ keep_mappers = False
+
+ def setup_mappers(self):
+ pass
+
+ def test_external_columns_bad(self):
+ """test that SA catches some common mis-configurations of external columns."""
+ f = (users.c.id * 2)
+ try:
+ mapper(User, users, properties={
+ 'concat': f,
+ })
+ class_mapper(User)
+ except exceptions.ArgumentError, e:
+ assert str(e) == "Column '%s' is not represented in mapper's table. Use the `column_property()` function to force this column to be mapped as a read-only attribute." % str(f)
+ else:
+ raise 'expected ArgumentError'
+ clear_mappers()
+ try:
+ mapper(User, users, properties={
+ 'concat': column_property(users.c.id * 2),
+ })
+ except exceptions.ArgumentError, e:
+ assert str(e) == 'ColumnProperties must be named for the mapper to work with them. Try .label() to fix this'
+ else:
+ raise 'expected ArgumentError'
+
+ def test_external_columns_good(self):
+ """test querying mappings that reference external columns or selectables."""
+ mapper(User, users, properties={
+ 'concat': column_property((users.c.id * 2).label('concat')),
+ 'count': column_property(select([func.count(addresses.c.id)], users.c.id==addresses.c.user_id).correlate(users).label('count'))
+ })
+
+ mapper(Address, addresses, properties={
+ 'user':relation(User, lazy=True)
+ })
+
+ sess = create_session()
+ l = sess.query(User).select()
+ assert [
+ User(id=7, concat=14, count=1),
+ User(id=8, concat=16, count=3),
+ User(id=9, concat=18, count=1),
+ User(id=10, concat=20, count=0),
+ ] == l
+
+ address_result = [
+ Address(id=1, user=User(id=7, concat=14, count=1)),
+ Address(id=2, user=User(id=8, concat=16, count=3)),
+ Address(id=3, user=User(id=8, concat=16, count=3)),
+ Address(id=4, user=User(id=8, concat=16, count=3)),
+ Address(id=5, user=User(id=9, concat=18, count=1))
+ ]
+
+ assert address_result == sess.query(Address).all()
+
+ # run the eager version twice to test caching of aliased clauses
+ for x in range(2):
+ sess.clear()
+ def go():
+ assert address_result == sess.query(Address).options(eagerload('user')).all()
+ self.assert_sql_count(testbase.db, go, 1)
+
+ tuple_address_result = [(address, address.user) for address in address_result]
+
+ tuple_address_result == sess.query(Address).join('user').add_entity(User).all()
+
+ assert tuple_address_result == sess.query(Address).join('user', aliased=True, id='ualias').add_entity(User, id='ualias').all()
if __name__ == '__main__':
testbase.main()
diff --git a/test/orm/relationships.py b/test/orm/relationships.py
index 7c9bbc898..9fca22b24 100644
--- a/test/orm/relationships.py
+++ b/test/orm/relationships.py
@@ -1,12 +1,12 @@
import testbase
-import unittest, sys, datetime
-
-db = testbase.db
-
+import datetime
from sqlalchemy import *
+from sqlalchemy.orm import *
+from sqlalchemy.orm import collections
+from sqlalchemy.orm.collections import collection
+from testlib import *
-
-class RelationTest(testbase.PersistTest):
+class RelationTest(PersistTest):
"""this is essentially an extension of the "dependency.py" topological sort test.
in this test, a table is dependent on two other tables that are otherwise unrelated to each other.
the dependency sort must insure that this childmost table is below both parent tables in the outcome
@@ -15,10 +15,8 @@ class RelationTest(testbase.PersistTest):
to subtle differences in program execution, this test case was exposing the bug whereas the simpler tests
were not."""
def setUpAll(self):
- global tbl_a
- global tbl_b
- global tbl_c
- global tbl_d
+ global metadata, tbl_a, tbl_b, tbl_c, tbl_d
+
metadata = MetaData()
tbl_a = Table("tbl_a", metadata,
Column("id", Integer, primary_key=True),
@@ -41,8 +39,8 @@ class RelationTest(testbase.PersistTest):
)
def setUp(self):
global session
- session = create_session(bind_to=testbase.db)
- conn = session.connect()
+ session = create_session(bind=testbase.db)
+ conn = testbase.db.connect()
conn.create(tbl_a)
conn.create(tbl_b)
conn.create(tbl_c)
@@ -80,14 +78,14 @@ class RelationTest(testbase.PersistTest):
session.save_or_update(b)
def tearDown(self):
- conn = session.connect()
+ conn = testbase.db.connect()
conn.drop(tbl_d)
conn.drop(tbl_c)
conn.drop(tbl_b)
conn.drop(tbl_a)
def tearDownAll(self):
- testbase.metadata.tables.clear()
+ metadata.drop_all(testbase.db)
def testDeleteRootTable(self):
session.flush()
@@ -99,7 +97,7 @@ class RelationTest(testbase.PersistTest):
session.delete(c) # fails
session.flush()
-class RelationTest2(testbase.PersistTest):
+class RelationTest2(PersistTest):
"""this test tests a relationship on a column that is included in multiple foreign keys,
as well as a self-referential relationship on a composite key where one column in the foreign key
is 'joined to itself'."""
@@ -216,7 +214,7 @@ class RelationTest2(testbase.PersistTest):
assert sess.query(Employee).get([c1.company_id, 3]).reports_to.name == 'emp1'
assert sess.query(Employee).get([c2.company_id, 3]).reports_to.name == 'emp5'
-class RelationTest3(testbase.PersistTest):
+class RelationTest3(PersistTest):
def setUpAll(self):
global jobs, pageversions, pages, metadata, Job, Page, PageVersion, PageComment
import datetime
@@ -350,7 +348,7 @@ class RelationTest3(testbase.PersistTest):
s.delete(j)
s.flush()
-class RelationTest4(testbase.ORMTest):
+class RelationTest4(ORMTest):
"""test syncrules on foreign keys that are also primary"""
def define_tables(self, metadata):
global tableA, tableB
@@ -498,7 +496,7 @@ class RelationTest4(testbase.ORMTest):
assert a1 not in sess
assert b1 not in sess
-class RelationTest5(testbase.ORMTest):
+class RelationTest5(ORMTest):
"""test a map to a select that relates to a map to the table"""
def define_tables(self, metadata):
global items
@@ -554,7 +552,7 @@ class RelationTest5(testbase.ORMTest):
assert old.id == new.id
-class TypeMatchTest(testbase.ORMTest):
+class TypeMatchTest(ORMTest):
"""test errors raised when trying to add items whose type is not handled by a relation"""
def define_tables(self, metadata):
global a, b, c, d
@@ -672,7 +670,7 @@ class TypeMatchTest(testbase.ORMTest):
except exceptions.AssertionError, err:
assert str(err) == "Attribute 'a' on class '%s' doesn't handle objects of type '%s'" % (D, B)
-class TypedAssociationTable(testbase.ORMTest):
+class TypedAssociationTable(ORMTest):
def define_tables(self, metadata):
global t1, t2, t3
@@ -722,7 +720,7 @@ class TypedAssociationTable(testbase.ORMTest):
assert t3.count().scalar() == 1
# TODO: move these tests to either attributes.py test or its own module
-class CustomCollectionsTest(testbase.ORMTest):
+class CustomCollectionsTest(ORMTest):
def define_tables(self, metadata):
global sometable, someothertable
sometable = Table('sometable', metadata,
@@ -745,7 +743,7 @@ class CustomCollectionsTest(testbase.ORMTest):
})
mapper(Bar, someothertable)
f = Foo()
- assert isinstance(f.bars.data, MyList)
+ assert isinstance(f.bars, MyList)
def testlazyload(self):
"""test that a 'set' can be used as a collection and can lazyload."""
class Foo(object):
@@ -769,23 +767,27 @@ class CustomCollectionsTest(testbase.ORMTest):
def testdict(self):
"""test that a 'dict' can be used as a collection and can lazyload."""
+
class Foo(object):
pass
class Bar(object):
pass
class AppenderDict(dict):
- def append(self, item):
+ @collection.appender
+ def set(self, item):
self[id(item)] = item
- def __iter__(self):
- return iter(self.values())
+ @collection.remover
+ def remove(self, item):
+ if id(item) in self:
+ del self[id(item)]
mapper(Foo, sometable, properties={
'bars':relation(Bar, collection_class=AppenderDict)
})
mapper(Bar, someothertable)
f = Foo()
- f.bars.append(Bar())
- f.bars.append(Bar())
+ f.bars.set(Bar())
+ f.bars.set(Bar())
sess = create_session()
sess.save(f)
sess.flush()
@@ -794,6 +796,44 @@ class CustomCollectionsTest(testbase.ORMTest):
assert len(list(f.bars)) == 2
f.bars.clear()
+ def testdictwrapper(self):
+ """test that the supplied 'dict' wrapper can be used as a collection and can lazyload."""
+
+ class Foo(object):
+ pass
+ class Bar(object):
+ def __init__(self, data): self.data = data
+
+ mapper(Foo, sometable, properties={
+ 'bars':relation(Bar,
+ collection_class=collections.column_mapped_collection(someothertable.c.data))
+ })
+ mapper(Bar, someothertable)
+
+ f = Foo()
+ col = collections.collection_adapter(f.bars)
+ col.append_with_event(Bar('a'))
+ col.append_with_event(Bar('b'))
+ sess = create_session()
+ sess.save(f)
+ sess.flush()
+ sess.clear()
+ f = sess.query(Foo).get(f.col1)
+ assert len(list(f.bars)) == 2
+
+ existing = set([id(b) for b in f.bars.values()])
+
+ col = collections.collection_adapter(f.bars)
+ col.append_with_event(Bar('b'))
+ f.bars['a'] = Bar('a')
+ sess.flush()
+ sess.clear()
+ f = sess.query(Foo).get(f.col1)
+ assert len(list(f.bars)) == 2
+
+ replaced = set([id(b) for b in f.bars.values()])
+ self.assert_(existing != replaced)
+
def testlist(self):
class Parent(object):
pass
@@ -811,13 +851,13 @@ class CustomCollectionsTest(testbase.ORMTest):
o = Child()
control.append(o)
p.children.append(o)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = [Child(), Child(), Child(), Child()]
control.extend(o)
p.children.extend(o)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
assert control[0] == p.children[0]
@@ -826,92 +866,92 @@ class CustomCollectionsTest(testbase.ORMTest):
del control[1]
del p.children[1]
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = [Child()]
control[1:3] = o
p.children[1:3] = o
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = [Child(), Child(), Child(), Child()]
control[1:3] = o
p.children[1:3] = o
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = [Child(), Child(), Child(), Child()]
control[-1:-2] = o
p.children[-1:-2] = o
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = [Child(), Child(), Child(), Child()]
control[4:] = o
p.children[4:] = o
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = Child()
control.insert(0, o)
p.children.insert(0, o)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = Child()
control.insert(3, o)
p.children.insert(3, o)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = Child()
control.insert(999, o)
p.children.insert(999, o)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
del control[0:1]
del p.children[0:1]
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
del control[1:1]
del p.children[1:1]
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
del control[1:3]
del p.children[1:3]
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
del control[7:]
del p.children[7:]
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
assert control.pop() == p.children.pop()
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
assert control.pop(0) == p.children.pop(0)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
assert control.pop(2) == p.children.pop(2)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
o = Child()
control.insert(2, o)
p.children.insert(2, o)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
control.remove(o)
p.children.remove(o)
- assert control == p.children.data
+ assert control == p.children
assert control == list(p.children)
def testobj(self):
@@ -922,9 +962,12 @@ class CustomCollectionsTest(testbase.ORMTest):
class MyCollection(object):
def __init__(self): self.data = []
+ @collection.appender
def append(self, value): self.data.append(value)
+ @collection.remover
+ def remove(self, value): self.data.remove(value)
+ @collection.iterator
def __iter__(self): return iter(self.data)
- def clear(self): self.data.clear()
mapper(Parent, sometable, properties={
'children':relation(Child, collection_class=MyCollection)
@@ -958,7 +1001,7 @@ class CustomCollectionsTest(testbase.ORMTest):
o = list(p2.children)
assert len(o) == 3
-class ViewOnlyTest(testbase.ORMTest):
+class ViewOnlyTest(ORMTest):
"""test a view_only mapping where a third table is pulled into the primary join condition,
using overlapping PK column names (should not produce "conflicting column" error)"""
def define_tables(self, metadata):
@@ -1009,7 +1052,7 @@ class ViewOnlyTest(testbase.ORMTest):
assert set([x.id for x in c1.t2s]) == set([c2a.id, c2b.id])
assert set([x.id for x in c1.t2_view]) == set([c2b.id])
-class ViewOnlyTest2(testbase.ORMTest):
+class ViewOnlyTest2(ORMTest):
"""test a view_only mapping where a third table is pulled into the primary join condition,
using non-overlapping PK column names (should not produce "mapper has no column X" error)"""
def define_tables(self, metadata):
diff --git a/test/orm/session.py b/test/orm/session.py
index 762722ecc..433279673 100644
--- a/test/orm/session.py
+++ b/test/orm/session.py
@@ -1,13 +1,9 @@
-from testbase import AssertMixin
import testbase
-import unittest, sys, datetime
-
-import tables
-from tables import *
-
-db = testbase.db
from sqlalchemy import *
-
+from sqlalchemy.orm import *
+from testlib import *
+from testlib.tables import *
+import testlib.tables as tables
class SessionTest(AssertMixin):
def setUpAll(self):
@@ -25,7 +21,7 @@ class SessionTest(AssertMixin):
c = testbase.db.connect()
class User(object):pass
mapper(User, users)
- s = create_session(bind_to=c)
+ s = create_session(bind=c)
s.save(User())
s.flush()
c.execute("select * from users")
@@ -38,6 +34,30 @@ class SessionTest(AssertMixin):
s.user_name = 'some other user'
s.flush()
+ def test_close_two(self):
+ c = testbase.db.connect()
+ try:
+ class User(object):pass
+ mapper(User, users)
+ s = create_session(bind=c)
+ s.begin()
+ tran = s.transaction
+ s.save(User())
+ s.flush()
+ c.execute("select * from users")
+ u = User()
+ s.save(u)
+ s.user_name = 'some user'
+ s.flush()
+ u = User()
+ s.save(u)
+ s.user_name = 'some other user'
+ s.flush()
+ assert s.transaction is tran
+ tran.close()
+ finally:
+ c.close()
+
def test_expunge_cascade(self):
tables.data()
mapper(Address, addresses)
@@ -52,49 +72,209 @@ class SessionTest(AssertMixin):
# then see if expunge fails
session.expunge(u)
-
+
+ @testing.unsupported('sqlite')
def test_transaction(self):
class User(object):pass
mapper(User, users)
- sess = create_session()
- transaction = sess.create_transaction()
+ conn1 = testbase.db.connect()
+ conn2 = testbase.db.connect()
+
+ sess = create_session(transactional=True, bind=conn1)
+ u = User()
+ sess.save(u)
+ sess.flush()
+ assert conn1.execute("select count(1) from users").scalar() == 1
+ assert conn2.execute("select count(1) from users").scalar() == 0
+ sess.commit()
+ assert conn1.execute("select count(1) from users").scalar() == 1
+ assert testbase.db.connect().execute("select count(1) from users").scalar() == 1
+
+ @testing.unsupported('sqlite')
+ def test_autoflush(self):
+ class User(object):pass
+ mapper(User, users)
+ conn1 = testbase.db.connect()
+ conn2 = testbase.db.connect()
+
+ sess = create_session(autoflush=True, bind=conn1)
+ u = User()
+ u.user_name='ed'
+ sess.save(u)
+ u2 = sess.query(User).filter_by(user_name='ed').one()
+ assert u2 is u
+ assert conn1.execute("select count(1) from users").scalar() == 1
+ assert conn2.execute("select count(1) from users").scalar() == 0
+ sess.commit()
+ assert conn1.execute("select count(1) from users").scalar() == 1
+ assert testbase.db.connect().execute("select count(1) from users").scalar() == 1
+
+ @testing.unsupported('sqlite')
+ def test_autoflush_unbound(self):
+ class User(object):pass
+ mapper(User, users)
+
try:
+ sess = create_session(autoflush=True)
u = User()
+ u.user_name='ed'
sess.save(u)
+ u2 = sess.query(User).filter_by(user_name='ed').one()
+ assert u2 is u
+ assert sess.execute("select count(1) from users", mapper=User).scalar() == 1
+ assert testbase.db.connect().execute("select count(1) from users").scalar() == 0
+ sess.commit()
+ assert sess.execute("select count(1) from users", mapper=User).scalar() == 1
+ assert testbase.db.connect().execute("select count(1) from users").scalar() == 1
+ except:
+ sess.rollback()
+ raise
+
+ def test_autoflush_2(self):
+ class User(object):pass
+ mapper(User, users)
+ conn1 = testbase.db.connect()
+ conn2 = testbase.db.connect()
+
+ sess = create_session(autoflush=True, bind=conn1)
+ u = User()
+ u.user_name='ed'
+ sess.save(u)
+ sess.commit()
+ assert conn1.execute("select count(1) from users").scalar() == 1
+ assert testbase.db.connect().execute("select count(1) from users").scalar() == 1
+
+ def test_external_joined_transaction(self):
+ class User(object):pass
+ mapper(User, users)
+ conn = testbase.db.connect()
+ trans = conn.begin()
+ sess = create_session(bind=conn)
+ sess.begin()
+ u = User()
+ sess.save(u)
+ sess.flush()
+ sess.commit() # commit does nothing
+ trans.rollback() # rolls back
+ assert len(sess.query(User).select()) == 0
+
+ @testing.supported('postgres', 'mysql')
+ def test_external_nested_transaction(self):
+ class User(object):pass
+ mapper(User, users)
+ try:
+ conn = testbase.db.connect()
+ trans = conn.begin()
+ sess = create_session(bind=conn)
+ u1 = User()
+ sess.save(u1)
sess.flush()
- sess.delete(u)
- sess.save(User())
+
+ sess.begin_nested()
+ u2 = User()
+ sess.save(u2)
sess.flush()
- # TODO: assertion ?
- transaction.commit()
+ sess.rollback()
+
+ trans.commit()
+ assert len(sess.query(User).select()) == 1
except:
- transaction.rollback()
+ conn.close()
+ raise
+
+ @testing.supported('postgres', 'mysql')
+ def test_twophase(self):
+ # TODO: mock up a failure condition here
+ # to ensure a rollback succeeds
+ class User(object):pass
+ class Address(object):pass
+ mapper(User, users)
+ mapper(Address, addresses)
+
+ engine2 = create_engine(testbase.db.url)
+ sess = create_session(twophase=True)
+ sess.bind_mapper(User, testbase.db)
+ sess.bind_mapper(Address, engine2)
+ sess.begin()
+ u1 = User()
+ a1 = Address()
+ sess.save(u1)
+ sess.save(a1)
+ sess.commit()
+ sess.close()
+ engine2.dispose()
+ assert users.count().scalar() == 1
+ assert addresses.count().scalar() == 1
+
+
+
+ def test_joined_transaction(self):
+ class User(object):pass
+ mapper(User, users)
+ sess = create_session()
+ sess.begin()
+ sess.begin()
+ u = User()
+ sess.save(u)
+ sess.flush()
+ sess.commit() # commit does nothing
+ sess.rollback() # rolls back
+ assert len(sess.query(User).select()) == 0
+ @testing.supported('postgres', 'mysql')
def test_nested_transaction(self):
class User(object):pass
mapper(User, users)
sess = create_session()
- transaction = sess.create_transaction()
- trans2 = sess.create_transaction()
+ sess.begin()
+
u = User()
sess.save(u)
sess.flush()
- trans2.commit()
- transaction.rollback()
- assert len(sess.query(User).select()) == 0
+
+ sess.begin_nested() # nested transaction
+
+ u2 = User()
+ sess.save(u2)
+ sess.flush()
+
+ sess.rollback()
+
+ sess.commit()
+ assert len(sess.query(User).select()) == 1
+
+ @testing.supported('postgres', 'mysql')
+ def test_nested_autotrans(self):
+ class User(object):pass
+ mapper(User, users)
+ sess = create_session(transactional=True)
+ u = User()
+ sess.save(u)
+ sess.flush()
+
+ sess.begin_nested() # nested transaction
+
+ u2 = User()
+ sess.save(u2)
+ sess.flush()
+
+ sess.rollback()
+
+ sess.commit()
+ assert len(sess.query(User).select()) == 1
def test_bound_connection(self):
class User(object):pass
mapper(User, users)
c = testbase.db.connect()
sess = create_session(bind=c)
- transaction = sess.create_transaction()
- trans2 = sess.create_transaction()
+ sess.create_transaction()
+ transaction = sess.transaction
u = User()
sess.save(u)
sess.flush()
- assert transaction.get_or_add(testbase.db) is trans2.get_or_add(testbase.db) is transaction.get_or_add(c) is trans2.get_or_add(c) is c
-
+ assert transaction.get_or_add(testbase.db) is transaction.get_or_add(c) is c
+
try:
transaction.add(testbase.db.connect())
assert False
@@ -112,34 +292,10 @@ class SessionTest(AssertMixin):
assert False
except exceptions.InvalidRequestError, e:
assert str(e) == "Session already has a Connection associated for the given Engine"
-
- trans2.commit()
+
transaction.rollback()
assert len(sess.query(User).select()) == 0
-
- def test_close_two(self):
- c = testbase.db.connect()
- try:
- class User(object):pass
- mapper(User, users)
- s = create_session(bind_to=c)
- tran = s.create_transaction()
- s.save(User())
- s.flush()
- c.execute("select * from users")
- u = User()
- s.save(u)
- s.user_name = 'some user'
- s.flush()
- u = User()
- s.save(u)
- s.user_name = 'some other user'
- s.flush()
- assert s.transaction is tran
- tran.close()
- finally:
- c.close()
-
+
def test_update(self):
"""test that the update() method functions and doesnet blow away changes"""
tables.delete()
@@ -164,7 +320,7 @@ class SessionTest(AssertMixin):
user = s.query(User).selectone()
assert user.user_name == 'fred'
- # insure its not dirty if no changes occur
+ # ensure its not dirty if no changes occur
s.clear()
assert user not in s
s.update(user)
diff --git a/test/orm/sessioncontext.py b/test/orm/sessioncontext.py
index 83bc2f2bf..7a60b47c7 100644
--- a/test/orm/sessioncontext.py
+++ b/test/orm/sessioncontext.py
@@ -1,15 +1,15 @@
-from testbase import PersistTest, AssertMixin
-import unittest, sys, os
+import testbase
+from sqlalchemy import *
+from sqlalchemy.orm import *
from sqlalchemy.ext.sessioncontext import SessionContext
from sqlalchemy.orm.session import object_session, Session
-from sqlalchemy import *
-import testbase
+from testlib import *
+
metadata = MetaData()
users = Table('users', metadata,
Column('user_id', Integer, Sequence('user_id_seq', optional=True), primary_key = True),
Column('user_name', String(40)),
- mysql_engine='innodb'
)
class SessionContextTest(AssertMixin):
diff --git a/test/orm/sharding/__init__.py b/test/orm/sharding/__init__.py
new file mode 100644
index 000000000..e69de29bb
--- /dev/null
+++ b/test/orm/sharding/__init__.py
diff --git a/test/orm/sharding/alltests.py b/test/orm/sharding/alltests.py
new file mode 100644
index 000000000..0cdb838a9
--- /dev/null
+++ b/test/orm/sharding/alltests.py
@@ -0,0 +1,18 @@
+import testbase
+import unittest
+
+def suite():
+ modules_to_test = (
+ 'orm.sharding.shard',
+ )
+ alltests = unittest.TestSuite()
+ for name in modules_to_test:
+ mod = __import__(name)
+ for token in name.split('.')[1:]:
+ mod = getattr(mod, token)
+ alltests.addTest(unittest.findTestCases(mod, suiteClass=None))
+ return alltests
+
+
+if __name__ == '__main__':
+ testbase.main(suite())
diff --git a/test/orm/sharding/shard.py b/test/orm/sharding/shard.py
new file mode 100644
index 000000000..faa980cc2
--- /dev/null
+++ b/test/orm/sharding/shard.py
@@ -0,0 +1,154 @@
+import testbase
+from sqlalchemy import *
+from sqlalchemy.orm import *
+
+from sqlalchemy.orm.shard import ShardedSession
+from sqlalchemy.sql import ColumnOperators
+import datetime, operator, os
+from testlib import PersistTest
+
+# TODO: ShardTest can be turned into a base for further subclasses
+
+class ShardTest(PersistTest):
+ def setUpAll(self):
+ global db1, db2, db3, db4, weather_locations, weather_reports
+
+ db1 = create_engine('sqlite:///shard1.db')
+ db2 = create_engine('sqlite:///shard2.db')
+ db3 = create_engine('sqlite:///shard3.db')
+ db4 = create_engine('sqlite:///shard4.db')
+
+ meta = MetaData()
+ ids = Table('ids', meta,
+ Column('nextid', Integer, nullable=False))
+
+ def id_generator(ctx):
+ # in reality, might want to use a separate transaction for this.
+ c = db1.connect()
+ nextid = c.execute(ids.select(for_update=True)).scalar()
+ c.execute(ids.update(values={ids.c.nextid : ids.c.nextid + 1}))
+ return nextid
+
+ weather_locations = Table("weather_locations", meta,
+ Column('id', Integer, primary_key=True, default=id_generator),
+ Column('continent', String(30), nullable=False),
+ Column('city', String(50), nullable=False)
+ )
+
+ weather_reports = Table("weather_reports", meta,
+ Column('id', Integer, primary_key=True),
+ Column('location_id', Integer, ForeignKey('weather_locations.id')),
+ Column('temperature', Float),
+ Column('report_time', DateTime, default=datetime.datetime.now),
+ )
+
+ for db in (db1, db2, db3, db4):
+ meta.create_all(db)
+
+ db1.execute(ids.insert(), nextid=1)
+
+ self.setup_session()
+ self.setup_mappers()
+
+ def tearDownAll(self):
+ for i in range(1,5):
+ os.remove("shard%d.db" % i)
+
+ def setup_session(self):
+ global create_session
+
+ shard_lookup = {
+ 'North America':'north_america',
+ 'Asia':'asia',
+ 'Europe':'europe',
+ 'South America':'south_america'
+ }
+
+ def shard_chooser(mapper, instance):
+ if isinstance(instance, WeatherLocation):
+ return shard_lookup[instance.continent]
+ else:
+ return shard_chooser(mapper, instance.location)
+
+ def id_chooser(ident):
+ return ['north_america', 'asia', 'europe', 'south_america']
+
+ def query_chooser(query):
+ ids = []
+
+ class FindContinent(sql.ClauseVisitor):
+ def visit_binary(self, binary):
+ if binary.left is weather_locations.c.continent:
+ if binary.operator == operator.eq:
+ ids.append(shard_lookup[binary.right.value])
+ elif binary.operator == ColumnOperators.in_op:
+ for bind in binary.right.clauses:
+ ids.append(shard_lookup[bind.value])
+
+ FindContinent().traverse(query._criterion)
+ if len(ids) == 0:
+ return ['north_america', 'asia', 'europe', 'south_america']
+ else:
+ return ids
+
+ def create_session():
+ s = ShardedSession(shard_chooser, id_chooser, query_chooser)
+ s.bind_shard('north_america', db1)
+ s.bind_shard('asia', db2)
+ s.bind_shard('europe', db3)
+ s.bind_shard('south_america', db4)
+ return s
+
+ def setup_mappers(self):
+ global WeatherLocation, Report
+
+ class WeatherLocation(object):
+ def __init__(self, continent, city):
+ self.continent = continent
+ self.city = city
+
+ class Report(object):
+ def __init__(self, temperature):
+ self.temperature = temperature
+
+ mapper(WeatherLocation, weather_locations, properties={
+ 'reports':relation(Report, backref='location')
+ })
+
+ mapper(Report, weather_reports)
+
+ def test_roundtrip(self):
+ tokyo = WeatherLocation('Asia', 'Tokyo')
+ newyork = WeatherLocation('North America', 'New York')
+ toronto = WeatherLocation('North America', 'Toronto')
+ london = WeatherLocation('Europe', 'London')
+ dublin = WeatherLocation('Europe', 'Dublin')
+ brasilia = WeatherLocation('South America', 'Brasila')
+ quito = WeatherLocation('South America', 'Quito')
+
+ tokyo.reports.append(Report(80.0))
+ newyork.reports.append(Report(75))
+ quito.reports.append(Report(85))
+
+ sess = create_session()
+ for c in [tokyo, newyork, toronto, london, dublin, brasilia, quito]:
+ sess.save(c)
+ sess.flush()
+
+ sess.clear()
+
+ t = sess.query(WeatherLocation).get(tokyo.id)
+ assert t.city == tokyo.city
+ assert t.reports[0].temperature == 80.0
+
+ north_american_cities = sess.query(WeatherLocation).filter(WeatherLocation.continent == 'North America')
+ assert set([c.city for c in north_american_cities]) == set(['New York', 'Toronto'])
+
+ asia_and_europe = sess.query(WeatherLocation).filter(WeatherLocation.continent.in_('Europe', 'Asia'))
+ assert set([c.city for c in asia_and_europe]) == set(['Tokyo', 'London', 'Dublin'])
+
+
+
+if __name__ == '__main__':
+ testbase.main()
+ \ No newline at end of file
diff --git a/test/orm/unitofwork.py b/test/orm/unitofwork.py
index 6ba3f8c4b..ae626db84 100644
--- a/test/orm/unitofwork.py
+++ b/test/orm/unitofwork.py
@@ -1,13 +1,14 @@
-from testbase import PersistTest, AssertMixin
-from sqlalchemy import *
import testbase
import pickleable
+from sqlalchemy import *
+from sqlalchemy.orm import *
from sqlalchemy.orm.mapper import global_extensions
from sqlalchemy.orm import util as ormutil
from sqlalchemy.ext.sessioncontext import SessionContext
import sqlalchemy.ext.assignmapper as assignmapper
-from tables import *
-import tables
+from testlib import *
+from testlib.tables import *
+from testlib import tables
"""tests unitofwork operations"""
@@ -26,6 +27,7 @@ class UnitOfWorkTest(AssertMixin):
class HistoryTest(UnitOfWorkTest):
def setUpAll(self):
+ tables.metadata.bind = testbase.db
UnitOfWorkTest.setUpAll(self)
users.create()
addresses.create()
@@ -61,7 +63,7 @@ class VersioningTest(UnitOfWorkTest):
UnitOfWorkTest.setUpAll(self)
ctx.current.clear()
global version_table
- version_table = Table('version_test', db,
+ version_table = Table('version_test', MetaData(testbase.db),
Column('id', Integer, Sequence('version_test_seq'), primary_key=True ),
Column('version_id', Integer, nullable=False),
Column('value', String(40), nullable=False)
@@ -253,9 +255,9 @@ class MutableTypesTest(UnitOfWorkTest):
ctx.current.flush()
def go():
ctx.current.flush()
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
f1.value = unicode('someothervalue')
- self.assert_sql(db, lambda: ctx.current.flush(), [
+ self.assert_sql(testbase.db, lambda: ctx.current.flush(), [
(
"UPDATE mutabletest SET value=:value WHERE mutabletest.id = :mutabletest_id",
{'mutabletest_id': f1.id, 'value': u'someothervalue'}
@@ -263,7 +265,7 @@ class MutableTypesTest(UnitOfWorkTest):
])
f1.value = unicode('hi')
f1.data.x = 9
- self.assert_sql(db, lambda: ctx.current.flush(), [
+ self.assert_sql(testbase.db, lambda: ctx.current.flush(), [
(
"UPDATE mutabletest SET data=:data, value=:value WHERE mutabletest.id = :mutabletest_id",
{'mutabletest_id': f1.id, 'value': u'hi', 'data':f1.data}
@@ -281,7 +283,7 @@ class MutableTypesTest(UnitOfWorkTest):
def go():
ctx.current.flush()
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
ctx.current.clear()
@@ -289,12 +291,12 @@ class MutableTypesTest(UnitOfWorkTest):
def go():
ctx.current.flush()
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
f2.data.y = 19
def go():
ctx.current.flush()
- self.assert_sql_count(db, go, 1)
+ self.assert_sql_count(testbase.db, go, 1)
ctx.current.clear()
f3 = ctx.current.query(Foo).get_by(id=f1.id)
@@ -303,7 +305,7 @@ class MutableTypesTest(UnitOfWorkTest):
def go():
ctx.current.flush()
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
def testunicode(self):
"""test that two equivalent unicode values dont get flagged as changed.
@@ -320,47 +322,42 @@ class MutableTypesTest(UnitOfWorkTest):
f1.value = u'hi'
def go():
ctx.current.flush()
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
class PKTest(UnitOfWorkTest):
def setUpAll(self):
UnitOfWorkTest.setUpAll(self)
- global table
- global table2
- global table3
+ global table, table2, table3, metadata
+ metadata = MetaData(testbase.db)
table = Table(
- 'multipk', db,
+ 'multipk', metadata,
Column('multi_id', Integer, Sequence("multi_id_seq", optional=True), primary_key=True),
Column('multi_rev', Integer, primary_key=True),
Column('name', String(50), nullable=False),
Column('value', String(100))
)
- table2 = Table('multipk2', db,
+ table2 = Table('multipk2', metadata,
Column('pk_col_1', String(30), primary_key=True),
Column('pk_col_2', String(30), primary_key=True),
Column('data', String(30), )
)
- table3 = Table('multipk3', db,
+ table3 = Table('multipk3', metadata,
Column('pri_code', String(30), key='primary', primary_key=True),
Column('sec_code', String(30), key='secondary', primary_key=True),
Column('date_assigned', Date, key='assigned', primary_key=True),
Column('data', String(30), )
)
- table.create()
- table2.create()
- table3.create()
+ metadata.create_all()
def tearDownAll(self):
- table.drop()
- table2.drop()
- table3.drop()
+ metadata.drop_all()
UnitOfWorkTest.tearDownAll(self)
# not support on sqlite since sqlite's auto-pk generation only works with
# single column primary keys
- @testbase.unsupported('sqlite')
+ @testing.unsupported('sqlite')
def testprimarykey(self):
class Entry(object):
pass
@@ -448,7 +445,7 @@ class ForeignPKTest(UnitOfWorkTest):
},
)
- assert list(m2.props['sites'].foreign_keys) == [peoplesites.c.person]
+ assert list(m2.get_property('sites').foreign_keys) == [peoplesites.c.person]
p = Person()
p.person = 'im the key'
p.firstname = 'asdf'
@@ -466,7 +463,7 @@ class PassiveDeletesTest(UnitOfWorkTest):
mytable = Table('mytable', metadata,
Column('id', Integer, primary_key=True),
Column('data', String(30)),
- mysql_engine='InnoDB'
+ test_needs_fk=True,
)
myothertable = Table('myothertable', metadata,
@@ -474,7 +471,7 @@ class PassiveDeletesTest(UnitOfWorkTest):
Column('parent_id', Integer),
Column('data', String(30)),
ForeignKeyConstraint(['parent_id'],['mytable.id'], ondelete="CASCADE"),
- mysql_engine='InnoDB'
+ test_needs_fk=True,
)
metadata.create_all()
@@ -482,7 +479,7 @@ class PassiveDeletesTest(UnitOfWorkTest):
metadata.drop_all()
UnitOfWorkTest.tearDownAll(self)
- @testbase.unsupported('sqlite')
+ @testing.unsupported('sqlite')
def testbasic(self):
class MyClass(object):
pass
@@ -519,6 +516,7 @@ class DefaultTest(UnitOfWorkTest):
defaults back from the engine."""
def setUpAll(self):
UnitOfWorkTest.setUpAll(self)
+ db = testbase.db
use_string_defaults = db.engine.__module__.endswith('postgres') or db.engine.__module__.endswith('oracle') or db.engine.__module__.endswith('sqlite')
if use_string_defaults:
@@ -529,21 +527,21 @@ class DefaultTest(UnitOfWorkTest):
hohotype = Integer
self.hohoval = 9
self.althohoval = 15
- self.table = Table('default_test', db,
+ global default_table
+ metadata = MetaData(db)
+ default_table = Table('default_test', metadata,
Column('id', Integer, Sequence("dt_seq", optional=True), primary_key=True),
Column('hoho', hohotype, PassiveDefault(str(self.hohoval))),
Column('counter', Integer, PassiveDefault("7")),
Column('foober', String(30), default="im foober", onupdate="im the update")
)
- self.table.create()
+ default_table.create()
def tearDownAll(self):
- self.table.drop()
+ default_table.drop()
UnitOfWorkTest.tearDownAll(self)
- def setUp(self):
- self.table = Table('default_test', db)
def testinsert(self):
class Hoho(object):pass
- assign_mapper(Hoho, self.table)
+ assign_mapper(Hoho, default_table)
h1 = Hoho(hoho=self.althohoval)
h2 = Hoho(counter=12)
h3 = Hoho(hoho=self.althohoval, counter=12)
@@ -571,7 +569,7 @@ class DefaultTest(UnitOfWorkTest):
def testinsertnopostfetch(self):
# populates the PassiveDefaults explicitly so there is no "post-update"
class Hoho(object):pass
- assign_mapper(Hoho, self.table)
+ assign_mapper(Hoho, default_table)
h1 = Hoho(hoho="15", counter="15")
ctx.current.flush()
self.assert_(h1.hoho=="15")
@@ -580,7 +578,7 @@ class DefaultTest(UnitOfWorkTest):
def testupdate(self):
class Hoho(object):pass
- assign_mapper(Hoho, self.table)
+ assign_mapper(Hoho, default_table)
h1 = Hoho()
ctx.current.flush()
self.assert_(h1.foober == 'im foober')
@@ -613,8 +611,7 @@ class OneToManyTest(UnitOfWorkTest):
a2 = Address()
a2.email_address = 'lala@test.org'
u.addresses.append(a2)
- self.echo( repr(u.addresses))
- self.echo( repr(u.addresses.added_items()))
+ print repr(u.addresses)
ctx.current.flush()
usertable = users.select(users.c.user_id.in_(u.user_id)).execute().fetchall()
@@ -664,7 +661,7 @@ class OneToManyTest(UnitOfWorkTest):
u2.user_name = 'user2modified'
u1.addresses.append(a3)
del u1.addresses[0]
- self.assert_sql(db, lambda: ctx.current.flush(),
+ self.assert_sql(testbase.db, lambda: ctx.current.flush(),
[
(
"UPDATE users SET user_name=:user_name WHERE users.user_id = :users_user_id",
@@ -836,7 +833,7 @@ class SaveTest(UnitOfWorkTest):
# assert the first one retreives the same from the identity map
nu = ctx.current.get(m, u.user_id)
- self.echo( "U: " + repr(u) + "NU: " + repr(nu))
+ print "U: " + repr(u) + "NU: " + repr(nu)
self.assert_(u is nu)
# clear out the identity map, so next get forces a SELECT
@@ -917,7 +914,7 @@ class SaveTest(UnitOfWorkTest):
u.user_name = ""
def go():
ctx.current.flush()
- self.assert_sql_count(db, go, 0)
+ self.assert_sql_count(testbase.db, go, 0)
def testmultitable(self):
"""tests a save of an object where each instance spans two tables. also tests
@@ -935,8 +932,7 @@ class SaveTest(UnitOfWorkTest):
u.email = 'multi@test.org'
ctx.current.flush()
- id = m.identity(u)
- print id
+ id = m.primary_key_from_instance(u)
ctx.current.clear()
@@ -1043,7 +1039,7 @@ class ManyToOneTest(UnitOfWorkTest):
objects[2].email_address = 'imnew@foo.bar'
objects[3].user = User()
objects[3].user.user_name = 'imnewlyadded'
- self.assert_sql(db, lambda: ctx.current.flush(), [
+ self.assert_sql(testbase.db, lambda: ctx.current.flush(), [
(
"INSERT INTO users (user_name) VALUES (:user_name)",
{'user_name': 'imnewlyadded'}
@@ -1213,7 +1209,7 @@ class ManyToManyTest(UnitOfWorkTest):
k = Keyword()
k.name = 'yellow'
objects[5].keywords.append(k)
- self.assert_sql(db, lambda:ctx.current.flush(), [
+ self.assert_sql(testbase.db, lambda:ctx.current.flush(), [
{
"UPDATE items SET item_name=:item_name WHERE items.item_id = :items_item_id":
{'item_name': 'item4updated', 'items_item_id': objects[4].item_id}
@@ -1242,7 +1238,7 @@ class ManyToManyTest(UnitOfWorkTest):
objects[2].keywords.append(k)
dkid = objects[5].keywords[1].keyword_id
del objects[5].keywords[1]
- self.assert_sql(db, lambda:ctx.current.flush(), [
+ self.assert_sql(testbase.db, lambda:ctx.current.flush(), [
(
"DELETE FROM itemkeywords WHERE itemkeywords.item_id = :item_id AND itemkeywords.keyword_id = :keyword_id",
[{'item_id': objects[5].item_id, 'keyword_id': dkid}]
@@ -1412,7 +1408,6 @@ class ManyToManyTest(UnitOfWorkTest):
k.user_name = 'keyworduser'
k.keyword_name = 'a keyword'
ctx.current.flush()
- print m.instance_key(k)
id = (k.user_id, k.keyword_id)
ctx.current.clear()
@@ -1427,7 +1422,7 @@ class SaveTest2(UnitOfWorkTest):
ctx.current.clear()
clear_mappers()
global meta, users, addresses
- meta = MetaData(db)
+ meta = MetaData(testbase.db)
users = Table('users', meta,
Column('user_id', Integer, Sequence('user_id_seq', optional=True), primary_key = True),
Column('user_name', String(20)),
@@ -1459,7 +1454,7 @@ class SaveTest2(UnitOfWorkTest):
a.user = User()
a.user.user_name = elem['user_name']
objects.append(a)
- self.assert_sql(db, lambda: ctx.current.flush(), [
+ self.assert_sql(testbase.db, lambda: ctx.current.flush(), [
(
"INSERT INTO users (user_name) VALUES (:user_name)",
{'user_name': 'thesub'}
@@ -1498,30 +1493,32 @@ class SaveTest2(UnitOfWorkTest):
]
)
-class SaveTest3(UnitOfWorkTest):
+class SaveTest3(UnitOfWorkTest):
def setUpAll(self):
+ global st3_metadata, t1, t2, t3
+
UnitOfWorkTest.setUpAll(self)
- global metadata, t1, t2, t3
- metadata = testbase.metadata
- t1 = Table('items', metadata,
+
+ st3_metadata = MetaData(testbase.db)
+ t1 = Table('items', st3_metadata,
Column('item_id', INT, Sequence('items_id_seq', optional=True), primary_key = True),
Column('item_name', VARCHAR(50)),
)
- t3 = Table('keywords', metadata,
+ t3 = Table('keywords', st3_metadata,
Column('keyword_id', Integer, Sequence('keyword_id_seq', optional=True), primary_key = True),
Column('name', VARCHAR(50)),
)
- t2 = Table('assoc', metadata,
+ t2 = Table('assoc', st3_metadata,
Column('item_id', INT, ForeignKey("items")),
Column('keyword_id', INT, ForeignKey("keywords")),
Column('foo', Boolean, default=True)
)
- metadata.create_all()
+ st3_metadata.create_all()
def tearDownAll(self):
- metadata.drop_all()
+ st3_metadata.drop_all()
UnitOfWorkTest.tearDownAll(self)
def setUp(self):
diff --git a/test/perf/cascade_speed.py b/test/perf/cascade_speed.py
index d2e741442..34d046381 100644
--- a/test/perf/cascade_speed.py
+++ b/test/perf/cascade_speed.py
@@ -1,5 +1,7 @@
import testbase
from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
from timeit import Timer
import sys
diff --git a/test/perf/masscreate.py b/test/perf/masscreate.py
index e603e2c00..346a725e3 100644
--- a/test/perf/masscreate.py
+++ b/test/perf/masscreate.py
@@ -1,8 +1,7 @@
# times how long it takes to create 26000 objects
-import sys
-sys.path.insert(0, './lib/')
+import testbase
-from sqlalchemy.attributes import *
+from sqlalchemy.orm.attributes import *
import time
import gc
diff --git a/test/perf/masscreate2.py b/test/perf/masscreate2.py
index 3a68f3612..2e29a6327 100644
--- a/test/perf/masscreate2.py
+++ b/test/perf/masscreate2.py
@@ -1,11 +1,9 @@
-import sys
-sys.path.insert(0, './lib/')
-
+import testbase
import gc
import random, string
-from sqlalchemy.attributes import *
+from sqlalchemy.orm.attributes import *
# with this test, run top. make sure the Python process doenst grow in size arbitrarily.
diff --git a/test/perf/masseagerload.py b/test/perf/masseagerload.py
index 9d77fed54..f1c0f292b 100644
--- a/test/perf/masseagerload.py
+++ b/test/perf/masseagerload.py
@@ -1,62 +1,54 @@
-from testbase import PersistTest, AssertMixin
-import unittest, sys, os
-from sqlalchemy import *
-import StringIO
import testbase
-import gc
-import time
-import hotshot
-import hotshot.stats
-
-db = testbase.db
+import hotshot, hotshot.stats
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
NUM = 500
DIVISOR = 50
-class LoadTest(AssertMixin):
- def setUpAll(self):
- global items, meta,subitems
- meta = MetaData(db)
- items = Table('items', meta,
- Column('item_id', Integer, primary_key=True),
- Column('value', String(100)))
- subitems = Table('subitems', meta,
- Column('sub_id', Integer, primary_key=True),
- Column('parent_id', Integer, ForeignKey('items.item_id')),
- Column('value', String(100)))
- meta.create_all()
- def tearDownAll(self):
- meta.drop_all()
- def setUp(self):
- clear_mappers()
+meta = MetaData(testbase.db)
+items = Table('items', meta,
+ Column('item_id', Integer, primary_key=True),
+ Column('value', String(100)))
+subitems = Table('subitems', meta,
+ Column('sub_id', Integer, primary_key=True),
+ Column('parent_id', Integer, ForeignKey('items.item_id')),
+ Column('value', String(100)))
+
+class Item(object):pass
+class SubItem(object):pass
+mapper(Item, items, properties={'subs':relation(SubItem, lazy=False)})
+mapper(SubItem, subitems)
+
+def load():
+ global l
+ l = []
+ for x in range(1,NUM/DIVISOR + 1):
+ l.append({'item_id':x, 'value':'this is item #%d' % x})
+ #print l
+ items.insert().execute(*l)
+ for x in range(1, NUM/DIVISOR + 1):
l = []
- for x in range(1,NUM/DIVISOR + 1):
- l.append({'item_id':x, 'value':'this is item #%d' % x})
+ for y in range(1, DIVISOR + 1):
+ z = ((x-1) * DIVISOR) + y
+ l.append({'sub_id':z,'value':'this is item #%d' % z, 'parent_id':x})
#print l
- items.insert().execute(*l)
- for x in range(1, NUM/DIVISOR + 1):
- l = []
- for y in range(1, DIVISOR + 1):
- z = ((x-1) * DIVISOR) + y
- l.append({'sub_id':z,'value':'this is iteim #%d' % z, 'parent_id':x})
- #print l
- subitems.insert().execute(*l)
- def testload(self):
- class Item(object):pass
- class SubItem(object):pass
- mapper(Item, items, properties={'subs':relation(SubItem, lazy=False)})
- mapper(SubItem, subitems)
- sess = create_session()
- prof = hotshot.Profile("masseagerload.prof")
- prof.start()
- query = sess.query(Item)
- l = query.select()
- print "loaded ", len(l), " items each with ", len(l[0].subs), "subitems"
- prof.stop()
- prof.close()
- stats = hotshot.stats.load("masseagerload.prof")
- stats.sort_stats('time', 'calls')
- stats.print_stats()
-
-if __name__ == "__main__":
- testbase.main()
+ subitems.insert().execute(*l)
+
+@profiling.profiled('masseagerload', always=True)
+def masseagerload(session):
+ query = session.query(Item)
+ l = query.select()
+ print "loaded ", len(l), " items each with ", len(l[0].subs), "subitems"
+
+def all():
+ meta.create_all()
+ try:
+ load()
+ masseagerload(create_session())
+ finally:
+ meta.drop_all()
+
+if __name__ == '__main__':
+ all()
diff --git a/test/perf/massload.py b/test/perf/massload.py
index 3530e4a65..92cf0fe92 100644
--- a/test/perf/massload.py
+++ b/test/perf/massload.py
@@ -1,13 +1,10 @@
-from testbase import PersistTest, AssertMixin
-import unittest, sys, os
-from sqlalchemy import *
-import sqlalchemy.orm.attributes as attributes
-import StringIO
import testbase
-import gc
import time
-
-db = testbase.db
+#import gc
+#import sqlalchemy.orm.attributes as attributes
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
NUM = 2500
@@ -20,7 +17,7 @@ for best results, dont run with sqlite :memory: database, and keep an eye on top
class LoadTest(AssertMixin):
def setUpAll(self):
global items, meta
- meta = MetaData(db)
+ meta = MetaData(testbase.db)
items = Table('items', meta,
Column('item_id', Integer, primary_key=True),
Column('value', String(100)))
@@ -28,8 +25,6 @@ class LoadTest(AssertMixin):
def tearDownAll(self):
items.drop()
def setUp(self):
- objectstore.clear()
- clear_mappers()
for x in range(1,NUM/500+1):
l = []
for y in range(x*500-500 + 1, x*500 + 1):
diff --git a/test/perf/massload2.py b/test/perf/massload2.py
index 1506ca503..d6424eb07 100644
--- a/test/perf/massload2.py
+++ b/test/perf/massload2.py
@@ -7,6 +7,7 @@ try:
except:
pass
from sqlalchemy import *
+from testbase import Table, Column
import time
metadata = create_engine('sqlite://', echo=True)
diff --git a/test/perf/masssave.py b/test/perf/masssave.py
index 5690eac3f..dd03f3962 100644
--- a/test/perf/masssave.py
+++ b/test/perf/masssave.py
@@ -1,20 +1,16 @@
-from testbase import PersistTest, AssertMixin
-import unittest, sys, os
-from sqlalchemy import *
-import sqlalchemy.attributes as attributes
-import StringIO
import testbase
-import gc
-import sqlalchemy.orm.session
import types
-db = testbase.db
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
+
NUM = 250000
class SaveTest(AssertMixin):
def setUpAll(self):
global items, metadata
- metadata = MetaData(db)
+ metadata = MetaData(testbase.db)
items = Table('items', metadata,
Column('item_id', Integer, primary_key=True),
Column('value', String(100)))
diff --git a/test/perf/ormsession.py b/test/perf/ormsession.py
new file mode 100644
index 000000000..a9d310ef6
--- /dev/null
+++ b/test/perf/ormsession.py
@@ -0,0 +1,225 @@
+import testbase
+import time
+from datetime import datetime
+
+from sqlalchemy import *
+from sqlalchemy.orm import *
+from testlib import *
+from testlib.profiling import profiled
+
+class Item(object):
+ def __repr__(self):
+ return 'Item<#%s "%s">' % (self.id, self.name)
+class SubItem(object):
+ def __repr__(self):
+ return 'SubItem<#%s "%s">' % (self.id, self.name)
+class Customer(object):
+ def __repr__(self):
+ return 'Customer<#%s "%s">' % (self.id, self.name)
+class Purchase(object):
+ def __repr__(self):
+ return 'Purchase<#%s "%s">' % (self.id, self.purchase_date)
+
+items, subitems, customers, purchases, purchaseitems = \
+ None, None, None, None, None
+
+metadata = MetaData()
+
+@profiled('table')
+def define_tables():
+ global items, subitems, customers, purchases, purchaseitems
+ items = Table('items', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('name', String(100)),
+ test_needs_acid=True)
+ subitems = Table('subitems', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('item_id', Integer, ForeignKey('items.id'),
+ nullable=False),
+ Column('name', String(100), PassiveDefault('no name')),
+ test_needs_acid=True)
+ customers = Table('customers', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('name', String(100)),
+ *[Column("col_%s" % chr(i), String(64), default=str(i))
+ for i in range(97,117)],
+ **dict(test_needs_acid=True))
+ purchases = Table('purchases', metadata,
+ Column('id', Integer, primary_key=True),
+ Column('customer_id', Integer,
+ ForeignKey('customers.id'), nullable=False),
+ Column('purchase_date', DateTime,
+ default=datetime.now),
+ test_needs_acid=True)
+ purchaseitems = Table('purchaseitems', metadata,
+ Column('purchase_id', Integer,
+ ForeignKey('purchases.id'),
+ nullable=False, primary_key=True),
+ Column('item_id', Integer, ForeignKey('items.id'),
+ nullable=False, primary_key=True),
+ test_needs_acid=True)
+
+@profiled('mapper')
+def setup_mappers():
+ mapper(Item, items, properties={
+ 'subitems': relation(SubItem, backref='item', lazy=True)
+ })
+ mapper(SubItem, subitems)
+ mapper(Customer, customers, properties={
+ 'purchases': relation(Purchase, lazy=True, backref='customer')
+ })
+ mapper(Purchase, purchases, properties={
+ 'items': relation(Item, lazy=True, secondary=purchaseitems)
+ })
+
+@profiled('inserts')
+def insert_data():
+ q_items = 1000
+ q_sub_per_item = 10
+ q_customers = 1000
+
+ con = testbase.db.connect()
+
+ transaction = con.begin()
+ data, subdata = [], []
+ for item_id in xrange(1, q_items + 1):
+ data.append({'name': "item number %s" % item_id})
+ for subitem_id in xrange(1, (item_id % q_sub_per_item) + 1):
+ subdata.append({'item_id': item_id,
+ 'name': "subitem number %s" % subitem_id})
+ if item_id % 100 == 0:
+ items.insert().execute(*data)
+ subitems.insert().execute(*subdata)
+ del data[:]
+ del subdata[:]
+ if data:
+ items.insert().execute(*data)
+ if subdata:
+ subitems.insert().execute(*subdata)
+ transaction.commit()
+
+ transaction = con.begin()
+ data = []
+ for customer_id in xrange(1, q_customers):
+ data.append({'name': "customer number %s" % customer_id})
+ if customer_id % 100 == 0:
+ customers.insert().execute(*data)
+ del data[:]
+ if data:
+ customers.insert().execute(*data)
+ transaction.commit()
+
+ transaction = con.begin()
+ data, subdata = [], []
+ order_t = int(time.time()) - (5000 * 5 * 60)
+ current = xrange(1, q_customers)
+ step, purchase_id = 1, 0
+ while current:
+ next = []
+ for customer_id in current:
+ order_t += 300
+ data.append({'customer_id': customer_id,
+ 'purchase_date': datetime.fromtimestamp(order_t)})
+ purchase_id += 1
+ for item_id in range(customer_id % 200, customer_id + 1, 200):
+ if item_id != 0:
+ subdata.append({'purchase_id': purchase_id,
+ 'item_id': item_id})
+ if customer_id % 10 > step:
+ next.append(customer_id)
+
+ if len(data) >= 100:
+ purchases.insert().execute(*data)
+ if subdata:
+ purchaseitems.insert().execute(*subdata)
+ del data[:]
+ del subdata[:]
+ step, current = step + 1, next
+
+ if data:
+ purchases.insert().execute(*data)
+ if subdata:
+ purchaseitems.insert().execute(*subdata)
+ transaction.commit()
+
+@profiled('queries')
+def run_queries():
+ session = create_session()
+ # no explicit transaction here.
+
+ # build a report of summarizing the last 50 purchases and
+ # the top 20 items from all purchases
+
+ q = session.query(Purchase). \
+ limit(50).order_by(desc(Purchase.purchase_date)). \
+ options(eagerload('items'), eagerload('items.subitems'),
+ eagerload('customer'))
+
+ report = []
+ # "write" the report. pretend it's going to a web template or something,
+ # the point is to actually pull data through attributes and collections.
+ for purchase in q:
+ report.append(purchase.customer.name)
+ report.append(purchase.customer.col_a)
+ report.append(purchase.purchase_date)
+ for item in purchase.items:
+ report.append(item.name)
+ report.extend([s.name for s in item.subitems])
+
+ # mix a little low-level with orm
+ # pull a report of the top 20 items of all time
+ _item_id = purchaseitems.c.item_id
+ top_20_q = select([func.distinct(_item_id).label('id')],
+ group_by=[purchaseitems.c.purchase_id, _item_id],
+ order_by=[desc(func.count(_item_id)), _item_id],
+ limit=20)
+ ids = [r.id for r in top_20_q.execute().fetchall()]
+ q2 = session.query(Item).filter(Item.id.in_(*ids))
+
+ for num, item in enumerate(q2):
+ report.append("number %s: %s" % (num + 1, item.name))
+
+@profiled('creating')
+def create_purchase():
+ # commit a purchase
+ customer_id = 100
+ item_ids = (10,22,34,46,58)
+
+ session = create_session()
+ session.begin()
+
+ customer = session.query(Customer).get(customer_id)
+ items = session.query(Item).filter(Item.id.in_(*item_ids))
+
+ purchase = Purchase()
+ purchase.customer = customer
+ purchase.items.extend(items)
+
+ session.flush()
+ session.commit()
+ session.expire(customer)
+
+def setup_db():
+ metadata.drop_all()
+ metadata.create_all()
+def cleanup_db():
+ metadata.drop_all()
+
+@profiled('default')
+def default():
+ run_queries()
+ create_purchase()
+
+@profiled('all')
+def main():
+ metadata.bind = testbase.db
+ try:
+ define_tables()
+ setup_mappers()
+ setup_db()
+ insert_data()
+ default()
+ finally:
+ cleanup_db()
+
+main()
diff --git a/test/perf/poolload.py b/test/perf/poolload.py
index d096f1c67..1a2ff6978 100644
--- a/test/perf/poolload.py
+++ b/test/perf/poolload.py
@@ -1,10 +1,11 @@
# load test of connection pool
+import testbase
from sqlalchemy import *
import sqlalchemy.pool as pool
import thread,time
-db = create_engine('mysql://scott:tiger@127.0.0.1/test', pool_timeout=30, echo_pool=True)
+db = create_engine(testbase.db.url, pool_timeout=30, echo_pool=True)
metadata = MetaData(db)
users_table = Table('users', metadata,
@@ -18,7 +19,7 @@ users_table.insert().execute([{'user_name':'user#%d' % i, 'password':'pw#%d' % i
def runfast():
while True:
- c = db.connection_provider._pool.connect()
+ c = db.pool.connect()
time.sleep(.5)
c.close()
# result = users_table.select(limit=100).execute()
diff --git a/test/perf/threaded_compile.py b/test/perf/threaded_compile.py
index eb9e2f669..13ec31fd6 100644
--- a/test/perf/threaded_compile.py
+++ b/test/perf/threaded_compile.py
@@ -2,9 +2,12 @@
when additional mappers are created while the existing
collection is being compiled."""
+import testbase
from sqlalchemy import *
+from sqlalchemy.orm import *
import thread, time
from sqlalchemy.orm import mapperlib
+from testlib import *
meta = MetaData('sqlite:///foo.db')
diff --git a/test/perf/wsgi.py b/test/perf/wsgi.py
index 365956dc7..d22eeb76a 100644
--- a/test/perf/wsgi.py
+++ b/test/perf/wsgi.py
@@ -1,53 +1,55 @@
#!/usr/bin/python
+"""Uses ``wsgiref``, standard in Python 2.5 and also in the cheeseshop."""
+import testbase
from sqlalchemy import *
-import sqlalchemy.pool as pool
+from sqlalchemy.orm import *
import thread
-from sqlalchemy import exceptions
+from testlib import *
+
+port = 8000
import logging
logging.basicConfig()
logging.getLogger('sqlalchemy.pool').setLevel(logging.INFO)
threadids = set()
-#meta = MetaData('postgres://scott:tiger@127.0.0.1/test')
-
-#meta = MetaData('mysql://scott:tiger@localhost/test', poolclass=pool.SingletonThreadPool)
-meta = MetaData('mysql://scott:tiger@localhost/test')
+meta = MetaData(testbase.db)
foo = Table('foo', meta,
Column('id', Integer, primary_key=True),
Column('data', String(30)))
-
-meta.drop_all()
-meta.create_all()
-
-data = []
-for x in range(1,500):
- data.append({'id':x,'data':"this is x value %d" % x})
-foo.insert().execute(data)
-
class Foo(object):
pass
-
mapper(Foo, foo)
-root = './'
-port = 8000
+def prep():
+ meta.drop_all()
+ meta.create_all()
+
+ data = []
+ for x in range(1,500):
+ data.append({'id':x,'data':"this is x value %d" % x})
+ foo.insert().execute(data)
def serve(environ, start_response):
+ start_response("200 OK", [('Content-type', 'text/plain')])
sess = create_session()
l = sess.query(Foo).select()
-
- start_response("200 OK", [('Content-type','text/plain')])
threadids.add(thread.get_ident())
- print "sending response on thread", thread.get_ident(), " total threads ", len(threadids)
- return ["\n".join([x.data for x in l])]
+
+ print ("sending response on thread", thread.get_ident(),
+ " total threads ", len(threadids))
+ return [str("\n".join([x.data for x in l]))]
if __name__ == '__main__':
- from wsgiutils import wsgiServer
- server = wsgiServer.WSGIServer (('localhost', port), {'/': serve})
- print "Server listening on port %d" % port
- server.serve_forever()
+ from wsgiref import simple_server
+ try:
+ prep()
+ server = simple_server.make_server('localhost', port, serve)
+ print "Server listening on port %d" % port
+ server.serve_forever()
+ finally:
+ meta.drop_all()
diff --git a/test/rundocs.py b/test/rundocs.py
deleted file mode 100644
index 1918e55be..000000000
--- a/test/rundocs.py
+++ /dev/null
@@ -1,242 +0,0 @@
-from sqlalchemy import *
-import sys
-sys.path.insert(0, './lib/')
-
-engine = create_engine('sqlite://')
-
-engine.echo = True
-
-# table metadata
-users = Table('users', engine,
- Column('user_id', Integer, primary_key = True),
- Column('user_name', String(16), nullable = False),
- Column('password', String(20), nullable = False)
-)
-users.create()
-users.insert().execute(
- dict(user_name = 'fred', password='45nfss')
-)
-
-
-# class definition
-class User(object):
- pass
-assign_mapper(User, users)
-
-# select
-user = User.get_by(user_name = 'fred')
-
-# modify
-user.user_name = 'fred jones'
-
-# commit
-objectstore.commit()
-
-objectstore.clear()
-
-
-
-addresses = Table('email_addresses', engine,
- Column('address_id', Integer, primary_key = True),
- Column('user_id', Integer, ForeignKey(users.c.user_id)),
- Column('email_address', String(20)),
-)
-addresses.create()
-addresses.insert().execute(
- dict(user_id = user.user_id, email_address='fred@bar.com')
-)
-
-# second class definition
-class Address(object):
- def __init__(self, email_address = None):
- self.email_address = email_address
-
- mapper = assignmapper(addresses)
-
-# obtain a Mapper. "private=True" means deletions of the user
-# will cascade down to the child Address objects
-User.mapper = assignmapper(users, properties = dict(
- addresses = relation(Address.mapper, lazy=True, private=True)
-))
-
-# select
-user = User.mapper.select(User.c.user_name == 'fred jones')[0]
-address = user.addresses[0]
-
-# modify
-user.user_name = 'fred'
-user.addresses[0].email_address = 'fredjones@foo.com'
-user.addresses.append(Address('freddy@hi.org'))
-
-# commit
-objectstore.commit()
-
-# going to change tables, etc., start over with a new engine
-objectstore.clear()
-engine = None
-engine = sqlite.engine(':memory:', {})
-engine.echo = True
-
-# a table to store a user's preferences for a site
-prefs = Table('user_prefs', engine,
- Column('pref_id', Integer, primary_key = True),
- Column('stylename', String(20)),
- Column('save_password', Boolean, nullable = False),
- Column('timezone', CHAR(3), nullable = False)
-)
-prefs.create()
-prefs.insert().execute(
- dict(pref_id=1, stylename='green', save_password=1, timezone='EST')
-)
-
-# user table gets 'preference_id' column added
-users = Table('users', engine,
- Column('user_id', Integer, primary_key = True),
- Column('user_name', String(16), nullable = False),
- Column('password', String(20), nullable = False),
- Column('preference_id', Integer, ForeignKey(prefs.c.pref_id))
-)
-users.drop()
-users.create()
-users.insert().execute(
- dict(user_name = 'fred', password='45nfss', preference_id=1)
-)
-
-
-addresses = Table('email_addresses', engine,
- Column('address_id', Integer, primary_key = True),
- Column('user_id', Integer, ForeignKey(users.c.user_id)),
- Column('email_address', String(20)),
-)
-addresses.drop()
-addresses.create()
-
-Address.mapper = assignmapper(addresses)
-
-# class definition for preferences
-class UserPrefs(object):
- mapper = assignmapper(prefs)
-
-# set a new Mapper on the user
-User.mapper = assignmapper(users, properties = dict(
- addresses = relation(Address.mapper, lazy=True, private=True),
- preferences = relation(UserPrefs.mapper, lazy=False, private=True),
-))
-
-# select
-user = User.mapper.select(User.c.user_name == 'fred')[0]
-save_password = user.preferences.save_password
-
-# modify
-user.preferences.stylename = 'bluesteel'
-user.addresses.append(Address('freddy@hi.org'))
-
-# commit
-objectstore.commit()
-
-
-
-articles = Table('articles', engine,
- Column('article_id', Integer, primary_key = True),
- Column('article_headline', String(150), key='headline'),
- Column('article_body', CLOB, key='body'),
-)
-
-keywords = Table('keywords', engine,
- Column('keyword_id', Integer, primary_key = True),
- Column('name', String(50))
-)
-
-itemkeywords = Table('article_keywords', engine,
- Column('article_id', Integer, ForeignKey(articles.c.article_id)),
- Column('keyword_id', Integer, ForeignKey(keywords.c.keyword_id))
-)
-
-articles.create()
-keywords.create()
-itemkeywords.create()
-
-# class definitions
-class Keyword(object):
- def __init__(self, name = None):
- self.name = name
- mapper = assignmapper(keywords)
-
-class Article(object):
- def __init__(self):
- self.keywords = []
- mapper = assignmapper(articles, properties = dict(
- keywords = relation(Keyword.mapper, itemkeywords, lazy=False)
- ))
-Article.mapper
-
-article = Article()
-article.headline = 'a headline'
-article.body = 'this is the body'
-article.keywords.append(Keyword('politics'))
-article.keywords.append(Keyword('entertainment'))
-objectstore.commit()
-
-# select articles based on some keywords. the extra selection criterion
-# won't get in the way of the separate eager load of all the article's keywords
-alist = Article.mapper.select(sql.and_(
- keywords.c.keyword_id==itemkeywords.c.keyword_id,
- itemkeywords.c.article_id==articles.c.article_id,
- keywords.c.name.in_('politics', 'entertainment')))
-
-# modify
-a = alist[0]
-del a.keywords[:]
-a.keywords.append(Keyword('topstories'))
-a.keywords.append(Keyword('government'))
-
-# commit. individual INSERT/DELETE operations will take place only for the list
-# elements that changed.
-objectstore.commit()
-
-
-clear_mappers()
-itemkeywords.drop()
-itemkeywords = Table('article_keywords', engine,
- Column('article_id', Integer, ForeignKey("articles.article_id")),
- Column('keyword_id', Integer, ForeignKey("keywords.keyword_id")),
- Column('attached_by', Integer, ForeignKey("users.user_id"))
-, redefine=True)
-itemkeywords.create()
-
-# define an association class
-class KeywordAssociation(object):pass
-
-# define the mapper. when we load an article, we always want to get the keywords via
-# eager loading. but the user who added each keyword, we usually dont need so specify
-# lazy loading for that.
-m = mapper(Article, articles, properties=dict(
- keywords = relation(KeywordAssociation, itemkeywords, lazy = False,
- primary_key=[itemkeywords.c.article_id, itemkeywords.c.keyword_id],
- properties=dict(
- keyword = relation(Keyword, keywords, lazy = False),
- user = relation(User, users, lazy = True)
- )
- )
- )
-)
-
-# bonus step - well, we do want to load the users in one shot,
-# so modify the mapper via an option.
-# this returns a new mapper with the option switched on.
-m2 = m.options(eagerload('keywords.user'))
-
-# select by keyword again
-alist = m2.select(
- sql.and_(
- keywords.c.keyword_id==itemkeywords.c.keyword_id,
- itemkeywords.c.article_id==articles.c.article_id,
- keywords.c.name == 'jacks_stories'
- ))
-
-# user is available
-for a in alist:
- for k in a.keywords:
- if k.keyword.name == 'jacks_stories':
- print k.user.user_name
-
diff --git a/test/sql/alltests.py b/test/sql/alltests.py
index 7be1a3ffb..a669a25f2 100644
--- a/test/sql/alltests.py
+++ b/test/sql/alltests.py
@@ -7,6 +7,8 @@ def suite():
'sql.testtypes',
'sql.constraints',
+ 'sql.generative',
+
# SQL syntax
'sql.select',
'sql.selectable',
@@ -30,7 +32,5 @@ def suite():
alltests.addTest(unittest.findTestCases(mod, suiteClass=None))
return alltests
-
-
if __name__ == '__main__':
testbase.main(suite())
diff --git a/test/sql/case_statement.py b/test/sql/case_statement.py
index 946279b9d..493545b22 100644
--- a/test/sql/case_statement.py
+++ b/test/sql/case_statement.py
@@ -1,13 +1,15 @@
-import sys
import testbase
+import sys
from sqlalchemy import *
+from testlib import *
-class CaseTest(testbase.PersistTest):
+class CaseTest(PersistTest):
def setUpAll(self):
+ metadata = MetaData(testbase.db)
global info_table
- info_table = Table('infos', testbase.db,
+ info_table = Table('infos', metadata,
Column('pk', Integer, primary_key=True),
Column('info', String(30)))
@@ -26,9 +28,9 @@ class CaseTest(testbase.PersistTest):
def testcase(self):
inner = select([case([
[info_table.c.pk < 3,
- literal('lessthan3', type=String)],
+ literal('lessthan3', type_=String)],
[and_(info_table.c.pk >= 3, info_table.c.pk < 7),
- literal('gt3', type=String)]]).label('x'),
+ literal('gt3', type_=String)]]).label('x'),
info_table.c.pk, info_table.c.info],
from_obj=[info_table]).alias('q_inner')
@@ -65,9 +67,9 @@ class CaseTest(testbase.PersistTest):
w_else = select([case([
[info_table.c.pk < 3,
- literal(3, type=Integer)],
+ literal(3, type_=Integer)],
[and_(info_table.c.pk >= 3, info_table.c.pk < 6),
- literal(6, type=Integer)]],
+ literal(6, type_=Integer)]],
else_ = 0).label('x'),
info_table.c.pk, info_table.c.info],
from_obj=[info_table]).alias('q_inner')
diff --git a/test/sql/constraints.py b/test/sql/constraints.py
index 7e1172850..3120185d5 100644
--- a/test/sql/constraints.py
+++ b/test/sql/constraints.py
@@ -1,8 +1,8 @@
import testbase
from sqlalchemy import *
-import sys
+from testlib import *
-class ConstraintTest(testbase.AssertMixin):
+class ConstraintTest(AssertMixin):
def setUp(self):
global metadata
@@ -52,7 +52,7 @@ class ConstraintTest(testbase.AssertMixin):
)
metadata.create_all()
- @testbase.unsupported('mysql')
+ @testing.unsupported('mysql')
def test_check_constraint(self):
foo = Table('foo', metadata,
Column('id', Integer, primary_key=True),
@@ -172,12 +172,13 @@ class ConstraintTest(testbase.AssertMixin):
capt = []
connection = testbase.db.connect()
- ex = connection._execute
+ # TODO: hacky, put a real connection proxy in
+ ex = connection._Connection__execute
def proxy(context):
capt.append(context.statement)
capt.append(repr(context.parameters))
ex(context)
- connection._execute = proxy
+ connection._Connection__execute = proxy
schemagen = testbase.db.dialect.schemagenerator(connection)
schemagen.traverse(events)
diff --git a/test/sql/defaults.py b/test/sql/defaults.py
index 10a3610f9..6c200232f 100644
--- a/test/sql/defaults.py
+++ b/test/sql/defaults.py
@@ -1,52 +1,59 @@
-from testbase import PersistTest
-import sqlalchemy.util as util
-import unittest, sys, os
-import sqlalchemy.schema as schema
import testbase
from sqlalchemy import *
-import sqlalchemy
-
-db = testbase.db
+import sqlalchemy.util as util
+import sqlalchemy.schema as schema
+from sqlalchemy.orm import mapper, create_session
+from testlib import *
+import datetime
class DefaultTest(PersistTest):
def setUpAll(self):
- global t, f, f2, ts, currenttime
+ global t, f, f2, ts, currenttime, metadata
+
+ db = testbase.db
+ metadata = MetaData(db)
x = {'x':50}
def mydefault():
x['x'] += 1
return x['x']
+ def mydefault_with_ctx(ctx):
+ return ctx.compiled_parameters['col1'] + 10
+
+ def myupdate_with_ctx(ctx):
+ return len(ctx.compiled_parameters['col2'])
+
use_function_defaults = db.engine.name == 'postgres' or db.engine.name == 'oracle'
is_oracle = db.engine.name == 'oracle'
# select "count(1)" returns different results on different DBs
# also correct for "current_date" compatible as column default, value differences
- currenttime = func.current_date(type=Date, engine=db);
+ currenttime = func.current_date(type_=Date, bind=db);
if is_oracle:
ts = db.func.trunc(func.sysdate(), literal_column("'DAY'")).scalar()
- f = select([func.count(1) + 5], engine=db).scalar()
- f2 = select([func.count(1) + 14], engine=db).scalar()
+ f = select([func.count(1) + 5], bind=db).scalar()
+ f2 = select([func.count(1) + 14], bind=db).scalar()
# TODO: engine propigation across nested functions not working
- currenttime = func.trunc(currenttime, literal_column("'DAY'"), engine=db)
+ currenttime = func.trunc(currenttime, literal_column("'DAY'"), bind=db)
def1 = currenttime
def2 = func.trunc(text("sysdate"), literal_column("'DAY'"))
deftype = Date
elif use_function_defaults:
- f = select([func.count(1) + 5], engine=db).scalar()
- f2 = select([func.count(1) + 14], engine=db).scalar()
+ f = select([func.count(1) + 5], bind=db).scalar()
+ f2 = select([func.count(1) + 14], bind=db).scalar()
def1 = currenttime
def2 = text("current_date")
deftype = Date
ts = db.func.current_date().scalar()
else:
- f = select([func.count(1) + 5], engine=db).scalar()
- f2 = select([func.count(1) + 14], engine=db).scalar()
+ f = select([func.count(1) + 5], bind=db).scalar()
+ f2 = select([func.count(1) + 14], bind=db).scalar()
def1 = def2 = "3"
ts = 3
deftype = Integer
- t = Table('default_test1', db,
+ t = Table('default_test1', metadata,
# python function
Column('col1', Integer, primary_key=True, default=mydefault),
@@ -66,7 +73,13 @@ class DefaultTest(PersistTest):
Column('col6', Date, default=currenttime, onupdate=currenttime),
Column('boolcol1', Boolean, default=True),
- Column('boolcol2', Boolean, default=False)
+ Column('boolcol2', Boolean, default=False),
+
+ # python function which uses ExecutionContext
+ Column('col7', Integer, default=mydefault_with_ctx, onupdate=myupdate_with_ctx),
+
+ # python builtin
+ Column('col8', Date, default=datetime.date.today, onupdate=datetime.date.today)
)
t.create()
@@ -75,9 +88,18 @@ class DefaultTest(PersistTest):
def tearDown(self):
t.delete().execute()
-
+
+ def testargsignature(self):
+ def mydefault(x, y):
+ pass
+ try:
+ c = ColumnDefault(mydefault)
+ assert False
+ except exceptions.ArgumentError, e:
+ assert str(e) == "ColumnDefault Python function takes zero or one positional arguments", str(e)
+
def teststandalone(self):
- c = db.engine.contextual_connect()
+ c = testbase.db.engine.contextual_connect()
x = c.execute(t.c.col1.default)
y = t.c.col2.default.execute()
z = c.execute(t.c.col3.default)
@@ -94,9 +116,10 @@ class DefaultTest(PersistTest):
t.insert().execute()
ctexec = currenttime.scalar()
- self.echo("Currenttime "+ repr(ctexec))
+ print "Currenttime "+ repr(ctexec)
l = t.select().execute()
- self.assert_(l.fetchall() == [(51, 'imthedefault', f, ts, ts, ctexec, True, False), (52, 'imthedefault', f, ts, ts, ctexec, True, False), (53, 'imthedefault', f, ts, ts, ctexec, True, False)])
+ today = datetime.date.today()
+ self.assert_(l.fetchall() == [(51, 'imthedefault', f, ts, ts, ctexec, True, False, 61, today), (52, 'imthedefault', f, ts, ts, ctexec, True, False, 62, today), (53, 'imthedefault', f, ts, ts, ctexec, True, False, 63, today)])
def testinsertvalues(self):
t.insert(values={'col3':50}).execute()
@@ -109,10 +132,10 @@ class DefaultTest(PersistTest):
pk = r.last_inserted_ids()[0]
t.update(t.c.col1==pk).execute(col4=None, col5=None)
ctexec = currenttime.scalar()
- self.echo("Currenttime "+ repr(ctexec))
+ print "Currenttime "+ repr(ctexec)
l = t.select(t.c.col1==pk).execute()
l = l.fetchone()
- self.assert_(l == (pk, 'im the update', f2, None, None, ctexec, True, False))
+ self.assert_(l == (pk, 'im the update', f2, None, None, ctexec, True, False, 13, datetime.date.today()))
# mysql/other db's return 0 or 1 for count(1)
self.assert_(14 <= f2 <= 15)
@@ -124,8 +147,35 @@ class DefaultTest(PersistTest):
l = l.fetchone()
self.assert_(l['col3'] == 55)
+ @testing.supported('postgres')
+ def testpassiveoverride(self):
+ """primarily for postgres, tests that when we get a primary key column back
+ from reflecting a table which has a default value on it, we pre-execute
+ that PassiveDefault upon insert, even though PassiveDefault says
+ "let the database execute this", because in postgres we must have all the primary
+ key values in memory before insert; otherwise we cant locate the just inserted row."""
+
+ try:
+ meta = MetaData(testbase.db)
+ testbase.db.execute("""
+ CREATE TABLE speedy_users
+ (
+ speedy_user_id SERIAL PRIMARY KEY,
+
+ user_name VARCHAR NOT NULL,
+ user_password VARCHAR NOT NULL
+ );
+ """, None)
+
+ t = Table("speedy_users", meta, autoload=True)
+ t.insert().execute(user_name='user', user_password='lala')
+ l = t.select().execute().fetchall()
+ self.assert_(l == [(1, 'user', 'lala')])
+ finally:
+ testbase.db.execute("drop table speedy_users", None)
+
class AutoIncrementTest(PersistTest):
- @testbase.supported('postgres', 'mysql')
+ @testing.supported('postgres', 'mysql')
def testnonautoincrement(self):
meta = MetaData(testbase.db)
nonai_table = Table("aitest", meta,
@@ -159,6 +209,9 @@ class AutoIncrementTest(PersistTest):
table.drop()
def testfetchid(self):
+
+ # TODO: what does this test do that all the various ORM tests dont ?
+
meta = MetaData(testbase.db)
table = Table("aitest", meta,
Column('id', Integer, primary_key=True),
@@ -186,7 +239,7 @@ class AutoIncrementTest(PersistTest):
class SequenceTest(PersistTest):
- @testbase.supported('postgres', 'oracle')
+ @testing.supported('postgres', 'oracle')
def setUpAll(self):
global cartitems, sometable, metadata
metadata = MetaData(testbase.db)
@@ -197,13 +250,13 @@ class SequenceTest(PersistTest):
)
sometable = Table( 'Manager', metadata,
Column( 'obj_id', Integer, Sequence('obj_id_seq'), ),
- Column( 'name', type= String, ),
+ Column( 'name', String, ),
Column( 'id', Integer, primary_key= True, ),
)
metadata.create_all()
- @testbase.supported('postgres', 'oracle')
+ @testing.supported('postgres', 'oracle')
def testseqnonpk(self):
"""test sequences fire off as defaults on non-pk columns"""
sometable.insert().execute(name="somename")
@@ -213,7 +266,7 @@ class SequenceTest(PersistTest):
(2, "someother", 2),
]
- @testbase.supported('postgres', 'oracle')
+ @testing.supported('postgres', 'oracle')
def testsequence(self):
cartitems.insert().execute(description='hi')
cartitems.insert().execute(description='there')
@@ -222,8 +275,8 @@ class SequenceTest(PersistTest):
cartitems.select().execute().fetchall()
- @testbase.supported('postgres', 'oracle')
- def teststandalone(self):
+ @testing.supported('postgres', 'oracle')
+ def test_implicit_sequence_exec(self):
s = Sequence("my_sequence", metadata=MetaData(testbase.db))
s.create()
try:
@@ -232,7 +285,7 @@ class SequenceTest(PersistTest):
finally:
s.drop()
- @testbase.supported('postgres', 'oracle')
+ @testing.supported('postgres', 'oracle')
def teststandalone_explicit(self):
s = Sequence("my_sequence")
s.create(bind=testbase.db)
@@ -242,12 +295,20 @@ class SequenceTest(PersistTest):
finally:
s.drop(testbase.db)
- @testbase.supported('postgres', 'oracle')
+ @testing.supported('postgres', 'oracle')
+ def test_checkfirst(self):
+ s = Sequence("my_sequence")
+ s.create(testbase.db, checkfirst=False)
+ s.create(testbase.db, checkfirst=True)
+ s.drop(testbase.db, checkfirst=False)
+ s.drop(testbase.db, checkfirst=True)
+
+ @testing.supported('postgres', 'oracle')
def teststandalone2(self):
x = cartitems.c.cart_id.sequence.execute()
self.assert_(1 <= x <= 4)
- @testbase.supported('postgres', 'oracle')
+ @testing.supported('postgres', 'oracle')
def tearDownAll(self):
metadata.drop_all()
diff --git a/test/sql/generative.py b/test/sql/generative.py
new file mode 100644
index 000000000..357a66fcd
--- /dev/null
+++ b/test/sql/generative.py
@@ -0,0 +1,275 @@
+import testbase
+from sql import select as selecttests
+from sqlalchemy import *
+from testlib import *
+
+class TraversalTest(AssertMixin):
+ """test ClauseVisitor's traversal, particularly its ability to copy and modify
+ a ClauseElement in place."""
+
+ def setUpAll(self):
+ global A, B
+
+ # establish two ficticious ClauseElements.
+ # define deep equality semantics as well as deep identity semantics.
+ class A(ClauseElement):
+ def __init__(self, expr):
+ self.expr = expr
+
+ def is_other(self, other):
+ return other is self
+
+ def __eq__(self, other):
+ return other.expr == self.expr
+
+ def __ne__(self, other):
+ return other.expr != self.expr
+
+ def __str__(self):
+ return "A(%s)" % repr(self.expr)
+
+ class B(ClauseElement):
+ def __init__(self, *items):
+ self.items = items
+
+ def is_other(self, other):
+ if other is not self:
+ return False
+ for i1, i2 in zip(self.items, other.items):
+ if i1 is not i2:
+ return False
+ return True
+
+ def __eq__(self, other):
+ for i1, i2 in zip(self.items, other.items):
+ if i1 != i2:
+ return False
+ return True
+
+ def __ne__(self, other):
+ for i1, i2 in zip(self.items, other.items):
+ if i1 != i2:
+ return True
+ return False
+
+ def _copy_internals(self):
+ self.items = [i._clone() for i in self.items]
+
+ def get_children(self, **kwargs):
+ return self.items
+
+ def __str__(self):
+ return "B(%s)" % repr([str(i) for i in self.items])
+
+ def test_test_classes(self):
+ a1 = A("expr1")
+ struct = B(a1, A("expr2"), B(A("expr1b"), A("expr2b")), A("expr3"))
+ struct2 = B(a1, A("expr2"), B(A("expr1b"), A("expr2b")), A("expr3"))
+ struct3 = B(a1, A("expr2"), B(A("expr1b"), A("expr2bmodified")), A("expr3"))
+
+ assert a1.is_other(a1)
+ assert struct.is_other(struct)
+ assert struct == struct2
+ assert struct != struct3
+ assert not struct.is_other(struct2)
+ assert not struct.is_other(struct3)
+
+ def test_clone(self):
+ struct = B(A("expr1"), A("expr2"), B(A("expr1b"), A("expr2b")), A("expr3"))
+
+ class Vis(ClauseVisitor):
+ def visit_a(self, a):
+ pass
+ def visit_b(self, b):
+ pass
+
+ vis = Vis()
+ s2 = vis.traverse(struct, clone=True)
+ assert struct == s2
+ assert not struct.is_other(s2)
+
+ def test_no_clone(self):
+ struct = B(A("expr1"), A("expr2"), B(A("expr1b"), A("expr2b")), A("expr3"))
+
+ class Vis(ClauseVisitor):
+ def visit_a(self, a):
+ pass
+ def visit_b(self, b):
+ pass
+
+ vis = Vis()
+ s2 = vis.traverse(struct, clone=False)
+ assert struct == s2
+ assert struct.is_other(s2)
+
+ def test_change_in_place(self):
+ struct = B(A("expr1"), A("expr2"), B(A("expr1b"), A("expr2b")), A("expr3"))
+ struct2 = B(A("expr1"), A("expr2modified"), B(A("expr1b"), A("expr2b")), A("expr3"))
+ struct3 = B(A("expr1"), A("expr2"), B(A("expr1b"), A("expr2bmodified")), A("expr3"))
+
+ class Vis(ClauseVisitor):
+ def visit_a(self, a):
+ if a.expr == "expr2":
+ a.expr = "expr2modified"
+ def visit_b(self, b):
+ pass
+
+ vis = Vis()
+ s2 = vis.traverse(struct, clone=True)
+ assert struct != s2
+ assert not struct.is_other(s2)
+ assert struct2 == s2
+
+ class Vis2(ClauseVisitor):
+ def visit_a(self, a):
+ if a.expr == "expr2b":
+ a.expr = "expr2bmodified"
+ def visit_b(self, b):
+ pass
+
+ vis2 = Vis2()
+ s3 = vis2.traverse(struct, clone=True)
+ assert struct != s3
+ assert struct3 == s3
+
+class ClauseTest(selecttests.SQLTest):
+ """test copy-in-place behavior of various ClauseElements."""
+
+ def setUpAll(self):
+ global t1, t2
+ t1 = table("table1",
+ column("col1"),
+ column("col2"),
+ column("col3"),
+ )
+ t2 = table("table2",
+ column("col1"),
+ column("col2"),
+ column("col3"),
+ )
+
+ def test_binary(self):
+ clause = t1.c.col2 == t2.c.col2
+ assert str(clause) == ClauseVisitor().traverse(clause, clone=True)
+
+ def test_join(self):
+ clause = t1.join(t2, t1.c.col2==t2.c.col2)
+ c1 = str(clause)
+ assert str(clause) == str(ClauseVisitor().traverse(clause, clone=True))
+
+ class Vis(ClauseVisitor):
+ def visit_binary(self, binary):
+ binary.right = t2.c.col3
+
+ clause2 = Vis().traverse(clause, clone=True)
+ assert c1 == str(clause)
+ assert str(clause2) == str(t1.join(t2, t1.c.col2==t2.c.col3))
+
+ def test_select(self):
+ s = t1.select()
+ s2 = select([s])
+ s2_assert = str(s2)
+ s3_assert = str(select([t1.select()], t1.c.col2==7))
+ class Vis(ClauseVisitor):
+ def visit_select(self, select):
+ select.append_whereclause(t1.c.col2==7)
+ s3 = Vis().traverse(s2, clone=True)
+ assert str(s3) == s3_assert
+ assert str(s2) == s2_assert
+ print str(s2)
+ print str(s3)
+ Vis().traverse(s2)
+ assert str(s2) == s3_assert
+
+ print "------------------"
+
+ s4_assert = str(select([t1.select()], and_(t1.c.col2==7, t1.c.col3==9)))
+ class Vis(ClauseVisitor):
+ def visit_select(self, select):
+ select.append_whereclause(t1.c.col3==9)
+ s4 = Vis().traverse(s3, clone=True)
+ print str(s3)
+ print str(s4)
+ assert str(s4) == s4_assert
+ assert str(s3) == s3_assert
+
+ print "------------------"
+ s5_assert = str(select([t1.select()], and_(t1.c.col2==7, t1.c.col1==9)))
+ class Vis(ClauseVisitor):
+ def visit_binary(self, binary):
+ if binary.left is t1.c.col3:
+ binary.left = t1.c.col1
+ binary.right = bindparam("table1_col1")
+ s5 = Vis().traverse(s4, clone=True)
+ print str(s4)
+ print str(s5)
+ assert str(s5) == s5_assert
+ assert str(s4) == s4_assert
+
+ def test_correlated_select(self):
+ s = select(['*'], t1.c.col1==t2.c.col1, from_obj=[t1, t2]).correlate(t2)
+ class Vis(ClauseVisitor):
+ def visit_select(self, select):
+ select.append_whereclause(t1.c.col2==7)
+
+ self.runtest(Vis().traverse(s, clone=True), "SELECT * FROM table1 WHERE table1.col1 = table2.col1 AND table1.col2 = :table1_col2")
+
+ def test_clause_adapter(self):
+ from sqlalchemy import sql_util
+
+ t1alias = t1.alias('t1alias')
+
+ vis = sql_util.ClauseAdapter(t1alias)
+ ff = vis.traverse(func.count(t1.c.col1).label('foo'), clone=True)
+ assert ff._get_from_objects() == [t1alias]
+
+ self.runtest(vis.traverse(select(['*'], from_obj=[t1]), clone=True), "SELECT * FROM table1 AS t1alias")
+ self.runtest(vis.traverse(select(['*'], t1.c.col1==t2.c.col2), clone=True), "SELECT * FROM table1 AS t1alias, table2 WHERE t1alias.col1 = table2.col2")
+ self.runtest(vis.traverse(select(['*'], t1.c.col1==t2.c.col2, from_obj=[t1, t2]), clone=True), "SELECT * FROM table1 AS t1alias, table2 WHERE t1alias.col1 = table2.col2")
+ self.runtest(vis.traverse(select(['*'], t1.c.col1==t2.c.col2, from_obj=[t1, t2]).correlate(t1), clone=True), "SELECT * FROM table2 WHERE t1alias.col1 = table2.col2")
+ self.runtest(vis.traverse(select(['*'], t1.c.col1==t2.c.col2, from_obj=[t1, t2]).correlate(t2), clone=True), "SELECT * FROM table1 AS t1alias WHERE t1alias.col1 = table2.col2")
+
+ ff = vis.traverse(func.count(t1.c.col1).label('foo'), clone=True)
+ self.runtest(ff, "count(t1alias.col1) AS foo")
+ assert ff._get_from_objects() == [t1alias]
+
+# TODO:
+# self.runtest(vis.traverse(select([func.count(t1.c.col1).label('foo')]), clone=True), "SELECT count(t1alias.col1) AS foo FROM table1 AS t1alias")
+
+ t2alias = t2.alias('t2alias')
+ vis.chain(sql_util.ClauseAdapter(t2alias))
+ self.runtest(vis.traverse(select(['*'], t1.c.col1==t2.c.col2), clone=True), "SELECT * FROM table1 AS t1alias, table2 AS t2alias WHERE t1alias.col1 = t2alias.col2")
+ self.runtest(vis.traverse(select(['*'], t1.c.col1==t2.c.col2, from_obj=[t1, t2]), clone=True), "SELECT * FROM table1 AS t1alias, table2 AS t2alias WHERE t1alias.col1 = t2alias.col2")
+ self.runtest(vis.traverse(select(['*'], t1.c.col1==t2.c.col2, from_obj=[t1, t2]).correlate(t1), clone=True), "SELECT * FROM table2 AS t2alias WHERE t1alias.col1 = t2alias.col2")
+ self.runtest(vis.traverse(select(['*'], t1.c.col1==t2.c.col2, from_obj=[t1, t2]).correlate(t2), clone=True), "SELECT * FROM table1 AS t1alias WHERE t1alias.col1 = t2alias.col2")
+
+
+
+class SelectTest(selecttests.SQLTest):
+ """tests the generative capability of Select"""
+
+ def setUpAll(self):
+ global t1, t2
+ t1 = table("table1",
+ column("col1"),
+ column("col2"),
+ column("col3"),
+ )
+ t2 = table("table2",
+ column("col1"),
+ column("col2"),
+ column("col3"),
+ )
+
+ def test_select(self):
+ self.runtest(t1.select().where(t1.c.col1==5).order_by(t1.c.col3), "SELECT table1.col1, table1.col2, table1.col3 FROM table1 WHERE table1.col1 = :table1_col1 ORDER BY table1.col3")
+
+ self.runtest(t1.select().select_from(select([t2], t2.c.col1==t1.c.col1)).order_by(t1.c.col3), "SELECT table1.col1, table1.col2, table1.col3 FROM table1, (SELECT table2.col1 AS col1, table2.col2 AS col2, table2.col3 AS col3 FROM table2 WHERE table2.col1 = table1.col1) ORDER BY table1.col3")
+
+ s = select([t2], t2.c.col1==t1.c.col1, correlate=False)
+ s = s.correlate(t1).order_by(t2.c.col3)
+ self.runtest(t1.select().select_from(s).order_by(t1.c.col3), "SELECT table1.col1, table1.col2, table1.col3 FROM table1, (SELECT table2.col1 AS col1, table2.col2 AS col2, table2.col3 AS col3 FROM table2 WHERE table2.col1 = table1.col1 ORDER BY table2.col3) ORDER BY table1.col3")
+
+
+if __name__ == '__main__':
+ testbase.main()
diff --git a/test/sql/labels.py b/test/sql/labels.py
index ee9fa6bc5..553a3a3bc 100644
--- a/test/sql/labels.py
+++ b/test/sql/labels.py
@@ -1,11 +1,12 @@
import testbase
-
from sqlalchemy import *
+from testlib import *
+
# TODO: either create a mock dialect with named paramstyle and a short identifier length,
# or find a way to just use sqlite dialect and make those changes
-class LabelTypeTest(testbase.PersistTest):
+class LabelTypeTest(PersistTest):
def test_type(self):
m = MetaData()
t = Table('sometable', m,
@@ -14,21 +15,26 @@ class LabelTypeTest(testbase.PersistTest):
assert isinstance(t.c.col1.label('hi').type, Integer)
assert isinstance(select([t.c.col2], scalar=True).label('lala').type, Float)
-class LongLabelsTest(testbase.PersistTest):
+class LongLabelsTest(PersistTest):
def setUpAll(self):
- global metadata, table1
- metadata = MetaData(engine=testbase.db)
+ global metadata, table1, maxlen
+ metadata = MetaData(testbase.db)
table1 = Table("some_large_named_table", metadata,
Column("this_is_the_primarykey_column", Integer, Sequence("this_is_some_large_seq"), primary_key=True),
Column("this_is_the_data_column", String(30))
)
metadata.create_all()
+
+ maxlen = testbase.db.dialect.max_identifier_length
+ testbase.db.dialect.max_identifier_length = lambda: 29
+
def tearDown(self):
table1.delete().execute()
def tearDownAll(self):
metadata.drop_all()
+ testbase.db.dialect.max_identifier_length = maxlen
def test_result(self):
table1.insert().execute(**{"this_is_the_primarykey_column":1, "this_is_the_data_column":"data1"})
@@ -88,7 +94,7 @@ class LongLabelsTest(testbase.PersistTest):
x = select([tt], use_labels=True, order_by=tt.oid_column).compile(dialect=dialect)
#print x
# assert it doesnt end with "ORDER BY foo.some_large_named_table_this_is_the_primarykey_column"
- assert str(x).endswith("""ORDER BY foo.some_large_named_table_t_1""")
+ assert str(x).endswith("""ORDER BY foo.some_large_named_table_t_2""")
if __name__ == '__main__':
testbase.main()
diff --git a/test/sql/query.py b/test/sql/query.py
index 8af5aafea..48a28a9a5 100644
--- a/test/sql/query.py
+++ b/test/sql/query.py
@@ -1,13 +1,9 @@
-from testbase import PersistTest
import testbase
-import unittest, sys, datetime
-
-import sqlalchemy.databases.sqlite as sqllite
-
-import tables
+import datetime
from sqlalchemy import *
-from sqlalchemy.engine import ResultProxy, RowProxy
from sqlalchemy import exceptions
+from testlib import *
+
class QueryTest(PersistTest):
@@ -24,25 +20,24 @@ class QueryTest(PersistTest):
Column('address', String(30)))
metadata.create_all()
- def setUp(self):
- self.users = users
def tearDown(self):
- self.users.delete().execute()
+ addresses.delete().execute()
+ users.delete().execute()
def tearDownAll(self):
metadata.drop_all()
def testinsert(self):
- self.users.insert().execute(user_id = 7, user_name = 'jack')
- print repr(self.users.select().execute().fetchall())
-
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ assert users.count().scalar() == 1
+
def testupdate(self):
- self.users.insert().execute(user_id = 7, user_name = 'jack')
- print repr(self.users.select().execute().fetchall())
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ assert users.count().scalar() == 1
- self.users.update(self.users.c.user_id == 7).execute(user_name = 'fred')
- print repr(self.users.select().execute().fetchall())
+ users.update(users.c.user_id == 7).execute(user_name = 'fred')
+ assert users.select(users.c.user_id==7).execute().fetchone()['user_name'] == 'fred'
def test_lastrow_accessor(self):
"""test the last_inserted_ids() and lastrow_has_id() functions"""
@@ -63,14 +58,15 @@ class QueryTest(PersistTest):
if result.lastrow_has_defaults():
criterion = and_(*[col==id for col, id in zip(table.primary_key, result.last_inserted_ids())])
row = table.select(criterion).execute().fetchone()
- ret.update(row)
+ for c in table.c:
+ ret[c.key] = row[c]
return ret
for supported, table, values, assertvalues in [
(
{'unsupported':['sqlite']},
Table("t1", metadata,
- Column('id', Integer, primary_key=True),
+ Column('id', Integer, Sequence('t1_id_seq', optional=True), primary_key=True),
Column('foo', String(30), primary_key=True)),
{'foo':'hi'},
{'id':1, 'foo':'hi'}
@@ -78,7 +74,7 @@ class QueryTest(PersistTest):
(
{'unsupported':['sqlite']},
Table("t2", metadata,
- Column('id', Integer, primary_key=True),
+ Column('id', Integer, Sequence('t2_id_seq', optional=True), primary_key=True),
Column('foo', String(30), primary_key=True),
Column('bar', String(30), PassiveDefault('hi'))
),
@@ -98,7 +94,7 @@ class QueryTest(PersistTest):
(
{'unsupported':[]},
Table("t4", metadata,
- Column('id', Integer, primary_key=True),
+ Column('id', Integer, Sequence('t4_id_seq', optional=True), primary_key=True),
Column('foo', String(30), primary_key=True),
Column('bar', String(30), PassiveDefault('hi'))
),
@@ -124,109 +120,94 @@ class QueryTest(PersistTest):
table.drop()
def testrowiteration(self):
- self.users.insert().execute(user_id = 7, user_name = 'jack')
- self.users.insert().execute(user_id = 8, user_name = 'ed')
- self.users.insert().execute(user_id = 9, user_name = 'fred')
- r = self.users.select().execute()
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ users.insert().execute(user_id = 8, user_name = 'ed')
+ users.insert().execute(user_id = 9, user_name = 'fred')
+ r = users.select().execute()
l = []
for row in r:
l.append(row)
self.assert_(len(l) == 3)
def test_fetchmany(self):
- self.users.insert().execute(user_id = 7, user_name = 'jack')
- self.users.insert().execute(user_id = 8, user_name = 'ed')
- self.users.insert().execute(user_id = 9, user_name = 'fred')
- r = self.users.select().execute()
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ users.insert().execute(user_id = 8, user_name = 'ed')
+ users.insert().execute(user_id = 9, user_name = 'fred')
+ r = users.select().execute()
l = []
for row in r.fetchmany(size=2):
l.append(row)
self.assert_(len(l) == 2, "fetchmany(size=2) got %s rows" % len(l))
def test_compiled_execute(self):
- s = select([self.users], self.users.c.user_id==bindparam('id')).compile()
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ s = select([users], users.c.user_id==bindparam('id')).compile()
c = testbase.db.connect()
- print repr(c.execute(s, id=7).fetchall())
-
- def test_global_metadata(self):
- t1 = Table('table1', Column('col1', Integer, primary_key=True),
- Column('col2', String(20)))
- t2 = Table('table2', Column('col1', Integer, primary_key=True),
- Column('col2', String(20)))
-
- assert t1.c.col1
- global_connect(testbase.db)
- default_metadata.create_all()
- try:
- assert t1.count().scalar() == 0
- finally:
- default_metadata.drop_all()
- default_metadata.clear()
-
+ assert c.execute(s, id=7).fetchall()[0]['user_id'] == 7
def test_repeated_bindparams(self):
"""test that a BindParam can be used more than once.
this should be run for dbs with both positional and named paramstyles."""
- self.users.insert().execute(user_id = 7, user_name = 'jack')
- self.users.insert().execute(user_id = 8, user_name = 'fred')
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ users.insert().execute(user_id = 8, user_name = 'fred')
u = bindparam('userid')
- s = self.users.select(or_(self.users.c.user_name==u, self.users.c.user_name==u))
+ s = users.select(or_(users.c.user_name==u, users.c.user_name==u))
r = s.execute(userid='fred').fetchall()
assert len(r) == 1
def test_bindparam_shortname(self):
"""test the 'shortname' field on BindParamClause."""
- self.users.insert().execute(user_id = 7, user_name = 'jack')
- self.users.insert().execute(user_id = 8, user_name = 'fred')
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ users.insert().execute(user_id = 8, user_name = 'fred')
u = bindparam('userid', shortname='someshortname')
- s = self.users.select(self.users.c.user_name==u)
+ s = users.select(users.c.user_name==u)
r = s.execute(someshortname='fred').fetchall()
assert len(r) == 1
def testdelete(self):
- self.users.insert().execute(user_id = 7, user_name = 'jack')
- self.users.insert().execute(user_id = 8, user_name = 'fred')
- print repr(self.users.select().execute().fetchall())
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ users.insert().execute(user_id = 8, user_name = 'fred')
+ print repr(users.select().execute().fetchall())
- self.users.delete(self.users.c.user_name == 'fred').execute()
+ users.delete(users.c.user_name == 'fred').execute()
- print repr(self.users.select().execute().fetchall())
+ print repr(users.select().execute().fetchall())
def testselectlimit(self):
- self.users.insert().execute(user_id=1, user_name='john')
- self.users.insert().execute(user_id=2, user_name='jack')
- self.users.insert().execute(user_id=3, user_name='ed')
- self.users.insert().execute(user_id=4, user_name='wendy')
- self.users.insert().execute(user_id=5, user_name='laura')
- self.users.insert().execute(user_id=6, user_name='ralph')
- self.users.insert().execute(user_id=7, user_name='fido')
- r = self.users.select(limit=3, order_by=[self.users.c.user_id]).execute().fetchall()
+ users.insert().execute(user_id=1, user_name='john')
+ users.insert().execute(user_id=2, user_name='jack')
+ users.insert().execute(user_id=3, user_name='ed')
+ users.insert().execute(user_id=4, user_name='wendy')
+ users.insert().execute(user_id=5, user_name='laura')
+ users.insert().execute(user_id=6, user_name='ralph')
+ users.insert().execute(user_id=7, user_name='fido')
+ r = users.select(limit=3, order_by=[users.c.user_id]).execute().fetchall()
self.assert_(r == [(1, 'john'), (2, 'jack'), (3, 'ed')], repr(r))
- @testbase.unsupported('mssql')
+ @testing.unsupported('mssql')
def testselectlimitoffset(self):
- self.users.insert().execute(user_id=1, user_name='john')
- self.users.insert().execute(user_id=2, user_name='jack')
- self.users.insert().execute(user_id=3, user_name='ed')
- self.users.insert().execute(user_id=4, user_name='wendy')
- self.users.insert().execute(user_id=5, user_name='laura')
- self.users.insert().execute(user_id=6, user_name='ralph')
- self.users.insert().execute(user_id=7, user_name='fido')
- r = self.users.select(limit=3, offset=2, order_by=[self.users.c.user_id]).execute().fetchall()
+ users.insert().execute(user_id=1, user_name='john')
+ users.insert().execute(user_id=2, user_name='jack')
+ users.insert().execute(user_id=3, user_name='ed')
+ users.insert().execute(user_id=4, user_name='wendy')
+ users.insert().execute(user_id=5, user_name='laura')
+ users.insert().execute(user_id=6, user_name='ralph')
+ users.insert().execute(user_id=7, user_name='fido')
+ r = users.select(limit=3, offset=2, order_by=[users.c.user_id]).execute().fetchall()
self.assert_(r==[(3, 'ed'), (4, 'wendy'), (5, 'laura')])
- r = self.users.select(offset=5, order_by=[self.users.c.user_id]).execute().fetchall()
+ r = users.select(offset=5, order_by=[users.c.user_id]).execute().fetchall()
self.assert_(r==[(6, 'ralph'), (7, 'fido')])
- @testbase.supported('mssql')
+ @testing.supported('mssql')
def testselectlimitoffset_mssql(self):
try:
- r = self.users.select(limit=3, offset=2, order_by=[self.users.c.user_id]).execute().fetchall()
+ r = users.select(limit=3, offset=2, order_by=[users.c.user_id]).execute().fetchall()
assert False # InvalidRequestError should have been raised
except exceptions.InvalidRequestError:
pass
- @testbase.unsupported('mysql')
+ @testing.unsupported('mysql')
def test_scalar_select(self):
"""test that scalar subqueries with labels get their type propigated to the result set."""
# mysql and/or mysqldb has a bug here, type isnt propigated for scalar subquery.
@@ -244,18 +225,26 @@ class QueryTest(PersistTest):
datetable.drop()
def test_column_accessor(self):
- self.users.insert().execute(user_id=1, user_name='john')
- self.users.insert().execute(user_id=2, user_name='jack')
- r = self.users.select(self.users.c.user_id==2).execute().fetchone()
- self.assert_(r.user_id == r['user_id'] == r[self.users.c.user_id] == 2)
- self.assert_(r.user_name == r['user_name'] == r[self.users.c.user_name] == 'jack')
-
- r = text("select * from query_users where user_id=2", engine=testbase.db).execute().fetchone()
- self.assert_(r.user_id == r['user_id'] == r[self.users.c.user_id] == 2)
- self.assert_(r.user_name == r['user_name'] == r[self.users.c.user_name] == 'jack')
+ users.insert().execute(user_id=1, user_name='john')
+ users.insert().execute(user_id=2, user_name='jack')
+ addresses.insert().execute(address_id=1, user_id=2, address='foo@bar.com')
+
+ r = users.select(users.c.user_id==2).execute().fetchone()
+ self.assert_(r.user_id == r['user_id'] == r[users.c.user_id] == 2)
+ self.assert_(r.user_name == r['user_name'] == r[users.c.user_name] == 'jack')
+
+ r = text("select * from query_users where user_id=2", bind=testbase.db).execute().fetchone()
+ self.assert_(r.user_id == r['user_id'] == r[users.c.user_id] == 2)
+ self.assert_(r.user_name == r['user_name'] == r[users.c.user_name] == 'jack')
+ # test slices
+ r = text("select * from query_addresses", bind=testbase.db).execute().fetchone()
+ self.assert_(r[0:1] == (1,))
+ self.assert_(r[1:] == (2, 'foo@bar.com'))
+ self.assert_(r[:-1] == (1, 2))
+
def test_ambiguous_column(self):
- self.users.insert().execute(user_id=1, user_name='john')
+ users.insert().execute(user_id=1, user_name='john')
r = users.outerjoin(addresses).select().execute().fetchone()
try:
print r['user_id']
@@ -264,18 +253,18 @@ class QueryTest(PersistTest):
assert str(e) == "Ambiguous column name 'user_id' in result set! try 'use_labels' option on select statement."
def test_keys(self):
- self.users.insert().execute(user_id=1, user_name='foo')
- r = self.users.select().execute().fetchone()
+ users.insert().execute(user_id=1, user_name='foo')
+ r = users.select().execute().fetchone()
self.assertEqual([x.lower() for x in r.keys()], ['user_id', 'user_name'])
def test_items(self):
- self.users.insert().execute(user_id=1, user_name='foo')
- r = self.users.select().execute().fetchone()
+ users.insert().execute(user_id=1, user_name='foo')
+ r = users.select().execute().fetchone()
self.assertEqual([(x[0].lower(), x[1]) for x in r.items()], [('user_id', 1), ('user_name', 'foo')])
def test_len(self):
- self.users.insert().execute(user_id=1, user_name='foo')
- r = self.users.select().execute().fetchone()
+ users.insert().execute(user_id=1, user_name='foo')
+ r = users.select().execute().fetchone()
self.assertEqual(len(r), 2)
r.close()
r = testbase.db.execute('select user_name, user_id from query_users', {}).fetchone()
@@ -295,7 +284,11 @@ class QueryTest(PersistTest):
x = testbase.db.func.current_date().execute().scalar()
y = testbase.db.func.current_date().select().execute().scalar()
z = testbase.db.func.current_date().scalar()
- assert x == y == z
+ assert (x == y == z) is True
+
+ x = testbase.db.func.current_date(type_=Date)
+ assert isinstance(x.type, Date)
+ assert isinstance(x.execute().scalar(), datetime.date)
def test_conn_functions(self):
conn = testbase.db.connect()
@@ -305,8 +298,8 @@ class QueryTest(PersistTest):
z = conn.scalar(func.current_date())
finally:
conn.close()
- assert x == y == z
-
+ assert (x == y == z) is True
+
def test_update_functions(self):
"""test sending functions and SQL expressions to the VALUES and SET clauses of INSERT/UPDATE instances,
and that column-level defaults get overridden"""
@@ -357,7 +350,7 @@ class QueryTest(PersistTest):
finally:
meta.drop_all()
- @testbase.supported('postgres')
+ @testing.supported('postgres')
def test_functions_with_cols(self):
# TODO: shouldnt this work on oracle too ?
x = testbase.db.func.current_date().execute().scalar()
@@ -366,7 +359,7 @@ class QueryTest(PersistTest):
w = select(['*'], from_obj=[testbase.db.func.current_date()]).scalar()
# construct a column-based FROM object out of a function, like in [ticket:172]
- s = select([column('date', type=DateTime)], from_obj=[testbase.db.func.current_date()])
+ s = select([column('date', type_=DateTime)], from_obj=[testbase.db.func.current_date()])
q = s.execute().fetchone()[s.c.date]
r = s.alias('datequery').select().scalar()
@@ -374,8 +367,8 @@ class QueryTest(PersistTest):
def test_column_order_with_simple_query(self):
# should return values in column definition order
- self.users.insert().execute(user_id=1, user_name='foo')
- r = self.users.select(self.users.c.user_id==1).execute().fetchone()
+ users.insert().execute(user_id=1, user_name='foo')
+ r = users.select(users.c.user_id==1).execute().fetchone()
self.assertEqual(r[0], 1)
self.assertEqual(r[1], 'foo')
self.assertEqual([x.lower() for x in r.keys()], ['user_id', 'user_name'])
@@ -383,14 +376,14 @@ class QueryTest(PersistTest):
def test_column_order_with_text_query(self):
# should return values in query order
- self.users.insert().execute(user_id=1, user_name='foo')
+ users.insert().execute(user_id=1, user_name='foo')
r = testbase.db.execute('select user_name, user_id from query_users', {}).fetchone()
self.assertEqual(r[0], 'foo')
self.assertEqual(r[1], 1)
self.assertEqual([x.lower() for x in r.keys()], ['user_name', 'user_id'])
self.assertEqual(r.values(), ['foo', 1])
- @testbase.unsupported('oracle', 'firebird')
+ @testing.unsupported('oracle', 'firebird')
def test_column_accessor_shadow(self):
meta = MetaData(testbase.db)
shadowed = Table('test_shadowed', meta,
@@ -420,7 +413,7 @@ class QueryTest(PersistTest):
finally:
shadowed.drop(checkfirst=True)
- @testbase.supported('mssql')
+ @testing.supported('mssql')
def test_fetchid_trigger(self):
meta = MetaData(testbase.db)
t1 = Table('t1', meta,
@@ -446,7 +439,7 @@ class QueryTest(PersistTest):
con.execute("""drop trigger paj""")
meta.drop_all()
- @testbase.supported('mssql')
+ @testing.supported('mssql')
def test_insertid_schema(self):
meta = MetaData(testbase.db)
con = testbase.db.connect()
@@ -459,7 +452,7 @@ class QueryTest(PersistTest):
tbl.drop()
con.execute('drop schema paj')
- @testbase.supported('mssql')
+ @testing.supported('mssql')
def test_insertid_reserved(self):
meta = MetaData(testbase.db)
table = Table(
@@ -476,51 +469,52 @@ class QueryTest(PersistTest):
def test_in_filtering(self):
- """test the 'shortname' field on BindParamClause."""
- self.users.insert().execute(user_id = 7, user_name = 'jack')
- self.users.insert().execute(user_id = 8, user_name = 'fred')
- self.users.insert().execute(user_id = 9, user_name = None)
+ """test the behavior of the in_() function."""
+
+ users.insert().execute(user_id = 7, user_name = 'jack')
+ users.insert().execute(user_id = 8, user_name = 'fred')
+ users.insert().execute(user_id = 9, user_name = None)
- s = self.users.select(self.users.c.user_name.in_())
+ s = users.select(users.c.user_name.in_())
r = s.execute().fetchall()
# No username is in empty set
assert len(r) == 0
- s = self.users.select(not_(self.users.c.user_name.in_()))
+ s = users.select(not_(users.c.user_name.in_()))
r = s.execute().fetchall()
# All usernames with a value are outside an empty set
assert len(r) == 2
- s = self.users.select(self.users.c.user_name.in_('jack','fred'))
+ s = users.select(users.c.user_name.in_('jack','fred'))
r = s.execute().fetchall()
assert len(r) == 2
- s = self.users.select(not_(self.users.c.user_name.in_('jack','fred')))
+ s = users.select(not_(users.c.user_name.in_('jack','fred')))
r = s.execute().fetchall()
# Null values are not outside any set
assert len(r) == 0
u = bindparam('search_key')
- s = self.users.select(u.in_())
+ s = users.select(u.in_())
r = s.execute(search_key='john').fetchall()
assert len(r) == 0
r = s.execute(search_key=None).fetchall()
assert len(r) == 0
- s = self.users.select(not_(u.in_()))
+ s = users.select(not_(u.in_()))
r = s.execute(search_key='john').fetchall()
assert len(r) == 3
r = s.execute(search_key=None).fetchall()
assert len(r) == 0
- s = self.users.select(self.users.c.user_name.in_() == True)
+ s = users.select(users.c.user_name.in_() == True)
r = s.execute().fetchall()
assert len(r) == 0
- s = self.users.select(self.users.c.user_name.in_() == False)
+ s = users.select(users.c.user_name.in_() == False)
r = s.execute().fetchall()
assert len(r) == 2
- s = self.users.select(self.users.c.user_name.in_() == None)
+ s = users.select(users.c.user_name.in_() == None)
r = s.execute().fetchall()
assert len(r) == 1
@@ -577,7 +571,7 @@ class CompoundTest(PersistTest):
assert u.execute().fetchall() == [('aaa', 'aaa'), ('bbb', 'bbb'), ('bbb', 'ccc'), ('ccc', 'aaa')]
assert u.alias('bar').select().execute().fetchall() == [('aaa', 'aaa'), ('bbb', 'bbb'), ('bbb', 'ccc'), ('ccc', 'aaa')]
- @testbase.unsupported('mysql')
+ @testing.unsupported('mysql')
def test_intersect(self):
i = intersect(
select([t2.c.col3, t2.c.col4]),
@@ -586,7 +580,7 @@ class CompoundTest(PersistTest):
assert i.execute().fetchall() == [('aaa', 'bbb'), ('bbb', 'ccc'), ('ccc', 'aaa')]
assert i.alias('bar').select().execute().fetchall() == [('aaa', 'bbb'), ('bbb', 'ccc'), ('ccc', 'aaa')]
- @testbase.unsupported('mysql', 'oracle')
+ @testing.unsupported('mysql', 'oracle')
def test_except_style1(self):
e = except_(union(
select([t1.c.col3, t1.c.col4]),
@@ -595,7 +589,7 @@ class CompoundTest(PersistTest):
), select([t2.c.col3, t2.c.col4]))
assert e.alias('bar').select().execute().fetchall() == [('aaa', 'aaa'), ('aaa', 'ccc'), ('bbb', 'aaa'), ('bbb', 'bbb'), ('ccc', 'bbb'), ('ccc', 'ccc')]
- @testbase.unsupported('mysql', 'oracle')
+ @testing.unsupported('mysql', 'oracle')
def test_except_style2(self):
e = except_(union(
select([t1.c.col3, t1.c.col4]),
@@ -605,7 +599,7 @@ class CompoundTest(PersistTest):
assert e.execute().fetchall() == [('aaa', 'aaa'), ('aaa', 'ccc'), ('bbb', 'aaa'), ('bbb', 'bbb'), ('ccc', 'bbb'), ('ccc', 'ccc')]
assert e.alias('bar').select().execute().fetchall() == [('aaa', 'aaa'), ('aaa', 'ccc'), ('bbb', 'aaa'), ('bbb', 'bbb'), ('ccc', 'bbb'), ('ccc', 'ccc')]
- @testbase.unsupported('sqlite', 'mysql', 'oracle')
+ @testing.unsupported('sqlite', 'mysql', 'oracle')
def test_except_style3(self):
# aaa, bbb, ccc - (aaa, bbb, ccc - (ccc)) = ccc
e = except_(
@@ -617,7 +611,7 @@ class CompoundTest(PersistTest):
)
self.assertEquals(e.execute().fetchall(), [('ccc',)])
- @testbase.unsupported('sqlite', 'mysql', 'oracle')
+ @testing.unsupported('sqlite', 'mysql', 'oracle')
def test_union_union_all(self):
e = union_all(
select([t1.c.col3]),
@@ -628,7 +622,7 @@ class CompoundTest(PersistTest):
)
self.assertEquals(e.execute().fetchall(), [('aaa',),('bbb',),('ccc',),('aaa',),('bbb',),('ccc',)])
- @testbase.unsupported('mysql')
+ @testing.unsupported('mysql')
def test_composite(self):
u = intersect(
select([t2.c.col3, t2.c.col4]),
diff --git a/test/sql/quote.py b/test/sql/quote.py
index bc40d52ee..2fdf9dba0 100644
--- a/test/sql/quote.py
+++ b/test/sql/quote.py
@@ -1,6 +1,7 @@
-from testbase import PersistTest
import testbase
from sqlalchemy import *
+from testlib import *
+
class QuoteTest(PersistTest):
def setUpAll(self):
@@ -78,7 +79,7 @@ class QuoteTest(PersistTest):
assert t1.c.UcCol.case_sensitive is False
assert t2.c.normalcol.case_sensitive is False
- @testbase.unsupported('oracle')
+ @testing.unsupported('oracle')
def testlabels(self):
"""test the quoting of labels.
diff --git a/test/sql/rowcount.py b/test/sql/rowcount.py
index df6a2a883..e0da96a81 100644
--- a/test/sql/rowcount.py
+++ b/test/sql/rowcount.py
@@ -1,7 +1,9 @@
-from sqlalchemy import *
import testbase
+from sqlalchemy import *
+from testlib import *
+
-class FoundRowsTest(testbase.AssertMixin):
+class FoundRowsTest(AssertMixin):
"""tests rowcount functionality"""
def setUpAll(self):
metadata = MetaData(testbase.db)
diff --git a/test/sql/select.py b/test/sql/select.py
index 4d3eb4ad7..a5cf061e2 100644
--- a/test/sql/select.py
+++ b/test/sql/select.py
@@ -1,8 +1,8 @@
-from testbase import PersistTest
import testbase
+import re, operator
from sqlalchemy import *
from sqlalchemy.databases import sqlite, postgres, mysql, oracle, firebird, mssql
-import unittest, re, operator
+from testlib import *
# the select test now tests almost completely with TableClause/ColumnClause objects,
@@ -10,21 +10,21 @@ import unittest, re, operator
# so SQLAlchemy's SQL construction engine can be used with no database dependencies at all.
table1 = table('mytable',
- column('myid'),
- column('name'),
- column('description'),
+ column('myid', Integer),
+ column('name', String),
+ column('description', String),
)
table2 = table(
'myothertable',
- column('otherid'),
- column('othername'),
+ column('otherid', Integer),
+ column('othername', String),
)
table3 = table(
'thirdtable',
- column('userid'),
- column('otherstuff'),
+ column('userid', Integer),
+ column('otherstuff', String),
)
metadata = MetaData()
@@ -54,7 +54,7 @@ addresses = table('addresses',
class SQLTest(PersistTest):
def runtest(self, clause, result, dialect = None, params = None, checkparams = None):
c = clause.compile(parameters=params, dialect=dialect)
- self.echo("\nSQL String:\n" + str(c) + repr(c.get_params()))
+ print "\nSQL String:\n" + str(c) + repr(c.get_params())
cc = re.sub(r'\n', '', str(c))
self.assert_(cc == result, "\n'" + cc + "'\n does not match \n'" + result + "'")
if checkparams is not None:
@@ -130,6 +130,15 @@ sq.myothertable_othername AS sq_myothertable_othername FROM (" + sqstring + ") A
crit = q.c.myid == table1.c.myid
self.runtest(select(['*'], crit), """SELECT * FROM (SELECT mytable.myid AS myid FROM mytable ORDER BY mytable.myid) AS foo, mytable WHERE foo.myid = mytable.myid""", dialect=sqlite.dialect())
self.runtest(select(['*'], crit), """SELECT * FROM (SELECT mytable.myid AS myid FROM mytable) AS foo, mytable WHERE foo.myid = mytable.myid""", dialect=mssql.dialect())
+
+ def testmssql_aliases_schemas(self):
+ self.runtest(table4.select(), "SELECT remotetable.rem_id, remotetable.datatype_id, remotetable.value FROM remote_owner.remotetable")
+
+ dialect = mssql.dialect()
+ self.runtest(table4.select(), "SELECT remotetable_1.rem_id, remotetable_1.datatype_id, remotetable_1.value FROM remote_owner.remotetable AS remotetable_1", dialect=dialect)
+
+ # TODO: this is probably incorrect; no "AS <foo>" is being applied to the table
+ self.runtest(table1.join(table4, table1.c.myid==table4.c.rem_id).select(), "SELECT mytable.myid, mytable.name, mytable.description, remotetable.rem_id, remotetable.datatype_id, remotetable.value FROM mytable JOIN remote_owner.remotetable ON remotetable.rem_id = mytable.myid")
def testdontovercorrelate(self):
self.runtest(select([table1], from_obj=[table1, table1.select()]), """SELECT mytable.myid, mytable.name, mytable.description FROM mytable, (SELECT mytable.myid AS myid, mytable.name AS name, mytable.description AS description FROM mytable)""")
@@ -142,6 +151,11 @@ sq.myothertable_othername AS sq_myothertable_othername FROM (" + sqstring + ") A
self.runtest(select([table1, exists([1], from_obj=[table2]).label('foo')]), "SELECT mytable.myid, mytable.name, mytable.description, EXISTS (SELECT 1 FROM myothertable) AS foo FROM mytable", params={})
def testwheresubquery(self):
+ s = select([addresses.c.street], addresses.c.user_id==users.c.user_id, correlate=True).alias('s')
+ self.runtest(
+ select([users, s.c.street], from_obj=[s]),
+ """SELECT users.user_id, users.user_name, users.password, s.street FROM users, (SELECT addresses.street AS street FROM addresses WHERE addresses.user_id = users.user_id) AS s""")
+
# TODO: this tests that you dont get a "SELECT column" without a FROM but its not working yet.
#self.runtest(
# table1.select(table1.c.myid == select([table1.c.myid], table1.c.name=='jack')), ""
@@ -194,7 +208,20 @@ sq.myothertable_othername AS sq_myothertable_othername FROM (" + sqstring + ") A
s = select([table1.c.myid], scalar=True)
self.runtest(select([table2, s]), "SELECT myothertable.otherid, myothertable.othername, (SELECT mytable.myid FROM mytable) FROM myothertable")
-
+
+ s = select([table1.c.myid]).correlate(None).as_scalar()
+ self.runtest(select([table1, s]), "SELECT mytable.myid, mytable.name, mytable.description, (SELECT mytable.myid FROM mytable) FROM mytable")
+
+ s = select([table1.c.myid]).as_scalar()
+ self.runtest(select([table2, s]), "SELECT myothertable.otherid, myothertable.othername, (SELECT mytable.myid FROM mytable) FROM myothertable")
+
+ # test expressions against scalar selects
+ self.runtest(select([s - literal(8)]), "SELECT (SELECT mytable.myid FROM mytable) - :literal")
+ self.runtest(select([select([table1.c.name]).as_scalar() + literal('x')]), "SELECT (SELECT mytable.name FROM mytable) || :literal")
+ self.runtest(select([s > literal(8)]), "SELECT (SELECT mytable.myid FROM mytable) > :literal")
+
+ self.runtest(select([select([table1.c.name]).label('foo')]), "SELECT (SELECT mytable.name FROM mytable) AS foo")
+
zips = table('zips',
column('zipcode'),
@@ -206,15 +233,17 @@ sq.myothertable_othername AS sq_myothertable_othername FROM (" + sqstring + ") A
column('nm')
)
zip = '12345'
- qlat = select([zips.c.latitude], zips.c.zipcode == zip, scalar=True, correlate=False)
- qlng = select([zips.c.longitude], zips.c.zipcode == zip, scalar=True, correlate=False)
+ qlat = select([zips.c.latitude], zips.c.zipcode == zip).correlate(None).as_scalar()
+ qlng = select([zips.c.longitude], zips.c.zipcode == zip).correlate(None).as_scalar()
q = select([places.c.id, places.c.nm, zips.c.zipcode, func.latlondist(qlat, qlng).label('dist')],
zips.c.zipcode==zip,
order_by = ['dist', places.c.nm]
)
- self.runtest(q,"SELECT places.id, places.nm, zips.zipcode, latlondist((SELECT zips.latitude FROM zips WHERE zips.zipcode = :zips_zipcode_1), (SELECT zips.longitude FROM zips WHERE zips.zipcode = :zips_zipcode_2)) AS dist FROM places, zips WHERE zips.zipcode = :zips_zipcode ORDER BY dist, places.nm")
+ self.runtest(q,"SELECT places.id, places.nm, zips.zipcode, latlondist((SELECT zips.latitude FROM zips WHERE "
+ "zips.zipcode = :zips_zipcode), (SELECT zips.longitude FROM zips WHERE zips.zipcode = :zips_zipcode_1)) AS dist "
+ "FROM places, zips WHERE zips.zipcode = :zips_zipcode_2 ORDER BY dist, places.nm")
zalias = zips.alias('main_zip')
qlat = select([zips.c.latitude], zips.c.zipcode == zalias.c.zipcode, scalar=True)
@@ -223,7 +252,7 @@ sq.myothertable_othername AS sq_myothertable_othername FROM (" + sqstring + ") A
order_by = ['dist', places.c.nm]
)
self.runtest(q, "SELECT places.id, places.nm, main_zip.zipcode, latlondist((SELECT zips.latitude FROM zips WHERE zips.zipcode = main_zip.zipcode), (SELECT zips.longitude FROM zips WHERE zips.zipcode = main_zip.zipcode)) AS dist FROM places, zips AS main_zip ORDER BY dist, places.nm")
-
+
a1 = table2.alias('t2alias')
s1 = select([a1.c.otherid], table1.c.myid==a1.c.otherid, scalar=True)
j1 = table1.join(table2, table1.c.myid==table2.c.otherid)
@@ -261,28 +290,20 @@ sq.myothertable_othername AS sq_myothertable_othername FROM (" + sqstring + ") A
)
def testoperators(self):
- self.runtest(
- table1.select((table1.c.myid != 12) & ~(table1.c.name=='john')),
- "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid != :mytable_myid AND NOT mytable.name = :mytable_name"
- )
-
- self.runtest(
- literal("a") + literal("b") * literal("c"), ":literal + :literal_1 * :literal_2"
- )
# exercise arithmetic operators
for (py_op, sql_op) in ((operator.add, '+'), (operator.mul, '*'),
(operator.sub, '-'), (operator.div, '/'),
):
for (lhs, rhs, res) in (
- ('a', table1.c.myid, ':mytable_myid %s mytable.myid'),
- ('a', literal('b'), ':literal %s :literal_1'),
+ (5, table1.c.myid, ':mytable_myid %s mytable.myid'),
+ (5, literal(5), ':literal %s :literal_1'),
(table1.c.myid, 'b', 'mytable.myid %s :mytable_myid'),
- (table1.c.myid, literal('b'), 'mytable.myid %s :literal'),
+ (table1.c.myid, literal(2.7), 'mytable.myid %s :literal'),
(table1.c.myid, table1.c.myid, 'mytable.myid %s mytable.myid'),
- (literal('a'), 'b', ':literal %s :literal_1'),
- (literal('a'), table1.c.myid, ':literal %s mytable.myid'),
- (literal('a'), literal('b'), ':literal %s :literal_1'),
+ (literal(5), 8, ':literal %s :literal_1'),
+ (literal(6), table1.c.myid, ':literal %s mytable.myid'),
+ (literal(7), literal(5.5), ':literal %s :literal_1'),
):
self.runtest(py_op(lhs, rhs), res % sql_op)
@@ -314,6 +335,25 @@ sq.myothertable_othername AS sq_myothertable_othername FROM (" + sqstring + ") A
"\n'" + compiled + "'\n does not match\n'" +
fwd_sql + "'\n or\n'" + rev_sql + "'")
+ self.runtest(
+ table1.select((table1.c.myid != 12) & ~(table1.c.name=='john')),
+ "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid != :mytable_myid AND mytable.name != :mytable_name"
+ )
+
+ self.runtest(
+ table1.select((table1.c.myid != 12) & ~and_(table1.c.name=='john', table1.c.name=='ed', table1.c.name=='fred')),
+ "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid != :mytable_myid AND NOT (mytable.name = :mytable_name AND mytable.name = :mytable_name_1 AND mytable.name = :mytable_name_2)"
+ )
+
+ self.runtest(
+ table1.select((table1.c.myid != 12) & ~table1.c.name),
+ "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid != :mytable_myid AND NOT mytable.name"
+ )
+
+ self.runtest(
+ literal("a") + literal("b") * literal("c"), ":literal || :literal_1 * :literal_2"
+ )
+
# test the op() function, also that its results are further usable in expressions
self.runtest(
table1.select(table1.c.myid.op('hoho')(12)==14),
@@ -374,13 +414,18 @@ sq.myothertable_othername AS sq_myothertable_othername FROM (" + sqstring + ") A
def testalias(self):
# test the alias for a table1. column names stay the same, table name "changes" to "foo".
self.runtest(
- select([alias(table1, 'foo')])
+ select([table1.alias('foo')])
,"SELECT foo.myid, foo.name, foo.description FROM mytable AS foo")
-
+
+ for dialect in (firebird.dialect(), oracle.dialect()):
+ self.runtest(
+ select([table1.alias('foo')])
+ ,"SELECT foo.myid, foo.name, foo.description FROM mytable foo"
+ ,dialect=dialect)
+
self.runtest(
- select([alias(table1, 'foo')])
- ,"SELECT foo.myid, foo.name, foo.description FROM mytable foo"
- ,dialect=firebird.dialect())
+ select([table1.alias()])
+ ,"SELECT mytable_1.myid, mytable_1.name, mytable_1.description FROM mytable AS mytable_1")
# create a select for a join of two tables. use_labels means the column names will have
# labels tablename_columnname, which become the column keys accessible off the Selectable object.
@@ -401,6 +446,12 @@ myothertable.otherid AS myothertable_otherid FROM mytable, myothertable \
WHERE mytable.myid = myothertable.otherid) AS t2view WHERE t2view.mytable_myid = :t2view_mytable_myid"
)
+
+ def test_prefixes(self):
+ self.runtest(table1.select().prefix_with("SQL_CALC_FOUND_ROWS").prefix_with("SQL_SOME_WEIRD_MYSQL_THING"),
+ "SELECT SQL_CALC_FOUND_ROWS SQL_SOME_WEIRD_MYSQL_THING mytable.myid, mytable.name, mytable.description FROM mytable"
+ )
+
def testtext(self):
self.runtest(
text("select * from foo where lala = bar") ,
@@ -429,7 +480,7 @@ WHERE mytable.myid = myothertable.otherid) AS t2view WHERE t2view.mytable_myid =
s.append_column("column2")
s.append_whereclause("column1=12")
s.append_whereclause("column2=19")
- s.order_by("column1")
+ s = s.order_by("column1")
s.append_from("table1")
self.runtest(s, "SELECT column1, column2 FROM table1 WHERE column1=12 AND column2=19 ORDER BY column1")
@@ -468,7 +519,14 @@ WHERE mytable.myid = myothertable.otherid) AS t2view WHERE t2view.mytable_myid =
checkparams={'bar':4, 'whee': 7},
params={'bar':4, 'whee': 7, 'hoho':10},
)
-
+
+ self.runtest(
+ text("select * from foo where clock='05:06:07'"),
+ "select * from foo where clock='05:06:07'",
+ checkparams={},
+ params={},
+ )
+
dialect = postgres.dialect()
self.runtest(
text("select * from foo where lala=:bar and hoho=:whee"),
@@ -477,6 +535,13 @@ WHERE mytable.myid = myothertable.otherid) AS t2view WHERE t2view.mytable_myid =
params={'bar':4, 'whee': 7, 'hoho':10},
dialect=dialect
)
+ self.runtest(
+ text("select * from foo where clock='05:06:07' and mork='\:mindy'"),
+ "select * from foo where clock='05:06:07' and mork=':mindy'",
+ checkparams={},
+ params={},
+ dialect=dialect
+ )
dialect = sqlite.dialect()
self.runtest(
@@ -509,7 +574,7 @@ FROM mytable, myothertable WHERE foo.id = foofoo(lala) AND datetime(foo) = Today
def testliteral(self):
self.runtest(select([literal("foo") + literal("bar")], from_obj=[table1]),
- "SELECT :literal + :literal_1 FROM mytable")
+ "SELECT :literal || :literal_1 FROM mytable")
def testcalculatedcolumns(self):
value_tbl = table('values',
@@ -663,7 +728,7 @@ FROM myothertable ORDER BY myid \
WHERE mytable.name = :mytable_name GROUP BY mytable.myid, mytable.name UNION SELECT mytable.myid, mytable.name, mytable.description \
FROM mytable WHERE mytable.name = :mytable_name_1"
)
-
+
def test_compound_select_grouping(self):
self.runtest(
union_all(
@@ -716,6 +781,7 @@ EXISTS (select yay from foo where boo = lar)",
dialect=postgres.dialect()
)
+
self.runtest(query,
"SELECT mytable.myid, mytable.name, mytable.description, myothertable.otherid, myothertable.othername \
FROM mytable, myothertable WHERE mytable.myid = myothertable.otherid(+) AND \
@@ -835,16 +901,16 @@ myothertable.othername != :myothertable_othername OR EXISTS (select yay from foo
self.runtest(select([table1], table1.c.myid.in_('a', literal('b'))),
"SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid IN (:mytable_myid, :literal)")
- self.runtest(select([table1], table1.c.myid.in_(literal('a') + 'a')),
+ self.runtest(select([table1], table1.c.myid.in_(literal(1) + 'a')),
"SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid = :literal + :literal_1")
self.runtest(select([table1], table1.c.myid.in_(literal('a') +'a', 'b')),
- "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid IN (:literal + :literal_1, :mytable_myid)")
+ "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid IN (:literal || :literal_1, :mytable_myid)")
self.runtest(select([table1], table1.c.myid.in_(literal('a') + literal('a'), literal('b'))),
- "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid IN (:literal + :literal_1, :literal_2)")
+ "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid IN (:literal || :literal_1, :literal_2)")
- self.runtest(select([table1], table1.c.myid.in_('a', literal('b') +'b')),
+ self.runtest(select([table1], table1.c.myid.in_(1, literal(3) + 4)),
"SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid IN (:mytable_myid, :literal + :literal_1)")
self.runtest(select([table1], table1.c.myid.in_(literal('a') < 'b')),
@@ -862,7 +928,7 @@ myothertable.othername != :myothertable_othername OR EXISTS (select yay from foo
self.runtest(select([table1], table1.c.myid.in_(literal('a'), table1.c.myid +'a')),
"SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid IN (:literal, mytable.myid + :mytable_myid)")
- self.runtest(select([table1], table1.c.myid.in_(literal('a'), 'a' + table1.c.myid)),
+ self.runtest(select([table1], table1.c.myid.in_(literal(1), 'a' + table1.c.myid)),
"SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid IN (:literal, :mytable_myid + mytable.myid)")
self.runtest(select([table1], table1.c.myid.in_(1, 2, 3)),
@@ -900,16 +966,6 @@ UNION SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE
"SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE (CASE WHEN (mytable.myid IS NULL) THEN NULL ELSE 0 END = 1)")
- def testlateargs(self):
- """tests that a SELECT clause will have extra "WHERE" clauses added to it at compile time if extra arguments
- are sent"""
-
- self.runtest(table1.select(), "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.name = :mytable_name AND mytable.myid = :mytable_myid", params={'myid':'3', 'name':'jack'})
-
- self.runtest(table1.select(table1.c.name=='jack'), "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid = :mytable_myid AND mytable.name = :mytable_name", params={'myid':'3'})
-
- self.runtest(table1.select(table1.c.name=='jack'), "SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE mytable.myid = :mytable_myid AND mytable.name = :mytable_name", params={'myid':'3', 'name':'fred'})
-
def testcast(self):
tbl = table('casttest',
column('id', Integer),
@@ -963,8 +1019,8 @@ UNION SELECT mytable.myid, mytable.name, mytable.description FROM mytable WHERE
"SELECT op.field FROM op WHERE :literal + (op.field IN (:op_field, :op_field_1))")
self.runtest(table.select((5 + table.c.field).in_(5,6)),
"SELECT op.field FROM op WHERE :op_field + op.field IN (:literal, :literal_1)")
- self.runtest(table.select(not_(table.c.field == 5)),
- "SELECT op.field FROM op WHERE NOT op.field = :op_field")
+ self.runtest(table.select(not_(and_(table.c.field == 5, table.c.field == 7))),
+ "SELECT op.field FROM op WHERE NOT (op.field = :op_field AND op.field = :op_field_1)")
self.runtest(table.select(not_(table.c.field) == 5),
"SELECT op.field FROM op WHERE (NOT op.field) = :literal")
self.runtest(table.select((table.c.field == table.c.field).between(False, True)),
@@ -1019,12 +1075,17 @@ class CRUDTest(SQLTest):
values = {
table1.c.name : table1.c.name + "lala",
table1.c.myid : func.do_stuff(table1.c.myid, literal('hoho'))
- }), "UPDATE mytable SET myid=do_stuff(mytable.myid, :literal_2), name=mytable.name + :mytable_name WHERE mytable.myid = hoho(:hoho) AND mytable.name = :literal + mytable.name + :literal_1")
+ }), "UPDATE mytable SET myid=do_stuff(mytable.myid, :literal), name=(mytable.name || :mytable_name) "
+ "WHERE mytable.myid = hoho(:hoho) AND mytable.name = :literal_1 || mytable.name || :literal_2")
def testcorrelatedupdate(self):
# test against a straight text subquery
- u = update(table1, values = {table1.c.name : text("select name from mytable where id=mytable.id")})
+ u = update(table1, values = {table1.c.name : text("(select name from mytable where id=mytable.id)")})
self.runtest(u, "UPDATE mytable SET name=(select name from mytable where id=mytable.id)")
+
+ mt = table1.alias()
+ u = update(table1, values = {table1.c.name : select([mt.c.name], mt.c.myid==table1.c.myid)})
+ self.runtest(u, "UPDATE mytable SET name=(SELECT mytable_1.name FROM mytable AS mytable_1 WHERE mytable_1.myid = mytable.myid)")
# test against a regular constructed subquery
s = select([table2], table2.c.otherid == table1.c.myid)
@@ -1043,7 +1104,18 @@ class CRUDTest(SQLTest):
def testdelete(self):
self.runtest(delete(table1, table1.c.myid == 7), "DELETE FROM mytable WHERE mytable.myid = :mytable_myid")
-
+
+ def testcorrelateddelete(self):
+ # test a non-correlated WHERE clause
+ s = select([table2.c.othername], table2.c.otherid == 7)
+ u = delete(table1, table1.c.name==s)
+ self.runtest(u, "DELETE FROM mytable WHERE mytable.name = (SELECT myothertable.othername FROM myothertable WHERE myothertable.otherid = :myothertable_otherid)")
+
+ # test one that is actually correlated...
+ s = select([table2.c.othername], table2.c.otherid == table1.c.myid)
+ u = table1.delete(table1.c.name==s)
+ self.runtest(u, "DELETE FROM mytable WHERE mytable.name = (SELECT myothertable.othername FROM myothertable WHERE myothertable.otherid = mytable.myid)")
+
class SchemaTest(SQLTest):
def testselect(self):
# these tests will fail with the MS-SQL compiler since it will alias schema-qualified tables
diff --git a/test/sql/selectable.py b/test/sql/selectable.py
index ecd8253b8..dcc855074 100755
--- a/test/sql/selectable.py
+++ b/test/sql/selectable.py
@@ -1,17 +1,13 @@
-"""tests that various From objects properly export their columns, as well as useable primary keys
-and foreign keys. Full relational algebra depends on every selectable unit behaving
-nicely with others.."""
-
+"""tests that various From objects properly export their columns, as well as
+useable primary keys and foreign keys. Full relational algebra depends on
+every selectable unit behaving nicely with others.."""
+
import testbase
-import unittest, sys, datetime
-
-
-db = testbase.db
-
from sqlalchemy import *
+from testlib import *
-
-table = Table('table1', db,
+metadata = MetaData()
+table = Table('table1', metadata,
Column('col1', Integer, primary_key=True),
Column('col2', String(20)),
Column('col3', Integer),
@@ -19,14 +15,14 @@ table = Table('table1', db,
)
-table2 = Table('table2', db,
+table2 = Table('table2', metadata,
Column('col1', Integer, primary_key=True),
Column('col2', Integer, ForeignKey('table1.col1')),
Column('col3', String(20)),
Column('coly', Integer),
)
-class SelectableTest(testbase.AssertMixin):
+class SelectableTest(AssertMixin):
def testdistance(self):
s = select([table.c.col1.label('c2'), table.c.col1, table.c.col1.label('c1')])
@@ -57,7 +53,7 @@ class SelectableTest(testbase.AssertMixin):
jj = select([ table.c.col1.label('bar_col1')],from_obj=[j]).alias('foo')
jjj = join(table, jj, table.c.col1==jj.c.bar_col1)
assert jjj.corresponding_column(jjj.c.table1_col1) is jjj.c.table1_col1
-
+
j2 = jjj.alias('foo')
print j2.corresponding_column(jjj.c.table1_col1)
assert j2.corresponding_column(jjj.c.table1_col1) is j2.c.table1_col1
@@ -170,8 +166,9 @@ class SelectableTest(testbase.AssertMixin):
print str(criterion)
print str(j.onclause)
self.assert_(criterion.compare(j.onclause))
+
-class PrimaryKeyTest(testbase.AssertMixin):
+class PrimaryKeyTest(AssertMixin):
def test_join_pk_collapse_implicit(self):
"""test that redundant columns in a join get 'collapsed' into a minimal primary key,
which is the root column along a chain of foreign key relationships."""
@@ -224,8 +221,7 @@ class PrimaryKeyTest(testbase.AssertMixin):
j.foreign_keys
assert list(j.primary_key) == [a.c.id]
-
-
+
if __name__ == "__main__":
testbase.main()
- \ No newline at end of file
+
diff --git a/test/sql/testtypes.py b/test/sql/testtypes.py
index ed9de0912..659033016 100644
--- a/test/sql/testtypes.py
+++ b/test/sql/testtypes.py
@@ -1,14 +1,11 @@
-from testbase import PersistTest, AssertMixin
import testbase
import pickleable
+import datetime, os
from sqlalchemy import *
-import string,datetime, re, sys, os
import sqlalchemy.engine.url as url
-
-import sqlalchemy.types
from sqlalchemy.databases import mssql, oracle, mysql
+from testlib import *
-db = testbase.db
class MyType(types.TypeEngine):
def get_col_spec(self):
@@ -107,7 +104,7 @@ class OverrideTest(PersistTest):
def setUpAll(self):
global users
- users = Table('type_users', db,
+ users = Table('type_users', MetaData(testbase.db),
Column('user_id', Integer, primary_key = True),
# totall custom type
Column('goofy', MyType, nullable = False),
@@ -138,11 +135,12 @@ class ColumnsTest(AssertMixin):
'float_column': 'float_column NUMERIC(25, 2)'
}
+ db = testbase.db
if not db.name=='sqlite' and not db.name=='oracle':
expectedResults['float_column'] = 'float_column FLOAT(25)'
print db.engine.__module__
- testTable = Table('testColumns', db,
+ testTable = Table('testColumns', MetaData(db),
Column('int_column', Integer),
Column('smallint_column', Smallinteger),
Column('varchar_column', String(20)),
@@ -157,7 +155,8 @@ class UnicodeTest(AssertMixin):
"""tests the Unicode type. also tests the TypeDecorator with instances in the types package."""
def setUpAll(self):
global unicode_table
- unicode_table = Table('unicode_table', db,
+ metadata = MetaData(testbase.db)
+ unicode_table = Table('unicode_table', metadata,
Column('id', Integer, Sequence('uni_id_seq', optional=True), primary_key=True),
Column('unicode_varchar', Unicode(250)),
Column('unicode_text', Unicode),
@@ -175,49 +174,49 @@ class UnicodeTest(AssertMixin):
unicode_text=unicodedata,
plain_varchar=rawdata)
x = unicode_table.select().execute().fetchone()
- self.echo(repr(x['unicode_varchar']))
- self.echo(repr(x['unicode_text']))
- self.echo(repr(x['plain_varchar']))
+ print repr(x['unicode_varchar'])
+ print repr(x['unicode_text'])
+ print repr(x['plain_varchar'])
self.assert_(isinstance(x['unicode_varchar'], unicode) and x['unicode_varchar'] == unicodedata)
self.assert_(isinstance(x['unicode_text'], unicode) and x['unicode_text'] == unicodedata)
if isinstance(x['plain_varchar'], unicode):
# SQLLite and MSSQL return non-unicode data as unicode
- self.assert_(db.name in ('sqlite', 'mssql'))
+ self.assert_(testbase.db.name in ('sqlite', 'mssql'))
self.assert_(x['plain_varchar'] == unicodedata)
- self.echo("it's %s!" % db.name)
+ print "it's %s!" % testbase.db.name
else:
self.assert_(not isinstance(x['plain_varchar'], unicode) and x['plain_varchar'] == rawdata)
def testengineparam(self):
"""tests engine-wide unicode conversion"""
- prev_unicode = db.engine.dialect.convert_unicode
+ prev_unicode = testbase.db.engine.dialect.convert_unicode
try:
- db.engine.dialect.convert_unicode = True
+ testbase.db.engine.dialect.convert_unicode = True
rawdata = 'Alors vous imaginez ma surprise, au lever du jour, quand une dr\xc3\xb4le de petit voix m\xe2\x80\x99a r\xc3\xa9veill\xc3\xa9. Elle disait: \xc2\xab S\xe2\x80\x99il vous pla\xc3\xaet\xe2\x80\xa6 dessine-moi un mouton! \xc2\xbb\n'
unicodedata = rawdata.decode('utf-8')
unicode_table.insert().execute(unicode_varchar=unicodedata,
unicode_text=unicodedata,
plain_varchar=rawdata)
x = unicode_table.select().execute().fetchone()
- self.echo(repr(x['unicode_varchar']))
- self.echo(repr(x['unicode_text']))
- self.echo(repr(x['plain_varchar']))
+ print repr(x['unicode_varchar'])
+ print repr(x['unicode_text'])
+ print repr(x['plain_varchar'])
self.assert_(isinstance(x['unicode_varchar'], unicode) and x['unicode_varchar'] == unicodedata)
self.assert_(isinstance(x['unicode_text'], unicode) and x['unicode_text'] == unicodedata)
self.assert_(isinstance(x['plain_varchar'], unicode) and x['plain_varchar'] == unicodedata)
finally:
- db.engine.dialect.convert_unicode = prev_unicode
+ testbase.db.engine.dialect.convert_unicode = prev_unicode
- @testbase.unsupported('oracle')
+ @testing.unsupported('oracle')
def testlength(self):
"""checks the database correctly understands the length of a unicode string"""
teststr = u'aaa\x1234'
- self.assert_(db.func.length(teststr).scalar() == len(teststr))
+ self.assert_(testbase.db.func.length(teststr).scalar() == len(teststr))
class BinaryTest(AssertMixin):
def setUpAll(self):
global binary_table
- binary_table = Table('binary_table', db,
+ binary_table = Table('binary_table', MetaData(testbase.db),
Column('primary_id', Integer, Sequence('binary_id_seq', optional=True), primary_key=True),
Column('data', Binary),
Column('data_slice', Binary(100)),
@@ -244,39 +243,31 @@ class BinaryTest(AssertMixin):
binary_table.insert().execute(primary_id=1, misc='binary_data_one.dat', data=stream1, data_slice=stream1[0:100], pickled=testobj1)
binary_table.insert().execute(primary_id=2, misc='binary_data_two.dat', data=stream2, data_slice=stream2[0:99], pickled=testobj2)
binary_table.insert().execute(primary_id=3, misc='binary_data_two.dat', data=None, data_slice=stream2[0:99], pickled=None)
- l = binary_table.select(order_by=binary_table.c.primary_id).execute().fetchall()
- print type(stream1), type(l[0]['data']), type(l[0]['data_slice'])
- print len(stream1), len(l[0]['data']), len(l[0]['data_slice'])
- self.assert_(list(stream1) == list(l[0]['data']))
- self.assert_(list(stream1[0:100]) == list(l[0]['data_slice']))
- self.assert_(list(stream2) == list(l[1]['data']))
- self.assert_(testobj1 == l[0]['pickled'])
- self.assert_(testobj2 == l[1]['pickled'])
+
+ for stmt in (
+ binary_table.select(order_by=binary_table.c.primary_id),
+ text("select * from binary_table order by binary_table.primary_id", typemap={'pickled':PickleType}, bind=testbase.db)
+ ):
+ l = stmt.execute().fetchall()
+ print type(stream1), type(l[0]['data']), type(l[0]['data_slice'])
+ print len(stream1), len(l[0]['data']), len(l[0]['data_slice'])
+ self.assert_(list(stream1) == list(l[0]['data']))
+ self.assert_(list(stream1[0:100]) == list(l[0]['data_slice']))
+ self.assert_(list(stream2) == list(l[1]['data']))
+ self.assert_(testobj1 == l[0]['pickled'])
+ self.assert_(testobj2 == l[1]['pickled'])
def load_stream(self, name, len=12579):
f = os.path.join(os.path.dirname(testbase.__file__), name)
# put a number less than the typical MySQL default BLOB size
return file(f).read(len)
- @testbase.supported('oracle')
- def test_oracle_autobinary(self):
- stream1 =self.load_stream('binary_data_one.dat')
- stream2 =self.load_stream('binary_data_two.dat')
- binary_table.insert().execute(primary_id=1, misc='binary_data_one.dat', data=stream1, data_slice=stream1[0:100])
- binary_table.insert().execute(primary_id=2, misc='binary_data_two.dat', data=stream2, data_slice=stream2[0:99])
- binary_table.insert().execute(primary_id=3, misc='binary_data_two.dat', data=None, data_slice=stream2[0:99], pickled=None)
- result = testbase.db.connect().execute("select primary_id, misc, data, data_slice from binary_table")
- l = result.fetchall()
- l[0]['data']
- self.assert_(list(stream1) == list(l[0]['data']))
- self.assert_(list(stream1[0:100]) == list(l[0]['data_slice']))
- self.assert_(list(stream2) == list(l[1]['data']))
-
class DateTest(AssertMixin):
def setUpAll(self):
global users_with_date, insert_data
+ db = testbase.db
if db.engine.name == 'oracle':
import sqlalchemy.databases.oracle as oracle
insert_data = [
@@ -314,13 +305,14 @@ class DateTest(AssertMixin):
if db.engine.name == 'mssql':
# MSSQL Datetime values have only a 3.33 milliseconds precision
insert_data[2] = [9, 'foo', datetime.datetime(2005, 11, 10, 11, 52, 35, 547000), datetime.date(1970,4,1), datetime.time(23,59,59,997000)]
-
+
fnames = ['user_id', 'user_name', 'user_datetime', 'user_date', 'user_time']
collist = [Column('user_id', INT, primary_key = True), Column('user_name', VARCHAR(20)), Column('user_datetime', DateTime(timezone=False)),
Column('user_date', Date), Column('user_time', Time)]
- users_with_date = Table('query_users_with_date', db, *collist)
+ users_with_date = Table('query_users_with_date',
+ MetaData(testbase.db), *collist)
users_with_date.create()
insert_dicts = [dict(zip(fnames, d)) for d in insert_data]
@@ -338,7 +330,7 @@ class DateTest(AssertMixin):
def testtextdate(self):
- x = db.text("select user_datetime from query_users_with_date", typemap={'user_datetime':DateTime}).execute().fetchall()
+ x = testbase.db.text("select user_datetime from query_users_with_date", typemap={'user_datetime':DateTime}).execute().fetchall()
print repr(x)
self.assert_(isinstance(x[0][0], datetime.datetime))
@@ -347,9 +339,13 @@ class DateTest(AssertMixin):
#print repr(x)
def testdate2(self):
- t = Table('testdate', testbase.metadata, Column('id', Integer, Sequence('datetest_id_seq', optional=True), primary_key=True),
+ meta = MetaData(testbase.db)
+ t = Table('testdate', meta,
+ Column('id', Integer,
+ Sequence('datetest_id_seq', optional=True),
+ primary_key=True),
Column('adate', Date), Column('adatetime', DateTime))
- t.create()
+ t.create(checkfirst=True)
try:
d1 = datetime.date(2007, 10, 30)
t.insert().execute(adate=d1, adatetime=d1)
@@ -361,8 +357,43 @@ class DateTest(AssertMixin):
self.assert_(x.adatetime.__class__ == datetime.datetime)
finally:
- t.drop()
+ t.drop(checkfirst=True)
+class NumericTest(AssertMixin):
+ def setUpAll(self):
+ global numeric_table, metadata
+ metadata = MetaData(testbase.db)
+ numeric_table = Table('numeric_table', metadata,
+ Column('id', Integer, Sequence('numeric_id_seq', optional=True), primary_key=True),
+ Column('numericcol', Numeric(asdecimal=False)),
+ Column('floatcol', Float),
+ Column('ncasdec', Numeric),
+ Column('fcasdec', Float(asdecimal=True))
+ )
+ metadata.create_all()
+
+ def tearDownAll(self):
+ metadata.drop_all()
+
+ def tearDown(self):
+ numeric_table.delete().execute()
+
+ def test_decimal(self):
+ from decimal import Decimal
+ numeric_table.insert().execute(numericcol=3.5, floatcol=5.6, ncasdec=12.4, fcasdec=15.78)
+ numeric_table.insert().execute(numericcol=Decimal("3.5"), floatcol=Decimal("5.6"), ncasdec=Decimal("12.4"), fcasdec=Decimal("15.78"))
+ l = numeric_table.select().execute().fetchall()
+ print l
+ rounded = [
+ (l[0][0], l[0][1], round(l[0][2], 5), l[0][3], l[0][4]),
+ (l[1][0], l[1][1], round(l[1][2], 5), l[1][3], l[1][4]),
+ ]
+ assert rounded == [
+ (1, 3.5, 5.6, Decimal("12.4"), Decimal("15.78")),
+ (2, 3.5, 5.6, Decimal("12.4"), Decimal("15.78")),
+ ]
+
+
class IntervalTest(AssertMixin):
def setUpAll(self):
global interval_table, metadata
diff --git a/test/sql/unicode.py b/test/sql/unicode.py
index 7ce42bf4c..f882c2a5f 100644
--- a/test/sql/unicode.py
+++ b/test/sql/unicode.py
@@ -1,23 +1,27 @@
# coding: utf-8
-import testbase
+"""verrrrry basic unicode column name testing"""
+import testbase
from sqlalchemy import *
+from sqlalchemy.orm import mapper, relation, create_session, eagerload
+from testlib import *
-"""verrrrry basic unicode column name testing"""
-class UnicodeSchemaTest(testbase.PersistTest):
+class UnicodeSchemaTest(PersistTest):
def setUpAll(self):
- global metadata, t1, t2
- metadata = MetaData(engine=testbase.db)
+ global unicode_bind, metadata, t1, t2
+
+ unicode_bind = self._unicode_bind()
+
+ metadata = MetaData(unicode_bind)
t1 = Table('unitable1', metadata,
Column(u'méil', Integer, primary_key=True),
- Column(u'éXXm', Integer),
+ Column(u'\u6e2c\u8a66', Integer),
)
- t2 = Table(u'unitéble2', metadata,
+ t2 = Table(u'Unitéble2', metadata,
Column(u'méil', Integer, primary_key=True, key="a"),
- Column(u'éXXm', Integer, ForeignKey(u'unitable1.méil'), key="b"),
-
+ Column(u'\u6e2c\u8a66', Integer, ForeignKey(u'unitable1.méil'), key="b"),
)
metadata.create_all()
@@ -26,24 +30,46 @@ class UnicodeSchemaTest(testbase.PersistTest):
t1.delete().execute()
def tearDownAll(self):
+ global unicode_bind
metadata.drop_all()
+ del unicode_bind
+
+ def _unicode_bind(self):
+ if testbase.db.name != 'mysql':
+ return testbase.db
+ else:
+ # most mysql installations don't default to utf8 connections
+ version = testbase.db.dialect.get_version_info(testbase.db)
+ if version < (4, 1):
+ raise AssertionError("Unicode not supported on MySQL < 4.1")
+
+ c = testbase.db.connect()
+ if not hasattr(c.connection.connection, 'set_character_set'):
+ raise AssertionError(
+ "Unicode not supported on this MySQL-python version")
+ else:
+ c.connection.set_character_set('utf8')
+ c.detach()
+
+ return c
def test_insert(self):
- t1.insert().execute({u'méil':1, u'éXXm':5})
+ t1.insert().execute({u'méil':1, u'\u6e2c\u8a66':5})
t2.insert().execute({'a':1, 'b':1})
assert t1.select().execute().fetchall() == [(1, 5)]
assert t2.select().execute().fetchall() == [(1, 1)]
def test_reflect(self):
- t1.insert().execute({u'méil':2, u'éXXm':7})
+ t1.insert().execute({u'méil':2, u'\u6e2c\u8a66':7})
t2.insert().execute({'a':2, 'b':2})
- meta = MetaData(testbase.db)
+ meta = MetaData(unicode_bind)
tt1 = Table(t1.name, meta, autoload=True)
tt2 = Table(t2.name, meta, autoload=True)
- tt1.insert().execute({u'méil':1, u'éXXm':5})
- tt2.insert().execute({u'méil':1, u'éXXm':1})
+
+ tt1.insert().execute({u'méil':1, u'\u6e2c\u8a66':5})
+ tt2.insert().execute({u'méil':1, u'\u6e2c\u8a66':1})
assert tt1.select(order_by=desc(u'méil')).execute().fetchall() == [(2, 7), (1, 5)]
assert tt2.select(order_by=desc(u'méil')).execute().fetchall() == [(2, 2), (1, 1)]
@@ -57,7 +83,7 @@ class UnicodeSchemaTest(testbase.PersistTest):
mapper(A, t1, properties={
't2s':relation(B),
'a':t1.c[u'méil'],
- 'b':t1.c[u'éXXm']
+ 'b':t1.c[u'\u6e2c\u8a66']
})
mapper(B, t2)
sess = create_session()
diff --git a/test/testbase.py b/test/testbase.py
index 7c5095d1a..1195db340 100644
--- a/test/testbase.py
+++ b/test/testbase.py
@@ -1,470 +1,14 @@
-import sys
-import os, unittest, StringIO, re, ConfigParser
-sys.path.insert(0, os.path.join(os.getcwd(), 'lib'))
-import sqlalchemy
-from sqlalchemy import sql, engine, pool
-import sqlalchemy.engine.base as base
-import optparse
-from sqlalchemy.schema import MetaData
-from sqlalchemy.orm import clear_mappers
-
-db = None
-metadata = None
-db_uri = None
-echo = True
-
-# redefine sys.stdout so all those print statements go to the echo func
-local_stdout = sys.stdout
-class Logger(object):
- def write(self, msg):
- if echo:
- local_stdout.write(msg)
- def flush(self):
- pass
-
-def echo_text(text):
- print text
-
-def parse_argv():
- # we are using the unittest main runner, so we are just popping out the
- # arguments we need instead of using our own getopt type of thing
- global db, db_uri, metadata
-
- DBTYPE = 'sqlite'
- PROXY = False
-
- base_config = """
-[db]
-sqlite=sqlite:///:memory:
-sqlite_file=sqlite:///querytest.db
-postgres=postgres://scott:tiger@127.0.0.1:5432/test
-mysql=mysql://scott:tiger@127.0.0.1:3306/test
-oracle=oracle://scott:tiger@127.0.0.1:1521
-oracle8=oracle://scott:tiger@127.0.0.1:1521/?use_ansi=0
-mssql=mssql://scott:tiger@SQUAWK\\SQLEXPRESS/test
-firebird=firebird://sysdba:s@localhost/tmp/test.fdb
-"""
- config = ConfigParser.ConfigParser()
- config.readfp(StringIO.StringIO(base_config))
- config.read(['test.cfg', os.path.expanduser('~/.satest.cfg')])
-
- parser = optparse.OptionParser(usage = "usage: %prog [options] [tests...]")
- parser.add_option("--dburi", action="store", dest="dburi", help="database uri (overrides --db)")
- parser.add_option("--db", action="store", dest="db", default="sqlite", help="prefab database uri (%s)" % ', '.join(config.options('db')))
- parser.add_option("--mockpool", action="store_true", dest="mockpool", help="use mock pool (asserts only one connection used)")
- parser.add_option("--verbose", action="store_true", dest="verbose", help="enable stdout echoing/printing")
- parser.add_option("--quiet", action="store_true", dest="quiet", help="suppress unittest output")
- parser.add_option("--log-info", action="append", dest="log_info", help="turn on info logging for <LOG> (multiple OK)")
- parser.add_option("--log-debug", action="append", dest="log_debug", help="turn on debug logging for <LOG> (multiple OK)")
- parser.add_option("--nothreadlocal", action="store_true", dest="nothreadlocal", help="dont use thread-local mod")
- parser.add_option("--enginestrategy", action="store", default=None, dest="enginestrategy", help="engine strategy (plain or threadlocal, defaults to plain)")
- parser.add_option("--coverage", action="store_true", dest="coverage", help="Dump a full coverage report after running")
- parser.add_option("--reversetop", action="store_true", dest="topological", help="Reverse the collection ordering for topological sorts (helps reveal dependency issues)")
- parser.add_option("--serverside", action="store_true", dest="serverside", help="Turn on server side cursors for PG")
- parser.add_option("--require", action="append", dest="require", help="Require a particular driver or module version", default=[])
-
- (options, args) = parser.parse_args()
- sys.argv[1:] = args
-
- if options.dburi:
- db_uri = param = options.dburi
- DBTYPE = db_uri[:db_uri.index(':')]
- elif options.db:
- DBTYPE = param = options.db
-
- if options.require or (config.has_section('require') and
- config.items('require')):
- try:
- import pkg_resources
- except ImportError:
- raise "setuptools is required for version requirements"
-
- cmdline = []
- for requirement in options.require:
- pkg_resources.require(requirement)
- cmdline.append(re.split('\s*(<!>=)', requirement, 1)[0])
-
- if config.has_section('require'):
- for label, requirement in config.items('require'):
- if not label == DBTYPE or label.startswith('%s.' % DBTYPE):
- continue
- seen = [c for c in cmdline if requirement.startswith(c)]
- if seen:
- continue
- pkg_resources.require(requirement)
-
- opts = {}
- if (None == db_uri):
- if DBTYPE not in config.options('db'):
- raise ("Could not create engine. specify --db <%s> to "
- "test runner." % '|'.join(config.options('db')))
-
- db_uri = config.get('db', DBTYPE)
-
- if not db_uri:
- raise "Could not create engine. specify --db <sqlite|sqlite_file|postgres|mysql|oracle|oracle8|mssql|firebird> to test runner."
-
- if not options.nothreadlocal:
- __import__('sqlalchemy.mods.threadlocal')
- sqlalchemy.mods.threadlocal.uninstall_plugin()
-
- global echo
- echo = options.verbose and not options.quiet
-
- global quiet
- quiet = options.quiet
-
- global with_coverage
- with_coverage = options.coverage
-
- if options.serverside:
- opts['server_side_cursors'] = True
-
- if options.enginestrategy is not None:
- opts['strategy'] = options.enginestrategy
- if options.mockpool:
- db = engine.create_engine(db_uri, poolclass=pool.AssertionPool, **opts)
- else:
- db = engine.create_engine(db_uri, **opts)
-
- # decorate the dialect's create_execution_context() method
- # to produce a wrapper
- create_context = db.dialect.create_execution_context
- def create_exec_context(*args, **kwargs):
- return ExecutionContextWrapper(create_context(*args, **kwargs))
- db.dialect.create_execution_context = create_exec_context
-
- global testdata
- testdata = TestData(db)
-
- if options.topological:
- from sqlalchemy.orm import unitofwork
- from sqlalchemy import topological
- class RevQueueDepSort(topological.QueueDependencySorter):
- def __init__(self, tuples, allitems):
- self.tuples = list(tuples)
- self.allitems = list(allitems)
- self.tuples.reverse()
- self.allitems.reverse()
- topological.QueueDependencySorter = RevQueueDepSort
- unitofwork.DependencySorter = RevQueueDepSort
-
- import logging
- logging.basicConfig()
- if options.log_info is not None:
- for elem in options.log_info:
- logging.getLogger(elem).setLevel(logging.INFO)
- if options.log_debug is not None:
- for elem in options.log_debug:
- logging.getLogger(elem).setLevel(logging.DEBUG)
- metadata = sqlalchemy.MetaData(db)
-
-def unsupported(*dbs):
- """a decorator that marks a test as unsupported by one or more database implementations"""
- def decorate(func):
- name = db.name
- for d in dbs:
- if d == name:
- def lala(self):
- echo_text("'" + func.__name__ + "' unsupported on DB implementation '" + name + "'")
- lala.__name__ = func.__name__
- return lala
- else:
- return func
- return decorate
-
-def supported(*dbs):
- """a decorator that marks a test as supported by one or more database implementations"""
- def decorate(func):
- name = db.name
- for d in dbs:
- if d == name:
- return func
- else:
- def lala(self):
- echo_text("'" + func.__name__ + "' unsupported on DB implementation '" + name + "'")
- lala.__name__ = func.__name__
- return lala
- return decorate
+"""First import for all test cases, sets sys.path and loads configuration."""
-
-class PersistTest(unittest.TestCase):
- """persist base class, provides default setUpAll, tearDownAll and echo functionality"""
- def __init__(self, *args, **params):
- unittest.TestCase.__init__(self, *args, **params)
- def echo(self, text):
- echo_text(text)
- def install_threadlocal(self):
- sqlalchemy.mods.threadlocal.install_plugin()
- def uninstall_threadlocal(self):
- sqlalchemy.mods.threadlocal.uninstall_plugin()
- def setUpAll(self):
- pass
- def tearDownAll(self):
- pass
- def shortDescription(self):
- """overridden to not return docstrings"""
- return None
+__all__ = 'db',
-class AssertMixin(PersistTest):
- """given a list-based structure of keys/properties which represent information within an object structure, and
- a list of actual objects, asserts that the list of objects corresponds to the structure."""
- def assert_result(self, result, class_, *objects):
- result = list(result)
- if echo:
- print repr(result)
- self.assert_list(result, class_, objects)
- def assert_list(self, result, class_, list):
- self.assert_(len(result) == len(list), "result list is not the same size as test list, for class " + class_.__name__)
- for i in range(0, len(list)):
- self.assert_row(class_, result[i], list[i])
- def assert_row(self, class_, rowobj, desc):
- self.assert_(rowobj.__class__ is class_, "item class is not " + repr(class_))
- for key, value in desc.iteritems():
- if isinstance(value, tuple):
- if isinstance(value[1], list):
- self.assert_list(getattr(rowobj, key), value[0], value[1])
- else:
- self.assert_row(value[0], getattr(rowobj, key), value[1])
- else:
- self.assert_(getattr(rowobj, key) == value, "attribute %s value %s does not match %s" % (key, getattr(rowobj, key), value))
- def assert_sql(self, db, callable_, list, with_sequences=None):
- global testdata
- testdata = TestData(db)
- if with_sequences is not None and (db.engine.name == 'postgres' or db.engine.name == 'oracle'):
- testdata.set_assert_list(self, with_sequences)
- else:
- testdata.set_assert_list(self, list)
- try:
- callable_()
- finally:
- testdata.set_assert_list(None, None)
-
- def assert_sql_count(self, db, callable_, count):
- global testdata
- testdata = TestData(db)
- try:
- callable_()
- finally:
- self.assert_(testdata.sql_count == count, "desired statement count %d does not match %d" % (count, testdata.sql_count))
-
- def capture_sql(self, db, callable_):
- global testdata
- testdata = TestData(db)
- buffer = StringIO.StringIO()
- testdata.buffer = buffer
- try:
- callable_()
- return buffer.getvalue()
- finally:
- testdata.buffer = None
-
-class ORMTest(AssertMixin):
- keep_mappers = False
- keep_data = False
- def setUpAll(self):
- global metadata
- metadata = MetaData(db)
- self.define_tables(metadata)
- metadata.create_all()
- self.insert_data()
- def define_tables(self, metadata):
- raise NotImplementedError()
- def insert_data(self):
- pass
- def get_metadata(self):
- return metadata
- def tearDownAll(self):
- metadata.drop_all()
- def tearDown(self):
- if not self.keep_mappers:
- clear_mappers()
- if not self.keep_data:
- for t in metadata.table_iterator(reverse=True):
- t.delete().execute().close()
-
-class TestData(object):
- def __init__(self, engine):
- self._engine = engine
- self.logger = engine.logger
- self.set_assert_list(None, None)
- self.sql_count = 0
- self.buffer = None
-
- def set_assert_list(self, unittest, list):
- self.unittest = unittest
- self.assert_list = list
- if list is not None:
- self.assert_list.reverse()
-
-class ExecutionContextWrapper(object):
- def __init__(self, ctx):
- self.__dict__['ctx'] = ctx
- def __getattr__(self, key):
- return getattr(self.ctx, key)
- def __setattr__(self, key, value):
- setattr(self.ctx, key, value)
-
- def post_exec(self):
- ctx = self.ctx
- statement = unicode(ctx.compiled)
- statement = re.sub(r'\n', '', ctx.statement)
- if testdata.buffer is not None:
- testdata.buffer.write(statement + "\n")
-
- if testdata.assert_list is not None:
- item = testdata.assert_list[-1]
- if not isinstance(item, dict):
- item = testdata.assert_list.pop()
- else:
- # asserting a dictionary of statements->parameters
- # this is to specify query assertions where the queries can be in
- # multiple orderings
- if not item.has_key('_converted'):
- for key in item.keys():
- ckey = self.convert_statement(key)
- item[ckey] = item[key]
- if ckey != key:
- del item[key]
- item['_converted'] = True
- try:
- entry = item.pop(statement)
- if len(item) == 1:
- testdata.assert_list.pop()
- item = (statement, entry)
- except KeyError:
- self.unittest.assert_(False, "Testing for one of the following queries: %s, received '%s'" % (repr([k for k in item.keys()]), statement))
-
- (query, params) = item
- if callable(params):
- params = params(ctx)
- if params is not None and isinstance(params, list) and len(params) == 1:
- params = params[0]
-
- if isinstance(ctx.compiled_parameters, sql.ClauseParameters):
- parameters = ctx.compiled_parameters.get_original_dict()
- elif isinstance(ctx.compiled_parameters, list):
- parameters = [p.get_original_dict() for p in ctx.compiled_parameters]
-
- query = self.convert_statement(query)
- if db.engine.name == 'mssql' and statement.endswith('; select scope_identity()'):
- statement = statement[:-25]
- testdata.unittest.assert_(statement == query and (params is None or params == parameters), "Testing for query '%s' params %s, received '%s' with params %s" % (query, repr(params), statement, repr(parameters)))
- testdata.sql_count += 1
- self.ctx.post_exec()
-
- def convert_statement(self, query):
- paramstyle = self.ctx.dialect.paramstyle
- if paramstyle == 'named':
- pass
- elif paramstyle =='pyformat':
- query = re.sub(r':([\w_]+)', r"%(\1)s", query)
- else:
- # positional params
- repl = None
- if paramstyle=='qmark':
- repl = "?"
- elif paramstyle=='format':
- repl = r"%s"
- elif paramstyle=='numeric':
- repl = None
- query = re.sub(r':([\w_]+)', repl, query)
- return query
-
-class TTestSuite(unittest.TestSuite):
- """override unittest.TestSuite to provide per-TestCase class setUpAll() and tearDownAll() functionality"""
- def __init__(self, tests=()):
- if len(tests) >0 and isinstance(tests[0], PersistTest):
- self._initTest = tests[0]
- else:
- self._initTest = None
- unittest.TestSuite.__init__(self, tests)
-
- def do_run(self, result):
- """nice job unittest ! you switched __call__ and run() between py2.3 and 2.4 thereby
- making straight subclassing impossible !"""
- for test in self._tests:
- if result.shouldStop:
- break
- test(result)
- return result
-
- def run(self, result):
- return self(result)
-
- def __call__(self, result):
- try:
- if self._initTest is not None:
- self._initTest.setUpAll()
- except:
- result.addError(self._initTest, self.__exc_info())
- pass
- try:
- return self.do_run(result)
- finally:
- try:
- if self._initTest is not None:
- self._initTest.tearDownAll()
- except:
- result.addError(self._initTest, self.__exc_info())
- pass
-
- def __exc_info(self):
- """Return a version of sys.exc_info() with the traceback frame
- minimised; usually the top level of the traceback frame is not
- needed.
- ripped off out of unittest module since its double __
- """
- exctype, excvalue, tb = sys.exc_info()
- if sys.platform[:4] == 'java': ## tracebacks look different in Jython
- return (exctype, excvalue, tb)
- return (exctype, excvalue, tb)
-
-unittest.TestLoader.suiteClass = TTestSuite
-
-parse_argv()
-
-
-def runTests(suite):
- sys.stdout = Logger()
- runner = unittest.TextTestRunner(verbosity = quiet and 1 or 2)
- if with_coverage:
- return cover(lambda:runner.run(suite))
- else:
- return runner.run(suite)
-
-def covered_files():
- for rec in os.walk(os.path.dirname(sqlalchemy.__file__)):
- for x in rec[2]:
- if x.endswith('.py'):
- yield os.path.join(rec[0], x)
-
-def cover(callable_):
- import coverage
- coverage_client = coverage.the_coverage
- coverage_client.get_ready()
- coverage_client.exclude('#pragma[: ]+[nN][oO] [cC][oO][vV][eE][rR]')
- coverage_client.erase()
- coverage_client.start()
- try:
- return callable_()
- finally:
- global echo
- echo=True
- coverage_client.stop()
- coverage_client.save()
- coverage_client.report(list(covered_files()), show_missing=False, ignore_errors=False)
-
-def main(suite=None):
-
- if not suite:
- if len(sys.argv[1:]):
- suite =unittest.TestLoader().loadTestsFromNames(sys.argv[1:], __import__('__main__'))
- else:
- suite = unittest.TestLoader().loadTestsFromModule(__import__('__main__'))
-
- result = runTests(suite)
- sys.exit(not result.wasSuccessful())
+import sys, os, logging
+sys.path.insert(0, os.path.join(os.getcwd(), 'lib'))
+logging.basicConfig()
+import testlib.config
+testlib.config.configure()
+from testlib.testing import main
+db = testlib.config.db
diff --git a/test/testlib/__init__.py b/test/testlib/__init__.py
new file mode 100644
index 000000000..ff5c4c125
--- /dev/null
+++ b/test/testlib/__init__.py
@@ -0,0 +1,11 @@
+"""Enhance unittest and instrument SQLAlchemy classes for testing.
+
+Load after sqlalchemy imports to use instrumented stand-ins like Table.
+"""
+
+import testlib.config
+from testlib.schema import Table, Column
+import testlib.testing as testing
+from testlib.testing import PersistTest, AssertMixin, ORMTest
+import testlib.profiling
+
diff --git a/test/testlib/config.py b/test/testlib/config.py
new file mode 100644
index 000000000..f05cda46d
--- /dev/null
+++ b/test/testlib/config.py
@@ -0,0 +1,255 @@
+import optparse, os, sys, ConfigParser, StringIO
+logging, require = None, None
+
+__all__ = 'parser', 'configure', 'options',
+
+db, db_uri, db_type, db_label = None, None, None, None
+
+options = None
+file_config = None
+
+base_config = """
+[db]
+sqlite=sqlite:///:memory:
+sqlite_file=sqlite:///querytest.db
+postgres=postgres://scott:tiger@127.0.0.1:5432/test
+mysql=mysql://scott:tiger@127.0.0.1:3306/test
+oracle=oracle://scott:tiger@127.0.0.1:1521
+oracle8=oracle://scott:tiger@127.0.0.1:1521/?use_ansi=0
+mssql=mssql://scott:tiger@SQUAWK\\SQLEXPRESS/test
+firebird=firebird://sysdba:s@localhost/tmp/test.fdb
+"""
+
+parser = optparse.OptionParser(usage = "usage: %prog [options] [tests...]")
+
+def configure():
+ global options, config
+ global getopts_options, file_config
+
+ file_config = ConfigParser.ConfigParser()
+ file_config.readfp(StringIO.StringIO(base_config))
+ file_config.read(['test.cfg', os.path.expanduser('~/.satest.cfg')])
+
+ # Opt parsing can fire immediate actions, like logging and coverage
+ (options, args) = parser.parse_args()
+ sys.argv[1:] = args
+
+ # Lazy setup of other options (post coverage)
+ for fn in post_configure:
+ fn(options, file_config)
+
+ return options, file_config
+
+def _log(option, opt_str, value, parser):
+ global logging
+ if not logging:
+ import logging
+ logging.basicConfig()
+
+ if opt_str.endswith('-info'):
+ logging.getLogger(value).setLevel(logging.INFO)
+ elif opt_str.endswith('-debug'):
+ logging.getLogger(value).setLevel(logging.DEBUG)
+
+def _start_coverage(option, opt_str, value, parser):
+ import sys, atexit, coverage
+ true_out = sys.stdout
+
+ def _iter_covered_files():
+ import sqlalchemy
+ for rec in os.walk(os.path.dirname(sqlalchemy.__file__)):
+ for x in rec[2]:
+ if x.endswith('.py'):
+ yield os.path.join(rec[0], x)
+ def _stop():
+ coverage.stop()
+ true_out.write("\nPreparing coverage report...\n")
+ coverage.report(list(_iter_covered_files()),
+ show_missing=False, ignore_errors=False,
+ file=true_out)
+ atexit.register(_stop)
+ coverage.erase()
+ coverage.start()
+
+def _list_dbs(*args):
+ print "Available --db options (use --dburi to override)"
+ for macro in sorted(file_config.options('db')):
+ print "%20s\t%s" % (macro, file_config.get('db', macro))
+ sys.exit(0)
+
+opt = parser.add_option
+opt("--verbose", action="store_true", dest="verbose",
+ help="enable stdout echoing/printing")
+opt("--quiet", action="store_true", dest="quiet", help="suppress output")
+opt("--log-info", action="callback", type="string", callback=_log,
+ help="turn on info logging for <LOG> (multiple OK)")
+opt("--log-debug", action="callback", type="string", callback=_log,
+ help="turn on debug logging for <LOG> (multiple OK)")
+opt("--require", action="append", dest="require", default=[],
+ help="require a particular driver or module version (multiple OK)")
+opt("--db", action="store", dest="db", default="sqlite",
+ help="Use prefab database uri")
+opt('--dbs', action='callback', callback=_list_dbs,
+ help="List available prefab dbs")
+opt("--dburi", action="store", dest="dburi",
+ help="Database uri (overrides --db)")
+opt("--mockpool", action="store_true", dest="mockpool",
+ help="Use mock pool (asserts only one connection used)")
+opt("--enginestrategy", action="store", dest="enginestrategy", default=None,
+ help="Engine strategy (plain or threadlocal, defaults toplain)")
+opt("--reversetop", action="store_true", dest="reversetop", default=False,
+ help="Reverse the collection ordering for topological sorts (helps "
+ "reveal dependency issues)")
+opt("--serverside", action="store_true", dest="serverside",
+ help="Turn on server side cursors for PG")
+opt("--mysql-engine", action="store", dest="mysql_engine", default=None,
+ help="Use the specified MySQL storage engine for all tables, default is "
+ "a db-default/InnoDB combo.")
+opt("--table-option", action="append", dest="tableopts", default=[],
+ help="Add a dialect-specific table option, key=value")
+opt("--coverage", action="callback", callback=_start_coverage,
+ help="Dump a full coverage report after running tests")
+opt("--profile", action="append", dest="profile_targets", default=[],
+ help="Enable a named profile target (multiple OK.)")
+opt("--profile-sort", action="store", dest="profile_sort", default=None,
+ help="Sort profile stats with this comma-separated sort order")
+opt("--profile-limit", type="int", action="store", dest="profile_limit",
+ default=None,
+ help="Limit function count in profile stats")
+
+class _ordered_map(object):
+ def __init__(self):
+ self._keys = list()
+ self._data = dict()
+
+ def __setitem__(self, key, value):
+ if key not in self._keys:
+ self._keys.append(key)
+ self._data[key] = value
+
+ def __iter__(self):
+ for key in self._keys:
+ yield self._data[key]
+
+post_configure = _ordered_map()
+
+def _engine_uri(options, file_config):
+ global db_label, db_uri
+ db_label = 'sqlite'
+ if options.dburi:
+ db_uri = options.dburi
+ db_label = db_uri[:db_uri.index(':')]
+ elif options.db:
+ db_label = options.db
+ db_uri = None
+
+ if db_uri is None:
+ if db_label not in file_config.options('db'):
+ raise RuntimeError(
+ "Unknown engine. Specify --dbs for known engines.")
+ db_uri = file_config.get('db', db_label)
+post_configure['engine_uri'] = _engine_uri
+
+def _require(options, file_config):
+ if not(options.require or
+ (file_config.has_section('require') and
+ file_config.items('require'))):
+ return
+
+ try:
+ import pkg_resources
+ except ImportError:
+ raise RuntimeError("setuptools is required for version requirements")
+
+ cmdline = []
+ for requirement in options.require:
+ pkg_resources.require(requirement)
+ cmdline.append(re.split('\s*(<!>=)', requirement, 1)[0])
+
+ if file_config.has_section('require'):
+ for label, requirement in file_config.items('require'):
+ if not label == db_label or label.startswith('%s.' % db_label):
+ continue
+ seen = [c for c in cmdline if requirement.startswith(c)]
+ if seen:
+ continue
+ pkg_resources.require(requirement)
+post_configure['require'] = _require
+
+def _create_testing_engine(options, file_config):
+ from sqlalchemy import engine
+ global db, db_type
+ engine_opts = {}
+ if options.serverside:
+ engine_opts['server_side_cursors'] = True
+
+ if options.enginestrategy is not None:
+ engine_opts['strategy'] = options.enginestrategy
+
+ if options.mockpool:
+ db = engine.create_engine(db_uri, poolclass=pool.AssertionPool,
+ **engine_opts)
+ else:
+ db = engine.create_engine(db_uri, **engine_opts)
+ db_type = db.name
+
+ # decorate the dialect's create_execution_context() method
+ # to produce a wrapper
+ from testlib.testing import ExecutionContextWrapper
+
+ create_context = db.dialect.create_execution_context
+ def create_exec_context(*args, **kwargs):
+ return ExecutionContextWrapper(create_context(*args, **kwargs))
+ db.dialect.create_execution_context = create_exec_context
+post_configure['create_engine'] = _create_testing_engine
+
+def _set_table_options(options, file_config):
+ import testlib.schema
+
+ table_options = testlib.schema.table_options
+ for spec in options.tableopts:
+ key, value = spec.split('=')
+ table_options[key] = value
+
+ if options.mysql_engine:
+ table_options['mysql_engine'] = options.mysql_engine
+post_configure['table_options'] = _set_table_options
+
+def _reverse_topological(options, file_config):
+ if options.reversetop:
+ from sqlalchemy.orm import unitofwork
+ from sqlalchemy import topological
+ class RevQueueDepSort(topological.QueueDependencySorter):
+ def __init__(self, tuples, allitems):
+ self.tuples = list(tuples)
+ self.allitems = list(allitems)
+ self.tuples.reverse()
+ self.allitems.reverse()
+ topological.QueueDependencySorter = RevQueueDepSort
+ unitofwork.DependencySorter = RevQueueDepSort
+post_configure['topological'] = _reverse_topological
+
+def _set_profile_targets(options, file_config):
+ from testlib import profiling
+
+ profile_config = profiling.profile_config
+
+ for target in options.profile_targets:
+ profile_config['targets'].add(target)
+
+ if options.profile_sort:
+ profile_config['sort'] = options.profile_sort.split(',')
+
+ if options.profile_limit:
+ profile_config['limit'] = options.profile_limit
+
+ if options.quiet:
+ profile_config['report'] = False
+
+ # magic "all" target
+ if 'all' in profiling.all_targets:
+ targets = profile_config['targets']
+ if 'all' in targets and len(targets) != 1:
+ targets.clear()
+ targets.add('all')
+post_configure['profile_targets'] = _set_profile_targets
diff --git a/test/coverage.py b/test/testlib/coverage.py
index 66e55e0c4..0203dbf7d 100644
--- a/test/coverage.py
+++ b/test/testlib/coverage.py
@@ -22,7 +22,8 @@
# interface and limitations. See [GDR 2001-12-04b] for requirements and
# design.
-r"""Usage:
+r"""\
+Usage:
coverage.py -x [-p] MODULE.py [ARG1 ARG2 ...]
Execute module, passing the given command-line arguments, collecting
@@ -54,18 +55,27 @@ coverage.py -a [-d dir] [-o dir1,dir2,...] FILE1 FILE2 ...
Coverage data is saved in the file .coverage by default. Set the
COVERAGE_FILE environment variable to save it somewhere else."""
-__version__ = "2.6.20060823" # see detailed history at the end of this file.
+__version__ = "2.75.20070722" # see detailed history at the end of this file.
import compiler
import compiler.visitor
+import glob
import os
import re
import string
+import symbol
import sys
import threading
+import token
import types
from socket import gethostname
+# Python version compatibility
+try:
+ strclass = basestring # new to 2.3
+except:
+ strclass = str
+
# 2. IMPLEMENTATION
#
# This uses the "singleton" pattern.
@@ -87,6 +97,9 @@ from socket import gethostname
# names to increase speed.
class StatementFindingAstVisitor(compiler.visitor.ASTVisitor):
+ """ A visitor for a parsed Abstract Syntax Tree which finds executable
+ statements.
+ """
def __init__(self, statements, excluded, suite_spots):
compiler.visitor.ASTVisitor.__init__(self)
self.statements = statements
@@ -95,7 +108,6 @@ class StatementFindingAstVisitor(compiler.visitor.ASTVisitor):
self.excluding_suite = 0
def doRecursive(self, node):
- self.recordNodeLine(node)
for n in node.getChildNodes():
self.dispatch(n)
@@ -131,12 +143,35 @@ class StatementFindingAstVisitor(compiler.visitor.ASTVisitor):
def doStatement(self, node):
self.recordLine(self.getFirstLine(node))
- visitAssert = visitAssign = visitAssTuple = visitDiscard = visitPrint = \
+ visitAssert = visitAssign = visitAssTuple = visitPrint = \
visitPrintnl = visitRaise = visitSubscript = visitDecorators = \
doStatement
+ def visitPass(self, node):
+ # Pass statements have weird interactions with docstrings. If this
+ # pass statement is part of one of those pairs, claim that the statement
+ # is on the later of the two lines.
+ l = node.lineno
+ if l:
+ lines = self.suite_spots.get(l, [l,l])
+ self.statements[lines[1]] = 1
+
+ def visitDiscard(self, node):
+ # Discard nodes are statements that execute an expression, but then
+ # discard the results. This includes function calls, so we can't
+ # ignore them all. But if the expression is a constant, the statement
+ # won't be "executed", so don't count it now.
+ if node.expr.__class__.__name__ != 'Const':
+ self.doStatement(node)
+
def recordNodeLine(self, node):
- return self.recordLine(node.lineno)
+ # Stmt nodes often have None, but shouldn't claim the first line of
+ # their children (because the first child might be an ignorable line
+ # like "global a").
+ if node.__class__.__name__ != 'Stmt':
+ return self.recordLine(self.getFirstLine(node))
+ else:
+ return 0
def recordLine(self, lineno):
# Returns a bool, whether the line is included or excluded.
@@ -145,7 +180,7 @@ class StatementFindingAstVisitor(compiler.visitor.ASTVisitor):
# keyword.
if lineno in self.suite_spots:
lineno = self.suite_spots[lineno][0]
- # If we're inside an exluded suite, record that this line was
+ # If we're inside an excluded suite, record that this line was
# excluded.
if self.excluding_suite:
self.excluded[lineno] = 1
@@ -197,6 +232,8 @@ class StatementFindingAstVisitor(compiler.visitor.ASTVisitor):
self.doSuite(node, node.body)
self.doElse(node.body, node)
+ visitWhile = visitFor
+
def visitIf(self, node):
# The first test has to be handled separately from the rest.
# The first test is credited to the line with the "if", but the others
@@ -206,10 +243,6 @@ class StatementFindingAstVisitor(compiler.visitor.ASTVisitor):
self.doSuite(t, n)
self.doElse(node.tests[-1][1], node)
- def visitWhile(self, node):
- self.doSuite(node, node.body)
- self.doElse(node.body, node)
-
def visitTryExcept(self, node):
self.doSuite(node, node.body)
for i in range(len(node.handlers)):
@@ -268,11 +301,13 @@ class coverage:
raise CoverageException, "Only one coverage object allowed."
self.usecache = 1
self.cache = None
+ self.parallel_mode = False
self.exclude_re = ''
self.nesting = 0
self.cstack = []
self.xstack = []
- self.relative_dir = os.path.normcase(os.path.abspath(os.curdir)+os.path.sep)
+ self.relative_dir = os.path.normcase(os.path.abspath(os.curdir)+os.sep)
+ self.exclude('# *pragma[: ]*[nN][oO] *[cC][oO][vV][eE][rR]')
# t(f, x, y). This method is passed to sys.settrace as a trace function.
# See [van Rossum 2001-07-20b, 9.2] for an explanation of sys.settrace and
@@ -280,23 +315,24 @@ class coverage:
# See [van Rossum 2001-07-20a, 3.2] for a description of frame and code
# objects.
- def t(self, f, w, a): #pragma: no cover
+ def t(self, f, w, unused): #pragma: no cover
if w == 'line':
+ #print "Executing %s @ %d" % (f.f_code.co_filename, f.f_lineno)
self.c[(f.f_code.co_filename, f.f_lineno)] = 1
for c in self.cstack:
c[(f.f_code.co_filename, f.f_lineno)] = 1
return self.t
- def help(self, error=None):
+ def help(self, error=None): #pragma: no cover
if error:
print error
print
print __doc__
sys.exit(1)
- def command_line(self, argv, help=None):
+ def command_line(self, argv, help_fn=None):
import getopt
- help = help or self.help
+ help_fn = help_fn or self.help
settings = {}
optmap = {
'-a': 'annotate',
@@ -327,12 +363,12 @@ class coverage:
pass # Can't get here, because getopt won't return anything unknown.
if settings.get('help'):
- help()
+ help_fn()
for i in ['erase', 'execute']:
for j in ['annotate', 'report', 'collect']:
if settings.get(i) and settings.get(j):
- help("You can't specify the '%s' and '%s' "
+ help_fn("You can't specify the '%s' and '%s' "
"options at the same time." % (i, j))
args_needed = (settings.get('execute')
@@ -342,18 +378,18 @@ class coverage:
or settings.get('collect')
or args_needed)
if not action:
- help("You must specify at least one of -e, -x, -c, -r, or -a.")
+ help_fn("You must specify at least one of -e, -x, -c, -r, or -a.")
if not args_needed and args:
- help("Unexpected arguments: %s" % " ".join(args))
+ help_fn("Unexpected arguments: %s" % " ".join(args))
- self.get_ready(settings.get('parallel-mode'))
- self.exclude('#pragma[: ]+[nN][oO] [cC][oO][vV][eE][rR]')
+ self.parallel_mode = settings.get('parallel-mode')
+ self.get_ready()
if settings.get('erase'):
self.erase()
if settings.get('execute'):
if not args:
- help("Nothing to do.")
+ help_fn("Nothing to do.")
sys.argv = args
self.start()
import __main__
@@ -387,13 +423,13 @@ class coverage:
def get_ready(self, parallel_mode=False):
if self.usecache and not self.cache:
self.cache = os.environ.get(self.cache_env, self.cache_default)
- if parallel_mode:
+ if self.parallel_mode:
self.cache += "." + gethostname() + "." + str(os.getpid())
self.restore()
self.analysis_cache = {}
def start(self, parallel_mode=False):
- self.get_ready(parallel_mode)
+ self.get_ready()
if self.nesting == 0: #pragma: no cover
sys.settrace(self.t)
if hasattr(threading, 'settrace'):
@@ -408,12 +444,12 @@ class coverage:
threading.settrace(None)
def erase(self):
+ self.get_ready()
self.c = {}
self.analysis_cache = {}
self.cexecuted = {}
if self.cache and os.path.exists(self.cache):
os.remove(self.cache)
- self.exclude_re = ""
def exclude(self, re):
if self.exclude_re:
@@ -464,11 +500,11 @@ class coverage:
def collect(self):
cache_dir, local = os.path.split(self.cache)
- for file in os.listdir(cache_dir):
- if not file.startswith(local):
+ for f in os.listdir(cache_dir or '.'):
+ if not f.startswith(local):
continue
- full_path = os.path.join(cache_dir, file)
+ full_path = os.path.join(cache_dir, f)
cexecuted = self.restore_file(full_path)
self.merge_data(cexecuted)
@@ -508,6 +544,9 @@ class coverage:
def canonicalize_filenames(self):
for filename, lineno in self.c.keys():
+ if filename == '<string>':
+ # Can't do anything useful with exec'd strings, so skip them.
+ continue
f = self.canonical_filename(filename)
if not self.cexecuted.has_key(f):
self.cexecuted[f] = {}
@@ -520,17 +559,19 @@ class coverage:
if isinstance(morf, types.ModuleType):
if not hasattr(morf, '__file__'):
raise CoverageException, "Module has no __file__ attribute."
- file = morf.__file__
+ f = morf.__file__
else:
- file = morf
- return self.canonical_filename(file)
+ f = morf
+ return self.canonical_filename(f)
# analyze_morf(morf). Analyze the module or filename passed as
# the argument. If the source code can't be found, raise an error.
# Otherwise, return a tuple of (1) the canonical filename of the
# source code for the module, (2) a list of lines of statements
- # in the source code, and (3) a list of lines of excluded statements.
-
+ # in the source code, (3) a list of lines of excluded statements,
+ # and (4), a map of line numbers to multi-line line number ranges, for
+ # statements that cross lines.
+
def analyze_morf(self, morf):
if self.analysis_cache.has_key(morf):
return self.analysis_cache[morf]
@@ -544,16 +585,53 @@ class coverage:
elif ext != '.py':
raise CoverageException, "File '%s' not Python source." % filename
source = open(filename, 'r')
- lines, excluded_lines = self.find_executable_statements(
+ lines, excluded_lines, line_map = self.find_executable_statements(
source.read(), exclude=self.exclude_re
)
source.close()
- result = filename, lines, excluded_lines
+ result = filename, lines, excluded_lines, line_map
self.analysis_cache[morf] = result
return result
+ def first_line_of_tree(self, tree):
+ while True:
+ if len(tree) == 3 and type(tree[2]) == type(1):
+ return tree[2]
+ tree = tree[1]
+
+ def last_line_of_tree(self, tree):
+ while True:
+ if len(tree) == 3 and type(tree[2]) == type(1):
+ return tree[2]
+ tree = tree[-1]
+
+ def find_docstring_pass_pair(self, tree, spots):
+ for i in range(1, len(tree)):
+ if self.is_string_constant(tree[i]) and self.is_pass_stmt(tree[i+1]):
+ first_line = self.first_line_of_tree(tree[i])
+ last_line = self.last_line_of_tree(tree[i+1])
+ self.record_multiline(spots, first_line, last_line)
+
+ def is_string_constant(self, tree):
+ try:
+ return tree[0] == symbol.stmt and tree[1][1][1][0] == symbol.expr_stmt
+ except:
+ return False
+
+ def is_pass_stmt(self, tree):
+ try:
+ return tree[0] == symbol.stmt and tree[1][1][1][0] == symbol.pass_stmt
+ except:
+ return False
+
+ def record_multiline(self, spots, i, j):
+ for l in range(i, j+1):
+ spots[l] = (i, j)
+
def get_suite_spots(self, tree, spots):
- import symbol, token
+ """ Analyze a parse tree to find suite introducers which span a number
+ of lines.
+ """
for i in range(1, len(tree)):
if type(tree[i]) == type(()):
if tree[i][0] == symbol.suite:
@@ -561,7 +639,9 @@ class coverage:
lineno_colon = lineno_word = None
for j in range(i-1, 0, -1):
if tree[j][0] == token.COLON:
- lineno_colon = tree[j][2]
+ # Colons are never executed themselves: we want the
+ # line number of the last token before the colon.
+ lineno_colon = self.last_line_of_tree(tree[j-1])
elif tree[j][0] == token.NAME:
if tree[j][1] == 'elif':
# Find the line number of the first non-terminal
@@ -583,8 +663,18 @@ class coverage:
if lineno_colon and lineno_word:
# Found colon and keyword, mark all the lines
# between the two with the two line numbers.
- for l in range(lineno_word, lineno_colon+1):
- spots[l] = (lineno_word, lineno_colon)
+ self.record_multiline(spots, lineno_word, lineno_colon)
+
+ # "pass" statements are tricky: different versions of Python
+ # treat them differently, especially in the common case of a
+ # function with a doc string and a single pass statement.
+ self.find_docstring_pass_pair(tree[i], spots)
+
+ elif tree[i][0] == symbol.simple_stmt:
+ first_line = self.first_line_of_tree(tree[i])
+ last_line = self.last_line_of_tree(tree[i])
+ if first_line != last_line:
+ self.record_multiline(spots, first_line, last_line)
self.get_suite_spots(tree[i], spots)
def find_executable_statements(self, text, exclude=None):
@@ -598,10 +688,13 @@ class coverage:
if reExclude.search(lines[i]):
excluded[i+1] = 1
+ # Parse the code and analyze the parse tree to find out which statements
+ # are multiline, and where suites begin and end.
import parser
tree = parser.suite(text+'\n\n').totuple(1)
self.get_suite_spots(tree, suite_spots)
-
+ #print "Suite spots:", suite_spots
+
# Use the compiler module to parse the text and find the executable
# statements. We add newlines to be impervious to final partial lines.
statements = {}
@@ -613,7 +706,7 @@ class coverage:
lines.sort()
excluded_lines = excluded.keys()
excluded_lines.sort()
- return lines, excluded_lines
+ return lines, excluded_lines, suite_spots
# format_lines(statements, lines). Format a list of line numbers
# for printing by coalescing groups of lines as long as the lines
@@ -646,7 +739,8 @@ class coverage:
return "%d" % start
else:
return "%d-%d" % (start, end)
- return string.join(map(stringify, pairs), ", ")
+ ret = string.join(map(stringify, pairs), ", ")
+ return ret
# Backward compatibility with version 1.
def analysis(self, morf):
@@ -654,13 +748,17 @@ class coverage:
return f, s, m, mf
def analysis2(self, morf):
- filename, statements, excluded = self.analyze_morf(morf)
+ filename, statements, excluded, line_map = self.analyze_morf(morf)
self.canonicalize_filenames()
if not self.cexecuted.has_key(filename):
self.cexecuted[filename] = {}
missing = []
for line in statements:
- if not self.cexecuted[filename].has_key(line):
+ lines = line_map.get(line, [line, line])
+ for l in range(lines[0], lines[1]+1):
+ if self.cexecuted[filename].has_key(l):
+ break
+ else:
missing.append(line)
return (filename, statements, excluded, missing,
self.format_lines(statements, missing))
@@ -698,6 +796,15 @@ class coverage:
def report(self, morfs, show_missing=1, ignore_errors=0, file=None, omit_prefixes=[]):
if not isinstance(morfs, types.ListType):
morfs = [morfs]
+ # On windows, the shell doesn't expand wildcards. Do it here.
+ globbed = []
+ for morf in morfs:
+ if isinstance(morf, strclass):
+ globbed.extend(glob.glob(morf))
+ else:
+ globbed.append(morf)
+ morfs = globbed
+
morfs = self.filter_by_prefix(morfs, omit_prefixes)
morfs.sort(self.morf_name_compare)
@@ -735,8 +842,8 @@ class coverage:
raise
except:
if not ignore_errors:
- type, msg = sys.exc_info()[0:2]
- print >>file, fmt_err % (name, type, msg)
+ typ, msg = sys.exc_info()[0:2]
+ print >>file, fmt_err % (name, typ, msg)
if len(morfs) > 1:
print >>file, "-" * len(header)
if total_statements > 0:
@@ -816,18 +923,41 @@ class coverage:
the_coverage = coverage()
# Module functions call methods in the singleton object.
-def use_cache(*args, **kw): return the_coverage.use_cache(*args, **kw)
-def start(*args, **kw): return the_coverage.start(*args, **kw)
-def stop(*args, **kw): return the_coverage.stop(*args, **kw)
-def erase(*args, **kw): return the_coverage.erase(*args, **kw)
-def begin_recursive(*args, **kw): return the_coverage.begin_recursive(*args, **kw)
-def end_recursive(*args, **kw): return the_coverage.end_recursive(*args, **kw)
-def exclude(*args, **kw): return the_coverage.exclude(*args, **kw)
-def analysis(*args, **kw): return the_coverage.analysis(*args, **kw)
-def analysis2(*args, **kw): return the_coverage.analysis2(*args, **kw)
-def report(*args, **kw): return the_coverage.report(*args, **kw)
-def annotate(*args, **kw): return the_coverage.annotate(*args, **kw)
-def annotate_file(*args, **kw): return the_coverage.annotate_file(*args, **kw)
+def use_cache(*args, **kw):
+ return the_coverage.use_cache(*args, **kw)
+
+def start(*args, **kw):
+ return the_coverage.start(*args, **kw)
+
+def stop(*args, **kw):
+ return the_coverage.stop(*args, **kw)
+
+def erase(*args, **kw):
+ return the_coverage.erase(*args, **kw)
+
+def begin_recursive(*args, **kw):
+ return the_coverage.begin_recursive(*args, **kw)
+
+def end_recursive(*args, **kw):
+ return the_coverage.end_recursive(*args, **kw)
+
+def exclude(*args, **kw):
+ return the_coverage.exclude(*args, **kw)
+
+def analysis(*args, **kw):
+ return the_coverage.analysis(*args, **kw)
+
+def analysis2(*args, **kw):
+ return the_coverage.analysis2(*args, **kw)
+
+def report(*args, **kw):
+ return the_coverage.report(*args, **kw)
+
+def annotate(*args, **kw):
+ return the_coverage.annotate(*args, **kw)
+
+def annotate_file(*args, **kw):
+ return the_coverage.annotate_file(*args, **kw)
# Save coverage data when Python exits. (The atexit module wasn't
# introduced until Python 2.0, so use sys.exitfunc when it's not
@@ -918,11 +1048,32 @@ if __name__ == '__main__':
#
# 2006-08-23 NMB Refactorings to improve testability. Fixes to command-line
# logic for parallel mode and collect.
+#
+# 2006-08-25 NMB "#pragma: nocover" is excluded by default.
+#
+# 2006-09-10 NMB Properly ignore docstrings and other constant expressions that
+# appear in the middle of a function, a problem reported by Tim Leslie.
+# Minor changes to avoid lint warnings.
+#
+# 2006-09-17 NMB coverage.erase() shouldn't clobber the exclude regex.
+# Change how parallel mode is invoked, and fix erase() so that it erases the
+# cache when called programmatically.
+#
+# 2007-07-21 NMB In reports, ignore code executed from strings, since we can't
+# do anything useful with it anyway.
+# Better file handling on Linux, thanks Guillaume Chazarain.
+# Better shell support on Windows, thanks Noel O'Boyle.
+# Python 2.2 support maintained, thanks Catherine Proulx.
+#
+# 2007-07-22 NMB Python 2.5 now fully supported. The method of dealing with
+# multi-line statements is now less sensitive to the exact line that Python
+# reports during execution. Pass statements are handled specially so that their
+# disappearance during execution won't throw off the measurement.
# C. COPYRIGHT AND LICENCE
#
# Copyright 2001 Gareth Rees. All rights reserved.
-# Copyright 2004-2006 Ned Batchelder. All rights reserved.
+# Copyright 2004-2007 Ned Batchelder. All rights reserved.
#
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are
@@ -949,4 +1100,4 @@ if __name__ == '__main__':
# USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH
# DAMAGE.
#
-# $Id: coverage.py 47 2006-08-24 01:08:48Z Ned $
+# $Id: coverage.py 67 2007-07-21 19:51:07Z nedbat $
diff --git a/test/testlib/profiling.py b/test/testlib/profiling.py
new file mode 100644
index 000000000..697df4ea2
--- /dev/null
+++ b/test/testlib/profiling.py
@@ -0,0 +1,74 @@
+"""Profiling support for unit and performance tests."""
+
+from testlib.config import parser, post_configure
+import testlib.config
+
+__all__ = 'profiled',
+
+all_targets = set()
+profile_config = { 'targets': set(),
+ 'report': True,
+ 'sort': ('time', 'calls'),
+ 'limit': None }
+
+def profiled(target, **target_opts):
+ """Optional function profiling.
+
+ @profiled('label')
+ or
+ @profiled('label', report=True, sort=('calls',), limit=20)
+
+ Enables profiling for a function when 'label' is targetted for
+ profiling. Report options can be supplied, and override the global
+ configuration and command-line options.
+ """
+
+ import time, hotshot, hotshot.stats
+
+ # manual or automatic namespacing by module would remove conflict issues
+ if target in all_targets:
+ print "Warning: redefining profile target '%s'" % target
+ all_targets.add(target)
+
+ filename = "%s.prof" % target
+
+ def decorator(fn):
+ def profiled(*args, **kw):
+ if (target not in profile_config['targets'] and
+ not target_opts.get('always', None)):
+ return fn(*args, **kw)
+
+ prof = hotshot.Profile(filename)
+ began = time.time()
+ prof.start()
+ try:
+ result = fn(*args, **kw)
+ finally:
+ prof.stop()
+ ended = time.time()
+ prof.close()
+
+ if not testlib.config.options.quiet:
+ print "Profiled target '%s', wall time: %.2f seconds" % (
+ target, ended - began)
+
+ report = target_opts.get('report', profile_config['report'])
+ if report:
+ sort_ = target_opts.get('sort', profile_config['sort'])
+ limit = target_opts.get('limit', profile_config['limit'])
+ print "Profile report for target '%s' (%s)" % (
+ target, filename)
+
+ stats = hotshot.stats.load(filename)
+ stats.sort_stats(*sort_)
+ if limit:
+ stats.print_stats(limit)
+ else:
+ stats.print_stats()
+ return result
+ try:
+ profiled.__name__ = fn.__name__
+ except:
+ pass
+ return profiled
+ return decorator
diff --git a/test/testlib/schema.py b/test/testlib/schema.py
new file mode 100644
index 000000000..a2fc91265
--- /dev/null
+++ b/test/testlib/schema.py
@@ -0,0 +1,28 @@
+import testbase
+from sqlalchemy import schema
+
+__all__ = 'Table', 'Column',
+
+table_options = {}
+
+def Table(*args, **kw):
+ """A schema.Table wrapper/hook for dialect-specific tweaks."""
+
+ test_opts = dict([(k,kw.pop(k)) for k in kw.keys()
+ if k.startswith('test_')])
+
+ kw.update(table_options)
+
+ if testbase.db.name == 'mysql':
+ if 'mysql_engine' not in kw and 'mysql_type' not in kw:
+ if 'test_needs_fk' in test_opts or 'test_needs_acid' in test_opts:
+ kw['mysql_engine'] = 'InnoDB'
+
+ return schema.Table(*args, **kw)
+
+def Column(*args, **kw):
+ """A schema.Column wrapper/hook for dialect-specific tweaks."""
+
+ # TODO: a Column that creates a Sequence automatically for PK columns,
+ # which would help Oracle tests
+ return schema.Column(*args, **kw)
diff --git a/test/tables.py b/test/testlib/tables.py
index 8e337999c..69c84c5b3 100644
--- a/test/tables.py
+++ b/test/testlib/tables.py
@@ -1,24 +1,24 @@
-
-from sqlalchemy import *
-import os
import testbase
+from sqlalchemy import *
+from testlib.schema import Table, Column
-ECHO = testbase.echo
-db = testbase.db
-metadata = MetaData(db)
+# these are older test fixtures, used primarily by test/orm/mapper.py and test/orm/unitofwork.py.
+# newer unit tests make usage of test/orm/fixtures.py.
+
+metadata = MetaData()
users = Table('users', metadata,
Column('user_id', Integer, Sequence('user_id_seq', optional=True), primary_key = True),
Column('user_name', String(40)),
- mysql_engine='innodb'
+ test_needs_acid=True,
+ test_needs_fk=True,
)
addresses = Table('email_addresses', metadata,
Column('address_id', Integer, Sequence('address_id_seq', optional=True), primary_key = True),
Column('user_id', Integer, ForeignKey(users.c.user_id)),
Column('email_address', String(40)),
-
)
orders = Table('orders', metadata,
@@ -26,20 +26,17 @@ orders = Table('orders', metadata,
Column('user_id', Integer, ForeignKey(users.c.user_id)),
Column('description', String(50)),
Column('isopen', Integer),
-
)
orderitems = Table('items', metadata,
Column('item_id', INT, Sequence('items_id_seq', optional=True), primary_key = True),
Column('order_id', INT, ForeignKey("orders")),
Column('item_name', VARCHAR(50)),
-
)
keywords = Table('keywords', metadata,
Column('keyword_id', Integer, Sequence('keyword_id_seq', optional=True), primary_key = True),
Column('name', VARCHAR(50)),
-
)
userkeywords = Table('userkeywords', metadata,
@@ -54,13 +51,19 @@ itemkeywords = Table('itemkeywords', metadata,
)
def create():
+ if not metadata.bind:
+ metadata.bind = testbase.db
metadata.create_all()
def drop():
+ if not metadata.bind:
+ metadata.bind = testbase.db
metadata.drop_all()
def delete():
for t in metadata.table_iterator(reverse=True):
t.delete().execute()
def user_data():
+ if not metadata.bind:
+ metadata.bind = testbase.db
users.insert().execute(
dict(user_id = 7, user_name = 'jack'),
dict(user_id = 8, user_name = 'ed'),
@@ -212,4 +215,4 @@ order_result = [
{'order_id' : 4, 'items':(Item, [])},
{'order_id' : 5, 'items':(Item, [])},
]
-#db.echo = True
+
diff --git a/test/testlib/testing.py b/test/testlib/testing.py
new file mode 100644
index 000000000..213772e9e
--- /dev/null
+++ b/test/testlib/testing.py
@@ -0,0 +1,363 @@
+"""TestCase and TestSuite artifacts and testing decorators."""
+
+# monkeypatches unittest.TestLoader.suiteClass at import time
+
+import unittest, re, sys, os
+from cStringIO import StringIO
+from sqlalchemy import MetaData, sql
+from sqlalchemy.orm import clear_mappers
+import testlib.config as config
+
+__all__ = 'PersistTest', 'AssertMixin', 'ORMTest'
+
+def unsupported(*dbs):
+ """Mark a test as unsupported by one or more database implementations"""
+
+ def decorate(fn):
+ fn_name = fn.__name__
+ def maybe(*args, **kw):
+ if config.db.name in dbs:
+ print "'%s' unsupported on DB implementation '%s'" % (
+ fn_name, config.db.name)
+ return True
+ else:
+ return fn(*args, **kw)
+ try:
+ maybe.__name__ = fn_name
+ except:
+ pass
+ return maybe
+ return decorate
+
+def supported(*dbs):
+ """Mark a test as supported by one or more database implementations"""
+
+ def decorate(fn):
+ fn_name = fn.__name__
+ def maybe(*args, **kw):
+ if config.db.name in dbs:
+ return fn(*args, **kw)
+ else:
+ print "'%s' unsupported on DB implementation '%s'" % (
+ fn_name, config.db.name)
+ return True
+ try:
+ maybe.__name__ = fn_name
+ except:
+ pass
+ return maybe
+ return decorate
+
+class TestData(object):
+ """Tracks SQL expressions as they are executed via an instrumented ExecutionContext."""
+
+ def __init__(self):
+ self.set_assert_list(None, None)
+ self.sql_count = 0
+ self.buffer = None
+
+ def set_assert_list(self, unittest, list):
+ self.unittest = unittest
+ self.assert_list = list
+ if list is not None:
+ self.assert_list.reverse()
+
+testdata = TestData()
+
+
+class ExecutionContextWrapper(object):
+ """instruments the ExecutionContext created by the Engine so that SQL expressions
+ can be tracked."""
+
+ def __init__(self, ctx):
+ self.__dict__['ctx'] = ctx
+ def __getattr__(self, key):
+ return getattr(self.ctx, key)
+ def __setattr__(self, key, value):
+ setattr(self.ctx, key, value)
+
+ def post_execution(self):
+ ctx = self.ctx
+ statement = unicode(ctx.compiled)
+ statement = re.sub(r'\n', '', ctx.statement)
+ if testdata.buffer is not None:
+ testdata.buffer.write(statement + "\n")
+
+ if testdata.assert_list is not None:
+ assert len(testdata.assert_list), "Received query but no more assertions: %s" % statement
+ item = testdata.assert_list[-1]
+ if not isinstance(item, dict):
+ item = testdata.assert_list.pop()
+ else:
+ # asserting a dictionary of statements->parameters
+ # this is to specify query assertions where the queries can be in
+ # multiple orderings
+ if not item.has_key('_converted'):
+ for key in item.keys():
+ ckey = self.convert_statement(key)
+ item[ckey] = item[key]
+ if ckey != key:
+ del item[key]
+ item['_converted'] = True
+ try:
+ entry = item.pop(statement)
+ if len(item) == 1:
+ testdata.assert_list.pop()
+ item = (statement, entry)
+ except KeyError:
+ assert False, "Testing for one of the following queries: %s, received '%s'" % (repr([k for k in item.keys()]), statement)
+
+ (query, params) = item
+ if callable(params):
+ params = params(ctx)
+ if params is not None and isinstance(params, list) and len(params) == 1:
+ params = params[0]
+
+ if isinstance(ctx.compiled_parameters, sql.ClauseParameters):
+ parameters = ctx.compiled_parameters.get_original_dict()
+ elif isinstance(ctx.compiled_parameters, list):
+ parameters = [p.get_original_dict() for p in ctx.compiled_parameters]
+
+ query = self.convert_statement(query)
+ if config.db.name == 'mssql' and statement.endswith('; select scope_identity()'):
+ statement = statement[:-25]
+ testdata.unittest.assert_(statement == query and (params is None or params == parameters), "Testing for query '%s' params %s, received '%s' with params %s" % (query, repr(params), statement, repr(parameters)))
+ testdata.sql_count += 1
+ self.ctx.post_execution()
+
+ def convert_statement(self, query):
+ paramstyle = self.ctx.dialect.paramstyle
+ if paramstyle == 'named':
+ pass
+ elif paramstyle =='pyformat':
+ query = re.sub(r':([\w_]+)', r"%(\1)s", query)
+ else:
+ # positional params
+ repl = None
+ if paramstyle=='qmark':
+ repl = "?"
+ elif paramstyle=='format':
+ repl = r"%s"
+ elif paramstyle=='numeric':
+ repl = None
+ query = re.sub(r':([\w_]+)', repl, query)
+ return query
+
+class PersistTest(unittest.TestCase):
+
+ def __init__(self, *args, **params):
+ unittest.TestCase.__init__(self, *args, **params)
+
+ def setUpAll(self):
+ pass
+
+ def tearDownAll(self):
+ pass
+
+ def shortDescription(self):
+ """overridden to not return docstrings"""
+ return None
+
+class AssertMixin(PersistTest):
+ """given a list-based structure of keys/properties which represent information within an object structure, and
+ a list of actual objects, asserts that the list of objects corresponds to the structure."""
+
+ def assert_result(self, result, class_, *objects):
+ result = list(result)
+ print repr(result)
+ self.assert_list(result, class_, objects)
+
+ def assert_list(self, result, class_, list):
+ self.assert_(len(result) == len(list),
+ "result list is not the same size as test list, " +
+ "for class " + class_.__name__)
+ for i in range(0, len(list)):
+ self.assert_row(class_, result[i], list[i])
+
+ def assert_row(self, class_, rowobj, desc):
+ self.assert_(rowobj.__class__ is class_,
+ "item class is not " + repr(class_))
+ for key, value in desc.iteritems():
+ if isinstance(value, tuple):
+ if isinstance(value[1], list):
+ self.assert_list(getattr(rowobj, key), value[0], value[1])
+ else:
+ self.assert_row(value[0], getattr(rowobj, key), value[1])
+ else:
+ self.assert_(getattr(rowobj, key) == value,
+ "attribute %s value %s does not match %s" % (
+ key, getattr(rowobj, key), value))
+
+ def assert_sql(self, db, callable_, list, with_sequences=None):
+ global testdata
+ testdata = TestData()
+ if with_sequences is not None and (config.db.name == 'postgres' or
+ config.db.name == 'oracle'):
+ testdata.set_assert_list(self, with_sequences)
+ else:
+ testdata.set_assert_list(self, list)
+ try:
+ callable_()
+ finally:
+ testdata.set_assert_list(None, None)
+
+ def assert_sql_count(self, db, callable_, count):
+ global testdata
+ testdata = TestData()
+ try:
+ callable_()
+ finally:
+ self.assert_(testdata.sql_count == count,
+ "desired statement count %d does not match %d" % (
+ count, testdata.sql_count))
+
+ def capture_sql(self, db, callable_):
+ global testdata
+ testdata = TestData()
+ buffer = StringIO()
+ testdata.buffer = buffer
+ try:
+ callable_()
+ return buffer.getvalue()
+ finally:
+ testdata.buffer = None
+
+_otest_metadata = None
+class ORMTest(AssertMixin):
+ keep_mappers = False
+ keep_data = False
+
+ def setUpAll(self):
+ global _otest_metadata
+ _otest_metadata = MetaData(config.db)
+ self.define_tables(_otest_metadata)
+ _otest_metadata.create_all()
+ self.insert_data()
+
+ def define_tables(self, _otest_metadata):
+ raise NotImplementedError()
+
+ def insert_data(self):
+ pass
+
+ def get_metadata(self):
+ return _otest_metadata
+
+ def tearDownAll(self):
+ clear_mappers()
+ _otest_metadata.drop_all()
+
+ def tearDown(self):
+ if not self.keep_mappers:
+ clear_mappers()
+ if not self.keep_data:
+ for t in _otest_metadata.table_iterator(reverse=True):
+ t.delete().execute().close()
+
+
+class TTestSuite(unittest.TestSuite):
+ """A TestSuite with once per TestCase setUpAll() and tearDownAll()"""
+
+ def __init__(self, tests=()):
+ if len(tests) >0 and isinstance(tests[0], PersistTest):
+ self._initTest = tests[0]
+ else:
+ self._initTest = None
+ unittest.TestSuite.__init__(self, tests)
+
+ def do_run(self, result):
+ # nice job unittest ! you switched __call__ and run() between py2.3
+ # and 2.4 thereby making straight subclassing impossible !
+ for test in self._tests:
+ if result.shouldStop:
+ break
+ test(result)
+ return result
+
+ def run(self, result):
+ return self(result)
+
+ def __call__(self, result):
+ try:
+ if self._initTest is not None:
+ self._initTest.setUpAll()
+ except:
+ result.addError(self._initTest, self.__exc_info())
+ pass
+ try:
+ return self.do_run(result)
+ finally:
+ try:
+ if self._initTest is not None:
+ self._initTest.tearDownAll()
+ except:
+ result.addError(self._initTest, self.__exc_info())
+ pass
+
+ def __exc_info(self):
+ """Return a version of sys.exc_info() with the traceback frame
+ minimised; usually the top level of the traceback frame is not
+ needed.
+ ripped off out of unittest module since its double __
+ """
+ exctype, excvalue, tb = sys.exc_info()
+ if sys.platform[:4] == 'java': ## tracebacks look different in Jython
+ return (exctype, excvalue, tb)
+ return (exctype, excvalue, tb)
+
+unittest.TestLoader.suiteClass = TTestSuite
+
+def _iter_covered_files():
+ import sqlalchemy
+ for rec in os.walk(os.path.dirname(sqlalchemy.__file__)):
+ for x in rec[2]:
+ if x.endswith('.py'):
+ yield os.path.join(rec[0], x)
+
+def cover(callable_, file_=None):
+ from testlib import coverage
+ coverage_client = coverage.the_coverage
+ coverage_client.get_ready()
+ coverage_client.exclude('#pragma[: ]+[nN][oO] [cC][oO][vV][eE][rR]')
+ coverage_client.erase()
+ coverage_client.start()
+ try:
+ return callable_()
+ finally:
+ coverage_client.stop()
+ coverage_client.save()
+ coverage_client.report(list(_iter_covered_files()),
+ show_missing=False, ignore_errors=False,
+ file=file_)
+
+class DevNullWriter(object):
+ def write(self, msg):
+ pass
+ def flush(self):
+ pass
+
+def runTests(suite):
+ verbose = config.options.verbose
+ quiet = config.options.quiet
+ orig_stdout = sys.stdout
+
+ try:
+ if not verbose or quiet:
+ sys.stdout = DevNullWriter()
+ runner = unittest.TextTestRunner(verbosity = quiet and 1 or 2)
+ return runner.run(suite)
+ finally:
+ if not verbose or quiet:
+ sys.stdout = orig_stdout
+
+def main(suite=None):
+ if not suite:
+ if len(sys.argv[1:]):
+ suite =unittest.TestLoader().loadTestsFromNames(
+ sys.argv[1:], __import__('__main__'))
+ else:
+ suite = unittest.TestLoader().loadTestsFromModule(
+ __import__('__main__'))
+
+ result = runTests(suite)
+ sys.exit(not result.wasSuccessful())
diff --git a/test/zblog/mappers.py b/test/zblog/mappers.py
index 244a53d0e..11eaf4fd0 100644
--- a/test/zblog/mappers.py
+++ b/test/zblog/mappers.py
@@ -4,6 +4,7 @@ import zblog.tables as tables
import zblog.user as user
from zblog.blog import *
from sqlalchemy import *
+from sqlalchemy.orm import *
import sqlalchemy.util as util
def zblog_mappers():
diff --git a/test/zblog/tables.py b/test/zblog/tables.py
index f01f18921..5b4054a19 100644
--- a/test/zblog/tables.py
+++ b/test/zblog/tables.py
@@ -1,13 +1,16 @@
+"""application table metadata objects are described here."""
+
from sqlalchemy import *
+from testlib import *
+
metadata = MetaData()
-"""application table metadata objects are described here."""
users = Table('users', metadata,
Column('user_id', Integer, primary_key=True),
Column('user_name', String(30), nullable=False),
Column('fullname', String(100), nullable=False),
- Column('password', String(30), nullable=False),
+ Column('password', String(40), nullable=False),
Column('groupname', String(20), nullable=False),
)
diff --git a/test/zblog/tests.py b/test/zblog/tests.py
index e538cff9d..ad6876937 100644
--- a/test/zblog/tests.py
+++ b/test/zblog/tests.py
@@ -1,20 +1,20 @@
-from testbase import AssertMixin
import testbase
-import unittest
-db = testbase.db
from sqlalchemy import *
-
+from sqlalchemy.orm import *
+from testlib import *
from zblog import mappers, tables
from zblog.user import *
from zblog.blog import *
+
class ZBlogTest(AssertMixin):
def create_tables(self):
- tables.metadata.create_all(connectable=db)
+ tables.metadata.drop_all(bind=testbase.db)
+ tables.metadata.create_all(bind=testbase.db)
def drop_tables(self):
- tables.metadata.drop_all(connectable=db)
+ tables.metadata.drop_all(bind=testbase.db)
def setUpAll(self):
self.create_tables()
@@ -31,7 +31,7 @@ class SavePostTest(ZBlogTest):
super(SavePostTest, self).setUpAll()
mappers.zblog_mappers()
global blog_id, user_id
- s = create_session(bind_to=db)
+ s = create_session(bind=testbase.db)
user = User('zbloguser', "Zblog User", "hello", group=administrator)
blog = Blog(owner=user)
blog.name = "this is a blog"
@@ -50,9 +50,9 @@ class SavePostTest(ZBlogTest):
"""test that a transient/pending instance has proper bi-directional behavior.
this requires that lazy loaders do not fire off for a transient/pending instance."""
- s = create_session(bind_to=db)
+ s = create_session(bind=testbase.db)
- trans = s.create_transaction()
+ s.begin()
try:
blog = s.query(Blog).get(blog_id)
post = Post(headline="asdf asdf", summary="asdfasfd")
@@ -61,14 +61,14 @@ class SavePostTest(ZBlogTest):
post.blog = blog
assert post in blog.posts
finally:
- trans.rollback()
+ s.rollback()
def testoptimisticorphans(self):
"""test that instances in the session with un-loaded parents will not
get marked as "orphans" and then deleted """
- s = create_session(bind_to=db)
+ s = create_session(bind=testbase.db)
- trans = s.create_transaction()
+ s.begin()
try:
blog = s.query(Blog).get(blog_id)
post = Post(headline="asdf asdf", summary="asdfasfd")
@@ -90,10 +90,10 @@ class SavePostTest(ZBlogTest):
assert s.query(Post).get(post.id) is not None
finally:
- trans.rollback()
+ s.rollback()
if __name__ == "__main__":
testbase.main()
- \ No newline at end of file
+
diff --git a/test/zblog/user.py b/test/zblog/user.py
index 1dca0328e..3e77fa842 100644
--- a/test/zblog/user.py
+++ b/test/zblog/user.py
@@ -1,13 +1,7 @@
"""user.py - handles user login and validation"""
import random, string
-try:
- from crypt import crypt
-except:
- try:
- from fcrypt import crypt
- except:
- raise "Need fcrypt module on non-Unix platform: http://home.clear.net.nz/pages/c.evans/sw/"
+from sha import sha
administrator = 'admin'
user = 'user'
@@ -16,7 +10,7 @@ groups = [user, administrator]
def cryptpw(password, salt=None):
if salt is None:
salt = string.join([chr(random.randint(ord('a'), ord('z'))), chr(random.randint(ord('a'), ord('z')))],'')
- return crypt(password, salt)
+ return sha(password + salt).hexdigest()
def checkpw(password, dbpw):
return cryptpw(password, dbpw[:2]) == dbpw