Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_insert.py631 linesDownload Raw Back to suite
1# testing/suite/test_insert.py
2# Copyright (C) 2005-2024 the SQLAlchemy authors and contributors
3# <see AUTHORS file>
4#
5# This module is part of SQLAlchemy and is released under
6# the MIT License: https://www.opensource.org/licenses/mit-license.php
7# mypy: ignore-errors
8
9from decimal import Decimal
10import uuid
11
12from . import testing
13from .. import fixtures
14from ..assertions import eq_
15from ..config import requirements
16from ..schema import Column
17from ..schema import Table
18from ... import Double
19from ... import Float
20from ... import Identity
21from ... import Integer
22from ... import literal
23from ... import literal_column
24from ... import Numeric
25from ... import select
26from ... import String
27from ...types import LargeBinary
28from ...types import UUID
29from ...types import Uuid
30
31
32class LastrowidTest(fixtures.TablesTest):
33    run_deletes = "each"
34
35    __backend__ = True
36
37    __requires__ = "implements_get_lastrowid", "autoincrement_insert"
38
39    @classmethod
40    def define_tables(cls, metadata):
41        Table(
42            "autoinc_pk",
43            metadata,
44            Column(
45                "id", Integer, primary_key=True, test_needs_autoincrement=True
46            ),
47            Column("data", String(50)),
48            implicit_returning=False,
49        )
50
51        Table(
52            "manual_pk",
53            metadata,
54            Column("id", Integer, primary_key=True, autoincrement=False),
55            Column("data", String(50)),
56            implicit_returning=False,
57        )
58
59    def _assert_round_trip(self, table, conn):
60        row = conn.execute(table.select()).first()
61        eq_(
62            row,
63            (
64                conn.dialect.default_sequence_base,
65                "some data",
66            ),
67        )
68
69    def test_autoincrement_on_insert(self, connection):
70        connection.execute(
71            self.tables.autoinc_pk.insert(), dict(data="some data")
72        )
73        self._assert_round_trip(self.tables.autoinc_pk, connection)
74
75    def test_last_inserted_id(self, connection):
76        r = connection.execute(
77            self.tables.autoinc_pk.insert(), dict(data="some data")
78        )
79        pk = connection.scalar(select(self.tables.autoinc_pk.c.id))
80        eq_(r.inserted_primary_key, (pk,))
81
82    @requirements.dbapi_lastrowid
83    def test_native_lastrowid_autoinc(self, connection):
84        r = connection.execute(
85            self.tables.autoinc_pk.insert(), dict(data="some data")
86        )
87        lastrowid = r.lastrowid
88        pk = connection.scalar(select(self.tables.autoinc_pk.c.id))
89        eq_(lastrowid, pk)
90
91
92class InsertBehaviorTest(fixtures.TablesTest):
93    run_deletes = "each"
94    __backend__ = True
95
96    @classmethod
97    def define_tables(cls, metadata):
98        Table(
99            "autoinc_pk",
100            metadata,
101            Column(
102                "id", Integer, primary_key=True, test_needs_autoincrement=True
103            ),
104            Column("data", String(50)),
105        )
106        Table(
107            "manual_pk",
108            metadata,
109            Column("id", Integer, primary_key=True, autoincrement=False),
110            Column("data", String(50)),
111        )
112        Table(
113            "no_implicit_returning",
114            metadata,
115            Column(
116                "id", Integer, primary_key=True, test_needs_autoincrement=True
117            ),
118            Column("data", String(50)),
119            implicit_returning=False,
120        )
121        Table(
122            "includes_defaults",
123            metadata,
124            Column(
125                "id", Integer, primary_key=True, test_needs_autoincrement=True
126            ),
127            Column("data", String(50)),
128            Column("x", Integer, default=5),
129            Column(
130                "y",
131                Integer,
132                default=literal_column("2", type_=Integer) + literal(2),
133            ),
134        )
135
136    @testing.variation("style", ["plain", "return_defaults"])
137    @testing.variation("executemany", [True, False])
138    def test_no_results_for_non_returning_insert(
139        self, connection, style, executemany
140    ):
141        """test another INSERT issue found during #10453"""
142
143        table = self.tables.no_implicit_returning
144
145        stmt = table.insert()
146        if style.return_defaults:
147            stmt = stmt.return_defaults()
148
149        if executemany:
150            data = [
151                {"data": "d1"},
152                {"data": "d2"},
153                {"data": "d3"},
154                {"data": "d4"},
155                {"data": "d5"},
156            ]
157        else:
158            data = {"data": "d1"}
159
160        r = connection.execute(stmt, data)
161        assert not r.returns_rows
162
163    @requirements.autoincrement_insert
164    def test_autoclose_on_insert(self, connection):
165        r = connection.execute(
166            self.tables.autoinc_pk.insert(), dict(data="some data")
167        )
168        assert r._soft_closed
169        assert not r.closed
170        assert r.is_insert
171
172        # new as of I8091919d45421e3f53029b8660427f844fee0228; for the moment
173        # an insert where the PK was taken from a row that the dialect
174        # selected, as is the case for mssql/pyodbc, will still report
175        # returns_rows as true because there's a cursor description.  in that
176        # case, the row had to have been consumed at least.
177        assert not r.returns_rows or r.fetchone() is None
178
179    @requirements.insert_returning
180    def test_autoclose_on_insert_implicit_returning(self, connection):
181        r = connection.execute(
182            # return_defaults() ensures RETURNING will be used,
183            # new in 2.0 as sqlite/mariadb offer both RETURNING and
184            # cursor.lastrowid
185            self.tables.autoinc_pk.insert().return_defaults(),
186            dict(data="some data"),
187        )
188        assert r._soft_closed
189        assert not r.closed
190        assert r.is_insert
191
192        # note we are experimenting with having this be True
193        # as of I8091919d45421e3f53029b8660427f844fee0228 .
194        # implicit returning has fetched the row, but it still is a
195        # "returns rows"
196        assert r.returns_rows
197
198        # and we should be able to fetchone() on it, we just get no row
199        eq_(r.fetchone(), None)
200
201        # and the keys, etc.
202        eq_(r.keys(), ["id"])
203
204        # but the dialect took in the row already.   not really sure
205        # what the best behavior is.
206
207    @requirements.empty_inserts
208    def test_empty_insert(self, connection):
209        r = connection.execute(self.tables.autoinc_pk.insert())
210        assert r._soft_closed
211        assert not r.closed
212
213        r = connection.execute(
214            self.tables.autoinc_pk.select().where(
215                self.tables.autoinc_pk.c.id != None
216            )
217        )
218        eq_(len(r.all()), 1)
219
220    @requirements.empty_inserts_executemany
221    def test_empty_insert_multiple(self, connection):
222        r = connection.execute(self.tables.autoinc_pk.insert(), [{}, {}, {}])
223        assert r._soft_closed
224        assert not r.closed
225
226        r = connection.execute(
227            self.tables.autoinc_pk.select().where(
228                self.tables.autoinc_pk.c.id != None
229            )
230        )
231
232        eq_(len(r.all()), 3)
233
234    @requirements.insert_from_select
235    def test_insert_from_select_autoinc(self, connection):
236        src_table = self.tables.manual_pk
237        dest_table = self.tables.autoinc_pk
238        connection.execute(
239            src_table.insert(),
240            [
241                dict(id=1, data="data1"),
242                dict(id=2, data="data2"),
243                dict(id=3, data="data3"),
244            ],
245        )
246
247        result = connection.execute(
248            dest_table.insert().from_select(
249                ("data",),
250                select(src_table.c.data).where(
251                    src_table.c.data.in_(["data2", "data3"])
252                ),
253            )
254        )
255
256        eq_(result.inserted_primary_key, (None,))
257
258        result = connection.execute(
259            select(dest_table.c.data).order_by(dest_table.c.data)
260        )
261        eq_(result.fetchall(), [("data2",), ("data3",)])
262
263    @requirements.insert_from_select
264    def test_insert_from_select_autoinc_no_rows(self, connection):
265        src_table = self.tables.manual_pk
266        dest_table = self.tables.autoinc_pk
267
268        result = connection.execute(
269            dest_table.insert().from_select(
270                ("data",),
271                select(src_table.c.data).where(
272                    src_table.c.data.in_(["data2", "data3"])
273                ),
274            )
275        )
276        eq_(result.inserted_primary_key, (None,))
277
278        result = connection.execute(
279            select(dest_table.c.data).order_by(dest_table.c.data)
280        )
281
282        eq_(result.fetchall(), [])
283
284    @requirements.insert_from_select
285    def test_insert_from_select(self, connection):
286        table = self.tables.manual_pk
287        connection.execute(
288            table.insert(),
289            [
290                dict(id=1, data="data1"),
291                dict(id=2, data="data2"),
292                dict(id=3, data="data3"),
293            ],
294        )
295
296        connection.execute(
297            table.insert()
298            .inline()
299            .from_select(
300                ("id", "data"),
301                select(table.c.id + 5, table.c.data).where(
302                    table.c.data.in_(["data2", "data3"])
303                ),
304            )
305        )
306
307        eq_(
308            connection.execute(
309                select(table.c.data).order_by(table.c.data)
310            ).fetchall(),
311            [("data1",), ("data2",), ("data2",), ("data3",), ("data3",)],
312        )
313
314    @requirements.insert_from_select
315    def test_insert_from_select_with_defaults(self, connection):
316        table = self.tables.includes_defaults
317        connection.execute(
318            table.insert(),
319            [
320                dict(id=1, data="data1"),
321                dict(id=2, data="data2"),
322                dict(id=3, data="data3"),
323            ],
324        )
325
326        connection.execute(
327            table.insert()
328            .inline()
329            .from_select(
330                ("id", "data"),
331                select(table.c.id + 5, table.c.data).where(
332                    table.c.data.in_(["data2", "data3"])
333                ),
334            )
335        )
336
337        eq_(
338            connection.execute(
339                select(table).order_by(table.c.data, table.c.id)
340            ).fetchall(),
341            [
342                (1, "data1", 5, 4),
343                (2, "data2", 5, 4),
344                (7, "data2", 5, 4),
345                (3, "data3", 5, 4),
346                (8, "data3", 5, 4),
347            ],
348        )
349
350
351class ReturningTest(fixtures.TablesTest):
352    run_create_tables = "each"
353    __requires__ = "insert_returning", "autoincrement_insert"
354    __backend__ = True
355
356    def _assert_round_trip(self, table, conn):
357        row = conn.execute(table.select()).first()
358        eq_(
359            row,
360            (
361                conn.dialect.default_sequence_base,
362                "some data",
363            ),
364        )
365
366    @classmethod
367    def define_tables(cls, metadata):
368        Table(
369            "autoinc_pk",
370            metadata,
371            Column(
372                "id", Integer, primary_key=True, test_needs_autoincrement=True
373            ),
374            Column("data", String(50)),
375        )
376
377    @requirements.fetch_rows_post_commit
378    def test_explicit_returning_pk_autocommit(self, connection):
379        table = self.tables.autoinc_pk
380        r = connection.execute(
381            table.insert().returning(table.c.id), dict(data="some data")
382        )
383        pk = r.first()[0]
384        fetched_pk = connection.scalar(select(table.c.id))
385        eq_(fetched_pk, pk)
386
387    def test_explicit_returning_pk_no_autocommit(self, connection):
388        table = self.tables.autoinc_pk
389        r = connection.execute(
390            table.insert().returning(table.c.id), dict(data="some data")
391        )
392
393        pk = r.first()[0]
394        fetched_pk = connection.scalar(select(table.c.id))
395        eq_(fetched_pk, pk)
396
397    def test_autoincrement_on_insert_implicit_returning(self, connection):
398        connection.execute(
399            self.tables.autoinc_pk.insert(), dict(data="some data")
400        )
401        self._assert_round_trip(self.tables.autoinc_pk, connection)
402
403    def test_last_inserted_id_implicit_returning(self, connection):
404        r = connection.execute(
405            self.tables.autoinc_pk.insert(), dict(data="some data")
406        )
407        pk = connection.scalar(select(self.tables.autoinc_pk.c.id))
408        eq_(r.inserted_primary_key, (pk,))
409
410    @requirements.insert_executemany_returning
411    def test_insertmanyvalues_returning(self, connection):
412        r = connection.execute(
413            self.tables.autoinc_pk.insert().returning(
414                self.tables.autoinc_pk.c.id
415            ),
416            [
417                {"data": "d1"},
418                {"data": "d2"},
419                {"data": "d3"},
420                {"data": "d4"},
421                {"data": "d5"},
422            ],
423        )
424        rall = r.all()
425
426        pks = connection.execute(select(self.tables.autoinc_pk.c.id))
427
428        eq_(rall, pks.all())
429
430    @testing.combinations(
431        (Double(), 8.5514716, True),
432        (
433            Double(53),
434            8.5514716,
435            True,
436            testing.requires.float_or_double_precision_behaves_generically,
437        ),
438        (Float(), 8.5514, True),
439        (
440            Float(8),
441            8.5514,
442            True,
443            testing.requires.float_or_double_precision_behaves_generically,
444        ),
445        (
446            Numeric(precision=15, scale=12, asdecimal=False),
447            8.5514716,
448            True,
449            testing.requires.literal_float_coercion,
450        ),
451        (
452            Numeric(precision=15, scale=12, asdecimal=True),
453            Decimal("8.5514716"),
454            False,
455        ),
456        argnames="type_,value,do_rounding",
457    )
458    @testing.variation("sort_by_parameter_order", [True, False])
459    @testing.variation("multiple_rows", [True, False])
460    def test_insert_w_floats(
461        self,
462        connection,
463        metadata,
464        sort_by_parameter_order,
465        type_,
466        value,
467        do_rounding,
468        multiple_rows,
469    ):
470        """test #9701.
471
472        this tests insertmanyvalues as well as decimal / floating point
473        RETURNING types
474
475        """
476
477        t = Table(
478            # Oracle backends seems to be getting confused if
479            # this table is named the same as the one
480            # in test_imv_returning_datatypes.  use a different name
481            "f_t",
482            metadata,
483            Column("id", Integer, Identity(), primary_key=True),
484            Column("value", type_),
485        )
486
487        t.create(connection)
488
489        result = connection.execute(
490            t.insert().returning(
491                t.c.id,
492                t.c.value,
493                sort_by_parameter_order=bool(sort_by_parameter_order),
494            ),
495            (
496                [{"value": value} for i in range(10)]
497                if multiple_rows
498                else {"value": value}
499            ),
500        )
501
502        if multiple_rows:
503            i_range = range(1, 11)
504        else:
505            i_range = range(1, 2)
506
507        # we want to test only that we are getting floating points back
508        # with some degree of the original value maintained, that it is not
509        # being truncated to an integer.  there's too much variation in how
510        # drivers return floats, which should not be relied upon to be
511        # exact, for us to just compare as is (works for PG drivers but not
512        # others) so we use rounding here.  There's precedent for this
513        # in suite/test_types.py::NumericTest as well
514
515        if do_rounding:
516            eq_(
517                {(id_, round(val_, 5)) for id_, val_ in result},
518                {(id_, round(value, 5)) for id_ in i_range},
519            )
520
521            eq_(
522                {
523                    round(val_, 5)
524                    for val_ in connection.scalars(select(t.c.value))
525                },
526                {round(value, 5)},
527            )
528        else:
529            eq_(
530                set(result),
531                {(id_, value) for id_ in i_range},
532            )
533
534            eq_(
535                set(connection.scalars(select(t.c.value))),
536                {value},
537            )
538
539    @testing.combinations(
540        (
541            "non_native_uuid",
542            Uuid(native_uuid=False),
543            uuid.uuid4(),
544        ),
545        (
546            "non_native_uuid_str",
547            Uuid(as_uuid=False, native_uuid=False),
548            str(uuid.uuid4()),
549        ),
550        (
551            "generic_native_uuid",
552            Uuid(native_uuid=True),
553            uuid.uuid4(),
554            testing.requires.uuid_data_type,
555        ),
556        (
557            "generic_native_uuid_str",
558            Uuid(as_uuid=False, native_uuid=True),
559            str(uuid.uuid4()),
560            testing.requires.uuid_data_type,
561        ),
562        ("UUID", UUID(), uuid.uuid4(), testing.requires.uuid_data_type),
563        (
564            "LargeBinary1",
565            LargeBinary(),
566            b"this is binary",
567        ),
568        ("LargeBinary2", LargeBinary(), b"7\xe7\x9f"),
569        argnames="type_,value",
570        id_="iaa",
571    )
572    @testing.variation("sort_by_parameter_order", [True, False])
573    @testing.variation("multiple_rows", [True, False])
574    @testing.requires.insert_returning
575    def test_imv_returning_datatypes(
576        self,
577        connection,
578        metadata,
579        sort_by_parameter_order,
580        type_,
581        value,
582        multiple_rows,
583    ):
584        """test #9739, #9808 (similar to #9701).
585
586        this tests insertmanyvalues in conjunction with various datatypes.
587
588        These tests are particularly for the asyncpg driver which needs
589        most types to be explicitly cast for the new IMV format
590
591        """
592        t = Table(
593            "d_t",
594            metadata,
595            Column("id", Integer, Identity(), primary_key=True),
596            Column("value", type_),
597        )
598
599        t.create(connection)
600
601        result = connection.execute(
602            t.insert().returning(
603                t.c.id,
604                t.c.value,
605                sort_by_parameter_order=bool(sort_by_parameter_order),
606            ),
607            (
608                [{"value": value} for i in range(10)]
609                if multiple_rows
610                else {"value": value}
611            ),
612        )
613
614        if multiple_rows:
615            i_range = range(1, 11)
616        else:
617            i_range = range(1, 2)
618
619        eq_(
620            set(result),
621            {(id_, value) for id_ in i_range},
622        )
623
624        eq_(
625            set(connection.scalars(select(t.c.value))),
626            {value},
627        )
628
629
630__all__ = ("LastrowidTest", "InsertBehaviorTest", "ReturningTest")
631 
codekingpro/portable-devtools · Team Ai