summaryrefslogtreecommitdiff
path: root/test
diff options
context:
space:
mode:
authormike bayer <mike_mp@zzzcomputing.com>2019-02-13 01:35:15 +0000
committerGerrit Code Review <gerrit@bbpush.zzzcomputing.com>2019-02-13 01:35:15 +0000
commit5cc5234e0b065907deb2765249dde1c526fe5c89 (patch)
tree43a1dff1a674080707b4ab6c4b0f91bc2b8b56f8 /test
parent104625941d5589ddb93da9c66217da6308c53a15 (diff)
parentc1f310df44033d943413170de878ce95fafa387e (diff)
downloadsqlalchemy-5cc5234e0b065907deb2765249dde1c526fe5c89.tar.gz
Merge "Allow SQL expression for ORM primary keys"
Diffstat (limited to 'test')
-rw-r--r--test/orm/test_unitofwork.py32
-rw-r--r--test/requirements.py4
-rw-r--r--test/sql/test_insert.py91
3 files changed, 127 insertions, 0 deletions
diff --git a/test/orm/test_unitofwork.py b/test/orm/test_unitofwork.py
index 6326f5f1a..dc8a818d2 100644
--- a/test/orm/test_unitofwork.py
+++ b/test/orm/test_unitofwork.py
@@ -10,6 +10,7 @@ from sqlalchemy import event
from sqlalchemy import ForeignKey
from sqlalchemy import func
from sqlalchemy import Integer
+from sqlalchemy import literal
from sqlalchemy import literal_column
from sqlalchemy import select
from sqlalchemy import String
@@ -499,6 +500,19 @@ class ClauseAttributesTest(fixtures.MappedTest):
Column("value", Boolean),
)
+ Table(
+ "pk_t",
+ metadata,
+ Column(
+ "p_id",
+ Integer,
+ key="id",
+ autoincrement=True,
+ primary_key=True,
+ ),
+ Column("data", String(30)),
+ )
+
@classmethod
def setup_classes(cls):
class User(cls.Comparable):
@@ -507,12 +521,17 @@ class ClauseAttributesTest(fixtures.MappedTest):
class HasBoolean(cls.Comparable):
pass
+ class PkDefault(cls.Comparable):
+ pass
+
@classmethod
def setup_mappers(cls):
User, users_t = cls.classes.User, cls.tables.users_t
HasBoolean, boolean_t = cls.classes.HasBoolean, cls.tables.boolean_t
+ PkDefault, pk_t = cls.classes.PkDefault, cls.tables.pk_t
mapper(User, users_t)
mapper(HasBoolean, boolean_t)
+ mapper(PkDefault, pk_t)
def test_update(self):
User = self.classes.User
@@ -568,6 +587,19 @@ class ClauseAttributesTest(fixtures.MappedTest):
assert (u.counter == 5) is True
+ @testing.requires.sql_expressions_inserted_as_primary_key
+ def test_insert_pk_expression(self):
+ PkDefault = self.classes.PkDefault
+
+ pk = PkDefault(id=literal(5) + 10, data="some data")
+ session = Session()
+ session.add(pk)
+ session.flush()
+
+ eq_(pk.id, 15)
+ session.commit()
+ eq_(pk.id, 15)
+
def test_update_special_comparator(self):
HasBoolean = self.classes.HasBoolean
diff --git a/test/requirements.py b/test/requirements.py
index c7c7206fc..ab3f01d04 100644
--- a/test/requirements.py
+++ b/test/requirements.py
@@ -350,6 +350,10 @@ class DefaultRequirements(SuiteRequirements):
)
@property
+ def sql_expressions_inserted_as_primary_key(self):
+ return only_if([self.returning, self.sqlite])
+
+ @property
def correlated_outer_joins(self):
"""Target must support an outer join to a subquery which
correlates to the parent."""
diff --git a/test/sql/test_insert.py b/test/sql/test_insert.py
index 066e54e93..cf6715a0e 100644
--- a/test/sql/test_insert.py
+++ b/test/sql/test_insert.py
@@ -16,6 +16,7 @@ from sqlalchemy import table
from sqlalchemy import text
from sqlalchemy.dialects import mysql
from sqlalchemy.dialects import postgresql
+from sqlalchemy.dialects import sqlite
from sqlalchemy.engine import default
from sqlalchemy.sql import crud
from sqlalchemy.testing import assert_raises
@@ -1196,6 +1197,96 @@ class MultirowTest(_InsertTestBase, fixtures.TablesTest, AssertsCompiledSQL):
dialect=postgresql.dialect(),
)
+ def test_sql_expression_pk_autoinc_lastinserted(self):
+ # test that postfetch isn't invoked for a SQL expression
+ # in a primary key column. the DB either needs to support a lastrowid
+ # that can return it, or RETURNING. [ticket:3133]
+ metadata = MetaData()
+ table = Table(
+ "sometable",
+ metadata,
+ Column("id", Integer, primary_key=True),
+ Column("data", String),
+ )
+
+ stmt = table.insert().return_defaults().values(id=func.foobar())
+ compiled = stmt.compile(dialect=sqlite.dialect(), column_keys=["data"])
+ eq_(compiled.postfetch, [])
+ eq_(compiled.returning, [])
+
+ self.assert_compile(
+ stmt,
+ "INSERT INTO sometable (id, data) VALUES " "(foobar(), ?)",
+ checkparams={"data": "foo"},
+ params={"data": "foo"},
+ dialect=sqlite.dialect(),
+ )
+
+ def test_sql_expression_pk_autoinc_returning(self):
+ # test that return_defaults() works with a primary key where we are
+ # sending a SQL expression, and we need to get the server-calculated
+ # value back. [ticket:3133]
+ metadata = MetaData()
+ table = Table(
+ "sometable",
+ metadata,
+ Column("id", Integer, primary_key=True),
+ Column("data", String),
+ )
+
+ stmt = table.insert().return_defaults().values(id=func.foobar())
+ returning_dialect = postgresql.dialect()
+ returning_dialect.implicit_returning = True
+ compiled = stmt.compile(
+ dialect=returning_dialect, column_keys=["data"]
+ )
+ eq_(compiled.postfetch, [])
+ eq_(compiled.returning, [table.c.id])
+
+ self.assert_compile(
+ stmt,
+ "INSERT INTO sometable (id, data) VALUES "
+ "(foobar(), %(data)s) RETURNING sometable.id",
+ checkparams={"data": "foo"},
+ params={"data": "foo"},
+ dialect=returning_dialect,
+ )
+
+ def test_sql_expression_pk_noautoinc_returning(self):
+ # test that return_defaults() works with a primary key where we are
+ # sending a SQL expression, and we need to get the server-calculated
+ # value back. [ticket:3133]
+ metadata = MetaData()
+ table = Table(
+ "sometable",
+ metadata,
+ Column(
+ "id",
+ Integer,
+ autoincrement=False,
+ primary_key=True,
+ ),
+ Column("data", String),
+ )
+
+ stmt = table.insert().return_defaults().values(id=func.foobar())
+ returning_dialect = postgresql.dialect()
+ returning_dialect.implicit_returning = True
+ compiled = stmt.compile(
+ dialect=returning_dialect, column_keys=["data"]
+ )
+ eq_(compiled.postfetch, [])
+ eq_(compiled.returning, [table.c.id])
+
+ self.assert_compile(
+ stmt,
+ "INSERT INTO sometable (id, data) VALUES "
+ "(foobar(), %(data)s) RETURNING sometable.id",
+ checkparams={"data": "foo"},
+ params={"data": "foo"},
+ dialect=returning_dialect,
+ )
+
def test_python_fn_default(self):
metadata = MetaData()
table = Table(