This is an automated email from the ASF dual-hosted git repository.

FreeOnePlus pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/doris-mcp-server.git


The following commit(s) were added to refs/heads/master by this push:
     new 10a0d30  refactor: unify bearer credentials (#115)
10a0d30 is described below

commit 10a0d30e45c3e50feb7548e15fba20a3634415ed
Author: Yijia Su <[email protected]>
AuthorDate: Wed Jul 29 20:28:14 2026 +0800

    refactor: unify bearer credentials (#115)
---
 CHANGELOG.md                                  |   3 +
 doris_mcp_server/auth/auth_middleware.py      |  90 +++++-----
 doris_mcp_server/auth/doris_oauth_provider.py |  36 ++--
 doris_mcp_server/auth/mcp_auth_middleware.py  |  26 ++-
 doris_mcp_server/utils/auth_credentials.py    |  95 +++++++++++
 doris_mcp_server/utils/security.py            | 199 +++++++++++----------
 test/auth/test_doris_oauth_routes.py          |  30 +++-
 test/security/test_auth_context.py            |   5 +-
 test/security/test_bearer_credentials.py      | 237 ++++++++++++++++++++++++++
 test/security/test_mcp_auth_middleware.py     |  14 +-
 10 files changed, 550 insertions(+), 185 deletions(-)

diff --git a/CHANGELOG.md b/CHANGELOG.md
index 70f4daa..ceca727 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -54,6 +54,9 @@ under **Unreleased** until a new version is selected and 
published.
   management token modes now require explicit high-entropy credentials.
 - Refused unauthenticated HTTP startup on non-loopback bind addresses unless
   the operator sets an explicit dangerous override.
+- Normalized bearer credentials once at the MCP authentication boundary and
+  passed the same redacted DTO to static token, JWT, external OAuth, and Doris
+  OAuth providers.
 - Released Doris connections on SQL profile, data freshness, and access
   analysis paths.
 - Improved Doris 4 role metadata compatibility and query recovery behavior.
diff --git a/doris_mcp_server/auth/auth_middleware.py 
b/doris_mcp_server/auth/auth_middleware.py
index c454295..e358063 100644
--- a/doris_mcp_server/auth/auth_middleware.py
+++ b/doris_mcp_server/auth/auth_middleware.py
@@ -20,12 +20,15 @@ Authentication Middleware Module
 Provides middleware for JWT authentication in HTTP and MCP contexts
 """
 
-from typing import Optional, Dict, Any, Callable, Awaitable
 from datetime import datetime
 
-from .jwt_manager import JWTManager
-from ..utils.security import AuthContext, SecurityLevel
+from ..utils.auth_credentials import (
+    BearerCredentials,
+    normalize_bearer_credentials,
+)
 from ..utils.logger import get_logger
+from ..utils.security import AuthContext, SecurityLevel
+from .jwt_manager import JWTManager
 
 logger = get_logger(__name__)
 
@@ -45,7 +48,7 @@ class AuthMiddleware:
         self.jwt_manager = jwt_manager
         logger.info("AuthMiddleware initialized")
     
-    def extract_token_from_header(self, authorization: str) -> Optional[str]:
+    def extract_token_from_header(self, authorization: str) -> str | None:
         """Extract JWT token from Authorization header
         
         Args:
@@ -54,24 +57,17 @@ class AuthMiddleware:
         Returns:
             JWT token string, or None if not found
         """
-        if not authorization:
-            return None
-        
-        # Support Bearer format
-        if authorization.startswith('Bearer '):
-            return authorization[7:]  # Remove "Bearer " prefix
-        
-        # Support direct token format
-        if not authorization.startswith('Basic '):
-            return authorization
-        
-        return None
+        credentials = BearerCredentials.from_authorization(authorization)
+        return credentials.token or None
     
-    async def authenticate_request(self, auth_info: Dict[str, Any]) -> 
AuthContext:
+    async def authenticate_request(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """Authenticate request and return authentication context
         
         Args:
-            auth_info: Authentication information dictionary
+            credentials: Normalized bearer credentials
             
         Returns:
             AuthContext authentication context
@@ -80,35 +76,26 @@ class AuthMiddleware:
             ValueError: Authentication failed
         """
         try:
-            auth_type = auth_info.get("type", "jwt")
-            
-            if auth_type == "jwt" or auth_type == "token":
-                return await self._authenticate_jwt(auth_info)
-            else:
-                raise ValueError(f"Unsupported authentication type: 
{auth_type}")
-                
+            return await self._authenticate_jwt(credentials)
         except Exception as e:
             logger.error(f"Request authentication failed: {e}")
             raise
     
-    async def _authenticate_jwt(self, auth_info: Dict[str, Any]) -> 
AuthContext:
+    async def _authenticate_jwt(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """JWT authentication processing
         
         Args:
-            auth_info: Authentication information containing JWT token
+            credentials: Normalized bearer credentials
             
         Returns:
             AuthContext authentication context
         """
-        # Get token
-        token = auth_info.get("token")
-        if not token:
-            # Try to get from Authorization header
-            authorization = auth_info.get("authorization")
-            token = self.extract_token_from_header(authorization)
-        
-        if not token:
+        if not credentials.is_bearer:
             raise ValueError("Missing JWT token")
+        token = credentials.token
         
         try:
             # Validate token
@@ -137,7 +124,10 @@ class AuthMiddleware:
             logger.error(f"JWT authentication failed: {e}")
             raise ValueError(f"JWT authentication failed: {str(e)}")
     
-    async def create_auth_response_headers(self, auth_context: AuthContext) -> 
Dict[str, str]:
+    async def create_auth_response_headers(
+        self,
+        auth_context: AuthContext,
+    ) -> dict[str, str]:
         """Create authentication response headers
         
         Args:
@@ -153,7 +143,7 @@ class AuthMiddleware:
             'X-Auth-Security-Level': auth_context.security_level.value
         }
     
-    def create_http_middleware(self, skip_paths: Optional[list] = None):
+    def create_http_middleware(self, skip_paths: list | None = None):
         """Create HTTP middleware function
         
         Args:
@@ -182,11 +172,8 @@ class AuthMiddleware:
             
             try:
                 # Perform authentication
-                auth_info = {
-                    'type': 'jwt',
-                    'authorization': authorization
-                }
-                auth_context = await self.authenticate_request(auth_info)
+                credentials = 
BearerCredentials.from_authorization(authorization)
+                auth_context = await self.authenticate_request(credentials)
                 
                 # Add authentication context to scope
                 scope['auth_context'] = auth_context
@@ -225,7 +212,10 @@ class AuthMiddleware:
         
         return middleware
     
-    async def authenticate_mcp_request(self, headers: Dict[str, str]) -> 
AuthContext:
+    async def authenticate_mcp_request(
+        self,
+        headers: dict[str, str],
+    ) -> AuthContext:
         """Authenticate MCP request
         
         Args:
@@ -243,12 +233,14 @@ class AuthMiddleware:
                 headers.get('x-auth-token')
             )
             
