| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114 |
- """Unit tests for DMS token acquisition; no live server or real secret required."""
- from __future__ import annotations
- import json
- import unittest
- from collections.abc import Callable
- from urllib.parse import parse_qs, urlparse
- from urllib.request import Request
- from step0_pre_prepare import DmsAuthSettings, DmsTokenManager
- def response(payload: dict[str, object]) -> bytes:
- return json.dumps(payload).encode("utf-8")
- class FakeTransport:
- def __init__(self, handler: Callable[[Request], bytes]) -> None:
- self.handler = handler
- self.requests: list[Request] = []
- def __call__(self, request: Request, timeout: float) -> bytes:
- self.requests.append(request)
- return self.handler(request)
- class DmsTokenManagerTests(unittest.TestCase):
- def settings(self, **overrides: object) -> DmsAuthSettings:
- values: dict[str, object] = {
- "dms_base_url": "https://dms.example/dms",
- "username": "user",
- "password": "password",
- "oauth_api_url": "https://oauth-api.example/oauth",
- }
- values.update(overrides)
- return DmsAuthSettings(**values) # type: ignore[arg-type]
- def test_reuses_valid_bootstrap_token_without_login(self) -> None:
- def handler(request: Request) -> bytes:
- self.assertEqual(urlparse(request.full_url).path, "/oauth/user/validateToken")
- self.assertEqual(parse_qs(urlparse(request.full_url).query)["serviceId"], ["2"])
- return response({"code": 200, "content": {}})
- transport = FakeTransport(handler)
- manager = DmsTokenManager(
- self.settings(bootstrap_token="existing-token"), transport=transport
- )
- self.assertEqual(manager.get_token(), "existing-token")
- self.assertEqual(manager.get_token(), "existing-token")
- self.assertEqual(len(transport.requests), 1)
- def test_invalid_bootstrap_token_falls_back_to_login(self) -> None:
- validated_tokens: list[str] = []
- def handler(request: Request) -> bytes:
- path = urlparse(request.full_url).path
- if path.endswith("/user/login"):
- form = parse_qs((request.data or b"").decode("utf-8"))
- self.assertEqual(form["userName"], ["user"])
- self.assertEqual(form["clientId"], ["0"])
- return response({"code": 200, "content": {}, "message": "fresh-token"})
- token = dict(request.header_items()).get("Token", "")
- validated_tokens.append(token)
- return response({"code": 200 if token == "fresh-token" else 212})
- manager = DmsTokenManager(
- self.settings(bootstrap_token="expired-token"),
- transport=FakeTransport(handler),
- )
- self.assertEqual(manager.get_token(), "fresh-token")
- self.assertEqual(validated_tokens, ["expired-token", "fresh-token"])
- def test_discovers_oauth_api_from_frontend_configuration(self) -> None:
- def handler(request: Request) -> bytes:
- if request.full_url == "https://dms.example/static/config/config.js":
- return b'window.config = { oauthUrl: "https://oauth-ui.example/login" }'
- if request.full_url == "https://oauth-ui.example/config.js":
- return b'window.systemConfig = { BASE_URL: "https://oauth-api.example/oauth" }'
- if urlparse(request.full_url).path.endswith("/user/login"):
- return response({"code": 200, "content": {}, "message": "fresh-token"})
- return response({"code": 200, "content": {}})
- manager = DmsTokenManager(
- self.settings(oauth_api_url=None, bootstrap_token=None),
- transport=FakeTransport(handler),
- )
- self.assertEqual(manager.get_token(), "fresh-token")
- def test_retries_once_after_dms_invalid_token_code(self) -> None:
- def handler(request: Request) -> bytes:
- if urlparse(request.full_url).path.endswith("/user/login"):
- return response({"code": 200, "content": {}, "message": "fresh-token"})
- return response({"code": 200, "content": {}})
- manager = DmsTokenManager(
- self.settings(bootstrap_token="old-token"), transport=FakeTransport(handler)
- )
- seen: list[str] = []
- def operation(token: str) -> dict[str, int]:
- seen.append(token)
- return {"code": 212 if len(seen) == 1 else 200}
- self.assertEqual(manager.call_with_refresh(operation), {"code": 200})
- self.assertEqual(seen, ["old-token", "fresh-token"])
- if __name__ == "__main__":
- unittest.main()
|