diff --git a/gns3server/api/routes/controller/images.py b/gns3server/api/routes/controller/images.py index a67c68351..65f534b8e 100644 --- a/gns3server/api/routes/controller/images.py +++ b/gns3server/api/routes/controller/images.py @@ -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) diff --git a/gns3server/api/routes/controller/llm_model_configs.py b/gns3server/api/routes/controller/llm_model_configs.py index 620f75c26..c36912a02 100644 --- a/gns3server/api/routes/controller/llm_model_configs.py +++ b/gns3server/api/routes/controller/llm_model_configs.py @@ -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, diff --git a/gns3server/api/routes/controller/pools.py b/gns3server/api/routes/controller/pools.py index a5ef8e8b5..882c05264 100644 --- a/gns3server/api/routes/controller/pools.py +++ b/gns3server/api/routes/controller/pools.py @@ -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( diff --git a/gns3server/api/routes/controller/projects.py b/gns3server/api/routes/controller/projects.py index 19c29eead..378a276b8 100644 --- a/gns3server/api/routes/controller/projects.py +++ b/gns3server/api/routes/controller/projects.py @@ -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, diff --git a/gns3server/api/routes/controller/templates.py b/gns3server/api/routes/controller/templates.py index 84d55d693..3c35eea16 100644 --- a/gns3server/api/routes/controller/templates.py +++ b/gns3server/api/routes/controller/templates.py @@ -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. diff --git a/gns3server/api/routes/controller/users.py b/gns3server/api/routes/controller/users.py index 1f834d712..673886195 100644 --- a/gns3server/api/routes/controller/users.py +++ b/gns3server/api/routes/controller/users.py @@ -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. diff --git a/gns3server/compute/virtualbox/virtualbox_vm.py b/gns3server/compute/virtualbox/virtualbox_vm.py index 6c9fb6154..8cb9f8016 100644 --- a/gns3server/compute/virtualbox/virtualbox_vm.py +++ b/gns3server/compute/virtualbox/virtualbox_vm.py @@ -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. diff --git a/gns3server/controller/appliance_manager.py b/gns3server/controller/appliance_manager.py index 3bf0a302c..53dacadcf 100644 --- a/gns3server/controller/appliance_manager.py +++ b/gns3server/controller/appliance_manager.py @@ -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. """ diff --git a/gns3server/controller/project.py b/gns3server/controller/project.py index 00906bc42..ad42e9323 100644 --- a/gns3server/controller/project.py +++ b/gns3server/controller/project.py @@ -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): """ diff --git a/gns3server/db/models/base.py b/gns3server/db/models/base.py index 1802c1771..4984b8773 100644 --- a/gns3server/db/models/base.py +++ b/gns3server/db/models/base.py @@ -16,12 +16,14 @@ # along with this program. If not, see . 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} diff --git a/gns3server/db/models/images.py b/gns3server/db/models/images.py index 9755d8470..626773aad 100644 --- a/gns3server/db/models/images.py +++ b/gns3server/db/models/images.py @@ -16,7 +16,7 @@ # along with this program. If not, see . 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) diff --git a/gns3server/db/models/llm_model_configs.py b/gns3server/db/models/llm_model_configs.py index 778d16bbc..4a95719f4 100644 --- a/gns3server/db/models/llm_model_configs.py +++ b/gns3server/db/models/llm_model_configs.py @@ -16,7 +16,7 @@ # along with this program. If not, see . 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") diff --git a/gns3server/db/models/pools.py b/gns3server/db/models/pools.py index d8c011cbc..cd6616429 100644 --- a/gns3server/db/models/pools.py +++ b/gns3server/db/models/pools.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 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") 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/models/users.py b/gns3server/db/models/users.py index d9fc9aa29..a66d20025 100644 --- a/gns3server/db/models/users.py +++ b/gns3server/db/models/users.py @@ -15,8 +15,11 @@ # You should have received a copy of the GNU General Public License # along with this program. If not, see . +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") diff --git a/gns3server/db/repositories/api_keys.py b/gns3server/db/repositories/api_keys.py index d43d44dd1..1cbfdf25b 100644 --- a/gns3server/db/repositories/api_keys.py +++ b/gns3server/db/repositories/api_keys.py @@ -15,9 +15,10 @@ # along with this program. If not, see . 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 diff --git a/gns3server/db/repositories/computes.py b/gns3server/db/repositories/computes.py index 5747fbf03..df9e5932d 100644 --- a/gns3server/db/repositories/computes.py +++ b/gns3server/db/repositories/computes.py @@ -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 diff --git a/gns3server/db/repositories/images.py b/gns3server/db/repositories/images.py index c709c9820..5f9b35520 100644 --- a/gns3server/db/repositories/images.py +++ b/gns3server/db/repositories/images.py @@ -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. """ diff --git a/gns3server/db/repositories/llm_model_configs.py b/gns3server/db/repositories/llm_model_configs.py index f65ca8cac..2fb9188f4 100644 --- a/gns3server/db/repositories/llm_model_configs.py +++ b/gns3server/db/repositories/llm_model_configs.py @@ -16,8 +16,9 @@ # along with this program. If not, see . 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] diff --git a/gns3server/db/repositories/pools.py b/gns3server/db/repositories/pools.py index 90615ee68..a9144166e 100644 --- a/gns3server/db/repositories/pools.py +++ b/gns3server/db/repositories/pools.py @@ -16,8 +16,9 @@ # along with this program. If not, see . 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()) 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: diff --git a/gns3server/db/repositories/users.py b/gns3server/db/repositories/users.py index f7e098dc4..556185272 100644 --- a/gns3server/db/repositories/users.py +++ b/gns3server/db/repositories/users.py @@ -16,8 +16,9 @@ # along with this program. If not, see . 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()) diff --git a/gns3server/db/tasks.py b/gns3server/db/tasks.py index 8b3395156..a9554c9cc 100644 --- a/gns3server/db/tasks.py +++ b/gns3server/db/tasks.py @@ -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) diff --git a/gns3server/schemas/controller/computes.py b/gns3server/schemas/controller/computes.py index 1c87d95a0..f8612f43d 100644 --- a/gns3server/schemas/controller/computes.py +++ b/gns3server/schemas/controller/computes.py @@ -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") diff --git a/gns3server/services/computes.py b/gns3server/services/computes.py index a97c8fcf4..c3e3dad80 100644 --- a/gns3server/services/computes.py +++ b/gns3server/services/computes.py @@ -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(