Skip to content
Open
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
3 changes: 3 additions & 0 deletions casbin/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,9 @@
from .distributed_enforcer import DistributedEnforcer
from .fast_enforcer import FastEnforcer
from .async_enforcer import AsyncEnforcer
from .cached_enforcer import CachedEnforcer
from .synced_cached_enforcer import SyncedCachedEnforcer
from .async_cached_enforcer import AsyncCachedEnforcer
from . import util
from .persist import *
from .effect import *
Expand Down
97 changes: 97 additions & 0 deletions casbin/async_cached_enforcer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
from casbin.cache import Cache, ErrNoSuchKey
from casbin.cache.default_cache import DefaultCache
from casbin.async_enforcer import AsyncEnforcer


class AsyncCachedEnforcer(AsyncEnforcer):
def __init__(self, *args, **kwargs):
self._enable_cache = True
self._cache = DefaultCache()
self._expire_time = None
super().__init__(*args, **kwargs)

def enable_cache(self, enable_cache):
self._enable_cache = enable_cache

def set_expire_time(self, expire_time):
self._expire_time = expire_time

def set_cache(self, cache):
self._cache = cache

def invalidate_cache(self):
self._cache.clear()

@staticmethod
def get_cache_key(*params):
parts = []
for p in params:
if isinstance(p, str):
parts.append(p)
elif hasattr(p, "get_cache_key"):
parts.append(p.get_cache_key())
else:
return None
parts.append("$$")
return "".join(parts)

def enforce(self, *rvals):
if not self._enable_cache:
return super().enforce(*rvals)

key = self.get_cache_key(*rvals)
if key is None:
return super().enforce(*rvals)

try:
return self._cache.get(key)
except ErrNoSuchKey:
pass

result = super().enforce(*rvals)
self._cache.set(key, result, self._expire_time)
return result

async def load_policy(self):
self.invalidate_cache()
return await super().load_policy()

async def clear_policy(self):
self.invalidate_cache()
return await super().clear_policy()

async def _add_policy(self, sec, ptype, rule):
self.invalidate_cache()
return await super()._add_policy(sec, ptype, rule)

async def _add_policies(self, sec, ptype, rules):
self.invalidate_cache()
return await super()._add_policies(sec, ptype, rules)

async def _update_policy(self, sec, ptype, old_rule, new_rule):
self.invalidate_cache()
return await super()._update_policy(sec, ptype, old_rule, new_rule)

async def _update_policies(self, sec, ptype, old_rules, new_rules):
self.invalidate_cache()
return await super()._update_policies(sec, ptype, old_rules, new_rules)

async def _update_filtered_policies(self, sec, ptype, new_rules, field_index, *field_values):
self.invalidate_cache()
return await super()._update_filtered_policies(sec, ptype, new_rules, field_index, *field_values)

async def _remove_policy(self, sec, ptype, rule):
self.invalidate_cache()
return await super()._remove_policy(sec, ptype, rule)

async def _remove_policies(self, sec, ptype, rules):
self.invalidate_cache()
return await super()._remove_policies(sec, ptype, rules)

async def _remove_filtered_policy(self, sec, ptype, field_index, *field_values):
self.invalidate_cache()
return await super()._remove_filtered_policy(sec, ptype, field_index, *field_values)

async def _remove_filtered_policy_returns_effects(self, sec, ptype, field_index, *field_values):
self.invalidate_cache()
return await super()._remove_filtered_policy_returns_effects(sec, ptype, field_index, *field_values)
23 changes: 23 additions & 0 deletions casbin/cache/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
from abc import ABC, abstractmethod


class ErrNoSuchKey(LookupError):
pass


class Cache(ABC):
@abstractmethod
def set(self, key, value, *extra):
pass

@abstractmethod
def get(self, key):
pass

@abstractmethod
def delete(self, key):
pass

@abstractmethod
def clear(self):
pass
28 changes: 28 additions & 0 deletions casbin/cache/default_cache.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
import time

from casbin.cache import Cache, ErrNoSuchKey


class DefaultCache(Cache):
def __init__(self):
self._store = {}

