definitions.py 7.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233
  1. """Secret-free Neo4j repository for external data-source definitions."""
  2. import json
  3. from dataclasses import dataclass
  4. from typing import Optional
  5. from app.core.common.identifiers import ensure_governance_uid
  6. from app.core.data_source.models import DataSourceDefinition
  7. NEO4J_SECRET_KEYS = {
  8. "username",
  9. "password",
  10. "passwd",
  11. "credential",
  12. "credentials",
  13. "encrypted_payload",
  14. "nonce",
  15. "api_key",
  16. "token",
  17. "authorization",
  18. "conn_str",
  19. "connection_string",
  20. "connection_url",
  21. }
  22. @dataclass(frozen=True)
  23. class DataSourceFilters:
  24. uid: Optional[str] = None
  25. name_en: Optional[str] = None
  26. name_zh: Optional[str] = None
  27. database_type: Optional[str] = None
  28. status: Optional[bool] = None
  29. class DataSourceDefinitionRepository:
  30. """Read and write DataSource nodes using stable UIDs only."""
  31. def __init__(self, session):
  32. self.session = session
  33. @staticmethod
  34. def _assert_secret_free(properties: dict) -> None:
  35. def visit(value):
  36. if isinstance(value, dict):
  37. for key, item in value.items():
  38. if str(key).strip().lower() in NEO4J_SECRET_KEYS:
  39. raise ValueError(
  40. "secret properties cannot be stored in DataSource"
  41. )
  42. visit(item)
  43. elif isinstance(value, (list, tuple)):
  44. for item in value:
  45. visit(item)
  46. visit(properties)
  47. @classmethod
  48. def _to_properties(cls, definition: DataSourceDefinition) -> dict:
  49. properties = {
  50. "uid": definition.uid,
  51. "name_en": definition.name_en,
  52. "name_zh": definition.name_zh,
  53. "type": definition.database_type,
  54. "host": definition.host,
  55. "port": definition.port,
  56. "database": definition.database,
  57. "schema": definition.schema,
  58. "credential_ref": definition.credential_ref,
  59. "credential_version": definition.credential_version,
  60. "pool_size": definition.pool_size,
  61. "max_overflow": definition.max_overflow,
  62. "tls_options_json": json.dumps(
  63. dict(definition.tls_options),
  64. sort_keys=True,
  65. separators=(",", ":"),
  66. ),
  67. "status": definition.status,
  68. "desc": definition.description,
  69. }
  70. properties.update(dict(definition.extra_properties))
  71. properties = {
  72. key: value for key, value in properties.items() if value is not None
  73. }
  74. cls._assert_secret_free(properties)
  75. return properties
  76. @staticmethod
  77. def _from_properties(properties: dict) -> DataSourceDefinition:
  78. values = dict(properties)
  79. tls_raw = values.pop("tls_options_json", "{}") or "{}"
  80. try:
  81. tls_options = json.loads(tls_raw)
  82. except (TypeError, ValueError):
  83. tls_options = {}
  84. known = {
  85. "uid",
  86. "name_en",
  87. "name_zh",
  88. "type",
  89. "host",
  90. "port",
  91. "database",
  92. "schema",
  93. "credential_ref",
  94. "credential_version",
  95. "pool_size",
  96. "max_overflow",
  97. "status",
  98. "desc",
  99. }
  100. return DataSourceDefinition(
  101. uid=values.get("uid"),
  102. name_en=values.get("name_en", ""),
  103. name_zh=values.get("name_zh"),
  104. database_type=values.get("type", ""),
  105. host=values.get("host", ""),
  106. port=values.get("port", 0),
  107. database=values.get("database", ""),
  108. schema=values.get("schema"),
  109. credential_ref=values.get("credential_ref"),
  110. credential_version=values.get("credential_version"),
  111. pool_size=values.get("pool_size"),
  112. max_overflow=values.get("max_overflow"),
  113. tls_options=tls_options if isinstance(tls_options, dict) else {},
  114. status=bool(values.get("status", True)),
  115. description=values.get("desc"),
  116. extra_properties={
  117. key: value
  118. for key, value in values.items()
  119. if key not in known and key != "_id"
  120. },
  121. )
  122. def get(self, uid: str) -> Optional[DataSourceDefinition]:
  123. record = self.session.run(
  124. """
  125. MATCH (n:DataSource {uid: $uid})
  126. RETURN properties(n) AS properties
  127. """,
  128. {"uid": str(uid)},
  129. ).single()
  130. if record is None:
  131. return None
  132. return self._from_properties(dict(record["properties"]))
  133. def list(self, filters: DataSourceFilters = None):
  134. filters = filters or DataSourceFilters()
  135. clauses = []
  136. parameters = {}
  137. mapping = {
  138. "uid": "n.uid",
  139. "name_en": "n.name_en",
  140. "name_zh": "n.name_zh",
  141. "database_type": "n.type",
  142. "status": "n.status",
  143. }
  144. for field_name, property_name in mapping.items():
  145. value = getattr(filters, field_name)
  146. if value is not None:
  147. clauses.append(f"{property_name} = ${field_name}")
  148. parameters[field_name] = value
  149. where = f"WHERE {' AND '.join(clauses)}" if clauses else ""
  150. records = self.session.run(
  151. f"""
  152. MATCH (n:DataSource)
  153. {where}
  154. RETURN properties(n) AS properties
  155. ORDER BY n.name_en
  156. """,
  157. parameters,
  158. )
  159. return [
  160. self._from_properties(dict(record["properties"]))
  161. for record in records
  162. ]
  163. def save(self, definition: DataSourceDefinition) -> DataSourceDefinition:
  164. uid_holder = {"uid": definition.uid} if definition.uid else {}
  165. uid = ensure_governance_uid(uid_holder)
  166. saved = definition.with_uid(uid)
  167. properties = self._to_properties(saved)
  168. self.session.run(
  169. """
  170. MERGE (n:DataSource {uid: $uid})
  171. SET n = $properties
  172. RETURN properties(n) AS properties
  173. """,
  174. {"uid": uid, "properties": properties},
  175. )
  176. return saved
  177. def delete(self, uid: str) -> bool:
  178. record = self.session.run(
  179. """
  180. MATCH (n:DataSource {uid: $uid})
  181. WITH n, count(n) AS found
  182. DETACH DELETE n
  183. RETURN found AS deleted_count
  184. """,
  185. {"uid": str(uid)},
  186. ).single()
  187. return bool(record and record["deleted_count"])
  188. def credential_references(self):
  189. records = self.session.run(
  190. """
  191. MATCH (n:DataSource)
  192. WHERE n.uid IS NOT NULL
  193. RETURN n.uid AS uid, n.credential_ref AS credential_ref,
  194. n.credential_version AS credential_version
  195. """
  196. )
  197. return {
  198. str(record["uid"]): (
  199. record["credential_ref"],
  200. int(record["credential_version"]),
  201. )
  202. for record in records
  203. if record.get("credential_ref") is not None
  204. and record.get("credential_version") is not None
  205. }
  206. def remove_legacy_secret_fields(self, uid: str) -> None:
  207. self.session.run(
  208. """
  209. MATCH (n:DataSource {uid: $uid})
  210. REMOVE n.username, n.password, n.conn_str,
  211. n.connection_string, n.connection_url
  212. """,
  213. {"uid": str(uid)},
  214. )