From a41163557316cbca7e435daa0dac292a94bdc166 Mon Sep 17 00:00:00 2001 From: Elijah Ben Izzy Date: Sun, 4 Oct 2026 13:22:59 -0700 Subject: [PATCH] fix(tracking): validate app ids in the local tracking client The local tracking client turned app ids (and the parent app ids used for fork/spawn links) straight into directory names under the storage directory. Apply the rule the project name already used -- letters, digits, '_', '-', ':' and '.', non-empty, not '.' or '..', at most 255 characters -- to every identifier that becomes a path, and check that joined paths stay inside the storage directory. The rule lives in burr.tracking.common.identifiers so the server side can share it. Errors surface as ValueError from ApplicationBuilder.build(). --- burr/tracking/client.py | 66 ++++++---- burr/tracking/common/identifiers.py | 92 +++++++++++++ tests/tracking/test_local_tracking_client.py | 131 +++++++++++++++++++ 3 files changed, 266 insertions(+), 23 deletions(-) create mode 100644 burr/tracking/common/identifiers.py diff --git a/burr/tracking/client.py b/burr/tracking/client.py index 5375cdb4b..5eb11b536 100644 --- a/burr/tracking/client.py +++ b/burr/tracking/client.py @@ -46,7 +46,6 @@ def flock(*args, **kwargs): import json import logging import os -import re import traceback from abc import ABC from typing import Any, Dict, Optional, Tuple @@ -68,6 +67,7 @@ def flock(*args, **kwargs): PreRunStepHook, PreStartSpanHook, ) +from burr.tracking.common.identifiers import join_within, validate_identifier from burr.tracking.common.models import ( ApplicationMetadataModel, ApplicationModel, @@ -109,13 +109,15 @@ def _filter_inputs(d: dict) -> dict: def _allowed_project_name(project_name: str, on_windows: bool) -> bool: - allowed_chars = r"a-zA-Z0-9_\-" - if not on_windows: - allowed_chars += ":" - pattern = f"^[{allowed_chars}]+$" + """Whether ``project_name`` passes the identifier rule shared with app ids. - # Use regular expression to check if the string is valid - return bool(re.match(pattern, project_name)) + Kept as a boolean wrapper around :func:`validate_identifier` for existing callers. + """ + try: + validate_identifier(project_name, "project", on_windows=on_windows) + except ValueError: + return False + return True @dataclasses.dataclass @@ -190,11 +192,7 @@ def __init__( """ self.f = None - if not _allowed_project_name(project, on_windows=system.IS_WINDOWS): - raise ValueError( - f"Project: {project} is not valid. Project name cannot contain non-alphanumeric (except _ and -) characters." - "We will be relaxing this restriction later but for now please rename your project!" - ) + validate_identifier(project, "project") self.raw_storage_dir = storage_dir self.storage_dir = LocalTrackingClient.get_storage_path(project, storage_dir) self.project_id = project @@ -255,7 +253,7 @@ def _log_child_relationships( ) ) for parent_id, child_of in parent_relationships: - parent_path = os.path.join(self.storage_dir, parent_id) + parent_path = self._application_directory(parent_id) if not os.path.exists(parent_path): # This currently makes the parent directory so that it does not fail # If the parent directory exists we'll just use that @@ -285,7 +283,22 @@ def copy(self) -> "LocalTrackingClient": @classmethod def get_storage_path(cls, project, storage_dir) -> str: - return str(os.path.join(os.path.expanduser(storage_dir), project)) + return str(join_within(os.path.expanduser(storage_dir), project)) + + @classmethod + def _application_log_path(cls, project: str, app_id: str, storage_dir: str) -> str: + """Path to an application's log file, with both identifiers validated and the + result kept inside the storage directory.""" + validate_identifier(project, "project") + validate_identifier(app_id, "app_id") + application_path = join_within(cls.get_storage_path(project, storage_dir), app_id) + return os.path.join(application_path, cls.LOG_FILENAME) + + def _application_directory(self, app_id: str) -> str: + """Directory for one application run, validated to sit inside this project's + storage directory. Nothing is created on disk.""" + validate_identifier(app_id, "app_id") + return join_within(self.storage_dir, app_id) @classmethod def app_log_exists( @@ -301,7 +314,7 @@ def app_log_exists( :param storage_dir: the storage directory. :return: True if state exists, False otherwise. """ - path = os.path.join(cls.get_storage_path(project, storage_dir), app_id, cls.LOG_FILENAME) + path = cls._application_log_path(project, app_id, storage_dir) if not os.path.exists(path): return False lines = open(path, "r", errors="replace", encoding="utf-8").readlines() @@ -336,7 +349,7 @@ def load_state( """ if sequence_id is None: sequence_id = -1 # get the last one - path = os.path.join(cls.get_storage_path(project, storage_dir), app_id, cls.LOG_FILENAME) + path = cls._application_log_path(project, app_id, storage_dir) if not os.path.exists(path): raise ValueError(f"No logs found for {project}/{app_id} under {storage_dir}") with open(path, "r", errors="replace", encoding="utf-8") as f: @@ -370,14 +383,16 @@ def load_state( prior_state["__SEQUENCE_ID"] = line_seq # add the sequence id back return prior_state, entry_point - def _ensure_dir_structure(self, app_id: str): + def _ensure_dir_structure(self, app_id: str) -> str: + # Validate the id before anything is created on disk + application_path = self._application_directory(app_id) if not os.path.exists(self.storage_dir): logger.info(f"Creating storage directory: {self.storage_dir}") os.makedirs(self.storage_dir) - application_path = os.path.join(self.storage_dir, app_id) if not os.path.exists(application_path): logger.info(f"Creating application directory: {application_path}") os.makedirs(application_path) + return application_path def __setstate__(self, state): self.__dict__.update(state) @@ -402,15 +417,20 @@ def post_application_create( spawning_parent_pointer: Optional[burr_types.ParentPointer], **future_kwargs: Any, ): - self._ensure_dir_structure(app_id) + # Every identifier that becomes a path is checked before the first write. A pointer + # with no app_id (plain resume, not a fork) never becomes a path, so it is skipped. + for pointer in (parent_pointer, spawning_parent_pointer): + if pointer is not None and pointer.app_id is not None: + validate_identifier(pointer.app_id, "app_id") + application_path = self._ensure_dir_structure(app_id) self.f = open( - os.path.join(self.storage_dir, app_id, self.LOG_FILENAME), + os.path.join(application_path, self.LOG_FILENAME), "a", encoding="utf-8", errors="replace", ) - graph_path = os.path.join(self.storage_dir, app_id, self.GRAPH_FILENAME) + graph_path = os.path.join(application_path, self.GRAPH_FILENAME) if os.path.exists(graph_path): logger.info(f"Graph already exists at {graph_path}. Not overwriting.") return @@ -420,7 +440,7 @@ def post_application_create( with open(graph_path, "w", encoding="utf-8", errors="replace") as f: json.dump(graph, f) - metadata_path = os.path.join(self.storage_dir, app_id, self.METADATA_FILENAME) + metadata_path = os.path.join(application_path, self.METADATA_FILENAME) if os.path.exists(metadata_path): logger.info(f"Metadata already exists at {metadata_path}. Not overwriting.") return @@ -617,7 +637,7 @@ def load( # TODO: if app_id is None: return # no application ID - path = os.path.join(self.storage_dir, app_id, self.LOG_FILENAME) + path = os.path.join(self._application_directory(app_id), self.LOG_FILENAME) if not os.path.exists(path): return None with open(path, "r", errors="replace", encoding="utf-8") as f: diff --git a/burr/tracking/common/identifiers.py b/burr/tracking/common/identifiers.py new file mode 100644 index 000000000..090231e27 --- /dev/null +++ b/burr/tracking/common/identifiers.py @@ -0,0 +1,92 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""Validation for the identifiers (project names, application ids) that the tracking +layer turns into filesystem paths, plus a helper that keeps the resulting paths inside +the storage directory.""" + +import os +import re +from typing import Optional + +from burr import system + +# Letters, digits, underscore, hyphen, colon and dot. This covers uuid4 (hex and +# hyphens) as well as the session/conversation style ids applications tend to pass +# through as app ids. +_ALLOWED_CHARACTERS = r"A-Za-z0-9_\-:." +IDENTIFIER_PATTERN = re.compile(f"^[{_ALLOWED_CHARACTERS}]+$") +MAX_IDENTIFIER_LENGTH = 255 + + +def validate_identifier(value: str, what: str, *, on_windows: Optional[bool] = None) -> str: + """Checks that ``value`` can be used as a single path component under the storage directory. + + The rule: a non-empty string of at most 255 characters drawn from letters, digits, ``_``, + ``-``, ``:`` and ``.``, and not the special directory names ``.`` or ``..``. On Windows + ``:`` is also refused, as it is a drive/stream separator there. + + :param value: the identifier to check + :param what: short label for the error message, e.g. ``"app_id"`` or ``"project"`` + :param on_windows: platform override, defaults to the current platform; exposed for tests + :return: ``value`` unchanged, so the call can be used inline + :raises ValueError: if the identifier does not meet the rule + """ + if on_windows is None: + on_windows = system.IS_WINDOWS + if not isinstance(value, str): + raise ValueError(f"{what} must be a string, got {type(value).__name__}: {value!r}") + if not value: + raise ValueError(f"{what} must not be empty") + if len(value) > MAX_IDENTIFIER_LENGTH: + raise ValueError( + f"{what} must be at most {MAX_IDENTIFIER_LENGTH} characters, got {len(value)}" + ) + if value in (".", ".."): + raise ValueError(f"{what} must not be '.' or '..', got {value!r}") + if not IDENTIFIER_PATTERN.match(value) or (on_windows and ":" in value): + allowed = "letters, digits, '_', '-', '.'" + ("" if on_windows else " and ':'") + raise ValueError(f"{what} may only contain {allowed}, got {value!r}") + return value + + +def join_within(base: str, *parts: str) -> str: + """Joins ``parts`` onto ``base`` and checks that the result stays strictly inside ``base``. + + Both sides are resolved with :func:`os.path.realpath` before comparing, so ``..`` segments + and symlinks are accounted for. The path returned is the plain join (not the resolved + form) so callers see the same spelling they would get from :func:`os.path.join`. + + :param base: the directory the result must stay inside + :param parts: path components to join onto ``base`` + :return: ``os.path.join(base, *parts)`` + :raises ValueError: if the joined path would land outside ``base`` + """ + joined = os.path.join(base, *parts) + base_resolved = os.path.realpath(base) + target_resolved = os.path.realpath(joined) + try: + inside = ( + target_resolved != base_resolved + and os.path.commonpath([base_resolved, target_resolved]) == base_resolved + ) + except ValueError: + # commonpath refuses to compare paths on different drives (Windows); treat as outside + inside = False + if not inside: + raise ValueError(f"Path {joined!r} is not inside the storage directory {base!r}") + return joined diff --git a/tests/tracking/test_local_tracking_client.py b/tests/tracking/test_local_tracking_client.py index a175e8d8f..ea94c8971 100644 --- a/tests/tracking/test_local_tracking_client.py +++ b/tests/tracking/test_local_tracking_client.py @@ -30,6 +30,7 @@ from burr.core.persistence import BaseStatePersister, PersistedStateData from burr.tracking import LocalTrackingClient from burr.tracking.client import _allowed_project_name +from burr.tracking.common.identifiers import join_within, validate_identifier from burr.tracking.common.models import ( ApplicationMetadataModel, ApplicationModel, @@ -665,3 +666,133 @@ def test_local_tracking_client_copy(): assert copy.project_id == tracking_client.project_id assert copy.serde_kwargs == tracking_client.serde_kwargs assert copy.storage_dir == tracking_client.storage_dir + + +# --- identifier validation: ids become directory names under the storage dir --------------- + + +@pytest.mark.parametrize( + "identifier", + [ + str(uuid.uuid4()), + "my-app_1", + "session:abc.123", + "a" * 255, + ], +) +def test_validate_identifier_accepts_reasonable_ids(identifier): + assert validate_identifier(identifier, "app_id", on_windows=False) == identifier + + +@pytest.mark.parametrize( + "identifier", + [ + "", + ".", + "..", + "../other", + "nested/app", + "/absolute", + "back\\slash", + "with space", + "a" * 256, + ], +) +def test_validate_identifier_rejects_ids_unfit_for_a_path_component(identifier): + with pytest.raises(ValueError): + validate_identifier(identifier, "app_id") + + +def test_validate_identifier_refuses_colon_on_windows(): + assert validate_identifier("a:b", "app_id", on_windows=False) == "a:b" + with pytest.raises(ValueError): + validate_identifier("a:b", "app_id", on_windows=True) + + +def test_join_within_keeps_paths_inside_base(tmp_path): + base = str(tmp_path) + assert join_within(base, "child") == os.path.join(base, "child") + assert join_within(base, "child", "log.jsonl") == os.path.join(base, "child", "log.jsonl") + for parts in [("..",), ("../sibling",), ("/absolute",), (".",), ("child", "..", "..")]: + with pytest.raises(ValueError): + join_within(base, *parts) + + +def _entries_under(root: str) -> set: + """Every file and directory below ``root``, as paths relative to ``root``.""" + found = set() + for dirpath, dirnames, filenames in os.walk(root): + for name in dirnames + filenames: + found.add(os.path.relpath(os.path.join(dirpath, name), root)) + return found + + +@pytest.mark.parametrize( + "app_id", + ["..", "../../escaped", "nested/app", "/absolute/path", "", "with space"], +) +def test_builder_rejects_app_id_that_leaves_storage_dir(tmpdir, app_id): + """The error surfaces from .build(), and nothing is written -- inside or outside the storage dir.""" + log_dir = os.path.join(str(tmpdir), "storage") + with pytest.raises(ValueError): + sample_application("test_builder_rejects_app_id", log_dir, app_id) + assert _entries_under(str(tmpdir)) == set() + + +def test_builder_accepts_uuid_and_simple_app_ids(tmpdir): + log_dir = str(tmpdir) + project_name = "test_builder_accepts_ids" + for app_id in [str(uuid.uuid4()), "my-app_1", "v1.2.3"]: + app = sample_application(project_name, log_dir, app_id) + app.run(halt_after=["result"]) + log_path = os.path.join(log_dir, project_name, app_id, LocalTrackingClient.LOG_FILENAME) + assert os.path.exists(log_path) + + +def test_builder_rejects_fork_parent_app_id_that_leaves_storage_dir(tmpdir): + """Forking goes through load(); the parent id is validated before the log is read.""" + log_dir = os.path.join(str(tmpdir), "storage") + tracker = LocalTrackingClient(project="test_fork_parent_id", storage_dir=log_dir) + with pytest.raises(ValueError): + ( + ApplicationBuilder() + .with_actions(counter=counter, result=Result("counter")) + .with_transitions(("counter", "result", default)) + .with_tracker(tracker) + .initialize_from( + tracker, + resume_at_next_action=True, + default_state={"counter": 0, "break_at": -1}, + default_entrypoint="counter", + fork_from_app_id="../escaped", + ) + .build() + ) + assert _entries_under(str(tmpdir)) == set() + + +def test_spawning_parent_app_id_is_validated_before_any_write(tmpdir): + """A bad parent id is caught before the child's own log/graph/metadata are written.""" + log_dir = os.path.join(str(tmpdir), "storage") + with pytest.raises(ValueError): + sample_application( + "test_spawn_parent_id", log_dir, str(uuid.uuid4()), spawn_from=("../escaped", 5) + ) + assert _entries_under(str(tmpdir)) == set() + + +def test_readers_reject_app_id_that_leaves_storage_dir(tmpdir): + project_name = "test_readers_reject" + tracker = LocalTrackingClient(project=project_name, storage_dir=str(tmpdir)) + with pytest.raises(ValueError): + tracker.load(partition_key=None, app_id="../escaped") + with pytest.raises(ValueError): + LocalTrackingClient.app_log_exists(project_name, "../escaped", storage_dir=str(tmpdir)) + with pytest.raises(ValueError): + LocalTrackingClient.load_state(project_name, "../escaped", storage_dir=str(tmpdir)) + + +@pytest.mark.parametrize("project", ["..", "../other", "a/b", "", "with space"]) +def test_project_name_follows_the_same_rule(tmpdir, project): + with pytest.raises(ValueError): + LocalTrackingClient(project=project, storage_dir=str(tmpdir))