mirror of
https://github.com/GNS3/gns3-server.git
synced 2026-09-30 07:40:12 +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 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.
|
||||
|
||||
|
||||
@ -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}")
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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.
|
||||
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user