import base64 import json import threading import time import unittest from quantux_auth import DownloadTokenStore, QuantUXSessionManager from quantux_client import QuantUXClient, QuantUXError def jwt(exp): def encode(value): raw = json.dumps(value, separators=(",", ":")).encode() return base64.urlsafe_b64encode(raw).rstrip(b"=").decode() return f"{encode({'alg': 'none'})}.{encode({'exp': exp})}.signature" class FakeClient: def __init__(self, token, now, login_delay=0): self.token = token self.now = now self.login_delay = login_delay self.login_calls = 0 def token_expires_at(self): return QuantUXClient("http://example", token=self.token).token_expires_at() def token_is_valid(self, leeway=60, now=None): current = self.now if now is None else now expires = self.token_expires_at() return bool(expires and expires > current + leeway) def login(self, email, password): self.login_calls += 1 if self.login_delay: time.sleep(self.login_delay) self.token = jwt(self.now + 3600) return {"token": self.token} class DummyResponse: def __init__(self, status, payload=None, text=""): self.status_code = status self._payload = payload or {} self.text = text def json(self): return self._payload class DummySession: def __init__(self, responses): self.responses = list(responses) self.headers = {} self.calls = 0 def request(self, method, url, **kwargs): self.calls += 1 return self.responses.pop(0) class AuthenticationTests(unittest.TestCase): def test_expired_token_is_refreshed(self): now = 1_000_000 client = FakeClient(jwt(now - 1), now) manager = QuantUXSessionManager(client, "admin@example.com", "secret", clock=lambda: now) manager.ensure_authenticated() self.assertEqual(1, client.login_calls) self.assertTrue(manager.status()["token_valid"]) def test_concurrent_expiry_causes_one_login(self): now = 1_000_000 client = FakeClient(jwt(now - 1), now, login_delay=0.03) manager = QuantUXSessionManager(client, "admin@example.com", "secret", clock=lambda: now) threads = [threading.Thread(target=manager.ensure_authenticated) for _ in range(12)] for thread in threads: thread.start() for thread in threads: thread.join() self.assertEqual(1, client.login_calls) def test_missing_credentials_is_clear_error(self): now = 1_000_000 manager = QuantUXSessionManager(FakeClient(jwt(now - 1), now), clock=lambda: now) with self.assertRaisesRegex(QuantUXError, "credentials are not configured"): manager.ensure_authenticated() def test_401_refreshes_and_retries_once(self): now = int(time.time()) client = QuantUXClient("http://example", token=jwt(now + 3600)) client.session = DummySession([ DummyResponse(401, text="unauthorized"), DummyResponse(200, payload={"ok": True}), ]) refreshes = [] def refresh(failed_token): refreshes.append(failed_token) client.token = jwt(now + 7200) client.set_auth_refresh_handler(refresh) self.assertEqual({"ok": True}, client._req("GET", "/rest/apps")) self.assertEqual(1, len(refreshes)) self.assertEqual(2, client.session.calls) def test_405_is_not_retried(self): now = int(time.time()) client = QuantUXClient("http://example", token=jwt(now + 3600)) client.session = DummySession([DummyResponse(405, text="not allowed")]) refreshes = [] client.set_auth_refresh_handler(lambda token: refreshes.append(token)) with self.assertRaisesRegex(QuantUXError, "HTTP 405"): client._req("POST", "/rest/apps", json={"name": "test"}) self.assertEqual([], refreshes) self.assertEqual(1, client.session.calls) def test_health_status_marks_expired_token_invalid(self): now = 1_000_000 manager = QuantUXSessionManager(FakeClient(jwt(now - 1), now), clock=lambda: now) status = manager.status() self.assertFalse(status["logged_in"]) self.assertFalse(status["token_valid"]) class DownloadTokenTests(unittest.TestCase): def test_token_is_scoped_and_expires(self): now = [1000] store = DownloadTokenStore(ttl_seconds=60, clock=lambda: now[0]) token = store.issue("prototype.html") self.assertTrue(store.validate(token, "prototype.html")) self.assertFalse(store.validate(token, "other.html")) now[0] += 61 self.assertFalse(store.validate(token, "prototype.html")) if __name__ == "__main__": unittest.main()