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
24 changes: 12 additions & 12 deletions .github/workflows/build.yml
Original file line number Diff line number Diff line change
Expand Up @@ -38,19 +38,19 @@ jobs:
COVERALLS_FLAG_NAME: ${{ matrix.os }} - ${{ matrix.python-version }}
COVERALLS_PARALLEL: true

lint:
name: Run Linters
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v2
lint:
name: Run Linters
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v2

- name: Super-Linter
uses: github/super-linter@v4.2.2
env:
VALIDATE_PYTHON_BLACK: true
DEFAULT_BRANCH: master
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
- name: Super-Linter
uses: github/super-linter@v4.2.2
env:
VALIDATE_PYTHON_BLACK: true
DEFAULT_BRANCH: master
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}

coveralls:
name: Indicate completion to coveralls.io
Expand Down
9 changes: 8 additions & 1 deletion casbin_async_sqlalchemy_adapter/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,14 @@ class Filter:
class Adapter(AsyncAdapter):
"""the interface for Casbin adapters."""

def __init__(self, engine, db_class=None, filtered=False, warning=True, db_session: Optional[AsyncSession] = None):
def __init__(
self,
engine,
db_class=None,
filtered=False,
warning=True,
db_session: Optional[AsyncSession] = None,
):
if isinstance(engine, str):
self._engine = create_async_engine(engine, future=True)
else:
Expand Down
12 changes: 9 additions & 3 deletions tests/test_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,9 @@ 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 @@ -71,7 +73,9 @@ 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 @@ -181,7 +185,9 @@ 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
168 changes: 92 additions & 76 deletions tests/test_external_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,48 +41,52 @@ class TestExternalSession(IsolatedAsyncioTestCase):
async def test_external_session_commit(self):
"""Test using external session with commit."""
# Create a temporary database file
db_file = tempfile.NamedTemporaryFile(suffix='.db', delete=False)
db_file = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
db_file.close()

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)

# Test with external session
async with async_session_factory() as external_session:
# Create adapter with external session
adapter = Adapter(engine, db_session=external_session)

# Create table
await adapter.create_table()

# Create enforcer
e = casbin.AsyncEnforcer(get_fixture("rbac_model.conf"), adapter)
await e.load_policy()

# Add permissions
await e.add_permission_for_user('alice', 'data1', 'read')
await e.add_permission_for_user('alice', 'data2', 'read')
await e.add_permission_for_user("alice", "data1", "read")
await e.add_permission_for_user("alice", "data2", "read")

# Verify permissions are available in current session
self.assertTrue(e.enforce('alice', 'data1', 'read'))
self.assertTrue(e.enforce('alice', 'data2', 'read'))
self.assertTrue(e.enforce("alice", "data1", "read"))
self.assertTrue(e.enforce("alice", "data2", "read"))

# Commit the transaction
await external_session.commit()

# 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'))
self.assertTrue(new_enforcer.enforce('alice', 'data2', 'read'))

self.assertTrue(new_enforcer.enforce("alice", "data1", "read"))
self.assertTrue(new_enforcer.enforce("alice", "data2", "read"))

finally:
# Clean up
if os.path.exists(db_file.name):
Expand All @@ -91,48 +95,52 @@ async def test_external_session_commit(self):
async def test_external_session_rollback(self):
"""Test using external session with rollback."""
# Create a temporary database file
db_file = tempfile.NamedTemporaryFile(suffix='.db', delete=False)
db_file = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
db_file.close()

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)

# Test with external session
async with async_session_factory() as external_session:
# Create adapter with external session
adapter = Adapter(engine, db_session=external_session)

# Create table
await adapter.create_table()

# Create enforcer
e = casbin.AsyncEnforcer(get_fixture("rbac_model.conf"), adapter)
await e.load_policy()

# Add permissions
await e.add_permission_for_user('alice', 'data1', 'read')
await e.add_permission_for_user('alice', 'data2', 'read')
await e.add_permission_for_user("alice", "data1", "read")
await e.add_permission_for_user("alice", "data2", "read")

# Verify permissions are available in current session
self.assertTrue(e.enforce('alice', 'data1', 'read'))
self.assertTrue(e.enforce('alice', 'data2', 'read'))
self.assertTrue(e.enforce("alice", "data1", "read"))
self.assertTrue(e.enforce("alice", "data2", "read"))

# Rollback the transaction
await external_session.rollback()

# 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'))
self.assertFalse(new_enforcer.enforce('alice', 'data2', 'read'))

self.assertFalse(new_enforcer.enforce("alice", "data1", "read"))
self.assertFalse(new_enforcer.enforce("alice", "data2", "read"))

finally:
# Clean up
if os.path.exists(db_file.name):
Expand All @@ -141,47 +149,51 @@ async def test_external_session_rollback(self):
async def test_external_session_with_save_policy(self):
"""Test save_policy with external session."""
# Create a temporary database file
db_file = tempfile.NamedTemporaryFile(suffix='.db', delete=False)
db_file = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
db_file.close()

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)

# Test with external session
async with async_session_factory() as external_session:
# Create adapter with external session
adapter = Adapter(engine, db_session=external_session)

# Create table
await adapter.create_table()

# Create enforcer
e = casbin.AsyncEnforcer(get_fixture("rbac_model.conf"), adapter)
await e.load_policy()

# Add permissions
await e.add_permission_for_user('alice', 'data1', 'read')
await e.add_permission_for_user('bob', 'data2', 'write')
await e.add_permission_for_user("alice", "data1", "read")
await e.add_permission_for_user("bob", "data2", "write")

# Save policy (should use external session)
await e.save_policy()

# Commit the transaction
await external_session.commit()

# 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'))
self.assertTrue(new_enforcer.enforce('bob', 'data2', 'write'))

self.assertTrue(new_enforcer.enforce("alice", "data1", "read"))
self.assertTrue(new_enforcer.enforce("bob", "data2", "write"))

finally:
# Clean up
if os.path.exists(db_file.name):
Expand All @@ -190,39 +202,43 @@ async def test_external_session_with_save_policy(self):
async def test_backward_compatibility(self):
"""Test that existing behavior is preserved."""
# Create a temporary database file
db_file = tempfile.NamedTemporaryFile(suffix='.db', delete=False)
db_file = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
db_file.close()

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)

# Create table
await adapter.create_table()

# Create enforcer
e = casbin.AsyncEnforcer(get_fixture("rbac_model.conf"), adapter)
await e.load_policy()

# Add permissions (should auto-commit)
await e.add_permission_for_user('alice', 'data1', 'read')
await e.add_permission_for_user('alice', 'data2', 'read')
await e.add_permission_for_user("alice", "data1", "read")
await e.add_permission_for_user("alice", "data2", "read")

# Verify permissions are committed automatically
self.assertTrue(e.enforce('alice', 'data1', 'read'))
self.assertTrue(e.enforce('alice', 'data2', 'read'))
self.assertTrue(e.enforce("alice", "data1", "read"))
self.assertTrue(e.enforce("alice", "data2", "read"))

# 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'))
self.assertTrue(new_enforcer.enforce('alice', 'data2', 'read'))

self.assertTrue(new_enforcer.enforce("alice", "data1", "read"))
self.assertTrue(new_enforcer.enforce("alice", "data2", "read"))

finally:
# Clean up
if os.path.exists(db_file.name):
Expand Down