asyncpg 0.32.0__cp315-cp315-win32.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (89) hide show
  1. asyncpg/__init__.py +24 -0
  2. asyncpg/_asyncio_compat.py +94 -0
  3. asyncpg/_testbase/__init__.py +552 -0
  4. asyncpg/_testbase/fuzzer.py +363 -0
  5. asyncpg/_version.py +17 -0
  6. asyncpg/cluster.py +729 -0
  7. asyncpg/compat.py +88 -0
  8. asyncpg/connect_utils.py +1373 -0
  9. asyncpg/connection.py +2828 -0
  10. asyncpg/connresource.py +44 -0
  11. asyncpg/cursor.py +323 -0
  12. asyncpg/exceptions/__init__.py +1211 -0
  13. asyncpg/exceptions/_base.py +299 -0
  14. asyncpg/introspection.py +296 -0
  15. asyncpg/pgproto/__init__.pxd +5 -0
  16. asyncpg/pgproto/__init__.py +5 -0
  17. asyncpg/pgproto/buffer.pxd +143 -0
  18. asyncpg/pgproto/buffer.pxi +3 -0
  19. asyncpg/pgproto/buffer.pyx +829 -0
  20. asyncpg/pgproto/codecs/__init__.pxd +159 -0
  21. asyncpg/pgproto/codecs/bits.pyx +47 -0
  22. asyncpg/pgproto/codecs/bytea.pyx +34 -0
  23. asyncpg/pgproto/codecs/context.pyx +26 -0
  24. asyncpg/pgproto/codecs/datetime.pyx +423 -0
  25. asyncpg/pgproto/codecs/float.pyx +34 -0
  26. asyncpg/pgproto/codecs/geometry.pyx +164 -0
  27. asyncpg/pgproto/codecs/hstore.pyx +73 -0
  28. asyncpg/pgproto/codecs/int.pyx +144 -0
  29. asyncpg/pgproto/codecs/json.pyx +57 -0
  30. asyncpg/pgproto/codecs/jsonpath.pyx +29 -0
  31. asyncpg/pgproto/codecs/misc.pyx +16 -0
  32. asyncpg/pgproto/codecs/network.pyx +139 -0
  33. asyncpg/pgproto/codecs/numeric.pyx +356 -0
  34. asyncpg/pgproto/codecs/pg_snapshot.pyx +63 -0
  35. asyncpg/pgproto/codecs/text.pyx +48 -0
  36. asyncpg/pgproto/codecs/tid.pyx +51 -0
  37. asyncpg/pgproto/codecs/uuid.pyx +27 -0
  38. asyncpg/pgproto/consts.pxi +9 -0
  39. asyncpg/pgproto/cpythonx.pxd +23 -0
  40. asyncpg/pgproto/debug.pxd +10 -0
  41. asyncpg/pgproto/frb.pxd +48 -0
  42. asyncpg/pgproto/frb.pyx +12 -0
  43. asyncpg/pgproto/hton.pxd +24 -0
  44. asyncpg/pgproto/pgproto.cp315-win32.pyd +0 -0
  45. asyncpg/pgproto/pgproto.pxd +19 -0
  46. asyncpg/pgproto/pgproto.pyi +20 -0
  47. asyncpg/pgproto/pgproto.pyx +49 -0
  48. asyncpg/pgproto/tohex.pxd +10 -0
  49. asyncpg/pgproto/types.py +435 -0
  50. asyncpg/pgproto/uuid.pyx +359 -0
  51. asyncpg/pool.py +1389 -0
  52. asyncpg/prepared_stmt.py +286 -0
  53. asyncpg/protocol/__init__.py +12 -0
  54. asyncpg/protocol/codecs/__init__.py +0 -0
  55. asyncpg/protocol/codecs/array.pyx +875 -0
  56. asyncpg/protocol/codecs/base.pxd +199 -0
  57. asyncpg/protocol/codecs/base.pyx +922 -0
  58. asyncpg/protocol/codecs/pgproto.pyx +485 -0
  59. asyncpg/protocol/codecs/range.pyx +207 -0
  60. asyncpg/protocol/codecs/record.pyx +71 -0
  61. asyncpg/protocol/codecs/textutils.pyx +99 -0
  62. asyncpg/protocol/consts.pxi +12 -0
  63. asyncpg/protocol/coreproto.pxd +193 -0
  64. asyncpg/protocol/coreproto.pyx +1239 -0
  65. asyncpg/protocol/cpythonx.pxd +19 -0
  66. asyncpg/protocol/encodings.pyx +63 -0
  67. asyncpg/protocol/pgtypes.pxi +270 -0
  68. asyncpg/protocol/prepared_stmt.pxd +39 -0
  69. asyncpg/protocol/prepared_stmt.pyx +395 -0
  70. asyncpg/protocol/protocol.cp315-win32.pyd +0 -0
  71. asyncpg/protocol/protocol.pxd +78 -0
  72. asyncpg/protocol/protocol.pyi +285 -0
  73. asyncpg/protocol/protocol.pyx +1076 -0
  74. asyncpg/protocol/record.cp315-win32.pyd +0 -0
  75. asyncpg/protocol/record.pyi +29 -0
  76. asyncpg/protocol/recordcapi.pxd +14 -0
  77. asyncpg/protocol/scram.pxd +31 -0
  78. asyncpg/protocol/scram.pyx +331 -0
  79. asyncpg/protocol/settings.pxd +30 -0
  80. asyncpg/protocol/settings.pyx +106 -0
  81. asyncpg/serverversion.py +70 -0
  82. asyncpg/transaction.py +246 -0
  83. asyncpg/types.py +223 -0
  84. asyncpg/utils.py +52 -0
  85. asyncpg-0.32.0.dist-info/METADATA +131 -0
  86. asyncpg-0.32.0.dist-info/RECORD +89 -0
  87. asyncpg-0.32.0.dist-info/WHEEL +5 -0
  88. asyncpg-0.32.0.dist-info/licenses/LICENSE +204 -0
  89. asyncpg-0.32.0.dist-info/top_level.txt +1 -0