-            auth_info = {
-                'type': 'jwt',
-                'authorization': authorization
-            }
-            
-            return await self.authenticate_request(auth_info)
+            credentials = normalize_bearer_credentials(
+                {
+                    "authorization": authorization,
+                    "token": headers.get("X-Auth-Token")
+                    or headers.get("x-auth-token"),
+                }
+            )
+            return await self.authenticate_request(credentials)
             
         except Exception as e:
             logger.error(f"MCP request authentication failed: {e}")
diff --git a/doris_mcp_server/auth/doris_oauth_provider.py 
b/doris_mcp_server/auth/doris_oauth_provider.py
index 4d78cd9..70be2fe 100644
--- a/doris_mcp_server/auth/doris_oauth_provider.py
+++ b/doris_mcp_server/auth/doris_oauth_provider.py
@@ -9,16 +9,21 @@ import time
 from datetime import UTC, datetime
 from urllib.parse import urlencode
 
+from ..utils.auth_credentials import BearerCredentials
 from ..utils.logger import get_logger
-from ..utils.security import AuthContext, RESERVED_DORIS_OAUTH_TOKEN_PREFIX, 
SecurityLevel
-from .doris_oauth_redirects import DorisOAuthRedirectPolicy, is_loopback_url
+from ..utils.security import (
+    RESERVED_DORIS_OAUTH_TOKEN_PREFIX,
+    AuthContext,
+    SecurityLevel,
+)
 from .doris_oauth_rate_limit import DorisOAuthRateLimiter
+from .doris_oauth_redirects import DorisOAuthRedirectPolicy, is_loopback_url
 from .doris_oauth_scope_policy import DorisOAuthScopePolicy
 from .doris_oauth_store import DorisOAuthStore
 from .doris_oauth_types import (
     AccessTokenRecord,
-    AuthTransactionRecord,
     AuthorizeError,
+    AuthTransactionRecord,
     ProtectedResourceAuthError,
     RefreshTokenRecord,
     RevocationEndpointError,
@@ -349,10 +354,16 @@ class DorisOAuthProvider:
             )
         return self._token_response(pair.access_token, pair.refresh_token, 
scopes)
 
