| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146 |
- from __future__ import annotations
- from collections.abc import Callable
- from dataclasses import replace
- from datetime import datetime
- from typing import Any
- from uuid import UUID
- from app.core.common.timezone_utils import now_china_naive
- from app.core.data_research.repository import (
- SqlAlchemyIngestionSourceRepository,
- )
- from app.models.data_research import DeviceAssetSourceMapping, IngestionSource
- MAX_SOURCE_BUSINESS_DOMAINS = 100
- class DeviceSourceScopeError(ValueError):
- code = "DEVICE_SOURCE_SCOPE_ERROR"
- http_status = 422
- class DeviceSourceScopeInvalid(DeviceSourceScopeError):
- code = "DEVICE_SOURCE_SCOPE_INVALID"
- class DeviceSourceScopeNotFound(DeviceSourceScopeError):
- code = "DEVICE_SOURCE_SCOPE_NOT_FOUND"
- http_status = 404
- class DeviceSourceScopeForbidden(DeviceSourceScopeError):
- code = "DEVICE_SOURCE_SCOPE_FORBIDDEN"
- http_status = 403
- def _business_domains(payload: Any) -> tuple[str, ...]:
- if not isinstance(payload, (list, tuple)):
- raise DeviceSourceScopeInvalid("business_domains must be an array")
- if len(payload) > MAX_SOURCE_BUSINESS_DOMAINS:
- raise DeviceSourceScopeInvalid(
- f"business_domains cannot exceed {MAX_SOURCE_BUSINESS_DOMAINS}"
- )
- normalized: set[str] = set()
- for value in payload:
- try:
- normalized.add(str(UUID(str(value).strip())))
- except (AttributeError, TypeError, ValueError) as exc:
- raise DeviceSourceScopeInvalid(
- "each business_domains value must be a UUID"
- ) from exc
- return tuple(sorted(normalized))
- def source_scope_access(permission_scope: Any) -> dict[str, Any]:
- scope = permission_scope if isinstance(permission_scope, dict) else {}
- domains = _business_domains(scope.get("business_domains", []))
- return {
- "business_domains": domains,
- "admin_only": not domains,
- }
- class SqlAlchemyDeviceSourceScopeRepository(
- SqlAlchemyIngestionSourceRepository
- ):
- def list_device_sources(self):
- source_uids = (
- self.session.query(DeviceAssetSourceMapping.source_uid)
- .distinct()
- .subquery()
- )
- models = (
- self.session.query(IngestionSource)
- .filter(IngestionSource.uid.in_(source_uids))
- .order_by(IngestionSource.name.asc(), IngestionSource.uid.asc())
- .all()
- )
- return tuple(self._record(model) for model in models)
- class DeviceSourceScopeService:
- def __init__(
- self,
- repository,
- *,
- clock: Callable[[], datetime] = now_china_naive,
- commit: Callable[[], Any] = lambda: None,
- rollback: Callable[[], Any] = lambda: None,
- ):
- self._repository = repository
- self._clock = clock
- self._commit = commit
- self._rollback = rollback
- def list(self) -> tuple[dict[str, Any], ...]:
- records = []
- for source in self._repository.list_device_sources():
- access = source_scope_access(source.permission_scope)
- records.append(
- {
- "uid": source.uid,
- "name": source.name,
- "source_type": source.source_type,
- "status": source.status,
- **access,
- "updated_at": (
- source.updated_at.isoformat()
- if source.updated_at is not None
- else None
- ),
- }
- )
- return tuple(records)
- def update(
- self,
- source_uid: str,
- payload: Any,
- *,
- actor_is_admin: bool,
- ):
- if not actor_is_admin:
- raise DeviceSourceScopeForbidden(
- "only administrators can update device source scope"
- )
- if not isinstance(payload, dict) or "business_domains" not in payload:
- raise DeviceSourceScopeInvalid("business_domains is required")
- source = self._repository.get(str(source_uid))
- if source is None:
- raise DeviceSourceScopeNotFound("device source was not found")
- domains = _business_domains(payload["business_domains"])
- permission_scope = dict(source.permission_scope or {})
- permission_scope["business_domains"] = list(domains)
- updated = replace(
- source,
- permission_scope=permission_scope,
- updated_at=self._clock(),
- )
- try:
- saved = self._repository.save(updated)
- self._commit()
- return saved
- except Exception:
- self._rollback()
- raise
|