121 lines
3.6 KiB
Python
121 lines
3.6 KiB
Python
import unittest
|
|||
|
|
|
||
|
|
from FlowerServices import Flower, InvalidCredentials, legacy_key
|
||
|
|
|
||
|
|
|
||
|
|
class FakeCursor:
|
||
|
|
def __init__(self, one=None, many=(), changed=1):
|
||
|
|
self.one = one
|
||
|
|
self.many = many
|
||
|
|
self.changed = changed
|
||
|
|
self.calls = []
|
||
|
|
|
||
|
|
def __enter__(self):
|
||
|
|
return self
|
||
|
|
|
||
|
|
def __exit__(self, *args):
|
||
|
|
return False
|
||
|
|
|
||
|
|
def execute(self, sql, parameters=None):
|
||
|
|
self.calls.append((" ".join(sql.split()), parameters))
|
||
|
|
return self.changed
|
||
|
|
|
||
|
|
def fetchone(self):
|
||
|
|
return self.one
|
||
|
|
|
||
|
|
def fetchall(self):
|
||
|
|
return self.many
|
||
|
|
|
||
|
|
|
||
|
|
class FakeConnection:
|
||
|
|
def __init__(self, cursor):
|
||
|
|
self.fake_cursor = cursor
|
||
|
|
self.commits = 0
|
||
|
|
self.rollbacks = 0
|
||
|
|
self.closed = 0
|
||
|
|
|
||
|
|
def cursor(self):
|
||
|
|
return self.fake_cursor
|
||
|
|
|
||
|
|
def commit(self):
|
||
|
|
self.commits += 1
|
||
|
|
|
||
|
|
def rollback(self):
|
||
|
|
self.rollbacks += 1
|
||
|
|
|
||
|
|
def close(self):
|
||
|
|
self.closed += 1
|
||
|
|
|
||
|
|
|
||
|
|
class FlowerCompatibilityTests(unittest.TestCase):
|
||
|
|
def make_vault(self, cursor):
|
||
|
|
connection = FakeConnection(cursor)
|
||
|
|
calls = []
|
||
|
|
|
||
|
|
def connect(**kwargs):
|
||
|
|
calls.append(kwargs)
|
||
|
|
return connection
|
||
|
|
|
||
|
|
return Flower("alice", "legacy_pwd", connect), connection, calls
|
||
|
|
|
||
|
|
def test_legacy_key_derivation_is_unchanged(self):
|
||
|
|
self.assertEqual(legacy_key("black"), "blackblackblackblackblackblackblackblackbla=")
|
||
|
|
|
||
|
|
def test_legacy_ciphertext_round_trip(self):
|
||
|
|
vault, _, _ = self.make_vault(FakeCursor())
|
||
|
|
token = vault.encode("existing database secret")
|
||
|
|
self.assertEqual(vault.decode(token), "existing database secret")
|
||
|
|
|
||
|
|
def test_existing_row_is_decrypted_without_schema_changes(self):
|
||
|
|
seed, _, _ = self.make_vault(FakeCursor())
|
||
|
|
row = (seed.encode("secret"), seed.encode("login"), seed.encode("Alice"))
|
||
|
|
cursor = FakeCursor(one=row)
|
||
|
|
vault, _, calls = self.make_vault(cursor)
|
||
|
|
|
||
|
|
item = vault.one("Acme", "2024-02-03 12:00:00")
|
||
|
|
|
||
|
|
self.assertEqual(item["mySecret"], "secret")
|
||
|
|
self.assertEqual(item["myID"], "login")
|
||
|
|
self.assertEqual(item["myName"], "Alice")
|
||
|
|
self.assertEqual(
|
||
|
|
cursor.calls[0][1], ("Acme", "2024-02-03 12:00:00")
|
||
|
|
)
|
||
|
|
self.assertEqual(calls[0]["user"], "alice_")
|
||
|
|
self.assertEqual(calls[0]["db"], "Flowers")
|
||
|
|
|
||
|
|
def test_writes_use_parameters_and_existing_columns(self):
|
||
|
|
cursor = FakeCursor()
|
||
|
|
vault, connection, _ = self.make_vault(cursor)
|
||
|
|
hostile_label = 'Acme"; DROP TABLE alice_; --'
|
||
|
|
|
||
|
|
result = vault.add(hostile_label, "Alice", "login", "secret")
|
||
|
|
|
||
|
|
sql, parameters = cursor.calls[0]
|
||
|
|
self.assertIn("organization, myID, myName, mySecret, deleted", sql)
|
||
|
|
self.assertNotIn(hostile_label, sql)
|
||
|
|
self.assertEqual(parameters[0], hostile_label)
|
||
|
|
self.assertEqual(result["dateCreated"], "STORED")
|
||
|
|
self.assertEqual(connection.commits, 1)
|
||
|
|
|
||
|
|
def test_latest_active_list_keeps_legacy_query_semantics(self):
|
||
|
|
seed, _, _ = self.make_vault(FakeCursor())
|
||
|
|
cursor = FakeCursor(
|
||
|
|
many=(("Acme", seed.encode("login"), "2024-02-03 12:00:00"),)
|
||
|
|
)
|
||
|
|
vault, _, _ = self.make_vault(cursor)
|
||
|
|
|
||
|
|
items = vault.all()
|
||
|
|
|
||
|
|
self.assertEqual(items[0]["organization"], "Acme")
|
||
|
|
self.assertEqual(items[0]["myID"], "login")
|
||
|
|
self.assertIn("MAX(dateCreated)", cursor.calls[0][0])
|
||
|
|
self.assertEqual(cursor.calls[0][1], (0,))
|
||
|
|
|
||
|
|
def test_unsafe_table_identifier_is_rejected(self):
|
||
|
|
with self.assertRaises(InvalidCredentials):
|
||
|
|
Flower("alice`; DROP DATABASE Flowers", "legacy_pwd")
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
unittest.main()
|