-    async def authenticate_access_token(self, auth_info: dict) -> AuthContext:
-        token = self._extract_bearer(auth_info)
-        if not token or not 
token.startswith(RESERVED_DORIS_OAUTH_TOKEN_PREFIX):
+    async def authenticate_access_token(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
+        if (
+            not credentials.is_bearer
+            or not 
credentials.token.startswith(RESERVED_DORIS_OAUTH_TOKEN_PREFIX)
+        ):
             raise ProtectedResourceAuthError("authentication_required", 
"Missing Doris OAuth access token")
+        token = credentials.token
         record = self.store.get_access_token(token)
         if not record or record.revoked_at is not None:
             raise ProtectedResourceAuthError("authentication_required", 
"Invalid Doris OAuth access token")
@@ -376,8 +387,8 @@ class DorisOAuthProvider:
             roles=["doris_oauth_user"],
             permissions=["read_data"],
             security_level=SecurityLevel.INTERNAL,
-            client_ip=auth_info.get("client_ip", "unknown"),
-            session_id=auth_info.get("session_id") or 
f"doris_oauth:{updated.token_id}",
+            client_ip=credentials.client_ip,
+            session_id=credentials.session_id or 
f"doris_oauth:{updated.token_id}",
             login_time=login_time,
             last_activity=last_activity,
             token="",
@@ -461,15 +472,6 @@ class DorisOAuthProvider:
         challenge = 
base64.urlsafe_b64encode(digest).decode("ascii").rstrip("=")
         return secrets.compare_digest(challenge, expected_challenge)
 
-    def _extract_bearer(self, auth_info: dict) -> str:
-        token = auth_info.get("token") or ""
-        if token:
-            return str(token)
-        authorization = auth_info.get("authorization") or ""
-        if authorization.startswith("Bearer "):
-            return authorization[7:]
-        return ""
-
     async def _cleanup_inactive_pools(self) -> None:
         if self.connection_manager and hasattr(self.connection_manager, 
"cleanup_idle_doris_user_pools"):
             await 
self.connection_manager.cleanup_idle_doris_user_pools(self.store.active_users())
diff --git a/doris_mcp_server/auth/mcp_auth_middleware.py 
b/doris_mcp_server/auth/mcp_auth_middleware.py
index 25ac8c3..fcf5140 100644
--- a/doris_mcp_server/auth/mcp_auth_middleware.py
+++ b/doris_mcp_server/auth/mcp_auth_middleware.py
@@ -22,6 +22,7 @@ from typing import Any
 
 from starlette.responses import JSONResponse
 
+from ..utils.auth_credentials import BearerCredentials
 from ..utils.config import EffectiveAuthConfig
 from ..utils.logger import get_logger
 from ..utils.security import (
@@ -40,23 +41,20 @@ ASGIApp = Callable[[dict[str, Any], Callable[..., 
Awaitable[Any]], Callable[...,
 logger = get_logger(__name__)
 
 
-async def extract_auth_info_from_scope(scope: dict[str, Any]) -> dict[str, 
Any]:
-    """Extract auth info from ASGI scope."""
+async def extract_bearer_credentials_from_scope(
+    scope: dict[str, Any],
+) -> BearerCredentials:
+    """Extract canonical bearer credentials from an ASGI scope."""
     headers = dict(scope.get("headers", []))
     authorization = headers.get(b"authorization", b"").decode("utf-8")
     client = scope.get("client") or ("unknown", 0)
     client_ip = client[0] if client else "unknown"
 
-    auth_info = {
-        "authorization": authorization,
-        "client_ip": client_ip,
-        "session_id": scope.get("session_id", ""),
-    }
-    if authorization.startswith("Bearer "):
-        auth_info["token"] = authorization[7:]
-    elif authorization.startswith("Token "):
-        auth_info["token"] = authorization[6:]
-    return auth_info
+    return BearerCredentials.from_authorization(
+        authorization,
+        client_ip=client_ip,
+        session_id=scope.get("session_id", ""),
+    )
 
 
 class MCPAuthASGIMiddleware:
@@ -69,8 +67,8 @@ class MCPAuthASGIMiddleware:
 
     async def __call__(self, scope, receive, send):
         try:
-            auth_info = await extract_auth_info_from_scope(scope)
-            auth_context = await 
self.security_manager.authenticate_request(auth_info)
+            credentials = await extract_bearer_credentials_from_scope(scope)
+            auth_context = await 
self.security_manager.authenticate_request(credentials)
         except Exception as exc:
             if self.effective_auth.oauth_discovery_mode == "doris_oauth":
                 response = protected_resource_error_response(
diff --git a/doris_mcp_server/utils/auth_credentials.py 
b/doris_mcp_server/utils/auth_credentials.py
new file mode 100644
index 0000000..80180ea
--- /dev/null
+++ b/doris_mcp_server/utils/auth_credentials.py
@@ -0,0 +1,95 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""Canonical request credentials for MCP bearer authentication."""
+
+from collections.abc import Mapping
+from dataclasses import dataclass, field
+from typing import Any
+
+
+@dataclass(frozen=True, slots=True)
+class BearerCredentials:
+    """Normalized credentials passed to every bearer authentication 
provider."""
+
+    scheme: str = ""
+    token: str = field(default="", repr=False)
+    client_ip: str = "unknown"
+    session_id: str = ""
+
+    @property
+    def is_bearer(self) -> bool:
+        return self.scheme == "bearer" and bool(self.token)
+
+    @property
+    def is_static_token(self) -> bool:
+        return self.scheme in {"bearer", "token"} and bool(self.token)
+
+    @classmethod
+    def from_authorization(
+        cls,
+        authorization: str | None,
+        *,
+        client_ip: str = "unknown",
+        session_id: str = "",
+    ) -> "BearerCredentials":
+        """Normalize an Authorization header without retaining the raw 
header."""
+        raw_header = str(authorization or "").strip()
+        scheme, separator, token = raw_header.partition(" ")
+        normalized_scheme = scheme.lower()
+        if (
+            not separator
+            or normalized_scheme not in {"bearer", "token"}
+            or not token.strip()
+        ):
+            token = ""
+        return cls(
+            scheme=normalized_scheme,
+            token=token.strip(),
+            client_ip=str(client_ip or "unknown"),
+            session_id=str(session_id or ""),
+        )
+
+
+def normalize_bearer_credentials(
+    auth_input: BearerCredentials | Mapping[str, Any],
+) -> BearerCredentials:
+    """Convert a legacy auth mapping once at the authentication boundary."""
+    if isinstance(auth_input, BearerCredentials):
+        return auth_input
+
+    client_ip = str(auth_input.get("client_ip") or "unknown")
+    session_id = str(auth_input.get("session_id") or "")
+    credentials = BearerCredentials.from_authorization(
+        auth_input.get("authorization"),
+        client_ip=client_ip,
+        session_id=session_id,
+    )
+    explicit_token = str(
+        auth_input.get("token") or auth_input.get("access_token") or ""
+    ).strip()
+    if not explicit_token:
+        return credentials
+
+    scheme = credentials.scheme
+    if scheme not in {"bearer", "token"}:
+        scheme = "bearer"
+    return BearerCredentials(
+        scheme=scheme,
+        token=explicit_token,
+        client_ip=client_ip,
+        session_id=session_id,
+    )
diff --git a/doris_mcp_server/utils/security.py 
b/doris_mcp_server/utils/security.py
index 6b99f08..9b490d0 100644
--- a/doris_mcp_server/utils/security.py
+++ b/doris_mcp_server/utils/security.py
@@ -22,7 +22,9 @@ Implements enterprise-level authentication, authorization, 
SQL security validati
 
 import logging
 import re
-from contextvars import ContextVar, Token as ContextToken
+from collections.abc import Mapping
+from contextvars import ContextVar
+from contextvars import Token as ContextToken
 from dataclasses import dataclass, field
 from datetime import datetime
 from enum import Enum
@@ -32,11 +34,12 @@ import sqlparse
 from sqlparse.sql import Statement
 from sqlparse.tokens import Keyword, Name
 
-from .logger import get_logger
+from .auth_credentials import BearerCredentials, normalize_bearer_credentials
 from .config import (
     DatabaseConfig,
     get_effective_auth_config,
 )
+from .logger import get_logger
 
 # Global ContextVar for auth_context - must be a single instance shared across 
all modules
 # This allows token-bound database configuration to work correctly in 
concurrent requests
@@ -257,7 +260,10 @@ class DorisSecurityManager:
 
         return default_rules
 
-    async def authenticate_request(self, auth_info: dict[str, Any]) -> 
AuthContext:
+    async def authenticate_request(
+        self,
+        auth_input: BearerCredentials | Mapping[str, Any],
+    ) -> AuthContext:
         """Validate request authentication information
         
         Tries authentication methods in normalized effective config order.
@@ -265,14 +271,14 @@ class DorisSecurityManager:
         If all methods are disabled, returns anonymous context
         """
         effective_auth = self._get_effective_auth_config()
-        bearer_token = str(auth_info.get("token") or "")
-        authorization = str(auth_info.get("authorization") or "")
-        if not bearer_token and authorization.startswith("Bearer "):
-            bearer_token = authorization[7:]
+        credentials = normalize_bearer_credentials(auth_input)
+        legacy_auth_type = ""
+        if isinstance(auth_input, Mapping):
+            legacy_auth_type = str(auth_input.get("type") or "")
 
         if not effective_auth.auth_methods:
-            if auth_info.get("type"):
-                return await self.auth_provider.authenticate(auth_info)
+            if legacy_auth_type:
+                return await self.auth_provider.authenticate(dict(auth_input))
             self.logger.debug("All authentication methods are disabled")
             return AuthContext(
                 token_id="anonymous",
@@ -280,7 +286,7 @@ class DorisSecurityManager:
                 roles=["anonymous"],
                 permissions=["read"],
                 security_level=SecurityLevel.PUBLIC,
-                client_ip=auth_info.get("client_ip", "unknown"),
+                client_ip=credentials.client_ip,
                 session_id="anonymous_session",
                 auth_method="anonymous",
                 pool_key="global",
@@ -291,22 +297,29 @@ class DorisSecurityManager:
         for auth_method in effective_auth.auth_methods:
             try:
                 if auth_method == "doris_oauth":
-                    return await 
self.auth_provider.authenticate_doris_oauth(auth_info)
+                    return await 
self.auth_provider.authenticate_doris_oauth(credentials)
                 if auth_method == "token":
-                    return await 
self.auth_provider.authenticate_token(auth_info)
+                    return await 
self.auth_provider.authenticate_token(credentials)
                 if auth_method == "jwt":
-                    return await self.auth_provider.authenticate_jwt(auth_info)
+                    return await 
self.auth_provider.authenticate_jwt(credentials)
                 if auth_method == "external_oauth":
-                    return await 
self.auth_provider.authenticate_oauth(auth_info)
+                    return await 
self.auth_provider.authenticate_oauth(credentials)
             except Exception as e:
                 self.logger.debug(f"{auth_method} authentication failed: {e}")
                 last_error = e
-                if auth_method == "doris_oauth" and 
bearer_token.startswith(RESERVED_DORIS_OAUTH_TOKEN_PREFIX):
+                if (
+                    auth_method == "doris_oauth"
+                    and credentials.token.startswith(
+                        RESERVED_DORIS_OAUTH_TOKEN_PREFIX
+                    )
+                ):
                     raise
         
         # All enabled authentication methods failed
         error_message = f"Authentication failed: {str(last_error)}" if 
last_error else "No authentication method succeeded"
-        self.logger.warning(f"Authentication failed for client 
{auth_info.get('client_ip', 'unknown')}: {error_message}")
+        self.logger.warning(
+            f"Authentication failed for client {credentials.client_ip}: 
{error_message}"
+        )
         raise ValueError(error_message)
 
     async def authorize_resource_access(
@@ -617,117 +630,125 @@ class AuthenticationProvider:
         that pass an explicit auth_info["type"].
         """
         auth_type = str(auth_info.get("type") or "").strip().lower()
+        credentials = normalize_bearer_credentials(auth_info)
         if auth_type == "token":
             if self.effective_auth.enable_token_auth and self.token_manager:
-                return await self.authenticate_token(auth_info)
-            return await self._authenticate_legacy_token(auth_info)
+                return await self.authenticate_token(credentials)
+            return await self._authenticate_legacy_token(credentials)
         if auth_type == "basic":
             return await self._authenticate_basic(auth_info)
         if auth_type == "jwt":
-            return await self.authenticate_jwt(auth_info)
+            return await self.authenticate_jwt(credentials)
         if auth_type == "oauth":
-            return await self.authenticate_oauth(auth_info)
+            if "code" in auth_info and "state" in auth_info:
+                if not self.effective_auth.enable_external_oauth_auth:
+                    raise ValueError("OAuth authentication is not enabled")
+                if not self.oauth_provider:
+                    raise ValueError("OAuth provider not initialized")
+                auth_context = await self.oauth_provider.handle_callback(
+                    auth_info["code"],
+                    auth_info["state"],
+                )
+                auth_context.auth_method = "external_oauth"
+                auth_context.token = ""
+                auth_context.pool_key = "global"
+                return auth_context
+            return await self.authenticate_oauth(credentials)
         if auth_type == "doris_oauth":
-            return await self.authenticate_doris_oauth(auth_info)
+            return await self.authenticate_doris_oauth(credentials)
         raise ValueError(f"Unsupported authentication type: {auth_type or 
'<missing>'}")
 
-    async def authenticate_token(self, auth_info: dict[str, Any]) -> 
AuthContext:
+    async def authenticate_token(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """Perform token authentication"""
         if not self.effective_auth.enable_token_auth:
             raise ValueError("Token authentication is not enabled")
-        return await self._authenticate_token(auth_info)
+        return await self._authenticate_token(credentials)
     
-    async def authenticate_jwt(self, auth_info: dict[str, Any]) -> AuthContext:
+    async def authenticate_jwt(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """Perform JWT authentication"""
         if not self.effective_auth.enable_jwt_auth:
             raise ValueError("JWT authentication is not enabled")
-        return await self._authenticate_jwt(auth_info)
+        return await self._authenticate_jwt(credentials)
     
-    async def authenticate_oauth(self, auth_info: dict[str, Any]) -> 
AuthContext:
+    async def authenticate_oauth(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """Perform OAuth authentication"""
         if not self.effective_auth.enable_external_oauth_auth:
             raise ValueError("OAuth authentication is not enabled")
-        return await self._authenticate_oauth(auth_info)
+        return await self._authenticate_oauth(credentials)
 
-    async def authenticate_doris_oauth(self, auth_info: dict[str, Any]) -> 
AuthContext:
+    async def authenticate_doris_oauth(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """Authenticate a Doris OAuth doa_ access token."""
         if not self.effective_auth.enable_doris_oauth_auth:
             raise ValueError("Doris OAuth authentication is not enabled")
-        token = auth_info.get("token")
-        if not token:
-            authorization = auth_info.get("authorization")
-            if authorization and authorization.startswith("Bearer "):
-                token = authorization[7:]
-        if not token or not token.startswith("doa_"):
+        if (
+            not credentials.is_bearer
+            or not 
credentials.token.startswith(RESERVED_DORIS_OAUTH_TOKEN_PREFIX)
+        ):
             raise ValueError("Missing Doris OAuth access token")
         if not self.doris_oauth_provider:
             raise ValueError("Doris OAuth provider is not initialized")
-        return await 
self.doris_oauth_provider.authenticate_access_token(auth_info)
+        return await 
self.doris_oauth_provider.authenticate_access_token(credentials)
 
-    async def _authenticate_jwt(self, auth_info: dict[str, Any]) -> 
AuthContext:
+    async def _authenticate_jwt(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """JWT authentication"""
         if not self.jwt_manager:
             raise ValueError("JWT manager not initialized")
-        
-        token = auth_info.get("token")
-        if not token:
-            # Try to extract from Authorization header
-            authorization = auth_info.get("authorization")
-            if authorization and authorization.startswith('Bearer '):
-                token = authorization[7:]
-            
-        if not token:
+        if not credentials.is_bearer:
             raise ValueError("Missing JWT token")
 
         try:
             # Use JWT middleware for authentication
             from ..auth.auth_middleware import AuthMiddleware
             middleware = AuthMiddleware(self.jwt_manager)
-            return await middleware.authenticate_request(auth_info)
+            return await middleware.authenticate_request(credentials)
             
         except Exception as e:
             self.logger.error(f"JWT authentication failed: {e}")
             raise ValueError(f"JWT authentication failed: {str(e)}")
 
-    async def _authenticate_oauth(self, auth_info: dict[str, Any]) -> 
AuthContext:
+    async def _authenticate_oauth(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """OAuth authentication"""
         if not self.oauth_provider:
             raise ValueError("OAuth provider not initialized")
-        
-        # Handle different OAuth authentication scenarios
-        if "access_token" in auth_info:
-            # Direct OAuth access token authentication
-            auth_context = await 
self.oauth_provider.authenticate_with_token(auth_info["access_token"])
-            auth_context.auth_method = "external_oauth"
-            auth_context.token = ""
-            auth_context.pool_key = "global"
-            return auth_context
-        elif "code" in auth_info and "state" in auth_info:
-            # OAuth callback authentication
-            auth_context = await 
self.oauth_provider.handle_callback(auth_info["code"], auth_info["state"])
-            auth_context.auth_method = "external_oauth"
-            auth_context.token = ""
-            auth_context.pool_key = "global"
-            return auth_context
-        else:
-            raise ValueError("OAuth authentication requires either 
access_token or code+state")
+        if not credentials.is_bearer:
+            raise ValueError("Missing external OAuth access token")
 
-    async def _authenticate_token(self, auth_info: dict[str, Any]) -> 
AuthContext:
+        auth_context = await self.oauth_provider.authenticate_with_token(
+            credentials.token
+        )
+        auth_context.auth_method = "external_oauth"
+        auth_context.token = ""
+        auth_context.pool_key = "global"
+        return auth_context
+
+    async def _authenticate_token(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """Token authentication"""
         if not self.token_manager:
             raise ValueError("Token manager not initialized")
-            
-        token = auth_info.get("token")
-        if not token:
-            # Try to extract from Authorization header
-            authorization = auth_info.get("authorization")
-            if authorization and authorization.startswith('Bearer '):
-                token = authorization[7:]
-            elif authorization and authorization.startswith('Token '):
-                token = authorization[6:]
-                
-        if not token:
+        if not credentials.is_static_token:
             raise ValueError("Missing authentication token")
+        token = credentials.token
 
         try:
             # Validate token using TokenManager
@@ -748,8 +769,8 @@ class AuthenticationProvider:
                 roles=["token_user"],  # Default role for token users
                 permissions=["read", "write"],  # Default permissions for 
token users
                 security_level=SecurityLevel.INTERNAL,
-                client_ip=auth_info.get("client_ip", "unknown"),
-                session_id=auth_info.get("session_id", 
f"session_{token_info.token_id}"),
+                client_ip=credentials.client_ip,
+                session_id=credentials.session_id or 
f"session_{token_info.token_id}",
                 login_time=datetime.utcnow(),
                 last_activity=token_info.last_used,
                 token=token,  # Store raw token for token-bound database 
configuration
@@ -761,18 +782,14 @@ class AuthenticationProvider:
             self.logger.error(f"Token authentication failed: {e}")
             raise ValueError(f"Token authentication failed: {str(e)}")
 
-    async def _authenticate_legacy_token(self, auth_info: dict[str, Any]) -> 
AuthContext:
+    async def _authenticate_legacy_token(
+        self,
+        credentials: BearerCredentials,
+    ) -> AuthContext:
         """Token authentication for legacy direct callers without 
TokenManager."""
-        token = auth_info.get("token")
-        if not token:
-            authorization = auth_info.get("authorization")
-            if authorization and authorization.startswith("Bearer "):
-                token = authorization[7:]
-            elif authorization and authorization.startswith("Token "):
-                token = authorization[6:]
-
-        if not token:
+        if not credentials.is_static_token:
             raise ValueError("Missing authentication token")
+        token = credentials.token
 
         user_info = await self._validate_token(token)
         return AuthContext(
@@ -781,8 +798,8 @@ class AuthenticationProvider:
             roles=user_info["roles"],
             permissions=user_info["permissions"],
             security_level=user_info["security_level"],
-            client_ip=auth_info.get("client_ip", "unknown"),
-            session_id=auth_info.get("session_id", 
f"session_{user_info['user_id']}"),
+            client_ip=credentials.client_ip,
+            session_id=credentials.session_id or 
f"session_{user_info['user_id']}",
             login_time=datetime.utcnow(),
             auth_method="token",
             token=token,
diff --git a/test/auth/test_doris_oauth_routes.py 
b/test/auth/test_doris_oauth_routes.py
index bc511dd..189e003 100644
--- a/test/auth/test_doris_oauth_routes.py
+++ b/test/auth/test_doris_oauth_routes.py
@@ -11,9 +11,16 @@ from starlette.applications import Starlette
 
 from doris_mcp_server.auth.doris_oauth_handlers import DorisOAuthHandlers
 from doris_mcp_server.auth.doris_oauth_provider import DorisOAuthProvider
-from doris_mcp_server.auth.doris_oauth_types import 
ProtectedResourceAuthError, TokenEndpointError
-from doris_mcp_server.utils.config import DorisConfig, _mark_source, 
normalize_effective_auth_config
-
+from doris_mcp_server.auth.doris_oauth_types import (
+    ProtectedResourceAuthError,
+    TokenEndpointError,
+)
+from doris_mcp_server.utils.auth_credentials import BearerCredentials
+from doris_mcp_server.utils.config import (
+    DorisConfig,
+    _mark_source,
+    normalize_effective_auth_config,
+)
 
 FULL_DORIS_OAUTH_SCOPE_SET = tuple(
     sorted(
@@ -525,7 +532,11 @@ async def 
test_full_login_code_exchange_auth_context_and_pool_missing_revocation
         assert token_json["access_token"].startswith("doa_")
         assert token_json["scope"] == "resource:list resource:read tool:list"
 
-    auth_context = await provider.authenticate_access_token({"token": 
token_json["access_token"]})
+    credentials = BearerCredentials(
+        scheme="bearer",
+        token=token_json["access_token"],
+    )
+    auth_context = await provider.authenticate_access_token(credentials)
     assert auth_context.auth_method == "doris_oauth"
     assert auth_context.doris_user == "alice"
     assert auth_context.oauth_client_id == client_id
@@ -535,14 +546,14 @@ async def 
test_full_login_code_exchange_auth_context_and_pool_missing_revocation
 
     cm.pools["alice"] = False
     with pytest.raises(ProtectedResourceAuthError) as exc:
-        await provider.authenticate_access_token({"token": 
token_json["access_token"]})
+        await provider.authenticate_access_token(credentials)
     assert exc.value.error == "login_required"
     assert exc.value.error_code == "DORIS_OAUTH_POOL_MISSING"
     assert cm.global_acquire_calls == 0
 
     cm.pools["alice"] = True
     with pytest.raises(ProtectedResourceAuthError):
-        await provider.authenticate_access_token({"token": 
token_json["access_token"]})
+        await provider.authenticate_access_token(credentials)
 
 
 @pytest.mark.asyncio
@@ -610,7 +621,12 @@ async def 
test_full_login_without_scope_grants_configured_rbac_capability_envelo
     assert "scope:monitoring:read" not in token_json["scope"].split()
     assert "scope:adbc:execute" not in token_json["scope"].split()
 
-    auth_context = await provider.authenticate_access_token({"token": 
token_json["access_token"]})
+    auth_context = await provider.authenticate_access_token(
+        BearerCredentials(
+            scheme="bearer",
+            token=token_json["access_token"],
+        )
+    )
     assert auth_context.auth_method == "doris_oauth"
     assert auth_context.doris_user == "alice"
     assert tuple(auth_context.oauth_scopes) == FULL_DORIS_OAUTH_SCOPE_SET
diff --git a/test/security/test_auth_context.py 
b/test/security/test_auth_context.py
index dbeded9..bef80c0 100644
--- a/test/security/test_auth_context.py
+++ b/test/security/test_auth_context.py
@@ -3,13 +3,14 @@ from datetime import datetime
 import pytest
 
 from doris_mcp_server.auth.auth_middleware import AuthMiddleware
+from doris_mcp_server.utils import sql_security_utils
+from doris_mcp_server.utils.auth_credentials import BearerCredentials
 from doris_mcp_server.utils.security import (
     AuthContext,
     get_current_auth_context,
     reset_auth_context,
     set_current_auth_context,
 )
-from doris_mcp_server.utils import sql_security_utils
 
 
 def test_sql_security_utils_uses_shared_contextvar():
@@ -45,7 +46,7 @@ async def test_jwt_auth_context_does_not_store_raw_token():
 
     middleware = AuthMiddleware(FakeJWTManager())
     auth_context = await middleware.authenticate_request(
-        {"authorization": "Bearer jwt.raw.token"}
+        BearerCredentials(scheme="bearer", token="jwt.raw.token")
     )
 
     assert auth_context.auth_method == "jwt"
diff --git a/test/security/test_bearer_credentials.py 
b/test/security/test_bearer_credentials.py
new file mode 100644
index 0000000..d22ff2d
--- /dev/null
+++ b/test/security/test_bearer_credentials.py
@@ -0,0 +1,237 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+
+import logging
+from datetime import datetime
+from types import SimpleNamespace
+
+import pytest
+
+from doris_mcp_server.utils.auth_credentials import (
+    BearerCredentials,
+    normalize_bearer_credentials,
+)
+from doris_mcp_server.utils.security import (
+    AuthContext,
+    AuthenticationProvider,
+    DorisSecurityManager,
+)
+
+
+def test_authorization_header_is_normalized_and_token_is_not_represented():
+    credentials = BearerCredentials.from_authorization(
+        "bEaReR secret-value",
+        client_ip="192.0.2.10",
+        session_id="session-1",
+    )
+
+    assert credentials.scheme == "bearer"
+    assert credentials.token == "secret-value"
+    assert credentials.client_ip == "192.0.2.10"
+    assert credentials.session_id == "session-1"
+    assert credentials.is_bearer is True
+    assert "secret-value" not in repr(credentials)
+
+
[email protected](
+    ("authorization", "scheme"),
+    [
+        ("", ""),
+        ("Bearer", "bearer"),
+        ("Bearer   ", "bearer"),
+        ("Basic abc", "basic"),
+    ],
+)
+def test_invalid_or_non_bearer_headers_do_not_produce_a_token(
+    authorization,
+    scheme,
+):
+    credentials = BearerCredentials.from_authorization(authorization)
+
+    assert credentials.scheme == scheme
+    assert credentials.token == ""
+    assert credentials.is_bearer is False
+
+
+def test_legacy_mapping_is_normalized_once_to_the_canonical_dto():
+    credentials = normalize_bearer_credentials(
+        {
+            "access_token": "external-token",
+            "client_ip": "198.51.100.4",
+            "session_id": "session-2",
+        }
+    )
+
+    assert credentials == BearerCredentials(
+        scheme="bearer",
+        token="external-token",
+        client_ip="198.51.100.4",
+        session_id="session-2",
+    )
+    assert normalize_bearer_credentials(credentials) is credentials
+
+
[email protected]
[email protected](
+    ("method", "provider_method", "token"),
+    [
+        ("token", "authenticate_token", "static-token"),
+        ("jwt", "authenticate_jwt", "jwt-token"),
+        ("external_oauth", "authenticate_oauth", "external-token"),
+        ("doris_oauth", "authenticate_doris_oauth", "doa_access-token"),
+    ],
+)
+async def test_security_manager_passes_the_same_dto_to_every_provider(
+    method,
+    provider_method,
+    token,
+):
+    credentials = BearerCredentials(
+        scheme="bearer",
+        token=token,
+        client_ip="203.0.113.8",
+        session_id="session-3",
+    )
+    expected_context = AuthContext(user_id=method, auth_method=method)
+    received = []
+
+    class Provider:
+        async def authenticate_token(self, value):
+            received.append(("authenticate_token", value))
+            return expected_context
+
+        async def authenticate_jwt(self, value):
+            received.append(("authenticate_jwt", value))
+            return expected_context
+
+        async def authenticate_oauth(self, value):
+            received.append(("authenticate_oauth", value))
+            return expected_context
+
+        async def authenticate_doris_oauth(self, value):
+            received.append(("authenticate_doris_oauth", value))
+            return expected_context
+
+    manager = object.__new__(DorisSecurityManager)
+    manager.auth_provider = Provider()
+    manager.logger = logging.getLogger(__name__)
+    manager._get_effective_auth_config = lambda: 
SimpleNamespace(auth_methods=(method,))
+
+    result = await manager.authenticate_request(credentials)
+
+    assert result is expected_context
+    assert received == [(provider_method, credentials)]
+    assert received[0][1] is credentials
+
+
[email protected]
+async def test_static_token_provider_uses_canonical_credentials():
+    token_info = SimpleNamespace(
+        token_id="static-id",
+        last_used=datetime.utcnow(),
+    )
+
+    class TokenManager:
+        async def validate_token(self, token):
+            assert token == "static-token"
+            return SimpleNamespace(is_valid=True, token_info=token_info)
+
+    provider = object.__new__(AuthenticationProvider)
+    provider.token_manager = TokenManager()
+    provider.security_manager = None
+    provider.logger = logging.getLogger(__name__)
+
+    context = await provider._authenticate_token(
+        BearerCredentials(
+            scheme="token",
+            token="static-token",
+            client_ip="192.0.2.20",
+            session_id="static-session",
+        )
+    )
+
+    assert context.auth_method == "token"
+    assert context.token_id == "static-id"
+    assert context.client_ip == "192.0.2.20"
+    assert context.session_id == "static-session"
+
+
[email protected]
+async def test_external_oauth_provider_receives_normalized_bearer_token():
+    received = []
+
+    class OAuthProvider:
+        async def authenticate_with_token(self, token):
+            received.append(token)
+            return AuthContext(user_id="oauth-user")
+
+    provider = object.__new__(AuthenticationProvider)
+    provider.oauth_provider = OAuthProvider()
+
+    context = await provider._authenticate_oauth(
+        BearerCredentials(scheme="bearer", token="external-token")
+    )
+
+    assert received == ["external-token"]
+    assert context.auth_method == "external_oauth"
+    assert context.token == ""
+    assert context.pool_key == "global"
+
+
[email protected]
+async def test_legacy_external_oauth_callback_keeps_normalized_auth_context():
+    class OAuthProvider:
+        async def handle_callback(self, code, state):
+            assert (code, state) == ("code-1", "state-1")
+            return AuthContext(user_id="oauth-user")
+
+    provider = object.__new__(AuthenticationProvider)
+    provider.effective_auth = SimpleNamespace(enable_external_oauth_auth=True)
+    provider.oauth_provider = OAuthProvider()
+
+    context = await provider.authenticate(
+        {
+            "type": "oauth",
+            "code": "code-1",
+            "state": "state-1",
+        }
+    )
+
+    assert context.auth_method == "external_oauth"
+    assert context.token == ""
+    assert context.pool_key == "global"
+
+
[email protected]
+async def test_doris_oauth_provider_receives_the_same_credentials_object():
+    credentials = BearerCredentials(scheme="bearer", token="doa_access-token")
+    received = []
+
+    class DorisOAuthProvider:
+        async def authenticate_access_token(self, value):
+            received.append(value)
+            return AuthContext(user_id="doris-user", auth_method="doris_oauth")
+
+    provider = object.__new__(AuthenticationProvider)
+    provider.effective_auth = SimpleNamespace(enable_doris_oauth_auth=True)
+    provider.doris_oauth_provider = DorisOAuthProvider()
+
+    context = await provider.authenticate_doris_oauth(credentials)
+
+    assert context.auth_method == "doris_oauth"
+    assert received == [credentials]
+    assert received[0] is credentials
diff --git a/test/security/test_mcp_auth_middleware.py 
b/test/security/test_mcp_auth_middleware.py
index a4ceffb..c601b2d 100644
--- a/test/security/test_mcp_auth_middleware.py
+++ b/test/security/test_mcp_auth_middleware.py
@@ -5,6 +5,7 @@ import pytest
 import doris_mcp_server.auth.mcp_auth_middleware as middleware_module
 from doris_mcp_server.auth.mcp_auth_middleware import MCPAuthASGIMiddleware
 from doris_mcp_server.auth.operation_policy import OperationAuthorizationError
+from doris_mcp_server.utils.auth_credentials import BearerCredentials
 from doris_mcp_server.utils.config import EffectiveAuthConfig
 from doris_mcp_server.utils.security import AuthContext, 
get_current_auth_context
 
@@ -41,8 +42,12 @@ async def 
test_mcp_auth_middleware_sets_scope_and_resets_context():
     auth_context = AuthContext(token_id="t1", user_id="u1", 
auth_method="token")
 
     class SecurityManager:
-        async def authenticate_request(self, auth_info):
-            assert auth_info["token"] == "abc"
+        async def authenticate_request(self, credentials):
+            assert credentials == BearerCredentials(
+                scheme="bearer",
+                token="abc",
+                client_ip="127.0.0.1",
+            )
             return auth_context
 
     async def downstream(scope, receive, send):
@@ -71,9 +76,8 @@ async def 
test_mcp_auth_middleware_sets_scope_and_resets_context():
 @pytest.mark.asyncio
 async def test_mcp_auth_middleware_rejects_query_string_token():
     class SecurityManager:
-        async def authenticate_request(self, auth_info):
-            assert auth_info["authorization"] == ""
-            assert "token" not in auth_info
+        async def authenticate_request(self, credentials):
+            assert credentials == BearerCredentials(client_ip="127.0.0.1")
             raise ValueError("missing bearer token")
 
     async def downstream(scope, receive, send):


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to