Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@

from ..base_validator import HasResult, SourceValidator
from ..validation_context import ValidationContext
from .validate_call_kwargs import UDFNodeMapping, ValidateCallKwargs
from .validate_call_kwargs import UDFNodeMapping, ValidateCallKwargs, udf_mapping_key


class FeatureNameToEntityTypeMapping(SourceValidator, HasResult[Dict[str, str]]):
Expand All @@ -21,7 +21,7 @@ def validate_source(self, source: 'grammar.Source') -> None:
and not statement.target.is_local
and isinstance(statement.value, grammar.Call)
):
_, args = self._udf_node_mapping[id(statement.value)]
_, args = self._udf_node_mapping[udf_mapping_key(statement.value)]
if isinstance(args, EntityArgumentsBase):
self._feature_name_to_entity_type[statement.target.identifier] = args.type.value

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from osprey.engine.utils.graph import CyclicDependencyError, Graph

from ..base_validator import BaseValidator, HasResult
from .validate_call_kwargs import UDFNodeMapping, ValidateCallKwargs
from .validate_call_kwargs import UDFNodeMapping, ValidateCallKwargs, udf_mapping_key

if TYPE_CHECKING:
from ..validation_context import ValidationContext
Expand Down Expand Up @@ -96,7 +96,7 @@ def validate_source(self, source: Source) -> None:
from osprey.engine.stdlib.udfs.import_ import Import

for call_node in filter_nodes(source.ast_root, Call):
udf, _ = self._udf_node_mapping[id(call_node)]
udf, _ = self._udf_node_mapping[udf_mapping_key(call_node)]
if not isinstance(udf, Import):
continue

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,13 @@
from ..validation_context import ValidationContext


UDFNodeMapping = Dict[int, Tuple[UDFBase[Any, Any], ArgumentsBase]]
# Keyed by (source_path, start_line, start_pos) — stable across Python object reuse
UDFNodeMapping = Dict[Tuple[str, int, int], Tuple[UDFBase[Any, Any], ArgumentsBase]]


def udf_mapping_key(node: Call) -> Tuple[str, int, int]:
"""Get stable (source_path, line, pos) key for a Call node."""
return (node.span.source.path, node.span.start_line, node.span.start_pos)


class ValidateCallKwargs(SourceValidator, HasResult[UDFNodeMapping]):
Expand Down Expand Up @@ -180,7 +186,7 @@ def validate_call_node(self, call_node: Call) -> None:
# If it's a dynamic UDF, we'll add the RValueTypeChecker in ValidateDynamicCallsHaveAnnotatedRValue
# Store the udf in the node mapping - this will be read in `CallExecutor` within the executor,
# to then pluck the udf + arguments that we parsed here to be executed upon.
self._udf_node_mapping[id(call_node)] = (udf, arguments)
self._udf_node_mapping[udf_mapping_key(call_node)] = (udf, arguments)

# Catch the ConstExprArgumentException - if any other exception is thrown, it's considered a
# compilation error!
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from osprey.engine.ast.grammar import Assign, Call, Name, Source
from osprey.engine.ast_validator.base_validator import SourceValidator
from osprey.engine.ast_validator.validation_utils import add_must_assign_to_variable_error
from osprey.engine.ast_validator.validators.validate_call_kwargs import ValidateCallKwargs
from osprey.engine.ast_validator.validators.validate_call_kwargs import ValidateCallKwargs, udf_mapping_key
from osprey.engine.udf.rvalue_type_checker import (
AnnotationConversionError,
convert_ast_annotation_to_type_checker,
Expand Down Expand Up @@ -106,5 +106,5 @@ def validate_call_node(self, call_node: Call) -> None:
),
)

udf, _ = self._udf_node_mapping[id(call_node)]
udf, _ = self._udf_node_mapping[udf_mapping_key(call_node)]
udf.set_rvalue_type_checker(rvalue_type_checker)
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from ..validation_context import ValidationContext
from .feature_name_to_entity_type_mapping import FeatureNameToEntityTypeMapping
from .validate_call_kwargs import UDFNodeMapping, ValidateCallKwargs
from .validate_call_kwargs import UDFNodeMapping, ValidateCallKwargs, udf_mapping_key

### meow

Expand All @@ -28,7 +28,7 @@ def __init__(self, context: 'ValidationContext'):

def validate_source(self, source: 'grammar.Source') -> None:
for call_node in filter_nodes(source.ast_root, grammar.Call):
_, arguments = self._udf_node_mapping[id(call_node)]
_, arguments = self._udf_node_mapping[udf_mapping_key(call_node)]
if isinstance(arguments, LabelArguments):
self._validate_label(call_node, arguments)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
from ..base_validator import HasInput, HasResult, SourceValidator
from .imports_must_not_have_cycles import ImportsMustNotHaveCycles
from .unique_stored_names import UniqueStoredNames
from .validate_call_kwargs import UDFNodeMapping, ValidateCallKwargs
from .validate_call_kwargs import UDFNodeMapping, ValidateCallKwargs, udf_mapping_key
from .validate_dynamic_calls_have_annotated_rvalue import ValidateDynamicCallsHaveAnnotatedRValue
from .variables_must_be_defined import VariablesMustBeDefined

Expand Down Expand Up @@ -301,7 +301,7 @@ def _validate_list(self, list_: grammar.List) -> type:
return List[child_type] # type: ignore # Doesn't like runtime types like this

def _validate_call(self, call: grammar.Call) -> type:
udf, arguments = self._udf_node_mapping[id(call)]
udf, arguments = self._udf_node_mapping[udf_mapping_key(call)]

