diff --git a/snappass/main.py b/snappass/main.py index da202596..a05e1382 100644 --- a/snappass/main.py +++ b/snappass/main.py @@ -1,4 +1,5 @@ import os +import re import sys import uuid @@ -54,6 +55,11 @@ def get_locale(): 'hour': 3600} DEFAULT_API_TTL = 1209600 MAX_TTL = DEFAULT_API_TTL +STORAGE_KEY_PATTERN = re.compile(r'^' + re.escape(REDIS_PREFIX) + r'[0-9a-f]{32}$') + + +def is_valid_storage_key(storage_key): + return bool(STORAGE_KEY_PATTERN.fullmatch(storage_key)) def check_redis_alive(fn): @@ -160,6 +166,9 @@ def get_password(token): If not, the password is simply returned as is. """ storage_key, decryption_key = parse_token(token) + if not is_valid_storage_key(storage_key): + return None + password = redis_client.get(storage_key) redis_client.delete(storage_key) @@ -173,7 +182,10 @@ def get_password(token): @check_redis_alive def password_exists(token): - storage_key, decryption_key = parse_token(token) + storage_key, _ = parse_token(token) + if not is_valid_storage_key(storage_key): + return False + return redis_client.exists(storage_key) diff --git a/tests.py b/tests.py index b4b089e9..e01a2330 100644 --- a/tests.py +++ b/tests.py @@ -59,11 +59,21 @@ def test_encryption_key_is_returned(self): def test_unencrypted_passwords_still_work(self): unencrypted_password = "trustevery1" - storage_key = uuid.uuid4().hex + storage_key = snappass.REDIS_PREFIX + uuid.uuid4().hex snappass.redis_client.setex(storage_key, 30, unencrypted_password) retrieved_password = snappass.get_password(storage_key) self.assertEqual(unencrypted_password, retrieved_password) + def test_get_password_rejects_non_snappass_storage_key(self): + snappass.redis_client.setex("session:42", 30, "sensitive-data") + retrieved_password = snappass.get_password("session:42") + self.assertIsNone(retrieved_password) + self.assertEqual("sensitive-data", snappass.redis_client.get("session:42").decode('utf-8')) + + def test_password_exists_rejects_non_snappass_storage_key(self): + snappass.redis_client.setex("session:42", 30, "sensitive-data") + self.assertFalse(snappass.password_exists("session:42")) + def test_password_is_decoded(self): password = "correct horse battery staple" key = snappass.set_password(password, 30) @@ -314,6 +324,11 @@ def test_check_password_api_v2_bad_keys(self): rvc = self.app.head('/api/v2/passwords/' + quote(key[::-1])) self.assertEqual(rvc.status_code, 404) + def test_check_password_api_v2_non_snappass_keys(self): + snappass.redis_client.setex("session:42", 30, "sensitive-data") + rvc = self.app.head('/api/v2/passwords/' + quote("session:42")) + self.assertEqual(rvc.status_code, 404) + def test_retrieve_password_api_v2(self): password = 'my name is my passport. verify me.' rv = self.app.post( @@ -352,6 +367,12 @@ def test_retrieve_password_api_v2_bad_keys(self): bad_token = invalid_params[0] self.assertEqual(bad_token['name'], 'token') + def test_retrieve_password_api_v2_non_snappass_keys(self): + snappass.redis_client.setex("session:42", 30, "sensitive-data") + rvc = self.app.get('/api/v2/passwords/' + quote("session:42")) + self.assertEqual(rvc.status_code, 404) + self.assertEqual("sensitive-data", snappass.redis_client.get("session:42").decode('utf-8')) + if __name__ == '__main__': unittest.main()