mirror of
https://github.com/GNS3/gns3-server.git
synced 2026-10-02 08:40:15 +03:00
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:
commit
2f8277aeb7
@ -23,6 +23,7 @@ from typing import Any, List, Union, Optional
|
|||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from gns3server.controller import Controller
|
from gns3server.controller import Controller
|
||||||
|
import gns3server.db.models as models
|
||||||
from gns3server.db.repositories.computes import ComputesRepository
|
from gns3server.db.repositories.computes import ComputesRepository
|
||||||
from gns3server.db.repositories.rbac import RbacRepository
|
from gns3server.db.repositories.rbac import RbacRepository
|
||||||
from gns3server.services.computes import ComputesService
|
from gns3server.services.computes import ComputesService
|
||||||
@ -31,7 +32,7 @@ from gns3server import schemas
|
|||||||
from .dependencies.database import get_repository
|
from .dependencies.database import get_repository
|
||||||
from .dependencies.rbac import has_privilege
|
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)
|
router = APIRouter(responses=responses)
|
||||||
|
|
||||||
@ -51,14 +52,14 @@ async def create_compute(
|
|||||||
compute_create: schemas.ComputeCreate,
|
compute_create: schemas.ComputeCreate,
|
||||||
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
|
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
|
||||||
connect: Optional[bool] = False,
|
connect: Optional[bool] = False,
|
||||||
) -> schemas.Compute:
|
) -> models.Compute:
|
||||||
"""
|
"""
|
||||||
Create a new compute on the controller.
|
Create a new compute on the controller.
|
||||||
|
|
||||||
Required privilege: Compute.Allocate
|
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(
|
@router.post(
|
||||||
@ -86,7 +87,7 @@ async def connect_compute(compute_id: Union[str, UUID]) -> None:
|
|||||||
)
|
)
|
||||||
async def get_compute(
|
async def get_compute(
|
||||||
compute_id: Union[str, UUID], computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository))
|
compute_id: Union[str, UUID], computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository))
|
||||||
) -> schemas.Compute:
|
) -> Union[models.Compute, dict]:
|
||||||
"""
|
"""
|
||||||
Return a compute from the controller.
|
Return a compute from the controller.
|
||||||
|
|
||||||
@ -104,7 +105,7 @@ async def get_compute(
|
|||||||
)
|
)
|
||||||
async def get_computes(
|
async def get_computes(
|
||||||
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
|
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
|
||||||
) -> List[schemas.Compute]:
|
) -> List[models.Compute]:
|
||||||
"""
|
"""
|
||||||
Return all computes known by the controller.
|
Return all computes known by the controller.
|
||||||
|
|
||||||
@ -124,7 +125,7 @@ async def update_compute(
|
|||||||
compute_id: Union[str, UUID],
|
compute_id: Union[str, UUID],
|
||||||
compute_update: schemas.ComputeUpdate,
|
compute_update: schemas.ComputeUpdate,
|
||||||
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
|
computes_repo: ComputesRepository = Depends(get_repository(ComputesRepository)),
|
||||||
) -> schemas.Compute:
|
) -> models.Compute:
|
||||||
"""
|
"""
|
||||||
Update a compute on the controller.
|
Update a compute on the controller.
|
||||||
|
|
||||||
|
|||||||
@ -20,7 +20,7 @@ import os
|
|||||||
import time
|
import time
|
||||||
import psutil
|
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.responses import StreamingResponse
|
||||||
from fastapi.encoders import jsonable_encoder
|
from fastapi.encoders import jsonable_encoder
|
||||||
from fastapi.routing import Mount
|
from fastapi.routing import Mount
|
||||||
@ -64,7 +64,7 @@ def get_version(request: Request) -> dict:
|
|||||||
# compute subapp
|
# compute subapp
|
||||||
controller_host = None
|
controller_host = None
|
||||||
for route in request.app.routes:
|
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
|
controller_host = route.app.state.controller_host
|
||||||
|
|
||||||
local_server = Config.instance().settings.Server.local
|
local_server = Config.instance().settings.Server.local
|
||||||
@ -129,7 +129,7 @@ async def shutdown() -> None:
|
|||||||
try:
|
try:
|
||||||
future.result()
|
future.result()
|
||||||
except Exception as e:
|
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
|
continue
|
||||||
|
|
||||||
# then shutdown the server itself
|
# then shutdown the server itself
|
||||||
@ -191,8 +191,8 @@ async def statistics() -> dict:
|
|||||||
# Node statistics - distinguish open vs closed project nodes
|
# Node statistics - distinguish open vs closed project nodes
|
||||||
open_project_nodes = []
|
open_project_nodes = []
|
||||||
closed_project_nodes = []
|
closed_project_nodes = []
|
||||||
node_by_type = {}
|
node_by_type: Dict[str, int] = {}
|
||||||
node_by_status = {}
|
node_by_status: Dict[str, int] = {}
|
||||||
|
|
||||||
for project in projects:
|
for project in projects:
|
||||||
nodes = project.nodes.values()
|
nodes = project.nodes.values()
|
||||||
@ -264,9 +264,8 @@ async def controller_http_notifications(request: Request) -> StreamingResponse:
|
|||||||
|
|
||||||
from gns3server.api.server import app
|
from gns3server.api.server import app
|
||||||
|
|
||||||
log.info(
|
client = f"{request.client.host}:{request.client.port}" if request.client else "unknown"
|
||||||
f"New client {request.client.host}:{request.client.port} has connected to controller HTTP notification stream"
|
log.info(f"New client {client} has connected to controller HTTP notification stream")
|
||||||
)
|
|
||||||
|
|
||||||
async def event_stream():
|
async def event_stream():
|
||||||
try:
|
try:
|
||||||
@ -275,10 +274,7 @@ async def controller_http_notifications(request: Request) -> StreamingResponse:
|
|||||||
msg = await queue.get_json(5)
|
msg = await queue.get_json(5)
|
||||||
yield f"{msg}\n".encode("utf-8")
|
yield f"{msg}\n".encode("utf-8")
|
||||||
finally:
|
finally:
|
||||||
log.info(
|
log.info(f"Client {client} has disconnected from controller HTTP notification stream")
|
||||||
f"Client {request.client.host}:{request.client.port} has disconnected from controller HTTP "
|
|
||||||
f"notification stream"
|
|
||||||
)
|
|
||||||
|
|
||||||
return StreamingResponse(event_stream(), media_type="application/json")
|
return StreamingResponse(event_stream(), media_type="application/json")
|
||||||
|
|
||||||
@ -294,14 +290,15 @@ async def controller_ws_notifications(
|
|||||||
if current_user is None:
|
if current_user is None:
|
||||||
return
|
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:
|
try:
|
||||||
with Controller.instance().notification.controller_queue() as queue:
|
with Controller.instance().notification.controller_queue() as queue:
|
||||||
while True:
|
while True:
|
||||||
notification = await queue.get_json(5)
|
notification = await queue.get_json(5)
|
||||||
await websocket.send_text(notification)
|
await websocket.send_text(notification)
|
||||||
except (ConnectionClosed, WebSocketDisconnect):
|
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:
|
except WebSocketException as e:
|
||||||
log.warning(f"Error while sending to controller event to WebSocket client: {e}")
|
log.warning(f"Error while sending to controller event to WebSocket client: {e}")
|
||||||
|
|
||||||
|
|||||||
@ -14,7 +14,7 @@
|
|||||||
# You should have received a copy of the GNU General Public License
|
# You should have received a copy of the GNU General Public License
|
||||||
# along with this program. If not, see <http://www.gnu.org/licenses/>.
|
# 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 fastapi import Depends
|
||||||
from starlette.requests import HTTPConnection
|
from starlette.requests import HTTPConnection
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
@ -22,7 +22,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
|||||||
from gns3server.db.repositories.base import BaseRepository
|
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:
|
async with AsyncSession(request.app.state._db_engine, expire_on_commit=False) as session:
|
||||||
try:
|
try:
|
||||||
@ -32,7 +32,7 @@ async def get_db_session(request: HTTPConnection) -> AsyncSession:
|
|||||||
|
|
||||||
|
|
||||||
def get_repository(repo: Type[BaseRepository]) -> Callable:
|
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 repo(db_session)
|
||||||
|
|
||||||
return get_repo
|
return get_repo
|
||||||
|
|||||||
@ -21,7 +21,7 @@ API routes for user groups.
|
|||||||
|
|
||||||
from fastapi import APIRouter, Depends, status
|
from fastapi import APIRouter, Depends, status
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
from typing import List
|
from typing import List, Optional
|
||||||
|
|
||||||
from gns3server import schemas
|
from gns3server import schemas
|
||||||
from gns3server.controller.controller_error import (
|
from gns3server.controller.controller_error import (
|
||||||
@ -31,6 +31,7 @@ from gns3server.controller.controller_error import (
|
|||||||
ControllerForbiddenError,
|
ControllerForbiddenError,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
import gns3server.db.models as models
|
||||||
from gns3server.db.repositories.users import UsersRepository
|
from gns3server.db.repositories.users import UsersRepository
|
||||||
from gns3server.db.repositories.rbac import RbacRepository
|
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"))])
|
@router.get("", response_model=List[schemas.UserGroup], dependencies=[Depends(has_privilege("Group.Audit"))])
|
||||||
async def get_user_groups(
|
async def get_user_groups(
|
||||||
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
|
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
|
||||||
) -> List[schemas.UserGroup]:
|
) -> List[models.UserGroup]:
|
||||||
"""
|
"""
|
||||||
Get all user groups.
|
Get all user groups.
|
||||||
|
|
||||||
@ -65,7 +66,7 @@ async def get_user_groups(
|
|||||||
)
|
)
|
||||||
async def create_user_group(
|
async def create_user_group(
|
||||||
user_group_create: schemas.UserGroupCreate, users_repo: UsersRepository = Depends(get_repository(UsersRepository))
|
user_group_create: schemas.UserGroupCreate, users_repo: UsersRepository = Depends(get_repository(UsersRepository))
|
||||||
) -> schemas.UserGroup:
|
) -> models.UserGroup:
|
||||||
"""
|
"""
|
||||||
Create a new user group.
|
Create a new user group.
|
||||||
|
|
||||||
@ -82,7 +83,7 @@ async def create_user_group(
|
|||||||
async def get_user_group(
|
async def get_user_group(
|
||||||
user_group_id: UUID,
|
user_group_id: UUID,
|
||||||
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
|
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
|
||||||
) -> schemas.UserGroup:
|
) -> models.UserGroup:
|
||||||
"""
|
"""
|
||||||
Get a user group.
|
Get a user group.
|
||||||
|
|
||||||
@ -100,7 +101,7 @@ async def update_user_group(
|
|||||||
user_group_id: UUID,
|
user_group_id: UUID,
|
||||||
user_group_update: schemas.UserGroupUpdate,
|
user_group_update: schemas.UserGroupUpdate,
|
||||||
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
|
users_repo: UsersRepository = Depends(get_repository(UsersRepository)),
|
||||||
) -> schemas.UserGroup:
|
) -> Optional[models.UserGroup]:
|
||||||
"""
|
"""
|
||||||
Update a user group.
|
Update a user group.
|
||||||
|
|
||||||
@ -148,7 +149,7 @@ async def delete_user_group(
|
|||||||
)
|
)
|
||||||
async def get_user_group_members(
|
async def get_user_group_members(
|
||||||
user_group_id: UUID, users_repo: UsersRepository = Depends(get_repository(UsersRepository))
|
user_group_id: UUID, users_repo: UsersRepository = Depends(get_repository(UsersRepository))
|
||||||
) -> List[schemas.User]:
|
) -> List[models.User]:
|
||||||
"""
|
"""
|
||||||
Get all user group members.
|
Get all user group members.
|
||||||
|
|
||||||
|
|||||||
@ -80,7 +80,7 @@ class UserGroupCreate(UserGroupBase):
|
|||||||
Properties to create a user group.
|
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):
|
class UserGroupUpdate(UserGroupBase):
|
||||||
|
|||||||
@ -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).
|
# 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.
|
# Remove modules from this list one small PR at a time. Never add new ones.
|
||||||
module = [
|
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.images", # 9
|
||||||
"gns3server.api.routes.controller.llm_model_configs", # 32
|
"gns3server.api.routes.controller.llm_model_configs", # 32
|
||||||
"gns3server.api.routes.controller.pools", # 8
|
"gns3server.api.routes.controller.pools", # 8
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user