Merge pull request #2922 from markparonyan/mypy-api-routes-controller-1

fix(typing): resolve mypy errors in gns3server.api.routes.controller
This commit is contained in:
Jeremy Grossmann 2026-09-29 13:28:43 +02:00 committed by GitHub
commit 2f8277aeb7
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
6 changed files with 29 additions and 35 deletions

View File

@ -23,6 +23,7 @@ from typing import Any, List, Union, Optional
from uuid import UUID
from gns3server.controller import Controller
import gns3server.db.models as models
from gns3server.db.repositories.computes import ComputesRepository
from gns3server.db.repositories.rbac import RbacRepository
from gns3server.services.computes import ComputesService
@ -31,7 +32,7 @@ from gns3server import schemas
from .dependencies.database import get_repository
from .dependencies.rbac import has_privilege
responses = {404: {"model": schemas.ErrorMessage, "description": "Compute not found"}}
responses: dict[int | str, dict[str, Any]] = {404: {"model": schemas.ErrorMessage, "description": "Compute not found"}}
router = APIRouter(responses=responses)
@ -51,14 +52,14 @@ async def create_compute(
compute_create: schemas.ComputeCreate,
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
connect: Optional[bool] = False,
) -> schemas.Compute:
) -> models.Compute:
"""
Create a new compute on the controller.
Required privilege: Compute.Allocate
"""
return await ComputesService(computes_repo).create_compute(compute_create, connect)
return await ComputesService(computes_repo).create_compute(compute_create, bool(connect))
@router.post(
@ -86,7 +87,7 @@ async def connect_compute(compute_id: Union[str, UUID]) -> None:
)
async def get_compute(
compute_id: Union[str, UUID], computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository))
) -> schemas.Compute:
) -> Union[models.Compute, dict]:
"""
Return a compute from the controller.
@ -104,7 +105,7 @@ async def get_compute(
)
async def get_computes(
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
) -> List[schemas.Compute]:
) -> List[models.Compute]:
"""
Return all computes known by the controller.
@ -124,7 +125,7 @@ async def update_compute(
compute_id: Union[str, UUID],
compute_update: schemas.ComputeUpdate,
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
) -> schemas.Compute:
) -> models.Compute:
"""
Update a compute on the controller.

View File

@ -20,7 +20,7 @@ import os
import time
import psutil
from fastapi import APIRouter, Request, Depends, WebSocket, WebSocketDisconnect, status
from fastapi import APIRouter, FastAPI, Request, Depends, WebSocket, WebSocketDisconnect, status
from fastapi.responses import StreamingResponse
from fastapi.encoders import jsonable_encoder
from fastapi.routing import Mount
@ -64,7 +64,7 @@ def get_version(request: Request) -> dict:
# compute subapp
controller_host = None
for route in request.app.routes:
if isinstance(route, Mount) and route.name == "compute":
if isinstance(route, Mount) and route.name == "compute" and isinstance(route.app, FastAPI):
controller_host = route.app.state.controller_host
local_server = Config.instance().settings.Server.local
@ -129,7 +129,7 @@ async def shutdown() -> None:
try:
future.result()
except Exception as e:
log.error(f"Could not close project: {e}", exc_info=1)
log.error(f"Could not close project: {e}", exc_info=True)
continue
# then shutdown the server itself
@ -191,8 +191,8 @@ async def statistics() -> dict:
# Node statistics - distinguish open vs closed project nodes
open_project_nodes = []
closed_project_nodes = []
node_by_type = {}
node_by_status = {}
node_by_type: Dict[str, int] = {}
node_by_status: Dict[str, int] = {}
for project in projects:
nodes = project.nodes.values()
@ -264,9 +264,8 @@ async def controller_http_notifications(request: Request) -> StreamingResponse:
from gns3server.api.server import app
log.info(
f"New client {request.client.host}:{request.client.port} has connected to controller HTTP notification stream"
)
client = f"{request.client.host}:{request.client.port}" if request.client else "unknown"
log.info(f"New client {client} has connected to controller HTTP notification stream")
async def event_stream():
try:
@ -275,10 +274,7 @@ async def controller_http_notifications(request: Request) -> StreamingResponse:
msg = await queue.get_json(5)
yield f"{msg}\n".encode("utf-8")
finally:
log.info(
f"Client {request.client.host}:{request.client.port} has disconnected from controller HTTP "
f"notification stream"
)
log.info(f"Client {client} has disconnected from controller HTTP notification stream")
return StreamingResponse(event_stream(), media_type="application/json")
@ -294,14 +290,15 @@ async def controller_ws_notifications(
if current_user is None:
return
log.info(f"New client {websocket.client.host}:{websocket.client.port} has connected to controller WebSocket")
client = f"{websocket.client.host}:{websocket.client.port}" if websocket.client else "unknown"
log.info(f"New client {client} has connected to controller WebSocket")
try:
with Controller.instance().notification.controller_queue() as queue:
while True:
notification = await queue.get_json(5)
await websocket.send_text(notification)
except (ConnectionClosed, WebSocketDisconnect):
log.info(f"Client {websocket.client.host}:{websocket.client.port} has disconnected from controller WebSocket")
log.info(f"Client {client} has disconnected from controller WebSocket")
except WebSocketException as e:
log.warning(f"Error while sending to controller event to WebSocket client: {e}")

View File

@ -14,7 +14,7 @@
# 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 typing import Callable, Type
from typing import AsyncGenerator, Callable, Type
from fastapi import Depends
from starlette.requests import HTTPConnection
from sqlalchemy.ext.asyncio import AsyncSession
@ -22,7 +22,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from gns3server.db.repositories.base import BaseRepository
async def get_db_session(request: HTTPConnection) -> AsyncSession:
async def get_db_session(request: HTTPConnection) -> AsyncGenerator[AsyncSession, None]:
async with AsyncSession(request.app.state._db_engine, expire_on_commit=False) as session:
try:
@ -32,7 +32,7 @@ async def get_db_session(request: HTTPConnection) -> AsyncSession:
def get_repository(repo: Type[BaseRepository]) -> Callable:
def get_repo(db_session: AsyncSession = Depends(get_db_session)) -> Type[BaseRepository]:
def get_repo(db_session: AsyncSession = Depends(get_db_session)) -> BaseRepository:
return repo(db_session)
return get_repo

View File

@ -21,7 +21,7 @@ API routes for user groups.
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 (
@ -31,6 +31,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
@ -47,7 +48,7 @@ router = APIRouter()
@router.get("", response_model=List[schemas.UserGroup], dependencies=[Depends(has_privilege("Group.Audit"))])
async def get_user_groups(
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
) -> List[schemas.UserGroup]:
) -> List[models.UserGroup]:
"""
Get all user groups.
@ -65,7 +66,7 @@ async def get_user_groups(
)
async def create_user_group(
user_group_create: schemas.UserGroupCreate, users_repo: UsersRepository = Depends(get_repository(UsersRepository))
) -> schemas.UserGroup:
) -> models.UserGroup:
"""
Create a new user group.
@ -82,7 +83,7 @@ async def create_user_group(
async def get_user_group(
user_group_id: UUID,
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
) -> schemas.UserGroup:
) -> models.UserGroup:
"""
Get a user group.
@ -100,7 +101,7 @@ async def update_user_group(
user_group_id: UUID,
user_group_update: schemas.UserGroupUpdate,
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
) -> schemas.UserGroup:
) -> Optional[models.UserGroup]:
"""
Update a user group.
@ -148,7 +149,7 @@ async def delete_user_group(
)
async def get_user_group_members(
user_group_id: UUID, users_repo: UsersRepository = Depends(get_repository(UsersRepository))
) -> List[schemas.User]:
) -> List[models.User]:
"""
Get all user group members.

View File

@ -80,7 +80,7 @@ class UserGroupCreate(UserGroupBase):
Properties to create a user group.
"""
name: Optional[str] = Field(..., min_length=3, pattern="[a-zA-Z0-9_-]+$")
name: str = Field(..., min_length=3, pattern="[a-zA-Z0-9_-]+$")
class UserGroupUpdate(UserGroupBase):

View File

@ -294,11 +294,6 @@ enable_error_code = ["ignore-without-code", "redundant-expr", "truthy-bool"]
# Baseline: modules with existing type errors (error count at baseline time).
# Remove modules from this list one small PR at a time. Never add new ones.
module = [
"gns3server.api.routes.controller.computes", # 6
"gns3server.api.routes.controller.controller", # 10
"gns3server.api.routes.controller.dependencies.authentication", # 11
"gns3server.api.routes.controller.dependencies.database", # 2
"gns3server.api.routes.controller.groups", # 6
"gns3server.api.routes.controller.images", # 9
"gns3server.api.routes.controller.llm_model_configs", # 32
"gns3server.api.routes.controller.pools", # 8