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
12 changes: 12 additions & 0 deletions examples/query-policy/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
# Query Policy example

This source-based example proves that a context-owned Query Policy receives an
independent query graph, governs root and nested queries once, preserves shared
nested-query identity, leaves the caller request reusable, and fails closed on
denial.

Run it from the repository root:

```bash
PYTHONPATH=src python examples/query-policy/main.py
```
36 changes: 36 additions & 0 deletions examples/query-policy/main.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
from teaql.core.query import FacetRequest, RelationLoad, SelectQuery
from teaql.runtime import UserContext


shared = SelectQuery("SchoolType").project("id")
original = SelectQuery("School").project("id")
original.relations.append(RelationLoad("schoolType", shared))
original.facets.append(FacetRequest("types", "schoolType", shared))
calls = []


def tenant_policy(query):
calls.append(query.entity)
if query.entity == "ForbiddenReport":
raise PermissionError("query policy denied ForbiddenReport")
query.project("tenant_id")


context = UserContext.new().with_request_policy(tenant_policy)
authorized = context.prepare_query(original)

assert authorized is not original
assert authorized.relations[0].query is authorized.facets[0].query
assert calls == ["School", "SchoolType"]
assert "tenant_id" in authorized.projection
assert "tenant_id" not in original.projection
assert "tenant_id" not in shared.projection

try:
context.prepare_query(SelectQuery("ForbiddenReport"))
except PermissionError as error:
assert str(error) == "query policy denied ForbiddenReport"
else:
raise AssertionError("query policy denial must fail closed")

print("PASS Python Query Policy example")
4 changes: 3 additions & 1 deletion scripts/verify-examples.sh
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
set -euo pipefail

repo="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
expected=(conformance mutation-policy order-management school-management task_board)
expected=(business-clock conformance mutation-policy order-management query-policy school-management task_board)
mapfile -t actual < <(find "$repo/examples" -mindepth 1 -maxdepth 1 -type d -printf '%f\n' | sort)
if [[ "${actual[*]}" != "${expected[*]}" ]]; then
echo "example inventory changed; update scripts/verify-examples.sh: ${actual[*]}" >&2
Expand All @@ -25,6 +25,8 @@ PYTHONPATH="$repo/src" python -m unittest discover -s "$repo/examples/conformanc
PYTHONPATH="$repo/examples/school-management:$repo/src" python -m app.main
PYTHONPATH="$repo/src" python -m unittest discover -s "$repo/examples/school-management" -p 'test_sql_log_intent.py' -v
PYTHONPATH="$repo/src" python "$repo/examples/mutation-policy/main.py"
PYTHONPATH="$repo/src" python "$repo/examples/business-clock/main.py"
PYTHONPATH="$repo/src" python "$repo/examples/query-policy/main.py"
order_management_tmp="$(mktemp -d)"
task_board_tmp="$(mktemp -d)"
trap 'rm -rf "$order_management_tmp" "$task_board_tmp"' EXIT
Expand Down
43 changes: 35 additions & 8 deletions src/teaql/runtime/context.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from typing import Dict, Any, Optional, List, Callable, TypeVar
import asyncio
import contextvars
from copy import deepcopy
from dataclasses import dataclass
from array import array
from datetime import date, datetime
Expand Down Expand Up @@ -327,17 +328,43 @@ def require_resource(self, resource_type: str) -> Any:
return res

def prepare_query(self, query: Any) -> Any:
"""Apply trusted request policy exactly once before execution."""
"""Clone a query graph and apply trusted policy once to every query node."""
prepared_root = deepcopy(query)
policy = self.get_resource("request_policy")
if policy is None:
return query
if callable(policy):
prepared = policy(query)
elif hasattr(policy, "apply"):
prepared = policy.apply(query)
else:
return prepared_root

apply_policy = policy if callable(policy) else getattr(policy, "apply", None)
if apply_policy is None or not callable(apply_policy):
raise TypeError("request_policy must be callable or expose apply(query)")
return query if prepared is None else prepared

prepared_nodes: Dict[int, Any] = {}

def prepare_node(node: Any) -> Any:
if node is None:
return None
existing = prepared_nodes.get(id(node))
if existing is not None:
return existing
replacement = apply_policy(node)
prepared = node if replacement is None else replacement
prepared_nodes[id(node)] = prepared
prepared_nodes[id(prepared)] = prepared

for relation in getattr(prepared, "relations", ()):
relation.query = prepare_node(relation.query)
for aggregate in getattr(prepared, "relation_aggregates", ()):
aggregate.query = prepare_node(aggregate.query)
for grouping in getattr(prepared, "object_group_bys", ()):
grouping.query = prepare_node(grouping.query)
for facet in getattr(prepared, "facets", ()):
facet.query = prepare_node(facet.query)
children = getattr(prepared, "child_enhancements", None)
if children is not None:
prepared.child_enhancements = [prepare_node(child) for child in children]
return prepared

return prepare_node(prepared_root)

def set_user_identifier(self, identifier: str):
self._user_identifier = identifier
Expand Down
74 changes: 74 additions & 0 deletions tests/runtime/test_query_policy.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
from teaql.core.query import FacetRequest, RelationLoad, SelectQuery
from teaql.runtime.context import UserContext


def test_prepare_query_clones_and_governs_every_query_node_once():
shared = SelectQuery("SchoolType").project("id")
root = SelectQuery("School").project("id")
root.relations.append(RelationLoad("schoolType", shared))
root.facets.append(FacetRequest("types", "schoolType", shared))
calls = []

def policy(query):
calls.append(query.entity)
query.project("policy_marker")

context = UserContext().with_request_policy(policy)
prepared = context.prepare_query(root)

assert prepared is not root
assert prepared.relations[0].query is prepared.facets[0].query
assert calls == ["School", "SchoolType"]
assert "policy_marker" in prepared.projection
assert "policy_marker" in prepared.relations[0].query.projection
assert "policy_marker" not in root.projection
assert "policy_marker" not in shared.projection


def test_prepare_query_without_policy_still_returns_independent_graph():
child = SelectQuery("SchoolType").project("id")
root = SelectQuery("School")
root.relations.append(RelationLoad("schoolType", child))

prepared = UserContext().prepare_query(root)
prepared.relations[0].query.project("code")

assert prepared is not root
assert prepared.relations[0].query is not child
assert child.projection == ["id"]


def test_prepare_query_propagates_policy_denial():
def deny(query):
if query.entity == "SchoolType":
raise PermissionError("query policy denied SchoolType")

root = SelectQuery("School")
root.relations.append(RelationLoad("schoolType", SelectQuery("SchoolType")))

context = UserContext().with_request_policy(deny)
try:
context.prepare_query(root)
except PermissionError as error:
assert str(error) == "query policy denied SchoolType"
else:
raise AssertionError("nested policy denial must fail closed")


def test_prepare_query_preserves_shared_nodes_when_policy_returns_replacement():
shared = SelectQuery("SchoolType")
root = SelectQuery("School")
root.relations.append(RelationLoad("schoolType", shared))
root.facets.append(FacetRequest("types", "schoolType", shared))

def replace(query):
replacement = SelectQuery(query.entity)
replacement.relations = query.relations
replacement.facets = query.facets
replacement.project("authorized")
return replacement

prepared = UserContext().with_request_policy(replace).prepare_query(root)

assert prepared.relations[0].query is prepared.facets[0].query
assert prepared.relations[0].query.projection == ["authorized"]
Loading