diff --git a/pkg-py/src/commons/__init__.py b/pkg-py/src/commons/__init__.py index 5ff40c2..d2dd98c 100644 --- a/pkg-py/src/commons/__init__.py +++ b/pkg-py/src/commons/__init__.py @@ -5,4 +5,12 @@ classification as to how much it can be trusted. """ -__all__: list[str] = [] +from ._measures import Injected, Measure, SemanticLayer, measure, semantic_layer + +__all__: list[str] = [ + "Injected", + "Measure", + "SemanticLayer", + "measure", + "semantic_layer", +] diff --git a/pkg-py/src/commons/_measures.py b/pkg-py/src/commons/_measures.py index d761137..98b210a 100644 --- a/pkg-py/src/commons/_measures.py +++ b/pkg-py/src/commons/_measures.py @@ -301,6 +301,55 @@ def _resolve_ref(node: dict[str, Any], defs: dict[str, Any]) -> dict[str, Any]: return defs[ref.removeprefix("#/$defs/")] +def resolve_injections( + measures: Mapping[str, Measure], + injectables: Mapping[str, Any], +) -> dict[str, dict[str, Any]]: + """Bind each measure's injected arguments to the agent's named sources. + + An argument named after a data source receives that source's connection, + even when the argument has a default. An argument matching no source keeps + its default; one with no default is an error, raised here so a measure + that can never run is caught at construction rather than mid-conversation. + """ + resolved: dict[str, dict[str, Any]] = {} + + for name, record in measures.items(): + defaults = inspect.signature(record.func).parameters + bound: dict[str, Any] = {} + unresolvable: list[str] = [] + + for argument in record.injected: + if argument in injectables: + bound[argument] = injectables[argument] + elif defaults[argument].default is inspect.Parameter.empty: + unresolvable.append(argument) + + if unresolvable: + raise ValueError(_unresolvable_message(name, unresolvable, injectables)) + + resolved[name] = bound + + return resolved + + +def _unresolvable_message( + name: str, + unresolvable: Sequence[str], + injectables: Mapping[str, Any], +) -> str: + available = ( + f"Available sources: {', '.join(injectables)}." + if injectables + else "The agent has no named data sources." + ) + return ( + f"Measure {name!r} has injected arguments matching no data source: " + f"{', '.join(unresolvable)}.\n" + f"{available}" + ) + + def as_measure(obj: Any) -> Measure | None: """Recognize a measure, whether decorated function or bare record.""" if isinstance(obj, Measure): diff --git a/pkg-py/tests/measure_sources/collision_a/uses_shared.py b/pkg-py/tests/measure_sources/collision_a/uses_shared.py index 6f48f21..8b09d8c 100644 --- a/pkg-py/tests/measure_sources/collision_a/uses_shared.py +++ b/pkg-py/tests/measure_sources/collision_a/uses_shared.py @@ -4,7 +4,7 @@ from shared_lib import value # type: ignore[missing-import] -from commons._measures import measure +from commons import measure @measure(description="From directory a.") diff --git a/pkg-py/tests/measure_sources/duplicate_helpers/a_file.py b/pkg-py/tests/measure_sources/duplicate_helpers/a_file.py index a6358dc..1fa7ff7 100644 --- a/pkg-py/tests/measure_sources/duplicate_helpers/a_file.py +++ b/pkg-py/tests/measure_sources/duplicate_helpers/a_file.py @@ -1,6 +1,6 @@ """First file, sorted before b_file.py in this directory.""" -from commons._measures import measure +from commons import measure def helper() -> int: diff --git a/pkg-py/tests/measure_sources/duplicate_helpers/b_file.py b/pkg-py/tests/measure_sources/duplicate_helpers/b_file.py index 6734b4c..9ccf566 100644 --- a/pkg-py/tests/measure_sources/duplicate_helpers/b_file.py +++ b/pkg-py/tests/measure_sources/duplicate_helpers/b_file.py @@ -4,7 +4,7 @@ helper name. """ -from commons._measures import measure +from commons import measure def helper() -> int: diff --git a/pkg-py/tests/measure_sources/nested/orders.py b/pkg-py/tests/measure_sources/nested/orders.py index a17d26b..31e1860 100644 --- a/pkg-py/tests/measure_sources/nested/orders.py +++ b/pkg-py/tests/measure_sources/nested/orders.py @@ -4,7 +4,7 @@ collide with the other orders.py in sys.modules. """ -from commons._measures import measure +from commons import measure @measure(description="Count of nested orders.") diff --git a/pkg-py/tests/measure_sources/orders.py b/pkg-py/tests/measure_sources/orders.py index 614fa97..06c6e64 100644 --- a/pkg-py/tests/measure_sources/orders.py +++ b/pkg-py/tests/measure_sources/orders.py @@ -4,7 +4,7 @@ from pydantic import Field -from commons._measures import Injected, measure +from commons import Injected, measure def double(x: int) -> int: diff --git a/pkg-py/tests/measure_sources/reentrant/composes_a_sibling.py b/pkg-py/tests/measure_sources/reentrant/composes_a_sibling.py index 0fd7e0a..320720a 100644 --- a/pkg-py/tests/measure_sources/reentrant/composes_a_sibling.py +++ b/pkg-py/tests/measure_sources/reentrant/composes_a_sibling.py @@ -7,7 +7,7 @@ from pathlib import Path -from commons._measures import measure, semantic_layer +from commons import measure, semantic_layer NESTED_LAYER = semantic_layer(Path(__file__).parent.parent / "nested" / "orders.py") diff --git a/pkg-py/tests/measure_sources/revenue.py b/pkg-py/tests/measure_sources/revenue.py index 78375b0..8d0598b 100644 --- a/pkg-py/tests/measure_sources/revenue.py +++ b/pkg-py/tests/measure_sources/revenue.py @@ -1,6 +1,6 @@ """A second file in the same directory, to prove directory loading.""" -from commons._measures import measure +from commons import measure @measure(description="Total revenue.") diff --git a/pkg-py/tests/measure_sources/sibling_imports/uses_helper.py b/pkg-py/tests/measure_sources/sibling_imports/uses_helper.py index 5ca6fe5..12ede19 100644 --- a/pkg-py/tests/measure_sources/sibling_imports/uses_helper.py +++ b/pkg-py/tests/measure_sources/sibling_imports/uses_helper.py @@ -4,7 +4,7 @@ from helper_lib import double # type: ignore[missing-import] -from commons._measures import measure +from commons import measure @measure(description="Doubled count.") diff --git a/pkg-py/tests/measure_sources/stdlib_collision/json.py b/pkg-py/tests/measure_sources/stdlib_collision/json.py index 6fa34b5..c0d47fb 100644 --- a/pkg-py/tests/measure_sources/stdlib_collision/json.py +++ b/pkg-py/tests/measure_sources/stdlib_collision/json.py @@ -2,7 +2,7 @@ before the directory ever goes on sys.path. """ -from commons._measures import measure +from commons import measure @measure(description="Should never load.") diff --git a/pkg-py/tests/test_measures.py b/pkg-py/tests/test_measures.py index 4f69865..1022abf 100644 --- a/pkg-py/tests/test_measures.py +++ b/pkg-py/tests/test_measures.py @@ -2,6 +2,7 @@ import enum import importlib +import inspect import sys import threading from dataclasses import FrozenInstanceError @@ -19,6 +20,7 @@ as_measure, measure, measure_schema_text, + resolve_injections, semantic_layer, ) @@ -754,6 +756,101 @@ def test_installed_but_unimported_module_is_a_construction_error( assert "already-importable" in message +def _region_revenue(default: Any = inspect.Parameter.empty) -> Measure: + if default is inspect.Parameter.empty: + + @measure(description="Revenue for a region.") + def region_revenue(warehouse: Injected[Any]) -> int: + return 0 + else: + + @measure(description="Revenue for a region.") + def region_revenue(warehouse: Injected[Any] = default) -> int: + return 0 + + return _as_measure(region_revenue) + + +def test_resolve_injections_binds_a_matching_source() -> None: + connection = object() + layer = semantic_layer(_region_revenue()) + + resolved = resolve_injections(layer.measures, {"warehouse": connection}) + + assert resolved == {"region_revenue": {"warehouse": connection}} + + +def test_resolve_injections_prefers_a_source_over_a_default() -> None: + connection = object() + layer = semantic_layer(_region_revenue(default="fallback")) + + resolved = resolve_injections(layer.measures, {"warehouse": connection}) + + assert resolved["region_revenue"]["warehouse"] is connection + + +def test_resolve_injections_leaves_an_unmatched_default_alone() -> None: + layer = semantic_layer(_region_revenue(default="fallback")) + + resolved = resolve_injections(layer.measures, {"finance": object()}) + + assert resolved == {"region_revenue": {}} + + +def test_resolve_injections_errors_on_an_unmatched_argument() -> None: + layer = semantic_layer(_region_revenue()) + + with pytest.raises(ValueError) as excinfo: + resolve_injections(layer.measures, {"finance": object()}) + + message = str(excinfo.value) + assert "region_revenue" in message + assert "warehouse" in message + assert "finance" in message + + +def test_resolve_injections_says_when_there_are_no_named_sources() -> None: + layer = semantic_layer(_region_revenue()) + + with pytest.raises(ValueError, match="no named data sources"): + resolve_injections(layer.measures, {}) + + +def test_resolve_injections_lists_every_unmatched_argument_at_once() -> None: + @measure(description="Joins two warehouses.") + def joined(left: Injected[Any], right: Injected[Any]) -> int: + return 0 + + layer = semantic_layer(joined) + + with pytest.raises(ValueError) as excinfo: + resolve_injections(layer.measures, {}) + + message = str(excinfo.value) + assert "left" in message + assert "right" in message + + +def test_resolve_injections_returns_an_entry_for_every_measure() -> None: + layer = semantic_layer(_count_measure()) + + assert resolve_injections(layer.measures, {}) == {"order_count": {}} + + +def test_public_api_exposes_the_semantic_layer() -> None: + import commons + + assert set(commons.__all__) >= { + "Injected", + "Measure", + "SemanticLayer", + "measure", + "semantic_layer", + } + assert commons.measure is measure + assert commons.semantic_layer is semantic_layer + + def test_semantic_layer_reenters_during_a_measure_files_import() -> None: # A non-reentrant lock deadlocks here rather than raising, so this runs # on a daemon thread with a timeout: a regression fails the test instead