Merge pull request #2930 from markparonyan/mypy-db-templates-images

fix(typing): resolve the remaining mypy errors (18 modules)
This commit is contained in:
Jeremy Grossmann 2026-09-29 21:37:14 +02:00 committed by GitHub
commit e90be1ebac
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
25 changed files with 201 additions and 147 deletions

View File

@ -38,6 +38,7 @@ from gns3server.utils.images import (
default_images_directory,
get_builtin_disks,
)
import gns3server.db.models as models
from gns3server.db.repositories.images import ImagesRepository
from gns3server.db.repositories.templates import TemplatesRepository
from gns3server.db.repositories.rbac import RbacRepository
@ -68,7 +69,7 @@ async def create_qemu_image(
image_path: str,
image_data: schemas.QemuDiskImageCreate,
images_repo: ImagesRepository = Depends(get_repository(ImagesRepository)),
) -> schemas.Image:
) -> models.Image:
"""
Create a new blank Qemu image.
@ -115,7 +116,7 @@ async def create_qemu_image(
async def get_images(
images_repo: ImagesRepository = Depends(get_repository(ImagesRepository)),
image_type: Optional[schemas.ImageType] = None,
) -> List[schemas.Image]:
) -> List[models.Image]:
"""
Return all images.
@ -139,7 +140,7 @@ async def upload_image(
current_user: schemas.User = Depends(get_current_active_user),
rbac_repo: RbacRepository = Depends(get_repository(RbacRepository)),
install_appliances: Optional[bool] = False,
) -> schemas.Image:
) -> models.Image:
"""
Upload an image.
@ -256,7 +257,7 @@ async def install_images(
async def get_image(
image_path: str,
images_repo: ImagesRepository = Depends(get_repository(ImagesRepository)),
) -> schemas.Image:
) -> models.Image:
"""
Return an image.
@ -299,7 +300,7 @@ async def delete_image(
templates = await images_repo.get_image_templates(image.image_id)
if templates:
template_names = ", ".join([template.name for template in templates])
template_names = ", ".join([str(template.name) for template in templates])
raise ControllerError(f"Image '{image_path}' is used by one or more templates: {template_names}")
project_names = Controller.instance().find_projects_using_image(image.filename)

View File

@ -206,7 +206,11 @@ async def create_user_llm_model_config(
# Extract config fields (excluding table-level fields)
config_fields = config_create.model_dump(exclude={"name", "model_type", "is_default"})
new_config = await llm_repo.create_user_config(
user_id, config_create.name, config_create.model_type, config_fields, is_default=config_create.is_default
user_id,
config_create.name,
config_create.model_type,
config_fields,
is_default=bool(config_create.is_default),
)
return schemas.LLMModelConfigResponse(
@ -346,6 +350,10 @@ async def set_user_default_llm_model_config(
# Get the updated config
config = await llm_repo.get_user_config(config_id)
if config is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail=f"LLM model configuration '{config_id}' not found"
)
return schemas.LLMModelConfigResponse(
config_id=config.config_id,
name=config.name,
@ -506,7 +514,11 @@ async def create_group_llm_model_config(
# Extract config fields (excluding table-level fields)
config_fields = config_create.model_dump(exclude={"name", "model_type", "is_default"})
new_config = await llm_repo.create_group_config(
group_id, config_create.name, config_create.model_type, config_fields, is_default=config_create.is_default
group_id,
config_create.name,
config_create.model_type,
config_fields,
is_default=bool(config_create.is_default),
)
return schemas.LLMModelConfigResponse(
@ -646,6 +658,10 @@ async def set_group_default_llm_model_config(
# Get the updated config
config = await llm_repo.get_group_config(config_id)
if config is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail=f"LLM model configuration '{config_id}' not found"
)
return schemas.LLMModelConfigResponse(
config_id=config.config_id,
name=config.name,

View File

@ -21,12 +21,13 @@ API routes for resource pools.
from fastapi import APIRouter, Depends, status
from uuid import UUID
from typing import List
from typing import List, Optional
from gns3server import schemas
from gns3server.controller.controller_error import ControllerError, ControllerBadRequestError, ControllerNotFoundError
from gns3server.controller import Controller
import gns3server.db.models as models
from gns3server.db.repositories.rbac import RbacRepository
from gns3server.db.repositories.pools import ResourcePoolsRepository
@ -43,7 +44,7 @@ router = APIRouter()
@router.get("", response_model=List[schemas.ResourcePool], dependencies=[Depends(has_privilege("Pool.Audit"))])
async def get_resource_pools(
pools_repo: ResourcePoolsRepository = Depends(get_repository(ResourcePoolsRepository)),
) -> List[schemas.ResourcePool]:
) -> List[models.ResourcePool]:
"""
Get all resource pools.
@ -62,7 +63,7 @@ async def get_resource_pools(
async def create_resource_pool(
resource_pool_create: schemas.ResourcePoolCreate,
pools_repo: ResourcePoolsRepository = Depends(get_repository(ResourcePoolsRepository)),
) -> schemas.ResourcePool:
) -> models.ResourcePool:
"""
Create a new resource pool
@ -80,7 +81,7 @@ async def create_resource_pool(
)
async def get_resource_pool(
resource_pool_id: UUID, pools_repo: ResourcePoolsRepository = Depends(get_repository(ResourcePoolsRepository))
) -> schemas.ResourcePool:
) -> models.ResourcePool:
"""
Get a resource pool.
@ -100,7 +101,7 @@ async def update_resource_pool(
resource_pool_id: UUID,
resource_pool_update: schemas.ResourcePoolUpdate,
pools_repo: ResourcePoolsRepository = Depends(get_repository(ResourcePoolsRepository)),
) -> schemas.ResourcePool:
) -> Optional[models.ResourcePool]:
"""
Update a resource pool.
@ -176,7 +177,7 @@ async def delete_resource_pool(
async def get_pool_resources(
resource_pool_id: UUID,
pools_repo: ResourcePoolsRepository = Depends(get_repository(ResourcePoolsRepository)),
) -> List[schemas.Resource]:
) -> List[models.Resource]:
"""
Get all resource in a pool.
@ -215,13 +216,13 @@ async def add_resource_to_pool(
# we only support projects in resource pools for now
project = Controller.instance().get_project(str(resource_id))
resource = await pools_repo.get_resource(resource_id)
if not resource:
db_resource = await pools_repo.get_resource(resource_id)
if not db_resource:
# the resource is not in the database yet, create it
resource_create = schemas.ResourceCreate(resource_id=resource_id, resource_type="project", name=project.name)
resource = await pools_repo.create_resource(resource_create)
db_resource = await pools_repo.create_resource(resource_create)
await pools_repo.add_resource_to_pool(resource_pool_id, resource)
await pools_repo.add_resource_to_pool(resource_pool_id, db_resource)
@router.delete(

View File

@ -34,7 +34,7 @@ from fastapi import APIRouter, Depends, Request, Body, HTTPException, status, We
from fastapi.encoders import jsonable_encoder
from fastapi.responses import StreamingResponse, FileResponse
from websockets.exceptions import ConnectionClosed, WebSocketException
from typing import List, Optional
from typing import Any, List, Optional
from uuid import UUID
from gns3server import schemas
@ -57,7 +57,9 @@ from .dependencies.rbac import has_privilege, has_privilege_on_websocket
from .dependencies.authentication import get_current_active_user
from .dependencies.database import get_repository
responses = {404: {"model": schemas.ErrorMessage, "description": "Could not find project"}}
responses: dict[int | str, dict[str, Any]] = {
404: {"model": schemas.ErrorMessage, "description": "Could not find project"}
}
router = APIRouter(responses=responses)
@ -636,8 +638,8 @@ async def export_project(
include_images: bool = False,
reset_mac_addresses: bool = False,
keep_compute_ids: bool = False,
compression: schemas.ProjectCompression = "zstd",
compression_level: int = None,
compression: schemas.ProjectCompression = schemas.ProjectCompression.zstd,
compression_level: Optional[int] = None,
) -> StreamingResponse:
"""
Export a project as a portable archive.
@ -650,19 +652,19 @@ async def export_project(
compression_query = compression.lower()
if compression_query == "zip":
compression = zipfile.ZIP_DEFLATED
zip_compression = zipfile.ZIP_DEFLATED
if compression_level is not None and (compression_level < 0 or compression_level > 9):
raise ControllerBadRequestError("Compression level must be between 0 and 9 for ZIP compression")
elif compression_query == "none":
compression = zipfile.ZIP_STORED
zip_compression = zipfile.ZIP_STORED
elif compression_query == "bzip2":
compression = zipfile.ZIP_BZIP2
zip_compression = zipfile.ZIP_BZIP2
if compression_level is not None and (compression_level < 1 or compression_level > 9):
raise ControllerBadRequestError("Compression level must be between 1 and 9 for BZIP2 compression")
elif compression_query == "lzma":
compression = zipfile.ZIP_LZMA
zip_compression = zipfile.ZIP_LZMA
elif compression_query == "zstd":
compression = zipfile.ZIP_ZSTANDARD
zip_compression = zipfile.ZIP_ZSTANDARD
if compression_level is not None and (compression_level < 1 or compression_level > 22):
raise ControllerBadRequestError("Compression level must be between 1 and 22 for Zstandard compression")
@ -681,7 +683,7 @@ async def export_project(
f"Exporting project '{project.name}' with '{compression_query}' compression (level {compression_level})"
)
with tempfile.TemporaryDirectory(dir=working_dir) as tmpdir:
with aiozipstream.ZipFile(compression=compression, compresslevel=compression_level) as zstream:
with aiozipstream.ZipFile(compression=zip_compression, compresslevel=compression_level) as zstream:
await export_controller_project(
zstream,
project,

View File

@ -27,7 +27,7 @@ import logging
log = logging.getLogger(__name__)
from fastapi import APIRouter, Request, HTTPException, Depends, Response, status, Query
from typing import List, Optional
from typing import Any, List, Optional, Union
from uuid import UUID
from gns3server import schemas
@ -43,7 +43,9 @@ from .dependencies.authentication import get_current_active_user
from .dependencies.rbac import has_privilege
from .dependencies.database import get_repository
responses = {404: {"model": schemas.ErrorMessage, "description": "Could not find template"}}
responses: dict[int | str, dict[str, Any]] = {
404: {"model": schemas.ErrorMessage, "description": "Could not find template"}
}
router = APIRouter(responses=responses)
@ -57,7 +59,7 @@ router = APIRouter(responses=responses)
async def create_template(
template_create: schemas.TemplateCreate,
templates_repo: TemplatesRepository = Depends(get_repository(TemplatesRepository)),
) -> schemas.Template:
) -> dict:
"""
Create a new template.
@ -80,7 +82,7 @@ async def get_template(
request: Request,
response: Response,
templates_repo: TemplatesRepository = Depends(get_repository(TemplatesRepository)),
) -> schemas.Template:
) -> Union[dict, Response]:
"""
Return a template.
@ -108,7 +110,7 @@ async def update_template(
template_id: UUID,
template_update: schemas.TemplateUpdate,
templates_repo: TemplatesRepository = Depends(get_repository(TemplatesRepository)),
) -> schemas.Template:
) -> dict:
"""
Update a template.
@ -125,7 +127,7 @@ async def delete_template(
template_id: UUID,
prune_images: Optional[bool] = False,
templates_repo: TemplatesRepository = Depends(get_repository(TemplatesRepository)),
images_repo: RbacRepository = Depends(get_repository(ImagesRepository)),
images_repo: ImagesRepository = Depends(get_repository(ImagesRepository)),
rbac_repo: RbacRepository = Depends(get_repository(RbacRepository)),
) -> None:
"""
@ -155,7 +157,7 @@ async def delete_template(
if str(template.template_id) != str(template_id)
]
if other_templates:
template_names = ", ".join([template.name for template in other_templates])
template_names = ", ".join([str(template.name) for template in other_templates])
raise ControllerError(f"Image '{image.path}' is used by one or more templates: {template_names}")
if referenced_filenames is None:
@ -193,7 +195,7 @@ async def get_templates(
templates_repo: TemplatesRepository = Depends(get_repository(TemplatesRepository)),
current_user: schemas.User = Depends(get_current_active_user),
tags: Optional[List[str]] = Query(None, description="Filter by tags (e.g. tags=vendor:cisco&tags=model:7200)"),
) -> List[schemas.Template]:
) -> List[dict]:
"""
Return all templates.
@ -244,7 +246,7 @@ async def get_templates(
)
async def duplicate_template(
template_id: UUID, templates_repo: TemplatesRepository = Depends(get_repository(TemplatesRepository))
) -> schemas.Template:
) -> dict:
"""
Duplicate a template.

View File

@ -22,7 +22,7 @@ API routes for users.
from fastapi import APIRouter, Depends, HTTPException, Response, status
from fastapi.security import OAuth2PasswordRequestForm
from uuid import UUID
from typing import List
from typing import List, Optional
from gns3server import schemas
from gns3server.controller.controller_error import (
@ -32,6 +32,7 @@ from gns3server.controller.controller_error import (
ControllerForbiddenError,
)
import gns3server.db.models as models
from gns3server.db.repositories.users import UsersRepository
from gns3server.db.repositories.rbac import RbacRepository
from gns3server.services import auth_service
@ -168,7 +169,7 @@ async def update_logged_in_user(
user_update: schemas.LoggedInUserUpdate,
current_user: schemas.User = Depends(get_current_active_user),
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
) -> schemas.User:
) -> Optional[models.User]:
"""
Update the current active user.
"""
@ -180,7 +181,7 @@ async def update_logged_in_user(
@router.get("", response_model=List[schemas.User], dependencies=[Depends(has_privilege("User.Audit"))])
async def get_users(users_repo: UsersRepository = Depends(get_repository(UsersRepository))) -> List[schemas.User]:
async def get_users(users_repo: UsersRepository = Depends(get_repository(UsersRepository))) -> List[models.User]:
"""
Get all users.
@ -198,7 +199,7 @@ async def get_users(users_repo: UsersRepository = Depends(get_repository(UsersRe
)
async def create_user(
user_create: schemas.UserCreate, users_repo: UsersRepository = Depends(get_repository(UsersRepository))
) -> schemas.User:
) -> models.User:
"""
Create a new user.
@ -218,7 +219,7 @@ async def create_user(
async def get_user(
user_id: UUID,
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
) -> schemas.User:
) -> models.User:
"""
Get a user.
@ -236,7 +237,7 @@ async def update_user(
user_id: UUID,
user_update: schemas.UserUpdate,
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
) -> schemas.User:
) -> Optional[models.User]:
"""
Update a user.
@ -287,7 +288,7 @@ async def delete_user(
)
async def get_user_memberships(
user_id: UUID, users_repo: UsersRepository = Depends(get_repository(UsersRepository))
) -> List[schemas.UserGroup]:
) -> List[models.UserGroup]:
"""
Get user memberships.

View File

@ -1050,7 +1050,11 @@ class VirtualBoxVM(BaseNode):
await self._stop_remote_console()
await self._start_console()
@BaseNode.console_type.setter
@property
def console_type(self):
return self._console_type
@console_type.setter
def console_type(self, new_console_type):
"""
Sets the console type for this VirtualBox VM.

View File

@ -21,7 +21,7 @@ import asyncio
import platformdirs
from typing import Tuple, List
from typing import List, Optional, Tuple
from aiohttp.client_exceptions import ClientError
from uuid import UUID
@ -210,7 +210,7 @@ class ApplianceManager:
log.info(f"Template '{template.get('name')}' has been created")
return template
async def _appliance_to_template(self, appliance: Appliance, version: str = None) -> dict:
async def _appliance_to_template(self, appliance: Appliance, version: Optional[dict] = None) -> dict:
"""
Get template data from appliance
"""
@ -314,7 +314,7 @@ class ApplianceManager:
templates_repo: TemplatesRepository,
rbac_repo: RbacRepository,
current_user: schemas.User,
) -> None:
) -> dict:
"""
Install a new appliance
"""
@ -362,7 +362,7 @@ class ApplianceManager:
template_data = await self._appliance_to_template(appliance)
return await self._create_template(template_data, templates_repo, rbac_repo, current_user)
def load_appliances(self, symbol_theme: str = None) -> None:
def load_appliances(self, symbol_theme: Optional[str] = None) -> None:
"""
Loads appliance files from disk.
"""
@ -403,7 +403,7 @@ class ApplianceManager:
print(f"Cannot load appliance file '{path}': {e}")
continue
def _get_default_symbol(self, appliance: dict, symbol_theme: str) -> str:
def _get_default_symbol(self, appliance: dict, symbol_theme: Optional[str]) -> str:
"""
Returns the default symbol for a given appliance.
"""

View File

@ -479,10 +479,6 @@ class Project:
def path(self):
return self._path
@property
def status(self):
return self._status
@path.setter
def path(self, path):
check_path_allowed(path)
@ -510,6 +506,10 @@ class Project:
self._path = path
@property
def status(self):
return self._status
@property
def captures_directory(self):
"""

View File

@ -16,12 +16,14 @@
# along with this program. If not, see <http://www.gnu.org/licenses/>.
import uuid
from datetime import datetime
from typing import Optional
from fastapi.encoders import jsonable_encoder
from sqlalchemy import Column, DateTime, func, inspect
from sqlalchemy.types import TypeDecorator, CHAR, VARCHAR
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy.orm import DeclarativeBase
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
class Base(DeclarativeBase):
@ -105,7 +107,9 @@ class BaseTable(Base):
__abstract__ = True
created_at = Column(DateTime, server_default=func.current_timestamp())
updated_at = Column(DateTime, server_default=func.current_timestamp(), onupdate=func.current_timestamp())
updated_at: Mapped[Optional[datetime]] = mapped_column(
DateTime, server_default=func.current_timestamp(), onupdate=func.current_timestamp()
)
__mapper_args__ = {"eager_defaults": True}

View File

@ -16,7 +16,7 @@
# along with this program. If not, see <http://www.gnu.org/licenses/>.
from sqlalchemy import Table, Column, String, ForeignKey, BigInteger, Integer
from sqlalchemy.orm import relationship
from sqlalchemy.orm import Mapped, mapped_column, relationship
from .base import Base, BaseTable, GUID
@ -32,9 +32,9 @@ image_template_map = Table(
class Image(BaseTable):
__tablename__ = "images"
image_id = Column(Integer, primary_key=True, autoincrement=True)
image_id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
filename = Column(String, index=True)
path = Column(String, unique=True)
path: Mapped[str] = mapped_column(String, unique=True, nullable=True)
image_type = Column(String)
image_size = Column(BigInteger)
checksum = Column(String, index=True)

View File

@ -16,7 +16,7 @@
# along with this program. If not, see <http://www.gnu.org/licenses/>.
from sqlalchemy import Column, Boolean, ForeignKey, CheckConstraint, Index, Integer, String, text, JSON
from sqlalchemy.orm import relationship
from sqlalchemy.orm import Mapped, mapped_column, relationship
from .base import BaseTable, generate_uuid, GUID
@ -36,11 +36,13 @@ class LLMModelConfig(BaseTable):
config_id = Column(GUID, primary_key=True, default=generate_uuid)
name = Column(String(100), nullable=False) # Configuration name (table-level for indexing)
model_type = Column(String(50), nullable=False) # Model type: text, vision, stt, tts, multimodal, etc.
config = Column(JSON, nullable=False) # Config fields: provider, base_url, model, temperature, api_key, etc.
config: Mapped[dict] = mapped_column(
JSON, nullable=False
) # Config fields: provider, base_url, model, temperature, api_key, etc.
user_id = Column(GUID, ForeignKey("users.user_id", ondelete="CASCADE"), nullable=True)
group_id = Column(GUID, ForeignKey("user_groups.user_group_id", ondelete="CASCADE"), nullable=True)
is_default = Column(Boolean, default=False, nullable=False)
version = Column(Integer, default=0, nullable=False) # Optimistic locking version
version: Mapped[int] = mapped_column(Integer, default=0, nullable=False) # Optimistic locking version
# Relationships
user = relationship("User", backref="llm_model_configs")

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 Table, Column, String, ForeignKey
from sqlalchemy.orm import relationship
from sqlalchemy.orm import Mapped, mapped_column, relationship
from .base import Base, BaseTable, generate_uuid, GUID
@ -36,7 +38,7 @@ resource_pool_map = Table(
class Resource(BaseTable):
__tablename__ = "resources"
resource_id = Column(GUID, primary_key=True)
resource_id: Mapped[uuid.UUID] = mapped_column(GUID, primary_key=True)
name = Column(String, unique=True, index=True)
resource_type = Column(String)
resource_pools = relationship("ResourcePool", secondary=resource_pool_map, back_populates="resources")
@ -45,6 +47,6 @@ class Resource(BaseTable):
class ResourcePool(BaseTable):
__tablename__ = "resource_pools"
resource_pool_id = Column(GUID, primary_key=True, default=generate_uuid)
resource_pool_id: Mapped[uuid.UUID] = mapped_column(GUID, primary_key=True, default=generate_uuid)
name = Column(String, unique=True, index=True)
resources = relationship("Resource", secondary=resource_pool_map, back_populates="resource_pools")

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

@ -15,8 +15,11 @@
# You should have received a copy of the GNU General Public License
# along with this program. If not, see <http://www.gnu.org/licenses/>.
from datetime import datetime
from typing import Optional
from sqlalchemy import Table, Boolean, Column, Integer, String, DateTime, ForeignKey, event
from sqlalchemy.orm import relationship
from sqlalchemy.orm import Mapped, mapped_column, relationship
from .base import Base, BaseTable, generate_uuid, GUID
@ -24,6 +27,7 @@ from gns3server.config import Config
from gns3server.services import auth_service
import logging
import uuid
log = logging.getLogger(__name__)
@ -43,8 +47,8 @@ class User(BaseTable):
email = Column(String, unique=True, index=True)
full_name = Column(String)
hashed_password = Column(String)
last_login = Column(DateTime)
token_version = Column(Integer, default=0, nullable=False, server_default="0")
last_login: Mapped[Optional[datetime]] = mapped_column(DateTime)
token_version: Mapped[int] = mapped_column(Integer, default=0, nullable=False, server_default="0")
is_active = Column(Boolean, default=True)
is_superadmin = Column(Boolean, default=False)
groups = relationship("UserGroup", secondary=user_group_map, back_populates="users")
@ -72,7 +76,7 @@ def create_default_super_admin(target, connection, **kw):
class UserGroup(BaseTable):
__tablename__ = "user_groups"
user_group_id = Column(GUID, primary_key=True, default=generate_uuid)
user_group_id: Mapped[uuid.UUID] = mapped_column(GUID, primary_key=True, default=generate_uuid)
name = Column(String, unique=True, index=True)
is_builtin = Column(Boolean, default=False)
users = relationship("User", secondary=user_group_map, back_populates="groups")

View File

@ -15,9 +15,10 @@
# along with this program. If not, see <http://www.gnu.org/licenses/>.
from uuid import UUID
from typing import Optional, List
from typing import Optional, List, cast
from datetime import datetime, timezone
from sqlalchemy import select, update, delete, func
from sqlalchemy.engine import CursorResult
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
@ -67,13 +68,13 @@ class ApiKeysRepository(BaseRepository):
query = update(models.ApiKey).where(models.ApiKey.api_key_id == api_key_id).values(revoked=True)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0
async def restore_api_key(self, api_key_id: UUID) -> bool:
query = update(models.ApiKey).where(models.ApiKey.api_key_id == api_key_id).values(revoked=False)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0
async def update_last_used(self, api_key_id: UUID) -> None:
query = update(models.ApiKey).where(models.ApiKey.api_key_id == api_key_id).values(last_used_at=func.now())
@ -84,4 +85,4 @@ class ApiKeysRepository(BaseRepository):
query = delete(models.ApiKey).where(models.ApiKey.api_key_id == api_key_id)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0

View File

@ -19,6 +19,8 @@ from uuid import UUID
from typing import Optional, List, Union
from sqlalchemy import select, update, delete
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.engine import CursorResult
from typing import cast
from .base import BaseRepository
@ -47,7 +49,7 @@ class ComputesRepository(BaseRepository):
query = select(models.Compute)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def create_compute(self, compute_create: schemas.ComputeCreate) -> models.Compute:
@ -58,7 +60,7 @@ class ComputesRepository(BaseRepository):
host=compute_create.host,
port=compute_create.port,
user=compute_create.user,
password=compute_create.password.get_secret_value(),
password=compute_create.password.get_secret_value() if compute_create.password else None,
)
self._db_session.add(db_compute)
await self._db_session.commit()
@ -87,4 +89,4 @@ class ComputesRepository(BaseRepository):
query = delete(models.Compute).where(models.Compute.compute_id == compute_id)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0

View File

@ -17,8 +17,9 @@
import os
from typing import Optional, List
from typing import Optional, List, cast
from sqlalchemy import select, delete, update
from sqlalchemy.engine import CursorResult
from sqlalchemy.ext.asyncio import AsyncSession
from .base import BaseRepository
@ -50,7 +51,7 @@ class ImagesRepository(BaseRepository):
result = await self._db_session.execute(query)
return result.scalars().one_or_none()
async def get_image_by_checksum(self, checksum: str, image_dir: str = None) -> Optional[models.Image]:
async def get_image_by_checksum(self, checksum: str, image_dir: Optional[str] = None) -> Optional[models.Image]:
"""
Get an image by its checksum.
"""
@ -76,9 +77,9 @@ class ImagesRepository(BaseRepository):
else:
query = select(models.Image)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def get_image_templates(self, image_id: int) -> Optional[List[models.Template]]:
async def get_image_templates(self, image_id: int) -> List[models.Template]:
"""
Get all templates that an image belongs to.
"""
@ -86,7 +87,7 @@ class ImagesRepository(BaseRepository):
query = select(models.Template).join(models.Template.images).filter(models.Image.image_id == image_id)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def add_image(self, image_name, image_type, image_size, path, checksum, checksum_algorithm) -> models.Image:
"""
@ -108,7 +109,7 @@ class ImagesRepository(BaseRepository):
await self._db_session.refresh(db_image)
return db_image
async def update_image(self, image_path: str, checksum: str, checksum_algorithm: str) -> models.Image:
async def update_image(self, image_path: str, checksum: str, checksum_algorithm: str) -> Optional[models.Image]:
"""
Update an image.
"""
@ -142,9 +143,9 @@ class ImagesRepository(BaseRepository):
query = delete(models.Image).where(models.Image.filename == image_name)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0
async def prune_images(self, skip_images: list[str] = None) -> int:
async def prune_images(self, skip_images: Optional[list[str]] = None) -> int:
"""
Prune images not attached to any template.
"""

View File

@ -16,8 +16,9 @@
# along with this program. If not, see <http://www.gnu.org/licenses/>.
from uuid import UUID
from typing import Optional, List, Dict, Any
from typing import Optional, List, Dict, Any, cast
from sqlalchemy import select, update, delete, and_
from sqlalchemy.engine import CursorResult
from datetime import datetime
import logging
@ -63,7 +64,7 @@ class LLMModelConfigsRepository(BaseRepository):
.order_by(models.LLMModelConfig.created_at)
)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def get_user_default_config(self, user_id: UUID) -> Optional[models.LLMModelConfig]:
"""Get a user's default LLM model configuration."""
@ -182,7 +183,7 @@ class LLMModelConfigsRepository(BaseRepository):
)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0
async def set_user_default_config(self, user_id: UUID, config_id: UUID) -> bool:
"""Set a user's default LLM model configuration."""
@ -202,7 +203,7 @@ class LLMModelConfigsRepository(BaseRepository):
)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0
# Group configuration methods
@ -222,7 +223,7 @@ class LLMModelConfigsRepository(BaseRepository):
.order_by(models.LLMModelConfig.created_at)
)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def get_group_default_config(self, group_id: UUID) -> Optional[models.LLMModelConfig]:
"""Get a group's default LLM model configuration."""
@ -341,7 +342,7 @@ class LLMModelConfigsRepository(BaseRepository):
)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0
async def set_group_default_config(self, group_id: UUID, config_id: UUID) -> bool:
"""Set a group's default LLM model configuration."""
@ -361,7 +362,7 @@ class LLMModelConfigsRepository(BaseRepository):
)
result = await self._db_session.execute(query)
await self._db_session.commit()
return result.rowcount > 0
return cast(CursorResult, result).rowcount > 0
# Inheritance methods
@ -441,21 +442,21 @@ class LLMModelConfigsRepository(BaseRepository):
)
# Select default_config with proper priority:
# 1. User's config marked with is_default: true
# 2. Group's config marked with is_default: true
# 3. First config in the list (user configs come first)
for config in configs_with_source:
if config["is_default"] and config["source"] == "user":
default_config = config
# 1. User's entry marked with is_default: true
# 2. Group's entry marked with is_default: true
# 3. First entry in the list (user configs come first)
for entry in configs_with_source:
if entry["is_default"] and entry["source"] == "user":
default_config = entry
break
if default_config is None:
for config in configs_with_source:
if config["is_default"] and config["source"] == "group":
default_config = config
for entry in configs_with_source:
if entry["is_default"] and entry["source"] == "group":
default_config = entry
break
# Fallback to first config if no default is marked
# Fallback to first entry if no default is marked
if default_config is None and configs_with_source:
default_config = configs_with_source[0]

