fix(typing): resolve mypy errors in db.repositories.rbac

This commit is contained in:
Mark Paronyan 2026-09-29 22:02:20 +03:00
parent 323055c486
commit fa9d1cde80
No known key found for this signature in database
GPG Key ID: 8DC7338A87FE22F8
2 changed files with 17 additions and 14 deletions

View File

@ -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)

View File

@ -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: