import jwt import pytest from monday_com_orca_backend.api.auth.utils import validate_jwt_token from monday_com_orca_backend.enums import JwtAlgorithm from monday_com_orca_backend.utils.custom_types import FrozenDict class TestValidateJwtToken: _secret = "test_secret" _algorithm = JwtAlgorithm.HS256.value _payload = FrozenDict({"user_id": 123, "is_admin": True, "exp": 9999999999}) def create_token(self, payload=None, secret=None, algorithm=None): """Helper function to create a mock JWT token for testing.""" if payload is None: payload = self._payload if secret is None: secret = self._secret if algorithm is None: algorithm = self._algorithm return jwt.encode(dict(payload), secret, algorithm=algorithm) def test_valid_token_returns_payload(self): """Test that a valid token returns the correct payload.""" token = self.create_token() result = validate_jwt_token(token, self._secret, self._algorithm) assert isinstance(result, FrozenDict) for k, v in self._payload.items(): assert result[k] == v def test_invalid_token_returns_none(self): """Test that an invalid token returns None.""" result = validate_jwt_token( "invalid.token.value", self._secret, self._algorithm ) assert result is None def test_wrong_secret_returns_none(self): """Test that using a wrong secret returns None.""" token = self.create_token() result = validate_jwt_token(token, "wrong_secret", self._algorithm) assert result is None def test_wrong_algorithm_returns_none(self): """Test that using a different algorithm returns None.""" token = self.create_token() result = validate_jwt_token(token, self._secret, JwtAlgorithm.HS512) assert result is None def test_read_only_result(self): """Ensure the result is read-only and cannot be modified by callers.""" token = self.create_token() result = validate_jwt_token(token, self._secret, self._algorithm) with pytest.raises(TypeError): result["new_key"] = "value" def test_cache_behavior(self, monkeypatch): """Test that the function uses caching to avoid repeated decoding.""" token = self.create_token() called = {"count": 0} original_decode = jwt.decode def fake_decode(*args, **kwargs): called["count"] += 1 return original_decode(*args, **kwargs) monkeypatch.setattr(jwt, "decode", fake_decode) validate_jwt_token.cache_clear() validate_jwt_token(token, self._secret, self._algorithm) validate_jwt_token(token, self._secret, self._algorithm) assert called["count"] == 1 # Only called once due to cache