asyncpg 0.30.0__cp313-cp313-win_amd64.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 (87) hide show
  1. asyncpg/__init__.py +24 -0
  2. asyncpg/_asyncio_compat.py +94 -0
  3. asyncpg/_testbase/__init__.py +543 -0
  4. asyncpg/_testbase/fuzzer.py +306 -0
  5. asyncpg/_version.py +17 -0
  6. asyncpg/cluster.py +729 -0
  7. asyncpg/compat.py +88 -0
  8. asyncpg/connect_utils.py +1139 -0
  9. asyncpg/connection.py +2749 -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 +298 -0
  15. asyncpg/pgproto/__init__.pxd +5 -0
  16. asyncpg/pgproto/__init__.py +5 -0
  17. asyncpg/pgproto/buffer.pxd +136 -0
  18. asyncpg/pgproto/buffer.pyx +817 -0
  19. asyncpg/pgproto/codecs/__init__.pxd +157 -0
  20. asyncpg/pgproto/codecs/bits.pyx +47 -0
  21. asyncpg/pgproto/codecs/bytea.pyx +34 -0
  22. asyncpg/pgproto/codecs/context.pyx +26 -0
  23. asyncpg/pgproto/codecs/datetime.pyx +423 -0
  24. asyncpg/pgproto/codecs/float.pyx +34 -0
  25. asyncpg/pgproto/codecs/geometry.pyx +164 -0
  26. asyncpg/pgproto/codecs/hstore.pyx +73 -0
  27. asyncpg/pgproto/codecs/int.pyx +144 -0
  28. asyncpg/pgproto/codecs/json.pyx +57 -0
  29. asyncpg/pgproto/codecs/jsonpath.pyx +29 -0
  30. asyncpg/pgproto/codecs/misc.pyx +16 -0
  31. asyncpg/pgproto/codecs/network.pyx +139 -0
  32. asyncpg/pgproto/codecs/numeric.pyx +356 -0
  33. asyncpg/pgproto/codecs/pg_snapshot.pyx +63 -0
  34. asyncpg/pgproto/codecs/text.pyx +48 -0
  35. asyncpg/pgproto/codecs/tid.pyx +51 -0
  36. asyncpg/pgproto/codecs/uuid.pyx +27 -0
  37. asyncpg/pgproto/consts.pxi +12 -0
  38. asyncpg/pgproto/cpythonx.pxd +23 -0
  39. asyncpg/pgproto/debug.pxd +10 -0
  40. asyncpg/pgproto/frb.pxd +48 -0
  41. asyncpg/pgproto/frb.pyx +12 -0
  42. asyncpg/pgproto/hton.pxd +24 -0
  43. asyncpg/pgproto/pgproto.cp313-win_amd64.pyd +0 -0
  44. asyncpg/pgproto/pgproto.pxd +19 -0
  45. asyncpg/pgproto/pgproto.pyi +13 -0
  46. asyncpg/pgproto/pgproto.pyx +49 -0
  47. asyncpg/pgproto/tohex.pxd +10 -0
  48. asyncpg/pgproto/types.py +423 -0
  49. asyncpg/pgproto/uuid.pyx +353 -0
  50. asyncpg/pool.py +1211 -0
  51. asyncpg/prepared_stmt.py +285 -0
  52. asyncpg/protocol/__init__.py +11 -0
  53. asyncpg/protocol/codecs/__init__.py +0 -0
  54. asyncpg/protocol/codecs/array.pyx +875 -0
  55. asyncpg/protocol/codecs/base.pxd +187 -0
  56. asyncpg/protocol/codecs/base.pyx +895 -0
  57. asyncpg/protocol/codecs/pgproto.pyx +484 -0
  58. asyncpg/protocol/codecs/range.pyx +207 -0
  59. asyncpg/protocol/codecs/record.pyx +71 -0
  60. asyncpg/protocol/codecs/textutils.pyx +99 -0
  61. asyncpg/protocol/consts.pxi +12 -0
  62. asyncpg/protocol/coreproto.pxd +192 -0
  63. asyncpg/protocol/coreproto.pyx +1233 -0
  64. asyncpg/protocol/cpythonx.pxd +19 -0
  65. asyncpg/protocol/encodings.pyx +63 -0
  66. asyncpg/protocol/pgtypes.pxi +266 -0
  67. asyncpg/protocol/prepared_stmt.pxd +39 -0
  68. asyncpg/protocol/prepared_stmt.pyx +395 -0
  69. asyncpg/protocol/protocol.cp313-win_amd64.pyd +0 -0
  70. asyncpg/protocol/protocol.pxd +77 -0
  71. asyncpg/protocol/protocol.pyi +300 -0
  72. asyncpg/protocol/protocol.pyx +1065 -0
  73. asyncpg/protocol/record/__init__.pxd +19 -0
  74. asyncpg/protocol/scram.pxd +31 -0
  75. asyncpg/protocol/scram.pyx +341 -0
  76. asyncpg/protocol/settings.pxd +30 -0
  77. asyncpg/protocol/settings.pyx +106 -0
  78. asyncpg/serverversion.py +70 -0
  79. asyncpg/transaction.py +246 -0
  80. asyncpg/types.py +223 -0
  81. asyncpg/utils.py +52 -0
  82. asyncpg-0.30.0.dist-info/AUTHORS +6 -0
  83. asyncpg-0.30.0.dist-info/LICENSE +204 -0
  84. asyncpg-0.30.0.dist-info/METADATA +142 -0
  85. asyncpg-0.30.0.dist-info/RECORD +87 -0
  86. asyncpg-0.30.0.dist-info/WHEEL +5 -0
  87. asyncpg-0.30.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 supresses 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,543 @@
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
+ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
105
+
106
+ loop = asyncio.new_event_loop()
107
+ asyncio.set_event_loop(None)
108
+ cls.loop = loop
109
+
110
+ @classmethod
111
+ def tearDownClass(cls):
112
+ cls.loop.close()
113
+ asyncio.set_event_loop(None)
114
+
115
+ def setUp(self):
116
+ self.loop.set_exception_handler(self.loop_exception_handler)
117
+ self.__unhandled_exceptions = []
118
+
119
+ def tearDown(self):
120
+ excs = []
121
+ for exc in self.__unhandled_exceptions:
122
+ if isinstance(exc, ConnectionResetError):
123
+ texc = traceback.TracebackException.from_exception(
124
+ exc, lookup_lines=False)
125
+ if texc.stack[-1].name == "_call_connection_lost":
126
+ # On Windows calling socket.shutdown may raise
127
+ # ConnectionResetError, which happens in the
128
+ # finally block of _call_connection_lost.
129
+ continue
130
+ excs.append(exc)
131
+
132
+ if excs:
133
+ formatted = []
134
+
135
+ for i, context in enumerate(excs):
136
+ formatted.append(self._format_loop_exception(context, i + 1))
137
+
138
+ self.fail(
139
+ 'unexpected exceptions in asynchronous code:\n' +
140
+ '\n'.join(formatted))
141
+
142
+ @contextlib.contextmanager
143
+ def assertRunUnder(self, delta):
144
+ st = time.monotonic()
145
+ try:
146
+ yield
147
+ finally:
148
+ elapsed = time.monotonic() - st
149
+ if elapsed > delta:
150
+ raise AssertionError(
151
+ 'running block took {:0.3f}s which is longer '
152
+ 'than the expected maximum of {:0.3f}s'.format(
153
+ elapsed, delta))
154
+
155
+ @contextlib.contextmanager
156
+ def assertLoopErrorHandlerCalled(self, msg_re: str):
157
+ contexts = []
158
+
159
+ def handler(loop, ctx):
160
+ contexts.append(ctx)
161
+
162
+ old_handler = self.loop.get_exception_handler()
163
+ self.loop.set_exception_handler(handler)
164
+ try:
165
+ yield
166
+
167
+ for ctx in contexts:
168
+ msg = ctx.get('message')
169
+ if msg and re.search(msg_re, msg):
170
+ return
171
+
172
+ raise AssertionError(
173
+ 'no message matching {!r} was logged with '
174
+ 'loop.call_exception_handler()'.format(msg_re))
175
+
176
+ finally:
177
+ self.loop.set_exception_handler(old_handler)
178
+
179
+ def loop_exception_handler(self, loop, context):
180
+ self.__unhandled_exceptions.append(context)
181
+ loop.default_exception_handler(context)
182
+
183
+ def _format_loop_exception(self, context, n):
184
+ message = context.get('message', 'Unhandled exception in event loop')
185
+ exception = context.get('exception')
186
+ if exception is not None:
187
+ exc_info = (type(exception), exception, exception.__traceback__)
188
+ else:
189
+ exc_info = None
190
+
191
+ lines = []
192
+ for key in sorted(context):
193
+ if key in {'message', 'exception'}:
194
+ continue
195
+ value = context[key]
196
+ if key == 'source_traceback':
197
+ tb = ''.join(traceback.format_list(value))
198
+ value = 'Object created at (most recent call last):\n'
199
+ value += tb.rstrip()
200
+ else:
201
+ try:
202
+ value = repr(value)
203
+ except Exception as ex:
204
+ value = ('Exception in __repr__ {!r}; '
205
+ 'value type: {!r}'.format(ex, type(value)))
206
+ lines.append('[{}]: {}\n\n'.format(key, value))
207
+
208
+ if exc_info is not None:
209
+ lines.append('[exception]:\n')
210
+ formatted_exc = textwrap.indent(
211
+ ''.join(traceback.format_exception(*exc_info)), ' ')
212
+ lines.append(formatted_exc)
213
+
214
+ details = textwrap.indent(''.join(lines), ' ')
215
+ return '{:02d}. {}:\n{}\n'.format(n, message, details)
216
+
217
+
218
+ _default_cluster = None
219
+
220
+
221
+ def _init_cluster(ClusterCls, cluster_kwargs, initdb_options=None):
222
+ cluster = ClusterCls(**cluster_kwargs)
223
+ cluster.init(**(initdb_options or {}))
224
+ cluster.trust_local_connections()
225
+ atexit.register(_shutdown_cluster, cluster)
226
+ return cluster
227
+
228
+
229
+ def _get_initdb_options(initdb_options=None):
230
+ if not initdb_options:
231
+ initdb_options = {}
232
+ else:
233
+ initdb_options = dict(initdb_options)
234
+
235
+ # Make the default superuser name stable.
236
+ if 'username' not in initdb_options:
237
+ initdb_options['username'] = 'postgres'
238
+
239
+ return initdb_options
240
+
241
+
242
+ def _init_default_cluster(initdb_options=None):
243
+ global _default_cluster
244
+
245
+ if _default_cluster is None:
246
+ pg_host = os.environ.get('PGHOST')
247
+ if pg_host:
248
+ # Using existing cluster, assuming it is initialized and running
249
+ _default_cluster = pg_cluster.RunningCluster()
250
+ else:
251
+ _default_cluster = _init_cluster(
252
+ pg_cluster.TempCluster,
253
+ cluster_kwargs={
254
+ "data_dir_suffix": ".apgtest",
255
+ },
256
+ initdb_options=_get_initdb_options(initdb_options),
257
+ )
258
+
259
+ return _default_cluster
260
+
261
+
262
+ def _shutdown_cluster(cluster):
263
+ if cluster.get_status() == 'running':
264
+ cluster.stop()
265
+ if cluster.get_status() != 'not-initialized':
266
+ cluster.destroy()
267
+
268
+
269
+ def create_pool(dsn=None, *,
270
+ min_size=10,
271
+ max_size=10,
272
+ max_queries=50000,
273
+ max_inactive_connection_lifetime=60.0,
274
+ connect=None,
275
+ setup=None,
276
+ init=None,
277
+ loop=None,
278
+ pool_class=pg_pool.Pool,
279
+ connection_class=pg_connection.Connection,
280
+ record_class=asyncpg.Record,
281
+ **connect_kwargs):
282
+ return pool_class(
283
+ dsn,
284
+ min_size=min_size,
285
+ max_size=max_size,
286
+ max_queries=max_queries,
287
+ loop=loop,
288
+ connect=connect,
289
+ setup=setup,
290
+ init=init,
291
+ max_inactive_connection_lifetime=max_inactive_connection_lifetime,
292
+ connection_class=connection_class,
293
+ record_class=record_class,
294
+ **connect_kwargs,
295
+ )
296
+
297
+
298
+ class ClusterTestCase(TestCase):
299
+ @classmethod
300
+ def get_server_settings(cls):
301
+ settings = {
302
+ 'log_connections': 'on'
303
+ }
304
+
305
+ if cls.cluster.get_pg_version() >= (11, 0):
306
+ # JITting messes up timing tests, and
307
+ # is not essential for testing.
308
+ settings['jit'] = 'off'
309
+
310
+ return settings
311
+
312
+ @classmethod
313
+ def new_cluster(cls, ClusterCls, *, cluster_kwargs={}, initdb_options={}):
314
+ cluster = _init_cluster(ClusterCls, cluster_kwargs,
315
+ _get_initdb_options(initdb_options))
316
+ cls._clusters.append(cluster)
317
+ return cluster
318
+
319
+ @classmethod
320
+ def start_cluster(cls, cluster, *, server_settings={}):
321
+ cluster.start(port='dynamic', server_settings=server_settings)
322
+
323
+ @classmethod
324
+ def setup_cluster(cls):
325
+ cls.cluster = _init_default_cluster()
326
+
327
+ if cls.cluster.get_status() != 'running':
328
+ cls.cluster.start(
329
+ port='dynamic', server_settings=cls.get_server_settings())
330
+
331
+ @classmethod
332
+ def setUpClass(cls):
333
+ super().setUpClass()
334
+ cls._clusters = []
335
+ cls.setup_cluster()
336
+
337
+ @classmethod
338
+ def tearDownClass(cls):
339
+ super().tearDownClass()
340
+ for cluster in cls._clusters:
341
+ if cluster is not _default_cluster:
342
+ cluster.stop()
343
+ cluster.destroy()
344
+ cls._clusters = []
345
+
346
+ @classmethod
347
+ def get_connection_spec(cls, kwargs={}):
348
+ conn_spec = cls.cluster.get_connection_spec()
349
+ if kwargs.get('dsn'):
350
+ conn_spec.pop('host')
351
+ conn_spec.update(kwargs)
352
+ if not os.environ.get('PGHOST') and not kwargs.get('dsn'):
353
+ if 'database' not in conn_spec:
354
+ conn_spec['database'] = 'postgres'
355
+ if 'user' not in conn_spec:
356
+ conn_spec['user'] = 'postgres'
357
+ return conn_spec
358
+
359
+ @classmethod
360
+ def connect(cls, **kwargs):
361
+ conn_spec = cls.get_connection_spec(kwargs)
362
+ return pg_connection.connect(**conn_spec, loop=cls.loop)
363
+
364
+ def setUp(self):
365
+ super().setUp()
366
+ self._pools = []
367
+
368
+ def tearDown(self):
369
+ super().tearDown()
370
+ for pool in self._pools:
371
+ pool.terminate()
372
+ self._pools = []
373
+
374
+ def create_pool(self, pool_class=pg_pool.Pool,
375
+ connection_class=pg_connection.Connection, **kwargs):
376
+ conn_spec = self.get_connection_spec(kwargs)
377
+ pool = create_pool(loop=self.loop, pool_class=pool_class,
378
+ connection_class=connection_class, **conn_spec)
379
+ self._pools.append(pool)
380
+ return pool
381
+
382
+
383
+ class ProxiedClusterTestCase(ClusterTestCase):
384
+ @classmethod
385
+ def get_server_settings(cls):
386
+ settings = dict(super().get_server_settings())
387
+ settings['listen_addresses'] = '127.0.0.1'
388
+ return settings
389
+
390
+ @classmethod
391
+ def get_proxy_settings(cls):
392
+ return {'fuzzing-mode': None}
393
+
394
+ @classmethod
395
+ def setUpClass(cls):
396
+ super().setUpClass()
397
+ conn_spec = cls.cluster.get_connection_spec()
398
+ host = conn_spec.get('host')
399
+ if not host:
400
+ host = '127.0.0.1'
401
+ elif host.startswith('/'):
402
+ host = '127.0.0.1'
403
+ cls.proxy = fuzzer.TCPFuzzingProxy(
404
+ backend_host=host,
405
+ backend_port=conn_spec['port'],
406
+ )
407
+ cls.proxy.start()
408
+
409
+ @classmethod
410
+ def tearDownClass(cls):
411
+ cls.proxy.stop()
412
+ super().tearDownClass()
413
+
414
+ @classmethod
415
+ def get_connection_spec(cls, kwargs):
416
+ conn_spec = super().get_connection_spec(kwargs)
417
+ conn_spec['host'] = cls.proxy.listening_addr
418
+ conn_spec['port'] = cls.proxy.listening_port
419
+ return conn_spec
420
+
421
+ def tearDown(self):
422
+ self.proxy.reset()
423
+ super().tearDown()
424
+
425
+
426
+ def with_connection_options(**options):
427
+ if not options:
428
+ raise ValueError('no connection options were specified')
429
+
430
+ def wrap(func):
431
+ func.__connect_options__ = options
432
+ return func
433
+
434
+ return wrap
435
+
436
+
437
+ class ConnectedTestCase(ClusterTestCase):
438
+
439
+ def setUp(self):
440
+ super().setUp()
441
+
442
+ # Extract options set up with `with_connection_options`.
443
+ test_func = getattr(self, self._testMethodName).__func__
444
+ opts = getattr(test_func, '__connect_options__', {})
445
+ self.con = self.loop.run_until_complete(self.connect(**opts))
446
+ self.server_version = self.con.get_server_version()
447
+
448
+ def tearDown(self):
449
+ try:
450
+ self.loop.run_until_complete(self.con.close())
451
+ self.con = None
452
+ finally:
453
+ super().tearDown()
454
+
455
+
456
+ class HotStandbyTestCase(ClusterTestCase):
457
+
458
+ @classmethod
459
+ def setup_cluster(cls):
460
+ cls.master_cluster = cls.new_cluster(pg_cluster.TempCluster)
461
+ cls.start_cluster(
462
+ cls.master_cluster,
463
+ server_settings={
464
+ 'max_wal_senders': 10,
465
+ 'wal_level': 'hot_standby'
466
+ }
467
+ )
468
+
469
+ con = None
470
+
471
+ try:
472
+ con = cls.loop.run_until_complete(
473
+ cls.master_cluster.connect(
474
+ database='postgres', user='postgres', loop=cls.loop))
475
+
476
+ cls.loop.run_until_complete(
477
+ con.execute('''
478
+ CREATE ROLE replication WITH LOGIN REPLICATION
479
+ '''))
480
+
481
+ cls.master_cluster.trust_local_replication_by('replication')
482
+
483
+ conn_spec = cls.master_cluster.get_connection_spec()
484
+
485
+ cls.standby_cluster = cls.new_cluster(
486
+ pg_cluster.HotStandbyCluster,
487
+ cluster_kwargs={
488
+ 'master': conn_spec,
489
+ 'replication_user': 'replication'
490
+ }
491
+ )
492
+ cls.start_cluster(
493
+ cls.standby_cluster,
494
+ server_settings={
495
+ 'hot_standby': True
496
+ }
497
+ )
498
+
499
+ finally:
500
+ if con is not None:
501
+ cls.loop.run_until_complete(con.close())
502
+
503
+ @classmethod
504
+ def get_cluster_connection_spec(cls, cluster, kwargs={}):
505
+ conn_spec = cluster.get_connection_spec()
506
+ if kwargs.get('dsn'):
507
+ conn_spec.pop('host')
508
+ conn_spec.update(kwargs)
509
+ if not os.environ.get('PGHOST') and not kwargs.get('dsn'):
510
+ if 'database' not in conn_spec:
511
+ conn_spec['database'] = 'postgres'
512
+ if 'user' not in conn_spec:
513
+ conn_spec['user'] = 'postgres'
514
+ return conn_spec
515
+
516
+ @classmethod
517
+ def get_connection_spec(cls, kwargs={}):
518
+ primary_spec = cls.get_cluster_connection_spec(
519
+ cls.master_cluster, kwargs
520
+ )
521
+ standby_spec = cls.get_cluster_connection_spec(
522
+ cls.standby_cluster, kwargs
523
+ )
524
+ return {
525
+ 'host': [primary_spec['host'], standby_spec['host']],
526
+ 'port': [primary_spec['port'], standby_spec['port']],
527
+ 'database': primary_spec['database'],
528
+ 'user': primary_spec['user'],
529
+ **kwargs
530
+ }
531
+
532
+ @classmethod
533
+ def connect_primary(cls, **kwargs):
534
+ conn_spec = cls.get_cluster_connection_spec(cls.master_cluster, kwargs)
535
+ return pg_connection.connect(**conn_spec, loop=cls.loop)
536
+
537
+ @classmethod
538
+ def connect_standby(cls, **kwargs):
539
+ conn_spec = cls.get_cluster_connection_spec(
540
+ cls.standby_cluster,
541
+ kwargs
542
+ )
543
+ return pg_connection.connect(**conn_spec, loop=cls.loop)