Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .github/workflows/build.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
20 changes: 5 additions & 15 deletions casbin_async_sqlalchemy_adapter/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand All @@ -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).

Expand Down Expand Up @@ -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()
Expand Down
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
[tool.black]
line-length = 150
16 changes: 4 additions & 12 deletions tests/test_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"))
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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"))

Expand Down Expand Up @@ -185,9 +179,7 @@ async def test_repr(self):
self.assertEqual(repr(rule), '<CasbinRule None: "p, alice, data1, read">')
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()
Expand Down
32 changes: 8 additions & 24 deletions tests/test_external_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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"))
Expand All @@ -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)
Expand Down Expand Up @@ -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"))
Expand All @@ -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)
Expand Down Expand Up @@ -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"))
Expand All @@ -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)
Expand All @@ -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"))
Expand Down