fix: avoid websocket token leakage in logs
This commit is contained in:
@@ -17,6 +17,8 @@ from utils.security import verify_access_token
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
WS_AUTH_PROTOCOL = "access_token"
|
||||||
|
|
||||||
|
|
||||||
class ConnectionManager:
|
class ConnectionManager:
|
||||||
"""WebSocket 连接管理器"""
|
"""WebSocket 连接管理器"""
|
||||||
@@ -27,9 +29,14 @@ class ConnectionManager:
|
|||||||
# 组连接: {group_name: {websocket1, websocket2, ...}}
|
# 组连接: {group_name: {websocket1, websocket2, ...}}
|
||||||
self.groups: Dict[str, Set[WebSocket]] = {}
|
self.groups: Dict[str, Set[WebSocket]] = {}
|
||||||
|
|
||||||
async def connect(self, websocket: WebSocket, user_id: str):
|
async def connect(
|
||||||
|
self,
|
||||||
|
websocket: WebSocket,
|
||||||
|
user_id: str,
|
||||||
|
subprotocol: Optional[str] = None,
|
||||||
|
):
|
||||||
"""添加连接"""
|
"""添加连接"""
|
||||||
await websocket.accept()
|
await websocket.accept(subprotocol=subprotocol)
|
||||||
if user_id not in self.active_connections:
|
if user_id not in self.active_connections:
|
||||||
self.active_connections[user_id] = set()
|
self.active_connections[user_id] = set()
|
||||||
self.active_connections[user_id].add(websocket)
|
self.active_connections[user_id].add(websocket)
|
||||||
@@ -101,27 +108,45 @@ class TokenAuthWebSocketConsumer:
|
|||||||
self.user_id: Optional[str] = None
|
self.user_id: Optional[str] = None
|
||||||
self.is_authenticated = False
|
self.is_authenticated = False
|
||||||
self._token: Optional[str] = None # 保存原始token用于心跳校验
|
self._token: Optional[str] = None # 保存原始token用于心跳校验
|
||||||
|
self._accept_subprotocol: Optional[str] = None
|
||||||
|
|
||||||
|
def _get_token_from_protocol(self) -> Optional[str]:
|
||||||
|
protocol_header = self.websocket.headers.get("sec-websocket-protocol") or ""
|
||||||
|
protocols = [
|
||||||
|
item.strip()
|
||||||
|
for item in protocol_header.split(",")
|
||||||
|
if item.strip()
|
||||||
|
]
|
||||||
|
for index, protocol in enumerate(protocols):
|
||||||
|
if protocol == WS_AUTH_PROTOCOL and index + 1 < len(protocols):
|
||||||
|
self._accept_subprotocol = WS_AUTH_PROTOCOL
|
||||||
|
return protocols[index + 1]
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _get_token_from_query(self) -> Optional[str]:
|
||||||
|
query_string = self.websocket.scope.get('query_string', b'').decode('utf-8')
|
||||||
|
if not query_string:
|
||||||
|
return None
|
||||||
|
|
||||||
|
query_params = parse_qs(query_string)
|
||||||
|
token_list = query_params.get('token', [])
|
||||||
|
return token_list[0] if token_list else None
|
||||||
|
|
||||||
|
async def _accept_then_close(self, code: int):
|
||||||
|
await self.websocket.accept(subprotocol=self._accept_subprotocol)
|
||||||
|
await self.websocket.close(code=code)
|
||||||
|
|
||||||
async def authenticate(self) -> bool:
|
async def authenticate(self) -> bool:
|
||||||
"""
|
"""
|
||||||
进行Token认证
|
进行Token认证
|
||||||
从查询参数中获取token并验证
|
优先从 WebSocket 子协议中获取 token,兼容旧 query 参数方式。
|
||||||
"""
|
"""
|
||||||
# 获取查询参数中的token
|
token = self._get_token_from_protocol() or self._get_token_from_query()
|
||||||
query_string = self.websocket.scope.get('query_string', b'').decode('utf-8')
|
|
||||||
token = None
|
|
||||||
|
|
||||||
if query_string:
|
|
||||||
query_params = parse_qs(query_string)
|
|
||||||
token_list = query_params.get('token', [])
|
|
||||||
if token_list:
|
|
||||||
token = token_list[0]
|
|
||||||
|
|
||||||
if not token:
|
if not token:
|
||||||
logger.warning("WebSocket connection rejected: No token provided")
|
logger.warning("WebSocket connection rejected: No token provided")
|
||||||
# 必须先accept才能close
|
# 必须先accept才能close
|
||||||
await self.websocket.accept()
|
await self._accept_then_close(code=4001)
|
||||||
await self.websocket.close(code=4001)
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# 验证token
|
# 验证token
|
||||||
@@ -130,15 +155,13 @@ class TokenAuthWebSocketConsumer:
|
|||||||
|
|
||||||
if not payload:
|
if not payload:
|
||||||
logger.warning("WebSocket connection rejected: Invalid token")
|
logger.warning("WebSocket connection rejected: Invalid token")
|
||||||
await self.websocket.accept()
|
await self._accept_then_close(code=4001)
|
||||||
await self.websocket.close(code=4001)
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
user_id = payload.get('sub')
|
user_id = payload.get('sub')
|
||||||
if not user_id:
|
if not user_id:
|
||||||
logger.warning("WebSocket connection rejected: Invalid token payload")
|
logger.warning("WebSocket connection rejected: Invalid token payload")
|
||||||
await self.websocket.accept()
|
await self._accept_then_close(code=4001)
|
||||||
await self.websocket.close(code=4001)
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
self.user_id = user_id
|
self.user_id = user_id
|
||||||
@@ -149,14 +172,17 @@ class TokenAuthWebSocketConsumer:
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"WebSocket authentication failed: {str(e)}")
|
logger.error(f"WebSocket authentication failed: {str(e)}")
|
||||||
await self.websocket.accept()
|
await self._accept_then_close(code=4001)
|
||||||
await self.websocket.close(code=4001)
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
async def connect(self):
|
async def connect(self):
|
||||||
"""连接时进行Token认证"""
|
"""连接时进行Token认证"""
|
||||||
if await self.authenticate():
|
if await self.authenticate():
|
||||||
await manager.connect(self.websocket, self.user_id)
|
await manager.connect(
|
||||||
|
self.websocket,
|
||||||
|
self.user_id,
|
||||||
|
subprotocol=self._accept_subprotocol,
|
||||||
|
)
|
||||||
|
|
||||||
async def disconnect(self, close_code: int = 1000):
|
async def disconnect(self, close_code: int = 1000):
|
||||||
"""断开连接"""
|
"""断开连接"""
|
||||||
|
|||||||
@@ -169,21 +169,22 @@ export class WebSocketManager {
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
// 构建带token的URL
|
const wsUrl = this.config.url;
|
||||||
const separator = this.config.url.includes('?') ? '&' : '?';
|
const protocols = Array.isArray(this.config.protocols)
|
||||||
const wsUrl = `${this.config.url}${separator}token=${encodeURIComponent(token)}`;
|
? [...this.config.protocols]
|
||||||
|
: this.config.protocols
|
||||||
|
? [this.config.protocols]
|
||||||
|
: [];
|
||||||
|
protocols.unshift('access_token', token);
|
||||||
|
|
||||||
console.log(
|
console.log('Connecting to WebSocket:', wsUrl);
|
||||||
'Connecting to WebSocket:',
|
|
||||||
wsUrl.replace(/token=[^&]+/, 'token=***'),
|
|
||||||
);
|
|
||||||
console.log('WebSocket URL详情:', {
|
console.log('WebSocket URL详情:', {
|
||||||
originalUrl: this.config.url,
|
originalUrl: this.config.url,
|
||||||
finalUrl: wsUrl.replace(/token=[^&]+/, 'token=***'),
|
finalUrl: wsUrl,
|
||||||
isDev: import.meta.env.DEV,
|
isDev: import.meta.env.DEV,
|
||||||
});
|
});
|
||||||
|
|
||||||
this.ws = new WebSocket(wsUrl, this.config.protocols);
|
this.ws = new WebSocket(wsUrl, protocols);
|
||||||
|
|
||||||
this.ws.addEventListener('open', (event) => {
|
this.ws.addEventListener('open', (event) => {
|
||||||
console.log('WebSocket连接已建立');
|
console.log('WebSocket连接已建立');
|
||||||
|
|||||||
Reference in New Issue
Block a user