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
66 changes: 43 additions & 23 deletions burr/tracking/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand All @@ -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()
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down
92 changes: 92 additions & 0 deletions burr/tracking/common/identifiers.py
Original file line number Diff line number Diff line change
@@ -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):

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

IDENTIFIER_PATTERN.match(value) with a pattern ending in $ accepts "x\n" and "..\n"; $ matches before a final newline. IDENTIFIER_PATTERN.fullmatch(value) (or \Z in the pattern) rejects them.

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
131 changes: 131 additions & 0 deletions tests/tracking/test_local_tracking_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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))
Loading