def set(self, key, value, *extra):
ttl = extra[0] if extra else None
expires_at = time.time() + ttl if ttl and ttl > 0 else None
self._store[key] = (value, expires_at)

def get(self, key):
if key not in self._store:
raise ErrNoSuchKey(key)
value, expires_at = self._store[key]
if expires_at and time.time() > expires_at:
del self._store[key]
raise ErrNoSuchKey(key)
return value

def delete(self, key):
self._store.pop(key, None)

def clear(self):
self._store.clear()
101 changes: 101 additions & 0 deletions casbin/cached_enforcer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
from casbin.cache import Cache, ErrNoSuchKey
from casbin.cache.default_cache import DefaultCache
from casbin.enforcer import Enforcer


class CachedEnforcer(Enforcer):
def __init__(self, *args, **kwargs):
self._enable_cache = True
self._cache = DefaultCache()
self._expire_time = None
super().__init__(*args, **kwargs)

def enable_cache(self, enable_cache):
self._enable_cache = enable_cache

def set_expire_time(self, expire_time):
self._expire_time = expire_time

def set_cache(self, cache):
self._cache = cache

def invalidate_cache(self):
self._cache.clear()

@staticmethod
def get_cache_key(*params):
parts = []
for p in params:
if isinstance(p, str):
parts.append(p)
elif hasattr(p, "get_cache_key"):
parts.append(p.get_cache_key())
else:
return None
parts.append("$$")
return "".join(parts)

def enforce(self, *rvals):
if not self._enable_cache:
return super().enforce(*rvals)

key = self.get_cache_key(*rvals)
if key is None:
return super().enforce(*rvals)

try:
return self._cache.get(key)
except ErrNoSuchKey:
pass

result = super().enforce(*rvals)
self._cache.set(key, result, self._expire_time)
return result

def load_policy(self):
self.invalidate_cache()
return super().load_policy()

def clear_policy(self):
self.invalidate_cache()
return super().clear_policy()

def _add_policy(self, sec, ptype, rule):
self.invalidate_cache()
return super()._add_policy(sec, ptype, rule)

def _add_policies(self, sec, ptype, rules):
self.invalidate_cache()
return super()._add_policies(sec, ptype, rules)

def _add_policies_ex(self, sec, ptype, rules):
self.invalidate_cache()
return super()._add_policies_ex(sec, ptype, rules)

def _update_policy(self, sec, ptype, old_rule, new_rule):
self.invalidate_cache()
return super()._update_policy(sec, ptype, old_rule, new_rule)

def _update_policies(self, sec, ptype, old_rules, new_rules):
self.invalidate_cache()
return super()._update_policies(sec, ptype, old_rules, new_rules)

def _update_filtered_policies(self, sec, ptype, new_rules, field_index, *field_values):
self.invalidate_cache()
return super()._update_filtered_policies(sec, ptype, new_rules, field_index, *field_values)

def _remove_policy(self, sec, ptype, rule):
self.invalidate_cache()
return super()._remove_policy(sec, ptype, rule)

def _remove_policies(self, sec, ptype, rules):
self.invalidate_cache()
return super()._remove_policies(sec, ptype, rules)

def _remove_filtered_policy(self, sec, ptype, field_index, *field_values):
self.invalidate_cache()
return super()._remove_filtered_policy(sec, ptype, field_index, *field_values)

def _remove_filtered_policy_returns_effects(self, sec, ptype, field_index, *field_values):
self.invalidate_cache()
return super()._remove_filtered_policy_returns_effects(sec, ptype, field_index, *field_values)
7 changes: 7 additions & 0 deletions casbin/core_enforcer.py
Original file line number Diff line number Diff line change
Expand Up @@ -313,6 +313,13 @@ def enable_auto_notify_watcher(self, auto_notify_watcher):
"""controls whether to save a policy rule automatically notify the watcher when it is added or removed."""
self.auto_notify_watcher = auto_notify_watcher

def enable_g_function_cache(self, enabled):
"""controls whether to cache g() function results."""
for rm in self.rm_map.values():
rm.enable_g_cache(enabled)
for crm in self.cond_rm_map.values():
crm.enable_g_cache(enabled)

def build_role_links(self):
"""manually rebuild the role inheritance relations."""

Expand Down
Loading