userver: en/testsuite/databases/pgsql/pytest_plugin.py Source File
Loading...
Searching...
No Matches
pytest_plugin.py
1import collections
3import concurrent.futures
4import contextlib
5import re
6import typing
7
8import pytest
9
10from . import connection, control, discover, exceptions, service, utils
11
12DB_FILE_RE_PATTERN = re.compile(r'/pg_(?P<pg_db_alias>\w+)(/?\w*)\.sql$')
13
14
15class ServiceLocalConfig(collections.abc.Mapping):
16 def __init__(
17 self,
18 databases: list[discover.PgShardedDatabase],
19 pgsql_control: control.PgControl,
20 cleanup_exclude_tables: frozenset[str],
21 ):
22 self._initialized = False
23 self._pgsql_control = pgsql_control
24 self._databases = databases
25 self._shard_connections = {
26 shard.pretty_name: pgsql_control.get_connection_cached(
27 shard.dbname,
28 )
29 for db in self._databases
30 for shard in db.shards
31 }
32 self._cleanup_exclude_tables = cleanup_exclude_tables
33
34 def __len__(self) -> int:
35 return len(self._shard_connections)
36
37 def __iter__(self) -> typing.Iterator[str]:
38 return iter(self._shard_connections)
39
40 def __getitem__(self, dbname: str) -> connection.PgConnectionInfo:
41 """Get
42 :py:class:`testsuite.databases.pgsql.connection.PgConnectionInfo`
43 instance by database name
44 """
45 return self._shard_connections[dbname].conninfo
46
47 def initialize(
48 self, parallel_init: bool
49 ) -> dict[str, control.ConnectionWrapper]:
50 if self._initialized:
51 return self._shard_connections
52
53 if self._databases:
54 self._pgsql_control.initialize()
55
56 def init_database(db):
57 self._pgsql_control.initialize_sharded_db(db)
58
59 for shard in db.shards:
60 self._shard_connections[shard.pretty_name].initialize(
62 )
63
64 if parallel_init:
65 with concurrent.futures.ThreadPoolExecutor() as executor:
66 init_db_futures = [
67 executor.submit(init_database, db) for db in self._databases
68 ]
69 for future in init_db_futures:
70 future.result()
71 else:
72 for database in self._databases:
73 init_database(database)
74
75 self._initialized = True
76 return self._shard_connections
77
78
79def pytest_addoption(parser):
80 """
81 :param parser: pytest's argument parser
82 """
83 group = parser.getgroup('postgresql')
84 group.addoption('--postgresql', help='PostgreSQL connection string')
85 group.addoption(
86 '--no-postgresql',
87 help='Disable use of PostgreSQL',
88 action='store_true',
89 )
90 group.addoption(
91 '--postgresql-keep-existing-db',
92 action='store_true',
93 help=(
94 'Keep existing databases with up-to-date schema. By default '
95 'testsuite will drop and create anew any existing database when '
96 'initializing databases.'
97 ),
98 )
99
100
101def pytest_report_header(config):
102 conninfo = _get_connection_info(config)
103 return [f'PostgreSQL: {conninfo.get_uri()}']
104
105
106def pytest_configure(config):
107 config.addinivalue_line(
108 'markers',
109 'pgsql: per-test PostgreSQL initialization',
110 )
111
112
113def pytest_service_register(register_service):
114 register_service('postgresql', service.create_pgsql_service)
115
116
117@pytest.fixture(scope='session')
118def pgsql_cleanup_exclude_tables() -> frozenset[str]:
119 return frozenset()
120
121
122@pytest.fixture
123def pgsql(_pgsql, pgsql_apply) -> dict[str, control.PgDatabaseWrapper]:
124 """
125 Returns str to
126 @ref testsuite.databases.pgsql.control.PgDatabaseWrapper dictionary
127
128 Example usage:
129
130 @code{.py}
131 def test_pg(pgsql):
132 cursor = pgsql['example_db'].cursor()
133 cursor.execute('SELECT ... FROM ...WHERE ...')
134 assert list(cusror) == [...]
135 @endcode
136
137 @ingroup userver_testsuite_fixtures
138 Part of the [yandex-taxi-testsuite](https://github.com/yandex/yandex-taxi-testsuite/blob/develop/testsuite/databases/pgsql/pytest_plugin.py#L123)
139 """
140 return {
141 dbname: control.PgDatabaseWrapper(connection)
142 for dbname, connection in _pgsql.items()
143 }
144
145
146@pytest.fixture(scope='session')
148 _pgsql_control,
149 pgsql_cleanup_exclude_tables,
150) -> typing.Callable[
151 [list[discover.PgShardedDatabase]],
152 ServiceLocalConfig,
153]:
154 """Creates pgsql configuration.
155
156 @param databases List of databases.
157 @returns @ref ServiceLocalConfig instance.
158
159 @ingroup userver_testsuite_fixtures
160 Part of the [yandex-taxi-testsuite](https://github.com/yandex/yandex-taxi-testsuite/blob/develop/testsuite/databases/pgsql/pytest_plugin.py#L144)
161 """
162
163 def _pgsql_local_create(databases):
164 return ServiceLocalConfig(
165 databases,
166 _pgsql_control,
167 pgsql_cleanup_exclude_tables,
168 )
169
170 return _pgsql_local_create
171
172
173@pytest.fixture(scope='session')
174def pgsql_disabled(pytestconfig) -> bool:
175 return pytestconfig.option.no_postgresql
176
177
178@pytest.fixture
179def pgsql_local(pgsql_local_create) -> ServiceLocalConfig:
180 """Configures local pgsql instance.
181
182 @returns @ref ServiceLocalConfig instance.
183
184 In order to use pgsql fixture you have to override pgsql_local()
185 in your local conftest.py file, example:
186
187 @code{.py}
188 @pytest.fixture(scope='session')
189 def pgsql_local(pgsql_local_create):
190 databases = discover.find_schemas(
191 'service_name', [PG_SCHEMAS_PATH])
192 return pgsql_local_create(list(databases.values()))
193 @endcode
194
195 Sometimes it is desirable to have tests-only database, maybe used in one
196 particular test or tests group. This can be achieved by by overriding
197 @c pgsql_local fixture in your test file:
198
199 @code{.py}
200 @pytest.fixture
201 def pgsql_local(pgsql_local_create):
202 databases = discover.find_schemas(
203 'testsuite', [pathlib.Path('custom/pgsql/schema/path')])
204 return pgsql_local_create(list(databases.values()))
205 @endcode
206
207 @c pgsql_local provides access to PostgreSQL connection parameters:
208
209 @code{.py}
210 def get_custom_connection_string(pgsql_local):
211 conninfo = pgsql_local['database_name']
212 custom_dsn: str = conninfo.replace(options='-c opt=val').get_dsn()
213 return custom_dsn
214 @endcode
215
216 @ingroup userver_testsuite_fixtures
217 Part of the [yandex-taxi-testsuite](https://github.com/yandex/yandex-taxi-testsuite/blob/develop/testsuite/databases/pgsql/pytest_plugin.py#L173)
218 """
219 return pgsql_local_create([])
220
221
222@pytest.fixture(scope='session')
223def pgsql_parallelization_enabled():
224 return True
225
226
227@pytest.fixture
228def _pgsql(
229 _pgsql_service,
230 _pgsql_control,
231 pgsql_local,
232 pgsql_cleanup_exclude_tables,
233 pgsql_disabled: bool,
234 pgsql_parallelization_enabled: bool,
235) -> dict[str, control.ConnectionWrapper]:
236 if pgsql_disabled:
237 pgsql_local = ServiceLocalConfig(
238 [],
239 _pgsql_control,
240 pgsql_cleanup_exclude_tables,
241 )
242 return pgsql_local.initialize(parallel_init=pgsql_parallelization_enabled)
243
244
245@pytest.fixture(scope='session')
246def pgsql_background_truncate_enabled():
247 return True
248
249
250@pytest.fixture
251def _pgsql_apply_queries(
252 request, _pgsql: ServiceLocalConfig, _pgsql_query_loader
253) -> dict[str, list[control.PgQuery]]:
254 def pgsql_default_queries(dbname):
255 return [
256 *_pgsql_query_loader.load(
257 f'pg_{dbname}.sql',
258 'pgsql.default_queries',
259 missing_ok=True,
260 ),
261 *_pgsql_query_loader.loaddir(
262 f'pg_{dbname}',
263 'pgsql.default_queries',
264 missing_ok=True,
265 ),
266 ]
267
268 def pgsql_mark(dbname, files=(), directories=(), queries=()):
269 result_queries = []
270
271 for path in files:
272 result_queries += _pgsql_query_loader.load(path, 'mark.pgsql.files')
273 for path in directories:
274 result_queries += _pgsql_query_loader.loaddir(
275 path,
276 'mark.pgsql.directories',
277 )
278 for query in queries:
279 queries_str: typing.Iterable = []
280 if isinstance(query, str):
281 queries_str = [query]
282 elif isinstance(query, (list, tuple)):
283 queries_str = query
284 else:
286 f'sql queries of type {type(query)} are not supported',
287 )
288 for query_str in queries_str:
289 result_queries.append(
291 body=query_str,
292 source='mark.pgsql.queries',
293 path=None,
294 ),
295 )
296 return dbname, result_queries
297
298 overrides: typing.DefaultDict[
299 str,
300 list[control.PgQuery],
301 ] = collections.defaultdict(list)
302 for mark in request.node.iter_markers('pgsql'):
303 dbname, queries = pgsql_mark(*mark.args, **mark.kwargs)
304 if dbname not in _pgsql:
306 'Unknown database {}'.format(dbname)
307 )
308 overrides[dbname].extend(queries)
309
310 queries = {}
311
312 for dbname in _pgsql.keys():
313 queries[dbname] = overrides.get(dbname, pgsql_default_queries(dbname))
314
315 return queries
316
317
318@pytest.fixture
319def pgsql_apply(
320 _pgsql: ServiceLocalConfig,
321 load,
322 pgsql_background_truncate_enabled: bool,
323 pgsql_parallelization_enabled: bool,
324 _pgsql_apply_queries,
325) -> None:
326 """Initialize PostgreSQL database with data.
327
328 By default pg_${DBNAME}.sql and pg_${DBNAME}/*.sql files are used
329 to fill PostgreSQL databases.
330
331 Use pytest.mark.pgsql to change this behaviour:
332
333 @code{.py}
334 @pytest.mark.pgsql(
335 'foo@0',
336 files=[
337 'pg_foo@0_alternative.sql'
338 ],
339 directories=[
340 'pg_foo@0_alternative_dir'
341 ],
342 queries=[
343 'INSERT INTO foo VALUES (1, 2, 3, 4)',
344 ]
345 )
346 @endcode
347
348 @ingroup userver_testsuite_fixtures
349 Part of the [yandex-taxi-testsuite](https://github.com/yandex/yandex-taxi-testsuite/blob/develop/testsuite/databases/pgsql/pytest_plugin.py#L310)
350 """
351
352 if pgsql_parallelization_enabled:
353 with concurrent.futures.ThreadPoolExecutor() as executor:
354 db_apply_queries_future = []
355 for dbname, pg_db in _pgsql.items():
356 db_apply_queries_future.append(
357 executor.submit(
358 pg_db.apply_queries, _pgsql_apply_queries[dbname]
359 )
360 )
361
362 for future in db_apply_queries_future:
363 future.result()
364
365 else:
366 for dbname, pg_db in _pgsql.items():
367 pg_db.apply_queries(_pgsql_apply_queries[dbname])
368
369 yield
370
371 if pgsql_background_truncate_enabled:
372 for pg_db in _pgsql.values():
373 pg_db.schedule_truncation()
374
375
376@pytest.fixture
377def _pgsql_query_loader(get_file_path, get_directory_path, mockserver_info):
378 def substitute_mockserver(str_val: str):
379 return str_val.replace(
380 '$mockserver',
381 f'http://{mockserver_info.host}:{mockserver_info.port}',
382 )
383
384 def load_pg_file(path, source):
385 query = substitute_mockserver(path.read_text())
386 return control.PgQuery(body=query, source=source, path=str(path))
387
388 class Loader:
389 @staticmethod
390 def load(path, source, missing_ok=False):
391 path = get_file_path(path, missing_ok=missing_ok)
392 if not path:
393 return []
394 return [load_pg_file(path, source)]
395
396 @staticmethod
397 def loaddir(directory, source, missing_ok=False):
398 result = []
399 directory = get_directory_path(directory, missing_ok=missing_ok)
400 if not directory:
401 return []
402 for path in utils.scan_sql_directory(directory):
403 result.append(load_pg_file(path, source))
404 return result
405
406 return Loader()
407
408
409@pytest.fixture
410def _pgsql_service(
411 pytestconfig,
412 pgsql_disabled: bool,
413 ensure_service_started,
414 pgsql_local: ServiceLocalConfig,
415 _pgsql_service_settings,
416) -> None:
417 if (
418 not pgsql_disabled
419 and pgsql_local
420 and not pytestconfig.option.postgresql
421 ):
422 ensure_service_started('postgresql', settings=_pgsql_service_settings)
423
424
425@pytest.fixture(scope='session')
426def _pgsql_control(pytestconfig, pgsql_disabled: bool):
427 if pgsql_disabled:
428 return {}
429 instance = control.PgControl(
430 _get_connection_info(pytestconfig),
431 verbose=pytestconfig.option.verbose,
432 skip_applied_schemas=(
433 pytestconfig.option.postgresql_keep_existing_db
434 or pytestconfig.option.service_wait
435 ),
436 )
437 with contextlib.closing(instance):
438 yield instance
439
440
441@pytest.fixture(scope='session')
442def _pgsql_service_settings() -> service.ServiceSettings:
443 return service.get_service_settings()
444
445
446@pytest.fixture(scope='session')
447def _pgsql_conninfo(
448 request,
449 _pgsql_service_settings,
450) -> connection.PgConnectionInfo:
451 return _get_connection_info(request.config)
452
453
454def _get_connection_info(config):
455 connstr = config.option.postgresql
456 if connstr:
457 return connection.parse_connection_string(connstr)
458 settings = service.get_service_settings()
459 return settings.get_conninfo()