test_datasource_lifecycle_api.py 7.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307
  1. from flask import Flask
  2. from app.core.data_source.models import (
  3. DataSourceCredential,
  4. DataSourceDefinition,
  5. SealedCredential,
  6. )
  7. UID = "01900000-0000-7000-8000-000000000010"
  8. class FakeSession:
  9. def __init__(self):
  10. self.commits = 0
  11. self.rollbacks = 0
  12. def commit(self):
  13. self.commits += 1
  14. def rollback(self):
  15. self.rollbacks += 1
  16. class FakeDefinitions:
  17. def __init__(self, existing=None, fail_save=False):
  18. self.existing = existing
  19. self.saved = None
  20. self.deleted = None
  21. self.fail_save = fail_save
  22. def get(self, uid):
  23. return self.existing if uid == UID else None
  24. def list(self, _filters):
  25. return [self.existing] if self.existing else []
  26. def save(self, definition):
  27. if self.fail_save:
  28. raise RuntimeError("neo4j unavailable")
  29. self.saved = definition
  30. self.existing = definition
  31. return definition
  32. def delete(self, uid):
  33. self.deleted = uid
  34. return True
  35. class FakeCredentials:
  36. def __init__(self):
  37. self.created = []
  38. self.compensations = []
  39. self.revoked = []
  40. def get_active(self, session, uid, version=None):
  41. del session, uid, version
  42. return DataSourceCredential("existing-reader", "existing-secret")
  43. def create_version(
  44. self,
  45. session,
  46. *,
  47. data_source_uid,
  48. credential,
  49. actor_uid,
  50. ):
  51. del session, actor_uid
  52. version = len(self.created) + 2
  53. self.created.append(credential)
  54. return SealedCredential(
  55. id=UID,
  56. data_source_uid=data_source_uid,
  57. credential_version=version,
  58. encrypted_payload=b"cipher",
  59. nonce=b"n" * 12,
  60. key_version="v1",
  61. )
  62. def compensate_failed_activation(
  63. self,
  64. session,
  65. *,
  66. data_source_uid,
  67. failed_version,
  68. restore_version,
  69. actor_uid,
  70. ):
  71. del session, actor_uid
  72. self.compensations.append(
  73. (data_source_uid, failed_version, restore_version)
  74. )
  75. def revoke_all(self, session, *, data_source_uid, actor_uid):
  76. del session, actor_uid
  77. self.revoked.append(data_source_uid)
  78. return 1
  79. class FakeAdapter:
  80. def __init__(self):
  81. self.tests = []
  82. def validate_options(self, options):
  83. return dict(options)
  84. def test_connection(
  85. self,
  86. definition,
  87. credential,
  88. *,
  89. query_timeout,
  90. ):
  91. self.tests.append((definition, credential, query_timeout))
  92. class FakeManager:
  93. def __init__(self):
  94. self.invalidations = []
  95. def invalidate(self, uid, reason):
  96. self.invalidations.append((uid, reason))
  97. def existing_definition():
  98. return DataSourceDefinition(
  99. uid=UID,
  100. name_en="warehouse",
  101. database_type="postgresql",
  102. host="old-db",
  103. port=5432,
  104. database="analytics",
  105. credential_ref=UID,
  106. credential_version=1,
  107. )
  108. def make_service(existing=None, fail_save=False):
  109. from app.core.data_source.service import DataSourceService
  110. session = FakeSession()
  111. definitions = FakeDefinitions(existing=existing, fail_save=fail_save)
  112. credentials = FakeCredentials()
  113. adapter = FakeAdapter()
  114. manager = FakeManager()
  115. events = []
  116. service = DataSourceService(
  117. definitions=definitions,
  118. credentials=credentials,
  119. platform_session=lambda: session,
  120. adapter_resolver=lambda _database_type: adapter,
  121. connection_manager=manager,
  122. settings_resolver=lambda _overrides: {"query_timeout": 30},
  123. outbox_enqueuer=lambda *_args, **kwargs: events.append(kwargs),
  124. )
  125. return (
  126. service,
  127. session,
  128. definitions,
  129. credentials,
  130. adapter,
  131. manager,
  132. events,
  133. )
  134. def test_create_stores_secret_only_in_credential_store():
  135. (
  136. service,
  137. session,
  138. definitions,
  139. credentials,
  140. adapter,
  141. manager,
  142. events,
  143. ) = make_service()
  144. result, created = service.save(
  145. {
  146. "name_en": "warehouse",
  147. "name_zh": "数仓",
  148. "type": "postgresql",
  149. "host": "source-postgres",
  150. "port": 5432,
  151. "database": "analytics",
  152. "schema": "public",
  153. "username": "reader",
  154. "password": "source-pass",
  155. },
  156. actor_uid=None,
  157. )
  158. assert created is True
  159. assert result.uid
  160. assert definitions.saved.credential_version == 2
  161. assert credentials.created[0].password == "source-pass"
  162. assert "password" not in repr(definitions.saved)
  163. assert adapter.tests[0][2] == 30
  164. assert manager.invalidations == []
  165. assert session.commits == 1
  166. assert events[0]["event_type"] == "datasource.credential_version_created"
  167. def test_update_reuses_omitted_credentials_and_invalidates_old_pool():
  168. (
  169. service,
  170. _session,
  171. _definitions,
  172. credentials,
  173. _adapter,
  174. manager,
  175. _events,
  176. ) = make_service(existing=existing_definition())
  177. result, created = service.save(
  178. {
  179. "uid": UID,
  180. "name_en": "warehouse",
  181. "type": "postgresql",
  182. "host": "new-db",
  183. "port": 5432,
  184. "database": "analytics",
  185. },
  186. actor_uid=None,
  187. )
  188. assert created is False
  189. assert result.credential_version == 2
  190. assert credentials.created[0] == DataSourceCredential(
  191. "existing-reader",
  192. "existing-secret",
  193. )
  194. assert manager.invalidations == [(UID, "configuration_changed")]
  195. def test_neo4j_failure_restores_prior_active_credential():
  196. (
  197. service,
  198. session,
  199. _definitions,
  200. credentials,
  201. _adapter,
  202. _manager,
  203. _events,
  204. ) = make_service(existing=existing_definition(), fail_save=True)
  205. try:
  206. service.save(
  207. {
  208. "uid": UID,
  209. "name_en": "warehouse",
  210. "type": "postgresql",
  211. "host": "new-db",
  212. "port": 5432,
  213. "database": "analytics",
  214. },
  215. actor_uid=None,
  216. )
  217. except Exception:
  218. pass
  219. else: # pragma: no cover - failure is required
  220. raise AssertionError("save should fail")
  221. assert credentials.compensations == [(UID, 2, 1)]
  222. assert session.commits == 2
  223. def test_routes_use_service_and_return_real_http_status(monkeypatch):
  224. from app.api.data_source import routes
  225. class FakeService:
  226. def save(self, payload, actor_uid):
  227. del payload, actor_uid
  228. return existing_definition(), True
  229. def serialize(self, definition):
  230. return {
  231. "uid": definition.uid,
  232. "name_en": definition.name_en,
  233. "credential_configured": True,
  234. }
  235. monkeypatch.setattr(
  236. routes,
  237. "get_data_source_service",
  238. lambda: FakeService(),
  239. )
  240. app = Flask(__name__)
  241. app.register_blueprint(routes.bp, url_prefix="/api/datasource")
  242. response = app.test_client().post(
  243. "/api/datasource/save",
  244. json={"name_en": "warehouse"},
  245. )
  246. assert response.status_code == 201
  247. assert response.get_json()["data"]["uid"] == UID
  248. def test_datasource_routes_do_not_construct_engines_directly():
  249. source = open(
  250. "app/api/data_source/routes.py",
  251. encoding="utf-8",
  252. ).read()
  253. assert "create_engine" not in source
  254. assert "URL.create" not in source