runtime.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257
  1. """Explicit and lazily constructed external data-source runtimes."""
  2. import threading
  3. from dataclasses import dataclass, field
  4. from typing import Mapping
  5. from neo4j import GraphDatabase
  6. from sqlalchemy import create_engine
  7. _lock = threading.Lock()
  8. _manager = None
  9. _factory = None
  10. @dataclass(frozen=True)
  11. class DataSourceRuntimeConfig:
  12. platform_database_url: str = field(repr=False)
  13. neo4j_uri: str
  14. neo4j_user: str
  15. neo4j_password: str = field(repr=False)
  16. credential_master_key: str = field(repr=False)
  17. credential_key_version: str
  18. certificate_dir: str = "/etc/dataops-platform/datasource-certs"
  19. pool_size: int = 1
  20. max_overflow: int = 1
  21. pool_timeout: int = 10
  22. pool_recycle: int = 1800
  23. idle_ttl: int = 900
  24. max_idle_pools: int = 4
  25. query_timeout: int = 30
  26. worker_count: int = 2
  27. connection_budget: int = 32
  28. @classmethod
  29. def from_mapping(cls, values: Mapping[str, object]):
  30. config = cls(**dict(values))
  31. if config.pool_size < 1 or config.max_overflow < 0:
  32. raise ValueError("data source pool size is invalid")
  33. if config.max_idle_pools < 1 or config.worker_count < 1:
  34. raise ValueError("data source worker configuration is invalid")
  35. if config.max_business_connections > config.connection_budget:
  36. raise ValueError("runner data source connection budget is exceeded")
  37. return config
  38. @property
  39. def max_business_connections(self):
  40. return (
  41. self.worker_count
  42. * self.max_idle_pools
  43. * (self.pool_size + self.max_overflow)
  44. )
  45. def pool_settings(self):
  46. return {
  47. "pool_size": self.pool_size,
  48. "max_overflow": self.max_overflow,
  49. "pool_timeout": self.pool_timeout,
  50. "pool_recycle": self.pool_recycle,
  51. "idle_ttl": self.idle_ttl,
  52. "max_idle_pools": self.max_idle_pools,
  53. "query_timeout": self.query_timeout,
  54. "drain_timeout": 30,
  55. }
  56. class _Neo4jDefinitionGateway:
  57. def __init__(self, driver=None):
  58. self._driver = driver
  59. def _call(self, method, *args, **kwargs):
  60. from app.core.data_source.definitions import (
  61. DataSourceDefinitionRepository,
  62. )
  63. if self._driver is None:
  64. from app.services.neo4j_driver import neo4j_driver
  65. session_context = neo4j_driver.get_session()
  66. else:
  67. session_context = self._driver.session()
  68. with session_context as session:
  69. repository = DataSourceDefinitionRepository(session)
  70. return getattr(repository, method)(*args, **kwargs)
  71. def get(self, *args, **kwargs):
  72. return self._call("get", *args, **kwargs)
  73. def list(self, *args, **kwargs):
  74. return self._call("list", *args, **kwargs)
  75. def save(self, *args, **kwargs):
  76. return self._call("save", *args, **kwargs)
  77. def delete(self, *args, **kwargs):
  78. return self._call("delete", *args, **kwargs)
  79. def credential_references(self, *args, **kwargs):
  80. return self._call("credential_references", *args, **kwargs)
  81. def remove_legacy_secret_fields(self, *args, **kwargs):
  82. return self._call(
  83. "remove_legacy_secret_fields",
  84. *args,
  85. **kwargs,
  86. )
  87. class _StandaloneCredentialGateway:
  88. def __init__(self, engine, repository):
  89. self._engine = engine
  90. self._repository = repository
  91. def get_active(self, _session, uid, version):
  92. with self._engine.connect() as connection:
  93. return self._repository.get_active(connection, uid, version)
  94. @dataclass
  95. class StandaloneDataSourceRuntime:
  96. manager: object
  97. platform_engine: object
  98. graph_driver: object
  99. def close(self):
  100. self.manager.close()
  101. self.graph_driver.close()
  102. self.platform_engine.dispose()
  103. def build_standalone_data_source_runtime(config):
  104. """Build a Runner-owned runtime without Flask globals or request state."""
  105. from app.core.data_source.adapters import adapter_for
  106. from app.core.data_source.credentials import (
  107. CredentialCodec,
  108. DataSourceCredentialRepository,
  109. )
  110. from app.core.data_source.manager import DataSourceConnectionManager
  111. from app.core.data_source.pool_registry import PoolRegistry
  112. platform_engine = create_engine(
  113. config.platform_database_url,
  114. pool_size=1,
  115. max_overflow=1,
  116. pool_pre_ping=True,
  117. pool_recycle=300,
  118. )
  119. graph_driver = GraphDatabase.driver(
  120. config.neo4j_uri,
  121. auth=(config.neo4j_user, config.neo4j_password),
  122. )
  123. codec = CredentialCodec.from_base64(
  124. config.credential_master_key,
  125. config.credential_key_version,
  126. )
  127. base_settings = config.pool_settings()
  128. def settings_resolver(overrides):
  129. values = dict(base_settings)
  130. for key in ("pool_size", "max_overflow"):
  131. if key in overrides:
  132. values[key] = int(overrides[key])
  133. if values["pool_size"] + values["max_overflow"] > (
  134. config.pool_size + config.max_overflow
  135. ):
  136. raise ValueError("data source override exceeds runner connection budget")
  137. return values
  138. manager = DataSourceConnectionManager(
  139. definitions=_Neo4jDefinitionGateway(graph_driver),
  140. credentials=_StandaloneCredentialGateway(
  141. platform_engine,
  142. DataSourceCredentialRepository(codec),
  143. ),
  144. platform_session=lambda: None,
  145. adapter_resolver=lambda database_type: adapter_for(
  146. database_type,
  147. certificate_dir=config.certificate_dir,
  148. ),
  149. registry=PoolRegistry(base_settings),
  150. settings_resolver=settings_resolver,
  151. )
  152. return StandaloneDataSourceRuntime(
  153. manager=manager,
  154. platform_engine=platform_engine,
  155. graph_driver=graph_driver,
  156. )
  157. def _build_default_manager():
  158. from flask import current_app
  159. from app import db
  160. from app.config.config import datasource_pool_settings
  161. from app.core.data_source.adapters import adapter_for
  162. from app.core.data_source.credentials import (
  163. CredentialCodec,
  164. DataSourceCredentialRepository,
  165. )
  166. from app.core.data_source.manager import DataSourceConnectionManager
  167. from app.core.data_source.pool_registry import PoolRegistry
  168. codec = CredentialCodec.from_base64(
  169. current_app.config.get("DATASOURCE_CREDENTIAL_MASTER_KEY", ""),
  170. current_app.config.get("DATASOURCE_CREDENTIAL_KEY_VERSION", "v1"),
  171. )
  172. base_settings = datasource_pool_settings()
  173. registry_settings = {
  174. **base_settings,
  175. "drain_timeout": 30,
  176. }
  177. certificate_dir = current_app.config.get("DATASOURCE_CERT_DIR")
  178. return DataSourceConnectionManager(
  179. definitions=_Neo4jDefinitionGateway(),
  180. credentials=DataSourceCredentialRepository(codec),
  181. platform_session=lambda: db.session,
  182. adapter_resolver=lambda database_type: adapter_for(
  183. database_type,
  184. certificate_dir=certificate_dir,
  185. ),
  186. registry=PoolRegistry(registry_settings),
  187. settings_resolver=datasource_pool_settings,
  188. )
  189. def configure_data_source_runtime(factory):
  190. """Replace the process-local factory and clear any current manager."""
  191. global _factory, _manager
  192. with _lock:
  193. if _manager is not None:
  194. _manager.close()
  195. _manager = None
  196. _factory = factory
  197. def get_data_source_manager():
  198. global _manager
  199. if _manager is not None:
  200. return _manager
  201. with _lock:
  202. if _manager is None:
  203. _manager = (_factory or _build_default_manager)()
  204. return _manager
  205. def peek_data_source_manager():
  206. """Return the current Worker manager without constructing one."""
  207. return _manager
  208. def close_data_source_runtime():
  209. global _manager
  210. with _lock:
  211. manager = _manager
  212. _manager = None
  213. if manager is not None:
  214. manager.close()