# Special case to handle import. Need to do this so we populate name types.
if isinstance(udf, Import):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
from ..base_validator import BaseValidator, HasInput, HasResult
from .imports_must_not_have_cycles import ImportsMustNotHaveCycles
from .unique_stored_names import UniqueStoredNames
from .validate_call_kwargs import ValidateCallKwargs
from .validate_call_kwargs import ValidateCallKwargs, udf_mapping_key


class VariablesMustBeDefined(BaseValidator, HasInput[Set[str]], HasResult[Mapping[Source, Set[str]]]):
Expand Down Expand Up @@ -131,7 +131,7 @@ def get_result(self) -> Mapping[Source, Set[str]]:

def get_imported_sources(self, node: Call) -> Sequence[Source]:
udf_mapping = self.context.get_validator_result(ValidateCallKwargs)
call_udf, _ = udf_mapping[id(node)]
call_udf, _ = udf_mapping[udf_mapping_key(node)]
assert isinstance(call_udf, Import)
return [s for s, _ in call_udf.sources_and_spans]

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,11 +22,14 @@ class CallExecutor(BaseNodeExecutor[Call, Any]):
_udf: 'UDFBase[Any, Any]'

def __init__(self, node: Call, sources: 'ValidatedSources'):
from osprey.engine.ast_validator.validators.validate_call_kwargs import ValidateCallKwargs
from osprey.engine.ast_validator.validators.validate_call_kwargs import (
ValidateCallKwargs,
udf_mapping_key,
)

super().__init__(node=node, sources=sources)
udf_map = sources.get_validator_result(ValidateCallKwargs)
self._udf, self.unresolved_arguments = udf_map[id(node)]
self._udf, self.unresolved_arguments = udf_map[udf_mapping_key(node)]
self.dependent_node_dict = self.unresolved_arguments.get_dependent_node_dict()

def set_tracing_tags(self, span: 'Span') -> None:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

from osprey.engine.ast import grammar
from osprey.engine.ast_validator.validation_context import ValidatedSources
from osprey.engine.ast_validator.validators.validate_call_kwargs import ValidateCallKwargs
from osprey.engine.ast_validator.validators.validate_call_kwargs import ValidateCallKwargs, udf_mapping_key
from osprey.engine.udf.base import QueryUdfBase
from osprey.engine.utils.osprey_unary_executor import OspreyUnaryExecutor

Expand Down Expand Up @@ -99,7 +99,7 @@ def transform_UnaryOperation(self, node: grammar.UnaryOperation) -> Dict[str, An
raise DruidQueryTransformException(node, 'Unknown Unary Operator')

def transform_Call(self, node: grammar.Call) -> Dict[str, Any]:
udf, _ = self._udf_node_mapping[id(node)]
udf, _ = self._udf_node_mapping[udf_mapping_key(node)]

if not isinstance(udf, QueryUdfBase):
raise DruidQueryTransformException(node, 'Unknown function call type')
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,66 @@ def test_parses_did_mutate_label(
assert check_json_output(transformed_query)


def test_parses_query_with_regex_match_on_username(
make_rules_sources: MakeRulesSourcesFunction,
) -> None:
validated_sources = parse_query_to_validated_ast(
"RegexMatch(item=UserName, regex='^jake')",
make_rules_sources([('UserName', '"some_user"')]),
)
transformed_query = DruidQueryTransformer(validated_sources=validated_sources).transform()

# Assert it returns a valid Druid query with the regex filter
assert isinstance(transformed_query, dict)
assert 'filter' in transformed_query
assert isinstance(transformed_query['filter'], dict)
assert transformed_query['filter'].get('type') == 'regex'
assert transformed_query['filter'].get('dimension') == 'UserName'
assert transformed_query['filter'].get('pattern') == '^jake'


def test_udf_node_mapping_uses_stable_keys(
make_rules_sources: MakeRulesSourcesFunction,
) -> None:
"""
Regression test for issue #158: UDF node mapping must use stable keys, not object identity.

The bug was that ValidateCallKwargs keyed the mapping by id(call_node), which is unstable
if the AST is reparsed or nodes are reconstructed. This caused KeyError in consumers like
DruidQueryTransformer, ValidateStaticTypes, etc. when they tried to look up the UDF.

The fix changes the key to be (source.path, start_line, start_pos), which is stable across
reparsing and object reconstruction.
"""
rules_sources = make_rules_sources([('UserName', '"some_user"')])

# Parse and validate the query
query = "RegexMatch(item=UserName, regex='^jake')"
validated_sources = parse_query_to_validated_ast(query, rules_sources)

# Get the UDF mapping
udf_mapping = validated_sources.get_validator_result(ValidateCallKwargs)
assert len(udf_mapping) > 0, 'Should have at least one UDF in mapping'

# Verify that the mapping uses stable (source_path, line, pos) keys (tuples)
# not id()-based keys (integers)
for key in udf_mapping.keys():
assert isinstance(key, tuple), (
f'Mapping should use stable (source_path, line, pos) tuple keys, not id() keys. Got {type(key)}: {key}'
)
assert len(key) == 3, f'Key should be (path, line, pos) tuple with 3 elements, got {len(key)}: {key}'
assert isinstance(key[0], str), f'First key element should be path (str), got {type(key[0])}'
assert isinstance(key[1], int), f'Second key element should be line (int), got {type(key[1])}'
assert isinstance(key[2], int), f'Third key element should be pos (int), got {type(key[2])}'

# Verify transformation works (consumes the stable keys)
transformer = DruidQueryTransformer(validated_sources=validated_sources)
transformed_query = transformer.transform()
assert transformed_query is not None
assert isinstance(transformed_query, dict)
assert 'filter' in transformed_query


@pytest.mark.parametrize(
'query',
[
Expand Down
Loading