asyncpg/__init__.py ADDED
@@ -0,0 +1,24 @@
1
+ # Copyright (C) 2016-present the asyncpg authors and contributors
2
+ # <see AUTHORS file>
3
+ #
4
+ # This module is part of asyncpg and is released under
5
+ # the Apache 2.0 License: http://www.apache.org/licenses/LICENSE-2.0
6
+
7
+ from __future__ import annotations
8
+
9
+ from .connection import connect, Connection # NOQA
10
+ from .exceptions import * # NOQA
11
+ from .pool import create_pool, Pool # NOQA
12
+ from .protocol import Record # NOQA
13
+ from .types import * # NOQA
14
+
15
+
16
+ from ._version import __version__ # NOQA
17
+
18
+ from . import exceptions
19
+
20
+
21
+ __all__: tuple[str, ...] = (
22
+ 'connect', 'create_pool', 'Pool', 'Record', 'Connection'
23
+ )
24
+ __all__ += exceptions.__all__ # NOQA
@@ -0,0 +1,94 @@
1
+ # Backports from Python/Lib/asyncio for older Pythons
2
+ #
3
+ # Copyright (c) 2001-2023 Python Software Foundation; All Rights Reserved
4
+ #
5
+ # SPDX-License-Identifier: PSF-2.0
6
+
7
+ from __future__ import annotations
8
+
9
+ import asyncio
10
+ import functools
11
+ import sys
12
+ import typing
13
+
14
+ if typing.TYPE_CHECKING:
15
+ from . import compat
16
+
17
+ if sys.version_info < (3, 11):
18
+ from async_timeout import timeout as timeout_ctx
19
+ else:
20
+ from asyncio import timeout as timeout_ctx
21
+
22
+ _T = typing.TypeVar('_T')
23
+
24
+
25
+ async def wait_for(fut: compat.Awaitable[_T], timeout: float | None) -> _T:
26
+ """Wait for the single Future or coroutine to complete, with timeout.
27
+
28
+ Coroutine will be wrapped in Task.
29
+
30
+ Returns result of the Future or coroutine. When a timeout occurs,
31
+ it cancels the task and raises TimeoutError. To avoid the task
32
+ cancellation, wrap it in shield().
33
+
34
+ If the wait is cancelled, the task is also cancelled.
35
+
36
+ If the task suppresses the cancellation and returns a value instead,
37
+ that value is returned.
38
+
39
+ This function is a coroutine.
40
+ """
41
+ # The special case for timeout <= 0 is for the following case:
42
+ #
43
+ # async def test_waitfor():
44
+ # func_started = False
45
+ #
46
+ # async def func():
47
+ # nonlocal func_started
48
+ # func_started = True
49
+ #
50
+ # try:
51
+ # await asyncio.wait_for(func(), 0)
52
+ # except asyncio.TimeoutError:
53
+ # assert not func_started
54
+ # else:
55
+ # assert False
56
+ #
57
+ # asyncio.run(test_waitfor())
58
+
59
+ if timeout is not None and timeout <= 0:
60
+ fut = asyncio.ensure_future(fut)
61
+
62
+ if fut.done():
63
+ return fut.result()
64
+
65
+ await _cancel_and_wait(fut)
66
+ try:
67
+ return fut.result()
68
+ except asyncio.CancelledError as exc:
69
+ raise TimeoutError from exc
70
+
71
+ async with timeout_ctx(timeout):
72
+ return await fut
73
+
74
+
75
+ async def _cancel_and_wait(fut: asyncio.Future[_T]) -> None:
76
+ """Cancel the *fut* future or task and wait until it completes."""
77
+
78
+ loop = asyncio.get_running_loop()
79
+ waiter = loop.create_future()
80
+ cb = functools.partial(_release_waiter, waiter)
81
+ fut.add_done_callback(cb)
82
+
83
+ try:
84
+ fut.cancel()
85
+ # We cannot wait on *fut* directly to make
86
+ # sure _cancel_and_wait itself is reliably cancellable.
87
+ await waiter
88
+ finally:
89
+ fut.remove_done_callback(cb)
90
+
91
+
92
+ def _release_waiter(waiter: asyncio.Future[typing.Any], *args: object) -> None:
93
+ if not waiter.done():
94
+ waiter.set_result(None)
@@ -0,0 +1,552 @@
1
+ # Copyright (C) 2016-present the asyncpg authors and contributors
2
+ # <see AUTHORS file>
3
+ #
4
+ # This module is part of asyncpg and is released under
5
+ # the Apache 2.0 License: http://www.apache.org/licenses/LICENSE-2.0
6
+
7
+
8
+ import asyncio
9
+ import atexit
10
+ import contextlib
11
+ import functools
12
+ import inspect
13
+ import logging
14
+ import os
15
+ import re
16
+ import textwrap
17
+ import time
18
+ import traceback
19
+ import unittest
20
+
21
+
22
+ import asyncpg
23
+ from asyncpg import cluster as pg_cluster
24
+ from asyncpg import connection as pg_connection
25
+ from asyncpg import pool as pg_pool
26
+
27
+ from . import fuzzer
28
+
29
+
30
+ @contextlib.contextmanager
31
+ def silence_asyncio_long_exec_warning():
32
+ def flt(log_record):
33
+ msg = log_record.getMessage()
34
+ return not msg.startswith('Executing ')
35
+
36
+ logger = logging.getLogger('asyncio')
37
+ logger.addFilter(flt)
38
+ try:
39
+ yield
40
+ finally:
41
+ logger.removeFilter(flt)
42
+
43
+
44
+ def with_timeout(timeout):
45
+ def wrap(func):
46
+ func.__timeout__ = timeout
47
+ return func
48
+
49
+ return wrap
50
+
51
+
52
+ class TestCaseMeta(type(unittest.TestCase)):
53
+ TEST_TIMEOUT = None
54
+
55
+ @staticmethod
56
+ def _iter_methods(bases, ns):
57
+ for base in bases:
58
+ for methname in dir(base):
59
+ if not methname.startswith('test_'):
60
+ continue
61
+
62
+ meth = getattr(base, methname)
63
+ if not inspect.iscoroutinefunction(meth):
64
+ continue
65
+
66
+ yield methname, meth
67
+
68
+ for methname, meth in ns.items():
69
+ if not methname.startswith('test_'):
70
+ continue
71
+
72
+ if not inspect.iscoroutinefunction(meth):
73
+ continue
74
+
75
+ yield methname, meth
76
+
77
+ def __new__(mcls, name, bases, ns):
78
+ for methname, meth in mcls._iter_methods(bases, ns):
79
+ @functools.wraps(meth)
80
+ def wrapper(self, *args, __meth__=meth, **kwargs):
81
+ coro = __meth__(self, *args, **kwargs)
82
+ timeout = getattr(__meth__, '__timeout__', mcls.TEST_TIMEOUT)
83
+ if timeout:
84
+ coro = asyncio.wait_for(coro, timeout)
85
+ try:
86
+ self.loop.run_until_complete(coro)
87
+ except asyncio.TimeoutError:
88
+ raise self.failureException(
89
+ 'test timed out after {} seconds'.format(
90
+ timeout)) from None
91
+ else:
92
+ self.loop.run_until_complete(coro)
93
+ ns[methname] = wrapper
94
+
95
+ return super().__new__(mcls, name, bases, ns)
96
+
97
+
98
+ class TestCase(unittest.TestCase, metaclass=TestCaseMeta):
99
+
100
+ @classmethod
101
+ def setUpClass(cls):
102
+ if os.environ.get('USE_UVLOOP'):
103
+ import uvloop
104
+ loop = uvloop.new_event_loop()
105
+ else:
106
+ loop = asyncio.new_event_loop()
107
+
108
+ asyncio.set_event_loop(None)
109
+ cls.loop = loop
110
+
111
+ @classmethod
112
+ def tearDownClass(cls):
113
+ cls.loop.close()
114
+ asyncio.set_event_loop(None)
115
+
116
+ def setUp(self):
117
+ self.loop.set_exception_handler(self.loop_exception_handler)
118
+ self.__unhandled_exceptions = []
119
+
120
+ def tearDown(self):
121
+ excs = []
122
+ for exc in self.__unhandled_exceptions:
123
+ if isinstance(exc, ConnectionResetError):
124
+ texc = traceback.TracebackException.from_exception(
125
+ exc, lookup_lines=False)
126
+ if texc.stack[-1].name == "_call_connection_lost":
127
+ # On Windows calling socket.shutdown may raise
128
+ # ConnectionResetError, which happens in the
129
+ # finally block of _call_connection_lost.
130
+ continue
131
+ excs.append(exc)
132
+
133
+ if excs:
134
+ formatted = []
135
+
136
+ for i, context in enumerate(excs):
137
+ formatted.append(self._format_loop_exception(context, i + 1))
138
+
139
+ self.fail(
140
+ 'unexpected exceptions in asynchronous code:\n' +
141
+ '\n'.join(formatted))
142
+
143
+ @contextlib.contextmanager
144
+ def assertRunUnder(self, delta):
145
+ st = time.monotonic()
146
+ try:
147
+ yield
148
+ finally:
149
+ elapsed = time.monotonic() - st
150
+ if elapsed > delta:
151
+ raise AssertionError(
152
+ 'running block took {:0.3f}s which is longer '
153
+ 'than the expected maximum of {:0.3f}s'.format(
154
+ elapsed, delta))
155
+
156
+ @contextlib.contextmanager
157
+ def assertLoopErrorHandlerCalled(self, msg_re: str):
158
+ contexts = []
159
+
160
+ def handler(loop, ctx):
161
+ contexts.append(ctx)
162
+
163
+ old_handler = self.loop.get_exception_handler()
164
+ self.loop.set_exception_handler(handler)
165
+ try:
166
+ yield
167
+
168
+ for ctx in contexts:
169
+ msg = ctx.get('message')
170
+ if msg and re.search(msg_re, msg):
171
+ return
172
+
173
+ raise AssertionError(
174
+ 'no message matching {!r} was logged with '
175
+ 'loop.call_exception_handler()'.format(msg_re))
176
+
177
+ finally:
178
+ self.loop.set_exception_handler(old_handler)
179
+
180
+ def loop_exception_handler(self, loop, context):
181
+ self.__unhandled_exceptions.append(context)
182
+ loop.default_exception_handler(context)
183
+
184
+ def _format_loop_exception(self, context, n):
185
+ message = context.get('message', 'Unhandled exception in event loop')
186
+ exception = context.get('exception')
187
+ if exception is not None:
188
+ exc_info = (type(exception), exception, exception.__traceback__)
189
+ else:
190
+ exc_info = None
191
+
192
+ lines = []
193
+ for key in sorted(context):
194
+ if key in {'message', 'exception'}:
195
+ continue
196
+ value = context[key]
197
+ if key == 'source_traceback':
198
+ tb = ''.join(traceback.format_list(value))
199
+ value = 'Object created at (most recent call last):\n'
200
+ value += tb.rstrip()
201
+ else:
202
+ try:
203
+ value = repr(value)
204
+ except Exception as ex:
205
+ value = ('Exception in __repr__ {!r}; '
206
+ 'value type: {!r}'.format(ex, type(value)))
207
+ lines.append('[{}]: {}\n\n'.format(key, value))
208
+
209
+ if exc_info is not None:
210
+ lines.append('[exception]:\n')
211
+ formatted_exc = textwrap.indent(
212
+ ''.join(traceback.format_exception(*exc_info)), ' ')
213
+ lines.append(formatted_exc)
214
+
215
+ details = textwrap.indent(''.join(lines), ' ')
216
+ return '{:02d}. {}:\n{}\n'.format(n, message, details)
217
+
218
+
219
+ _default_cluster = None
220
+
221
+
222
+ def _init_cluster(ClusterCls, cluster_kwargs, initdb_options=None):
223
+ cluster = ClusterCls(**cluster_kwargs)
224
+ cluster.init(**(initdb_options or {}))
225
+ cluster.trust_local_connections()
226
+ atexit.register(_shutdown_cluster, cluster)
227
+ return cluster
228
+
229
+
230
+ def _get_initdb_options(initdb_options=None):
231
+ if not initdb_options:
232
+ initdb_options = {}
233
+ else:
234
+ initdb_options = dict(initdb_options)
235
+
236
+ # Make the default superuser name stable.
237
+ if 'username' not in initdb_options:
238
+ initdb_options['username'] = 'postgres'
239
+
240
+ return initdb_options
241
+
242
+
243
+ def _init_default_cluster(initdb_options=None):
244
+ global _default_cluster
245
+
246
+ if _default_cluster is None:
247
+ pg_host = os.environ.get('PGHOST')
248
+ if pg_host:
249
+ # Using existing cluster, assuming it is initialized and running
250
+ _default_cluster = pg_cluster.RunningCluster()
251
+ else:
252
+ _default_cluster = _init_cluster(
253
+ pg_cluster.TempCluster,
254
+ cluster_kwargs={
255
+ "data_dir_suffix": ".apgtest",
256
+ },
257
+ initdb_options=_get_initdb_options(initdb_options),
258
+ )
259
+
260
+ return _default_cluster
261
+
262
+
263
+ def _shutdown_cluster(cluster):
264
+ if cluster.get_status() == 'running':
265
+ cluster.stop()
266
+ if cluster.get_status() != 'not-initialized':
267
+ cluster.destroy()
268
+
269
+
270
+ def create_pool(dsn=None, *,
271
+ init_size=None,
272
+ min_size=10,
273
+ max_size=10,
274
+ max_queries=50000,
275
+ max_inactive_connection_lifetime=60.0,
276
+ connect=None,
277
+ setup=None,
278
+ init=None,
279
+ loop=None,
280
+ pool_class=pg_pool.Pool,
281
+ connection_class=pg_connection.Connection,
282
+ record_class=asyncpg.Record,
283
+ **connect_kwargs):
284
+ return pool_class(
285
+ dsn,
286
+ init_size=init_size,
287
+ min_size=min_size,
288
+ max_size=max_size,
289
+ max_queries=max_queries,
290
+ loop=loop,
291
+ connect=connect,
292
+ setup=setup,
293
+ init=init,
294
+ max_inactive_connection_lifetime=max_inactive_connection_lifetime,
295
+ connection_class=connection_class,
296
+ record_class=record_class,
297
+ **connect_kwargs,
298
+ )
299
+
300
+
301
+ class ClusterTestCase(TestCase):
302
+ @classmethod
303
+ def get_server_settings(cls):
304
+ settings = {
305
+ 'log_connections': 'on'
306
+ }
307
+
308
+ if cls.cluster.get_pg_version() >= (11, 0):
309
+ # JITting messes up timing tests, and
310
+ # is not essential for testing.
311
+ settings['jit'] = 'off'
312
+
313
+ return settings
314
+
315
+ @classmethod
316
+ def new_cluster(cls, ClusterCls, *, cluster_kwargs={}, initdb_options={}):
317
+ cluster = _init_cluster(ClusterCls, cluster_kwargs,
318
+ _get_initdb_options(initdb_options))
319
+ cls._clusters.append(cluster)
320
+ return cluster
321
+
322
+ @classmethod
323
+ def start_cluster(cls, cluster, *, server_settings={}):
324
+ cluster.start(port='dynamic', server_settings=server_settings)
325
+
326
+ @classmethod
327
+ def setup_cluster(cls):
328
+ cls.cluster = _init_default_cluster()
329
+
330
+ if cls.cluster.get_status() != 'running':
331
+ cls.cluster.start(
332
+ port='dynamic', server_settings=cls.get_server_settings())
333
+
334
+ @classmethod
335
+ def setUpClass(cls):
336
+ super().setUpClass()
337
+ cls._clusters = []
338
+ cls.setup_cluster()
339
+
340
+ @classmethod
341
+ def tearDownClass(cls):
342
+ super().tearDownClass()
343
+ for cluster in cls._clusters:
344
+ if cluster is not _default_cluster:
345
+ cluster.stop()
346
+ cluster.destroy()
347
+ cls._clusters = []
348
+
349
+ @classmethod
350
+ def get_connection_spec(cls, kwargs={}):
351
+ conn_spec = cls.cluster.get_connection_spec()
352
+ if kwargs.get('dsn'):
353
+ conn_spec.pop('host')
354
+ conn_spec.update(kwargs)
355
+ if not os.environ.get('PGHOST') and not kwargs.get('dsn'):
356
+ if 'database' not in conn_spec:
357
+ conn_spec['database'] = 'postgres'
358
+ if 'user' not in conn_spec:
359
+ conn_spec['user'] = 'postgres'
360
+ return conn_spec
361
+
362
+ @classmethod
363
+ def connect(cls, **kwargs):
364
+ conn_spec = cls.get_connection_spec(kwargs)
365
+ return pg_connection.connect(**conn_spec, loop=cls.loop)
366
+
367
+ def setUp(self):
368
+ super().setUp()
369
+ self._pools = []
370
+
371
+ def tearDown(self):
372
+ maintenance_tasks = []
373
+ for pool in self._pools:
374
+ pool.terminate()
375
+ if pool._maintenance_task is not None:
376
+ maintenance_tasks.append(pool._maintenance_task)
377
+ if maintenance_tasks:
378
+ self.loop.run_until_complete(asyncio.gather(
379
+ *maintenance_tasks, return_exceptions=True))
380
+ self._pools = []
381
+ super().tearDown()
382
+
383
+ def create_pool(self, pool_class=pg_pool.Pool,
384
+ connection_class=pg_connection.Connection, **kwargs):
385
+ conn_spec = self.get_connection_spec(kwargs)
386
+ pool = create_pool(loop=self.loop, pool_class=pool_class,
387
+ connection_class=connection_class, **conn_spec)
388
+ self._pools.append(pool)
389
+ return pool
390
+
391
+
392
+ class ProxiedClusterTestCase(ClusterTestCase):
393
+ @classmethod
394
+ def get_server_settings(cls):
395
+ settings = dict(super().get_server_settings())
396
+ settings['listen_addresses'] = '127.0.0.1'
397
+ return settings
398
+
399
+ @classmethod
400
+ def get_proxy_settings(cls):
401
+ return {'fuzzing-mode': None}
402
+
403
+ @classmethod
404
+ def setUpClass(cls):
405
+ super().setUpClass()
406
+ conn_spec = cls.cluster.get_connection_spec()
407
+ host = conn_spec.get('host')
408
+ if not host:
409
+ host = '127.0.0.1'
410
+ elif host.startswith('/'):
411
+ host = '127.0.0.1'
412
+ cls.proxy = fuzzer.TCPFuzzingProxy(
413
+ backend_host=host,
414
+ backend_port=int(conn_spec['port']),
415
+ )
416
+ cls.proxy.start()
417
+
418
+ @classmethod
419
+ def tearDownClass(cls):
420
+ cls.proxy.stop()
421
+ super().tearDownClass()
422
+
423
+ @classmethod
424
+ def get_connection_spec(cls, kwargs):
425
+ conn_spec = super().get_connection_spec(kwargs)
426
+ conn_spec['host'] = cls.proxy.listening_addr
427
+ conn_spec['port'] = cls.proxy.listening_port
428
+ return conn_spec
429
+
430
+ def tearDown(self):
431
+ self.proxy.reset()
432
+ super().tearDown()
433
+
434
+
435
+ def with_connection_options(**options):
436
+ if not options:
437
+ raise ValueError('no connection options were specified')
438
+
439
+ def wrap(func):
440
+ func.__connect_options__ = options
441
+ return func
442
+
443
+ return wrap
444
+
445
+
446
+ class ConnectedTestCase(ClusterTestCase):
447
+
448
+ def setUp(self):
449
+ super().setUp()
450
+
451
+ # Extract options set up with `with_connection_options`.
452
+ test_func = getattr(self, self._testMethodName).__func__
453
+ opts = getattr(test_func, '__connect_options__', {})
454
+ self.con = self.loop.run_until_complete(self.connect(**opts))
455
+ self.server_version = self.con.get_server_version()
456
+
457
+ def tearDown(self):
458
+ try:
459
+ self.loop.run_until_complete(self.con.close())
460
+ self.con = None
461
+ finally:
462
+ super().tearDown()
463
+
464
+
465
+ class HotStandbyTestCase(ClusterTestCase):
466
+
467
+ @classmethod
468
+ def setup_cluster(cls):
469
+ cls.master_cluster = cls.new_cluster(pg_cluster.TempCluster)
470
+ cls.start_cluster(
471
+ cls.master_cluster,
472
+ server_settings={
473
+ 'max_wal_senders': 10,
474
+ 'wal_level': 'hot_standby'
475
+ }
476
+ )
477
+
478
+ con = None
479
+
480
+ try:
481
+ con = cls.loop.run_until_complete(
482
+ cls.master_cluster.connect(
483
+ database='postgres', user='postgres', loop=cls.loop))
484
+
485
+ cls.loop.run_until_complete(
486
+ con.execute('''
487
+ CREATE ROLE replication WITH LOGIN REPLICATION
488
+ '''))
489
+
490
+ cls.master_cluster.trust_local_replication_by('replication')
491
+
492
+ conn_spec = cls.master_cluster.get_connection_spec()
493
+
494
+ cls.standby_cluster = cls.new_cluster(
495
+ pg_cluster.HotStandbyCluster,
496
+ cluster_kwargs={
497
+ 'master': conn_spec,
498
+ 'replication_user': 'replication'
499
+ }
500
+ )
501
+ cls.start_cluster(
502
+ cls.standby_cluster,
503
+ server_settings={
504
+ 'hot_standby': True
505
+ }
506
+ )
507
+
508
+ finally:
509
+ if con is not None:
510
+ cls.loop.run_until_complete(con.close())
511
+
512
+ @classmethod
513
+ def get_cluster_connection_spec(cls, cluster, kwargs={}):
514
+ conn_spec = cluster.get_connection_spec()
515
+ if kwargs.get('dsn'):
516
+ conn_spec.pop('host')
517
+ conn_spec.update(kwargs)
518
+ if not os.environ.get('PGHOST') and not kwargs.get('dsn'):
519
+ if 'database' not in conn_spec:
520
+ conn_spec['database'] = 'postgres'
521
+ if 'user' not in conn_spec:
522
+ conn_spec['user'] = 'postgres'
523
+ return conn_spec
524
+
525
+ @classmethod
526
+ def get_connection_spec(cls, kwargs={}):
527
+ primary_spec = cls.get_cluster_connection_spec(
528
+ cls.master_cluster, kwargs
529
+ )
530
+ standby_spec = cls.get_cluster_connection_spec(
531
+ cls.standby_cluster, kwargs
532
+ )
533
+ return {
534
+ 'host': [primary_spec['host'], standby_spec['host']],
535
+ 'port': [primary_spec['port'], standby_spec['port']],
536
+ 'database': primary_spec['database'],
537
+ 'user': primary_spec['user'],
538
+ **kwargs
539
+ }
540
+
541
+ @classmethod
542
+ def connect_primary(cls, **kwargs):
543
+ conn_spec = cls.get_cluster_connection_spec(cls.master_cluster, kwargs)
544
+ return pg_connection.connect(**conn_spec, loop=cls.loop)
545
+
546
+ @classmethod
547
+ def connect_standby(cls, **kwargs):
548
+ conn_spec = cls.get_cluster_connection_spec(
549
+ cls.standby_cluster,
550
+ kwargs
551
+ )
552
+ return pg_connection.connect(**conn_spec, loop=cls.loop)