diff --git a/burr/cli/__main__.py b/burr/cli/__main__.py index 8967602d9..9205a95d2 100644 --- a/burr/cli/__main__.py +++ b/burr/cli/__main__.py @@ -212,8 +212,9 @@ def _run_server( @click.option( "--host", default="127.0.0.1", - help="Host to run the server on -- use 0.0.0.0 if you want " - "to expose it to the network (E.G. in a docker image)", + help="Host to run the server on -- defaults to 127.0.0.1 (local only). Use 0.0.0.0 if you " + "want to expose it to the network (E.G. in a docker image); the server has no built-in " + "authentication, so put it behind an authenticating proxy when doing so.", ) @click.option( "--backend", 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/burr/tracking/server/backend.py b/burr/tracking/server/backend.py index 5e9e1c771..5dc250cc9 100644 --- a/burr/tracking/server/backend.py +++ b/burr/tracking/server/backend.py @@ -31,6 +31,7 @@ from pydantic_settings import BaseSettings, SettingsConfigDict from burr.tracking.common import models +from burr.tracking.common.identifiers import join_within, validate_identifier from burr.tracking.common.models import ChildApplicationModel from burr.tracking.server import schema from burr.tracking.server.schema import ( @@ -293,6 +294,42 @@ def get_uri(project_id: str) -> str: DEFAULT_PATH = os.path.expanduser("~/.burr") +def _validate_identifier(value: str, name: str = "identifier") -> str: + """Validate a project/app identifier that is used as a path component under the storage + directory. The rule itself lives in :mod:`burr.tracking.common.identifiers` and is shared + with the tracking client; this wrapper reports failures as HTTP 400. + + :param value: the identifier to validate + :param name: label for the error message, e.g. ``"project_id"`` + :return: the identifier if valid + :raises fastapi.HTTPException: 400 if the identifier does not meet the rule + """ + try: + return validate_identifier(value, name) + except ValueError as e: + raise fastapi.HTTPException(status_code=400, detail=str(e)) from e + + +def _safe_join(base: str, *parts: str) -> str: + """Join path components and ensure the result stays inside ``base``. + + Delegates to :func:`burr.tracking.common.identifiers.join_within`, which resolves symlinks + on both sides before comparing, so a link under the storage directory that points elsewhere + is rejected too. Failures are reported as HTTP 400 without echoing the storage location. + + :param base: the allowed base directory + :param parts: path components to join + :return: the joined path + :raises fastapi.HTTPException: 400 if the resolved path is not inside the base directory + """ + try: + return join_within(base, *parts) + except ValueError as e: + raise fastapi.HTTPException( + status_code=400, detail="Invalid path: must be inside the storage directory." + ) from e + + class LocalBackend(BackendBase, AnnotationsBackendMixin): """Quick implementation of a local backend for testing purposes. This is not a production backend. @@ -303,7 +340,8 @@ def __init__(self, path: str = DEFAULT_PATH): self.path = path def _get_annotation_path(self, project_id: str) -> str: - return os.path.join(self.path, project_id, "annotations.jsonl") + _validate_identifier(project_id, "project_id") + return _safe_join(self.path, project_id, "annotations.jsonl") async def _load_project_annotations(self, project_id: str): annotations_path = self._get_annotation_path(project_id) @@ -464,7 +502,8 @@ async def list_apps( limit: int = 100, offset: int = 0, ) -> Tuple[Sequence[ApplicationSummary], int]: - project_filepath = os.path.join(self.path, project_id) + _validate_identifier(project_id, "project_id") + project_filepath = _safe_join(self.path, project_id) if not os.path.exists(project_filepath): return [], 0 # raise fastapi.HTTPException(status_code=404, detail=f"Project: {project_id} not found") @@ -506,7 +545,9 @@ async def get_application_logs( ) -> ApplicationLogs: # TODO -- handle partition key here # This currently assumes uniqueness - app_filepath = os.path.join(self.path, project_id, app_id) + _validate_identifier(project_id, "project_id") + _validate_identifier(app_id, "app_id") + app_filepath = _safe_join(self.path, project_id, app_id) if not os.path.exists(app_filepath): raise fastapi.HTTPException( status_code=404, detail=f"App: {app_id} from project: {project_id} not found" diff --git a/burr/tracking/server/run.py b/burr/tracking/server/run.py index da5b8b965..7860f1859 100644 --- a/burr/tracking/server/run.py +++ b/burr/tracking/server/run.py @@ -195,6 +195,10 @@ def create_burr_ui_app(serve_static: bool = SERVE_STATIC) -> FastAPI: This factory creates a new FastAPI instance with all Burr UI routes, demo routers, and (optionally) static file serving configured. + The app does not authenticate requests. Serve it on localhost (the default for the + ``burr`` CLI and for ``python -m burr.tracking.server.run``) or put it behind an + authenticating reverse proxy when exposing it beyond the local machine. + :param serve_static: Whether to serve the React UI static files. Defaults to the BURR_SERVE_STATIC environment variable (true by default). :return: A fully-configured FastAPI application. @@ -450,4 +454,8 @@ def mount_burr_ui( if __name__ == "__main__": port = int(os.getenv("PORT", 8000)) # Default to 8000 if no PORT environment variable is set - uvicorn.run(app, host="0.0.0.0", port=port) + # Bind to localhost unless BURR_SERVER_HOST says otherwise. The server has no built-in + # authentication, so only expose it beyond the local machine (e.g. BURR_SERVER_HOST=0.0.0.0) + # behind an authenticating reverse proxy. + host = os.getenv("BURR_SERVER_HOST", "127.0.0.1") + uvicorn.run(app, host=host, port=port) diff --git a/tests/tracking/test_local_backend_paths.py b/tests/tracking/test_local_backend_paths.py new file mode 100644 index 000000000..58efdcbb6 --- /dev/null +++ b/tests/tracking/test_local_backend_paths.py @@ -0,0 +1,324 @@ +# 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. + +"""Identifier validation and path containment for the local tracking backend. + +Project and app identifiers arrive as URL path parameters and are used as path +components under the storage directory. These tests check that only plain +identifiers are accepted, that every resolved path stays inside the storage +directory, and that the server reports rejected identifiers as HTTP 400. +""" + +import importlib +import os +import uuid +from pathlib import Path + +import pytest +from fastapi import HTTPException + +from burr import system +from burr.tracking.server.backend import LocalBackend, _safe_join, _validate_identifier + +# Minimal valid graph.json matching the ApplicationModel schema +GRAPH_JSON = ( + '{"entrypoint": "counter", "actions": [{"name": "counter", "reads": [], ' + '"writes": ["counter"], "code": "pass"}], "transitions": []}' +) + +ANNOTATION = { + "span_id": None, + "step_name": "counter", + "tags": ["review"], + "observations": [ + {"data_fields": {"note": "looks fine"}, "thumbs_up_thumbs_down": True, "data_pointers": []} + ], +} + +# ".." sent percent-encoded: the HTTP client collapses a literal ".." segment before the +# request leaves, while the server decodes %2E%2E back to ".." and hands it to the handler. +PARENT = "%2E%2E" + + +def _make_app(storage: Path, project_id: str, app_id: str) -> Path: + """Creates the on-disk layout the local tracking client writes for one app.""" + app_dir = storage / project_id / app_id + app_dir.mkdir(parents=True) + (app_dir / "graph.json").write_text(GRAPH_JSON) + (app_dir / "log.jsonl").write_text("") + (app_dir / "metadata.json").write_text("{}") + return app_dir + + +class TestValidateIdentifier: + def test_valid_identifiers(self): + assert _validate_identifier("hello_world") == "hello_world" + assert _validate_identifier("hello-world") == "hello-world" + assert _validate_identifier("Hello:World_123") == "Hello:World_123" + + def test_accepts_everything_the_client_produces(self): + # default app ids are uuid4; project names are [a-zA-Z0-9_-] plus ":" off Windows + app_id = str(uuid.uuid4()) + assert _validate_identifier(app_id, "app_id") == app_id + assert _validate_identifier("my-project_1", "project_id") == "my-project_1" + assert _validate_identifier("a" * 255) == "a" * 255 + # dots are ordinary filename characters as long as the value is not "." or ".." + assert _validate_identifier("v1.2.3", "app_id") == "v1.2.3" + assert _validate_identifier("hello..world", "app_id") == "hello..world" + if not system.IS_WINDOWS: + assert _validate_identifier("demo:chatbot", "project_id") == "demo:chatbot" + + def test_invalid_identifiers(self): + with pytest.raises(HTTPException) as exc: + _validate_identifier("../etc/passwd") + assert exc.value.status_code == 400 + + with pytest.raises(HTTPException) as exc: + _validate_identifier("hello/world") + assert exc.value.status_code == 400 + + with pytest.raises(HTTPException) as exc: + _validate_identifier("hello\\world") + assert exc.value.status_code == 400 + + @pytest.mark.parametrize("value", ["", ".", "..", " ", "a b", "/etc/passwd", "C:\\Windows"]) + def test_rejects_empty_dot_and_absolute_identifiers(self, value): + with pytest.raises(HTTPException) as exc: + _validate_identifier(value, "project_id") + assert exc.value.status_code == 400 + assert "project_id" in exc.value.detail + + def test_rejects_over_long_and_non_string_identifiers(self): + with pytest.raises(HTTPException) as exc: + _validate_identifier("a" * 256) + assert exc.value.status_code == 400 + + with pytest.raises(HTTPException) as exc: + _validate_identifier(None) # type: ignore[arg-type] + assert exc.value.status_code == 400 + + +class TestSafeJoin: + def test_safe_join_within_base(self, tmp_path): + base = str(tmp_path) + assert _safe_join(base, "project1") == str(tmp_path / "project1") + assert _safe_join(base, "project1", "app1") == str(tmp_path / "project1" / "app1") + + def test_safe_join_rejects_parent_directory(self, tmp_path): + base = str(tmp_path) + with pytest.raises(HTTPException) as exc: + _safe_join(base, "..", "etc") + assert exc.value.status_code == 400 + + with pytest.raises(HTTPException) as exc: + _safe_join(base, "project", "..", "..", "etc") + assert exc.value.status_code == 400 + + def test_safe_join_rejects_absolute_parts(self, tmp_path): + # os.path.join discards everything before an absolute component + with pytest.raises(HTTPException) as exc: + _safe_join(str(tmp_path), "/etc/passwd") + assert exc.value.status_code == 400 + + with pytest.raises(HTTPException) as exc: + _safe_join(str(tmp_path), "project", os.sep + "etc") + assert exc.value.status_code == 400 + + def test_safe_join_refuses_exact_base(self, tmp_path): + # identifier paths must be strictly inside the storage directory, never the root itself + with pytest.raises(HTTPException) as exc: + _safe_join(str(tmp_path)) + assert exc.value.status_code == 400 + + def test_safe_join_follows_symlinks(self, tmp_path): + storage = tmp_path / "storage" + storage.mkdir() + outside = tmp_path / "outside" + outside.mkdir() + (storage / "inside").mkdir() + (storage / "to-outside").symlink_to(outside, target_is_directory=True) + (storage / "to-inside").symlink_to(storage / "inside", target_is_directory=True) + + # a link under the storage directory that resolves elsewhere is rejected ... + with pytest.raises(HTTPException) as exc: + _safe_join(str(storage), "to-outside") + assert exc.value.status_code == 400 + with pytest.raises(HTTPException) as exc: + _safe_join(str(storage), "to-outside", "annotations.jsonl") + assert exc.value.status_code == 400 + + # ... while one that stays inside is accepted (the plain join is returned) + assert _safe_join(str(storage), "to-inside", "app") == str(storage / "to-inside" / "app") + + def test_safe_join_accepts_symlinked_base(self, tmp_path): + real_base = tmp_path / "real" + real_base.mkdir() + linked_base = tmp_path / "linked" + linked_base.symlink_to(real_base, target_is_directory=True) + assert _safe_join(str(linked_base), "project") == str(linked_base / "project") + + +class TestLocalBackendPathContainment: + def test_get_annotation_path_rejects_parent_directory(self, tmp_path): + backend = LocalBackend(path=str(tmp_path)) + with pytest.raises(HTTPException) as exc: + backend._get_annotation_path("../etc") + assert exc.value.status_code == 400 + + @pytest.mark.asyncio + async def test_list_apps_rejects_parent_directory(self, tmp_path): + backend = LocalBackend(path=str(tmp_path)) + with pytest.raises(HTTPException) as exc: + await backend.list_apps(None, "../../../etc", None) + assert exc.value.status_code == 400 + + @pytest.mark.asyncio + async def test_get_application_logs_rejects_parent_directory_in_project(self, tmp_path): + backend = LocalBackend(path=str(tmp_path)) + with pytest.raises(HTTPException) as exc: + await backend.get_application_logs(None, "../etc", "app1", None) + assert exc.value.status_code == 400 + + @pytest.mark.asyncio + async def test_get_application_logs_rejects_parent_directory_in_app(self, tmp_path): + backend = LocalBackend(path=str(tmp_path)) + with pytest.raises(HTTPException) as exc: + await backend.get_application_logs(None, "project1", "../etc", None) + assert exc.value.status_code == 400 + + @pytest.mark.asyncio + async def test_get_application_logs_allows_valid(self, tmp_path): + backend = LocalBackend(path=str(tmp_path)) + _make_app(tmp_path, "project1", "app1") + result = await backend.get_application_logs(None, "project1", "app1", None) + assert result is not None + + +@pytest.fixture +def server(tmp_path, monkeypatch): + """A TestClient for the tracking server with a LocalBackend on a fresh directory. + + ``burr.tracking.server.run`` builds its backend and FastAPI app at import time from the + environment, so the environment is prepared before the import and the module-level + backend is swapped afterwards (the endpoints look it up on every call). + """ + storage = tmp_path / "storage" + storage.mkdir() + monkeypatch.setenv("BURR_SERVE_STATIC", "false") + monkeypatch.setenv("BURR_BACKEND_IMPL", "burr.tracking.server.backend.LocalBackend") + monkeypatch.setenv("burr_path", str(storage)) + run = importlib.import_module("burr.tracking.server.run") + monkeypatch.setattr(run, "backend", LocalBackend(path=str(storage))) + + from fastapi.testclient import TestClient + + return TestClient(run.app), storage + + +class TestServerEndpoints: + def test_list_apps_rejects_parent_directory(self, server): + client, _ = server + response = client.get(f"/api/v0/{PARENT}/__none__/apps") + assert response.status_code == 400 + assert "project_id" in response.json()["detail"] + + def test_application_logs_reject_parent_directory(self, server): + client, storage = server + _make_app(storage, "proj", "app-1") + assert client.get(f"/api/v0/{PARENT}/app-1/__none__/apps").status_code == 400 + assert client.get(f"/api/v0/proj/{PARENT}/__none__/apps").status_code == 400 + + @pytest.mark.parametrize("project_id", ["..%5C..%5Cetc", "my%20project", "a" * 256]) + def test_rejects_other_malformed_identifiers(self, server, project_id): + client, _ = server + assert client.get(f"/api/v0/{project_id}/__none__/apps").status_code == 400 + + def test_separators_never_reach_a_handler(self, server): + # The router splits on "/" whether or not it is percent-encoded, so an identifier + # containing one cannot match any identifier route; the function-level tests above + # cover these values directly. + client, _ = server + assert client.get("/api/v0/..%2F..%2Fetc/__none__/apps").status_code == 404 + assert client.get("/api/v0/%2Fetc%2Fpasswd/__none__/apps").status_code == 404 + + def test_annotation_endpoints_reject_parent_directory(self, server): + client, storage = server + assert client.get(f"/api/v0/{PARENT}/annotations").status_code == 400 + response = client.post(f"/api/v0/{PARENT}/app-1/__none__/0/annotations", json=ANNOTATION) + assert response.status_code == 400 + response = client.put(f"/api/v0/{PARENT}/0/update_annotations", json=ANNOTATION) + assert response.status_code == 400 + # nothing was written next to the storage directory + assert not (storage.parent / "annotations.jsonl").exists() + + def test_symlinked_project_outside_storage_is_rejected(self, server, tmp_path): + client, storage = server + outside = tmp_path / "outside" + outside.mkdir() + (storage / "linked").symlink_to(outside, target_is_directory=True) + + response = client.get("/api/v0/linked/__none__/apps") + assert response.status_code == 400 + # the error does not reveal where the storage directory lives + assert str(storage) not in response.json()["detail"] + assert client.get("/api/v0/linked/app-1/__none__/apps").status_code == 400 + assert client.get("/api/v0/linked/annotations").status_code == 400 + response = client.post("/api/v0/linked/app-1/__none__/0/annotations", json=ANNOTATION) + assert response.status_code == 400 + assert not (outside / "annotations.jsonl").exists() + + def test_accepts_identifiers_the_client_produces(self, server): + client, storage = server + project_id = "my-project_1" + app_id = str(uuid.uuid4()) + _make_app(storage, project_id, app_id) + + response = client.get(f"/api/v0/{project_id}/__none__/apps") + assert response.status_code == 200 + assert [a["app_id"] for a in response.json()["applications"]] == [app_id] + + response = client.get(f"/api/v0/{project_id}/{app_id}/__none__/apps") + assert response.status_code == 200 + assert response.json()["application"]["entrypoint"] == "counter" + + # ":" is used by the demo projects on non-Windows hosts; an unknown project is an + # empty listing rather than a rejection + if not system.IS_WINDOWS: + response = client.get("/api/v0/demo:chatbot/__none__/apps") + assert response.status_code == 200 + assert response.json()["applications"] == [] + + def test_valid_but_missing_app_is_not_found(self, server): + client, storage = server + _make_app(storage, "proj", "app-1") + assert client.get(f"/api/v0/proj/{uuid.uuid4()}/__none__/apps").status_code == 404 + + def test_annotation_round_trip_with_valid_identifiers(self, server): + client, storage = server + app_id = str(uuid.uuid4()) + _make_app(storage, "proj", app_id) + + response = client.post(f"/api/v0/proj/{app_id}/__none__/0/annotations", json=ANNOTATION) + assert response.status_code == 200 + assert (storage / "proj" / "annotations.jsonl").exists() + + response = client.put("/api/v0/proj/0/update_annotations", json=ANNOTATION) + assert response.status_code == 200 + + response = client.get(f"/api/v0/proj/annotations?app_id={app_id}") + assert response.status_code == 200 + assert [a["app_id"] for a in response.json()] == [app_id]