test_datasource_lifecycle_api.py 8.2 KB

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