Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_rowcount.py259 linesDownload Raw Back to suite
1# testing/suite/test_rowcount.py
2# Copyright (C) 2005-2026 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 sqlalchemy import bindparam
10from sqlalchemy import Column
11from sqlalchemy import Integer
12from sqlalchemy import MetaData
13from sqlalchemy import select
14from sqlalchemy import String
15from sqlalchemy import Table
16from sqlalchemy import testing
17from sqlalchemy import text
18from sqlalchemy.testing import eq_
19from sqlalchemy.testing import fixtures
20
21
22class RowCountTest(fixtures.TablesTest):
23    """test rowcount functionality"""
24
25    __requires__ = ("sane_rowcount",)
26    __backend__ = True
27
28    @classmethod
29    def define_tables(cls, metadata):
30        Table(
31            "employees",
32            metadata,
33            Column(
34                "employee_id",
35                Integer,
36                autoincrement=False,
37                primary_key=True,
38            ),
39            Column("name", String(50)),
40            Column("department", String(1)),
41        )
42
43    @classmethod
44    def insert_data(cls, connection):
45        cls.data = data = [
46            ("Angela", "A"),
47            ("Andrew", "A"),
48            ("Anand", "A"),
49            ("Bob", "B"),
50            ("Bobette", "B"),
51            ("Buffy", "B"),
52            ("Charlie", "C"),
53            ("Cynthia", "C"),
54            ("Chris", "C"),
55        ]
56
57        employees_table = cls.tables.employees
58        connection.execute(
59            employees_table.insert(),
60            [
61                {"employee_id": i, "name": n, "department": d}
62                for i, (n, d) in enumerate(data)
63            ],
64        )
65
66    def test_basic(self, connection):
67        employees_table = self.tables.employees
68        s = select(
69            employees_table.c.name, employees_table.c.department
70        ).order_by(employees_table.c.employee_id)
71        rows = connection.execute(s).fetchall()
72
73        eq_(rows, self.data)
74
75    @testing.variation("statement", ["update", "delete", "insert", "select"])
76    @testing.variation("close_first", [True, False])
77    def test_non_rowcount_scenarios_no_raise(
78        self, connection, statement, close_first
79    ):
80        employees_table = self.tables.employees
81
82        # WHERE matches 3, 3 rows changed
83        department = employees_table.c.department
84
85        if statement.update:
86            r = connection.execute(
87                employees_table.update().where(department == "C"),
88                {"department": "Z"},
89            )
90        elif statement.delete:
91            r = connection.execute(
92                employees_table.delete().where(department == "C"),
93                {"department": "Z"},
94            )
95        elif statement.insert:
96            r = connection.execute(
97                employees_table.insert(),
98                [
99                    {"employee_id": 25, "name": "none 1", "department": "X"},
100                    {"employee_id": 26, "name": "none 2", "department": "Z"},
101                    {"employee_id": 27, "name": "none 3", "department": "Z"},
102                ],
103            )
104        elif statement.select:
105            s = select(
106                employees_table.c.name, employees_table.c.department
107            ).where(employees_table.c.department == "C")
108            r = connection.execute(s)
109            r.all()
110        else:
111            statement.fail()
112
113        if close_first:
114            r.close()
115
116        assert r.rowcount in (-1, 3)
117
118    def test_update_rowcount1(self, connection):
119        employees_table = self.tables.employees
120
121        # WHERE matches 3, 3 rows changed
122        department = employees_table.c.department
123        r = connection.execute(
124            employees_table.update().where(department == "C"),
125            {"department": "Z"},
126        )
127        assert r.rowcount == 3
128
129    def test_update_rowcount2(self, connection):
130        employees_table = self.tables.employees
131
132        # WHERE matches 3, 0 rows changed
133        department = employees_table.c.department
134
135        r = connection.execute(
136            employees_table.update().where(department == "C"),
137            {"department": "C"},
138        )
139        eq_(r.rowcount, 3)
140
141    @testing.variation("implicit_returning", [True, False])
142    @testing.variation(
143        "dml",
144        [
145            ("update", testing.requires.update_returning),
146            ("delete", testing.requires.delete_returning),
147        ],
148    )
149    def test_update_delete_rowcount_return_defaults(
150        self, connection, implicit_returning, dml
151    ):
152        """note this test should succeed for all RETURNING backends
153        as of 2.0.  In
154        Idf28379f8705e403a3c6a937f6a798a042ef2540 we changed rowcount to use
155        len(rows) when we have implicit returning
156
157        """
158
159        if implicit_returning:
160            employees_table = self.tables.employees
161        else:
162            employees_table = Table(
163                "employees",
164                MetaData(),
165                Column(
166                    "employee_id",
167                    Integer,
168                    autoincrement=False,
169                    primary_key=True,
170                ),
171                Column("name", String(50)),
172                Column("department", String(1)),
173                implicit_returning=False,
174            )
175
176        department = employees_table.c.department
177
178        if dml.update:
179            stmt = (
180                employees_table.update()
181                .where(department == "C")
182                .values(name=employees_table.c.department + "Z")
183                .return_defaults()
184            )
185        elif dml.delete:
186            stmt = (
187                employees_table.delete()
188                .where(department == "C")
189                .return_defaults()
190            )
191        else:
192            dml.fail()
193
194        r = connection.execute(stmt)
195        eq_(r.rowcount, 3)
196
197    def test_raw_sql_rowcount(self, connection):
198        # test issue #3622, make sure eager rowcount is called for text
199        result = connection.exec_driver_sql(
200            "update employees set department='Z' where department='C'"
201        )
202        eq_(result.rowcount, 3)
203
204    def test_text_rowcount(self, connection):
205        # test issue #3622, make sure eager rowcount is called for text
206        result = connection.execute(
207            text("update employees set department='Z' where department='C'")
208        )
209        eq_(result.rowcount, 3)
210
211    def test_delete_rowcount(self, connection):
212        employees_table = self.tables.employees
213
214        # WHERE matches 3, 3 rows deleted
215        department = employees_table.c.department
216        r = connection.execute(
217            employees_table.delete().where(department == "C")
218        )
219        eq_(r.rowcount, 3)
220
221    @testing.requires.sane_multi_rowcount
222    def test_multi_update_rowcount(self, connection):
223        employees_table = self.tables.employees
224        stmt = (
225            employees_table.update()
226            .where(employees_table.c.name == bindparam("emp_name"))
227            .values(department="C")
228        )
229
230        r = connection.execute(
231            stmt,
232            [
233                {"emp_name": "Bob"},
234                {"emp_name": "Cynthia"},
235                {"emp_name": "nonexistent"},
236            ],
237        )
238
239        eq_(r.rowcount, 2)
240
241    @testing.requires.sane_multi_rowcount
242    def test_multi_delete_rowcount(self, connection):
243        employees_table = self.tables.employees
244
245        stmt = employees_table.delete().where(
246            employees_table.c.name == bindparam("emp_name")
247        )
248
249        r = connection.execute(
250            stmt,
251            [
252                {"emp_name": "Bob"},
253                {"emp_name": "Cynthia"},
254                {"emp_name": "nonexistent"},
255            ],
256        )
257
258        eq_(r.rowcount, 2)
259