manager.py 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191
  1. """Public context-managed access to external business data sources."""
  2. import inspect
  3. import time
  4. from contextlib import contextmanager, suppress
  5. from sqlalchemy.exc import DBAPIError, OperationalError, TimeoutError
  6. from app.core.data_source.adapters.base import PURPOSES
  7. from app.core.data_source.errors import (
  8. DataSourceConfigurationInvalid,
  9. DataSourceConnectionFailed,
  10. DataSourceNotFound,
  11. DataSourcePoolTimeout,
  12. DataSourceQueryTimeout,
  13. DataSourceReadOnlyViolation,
  14. DataSourceWriteOutcomeUnknown,
  15. )
  16. from app.core.data_source.models import PoolKey
  17. def _driver_error_code(error):
  18. original = getattr(error, "orig", error)
  19. code = getattr(original, "pgcode", None)
  20. if code:
  21. return str(code)
  22. args = getattr(original, "args", ())
  23. return str(args[0]) if args else ""
  24. def _translate_query_error(error):
  25. code = _driver_error_code(error)
  26. if code in {"57014", "3024", "1317"}:
  27. return DataSourceQueryTimeout()
  28. if code in {"25006", "1792"}:
  29. return DataSourceReadOnlyViolation()
  30. return error
  31. class ManagedConnection:
  32. def __init__(self, connection, *, key, registry, clock=None):
  33. self._connection = connection
  34. self._key = key
  35. self._registry = registry
  36. self._clock = clock or time.monotonic
  37. def execute(self, statement, *args, **kwargs):
  38. started = self._clock()
  39. try:
  40. return self._connection.execute(statement, *args, **kwargs)
  41. except DBAPIError as exc:
  42. translated = _translate_query_error(exc)
  43. if translated is exc:
  44. raise
  45. raise translated from exc
  46. finally:
  47. self._registry.record_query_duration(
  48. self._key,
  49. (self._clock() - started) * 1000,
  50. )
  51. def __getattr__(self, name):
  52. return getattr(self._connection, name)
  53. class DataSourceConnectionManager:
  54. def __init__(
  55. self,
  56. *,
  57. definitions,
  58. credentials,
  59. platform_session,
  60. adapter_resolver,
  61. registry,
  62. settings_resolver,
  63. clock=None,
  64. ):
  65. self.definitions = definitions
  66. self.credentials = credentials
  67. self.platform_session = platform_session
  68. self.adapter_resolver = adapter_resolver
  69. self.registry = registry
  70. self.settings_resolver = settings_resolver
  71. self._clock = clock or time.monotonic
  72. @contextmanager
  73. def connect(
  74. self,
  75. data_source_uid,
  76. purpose,
  77. *,
  78. environment="production",
  79. allow_insecure_development=False,
  80. ):
  81. if purpose not in PURPOSES or purpose == "connection_test":
  82. raise DataSourceConfigurationInvalid(
  83. "unsupported pooled connection purpose"
  84. )
  85. if environment not in {"development", "staging", "production"}:
  86. raise DataSourceConfigurationInvalid(
  87. "trusted connector environment is invalid"
  88. )
  89. definition = self.definitions.get(data_source_uid)
  90. if not definition:
  91. raise DataSourceNotFound()
  92. credential = self.credentials.get_active(
  93. self.platform_session(),
  94. data_source_uid,
  95. definition.credential_version,
  96. )
  97. adapter = self.adapter_resolver(definition.database_type)
  98. settings = self.settings_resolver(definition.pool_overrides())
  99. key = PoolKey(
  100. data_source_uid=str(data_source_uid),
  101. credential_version=int(definition.credential_version),
  102. config_fingerprint=definition.connection_fingerprint(),
  103. )
  104. create_parameters = inspect.signature(adapter.create_pooled_engine).parameters
  105. supports_security_context = (
  106. "trusted_environment" in create_parameters
  107. or any(
  108. item.kind is inspect.Parameter.VAR_KEYWORD
  109. for item in create_parameters.values()
  110. )
  111. )
  112. security_context = (
  113. {
  114. "trusted_environment": environment,
  115. "allow_insecure_development": bool(allow_insecure_development),
  116. }
  117. if supports_security_context
  118. else {}
  119. )
  120. with self.registry.lease(
  121. key,
  122. lambda: adapter.create_pooled_engine(
  123. definition,
  124. credential,
  125. settings,
  126. **security_context,
  127. ),
  128. ) as engine:
  129. breaker = self.registry.breaker_for(key)
  130. try:
  131. with breaker.probe_permission():
  132. started = self._clock()
  133. connection = engine.connect()
  134. self.registry.add_checkout_wait(
  135. key,
  136. (self._clock() - started) * 1000,
  137. )
  138. self.registry.record_connection_success(key)
  139. except TimeoutError as exc:
  140. self.registry.record_pool_timeout(key)
  141. self.registry.record_connection_failure(key)
  142. raise DataSourcePoolTimeout() from exc
  143. except (OperationalError, DBAPIError) as exc:
  144. self.registry.record_connection_failure(key)
  145. raise DataSourceConnectionFailed() from exc
  146. with connection:
  147. transaction = connection.begin()
  148. try:
  149. adapter.configure_transaction(connection, purpose)
  150. yield ManagedConnection(
  151. connection,
  152. key=key,
  153. registry=self.registry,
  154. clock=self._clock,
  155. )
  156. if purpose == "dataflow_write":
  157. try:
  158. transaction.commit()
  159. except (OperationalError, DBAPIError) as exc:
  160. raise DataSourceWriteOutcomeUnknown() from exc
  161. else:
  162. transaction.rollback()
  163. except Exception:
  164. with suppress(Exception):
  165. transaction.rollback()
  166. raise
  167. def invalidate(self, data_source_uid, reason):
  168. self.registry.invalidate(data_source_uid, reason)
  169. def snapshot(self):
  170. return self.registry.snapshot()
  171. def close(self):
  172. self.registry.close_all()