diff --git a/py/src/braintrust/__init__.py b/py/src/braintrust/__init__.py index 1e66fced..f568fdb0 100644 --- a/py/src/braintrust/__init__.py +++ b/py/src/braintrust/__init__.py @@ -50,6 +50,17 @@ def is_equal(expected, output): # Check env var at import time for auto-instrumentation import os +from typing import TYPE_CHECKING + + +if TYPE_CHECKING: + # These names must precede the generated wildcard so type checkers keep the + # runtime resource classes for the four names shared by both modules. + from .logger import Dataset as Dataset # noqa: I001 + from .logger import Experiment as Experiment + from .logger import Project as Project + from .logger import Prompt as Prompt + from .generated_types import * if os.getenv("BRAINTRUST_INSTRUMENT_THREADS", "").lower() in ("true", "1", "yes"): @@ -60,6 +71,7 @@ def is_equal(expected, output): except Exception: pass # Never break on import +from . import generated_types as _generated_types # noqa: I001 from .audit import * from .auto import auto_instrument as auto_instrument from .dataset_pipeline import * @@ -67,8 +79,14 @@ def is_equal(expected, output): from .framework2 import * from .functions.invoke import * from .functions.stream import * -from .generated_types import * -from .integrations.ai_sdk import setup_ai_sdk as setup_ai_sdk + +# Keep this before the logger wildcard so its existing runtime collision +# precedence remains unchanged while new generated names are picked up. +for _name in _generated_types.__all__: + if _name not in {"Dataset", "Experiment", "Project", "Prompt"}: + globals()[_name] = getattr(_generated_types, _name) + +from .integrations.ai_sdk import setup_ai_sdk as setup_ai_sdk # noqa: I001 from .integrations.anthropic import wrap_anthropic as wrap_anthropic from .integrations.instructor import wrap_instructor as wrap_instructor from .integrations.litellm import wrap_litellm as wrap_litellm @@ -77,6 +95,10 @@ def is_equal(expected, output): from .integrations.pydantic_ai import setup_pydantic_ai as setup_pydantic_ai from .logger import * from .logger import ( + Dataset as Dataset, + Experiment as Experiment, + Project as Project, + Prompt as Prompt, _internal_get_global_state, # noqa: F401 # type: ignore[reportUnusedImport] _internal_reset_global_state, # noqa: F401 # type: ignore[reportUnusedImport] _internal_with_custom_background_logger, # noqa: F401 # type: ignore[reportUnusedImport] diff --git a/py/src/braintrust/type_tests/test_public_exports.py b/py/src/braintrust/type_tests/test_public_exports.py index 928f6716..0defbbf1 100644 --- a/py/src/braintrust/type_tests/test_public_exports.py +++ b/py/src/braintrust/type_tests/test_public_exports.py @@ -1,15 +1,18 @@ -"""Regression test for pyright's ``reportPrivateImportUsage`` on top-level ``braintrust`` symbols. +"""Regression tests for top-level ``braintrust`` symbols. -Without PEP 484 ``as``-aliasing in ``braintrust/__init__.py``, pyright flags -``from braintrust import auto_instrument`` (and peers) as private in a -``py.typed`` consumer. The local ``pyrightconfig.json`` turns the rule into -an error so this file breaks ``nox -s test_types`` if someone regresses the -aliasing pattern. +The static resource check keeps mypy from resolving generated ``TypedDict`` +names instead of the public runtime classes. The runtime checks cover PEP 484 +aliasing for pyright's ``reportPrivateImportUsage`` rule and keep generated +exports synchronized with the package root. """ +import subprocess +import sys + import braintrust import pytest from braintrust import ( + Acl, auto_instrument, setup_ai_sdk, setup_pydantic_ai, @@ -31,7 +34,42 @@ ] +def accepts_public_resource_types( + experiment: braintrust.Experiment, + dataset: braintrust.Dataset, + project: braintrust.Project, + prompt: braintrust.Prompt, + acl: Acl, +) -> None: + experiment.fetch() + dataset.fetch() + _ = project.name + prompt.build() + _ = acl["id"] + + @pytest.mark.parametrize("name,imported", _PUBLIC_SYMBOLS) def test_top_level_public_symbol(name: str, imported: object) -> None: assert callable(imported) assert callable(getattr(braintrust, name)) + + +def test_generated_exports_follow_generated_all() -> None: + script = """ +import importlib + +import braintrust +from braintrust import generated_types + +future_type = type("FutureGeneratedType", (), {}) +generated_types.FutureGeneratedType = future_type +generated_types.__all__.append("FutureGeneratedType") +try: + importlib.reload(braintrust) + assert braintrust.FutureGeneratedType is future_type +finally: + generated_types.__all__.remove("FutureGeneratedType") + del generated_types.FutureGeneratedType +""" + + subprocess.run([sys.executable, "-c", script], check=True)