codekingpro/portable-devtools
114k
1# testing/suite/test_rowcount.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 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 