From dca6bada9d5beea2eeea6a7fa71857e16cef2c05 Mon Sep 17 00:00:00 2001 From: Caidan Date: Mon, 27 Oct 2025 14:08:27 -0700 Subject: [PATCH] fix: engine tests for labels (#27) --- .../executor/tests/test_render_graph.py | 12 +- .../osprey/engine/query_language/__init__.py | 7 +- .../udfs/did_declare_verdict.py | 5 +- .../engine/query_language/udfs/regex_match.py | 3 +- .../engine/stdlib/udfs/tests/test_labels.py | 139 ++++++++++-------- .../engine/stdlib/udfs/tests/test_rules.py | 5 +- .../src/osprey/engine/udf/registry.py | 8 +- .../worker/_stdlibplugin/udf_register.py | 4 + 8 files changed, 98 insertions(+), 85 deletions(-) diff --git a/osprey_worker/src/osprey/engine/executor/tests/test_render_graph.py b/osprey_worker/src/osprey/engine/executor/tests/test_render_graph.py index 7d7a2db..b92c182 100644 --- a/osprey_worker/src/osprey/engine/executor/tests/test_render_graph.py +++ b/osprey_worker/src/osprey/engine/executor/tests/test_render_graph.py @@ -4,19 +4,15 @@ from typing import Any, Dict, Set import pytest from osprey.engine.ast_validator.validation_context import ValidatedSources -from osprey.engine.ast_validator.validator_registry import ValidatorRegistry from osprey.engine.ast_validator.validators.unique_stored_names import UniqueStoredNames from osprey.engine.ast_validator.validators.validate_call_kwargs import ValidateCallKwargs -from osprey.engine.ast_validator.validators.validate_dynamic_calls_have_annotated_rvalue import ( - ValidateDynamicCallsHaveAnnotatedRValue, -) from osprey.engine.conftest import RunValidationFunction from osprey.engine.executor.execution_graph import ExecutionGraph, compile_execution_graph from osprey.engine.executor.execution_visualizer import _render_graph -from osprey.engine.stdlib import get_config_registry pytestmark = [ - pytest.mark.use_validators([ValidateCallKwargs, ValidateDynamicCallsHaveAnnotatedRValue, UniqueStoredNames]), + pytest.mark.use_validators([ValidateCallKwargs, UniqueStoredNames]), + pytest.mark.use_standard_rules_validators, pytest.mark.use_osprey_stdlib, ] @@ -147,9 +143,7 @@ def execution_graph(run_validation: RunValidationFunction) -> ExecutionGraph: """ Compiles an ExecutionGraph based on the above Osprey Rules configs """ - config_validator = get_config_registry().get_validator() - validator_registry = ValidatorRegistry.get_instance().instance_with_additional_validators(config_validator) - validated_sources: ValidatedSources = run_validation(config, validator_registry=validator_registry) + validated_sources: ValidatedSources = run_validation(config) execution_graph = compile_execution_graph(validated_sources) return execution_graph diff --git a/osprey_worker/src/osprey/engine/query_language/__init__.py b/osprey_worker/src/osprey/engine/query_language/__init__.py index 820b6fe..c21cbc4 100644 --- a/osprey_worker/src/osprey/engine/query_language/__init__.py +++ b/osprey_worker/src/osprey/engine/query_language/__init__.py @@ -3,12 +3,11 @@ from osprey.engine.ast_validator.validation_context import ValidatedSources, Val from osprey.engine.ast_validator.validators.unique_stored_names import UniqueStoredNames from osprey.engine.ast_validator.validators.validate_static_types import ValidateStaticTypes from osprey.engine.ast_validator.validators.variables_must_be_defined import VariablesMustBeDefined +from osprey.engine.query_language import udfs +from osprey.engine.query_language.ast_validator import REGISTRY +from osprey.engine.query_language.udfs.registry import UDF_REGISTRY from osprey.engine.utils.imports import import_all_direct_children -from . import udfs -from .ast_validator import REGISTRY -from .udfs.registry import UDF_REGISTRY - def parse_query_to_validated_ast(query: str, rules_sources: ValidatedSources) -> ValidatedSources: """ diff --git a/osprey_worker/src/osprey/engine/query_language/udfs/did_declare_verdict.py b/osprey_worker/src/osprey/engine/query_language/udfs/did_declare_verdict.py index 3934b82..d740a9f 100644 --- a/osprey_worker/src/osprey/engine/query_language/udfs/did_declare_verdict.py +++ b/osprey_worker/src/osprey/engine/query_language/udfs/did_declare_verdict.py @@ -1,12 +1,11 @@ from typing import Dict +from osprey.engine import shared_constants from osprey.engine.ast_validator.validation_context import ValidationContext +from osprey.engine.query_language.udfs.registry import register from osprey.engine.udf.arguments import ArgumentsBase, ConstExpr from osprey.engine.udf.base import QueryUdfBase -from ... import shared_constants -from .registry import register - class Arguments(ArgumentsBase): verdict: ConstExpr[str] diff --git a/osprey_worker/src/osprey/engine/query_language/udfs/regex_match.py b/osprey_worker/src/osprey/engine/query_language/udfs/regex_match.py index ae84bdd..c8cce62 100644 --- a/osprey_worker/src/osprey/engine/query_language/udfs/regex_match.py +++ b/osprey_worker/src/osprey/engine/query_language/udfs/regex_match.py @@ -3,11 +3,10 @@ from typing import Dict from osprey.engine.ast import grammar from osprey.engine.ast_validator.validation_context import ValidationContext +from osprey.engine.query_language.udfs.registry import register from osprey.engine.udf.arguments import ArgumentsBase, ConstExpr from osprey.engine.udf.base import QueryUdfBase -from .registry import register - class Arguments(ArgumentsBase): item: str diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_labels.py b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_labels.py index 80c5929..0416a8e 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_labels.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_labels.py @@ -1,6 +1,6 @@ import json from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set +from typing import Any, Callable, Dict, List, Optional, Sequence, Set import gevent import pytest @@ -20,18 +20,23 @@ from osprey.engine.conftest import ( ) from osprey.engine.executor.udf_execution_helpers import UDFHelpers from osprey.engine.language_types.entities import EntityT +from osprey.engine.language_types.labels import LabelStatus from osprey.engine.stdlib import get_config_registry from osprey.engine.stdlib.udfs.entity import Entity from osprey.engine.stdlib.udfs.labels import HasLabel, LabelAdd, LabelRemove from osprey.engine.stdlib.udfs.rules import Rule, WhenRules from osprey.engine.stdlib.udfs.time_delta import TimeDelta from osprey.engine.udf.registry import UDFRegistry -from osprey.engine.utils.proto_utils import datetime_to_timestamp -from osprey.rpc.labels.v1.service_pb2 import LabelReason, Labels, LabelState, LabelStatus +from osprey.worker.lib.osprey_shared.labels import ( + EntityLabelMutation, + EntityLabelMutationsResult, + EntityLabels, + LabelReason, + LabelReasons, + LabelState, +) from osprey.worker.lib.storage.labels import LabelsProvider - -if TYPE_CHECKING: - from osprey.rpc.labels.v1.service_pb2 import LabelStatusValue +from result import Result pytestmark: List[Callable[[Any], Any]] = [ pytest.mark.use_validators( @@ -50,20 +55,28 @@ pytestmark: List[Callable[[Any], Any]] = [ class StaticLabelProvider(LabelsProvider): - def __init__(self, entity_labels: Dict[EntityT[Any], Labels]) -> None: + def __init__(self, entity_labels: Dict[EntityT[Any], EntityLabels]) -> None: self._entity_labels = entity_labels - def get_from_service(self, key: EntityT[Any]) -> Labels: + def get_from_service(self, key: EntityT[Any]) -> EntityLabels: return self._entity_labels[key] + def batch_get_from_service(self, keys: Sequence[EntityT[Any]]) -> Sequence[Result[EntityLabels, Exception]]: + return [Result.Ok(self.get_from_service(key)) for key in keys] + + def apply_entity_mutation( + self, entity_key: EntityT[Any], mutations: List[EntityLabelMutation] + ) -> EntityLabelMutationsResult: + return self.apply_entity_label_mutations(entity_key, mutations) + class BlockingLabelProvider(StaticLabelProvider): - def __init__(self, entity_labels: Dict[EntityT[Any], Labels]) -> None: + def __init__(self, entity_labels: Dict[EntityT[Any], EntityLabels]) -> None: super().__init__(entity_labels) self.blocking_events: List[Event] = [] self.calls: List[EntityT[Any]] = [] - def get_from_service(self, key: EntityT[Any]) -> Labels: + def get_from_service(self, key: EntityT[Any]) -> EntityLabels: event = Event() self.blocking_events.append(event) event.wait() @@ -81,84 +94,88 @@ def source_with_labels_config(source: str, labels: Set[str]) -> Dict[str, str]: @pytest.mark.parametrize( 'checking_status, manual, actual_status, reasons, result', ( - ('added', None, LabelStatus.ADDED, {'TestReason': LabelReason()}, True), + ('added', None, LabelStatus.ADDED, LabelReasons({'TestReason': LabelReason()}), True), ( 'added', None, LabelStatus.ADDED, - {'ExpiredReason': LabelReason(expires_at=datetime_to_timestamp(datetime.now() - timedelta(hours=1)))}, + LabelReasons({'ExpiredReason': LabelReason(expires_at=(datetime.now() - timedelta(hours=1)))}), False, ), ( 'added', None, LabelStatus.ADDED, - { - 'ExpiredReason': LabelReason(expires_at=datetime_to_timestamp(datetime.now() - timedelta(hours=1))), - 'TestReason': LabelReason(), - }, + LabelReasons( + { + 'ExpiredReason': LabelReason(expires_at=(datetime.now() - timedelta(hours=1))), + 'TestReason': LabelReason(), + } + ), True, ), ( 'added', None, LabelStatus.ADDED, - { - 'ExpiredReason': LabelReason(expires_at=datetime_to_timestamp(datetime.now() - timedelta(hours=1))), - 'ExpiringReason': LabelReason(expires_at=datetime_to_timestamp(datetime.now() + timedelta(hours=1))), - }, + LabelReasons( + { + 'ExpiredReason': LabelReason(expires_at=(datetime.now() - timedelta(hours=1))), + 'ExpiringReason': LabelReason(expires_at=(datetime.now() + timedelta(hours=1))), + } + ), True, ), ( 'added', None, LabelStatus.ADDED, - {'ExpiringReason': LabelReason(expires_at=datetime_to_timestamp(datetime.now() + timedelta(hours=1)))}, + LabelReasons({'ExpiringReason': LabelReason(expires_at=(datetime.now() + timedelta(hours=1)))}), True, ), - ('added', None, LabelStatus.MANUALLY_ADDED, {'TestReason': LabelReason()}, True), - ('added', None, LabelStatus.REMOVED, {'TestReason': LabelReason()}, False), - ('added', None, LabelStatus.MANUALLY_REMOVED, {'TestReason': LabelReason()}, False), - ('added', None, None, {'TestReason': LabelReason()}, False), - ('added', True, LabelStatus.ADDED, {'TestReason': LabelReason()}, False), - ('added', True, LabelStatus.MANUALLY_ADDED, {'TestReason': LabelReason()}, True), - ('added', True, LabelStatus.REMOVED, {'TestReason': LabelReason()}, False), - ('added', True, LabelStatus.MANUALLY_REMOVED, {'TestReason': LabelReason()}, False), - ('added', True, None, {'TestReason': LabelReason()}, False), - ('added', False, LabelStatus.ADDED, {'TestReason': LabelReason()}, True), - ('added', False, LabelStatus.MANUALLY_ADDED, {'TestReason': LabelReason()}, False), - ('added', False, LabelStatus.REMOVED, {'TestReason': LabelReason()}, False), - ('added', False, LabelStatus.MANUALLY_REMOVED, {'TestReason': LabelReason()}, False), - ('added', False, None, {'TestReason': LabelReason()}, False), - ('removed', None, LabelStatus.ADDED, {'TestReason': LabelReason()}, False), - ('removed', None, LabelStatus.MANUALLY_ADDED, {'TestReason': LabelReason()}, False), - ('removed', None, LabelStatus.REMOVED, {'TestReason': LabelReason()}, True), - ('removed', None, LabelStatus.MANUALLY_REMOVED, {'TestReason': LabelReason()}, True), - ('removed', None, None, {'TestReason': LabelReason()}, True), - ('removed', True, LabelStatus.ADDED, {'TestReason': LabelReason()}, False), - ('removed', True, LabelStatus.MANUALLY_ADDED, {'TestReason': LabelReason()}, False), - ('removed', True, LabelStatus.REMOVED, {'TestReason': LabelReason()}, False), - ('removed', True, LabelStatus.MANUALLY_REMOVED, {'TestReason': LabelReason()}, True), - ('removed', True, None, {'TestReason': LabelReason()}, False), - ('removed', False, LabelStatus.ADDED, {'TestReason': LabelReason()}, False), - ('removed', False, LabelStatus.MANUALLY_ADDED, {'TestReason': LabelReason()}, False), - ('removed', False, LabelStatus.REMOVED, {'TestReason': LabelReason()}, True), - ('removed', False, LabelStatus.MANUALLY_REMOVED, {'TestReason': LabelReason()}, False), - ('removed', False, None, {'TestReason': LabelReason()}, True), + ('added', None, LabelStatus.MANUALLY_ADDED, LabelReasons({'TestReason': LabelReason()}), True), + ('added', None, LabelStatus.REMOVED, LabelReasons({'TestReason': LabelReason()}), False), + ('added', None, LabelStatus.MANUALLY_REMOVED, LabelReasons({'TestReason': LabelReason()}), False), + ('added', None, None, LabelReasons({'TestReason': LabelReason()}), False), + ('added', True, LabelStatus.ADDED, LabelReasons({'TestReason': LabelReason()}), False), + ('added', True, LabelStatus.MANUALLY_ADDED, LabelReasons({'TestReason': LabelReason()}), True), + ('added', True, LabelStatus.REMOVED, LabelReasons({'TestReason': LabelReason()}), False), + ('added', True, LabelStatus.MANUALLY_REMOVED, LabelReasons({'TestReason': LabelReason()}), False), + ('added', True, None, LabelReasons({'TestReason': LabelReason()}), False), + ('added', False, LabelStatus.ADDED, LabelReasons({'TestReason': LabelReason()}), True), + ('added', False, LabelStatus.MANUALLY_ADDED, LabelReasons({'TestReason': LabelReason()}), False), + ('added', False, LabelStatus.REMOVED, LabelReasons({'TestReason': LabelReason()}), False), + ('added', False, LabelStatus.MANUALLY_REMOVED, LabelReasons({'TestReason': LabelReason()}), False), + ('added', False, None, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', None, LabelStatus.ADDED, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', None, LabelStatus.MANUALLY_ADDED, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', None, LabelStatus.REMOVED, LabelReasons({'TestReason': LabelReason()}), True), + ('removed', None, LabelStatus.MANUALLY_REMOVED, LabelReasons({'TestReason': LabelReason()}), True), + ('removed', None, None, LabelReasons({'TestReason': LabelReason()}), True), + ('removed', True, LabelStatus.ADDED, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', True, LabelStatus.MANUALLY_ADDED, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', True, LabelStatus.REMOVED, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', True, LabelStatus.MANUALLY_REMOVED, LabelReasons({'TestReason': LabelReason()}), True), + ('removed', True, None, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', False, LabelStatus.ADDED, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', False, LabelStatus.MANUALLY_ADDED, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', False, LabelStatus.REMOVED, LabelReasons({'TestReason': LabelReason()}), True), + ('removed', False, LabelStatus.MANUALLY_REMOVED, LabelReasons({'TestReason': LabelReason()}), False), + ('removed', False, None, LabelReasons({'TestReason': LabelReason()}), True), ), ) def test_get_labels_retrieves_data( execute: ExecuteFunction, checking_status: str, manual: Optional[bool], - actual_status: Optional['LabelStatusValue'], - reasons: Dict[str, LabelReason], + actual_status: Optional[LabelStatus], + reasons: LabelReasons, result: bool, ) -> None: if actual_status is None: - labels = Labels(labels={}) + labels = EntityLabels(labels={}) else: - labels = Labels(labels={'my_label': LabelState(status=actual_status, reasons=reasons)}) + labels = EntityLabels(labels={'my_label': LabelState(status=actual_status, reasons=reasons)}) label_provider = StaticLabelProvider({EntityT('MyEntity', 'my_id'): labels}) data = execute( source_with_labels_config( @@ -184,7 +201,7 @@ def test_get_labels_retrieves_data( 'added', None, LabelStatus.ADDED, - {'TestReason': LabelReason(created_at=datetime_to_timestamp(datetime.now() - timedelta(days=1)))}, + LabelReasons({'TestReason': LabelReason(created_at=(datetime.now() - timedelta(days=1)))}), timedelta(days=1), True, ), @@ -192,7 +209,7 @@ def test_get_labels_retrieves_data( 'added', None, LabelStatus.ADDED, - {'TestReason': LabelReason(created_at=datetime_to_timestamp(datetime.now()))}, + LabelReasons({'TestReason': LabelReason(created_at=(datetime.now()))}), timedelta(days=1), False, ), @@ -202,15 +219,15 @@ def test_get_labels_retrieves_data_added_after( execute: ExecuteFunction, checking_status: str, manual: Optional[bool], - actual_status: Optional['LabelStatusValue'], - reasons: Dict[str, LabelReason], + actual_status: Optional[LabelStatus], + reasons: LabelReasons, min_label_age: timedelta, result: bool, ) -> None: if actual_status is None: - labels = Labels(labels={}) + labels = EntityLabels(labels={}) else: - labels = Labels(labels={'my_label': LabelState(status=actual_status, reasons=reasons)}) + labels = EntityLabels(labels={'my_label': LabelState(status=actual_status, reasons=reasons)}) label_provider = StaticLabelProvider({EntityT('MyEntity', 'my_id'): labels}) data = execute( source_with_labels_config( @@ -268,7 +285,7 @@ def test_get_labels_retrieves_data_added_after( def test_gets_only_debounces_single_execution(execute: ExecuteFunction) -> None: - label_provider = BlockingLabelProvider({EntityT('User', 123): Labels(labels={})}) + label_provider = BlockingLabelProvider({EntityT('User', 123): EntityLabels(labels={})}) def do_execute() -> None: execute( diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_rules.py b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_rules.py index 9c25b7c..64b686a 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_rules.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_rules.py @@ -13,7 +13,6 @@ from osprey.engine.conftest import ( RunValidationFunction, ) from osprey.engine.executor.execution_context import ( - EntityLabelMutation, ExecutionContext, ) from osprey.engine.language_types.entities import EntityT @@ -24,11 +23,9 @@ from osprey.engine.stdlib.udfs.time_delta import TimeDelta from osprey.engine.udf.arguments import ArgumentsBase from osprey.engine.udf.base import UDFBase from osprey.engine.udf.registry import UDFRegistry -from osprey.rpc.labels.v1.service_pb2 import LabelStatus +from osprey.worker.lib.osprey_shared.labels import EntityLabelMutation, LabelStatus from osprey.worker.sinks.sink.output_sink import _get_label_effects_from_result -# Moved here because WhenRules is not included in the MVP yet - class FailingUdf(UDFBase[ArgumentsBase, bool]): def execute(self, execution_context: ExecutionContext, arguments: ArgumentsBase) -> bool: diff --git a/osprey_worker/src/osprey/engine/udf/registry.py b/osprey_worker/src/osprey/engine/udf/registry.py index c1b69ba..541d127 100644 --- a/osprey_worker/src/osprey/engine/udf/registry.py +++ b/osprey_worker/src/osprey/engine/udf/registry.py @@ -21,8 +21,12 @@ class UDFRegistry: return instance def register(self, func: Type[UDFBase[Any, Any]]) -> Type[UDFBase[Any, Any]]: - if func.__name__ in self._functions: - raise Exception(f'A function with the name {func.__name__} is already registered.') + # Allow idempotent re-registration of the exact same class. + existing = self.get(func.__name__) + if existing is not None: + if existing is func: + return existing + raise Exception(f'A function with the name {func.__name__} is already registered with {func}.') try: rvalue_type = func.get_rvalue_type() diff --git a/osprey_worker/src/osprey/worker/_stdlibplugin/udf_register.py b/osprey_worker/src/osprey/worker/_stdlibplugin/udf_register.py index 4b5d91f..caa3ab7 100644 --- a/osprey_worker/src/osprey/worker/_stdlibplugin/udf_register.py +++ b/osprey_worker/src/osprey/worker/_stdlibplugin/udf_register.py @@ -15,6 +15,7 @@ from osprey.engine.stdlib.udfs.get_action_name import GetActionName from osprey.engine.stdlib.udfs.import_ import Import from osprey.engine.stdlib.udfs.ip_network import IpNetwork from osprey.engine.stdlib.udfs.json_data import JsonData +from osprey.engine.stdlib.udfs.labels import HasLabel, LabelAdd, LabelRemove from osprey.engine.stdlib.udfs.list_length import ListLength from osprey.engine.stdlib.udfs.list_read import ListRead from osprey.engine.stdlib.udfs.list_sort import ListSort @@ -82,9 +83,12 @@ def register_udfs() -> Sequence[Type[UDFBase[Any, Any]]]: ExperimentsBucketAssignment, ExtractCookie, GetActionName, + HasLabel, Import, IpNetwork, JsonData, + LabelAdd, + LabelRemove, ListLength, ListRead, ListSort, -- 2.51.2