Skip to content
Draft
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
10 changes: 9 additions & 1 deletion pkg-py/src/commons/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
]
49 changes: 49 additions & 0 deletions pkg-py/src/commons/_measures.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
2 changes: 1 addition & 1 deletion pkg-py/tests/measure_sources/collision_a/uses_shared.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.")
Expand Down
2 changes: 1 addition & 1 deletion pkg-py/tests/measure_sources/duplicate_helpers/a_file.py
Original file line number Diff line number Diff line change
@@ -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:
Expand Down
2 changes: 1 addition & 1 deletion pkg-py/tests/measure_sources/duplicate_helpers/b_file.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
helper name.
"""

from commons._measures import measure
from commons import measure


def helper() -> int:
Expand Down
2 changes: 1 addition & 1 deletion pkg-py/tests/measure_sources/nested/orders.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.")
Expand Down
2 changes: 1 addition & 1 deletion pkg-py/tests/measure_sources/orders.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@

from pydantic import Field

from commons._measures import Injected, measure
from commons import Injected, measure


def double(x: int) -> int:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down
2 changes: 1 addition & 1 deletion pkg-py/tests/measure_sources/revenue.py
Original file line number Diff line number Diff line change
@@ -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.")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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.")
Expand Down
2 changes: 1 addition & 1 deletion pkg-py/tests/measure_sources/stdlib_collision/json.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.")
Expand Down
97 changes: 97 additions & 0 deletions pkg-py/tests/test_measures.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import enum
import importlib
import inspect
import sys
import threading
from dataclasses import FrozenInstanceError
Expand All @@ -19,6 +20,7 @@
as_measure,
measure,
measure_schema_text,
resolve_injections,
semantic_layer,
)

Expand Down Expand Up @@ -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
Expand Down