from __future__ import annotations
import re
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Self
import httpx
from sqlalchemy import and_, false, func, not_, or_, true
from sqlalchemy_utils import LtreeType
from fsh_lib.rbac import (
RoleBindingScopeMixin, # noqa: F401 (re-export; was defined here)
)
if TYPE_CHECKING:
from collections.abc import Mapping, Sequence
from sqlalchemy.sql.elements import ColumnElement
from fsh_lib.actions import ActionSpec
from fsh_lib.resource_registry import ResourceEntry, ResourceRegistry
_DEFAULT_TIMEOUT = 5.0
[docs]
class OpaError(RuntimeError):
pass
[docs]
@dataclass(frozen=True)
class Subject:
type: str
id: str
user: Mapping[str, Any] | None = None
def as_input(self) -> dict[str, Any]:
subject: dict[str, Any] = {"type": self.type, "id": self.id}
if self.user is not None:
subject["user"] = dict(self.user)
return subject
[docs]
@dataclass(frozen=True)
class ResourceRef:
type: str
id: str | None = None
attributes: Mapping[str, Any] = field(default_factory=dict)
def as_input(self) -> dict[str, Any]:
resource: dict[str, Any] = dict(self.attributes)
resource["type"] = self.type
if self.id is not None:
resource["id"] = self.id
return resource
[docs]
class ScopeError(ValueError):
pass
def scope_fields(opa_rep_class: type) -> dict[str, type]:
json_py: dict[str, type] = {
"string": str,
"integer": int,
"number": float,
"boolean": bool,
}
fields: dict[str, type] = {}
properties = opa_rep_class.model_json_schema().get("properties", {})
for name, prop in properties.items():
if name == "type": # the resource-type discriminator, not a scope field
continue
json_type = _json_boundary_type(prop)
if json_type in json_py:
fields[name] = json_py[json_type]
return fields
def _json_boundary_type(prop: Mapping[str, Any]) -> str | None:
if "type" in prop:
return prop["type"]
for sub in prop.get("anyOf", []):
sub_type = sub.get("type")
if sub_type not in (None, "null"):
return sub_type
return None
def parse_scope_attrs(
attrs: Mapping[str, Any],
allowed: Mapping[str, type],
) -> dict[str, Any]:
if not attrs:
msg = "attribute scope is empty; use a type-scoped binding instead"
raise ScopeError(msg)
validated: dict[str, Any] = {}
for key, value in attrs.items():
if key not in allowed:
msg = f"unknown scope field {key!r}; allowed: {sorted(allowed)}"
raise ScopeError(msg)
expected = allowed[key]
if expected is not object and type(value) is not expected:
msg = (
f"scope field {key!r} expects {expected.__name__}, "
f"got {type(value).__name__} ({value!r})"
)
raise ScopeError(msg)
validated[key] = value
return validated
[docs]
@dataclass(frozen=True)
class RoleBinding:
role: str
scope: ResourceRef | None = None
def as_input(self) -> dict[str, Any]:
scope: dict[str, Any] | None = None
if self.scope is not None:
scope = {"type": self.scope.type}
if self.scope.id is not None:
scope["id"] = self.scope.id
if self.scope.attributes:
scope["attrs"] = dict(self.scope.attributes)
return {"role": self.role, "object": scope}
@classmethod
def from_row(
cls,
role: str,
*,
object_type: str | None = None,
object_id: str | None = None,
scope_attrs: Mapping[str, Any] | None = None,
allowed: Mapping[str, type] | None = None,
) -> Self:
if object_type is None:
return cls(role)
attrs = {k: v for k, v in (scope_attrs or {}).items() if v is not None}
if attrs and allowed is not None:
attrs = parse_scope_attrs(attrs, allowed)
return cls(role, ResourceRef(object_type, object_id, attributes=attrs))
def _subject_input(
subject: Subject,
bindings: Sequence[RoleBinding],
) -> dict[str, Any]:
doc = subject.as_input()
doc["bindings"] = [binding.as_input() for binding in bindings]
return doc
[docs]
@dataclass(frozen=True)
class Decision:
permit: bool
allow: bool
deny: bool
raw: Mapping[str, Any]
@classmethod
def from_result(cls, result: Mapping[str, Any]) -> Decision:
return cls(
permit=bool(result.get("permit", False)),
allow=bool(result.get("allow", False)),
deny=bool(result.get("deny", False)),
raw=result,
)
[docs]
@dataclass(frozen=True)
class ActionPermissions:
object_scoped: Mapping[str, set[str]]
collection_scoped: set[str]
type ObjectAttributes = tuple[str, Mapping[str, Any]]
# OPA returns verdicts as a list, need to keep track of the
# indices of the results/verdicts
[docs]
@dataclass(frozen=True)
class VerdictRange:
view: int
actions_start: int
[docs]
@dataclass(frozen=True)
class VerdictIndices:
objects: Mapping[str, VerdictRange]
collection: VerdictRange
_ACTION_VERB_OVERRIDE = {"get": "view"}
def action_string(resource_type: str, action_name: str) -> str:
verb = _ACTION_VERB_OVERRIDE.get(action_name, action_name)
return f"{resource_type.lower()}:{verb}"
def _empty_action_permissions(
objects: Sequence[ObjectAttributes],
) -> ActionPermissions:
return ActionPermissions(
object_scoped={object_id: set() for object_id, _ in objects},
collection_scoped=set(),
)
def _full_action_permissions(
*,
objects: Sequence[ObjectAttributes],
object_actions: Sequence[ActionSpec],
collection_actions: Sequence[ActionSpec],
) -> ActionPermissions:
return ActionPermissions(
object_scoped={
object_id: {action.name for action in object_actions}
for object_id, _ in objects
},
collection_scoped={action.name for action in collection_actions},
)
def _build_action_checks(
*,
scopes: Mapping[str, Sequence[ObjectAttributes]],
registry: ResourceRegistry[Any],
) -> tuple[list[tuple[str, ResourceRef]], dict[str, VerdictIndices]]:
checks: list[tuple[str, ResourceRef]] = []
indices: dict[str, VerdictIndices] = {}
for resource_type, objects in scopes.items():
resource = registry[resource_type].resource
view_action = action_string(resource_type, "view")
object_indices = {
object_id: _append_object_checks(
checks=checks,
resource_type=resource_type,
object_id=object_id,
attributes=attributes,
view_action=view_action,
actions=resource.object_actions,
)
for object_id, attributes in objects
}
collection_indices = _append_collection_checks(
checks=checks,
resource_type=resource_type,
view_action=view_action,
actions=resource.collection_actions,
)
indices[resource_type] = VerdictIndices(
objects=object_indices,
collection=collection_indices,
)
return checks, indices
def _append_object_checks(
*,
checks: list[tuple[str, ResourceRef]],
resource_type: str,
object_id: str,
attributes: Mapping[str, Any],
view_action: str,
actions: Sequence[ActionSpec],
) -> VerdictRange:
resource = ResourceRef(resource_type, object_id, attributes)
view_index = len(checks)
checks.append((view_action, resource))
actions_start = len(checks)
checks.extend(
(action_string(resource_type, action.name), resource)
for action in actions
)
return VerdictRange(
view=view_index,
actions_start=actions_start,
)
def _append_collection_checks(
*,
checks: list[tuple[str, ResourceRef]],
resource_type: str,
view_action: str,
actions: Sequence[ActionSpec],
) -> VerdictRange:
resource = ResourceRef(resource_type)
view_index = len(checks)
checks.append((view_action, resource))
actions_start = len(checks)
checks.extend(
(action_string(resource_type, action.name), resource)
for action in actions
)
return VerdictRange(
view=view_index,
actions_start=actions_start,
)
def _resolve_actions(
*,
actions: Sequence[ActionSpec],
verdicts: Sequence[bool],
check_range: VerdictRange,
) -> set[str]:
if not verdicts[check_range.view]:
return set()
return {
action.name
for offset, action in enumerate(actions)
if verdicts[check_range.actions_start + offset]
}
def _resolve_action_permissions(
*,
scopes: Mapping[str, Sequence[ObjectAttributes]],
indices: Mapping[str, VerdictIndices],
verdicts: Sequence[bool],
registry: ResourceRegistry[Any],
) -> dict[str, ActionPermissions]:
return {
resource_type: _resolve_resource_permissions(
objects=objects,
indices=indices[resource_type],
resource=registry[resource_type].resource,
verdicts=verdicts,
)
for resource_type, objects in scopes.items()
}
def _resolve_resource_permissions(
*,
objects: Sequence[ObjectAttributes],
indices: VerdictIndices,
resource: ResourceEntry,
verdicts: Sequence[bool],
) -> ActionPermissions:
return ActionPermissions(
object_scoped={
object_id: _resolve_actions(
actions=resource.object_actions,
verdicts=verdicts,
check_range=indices.objects[object_id],
)
for object_id, _ in objects
},
collection_scoped=_resolve_actions(
actions=resource.collection_actions,
verdicts=verdicts,
check_range=indices.collection,
),
)
def _fallback_action_permissions(
*,
scopes: Mapping[str, Sequence[ObjectAttributes]],
registry: ResourceRegistry[Any],
allow: bool,
) -> dict[str, ActionPermissions]:
return {
resource_type: (
_full_action_permissions(
objects=objects,
object_actions=registry[resource_type].resource.object_actions,
collection_actions=registry[
resource_type
].resource.collection_actions,
)
if allow
else _empty_action_permissions(objects)
)
for resource_type, objects in scopes.items()
}
def _opa_rep_unknowns(opa_rep_class: type | None) -> list[str]:
if opa_rep_class is None:
return ["input.resource.id"]
return [
*(
f"input.resource.{name}"
for name in opa_rep_class.model_fields
if name not in ("id", "type")
),
"input.resource.id",
]
async def resolve_action_permissions(
*,
client: OpaClient,
subject: Subject,
roles: Mapping[str, Any],
bindings: Sequence[RoleBinding],
scopes: Mapping[str, Sequence[ObjectAttributes]],
registry: ResourceRegistry[Any],
fail_open: bool,
) -> dict[str, ActionPermissions]:
checks, indices = _build_action_checks(
scopes=scopes,
registry=registry,
)
try:
verdicts = list(
await client.check_many(
subject=subject,
roles=roles,
bindings=bindings,
items=checks,
)
)
except OpaError:
return _fallback_action_permissions(
scopes=scopes,
registry=registry,
allow=fail_open,
)
# A bare collection point-check carries no row attributes, so an
# attribute-scoped grant (a subtree) can never satisfy it. Derive
# collection "view" from the compiled row filter instead: viewable
# means the filter is not deny-all (some row passes). The view
# verdict shows up twice -- once as the gate, once as the `search`
# action (whose operation is `view`) -- so override both.
for resource_type in scopes:
resource = registry[resource_type].resource
collection = indices[resource_type].collection
view_action = action_string(resource_type, "view")
try:
result = await client.compile_filter(
subject=subject,
action=view_action,
resource_type=resource_type,
roles=roles,
bindings=bindings,
unknowns=_opa_rep_unknowns(resource.opa_rep_class),
)
viewable = not result.always_deny
except OpaError:
# Leave the point-check verdicts as the fallback.
continue
verdicts[collection.view] = viewable
for offset, spec in enumerate(resource.collection_actions):
if action_string(resource_type, spec.name) == view_action:
verdicts[collection.actions_start + offset] = viewable
return _resolve_action_permissions(
scopes=scopes,
indices=indices,
verdicts=verdicts,
registry=registry,
)
[docs]
@dataclass(frozen=True)
class Condition:
field: str
op: str
value: Any
negated: bool = False
_LTREE_LABEL = re.compile(r"[A-Za-z0-9_]+")
def _eq_clause(
col: ColumnElement[Any],
value: Any, # noqa: ANN401
) -> ColumnElement[bool]:
# An ltree column compared to a string literal must cast the
# literal to ltree on the server: a raw Python str binds as text
# and `ltree = text` does not resolve.
if isinstance(col.type, LtreeType) and isinstance(value, str):
if not all(_LTREE_LABEL.fullmatch(lbl) for lbl in value.split(".")):
return false()
return col == func.text2ltree(value)
return col == value
def _startswith_clause(
col: ColumnElement[Any],
value: Any, # noqa: ANN401
) -> ColumnElement[bool]:
# A trailing dot marks the label-boundary form
# (``startswith(p, root + ".")``): the policy expresses it as
# strict ltree descent. ``root`` itself has no trailing dot, so
# ``descendant_of(root)`` must exclude ``root`` for the residual
# to stay faithful under negation; the rule's separate
# ``== root`` clause covers the root row. Anything else stays a
# plain string prefix so the SQL matches OPA's string semantics.
if (
isinstance(col.type, LtreeType)
and isinstance(value, str)
and value.endswith(".")
):
root = value[:-1]
if not all(_LTREE_LABEL.fullmatch(lbl) for lbl in root.split(".")):
# A malformed root (empty label) can match no row's
# ltree; fail only this disjunct, not the whole filter.
return false()
ltree_root = func.text2ltree(root)
return and_(col.descendant_of(ltree_root), col != ltree_root)
return col.startswith(value)
#: Builds a SQLAlchemy clause from a :class:`Condition`'s op.
_CLAUSE_BUILDERS = {
"eq": _eq_clause,
"ne": lambda col, value: col != value,
"lt": lambda col, value: col < value,
"le": lambda col, value: col <= value,
"gt": lambda col, value: col > value,
"ge": lambda col, value: col >= value,
"in": lambda col, value: col.in_(value),
"startswith": _startswith_clause,
}
[docs]
@dataclass(frozen=True)
class FilterResult:
always_allow: bool
always_deny: bool
conjunctions: tuple[tuple[Condition, ...], ...]
def to_sqlalchemy(
self,
columns: Mapping[str, ColumnElement[Any]],
) -> ColumnElement[bool]:
if self.always_deny:
return false()
if self.always_allow:
return true()
disjuncts = [
and_(*(self._clause(cond, columns) for cond in conjunction))
for conjunction in self.conjunctions
]
return or_(*disjuncts)
@staticmethod
def _clause(
cond: Condition,
columns: Mapping[str, ColumnElement[Any]],
) -> ColumnElement[bool]:
column = columns.get(cond.field)
if column is None:
msg = (
f"residual constrains resource field {cond.field!r}, "
f"which is not in the supplied column map "
f"{sorted(columns)}"
)
raise OpaError(msg)
clause = _CLAUSE_BUILDERS[cond.op](column, cond.value)
return not_(clause) if cond.negated else clause
class OpaClient:
def __init__(
self,
base_url: str,
*,
opa_package: str = "authz",
client: httpx.AsyncClient | None = None,
timeout: float = _DEFAULT_TIMEOUT,
) -> None:
self._opa_package = opa_package
self._decision_path = f"/v1/data/{opa_package}/decision"
self._bulk_path = f"/v1/data/{opa_package}/decisions"
self._owns_client = client is None
self._client = client or httpx.AsyncClient(
base_url=base_url,
timeout=timeout,
)
async def __aenter__(self) -> Self:
return self
async def __aexit__(self, *_exc: object) -> None:
await self.aclose()
async def aclose(self) -> None:
if self._owns_client:
await self._client.aclose()
async def check(
self,
*,
subject: Subject,
action: str,
resource: ResourceRef,
roles: Mapping[str, Any],
bindings: Sequence[RoleBinding],
format: str | None = None, # noqa: A002
representation: str | None = None,
) -> Decision:
document: dict[str, Any] = {
"subject": _subject_input(subject, bindings),
"action": action,
"resource": resource.as_input(),
"roles": dict(roles),
}
if format is not None:
document["format"] = format
if representation is not None:
document["representation"] = representation
response = await self._post(self._decision_path, {"input": document})
return Decision.from_result(response.get("result", {}))
async def check_many(
self,
*,
subject: Subject,
roles: Mapping[str, Any],
bindings: Sequence[RoleBinding],
items: Sequence[tuple[str, ResourceRef]],
format: str | None = None, # noqa: A002
representation: str | None = None,
) -> list[bool]:
subject_input = _subject_input(subject, bindings)
roles_input = dict(roles)
queries: list[dict[str, Any]] = []
for action, resource in items:
query: dict[str, Any] = {
"subject": subject_input,
"action": action,
"resource": resource.as_input(),
"roles": roles_input,
}
if format is not None:
query["format"] = format
if representation is not None:
query["representation"] = representation
queries.append(query)
response = await self._post(
self._bulk_path,
{"input": {"queries": queries}},
)
return _parse_bulk_result(response.get("result", []), len(items))
async def compile_filter(
self,
*,
subject: Subject,
action: str,
resource_type: str,
roles: Mapping[str, Any],
bindings: Sequence[RoleBinding],
unknowns: Sequence[str] = ("input.resource.id",),
format: str | None = None, # noqa: A002
representation: str | None = None,
) -> FilterResult:
input_doc: dict[str, Any] = {
"subject": _subject_input(subject, bindings),
"action": action,
"resource": {"type": resource_type},
"roles": dict(roles),
}
if format is not None:
input_doc["format"] = format
if representation is not None:
input_doc["representation"] = representation
body = {
"query": f"data.{self._opa_package}.allow == true",
"input": input_doc,
"unknowns": list(unknowns),
}
response = await self._post("/v1/compile", body)
return _parse_compile_result(
response.get("result", {}),
opa_package=self._opa_package,
)
async def _post(
self,
path: str,
body: Mapping[str, Any],
) -> dict[str, Any]:
try:
response = await self._client.post(path, json=dict(body))
response.raise_for_status()
except httpx.HTTPError as exc:
msg = f"permissions service call failed: {exc}"
raise OpaError(msg) from exc
return response.json()
def _parse_bulk_result(result: Any, count: int) -> list[bool]: # noqa: ANN401
by_index: dict[int, bool] = {}
if isinstance(result, list):
for entry in result:
index = entry.get("index")
if isinstance(index, int):
by_index[index] = bool(entry.get("permit", False))
return [by_index.get(position, False) for position in range(count)]
_OPERATORS: dict[str, tuple[str, str]] = {
"equal": ("eq", "eq"),
"eq": ("eq", "eq"),
"neq": ("ne", "ne"),
"lt": ("lt", "gt"),
"lte": ("le", "ge"),
"gt": ("gt", "lt"),
"gte": ("ge", "le"),
"internal.member_2": ("in", "in"),
"startswith": ("startswith", "startswith"),
}
#: Prefix every residual reference this translator understands
#: starts with. Generic RBAC list filtering leaves the resource
#: unknown, so every surviving constraint is on a resource field.
_RESOURCE_PREFIX = ("input", "resource")
#: A FilterResult that passes every row.
_ALWAYS_ALLOW = FilterResult(
always_allow=True,
always_deny=False,
conjunctions=(),
)
#: A FilterResult that rejects every row.
_ALWAYS_DENY = FilterResult(
always_allow=False,
always_deny=True,
conjunctions=(),
)
def _parse_compile_result(
result: Mapping[str, Any],
*,
opa_package: str | None = None,
) -> FilterResult:
queries = result.get("queries")
# No `queries` key at all -- the query is unsatisfiable, so no
# row can ever pass.
if queries is None:
return _ALWAYS_DENY
walker = _SupportWalker(result.get("support") or [], opa_package)
disjuncts: list[tuple[Condition, ...]] = []
for query in queries:
disjuncts.extend(walker.query_dnf(query))
return _from_dnf(disjuncts)
def _from_dnf(disjuncts: list[tuple[Condition, ...]]) -> FilterResult:
if any(not conjunction for conjunction in disjuncts):
return _ALWAYS_ALLOW
if not disjuncts:
return _ALWAYS_DENY
# The same conjunction can arrive through several support
# paths; keep the first of each. Deduped by equality, not
# hash -- an ``in`` condition holds an unhashable list.
unique: list[tuple[Condition, ...]] = []
for conjunction in disjuncts:
if conjunction not in unique:
unique.append(conjunction)
return FilterResult(
always_allow=False,
always_deny=False,
conjunctions=tuple(unique),
)
class _SupportWalker:
def __init__(
self,
support: list[Mapping[str, Any]],
opa_package: str | None,
) -> None:
self._rules: dict[tuple[str, ...], list[Mapping[str, Any]]] = {}
for entry in support:
package_path = _support_package_path(entry)
for rule in entry.get("rules", []):
name = _support_rule_name(rule)
if name is None:
continue
key = (*package_path, name)
self._rules.setdefault(key, []).append(rule)
self._rule_package_prefix = (
("data", opa_package, "rule") if opa_package is not None else None
)
self._expanding: set[tuple[str, ...]] = set()
def query_dnf(
self,
exprs: list[Mapping[str, Any]],
) -> list[tuple[Condition, ...]]:
acc: list[tuple[Condition, ...]] = [()]
for expr in exprs:
expr_disjuncts = self._expr_dnf(expr)
acc = [(*left, *right) for left in acc for right in expr_disjuncts]
return acc
def _expr_dnf(
self,
expr: Mapping[str, Any],
) -> list[tuple[Condition, ...]]:
literal = _boolean_literal(expr)
if literal is not None:
if expr.get("negated"):
literal = not literal
return [()] if literal else []
terms = expr.get("terms")
# A single term is a bare reference (a support rule or an
# existence check), not a comparison call.
if isinstance(terms, dict):
return self._ref_dnf(expr, terms)
presence = self._presence_dnf(expr, terms)
if presence is not None:
return presence
return [(_parse_expr(expr),)]
def _presence_dnf(
self,
expr: Mapping[str, Any],
terms: Any, # noqa: ANN401
) -> list[tuple[Condition, ...]] | None:
if not isinstance(terms, list) or len(terms) != _BINARY_TERMS:
return None
operator_term, left, right = terms
if (
operator_term.get("type") != "ref"
or _operator_name(operator_term) != "neq"
):
return None
if left.get("type") == "null":
ref_term = right
elif right.get("type") == "null":
ref_term = left
else:
return None
path = _ref_path(ref_term)
if path is None:
return None
rules = self._rules.get(tuple(path))
if rules is None:
return None
if not any(rule.get("default") for rule in rules):
# Without a default the rule is undefined at runtime
# whenever no clause fires, so whether it ``!= null``
# depends on the unknowns -- refuse rather than guess.
msg = (
f"presence check on support rule {'.'.join(path)!r} "
f"without a default clause"
)
raise OpaError(msg)
return [] if expr.get("negated") else [()]
def _ref_dnf(
self,
expr: Mapping[str, Any],
term: Mapping[str, Any],
) -> list[tuple[Condition, ...]]:
path = _ref_path(term)
negated = bool(expr.get("negated", False))
if path is None:
msg = (
f"cannot translate residual expression {expr!r}: "
f"unsupported bare reference"
)
raise OpaError(msg)
key = tuple(path)
if key in self._rules:
if negated:
expansion = self._rule_dnf(key)
if not expansion:
return [()]
if any(not clause for clause in expansion):
return []
msg = (
f"unsupported negated support-rule reference "
f"{'.'.join(path)!r}"
)
raise OpaError(msg)
return self._rule_dnf(key)
prefix = self._rule_package_prefix
if prefix is not None and key[: len(prefix)] == prefix:
return [] if negated else [()]
msg = f"unsupported residual reference {'.'.join(path)!r}"
raise OpaError(msg)
def _rule_dnf(self, key: tuple[str, ...]) -> list[tuple[Condition, ...]]:
if key in self._expanding:
msg = f"circular support-rule reference at {'.'.join(key)!r}"
raise OpaError(msg)
self._expanding.add(key)
try:
disjuncts: list[tuple[Condition, ...]] = []
for rule in self._rules[key]:
if rule.get("default"):
# ``default x := false`` contributes nothing; a
# true default makes the rule unconditional.
head_value = rule.get("head", {}).get("value", {})
if head_value.get("value"):
return [()]
continue
disjuncts.extend(self.query_dnf(rule.get("body", [])))
return disjuncts
finally:
self._expanding.discard(key)
def _boolean_literal(expr: Mapping[str, Any]) -> bool | None:
terms = expr.get("terms")
if isinstance(terms, dict) and terms.get("type") == "boolean":
return bool(terms.get("value"))
return None
def _support_package_path(entry: Mapping[str, Any]) -> tuple[str, ...]:
parts = entry.get("package", {}).get("path", [])
return tuple(part.get("value") for part in parts)
def _support_rule_name(rule: Mapping[str, Any]) -> str | None:
return rule.get("head", {}).get("name")
def _ref_path(term: Mapping[str, Any]) -> list[str] | None:
if term.get("type") != "ref":
return None
return [part.get("value") for part in term.get("value", [])]
def _parse_expr(expr: Mapping[str, Any]) -> Condition:
terms = expr.get("terms")
negated = bool(expr.get("negated", False))
if not isinstance(terms, list) or len(terms) != _BINARY_TERMS:
msg = (
f"cannot translate residual expression {expr!r}: only "
f"binary comparisons of a resource field against a "
f"literal are supported"
)
raise OpaError(msg)
operator_term, left, right = terms
builtin = _operator_name(operator_term)
ops = _OPERATORS.get(builtin)
if ops is None:
msg = f"unsupported residual operator {builtin!r}"
raise OpaError(msg)
field, value, flipped = _split_operands(left, right)
if flipped and builtin in {"internal.member_2", "startswith"}:
name = "`in`" if builtin == "internal.member_2" else "startswith"
msg = f"unsupported residual: {name} with the field on the right"
raise OpaError(msg)
op = ops[1] if flipped else ops[0]
return Condition(field=field, op=op, value=value, negated=negated)
_BINARY_TERMS = 3
def _operator_name(term: Mapping[str, Any]) -> str:
if term.get("type") != "ref":
msg = f"expected an operator reference, got {term!r}"
raise OpaError(msg)
return ".".join(part["value"] for part in term["value"])
def _split_operands(
left: Mapping[str, Any],
right: Mapping[str, Any],
) -> tuple[str, Any, bool]:
left_field = _resource_field(left)
right_field = _resource_field(right)
if left_field is not None and right_field is None:
return left_field, _literal(right), False
if right_field is not None and left_field is None:
return right_field, _literal(left), True
msg = (
"cannot translate residual: a comparison must be between "
"exactly one input.resource field and one literal"
)
raise OpaError(msg)
def _resource_field(term: Mapping[str, Any]) -> str | None:
if term.get("type") != "ref":
return None
parts = term["value"]
head = tuple(part.get("value") for part in parts[: len(_RESOURCE_PREFIX)])
if head != _RESOURCE_PREFIX or len(parts) <= len(_RESOURCE_PREFIX):
return None
return ".".join(part["value"] for part in parts[len(_RESOURCE_PREFIX) :])
def _literal(term: Mapping[str, Any]) -> Any: # noqa: ANN401
kind = term.get("type")
if kind in {"string", "number", "boolean", "null"}:
return term.get("value")
if kind in {"array", "set"}:
return [_literal(element) for element in term["value"]]
msg = f"cannot translate residual operand of type {kind!r}"
raise OpaError(msg)