Skip to content

Commit e7aa07f

Browse files
authored
feat: improve performance by doing bulk insert in add_policies() API (#9)
1 parent 0abab6a commit e7aa07f

2 files changed

Lines changed: 52 additions & 9 deletions

File tree

casbin_async_sqlalchemy_adapter/adapter.py

Lines changed: 15 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717

1818
from casbin import persist
1919
from casbin.persist.adapters.asyncio import AsyncAdapter
20-
from sqlalchemy import Column, Integer, String, delete
20+
from sqlalchemy import Column, Integer, String, delete, insert
2121
from sqlalchemy import or_
2222
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession
2323
from sqlalchemy.future import select
@@ -190,14 +190,20 @@ async def add_policy(self, sec, ptype, rule):
190190

191191
async def add_policies(self, sec, ptype, rules):
192192
"""adds a policy rules to the storage."""
193-
if self._external_session is not None:
194-
# Use external session to add all rules in the same transaction
195-
for rule in rules:
196-
await self._save_policy_line(ptype, rule, self._external_session)
197-
else:
198-
# Use individual sessions for each rule (original behavior)
199-
for rule in rules:
200-
await self._save_policy_line(ptype, rule)
193+
if not rules:
194+
return
195+
196+
# Build rows for executemany bulk insert
197+
rows = []
198+
for rule in rules:
199+
row = {"ptype": ptype}
200+
for i, v in enumerate(rule):
201+
row[f"v{i}"] = v
202+
rows.append(row)
203+
204+
async with self._session_scope() as session:
205+
stmt = insert(self._db_class)
206+
await session.execute(stmt, rows)
201207

202208
async def remove_policy(self, sec, ptype, rule):
203209
"""removes a policy rule from the storage."""

tests/test_adapter.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -399,5 +399,42 @@ async def test_update_filtered_policies(self):
399399
self.assertTrue(e.enforce("bob", "data2", "read"))
400400

401401

402+
class TestBulkInsert(IsolatedAsyncioTestCase):
403+
async def test_add_policies_bulk_internal_session(self):
404+
engine = create_async_engine("sqlite+aiosqlite://", future=True)
405+
adapter = Adapter(engine)
406+
await adapter.create_table()
407+
408+
rules = [
409+
("u1", "obj1", "read"),
410+
("u2", "obj2", "write"),
411+
("u3", "obj3", "read"),
412+
]
413+
await adapter.add_policies("p", "p", rules)
414+
415+
async_session = async_sessionmaker(
416+
engine, expire_on_commit=False, class_=AsyncSession
417+
)
418+
async with async_session() as s:
419+
# count inserted rows
420+
from sqlalchemy import select, func
421+
422+
cnt = await s.execute(
423+
select(func.count())
424+
.select_from(CasbinRule)
425+
.where(CasbinRule.ptype == "p")
426+
)
427+
assert cnt.scalar_one() == len(rules)
428+
429+
rows = (
430+
(await s.execute(select(CasbinRule).order_by(CasbinRule.id)))
431+
.scalars()
432+
.all()
433+
)
434+
tuples = [(r.v0, r.v1, r.v2) for r in rows]
435+
for r in rules:
436+
assert r in tuples
437+
438+
402439
if __name__ == "__main__":
403440
unittest.main()

0 commit comments

Comments
 (0)