fix: refresh wechat mp token on invalid credential
This commit is contained in:
@@ -23,6 +23,7 @@ from utils.logger import session_logger
|
|||||||
|
|
||||||
WECHAT_MP_ACCOUNT_TYPE = "service_account"
|
WECHAT_MP_ACCOUNT_TYPE = "service_account"
|
||||||
WECHAT_MP_NOTIFICATION_TOKEN_VERSION = 1
|
WECHAT_MP_NOTIFICATION_TOKEN_VERSION = 1
|
||||||
|
WECHAT_MP_ACCESS_TOKEN_INVALID_ERRCODES = {40001, 40014, 42001}
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -440,7 +441,6 @@ class WechatMpNotificationService(DatabaseServiceBase):
|
|||||||
condition: str,
|
condition: str,
|
||||||
page: str,
|
page: str,
|
||||||
) -> Mapping[str, Any]:
|
) -> Mapping[str, Any]:
|
||||||
access_token = await self._get_access_token()
|
|
||||||
payload: dict[str, Any] = {
|
payload: dict[str, Any] = {
|
||||||
"touser": openid,
|
"touser": openid,
|
||||||
"template_id": template_id,
|
"template_id": template_id,
|
||||||
@@ -456,16 +456,36 @@ class WechatMpNotificationService(DatabaseServiceBase):
|
|||||||
"pagepath": page,
|
"pagepath": page,
|
||||||
}
|
}
|
||||||
url = f"{settings.wechat_mp_api_base_url.rstrip('/')}/cgi-bin/message/template/send"
|
url = f"{settings.wechat_mp_api_base_url.rstrip('/')}/cgi-bin/message/template/send"
|
||||||
params = {"access_token": access_token}
|
data = await self._post_template_message(url=url, payload=payload)
|
||||||
async with httpx.AsyncClient(timeout=settings.wechat_mp_http_timeout_seconds) as client:
|
errcode = int(data.get("errcode") or 0)
|
||||||
response = await client.post(url, params=params, json=payload)
|
if errcode in WECHAT_MP_ACCESS_TOKEN_INVALID_ERRCODES:
|
||||||
response.raise_for_status()
|
session_logger.warning(
|
||||||
data = response.json()
|
"system",
|
||||||
|
"wechat_mp",
|
||||||
|
f"公众号access_token失效,强制刷新后重试: errcode={errcode}, errmsg={data.get('errmsg')}",
|
||||||
|
)
|
||||||
|
self._clear_access_token_cache()
|
||||||
|
data = await self._post_template_message(url=url, payload=payload, force_refresh_token=True)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
**data,
|
**data,
|
||||||
"ok": int(data.get("errcode") or 0) == 0,
|
"ok": int(data.get("errcode") or 0) == 0,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async def _post_template_message(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
url: str,
|
||||||
|
payload: Mapping[str, Any],
|
||||||
|
force_refresh_token: bool = False,
|
||||||
|
) -> Mapping[str, Any]:
|
||||||
|
access_token = await self._get_access_token(force_refresh=force_refresh_token)
|
||||||
|
params = {"access_token": access_token}
|
||||||
|
async with httpx.AsyncClient(timeout=settings.wechat_mp_http_timeout_seconds) as client:
|
||||||
|
response = await client.post(url, params=params, json=payload)
|
||||||
|
response.raise_for_status()
|
||||||
|
return response.json()
|
||||||
|
|
||||||
async def _list_device_family_mp_openids(self, *, device_id: str) -> list[str]:
|
async def _list_device_family_mp_openids(self, *, device_id: str) -> list[str]:
|
||||||
recipients = await self._list_device_family_mp_recipients(device_id=device_id)
|
recipients = await self._list_device_family_mp_recipients(device_id=device_id)
|
||||||
return [recipient.openid for recipient in recipients]
|
return [recipient.openid for recipient in recipients]
|
||||||
@@ -677,18 +697,23 @@ class WechatMpNotificationService(DatabaseServiceBase):
|
|||||||
finally:
|
finally:
|
||||||
await db_session.close()
|
await db_session.close()
|
||||||
|
|
||||||
async def _get_access_token(self) -> str:
|
def _clear_access_token_cache(self) -> None:
|
||||||
|
self._access_token = None
|
||||||
|
self._access_token_expire_at = 0.0
|
||||||
|
|
||||||
|
async def _get_access_token(self, *, force_refresh: bool = False) -> str:
|
||||||
now = time.time()
|
now = time.time()
|
||||||
if self._access_token and now < self._access_token_expire_at - 60:
|
if not force_refresh and self._access_token and now < self._access_token_expire_at - 60:
|
||||||
return self._access_token
|
return self._access_token
|
||||||
url = f"{settings.wechat_mp_api_base_url.rstrip('/')}/cgi-bin/token"
|
url = f"{settings.wechat_mp_api_base_url.rstrip('/')}/cgi-bin/stable_token"
|
||||||
params = {
|
payload = {
|
||||||
"grant_type": "client_credential",
|
|
||||||
"appid": settings.wechat_mp_app_id,
|
"appid": settings.wechat_mp_app_id,
|
||||||
"secret": settings.wechat_mp_app_secret,
|
"secret": settings.wechat_mp_app_secret,
|
||||||
|
"grant_type": "client_credential",
|
||||||
|
"force_refresh": force_refresh,
|
||||||
}
|
}
|
||||||
async with httpx.AsyncClient(timeout=settings.wechat_mp_http_timeout_seconds) as client:
|
async with httpx.AsyncClient(timeout=settings.wechat_mp_http_timeout_seconds) as client:
|
||||||
response = await client.get(url, params=params)
|
response = await client.post(url, json=payload)
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
data = response.json()
|
data = response.json()
|
||||||
access_token = str(data.get("access_token") or "").strip()
|
access_token = str(data.get("access_token") or "").strip()
|
||||||
|
|||||||
@@ -327,7 +327,8 @@ async def test_template_payload_uses_service_template_fields(monkeypatch):
|
|||||||
captured.update({"url": url, "params": params, "json": json, "timeout": self.timeout})
|
captured.update({"url": url, "params": params, "json": json, "timeout": self.timeout})
|
||||||
return FakeResponse()
|
return FakeResponse()
|
||||||
|
|
||||||
async def fake_get_access_token():
|
async def fake_get_access_token(*, force_refresh=False):
|
||||||
|
assert force_refresh is False
|
||||||
return "access-token"
|
return "access-token"
|
||||||
|
|
||||||
monkeypatch.setattr(service, "_get_access_token", fake_get_access_token)
|
monkeypatch.setattr(service, "_get_access_token", fake_get_access_token)
|
||||||
@@ -384,7 +385,8 @@ async def test_template_payload_can_disable_miniprogram_link(monkeypatch):
|
|||||||
captured.update({"url": url, "params": params, "json": json, "timeout": self.timeout})
|
captured.update({"url": url, "params": params, "json": json, "timeout": self.timeout})
|
||||||
return FakeResponse()
|
return FakeResponse()
|
||||||
|
|
||||||
async def fake_get_access_token():
|
async def fake_get_access_token(*, force_refresh=False):
|
||||||
|
assert force_refresh is False
|
||||||
return "access-token"
|
return "access-token"
|
||||||
|
|
||||||
monkeypatch.setattr(service, "_get_access_token", fake_get_access_token)
|
monkeypatch.setattr(service, "_get_access_token", fake_get_access_token)
|
||||||
@@ -407,3 +409,70 @@ async def test_template_payload_can_disable_miniprogram_link(monkeypatch):
|
|||||||
"thing5": {"value": "设备电量低于20%"},
|
"thing5": {"value": "设备电量低于20%"},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_template_send_refreshes_stable_token_once_when_cached_token_is_invalid(monkeypatch):
|
||||||
|
configure_enabled_wechat_mp(monkeypatch)
|
||||||
|
service = WechatMpNotificationService()
|
||||||
|
service._access_token = "cached-token"
|
||||||
|
service._access_token_expire_at = 9999999999
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
class FakeResponse:
|
||||||
|
def __init__(self, payload):
|
||||||
|
self._payload = payload
|
||||||
|
|
||||||
|
def raise_for_status(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def json(self):
|
||||||
|
return self._payload
|
||||||
|
|
||||||
|
class FakeAsyncClient:
|
||||||
|
def __init__(self, timeout):
|
||||||
|
self.timeout = timeout
|
||||||
|
|
||||||
|
async def __aenter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def __aexit__(self, exc_type, exc, tb):
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def post(self, url, params=None, json=None):
|
||||||
|
calls.append({"url": url, "params": params, "json": json})
|
||||||
|
if url.endswith("/cgi-bin/message/template/send") and params == {"access_token": "cached-token"}:
|
||||||
|
return FakeResponse(
|
||||||
|
{
|
||||||
|
"errcode": 40001,
|
||||||
|
"errmsg": "invalid credential, access_token is invalid or not latest",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
if url.endswith("/cgi-bin/stable_token"):
|
||||||
|
assert json == {
|
||||||
|
"appid": "wx-service",
|
||||||
|
"secret": "secret",
|
||||||
|
"grant_type": "client_credential",
|
||||||
|
"force_refresh": True,
|
||||||
|
}
|
||||||
|
return FakeResponse({"access_token": "fresh-token", "expires_in": 7200})
|
||||||
|
if url.endswith("/cgi-bin/message/template/send") and params == {"access_token": "fresh-token"}:
|
||||||
|
return FakeResponse({"errcode": 0, "msgid": "msg-id"})
|
||||||
|
return FakeResponse({"errcode": 99999, "errmsg": "unexpected request"})
|
||||||
|
|
||||||
|
monkeypatch.setattr("banban.service.wechat_mp_notification.httpx.AsyncClient", FakeAsyncClient)
|
||||||
|
|
||||||
|
result = await service._send_template_message(
|
||||||
|
openid="openid-1",
|
||||||
|
template_id="template-id",
|
||||||
|
title="小明",
|
||||||
|
condition="收到新的设备留言",
|
||||||
|
page="pages/chat/detail/index",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == {"errcode": 0, "msgid": "msg-id", "ok": True}
|
||||||
|
assert [call["params"] for call in calls if call["url"].endswith("/cgi-bin/message/template/send")] == [
|
||||||
|
{"access_token": "cached-token"},
|
||||||
|
{"access_token": "fresh-token"},
|
||||||
|
]
|
||||||
|
assert service._access_token == "fresh-token"
|
||||||
|
|||||||
Reference in New Issue
Block a user