Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 18 additions & 3 deletions lark_oapi/ws/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import random
import time
from typing import Callable, Dict, Mapping, Optional
from urllib.parse import urlparse, parse_qs
from urllib.parse import parse_qs, unquote_plus, urlparse, urlsplit, urlunsplit

import requests
import websockets
Expand Down Expand Up @@ -65,6 +65,21 @@ def _ordinal(n: int):
return str(n) + suffix


def _redact_ws_url(url: str) -> str:
parsed = urlsplit(url)
if not parsed.query:
return url

redacted_query = []
for parameter in parsed.query.split("&"):
key, separator, _ = parameter.partition("=")
if separator and unquote_plus(key) in {"access_key", "ticket"}:
parameter = f"{key}=REDACTED"
redacted_query.append(parameter)

return urlunsplit(parsed._replace(query="&".join(redacted_query)))


async def _select():
while True:
await asyncio.sleep(3600)
Expand Down Expand Up @@ -207,7 +222,7 @@ async def _connect(self) -> None:
self._conn_id = conn_id
self._service_id = service_id

logger.info(self._fmt_log("connected to {}", conn_url))
logger.info(self._fmt_log("connected to {}", _redact_ws_url(conn_url)))
loop.create_task(self._receive_message_loop())
except InvalidHandshake as e:
_parse_ws_conn_exception(e)
Expand Down Expand Up @@ -413,7 +428,7 @@ async def _disconnect(self):
if self._conn is None:
return
await self._conn.close()
logger.info(self._fmt_log("disconnected to {}", self._conn_url))
logger.info(self._fmt_log("disconnected to {}", _redact_ws_url(self._conn_url)))
finally:
self._conn = None
self._conn_url = ""
Expand Down
47 changes: 47 additions & 0 deletions lark_oapi/ws/tests/test_websockets_compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,20 @@ async def close(self):
pass


def test_redact_ws_url_preserves_diagnostic_context_and_query_shape():
url = (
"wss://msg-frontier.example.test/ws/v2?fpid=493&"
"access%5Fkey=access%2Fsecret&ticket=first&ticket=second&"
"empty=&value-less#fragment"
)

assert ws_client._redact_ws_url(url) == (
"wss://msg-frontier.example.test/ws/v2?fpid=493&"
"access%5Fkey=REDACTED&ticket=REDACTED&ticket=REDACTED&"
"empty=&value-less#fragment"
)


def test_parse_ws_connection_exception_reads_new_invalid_status_response_headers():
exc = RuntimeError("handshake failed")
exc.response = SimpleNamespace(
Expand Down Expand Up @@ -113,6 +127,39 @@ async def fake_connect(uri, *, proxy=True):
}


@pytest.mark.asyncio
async def test_connection_lifecycle_logs_redact_credentials_without_changing_connection_url(monkeypatch):
original_url = (
"wss://msg-frontier.example.test/ws/v2?device_id=device&service_id=42&"
"access_key=access-key-test-value&ticket=ticket-test-value"
)
captured = {"logs": []}

async def fake_connect(uri, **kwargs):
captured["uri"] = uri
return _FakeConn()

client = ws_client.Client("app_id", "app_secret")
monkeypatch.setattr(client, "_get_conn_url", lambda: original_url)
monkeypatch.setattr(ws_client.websockets, "connect", fake_connect)
monkeypatch.setattr(ws_client.logger, "info", captured["logs"].append)
monkeypatch.setattr(
ws_client.loop,
"create_task",
lambda coro: coro.close() if hasattr(coro, "close") else None,
)

await client._connect()
await client._disconnect()

assert captured["uri"] == original_url
assert len(captured["logs"]) == 2
assert all("access-key-test-value" not in log for log in captured["logs"])
assert all("ticket-test-value" not in log for log in captured["logs"])
assert all("device_id=device" in log for log in captured["logs"])
assert all(log.count("REDACTED") == 2 for log in captured["logs"])


@pytest.mark.asyncio
async def test_connect_does_not_pass_proxy_to_older_websockets(monkeypatch):
captured = {}
Expand Down