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
75 changes: 61 additions & 14 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 urlparse, parse_qs, urlunparse, urlencode

import requests
import websockets
Expand All @@ -28,12 +28,6 @@
from lark_oapi.ws.pb.google.protobuf.internal.containers import RepeatedCompositeFieldContainer
from lark_oapi.ws.pb.pbbp2_pb2 import Frame

try:
loop = asyncio.get_event_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)


def _get_by_key(headers: RepeatedCompositeFieldContainer, key: str) -> str:
for header in headers:
Expand All @@ -43,6 +37,31 @@ def _get_by_key(headers: RepeatedCompositeFieldContainer, key: str) -> str:
raise HeaderNotFoundException(key)


# Query parameters on the WS endpoint URL whose values are credentials
# (access_key, ticket) and must never appear in logs.
_SENSITIVE_QUERY_KEYS = ("access_key", "ticket")


def _redact_conn_url(url: Optional[str]) -> Optional[str]:
"""Return a copy of ``url`` safe for logging.

The WS endpoint URL carries ``access_key`` and ``ticket`` credentials as
query parameters; those values are masked while the rest of the URL
(scheme, host, path, other params) is preserved. The returned URL must
only be used for logging - never for connecting.
"""
if not url:
return url
u = urlparse(url)
if not u.query:
return url
q = parse_qs(u.query, keep_blank_values=True)
for key in _SENSITIVE_QUERY_KEYS:
if key in q:
q[key] = ["***"]
return urlunparse(u._replace(query=urlencode(q, doseq=True)))


def _new_ping_frame(service_id: int) -> Frame:
frame = Frame()
header = frame.headers.add()
Expand Down Expand Up @@ -151,7 +170,13 @@ def __init__(self,
self._reconnect_interval: int = 120
self._ping_interval: int = 120
self._cache: ExpiringCache = ExpiringCache(clear_interval=30)
self._lock = asyncio.Lock()
# Event loop owned by this client instance (created lazily on first
# use). Previously a module-level global was shared by every client,
# which broke multi-bot setups (issues #119 / #133).
self._loop: Optional[asyncio.AbstractEventLoop] = None
# asyncio primitives bind to an event loop; created lazily on the
# client's own loop so it never binds to a foreign loop.
self._lock: Optional[asyncio.Lock] = None
# Observer hooks for higher-level wrappers (e.g. FeishuChannel) to
# react to reconnect lifecycle. ``on_reconnecting`` fires when the
# client decides a connection was lost and starts retrying;
Expand All @@ -161,7 +186,28 @@ def __init__(self,
self.on_reconnected: Callable[[], None] = lambda: None
logger.setLevel(log_level.value)

def _get_loop(self) -> asyncio.AbstractEventLoop:
"""Return the event loop this client runs on.

Every client owns a dedicated loop, created lazily on first use, so
multiple clients (multi-bot, one thread per bot) never share a loop,
and a client constructed inside an already-running loop (e.g. inside
``asyncio.run``) is not tied to it. See issues #119 / #133.
"""
if self._loop is None or self._loop.is_closed():
self._loop = asyncio.new_event_loop()
return self._loop

def _get_lock(self) -> asyncio.Lock:
# asyncio primitives bind to the running loop when created; create the
# lock lazily (from a coroutine running on the client's own loop) so
# it never binds to a foreign loop.
if self._lock is None:
self._lock = asyncio.Lock()
return self._lock

def start(self) -> None:
loop = self._get_loop()
try:
loop.run_until_complete(self._connect())
except ClientException as e:
Expand Down Expand Up @@ -191,7 +237,8 @@ async def _ping_loop(self):
await asyncio.sleep(self._ping_interval)

async def _connect(self) -> None:
await self._lock.acquire()
lock = self._get_lock()
await lock.acquire()
if self._conn is not None:
return
try:
Expand All @@ -207,20 +254,20 @@ async def _connect(self) -> None:
self._conn_id = conn_id
self._service_id = service_id

logger.info(self._fmt_log("connected to {}", conn_url))
loop.create_task(self._receive_message_loop())
logger.info(self._fmt_log("connected to {}", _redact_conn_url(conn_url)))
asyncio.get_running_loop().create_task(self._receive_message_loop())
except InvalidHandshake as e:
_parse_ws_conn_exception(e)
finally:
self._lock.release()
lock.release()

async def _receive_message_loop(self):
try:
while True:
if self._conn is None:
raise ConnectionClosedException("connection is closed")
msg = await self._conn.recv()
loop.create_task(self._handle_message(msg))
asyncio.get_running_loop().create_task(self._handle_message(msg))
except Exception as e:
logger.error(self._fmt_log("receive message loop exit, err: {}", e))
await self._disconnect()
Expand Down Expand Up @@ -413,7 +460,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_conn_url(self._conn_url)))
finally:
self._conn = None
self._conn_url = ""
Expand Down
30 changes: 30 additions & 0 deletions lark_oapi/ws/tests/test_redact_conn_url.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
from lark_oapi.ws.client import _redact_conn_url

FAKE_URL = (
"wss://msg-frontier.example.test/ws/v2"
"?fpid=493&access_key=access-key-test-value"
"&service_id=33554678&ticket=ticket-test-value"
)


def test_redacts_access_key_and_ticket_values():
redacted = _redact_conn_url(FAKE_URL)
assert "access-key-test-value" not in redacted
assert "ticket-test-value" not in redacted
# Values are masked (urlencode may encode '*' as %2A).
assert "access_key=" in redacted and "ticket=" in redacted
assert "*" not in redacted.replace("%2A", "")
# Non-sensitive parts are preserved (host, path, other query params).
assert redacted.startswith("wss://msg-frontier.example.test/ws/v2?")
assert "fpid=493" in redacted
assert "service_id=33554678" in redacted


def test_keeps_url_without_query():
url = "wss://msg-frontier.example.test/ws/v2"
assert _redact_conn_url(url) == url


def test_handles_none():
assert _redact_conn_url(None) is None
assert _redact_conn_url("") == ""