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
42 changes: 25 additions & 17 deletions entitled/rules.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,15 @@

Actor = TypeVar("Actor", contravariant=True)

_POSITIONAL_KINDS = (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
_KEYWORD_KINDS = (
inspect.Parameter.KEYWORD_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)


class RuleProto(Protocol[Actor]):
async def __call__(
Expand All @@ -17,34 +26,33 @@ async def __call__(
class Rule[Actor]:
name: str
callable: RuleProto[Actor]
_signature: inspect.Signature
_positional_count: int
_keyword_names: frozenset[str]

def __init__(self, name: str, callable: RuleProto[Actor]) -> None:
self.name = name
self.callable = callable
self._signature = inspect.signature(callable)
params = self._signature.parameters
self._positional_count = sum(
1 for param in params.values() if param.kind in _POSITIONAL_KINDS
)
self._keyword_names = frozenset(
param_name
for param_name, param in params.items()
if param.kind in _KEYWORD_KINDS
)

async def __call__(
self,
actor: Actor,
*args: Any,
**kwargs: Any,
) -> Response | bool:
sig = inspect.signature(self.callable)
args_count = len(
[
p
for p in sig.parameters.values()
if p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD)
]
)
valid_positionals = (actor,) + args[: args_count - 1]
valid_kwargs = {
k: v
for k, v in kwargs.items()
if k in sig.parameters
and sig.parameters[k].kind
in (inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD)
}
bound = sig.bind_partial(*valid_positionals, **valid_kwargs)
valid_positionals = (actor,) + args[: self._positional_count - 1]
valid_kwargs = {k: v for k, v in kwargs.items() if k in self._keyword_names}
bound = self._signature.bind_partial(*valid_positionals, **valid_kwargs)

return await self.callable(*bound.args, **bound.kwargs)

Expand Down
26 changes: 26 additions & 0 deletions tests/test_rules.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,11 +23,37 @@ async def is_owner(
return Ok() if resource.owner == actor else Err("Not owner on the tenant")


async def can_manage(
actor: User,
*,
access_level: str,
) -> bool:
return actor.tenant is not None and access_level == "admin"


def test_define():
rule = Rule[User]("is_member", is_member)
assert rule.callable == is_member


async def test_binds_keyword_only_arguments():
user = UserFactory(tenant=TenantFactory())
rule = Rule[User]("can_manage", can_manage)

assert await rule.allows(user, access_level="admin")
assert await rule.denies(user, access_level="viewer")


async def test_drops_arguments_the_rule_does_not_declare():
tenant = TenantFactory()
user = UserFactory(tenant=tenant)

assert await Rule[User]("can_manage", can_manage).allows(
user, access_level="admin", unexpected="ignored"
)
assert await Rule[User]("is_member", is_member).allows(user, tenant, "extra", 42)


async def test_allows():
tenant1 = TenantFactory()
tenant2 = TenantFactory()
Expand Down