mirror of
https://github.com/GNS3/gns3-server.git
synced 2026-10-01 00:02:34 +03:00
fix(typing): resolve mypy errors in db.repositories.rbac
This commit is contained in:
parent
323055c486
commit
fa9d1cde80
@ -15,8 +15,10 @@
|
||||
# You should have received a copy of the GNU General Public License
|
||||
# along with this program. If not, see <http://www.gnu.org/licenses/>.
|
||||
|
||||
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)
|
||||
|
||||
@ -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:
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user