Skip to content
Closed
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
194 changes: 194 additions & 0 deletions authentik/blueprints/tests/test_v1_tasks.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Test blueprints v1 tasks"""

from hashlib import sha512
from pathlib import Path
from tempfile import NamedTemporaryFile, mkdtemp

from django.test import TransactionTestCase
Expand Down Expand Up @@ -156,3 +157,196 @@ def test_valid_disabled(self):
instance.status,
BlueprintInstanceStatus.UNKNOWN,
)

@CONFIG.patch("blueprints_dir", TMP)
def test_valid_crlf(self):
"""Test discovered hash matches the applied hash for a file with CRLF line endings"""
blueprint_id = generate_id()
with NamedTemporaryFile(suffix=".yaml", dir=TMP) as file:
file.write(
f"version: 1\r\nentries: []\r\nmetadata:\r\n name: {blueprint_id}\r\n".encode()
)
file.flush()
file_hash = sha512(Path(file.name).read_text(encoding="utf-8").encode()).hexdigest()
for _ in range(2):
blueprints_discovery.send()
instance = BlueprintInstance.objects.filter(name=blueprint_id).first()
self.assertEqual(instance.status, BlueprintInstanceStatus.SUCCESSFUL)
found = next(
found for found in blueprints_find() if found.path == Path(file.name).name
)
self.assertEqual(found.hash, file_hash)
self.assertEqual(instance.last_applied_hash, found.hash)
file.seek(0)
self.assertIn(b"\r\n", file.read())

def write_blueprint(self, file, value: str):
file.seek(0)
file.truncate()
file.write(f"version: 1\nentries: []\ncontext:\n secret: {value}\n")
file.flush()
return next(found for found in blueprints_find() if found.path == Path(file.name).name).hash

def write_secret(self, file, value: str):
file.seek(0)
file.truncate()
file.write(value)
file.flush()

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_content_changed(self):
"""Test hash changes when the contents of a referenced `!File` change"""
with NamedTemporaryFile(mode="w+", dir=TMP) as secret:
for label, reference in (
("direct", f"!File {secret.name}"),
("argument of another tag", f'!Format ["client-%s", !File {secret.name}]'),
("reached through a cycle", f"&anchor [*anchor, !File {secret.name}]"),
("reached through an alias", f"&anchor [!File {secret.name}]\n other: *anchor"),
):
with (
self.subTest(label),
NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as file,
):
self.write_secret(secret, "initial")
before = self.write_blueprint(file, reference)
self.write_secret(secret, "rotated")
self.assertNotEqual(before, self.write_blueprint(file, reference))

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_created(self):
"""Test hash changes when a referenced `!File` that was missing appears"""
with NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as file:
secret_path = Path(TMP) / generate_id()
reference = f"!File {secret_path}"
before = self.write_blueprint(file, reference)
secret_path.write_text("created")
try:
after = self.write_blueprint(file, reference)
finally:
secret_path.unlink()
self.assertNotEqual(before, after)

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_removed(self):
"""Test hash changes when a referenced `!File` that existed disappears"""
with NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as file:
secret_path = Path(TMP) / generate_id()
secret_path.write_text("present")
reference = f"!File {secret_path}"
try:
before = self.write_blueprint(file, reference)
finally:
secret_path.unlink()
self.assertNotEqual(before, self.write_blueprint(file, reference))

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_contents_swapped(self):
"""Test hash changes when two referenced `!File`s exchange their contents"""
with (
NamedTemporaryFile(mode="w+", dir=TMP) as first,
NamedTemporaryFile(mode="w+", dir=TMP) as second,
):
self.write_secret(first, "alpha")
self.write_secret(second, "beta")
with NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as file:
reference = f"[!File {first.name}, !File {second.name}]"
before = self.write_blueprint(file, reference)
self.write_secret(first, "beta")
self.write_secret(second, "alpha")
self.assertNotEqual(before, self.write_blueprint(file, reference))

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_hashed_once_per_route(self):
"""Test a referenced `!File` is folded into the hash once for each route to it"""
with NamedTemporaryFile(mode="w+", dir=TMP) as secret:
self.write_secret(secret, "initial")
for label, reference, routes in (
("cycle", f"&anchor [*anchor, !File {secret.name}]", 1),
("alias", f"&anchor [!File {secret.name}]\n other: *anchor", 2),
):
with (
self.subTest(label),
NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as file,
):
content = f"version: 1\nentries: []\ncontext:\n secret: {reference}\n"
expected = sha512(content.encode())
for _ in range(routes):
expected.update(sha512(b"initial").digest())
self.assertEqual(self.write_blueprint(file, reference), expected.hexdigest())

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_content_unchanged(self):
"""Test hash is stable when a referenced `!File` does not change"""
with NamedTemporaryFile(mode="w+", dir=TMP) as secret:
self.write_secret(secret, "initial")
with NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as file:
reference = f"!File {secret.name}"
self.assertEqual(
self.write_blueprint(file, reference),
self.write_blueprint(file, reference),
)

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_unreadable_hash_stable(self):
"""Test hash is stable when a referenced `!File` cannot be read"""
for label, reference in (
("missing file", f"!File {Path(TMP) / generate_id()}"),
("path from a tag", f'!File [!Env [{generate_id()}, "{TMP}/fallback"], "default"]'),
("path from a mapping", f'!File {{path: "{TMP}/fallback"}}'),
("path no syscall can take", '!File "\\0"'),
("deeply nested", "[" * 50 + f'!File "{TMP}/fallback"' + "]" * 50),
):
with (
self.subTest(label),
NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as file,
):
self.assertEqual(
self.write_blueprint(file, reference),
self.write_blueprint(file, reference),
)

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_unreadable_discovery_continues(self):
"""Test a blueprint that cannot be hashed does not stop others being discovered"""
for label, reference in (
("path from a mapping", f'!File {{path: "{TMP}/fallback"}}'),
("path no syscall can take", '!File "\\0"'),
("path outside the filesystem encoding", '!File "\\ud800"'),
("sequence containing itself", "&anchor [*anchor]"),
("mapping containing itself", "&anchor {key: *anchor}"),
("two anchors containing each other", "&outer [{inner: &inner [*outer]}, *inner]"),
):
with (
self.subTest(label),
NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as broken,
NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as healthy,
):
broken.write(f"version: 1\nentries: []\ncontext:\n secret: {reference}\n")
broken.flush()
healthy.write(f"version: 1\nentries: []\nmetadata:\n name: {generate_id()}\n")
healthy.flush()
found = [blueprint.path for blueprint in blueprints_find()]
self.assertIn(Path(healthy.name).name, found)
self.assertIn(Path(broken.name).name, found)

@CONFIG.patch("blueprints_dir", TMP)
def test_file_tag_applied_on_change(self):
"""Test blueprint is re-applied when the contents of a referenced `!File` change"""
blueprint_id = generate_id()
with NamedTemporaryFile(mode="w+", dir=TMP) as secret:
self.write_secret(secret, "initial")
with NamedTemporaryFile(mode="w+", suffix=".yaml", dir=TMP) as file:
file.write(
f"version: 1\nentries: []\n"
f"metadata:\n name: {blueprint_id}\n"
f"context:\n secret: !File {secret.name}\n"
)
file.flush()
blueprints_discovery.send()
instance = BlueprintInstance.objects.filter(name=blueprint_id).first()
before = instance.last_applied_hash
self.assertEqual(instance.status, BlueprintInstanceStatus.SUCCESSFUL)
self.write_secret(secret, "rotated")
blueprints_discovery.send()
instance.refresh_from_db()
self.assertNotEqual(instance.last_applied_hash, before)
60 changes: 56 additions & 4 deletions authentik/blueprints/v1/tasks.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
"""v1 blueprints tasks"""

from collections.abc import Generator
from dataclasses import asdict, dataclass, field
from hashlib import sha512
from pathlib import Path
from sys import platform
from typing import Any
from uuid import UUID

from dacite.core import from_dict
Expand All @@ -30,7 +32,13 @@
BlueprintInstanceStatus,
BlueprintRetrievalFailed,
)
from authentik.blueprints.v1.common import BlueprintLoader, BlueprintMetadata, EntryInvalidError
from authentik.blueprints.v1.common import (
BlueprintLoader,
BlueprintMetadata,
EntryInvalidError,
File,
YAMLTag,
)
from authentik.blueprints.v1.importer import Importer
from authentik.blueprints.v1.labels import LABEL_AUTHENTIK_INSTANTIATE
from authentik.blueprints.v1.oci import OCI_PREFIX
Expand All @@ -56,6 +64,49 @@ class BlueprintFile:
meta: BlueprintMetadata | None = field(default=None)


def iter_file_tags(value: Any, ancestors: frozenset[int] = frozenset()) -> Generator[File]:
"""Find all `!File` tags in a loaded blueprint, including tags used as arguments
of other tags. A node is not descended into again below itself; a node reached by
several routes is visited once per route."""
if id(value) in ancestors:
return
ancestors = ancestors | {id(value)}
if isinstance(value, File):
yield value
if isinstance(value, dict):
children = value.values()
elif isinstance(value, list | tuple):
children = value
elif isinstance(value, YAMLTag):
children = vars(value).values()
else:
return
for child in children:
yield from iter_file_tags(child, ancestors)


def blueprint_hash(content: str) -> str:
"""Hash a blueprint's content and the contents of the files it references with
`!File` tags"""
hasher = sha512(content.encode())
try:
raw_blueprint = load(content, BlueprintLoader)
except YAMLError:
return hasher.hexdigest()
for tag in iter_file_tags(raw_blueprint):
# Mapping-node tags have no path; nested tags cannot be resolved here
path = getattr(tag, "path", None)
if not isinstance(path, str):
continue
try:
referenced = Path(path).read_bytes()
except OSError, ValueError:
# Unreadable references contribute only their blueprint source text
continue
hasher.update(sha512(referenced).digest())
return hasher.hexdigest()


class BlueprintWatcherMiddleware(Middleware):
def start_blueprint_watcher(self):
"""Start blueprint watcher"""
Expand Down Expand Up @@ -129,8 +180,9 @@ def blueprints_find() -> list[BlueprintFile]:
if any(part for part in rel_path.parts if part.startswith(".")):
continue
with open(path, encoding="utf-8") as blueprint_file:
content = blueprint_file.read()
try:
raw_blueprint = load(blueprint_file.read(), BlueprintLoader)
raw_blueprint = load(content, BlueprintLoader)
except YAMLError as exc:
raw_blueprint = None
LOGGER.warning("failed to parse blueprint", exc=exc, path=str(rel_path))
Expand All @@ -141,7 +193,7 @@ def blueprints_find() -> list[BlueprintFile]:
if version != 1:
LOGGER.warning("invalid blueprint version", version=version, path=str(rel_path))
continue
file_hash = sha512(path.read_bytes()).hexdigest()
file_hash = blueprint_hash(content)
blueprint = BlueprintFile(str(rel_path), version, file_hash, int(path.stat().st_mtime))
blueprint.meta = from_dict(BlueprintMetadata, metadata) if metadata else None
blueprints.append(blueprint)
Expand Down Expand Up @@ -202,7 +254,7 @@ def apply_blueprint(instance_pk: UUID):
self.info(f"Blueprint {instance.name} is disabled, skipping")
return
blueprint_content = instance.retrieve()
file_hash = sha512(blueprint_content.encode()).hexdigest()
file_hash = blueprint_hash(blueprint_content)
importer = Importer.from_string(blueprint_content, instance.context)
if importer.blueprint.metadata:
instance.metadata = asdict(importer.blueprint.metadata)
Expand Down
Loading