From e6d444fb9fe4eb4f453520e24de1b7a6199e84e4 Mon Sep 17 00:00:00 2001 From: Theo Date: Wed, 17 Sep 2025 10:59:12 +0800 Subject: [PATCH] feat: add configuration for black --- .github/workflows/build.yml | 2 ++ casbin_async_sqlalchemy_adapter/adapter.py | 20 ++++---------- pyproject.toml | 2 ++ tests/test_adapter.py | 16 +++-------- tests/test_external_session.py | 32 ++++++---------------- 5 files changed, 21 insertions(+), 51 deletions(-) create mode 100644 pyproject.toml diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 766fb3a..06b2f95 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -51,6 +51,8 @@ jobs: VALIDATE_PYTHON_BLACK: true DEFAULT_BRANCH: master GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + LINTER_RULES_PATH: / + PYTHON_BLACK_CONFIG_FILE: pyproject.toml coveralls: name: Indicate completion to coveralls.io diff --git a/casbin_async_sqlalchemy_adapter/adapter.py b/casbin_async_sqlalchemy_adapter/adapter.py index e16d365..b5882c7 100644 --- a/casbin_async_sqlalchemy_adapter/adapter.py +++ b/casbin_async_sqlalchemy_adapter/adapter.py @@ -101,9 +101,7 @@ def __init__( self._db_class = db_class self._external_session = db_session - self.session_local = sessionmaker( - self._engine, expire_on_commit=False, class_=AsyncSession - ) + self.session_local = sessionmaker(self._engine, expire_on_commit=False, class_=AsyncSession) self._filtered = filtered @@ -151,9 +149,7 @@ async def load_filtered_policy(self, model, filter) -> None: def filter_query(self, stmt, filter): for attr in ("ptype", "v0", "v1", "v2", "v3", "v4", "v5"): if len(getattr(filter, attr)) > 0: - stmt = stmt.where( - getattr(self._db_class, attr).in_(getattr(filter, attr)) - ) + stmt = stmt.where(getattr(self._db_class, attr).in_(getattr(filter, attr))) return stmt.order_by(self._db_class.id) async def _save_policy_line(self, ptype, rule, session=None): @@ -217,9 +213,7 @@ async def remove_policies(self, sec, ptype, rules): 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) - ) + stmt = stmt.where(or_(getattr(self._db_class, "v{}".format(i)) == v for v in rule)) await session.execute(stmt) async def remove_filtered_policy(self, sec, ptype, field_index, *field_values): @@ -241,9 +235,7 @@ async def remove_filtered_policy(self, sec, ptype, field_index, *field_values): return True if r.rowcount > 0 else False - async def update_policy( - self, sec: str, ptype: str, old_rule: List[str], new_rule: List[str] - ) -> None: + async def update_policy(self, sec: str, ptype: str, old_rule: List[str], new_rule: List[str]) -> None: """ Update the old_rule with the new_rule in the database (storage). @@ -295,9 +287,7 @@ async def update_policies( for i in range(len(old_rules)): await self.update_policy(sec, ptype, old_rules[i], new_rules[i]) - async def update_filtered_policies( - self, sec, ptype, new_rules: List[List[str]], field_index, *field_values - ) -> List[List[str]]: + async def update_filtered_policies(self, sec, ptype, new_rules: List[List[str]], field_index, *field_values) -> List[List[str]]: """update_filtered_policies updates all the policies on the basis of the filter.""" filter = Filter() diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..6b313bc --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,2 @@ +[tool.black] +line-length = 150 diff --git a/tests/test_adapter.py b/tests/test_adapter.py index 960f632..2aee1dc 100644 --- a/tests/test_adapter.py +++ b/tests/test_adapter.py @@ -38,9 +38,7 @@ async def get_enforcer(): adapter = Adapter(engine) await adapter.create_table() - async_session = async_sessionmaker( - engine, expire_on_commit=False, class_=AsyncSession - ) + async_session = async_sessionmaker(engine, expire_on_commit=False, class_=AsyncSession) async with async_session() as s: s.add(CasbinRule(ptype="p", v0="alice", v1="data1", v2="read")) s.add(CasbinRule(ptype="p", v0="bob", v1="data2", v2="write")) @@ -73,9 +71,7 @@ class CustomRule(Base): async with engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) - session = async_sessionmaker( - engine, expire_on_commit=False, class_=AsyncSession - ) + session = async_sessionmaker(engine, expire_on_commit=False, class_=AsyncSession) async with session() as s: s.add(CustomRule(not_exist="NotNone")) await s.commit() @@ -140,9 +136,7 @@ async def test_remove_policies(self): await e.add_policies((("alice", "data5", "read"), ("alice", "data6", "read"))) self.assertTrue(e.enforce("alice", "data5", "read")) self.assertTrue(e.enforce("alice", "data6", "read")) - await e.remove_policies( - (("alice", "data5", "read"), ("alice", "data6", "read")) - ) + await e.remove_policies((("alice", "data5", "read"), ("alice", "data6", "read"))) self.assertFalse(e.enforce("alice", "data5", "read")) self.assertFalse(e.enforce("alice", "data6", "read")) @@ -185,9 +179,7 @@ async def test_repr(self): self.assertEqual(repr(rule), '') engine = create_async_engine("sqlite+aiosqlite://", future=True) - session = async_sessionmaker( - engine, expire_on_commit=False, class_=AsyncSession - ) + session = async_sessionmaker(engine, expire_on_commit=False, class_=AsyncSession) async with engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) s = session() diff --git a/tests/test_external_session.py b/tests/test_external_session.py index 13e82a4..4a85893 100644 --- a/tests/test_external_session.py +++ b/tests/test_external_session.py @@ -46,9 +46,7 @@ async def test_external_session_commit(self): try: # Create async engine - engine = create_async_engine( - f"sqlite+aiosqlite:///{db_file.name}", future=True - ) + engine = create_async_engine(f"sqlite+aiosqlite:///{db_file.name}", future=True) # Create session factory async_session_factory = async_sessionmaker(engine, expire_on_commit=False) @@ -79,9 +77,7 @@ async def test_external_session_commit(self): # Verify permissions persist after commit with new session async with async_session_factory() as new_session: new_adapter = Adapter(engine, db_session=new_session) - new_enforcer = casbin.AsyncEnforcer( - get_fixture("rbac_model.conf"), new_adapter - ) + new_enforcer = casbin.AsyncEnforcer(get_fixture("rbac_model.conf"), new_adapter) await new_enforcer.load_policy() self.assertTrue(new_enforcer.enforce("alice", "data1", "read")) @@ -100,9 +96,7 @@ async def test_external_session_rollback(self): try: # Create async engine - engine = create_async_engine( - f"sqlite+aiosqlite:///{db_file.name}", future=True - ) + engine = create_async_engine(f"sqlite+aiosqlite:///{db_file.name}", future=True) # Create session factory async_session_factory = async_sessionmaker(engine, expire_on_commit=False) @@ -133,9 +127,7 @@ async def test_external_session_rollback(self): # Verify permissions do not persist after rollback with new session async with async_session_factory() as new_session: new_adapter = Adapter(engine, db_session=new_session) - new_enforcer = casbin.AsyncEnforcer( - get_fixture("rbac_model.conf"), new_adapter - ) + new_enforcer = casbin.AsyncEnforcer(get_fixture("rbac_model.conf"), new_adapter) await new_enforcer.load_policy() self.assertFalse(new_enforcer.enforce("alice", "data1", "read")) @@ -154,9 +146,7 @@ async def test_external_session_with_save_policy(self): try: # Create async engine - engine = create_async_engine( - f"sqlite+aiosqlite:///{db_file.name}", future=True - ) + engine = create_async_engine(f"sqlite+aiosqlite:///{db_file.name}", future=True) # Create session factory async_session_factory = async_sessionmaker(engine, expire_on_commit=False) @@ -186,9 +176,7 @@ async def test_external_session_with_save_policy(self): # Verify policies persist after commit with new session async with async_session_factory() as new_session: new_adapter = Adapter(engine, db_session=new_session) - new_enforcer = casbin.AsyncEnforcer( - get_fixture("rbac_model.conf"), new_adapter - ) + new_enforcer = casbin.AsyncEnforcer(get_fixture("rbac_model.conf"), new_adapter) await new_enforcer.load_policy() self.assertTrue(new_enforcer.enforce("alice", "data1", "read")) @@ -207,9 +195,7 @@ async def test_backward_compatibility(self): try: # Create async engine - engine = create_async_engine( - f"sqlite+aiosqlite:///{db_file.name}", future=True - ) + engine = create_async_engine(f"sqlite+aiosqlite:///{db_file.name}", future=True) # Create adapter without external session (original way) adapter = Adapter(engine) @@ -231,9 +217,7 @@ async def test_backward_compatibility(self): # Create new adapter to verify persistence new_adapter = Adapter(engine) - new_enforcer = casbin.AsyncEnforcer( - get_fixture("rbac_model.conf"), new_adapter - ) + new_enforcer = casbin.AsyncEnforcer(get_fixture("rbac_model.conf"), new_adapter) await new_enforcer.load_policy() self.assertTrue(new_enforcer.enforce("alice", "data1", "read"))