mirror of
https://github.com/GNS3/gns3-server.git
synced 2026-10-01 00:02:34 +03:00
Merge pull request #2930 from markparonyan/mypy-db-templates-images
fix(typing): resolve the remaining mypy errors (18 modules)
This commit is contained in:
commit
e90be1ebac
@ -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)
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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.
|
||||
|
||||
|
||||
@ -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.
|
||||
|
||||
|
||||
@ -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.
|
||||
|
||||
@ -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.
|
||||
"""
|
||||
|
||||
@ -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):
|
||||
"""
|
||||
|
||||
@ -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}
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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.
|
||||
"""
|
||||
|
||||
@ -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]
|
||||
|
||||
|
||||
@ -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())
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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())
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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(
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user