View File

@ -16,8 +16,9 @@
# along with this program. If not, see <http://www.gnu.org/licenses/>.
from uuid import UUID
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
@ -52,7 +53,7 @@ class ResourcePoolsRepository(BaseRepository):
query = select(models.Resource)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def create_resource(self, resource: schemas.ResourceCreate) -> models.Resource:
"""
@ -75,7 +76,7 @@ class ResourcePoolsRepository(BaseRepository):
query = delete(models.Resource).where(models.Resource.resource_id == resource_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 get_resource_memberships(self, resource_id: UUID) -> List[models.ResourcePool]:
"""
@ -89,7 +90,7 @@ class ResourcePoolsRepository(BaseRepository):
)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def get_resource_pool(self, resource_pool_id: UUID) -> Optional[models.ResourcePool]:
"""
@ -116,7 +117,7 @@ class ResourcePoolsRepository(BaseRepository):
query = select(models.ResourcePool)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def create_resource_pool(self, resource_pool: schemas.ResourcePoolCreate) -> models.ResourcePool:
"""
@ -166,7 +167,7 @@ class ResourcePoolsRepository(BaseRepository):
query = delete(models.ResourcePool).where(models.ResourcePool.resource_pool_id == resource_pool_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_resource_to_pool(
self, resource_pool_id: UUID, resource: models.Resource
@ -225,4 +226,4 @@ class ResourcePoolsRepository(BaseRepository):
)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())

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:

