device_scope.py 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146
  1. from __future__ import annotations
  2. from collections.abc import Callable
  3. from dataclasses import replace
  4. from datetime import datetime
  5. from typing import Any
  6. from uuid import UUID
  7. from app.core.common.timezone_utils import now_china_naive
  8. from app.core.data_research.repository import (
  9. SqlAlchemyIngestionSourceRepository,
  10. )
  11. from app.models.data_research import DeviceAssetSourceMapping, IngestionSource
  12. MAX_SOURCE_BUSINESS_DOMAINS = 100
  13. class DeviceSourceScopeError(ValueError):
  14. code = "DEVICE_SOURCE_SCOPE_ERROR"
  15. http_status = 422
  16. class DeviceSourceScopeInvalid(DeviceSourceScopeError):
  17. code = "DEVICE_SOURCE_SCOPE_INVALID"
  18. class DeviceSourceScopeNotFound(DeviceSourceScopeError):
  19. code = "DEVICE_SOURCE_SCOPE_NOT_FOUND"
  20. http_status = 404
  21. class DeviceSourceScopeForbidden(DeviceSourceScopeError):
  22. code = "DEVICE_SOURCE_SCOPE_FORBIDDEN"
  23. http_status = 403
  24. def _business_domains(payload: Any) -> tuple[str, ...]:
  25. if not isinstance(payload, (list, tuple)):
  26. raise DeviceSourceScopeInvalid("business_domains must be an array")
  27. if len(payload) > MAX_SOURCE_BUSINESS_DOMAINS:
  28. raise DeviceSourceScopeInvalid(
  29. f"business_domains cannot exceed {MAX_SOURCE_BUSINESS_DOMAINS}"
  30. )
  31. normalized: set[str] = set()
  32. for value in payload:
  33. try:
  34. normalized.add(str(UUID(str(value).strip())))
  35. except (AttributeError, TypeError, ValueError) as exc:
  36. raise DeviceSourceScopeInvalid(
  37. "each business_domains value must be a UUID"
  38. ) from exc
  39. return tuple(sorted(normalized))
  40. def source_scope_access(permission_scope: Any) -> dict[str, Any]:
  41. scope = permission_scope if isinstance(permission_scope, dict) else {}
  42. domains = _business_domains(scope.get("business_domains", []))
  43. return {
  44. "business_domains": domains,
  45. "admin_only": not domains,
  46. }
  47. class SqlAlchemyDeviceSourceScopeRepository(
  48. SqlAlchemyIngestionSourceRepository
  49. ):
  50. def list_device_sources(self):
  51. source_uids = (
  52. self.session.query(DeviceAssetSourceMapping.source_uid)
  53. .distinct()
  54. .subquery()
  55. )
  56. models = (
  57. self.session.query(IngestionSource)
  58. .filter(IngestionSource.uid.in_(source_uids))
  59. .order_by(IngestionSource.name.asc(), IngestionSource.uid.asc())
  60. .all()
  61. )
  62. return tuple(self._record(model) for model in models)
  63. class DeviceSourceScopeService:
  64. def __init__(
  65. self,
  66. repository,
  67. *,
  68. clock: Callable[[], datetime] = now_china_naive,
  69. commit: Callable[[], Any] = lambda: None,
  70. rollback: Callable[[], Any] = lambda: None,
  71. ):
  72. self._repository = repository
  73. self._clock = clock
  74. self._commit = commit
  75. self._rollback = rollback
  76. def list(self) -> tuple[dict[str, Any], ...]:
  77. records = []
  78. for source in self._repository.list_device_sources():
  79. access = source_scope_access(source.permission_scope)
  80. records.append(
  81. {
  82. "uid": source.uid,
  83. "name": source.name,
  84. "source_type": source.source_type,
  85. "status": source.status,
  86. **access,
  87. "updated_at": (
  88. source.updated_at.isoformat()
  89. if source.updated_at is not None
  90. else None
  91. ),
  92. }
  93. )
  94. return tuple(records)
  95. def update(
  96. self,
  97. source_uid: str,
  98. payload: Any,
  99. *,
  100. actor_is_admin: bool,
  101. ):
  102. if not actor_is_admin:
  103. raise DeviceSourceScopeForbidden(
  104. "only administrators can update device source scope"
  105. )
  106. if not isinstance(payload, dict) or "business_domains" not in payload:
  107. raise DeviceSourceScopeInvalid("business_domains is required")
  108. source = self._repository.get(str(source_uid))
  109. if source is None:
  110. raise DeviceSourceScopeNotFound("device source was not found")
  111. domains = _business_domains(payload["business_domains"])
  112. permission_scope = dict(source.permission_scope or {})
  113. permission_scope["business_domains"] = list(domains)
  114. updated = replace(
  115. source,
  116. permission_scope=permission_scope,
  117. updated_at=self._clock(),
  118. )
  119. try:
  120. saved = self._repository.save(updated)
  121. self._commit()
  122. return saved
  123. except Exception:
  124. self._rollback()
  125. raise