diff --git a/README.md b/README.md index 2c89798..ff1f96a 100644 --- a/README.md +++ b/README.md @@ -131,7 +131,24 @@ async with async_session() as session: # Commit the transaction await session.commit() ``` +## Soft Delete +Soft Delete for casbin rules is supported, only when using a custom casbin rule model. +The Soft Delete mechanism is enabled by passing the attribute of the flag indicating whether +a rule is deleted to `soft_delete`. +That attribute needs to be of type `sqlalchemy.Boolean`. + +```python +adapter = Adapter( + engine, + db_class=MyCustomCasbinRuleModel, + soft_delete=MyCustomCasbinRuleModel.is_deleted +) +``` + +Please be aware that this adapter only sets a flag like `is_deleted` to `True`. +The provided model needs to handle the update of fields like `deleted_by`, `deleted_at`, etc. +An example for this is given in [softdelete.py](https://github.com/pycasbin/sqlalchemy-adapter/blob/master/examples/softdelete.py). ### Getting Help diff --git a/casbin_async_sqlalchemy_adapter/adapter.py b/casbin_async_sqlalchemy_adapter/adapter.py index e919a97..3feb66d 100644 --- a/casbin_async_sqlalchemy_adapter/adapter.py +++ b/casbin_async_sqlalchemy_adapter/adapter.py @@ -17,8 +17,8 @@ from casbin import persist from casbin.persist.adapters.asyncio import AsyncAdapter -from sqlalchemy import Column, Integer, String, delete, insert -from sqlalchemy import or_ +from sqlalchemy import Column, Integer, String, Boolean, delete, insert, update, func +from sqlalchemy import or_, not_, and_ from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession from sqlalchemy.future import select from sqlalchemy.orm import declarative_base, sessionmaker @@ -27,6 +27,8 @@ class CasbinRule(Base): + """default casbinrule (not support soft delete)""" + __tablename__ = "casbin_rule" id = Column(Integer, primary_key=True) @@ -67,6 +69,7 @@ def __init__( self, engine, db_class=None, + soft_delete=None, filtered=False, warning=True, db_session: Optional[AsyncSession] = None, @@ -76,6 +79,8 @@ def __init__( else: self._engine = engine + self.softdelete_attribute = None + if db_class is None: db_class = CasbinRule if warning: @@ -85,6 +90,12 @@ def __init__( RuntimeWarning, ) else: + if soft_delete is not None and not isinstance(soft_delete.type, Boolean): + msg = f"The type of db_class_softdelete_attribute needs to be {str(Boolean)!r}. " + msg += f"An attribute of type {str(type(soft_delete.type))!r} was given." + raise ValueError(msg) + # Softdelete is only supported when using custom class + self.softdelete_attribute = soft_delete for attr in ( "id", "ptype", @@ -130,6 +141,8 @@ async def load_policy(self, model): """loads all policy rules from the storage.""" async with self._session_scope() as session: lines = await session.execute(select(self._db_class)) + if self.softdelete_attribute: + lines = await session.execute(select(self._db_class).where(not_(self.softdelete_attribute))) for line in lines.scalars(): persist.load_policy_line(str(line), model) @@ -141,6 +154,8 @@ async def load_filtered_policy(self, model, filter) -> None: async with self._session_scope() as session: stmt = select(self._db_class) stmt = self.filter_query(stmt, filter) + if self.softdelete_attribute: + stmt = stmt.where(not_(self.softdelete_attribute)) result = await session.execute(stmt) for line in result.scalars(): persist.load_policy_line(str(line), model) @@ -169,16 +184,54 @@ async def _save_policy_line(self, ptype, rule, session=None): async def save_policy(self, model): """saves all policy rules to the storage.""" - async with self._session_scope() as session: - stmt = delete(self._db_class) - await session.execute(stmt) - for sec in ["p", "g"]: - if sec not in model.model.keys(): - continue - for ptype, ast in model.model[sec].items(): - for rule in ast.policy: - await self._save_policy_line(ptype, rule, session) - return True + if self.softdelete_attribute: + async with self._session_scope() as session: + # Fetch all currently active rules from the database. + soft_delete_column = getattr(self._db_class, self.softdelete_attribute.name) + stmt = select(self._db_class).where(not_(soft_delete_column)) + result = await session.execute(stmt) + lines_before_changes = result.scalars().all() + + print(lines_before_changes) + + # Add new rules. + for sec in ["p", "g"]: + if sec not in model.model.keys(): + continue + for ptype, ast in model.model[sec].items(): + for rule in ast.policy: + filter_stmt = select(self._db_class).where(self._db_class.ptype == ptype) + for index, value in enumerate(rule): + v_column = getattr(self._db_class, f"v{index}") + filter_stmt = filter_stmt.where(v_column == value) + + count_stmt = select(func.count()).select_from(filter_stmt.subquery()) + count_result = await session.execute(count_stmt) + + if count_result.scalar_one() == 0: + await self._save_policy_line(ptype, rule, session=session) + + # Mark old rules as deleted. + for line in lines_before_changes: + ptype = line.ptype + sec = ptype[0] + rule_parts = [v for v in (line.v0, line.v1, line.v2, line.v3, line.v4, line.v5) if v is not None] + + if not model.has_policy(sec, ptype, rule_parts): + setattr(line, self.softdelete_attribute.name, True) + + return True + else: + async with self._session_scope() as session: + stmt = delete(self._db_class) + await session.execute(stmt) + for sec in ["p", "g"]: + if sec not in model.model.keys(): + continue + for ptype, ast in model.model[sec].items(): + for rule in ast.policy: + await self._save_policy_line(ptype, rule, session) + return True async def add_policy(self, sec, ptype, rule): """adds a policy rule to the storage.""" @@ -204,7 +257,11 @@ async def add_policies(self, sec, ptype, rules): async def remove_policy(self, sec, ptype, rule): """removes a policy rule from the storage.""" async with self._session_scope() as session: - stmt = delete(self._db_class).where(self._db_class.ptype == ptype) + if self.softdelete_attribute: + soft_delete_column = self.softdelete_attribute + stmt = update(self._db_class).where(self._db_class.ptype == ptype).values({soft_delete_column.name: True}) + else: + stmt = delete(self._db_class).where(self._db_class.ptype == ptype) for i, v in enumerate(rule): stmt = stmt.where(getattr(self._db_class, "v{}".format(i)) == v) r = await session.execute(stmt) @@ -216,18 +273,39 @@ async def remove_policies(self, sec, ptype, rules): if not rules: return async with self._session_scope() as session: - stmt = delete(self._db_class).where(self._db_class.ptype == ptype) - rules = zip(*rules) - for i, rule in enumerate(rules): - stmt = stmt.where(or_(getattr(self._db_class, "v{}".format(i)) == v for v in rule)) - await session.execute(stmt) + conditions = [] + for rule in rules: + rule_conditions = [getattr(self._db_class, f"v{i}") == v for i, v in enumerate(rule)] + conditions.append(and_(*rule_conditions)) + + if self.softdelete_attribute: + soft_delete_column = self.softdelete_attribute + stmt = ( + update(self._db_class) + .where(self._db_class.ptype == ptype, not_(soft_delete_column), or_(*conditions)) + .values({self.softdelete_attribute.name: True}) + ) + else: + stmt = delete(self._db_class).where(self._db_class.ptype == ptype, or_(*conditions)) + + result = await session.execute(stmt) + return result.rowcount > 0 async def remove_filtered_policy(self, sec, ptype, field_index, *field_values): """removes policy rules that match the filter from the storage. This is part of the Auto-Save feature. """ async with self._session_scope() as session: - stmt = delete(self._db_class).where(self._db_class.ptype == ptype) + if self.softdelete_attribute: + soft_delete_column = self.softdelete_attribute + stmt = ( + update(self._db_class) + .where(self._db_class.ptype == ptype) + .where(not_(soft_delete_column)) + .values({soft_delete_column.name: True}) + ) + else: + stmt = delete(self._db_class).where(self._db_class.ptype == ptype) if not (0 <= field_index <= 5): return False @@ -238,8 +316,7 @@ async def remove_filtered_policy(self, sec, ptype, field_index, *field_values): v_value = getattr(self._db_class, "v{}".format(field_index + i)) stmt = stmt.where(v_value == v) r = await session.execute(stmt) - - return True if r.rowcount > 0 else False + return r.rowcount > 0 async def update_policy(self, sec: str, ptype: str, old_rule: List[str], new_rule: List[str]) -> None: """ diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/test_adapter.py b/tests/test_adapter.py index 72e2f09..0164181 100644 --- a/tests/test_adapter.py +++ b/tests/test_adapter.py @@ -17,7 +17,7 @@ from unittest import IsolatedAsyncioTestCase import casbin -from sqlalchemy import Column, Integer, String, select +from sqlalchemy import Column, Integer, String, Boolean, select from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from casbin_async_sqlalchemy_adapter import Adapter @@ -56,6 +56,7 @@ class TestConfig(IsolatedAsyncioTestCase): async def test_custom_db_class(self): class CustomRule(Base): __tablename__ = "casbin_rule2" + __table_args__ = {"extend_existing": True} id = Column(Integer, primary_key=True) ptype = Column(String(255)) @@ -66,6 +67,7 @@ class CustomRule(Base): v4 = Column(String(255)) v5 = Column(String(255)) not_exist = Column(String(255)) + is_deleted = Column(Boolean, default=False, nullable=False) engine = create_async_engine("sqlite+aiosqlite://", future=True) async with engine.begin() as conn: diff --git a/tests/test_adapter_softdelete.py b/tests/test_adapter_softdelete.py new file mode 100644 index 0000000..385442a --- /dev/null +++ b/tests/test_adapter_softdelete.py @@ -0,0 +1,201 @@ +import os +from pathlib import Path + +import casbin +from sqlalchemy import Column, Boolean, Integer, String, select +from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker + +from casbin_async_sqlalchemy_adapter import Adapter +from casbin_async_sqlalchemy_adapter import Base +from casbin_async_sqlalchemy_adapter.adapter import Filter +from casbin_async_sqlalchemy_adapter import CasbinRule + +from tests.test_adapter import TestConfig + + +class CasbinRuleSoftDelete(Base): + __tablename__ = "casbin_rule_soft_delete" + + id = Column(Integer, primary_key=True) + ptype = Column(String(255)) + v0 = Column(String(255)) + v1 = Column(String(255)) + v2 = Column(String(255)) + v3 = Column(String(255)) + v4 = Column(String(255)) + v5 = Column(String(255)) + + is_deleted = Column(Boolean, default=False, nullable=False) + + def __str__(self): + arr = [self.ptype] + for v in (self.v0, self.v1, self.v2, self.v3, self.v4, self.v5): + if v is None: + break + arr.append(v) + return ", ".join(arr) + + def __repr__(self): + return ''.format(self.id, str(self)) + + +def query_for_rule(adaper, ptype, v0, v1, v2): + rule_filter = Filter() + rule_filter.ptype = [ptype] + rule_filter.v0 = [v0] + rule_filter.v1 = [v1] + rule_filter.v2 = [v2] + + stmt = select(CasbinRuleSoftDelete) + stmt = adaper.filter_query(stmt, rule_filter) + return stmt + + +class TestConfigSoftDelete(TestConfig): + def setUp(self): + """ensure a clean state by deleting the old database file""" + db_file = "./test.db" + if os.path.exists(db_file): + os.remove(db_file) + + def tearDown(self): + """clean up by deleting the database file""" + db_file = "./test.db" + if os.path.exists(db_file): + os.remove(db_file) + + async def get_enforcer(self): + engine = create_async_engine("sqlite+aiosqlite:///./test.db", future=True) + adapter = Adapter(engine, db_class=CasbinRuleSoftDelete, soft_delete=CasbinRuleSoftDelete.is_deleted) + await adapter.create_table() + + async_session_maker = async_sessionmaker(engine, expire_on_commit=False, class_=AsyncSession) + async with async_session_maker() as s: + s.add(CasbinRuleSoftDelete(ptype="p", v0="alice", v1="data1", v2="read")) + s.add(CasbinRuleSoftDelete(ptype="p", v0="bob", v1="data2", v2="write")) + s.add(CasbinRuleSoftDelete(ptype="p", v0="data2_admin", v1="data2", v2="read")) + s.add(CasbinRuleSoftDelete(ptype="p", v0="data2_admin", v1="data2", v2="write")) + s.add(CasbinRuleSoftDelete(ptype="g", v0="alice", v1="data2_admin")) + await s.commit() + + scriptdir = Path(os.path.dirname(os.path.realpath(__file__))) + model_path = scriptdir / "rbac_model.conf" + e = casbin.AsyncEnforcer(str(model_path), adapter) + await e.load_policy() + + return e + + async def test_custom_db_class(self): + class CustomRule(Base): + __tablename__ = "casbin_rule3" + __table_args__ = {"extend_existing": True} + + id = Column(Integer, primary_key=True) + ptype = Column(String(255)) + v0 = Column(String(255)) + v1 = Column(String(255)) + v2 = Column(String(255)) + v3 = Column(String(255)) + v4 = Column(String(255)) + v5 = Column(String(255)) + is_deleted = Column(Boolean, default=False) + not_exist = Column(String(255)) + + engine = create_async_engine("sqlite+aiosqlite://", future=True) + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + + async_session_maker = async_sessionmaker(engine, expire_on_commit=False, class_=AsyncSession) + async with async_session_maker() as s: + s.add(CustomRule(not_exist="NotNone")) + await s.commit() + a = await s.execute(select(CustomRule)) + self.assertEqual(a.scalars().all()[0].not_exist, "NotNone") + + async def test_softdelete_flag(self): + e = await self.get_enforcer() + async_session_maker = e.adapter.session_local + + self.assertFalse(e.enforce("alice", "data5", "read")) + + async with async_session_maker() as session: + stmt = query_for_rule(e.adapter, "p", "alice", "data5", "read") + result = await session.execute(stmt) + rule = result.scalars().first() + self.assertIsNone(rule) + + await e.add_permission_for_user("alice", "data5", "read") + self.assertTrue(e.enforce("alice", "data5", "read")) + async with async_session_maker() as session: + stmt = query_for_rule(e.adapter, "p", "alice", "data5", "read") + result = await session.execute(stmt) + rule = result.scalars().first() + self.assertIsNotNone(rule) + self.assertFalse(rule.is_deleted) + + await e.delete_permission_for_user("alice", "data5", "read") + self.assertFalse(e.enforce("alice", "data5", "read")) + async with async_session_maker() as session: + stmt = query_for_rule(e.adapter, "p", "alice", "data5", "read") + result = await session.execute(stmt) + rule = result.scalars().first() + self.assertIsNotNone(rule) + self.assertTrue(rule.is_deleted) + + async def test_save_policy_softdelete(self): + e = await self.get_enforcer() + async_session_maker = e.adapter.session_local + + # Turn off auto save + e.enable_auto_save(auto_save=False) + + # Delete some preexisting rules + await e.delete_permission_for_user("alice", "data1", "read") + await e.delete_permission_for_user("bob", "data2", "write") + # Delete a non existing rule + await e.delete_permission_for_user("bob", "data100", "read") + # Add some new rules + await e.add_permission_for_user("alice", "data100", "read") + await e.add_permission_for_user("bob", "data100", "write") + + # Write changes to database + await e.save_policy() + + # Check1: ("alice", "data1", "read") should be marked as deleted + async with async_session_maker() as session: + stmt = query_for_rule(e.adapter, "p", "alice", "data1", "read") + result = await session.execute(stmt) + rule = result.scalars().first() + self.assertIsNotNone(rule) + self.assertTrue(rule.is_deleted) + + # Check2: ("bob", "data2", "write") should be marked as deleted + async with async_session_maker() as session: + stmt = query_for_rule(e.adapter, "p", "bob", "data2", "write") + result = await session.execute(stmt) + rule = result.scalars().first() + self.assertIsNotNone(rule) + self.assertTrue(rule.is_deleted) + + # Check3: ("bob", "data100", "read") should not exist + async with async_session_maker() as session: + stmt = query_for_rule(e.adapter, "p", "bob", "data100", "read") + result = await session.execute(stmt) + rule = result.scalars().first() + self.assertIsNone(rule) + + # Check4: ("alice", "data100", "read") should exist and not be deleted + async with async_session_maker() as session: + stmt = query_for_rule(e.adapter, "p", "alice", "data100", "read") + result = await session.execute(stmt) + rule = result.scalars().first() + self.assertIsNotNone(rule) + self.assertFalse(rule.is_deleted) + + # Check5: ("bob", "data100", "write") should exist and not be deleted + async with async_session_maker() as session: + stmt = query_for_rule(e.adapter, "p", "bob", "data100", "write") + result = await session.execute(stmt) + rule = result.scalars().first() + self.assertIsNotNone(rule) + self.assertFalse(rule.is_deleted)