139 lines
4.8 KiB
Python
139 lines
4.8 KiB
Python
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()
|