diff --git a/gns3server/api/routes/controller/dependencies/authentication.py b/gns3server/api/routes/controller/dependencies/authentication.py index 1043db717..1517e0572 100644 --- a/gns3server/api/routes/controller/dependencies/authentication.py +++ b/gns3server/api/routes/controller/dependencies/authentication.py @@ -55,7 +55,7 @@ async def get_user_from_token( user_repo: UsersRepository = Depends(get_repository(UsersRepository)), api_keys_repo: ApiKeysRepository = Depends(get_repository(ApiKeysRepository)), token: Optional[str] = Query(None, include_in_schema=False), -) -> schemas.User: +) -> models.User: if bearer_token: # bearer token is used first, then any token passed as a URL parameter @@ -169,7 +169,7 @@ async def get_current_active_user_from_websocket( websocket: WebSocket, token: str = Query(...), user_repo: UsersRepository = Depends(get_repository(UsersRepository)), -) -> Optional[schemas.User]: +) -> Optional[models.User]: # Extract requested subprotocols from headers for proper WebSocket negotiation # This is critical for protocols like xpra that require specific subprotocols @@ -238,4 +238,5 @@ async def get_current_active_user_from_websocket( websocket_error = {"action": "log.error", "event": {"message": err_msg}} await websocket.send_json(websocket_error) log.error(err_msg) - return await websocket.close(code=1008) + await websocket.close(code=1008) + return None diff --git a/gns3server/db/models/api_keys.py b/gns3server/db/models/api_keys.py index 5cc40eb0b..ec320f9dd 100644 --- a/gns3server/db/models/api_keys.py +++ b/gns3server/db/models/api_keys.py @@ -14,7 +14,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, DateTime, ForeignKey, func +from sqlalchemy.orm import Mapped, mapped_column from .base import BaseTable, GUID @@ -22,8 +25,10 @@ from .base import BaseTable, GUID class ApiKey(BaseTable): __tablename__ = "api_keys" - api_key_id = Column(GUID, primary_key=True) - user_id = Column(GUID, ForeignKey("users.user_id", ondelete="CASCADE"), nullable=False, index=True) + api_key_id: Mapped[uuid.UUID] = mapped_column(GUID, primary_key=True) + user_id: Mapped[uuid.UUID] = mapped_column( + GUID, ForeignKey("users.user_id", ondelete="CASCADE"), nullable=False, index=True + ) name = Column(String(128), nullable=False) key_hash = Column(String(128), nullable=False) key_prefix = Column(String(8), nullable=False) diff --git a/gns3server/db/models/templates.py b/gns3server/db/models/templates.py index 62ea74a9e..0450be881 100644 --- a/gns3server/db/models/templates.py +++ b/gns3server/db/models/templates.py @@ -16,19 +16,26 @@ # along with this program. If not, see . +import uuid +from typing import Optional + from sqlalchemy import Boolean, Column, String, Integer, Float, ForeignKey, JSON -from sqlalchemy.orm import relationship +from sqlalchemy.orm import Mapped, mapped_column, relationship from .base import BaseTable, generate_uuid, GUID from .images import image_template_map +def template_id_column() -> Mapped[uuid.UUID]: + return mapped_column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + + class Template(BaseTable): __tablename__ = "templates" - template_id = Column(GUID, primary_key=True, default=generate_uuid) + template_id: Mapped[uuid.UUID] = mapped_column(GUID, primary_key=True, default=generate_uuid) name = Column(String, index=True) - version = Column(String) + version: Mapped[Optional[str]] = mapped_column(String) category = Column(String) default_name_format = Column(String) symbol = Column(String) @@ -36,7 +43,7 @@ class Template(BaseTable): usage = Column(String) netmiko_device_type = Column(String) appliance_metadata = Column(JSON) - template_type = Column(String) + template_type: Mapped[str] = mapped_column(String, nullable=True) tags = Column(JSON) compute_id = Column(String) images = relationship("Image", secondary=image_template_map, back_populates="templates") @@ -50,7 +57,7 @@ class Template(BaseTable): class CloudTemplate(Template): __tablename__ = "cloud_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + template_id: Mapped[uuid.UUID] = template_id_column() ports_mapping = Column(JSON) remote_console_host = Column(String) remote_console_port = Column(Integer) @@ -63,7 +70,7 @@ class CloudTemplate(Template): class DockerTemplate(Template): __tablename__ = "docker_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + template_id: Mapped[uuid.UUID] = template_id_column() image = Column(String) adapters = Column(Integer) mac_address = Column(String) @@ -88,10 +95,10 @@ class DockerTemplate(Template): class DynamipsTemplate(Template): __tablename__ = "dynamips_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) - platform = Column(String) + template_id: Mapped[uuid.UUID] = template_id_column() + platform: Mapped[str] = mapped_column(String, nullable=True) chassis = Column(String) - image = Column(String) + image: Mapped[str] = mapped_column(String, nullable=True) exec_area = Column(Integer) mmap = Column(Boolean) mac_addr = Column(String) @@ -130,7 +137,7 @@ class DynamipsTemplate(Template): class EthernetHubTemplate(Template): __tablename__ = "ethernet_hub_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + template_id: Mapped[uuid.UUID] = template_id_column() ports_mapping = Column(JSON) __mapper_args__ = {"polymorphic_identity": "ethernet_hub", "polymorphic_load": "selectin"} @@ -139,7 +146,7 @@ class EthernetHubTemplate(Template): class EthernetSwitchTemplate(Template): __tablename__ = "ethernet_switch_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + template_id: Mapped[uuid.UUID] = template_id_column() ports_mapping = Column(JSON) console_type = Column(String) @@ -149,8 +156,8 @@ class EthernetSwitchTemplate(Template): class IOUTemplate(Template): __tablename__ = "iou_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) - path = Column(String) + template_id: Mapped[uuid.UUID] = template_id_column() + path: Mapped[str] = mapped_column(String, nullable=True) ethernet_adapters = Column(Integer) serial_adapters = Column(Integer) ram = Column(Integer) @@ -168,7 +175,7 @@ class IOUTemplate(Template): class QemuTemplate(Template): __tablename__ = "qemu_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + template_id: Mapped[uuid.UUID] = template_id_column() qemu_path = Column(String) platform = Column(String) linked_clone = Column(Boolean) @@ -214,7 +221,7 @@ class QemuTemplate(Template): class VirtualBoxTemplate(Template): __tablename__ = "virtualbox_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + template_id: Mapped[uuid.UUID] = template_id_column() vmname = Column(String) ram = Column(Integer) linked_clone = Column(Boolean) @@ -236,7 +243,7 @@ class VirtualBoxTemplate(Template): class VMwareTemplate(Template): __tablename__ = "vmware_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + template_id: Mapped[uuid.UUID] = template_id_column() vmx_path = Column(String) linked_clone = Column(Boolean) first_port_name = Column(String) @@ -257,7 +264,7 @@ class VMwareTemplate(Template): class VPCSTemplate(Template): __tablename__ = "vpcs_templates" - template_id = Column(GUID, ForeignKey("templates.template_id", ondelete="CASCADE"), primary_key=True) + template_id: Mapped[uuid.UUID] = template_id_column() base_script_file = Column(String) console_type = Column(String) console_auto_start = Column(Boolean, default=False) diff --git a/gns3server/db/repositories/templates.py b/gns3server/db/repositories/templates.py index 582404301..9d5c1a9af 100644 --- a/gns3server/db/repositories/templates.py +++ b/gns3server/db/repositories/templates.py @@ -16,11 +16,13 @@ # along with this program. If not, see . import os +import uuid import logging from uuid import UUID -from typing import List, Union, Optional +from typing import List, Union, Optional, cast from sqlalchemy import select, delete +from sqlalchemy.engine import CursorResult from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import selectinload from sqlalchemy.orm.session import make_transient @@ -64,7 +66,7 @@ class TemplatesRepository(BaseRepository): result = await self._db_session.execute(query) return result.scalars().first() - async def get_template_by_name_and_version(self, name: str, version: str) -> Union[None, models.Template]: + async def get_template_by_name_and_version(self, name: str, version: Optional[str]) -> Union[None, models.Template]: query = ( select(models.Template) @@ -89,7 +91,7 @@ class TemplatesRepository(BaseRepository): query = select(models.Template).options(selectinload(models.Template.images)) result = await self._db_session.execute(query) - return result.scalars().all() + return list(result.scalars().all()) async def create_template(self, template_type: str, template_settings: dict) -> models.Template: @@ -100,7 +102,7 @@ class TemplatesRepository(BaseRepository): await self._db_session.refresh(db_template) return db_template - async def update_template(self, db_template: models.Template, template_settings: dict) -> schemas.Template: + async def update_template(self, db_template: models.Template, template_settings: dict) -> models.Template: # update the fields directly because update() query couldn't work for key, value in template_settings.items(): @@ -114,9 +116,9 @@ class TemplatesRepository(BaseRepository): query = delete(models.Template).where(models.Template.template_id == template_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 duplicate_template(self, template_id: UUID) -> Optional[schemas.Template]: + async def duplicate_template(self, template_id: UUID) -> Optional[models.Template]: query = ( select(models.Template) @@ -128,7 +130,7 @@ class TemplatesRepository(BaseRepository): # duplicate db object with new primary key (template_id) self._db_session.expunge(db_template) make_transient(db_template) - db_template.template_id = None + db_template.template_id = uuid.uuid4() self._db_session.add(db_template) await self._db_session.commit() await self._db_session.refresh(db_template) @@ -205,4 +207,4 @@ class TemplatesRepository(BaseRepository): query = select(models.Image).join(models.Image.templates).filter(models.Template.template_id == template_id) result = await self._db_session.execute(query) - return result.scalars().all() + return list(result.scalars().all()) diff --git a/gns3server/schemas/controller/tokens.py b/gns3server/schemas/controller/tokens.py index 29197acfa..2ac14dda6 100644 --- a/gns3server/schemas/controller/tokens.py +++ b/gns3server/schemas/controller/tokens.py @@ -25,7 +25,7 @@ class Token(BaseModel): class TokenData(BaseModel): - username: Optional[str] = None + username: str token_version: int = 0 token_use: str = "access" diff --git a/gns3server/services/templates.py b/gns3server/services/templates.py index 10a167c89..00f02b4b0 100644 --- a/gns3server/services/templates.py +++ b/gns3server/services/templates.py @@ -20,7 +20,7 @@ import pydantic from uuid import UUID from fastapi.encoders import jsonable_encoder -from typing import List +from typing import List, Optional from gns3server import schemas from gns3server.config import Config @@ -34,7 +34,7 @@ from gns3server.controller.controller_error import ( ) -TEMPLATE_TYPE_TO_SCHEMA = { +TEMPLATE_TYPE_TO_SCHEMA: dict[str, type[pydantic.BaseModel]] = { "cloud": schemas.CloudTemplate, "ethernet_hub": schemas.EthernetHubTemplate, "ethernet_switch": schemas.EthernetSwitchTemplate, @@ -47,7 +47,7 @@ TEMPLATE_TYPE_TO_SCHEMA = { "qemu": schemas.QemuTemplate, } -TEMPLATE_TYPE_TO_UPDATE_SCHEMA = { +TEMPLATE_TYPE_TO_UPDATE_SCHEMA: dict[str, type[pydantic.BaseModel]] = { "cloud": schemas.CloudTemplateUpdate, "ethernet_hub": schemas.EthernetHubTemplateUpdate, "ethernet_switch": schemas.EthernetSwitchTemplateUpdate, @@ -59,7 +59,7 @@ TEMPLATE_TYPE_TO_UPDATE_SCHEMA = { "qemu": schemas.QemuTemplateUpdate, } -DYNAMIPS_PLATFORM_TO_SCHEMA = { +DYNAMIPS_PLATFORM_TO_SCHEMA: dict[str, type[pydantic.BaseModel]] = { "c7200": schemas.C7200DynamipsTemplate, "c3745": schemas.C3745DynamipsTemplate, "c3725": schemas.C3725DynamipsTemplate, @@ -69,7 +69,7 @@ DYNAMIPS_PLATFORM_TO_SCHEMA = { "c1700": schemas.C1700DynamipsTemplate, } -DYNAMIPS_PLATFORM_TO_UPDATE_SCHEMA = { +DYNAMIPS_PLATFORM_TO_UPDATE_SCHEMA: dict[str, type[pydantic.BaseModel]] = { "c7200": schemas.C7200DynamipsTemplateUpdate, "c3745": schemas.C3745DynamipsTemplateUpdate, "c3725": schemas.C3725DynamipsTemplateUpdate, @@ -169,11 +169,12 @@ class TemplatesService: for builtin_template in BUILTIN_TEMPLATES: builtin_template["symbol"] = self._controller.symbols.resolve_symbol(builtin_template["symbol"]) - def get_builtin_template(self, template_id: UUID) -> dict: + def get_builtin_template(self, template_id: UUID) -> Optional[dict]: for builtin_template in BUILTIN_TEMPLATES: if builtin_template["template_id"] == template_id: return jsonable_encoder(builtin_template) + return None def _base_path(self): return self._templates_repo.configs_path() @@ -303,7 +304,7 @@ class TemplatesService: try: # validate the update settings update_settings = jsonable_encoder(template_update, exclude_unset=True) - if db_template.template_type == "dynamips": + if isinstance(db_template, models.DynamipsTemplate): template_schema = DYNAMIPS_PLATFORM_TO_UPDATE_SCHEMA[db_template.platform] else: template_schema = TEMPLATE_TYPE_TO_UPDATE_SCHEMA[db_template.template_type] @@ -312,9 +313,9 @@ class TemplatesService: raise ControllerBadRequestError(f"JSON schema error received while updating template: {e}") images_to_add_to_template = await self._find_images(db_template.template_type, template_settings) - if db_template.template_type == "dynamips" and "image" in template_settings: + if isinstance(db_template, models.DynamipsTemplate) and "image" in template_settings: await self._remove_image(db_template.template_id, db_template.image) - elif db_template.template_type == "iou" and "path" in template_settings: + elif isinstance(db_template, models.IOUTemplate) and "path" in template_settings: await self._remove_image(db_template.template_id, db_template.path) elif db_template.template_type == "qemu": for key in template_update.model_dump().keys():