View File

@ -16,8 +16,9 @@
# along with this program. If not, see <http://www.gnu.org/licenses/>.
from uuid import UUID
from typing import Optional, List, Union
from typing import Optional, List, Union, cast
from sqlalchemy import select, update, delete, func
from sqlalchemy.engine import CursorResult
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
@ -72,7 +73,7 @@ class UsersRepository(BaseRepository):
query = select(models.User)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def create_user(self, user: schemas.UserCreate) -> models.User:
"""
@ -88,7 +89,9 @@ class UsersRepository(BaseRepository):
await self._db_session.refresh(db_user)
return db_user
async def update_user(self, user_id: UUID, user_update: schemas.UserUpdate) -> Optional[models.User]:
async def update_user(
self, user_id: UUID, user_update: Union[schemas.UserUpdate, schemas.LoggedInUserUpdate]
) -> Optional[models.User]:
"""
Update a user.
"""
@ -129,7 +132,7 @@ class UsersRepository(BaseRepository):
query = delete(models.User).where(models.User.user_id == user_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 authenticate_user(self, username: str, password: str) -> Optional[models.User]:
"""
@ -168,7 +171,7 @@ class UsersRepository(BaseRepository):
query = select(models.UserGroup).join(models.UserGroup.users).filter(models.User.user_id == user_id)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def get_user_group(self, user_group_id: UUID) -> Optional[models.UserGroup]:
"""
@ -195,7 +198,7 @@ class UsersRepository(BaseRepository):
query = select(models.UserGroup)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())
async def create_user_group(self, user_group: schemas.UserGroupCreate) -> models.UserGroup:
"""
@ -233,7 +236,7 @@ class UsersRepository(BaseRepository):
query = delete(models.UserGroup).where(models.UserGroup.user_group_id == user_group_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_member_to_user_group(self, user_group_id: UUID, user: models.User) -> Union[None, models.UserGroup]:
"""
@ -285,4 +288,4 @@ class UsersRepository(BaseRepository):
query = select(models.User).join(models.User.groups).filter(models.UserGroup.user_group_id == user_group_id)
result = await self._db_session.execute(query)
return result.scalars().all()
return list(result.scalars().all())

View File

@ -171,7 +171,7 @@ async def disconnect_from_db(app: FastAPI) -> None:
log.info(f"Disconnected from database")
async def get_computes(app: FastAPI) -> List[dict]:
async def get_computes(app: FastAPI) -> List[schemas.Compute]:
computes = []
async with AsyncSession(app.state._db_engine) as db_session:
@ -201,12 +201,12 @@ async def discover_images_on_filesystem(app: FastAPI) -> None:
continue
for image_type in ("qemu", "ios", "iou"):
discovered_images = await discover_images(image_type, existing_image_paths)
for image in discovered_images:
log.info(f"Adding discovered image '{image['path']}' to the database")
for image_info in discovered_images:
log.info(f"Adding discovered image '{image_info['path']}' to the database")
try:
await images_repository.add_image(**image)
await images_repository.add_image(**image_info)
except SQLAlchemyError as e:
log.warning(f"Error while adding image '{image['path']}' to the database: {e}")
log.warning(f"Error while adding image '{image_info['path']}' to the database: {e}")
# monitor if images have been manually added
asyncio.create_task(monitor_images_on_filesystem(app))
@ -237,7 +237,7 @@ class EventHandler(PatternMatchingEventHandler):
Watchdog event handler.
"""
def __init__(self, queue: asyncio.Queue, loop: asyncio.BaseEventLoop, **kwargs):
def __init__(self, queue: asyncio.Queue, loop: asyncio.AbstractEventLoop, **kwargs):
self._loop = loop
self._queue = queue
@ -274,7 +274,7 @@ class EventIterator(object):
async def monitor_images_on_filesystem(app: FastAPI):
def watchdog(
path: str, queue: asyncio.Queue, loop: asyncio.BaseEventLoop, app: FastAPI, recursive: bool = False
path: str, queue: asyncio.Queue, loop: asyncio.AbstractEventLoop, app: FastAPI, recursive: bool = False
) -> None:
"""
Thread to monitor a directory for new images.
@ -295,7 +295,7 @@ async def monitor_images_on_filesystem(app: FastAPI):
loop.call_soon_threadsafe(queue.put_nowait, None)
break
queue = asyncio.Queue()
queue: asyncio.Queue = asyncio.Queue()
loop = asyncio.get_event_loop()
server_config = Config.instance().settings.Server
image_dir = os.path.expanduser(server_config.images_path)
@ -348,7 +348,7 @@ async def get_user_llm_config_full(user_id: str, app: FastAPI) -> Optional[dict]
from gns3server.utils.encryption import decrypt, is_encrypted
try:
user_uuid = UUID(user_id) if isinstance(user_id, str) else user_id
user_uuid = UUID(user_id)
async with AsyncSession(app.state._db_engine, expire_on_commit=False) as session:
repo = LLMModelConfigsRepository(session)

View File

@ -38,10 +38,10 @@ class ComputeBase(BaseModel):
Data to create a compute.
"""
protocol: Protocol
host: str
port: int = Field(..., gt=0, le=65535)
user: str = None
protocol: Optional[Protocol] = None
host: Optional[str] = None
port: Optional[int] = Field(None, gt=0, le=65535)
user: Optional[str] = None
password: Optional[SecretStr] = None
name: Optional[str] = None
model_config = ConfigDict(use_enum_values=True)
@ -52,7 +52,10 @@ class ComputeCreate(ComputeBase):
Data to create a compute.
"""
compute_id: Union[str, uuid.UUID] = None
protocol: Protocol
host: str
port: int = Field(..., gt=0, le=65535)
compute_id: Optional[Union[str, uuid.UUID]] = None
model_config = ConfigDict(
json_schema_extra={
"example": {"name": "My compute", "host": "127.0.0.1", "port": 3080, "user": "user", "password": "password"}
@ -77,9 +80,6 @@ class ComputeUpdate(ComputeBase):
Data to update a compute.
"""
protocol: Optional[Protocol] = None
host: Optional[str] = None
port: Optional[int] = Field(None, gt=0, le=65535)
user: Optional[str] = None
password: Optional[SecretStr] = None
model_config = ConfigDict(
@ -110,6 +110,9 @@ class Compute(DateTimeModelMixin, ComputeBase):
Data returned for a compute.
"""
protocol: Protocol
host: str
port: int = Field(..., gt=0, le=65535)
compute_id: Union[str, uuid.UUID]
name: str
connected: Optional[bool] = Field(None, description="Whether the controller is connected to the compute or not")

View File

@ -43,7 +43,7 @@ class ComputesService:
async def create_compute(self, compute_create: schemas.ComputeCreate, connect: bool = False) -> models.Compute:
if await self._computes_repo.get_compute(compute_create.compute_id):
if compute_create.compute_id and await self._computes_repo.get_compute(compute_create.compute_id):
raise ControllerBadRequestError(f"Compute '{compute_create.compute_id}' is already registered")
db_compute = await self._computes_repo.create_compute(compute_create)
compute = await self._controller.add_compute(