diff --git a/gns3server/db/models/roles.py b/gns3server/db/models/roles.py index 73c29dacb..b56b6ef95 100644 --- a/gns3server/db/models/roles.py +++ b/gns3server/db/models/roles.py @@ -15,8 +15,10 @@ # You should have received a copy of the GNU General Public License # along with this program. If not, see . +import uuid + from sqlalchemy import Column, String, Boolean, event -from sqlalchemy.orm import relationship +from sqlalchemy.orm import Mapped, mapped_column, relationship from .base import BaseTable, generate_uuid, GUID from .privileges import privilege_role_map @@ -29,7 +31,7 @@ log = logging.getLogger(__name__) class Role(BaseTable): __tablename__ = "roles" - role_id = Column(GUID, primary_key=True, default=generate_uuid) + role_id: Mapped[uuid.UUID] = mapped_column(GUID, primary_key=True, default=generate_uuid) name = Column(String, unique=True, index=True) description = Column(String) is_builtin = Column(Boolean, default=False) diff --git a/gns3server/db/repositories/rbac.py b/gns3server/db/repositories/rbac.py index 13bfd067c..72cb1f954 100644 --- a/gns3server/db/repositories/rbac.py +++ b/gns3server/db/repositories/rbac.py @@ -17,8 +17,9 @@ from uuid import UUID from urllib.parse import urlparse -from typing import Optional, List, Union +from typing import Optional, List, Union, cast from sqlalchemy import select, update, delete +from sqlalchemy.engine import CursorResult from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import selectinload @@ -62,9 +63,9 @@ class RbacRepository(BaseRepository): query = select(models.Role).options(selectinload(models.Role.privileges)) result = await self._db_session.execute(query) - return result.scalars().all() + return list(result.scalars().all()) - async def create_role(self, role_create: schemas.RoleCreate) -> models.Role: + async def create_role(self, role_create: schemas.RoleCreate) -> Optional[models.Role]: """ Create a new role. """ @@ -100,7 +101,7 @@ class RbacRepository(BaseRepository): query = delete(models.Role).where(models.Role.role_id == role_id) result = await self._db_session.execute(query) await self._db_session.commit() - return result.rowcount > 0 + return cast(CursorResult, result).rowcount > 0 async def add_privilege_to_role(self, role_id: UUID, privilege: models.Privilege) -> Union[None, models.Role]: """ @@ -149,7 +150,7 @@ class RbacRepository(BaseRepository): query = select(models.Privilege).join(models.Privilege.roles).filter(models.Role.role_id == role_id) result = await self._db_session.execute(query) - return result.scalars().all() + return list(result.scalars().all()) async def get_privilege(self, privilege_id: UUID) -> Optional[models.Privilege]: """ @@ -176,7 +177,7 @@ class RbacRepository(BaseRepository): query = select(models.Privilege) result = await self._db_session.execute(query) - return result.scalars().all() + return list(result.scalars().all()) async def get_ace(self, ace_id: UUID) -> Optional[models.ACE]: """ @@ -203,7 +204,7 @@ class RbacRepository(BaseRepository): query = select(models.ACE) result = await self._db_session.execute(query) - return result.scalars().all() + return list(result.scalars().all()) async def get_aces_for_path(self, path: str) -> List[models.ACE]: """ @@ -218,7 +219,7 @@ class RbacRepository(BaseRepository): .options(selectinload(models.ACE.user), selectinload(models.ACE.group), selectinload(models.ACE.role)) ) result = await self._db_session.execute(query) - return result.scalars().all() + return list(result.scalars().all()) async def check_ace_exists(self, path: str) -> bool: """ @@ -264,7 +265,7 @@ class RbacRepository(BaseRepository): query = delete(models.ACE).where(models.ACE.ace_id == ace_id) result = await self._db_session.execute(query) await self._db_session.commit() - return result.rowcount > 0 + return cast(CursorResult, result).rowcount > 0 async def delete_all_ace_starting_with_path(self, path: str) -> None: """ @@ -273,7 +274,7 @@ class RbacRepository(BaseRepository): query = delete(models.ACE).where(models.ACE.path.startswith(path)).execution_options(synchronize_session=False) result = await self._db_session.execute(query) - log.debug(f"{result.rowcount} ACE(s) have been deleted") + log.debug(f"{cast(CursorResult, result).rowcount} ACE(s) have been deleted") @staticmethod def _check_path_with_aces(path: str, aces) -> bool: @@ -297,7 +298,7 @@ class RbacRepository(BaseRepository): return True # only allow if the path is the original path or the ACE is set to propagate return False - async def _get_resources_in_pools(self, aces, path: str = None) -> List[models.Resource]: + async def _get_resources_in_pools(self, aces, path: Optional[str] = None) -> List[models.Resource]: """ Get all resources in pools. """ @@ -392,7 +393,7 @@ class RbacRepository(BaseRepository): all_resources = result.scalars().all() # Precompute pool_id -> set of project_ids - pool_to_projects = {} + pool_to_projects: dict[str, set[str]] = {} for r in all_resources: if r.resource_type == "project": for pool in r.resource_pools: