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