diff --git a/network/authentication.py b/network/authentication.py index d676cf4..a7a96e1 100644 --- a/network/authentication.py +++ b/network/authentication.py @@ -1,6 +1,28 @@ import fastapi +def get_bearer_token( + request: fastapi.Request | fastapi.WebSocket, +) -> str | None: + """ + 大小写不敏感地获取 Authorization header 值 + + uvicorn 按 ASGI 规范会把 header 名转为小写,但不同服务器/客户端可能保留 + 原始大小写,因此逐项比较 key.lower() 最稳妥(starlette 的 get() 依赖内部 + 存储大小写,两种场景下行为不一致) + + Args: + request (fastapi.Request | fastapi.WebSocket): 请求信息 + + Returns: + str | None: Authorization header 值,不存在时为 None + """ + for key, value in request.headers.items(): + if key.lower() == "authorization": + return value + return None + + def verify_access_token( request: fastapi.Request | fastapi.WebSocket, access_token: str | None ) -> bool: @@ -16,6 +38,7 @@ def verify_access_token( """ if access_token is None: return True - if "Authorization" in request.headers.keys(): - return request.headers["Authorization"] == f"Bearer {access_token}" + authorization = get_bearer_token(request) + if authorization is not None: + return authorization == f"Bearer {access_token}" return request.query_params.get("access_token") == access_token diff --git a/network/v11/http.py b/network/v11/http.py index 44f2b43..20170c8 100644 --- a/network/v11/http.py +++ b/network/v11/http.py @@ -1,4 +1,4 @@ -from network.authentication import verify_access_token +from network.authentication import get_bearer_token, verify_access_token import utils.uvicorn_server as uvicorn_server import fastapi from utils.logger import get_logger @@ -17,7 +17,7 @@ def __init__(self, config: dict) -> None: self.check_access_token() def check_access_token(self) -> None: - if self.config["host"] == "0.0.0.0" or self.config["access_token"]: + if not self.config["access_token"]: logger.warning( f'[{self.config["host"]}:{self.config["port"]}] 未配置 Access Token !' ) @@ -31,9 +31,7 @@ async def handle_request( self, request: fastapi.Request ) -> fastapi.responses.JSONResponse: if not verify_access_token(request, self.config["access_token"]): - if "Authorization" in request.headers.keys() or request.query_params.get( - "access_token" - ): + if get_bearer_token(request) or request.query_params.get("access_token"): raise fastapi.HTTPException(status_code=403, detail="Forbidden") else: raise fastapi.HTTPException(status_code=401, detail="Unauthorized") diff --git a/network/v11/ws.py b/network/v11/ws.py index 5e3e5f3..78b64ba 100644 --- a/network/v11/ws.py +++ b/network/v11/ws.py @@ -1,6 +1,6 @@ import utils.translator as translator import utils.event as event -from network.authentication import verify_access_token +from network.authentication import get_bearer_token, verify_access_token from utils.logger import get_logger import utils.uvicorn_server as uvicorn_server import fastapi @@ -22,7 +22,7 @@ def __init__(self, config: dict) -> None: self.check_access_token() def check_access_token(self) -> None: - if self.config["host"] == "0.0.0.0" or self.config["access_token"]: + if not self.config["access_token"]: logger.warning( f'[{self.config["host"]}:{self.config["port"]}] 未配置 Access Token !' ) @@ -35,7 +35,7 @@ async def start(self) -> None: async def handle_event_route(self, websocket: fastapi.WebSocket) -> None: if not verify_access_token(websocket, self.config["access_token"]): if ( - "Authorization" in websocket.headers.keys() + get_bearer_token(websocket) or websocket.query_params.get("access_token") ): await websocket.close(403, "Invalid access token") @@ -53,7 +53,7 @@ async def handle_event_route(self, websocket: fastapi.WebSocket) -> None: async def handle_api_route(self, websocket: fastapi.WebSocket) -> None: if not verify_access_token(websocket, self.config["access_token"]): if ( - "Authorization" in websocket.headers.keys() + get_bearer_token(websocket) or websocket.query_params.get("access_token") ): await websocket.close(403, "Invalid access token") @@ -73,7 +73,7 @@ async def handle_api_route(self, websocket: fastapi.WebSocket) -> None: async def handle_root_route(self, websocket: fastapi.WebSocket) -> None: if not verify_access_token(websocket, self.config["access_token"]): if ( - "Authorization" in websocket.headers.keys() + get_bearer_token(websocket) or websocket.query_params.get("access_token") ): await websocket.close(403, "Invalid access token") diff --git a/network/v12/http.py b/network/v12/http.py index a7f9bb7..b5c9c59 100644 --- a/network/v12/http.py +++ b/network/v12/http.py @@ -40,7 +40,7 @@ def __init__(self, config: dict) -> None: self.check_access_token() def check_access_token(self) -> None: - if self.config["host"] == "0.0.0.0" or self.config["access_token"]: + if not self.config["access_token"]: logger.warning( f'[HTTP {self.config["host"]}:{self.config["port"]}] 未配置 Access Token !' ) @@ -61,7 +61,7 @@ async def handle_http_connection( dict: 返回值 """ logger.debug(request) - if verify_access_token(request, self.config["access_token"]): + if not verify_access_token(request, self.config["access_token"]): raise fastapi.HTTPException(fastapi.status.HTTP_401_UNAUTHORIZED) logger.debug(await request.body()) return fastapi.responses.JSONResponse( diff --git a/network/v12/ws.py b/network/v12/ws.py index f001843..f2970d2 100644 --- a/network/v12/ws.py +++ b/network/v12/ws.py @@ -27,7 +27,7 @@ def __init__(self, config: dict) -> None: self.check_access_token() def check_access_token(self) -> None: - if self.config["host"] == "0.0.0.0" or self.config["access_token"]: + if not self.config["access_token"]: logger.warning( f'[{self.config["host"]}:{self.config["port"]}] 未配置 Access Token !' ) diff --git a/pyproject.toml b/pyproject.toml index 662c7fe..c9a9ccf 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "onedisc" -version = "1.0.0" +version = "1.0.1" description = "OneBot implement for Discord" authors = [ {name = "XiaoDeng3386",email = "1744793737@qq.com"}