diff --git a/tests/code_coverage_tests/test_e2e_metadata.py b/tests/code_coverage_tests/test_e2e_metadata.py index a3e1b85148a..042a89610ae 100644 --- a/tests/code_coverage_tests/test_e2e_metadata.py +++ b/tests/code_coverage_tests/test_e2e_metadata.py @@ -12,6 +12,7 @@ from __future__ import annotations import ast import inspect +import re import string import threading import warnings @@ -136,7 +137,7 @@ def _field_type(annotation: object) -> object: return present[0] if isinstance(annotation, UnionType) and len(present) == 1 else annotation -def _dotted_placeholders(owner: type) -> Iterator[tuple[str, str]]: +def _placeholders(owner: type) -> Iterator[tuple[str, str]]: tree: Final = ast.parse(inspect.getsource(owner)) for node in ast.walk(tree): if not isinstance(node, ast.FunctionDef): @@ -145,12 +146,16 @@ def _dotted_placeholders(owner: type) -> Iterator[tuple[str, str]]: match decorator: case ast.Call(func=ast.Name(id="step"), args=[ast.Constant(value=str(label))]): for _, field, _, _ in string.Formatter().parse(label): - if field is not None and "." in field: + if field is not None: yield node.name, field case _: pass +def _dotted_placeholders(owner: type) -> Iterator[tuple[str, str]]: + return ((method, field) for method, field in _placeholders(owner) if "." in field) + + def _resolves(owner: type, method: str, field: str) -> bool: root, *attributes = field.split(".") wrapped: Final = cast("Callable[..., object]", getattr(owner, method)) @@ -163,6 +168,35 @@ def _resolves(owner: type, method: str, field: str) -> bool: return True +SECRET_NAME: Final = re.compile( + r"secret|password|api_key|access_key|private_key|credential_values|^token$|(access|auth|bearer|refresh|session)_token$" +) + + +def _models_in(annotation: object, seen: frozenset[type] = frozenset()) -> frozenset[type[BaseModel]]: + """Every request model a value of this type can print, however deeply nested.""" + if isinstance(annotation, type) and issubclass(annotation, BaseModel): + if annotation in seen: + return frozenset() + nested: Final = ( + _models_in(field.annotation, seen | {annotation}) for field in annotation.model_fields.values() + ) + return frozenset({annotation}).union(*nested) + args: Final = cast("tuple[object, ...]", get_args(annotation)) + return frozenset[type[BaseModel]]().union(*(_models_in(arg, seen) for arg in args)) + + +def _printed_models(owner: type) -> frozenset[type[BaseModel]]: + def hint(method: str, field: str) -> object: + wrapped: Final = cast("Callable[..., object]", getattr(owner, method)) + hints: Final = cast("Mapping[str, object]", get_type_hints(inspect.unwrap(wrapped))) + return hints[field.split(".")[0]] + + return frozenset[type[BaseModel]]().union( + *(_models_in(hint(method, field)) for method, field in _placeholders(owner)) + ) + + class TestLabelTemplates: """A label's `{placeholders}` are filled from the call's own arguments, so the story says what the test asked for in words, and nothing the label doesn't name @@ -227,6 +261,20 @@ class TestLabelTemplates: assert placeholders assert [f"{method}: {field}" for method, field in placeholders if not _resolves(owner, method, field)] == [] + @pytest.mark.parametrize("owner", [ProxyClient], ids=["ProxyClient"]) + def test_every_secret_field_a_label_can_print_is_hidden(self, owner: type) -> None: + """A `{body}` label prints nested models too, so a callback's credentials + inside key metadata would land in the public report unless marked `repr=False`.""" + models: Final = _printed_models(owner) + assert models + exposed: Final = sorted( + f"{model.__name__}.{name}" + for model in models + for name, field in model.model_fields.items() + if field.repr and SECRET_NAME.search(name) + ) + assert exposed == [] + def test_escaped_braces_stay_literal(self) -> None: @step("GET /v1/batches/{{id}}") def retrieve_batch(batch_id: str) -> None: diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 2f5c4e89f0c..65b5ac8078b 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -42,10 +42,10 @@ class BudgetWindowState(BudgetWindow): class KeyLoggingCallbackVars(BaseModel): - langfuse_public_key: str | None = None - langfuse_secret_key: str | None = None + langfuse_public_key: str | None = Field(default=None, repr=False) + langfuse_secret_key: str | None = Field(default=None, repr=False) langfuse_host: str | None = None - wandb_api_key: str | None = None + wandb_api_key: str | None = Field(default=None, repr=False) weave_project_id: str | None = None