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