Skip to content

Commit d8ef52a

Browse files
authored
Merge pull request #119 from OpenSemanticLab/dev
fix: handle expired MW sessions in ApiGateway transport
2 parents 271ba6e + 00bad82 commit d8ef52a

2 files changed

Lines changed: 45 additions & 5 deletions

File tree

src/osw/utils/_httpx_gateway.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,9 +37,18 @@ def _ensure_initialized(self):
3737
from osw.utils.workflow import ApiGatewayTransport, connect
3838

3939
osw_instance = connect()
40+
mw_site = osw_instance.site.mw_site
41+
42+
def _relogin():
43+
cred = osw_instance.site._cred_mngr.get_credential(
44+
osw_instance.site._iri
45+
)
46+
mw_site.login(username=cred.username, password=cred.password)
47+
4048
self._inner = ApiGatewayTransport(
4149
gateway_url=self._gateway_url,
42-
mw_site=osw_instance.site.mw_site,
50+
mw_site=mw_site,
51+
relogin_cb=_relogin,
4352
)
4453

4554
async def handle_async_request(self, request):

src/osw/utils/workflow.py

Lines changed: 35 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -116,7 +116,13 @@ class ApiGatewayTransport(httpx.AsyncBaseTransport):
116116
and injects MediaWiki session cookies + CSRF tokens.
117117
"""
118118

119-
def __init__(self, gateway_url: str, mw_site, csrf_required: bool = True):
119+
def __init__(
120+
self,
121+
gateway_url: str,
122+
mw_site,
123+
csrf_required: bool = True,
124+
relogin_cb=None,
125+
):
120126
"""
121127
Parameters
122128
----------
@@ -127,11 +133,14 @@ def __init__(self, gateway_url: str, mw_site, csrf_required: bool = True):
127133
Authenticated mwclient Site instance.
128134
csrf_required
129135
Whether to send MW CSRF token for write methods.
136+
relogin_cb
137+
Callable that re-authenticates ``mw_site`` when the session expires.
130138
"""
131139
self._gateway_url = gateway_url.rstrip("/")
132140
self._mw_site = mw_site
133141
self._csrf_token = None
134142
self._csrf_required = csrf_required
143+
self._relogin_cb = relogin_cb
135144

136145
def _get_csrf_token(self) -> str:
137146
if self._csrf_token is None:
@@ -230,11 +239,19 @@ async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
230239
response = await httpx.AsyncHTTPTransport().handle_async_request(
231240
redirect_req
232241
)
233-
# Refresh CSRF token and retry once on 403
242+
# Refresh CSRF token and retry on 403; if still 403, re-login
234243
if response.status_code == 403:
235244
self._csrf_token = None
236245
rewritten = self._rewrite_request(request)
237246
response = await httpx.AsyncHTTPTransport().handle_async_request(rewritten)
247+
if response.status_code == 403 and self._relogin_cb:
248+
log.warning("ApiGateway session expired, re-authenticating")
249+
self._relogin_cb()
250+
self._csrf_token = None
251+
rewritten = self._rewrite_request(request)
252+
response = await httpx.AsyncHTTPTransport().handle_async_request(
253+
rewritten
254+
)
238255
return response
239256

240257

@@ -250,10 +267,17 @@ def get_gateway_httpx_settings(gateway_url: str, osw_instance: OSW) -> dict:
250267
osw_instance
251268
A connected OSW instance (provides mwclient session).
252269
"""
270+
mw_site = osw_instance.site.mw_site
271+
272+
def _relogin():
273+
cred = osw_instance.site._cred_mngr.get_credential(osw_instance.site._iri)
274+
mw_site.login(username=cred.username, password=cred.password)
275+
253276
transport = ApiGatewayTransport(
254277
gateway_url=gateway_url,
255-
mw_site=osw_instance.site.mw_site,
278+
mw_site=mw_site,
256279
csrf_required=False,
280+
relogin_cb=_relogin,
257281
)
258282
return {"transport": transport, "base_url": gateway_url}
259283

@@ -640,9 +664,16 @@ async def _deploy(param: DeployParam):
640664
if _is_apigateway_url(gateway_url) and param.osw is not None:
641665
_original_api_url = environ.get("PREFECT_API_URL")
642666
environ["PREFECT_API_URL"] = gateway_url
667+
_mw_site = param.osw.site.mw_site
668+
669+
def _relogin():
670+
cred = param.osw.site._cred_mngr.get_credential(param.osw.site._iri)
671+
_mw_site.login(username=cred.username, password=cred.password)
672+
643673
_gw_transport = ApiGatewayTransport(
644674
gateway_url=gateway_url,
645-
mw_site=param.osw.site.mw_site,
675+
mw_site=_mw_site,
676+
relogin_cb=_relogin,
646677
)
647678
# Patch httpx.AsyncClient to auto-inject our transport when
648679
# the base_url is an ApiGateway URL. One patch covers ALL

0 commit comments

Comments
 (0)