Files
quantux_mcp/tests/test_auth.py
T

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()