userver: en/testsuite/databases/mysql/pytest_plugin.py Source File
Loading...
Searching...
No Matches
pytest_plugin.py
1import collections
2
3import pytest
4
5from . import classes, control, service, utils
6
7
8def pytest_addoption(parser):
9 """
10 :param parser: pytest's argument parser
11 """
12 group = parser.getgroup('mysql')
13 group.addoption('--mysql')
14 group.addoption(
15 '--no-mysql',
16 help='Disable use of MySQL',
17 action='store_true',
18 )
19
20
21def pytest_configure(config):
22 config.addinivalue_line('markers', 'mysql: per-test MySQL initialization')
23
24
25def pytest_service_register(register_service):
26 register_service('mysql', service.create_service)
27
28
29@pytest.fixture
30def mysql(_mysql, _mysql_apply) -> dict[str, control.ConnectionWrapper]:
31 """MySQL fixture.
32
33 Returns dictionary where key is database alias and value is
34 @ref control.ConnectionWrapper
35
36 @ingroup userver_testsuite_fixtures
37 Part of the [yandex-taxi-testsuite](https://github.com/yandex/yandex-taxi-testsuite/blob/develop/testsuite/databases/mysql/pytest_plugin.py#L30)
38 """
39 return _mysql.get_wrappers()
40
41
42@pytest.fixture(scope='session')
43def mysql_disabled(pytestconfig) -> bool:
44 return pytestconfig.option.no_mysql
45
46
47@pytest.fixture(scope='session')
48def mysql_conninfo(pytestconfig, _mysql_service_settings):
49 if pytestconfig.option.mysql:
50 return service.parse_connection_url(pytestconfig.option.mysql)
51 return _mysql_service_settings.get_conninfo()
52
53
54@pytest.fixture(scope='session')
55def mysql_local() -> classes.DatabasesDict:
56 """Use to override databases configuration.
57
58 @ingroup userver_testsuite_fixtures
59 Part of the [yandex-taxi-testsuite](https://github.com/yandex/yandex-taxi-testsuite/blob/develop/testsuite/databases/mysql/pytest_plugin.py#L52)
60 """
61 return {}
62
63
64@pytest.fixture
65def _mysql(mysql_local, _mysql_service, _mysql_state):
66 if not _mysql_service:
67 mysql_local = {}
68 dbcontrol = control.Control(mysql_local, _mysql_state)
69 dbcontrol.run_migrations()
70 return dbcontrol
71
72
73@pytest.fixture
74def _mysql_apply(
75 mysql_local,
76 _mysql_state,
77 _mysql_query_loader,
78 request,
79):
80 def load_default_queries(dbname):
81 return [
82 *_mysql_query_loader.load(
83 f'my_{dbname}.sql', 'mysql.default_queries', missing_ok=True
84 ),
85 *_mysql_query_loader.loaddir(
86 f'my_{dbname}', 'mysql.default_queries', missing_ok=True
87 ),
88 ]
89
90 def mysql_mark(dbname, *, files=(), directories=(), queries=()):
91 result_queries = []
92 for path in files:
93 result_queries += _mysql_query_loader.load(path, 'mark.mysql.files')
94 for path in directories:
95 result_queries += _mysql_query_loader.loaddir(
96 path, 'mark.mysql.directories'
97 )
98 for query in queries:
99 result_queries.append(
100 control.MysqlQuery(
101 body=query,
102 source='mark.mysql.queries',
103 path=None,
104 ),
105 )
106 return dbname, result_queries
107
108 overrides = collections.defaultdict(list)
109
110 for mark in request.node.iter_markers('mysql'):
111 dbname, queries = mysql_mark(*mark.args, **mark.kwargs)
112 if dbname not in mysql_local:
113 raise RuntimeError(f'Unknown mysql database {dbname}')
114 overrides[dbname].extend(queries)
115
116 for alias, dbconfig in mysql_local.items():
117 if alias in overrides:
118 queries = overrides[alias]
119 else:
120 queries = load_default_queries(alias)
121 connection_wrapper = _mysql_state.wrapper_for(dbconfig.dbname)
122 connection_wrapper.apply_queries(
123 queries,
124 keep_tables=dbconfig.keep_tables,
125 truncate_non_empty=dbconfig.truncate_non_empty,
126 )
127
128
129@pytest.fixture
130def _mysql_query_loader(get_file_path, get_directory_path):
131 def load_query(path, source):
132 return control.MysqlQuery(
133 body=path.read_text(),
134 source=source,
135 path=str(path),
136 )
137
138 class Loader:
139 @staticmethod
140 def load(path, source, missing_ok=False):
141 data = get_file_path(path, missing_ok=missing_ok)
142 if not data:
143 return []
144 return [load_query(data, source)]
145
146 @staticmethod
147 def loaddir(directory, source, missing_ok=False):
148 result = []
149 directory = get_directory_path(directory, missing_ok=missing_ok)
150 if not directory:
151 return []
152 for path in utils.scan_sql_directory(directory):
153 result.append(load_query(path, source))
154 return result
155
156 return Loader()
157
158
159@pytest.fixture(scope='session')
160def _mysql_service_settings():
161 return service.get_service_settings()
162
163
164@pytest.fixture
165def _mysql_service(
166 ensure_service_started,
167 mysql_local,
168 mysql_disabled,
169 pytestconfig,
170 _mysql_service_settings,
171):
172 if not mysql_local or mysql_disabled:
173 return False
174 if not pytestconfig.option.mysql:
175 ensure_service_started('mysql', settings=_mysql_service_settings)
176 return True
177
178
179@pytest.fixture(scope='session')
180def _mysql_state(pytestconfig, mysql_conninfo):
181 return control.DatabasesState(
182 connections=control.ConnectionCache(mysql_conninfo),
183 verbose=pytestconfig.option.verbose,
184 )