service.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392
  1. """Secure cross-store lifecycle orchestration for external data sources."""
  2. from dataclasses import replace
  3. from app.core.common.identifiers import new_governance_uid
  4. from app.core.data_source.definitions import DataSourceFilters
  5. from app.core.data_source.errors import (
  6. DataSourceConfigurationInvalid,
  7. DataSourceError,
  8. DataSourceNotFound,
  9. )
  10. from app.core.data_source.models import (
  11. DataSourceCredential,
  12. DataSourceDefinition,
  13. )
  14. TYPE_ALIASES = {
  15. "postgres": "postgresql",
  16. "postgresql": "postgresql",
  17. "mysql": "mysql",
  18. "oracle": "oracle",
  19. "sqlserver": "sqlserver",
  20. "mssql": "sqlserver",
  21. }
  22. POOL_INVALIDATION_REASONS = {
  23. "admin_reset",
  24. "configuration_changed",
  25. "credential_rotated",
  26. }
  27. class DataSourceService:
  28. def __init__(
  29. self,
  30. *,
  31. definitions,
  32. credentials,
  33. platform_session,
  34. adapter_resolver,
  35. connection_manager,
  36. settings_resolver,
  37. outbox_enqueuer,
  38. ):
  39. self.definitions = definitions
  40. self.credentials = credentials
  41. self.platform_session = platform_session
  42. self.adapter_resolver = adapter_resolver
  43. self.connection_manager = connection_manager
  44. self.settings_resolver = settings_resolver
  45. self.outbox_enqueuer = outbox_enqueuer
  46. @staticmethod
  47. def _database_type(payload):
  48. requested = str(payload.get("type") or "").strip().lower()
  49. database_type = TYPE_ALIASES.get(requested)
  50. if database_type is None:
  51. raise DataSourceConfigurationInvalid(
  52. "only registered PostgreSQL, MySQL, Oracle and SQL Server data sources are supported"
  53. )
  54. return database_type
  55. @staticmethod
  56. def _required_text(payload, name):
  57. value = str(payload.get(name) or "").strip()
  58. if not value:
  59. raise DataSourceConfigurationInvalid(f"{name} is required")
  60. return value
  61. @classmethod
  62. def _definition(
  63. cls,
  64. payload,
  65. *,
  66. uid,
  67. credential_version,
  68. existing=None,
  69. ):
  70. database_type = cls._database_type(payload)
  71. try:
  72. port = int(payload.get("port"))
  73. except (TypeError, ValueError) as exc:
  74. raise DataSourceConfigurationInvalid("port is invalid") from exc
  75. if port < 1 or port > 65535:
  76. raise DataSourceConfigurationInvalid("port is invalid")
  77. tls_options = payload.get("tls_options") or {}
  78. if not isinstance(tls_options, dict):
  79. raise DataSourceConfigurationInvalid("tls_options is invalid")
  80. return DataSourceDefinition(
  81. uid=uid,
  82. name_en=cls._required_text(payload, "name_en"),
  83. name_zh=payload.get("name_zh"),
  84. database_type=database_type,
  85. host=cls._required_text(payload, "host"),
  86. port=port,
  87. database=cls._required_text(payload, "database"),
  88. schema=payload.get("schema"),
  89. credential_ref=uid,
  90. credential_version=credential_version,
  91. pool_size=payload.get("pool_size"),
  92. max_overflow=payload.get("max_overflow"),
  93. tls_options=tls_options,
  94. status=bool(payload.get("status", True)),
  95. description=payload.get("desc"),
  96. extra_properties=(
  97. dict(existing.extra_properties) if existing else {}
  98. ),
  99. )
  100. def _effective_credential(self, payload, existing):
  101. username = payload.get("username")
  102. password = payload.get("password")
  103. current = None
  104. if existing is not None and (not username or not password):
  105. current = self.credentials.get_active(
  106. self.platform_session(),
  107. existing.uid,
  108. existing.credential_version,
  109. )
  110. username = str(username or (current.username if current else "")).strip()
  111. password = str(password or (current.password if current else ""))
  112. if not username or not password:
  113. raise DataSourceConfigurationInvalid(
  114. "username and password are required"
  115. )
  116. options = payload.get("credential_options") or (
  117. dict(current.options) if current else {}
  118. )
  119. if not isinstance(options, dict):
  120. raise DataSourceConfigurationInvalid(
  121. "credential options are invalid"
  122. )
  123. return DataSourceCredential(username, password, options)
  124. def save(self, payload, *, actor_uid):
  125. if not isinstance(payload, dict):
  126. raise DataSourceConfigurationInvalid()
  127. requested_uid = str(payload.get("uid") or "").strip() or None
  128. existing = (
  129. self.definitions.get(requested_uid)
  130. if requested_uid is not None
  131. else None
  132. )
  133. if requested_uid is not None and existing is None:
  134. raise DataSourceNotFound()
  135. uid = requested_uid or new_governance_uid()
  136. credential = self._effective_credential(payload, existing)
  137. candidate = self._definition(
  138. payload,
  139. uid=uid,
  140. credential_version=(
  141. existing.credential_version if existing else 1
  142. ),
  143. existing=existing,
  144. )
  145. adapter = self.adapter_resolver(candidate.database_type)
  146. adapter.validate_options(candidate.tls_options)
  147. settings = self.settings_resolver(candidate.pool_overrides())
  148. adapter.test_connection(
  149. candidate,
  150. credential,
  151. query_timeout=settings["query_timeout"],
  152. )
  153. session = self.platform_session()
  154. previous_version = (
  155. existing.credential_version if existing is not None else None
  156. )
  157. try:
  158. sealed = self.credentials.create_version(
  159. session,
  160. data_source_uid=uid,
  161. credential=credential,
  162. actor_uid=actor_uid,
  163. )
  164. self.outbox_enqueuer(
  165. session,
  166. aggregate_type="datasource",
  167. aggregate_id=uid,
  168. event_type="datasource.credential_version_created",
  169. payload={
  170. "uid": uid,
  171. "database_type": candidate.database_type,
  172. "credential_version": sealed.credential_version,
  173. },
  174. )
  175. session.commit()
  176. except Exception:
  177. session.rollback()
  178. raise
  179. candidate = replace(
  180. candidate,
  181. credential_version=sealed.credential_version,
  182. )
  183. try:
  184. saved = self.definitions.save(candidate)
  185. except Exception as graph_error:
  186. try:
  187. self.credentials.compensate_failed_activation(
  188. session,
  189. data_source_uid=uid,
  190. failed_version=sealed.credential_version,
  191. restore_version=previous_version,
  192. actor_uid=actor_uid,
  193. )
  194. session.commit()
  195. except Exception as compensation_error:
  196. session.rollback()
  197. raise DataSourceError(
  198. "data source save requires reconciliation"
  199. ) from compensation_error
  200. raise DataSourceError(
  201. "data source definition could not be saved"
  202. ) from graph_error
  203. if existing is not None:
  204. self.connection_manager.invalidate(
  205. uid,
  206. "configuration_changed",
  207. )
  208. return saved, existing is None
  209. def list(self, payload):
  210. payload = payload if isinstance(payload, dict) else {}
  211. def optional_text(name):
  212. value = payload.get(name)
  213. if value is None:
  214. return None
  215. normalized = str(value).strip()
  216. return normalized or None
  217. filters = DataSourceFilters(
  218. uid=optional_text("uid"),
  219. name_en=optional_text("name_en"),
  220. name_zh=optional_text("name_zh"),
  221. database_type=optional_text("type"),
  222. status=payload.get("status"),
  223. )
  224. return self.definitions.list(filters)
  225. def test_connection(self, payload):
  226. if not isinstance(payload, dict):
  227. raise DataSourceConfigurationInvalid()
  228. uid = str(payload.get("uid") or "").strip() or None
  229. existing = self.definitions.get(uid) if uid else None
  230. if uid and existing is None:
  231. raise DataSourceNotFound()
  232. effective_uid = uid or new_governance_uid()
  233. credential = self._effective_credential(payload, existing)
  234. candidate = self._definition(
  235. payload,
  236. uid=effective_uid,
  237. credential_version=(
  238. existing.credential_version if existing else 1
  239. ),
  240. existing=existing,
  241. )
  242. adapter = self.adapter_resolver(candidate.database_type)
  243. settings = self.settings_resolver(candidate.pool_overrides())
  244. adapter.test_connection(
  245. candidate,
  246. credential,
  247. query_timeout=settings["query_timeout"],
  248. )
  249. return {"connected": True, "message": "连接测试成功"}
  250. def delete(self, uid, *, actor_uid):
  251. uid = str(uid or "").strip()
  252. existing = self.definitions.get(uid)
  253. if existing is None:
  254. raise DataSourceNotFound()
  255. self.connection_manager.invalidate(uid, "datasource_deleted")
  256. if not self.definitions.delete(uid):
  257. raise DataSourceNotFound()
  258. session = self.platform_session()
  259. try:
  260. revoked = self.credentials.revoke_all(
  261. session,
  262. data_source_uid=uid,
  263. actor_uid=actor_uid,
  264. )
  265. self.outbox_enqueuer(
  266. session,
  267. aggregate_type="datasource",
  268. aggregate_id=uid,
  269. event_type="datasource.deleted",
  270. payload={"uid": uid},
  271. )
  272. session.commit()
  273. except Exception as exc:
  274. session.rollback()
  275. raise DataSourceError(
  276. "data source deletion requires reconciliation"
  277. ) from exc
  278. return {"uid": uid, "revoked_credential_versions": revoked}
  279. def pool_statuses(self):
  280. return self.connection_manager.snapshot()
  281. def pool_status(self, uid):
  282. uid = str(uid or "").strip()
  283. for status in self.pool_statuses():
  284. if status.data_source_uid == uid:
  285. return status
  286. raise DataSourceNotFound()
  287. def invalidate_pool(self, uid, *, reason, actor_uid):
  288. uid = str(uid or "").strip()
  289. if reason not in POOL_INVALIDATION_REASONS:
  290. raise DataSourceConfigurationInvalid(
  291. "pool invalidation reason is invalid"
  292. )
  293. if self.definitions.get(uid) is None:
  294. raise DataSourceNotFound()
  295. self.connection_manager.invalidate(uid, reason)
  296. session = self.platform_session()
  297. try:
  298. self.credentials.record_pool_event(
  299. session,
  300. data_source_uid=uid,
  301. actor_uid=actor_uid,
  302. event_type="pool_invalidated",
  303. safe_detail=f"pool invalidated: {reason}",
  304. )
  305. session.commit()
  306. except Exception:
  307. session.rollback()
  308. raise
  309. return {"uid": uid, "invalidated": True, "reason": reason}
  310. @staticmethod
  311. def serialize(definition):
  312. return {
  313. "uid": definition.uid,
  314. "name_en": definition.name_en,
  315. "name_zh": definition.name_zh,
  316. "type": definition.database_type,
  317. "host": definition.host,
  318. "port": definition.port,
  319. "database": definition.database,
  320. "schema": definition.schema,
  321. "pool_size": definition.pool_size,
  322. "max_overflow": definition.max_overflow,
  323. "status": definition.status,
  324. "desc": definition.description,
  325. "credential_configured": bool(
  326. definition.credential_version
  327. ),
  328. }
  329. @staticmethod
  330. def serialize_pool_status(status):
  331. return {
  332. "data_source_uid": status.data_source_uid,
  333. "credential_version": status.credential_version,
  334. "pool_state": status.pool_state,
  335. "pool_size": status.pool_size,
  336. "checked_out": status.checked_out,
  337. "checked_out_peak": status.checked_out_peak,
  338. "checked_in": status.checked_in,
  339. "overflow": status.overflow,
  340. "leases": status.leases,
  341. "checkout_wait_ms": status.checkout_wait_ms,
  342. "last_used_at": status.last_used_at,
  343. "connection_created_total": status.connection_created_total,
  344. "connection_failed_total": status.connection_failed_total,
  345. "pool_timeout_total": status.pool_timeout_total,
  346. "invalidated_total": status.invalidated_total,
  347. "query_total": status.query_total,
  348. "last_query_duration_ms": status.last_query_duration_ms,
  349. "consecutive_failures": status.consecutive_failures,
  350. "circuit_open_until": status.circuit_open_until,
  351. }
  352. def build_data_source_service(manager):
  353. from app.core.events.outbox import enqueue_outbox
  354. return DataSourceService(
  355. definitions=manager.definitions,
  356. credentials=manager.credentials,
  357. platform_session=manager.platform_session,
  358. adapter_resolver=manager.adapter_resolver,
  359. connection_manager=manager,
  360. settings_resolver=manager.settings_resolver,
  361. outbox_enqueuer=enqueue_outbox,
  362. )