diff --git a/docker-compose.async-worker.yaml b/docker-compose.async-worker.yaml new file mode 100644 index 0000000..9a8733b --- /dev/null +++ b/docker-compose.async-worker.yaml @@ -0,0 +1,24 @@ +# EXPERIMENTAL — opt-in asyncio-native worker (Phase 0: static rules + stdout sink, +# no Kafka input). Kept in its own compose file so it is never built or started by the +# default stack (or the `test` profile). Run it explicitly with: +# docker compose -f docker-compose.async-worker.yaml up --build osprey-async-worker +services: + osprey-async-worker: + container_name: osprey-async-worker + build: + context: . + dockerfile: osprey_async_worker/Dockerfile + command: ["osprey-async-worker"] + environment: + - PYTHONPATH=/osprey + # Bundled stdlib example rules + sample actions so the worker runs out of the box. + - OSPREY_RULES_PATH=/osprey/osprey_async_worker/example_rules + - OSPREY_INPUT_FILE=/osprey/osprey_async_worker/example_rules/sample_actions.jsonl + - DD_TRACE_ENABLED=False + - DD_DOGSTATSD_DISABLE=True + volumes: + - ./osprey_async_worker:/osprey/osprey_async_worker + - ./osprey_worker:/osprey/osprey_worker + - ./osprey_rpc:/osprey/osprey_rpc + - ./example_rules:/osprey/example_rules + - ./entrypoint.sh:/osprey/entrypoint.sh diff --git a/docs/development/README.md b/docs/development/README.md index 80f2dce..80d4956 100644 --- a/docs/development/README.md +++ b/docs/development/README.md @@ -125,6 +125,20 @@ def register_ast_validators() -> None: # Register AST validators ``` +### Available hooks + +Implement any subset of these in your plugin's `register_plugins.py`: + +| Hook | Returns | Notes | +| --- | --- | --- | +| `register_udfs` | `Sequence[Type[UDFBase]]` | Custom user-defined functions. | +| `register_output_sinks` | `Sequence[BaseOutputSink]` | Where execution results go. | +| `register_ast_validators` | `Sequence[Type[BaseValidator]]` | Extra SML validators. | +| `register_action_proto_deserializer` | `ActionProtoDeserializer \| None` | Custom action proto → JSON. | +| `register_input_stream` | `BaseInputStream` | Single-provider (`firstresult`). | +| `register_execution_result_store` | `ExecutionResultStore` | Single-provider (`firstresult`). | +| `register_labels_service_or_provider` | `LabelsServiceBase \| LabelsProvider` | Single-provider (`firstresult`). | + ## Rules Rules are written in SML, some examples are provided in `example_rules/` with YAML config, the rules are mounted to the worker processes when the containers start via environment variables. ex: diff --git a/docs/development/workflow.md b/docs/development/workflow.md index 0e55c24..23ed6ba 100644 --- a/docs/development/workflow.md +++ b/docs/development/workflow.md @@ -17,48 +17,25 @@ Every commit automatically runs: 3. **YAML/JSON/TOML validation** 4. **Ruff linting and formatting** -## Making Changes - -1. **Sync Dependencies** - ```bash - uv sync --frozen - ``` - -2. **Create a new branch:** - - ```bash - git checkout -b username/feature-name - ``` - -3. **Make your changes** - -4. **Run quality checks:** - - ```bash - # This prevents uv.lock from being modified - # this is preferred. - uv sync --frozen - - uv run ruff check --fix - uv run ruff format - ``` - -5. **Test your changes** (if tests exist) +### Manual Checks -6. **Commit your changes:** +Before pushing, run: - ```bash - git add . - git commit -m "feat: descriptive commit message" - ``` +```bash +# Comprehensive linting check +uv run ruff check - Pre-commit hooks will run automatically and may fix formatting issues. +# Format all code +uv run ruff format -7. **Push your branch:** +# Type checking (on specific files/modules) +uv run mypy osprey_worker/src/osprey_worker/lib +# Or you can type check every module (this will happen in CI) +uv run mypy . - ```bash - git push origin username/feature-name - ``` +# Run all pre-commit hooks +uv run pre-commit run --all-files +``` ## Commit Standards @@ -80,22 +57,36 @@ refactor: simplify rule evaluation logic - `test:` - Adding or updating tests - `chore:` - Maintenance tasks -### Manual Checks +## Making Changes -Before pushing, run: +1. **Create a new branch:** -```bash -# Comprehensive linting check -uv run ruff check + ```bash + git checkout -b username/feature-name + ``` -# Format all code -uv run ruff format +2. **Make your changes** -# Type checking (on specific files/modules) -uv run mypy osprey_worker/src/osprey_worker/lib -# Or you can type check every module (this will happen in CI) -uv run mypy . +3. **Run quality checks:** -# Run all pre-commit hooks -uv run pre-commit run --all-files -``` + ```bash + uv run ruff check --fix + uv run ruff format + ``` + +4. **Test your changes** (if tests exist) + +5. **Commit your changes:** + + ```bash + git add . + git commit -m "feat: descriptive commit message" + ``` + + Pre-commit hooks will run automatically and may fix formatting issues. + +6. **Push your branch:** + + ```bash + git push origin username/feature-name + ``` diff --git a/entrypoint.sh b/entrypoint.sh index c798b5c..4d937bb 100755 --- a/entrypoint.sh +++ b/entrypoint.sh @@ -8,6 +8,8 @@ Osprey docker entrypoint. Commands: osprey-worker Runs the worker + osprey-async-worker + [EXPERIMENTAL] Runs the asyncio-native worker (Phase 0; async image only) osprey-ui-api Runs the Osprey UI API run-tests @@ -35,6 +37,20 @@ cli-osprey-worker() { exec uv run python3.11 osprey_worker/src/osprey/worker/cli/sinks.py run-rules-sink } +cli-osprey-async-worker() { + # EXPERIMENTAL: asyncio-native worker (Phase 0 — static/JSONL input, stdout sink; + # no Kafka/coordinator input wired upstream yet). Only functional in the + # osprey_async_worker/Dockerfile image, which installs osprey_async_worker. + # Defaults to the bundled stdlib example rules so the image runs out of the box; + # set OSPREY_INPUT_FILE to feed a JSONL action file (otherwise input is empty). + local rules_path="${OSPREY_RULES_PATH:-/osprey/osprey_async_worker/example_rules}" + local args=(run --rules-path "${rules_path}") + if [[ -n "${OSPREY_INPUT_FILE:-}" ]]; then + args+=(--input-file "${OSPREY_INPUT_FILE}") + fi + exec uv run osprey-async-cli "${args[@]}" "$@" +} + cli-run-tests() { # Only use in CI via harbormaster buildkite run_tests VARIANT PROJECT [directories] # Docker command will be run-tests --junitxml=/osprey/junit-pytest.xml [directory] diff --git a/example_plugins/pyproject.toml b/example_plugins/pyproject.toml index 0eaed7d..af2c977 100644 --- a/example_plugins/pyproject.toml +++ b/example_plugins/pyproject.toml @@ -4,7 +4,7 @@ version = "0.1.0" description = "Example plugins for Osprey" requires-python = ">=3.11" dependencies = [ - "pluggy==1.5.0" + "pluggy==1.5.0", ] [tool.setuptools] @@ -15,3 +15,6 @@ where = ["src"] [project.entry-points.osprey_plugin] register_plugins = "register_plugins" + +[project.entry-points.osprey_async_plugin] +register_async_plugins = "register_async_plugins" diff --git a/example_plugins/src/async_sinks/__init__.py b/example_plugins/src/async_sinks/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/example_plugins/src/async_sinks/example_async_output_sink.py b/example_plugins/src/async_sinks/example_async_output_sink.py new file mode 100644 index 0000000..0b52fb6 --- /dev/null +++ b/example_plugins/src/async_sinks/example_async_output_sink.py @@ -0,0 +1,31 @@ +"""Reference async output sink for the experimental asyncio worker. + +Demonstrates the ``register_async_output_sinks`` hook: the async worker awaits +``push(result)`` for each execution result. Output sinks for the async worker +must subclass ``AsyncBaseOutputSink`` (not the sync ``BaseOutputSink``), since +the async worker runs on asyncio without gevent. +""" + +import logging + +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.engine.executor.execution_context import ExecutionResult + +logger = logging.getLogger(__name__) + + +class ExampleAsyncOutputSink(AsyncBaseOutputSink): + """Logs each result's extracted features and verdicts.""" + + def will_do_work(self, result: ExecutionResult) -> bool: + return True + + async def push(self, result: ExecutionResult) -> None: + logger.info( + 'example async output sink: features=%s verdicts=%s', + result.extracted_features_json, + result.verdicts, + ) + + async def stop(self) -> None: + pass diff --git a/example_plugins/src/register_async_plugins.py b/example_plugins/src/register_async_plugins.py new file mode 100644 index 0000000..926adef --- /dev/null +++ b/example_plugins/src/register_async_plugins.py @@ -0,0 +1,31 @@ +"""Example plugin registrations for the experimental asyncio worker. + +This is the async counterpart to ``register_plugins`` (which targets the sync +gevent worker). It is discovered via the ``osprey_async_plugin`` entry-point +group and loaded by the async worker's plugin manager, so it never runs in the +sync worker. + +Registers a pure-computation UDF (``TextContains`` runs inline in the async +executor — no I/O, so it needs no async variant) and an example async output +sink. UDFs that perform I/O must subclass ``AsyncUDFBase`` instead; see +``osprey.async_worker.stdlib_udfs.async_mx_lookup`` for that pattern. +""" + +from typing import Any, Sequence, Type + +from async_sinks.example_async_output_sink import ExampleAsyncOutputSink +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.async_worker.adaptor.plugin_manager import hookimpl_osprey_async +from osprey.engine.udf.base import UDFBase +from osprey.worker.lib.config import Config +from udfs.text_contains import TextContains + + +@hookimpl_osprey_async +def register_udfs() -> Sequence[Type[UDFBase[Any, Any]]]: + return [TextContains] + + +@hookimpl_osprey_async +def register_async_output_sinks(config: Config) -> Sequence[AsyncBaseOutputSink]: + return [ExampleAsyncOutputSink()] diff --git a/example_plugins/src/tests/__init__.py b/example_plugins/src/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/Dockerfile b/osprey_async_worker/Dockerfile new file mode 100644 index 0000000..43fb3b7 --- /dev/null +++ b/osprey_async_worker/Dockerfile @@ -0,0 +1,72 @@ +# syntax=docker/dockerfile:1 +# +# EXPERIMENTAL — asyncio-native Osprey worker (opt-in, NOT production-ready). +# +# The asyncio worker is a Phase-0 prototype. Its CLI (`osprey-async-cli run`) +# reads actions from a static/JSONL source and writes to a stdout sink — there +# is no Kafka or coordinator input wired upstream yet. Build/run this image only +# to try the asyncio engine. The stable gevent worker lives in +# osprey_worker/Dockerfile and is unaffected by this image. +# +FROM python:3.11-slim +ARG TARGETPLATFORM + +WORKDIR /osprey +# Image-structural only: the editable workspace packages resolve from /osprey. +ENV PYTHONPATH=/osprey +# Runtime config (rules path, input file, etc.) is supplied via env at run time +# (docker run -e / compose), not baked here. The Phase-0 async worker reads rules +# from --rules-path; it does not use etcd or serve a port. + +RUN set -ex && \ + apt-get update && \ + apt-get install -yqq --no-install-recommends libjemalloc-dev git gcc g++ wget ssh && \ + apt-get clean && \ + apt-get -y purge wget && \ + apt-get -y autoremove && \ + rm -fr /var/cache/apt/archives/* && \ + rm -rf /var/lib/apt/lists/* + +# Hack to get the right libjemalloc location +ENV LIB_ARCH=${TARGETPLATFORM##*/} +ENV LIB_ARCH=${LIB_ARCH/amd64/x86_64-linux-gnu} +ENV LIB_ARCH=${LIB_ARCH/arm64/aarch64-linux-gnu} +ENV LD_PRELOAD="/usr/lib/${LIB_ARCH}/libjemalloc.so" +ENV MALLOC_CONF="narenas:4" + +# Set up python environment +ADD uv.lock /osprey/uv.lock +ADD pyproject.toml /osprey/pyproject.toml +ADD README.md /osprey/README.md +ADD LICENSE.md /osprey/LICENSE.md + +# Create workspace structure with pyproject.toml files for uv sync to work +ADD osprey_rpc/pyproject.toml /osprey/osprey_rpc/pyproject.toml +ADD osprey_worker/pyproject.toml /osprey/osprey_worker/pyproject.toml +ADD osprey_async_worker/pyproject.toml /osprey/osprey_async_worker/pyproject.toml +ADD example_plugins/pyproject.toml /osprey/example_plugins/pyproject.toml + +# Create minimal package structure required by uv +RUN mkdir -p /osprey/osprey_worker /osprey/osprey_async_worker /osprey/osprey_rpc /osprey/example_plugins/src && \ + touch /osprey/osprey_worker/__init__.py /osprey/osprey_async_worker/__init__.py /osprey/osprey_rpc/__init__.py /osprey/example_plugins/src/__init__.py + +# Install the full workspace, including osprey_async_worker (this is the one image +# that does). This layer is cached when only source code changes. +RUN pip install --upgrade pip uv && \ + uv sync --locked --python=$(which python3.11) && \ + pip cache purge && uv cache clean + +# https://tld.readthedocs.io/en/latest/#update-the-list-of-tld-names +RUN . .venv/bin/activate && update-tld-names + +# Add source code after dependencies are installed +ADD example_rules /osprey/example_rules +ADD osprey_worker /osprey/osprey_worker +ADD osprey_async_worker /osprey/osprey_async_worker +ADD osprey_rpc /osprey/osprey_rpc +ADD example_plugins /osprey/example_plugins + +COPY entrypoint.sh /osprey/entrypoint.sh + +ENTRYPOINT ["/osprey/entrypoint.sh"] +CMD ["osprey-async-worker"] diff --git a/osprey_async_worker/example_rules/config/example_config.yaml b/osprey_async_worker/example_rules/config/example_config.yaml new file mode 100644 index 0000000..e5a5949 --- /dev/null +++ b/osprey_async_worker/example_rules/config/example_config.yaml @@ -0,0 +1,5 @@ +# Minimal config so the rules directory loads — Osprey requires a non-empty +# config alongside the rules. The asyncio Phase-0 example does not use it. +example: + some_str: "demo" + some_int: 1 diff --git a/osprey_async_worker/example_rules/main.sml b/osprey_async_worker/example_rules/main.sml new file mode 100644 index 0000000..2c79a61 --- /dev/null +++ b/osprey_async_worker/example_rules/main.sml @@ -0,0 +1,7 @@ +# Example rules for the EXPERIMENTAL asyncio worker. +# +# These use only stdlib UDFs, so they compile with `osprey-async-cli run` +# (the Phase-0 stdlib engine, no plugins required). The top-level ./example_rules +# is for the gevent worker and depends on plugin UDFs (TextContains, BanUser), +# which the stdlib async engine does not provide. +Require(rule='rules/long_message.sml') diff --git a/osprey_async_worker/example_rules/rules/long_message.sml b/osprey_async_worker/example_rules/rules/long_message.sml new file mode 100644 index 0000000..87271f8 --- /dev/null +++ b/osprey_async_worker/example_rules/rules/long_message.sml @@ -0,0 +1,17 @@ +# Flag actions whose message text is longer than 100 characters. +# +# Demonstrates stdlib-only UDFs the async Phase-0 engine supports: +# EntityJson, JsonData, GetActionName, StringLength. +UserId: Entity[str] = EntityJson(type='User', path='$.user_id', coerce_type=True) +EventType: str = JsonData(path='$.event_type', coerce_type=True) +ActionName = GetActionName() + +MessageText: str = JsonData(path='$.message', coerce_type=True) +MessageLength = StringLength(s=MessageText) + +LongMessage = Rule( + when_all=[ + MessageLength > 100, + ], + description='Message text is longer than 100 characters', +) diff --git a/osprey_async_worker/example_rules/sample_actions.jsonl b/osprey_async_worker/example_rules/sample_actions.jsonl new file mode 100644 index 0000000..78f16ee --- /dev/null +++ b/osprey_async_worker/example_rules/sample_actions.jsonl @@ -0,0 +1,2 @@ +{"id": 1, "name": "create_post", "data": {"user_id": "u1", "event_type": "create_post", "message": "short message"}} +{"id": 2, "name": "create_post", "data": {"user_id": "u2", "event_type": "create_post", "message": "hello world this is a deliberately long message that exceeds one hundred characters so the LongMessage rule fires"}} diff --git a/osprey_async_worker/pyproject.toml b/osprey_async_worker/pyproject.toml new file mode 100644 index 0000000..aea1f4a --- /dev/null +++ b/osprey_async_worker/pyproject.toml @@ -0,0 +1,28 @@ +[build-system] +requires = ["setuptools>=61.0", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "osprey-async-worker" +version = "0.1.0" +description = "Asyncio-based Osprey worker" +readme = "README.md" +license = "Apache-2.0" +requires-python = ">=3.11" +dependencies = [ + "osprey_rpc", + "osprey_worker", + "aiodns", + "pycares", +] + +[project.scripts] +osprey-async-cli = "osprey.async_worker.cli.main:cli" + +[project.entry-points."osprey_async_plugin"] + +[tool.setuptools] +include-package-data = true + +[tool.setuptools.package-data] +"*" = ["*.json", "*.yaml", "*.yml"] diff --git a/osprey_async_worker/src/osprey/async_worker/__init__.py b/osprey_async_worker/src/osprey/async_worker/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/adaptor/__init__.py b/osprey_async_worker/src/osprey/async_worker/adaptor/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/adaptor/constants.py b/osprey_async_worker/src/osprey/async_worker/adaptor/constants.py new file mode 100644 index 0000000..0ff5dce --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/adaptor/constants.py @@ -0,0 +1 @@ +OSPREY_ASYNC_ADAPTOR = 'osprey_async_plugin' diff --git a/osprey_async_worker/src/osprey/async_worker/adaptor/hookspecs.py b/osprey_async_worker/src/osprey/async_worker/adaptor/hookspecs.py new file mode 100644 index 0000000..c60e0af --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/adaptor/hookspecs.py @@ -0,0 +1,92 @@ +"""Hook specifications for the async worker plugin system. + +Mirrors osprey.worker.adaptor.hookspecs but uses the 'osprey_async_plugin' +entry_point group. Plugins register async output sinks and UDFs here. + +UDFs are shared with the sync worker (they're registered via the existing +'osprey_plugin' hooks and wrapped with SyncUDFAdapter). Async-native UDFs +can also be registered here. + +Output sinks MUST be async (AsyncBaseOutputSink) since the async worker +doesn't use gevent. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, Sequence, Tuple, Type + +import pluggy +from osprey.async_worker.adaptor.constants import OSPREY_ASYNC_ADAPTOR +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.engine.ast_validator.base_validator import BaseValidator +from osprey.engine.udf.base import UDFBase + +if TYPE_CHECKING: + from osprey.worker.lib.action_proto_deserializer import ActionProtoDeserializer + from osprey.worker.lib.config import Config + from osprey.worker.lib.data_exporters.validation_result_exporter import BaseValidationResultExporter + +hookspec: pluggy.HookspecMarker = pluggy.HookspecMarker(OSPREY_ASYNC_ADAPTOR) + + +@hookspec +def register_async_output_sinks(config: Config) -> Sequence[AsyncBaseOutputSink]: + """Register async output sinks for the async worker. + + These must be AsyncBaseOutputSink instances (not sync BaseOutputSink). + The async worker will call `await sink.push(result)` for each result. + """ + raise NotImplementedError + + +@hookspec +def register_udfs() -> Sequence[Type[UDFBase[Any, Any]]]: + """Register UDFs for the async worker. + + These are the same UDFBase types as the sync worker. The async executor + runs them in a thread pool via run_in_executor. Async-native UDFs can + also be registered here in the future. + """ + raise NotImplementedError + + +@hookspec +def register_ast_validators() -> Sequence[Type[BaseValidator]]: + """Register AST validators. Same interface as the sync worker.""" + raise NotImplementedError + + +@hookspec +def register_action_proto_deserializer() -> 'ActionProtoDeserializer': + """Register an action proto deserializer. + + Same interface as the sync worker's register_action_proto_deserializer. + """ + raise NotImplementedError + + +@hookspec +def register_udf_helpers(config: 'Config') -> Sequence[Tuple[Type[UDFBase[Any, Any]], Any]]: + """Register `(udf_class, helper)` bindings for UDFs that need a runtime helper + not constructible from the UDF class alone (typically because the helper + depends on `config` or wraps an external service client). + + Plugins return a sequence of pairs; the framework calls + ``udf_helpers.set_udf_helper(udf_class, helper)`` for each. The plugin owns + the UDF-class import and the helper construction, so the framework never + needs to know about plugin-provided UDF types. + + UDFs that extend :class:`HasHelper` are wired automatically and should not + appear here. + """ + raise NotImplementedError + + +@hookspec(firstresult=True) +def register_validation_exporter(config: 'Config') -> 'BaseValidationResultExporter': + """Register a validation result exporter. + + Called after rule compilation to export experiment metadata + and other validation results to analytics. + """ + raise NotImplementedError diff --git a/osprey_async_worker/src/osprey/async_worker/adaptor/interfaces.py b/osprey_async_worker/src/osprey/async_worker/adaptor/interfaces.py new file mode 100644 index 0000000..ca96123 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/adaptor/interfaces.py @@ -0,0 +1,129 @@ +"""Async plugin interfaces for the osprey async worker. + +AsyncUDFBase extends UDFBase so it works with the existing engine machinery +(UDFRegistry, CallExecutor, validation, type checking, argument resolution). + +All I/O UDFs in the async worker MUST be AsyncUDFBase subclasses — existing +sync UDFs that use gevent primitives will not work without monkey patching. +Pure-computation UDFs (no I/O) can remain as regular UDFBase and run inline. +""" + +import abc +from typing import Any, ClassVar, Sequence, TypeVar + +from osprey.engine.executor.execution_context import ExecutionContext, ExecutionResult +from osprey.engine.udf.base import ( + Arguments, + BatchableArguments, + BatchableUDFBase, + RValue, + UDFBase, +) +from result import Result + +_T = TypeVar('_T') + + +class AsyncUDFBase(UDFBase[Arguments, RValue]): + """Native async UDF base class. + + Extends UDFBase so it integrates with UDFRegistry, CallExecutor, and + the full validation/type-checking pipeline. The async executor detects + AsyncUDFBase instances via isinstance() and awaits async_execute() + directly on the event loop — no thread pool. + + The sync execute() raises so it can't accidentally be called in the + async executor's sync path. + """ + + execute_async: ClassVar[bool] = True + is_native_async: ClassVar[bool] = True + + def __init__(self, validation_context, arguments): + super().__init__(validation_context, arguments) + + @classmethod + def _get_udf_base_args(cls): + """Override to include AsyncUDFBase in the generic origin check.""" + import typing_inspect + + for base in cls.__mro__: + for generic_base in typing_inspect.get_generic_bases(base): + origin = typing_inspect.get_origin(generic_base) + if origin in (UDFBase, AsyncUDFBase) or (hasattr(origin, '__mro__') and UDFBase in origin.__mro__): + args = typing_inspect.get_args(generic_base) + # Only return if args are concrete (not TypeVars) + if args and not any(isinstance(a, TypeVar) for a in args): + return args + + # Fallback to parent + return super()._get_udf_base_args() + + def execute(self, execution_context: ExecutionContext, arguments: Arguments) -> RValue: + raise RuntimeError( + f'{self.__class__.__name__} is a native async UDF. Use async_execute() instead of execute().' + ) + + @abc.abstractmethod + async def async_execute(self, execution_context: ExecutionContext, arguments: Arguments) -> RValue: + """Override this to implement the UDF's async execution logic.""" + raise NotImplementedError + + +class AsyncBatchableUDFBase(BatchableUDFBase[Arguments, RValue, BatchableArguments]): + """Native async batchable UDF base class. + + Same as AsyncUDFBase but for batchable UDFs. The async executor detects + these and awaits async_execute_batch() directly. + """ + + is_native_async: ClassVar[bool] = True + + def execute(self, execution_context: ExecutionContext, arguments: Arguments) -> RValue: + raise RuntimeError( + f'{self.__class__.__name__} is a native async UDF. Use async_execute() instead of execute().' + ) + + def execute_batch( + self, + execution_context: ExecutionContext, + udfs: Sequence[UDFBase[Any, Any]], + arguments: Sequence[BatchableArguments], + ) -> Sequence[Result[RValue, Exception]]: + raise RuntimeError( + f'{self.__class__.__name__} is a native async UDF. Use async_execute_batch() instead of execute_batch().' + ) + + @abc.abstractmethod + async def async_execute(self, execution_context: ExecutionContext, arguments: Arguments) -> RValue: + raise NotImplementedError + + @abc.abstractmethod + async def async_execute_batch( + self, + execution_context: ExecutionContext, + udfs: Sequence[UDFBase[Any, Any]], + arguments: Sequence[BatchableArguments], + ) -> Sequence[Result[RValue, Exception]]: + raise NotImplementedError + + +# --- Output sinks and input streams (unchanged) --- + + +class AsyncBaseOutputSink(abc.ABC): + """Async output sink.""" + + timeout: float = 2.0 + max_retries: int = 0 + + @abc.abstractmethod + def will_do_work(self, result: ExecutionResult) -> bool: + raise NotImplementedError + + @abc.abstractmethod + async def push(self, result: ExecutionResult) -> None: + raise NotImplementedError + + async def stop(self) -> None: + pass diff --git a/osprey_async_worker/src/osprey/async_worker/adaptor/plugin_manager.py b/osprey_async_worker/src/osprey/async_worker/adaptor/plugin_manager.py new file mode 100644 index 0000000..38c41bf --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/adaptor/plugin_manager.py @@ -0,0 +1,170 @@ +"""Plugin manager for the async worker. + +Discovers plugins via the 'osprey_async_plugin' setuptools entry_point group. +All UDFs with I/O must have native async implementations — no sync fallbacks. +""" + +from __future__ import annotations + +import logging +from functools import lru_cache +from typing import TYPE_CHECKING, Any, List, Type, cast + +import pluggy +from osprey.async_worker.adaptor import hookspecs as async_hookspecs +from osprey.async_worker.adaptor.constants import OSPREY_ASYNC_ADAPTOR +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.async_worker.sinks.sink.output_sink import AsyncMultiOutputSink +from osprey.engine.ast_validator import ValidatorRegistry +from osprey.engine.executor.udf_execution_helpers import HasHelper, UDFHelpers +from osprey.engine.udf.base import UDFBase +from osprey.engine.udf.registry import UDFRegistry +from osprey.worker.lib.action_proto_deserializer import ActionProtoDeserializer + +if TYPE_CHECKING: + from osprey.worker.lib.config import Config + +hookimpl_osprey_async: pluggy.HookimplMarker = pluggy.HookimplMarker(OSPREY_ASYNC_ADAPTOR) + +plugin_manager = pluggy.PluginManager(OSPREY_ASYNC_ADAPTOR) +plugin_manager.add_hookspecs(async_hookspecs) + + +def _flatten(seq: List[List[Any]]) -> List[Any]: + return sum(seq, []) + + +@lru_cache(maxsize=1) +def load_all_async_plugins() -> None: + """Load the first-party async-stdlib plugin and any third-party plugins. + + The first-party plugin (osprey.async_worker.stdlib_udfs._async_stdlib_plugin) + contributes async-native replacements for sync stdlib UDFs (e.g. MXLookup). + Third-party plugins are discovered via the 'osprey_async_plugin' setuptools + entry_point group. + """ + from osprey.async_worker.stdlib_udfs import _async_stdlib_plugin + + plugin_manager.register(_async_stdlib_plugin) + plugin_manager.load_setuptools_entrypoints(OSPREY_ASYNC_ADAPTOR) + plugin_manager.check_pending() + + +def _deduplicate_udfs( + stdlib_udfs: List[Type[UDFBase[Any, Any]]], + plugin_udfs: List[Type[UDFBase[Any, Any]]], +) -> List[Type[UDFBase[Any, Any]]]: + """Merge stdlib and plugin UDFs, with plugin UDFs winning on name conflicts. + + Async plugin UDFs shadow their sync stdlib counterparts by class name. + This lets async plugins register e.g. `HasLabel` or `MXLookup` without + needing a separate replacement table — the plugin version just wins. + """ + plugin_names = {udf.__name__ for udf in plugin_udfs} + deduplicated = [udf for udf in stdlib_udfs if udf.__name__ not in plugin_names] + deduplicated.extend(plugin_udfs) + return deduplicated + + +def bootstrap_async_udfs(config: 'Config | None' = None) -> tuple[UDFRegistry, UDFHelpers]: + """Bootstrap UDFs from async plugins + stdlib. + + Loads stdlib UDFs (JsonData, StringLength, Rule, etc.) and async plugin UDFs. + Plugin UDFs override stdlib UDFs with the same name — this is how async + replacements (HasLabel, MXLookup, etc.) shadow their sync counterparts. + Async-native replacements for sync stdlib UDFs come from the first-party + `_async_stdlib_plugin`, which registers through the same `register_udfs` + hook as third-party plugins. No sync fallbacks — all I/O UDFs must be + native async. + """ + from osprey.worker._stdlibplugin.udf_register import register_udfs as stdlib_register_udfs + + load_all_async_plugins() + udf_helpers = UDFHelpers() + + stdlib_udfs = list(stdlib_register_udfs()) + plugin_udfs = _flatten(plugin_manager.hook.register_udfs()) + all_udfs = _deduplicate_udfs(stdlib_udfs, plugin_udfs) + + # Auto-register helpers for UDFs that extend HasHelper + for udf in all_udfs: + if issubclass(udf, HasHelper): + udf_helpers.set_udf_helper(udf, udf.create_provider()) + + # Plugin-provided helper bindings for UDFs whose helper depends on `config` + # or wraps an external service. Each plugin returns `(udf_class, helper)` + # pairs; the framework binds them without needing to import the UDF class. + if config is not None: + for udf_class, helper in _iter_plugin_udf_helpers(config): + # Plugin-provided pairs are helper-bearing UDFs by contract, but the + # collection is typed as the looser UDFBase; narrow for set_udf_helper. + udf_helpers.set_udf_helper(cast('Type[HasHelper[Any]]', udf_class), helper) + + udf_registry = UDFRegistry.with_udfs(*all_udfs) + return udf_registry, udf_helpers + + +def _iter_plugin_udf_helpers(config: Config) -> List[tuple[Type[UDFBase[Any, Any]], Any]]: + """Collect `(udf_class, helper)` bindings from every plugin that implements + the `register_udf_helpers` hook. Failures are logged and skipped so one + broken plugin doesn't take down bootstrap. + """ + pairs: List[tuple[Type[UDFBase[Any, Any]], Any]] = [] + if not hasattr(plugin_manager.hook, 'register_udf_helpers'): + return pairs + try: + for plugin_result in plugin_manager.hook.register_udf_helpers(config=config): + pairs.extend(plugin_result) + except Exception: + logging.exception('Failed to collect UDF helpers from plugins') + return pairs + + +def bootstrap_async_action_proto_deserializer() -> ActionProtoDeserializer | None: + """Bootstrap action proto deserializer from async plugins.""" + load_all_async_plugins() + try: + [deserializer] = plugin_manager.hook.register_action_proto_deserializer() + return deserializer + except Exception: + return None + + +def bootstrap_validation_exporter(config: Config) -> Any: + """Bootstrap validation result exporter from async plugins. + + Returns the exporter or None if not registered. + """ + load_all_async_plugins() + if not hasattr(plugin_manager.hook, 'register_validation_exporter'): + return None + try: + return plugin_manager.hook.register_validation_exporter(config=config) + except Exception: + logging.exception('Failed to bootstrap validation exporter') + return None + + +def bootstrap_async_output_sinks(config: Config) -> AsyncMultiOutputSink: + """Bootstrap async output sinks from async plugins only. + + Does NOT load sync output sinks — the async worker uses only async sinks. + """ + load_all_async_plugins() + sinks: List[AsyncBaseOutputSink] = _flatten(plugin_manager.hook.register_async_output_sinks(config=config)) + return AsyncMultiOutputSink(sinks) + + +def bootstrap_async_ast_validators() -> None: + """Bootstrap AST validators from async plugins + stdlib.""" + from osprey.worker._stdlibplugin.validator_regsiter import register_ast_validators as stdlib_register_validators + + load_all_async_plugins() + validators = list(stdlib_register_validators()) + _flatten(plugin_manager.hook.register_ast_validators()) + + registry = ValidatorRegistry.get_instance() + seen = set() + for validator in validators: + if validator not in seen: + seen.add(validator) + registry.register_to_instance(validator) diff --git a/osprey_async_worker/src/osprey/async_worker/cli/__init__.py b/osprey_async_worker/src/osprey/async_worker/cli/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/cli/main.py b/osprey_async_worker/src/osprey/async_worker/cli/main.py new file mode 100644 index 0000000..06663f7 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/cli/main.py @@ -0,0 +1,416 @@ +"""Minimal async worker CLI for Phase 0 validation. + +No monkey patching. No gevent. Uses asyncio event loop. +Supports static rules files for testing without etcd. +""" + +import asyncio +import json +import logging +import signal +from datetime import datetime, timezone +from pathlib import Path +from typing import AsyncIterator, Optional, Tuple + +import click +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.async_worker.engine import AsyncOspreyEngine +from osprey.async_worker.sinks.sink.input_stream import ( + AsyncBaseInputStream, + AsyncKafkaInputStream, + AsyncStaticInputStream, +) +from osprey.async_worker.sinks.sink.output_sink import AsyncStdoutOutputSink +from osprey.async_worker.sinks.sink.rules_sink import AsyncRulesSink +from osprey.engine.ast.sources import Sources +from osprey.engine.executor.execution_context import Action +from osprey.engine.executor.udf_execution_helpers import UDFHelpers +from osprey.engine.udf.registry import UDFRegistry +from osprey.worker.lib.config import Config +from osprey.worker.lib.instruments import set_worker_type_tag +from osprey.worker.lib.osprey_shared.logging import get_logger +from osprey.worker.lib.singletons import CONFIG +from osprey.worker.lib.sources_provider_base import StaticSourcesProvider +from osprey.worker.sinks.utils.acking_contexts_base import BaseAckingContext, NoopAckingContext + +logger = get_logger(__name__) + +# How long to let the sink drain the in-flight action on shutdown before cancelling. +_GRACEFUL_SHUTDOWN_SECONDS = 10.0 + + +def init_config() -> Config: + config = CONFIG.instance() + config.configure_from_env() + set_worker_type_tag('async') + return config + + +def bootstrap_stdlib_engine(rules_path: str) -> Tuple[AsyncOspreyEngine, UDFHelpers]: + """Bootstrap engine with only stdlib UDFs — no external plugins, no Postgres, no labels. + + This avoids loading example_plugins or any third-party plugins that require database connections. + """ + from osprey.engine.ast_validator import ValidatorRegistry + from osprey.worker._stdlibplugin.udf_register import register_udfs as stdlib_register_udfs + from osprey.worker._stdlibplugin.validator_regsiter import register_ast_validators as stdlib_register_validators + + udf_helpers = UDFHelpers() + udfs = stdlib_register_udfs() + udf_registry = UDFRegistry.with_udfs(*udfs) + + validators = stdlib_register_validators() + registry = ValidatorRegistry.get_instance() + for validator in validators: + registry.register_to_instance(validator) + + sources_provider = StaticSourcesProvider(sources=Sources.from_path(Path(rules_path))) + + engine = AsyncOspreyEngine( + sources_provider=sources_provider, + udf_registry=udf_registry, + ) + + return engine, udf_helpers + + +def bootstrap_plugin_engine(rules_path: str) -> Tuple[AsyncOspreyEngine, UDFHelpers, AsyncBaseOutputSink]: + """Bootstrap the async engine with all async plugins loaded. + + Loads UDFs, AST validators, and output sinks from the ``osprey_async_plugin`` + entry-point group (plus the first-party async stdlib plugin). No gevent and + no sync engine — the output sink is whatever the plugins register. + """ + from osprey.async_worker.adaptor.plugin_manager import ( + bootstrap_async_ast_validators, + bootstrap_async_output_sinks, + bootstrap_async_udfs, + ) + + config = CONFIG.instance() + bootstrap_async_ast_validators() + udf_registry, udf_helpers = bootstrap_async_udfs(config) + + sources_provider = StaticSourcesProvider(sources=Sources.from_path(Path(rules_path))) + engine = AsyncOspreyEngine(sources_provider=sources_provider, udf_registry=udf_registry) + + output_sink: AsyncBaseOutputSink = bootstrap_async_output_sinks(config) + if not getattr(output_sink, '_sinks', None): + logger.warning( + 'No async output sink was registered by any osprey_async_plugin; ' + 'falling back to the stdout sink so execution results are not silently dropped.' + ) + output_sink = AsyncStdoutOutputSink() + return engine, udf_helpers, output_sink + + +class AsyncFileInputStream(AsyncBaseInputStream[BaseAckingContext[Action]]): + """Read actions from a JSON file. Each line is a JSON action object.""" + + def __init__(self, path: str): + self._path = path + + async def _gen(self) -> AsyncIterator[BaseAckingContext[Action]]: + with open(self._path) as f: + for line in f: + line = line.strip() + if not line: + continue + data = json.loads(line) + action = Action( + action_id=data.get('id', 0), + action_name=data.get('name', 'unknown'), + data=data.get('data', {}), + timestamp=datetime.now(timezone.utc), + ) + yield NoopAckingContext(action) + + +def build_kafka_input_stream( + topic: str, + bootstrap_servers: str, + group_id: Optional[str], + offset_reset: str, +) -> AsyncKafkaInputStream: + """Build a Kafka-backed async input stream. + + Uses the plain kafka-python consumer (not the gevent-patched one); the async + worker polls it off the event loop, so the FairRLock patch is unnecessary. + The envelope shape matches the gevent worker's KafkaInputStream. + """ + from kafka import KafkaConsumer + + # Manual commit: AsyncKafkaInputStream commits offsets only after a polled + # batch has been handed to (and processed by) the rules sink, so a crash + # mid-batch reprocesses rather than silently dropping actions (at-least-once). + consumer = KafkaConsumer( + topic, + bootstrap_servers=[s.strip() for s in bootstrap_servers.split(',') if s.strip()], + group_id=group_id, + auto_offset_reset=offset_reset, + enable_auto_commit=False, + ) + return AsyncKafkaInputStream(consumer) + + +@click.group() +def cli() -> None: + pass + + +@cli.command() +@click.option('--rules-path', type=click.Path(exists=True), required=True, help='Path to rules directory') +@click.option('--input-file', type=click.Path(exists=True), default=None, help='Path to JSONL input file') +@click.option('--max-concurrent', type=int, default=12, help='Max concurrent async UDF executions') +@click.option('--with-plugins', is_flag=True, default=False, help='Load all plugins (requires external services)') +@click.option( + '--input-source', + type=click.Choice(['file', 'kafka']), + default='file', + help='Where actions come from: a JSONL file or a Kafka topic.', +) +@click.option('--kafka-topic', default='osprey.actions_input', help='Kafka topic to consume (--input-source kafka).') +@click.option( + '--kafka-bootstrap-servers', + default='localhost:9092', + help='Comma-separated Kafka bootstrap servers (--input-source kafka).', +) +@click.option('--kafka-group-id', default=None, help='Kafka consumer group id (--input-source kafka).') +@click.option( + '--kafka-offset-reset', + type=click.Choice(['latest', 'earliest']), + default='latest', + help='Where to start consuming when no committed offset exists.', +) +def run( + rules_path: str, + input_file: Optional[str], + max_concurrent: int, + with_plugins: bool, + input_source: str, + kafka_topic: str, + kafka_bootstrap_servers: str, + kafka_group_id: Optional[str], + kafka_offset_reset: str, +) -> None: + """Run the async rules worker with a static rules file and optional input file.""" + logging.basicConfig(level=logging.INFO, format='%(asctime)s %(levelname)s %(name)s: %(message)s') + logger.info('Starting async osprey worker (Phase 0)') + logger.info(f'Rules path: {rules_path}') + logger.info(f'Max concurrent UDFs: {max_concurrent}') + + # Side-effecting: configures the global CONFIG singleton and worker-type tag. + init_config() + + # --with-plugins loads the async plugin system (osprey_async_plugin entry + # points plus the first-party async stdlib plugin): plugin UDFs, validators, + # and output sinks. The default branch uses stdlib UDFs and a stdout sink. + engine: AsyncOspreyEngine + output_sink: AsyncBaseOutputSink + if with_plugins: + engine, udf_helpers, output_sink = bootstrap_plugin_engine(rules_path) + else: + engine, udf_helpers = bootstrap_stdlib_engine(rules_path) + output_sink = AsyncStdoutOutputSink() + + # Input stream + input_stream: AsyncBaseInputStream[BaseAckingContext[Action]] + if input_source == 'kafka': + logger.info(f'Consuming from Kafka topic {kafka_topic!r} at {kafka_bootstrap_servers}') + input_stream = build_kafka_input_stream( + topic=kafka_topic, + bootstrap_servers=kafka_bootstrap_servers, + group_id=kafka_group_id, + offset_reset=kafka_offset_reset, + ) + elif input_file: + input_stream = AsyncFileInputStream(input_file) + else: + # No input — just validate the worker boots correctly + input_stream = AsyncStaticInputStream([]) + + rules_sink = AsyncRulesSink( + engine=engine, + input_stream=input_stream, + output_sink=output_sink, + udf_helpers=udf_helpers, + max_concurrent_udfs=max_concurrent, + ) + + async def _run(): + loop = asyncio.get_running_loop() + stop_event = asyncio.Event() + + def _signal_handler(): + logger.info('Received shutdown signal') + stop_event.set() + + for sig in (signal.SIGTERM, signal.SIGINT): + loop.add_signal_handler(sig, _signal_handler) + + sink_task = asyncio.create_task(rules_sink.run()) + + # Wait for either the sink to finish or a shutdown signal + stop_task = asyncio.create_task(stop_event.wait()) + await asyncio.wait([sink_task, stop_task], return_when=asyncio.FIRST_COMPLETED) + + if not stop_task.done(): + stop_task.cancel() + + sink_error: Optional[BaseException] = None + if sink_task.done(): + # The sink finished on its own rather than via a shutdown signal — + # surface any fatal error instead of reporting a clean shutdown. + sink_error = sink_task.exception() + await rules_sink.stop() + else: + # Graceful shutdown under a single bounded budget covering BOTH + # stopping the input stream (which can itself block, e.g. Kafka + # consumer.close()) and draining the in-flight action. The input + # stream is stopped first so the action can finalize through the + # stream's own shutdown path (the coordinator stream acks/nacks and + # graceful-disconnects after the current yield resumes); we fall back + # to cancellation only if the whole thing overruns the budget. + async def _stop_and_drain() -> None: + await rules_sink.stop() + await sink_task + + drain = asyncio.ensure_future(_stop_and_drain()) + try: + await asyncio.wait_for(asyncio.shield(drain), timeout=_GRACEFUL_SHUTDOWN_SECONDS) + except asyncio.TimeoutError: + logger.warning('Graceful shutdown exceeded %ss; cancelling', _GRACEFUL_SHUTDOWN_SECONDS) + drain.cancel() + sink_task.cancel() + for task in (drain, sink_task): + try: + await task + except asyncio.CancelledError: + pass + except Exception: + pass # a sink failure is surfaced via sink_error below + if sink_task.done() and not sink_task.cancelled(): + sink_error = sink_task.exception() + + if sink_error is not None: + logger.error('Async worker sink task failed', exc_info=sink_error) + raise sink_error + + logger.info('Async worker shutdown complete') + + asyncio.run(_run()) + + +@cli.command() +@click.option('--rules-path', type=click.Path(exists=True), required=True, help='Path to rules directory') +@click.option('--input-file', type=click.Path(exists=True), required=True, help='Path to JSONL input file') +@click.option('--max-concurrent', type=int, default=12, help='Max concurrent async UDF executions') +@click.option('--iterations', type=int, default=1000, help='Number of iterations to run') +@click.option('--warmup', type=int, default=50, help='Warmup iterations (not counted)') +def benchmark(rules_path: str, input_file: str, max_concurrent: int, iterations: int, warmup: int) -> None: + """Benchmark the async executor vs the gevent executor. + + Runs both executors against the same rules and input data, then compares + throughput and latency. + """ + import time + + from osprey.async_worker.executor import execute as async_execute + + logging.basicConfig(level=logging.WARNING) + # Side-effecting: configures the global CONFIG singleton and worker-type tag. + init_config() + engine, udf_helpers = bootstrap_stdlib_engine(rules_path) + + # Load actions + actions = [] + with open(input_file) as f: + for line in f: + line = line.strip() + if not line: + continue + data = json.loads(line) + actions.append( + Action( + action_id=data.get('id', 0), + action_name=data.get('name', 'unknown'), + data=data.get('data', {}), + timestamp=datetime.now(timezone.utc), + ) + ) + + if not actions: + click.echo('No actions found in input file') + return + + click.echo(f'Loaded {len(actions)} actions, {iterations} iterations (+ {warmup} warmup)') + click.echo(f'Rules: {rules_path}') + click.echo() + + # --- Gevent executor (optional, for comparison) --- + try: + import gevent.pool + from osprey.engine.executor.executor import execute as gevent_execute + + pool = gevent.pool.Pool(max_concurrent) + + for i in range(warmup): + action = actions[i % len(actions)] + gevent_execute(engine.execution_graph, udf_helpers, action, pool) + + start = time.perf_counter() + for i in range(iterations): + action = actions[i % len(actions)] + gevent_execute(engine.execution_graph, udf_helpers, action, pool) + gevent_elapsed = time.perf_counter() - start + + gevent_throughput = iterations / gevent_elapsed + gevent_latency_ms = (gevent_elapsed / iterations) * 1000 + + click.echo('Gevent Executor:') + click.echo(f' Total time: {gevent_elapsed:.3f}s') + click.echo(f' Throughput: {gevent_throughput:.1f} actions/sec') + click.echo(f' Avg latency: {gevent_latency_ms:.3f}ms') + click.echo() + except ImportError: + click.echo('Gevent not available, skipping gevent benchmark') + click.echo() + gevent_throughput = None + + # --- Async executor --- + async def run_async(): + for i in range(warmup): + action = actions[i % len(actions)] + await async_execute(engine.execution_graph, udf_helpers, action, max_concurrent=max_concurrent) + + start = time.perf_counter() + for i in range(iterations): + action = actions[i % len(actions)] + await async_execute(engine.execution_graph, udf_helpers, action, max_concurrent=max_concurrent) + return time.perf_counter() - start + + async_elapsed = asyncio.run(run_async()) + async_throughput = iterations / async_elapsed + async_latency_ms = (async_elapsed / iterations) * 1000 + + click.echo('Async Executor:') + click.echo(f' Total time: {async_elapsed:.3f}s') + click.echo(f' Throughput: {async_throughput:.1f} actions/sec') + click.echo(f' Avg latency: {async_latency_ms:.3f}ms') + click.echo() + + # --- Comparison --- + if gevent_throughput: + ratio = async_throughput / gevent_throughput + click.echo('Comparison:') + click.echo(f' Async/Gevent ratio: {ratio:.2f}x') + if ratio > 1: + click.echo(f' Async is {((ratio - 1) * 100):.1f}% faster') + elif ratio < 1: + click.echo(f' Async is {((1 - ratio) * 100):.1f}% slower') + else: + click.echo(' Same performance') + + +if __name__ == '__main__': + cli() diff --git a/osprey_async_worker/src/osprey/async_worker/engine.py b/osprey_async_worker/src/osprey/async_worker/engine.py new file mode 100644 index 0000000..45c318f --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/engine.py @@ -0,0 +1,365 @@ +"""Async Osprey engine — no gevent dependency. + +Replaces OspreyEngine's gevent.pool.ThreadPool with stdlib +concurrent.futures.ThreadPoolExecutor for rule compilation. +Provides async execute() that calls the async executor directly. +""" + +import asyncio +import gc +import logging +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from time import time +from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Set, Type, TypedDict, cast + +if TYPE_CHECKING: + from osprey.worker.lib.data_exporters.validation_result_exporter import BaseValidationResultExporter + +from ddtrace.span import Span as TracerSpan +from osprey.async_worker.executor import execute as async_execute +from osprey.engine.ast.ast_utils import iter_nodes +from osprey.engine.ast.grammar import Assign, Span, parsed_ast_root_cache +from osprey.engine.ast.sources import SourcesConfig +from osprey.engine.ast_validator import validate_sources +from osprey.engine.ast_validator.validator_registry import ValidatorRegistry +from osprey.engine.ast_validator.validators.feature_name_to_entity_type_mapping import ( + FeatureNameToEntityTypeMapping, +) +from osprey.engine.ast_validator.validators.rule_name_to_description_mapping import ( + RuleNameToDescriptionMapping, +) +from osprey.engine.ast_validator.validators.unique_stored_names import ( + IdentifierIndex, + UniqueStoredNames, +) +from osprey.engine.ast_validator.validators.validate_static_types import ( + ValidateStaticTypes, +) +from osprey.engine.config.config_subkey_handler import ConfigSubkeyHandler, ModelT +from osprey.engine.executor.execution_context import Action, ExecutionResult +from osprey.engine.executor.execution_graph import ExecutionGraph, compile_execution_graph +from osprey.engine.executor.udf_execution_helpers import UDFHelpers +from osprey.engine.udf.registry import UDFRegistry +from osprey.engine.utils.periodic_execution_yielder import periodic_execution_yield +from osprey.worker.lib.instruments import metrics +from osprey.worker.lib.singletons import CONFIG +from osprey.worker.lib.sources_config import get_config_registry +from osprey.worker.lib.sources_provider_base import BaseSourcesProvider + +log = logging.getLogger(__name__) + +_DEFAULT_MAX_ASYNC_PER_EXECUTION = 12 + + +class FeatureLocation(TypedDict): + """Where a stored-name identifier is declared in the rule sources. + + Returned by :meth:`AsyncOspreyEngine.get_known_feature_locations` so + out-of-process consumers can render the rules-engine feature catalog + without re-deriving these positions. + """ + + name: str + source_path: str + source_line: int + source_snippet: str + + +def _extract_source_snippet(span: Span) -> str: + """Three-line snippet around `span`. Mirrors osprey_engine.extract_source_snippet.""" + src = span.source.contents + lines = src.splitlines() + start = max(span.start_line - 2, 0) + end = min(span.start_line + 1, len(lines)) + return '\n'.join(lines[start:end]) + + +class AsyncOspreyEngine: + """Async rules engine — no gevent dependency. + + Uses concurrent.futures.ThreadPoolExecutor for compilation (CPU-bound) + and the async executor for action execution. + """ + + def __init__( + self, + sources_provider: BaseSourcesProvider, + udf_registry: UDFRegistry, + should_yield_during_compilation: bool = False, + validation_exporter: Optional['BaseValidationResultExporter'] = None, + ): + self._sources_provider = sources_provider + self._udf_registry = udf_registry + self._should_yield_during_compilation = should_yield_during_compilation + config_registry = get_config_registry() + self._validator_registry = ValidatorRegistry.instance_with_additional_validators( + config_registry.get_validator() + ) + self._thread_pool = ThreadPoolExecutor(max_workers=1) + # Initial compile runs without periodic yields — there is no in-flight + # work yet to protect, and we want fast cold-start. + self._execution_graph = self._compile_execution_graph_sync(yield_during_compile=False) + # BaseSourcesProvider.set_sources_watcher is typed Callable[[], None], but the + # async provider (AsyncEtcdSourcesProvider) awaits an awaitable result, so a + # coroutine callback is valid at runtime. Cast to satisfy the narrower base type. + self._sources_provider.set_sources_watcher(cast(Callable[[], None], self._handle_updated_sources)) + self._config_subkey_handler = ConfigSubkeyHandler(config_registry, self._execution_graph.validated_sources) + self._validation_result_exporter = validation_exporter + # Freeze the boot graph out of gen-2 GC (see _freeze_resident_graph). + self._freeze_resident_graph() + + def _compile_execution_graph_sync(self, yield_during_compile: bool = True) -> ExecutionGraph: + """Compile the execution graph synchronously. + + When ``yield_during_compile`` is True (the default for recompiles), + wraps the work in ``periodic_execution_yield`` which causes the + compile thread to ``time.sleep`` periodically, releasing the GIL so + the asyncio event loop can keep servicing in-flight tasks. This + mirrors what the gevent engine does via the same context manager. + + Without this, the compile thread holds the GIL contiguously for ~7s, + and the asyncio main thread (running the event loop) gets starved + of CPU even though it's only ~5ms of bytecode releases away from + running. With CFS throttling on the asyncio worker pods, that + starvation compounds — compile pegs CPU, CFS throttles, and every + in-flight coroutine stalls until compile finishes. + + With yields: compile takes ~6× longer wall-clock (~42s) but uses + only ~16% of one core's CPU duty cycle, leaving plenty of headroom + for the rest of the worker. + """ + with periodic_execution_yield( + on=yield_during_compile and self._should_yield_during_compilation, + execution_time_ms=5, + yield_time_ms=25, + ): + sources = self._sources_provider.get_current_sources() + + start_time = time() + validated_sources = validate_sources( + sources, udf_registry=self._udf_registry, validator_registry=self._validator_registry + ) + validation_time = time() - start_time + + start_time = time() + execution_graph = compile_execution_graph(validated_sources) + compile_time = time() - start_time + + log.debug( + 'execution graph compiled: validation %.2fs, compilation %.2fs, total %.2fs', + validation_time, + compile_time, + validation_time + compile_time, + ) + return execution_graph + + async def compile_execution_graph(self) -> ExecutionGraph: + """Compile the execution graph in a thread pool (CPU-bound work).""" + loop = asyncio.get_running_loop() + return await loop.run_in_executor(self._thread_pool, self._compile_execution_graph_sync) + + async def _handle_updated_sources(self) -> None: + """Called by the sources provider when rules change in etcd. + + Runs the compile in self._thread_pool rather than on the event loop + so the loop stays free during compile and in-flight gRPC tasks can + drain and release their pinned response buffers. + """ + desired_hash = self._sources_provider.get_current_sources().hash() + try: + new_graph = await self.compile_execution_graph() + except Exception: + log.exception(f'Failed to compile execution graph for sources={desired_hash}') + metrics.increment( + 'osprey.rules_compile_failed', + tags=[f'desired:{desired_hash[:16]}'], + ) + return + + # Atomic swap. In-flight actions captured the old graph by reference at + # rules_sink.classify_one start and continue to use it until they finish; + # only newly-arriving actions read the new graph. Safe regardless of + # whether the input stream is paused. + # + # After the swap we call _freeze_resident_graph(): gc.collect() then gc.freeze(). + # The collect promotes survivors into gen 2, but the freeze immediately moves the + # resident graph into the permanent generation (excluded from automatic collection), + # so per-message gen-2 scans stay cheap as rules recompile. + # + # The old graph contains refcycles between AST roots and their children + # via ASTNode.parent back-pointers (see osprey/engine/ast/grammar.py), + # so plain refcount alone cannot reclaim it. Without intervention the + # old graph persists until gen-2 GC catches the cycle, and across many + # rule recompiles the resulting GC pressure raises per-message CPU. + # + # Fix: after the swap, walk the OLD graph's AST and null `parent` + # pointers — but ONLY on sources whose ast_root is not shared with the + # NEW graph. `parsed_ast_root_cache` in osprey/engine/ast/grammar.py + # memoizes ast_root by Source content, so unchanged source files share + # the same ast_root between graphs. Nulling parents on shared nodes + # would corrupt the new graph's AST. + old_graph = self._execution_graph + self._execution_graph = new_graph + + # Confirm to the provider which sources are now actually live so it dedups + # future no-op re-deliveries against what we APPLIED (not just received). + # The compile-failure path above returns early without marking, so a + # transient failure self-heals on the next etcd re-delivery. + self._sources_provider.mark_sources_applied(new_graph.validated_sources.sources.hash()) + + log.info(f'Compiled new execution graph for sources={self._sources_provider.get_current_sources().hash()}') + self._config_subkey_handler.dispatch_config(self._execution_graph.validated_sources) + + if self._validation_result_exporter is not None: + try: + self._validation_result_exporter.send(self._execution_graph.validated_sources) + except Exception: + log.exception('Failed to export validation results') + + self._break_old_graph_cycles(old_graph, new_graph) + self._freeze_resident_graph() + + @staticmethod + def _break_old_graph_cycles(old_graph: ExecutionGraph, new_graph: ExecutionGraph) -> None: + """Null `parent` back-pointers on every AST node in the discarded graph + so plain refcount can reclaim it without waiting for gen-2 GC. + + Skips any source whose ast_root is shared with the new graph — those + come from the module-level ``parsed_ast_root_cache`` and mutating them + would corrupt the new graph the engine just swapped in. + + Best-effort: any exception here is logged and swallowed. A leaked old + graph is wasteful but not incorrect. + """ + try: + new_root_ids = {id(s.ast_root) for s in new_graph.validated_sources.sources} + count = 0 + for source in old_graph.validated_sources.sources: + root = source.ast_root + if id(root) in new_root_ids: + continue + # Evict this discarded content from parsed_ast_root_cache BEFORE nulling + # its parents. That cache memoizes ast_root by source content and is never + # evicted, so a later graph that re-uses this exact content (e.g. a rule + # revert) would otherwise be handed back this same parent-nulled Root and + # fail validation with "`Rule(...)` must be assigned to a variable", + # wedging the worker on stale rules. Eviction forces a fresh re-parse on + # any future recurrence. + parsed_ast_root_cache.pop(source, None) + for node in iter_nodes(root): + node.parent = None + count += 1 + log.debug('broke parent pointers on %d AST nodes from old graph', count) + except Exception: + log.exception('failed to break cycles on old execution graph') + + def _freeze_resident_graph(self) -> None: + """Move the resident rule graph + other long-lived objects into the permanent + GC generation so per-message gen-2 collections stay cheap as rules recompile. + + osprey rebuilds the graph on each etcd update; without freezing, refcycle-held + AST objects (parsed_ast_root_cache) accumulate as permanent gen-2 survivors and + per-message GC CPU climbs with uptime. Called after the swap + cycle-break so the + freeze captures the new resident graph (and at boot for the initial graph). + Best-effort: any exception is logged and swallowed. + """ + try: + gc.collect() + gc.freeze() + except Exception: + log.exception('gc freeze of resident execution graph failed') + + @property + def execution_graph(self) -> ExecutionGraph: + return self._execution_graph + + @property + def config(self) -> SourcesConfig: + return self._execution_graph.validated_sources.sources.config + + async def execute( + self, + udf_helpers: UDFHelpers, + action: Action, + max_concurrent: Optional[int] = None, + sample_rate: int = 100, + parent_tracer_span: Optional[TracerSpan] = None, + ) -> ExecutionResult: + """Execute an action against the rules using the async executor.""" + if max_concurrent is None: + max_concurrent = CONFIG.instance().get_int( + 'OSPREY_MAX_ASYNC_PER_EXECUTION', _DEFAULT_MAX_ASYNC_PER_EXECUTION + ) + return await async_execute( + self._execution_graph, + udf_helpers, + action, + max_concurrent=max_concurrent, + sample_rate=sample_rate, + parent_tracer_span=parent_tracer_span, + ) + + def get_config_subkey(self, model_class: Type[ModelT]) -> ModelT: + return self._config_subkey_handler.get_config_subkey(model_class) + + def watch_config_subkey(self, model_class: Type[ModelT], update_callback: Callable[[ModelT], None]) -> None: + self._config_subkey_handler.watch_config_subkey(model_class, update_callback) + + def get_known_feature_locations(self) -> List[FeatureLocation]: + """Return locations of named identifiers that the rules engine extracts. + + Mirrors osprey.worker.lib.osprey_engine.OspreyEngine.get_known_feature_locations: + filters the UniqueStoredNames result by the parent Assign node's + should_extract flag so only identifiers the engine actually extracts make + it into the result. + + Returns FeatureLocation TypedDicts (JSON-shaped) — the gevent engine + used a dataclass; this surface stays a plain dict so consumers can + serialize without conversion while still getting precise types. + """ + + def _should_extract(span: Span) -> bool: + maybe_assign = span.parent_ast_node() + return bool(maybe_assign.should_extract if isinstance(maybe_assign, Assign) else True) + + identifier_index: IdentifierIndex = self._execution_graph.validated_sources.get_validator_result( + UniqueStoredNames + ) + return [ + FeatureLocation( + name=name, + source_path=span.source.path, + source_line=span.start_line, + source_snippet=_extract_source_snippet(span), + ) + for name, span in identifier_index.items() + if _should_extract(span) + ] + + def get_known_action_names(self) -> Set[str]: + return { + Path(source.path).stem for source in self.execution_graph.validated_sources.sources.glob('actions/*.sml') + } + + def get_rule_to_info_mapping(self) -> Dict[str, str]: + return self._execution_graph.validated_sources.get_validator_result(RuleNameToDescriptionMapping) + + def get_feature_name_to_entity_type_mapping(self) -> Dict[str, str]: + """Returns a mapping from 'feature name' -> 'entity type' for each feature that holds an entity.""" + return self._execution_graph.validated_sources.get_validator_result(FeatureNameToEntityTypeMapping) + + def get_post_execution_feature_name_to_value_type_mapping(self) -> Dict[str, type]: + """Returns a mapping from 'feature name' -> 'value type' for each feature.""" + post_execution_name_to_type_and_span = ValidateStaticTypes.to_post_execution_types( + self._execution_graph.validated_sources.get_validator_result(ValidateStaticTypes) + ) + return { + name: type_and_span.type + for name, type_and_span in post_execution_name_to_type_and_span.items() + if type_and_span.should_extract + } + + def shutdown(self) -> None: + """Shutdown the compilation thread pool.""" + self._thread_pool.shutdown(wait=True) diff --git a/osprey_async_worker/src/osprey/async_worker/executor.py b/osprey_async_worker/src/osprey/async_worker/executor.py new file mode 100644 index 0000000..bb7d1be --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/executor.py @@ -0,0 +1,525 @@ +"""Async executor for Osprey rules engine. + +Two execution paths: +- Sync UDFs (execute_async=False): pure computation, run inline on the event loop +- Async UDFs (AsyncUDFBase): native async I/O, awaited as tasks with semaphore + +No thread pool. No run_in_executor. Existing sync UDFs that use gevent +primitives will NOT work here — they must be ported to AsyncUDFBase. +""" + +import asyncio +import os +from collections import defaultdict +from typing import Any, Dict, List, Optional, Sequence, Set, Tuple + +import sentry_sdk +from ddtrace import tracer +from ddtrace.span import Span as TracerSpan +from osprey.async_worker.adaptor.interfaces import AsyncBatchableUDFBase, AsyncUDFBase +from osprey.async_worker.lib.pigeon.exceptions import RPCException +from osprey.engine.ast.grammar import ASTNode +from osprey.engine.executor.custom_extracted_features import ( + ActionIdExtractedFeature, + ErrorCountExtractedFeature, + SampleRateExtractedFeature, + TimestampExtractedFeature, +) +from osprey.engine.executor.dependency_chain import DependencyChain +from osprey.engine.executor.execution_context import ( + Action, + ExecutionContext, + ExecutionResult, + ExpectedUdfException, + NodeErrorInfo, + NodeFailurePropagationException, + NodeResult, +) +from osprey.engine.executor.execution_graph import ExecutionGraph +from osprey.engine.executor.node_executor.call_executor import CallExecutor +from osprey.engine.executor.udf_execution_helpers import UDFHelpers +from osprey.engine.stdlib.udfs.json_utils import MissingJsonPath +from osprey.engine.udf.base import BatchableUDFBase +from osprey.worker.lib.instruments import metrics +from osprey.worker.lib.osprey_shared.logging import get_logger +from result import Err, Ok + +logger = get_logger(__name__) + +_DEFAULT_MAX_ASYNC_PER_EXECUTION = 12 + + +def _get_ready_sync_and_async( + allow_async: bool, context: ExecutionContext +) -> Tuple[Sequence[DependencyChain], Sequence[DependencyChain]]: + _ready_sync = [] + _ready_async = [] + for ready_chain in context.get_ready_to_execute(): + if ready_chain.executor.execute_async and allow_async: + _ready_async.append(ready_chain) + else: + _ready_sync.append(ready_chain) + return _ready_sync, _ready_async + + +_SUPPRESS_IN_PROD = (ExpectedUdfException, NodeFailurePropagationException, MissingJsonPath, TypeError) + + +def _is_spammy_exception(e: Optional[Exception]) -> bool: + if e is None: + return True + if os.environ.get('ENVIRONMENT') in ('staging', 'development'): + return False + return isinstance(e, _SUPPRESS_IN_PROD) + + +def _get_metric_tags( + context: ExecutionContext, batchable_udf: Optional[BatchableUDFBase[Any, Any, Any]] = None +) -> List[str]: + return [ + f'action:{context.get_action_name()}', + f'encoding:{context.get_data_encoding()}', + f'batch_type:{batchable_udf.get_batchable_arguments_type().__name__}' + if batchable_udf is not None + else 'batch_type:none', + 'host:none', + 'kube_node:none', + 'instance-id:none', + 'internal-hostname:none', + 'name:none', + ] + + +def _record_udf_metric( + metric_tags: List[str], + execution_result: NodeResult, + caught_exception: Optional[Exception], +) -> None: + """Emit UDF execution metrics. Shared by sync and async paths.""" + if execution_result.is_ok(): + metrics.increment('udf_execution', tags=metric_tags + ['exc_name:none', 'result:success']) + elif not _is_spammy_exception(caught_exception): + exc_name = caught_exception.__class__.__name__ + if isinstance(caught_exception, RPCException): + exc_name = exc_name + f'.{caught_exception.code().name.lower()}' + metrics.increment( + 'udf_execution', + tags=metric_tags + [f'exc_name:{exc_name}', 'result:unexpected_failure'], + ) + sentry_sdk.capture_exception(caught_exception) + + +# --- Sync execution (inline, for pure-computation UDFs) --- + + +def _execute_sync( + chain: DependencyChain, + context: ExecutionContext, + error_info_: List[NodeErrorInfo], +) -> NodeResult: + """Execute a sync UDF inline. For pure computation only — no I/O.""" + execution_result: NodeResult = Err(None) + try: + execution_result = Ok(chain.executor.execute(execution_context=context)) + except Exception as e: + if not isinstance(e, NodeFailurePropagationException): + error_info_.append(NodeErrorInfo(e, chain.executor.node)) + execution_result = Err(None) + return execution_result + + +# --- Legacy sync UDF fallback (thread pool, for execute_async=True UDFs not yet ported) --- + + +def _execute_legacy_sync( + chain: DependencyChain, + context: ExecutionContext, + error_info_: List[NodeErrorInfo], +) -> NodeResult: + """Execute a legacy sync UDF that has execute_async=True but is not AsyncUDFBase. + + Runs in a thread pool via run_in_executor. May fail if the UDF uses gevent + primitives (no monkey patching), but errors are captured gracefully. + """ + caught_exception: Optional[Exception] = None + metric_tags = _get_metric_tags(context) + if isinstance(chain.executor, CallExecutor): + call_node: CallExecutor = chain.executor + metric_tags += [f'udf:{call_node._udf.__class__.__name__}'] + + execution_result: NodeResult = Err(None) + try: + with metrics.timed('udf_execution_duration', tags=metric_tags, sample_rate=0.01): + execution_result = Ok(chain.executor.execute(execution_context=context)) + except Exception as e: + if not isinstance(e, NodeFailurePropagationException): + error_info_.append(NodeErrorInfo(e, chain.executor.node)) + execution_result = Err(None) + caught_exception = e + finally: + if isinstance(chain.executor, CallExecutor): + _record_udf_metric(metric_tags, execution_result, caught_exception) + return execution_result + + +async def _execute_legacy_in_executor( + loop: asyncio.AbstractEventLoop, + semaphore: asyncio.Semaphore, + chain: DependencyChain, + context: ExecutionContext, + error_info_: List[NodeErrorInfo], +) -> NodeResult: + """Run a legacy sync UDF in the thread pool with semaphore.""" + async with semaphore: + return await loop.run_in_executor(None, _execute_legacy_sync, chain, context, error_info_) + + +def _execute_legacy_batch_sync( + udfs: Sequence[BatchableUDFBase[Any, Any, Any]], + nodes: Sequence[ASTNode], + batchable_args: Sequence[Any], + context: ExecutionContext, + error_info_: List[NodeErrorInfo], +) -> Sequence[NodeResult]: + """Execute a batch of legacy sync batchable UDFs in thread pool.""" + assert len(udfs) == len(nodes) == len(batchable_args) + num_executions = len(udfs) + metric_tags = _get_metric_tags(context, udfs[0]) + + try: + with metrics.timed('udf_execution_batch_duration', tags=metric_tags, sample_rate=0.01): + results = udfs[0].execute_batch(context, udfs, batchable_args) + assert len(results) == num_executions + except Exception as e: + if not isinstance(e, NodeFailurePropagationException): + for n in nodes: + error_info_.append(NodeErrorInfo(e, n)) + if not _is_spammy_exception(e): + metrics.increment( + 'udf_execution_batch', + tags=metric_tags + [f'exc_name:{e.__class__.__name__}', 'result:unexpected_failure'], + ) + return [Err(None)] * num_executions + + type_checked_results = [] + for udf, node, result in zip(udfs, nodes, results): + if result.is_err(): + if not isinstance(result.value, NodeFailurePropagationException): + error_info_.append(NodeErrorInfo(result.value, node)) + type_checked_results.append(Err(None)) + continue + try: + type_checked_results.append(Ok(udf.check_result_type(result.value))) + except Exception as e: + if not isinstance(e, NodeFailurePropagationException): + error_info_.append(NodeErrorInfo(e, node)) + type_checked_results.append(Err(None)) + + return type_checked_results + + +# --- Async execution (native async UDFs, awaited on event loop) --- + + +async def _execute_async_udf( + semaphore: asyncio.Semaphore, + chain: DependencyChain, + context: ExecutionContext, + error_info_: List[NodeErrorInfo], +) -> NodeResult: + """Execute a native async UDF. Awaited directly on the event loop.""" + async with semaphore: + call_executor: CallExecutor = chain.executor # type: ignore + udf: AsyncUDFBase[Any, Any] = call_executor._udf # type: ignore + metric_tags = _get_metric_tags(context) + [f'udf:{udf.__class__.__name__}'] + + caught_exception: Optional[Exception] = None + execution_result: NodeResult = Err(None) + try: + resolved_arguments = udf.resolve_arguments(context, call_executor) + with metrics.timed('udf_execution_duration', tags=metric_tags, sample_rate=0.01): + result = await udf.async_execute(context, resolved_arguments) + execution_result = Ok(udf.check_result_type(result)) + except Exception as e: + if not isinstance(e, NodeFailurePropagationException): + error_info_.append(NodeErrorInfo(e, call_executor.node)) + execution_result = Err(None) + caught_exception = e + finally: + _record_udf_metric(metric_tags, execution_result, caught_exception) + return execution_result + + +async def _execute_async_batch( + semaphore: asyncio.Semaphore, + udfs: Sequence[AsyncBatchableUDFBase[Any, Any, Any]], + nodes: Sequence[ASTNode], + batchable_args: Sequence[Any], + context: ExecutionContext, + error_info_: List[NodeErrorInfo], +) -> Sequence[NodeResult]: + """Execute a batch of native async batchable UDFs.""" + async with semaphore: + assert len(udfs) == len(nodes) == len(batchable_args) + num_executions = len(udfs) + metric_tags = _get_metric_tags(context, udfs[0]) + + try: + with metrics.timed('udf_execution_batch_duration', tags=metric_tags, sample_rate=0.01): + results = await udfs[0].async_execute_batch(context, udfs, batchable_args) + assert len(results) == num_executions + except Exception as e: + if not isinstance(e, NodeFailurePropagationException): + for n in nodes: + error_info_.append(NodeErrorInfo(e, n)) + if not _is_spammy_exception(e): + metrics.increment( + 'udf_execution_batch', + tags=metric_tags + [f'exc_name:{e.__class__.__name__}', 'result:unexpected_failure'], + ) + return [Err(None)] * num_executions + + type_checked_results = [] + for udf, node, result in zip(udfs, nodes, results): + if result.is_err(): + if not isinstance(result.value, NodeFailurePropagationException): + error_info_.append(NodeErrorInfo(result.value, node)) + if not _is_spammy_exception(result.value): + exc_name = result.value.__class__.__name__ + if isinstance(result.value, RPCException): + exc_name = exc_name + f'.{result.value.code().name.lower()}' + metrics.increment( + 'udf_execution', + tags=metric_tags + + [f'udf:{udf.__class__.__name__}', f'exc_name:{exc_name}', 'result:unexpected_failure'], + ) + type_checked_results.append(Err(None)) + continue + try: + type_checked_results.append(Ok(udf.check_result_type(result.value))) + metrics.increment( + 'udf_execution', + tags=metric_tags + [f'udf:{udf.__class__.__name__}', 'exc_name:none', 'result:success'], + ) + except Exception as e: + if not isinstance(e, NodeFailurePropagationException): + error_info_.append(NodeErrorInfo(e, node)) + type_checked_results.append(Err(None)) + + metrics.increment('udf_execution_batch', tags=metric_tags + ['exc_name:none', 'result:success']) + return type_checked_results + + +# --- Batching logic --- + + +async def _enqueue_batches( + loop: asyncio.AbstractEventLoop, + semaphore: asyncio.Semaphore, + context: ExecutionContext, + error_infos: List[NodeErrorInfo], + ready_async: Sequence[DependencyChain], +) -> Tuple[Sequence[DependencyChain], Dict[asyncio.Task[Sequence[NodeResult]], Sequence[DependencyChain]]]: + """Collect batchable async chains and launch them as tasks. + + Returns (remaining non-batched chains, dict of batch tasks -> chains). + """ + batch_chains: Dict[Tuple[type, str], List[Tuple[DependencyChain, Any]]] = defaultdict(list) + chains_to_remove: List[DependencyChain] = [] + + for async_chain in ready_async: + if not isinstance(async_chain.executor, CallExecutor): + continue + call_executor: CallExecutor = async_chain.executor + if not isinstance(call_executor._udf, (AsyncBatchableUDFBase, BatchableUDFBase)): + continue + udf = call_executor._udf + + batch_type = udf.get_batchable_arguments_type() + try: + resolved_arguments = udf.resolve_arguments(context, call_executor) + batchable_arguments = udf.get_batchable_arguments(resolved_arguments) + routing_key = udf.get_batch_routing_key(batchable_arguments) + batch_chains[(batch_type, routing_key)].append((async_chain, batchable_arguments)) + except Exception as e: + if not isinstance(e, NodeFailurePropagationException): + error_infos.append(NodeErrorInfo(e, call_executor.node)) + chains_to_remove.append(async_chain) + context.set_resolved_value(async_chain, Err(None)) + + new_batch_tasks: Dict[asyncio.Task[Sequence[NodeResult]], Sequence[DependencyChain]] = {} + + for _, chains_and_args in batch_chains.items(): + if len(chains_and_args) < 2: + continue + + chains, args = zip(*chains_and_args) + chains_to_remove.extend(chains) + + batch_udfs = [chain.executor._udf for chain in chains] + batch_nodes = [chain.executor.node for chain in chains] + + if isinstance(batch_udfs[0], AsyncBatchableUDFBase): + task = asyncio.create_task( + _execute_async_batch(semaphore, batch_udfs, batch_nodes, args, context, error_infos) + ) + else: + # Legacy sync batchable UDF — run in thread pool + async def _run_legacy_batch(s, u, n, a, c, e): + async with s: + return await loop.run_in_executor(None, _execute_legacy_batch_sync, u, n, a, c, e) + + task = asyncio.create_task( + _run_legacy_batch(semaphore, batch_udfs, batch_nodes, args, context, error_infos) + ) + new_batch_tasks[task] = chains + + remaining = [chain for chain in ready_async if chain not in chains_to_remove] + return remaining, new_batch_tasks + + +# --- Main executor --- + + +async def execute( + execution_graph: ExecutionGraph, + udf_helpers: UDFHelpers, + action: Action, + max_concurrent: int = _DEFAULT_MAX_ASYNC_PER_EXECUTION, + sample_rate: int = 100, + parent_tracer_span: Optional[TracerSpan] = None, +) -> ExecutionResult: + """Async executor for the osprey rules engine. + + Three paths: + - Sync UDFs (execute_async=False): run inline, pure computation only + - AsyncUDFBase: awaited as tasks directly on event loop (native async) + - Legacy UDFBase with execute_async=True: run in thread pool via run_in_executor + (may fail on gevent calls, errors captured gracefully) + """ + if parent_tracer_span: + parent_tracer_span.set_tag('action-name', action.action_name) + + context = ExecutionContext(execution_graph=execution_graph, helpers=udf_helpers, action=action) + allow_async = max_concurrent > 0 + semaphore = asyncio.Semaphore(max_concurrent) + loop = asyncio.get_running_loop() + error_infos: List[NodeErrorInfo] = [] + + in_progress_singlets: Dict[asyncio.Task[NodeResult], DependencyChain] = {} + in_progress_batches: Dict[asyncio.Task[Sequence[NodeResult]], Sequence[DependencyChain]] = {} + + ready_sync, ready_async = _get_ready_sync_and_async(allow_async, context) + + while ready_sync or ready_async or in_progress_singlets or in_progress_batches: + # Check for already-finished tasks (non-blocking) + finished_singlets = [t for t in in_progress_singlets if t.done()] + finished_batches = [t for t in in_progress_batches if t.done()] + + if not ready_sync and not ready_async and not finished_singlets and not finished_batches: + # Block until at least one async task finishes + all_pending: Set[asyncio.Task[Any]] = set(in_progress_singlets.keys()) | set(in_progress_batches.keys()) + if all_pending: + with tracer.start_span('osprey.rules.async_wait_nodes', child_of=parent_tracer_span): + done, _ = await asyncio.wait(all_pending, return_when=asyncio.FIRST_COMPLETED) + finished_singlets = [t for t in done if t in in_progress_singlets] + finished_batches = [t for t in done if t in in_progress_batches] + + # Process finished singlets + for task in finished_singlets: + chain = in_progress_singlets.pop(task) + context.set_resolved_value(chain, task.result()) + + # Process finished batches + # Distinct loop variable from the singlet `task` above: the two dicts hold + # tasks with different result types (Task[NodeResult] vs Task[Sequence[NodeResult]]), + # and reusing one name would unify them to the singlet type. + for batch_task in finished_batches: + chains = in_progress_batches.pop(batch_task) + results = batch_task.result() + for i, chain in enumerate(chains): + context.set_resolved_value(chain, results[i]) + + # Enqueue async tasks + if allow_async and ready_async: + with tracer.start_span('osprey.rules.try_enqueue_batches', child_of=parent_tracer_span): + remaining_ready_async, new_batch_tasks = await _enqueue_batches( + loop, semaphore, context, error_infos, ready_async + ) + in_progress_batches.update(new_batch_tasks) + + for async_chain in remaining_ready_async: + # Native async UDF → await on event loop + if isinstance(async_chain.executor, CallExecutor) and isinstance( + async_chain.executor._udf, (AsyncUDFBase, AsyncBatchableUDFBase) + ): + task = asyncio.create_task(_execute_async_udf(semaphore, async_chain, context, error_infos)) + else: + # Legacy sync UDF with execute_async=True → thread pool fallback + task = asyncio.create_task( + _execute_legacy_in_executor(loop, semaphore, async_chain, context, error_infos) + ) + in_progress_singlets[task] = async_chain + + # Execute sync chains inline (pure computation, fast, no I/O). + # Only yield deep into a long sync round when async tasks are in flight. + # Each sleep(0) triggers a full event loop cycle including gRPC C-core polling, + # so unnecessary yields cause significant context-switch overhead. + # Short rounds (<100 chains) skip yielding entirely — the asyncio.wait() at the + # top of the loop provides natural yield points between rounds. + for i, sync_chain in enumerate(ready_sync): + if (in_progress_singlets or in_progress_batches) and i > 0 and i % 100 == 0: + await asyncio.sleep(0) + result = _execute_sync(sync_chain, context, error_infos) + context.set_resolved_value(sync_chain, result) + + ready_sync, ready_async = _get_ready_sync_and_async(allow_async, context) + + # --- Build result --- + + unexpected_error_infos = [ + error_info for error_info in error_infos if not isinstance(error_info.error, ExpectedUdfException) + ] + validator_results = execution_graph.validated_sources.validation_results + + context.add_custom_extracted_features( + [ + ActionIdExtractedFeature(action_id=action.action_id), + TimestampExtractedFeature(timestamp=action.timestamp), + ErrorCountExtractedFeature(error_count=len(unexpected_error_infos)), + SampleRateExtractedFeature(sample_rate=sample_rate), + ] + ) + + effects = context.get_effects() + + actionable_error_infos = [ + error_info + for error_info in error_infos + if isinstance(error_info.error, Exception) and not _is_spammy_exception(error_info.error) + ] + has_effects = len(effects) > 0 + has_actionable_errors = len(actionable_error_infos) > 0 + action_tags = [ + f'action:{action.action_name}', + f'had_actionable_errors:{has_actionable_errors}', + f'had_effects:{has_effects}', + ] + metrics.increment('osprey.action_health', tags=action_tags) + if has_actionable_errors: + metrics.histogram( + 'osprey.action_error_count', len(actionable_error_infos), tags=[f'action:{action.action_name}'] + ) + + trace_id = str(parent_tracer_span.trace_id) if parent_tracer_span else None + + return ExecutionResult( + extracted_features=context.get_extracted_features(), + action=action, + effects=effects, + validator_results=validator_results, + error_infos=unexpected_error_infos, + sample_rate=sample_rate, + trace_id=trace_id, + rule_audit_entries=context.get_rule_audit_entries(), + ) diff --git a/osprey_async_worker/src/osprey/async_worker/lib/__init__.py b/osprey_async_worker/src/osprey/async_worker/lib/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/lib/coordinator_input_stream.py b/osprey_async_worker/src/osprey/async_worker/lib/coordinator_input_stream.py new file mode 100644 index 0000000..5061486 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/coordinator_input_stream.py @@ -0,0 +1,523 @@ +import asyncio +import json +import random +import time +from typing import TYPE_CHECKING, Any, AsyncIterator, Callable, Dict, Optional, Tuple, Union + +import grpc +import grpc.aio +import pytz +import sentry_sdk +from osprey.async_worker.lib.discovery.async_directory import AsyncServiceWatcher +from osprey.async_worker.lib.etcd.sources_provider import AsyncInputStreamReadySignaler +from osprey.async_worker.sinks.sink.input_stream import AsyncBaseInputStream +from osprey.engine.executor.execution_context import Action as OspreyEngineAction +from osprey.engine.executor.execution_context import ExecutionResult +from osprey.rpc.common.v1.verdicts_pb2 import Verdicts +from osprey.rpc.osprey_coordinator.bidirectional_stream.v1.service_pb2 import ( + Ack, + AckOrNack, + ActionRequest, + ClientDetails, + Disconnect, + Nack, + OspreyCoordinatorAction, + Request, +) +from osprey.rpc.osprey_coordinator.bidirectional_stream.v1.service_pb2_grpc import ( + OspreyCoordinatorServiceStub, +) +from osprey.worker.lib.discovery.service import Service +from osprey.worker.lib.instruments import metrics +from osprey.worker.lib.osprey_shared.logging import get_logger, info_log_osprey_action +from osprey.worker.sinks.utils.acking_contexts_base import BaseAckingContext, VerdictsAckingContext + +if TYPE_CHECKING: + # Type-only: the sync ServiceWatcher module imports gevent, which the async + # worker must not import at runtime. The sync watcher is only constructed in + # the gevent-discovery __init__ path (which lazily imports Directory). + from osprey.worker.lib.discovery.service_watcher import ServiceWatcher + +logger = get_logger() + +MIN_SECONDS_BEFORE_RECONNECT = 60 +SECONDS_BEFORE_RECONNECT_JITTER = 60 + + +class AsyncVerdictsAckingContext(VerdictsAckingContext[OspreyEngineAction]): + """Async-compatible verdicts acking context that sends ack/nack back through the bidirectional stream. + + Holds a reference to the stream and ack_id so the rules sink can use it as a normal + context manager, and the ack is sent when the context exits. + """ + + def __init__( + self, + item: OspreyEngineAction, + stream: 'OspreyCoordinatorBiDirectionalStream', + ack_id: int, + ) -> None: + super().__init__(item) + self._stream = stream + self._ack_id = ack_id + + +class GrpcConnectionDiscoveryPool: + """Maintains a pool of async gRPC channels discovered via etcd service discovery.""" + + # The watcher is a gevent ServiceWatcher (__init__), an AsyncServiceWatcher + # (from_async_discovery), or None (from_static — no discovery). The async-only + # methods (ensure_initialized) are only reached when _needs_async_init is True, + # which is only set by from_async_discovery (AsyncServiceWatcher). + _service_watcher: Union['ServiceWatcher', AsyncServiceWatcher, None] + _handle_service_change_fn: Optional[Callable[[str, Service], None]] + + def __init__(self, service_name: str) -> None: + # Lazy import: Directory uses gevent-based service discovery. The async worker + # bypasses this constructor entirely via OspreyCoordinatorInputStream.from_direct_address(). + from osprey.worker.lib.discovery.directory import Directory + + self._service_name = service_name + self._needs_async_init = False + directory = Directory.instance(secure=False) + + self._grpc_channels: Dict[Service, Tuple[grpc.aio.Channel, Service]] = { + service: (self._create_async_channel(service), service) + for service in directory.select_all(self._service_name) + } + + self._service_watcher = directory.get_watcher(self._service_name) + self._handle_service_change_fn = self._handle_service_change + self._service_watcher.add_lazy_listener(self._handle_service_change_fn) + + @classmethod + def from_static(cls, address: str, service_name: str) -> 'GrpcConnectionDiscoveryPool': + """Create a pool with a single static address (no etcd discovery).""" + host, port_str = address.rsplit(':', 1) + service = Service( + name=service_name, + address=host, + port=int(port_str), + ports={'grpc': int(port_str)}, + metadata={}, + ) + instance = object.__new__(cls) + instance._service_name = service_name + instance._needs_async_init = False + instance._grpc_channels = {service: (grpc.aio.insecure_channel(address), service)} + instance._service_watcher = None + instance._handle_service_change_fn = None + return instance + + @classmethod + def from_async_discovery(cls, service_name: str) -> 'GrpcConnectionDiscoveryPool': + """Create a pool using async etcd service discovery (no gevent dependency). + + The initial etcd load is deferred to the first ``get_connection()`` call + (which is async) since the etcd watcher requires ``await ensure_initialized()``. + After initialization, the watcher keeps the pool updated as pods scale up/down. + """ + from osprey.async_worker.lib.discovery.async_directory import AsyncDirectory + + instance = object.__new__(cls) + instance._service_name = service_name + instance._needs_async_init = True + instance._grpc_channels = {} + + directory = AsyncDirectory.instance(secure=False) + instance._service_watcher = directory.get_watcher(service_name) + instance._handle_service_change_fn = instance._handle_service_change + instance._service_watcher.add_lazy_listener(instance._handle_service_change_fn) + + return instance + + async def _initialize_from_etcd(self) -> None: + """Load coordinator instances from etcd on first async call.""" + self._needs_async_init = False + # Only the async-discovery constructor sets _needs_async_init, and it always + # installs an AsyncServiceWatcher. Narrow for the async-only methods below. + assert isinstance(self._service_watcher, AsyncServiceWatcher) + try: + await self._service_watcher.ensure_initialized() + for service in self._service_watcher.select_all(): + if service not in self._grpc_channels: + self._grpc_channels[service] = (self._create_async_channel(service), service) + logger.info( + 'async discovery initialized for %s: %d instances', + self._service_name, + len(self._grpc_channels), + ) + except Exception: + self._needs_async_init = True + logger.exception('failed to initialize async discovery for %s', self._service_name) + + @staticmethod + def _create_async_channel(service: Service) -> grpc.aio.Channel: + return grpc.aio.insecure_channel(target=f'{service.connection_address}:{service.grpc_port}') + + def _handle_service_change(self, service_state: str, service: Service) -> None: + if service_state == 'up': + if not self._grpc_channels.get(service): + self._grpc_channels[service] = (self._create_async_channel(service), service) + elif service_state == 'down': + if self._grpc_channels.get(service): + del self._grpc_channels[service] + + async def get_connection(self) -> Tuple[grpc.aio.Channel, Service]: + """Gets an async gRPC channel to a coordinator instance. + + If no services are registered, polls with exponential backoff until one becomes available. + """ + if self._needs_async_init: + await self._initialize_from_etcd() + + channels = list(self._grpc_channels.values()) + backoff = 1.0 + + while len(channels) == 0: + logger.info('all %s instances offline... retrying in %.1fs', self._service_name, backoff) + await asyncio.sleep(backoff) + backoff = min(backoff * 2, 30.0) + channels = list(self._grpc_channels.values()) + + return random.choice(channels) + + async def close(self) -> None: + """Close all gRPC channels.""" + for channel, _ in list(self._grpc_channels.values()): + try: + await channel.close() + except Exception: + pass + self._grpc_channels.clear() + + +class OspreyCoordinatorBiDirectionalStream: + """Manages a single bidirectional gRPC stream with the osprey coordinator. + + Outgoing requests (initial ClientDetails, then ack/nack) are fed through an + asyncio.Queue and yielded as an async iterator to the gRPC call. Incoming + OspreyCoordinatorAction messages are exposed via ``async for``. + """ + + HEARTBEAT_INTERVAL_SECONDS = 30 + + def __init__(self, client_id: str, channel: grpc.aio.Channel, service: Service) -> None: + self._client_id = client_id + self._outgoing_queue: asyncio.Queue[Optional[Request]] = asyncio.Queue() + self._stub = OspreyCoordinatorServiceStub(channel=channel) + self._tags = [f'coordinator_connection_address:{service.connection_address}'] + self._connect_time: Optional[float] = None + self._last_action_request_time: float = 0.0 + self._stopped = False + + # -- outgoing request helpers -------------------------------------------------- + + async def _outgoing_iterator(self) -> AsyncIterator[Request]: + """Async generator that drains the outgoing queue for the gRPC call.""" + while True: + request = await self._outgoing_queue.get() + if request is None: + # Sentinel value — stop the outgoing side of the stream + return + yield request + + async def _send(self, request: Request) -> None: + await self._outgoing_queue.put(request) + + async def _enqueue_stop_signal(self) -> None: + await self._outgoing_queue.put(None) + + async def send_graceful_disconnect( + self, ack_id: int, ack: bool = True, verdicts: Optional[Verdicts] = None + ) -> None: + ack_or_nack = ( + AckOrNack(ack_id=ack_id, ack=Ack(verdicts=verdicts if verdicts else None)) + if ack + else AckOrNack(ack_id=ack_id, nack=Nack()) + ) + metrics.increment('ack_or_nack.disconnect', tags=[f'ack:{ack}', f'verdicts:{verdicts is not None}']) + logger.debug('submitting acking disconnect') + await self._outgoing_queue.put(Request(disconnect=Disconnect(ack_or_nack=ack_or_nack))) + await self._enqueue_stop_signal() + + def send_ack_or_nack(self, ack_id: int, ack: bool = True, verdicts: Optional[Verdicts] = None) -> None: + """Fire-and-forget ack — uses put_nowait to match gevent's non-blocking Queue.put().""" + ack_or_nack = ( + AckOrNack(ack_id=ack_id, ack=Ack(verdicts=verdicts if verdicts else None)) + if ack + else AckOrNack(ack_id=ack_id, nack=Nack()) + ) + req = Request(action_request=ActionRequest(ack_or_nack=ack_or_nack)) + metrics.increment('ack_or_nack', tags=[f'ack:{ack}', f'verdicts:{verdicts is not None}']) + logger.debug('submitting acking action request') + self._last_action_request_time = time.time() + self._outgoing_queue.put_nowait(req) + + def get_uptime(self) -> float: + assert self._connect_time is not None, 'This was called before a connection was established' + return time.time() - self._connect_time + + # -- incoming action iteration ------------------------------------------------- + + async def __aiter__(self) -> AsyncIterator[OspreyCoordinatorAction]: + async for action in self._gen(): + yield action + + async def _gen(self) -> AsyncIterator[OspreyCoordinatorAction]: + logger.info( + 'bidi stream connecting to coordinator %s (client_id=%s)', + self._tags, + self._client_id, + ) + await self._send(Request(action_request=ActionRequest(initial=ClientDetails(id=self._client_id)))) + self._last_action_request_time = time.time() + self._connect_time = time.time() + metrics.increment('osprey_coordinator_input_stream.connect', tags=self._tags) + + try: + incoming_stream = self._stub.OspreyBidirectionalStream(self._outgoing_iterator(), timeout=None) + async for osprey_coordinator_action in incoming_stream: + elapsed_time_since_last_action_request = time.time() - self._last_action_request_time + metrics.histogram( + 'osprey_coordinator_input_stream.elapsed_time_since_action_request', + elapsed_time_since_last_action_request, + tags=self._tags, + ) + yield osprey_coordinator_action + except grpc.aio.AioRpcError as e: + if e.code() != grpc.StatusCode.CANCELLED: + logger.exception(e) + sentry_sdk.capture_exception() + logger.error('Received failure from stream...closing stream') + metrics.increment( + 'osprey_coordinator_input_stream.stream_error', + tags=self._tags + [f'rpc_error_code:{e.code().name.lower()}'], + ) + + +class OspreyCoordinatorInputStream(AsyncBaseInputStream[BaseAckingContext[OspreyEngineAction]]): + """Async input stream for the coordinator bidirectional gRPC transport. + + Wraps ``OspreyCoordinatorBiDirectionalStream`` and handles: + * Reconnecting on a jittered interval + * Deserializing OspreyCoordinatorAction -> OspreyEngineAction + * Graceful shutdown via ``asyncio.Event`` + * Acking / nacking actions back through the stream + """ + + def __init__( + self, + client_id: str, + coordinator_service_name: str = 'osprey_coordinator', + input_stream_ready_signaler: Optional[AsyncInputStreamReadySignaler] = None, + ) -> None: + self._client_id = client_id + self._channel_pool = GrpcConnectionDiscoveryPool(coordinator_service_name) + self._shutdown_event = asyncio.Event() + self._current_execution_result: Optional[ExecutionResult] = None + self._input_stream_ready_signaler = input_stream_ready_signaler + + @classmethod + def from_direct_address( + cls, + client_id: str, + address: str, + service_name: str = 'osprey_coordinator', + input_stream_ready_signaler: Optional[AsyncInputStreamReadySignaler] = None, + ) -> 'OspreyCoordinatorInputStream': + """Create an input stream connected directly to a coordinator address. + + Bypasses etcd service discovery. Each instance gets its own gRPC channel. + """ + instance = object.__new__(cls) + instance._client_id = client_id + instance._shutdown_event = asyncio.Event() + instance._current_execution_result = None + instance._channel_pool = GrpcConnectionDiscoveryPool.from_static(address, service_name) + instance._input_stream_ready_signaler = input_stream_ready_signaler + return instance + + @classmethod + def from_async_discovery( + cls, + client_id: str, + service_name: str = 'osprey_coordinator', + input_stream_ready_signaler: Optional[AsyncInputStreamReadySignaler] = None, + ) -> 'OspreyCoordinatorInputStream': + """Create an input stream using async etcd discovery (no gevent). + + Discovers all coordinator instances from etcd. Each ``get_connection()`` call + randomly selects a coordinator, distributing streams across all pods. + """ + instance = object.__new__(cls) + instance._client_id = client_id + instance._shutdown_event = asyncio.Event() + instance._current_execution_result = None + instance._channel_pool = GrpcConnectionDiscoveryPool.from_async_discovery(service_name) + instance._input_stream_ready_signaler = input_stream_ready_signaler + return instance + + async def stop(self) -> None: + logger.info('Received shutdown signal... safely shutting down') + self._shutdown_event.set() + await self._channel_pool.close() + + # -- deserialization ----------------------------------------------------------- + + def _create_osprey_engine_action( + self, osprey_coordinator_action: OspreyCoordinatorAction + ) -> Optional[OspreyEngineAction]: + try: + tags = [f'action_name:{osprey_coordinator_action.action_name}'] + + secret_data: Dict[str, Any] = {} + which_of_action_data = osprey_coordinator_action.WhichOneof('action_data') + if which_of_action_data == 'json_action_data': + info_log_osprey_action( + osprey_coordinator_action.action_id, + osprey_coordinator_action.action_name, + 'received json-encoded action', + ) + with metrics.timed( + 'osprey_coordinator_input_stream.deserialize_message', + tags=tags + ['serialization_type:json'], + use_ms=True, + ): + data = json.loads(osprey_coordinator_action.json_action_data) + if osprey_coordinator_action.HasField('json_secret_data'): + secret_data = json.loads(osprey_coordinator_action.json_secret_data) + encoding = 'json' + + elif which_of_action_data == 'proto_action_data': + info_log_osprey_action( + osprey_coordinator_action.action_id, + osprey_coordinator_action.action_name, + 'received proto-encoded action', + ) + with metrics.timed( + 'osprey_coordinator_input_stream.deserialize_message', + tags=tags + ['serialization_type:proto'], + use_ms=True, + ): + from osprey.async_worker.adaptor.plugin_manager import bootstrap_async_action_proto_deserializer + + deserializer = bootstrap_async_action_proto_deserializer() + if deserializer is not None: + res = deserializer.proto_bytes_to_dict(osprey_coordinator_action.proto_action_data) + data = res.data + encoding = 'proto' + else: + logger.warning('Proto deserializer plugin not available, falling back to JSON processing') + data = json.loads(osprey_coordinator_action.proto_action_data) + encoding = 'json' + else: + metrics.increment( + 'osprey_coordinator_input_stream.deserialize_message_failure', + tags=tags + ['failure:invalid_serialization_type'], + ) + return None + + if not isinstance(data, dict): + raise ValueError('json_action_data was not a dict') + if not isinstance(secret_data, dict): + raise ValueError('json_secret_data was not a dict') + if not osprey_coordinator_action.action_name: + raise ValueError('action name must never be empty') + + return OspreyEngineAction( + action_id=osprey_coordinator_action.action_id, + action_name=osprey_coordinator_action.action_name, + data=data, + secret_data=secret_data, + timestamp=osprey_coordinator_action.timestamp.ToDatetime(tzinfo=pytz.utc), + encoding=encoding, + ) + except Exception: + logger.exception('Error while generating input message') + sentry_sdk.capture_exception() + metrics.increment( + 'osprey_coordinator_input_stream.deserialize_message_failure', + tags=tags + ['failure:unknown_exc'], + ) + return None + + # -- main loop ----------------------------------------------------------------- + + async def _gen(self) -> AsyncIterator[BaseAckingContext[OspreyEngineAction]]: + while not self._shutdown_event.is_set(): + channel, service = await self._channel_pool.get_connection() + bidirectional_stream = OspreyCoordinatorBiDirectionalStream( + client_id=self._client_id, channel=channel, service=service + ) + max_uptime_allowed = MIN_SECONDS_BEFORE_RECONNECT + random.uniform(0, SECONDS_BEFORE_RECONNECT_JITTER) + actions_handled = 0 + + async for osprey_coordinator_action in bidirectional_stream: + actions_handled += 1 + ack_id = osprey_coordinator_action.ack_id + osprey_engine_action = self._create_osprey_engine_action(osprey_coordinator_action) + + if not osprey_engine_action: + info_log_osprey_action( + osprey_coordinator_action.action_id, + osprey_coordinator_action.action_name, + "nacking (couldn't create OspreyEngineAction)", + ) + bidirectional_stream.send_ack_or_nack(ack_id, ack=False) + continue + + context: AsyncVerdictsAckingContext = AsyncVerdictsAckingContext( + osprey_engine_action, bidirectional_stream, ack_id + ) + with metrics.timed( + 'osprey_coordinator_input_stream.action_handle_time', + tags=[f'action_name:{osprey_engine_action.action_name}'], + use_ms=True, + ): + yield context + + # The sink may have marked this action for nack. Ack (carrying its + # verdicts) and nack are mutually exclusive, and this decision is the + # same on every finalize path below — graceful disconnect or normal. + ack = not context.should_nack + verdicts = context.get_verdicts() if ack else None + + # Prioritize shutdown so we can finalize the last action and disconnect gracefully + if self._shutdown_event.is_set(): + await bidirectional_stream.send_graceful_disconnect(ack_id, ack=ack, verdicts=verdicts) + break + + # Pause-and-rotate on rule reload. Mirrors the gevent equivalent at + # osprey/worker/sinks/sink/osprey_coordinator_input_stream.py:325-331: + # disconnect the bidi stream so the coordinator routes work elsewhere, + # wait for the reload to finish, then let the outer loop reconnect. + if ( + self._input_stream_ready_signaler is not None + and self._input_stream_ready_signaler.should_pause_input_stream() + ): + logger.info('Disconnecting due to input stream ready signaler') + await bidirectional_stream.send_graceful_disconnect(ack_id, ack=ack, verdicts=verdicts) + await self._input_stream_ready_signaler.wait_until_resume() + break + + # Reconnect after the jittered uptime threshold + uptime = bidirectional_stream.get_uptime() + if uptime > max_uptime_allowed: + logger.debug(f'Reconnecting because {uptime} seconds have passed') + await bidirectional_stream.send_graceful_disconnect(ack_id, ack=ack, verdicts=verdicts) + break + + # Normal path: ack (or nack) the last action and request the next one. + bidirectional_stream.send_ack_or_nack(ack_id, ack=ack, verdicts=verdicts) + + info_log_osprey_action( + osprey_coordinator_action.action_id, osprey_coordinator_action.action_name, 'acking' + ) + + if not self._shutdown_event.is_set(): + metrics.gauge('osprey_coordinator_input_stream.actions_handled', actions_handled) + logger.debug(f'Reconnecting due to stream ending, actions handled: {actions_handled}') + else: + logger.info('shutting down') + metrics.increment('osprey_coordinator_input_stream.shutdown') diff --git a/osprey_async_worker/src/osprey/async_worker/lib/discovery/__init__.py b/osprey_async_worker/src/osprey/async_worker/lib/discovery/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/lib/discovery/async_directory.py b/osprey_async_worker/src/osprey/async_worker/lib/discovery/async_directory.py new file mode 100644 index 0000000..bf19373 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/discovery/async_directory.py @@ -0,0 +1,382 @@ +"""Async service discovery via etcd. + +Pure asyncio replacement for osprey.worker.lib.discovery that avoids gevent. +Uses run_in_executor for the sync EtcdClient (etcd updates are infrequent). +The hash ring and service list are maintained in-memory with asyncio tasks +for background watching. +""" + +import asyncio +import collections +import json +import logging +from random import randint, uniform +from time import time +from typing import Any, Callable, ClassVar, Deque, Dict, List, Optional, Tuple + +from osprey.async_worker.lib.discovery.hash_ring import HashRing, HashRingNode +from osprey.worker.lib import etcd +from osprey.worker.lib.discovery.exceptions import ServiceUnavailable +from osprey.worker.lib.discovery.service import Service +from osprey.worker.lib.etcd import EtcdClient + +logger = logging.getLogger(__name__) + +UP = 'up' +DOWN = 'down' + +DEFAULT_SECONDARIES = 2 +VISIBILITY_PERIOD_MAX_SEC = 15.0 + +ListenerFn = Callable[[str, Service], None] + + +class _AsyncHashRing: + """Async-compatible hash ring backed by etcd. + + Reads the ring topology from etcd once, then watches for changes + via a background asyncio task (sync etcd calls in run_in_executor). + """ + + def __init__(self, etcd_client: EtcdClient, key: str, default_num_replicas: int = 512) -> None: + self._etcd_client = etcd_client + self._key = key + self._default_num_replicas = default_num_replicas + self._ring: HashRing = HashRing(default_num_replicas) + self._members: Dict[bytes, HashRingNode] = {} + self._overrides: Dict[Any, Any] = {} + self._initialized = False + self._watch_task: Optional[asyncio.Task[None]] = None + + async def ensure_initialized(self) -> None: + if self._initialized: + return + # Set early to prevent concurrent coroutines from double-initializing. + self._initialized = True + try: + loop = asyncio.get_running_loop() + watcher = await loop.run_in_executor(None, self._etcd_client.get_watcher, self._key, False) + event = await loop.run_in_executor(None, watcher.begin_watching) + self._handle_event(event) + self._watch_task = asyncio.create_task(self._watch_loop(watcher)) + except Exception: + self._initialized = False + raise + + def select(self, key: Any, secondaries: int = 0) -> Any: + if isinstance(key, str): + key = key.encode() + + if secondaries == 0: + override = self._overrides.get(key) + if override and len(override) > 0: + return override[0] + return self._ring.find_node(key) + + nodes = self._ring.find_nodes(key, secondaries + 1) + overrides = self._overrides.get(key) + if overrides is not None: + result = list(overrides) + for node in nodes: + if node not in result: + result.append(node) + return result[: secondaries + 1] + return nodes + + def _handle_event(self, event: Any) -> None: + if isinstance(event, etcd.FullSyncOne): + self._update_from_value(event.value) + elif isinstance(event, etcd.FullSyncOneNoKey): + self._ring = HashRing(self._default_num_replicas) + self._members = {} + self._overrides = {} + + def _update_from_value(self, value: str) -> None: + data = json.loads(value) + members: List[HashRingNode] = [] + overrides: Dict[Any, Any] = {} + + if isinstance(data, list): + members = self._read_members(data) + elif isinstance(data, dict) and data.get('schema') == 'v1': + members = self._read_members(data['members']) + overrides = self._read_overrides(data.get('overrides', [])) + + ring = HashRing(self._default_num_replicas) + ring.add_nodes(members) + self._ring = ring + self._members = {m.name: m for m in members} + self._overrides = overrides + + def _read_members(self, raw: List[Any]) -> List[HashRingNode]: + members = [] + for item in raw: + if isinstance(item, str): + members.append(HashRingNode(item.encode(), self._default_num_replicas)) + elif isinstance(item, dict): + name = item['name'] + if isinstance(name, str): + name = name.encode() + members.append(HashRingNode(name, int(item['num_replicas']))) + return members + + @staticmethod + def _read_overrides(raw: List[Any]) -> Dict[Any, Any]: + result = {} + for items in raw: + if len(items) >= 2: + key = items[0].encode() if isinstance(items[0], str) else items[0] + values = [v.encode() if isinstance(v, str) else v for v in items[1:]] + result[key] = values + return result + + async def _watch_loop(self, watcher: Any) -> None: + loop = asyncio.get_running_loop() + try: + while True: + event = await loop.run_in_executor(None, self._blocking_next, watcher) + if event is None: + break + self._handle_event(event) + except asyncio.CancelledError: + pass + except Exception: + logger.exception('hash ring watcher failed') + + @staticmethod + def _blocking_next(watcher: Any) -> Any: + """Blocking call — runs in thread pool.""" + try: + for event in watcher.continue_watching(): + return event + except Exception: + return None + return None + + async def stop(self) -> None: + if self._watch_task: + self._watch_task.cancel() + try: + await self._watch_task + except asyncio.CancelledError: + pass + + +class _ServiceWrapper: + __slots__ = ('service', 'visible_at') + + def __init__(self, service: Service, visible_at: Optional[float]) -> None: + self.service = service + self.visible_at = visible_at + + def is_visible(self, tolerate_draining: bool = False) -> bool: + if self.service.draining and not tolerate_draining: + return False + if self.visible_at is None: + return True + if self.visible_at < time(): + self.visible_at = None + return True + return False + + +class AsyncServiceWatcher: + """Async replacement for ServiceWatcher. + + Watches etcd for service instance changes using asyncio tasks. + Supports SCALAR routing (via hash ring) and ROUND_ROBIN (via rotation). + """ + + def __init__(self, etcd_client: EtcdClient, base_key: str, service_name: str) -> None: + self._etcd_client = etcd_client + self._key = f'{base_key}/{service_name}/instances' + self._service_name = service_name + self._instances: Dict[str, _ServiceWrapper] = {} + self._rotation: Deque[str] = collections.deque() + self._ring = _AsyncHashRing(etcd_client, f'{base_key}/{service_name}/ring') + self._listeners: List[ListenerFn] = [] + self._initialized = False + self._watch_task: Optional[asyncio.Task[None]] = None + + async def ensure_initialized(self) -> None: + if self._initialized: + return + # Set early to prevent concurrent coroutines from double-initializing. + self._initialized = True + try: + loop = asyncio.get_running_loop() + watcher = await loop.run_in_executor(None, self._etcd_client.get_watcher, self._key, True) + event = await loop.run_in_executor(None, watcher.begin_watching) + self._handle_full_sync(event, delay_visibility=False) + logger.info( + 'async watcher %s: %d instances loaded: %s', + self._service_name, + len(self._instances), + list(self._instances.keys()), + ) + await self._ring.ensure_initialized() + self._watch_task = asyncio.create_task(self._watch_loop(watcher)) + except Exception: + self._initialized = False + raise + + def select( + self, + selector: Any = None, + secondaries: int = DEFAULT_SECONDARIES, + instances_to_skip: int = 0, + tolerate_draining: bool = False, + ) -> Service: + if selector is None or callable(selector): + self._rotation.rotate() + not_yet_visible = [] + for service_id in self._rotation: + wrapper = self._instances[service_id] + if not selector or selector(wrapper.service): + if wrapper.is_visible(tolerate_draining): + return wrapper.service + elif not wrapper.service.draining: + not_yet_visible.append(wrapper.service) + if not_yet_visible: + from random import choice + + return choice(not_yet_visible) + raise ServiceUnavailable(f'No service for {self._service_name}') + else: + members = self._ring.select(selector, secondaries) + if not isinstance(members, list): + members = [members] + for member in members[instances_to_skip:]: + member_str = member.decode() if isinstance(member, bytes) else str(member) + # Distinct name from the `wrapper` above (which is non-optional from + # __getitem__); .get() returns Optional and the guard below narrows it. + maybe_wrapper = self._instances.get(member_str) + if maybe_wrapper and (tolerate_draining or not maybe_wrapper.service.draining): + return maybe_wrapper.service + raise ServiceUnavailable( + f'No service for {self._service_name} key={selector} ' + f'ring_members={[m.decode() if isinstance(m, bytes) else m for m in members]} ' + f'instances={list(self._instances.keys())}' + ) + + def select_all(self, tolerate_draining: bool = False) -> List[Service]: + services = [] + for service_id in self._rotation: + wrapper = self._instances[service_id] + if wrapper.is_visible(tolerate_draining): + services.append(wrapper.service) + return services if services else [w.service for w in self._instances.values() if not w.service.draining] + + def add_lazy_listener(self, listener: ListenerFn) -> None: + self._listeners.append(listener) + + def _handle_full_sync(self, event: Any, delay_visibility: bool = True) -> None: + if not isinstance(event, etcd.FullSyncRecursive): + return + latest = {} + for value in event.values(): + instance = Service.deserialize(value) + latest[instance.id] = instance + for instance_id in list(self._instances.keys()): + if instance_id not in latest: + self._remove_instance(instance_id) + for instance in latest.values(): + self._add_instance(instance, delay_visibility) + + def _add_instance(self, new_instance: Service, delay_visibility: bool = True) -> None: + if new_instance.id not in self._instances: + idx = randint(0, len(self._rotation)) + self._rotation.insert(idx, new_instance.id) + self._instances[new_instance.id] = _ServiceWrapper( + service=new_instance, + visible_at=(time() + uniform(0, VISIBILITY_PERIOD_MAX_SEC)) if delay_visibility else None, + ) + for listener in self._listeners: + listener(UP, new_instance) + logger.debug('async discovery: + %s@%s', new_instance.name, new_instance.id) + else: + existing = self._instances[new_instance.id] + existing.service.merge(new_instance) + if existing.service.id not in self._rotation: + self._rotation.insert(randint(0, len(self._rotation)), new_instance.id) + + def _remove_instance(self, instance_id: str) -> None: + wrapper = self._instances.pop(instance_id, None) + if wrapper is None: + return + if instance_id in self._rotation: + self._rotation.remove(instance_id) + for listener in self._listeners: + listener(DOWN, wrapper.service) + logger.debug('async discovery: - %s@%s', wrapper.service.name, instance_id) + + async def _watch_loop(self, watcher: Any) -> None: + loop = asyncio.get_running_loop() + try: + while True: + event = await loop.run_in_executor(None, self._blocking_next, watcher) + if event is None: + break + if isinstance(event, etcd.IncrementalSyncUpsert): + self._add_instance(Service.deserialize(event.value)) + elif isinstance(event, etcd.IncrementalSyncDelete): + self._remove_instance(Service.deserialize(event.prev_value).id) + elif isinstance(event, etcd.FullSyncRecursive): + self._handle_full_sync(event) + except asyncio.CancelledError: + pass + except Exception: + logger.exception('async service watcher failed for %s', self._service_name) + + @staticmethod + def _blocking_next(watcher: Any) -> Any: + try: + for event in watcher.continue_watching(): + return event + except Exception: + return None + return None + + async def stop(self) -> None: + if self._watch_task: + self._watch_task.cancel() + try: + await self._watch_task + except asyncio.CancelledError: + pass + await self._ring.stop() + + +class AsyncDirectory: + """Async replacement for Directory. + + Singleton that creates AsyncServiceWatcher instances per service name. + Uses the sync EtcdClient via run_in_executor for etcd reads. + """ + + _instances: ClassVar[Dict[Tuple[Any, ...], 'AsyncDirectory']] = {} + + @classmethod + def instance(cls, secure: bool = False) -> 'AsyncDirectory': + key = (secure,) + if key not in cls._instances: + cls._instances[key] = cls(secure=secure) + return cls._instances[key] + + def __init__(self, base_key: str = '/discovery', secure: bool = False) -> None: + self._base_key = base_key + self._etcd_client = EtcdClient(secure=secure) + self._watchers: Dict[str, AsyncServiceWatcher] = {} + + def get_watcher(self, service_name: str) -> AsyncServiceWatcher: + if service_name not in self._watchers: + self._watchers[service_name] = AsyncServiceWatcher(self._etcd_client, self._base_key, service_name) + return self._watchers[service_name] + + def select_all(self, service_name: str) -> List[Service]: + watcher = self.get_watcher(service_name) + return watcher.select_all() + + async def stop(self) -> None: + for watcher in self._watchers.values(): + await watcher.stop() diff --git a/osprey_async_worker/src/osprey/async_worker/lib/discovery/hash_ring.py b/osprey_async_worker/src/osprey/async_worker/lib/discovery/hash_ring.py new file mode 100644 index 0000000..d2434ba --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/discovery/hash_ring.py @@ -0,0 +1,80 @@ +"""Self-contained consistent hash ring for async-worker service discovery. + +A small, dependency-free (stdlib-only) consistent hashing implementation used by +the async worker's etcd-backed service discovery. The API mirrors the subset the +discovery layer relies on: + + - ``HashRingNode(name: bytes, num_replicas: int)`` with ``.name`` / ``.num_replicas`` + - ``HashRing(default_num_replicas).add_nodes([...])`` + - ``HashRing.find_node(key: bytes) -> Optional[bytes]`` (the node *name*) + - ``HashRing.find_nodes(key: bytes, count: int) -> List[bytes]`` (distinct names) + +Placement is deterministic (MD5 over ``:`` virtual points) so every +worker computes the same ring for the same membership. Exact placement is an +internal detail — only self-consistency across a fleet matters. +""" + +from __future__ import annotations + +import bisect +import hashlib +from typing import Dict, List, Optional, Sequence, Set + + +class HashRingNode: + __slots__ = ('name', 'num_replicas') + + def __init__(self, name: bytes, num_replicas: int) -> None: + self.name = name + self.num_replicas = num_replicas + + +def _hash(data: bytes) -> int: + # 64 bits of MD5 is plenty of spread for ring placement. + return int.from_bytes(hashlib.md5(data).digest()[:8], 'big') + + +class HashRing: + def __init__(self, default_num_replicas: int = 512) -> None: + self._default_num_replicas = default_num_replicas + self._points: List[int] = [] # sorted virtual-point hashes + self._owner: Dict[int, bytes] = {} # point hash -> node name + self._names: Set[bytes] = set() + + def add_node(self, node: HashRingNode) -> None: + if node.name in self._names: + return + self._names.add(node.name) + replicas = node.num_replicas or self._default_num_replicas + for i in range(replicas): + point = _hash(node.name + b':' + str(i).encode()) + # Skip the rare collision rather than silently reassign ownership. + if point not in self._owner: + self._owner[point] = node.name + bisect.insort(self._points, point) + + def add_nodes(self, nodes: Sequence[HashRingNode]) -> None: + for node in nodes: + self.add_node(node) + + def find_node(self, key: bytes) -> Optional[bytes]: + if not self._points: + return None + idx = bisect.bisect(self._points, _hash(key)) % len(self._points) + return self._owner[self._points[idx]] + + def find_nodes(self, key: bytes, count: int) -> List[bytes]: + if not self._points or count <= 0: + return [] + start = bisect.bisect(self._points, _hash(key)) + total = len(self._points) + result: List[bytes] = [] + seen: Set[bytes] = set() + for offset in range(total): + name = self._owner[self._points[(start + offset) % total]] + if name not in seen: + seen.add(name) + result.append(name) + if len(result) >= count: + break + return result diff --git a/osprey_async_worker/src/osprey/async_worker/lib/etcd/__init__.py b/osprey_async_worker/src/osprey/async_worker/lib/etcd/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/lib/etcd/sources_provider.py b/osprey_async_worker/src/osprey/async_worker/lib/etcd/sources_provider.py new file mode 100644 index 0000000..7ea59ce --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/etcd/sources_provider.py @@ -0,0 +1,217 @@ +"""Async sources provider for the async worker. + +Port of osprey.worker.lib.sources_provider with asyncio instead of gevent. +The etcd watcher runs in a thread pool (run_in_executor) since the underlying +etcd client is synchronous. This is acceptable because etcd updates are +infrequent (rule deployments, not per-request). +""" + +import asyncio +import inspect +import json +import logging +import random +from typing import Any, Awaitable, Callable, Dict, Iterator, Optional, Union + +from osprey.engine.ast.sources import Sources +from osprey.worker.lib.etcd import BaseWatcher, EtcdClient, FullSyncOne, FullSyncOneNoKey +from osprey.worker.lib.sources_provider_base import BaseSourcesProvider + +# The async engine's _handle_updated_sources is a coroutine function so the +# compile can run in a thread pool while the event loop continues servicing +# in-flight tasks. We accept either a sync or async callable for back-compat. +SourcesWatcherCallback = Callable[[], Union[None, Awaitable[None]]] + + +class AsyncInputStreamReadySignaler: + """Async version of InputStreamReadySignaler. + + Uses asyncio.Event instead of gevent.event.Event for pause/resume signaling. + """ + + def __init__(self) -> None: + self._event = asyncio.Event() + self._event.set() # Start in "ready" state + + def should_pause_input_stream(self) -> bool: + return not self._event.is_set() + + async def pause_input_stream(self) -> None: + # Match the gevent worker's nominal jitter range. The compile pauses + # input on the worker for ~28s; spreading the pause start uniformly + # across the fleet over 10 minutes keeps the fraction of paused pods + # to ~5% at any moment, avoiding the throughput cliff that happens + # when the whole fleet pauses together. + await asyncio.sleep(random.uniform(0, 600)) + self._event.clear() + + def resume_input_stream(self) -> None: + self._event.set() + + async def wait_until_resume(self) -> None: + await self._event.wait() + + +class AsyncEtcdSourcesProvider(BaseSourcesProvider): + """Provides sources dynamically updated by etcd, using asyncio. + + The etcd client is synchronous, so watch operations are offloaded to + a thread pool via run_in_executor. This is fine because etcd updates + happen infrequently (rule deployments). + """ + + def __init__( + self, + etcd_key: str, + etcd_client: Optional[EtcdClient] = None, + input_stream_ready_signaler: Optional[AsyncInputStreamReadySignaler] = None, + ): + self._etcd_key = etcd_key + self._client = etcd_client or EtcdClient() + self._current_sources: Optional[Sources] = None + # The hash the engine has actually compiled and swapped in. Dedup keys off + # this (not the last *received* sources) so a failed/dropped compile leaves + # it behind and the next identical etcd re-delivery re-fires the recompile + # instead of being suppressed (which would wedge the worker on stale rules). + # Advanced only from mark_sources_applied(); seeded None so every event + # before the first successful apply triggers a compile attempt. + self._applied_sources_hash: Optional[str] = None + self._sources_watcher_callback: Optional[SourcesWatcherCallback] = None + self._input_stream_ready_signaler = input_stream_ready_signaler + self._watcher: Optional[BaseWatcher] = None + # Long-lived iterator over the watcher's event stream. continue_watching() + # is a generator function — every call creates a new generator with a + # fresh WatchMux and a reset _index, which defeats the watcher's built-in + # dedup of redundant FullSyncOne events. Iterate one generator persistently + # to match how ReadOnlyEtcdDict drives the gevent watcher. + self._watcher_iter: Optional[Iterator[Any]] = None + self._watcher_task: Optional[asyncio.Task[None]] = None + + async def start(self) -> None: + """Initialize sources from etcd and start watching for changes.""" + loop = asyncio.get_running_loop() + + # Initial load in thread pool (sync etcd client) + initial_dict = await loop.run_in_executor(None, self._load_initial) + self._current_sources = Sources.from_dict(initial_dict) + + # Start watcher loop as an async task + self._watcher_task = asyncio.create_task(self._watch_loop()) + + def _load_initial(self) -> Dict[str, str]: + """Load initial sources from etcd. Runs in thread pool.""" + watcher = self._client.get_watcher(self._etcd_key, recursive=False) + initial_event = watcher.begin_watching() + self._watcher = watcher + return self._parse_event(initial_event) + + def _parse_event(self, event) -> Dict[str, str]: + """Parse an etcd event into a sources dict.""" + if isinstance(event, FullSyncOne): + return json.loads(str(event.value)) + elif isinstance(event, FullSyncOneNoKey): + return {} + return {} + + async def _watch_loop(self) -> None: + """Watch for etcd changes, running the sync watcher in a thread pool.""" + loop = asyncio.get_running_loop() + backoff = 1.0 + try: + while True: + if self._watcher is None: + self._watcher = await loop.run_in_executor(None, self._client.get_watcher, self._etcd_key, False) + self._watcher_iter = None + if self._watcher_iter is None: + assert self._watcher is not None + self._watcher_iter = self._watcher.continue_watching() + + # Block in thread pool waiting for next etcd event. Drive the + # SAME generator each iteration — the watcher's WatchMux dedups + # consecutive identical FullSyncOne events (which are common + # post-rule-deploy as etcd's wait API re-syncs), but the dedup + # state lives on the generator. Recreating the generator per + # event would defeat the dedup and put the loop into a tight + # SYNC re-fetch loop on the main asyncio thread. + watcher_iter = self._watcher_iter + try: + event = await loop.run_in_executor(None, lambda: next(watcher_iter)) + backoff = 1.0 # Reset on success + except StopIteration: + # Generator exhausted (e.g. transient etcd error inside + # continue_watching). Reopen the watcher fresh. + self._watcher = None + self._watcher_iter = None + continue + except Exception: + logging.exception('Error in etcd watcher loop, retrying in %.1fs', backoff) + self._watcher = None + self._watcher_iter = None + await asyncio.sleep(backoff) + backoff = min(backoff * 2, 30.0) + continue + + if event is not None: + await self._handle_event(event) + except asyncio.CancelledError: + return + finally: + self._watcher = None + self._watcher_iter = None + + async def _handle_event(self, event) -> None: + """Handle an etcd event by updating sources and notifying watchers.""" + sources_dict = self._parse_event(event) + new_sources = Sources.from_dict(sources_dict) + + # Etcd watcher reconnects and session refreshes re-deliver the current + # value as a FullSyncOne event, so we see many events where the + # content is unchanged. Skip the (peak-memory-doubling) recompile only + # when the engine has ALREADY applied this exact hash. Comparing against + # the last *applied* hash — not merely the last *received* one — keeps + # this self-healing: if a recompile fails the engine keeps its old graph + # and _applied_sources_hash stays behind, so the next re-delivery re-fires + # the recompile instead of leaving the worker wedged on stale rules. + if self._applied_sources_hash is not None and new_sources.hash() == self._applied_sources_hash: + return + + if self._input_stream_ready_signaler is not None: + logging.info('Pausing input streams') + await self._input_stream_ready_signaler.pause_input_stream() + + self._current_sources = new_sources + if self._sources_watcher_callback: + result = self._sources_watcher_callback() + if inspect.isawaitable(result): + await result + + if self._input_stream_ready_signaler is not None: + logging.info('Restarting input streams') + self._input_stream_ready_signaler.resume_input_stream() + + # NOTE: BaseSourcesProvider.get_current_sources is typed -> Sources, but this + # provider legitimately returns None before start() loads from etcd (see + # test_provider_get_current_sources_default_none). Widening the base return type + # to Optional[Sources] is the correct fix but lives in sources_provider_base.py + # (a shared file outside this change). Keep the precise return type here. + def get_current_sources(self) -> Optional[Sources]: # type: ignore[override] + return self._current_sources + + def mark_sources_applied(self, sources_hash: str) -> None: + # Advances the dedup baseline only once the engine confirms it compiled + # and swapped these sources, so failed/dropped applies retry on the next + # etcd re-delivery rather than being deduped away. + self._applied_sources_hash = sources_hash + + def set_sources_watcher(self, callback: SourcesWatcherCallback) -> None: + self._sources_watcher_callback = callback + + async def stop(self) -> None: + """Stop watching for etcd changes.""" + if self._watcher_task is not None: + self._watcher_task.cancel() + try: + await self._watcher_task + except asyncio.CancelledError: + pass + self._watcher_task = None diff --git a/osprey_async_worker/src/osprey/async_worker/lib/external_service.py b/osprey_async_worker/src/osprey/async_worker/lib/external_service.py new file mode 100644 index 0000000..93f8f59 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/external_service.py @@ -0,0 +1,148 @@ +"""Async external service utilities for the async worker. + +Port of osprey.engine.executor.external_service_utils with asyncio instead of gevent. +Uses asyncio.Future instead of gevent.event.AsyncResult for cache entries. +""" + +import asyncio +from abc import ABC, abstractmethod +from datetime import datetime, timedelta +from typing import Dict, Generic, Hashable, Optional, Sequence, Tuple, TypeVar, cast + +from result import Err, Ok, Result + +KeyT = TypeVar('KeyT', bound=Hashable) +ValueT = TypeVar('ValueT') + + +class AsyncExternalService(ABC, Generic[KeyT, ValueT]): + @abstractmethod + async def get_from_service(self, key: KeyT) -> ValueT: + raise NotImplementedError + + # Not abstract because not all services support batching multiple keys + async def batch_get_from_service(self, keys: Sequence[KeyT]) -> Sequence[Result[ValueT, Exception]]: + raise NotImplementedError + + def cache_ttl(self) -> Optional[timedelta]: + """ + Returns a time to live for items in the cache. By default, KVs are cached indefinitely. + + To have cache entries auto-expire, override this method in your external service definition. + + Note that timedeltas can accept negative values to represent the past, but only on the days field. + You *can* use timedelta(seconds=0) to disable caching, but a negative time delta *ensures* that even + if a time shift occurs (such as daylight savings), the cache_ttl will still be immediate. + + Therefore, to disable the read cache, it is recommended to set this to `timedelta(days=-1)` + """ + return None + + def count_error_once(self) -> bool: + """ + When True, only the caller that initiated the external service call + receives the exception. Subsequent callers that would hit the cached + error receive None instead. + + Only enable this when ValueT is Optional and None is a safe fallback. + """ + return False + + +class ExternalServiceAccessor(Generic[KeyT, ValueT]): + """Facilitates accessing an async external service in a way that caches and debounces requests based on a key.""" + + def __init__(self, service: AsyncExternalService[KeyT, ValueT]): + self._service = service + # Key -> Tuple[ Future[ValueT], Expiration datetime ] + self._cache: Dict[KeyT, Tuple[asyncio.Future[ValueT], Optional[datetime]]] = {} + + def _is_past_cache_expiration(self, cache_expiration: Optional[datetime]) -> bool: + """ + Helper method to perform a time check on an optional datetime. + """ + if cache_expiration is None: + return False + return datetime.now() > cache_expiration + + def _get_cache_expiration_datetime(self) -> Optional[datetime]: + """ + Helper method to generate an optional cache expiration datetime based on the cache TTL. + """ + ttl = self._service.cache_ttl() + return datetime.now() + ttl if ttl is not None else None + + def _make_future(self) -> asyncio.Future[ValueT]: + """Create a new Future on the running event loop.""" + return asyncio.get_running_loop().create_future() + + async def get_without_cache(self, key: KeyT) -> ValueT: + """ + Ignores any cached values and performs a read-through `get` to the external service. + The new value is then used to update the cache entry for subsequent `get` calls. + """ + future: asyncio.Future[ValueT] = self._make_future() + cache_entry: Tuple[asyncio.Future[ValueT], Optional[datetime]] = ( + future, + self._get_cache_expiration_datetime(), + ) + self._cache[key] = cache_entry + try: + result = await self._service.get_from_service(key) + future.set_result(result) + except Exception as e: + future.set_exception(e) + + return await future + + async def get(self, key: KeyT) -> ValueT: + cache_entry = self._cache.get(key) + if cache_entry is not None and not self._is_past_cache_expiration(cache_entry[1]): + # Cache hit — await the existing future (may still be in-flight from another caller) + return await cache_entry[0] + + future: asyncio.Future[ValueT] = self._make_future() + cache_entry = (future, self._get_cache_expiration_datetime()) + self._cache[key] = cache_entry + try: + result = await self._service.get_from_service(key) + future.set_result(result) + except Exception as e: + if self._service.count_error_once(): + future.set_result(cast(ValueT, None)) + else: + future.set_exception(e) + raise + + return await future + + async def batch_get(self, keys: Sequence[KeyT]) -> Sequence[Result[ValueT, Exception]]: + cached_entries = [self._cache.get(key) for key in keys] + non_cached_keys = [ + key + for key, cache_entry in zip(keys, cached_entries) + if cache_entry is None or self._is_past_cache_expiration(cache_entry[1]) + ] + if non_cached_keys: + for key in non_cached_keys: + self._cache[key] = (self._make_future(), self._get_cache_expiration_datetime()) + try: + result = await self._service.batch_get_from_service(non_cached_keys) + for i, key in enumerate(non_cached_keys): + if result[i].is_ok(): + self._cache[key][0].set_result(result[i].unwrap()) + else: + self._cache[key][0].set_exception(cast(BaseException, result[i].value)) + except Exception as e: + for key in non_cached_keys: + self._cache[key][0].set_exception(e) + + results: list[Result[ValueT, Exception]] = [] + for key in keys: + future = self._cache[key][0] + try: + value = await future + results.append(Ok(value)) + except Exception as e: + results.append(Err(e)) + return results diff --git a/osprey_async_worker/src/osprey/async_worker/lib/pigeon/__init__.py b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/lib/pigeon/client.py b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/client.py new file mode 100644 index 0000000..55d7b8c --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/client.py @@ -0,0 +1,608 @@ +import asyncio +import copy +import itertools +import logging +import weakref +from collections import defaultdict +from collections.abc import Mapping +from dataclasses import dataclass +from time import time_ns +from typing import Any, Dict, Generic, List, Optional, Set, Tuple, Type, TypeVar, Union, cast + +import grpc +import grpc.aio +from ddtrace.constants import ERROR_MSG +from ddtrace.contrib.grpc.constants import GRPC_STATUS_CODE_KEY +from ddtrace.ext.http import STATUS_CODE +from ddtrace.span import Span +from google.protobuf.message import Message +from osprey.async_worker.lib.discovery.async_directory import AsyncDirectory +from osprey.async_worker.lib.pigeon.exceptions import InvalidRoutingValueException, NoResponsesException, RPCException +from osprey.async_worker.lib.pigeon.interceptors.baggage import BaggageInterceptor +from osprey.async_worker.lib.pigeon.interceptors.metadata import MetadataInterceptor +from osprey.async_worker.lib.pigeon.skip_rate_limit import skip_rate_limit_context +from osprey.worker.lib.ddtrace_utils import current_span, noop_span, pin_override, trace +from osprey.worker.lib.discovery.exceptions import ServiceUnavailable +from osprey.worker.lib.discovery.service import Service +from osprey.worker.lib.instruments import metrics +from typing_extensions import TypedDict + +logger = logging.getLogger(__name__) + +T = TypeVar('T') + +# String constant matching osprey.worker.lib.discovery.service_watcher.DOWN +# Defined locally to avoid importing service_watcher which pulls in gevent. +DOWN = 'down' + + +class ServiceDefinition(TypedDict): + address: str + ip: Optional[str] + port: int + + +# This mapping follows the recommendation at https://cloud.google.com/apis/design/errors#handling_errors. +# We *could* use `http.HTTPStatus` here, but it's arguably more readable without the abstraction in this case +# and wouldn't be complete anyway as it doesn't contain anything for 499. +_GRPC_HTTP_CODE_TRANSLATIONS = { + grpc.StatusCode.OK: 200, + grpc.StatusCode.CANCELLED: 499, + grpc.StatusCode.UNKNOWN: 500, + grpc.StatusCode.INVALID_ARGUMENT: 400, + grpc.StatusCode.DEADLINE_EXCEEDED: 504, + grpc.StatusCode.NOT_FOUND: 404, + grpc.StatusCode.ALREADY_EXISTS: 409, + grpc.StatusCode.PERMISSION_DENIED: 403, + grpc.StatusCode.RESOURCE_EXHAUSTED: 429, + grpc.StatusCode.FAILED_PRECONDITION: 400, + grpc.StatusCode.ABORTED: 409, + grpc.StatusCode.OUT_OF_RANGE: 400, + grpc.StatusCode.UNIMPLEMENTED: 501, + grpc.StatusCode.INTERNAL: 500, + grpc.StatusCode.UNAVAILABLE: 503, + grpc.StatusCode.DATA_LOSS: 500, + grpc.StatusCode.UNAUTHENTICATED: 401, +} +_GRPC_CODE_FALLBACK = grpc.StatusCode.UNKNOWN + +# This is the name of the span that the pigeon client uses. The Service RPC Guard uses this name +# to avoid creating a new span for the grpc request and just adds info to the existing guard span +PIGEON_REQUEST_SPAN_NAME = 'pigeon.request' + + +class RoutingType: + CHUNKED = 1 + SCALAR = 2 + ROUND_ROBIN = 3 + ENVOY = 4 + + ALL = {CHUNKED, SCALAR, ROUND_ROBIN, ENVOY} + + +class RoutedClient(Generic[T]): + def __init__( + self, + service_name, + read_timeout, + stub_cls: Type[T], + request_field=None, + request_field_routing_value_transform=None, + routing_type=RoutingType.CHUNKED, + secondaries=1, + pool_size=200, + chunk_size=250, + secure_etcd=False, + metadata=None, + interceptors=None, + grpc_options=None, + acceptable_duration_ms=None, + baggage_header=None, + baggage=None, + envoy_endpoint: Optional[ServiceDefinition] = None, + use_peer_service_name=False, + default_retry_policy: Optional['RetryPolicy'] = None, + ): + self._service_name = service_name + self._peer_service = f'{service_name}-client' if use_peer_service_name else service_name + self._stub_cls: Type[T] = stub_cls + self._request_field = request_field + self._request_field_routing_value_transform = request_field_routing_value_transform + self._routing_type = routing_type + self._open_channels: Dict[Tuple[Tuple[str, Optional[str]], int], weakref.ReferenceType[grpc.aio.Channel]] = {} + self._clients: Dict[Tuple[Tuple[str, Optional[str]], int], T] = {} + self._secondaries = secondaries + self._chunk_size = chunk_size + self._semaphore = asyncio.Semaphore(pool_size) + self._read_timeout = read_timeout + grpc_options = {'grpc.keepalive_time_ms': 300_000, **(grpc_options or dict())} + self._grpc_options = list(grpc_options.items()) + self._connect_eagerly = False + self._acceptable_duration_ms: Optional[int] = acceptable_duration_ms + self._default_retry_policy: Optional['RetryPolicy'] = default_retry_policy + self._interceptors: List[Any] = [BaggageInterceptor(baggage_header=baggage_header, baggage=baggage)] + + if metadata: + self._interceptors.append(MetadataInterceptor(metadata)) + + if interceptors: + self._interceptors.extend(interceptors) + + if RoutingType.ENVOY == self._routing_type: + if envoy_endpoint is None: + raise ValueError('RoutingType.ENVOY could not create service') + self._envoy_service = Service( + name=service_name, + address=envoy_endpoint['address'], + port=int(envoy_endpoint['port']), + ports={'grpc': envoy_endpoint['port']}, + metadata={}, + ip=envoy_endpoint.get('ip'), + ) + else: + self._async_directory = AsyncDirectory.instance(secure=secure_etcd) + self._service_watcher = self._async_directory.get_watcher(service_name) + self._service_watcher_initialized = False + self._handle_service_change_fn = self._handle_service_change + self._service_watcher.add_lazy_listener(self._handle_service_change_fn) + + def __getattr__(self, method_name) -> 'AsyncUnaryUnaryRpcCallable[T, Any, Any]': + return AsyncUnaryUnaryRpcCallable(service_name=self._service_name, method_name=method_name, client=self) + + @property + def acceptable_duration_ms(self) -> Optional[int]: + return self._acceptable_duration_ms + + async def _ensure_watcher_initialized(self) -> None: + """Lazily initialize the async service watcher on first use. + + Must be called from an async context (inside the event loop). + Safe to call multiple times — only initializes once. + """ + if hasattr(self, '_service_watcher_initialized') and not self._service_watcher_initialized: + # Set flag before awaiting to prevent concurrent coroutines from + # double-initializing. Reset on failure so the next call retries. + self._service_watcher_initialized = True + try: + await self._service_watcher.ensure_initialized() + logger.info( + 'async service watcher initialized for %s: %d instances', + self._service_name, + len(self._service_watcher._instances), + ) + except Exception: + self._service_watcher_initialized = False + logger.exception('failed to initialize async service watcher for %s', self._service_name) + raise + + async def request( + self, + method_name: str, + message: Message, + request_field: Optional[str] = None, + routing_type: Optional[int] = None, + timeout: Optional[float] = None, + metadata: Optional[List[Tuple[str, str]]] = None, + instances_to_skip: int = 0, + ): + await self._ensure_watcher_initialized() + routing_type = routing_type if routing_type is not None else self._routing_type + request_field = request_field if request_field is not None else self._request_field + timeout = timeout or self._read_timeout + + if skip_rate_limit_context.skip: + metadata = metadata or [] + metadata.append(('skip-rate-limit', 'true')) + + if routing_type == RoutingType.CHUNKED: + return await self._chunked_request(method_name, message, request_field, timeout=timeout, metadata=metadata) + else: + return await self._request( + method_name, + message, + request_field, + routing_type, + timeout=timeout, + metadata=metadata, + instances_to_skip=instances_to_skip, + ) + + async def _chunked_request( + self, + method_name: str, + message: Message, + request_field: str, + timeout: Optional[float] = None, + metadata: Optional[List[Tuple[str, str]]] = None, + ): + """Call a remote service concurrently. Route based on a routing key.""" + calls = self._generate_routed_calls(request_field, message) + num_chunks = len(calls) + + current_span().set_tag('num_request_chunks', str(num_chunks)) + + # Avoid spawning a task for singular calls. + if num_chunks == 1: + return await self._do_routed_request( + method_name, message, request_field, timeout, metadata, next(iter(calls.items())) + ) + + async def _bounded_request(item): + async with self._semaphore: + return await self._do_routed_request(method_name, message, request_field, timeout, metadata, item) + + responses = await asyncio.gather(*[_bounded_request(item) for item in calls.items()]) + + final_response = None + for response in responses: + if response: + if not final_response: + final_response = response + else: + final_response.MergeFrom(response) + + if final_response is None: + raise NoResponsesException() + + return final_response + + async def _request( + self, + method_name: str, + message: Message, + request_field: str, + routing_type: int, + timeout: Optional[float] = None, + metadata: Optional[List[Tuple[str, str]]] = None, + instances_to_skip: int = 0, + ): + """Request from a remote service.""" + service = self._select_service(message, request_field, routing_type, instances_to_skip) + client = self._get_client(service) + method = getattr(client, method_name) + return await method(message, timeout=timeout, metadata=metadata) + + async def _do_routed_request( + self, + method_name: str, + message_template: Message, + request_field: str, + timeout: Optional[float], + metadata: Optional[List[Tuple[str, str]]], + service_and_routing_values, + ) -> Optional[Message]: + """Request from remote service.""" + (service, routing_values) = service_and_routing_values + with maybe_start_span('pigeon.routed_request', self._peer_service, method_name): + span = current_span() + set_protocol(span) + client = self._get_client(service) + method = getattr(client, method_name) + routing_values_iter = iter(routing_values) + final_response = None + while True: + routing_values_chunk = list(itertools.islice(routing_values_iter, 0, self._chunk_size)) + if not routing_values_chunk: + break + + next_message = _make_message(message_template, request_field, routing_values, routing_values_chunk) + response = await method(next_message, timeout=timeout, metadata=metadata) + if not final_response: + final_response = response.__class__() + final_response.MergeFrom(response) + + return final_response + + def _select_service( + self, message: Message, request_field: str, routing_type: int, instances_to_skip: int = 0 + ) -> Service: + """Select a service based on the client routing.""" + if routing_type == RoutingType.SCALAR: + routing_value = getattr(message, request_field) + if self._request_field_routing_value_transform: + routing_value = self._request_field_routing_value_transform(routing_value) + + return self._service_watcher.select( + routing_value, secondaries=self._secondaries, instances_to_skip=instances_to_skip + ) + elif routing_type == RoutingType.ROUND_ROBIN: + return self._service_watcher.select(secondaries=self._secondaries) + elif routing_type == RoutingType.ENVOY: + return self._envoy_service + else: + # RoutingType.CHUNKED is handled in another code path. + raise RuntimeError(f'RoutingType {routing_type} not supported') + + def _generate_routed_calls(self, request_field, message): + """Generate the routed calls. XXX: Modifies message by removing its routed values.""" + field = getattr(message, request_field) + if not field: + raise InvalidRoutingValueException(request_field, field) + + request_field_routing_value_transform = self._request_field_routing_value_transform + + if isinstance(field, Mapping): + groups = defaultdict(dict) + for key, value in field.items(): + if request_field_routing_value_transform: + key = request_field_routing_value_transform(key) + + groups[self._service_watcher.select(key, self._secondaries)][key] = value + else: + groups = defaultdict(list) + for value in field: + if request_field_routing_value_transform: + value = request_field_routing_value_transform(value) + + groups[self._service_watcher.select(value, self._secondaries)].append(value) + + message.ClearField(request_field) + return groups + + def connect_eagerly(self) -> None: + self._connect_eagerly = True + for service in self._service_watcher.select_all(): + self._get_client(service) + + def _handle_service_change(self, status: str, service: Service) -> None: + # Clean up the client when the service marks itself as DOWN. + if status != DOWN: + return + service_key = self._get_service_key(service) + self._cleanup_client(service_key) + + def _cleanup_client(self, service_key): + if service_key in self._open_channels: + channel = self._open_channels.pop(service_key)() + if channel is not None and hasattr(channel, 'close'): + # grpc.aio channels have an async close, but we fire-and-forget here + # since this is called from a sync callback. + try: + asyncio.get_event_loop().create_task(channel.close()) + except RuntimeError: + pass # No running event loop (shutdown or non-main thread) + if service_key in self._clients: + del self._clients[service_key] + + @staticmethod + def _get_service_key(service: Service) -> Tuple[Tuple[str, Optional[str]], int]: + return (service.connection_key, service.grpc_port) + + def _get_client(self, service: Service) -> T: + """Get the client for the service""" + key = self._get_service_key(service) + try: + return self._clients[key] + except KeyError: + addr_port = f'{service.connection_address}:{service.grpc_port}' + + # grpc.aio requires async interceptors (grpc.aio.UnaryUnaryClientInterceptor). + # The BaggageInterceptor/MetadataInterceptor are sync interceptors that only + # work with ENVOY routing (pre-created channels). For discovered services + # (SCALAR/ROUND_ROBIN), create channels without interceptors. + pin_override(grpc.Channel, f'{service.name}-grpc-client') + channel = grpc.aio.insecure_channel(addr_port, options=self._grpc_options) + self._open_channels[key] = weakref.ref(channel) + + pin_override(grpc.Channel, None) + + client = self._stub_cls(channel) # type: ignore + self._clients[key] = client + return client + + +def _make_message( + message_template: Message, + request_field: str, + routing_values: Union[List[Any], Dict[Any, Any]], + routing_values_chunk: List[Any], +): + message = copy.copy(message_template) + field = getattr(message, request_field) + if isinstance(field, Mapping): + for key in routing_values_chunk: + value = routing_values[key] + if isinstance(value, Message): + field[key].CopyFrom(value) + else: + field[key] = value # type: ignore + else: + field.extend(routing_values_chunk) + return message + + +Request = TypeVar('Request') +Response = TypeVar('Response') + + +class RetryPolicy(TypedDict): + retryable_grpc_status_codes: Set[grpc.StatusCode] + max_secondaries_to_retry: int + + +@dataclass +class AsyncUnaryUnaryRpcCallable(Generic[T, Request, Response]): + service_name: str + method_name: str + client: RoutedClient[T] + + async def __call__( + self, + message: Request, + request_field: Optional[str] = None, + routing_type: Optional[int] = None, + timeout: Optional[float] = None, + acceptable_duration_ms: Optional[int] = None, + metadata: Optional[List[Tuple[str, str]]] = None, + retry_policy: Optional[RetryPolicy] = None, + ) -> Response: + retry_policy = retry_policy or self.client._default_retry_policy + try_count = 0 + last_exception = None + while True: + try: + instances_to_skip = try_count + return await self.request( + message, + request_field, + routing_type, + timeout, + acceptable_duration_ms, + metadata, + instances_to_skip, + ) + except ServiceUnavailable as e: + if last_exception: + # ran out of secondaries to try + raise last_exception + + if retry_policy and try_count < retry_policy['max_secondaries_to_retry']: + # Retry ServiceUnavailable (empty ring at startup) with backoff + try_count += 1 + await asyncio.sleep(0.5 * try_count) + continue + + raise e + except RPCException as e: + last_exception = e + error_code = e.code() + if ( + retry_policy + and error_code in retry_policy['retryable_grpc_status_codes'] + and try_count < retry_policy['max_secondaries_to_retry'] + ): + try_count += 1 + await asyncio.sleep(0.5 * try_count) + continue + + raise e + + async def request( + self, + message: Request, + request_field: Optional[str] = None, + routing_type: Optional[int] = None, + timeout: Optional[float] = None, + acceptable_duration_ms: Optional[int] = None, + metadata: Optional[List[Tuple[str, str]]] = None, + instances_to_skip: int = 0, + ) -> Response: + pb2_message = self._to_proto(message) + start = time_ns() + tags = [f'service:{self.service_name}', f'resource_name:{self.method_name}'] + grpc_code = _GRPC_CODE_FALLBACK + acceptable_duration_ms = acceptable_duration_ms or self.client.acceptable_duration_ms + + with maybe_start_span(PIGEON_REQUEST_SPAN_NAME, self.client._peer_service, self.method_name): + span = current_span() + set_protocol(span) + + try: + response = await self.client.request( + self.method_name, + pb2_message, + request_field=request_field, + routing_type=routing_type, + timeout=timeout, + metadata=metadata, + instances_to_skip=instances_to_skip, + ) + + grpc_code = grpc.StatusCode.OK + duration_tag = 'classification:acceptable' + + duration_ms = round((time_ns() - start) / 1000000) + if acceptable_duration_ms and duration_ms > acceptable_duration_ms: + duration_tag = 'classification:unacceptable' + + tags.append(duration_tag) + + return self._from_proto(response) + except grpc.RpcError as e: + error = RPCException(self.service_name, self.method_name, e) + grpc_code = error.code() + + span.set_tag(ERROR_MSG, str(error)) + + raise error + except ServiceUnavailable as e: + grpc_code = grpc.StatusCode.UNAVAILABLE + + span.set_tag(ERROR_MSG, str(e)) + + raise e + finally: + http_code = _GRPC_HTTP_CODE_TRANSLATIONS[grpc_code] + + span.set_tag(GRPC_STATUS_CODE_KEY, grpc_code) + # Setting the HTTP status code on the tag is a hack to generate more metrics from APM, + # since Datadog automatically generates metrics based on HTTP status but not based on gRPC. + span.set_tag(STATUS_CODE, http_code) + + if http_code >= 500: + span.error = 1 + + # Group the codes so they're easier to use, e.g., in SLOs. + # NOTE: At the time of implementation, tracking the actual status didn't seem super useful. + # If you stumble upon this and feel some reason to add it, by all means do! + if http_code >= 200 and http_code < 300: + tags.append('status_class:ok') + tags.append('http.status_class:2xx') + elif http_code >= 400 and http_code < 500: + tags.append('status_class:client_error') + tags.append('http.status_class:4xx') + # This deliberately captures all unexpected codes (1xx, 3xx, and anything else that might + # somehow occur) as "5xx"/"server_error" so that they can be grouped as server errors + # for free, e.g., for use in an SLO denominator. + # NOTE: Based on the gRPC <> HTTP mapping, such codes should never happen. + else: + tags.append('status_class:server_error') + tags.append('http.status_class:5xx') + + if http_code < 500 or http_code >= 600: + metrics.increment( + 'pigeon.requests.unexpected_code', + tags=[f'http.status_code:{http_code}', f'grpc.status_code:{grpc_code}'], + ) + + # In most cases, the generated "HTTP" metrics mentioned above should be sufficient to track + # the availability of a gRPC endpoint from the client's perspective - unless we also want to + # directly account for latency, i.e., whether or not the response was fast enough. + # We only explicitly track metrics independent of APM in these cases, as otherwise the data + # is effectively redundant and we don't want to unnecessarily create a zillion unnecessary + # but expensive distinct metrics. + if acceptable_duration_ms: + metrics.increment('pigeon.requests', tags=tags) + + def _to_proto(self, message: Request) -> Message: + return cast(Message, message) + + def _from_proto(self, message: Message) -> Response: + return cast(Response, message) + + +# Keep the sync name as an alias for backward compatibility in type annotations +UnaryUnaryRpcCallable = AsyncUnaryUnaryRpcCallable + + +def set_protocol(span): + if span: + span.set_tag('protocol', 'grpc') + + +def maybe_start_span(name: str, service: str, resource: str) -> Span: + existing_span = current_span() + if existing_span.name == name: + # Add tags to the existing span to indicate the inner service and resource that would have been recorded + # if we had started a new span (otherwise we'd just lose this information). + existing_span.set_tag('pigeon.request.subsumed', True) + existing_span.set_tag('pigeon.request.service', service) + existing_span.set_tag('pigeon.request.resource', resource) + + # Don't start a new span if we're already in a pigeon. span + # This is useful for the Service RPC Guard, which cheats by using the same span name to subsume ownership of + # the whole request, including the pigeon.request span. + # But don't return the parent span because we don't want close it at the end of the 'with' block + return noop_span() + # No existing span, so start a new one + return trace(name, service=service, resource=resource) diff --git a/osprey_async_worker/src/osprey/async_worker/lib/pigeon/exceptions.py b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/exceptions.py new file mode 100644 index 0000000..5274c8e --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/exceptions.py @@ -0,0 +1,68 @@ +from typing import Optional, cast + +import grpc + + +class PigeonException(Exception): + pass + + +class InvalidRoutingValueException(PigeonException): + def __init__(self, field: str, value): + self.field = field + self.value = value + + def __str__(self): + return f'Field {self.field} is empty. value={self.value}' + + +class NoResponsesException(PigeonException): + pass + + +class RPCException(PigeonException, grpc.RpcError, grpc.Call): # type: ignore[misc] + # https://grpc.github.io/grpc/python/grpc.html#grpc.Call + def __init__(self, service: str, method: str, error: grpc.RpcError): + # grpc.RpcError is always a grpc.Call + inner = cast(grpc.Call, error) + message = f'{service}.{method} failed. status={inner.code()} details="{inner.details()}" inner={inner}' + super().__init__(message) + self._service = service + self._method = method + self._inner = inner + + @property + def service(self) -> str: + return self._service + + @property + def method(self) -> str: + return self._method + + @property + def inner(self) -> grpc.Call: + return self._inner + + def initial_metadata(self): + return self.inner.initial_metadata() + + def trailing_metadata(self): + return self.inner.trailing_metadata() + + def code(self) -> grpc.StatusCode: + return cast(grpc.StatusCode, self.inner.code()) + + def details(self) -> str: + return cast(str, self.inner.details()) + + def is_active(self) -> bool: + return cast(bool, self.inner.is_active()) + + def time_remaining(self) -> Optional[float]: + return self.inner.time_remaining() + + def cancel(self): + return self.inner.cancel() + + def add_callback(self, callback): + return self.inner.add_callback(callback) diff --git a/osprey_async_worker/src/osprey/async_worker/lib/pigeon/interceptors/__init__.py b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/interceptors/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/lib/pigeon/interceptors/baggage.py b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/interceptors/baggage.py new file mode 100644 index 0000000..99ceb05 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/interceptors/baggage.py @@ -0,0 +1,74 @@ +import collections +from typing import Dict, Optional + +import grpc +from osprey.worker.lib.ddtrace_utils import ( + baggage_propagator, + current_span, + get_baggage, +) + +BAGGAGE_HEADER = 'baggage' + + +class BaggageInterceptor( + grpc.UnaryUnaryClientInterceptor, # type: ignore[misc] + grpc.UnaryStreamClientInterceptor, # type: ignore[misc] + grpc.StreamUnaryClientInterceptor, # type: ignore[misc] + grpc.StreamStreamClientInterceptor, # type: ignore[misc] +): + """ + Propagates tracing "baggage" to downstream grpc services. + + This is only a partial, simple implementation of the W3C spec, skipping checks + such as limits on the maximum number of baggage keys and the lengths of keys + or values. + + All baggage keys and headers should already be strings, as Datadog expects + tags to all be pre-stringified. + + For more info: + https://opentelemetry.io/docs/reference/specification/baggage/api/ + """ + + def __init__(self, baggage_header: Optional[str] = None, baggage: Optional[Dict[str, str]] = None): + self._baggage_header = baggage_header or BAGGAGE_HEADER + self._baggage = baggage or {} + + def intercept_unary_unary(self, continuation, client_call_details, request): + return self._intercept_call(continuation, client_call_details, request) + + def intercept_unary_stream(self, continuation, client_call_details, request): + return self._intercept_call(continuation, client_call_details, request) + + def intercept_stream_unary(self, continuation, client_call_details, request_iterator): + return self._intercept_call(continuation, client_call_details, request_iterator) + + def intercept_stream_stream(self, continuation, client_call_details, request_iterator): + return self._intercept_call(continuation, client_call_details, request_iterator) + + def _intercept_call(self, continuation, client_call_details, request_or_iterator): + # The provided client_call_details may a) be a namedtuple and b) not yet have metadata + # (i.e., client_call_details.metadata == None). + # Since namedtuple fields are immutable, we need to provide our own namedtuple wrapper + # in order to handle the case in which metadata is not yet set. + metadata = list(client_call_details.metadata) if client_call_details.metadata else [] + + # The baggage propagator expects a header-style dictionary but python's grpc metadata + # is a list of tuples, so we need to convert + baggage_headers = {} + baggage_propagator.inject(get_baggage(current_span()), baggage_headers) + metadata.extend([(k, v) for k, v in baggage_headers.items()]) + + client_call_details = _ClientCallDetails( + client_call_details.method, client_call_details.timeout, metadata, client_call_details.credentials + ) + + return continuation(client_call_details, request_or_iterator) + + +class _ClientCallDetails( + collections.namedtuple('_ClientCallDetails', ('method', 'timeout', 'metadata', 'credentials')), + grpc.ClientCallDetails, # type: ignore[misc] +): + pass diff --git a/osprey_async_worker/src/osprey/async_worker/lib/pigeon/interceptors/metadata.py b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/interceptors/metadata.py new file mode 100644 index 0000000..de0d9e4 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/interceptors/metadata.py @@ -0,0 +1,49 @@ +import collections +from typing import Dict, Text + +import grpc + + +class MetadataInterceptor( + grpc.UnaryUnaryClientInterceptor, # type: ignore[misc] + grpc.UnaryStreamClientInterceptor, # type: ignore[misc] + grpc.StreamUnaryClientInterceptor, # type: ignore[misc] + grpc.StreamStreamClientInterceptor, # type: ignore[misc] +): + def __init__(self, metadata: Dict[Text, Text]): + self._metadata = metadata + + def intercept_unary_unary(self, continuation, client_call_details, request): + return self._intercept_call(continuation, client_call_details, request) + + def intercept_unary_stream(self, continuation, client_call_details, request): + return self._intercept_call(continuation, client_call_details, request) + + def intercept_stream_unary(self, continuation, client_call_details, request_iterator): + return self._intercept_call(continuation, client_call_details, request_iterator) + + def intercept_stream_stream(self, continuation, client_call_details, request_iterator): + return self._intercept_call(continuation, client_call_details, request_iterator) + + def _intercept_call(self, continuation, client_call_details, request_or_iterator): + # The provided client_call_details may a) be a namedtuple and b) not yet have metadata + # (i.e., client_call_details.metadata == None). + # Since namedtuple fields are immutable, we need to provide our own namedtuple wrapper + # in order to handle the case in which metadata is not yet set. + metadata = list(client_call_details.metadata) if client_call_details.metadata else [] + + for key, value in self._metadata.items(): + metadata.insert(0, (key, value)) + + client_call_details = _ClientCallDetails( + client_call_details.method, client_call_details.timeout, metadata, client_call_details.credentials + ) + + return continuation(client_call_details, request_or_iterator) + + +class _ClientCallDetails( + collections.namedtuple('_ClientCallDetails', ('method', 'timeout', 'metadata', 'credentials')), + grpc.ClientCallDetails, # type: ignore[misc] +): + pass diff --git a/osprey_async_worker/src/osprey/async_worker/lib/pigeon/skip_rate_limit.py b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/skip_rate_limit.py new file mode 100644 index 0000000..d29fa93 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/pigeon/skip_rate_limit.py @@ -0,0 +1,19 @@ +from contextvars import ContextVar + +skip_rate_limit: ContextVar[bool] = ContextVar('skip_rate_limit', default=False) + + +class SkipRateLimitContext: + """Provides the same attribute-based API as the gevent.local version + but backed by a ContextVar so it is safe for use with asyncio tasks.""" + + @property + def skip(self) -> bool: + return skip_rate_limit.get() + + @skip.setter + def skip(self, value: bool) -> None: + skip_rate_limit.set(value) + + +skip_rate_limit_context = SkipRateLimitContext() diff --git a/osprey_async_worker/src/osprey/async_worker/lib/publisher.py b/osprey_async_worker/src/osprey/async_worker/lib/publisher.py new file mode 100644 index 0000000..f6ebb2f --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/lib/publisher.py @@ -0,0 +1,147 @@ +"""Async Pub/Sub publisher with batching for the async worker. + +Uses asyncio.Queue for buffering and a background task for periodic +flushing. No threading locks, no gevent, no background threads. +""" + +import asyncio +import logging +from typing import List, Optional + +from google.api_core.retry import Retry +from google.cloud import pubsub_v1 +from osprey.worker.lib.instruments import metrics +from pydantic import BaseModel + +logger = logging.getLogger(__name__) + +# Retry policy passed to PublisherClient.publish(). Google's Retry already +# classifies transient gRPC codes (UNAVAILABLE, DEADLINE_EXCEEDED, +# RESOURCE_EXHAUSTED, ABORTED, INTERNAL) and TimeoutError as retryable. +# The 30-second deadline gives the internal retry loop enough headroom for a +# few backoff attempts before we give up. The future.result() timeout below +# is set slightly above deadline so the Retry loop, not the wall-clock cap, +# decides when to stop. +_PUBLISH_RETRY = Retry(initial=0.5, maximum=10.0, multiplier=2.0, deadline=30.0) + + +class AsyncPubSubPublisher: + """Publishes Pydantic models to a Pub/Sub topic with async batching. + + Messages are buffered in an asyncio.Queue and flushed either when + the batch reaches max_messages or after max_latency_seconds. + """ + + def __init__( + self, + project_id: str, + topic_id: str, + max_messages: int = 250, + max_latency_seconds: float = 1.0, + ): + self._topic_path = f'projects/{project_id}/topics/{topic_id}' + self._client = pubsub_v1.PublisherClient( + batch_settings=pubsub_v1.types.BatchSettings(max_messages=1), + ) + self._queue: asyncio.Queue[bytes] = asyncio.Queue() + self._max_messages = max_messages + self._max_latency = max_latency_seconds + self._flush_task: Optional[asyncio.Task[None]] = None + self._started = False + self._metric_tags = [f'project:{project_id}', f'topic:{topic_id}'] + + def _ensure_started(self) -> None: + """Start the background flush task on first publish.""" + if not self._started: + self._started = True + try: + loop = asyncio.get_running_loop() + self._flush_task = loop.create_task(self._flush_loop()) + except RuntimeError: + # No event loop. Messages will be enqueued but never flushed — + # surface that explicitly so dashboards can catch it. + metrics.increment('async_pubsub_publisher.no_event_loop', tags=self._metric_tags) + + async def _flush_loop(self) -> None: + """Background task that flushes the buffer periodically.""" + while True: + try: + batch: List[bytes] = [] + + # Wait for first message or timeout + try: + msg = await asyncio.wait_for(self._queue.get(), timeout=self._max_latency) + batch.append(msg) + except TimeoutError: + continue + + # Drain up to max_messages + while len(batch) < self._max_messages: + try: + msg = self._queue.get_nowait() + batch.append(msg) + except asyncio.QueueEmpty: + break + + if batch: + await self._flush_batch(batch) + + except asyncio.CancelledError: + # Flush remaining on shutdown + remaining: List[bytes] = [] + while not self._queue.empty(): + try: + remaining.append(self._queue.get_nowait()) + except asyncio.QueueEmpty: + break + if remaining: + await self._flush_batch(remaining) + return + except Exception: + logger.exception('Error in flush loop') + + async def _flush_batch(self, batch: List[bytes]) -> None: + """Publish a batch of messages. Runs sync publishes in executor.""" + loop = asyncio.get_running_loop() + await loop.run_in_executor(None, self._sync_flush, batch) + + def _sync_flush(self, batch: List[bytes]) -> None: + """Synchronous batch publish.""" + futures = [] + for data in batch: + metrics.increment('async_pubsub_publisher.publish.attempt', tags=self._metric_tags) + futures.append(self._client.publish(self._topic_path, data, retry=_PUBLISH_RETRY)) + for future in futures: + try: + # deadline=30s above; 35s here ensures Retry's own deadline, + # not this wall-clock cap, is what terminates failed attempts. + future.result(timeout=35) + metrics.increment('async_pubsub_publisher.publish.success', tags=self._metric_tags) + except Exception as e: + logger.exception('Failed to publish message') + metrics.increment( + 'async_pubsub_publisher.publish.failure', + tags=self._metric_tags + [f'error:{e.__class__.__name__}'], + ) + + def publish(self, data: BaseModel) -> None: + """Queue a Pydantic model for async batched publishing.""" + self.publish_bytes(data.json(exclude_none=True).encode()) + + def publish_bytes(self, data: bytes) -> None: + """Queue raw bytes for async batched publishing.""" + self._ensure_started() + try: + self._queue.put_nowait(data) + except asyncio.QueueFull: + logger.warning('Publisher queue full, dropping message') + metrics.increment('async_pubsub_publisher.queue_full', tags=self._metric_tags) + + async def stop(self) -> None: + """Flush remaining messages and stop.""" + if self._flush_task is not None: + self._flush_task.cancel() + try: + await self._flush_task + except asyncio.CancelledError: + pass diff --git a/osprey_async_worker/src/osprey/async_worker/lib/utils/__init__.py b/osprey_async_worker/src/osprey/async_worker/lib/utils/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/singletons.py b/osprey_async_worker/src/osprey/async_worker/singletons.py new file mode 100644 index 0000000..ef387b3 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/singletons.py @@ -0,0 +1,77 @@ +"""Singletons for the async worker. + +Mirrors osprey.worker.lib.singletons but for the async engine. Services that +migrate from the gevent engine to the async engine often rely on a process-wide +ENGINE.instance() accessor; this module supplies the async equivalent. + +The gevent ENGINE singleton lives at osprey.worker.lib.singletons.ENGINE and +holds an OspreyEngine. This module's ENGINE holds an AsyncOspreyEngine and is +distinct — they share the underlying execution graph machinery, but each owns +its own compile thread pool, sources provider, and config subkey state. +""" + +from pathlib import Path + +from osprey.async_worker.adaptor.plugin_manager import ( + bootstrap_async_ast_validators, + bootstrap_async_udfs, + bootstrap_validation_exporter, +) +from osprey.async_worker.engine import AsyncOspreyEngine +from osprey.engine.ast.sources import Sources +from osprey.worker.lib.singleton import Singleton +from osprey.worker.lib.singletons import CONFIG +from osprey.worker.lib.sources_provider_base import BaseSourcesProvider, StaticSourcesProvider + + +def _resolve_sources_provider() -> BaseSourcesProvider: + """Build a sources provider from config. + + Mirrors osprey.worker.lib.osprey_engine.get_sources_provider. If + OSPREY_RULES_PATH is set, returns a StaticSourcesProvider; otherwise + returns a (sync) EtcdSourcesProvider keyed by OSPREY_ETCD_SOURCES_PROVIDER_KEY. + + Note: we reuse the *sync* EtcdSourcesProvider rather than the async one + because the factory is called synchronously on first ENGINE.instance(). + The async etcd provider requires `await provider.start()`, which can't + run from a sync singleton initializer. For services that need the async + etcd provider with a live watch (e.g. the async worker itself), construct + AsyncOspreyEngine directly and don't go through this singleton. + """ + config = CONFIG.instance() + rules_path_str = config.get_optional_str('OSPREY_RULES_PATH') + if rules_path_str: + return StaticSourcesProvider(sources=Sources.from_path(Path(rules_path_str))) + + # EtcdSourcesProvider is imported locally because its module chain + # (osprey.worker.lib.etcd.*) does `import gevent`. Keeping it off the + # module-level import set means consumers that only ever hit the + # OSPREY_RULES_PATH branch never pull gevent into sys.modules. + from osprey.worker.lib.sources_provider import EtcdSourcesProvider + + etcd_key = config.get_str('OSPREY_ETCD_SOURCES_PROVIDER_KEY', '/config/osprey/rules-sink-sources') + return EtcdSourcesProvider(etcd_key=etcd_key) + + +def _init_engine() -> AsyncOspreyEngine: + """Factory for the ENGINE singleton. + + Bootstraps async UDFs + AST validators, builds a sources provider from + config, and constructs an AsyncOspreyEngine. Mirrors the gevent + bootstrap_engine_with_helpers wiring. + """ + config = CONFIG.instance() + udf_registry, _udf_helpers = bootstrap_async_udfs(config=config) + bootstrap_async_ast_validators() + + validation_exporter = bootstrap_validation_exporter(config) + + return AsyncOspreyEngine( + sources_provider=_resolve_sources_provider(), + udf_registry=udf_registry, + should_yield_during_compilation=config.get_bool('OSPREY_PERIODIC_YIELD_DURING_COMPILATION', False), + validation_exporter=validation_exporter, + ) + + +ENGINE: Singleton[AsyncOspreyEngine] = Singleton(_init_engine) diff --git a/osprey_async_worker/src/osprey/async_worker/sinks/__init__.py b/osprey_async_worker/src/osprey/async_worker/sinks/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/sinks/sink/__init__.py b/osprey_async_worker/src/osprey/async_worker/sinks/sink/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/sinks/sink/input_stream.py b/osprey_async_worker/src/osprey/async_worker/sinks/sink/input_stream.py new file mode 100644 index 0000000..a189c45 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/sinks/sink/input_stream.py @@ -0,0 +1,139 @@ +"""Async input streams for the async worker.""" + +import abc +import asyncio +import functools +import json +import logging +from collections import deque +from concurrent.futures import ThreadPoolExecutor +from typing import Any, AsyncIterator, Callable, Generic, Mapping, Protocol, Sequence, TypeVar + +from osprey.engine.executor.execution_context import Action +from osprey.worker.lib.utils.dates import parse_go_timestamp +from osprey.worker.sinks.utils.acking_contexts_base import BaseAckingContext, NoopAckingContext + +_T = TypeVar('_T') + +logger = logging.getLogger(__name__) + + +class AsyncBaseInputStream(abc.ABC, Generic[_T]): + """Async version of BaseInputStream. Uses async iteration.""" + + def __aiter__(self) -> AsyncIterator[_T]: + return self._gen() + + @abc.abstractmethod + async def _gen(self) -> AsyncIterator[_T]: + raise NotImplementedError + yield # make this an async generator + + async def stop(self) -> None: + pass + + +class AsyncStaticInputStream(AsyncBaseInputStream[_T]): + """An async input stream that returns a static list, until exhausted. For testing.""" + + def __init__(self, items: Sequence[_T]): + self._items = deque(items) + + async def _gen(self) -> AsyncIterator[_T]: + for item in self._items: + yield item + + +class _PollableConsumer(Protocol): + """The slice of kafka.KafkaConsumer this stream needs. + + Declaring it as a Protocol keeps this module free of a hard `kafka` import + (the async worker forbids gevent, and kafka-python's patched consumer pulls + it in) and makes the stream trivially testable with a fake consumer. + """ + + def poll(self, timeout_ms: int = ..., max_records: int = ...) -> Mapping[Any, Sequence[Any]]: ... + + def commit(self) -> None: ... + + def close(self) -> None: ... + + +class AsyncKafkaInputStream(AsyncBaseInputStream[BaseAckingContext[Action]]): + """Consume Osprey-format actions from Kafka on the asyncio event loop. + + Decodes the same envelope the gevent ``KafkaInputStream`` reads: + ``{"send_time": ..., "data": {"action_id", "action_name", "data": {...}}}``. + + kafka-python's consumer is blocking, so each poll runs in a worker thread + via ``asyncio.to_thread``; the loop stays responsive and the bounded poll + timeout lets ``stop()`` interrupt between polls. The consumer is injected so + this module needs neither the kafka client nor any gevent-tainted helper. + + Offsets are committed manually, once per polled batch, only after every + record in that batch has been handed back to (and processed by) the + consumer of this generator. With ``enable_auto_commit=False`` this gives + at-least-once delivery: a crash mid-batch reprocesses the batch rather than + losing actions to an offset that was auto-committed before processing. + """ + + def __init__( + self, + consumer: _PollableConsumer, + poll_timeout_ms: int = 1000, + max_records: int = 500, + ): + self._consumer = consumer + self._poll_timeout_ms = poll_timeout_ms + self._max_records = max_records + self._stopped = False + # kafka-python consumers are NOT safe for concurrent calls. Funnel every + # consumer operation (poll/commit/close) through a single worker thread so + # stop()'s close() can never race an in-flight poll() on another thread. + self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix='async-kafka-consumer') + + async def _run_on_consumer(self, fn: Callable[[], Any]) -> Any: + return await asyncio.get_running_loop().run_in_executor(self._executor, fn) + + @staticmethod + def decode(record_value: Any) -> Action: + """Decode a Kafka record value (bytes/str) into an Action.""" + data = json.loads(record_value) + action_data = data['data'] + return Action( + action_id=int(action_data['action_id']), + action_name=action_data['action_name'], + data=action_data['data'], + timestamp=parse_go_timestamp(data['send_time']), + ) + + async def _gen(self) -> AsyncIterator[BaseAckingContext[Action]]: + while not self._stopped: + batches = await self._run_on_consumer( + functools.partial( + self._consumer.poll, + timeout_ms=self._poll_timeout_ms, + max_records=self._max_records, + ) + ) + for records in batches.values(): + for record in records: + try: + action = self.decode(record.value) + except Exception: + logger.exception('Error decoding Kafka record; skipping') + continue + yield NoopAckingContext(action) + # Reached only after the generator resumes past the last yielded + # record of this batch, i.e. once the batch has been processed. + if batches and not self._stopped: + await self._run_on_consumer(self._consumer.commit) + + async def stop(self) -> None: + self._stopped = True + # close() is queued on the single worker thread, so it runs only after any + # in-flight poll/commit returns — never concurrently with them. + try: + await self._run_on_consumer(self._consumer.close) + finally: + self._executor.shutdown(wait=False) diff --git a/osprey_async_worker/src/osprey/async_worker/sinks/sink/output_sink.py b/osprey_async_worker/src/osprey/async_worker/sinks/sink/output_sink.py new file mode 100644 index 0000000..f922ab4 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/sinks/sink/output_sink.py @@ -0,0 +1,74 @@ +"""Async output sink with timeout and retry support.""" + +import asyncio +import logging +from typing import Sequence + +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.engine.executor.execution_context import ExecutionResult +from osprey.worker.lib.instruments import metrics + +logger = logging.getLogger(__name__) + + +class AsyncMultiOutputSink(AsyncBaseOutputSink): + """Tees execution results to multiple async output sinks with timeout and retry.""" + + def __init__(self, sinks: Sequence[AsyncBaseOutputSink]): + self._sinks = sinks + + def will_do_work(self, result: ExecutionResult) -> bool: + return any(sink.will_do_work(result) for sink in self._sinks) + + async def push(self, result: ExecutionResult) -> None: + tasks = [] + for sink in self._sinks: + if sink.will_do_work(result): + tasks.append(self._push_one(sink, result)) + if tasks: + await asyncio.gather(*tasks) + + async def _push_one(self, sink: AsyncBaseOutputSink, result: ExecutionResult) -> None: + """Push to a single sink with timeout and retry. Runs concurrently via gather().""" + sink_name = sink.__class__.__name__ + attempts = sink.max_retries + 1 + + for attempt in range(1, attempts + 1): + try: + start = asyncio.get_running_loop().time() + async with asyncio.timeout(sink.timeout): + await sink.push(result) + metrics.timing( + 'handled_message_output', + (asyncio.get_running_loop().time() - start) * 1000, + tags=[f'sink:{sink_name}'], + ) + break + except TimeoutError: + logger.warning(f'Timeout pushing to {sink_name} (attempt {attempt}/{attempts})') + metrics.increment('output_sink.timeout', tags=[f'sink:{sink_name}']) + if attempt == attempts: + metrics.increment('output_sink.timeout_exhausted', tags=[f'sink:{sink_name}']) + except Exception as exc: + logger.exception(f'Error pushing to {sink_name}: {exc}') + metrics.increment('output_sink.error', tags=[f'sink:{sink_name}', f'error:{exc.__class__.__name__}']) + if attempt == attempts: + break + await asyncio.sleep(0.5 * attempt) + + async def stop(self) -> None: + for sink in self._sinks: + await sink.stop() + + +class AsyncStdoutOutputSink(AsyncBaseOutputSink): + """Debug output sink that prints to stdout.""" + + def will_do_work(self, result: ExecutionResult) -> bool: + return True + + async def push(self, result: ExecutionResult) -> None: + logger.info(f'result: {result.extracted_features_json} {result.verdicts}') + + async def stop(self) -> None: + pass diff --git a/osprey_async_worker/src/osprey/async_worker/sinks/sink/rules_sink.py b/osprey_async_worker/src/osprey/async_worker/sinks/sink/rules_sink.py new file mode 100644 index 0000000..fc72e9d --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/sinks/sink/rules_sink.py @@ -0,0 +1,167 @@ +"""Async rules sink — the main processing loop for the async worker.""" + +import asyncio +import logging +from dataclasses import dataclass +from random import randint +from typing import Optional + +import sentry_sdk +from ddtrace import tracer +from ddtrace.span import Span as TracerSpan +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.async_worker.engine import AsyncOspreyEngine +from osprey.async_worker.executor import execute as async_execute +from osprey.async_worker.sinks.sink.input_stream import AsyncBaseInputStream +from osprey.engine.executor.execution_context import Action, ExecutionResult +from osprey.engine.executor.udf_execution_helpers import UDFHelpers +from osprey.worker.lib.instruments import metrics +from osprey.worker.lib.osprey_shared.logging import info_log_osprey_action +from osprey.worker.lib.snowflake import generate_snowflake +from osprey.worker.lib.sources_config.subkeys.action_config import ActionConfigs +from osprey.worker.sinks.utils.acking_contexts_base import BaseAckingContext, VerdictsAckingContext + +logger = logging.getLogger(__name__) + + +@dataclass +class SampleDecision: + sample_rate: int + drop: bool + + +_SAMPLE_NEVER = SampleDecision(sample_rate=100, drop=False) +_SAMPLE_ALWAYS = SampleDecision(sample_rate=0, drop=True) + + +class ActionSampler: + """Checks whether an action should be sampled. No gevent dependency.""" + + def __init__(self, engine: AsyncOspreyEngine): + self._engine = engine + + def sample(self, action: Action) -> SampleDecision: + action_configs = self._engine.get_config_subkey(ActionConfigs) + action_config = action_configs.get_action_config(action.action_name) + + if not action_config or action_config.sample_rate == 100: + return _SAMPLE_NEVER + if action_config.sample_rate == 0: + return _SAMPLE_ALWAYS + + p = randint(0, 99) + should_drop = p < action_config.sample_rate + return SampleDecision(sample_rate=action_config.sample_rate, drop=should_drop) + + +class AsyncRulesRunner: + """Async version of RulesRunner — classifies one action and pushes to output sink.""" + + def __init__( + self, + engine: AsyncOspreyEngine, + output_sink: AsyncBaseOutputSink, + udf_helpers: UDFHelpers, + max_concurrent_udfs: int = 12, + ) -> None: + self._engine = engine + self._sampler = ActionSampler(engine) + self._output_sink = output_sink + self._udf_helpers = udf_helpers + self._max_concurrent_udfs = max_concurrent_udfs + + async def classify_one( + self, + action: Action, + tag: str, + parent_tracer_span: Optional[TracerSpan] = None, + ) -> Optional[ExecutionResult]: + sample_config = self._sampler.sample(action) + tags = [ + tag, + f'action:{action.action_name}', + f'sample_rate:{sample_config.sample_rate}', + f'rules_hash:{self._engine.execution_graph.validated_sources.sources.hash()}', + ] + + if sample_config.drop: + metrics.increment('dropped_message', tags=tags) + return None + + result: Optional[ExecutionResult] = None + try: + with metrics.timed('handled_message', tags=tags, use_ms=True): + result = await async_execute( + self._engine.execution_graph, + self._udf_helpers, + action, + max_concurrent=self._max_concurrent_udfs, + sample_rate=sample_config.sample_rate, + parent_tracer_span=parent_tracer_span, + ) + with metrics.timed('handled_output', tags=tags, use_ms=True): + await self._output_sink.push(result) + info_log_osprey_action(action.action_id, action.action_name, 'pushed to output sink') + return result + except Exception: + logging.exception('Error in classify_one for action %s', action.action_name) + metrics.increment('rules_runner.classify_error', tags=tags) + sentry_sdk.capture_exception() + return result + + +class AsyncRulesSink: + """Async rules sink — iterates an async input stream, executes rules, pushes to output sinks.""" + + def __init__( + self, + engine: AsyncOspreyEngine, + input_stream: AsyncBaseInputStream[BaseAckingContext[Action]], + output_sink: AsyncBaseOutputSink, + udf_helpers: UDFHelpers, + max_concurrent_udfs: int = 12, + ): + self._input_stream = input_stream + self._rules_runner = AsyncRulesRunner(engine, output_sink, udf_helpers, max_concurrent_udfs) + + async def run(self) -> None: + async for message_context in self._input_stream: + try: + with message_context as action: + action_tags = [f'action:{action.action_name}'] + metrics.increment('rules_sink.input_action_received', tags=action_tags) + + if action.data.get('osprey_skip_async', False): + metrics.increment('rules_sink.skipped', tags=action_tags) + continue + + with tracer.start_span('osprey.async.classify_one', child_of=None) as span: + tracer.context_provider.activate(span.context) + + if not action.action_id and action.action_id != 0: + action.action_id = generate_snowflake(retries=3).to_int() + + info_log_osprey_action(action.action_id, action.action_name, 'beginning async classify_one') + result = await self._rules_runner.classify_one( + action, + tag='sink:async-rules-sink', + parent_tracer_span=span, + ) + + if isinstance(message_context, VerdictsAckingContext): + if result is None: + metrics.increment('rules_sink.missing_result') + else: + message_context.set_verdicts(result.get_verdicts_pb2_proto()) + metrics.increment('rules_sink.captured_verdicts') + + info_log_osprey_action(action.action_id, action.action_name, 'async classify_one complete') + except asyncio.CancelledError: + return + except Exception as e: + logging.exception('Unexpected error in async rules sink') + metrics.increment('rules_sink.unexpected_error', tags=[f'err:{e.__class__.__name__}']) + sentry_sdk.capture_exception(e) + + async def stop(self) -> None: + await self._input_stream.stop() diff --git a/osprey_async_worker/src/osprey/async_worker/stdlib_udfs/__init__.py b/osprey_async_worker/src/osprey/async_worker/stdlib_udfs/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/stdlib_udfs/_async_stdlib_plugin.py b/osprey_async_worker/src/osprey/async_worker/stdlib_udfs/_async_stdlib_plugin.py new file mode 100644 index 0000000..aa68bd1 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/stdlib_udfs/_async_stdlib_plugin.py @@ -0,0 +1,25 @@ +"""First-party async-stdlib plugin. + +osprey_async_worker registers itself as an internal pluggy plugin so that +async-native replacements for sync stdlib UDFs (e.g. AsyncMXLookup) flow +through the same `register_udfs` hook used by third-party async plugins. +This keeps `bootstrap_async_udfs` free of hardcoded override lists — adding +a new async stdlib override means appending a class here. + +Override-by-class-name is handled by `_deduplicate_udfs` in plugin_manager: +each class returned here shadows any sync stdlib UDF with the same +`__name__`. +""" + +from __future__ import annotations + +from typing import Any, Sequence, Type + +from osprey.async_worker.adaptor.plugin_manager import hookimpl_osprey_async +from osprey.async_worker.stdlib_udfs.async_mx_lookup import MXLookup +from osprey.engine.udf.base import UDFBase + + +@hookimpl_osprey_async +def register_udfs() -> Sequence[Type[UDFBase[Any, Any]]]: + return [MXLookup] diff --git a/osprey_async_worker/src/osprey/async_worker/stdlib_udfs/async_mx_lookup.py b/osprey_async_worker/src/osprey/async_worker/stdlib_udfs/async_mx_lookup.py new file mode 100644 index 0000000..666c835 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/stdlib_udfs/async_mx_lookup.py @@ -0,0 +1,60 @@ +"""Native async MXLookup UDF using aiodns. + +Replaces the sync MXLookup which uses blocking dns.resolver calls. +Uses aiodns (c-ares) for fully async DNS resolution on the event loop +without consuming thread pool threads. + +Class named `MXLookup` to shadow the sync version in the UDF registry. +""" + +import asyncio +from typing import Optional + +import aiodns +import pycares +from osprey.async_worker.adaptor.interfaces import AsyncUDFBase +from osprey.engine.executor.execution_context import ExecutionContext, ExpectedUdfException +from osprey.engine.stdlib.udfs.mx_lookup import Arguments +from osprey.engine.stdlib.udfs.mx_lookup import MXLookup as SyncMXLookup + +_DNS_TIMEOUT = 5.0 +_resolver: Optional[aiodns.DNSResolver] = None + + +def _get_resolver() -> aiodns.DNSResolver: + """Lazily create the resolver on the running event loop.""" + global _resolver + loop = asyncio.get_running_loop() + if _resolver is None or _resolver.loop is not loop: + _resolver = aiodns.DNSResolver(timeout=_DNS_TIMEOUT, loop=loop) + return _resolver + + +class MXLookup(AsyncUDFBase[Arguments, str]): # type: ignore[misc] + """Async MXLookup — uses aiodns for non-blocking DNS resolution.""" + + category = SyncMXLookup.category + + @classmethod + def _get_udf_base_args(cls): + return (Arguments, str) + + async def async_execute(self, execution_context: ExecutionContext, arguments: Arguments) -> str: + resolver = _get_resolver() + try: + mx_result = await resolver.query_dns(arguments.domain, 'MX') + mx_records = [r for r in mx_result.answer if hasattr(r.data, 'priority')] + if not mx_records: + raise ExpectedUdfException() + # hasattr filter above guarantees these are MX records; pycares' record + # union doesn't narrow on hasattr, so suppress union-attr here. + best_mx = sorted(mx_records, key=lambda r: r.data.priority)[0].data.exchange # type: ignore[union-attr] + a_result = await resolver.query_dns(best_mx, 'A') + except (aiodns.error.DNSError, pycares.AresError): + raise ExpectedUdfException() + + a_records = [r for r in a_result.answer if hasattr(r.data, 'addr')] + if not a_records: + raise ExpectedUdfException() + # hasattr filter guarantees A records; pycares union doesn't narrow on hasattr. + return min(r.data.addr for r in a_records) # type: ignore[union-attr] diff --git a/osprey_async_worker/src/osprey/async_worker/tests/__init__.py b/osprey_async_worker/src/osprey/async_worker/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/osprey_async_worker/src/osprey/async_worker/tests/conftest.py b/osprey_async_worker/src/osprey/async_worker/tests/conftest.py new file mode 100644 index 0000000..071836c --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/conftest.py @@ -0,0 +1,137 @@ +"""Shared fixtures for async worker tests. + +Provides async equivalents of the osprey engine conftest fixtures, +using the async executor instead of the gevent one. +""" + +from datetime import datetime +from textwrap import dedent +from typing import Dict, Optional, Union + +import pytest +from osprey.async_worker.executor import execute as async_execute +from osprey.engine.ast.sources import SOURCE_ENTRY_POINT_PATH, Sources +from osprey.engine.ast_validator import validate_sources +from osprey.engine.ast_validator.validation_context import ValidationFailed +from osprey.engine.ast_validator.validator_registry import ValidatorRegistry +from osprey.engine.executor.execution_context import Action, ExecutionResult +from osprey.engine.executor.execution_graph import compile_execution_graph +from osprey.engine.executor.udf_execution_helpers import UDFHelpers +from osprey.engine.stdlib import get_config_registry +from osprey.engine.udf.registry import UDFRegistry +from osprey.worker.lib.singletons import CONFIG + +SourcesDict = Union[Sources, str, Dict[str, str]] + + +def _into_sources(sources_dict: SourcesDict) -> Sources: + if isinstance(sources_dict, Sources): + return sources_dict + if isinstance(sources_dict, str): + sources_dict = {SOURCE_ENTRY_POINT_PATH: sources_dict} + for k, v in sources_dict.items(): + sources_dict[k] = dedent(v) + return Sources.from_dict(sources_dict) + + +@pytest.fixture(autouse=True) +def config_setup(): + CONFIG.instance().configure_from_env() + yield + CONFIG.instance().unconfigure_for_tests() + + +@pytest.fixture() +def stdlib_udf_registry() -> UDFRegistry: + """UDF registry with stdlib UDFs only.""" + from osprey.worker._stdlibplugin.udf_register import register_udfs + from osprey.worker._stdlibplugin.validator_regsiter import register_ast_validators + + # Register standard validators (needed for compile_execution_graph) + registry = ValidatorRegistry.get_instance() + for validator in register_ast_validators(): + registry.register_to_instance(validator) + + return UDFRegistry.with_udfs(*register_udfs()) + + +@pytest.fixture() +def async_execute_with_result(stdlib_udf_registry: UDFRegistry): + """Execute .sml rules using the async executor. Returns full ExecutionResult.""" + + async def _execute( + sources_dict: SourcesDict, + data: Optional[Dict[str, object]] = None, + action_name: str = 'test', + action_id: int = 1, + udf_helpers: Optional[UDFHelpers] = None, + udf_registry: Optional[UDFRegistry] = None, + max_concurrent: int = 12, + action_time: Optional[datetime] = None, + ) -> ExecutionResult: + registry = udf_registry or stdlib_udf_registry + sources = _into_sources(sources_dict) + + config_validator = get_config_registry().get_validator() + validator_registry = ValidatorRegistry.get_instance().instance_with_additional_validators(config_validator) + + try: + validated_sources = validate_sources(sources, registry, validator_registry) + except ValidationFailed as e: + print(e.rendered()) + raise + + execution_graph = compile_execution_graph(validated_sources) + action = Action( + action_id=action_id, + data=data or {}, + action_name=action_name, + timestamp=action_time or datetime.utcnow(), + ) + return await async_execute(execution_graph, udf_helpers or UDFHelpers(), action, max_concurrent=max_concurrent) + + return _execute + + +@pytest.fixture() +def async_execute_fn(async_execute_with_result): + """Execute .sml rules using the async executor. Returns extracted features dict.""" + + async def _execute( + sources_dict: SourcesDict, + data: Optional[Dict[str, object]] = None, + action_name: str = 'test', + action_id: int = 1, + udf_helpers: Optional[UDFHelpers] = None, + udf_registry: Optional[UDFRegistry] = None, + max_concurrent: int = 12, + allow_errors: bool = False, + ) -> Dict[str, object]: + result = await async_execute_with_result( + sources_dict=sources_dict, + data=data, + action_name=action_name, + action_id=action_id, + udf_helpers=udf_helpers, + udf_registry=udf_registry, + max_concurrent=max_concurrent, + ) + if not allow_errors and len(result.error_infos) > 0: + raise result.error_infos[0].error + + features = result.extracted_features + # Remove internal features like the gevent conftest does + for key in [ + '__timestamp', + '__action_id', + '__error_count', + '__sample_rate', + '__entity_label_mutations', + '__classifications', + '__signals', + '__verdicts', + ]: + features.pop(key, None) + return features + + return _execute diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_async_executor.py b/osprey_async_worker/src/osprey/async_worker/tests/test_async_executor.py new file mode 100644 index 0000000..0bc1d6b --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_async_executor.py @@ -0,0 +1,190 @@ +"""Tests for the async executor. + +Validates that the async executor produces the same results as the gevent +executor for stdlib UDFs (pure computation, no I/O). +""" + +import pytest + + +@pytest.mark.asyncio +async def test_execute_sync_udfs(async_execute_fn): + """Sync stdlib UDFs run inline and produce correct results.""" + result = await async_execute_fn( + """ + Msg: str = JsonData(path="$.message", coerce_type=True) + MessageLength = StringLength(s=Msg) + """, + data={'message': 'hello world'}, + ) + assert result['MessageLength'] == 11 + + +@pytest.mark.asyncio +async def test_execute_json_data(async_execute_fn): + """JsonData UDF extracts values from action data.""" + result = await async_execute_fn( + 'Username: str = JsonData(path="$.user.name", coerce_type=True)', + data={'user': {'name': 'alice'}}, + ) + assert result['Username'] == 'alice' + + +@pytest.mark.asyncio +async def test_execute_multiple_udfs(async_execute_fn): + """Multiple UDFs in a single execution graph resolve correctly.""" + result = await async_execute_fn( + """ + Name: str = JsonData(path="$.name", coerce_type=True) + NameLength = StringLength(s=Name) + NameLower = StringToLower(s=Name) + """, + data={'name': 'HELLO'}, + ) + assert result['Name'] == 'HELLO' + assert result['NameLength'] == 5 + assert result['NameLower'] == 'hello' + + +@pytest.mark.asyncio +async def test_execute_dependent_chain(async_execute_fn): + """UDFs with dependencies resolve in correct order.""" + result = await async_execute_fn( + """ + Raw: str = JsonData(path="$.text", coerce_type=True) + Stripped = StringStrip(s=Raw) + Lower = StringToLower(s=Stripped) + Length = StringLength(s=Lower) + """, + data={'text': ' Hello World '}, + ) + assert result['Raw'] == ' Hello World ' + assert result['Stripped'] == 'Hello World' + assert result['Lower'] == 'hello world' + assert result['Length'] == 11 + + +@pytest.mark.asyncio +async def test_execute_with_rules(async_execute_fn): + """Rule evaluation works correctly.""" + result = await async_execute_fn( + """ + Txt: str = JsonData(path="$.text", coerce_type=True) + Length = StringLength(s=Txt) + IsLong = Rule( + when_all=[Length > 10], + description="Text is long", + ) + """, + data={'text': 'short'}, + ) + assert result['Length'] == 5 + assert result['IsLong'] is False + + result = await async_execute_fn( + """ + Txt: str = JsonData(path="$.text", coerce_type=True) + Length = StringLength(s=Txt) + IsLong = Rule( + when_all=[Length > 10], + description="Text is long", + ) + """, + data={'text': 'this is a longer text'}, + ) + assert result['IsLong'] is True + + +@pytest.mark.asyncio +async def test_execute_empty_rules(async_execute_with_result): + """Empty rules produce a valid ExecutionResult with no errors.""" + result = await async_execute_with_result( + '# empty rules file', + data={}, + ) + assert result is not None + assert len(result.error_infos) == 0 + + +@pytest.mark.asyncio +async def test_execute_missing_json_path(async_execute_fn): + """Missing JSON path returns None, not an error.""" + result = await async_execute_fn( + 'Value: str = JsonData(path="$.nonexistent", coerce_type=True)', + data={'something': 'else'}, + allow_errors=True, + ) + assert result['Value'] is None + + +@pytest.mark.asyncio +async def test_execute_sync_only_mode(async_execute_with_result): + """With max_concurrent=0, everything runs synchronously.""" + result = await async_execute_with_result( + """ + Txt: str = JsonData(path="$.text", coerce_type=True) + Value = StringLength(s=Txt) + """, + data={'text': 'test'}, + max_concurrent=0, + ) + assert result.extracted_features['Value'] == 4 + assert len(result.error_infos) == 0 + + +@pytest.mark.asyncio +async def test_execution_result_has_expected_fields(async_execute_with_result): + """ExecutionResult contains all expected fields.""" + result = await async_execute_with_result( + 'Name: str = JsonData(path="$.name", coerce_type=True)', + data={'name': 'test'}, + action_name='test_action', + action_id=42, + ) + assert result.action.action_name == 'test_action' + assert result.action.action_id == 42 + assert '__action_id' in result.extracted_features + assert '__timestamp' in result.extracted_features + assert '__error_count' in result.extracted_features + assert result.extracted_features['__error_count'] == 0 + + +@pytest.mark.asyncio +async def test_string_operations(async_execute_fn): + """Various string UDFs work correctly.""" + result = await async_execute_fn( + """ + Text: str = JsonData(path="$.text", coerce_type=True) + Upper = StringToUpper(s=Text) + StartsWith = StringStartsWith(s=Text, start="hello") + EndsWith = StringEndsWith(s=Text, end="world") + """, + data={'text': 'hello world'}, + ) + assert result['Upper'] == 'HELLO WORLD' + assert result['StartsWith'] is True + assert result['EndsWith'] is True + + +@pytest.mark.asyncio +async def test_parity_complex_graph(async_execute_fn): + """Complex dependency graph produces correct results.""" + result = await async_execute_fn( + """ + A: str = JsonData(path="$.a", coerce_type=True) + B: str = JsonData(path="$.b", coerce_type=True) + LenA = StringLength(s=A) + LenB = StringLength(s=B) + ALower = StringToLower(s=A) + BUpper = StringToUpper(s=B) + RuleA = Rule(when_all=[LenA > 3], description="A is long") + RuleB = Rule(when_all=[LenB > 3], description="B is long") + """, + data={'a': 'Hello', 'b': 'Hi'}, + ) + assert result['LenA'] == 5 + assert result['LenB'] == 2 + assert result['ALower'] == 'hello' + assert result['BUpper'] == 'HI' + assert result['RuleA'] is True + assert result['RuleB'] is False diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_async_sinks.py b/osprey_async_worker/src/osprey/async_worker/tests/test_async_sinks.py new file mode 100644 index 0000000..893d483 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_async_sinks.py @@ -0,0 +1,189 @@ +"""Tests for async sink infrastructure.""" + +import asyncio +from datetime import datetime +from typing import List + +import pytest +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.async_worker.sinks.sink.input_stream import AsyncStaticInputStream +from osprey.async_worker.sinks.sink.output_sink import AsyncMultiOutputSink, AsyncStdoutOutputSink +from osprey.engine.executor.execution_context import Action, ExecutionResult + + +def _make_result(action_id: int = 1, action_name: str = 'test') -> ExecutionResult: + return ExecutionResult( + extracted_features={}, + action=Action( + action_id=action_id, + action_name=action_name, + data={}, + timestamp=datetime.utcnow(), + ), + effects={}, + error_infos=[], + validator_results=None, + sample_rate=100, + ) + + +# --- AsyncStaticInputStream --- + + +@pytest.mark.asyncio +async def test_static_input_stream_yields_all_items(): + items = ['a', 'b', 'c'] + stream = AsyncStaticInputStream(items) + collected = [] + async for item in stream: + collected.append(item) + assert collected == items + + +@pytest.mark.asyncio +async def test_static_input_stream_empty(): + stream = AsyncStaticInputStream([]) + collected = [] + async for item in stream: + collected.append(item) + assert collected == [] + + +# --- AsyncMultiOutputSink --- + + +class RecordingSink(AsyncBaseOutputSink): + """Test sink that records all pushed results.""" + + def __init__(self): + self.results: List[ExecutionResult] = [] + self.stopped = False + + def will_do_work(self, result: ExecutionResult) -> bool: + return True + + async def push(self, result: ExecutionResult) -> None: + self.results.append(result) + + async def stop(self) -> None: + self.stopped = True + + +class SelectiveSink(AsyncBaseOutputSink): + """Test sink that only processes specific action names.""" + + def __init__(self, allowed: str): + self._allowed = allowed + self.results: List[ExecutionResult] = [] + + def will_do_work(self, result: ExecutionResult) -> bool: + return result.action.action_name == self._allowed + + async def push(self, result: ExecutionResult) -> None: + self.results.append(result) + + async def stop(self) -> None: + pass + + +class FailingSink(AsyncBaseOutputSink): + """Test sink that always raises.""" + + def will_do_work(self, result: ExecutionResult) -> bool: + return True + + async def push(self, result: ExecutionResult) -> None: + raise RuntimeError('sink failure') + + async def stop(self) -> None: + pass + + +class SlowSink(AsyncBaseOutputSink): + """Test sink that takes too long.""" + + timeout = 0.05 + + def __init__(self): + self.attempted = False + + def will_do_work(self, result: ExecutionResult) -> bool: + return True + + async def push(self, result: ExecutionResult) -> None: + self.attempted = True + await asyncio.sleep(1.0) # way longer than timeout + + async def stop(self) -> None: + pass + + +@pytest.mark.asyncio +async def test_multi_sink_pushes_to_all(): + sink_a = RecordingSink() + sink_b = RecordingSink() + multi = AsyncMultiOutputSink([sink_a, sink_b]) + + result = _make_result() + await multi.push(result) + + assert len(sink_a.results) == 1 + assert len(sink_b.results) == 1 + + +@pytest.mark.asyncio +async def test_multi_sink_respects_will_do_work(): + sink_a = SelectiveSink('action_a') + sink_b = SelectiveSink('action_b') + multi = AsyncMultiOutputSink([sink_a, sink_b]) + + await multi.push(_make_result(action_name='action_a')) + await multi.push(_make_result(action_name='action_b')) + await multi.push(_make_result(action_name='action_c')) + + assert len(sink_a.results) == 1 + assert len(sink_b.results) == 1 + + +@pytest.mark.asyncio +async def test_multi_sink_continues_after_failure(): + """A failing sink doesn't prevent other sinks from receiving results.""" + failing = FailingSink() + recording = RecordingSink() + multi = AsyncMultiOutputSink([failing, recording]) + + await multi.push(_make_result()) + + # Recording sink still got the result despite failing sink + assert len(recording.results) == 1 + + +@pytest.mark.asyncio +async def test_multi_sink_handles_timeout(): + """A slow sink times out without blocking other sinks.""" + slow = SlowSink() + recording = RecordingSink() + multi = AsyncMultiOutputSink([slow, recording]) + + await multi.push(_make_result()) + + assert slow.attempted is True + assert len(recording.results) == 1 + + +@pytest.mark.asyncio +async def test_multi_sink_stop(): + sink_a = RecordingSink() + sink_b = RecordingSink() + multi = AsyncMultiOutputSink([sink_a, sink_b]) + + await multi.stop() + + assert sink_a.stopped is True + assert sink_b.stopped is True + + +@pytest.mark.asyncio +async def test_stdout_sink_will_do_work(): + sink = AsyncStdoutOutputSink() + assert sink.will_do_work(_make_result()) is True diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_coordinator_input_stream.py b/osprey_async_worker/src/osprey/async_worker/tests/test_coordinator_input_stream.py new file mode 100644 index 0000000..2cae831 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_coordinator_input_stream.py @@ -0,0 +1,127 @@ +"""Tests for the async coordinator input stream.""" + +import asyncio +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from osprey.async_worker.lib.coordinator_input_stream import ( + GrpcConnectionDiscoveryPool, + OspreyCoordinatorBiDirectionalStream, + OspreyCoordinatorInputStream, +) + +# --- GrpcConnectionDiscoveryPool --- + + +@pytest.mark.asyncio +async def test_discovery_pool_creates_channels(): + """Pool creates grpc.aio channels from service discovery. + + Async so a running event loop exists: grpc.aio.insecure_channel() binds to + the running loop, and constructing it without one raises 'no running event + loop' (order-dependent when other async tests have closed their loop).""" + mock_service = MagicMock() + mock_service.connection_address = 'localhost' + mock_service.grpc_port = 50051 + + mock_watcher = MagicMock() + + mock_directory = MagicMock() + mock_directory.select_all.return_value = [mock_service] + mock_directory.get_watcher.return_value = mock_watcher + + with patch('osprey.worker.lib.discovery.directory.Directory') as MockDirectory: + MockDirectory.instance.return_value = mock_directory + pool = GrpcConnectionDiscoveryPool('test_coordinator') + assert len(pool._grpc_channels) == 1 + + +# --- OspreyCoordinatorBiDirectionalStream --- + + +@pytest.mark.asyncio +async def test_bidirectional_stream_queue_based(): + """Stream uses asyncio.Queue for sending requests.""" + stream = OspreyCoordinatorBiDirectionalStream.__new__(OspreyCoordinatorBiDirectionalStream) + stream._request_queue = asyncio.Queue() + stream._should_run = True + + await stream._request_queue.put('test_request') + item = await stream._request_queue.get() + assert item == 'test_request' + + +def test_acking_context_should_nack_contract(): + """The out-of-band ack path keys off should_nack, so mark_as_nack must flip it.""" + from osprey.worker.sinks.utils.acking_contexts_base import NoopAckingContext + + ctx = NoopAckingContext(item='x') + assert ctx.should_nack is False + ctx.mark_as_nack() + assert ctx.should_nack is True + + +def test_send_ack_or_nack_emits_nack_when_not_ack(): + """ack=False must enqueue a Nack (not an Ack), so a nacked context isn't acked.""" + stream = OspreyCoordinatorBiDirectionalStream.__new__(OspreyCoordinatorBiDirectionalStream) + stream._outgoing_queue = asyncio.Queue() + + stream.send_ack_or_nack(123, ack=True) + stream.send_ack_or_nack(456, ack=False) + + ack_req = stream._outgoing_queue.get_nowait() + nack_req = stream._outgoing_queue.get_nowait() + assert ack_req.action_request.ack_or_nack.HasField('ack') + assert nack_req.action_request.ack_or_nack.HasField('nack') + + +@pytest.mark.asyncio +async def test_send_graceful_disconnect_emits_nack_when_not_ack(): + """Graceful-disconnect finalize paths must nack a nacked context, not ack it.""" + stream = OspreyCoordinatorBiDirectionalStream.__new__(OspreyCoordinatorBiDirectionalStream) + stream._outgoing_queue = asyncio.Queue() + + await stream.send_graceful_disconnect(99, ack=False) + + req = stream._outgoing_queue.get_nowait() + assert req.disconnect.ack_or_nack.HasField('nack') + + +# --- OspreyCoordinatorInputStream --- + + +@pytest.mark.asyncio +async def test_input_stream_stop(): + """Stop sets the shutdown event.""" + stream = OspreyCoordinatorInputStream.__new__(OspreyCoordinatorInputStream) + stream._shutdown_event = asyncio.Event() + stream._channel_pool = AsyncMock() # stop() awaits channel_pool.close() + + assert not stream._shutdown_event.is_set() + await stream.stop() + assert stream._shutdown_event.is_set() + stream._channel_pool.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_input_stream_shutdown_event_unblocks(): + """Setting shutdown event should unblock any waiters.""" + stream = OspreyCoordinatorInputStream.__new__(OspreyCoordinatorInputStream) + stream._shutdown_event = asyncio.Event() + stream._channel_pool = AsyncMock() # stop() awaits channel_pool.close() + + unblocked = False + + async def waiter(): + nonlocal unblocked + await stream._shutdown_event.wait() + unblocked = True + + task = asyncio.create_task(waiter()) + await asyncio.sleep(0.01) + assert not unblocked + + await stream.stop() + await asyncio.sleep(0.01) + assert unblocked + await task diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_engine.py b/osprey_async_worker/src/osprey/async_worker/tests/test_engine.py new file mode 100644 index 0000000..6e38575 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_engine.py @@ -0,0 +1,229 @@ +"""Tests for AsyncOspreyEngine recompile behavior on etcd updates.""" + +import threading +from unittest.mock import MagicMock, patch + +import pytest +from osprey.async_worker.engine import AsyncOspreyEngine + + +def _make_engine_with_stub_compile(initial_graph, recompile_graph): + """Build an AsyncOspreyEngine without exercising the real compile path. + + Patches _compile_execution_graph_sync so __init__ uses initial_graph and + later recompiles return recompile_graph. Returns the constructed engine. + """ + sources_provider = MagicMock() + sources_provider.get_current_sources.return_value = MagicMock(hash=MagicMock(return_value='abc')) + + udf_registry = MagicMock() + + with ( + patch.object(AsyncOspreyEngine, '_compile_execution_graph_sync', return_value=initial_graph), + patch('osprey.async_worker.engine.ConfigSubkeyHandler'), + ): + engine = AsyncOspreyEngine( + sources_provider=sources_provider, + udf_registry=udf_registry, + ) + + # Register what subsequent compiles should return. + engine._stub_recompile_graph = recompile_graph + return engine + + +@pytest.mark.asyncio +async def test_handle_updated_sources_runs_compile_off_event_loop(): + """The compile must run in self._thread_pool, not on the event loop thread. + + When the loop is blocked by the compile, in-flight gRPC tasks cannot drain + and their pinned response buffers stack up alongside the compile-transient + memory. Routing through self._thread_pool keeps the loop responsive. + """ + initial = MagicMock(name='initial_graph') + new = MagicMock(name='new_graph') + engine = _make_engine_with_stub_compile(initial, new) + + loop_thread = threading.get_ident() + compile_thread = {} + + def fake_sync_compile(): + compile_thread['tid'] = threading.get_ident() + return new + + with patch.object(engine, '_compile_execution_graph_sync', side_effect=fake_sync_compile): + await engine._handle_updated_sources() + + assert 'tid' in compile_thread, 'compile did not run' + assert compile_thread['tid'] != loop_thread, ( + 'compile ran on the event-loop thread; it must run in the thread pool ' + 'so in-flight tasks can drain during compile' + ) + assert engine._execution_graph is new + + +@pytest.mark.asyncio +async def test_handle_updated_sources_does_not_force_gc(): + """We deliberately do NOT call gc.collect() after swap. + + Forcing gen-2 collection promotes every surviving object to gen 2, which + makes subsequent automatic collections during action processing more + expensive. We let CPython's reference counting reclaim the old graph + naturally — same as the gevent engine. + """ + import gc as gc_module + + initial = MagicMock(name='initial_graph') + new = MagicMock(name='new_graph') + engine = _make_engine_with_stub_compile(initial, new) + + with ( + patch.object(engine, '_compile_execution_graph_sync', return_value=new), + patch.object(gc_module, 'collect') as mock_collect, + ): + await engine._handle_updated_sources() + + assert mock_collect.call_count == 0 + + +@pytest.mark.asyncio +async def test_handle_updated_sources_keeps_old_graph_on_compile_failure(): + """If the new compile raises, _execution_graph must still point at the old + graph. Otherwise in-flight actions would crash on a missing graph.""" + initial = MagicMock(name='initial_graph') + new = MagicMock(name='new_graph') + engine = _make_engine_with_stub_compile(initial, new) + pre_swap_graph = engine._execution_graph + + with patch.object(engine, '_compile_execution_graph_sync', side_effect=RuntimeError('boom')): + await engine._handle_updated_sources() + + assert engine._execution_graph is pre_swap_graph, 'failed compile should not have replaced the existing graph' + + +@pytest.mark.asyncio +async def test_handle_updated_sources_dispatches_config_after_swap(): + """After a successful recompile, the config subkey handler must be notified + with the validated_sources of the new graph.""" + initial = MagicMock(name='initial_graph') + new = MagicMock(name='new_graph') + engine = _make_engine_with_stub_compile(initial, new) + engine._config_subkey_handler = MagicMock() + + with patch.object(engine, '_compile_execution_graph_sync', return_value=new): + await engine._handle_updated_sources() + + engine._config_subkey_handler.dispatch_config.assert_called_once_with(new.validated_sources) + + +@pytest.mark.asyncio +async def test_handle_updated_sources_nulls_parents_on_old_graph(): + """After swap, parent pointers on every AST node in the OLD graph must be + nulled so refcount reclaims the old graph without waiting for gen-2 GC. + The NEW graph's parents must be left intact.""" + + # Old graph: two mock sources, each with a fake ast_root whose iter_nodes + # walk yields three nodes carrying a `parent` attribute. + def make_nodes(n): + return [MagicMock(parent=MagicMock(name=f'parent_{i}')) for i in range(n)] + + old_nodes_per_source = [make_nodes(3), make_nodes(4)] + new_nodes_per_source = [make_nodes(2)] + + def fake_iter_nodes(root): + # The patched iter_nodes is called with source.ast_root; we look it up + # via the identity of the source it came from (mock attribute chain). + return iter(root._test_nodes) + + def make_sources(nodes_per_source): + sources = [] + for nodes in nodes_per_source: + src = MagicMock() + src.ast_root._test_nodes = nodes + sources.append(src) + validated = MagicMock() + validated.sources = sources + graph = MagicMock(validated_sources=validated) + return graph + + old_graph = make_sources(old_nodes_per_source) + new_graph = make_sources(new_nodes_per_source) + + engine = _make_engine_with_stub_compile(old_graph, new_graph) + + with ( + patch.object(engine, '_compile_execution_graph_sync', return_value=new_graph), + patch('osprey.async_worker.engine.iter_nodes', side_effect=fake_iter_nodes), + ): + await engine._handle_updated_sources() + + # OLD graph nodes: every parent nulled. + for nodes in old_nodes_per_source: + for n in nodes: + assert n.parent is None, 'old-graph AST node still has a parent' + + # NEW graph nodes: parents untouched. + for nodes in new_nodes_per_source: + for n in nodes: + assert n.parent is not None, 'new-graph AST node parent was wrongly nulled' + + +@pytest.mark.asyncio +async def test_handle_updated_sources_swallows_cycle_break_errors(): + """If cycle-breaking raises, the swap must still have succeeded — the new + graph stays installed and the engine continues to function.""" + initial = MagicMock(name='initial_graph') + new = MagicMock(name='new_graph') + engine = _make_engine_with_stub_compile(initial, new) + + with ( + patch.object(engine, '_compile_execution_graph_sync', return_value=new), + patch('osprey.async_worker.engine.iter_nodes', side_effect=RuntimeError('walker boom')), + ): + await engine._handle_updated_sources() + + assert engine._execution_graph is new, 'swap must complete even if cycle-break raises' + + +def test_break_old_graph_cycles_evicts_reverted_content_from_ast_cache(): + """_break_old_graph_cycles must evict a discarded source's content from the + never-evicted module-level parsed_ast_root_cache BEFORE nulling its AST parents. + + Otherwise a later graph that re-uses that exact content (e.g. a rule revert) + is handed back the same parent-nulled Root and fails validation with + "`Rule(...)` must be assigned to a variable", wedging the worker on stale rules. + Regression test: fails if the cache eviction is removed. + """ + from osprey.engine.ast.ast_utils import iter_nodes + from osprey.engine.ast.grammar import Source, parsed_ast_root_cache + + def rule_parent_type(root): + for node in iter_nodes(root): + func = getattr(node, 'func', None) + if type(node).__name__ == 'Call' and getattr(func, 'identifier', None) == 'Rule': + return type(node.parent).__name__ if node.parent is not None else None + return None + + content = "MyRule = Rule(\n when_all=[True],\n description='x',\n)\n" + path = 'actions/regression_evict.sml' + before = set(parsed_ast_root_cache.keys()) + try: + old_src = Source(path=path, contents=content) + assert rule_parent_type(old_src.ast_root) == 'Assign' # fresh parse: Rule is assigned + assert old_src in parsed_ast_root_cache + + # New graph changes this file (different content) -> old content is discarded. + new_src = Source(path=path, contents="Other = Rule(\n when_all=[False],\n description='y',\n)\n") + old_graph = MagicMock(validated_sources=MagicMock(sources=[old_src])) + new_graph = MagicMock(validated_sources=MagicMock(sources=[new_src])) + + AsyncOspreyEngine._break_old_graph_cycles(old_graph, new_graph) + + # Discarded content must be evicted so a later recurrence re-parses fresh. + assert old_src not in parsed_ast_root_cache, 'reverted content left in cache -> parent-nulled Root reused' + # Re-using the exact content now yields an intact AST (Rule still assigned). + assert rule_parent_type(Source(path=path, contents=content).ast_root) == 'Assign' + finally: + for key in list(parsed_ast_root_cache.keys()): + if key not in before: + parsed_ast_root_cache.pop(key, None) diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_etcd_sources_provider.py b/osprey_async_worker/src/osprey/async_worker/tests/test_etcd_sources_provider.py new file mode 100644 index 0000000..a781f51 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_etcd_sources_provider.py @@ -0,0 +1,199 @@ +"""Tests for the async etcd sources provider and input stream signaler.""" + +import asyncio +from unittest.mock import MagicMock, patch + +import pytest +from osprey.async_worker.lib.etcd.sources_provider import ( + AsyncEtcdSourcesProvider, + AsyncInputStreamReadySignaler, +) + +# --- AsyncInputStreamReadySignaler --- + + +@pytest.mark.asyncio +async def test_signaler_starts_ready(): + signaler = AsyncInputStreamReadySignaler() + assert not signaler.should_pause_input_stream() + + +@pytest.mark.asyncio +async def test_signaler_pause_clears_event(): + signaler = AsyncInputStreamReadySignaler() + # Mock the jitter sleep so the test is fast + with patch('osprey.async_worker.lib.etcd.sources_provider.asyncio.sleep', return_value=None): + await signaler.pause_input_stream() + assert signaler.should_pause_input_stream() + + +@pytest.mark.asyncio +async def test_signaler_resume_sets_event(): + signaler = AsyncInputStreamReadySignaler() + with patch('osprey.async_worker.lib.etcd.sources_provider.asyncio.sleep', return_value=None): + await signaler.pause_input_stream() + signaler.resume_input_stream() + assert not signaler.should_pause_input_stream() + + +@pytest.mark.asyncio +async def test_signaler_wait_blocks_when_paused(): + signaler = AsyncInputStreamReadySignaler() + with patch('osprey.async_worker.lib.etcd.sources_provider.asyncio.sleep', return_value=None): + await signaler.pause_input_stream() + + # wait_until_resume should block until resume is called + resumed = False + + async def waiter(): + nonlocal resumed + await signaler.wait_until_resume() + resumed = True + + task = asyncio.create_task(waiter()) + await asyncio.sleep(0.01) + assert not resumed + + signaler.resume_input_stream() + await asyncio.sleep(0.01) + assert resumed + await task + + +@pytest.mark.asyncio +async def test_signaler_wait_returns_immediately_when_ready(): + signaler = AsyncInputStreamReadySignaler() + # Should not block + await asyncio.wait_for(signaler.wait_until_resume(), timeout=0.1) + + +# --- AsyncEtcdSourcesProvider --- + + +@pytest.mark.asyncio +async def test_provider_get_current_sources_default_none(): + """Before start(), sources should be None.""" + provider = AsyncEtcdSourcesProvider(etcd_key='/test/key', etcd_client=MagicMock()) + sources = provider.get_current_sources() + assert sources is None + + +@pytest.mark.asyncio +async def test_provider_set_sources_watcher(): + """Watcher callback can be set.""" + provider = AsyncEtcdSourcesProvider(etcd_key='/test/key', etcd_client=MagicMock()) + callback = MagicMock() + provider.set_sources_watcher(callback) + assert provider._sources_watcher_callback is callback + + +@pytest.mark.asyncio +async def test_provider_stop_without_start(): + """Stop without start should be safe.""" + provider = AsyncEtcdSourcesProvider(etcd_key='/test/key', etcd_client=MagicMock()) + await provider.stop() # Should not raise + + +# --- _handle_event callback path --- + + +def _make_full_sync_event(payload: dict): + from osprey.worker.lib.etcd import FullSyncOne + + event = MagicMock(spec=FullSyncOne) + import json as _json + + event.value = _json.dumps(payload) + return event + + +@pytest.mark.asyncio +async def test_handle_event_awaits_async_callback(): + """An async (coroutine-returning) sources_watcher_callback must be awaited. + + The async engine's _handle_updated_sources is a coroutine function that + runs compile in a thread pool. If _handle_event calls it without awaiting, + the compile coroutine is dropped and never executes. + """ + provider = AsyncEtcdSourcesProvider(etcd_key='/test/key', etcd_client=MagicMock()) + awaited = asyncio.Event() + + async def async_callback() -> None: + awaited.set() + + provider.set_sources_watcher(async_callback) + + with patch('osprey.async_worker.lib.etcd.sources_provider.asyncio.sleep', return_value=None): + await provider._handle_event(_make_full_sync_event({'main.sml': '# noop'})) + + assert awaited.is_set(), 'async callback was not awaited' + + +@pytest.mark.asyncio +async def test_handle_event_calls_sync_callback(): + """A plain (non-coroutine) callback continues to work for back-compat.""" + provider = AsyncEtcdSourcesProvider(etcd_key='/test/key', etcd_client=MagicMock()) + callback = MagicMock() + provider.set_sources_watcher(callback) + + with patch('osprey.async_worker.lib.etcd.sources_provider.asyncio.sleep', return_value=None): + await provider._handle_event(_make_full_sync_event({'main.sml': '# noop'})) + + callback.assert_called_once() + + +@pytest.mark.asyncio +async def test_handle_event_skips_callback_on_already_applied_hash(): + """Dedup short-circuit: if the new sources match what the engine has APPLIED, + no callback. Dedup keys off the applied hash, not merely the received one.""" + from osprey.engine.ast.sources import Sources + + provider = AsyncEtcdSourcesProvider(etcd_key='/test/key', etcd_client=MagicMock()) + payload = {'main.sml': '# noop'} + + # The engine has already compiled & applied this exact content. + provider._applied_sources_hash = Sources.from_dict(payload).hash() + + callback = MagicMock() + provider.set_sources_watcher(callback) + + await provider._handle_event(_make_full_sync_event(payload)) + + callback.assert_not_called() + + +@pytest.mark.asyncio +async def test_handle_event_reapplies_after_failed_compile(): + """Self-heal: a recompile that never reports applied must NOT suppress the + next identical re-delivery, or the worker wedges on stale rules. + + Models a failed/dropped compile by invoking the callback (which would + compile) but never calling mark_sources_applied(). The same payload + redelivered must fire the callback again. + """ + from osprey.engine.ast.sources import Sources + + provider = AsyncEtcdSourcesProvider(etcd_key='/test/key', etcd_client=MagicMock()) + good = {'main.sml': '# good'} + bad = {'main.sml': '# bad'} + + callback = MagicMock() + provider.set_sources_watcher(callback) + + with patch('osprey.async_worker.lib.etcd.sources_provider.asyncio.sleep', return_value=None): + # First payload applies cleanly (engine confirms via mark_sources_applied). + await provider._handle_event(_make_full_sync_event(good)) + provider.mark_sources_applied(Sources.from_dict(good).hash()) + assert callback.call_count == 1 + + # Re-delivery of the applied payload is deduped. + await provider._handle_event(_make_full_sync_event(good)) + assert callback.call_count == 1 + + # New payload whose compile "fails" — callback fires but apply is never marked. + await provider._handle_event(_make_full_sync_event(bad)) + assert callback.call_count == 2 + + # Identical re-delivery must re-fire (self-heal), not be suppressed. + await provider._handle_event(_make_full_sync_event(bad)) + assert callback.call_count == 3 diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_external_service.py b/osprey_async_worker/src/osprey/async_worker/tests/test_external_service.py new file mode 100644 index 0000000..26e2ccd --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_external_service.py @@ -0,0 +1,221 @@ +"""Tests for the async external service cache.""" + +import asyncio +from datetime import timedelta +from typing import Optional, Sequence + +import pytest +from osprey.async_worker.lib.external_service import AsyncExternalService, ExternalServiceAccessor +from result import Ok, Result + + +class FakeService(AsyncExternalService[str, str]): + """Test service that records calls and returns predictable results.""" + + def __init__(self, delay: float = 0.0): + self.call_count = 0 + self.delay = delay + + async def get_from_service(self, key: str) -> str: + self.call_count += 1 + if self.delay > 0: + await asyncio.sleep(self.delay) + return f'value_{key}' + + +class FailingService(AsyncExternalService[str, str]): + """Test service that always raises.""" + + async def get_from_service(self, key: str) -> str: + raise ValueError(f'service error for {key}') + + +class FailOnceService(AsyncExternalService[str, Optional[str]]): + """Raises on first call per key, succeeds after.""" + + def __init__(self): + self.seen = set() + + def count_error_once(self) -> bool: + return True + + async def get_from_service(self, key: str) -> Optional[str]: + if key not in self.seen: + self.seen.add(key) + raise ValueError('first call fails') + return f'value_{key}' + + +class TTLService(AsyncExternalService[str, str]): + def __init__(self, ttl: timedelta): + self._ttl = ttl + self.call_count = 0 + + def cache_ttl(self) -> Optional[timedelta]: + return self._ttl + + async def get_from_service(self, key: str) -> str: + self.call_count += 1 + return f'value_{key}_{self.call_count}' + + +class BatchService(AsyncExternalService[str, str]): + """Test service that supports batch operations.""" + + def __init__(self): + self.batch_call_count = 0 + + async def get_from_service(self, key: str) -> str: + return f'value_{key}' + + async def batch_get_from_service(self, keys: Sequence[str]) -> Sequence[Result[str, Exception]]: + self.batch_call_count += 1 + return [Ok(f'batch_{key}') for key in keys] + + +# --- Cache tests --- + + +@pytest.mark.asyncio +async def test_get_returns_value(): + service = FakeService() + accessor = ExternalServiceAccessor(service) + result = await accessor.get('foo') + assert result == 'value_foo' + + +@pytest.mark.asyncio +async def test_get_caches_result(): + service = FakeService() + accessor = ExternalServiceAccessor(service) + await accessor.get('foo') + await accessor.get('foo') + assert service.call_count == 1 + + +@pytest.mark.asyncio +async def test_get_different_keys_not_cached(): + service = FakeService() + accessor = ExternalServiceAccessor(service) + await accessor.get('foo') + await accessor.get('bar') + assert service.call_count == 2 + + +@pytest.mark.asyncio +async def test_get_without_cache_bypasses(): + service = FakeService() + accessor = ExternalServiceAccessor(service) + await accessor.get('foo') + await accessor.get_without_cache('foo') + assert service.call_count == 2 + + +@pytest.mark.asyncio +async def test_get_without_cache_updates_cache(): + service = FakeService() + accessor = ExternalServiceAccessor(service) + await accessor.get_without_cache('foo') + await accessor.get('foo') + assert service.call_count == 1 # Second get hits cache + + +# --- Concurrent access (future dedup) --- + + +@pytest.mark.asyncio +async def test_concurrent_get_deduplicates(): + """Multiple concurrent gets for the same key should only call service once.""" + service = FakeService(delay=0.05) + accessor = ExternalServiceAccessor(service) + results = await asyncio.gather( + accessor.get('foo'), + accessor.get('foo'), + accessor.get('foo'), + ) + assert all(r == 'value_foo' for r in results) + assert service.call_count == 1 + + +# --- Error handling --- + + +@pytest.mark.asyncio +async def test_get_propagates_error(): + service = FailingService() + accessor = ExternalServiceAccessor(service) + with pytest.raises(ValueError, match='service error for foo'): + await accessor.get('foo') + + +@pytest.mark.asyncio +async def test_get_error_cached(): + """Errors are cached — second get raises the same error.""" + service = FailingService() + accessor = ExternalServiceAccessor(service) + with pytest.raises(ValueError): + await accessor.get('foo') + with pytest.raises(ValueError): + await accessor.get('foo') + + +@pytest.mark.asyncio +async def test_count_error_once(): + """With count_error_once, subsequent callers get None instead of the error.""" + service = FailOnceService() + accessor = ExternalServiceAccessor(service) + with pytest.raises(ValueError): + await accessor.get('foo') + # Second get should return None (cached as None due to count_error_once) + result = await accessor.get('foo') + assert result is None + + +# --- TTL --- + + +@pytest.mark.asyncio +async def test_ttl_expires_cache(): + """Expired TTL causes a re-fetch.""" + service = TTLService(ttl=timedelta(days=-1)) # Immediately expired + accessor = ExternalServiceAccessor(service) + r1 = await accessor.get('foo') + r2 = await accessor.get('foo') + assert r1 != r2 # Different values = two service calls + assert service.call_count == 2 + + +@pytest.mark.asyncio +async def test_no_ttl_caches_forever(): + service = FakeService() + accessor = ExternalServiceAccessor(service) + await accessor.get('foo') + await accessor.get('foo') + await accessor.get('foo') + assert service.call_count == 1 + + +# --- Batch --- + + +@pytest.mark.asyncio +async def test_batch_get(): + service = BatchService() + accessor = ExternalServiceAccessor(service) + results = await accessor.batch_get(['a', 'b', 'c']) + assert len(results) == 3 + assert results[0] == Ok('batch_a') + assert results[1] == Ok('batch_b') + assert results[2] == Ok('batch_c') + assert service.batch_call_count == 1 + + +@pytest.mark.asyncio +async def test_batch_get_uses_cache(): + service = BatchService() + accessor = ExternalServiceAccessor(service) + await accessor.batch_get(['a', 'b']) + # Second batch with overlap — 'a' and 'b' cached, only 'c' fetched + results = await accessor.batch_get(['a', 'b', 'c']) + assert len(results) == 3 + assert service.batch_call_count == 2 diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_kafka_input_stream.py b/osprey_async_worker/src/osprey/async_worker/tests/test_kafka_input_stream.py new file mode 100644 index 0000000..7aa5720 --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_kafka_input_stream.py @@ -0,0 +1,190 @@ +"""Tests for AsyncKafkaInputStream. + +These exercise the decode contract against the Osprey-format envelope that an +upstream producer (e.g. at-kafka in Osprey-compatible mode) writes, plus the +poll/skip/stop behaviour of the async stream, without a real Kafka broker. +""" + +import asyncio +import json +import threading +from typing import Any, List, Mapping, Sequence + +import pytest +from osprey.async_worker.sinks.sink.input_stream import AsyncKafkaInputStream +from osprey.worker.sinks.utils.acking_contexts_base import NoopAckingContext + + +class _FakeRecord: + def __init__(self, value: Any): + self.value = value + + +class _FakeConsumer: + """Returns each queued batch once per poll(), then empty batches. + + Once batches are exhausted, flips the stream's stop flag on the next poll so + the async generator terminates instead of spinning on empty polls. + """ + + def __init__(self, batches: Sequence[Sequence[_FakeRecord]]): + self._batches = list(batches) + self.closed = False + self.commits = 0 + self.stream: object = None # set by the test to enable auto-stop + + def poll(self, timeout_ms: int = 0, max_records: int = 0) -> Mapping[str, Sequence[_FakeRecord]]: + if self._batches: + return {'tp-0': self._batches.pop(0)} + if self.stream is not None: + self.stream._stopped = True # type: ignore[attr-defined] + return {} + + def commit(self) -> None: + self.commits += 1 + + def close(self) -> None: + self.closed = True + + +def _osprey_envelope(action_id: int, action_name: str, payload: dict) -> bytes: + return json.dumps( + { + 'send_time': '2024-01-01T12:00:00Z', + 'data': { + 'action_id': action_id, + 'action_name': action_name, + 'data': payload, + 'secret_data': {}, + 'encoding': 'UTF8', + }, + } + ).encode('utf-8') + + +def test_decode_maps_osprey_envelope_to_action() -> None: + value = _osprey_envelope(123, 'operation#create', {'did': 'did:plc:abc', 'message': 'hello'}) + + action = AsyncKafkaInputStream.decode(value) + + assert action.action_id == 123 + assert action.action_name == 'operation#create' + assert action.data == {'did': 'did:plc:abc', 'message': 'hello'} + # send_time parsed into a tz-aware datetime + assert action.timestamp.year == 2024 + + +def test_decode_coerces_string_action_id() -> None: + value = _osprey_envelope('456', 'identity', {}) + # action_id arrives as a JSON string; decode must coerce to int. + raw = json.loads(value) + raw['data']['action_id'] = '456' + action = AsyncKafkaInputStream.decode(json.dumps(raw).encode('utf-8')) + assert action.action_id == 456 + + +async def test_gen_yields_decoded_actions_then_stops() -> None: + records = [ + _FakeRecord(_osprey_envelope(1, 'operation#create', {'message': 'a'})), + _FakeRecord(_osprey_envelope(2, 'operation#create', {'message': 'b'})), + ] + consumer = _FakeConsumer([records]) + stream = AsyncKafkaInputStream(consumer, poll_timeout_ms=1, max_records=10) + + collected: List[int] = [] + async for ctx in stream: + assert isinstance(ctx, NoopAckingContext) + with ctx as action: + collected.append(action.action_id) + if len(collected) == 2: + await stream.stop() + break + + assert collected == [1, 2] + assert consumer.closed is True + + +async def test_commits_once_after_batch_is_processed() -> None: + records = [ + _FakeRecord(_osprey_envelope(1, 'operation#create', {'message': 'a'})), + _FakeRecord(_osprey_envelope(2, 'operation#create', {'message': 'b'})), + ] + consumer = _FakeConsumer([records]) + stream = AsyncKafkaInputStream(consumer, poll_timeout_ms=1, max_records=10) + consumer.stream = stream # auto-stop once polls go empty + + collected: List[int] = [] + async for ctx in stream: + with ctx as action: + collected.append(action.action_id) + + assert collected == [1, 2] + # Exactly one commit, and only after both records in the batch were processed + # (manual-commit at-least-once, not auto-commit-before-processing). + assert consumer.commits == 1 + + +async def test_stop_does_not_close_consumer_during_inflight_poll() -> None: + """stop() must not close() the (thread-unsafe) consumer while a poll() is in + flight. All consumer ops share one worker thread, so close() is queued behind + the active poll and can only run once it returns.""" + poll_entered = threading.Event() + release_poll = threading.Event() + order: List[str] = [] + + class _BlockingConsumer: + def __init__(self) -> None: + self.closed = False + + def poll(self, timeout_ms: int = 0, max_records: int = 0) -> Mapping[str, Sequence[Any]]: + order.append('poll_start') + poll_entered.set() + release_poll.wait(2.0) + order.append('poll_end') + return {} + + def commit(self) -> None: + order.append('commit') + + def close(self) -> None: + order.append('close') + self.closed = True + + consumer = _BlockingConsumer() + stream = AsyncKafkaInputStream(consumer, poll_timeout_ms=1, max_records=1) + anext_task = asyncio.create_task(stream.__aiter__().__anext__()) + + # Wait until poll() is actually executing on the consumer's worker thread. + await asyncio.get_running_loop().run_in_executor(None, poll_entered.wait, 2.0) + + # Request stop while poll is still blocked; close() must stay queued. + stop_task = asyncio.create_task(stream.stop()) + await asyncio.sleep(0.05) + assert 'close' not in order, 'close() ran concurrently with an in-flight poll()' + + release_poll.set() + with pytest.raises(StopAsyncIteration): + await anext_task + await stop_task + + assert consumer.closed + assert order.index('poll_end') < order.index('close') + + +async def test_gen_skips_malformed_record_and_continues() -> None: + records = [ + _FakeRecord(b'not-json'), + _FakeRecord(_osprey_envelope(7, 'operation#create', {'message': 'ok'})), + ] + consumer = _FakeConsumer([records]) + stream = AsyncKafkaInputStream(consumer, poll_timeout_ms=1, max_records=10) + + collected: List[int] = [] + async for ctx in stream: + with ctx as action: + collected.append(action.action_id) + await stream.stop() + break + + # The malformed record is skipped; the next valid one is yielded. + assert collected == [7] diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_no_gevent_imports.py b/osprey_async_worker/src/osprey/async_worker/tests/test_no_gevent_imports.py new file mode 100644 index 0000000..6a840ed --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_no_gevent_imports.py @@ -0,0 +1,95 @@ +"""Verify that core async worker modules do not import gevent through osprey code. + +The async worker runs on pure asyncio — gevent should not be loaded through +our code. Third-party libraries (sentry_sdk, ddtrace) may import gevent to +detect monkey-patching, which is outside our control. + +This test checks that osprey's own import chains don't pull in gevent. +It ignores gevent loaded by third-party libraries. + +NOTE: osprey.async_worker.lib.etcd.sources_provider is excluded because it +genuinely needs the sync etcd client (which imports gevent). That module is +only used at startup for rule source watching, not in the hot path. +""" + +import subprocess +import sys + +import pytest + +# Modules that must not pull in gevent through osprey code. +_ASYNC_WORKER_MODULES = [ + 'osprey.async_worker.engine', + 'osprey.async_worker.executor', + 'osprey.async_worker.sinks.sink.rules_sink', + 'osprey.async_worker.sinks.sink.input_stream', + 'osprey.async_worker.lib.coordinator_input_stream', + 'osprey.async_worker.lib.pigeon.client', +] + +# Third-party modules that legitimately import gevent (to detect monkey-patching). +# These are not our code and we can't control them. +_ALLOWED_GEVENT_IMPORTERS = frozenset( + { + 'sentry_sdk', + 'ddtrace', + } +) + + +@pytest.mark.parametrize('module', _ASYNC_WORKER_MODULES) +def test_no_osprey_gevent_import(module: str) -> None: + """Importing an async worker module must not load gevent through osprey code. + + Third-party libraries (sentry_sdk, ddtrace) may import gevent for + detection purposes — those are allowed. + """ + # Use a script that tracks WHO imports gevent first. + script = ( + f'import sys\n' + f'class _Tracker:\n' + f' importer = None\n' + f' def find_module(self, name, path=None):\n' + f' if name == "gevent" and "gevent" not in sys.modules and self.importer is None:\n' + f' import traceback\n' + f' stack = traceback.format_stack()\n' + f' for line in reversed(stack):\n' + f' if "/site-packages/" in line:\n' + f' # Extract package name from site-packages path\n' + f' parts = line.split("/site-packages/")[1].split("/")[0]\n' + f' self.importer = parts.split(".")[0]\n' + f' break\n' + f' elif "/osprey" in line.lower():\n' + f' self.importer = "osprey"\n' + f' break\n' + f' if self.importer is None:\n' + f' self.importer = "unknown"\n' + f' return None\n' + f'tracker = _Tracker()\n' + f'sys.meta_path.insert(0, tracker)\n' + f'try:\n' + f' import {module}\n' + f'except Exception as e:\n' + f' print("IMPORT_ERROR:", type(e).__name__, e)\n' + f' sys.exit(2)\n' + f'gevent_loaded = any(k.startswith("gevent") for k in sys.modules)\n' + f'allowed = {{"sentry_sdk", "ddtrace"}}\n' + f'if gevent_loaded and tracker.importer not in allowed:\n' + f' print(f"GEVENT_LOADED_BY: {{tracker.importer}}")\n' + f' sys.exit(1)\n' + f'print(f"CLEAN (gevent_loaded={{gevent_loaded}}, importer={{tracker.importer}})")\n' + ) + result = subprocess.run( + [sys.executable, '-c', script], + capture_output=True, + text=True, + timeout=30, + ) + if result.returncode == 2: + pytest.skip(f'Could not import {module}: {result.stdout.strip()}') + elif result.returncode == 1: + pytest.fail( + f'Importing {module} loaded gevent through osprey code.\n' + f'{result.stdout.strip()}\n' + f'stderr: {result.stderr[:500]}' + ) diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_pigeon_client.py b/osprey_async_worker/src/osprey/async_worker/tests/test_pigeon_client.py new file mode 100644 index 0000000..b00403d --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_pigeon_client.py @@ -0,0 +1,75 @@ +"""Tests for the async pigeon client.""" + +import grpc +from osprey.async_worker.lib.pigeon.skip_rate_limit import skip_rate_limit_context + +# --- skip_rate_limit contextvars --- + + +def test_skip_rate_limit_default_false(): + assert skip_rate_limit_context.skip is False + + +def test_skip_rate_limit_set_and_get(): + skip_rate_limit_context.skip = True + assert skip_rate_limit_context.skip is True + skip_rate_limit_context.skip = False + assert skip_rate_limit_context.skip is False + + +def test_skip_rate_limit_property_api(): + """Uses .skip property, matching the gevent.local API.""" + skip_rate_limit_context.skip = True + assert skip_rate_limit_context.skip is True + skip_rate_limit_context.skip = False + + +# --- RoutingType constants --- + + +def test_routing_type_values(): + from osprey.async_worker.lib.pigeon.client import RoutingType + + assert RoutingType.CHUNKED == 1 + assert RoutingType.SCALAR == 2 + assert RoutingType.ROUND_ROBIN == 3 + assert RoutingType.ENVOY == 4 + assert len(RoutingType.ALL) == 4 + + +# --- GRPC HTTP code translation --- + + +def test_grpc_http_translations(): + from osprey.async_worker.lib.pigeon.client import _GRPC_HTTP_CODE_TRANSLATIONS + + assert _GRPC_HTTP_CODE_TRANSLATIONS[grpc.StatusCode.OK] == 200 + assert _GRPC_HTTP_CODE_TRANSLATIONS[grpc.StatusCode.NOT_FOUND] == 404 + assert _GRPC_HTTP_CODE_TRANSLATIONS[grpc.StatusCode.INTERNAL] == 500 + assert _GRPC_HTTP_CODE_TRANSLATIONS[grpc.StatusCode.UNAVAILABLE] == 503 + assert _GRPC_HTTP_CODE_TRANSLATIONS[grpc.StatusCode.DEADLINE_EXCEEDED] == 504 + + +# --- RetryPolicy --- + + +def test_retry_policy_type(): + from osprey.async_worker.lib.pigeon.client import RetryPolicy + + policy: RetryPolicy = { + 'retryable_grpc_status_codes': {grpc.StatusCode.UNAVAILABLE}, + 'max_secondaries_to_retry': 2, + } + assert grpc.StatusCode.UNAVAILABLE in policy['retryable_grpc_status_codes'] + assert policy['max_secondaries_to_retry'] == 2 + + +# --- ServiceDefinition --- + + +def test_service_definition_type(): + from osprey.async_worker.lib.pigeon.client import ServiceDefinition + + sd: ServiceDefinition = {'address': 'localhost', 'ip': '127.0.0.1', 'port': 5000} + assert sd['address'] == 'localhost' + assert sd['port'] == 5000 diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_plugin_manager.py b/osprey_async_worker/src/osprey/async_worker/tests/test_plugin_manager.py new file mode 100644 index 0000000..0d4b84c --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_plugin_manager.py @@ -0,0 +1,173 @@ +"""Tests for the async worker plugin manager. + +Locks down the behavior that bootstrap_async_udfs: +1. Resolves MXLookup to the async-native class (not sync stdlib). +2. Goes through the same register_udfs hook that third-party plugins use. +3. Doesn't drop other stdlib UDFs in the process. +""" + +from __future__ import annotations + +import pytest +from osprey.async_worker.adaptor import plugin_manager as pm +from osprey.async_worker.stdlib_udfs import _async_stdlib_plugin +from osprey.async_worker.stdlib_udfs.async_mx_lookup import MXLookup as AsyncMXLookup +from osprey.engine.stdlib.udfs.json_data import JsonData +from osprey.engine.stdlib.udfs.mx_lookup import MXLookup as SyncMXLookup +from osprey.engine.stdlib.udfs.rules import Rule + + +@pytest.fixture(autouse=True) +def reset_plugin_manager(): + """Clear lru_cache and unregister any plugins between tests. + + plugin_manager is a module-level singleton. Without this, state from + one test (e.g. a registered plugin) leaks into the next. + """ + pm.load_all_async_plugins.cache_clear() + yield + pm.load_all_async_plugins.cache_clear() + if pm.plugin_manager.is_registered(_async_stdlib_plugin): + pm.plugin_manager.unregister(_async_stdlib_plugin) + + +def test_async_stdlib_plugin_returns_async_mx_lookup() -> None: + """The first-party plugin's register_udfs returns the async MXLookup directly.""" + udfs = list(_async_stdlib_plugin.register_udfs()) + assert AsyncMXLookup in udfs + assert SyncMXLookup not in udfs + + +def test_async_stdlib_plugin_overrides_share_class_name() -> None: + """Overrides shadow stdlib by class name — verify the assumption holds. + + _deduplicate_udfs matches by __name__, so the async override class must + have the same __name__ as the sync class it replaces. + """ + for async_udf in _async_stdlib_plugin.register_udfs(): + assert async_udf.__name__ == 'MXLookup' # currently the only override + + +def test_bootstrap_resolves_mx_lookup_to_async_version() -> None: + registry, _helpers = pm.bootstrap_async_udfs(config=None) + resolved = registry.get('MXLookup') + assert resolved is AsyncMXLookup, ( + f'Expected MXLookup to resolve to AsyncMXLookup, got {resolved!r} ' + f'from module {resolved.__module__ if resolved else None}' + ) + + +def test_bootstrap_does_not_register_sync_mx_lookup() -> None: + """Sync MXLookup must not appear in the merged registry under any name.""" + registry, _helpers = pm.bootstrap_async_udfs(config=None) + for udf in registry.iter_functions(): + assert udf is not SyncMXLookup, 'Sync MXLookup leaked into the async registry' + + +def test_bootstrap_preserves_non_overridden_stdlib_udfs() -> None: + """Stdlib UDFs without an async override should still be registered as-is.""" + registry, _helpers = pm.bootstrap_async_udfs(config=None) + assert registry.get('JsonData') is JsonData + assert registry.get('Rule') is Rule + + +def test_bootstrap_registers_internal_plugin() -> None: + """The internal async-stdlib plugin must be registered after bootstrap. + + This confirms the override flows through the pluggy hook system rather + than a hardcoded path inside bootstrap_async_udfs. + """ + pm.bootstrap_async_udfs(config=None) + assert pm.plugin_manager.is_registered(_async_stdlib_plugin) + + +def test_bootstrap_register_udfs_hook_emits_async_mx_lookup() -> None: + """The register_udfs hook itself returns AsyncMXLookup via the internal plugin.""" + pm.load_all_async_plugins() + flattened: list = [] + for udfs in pm.plugin_manager.hook.register_udfs(): + flattened.extend(udfs) + assert AsyncMXLookup in flattened + + +class _StubUDF: + """A stand-in UDF class used to verify helper binding without depending on + any concrete UDFBase subclass. Helper binding only stores the class as a + dict key, so any hashable type works here.""" + + +class _UDFHelpersPlugin: + """A pluggy plugin that returns one (udf_class, helper) pair when + register_udf_helpers is called.""" + + def __init__(self, udf_class, helper, capture): + self._udf_class = udf_class + self._helper = helper + self._capture = capture + + @pm.hookimpl_osprey_async + def register_udf_helpers(self, config): + self._capture.append(config) + return [(self._udf_class, self._helper)] + + +def test_bootstrap_applies_register_udf_helpers_bindings() -> None: + """A plugin that implements register_udf_helpers should have its (udf, helper) + pair set on UDFHelpers during bootstrap. The framework must not need to + import the plugin's UDF class to bind the helper.""" + helper = object() + captured: list = [] + plugin = _UDFHelpersPlugin(_StubUDF, helper, captured) + pm.plugin_manager.register(plugin) + try: + fake_config = object() + _registry, helpers = pm.bootstrap_async_udfs(config=fake_config) # type: ignore[arg-type] + assert captured == [fake_config], 'register_udf_helpers must receive the config' + # UDFHelpers.get_udf_helper expects an instance (it calls type()). + # Inspect the underlying dict directly since _StubUDF is not instantiable. + assert helpers._helpers[_StubUDF] is helper + finally: + pm.plugin_manager.unregister(plugin) + + +def test_bootstrap_skips_helper_wiring_when_config_is_none() -> None: + """register_udf_helpers depends on `config`; if no config is supplied, + bootstrap must still succeed without invoking the hook.""" + captured: list = [] + plugin = _UDFHelpersPlugin(_StubUDF, object(), captured) + pm.plugin_manager.register(plugin) + try: + _registry, helpers = pm.bootstrap_async_udfs(config=None) + assert captured == [], 'hook must not be called when config is None' + assert _StubUDF not in helpers._helpers + finally: + pm.plugin_manager.unregister(plugin) + + +def test_bootstrap_swallows_exceptions_from_register_udf_helpers() -> None: + """A misbehaving plugin must not take down bootstrap. The exception is + logged and other UDFs/helpers still load.""" + + class _BrokenPlugin: + @pm.hookimpl_osprey_async + def register_udf_helpers(self, config): + raise RuntimeError('plugin boom') + + plugin = _BrokenPlugin() + pm.plugin_manager.register(plugin) + try: + fake_config = object() + registry, _helpers = pm.bootstrap_async_udfs(config=fake_config) # type: ignore[arg-type] + # Standard UDFs still resolved despite the broken hook. + assert registry.get('JsonData') is JsonData + finally: + pm.plugin_manager.unregister(plugin) + + +def test_no_residual_register_labels_service_or_provider_hookspec() -> None: + """The legacy labels-specific hookspec has been removed in favor of the + generic register_udf_helpers hook.""" + pm.load_all_async_plugins() + assert not hasattr(pm.plugin_manager.hook, 'register_labels_service_or_provider'), ( + 'register_labels_service_or_provider should be removed in favor of register_udf_helpers' + ) diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_publisher.py b/osprey_async_worker/src/osprey/async_worker/tests/test_publisher.py new file mode 100644 index 0000000..e589f4c --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_publisher.py @@ -0,0 +1,50 @@ +"""Tests for AsyncPubSubPublisher.""" + +from unittest.mock import MagicMock, patch + +from google.api_core.exceptions import NotFound +from osprey.async_worker.lib.publisher import _PUBLISH_RETRY, AsyncPubSubPublisher + + +def _make_publisher(): + """Return an AsyncPubSubPublisher with a mocked PublisherClient.""" + with patch('osprey.async_worker.lib.publisher.pubsub_v1.PublisherClient'): + publisher = AsyncPubSubPublisher(project_id='proj', topic_id='topic') + publisher._client = MagicMock() + return publisher + + +def _make_future(result=None, exc=None): + future = MagicMock() + if exc is not None: + future.result.side_effect = exc + else: + future.result.return_value = result + return future + + +@patch('osprey.async_worker.lib.publisher.metrics') +def test_single_attempt_success(mock_metrics): + publisher = _make_publisher() + publisher._client.publish.return_value = _make_future(result='msg-id-1') + + publisher._sync_flush([b'hello']) + + publisher._client.publish.assert_called_once_with(publisher._topic_path, b'hello', retry=_PUBLISH_RETRY) + mock_metrics.increment.assert_any_call('async_pubsub_publisher.publish.success', tags=publisher._metric_tags) + failure_calls = [c for c in mock_metrics.increment.call_args_list if 'failure' in c[0][0]] + assert failure_calls == [] + + +@patch('osprey.async_worker.lib.publisher.metrics') +def test_permanent_failure_metric_fires(mock_metrics): + publisher = _make_publisher() + exc = NotFound('topic not found') + publisher._client.publish.return_value = _make_future(exc=exc) + + publisher._sync_flush([b'data']) + + failure_calls = [c for c in mock_metrics.increment.call_args_list if 'failure' in c[0][0]] + assert len(failure_calls) == 1 + assert failure_calls[0][0][0] == 'async_pubsub_publisher.publish.failure' + assert f'error:{exc.__class__.__name__}' in failure_calls[0][1]['tags'] diff --git a/osprey_async_worker/src/osprey/async_worker/tests/test_register_async_plugins.py b/osprey_async_worker/src/osprey/async_worker/tests/test_register_async_plugins.py new file mode 100644 index 0000000..bf5933a --- /dev/null +++ b/osprey_async_worker/src/osprey/async_worker/tests/test_register_async_plugins.py @@ -0,0 +1,33 @@ +"""Tests for the example async plugin registrations. + +Mirrors register_plugins (sync) for the experimental asyncio worker: verifies +the osprey_async_plugin hooks return the expected UDFs and async output sink, +and that the entry point is discoverable so the async worker can load it. +""" + +from importlib.metadata import entry_points +from typing import cast + +import register_async_plugins +from async_sinks.example_async_output_sink import ExampleAsyncOutputSink +from osprey.async_worker.adaptor.interfaces import AsyncBaseOutputSink +from osprey.worker.lib.config import Config +from udfs.text_contains import TextContains + + +def test_register_udfs_returns_text_contains() -> None: + assert TextContains in register_async_plugins.register_udfs() + + +def test_register_async_output_sinks_returns_example_sink() -> None: + sinks = register_async_plugins.register_async_output_sinks(config=cast(Config, None)) + assert len(sinks) == 1 + assert isinstance(sinks[0], ExampleAsyncOutputSink) + assert isinstance(sinks[0], AsyncBaseOutputSink) + + +def test_example_async_plugin_entry_point_is_registered() -> None: + eps = entry_points(group='osprey_async_plugin') + assert any(ep.value == 'register_async_plugins' for ep in eps), ( + 'example async plugin must be discoverable via the osprey_async_plugin entry-point group' + ) diff --git a/osprey_async_worker/test_data/input.jsonl b/osprey_async_worker/test_data/input.jsonl new file mode 100644 index 0000000..122a925 --- /dev/null +++ b/osprey_async_worker/test_data/input.jsonl @@ -0,0 +1,3 @@ +{"id": 1, "name": "test_action", "data": {"user_id": "123", "event_type": "create_post", "message": "hello world"}} +{"id": 2, "name": "test_action", "data": {"user_id": "456", "event_type": "create_post", "message": "this is a much longer message that definitely exceeds one hundred characters in total length because we need to test the rule evaluation properly with the async executor"}} +{"id": 3, "name": "test_action", "data": {"user_id": "789", "event_type": "update_profile", "message": "short"}} diff --git a/osprey_async_worker/test_data/rules/main.sml b/osprey_async_worker/test_data/rules/main.sml new file mode 100644 index 0000000..66053d7 --- /dev/null +++ b/osprey_async_worker/test_data/rules/main.sml @@ -0,0 +1 @@ +Require(rule='rules/test_rule.sml') diff --git a/osprey_async_worker/test_data/rules/rules/test_rule.sml b/osprey_async_worker/test_data/rules/rules/test_rule.sml new file mode 100644 index 0000000..f8a6de3 --- /dev/null +++ b/osprey_async_worker/test_data/rules/rules/test_rule.sml @@ -0,0 +1,16 @@ +UserId: Entity[str] = EntityJson( + type='User', + path='$.user_id', + coerce_type=True +) + +EventType: str = JsonData(path='$.event_type', coerce_type=True) +ActionName = GetActionName() +MessageLength = StringLength(value=JsonData(path='$.message', coerce_type=True)) + +IsLongMessage = Rule( + when_all=[ + MessageLength > 100, + ], + description='Message is longer than 100 characters', +) diff --git a/osprey_coordinator/src/consumer/pubsub.rs b/osprey_coordinator/src/consumer/pubsub.rs index 9b5331c..b89ac0e 100644 --- a/osprey_coordinator/src/consumer/pubsub.rs +++ b/osprey_coordinator/src/consumer/pubsub.rs @@ -117,7 +117,11 @@ async fn create_pubsub_subscription_client( ) -> SubscriberClient> { let emulator_host = std::env::var("PUBSUB_EMULATOR_HOST").ok(); - let timeout = Duration::from_secs(5); + let timeout_secs: u64 = std::env::var("PUBSUB_CHANNEL_TIMEOUT_SECS") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(5); + let timeout = Duration::from_secs(timeout_secs); if let Some(emulator_host) = emulator_host { tracing::info!("Creating subscription client to emulator"); @@ -179,6 +183,14 @@ pub async fn start_pubsub_subscriber( .unwrap_or("5000".to_string()) .parse::() .unwrap(); + let min_lease_extension_secs = std::env::var("PUBSUB_MIN_LEASE_EXTENSION_SECS") + .unwrap_or_else(|_| "30".to_string()) + .parse::() + .unwrap(); + let max_lease_extension_secs = std::env::var("PUBSUB_MAX_LEASE_EXTENSION_SECS") + .unwrap_or_else(|_| "600".to_string()) + .parse::() + .unwrap(); let config = ConsumerConfig::default(); let max_time_to_send_to_async_queue = config.max_time_to_send_to_async_queue; @@ -191,7 +203,8 @@ pub async fn start_pubsub_subscriber( let flow_control = FlowControl::default() .set_max_messages(max_messages) .set_max_processing_messages(max_processing_messages) - .set_max_bytes(1024 * 1024 * 1024); + .set_max_bytes(1024 * 1024 * 1024) + .set_duration_per_lease_extension(min_lease_extension_secs, max_lease_extension_secs); StreamingPullManager::new( subscriber_client, subscription_name, diff --git a/osprey_coordinator/src/coordinator_metrics.rs b/osprey_coordinator/src/coordinator_metrics.rs index 409948b..2bb67a3 100644 --- a/osprey_coordinator/src/coordinator_metrics.rs +++ b/osprey_coordinator/src/coordinator_metrics.rs @@ -35,6 +35,8 @@ define_metrics!(OspreyCoordinatorMetrics, [ sync_classification_failure_oneshot_dropped => StaticCounter("sync_classification_failure", ["error" => "oneshot_dropped"]), // How many sync actions have failed because we couldn't send to the priority queue sync_classification_failure_pq_send => StaticCounter("sync_classification_failure", ["error" => "pq_send"]), + // How many sync actions were fast-rejected because this pod is shutting down + sync_classification_failure_shutting_down => StaticCounter("sync_classification_failure", ["error" => "shutting_down"]), sync_classification_failure_label_service => StaticCounter("sync_classification_failure", ["error" => "label_service"]), // How many acks for sync actions @@ -63,4 +65,8 @@ define_metrics!(OspreyCoordinatorMetrics, [ // How many times action ID generation from snowflake is used (when pubsub_action.id is None) action_id_snowflake_generation_json => StaticCounter("action_id_snowflake_generation", ["proto"=>"false"]), action_id_snowflake_generation_proto => StaticCounter("action_id_snowflake_generation", ["proto"=>"true"]), + + // Bidi stream dispatch counters (for throughput investigation) + bidi_actions_sent => StaticCounter("bidi_stream.actions_sent"), + bidi_acks_received => StaticCounter("bidi_stream.acks_received"), ]); diff --git a/osprey_coordinator/src/main.rs b/osprey_coordinator/src/main.rs index 0c61d0b..0673d1d 100644 --- a/osprey_coordinator/src/main.rs +++ b/osprey_coordinator/src/main.rs @@ -24,6 +24,7 @@ mod tonic_mock; use anyhow::Result; use clap::Parser; use proto::osprey_coordinator_sync_action::osprey_coordinator_sync_action_service_server::OspreyCoordinatorSyncActionServiceServer; +use std::sync::atomic::AtomicBool; use std::sync::Arc; use std::time::Duration; @@ -61,6 +62,12 @@ struct CliOptions { env = "SNOWFLAKE_API_ENDPOINT" )] snowflake_api_endpoint: String, + #[arg( + long, + default_value = "osprey_coordinator", + env = "OSPREY_COORDINATOR_SERVICE_NAME" + )] + service_name: String, } #[tokio::main] @@ -87,11 +94,14 @@ async fn main() -> Result<()> { metrics.clone(), )); + let is_shutting_down = Arc::new(AtomicBool::new(false)); + let osprey_coordinator_sync_action_service = OspreyCoordinatorSyncActionServiceServer::new(sync_action_rpc::SyncActionServer::new( snowflake_client.clone(), priority_queue_sender.clone(), metrics.clone(), + is_shutting_down.clone(), )); let consumer_type = std::env::var("OSPREY_COORDINATOR_CONSUMER_TYPE").ok(); @@ -134,15 +144,23 @@ async fn main() -> Result<()> { } }; + let bidi_service_name = opts.service_name.clone(); + let sync_action_service_name = format!("{}_sync_action", opts.service_name); + tracing::info!( + bidi_service_name = %bidi_service_name, + sync_action_service_name = %sync_action_service_name, + "registering coordinator services in etcd" + ); + let grpc_bidi_stream_service_fut = pigeon::serve( osprey_coordinator_grpc_bidi_stream_service, - "osprey_coordinator", + &bidi_service_name, opts.bidi_stream_port, Duration::from_secs(30), ); let sync_action_service_fut = pigeon::serve( osprey_coordinator_sync_action_service, - "osprey_coordinator_sync_action", + &sync_action_service_name, opts.sync_action_port, Duration::from_secs(60), ); @@ -154,6 +172,7 @@ async fn main() -> Result<()> { shutdown_handler::spawn_shutdown_handler( priority_queue_sender.clone(), priority_queue_receiver.clone(), + is_shutting_down.clone(), ); tracing::info!("starting consumer/bidi stream/sync classification rpc"); diff --git a/osprey_coordinator/src/osprey_bidirectional_stream.rs b/osprey_coordinator/src/osprey_bidirectional_stream.rs index 6ba86cf..b6200ae 100644 --- a/osprey_coordinator/src/osprey_bidirectional_stream.rs +++ b/osprey_coordinator/src/osprey_bidirectional_stream.rs @@ -87,6 +87,7 @@ fn handle_action_request( ack_or_nack )), (ActionRequest::AckOrNack(ack_or_nack), ClientState::OutstandingAction(state)) => { + metrics.bidi_acks_received.incr(); let duration = Instant::now().duration_since(state.send_time); metrics.action_outstanding_duration.record(duration); state.action_acker.ack_or_nack( @@ -145,6 +146,7 @@ async fn handle_request( .record(Instant::now().duration_since(priority_queue_receive_start_time)); let (action, action_acker) = ackable_action.into_action(); sender.send(Ok(action)).await?; + metrics.bidi_actions_sent.incr(); Ok(UpdateClientStateOrDisconnect::UpdateClientState( ClientState::OutstandingAction(OutstandingActionState { action_acker, diff --git a/osprey_coordinator/src/pigeon/mod.rs b/osprey_coordinator/src/pigeon/mod.rs index c5a8ccc..09e4359 100644 --- a/osprey_coordinator/src/pigeon/mod.rs +++ b/osprey_coordinator/src/pigeon/mod.rs @@ -50,7 +50,7 @@ pub use tonic::server::NamedService; pub async fn serve( grpc_service: GS, - service_name: &'static str, + service_name: &str, service_port: u16, announce_delay: Duration, ) -> Result<()> diff --git a/osprey_coordinator/src/priority_queue.rs b/osprey_coordinator/src/priority_queue.rs index 9d3dcdb..de10a68 100644 --- a/osprey_coordinator/src/priority_queue.rs +++ b/osprey_coordinator/src/priority_queue.rs @@ -197,8 +197,23 @@ impl PriorityQueueReceiver { } pub fn nack_all_async(&self) { + Self::nack_all(&self.async_receiver); + } + + /// Drain the sync queue and nack each pending action. Surfaces to the sync + /// RPC handler as `AckOrNack::Nack`, which it maps to + /// `tonic::Status::aborted("action nacked")` — retryable by the client on a + /// different coordinator pod. Called on shutdown so queued-but-undispatched + /// sync requests don't hang to the per-request timeout and then return + /// `internal("acking onshot dropped")` when the oneshot is finally torn + /// down. + pub fn nack_all_sync(&self) { + Self::nack_all(&self.sync_receiver); + } + + fn nack_all(receiver: &async_channel::Receiver) { loop { - match self.async_receiver.try_recv() { + match receiver.try_recv() { Ok(action) => match action.acking_oneshot_sender.send(AckOrNack::Nack) { Ok(_) => (), Err(_) => println!( diff --git a/osprey_coordinator/src/pub_sub_streaming_pull/flow_control.rs b/osprey_coordinator/src/pub_sub_streaming_pull/flow_control.rs index 7da0420..7498747 100644 --- a/osprey_coordinator/src/pub_sub_streaming_pull/flow_control.rs +++ b/osprey_coordinator/src/pub_sub_streaming_pull/flow_control.rs @@ -1,7 +1,15 @@ -/// The minimum and default value for [`FlowControl::min_duration_per_please_extension`] +/// The minimum allowed value for [`FlowControl::min_duration_per_lease_extension`]. +/// Below this the streaming-pull `stream_ack_deadline_seconds` becomes too aggressive +/// and redelivery cascades trip on brief dispatch-queue stalls. const MIN_LEASE_EXTENSION_DURATION_SECS: u32 = 10; -/// The maximum and default value for [`FlowControl::max_duration_per_please_extension`] +/// The default value for [`FlowControl::min_duration_per_lease_extension`]. +/// This sets both the streaming pull's initial ack deadline and the modack interval — +/// keeping it well above the floor avoids the redelivery loop that the previous 10s +/// default exposed during high-load periods. +const DEFAULT_MIN_LEASE_EXTENSION_DURATION_SECS: u32 = 30; + +/// The maximum and default value for [`FlowControl::max_duration_per_lease_extension`] const MAX_LEASE_EXTENSION_DURATION_SECS: u32 = 600; /// Flow control is used to control the buffering of messages for a given [`crate::StreamingPullManager`]. @@ -26,7 +34,7 @@ impl FlowControl { max_bytes: 100 * 1024 * 1024, max_messages: 1000, max_processing_messages: 1000, - min_duration_per_lease_extension: MIN_LEASE_EXTENSION_DURATION_SECS, + min_duration_per_lease_extension: DEFAULT_MIN_LEASE_EXTENSION_DURATION_SECS, max_duration_per_lease_extension: MAX_LEASE_EXTENSION_DURATION_SECS, } } diff --git a/osprey_coordinator/src/shutdown_handler.rs b/osprey_coordinator/src/shutdown_handler.rs index 9ad7b09..e041ac6 100644 --- a/osprey_coordinator/src/shutdown_handler.rs +++ b/osprey_coordinator/src/shutdown_handler.rs @@ -1,17 +1,36 @@ use crate::priority_queue::{PriorityQueueReceiver, PriorityQueueSender}; use crate::signals; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; pub fn spawn_shutdown_handler( priority_queue_sender: PriorityQueueSender, priority_queue_receiver: PriorityQueueReceiver, + is_shutting_down: Arc, ) { tokio::spawn(async move { tracing::info!("shutdown handler spawned - waiting on exit signal"); signals::exit_signal().await; tracing::info!("got exit signal"); + // Flip the shutting-down flag first. The sync RPC handler reads this + // and fast-rejects new requests with Status::unavailable, which the + // gRPC client retries against a different coordinator pod. Setting + // this before anything else means we stop taking on new in-flight + // work the moment SIGTERM arrives, even while load-balancer health + // propagation lags. + is_shutting_down.store(true, Ordering::Release); + // Drain everything queued-but-undispatched. Sync nacks bubble up to + // the sync RPC handler as Status::aborted, which the client can + // retry on a different coordinator pod. Async nacks trigger immediate + // pubsub redelivery rather than waiting for the lease to expire. + priority_queue_receiver.nack_all_sync(); priority_queue_receiver.nack_all_async(); - tracing::info!("nacked all outstanding async actions"); - tokio::time::sleep(tokio::time::Duration::from_secs(15)).await; + tracing::info!("nacked all queued sync + async actions"); + // Hold the channel open while workers ack dispatched-but-not-yet-acked + // actions over bidi. At typical worker latencies of ~150ms p95, 30s + // gives ~200x the processing window for in-flight actions to drain + // naturally before the channel close tears down the bidi streams. + tokio::time::sleep(tokio::time::Duration::from_secs(30)).await; priority_queue_sender.close(); tracing::info!("closed priority queue"); }); diff --git a/osprey_coordinator/src/sync_action_rpc.rs b/osprey_coordinator/src/sync_action_rpc.rs index 8b6678e..4c61af7 100644 --- a/osprey_coordinator/src/sync_action_rpc.rs +++ b/osprey_coordinator/src/sync_action_rpc.rs @@ -8,6 +8,7 @@ use crate::{ proto::{self, osprey_coordinator_sync_action}, }; use anyhow::{anyhow, Context, Result}; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; use tokio::time::Instant; @@ -19,6 +20,7 @@ pub(crate) struct SyncActionServer { snowflake_client: Arc, priority_queue_sender: PriorityQueueSender, metrics: Arc, + is_shutting_down: Arc, } impl SyncActionServer { @@ -26,11 +28,13 @@ impl SyncActionServer { snowflake_client: Arc, priority_queue_sender: PriorityQueueSender, metrics: Arc, + is_shutting_down: Arc, ) -> SyncActionServer { SyncActionServer { snowflake_client, priority_queue_sender, metrics, + is_shutting_down, } } } @@ -81,6 +85,20 @@ impl SyncActionServer { action_request: &ProcessActionRequest, ) -> Result, tonic::Status> { + // Fast-reject new RPCs once this pod is draining. Returns `Unavailable` + // before the request is enqueued or a snowflake is allocated. The + // gRPC client retries `Unavailable`, routing the retry to a healthier + // coordinator pod. Without this, + // requests landing during the shutdown window queue up, await workers + // whose bidi streams are about to tear down, and end up hitting the + // client-side deadline as `DEADLINE_EXCEEDED`. + if self.is_shutting_down.load(Ordering::Acquire) { + self.metrics + .sync_classification_failure_shutting_down + .incr(); + return Err(tonic::Status::unavailable("coordinator draining")); + } + let unvalidated_action_id = action_request.action_id; let osprey_coordinator_action = match create_osprey_coordinator_action( diff --git a/osprey_ui/src/components/event_stream/FeatureSelectModal.tsx b/osprey_ui/src/components/event_stream/FeatureSelectModal.tsx index 06b7ef1..cbece04 100644 --- a/osprey_ui/src/components/event_stream/FeatureSelectModal.tsx +++ b/osprey_ui/src/components/event_stream/FeatureSelectModal.tsx @@ -76,12 +76,6 @@ const FeatureSelectModal = () => { return featureCategory.substring(0, featureCategory.lastIndexOf('.')).replace(/_|\//g, ' '); }; - const isSetSelected = (feature: string) => { - if (!useCustomFeatures) return true; - - return selectedFeatures.has(feature as string); - }; - const renderModalTitle = () => { return ( <> @@ -136,7 +130,7 @@ const FeatureSelectModal = () => { value={feature} disabled={!useCustomFeatures} onChange={handleSelectFeature} - checked={isSetSelected(feature)} + checked={selectedFeatures.has(feature as string)} className={styles.featureCheckboxRow} > {feature} diff --git a/osprey_worker/Dockerfile b/osprey_worker/Dockerfile index 39f9aa2..725f081 100644 --- a/osprey_worker/Dockerfile +++ b/osprey_worker/Dockerfile @@ -31,24 +31,33 @@ ADD pyproject.toml /osprey/pyproject.toml ADD README.md /osprey/README.md ADD LICENSE.md /osprey/LICENSE.md -# Create workspace structure with pyproject.toml files for uv sync to work +# Create workspace structure with pyproject.toml files for uv sync to work. +# osprey_async_worker is a uv workspace member, so its pyproject must be present +# for `uv sync --locked` to validate the lockfile. This (stable, gevent) image +# does NOT install it — see --no-install-package below. The experimental asyncio +# worker ships in its own image: osprey_async_worker/Dockerfile. ADD osprey_rpc/pyproject.toml /osprey/osprey_rpc/pyproject.toml ADD osprey_worker/pyproject.toml /osprey/osprey_worker/pyproject.toml +ADD osprey_async_worker/pyproject.toml /osprey/osprey_async_worker/pyproject.toml ADD example_plugins/pyproject.toml /osprey/example_plugins/pyproject.toml # Create minimal package structure required by uv RUN mkdir -p /osprey/osprey_worker /osprey/osprey_rpc /osprey/example_plugins/src && \ touch /osprey/osprey_worker/__init__.py /osprey/osprey_rpc/__init__.py /osprey/example_plugins/src/__init__.py -# Install dependencies first (this layer will be cached when only source code changes) +# Install dependencies first (this layer will be cached when only source code changes). +# Exclude osprey_async_worker: the asyncio worker is experimental and lives in its own +# image, so the gevent worker stays free of it (osprey.async_worker is not importable +# here, so execution_context's optional async import cleanly falls back to None). RUN pip install --upgrade pip uv && \ - uv sync --locked --python=$(which python3.11) && \ + uv sync --locked --no-install-package osprey-async-worker --python=$(which python3.11) && \ pip cache purge && uv cache clean # https://tld.readthedocs.io/en/latest/#update-the-list-of-tld-names RUN . .venv/bin/activate && update-tld-names -# Add source code after dependencies are installed +# Add source code after dependencies are installed (osprey_async_worker is +# intentionally omitted — it is not installed in this image). ADD example_rules /osprey/example_rules ADD osprey_worker /osprey/osprey_worker ADD osprey_rpc /osprey/osprey_rpc diff --git a/osprey_worker/src/osprey/engine/ast/grammar.py b/osprey_worker/src/osprey/engine/ast/grammar.py index 34ba793..88ec69d 100644 --- a/osprey_worker/src/osprey/engine/ast/grammar.py +++ b/osprey_worker/src/osprey/engine/ast/grammar.py @@ -4,10 +4,9 @@ from collections import defaultdict from dataclasses import dataclass, field, replace from enum import Enum from pathlib import Path +from threading import Lock from typing import ClassVar, Dict, Optional, Sequence, TypeVar, Union -from gevent.lock import Semaphore - # TODO: Uncomment logging when we have a logging system # from osprey.worker.ui_api.lib.osprey_shared.logging import get_logger from osprey.engine.utils.types import add_slots, cached_property @@ -16,7 +15,7 @@ from osprey.engine.utils.types import add_slots, cached_property # Will this leak memory? Maybe # TODO(old): put this stuff back in cached_property parsed_ast_root_cache: Dict['Source', 'Root'] = {} -ast_root_lock_cache: Dict['Source', Semaphore] = defaultdict(lambda: Semaphore()) +ast_root_lock_cache: Dict['Source', Lock] = defaultdict(Lock) # logger = get_logger() diff --git a/osprey_worker/src/osprey/engine/executor/execution_context.py b/osprey_worker/src/osprey/engine/executor/execution_context.py index 88c00fb..35e3c43 100644 --- a/osprey_worker/src/osprey/engine/executor/execution_context.py +++ b/osprey_worker/src/osprey/engine/executor/execution_context.py @@ -4,6 +4,7 @@ import traceback from collections import defaultdict from dataclasses import dataclass, field from datetime import datetime +from functools import lru_cache from typing import ( TYPE_CHECKING, Any, @@ -12,11 +13,13 @@ from typing import ( Iterable, List, Mapping, + Optional, Sequence, Set, Type, TypeAlias, TypeVar, + cast, ) from google.protobuf.timestamp_pb2 import Timestamp @@ -27,13 +30,33 @@ from osprey.engine.executor.custom_extracted_features import ( ) from osprey.engine.executor.dependency_chain import DependencyChain from osprey.engine.executor.execution_graph import ExecutionGraph -from osprey.engine.executor.external_service_utils import ExternalService, ExternalServiceAccessor, KeyT, ValueT +from osprey.engine.executor.external_service_utils_base import ( + ExternalService, + KeyT, + PlainExternalServiceAccessor, + ValueT, +) + +if TYPE_CHECKING: + from osprey.engine.executor.external_service_utils import ExternalServiceAccessor + +try: + from osprey.async_worker.lib.external_service import ( + AsyncExternalService, + ) + from osprey.async_worker.lib.external_service import ( + ExternalServiceAccessor as AsyncExternalServiceAccessor, + ) +except ImportError: + AsyncExternalService = None # type: ignore[assignment,misc] + AsyncExternalServiceAccessor = None # type: ignore[assignment,misc] from osprey.engine.executor.topological_sorter import TopologicalSorter from osprey.engine.executor.udf_execution_helpers import HasHelperInternal, HelperT, UDFHelpers from osprey.engine.language_types.effects import ( EffectBase, EffectToCustomExtractedFeatureBase, ) +from osprey.engine.language_types.labels import LabelEffect, LabelStatus from osprey.engine.language_types.post_execution_convertible import PostExecutionConvertible from osprey.engine.language_types.verdicts import VerdictEffect from osprey.engine.utils.types import add_slots, cached_property @@ -45,6 +68,32 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) + +@lru_cache(maxsize=1) +def _is_gevent_patched() -> bool: + try: + from gevent.monkey import is_module_patched + + return is_module_patched('socket') + except ImportError: + return False + + +def _label_effect_takes_effect(effect: LabelEffect) -> bool: + """True iff a LabelAdd effect would actually be applied for this action. + + A LabelAdd is filtered out when it has status REMOVED, was suppressed + (e.g. an `apply_if` AST that failed to evaluate), or is gated by a + `dependent_rule` whose value is falsy. Filtered-out effects are not + surfaced to sync callers as entity verdicts. + """ + return ( + effect.status == LabelStatus.ADDED + and not effect.suppressed + and (effect.dependent_rule is None or effect.dependent_rule.value) + ) + + NodeResult: TypeAlias = Result[object, None] @@ -65,6 +114,19 @@ class ExternalServiceException(Exception): """Indicates that an external service call failed or returned unexpected data.""" +@add_slots +@dataclass +class WhenRulesAuditEntry: + """Audit record for a single WhenRules evaluation.""" + + rules_evaluated: List[str] + rules_matched: List[str] + rules_failed: List[str] + effects_emitted: List[str] + effects_failed: int + is_degraded: bool + + class ExecutionContext: """The execution context stores any outputs or intermediate state of an execution.""" @@ -80,9 +142,11 @@ class ExecutionContext: '_effects', '_udf_helpers', '_external_service_accessors_by_getter_id', + '_async_external_service_accessors_by_getter_id', '_dependency_dag', '_chain_by_id', '_custom_extracted_features', + '_rule_audit_entries', ) def __init__(self, execution_graph: ExecutionGraph, action: 'Action', helpers: UDFHelpers): @@ -97,11 +161,13 @@ class ExecutionContext: self._visited_executions: Set[DependencyChain] = set() # a k/v store of effects, by effect type self._effects: DefaultDict[Type[EffectBase], List[EffectBase]] = defaultdict(list) - self._external_service_accessors_by_getter_id: Dict[int, ExternalServiceAccessor[Any, Any]] = {} + self._external_service_accessors_by_getter_id: Dict[int, Any] = {} + self._async_external_service_accessors_by_getter_id: Dict[int, Any] = {} self._dependency_dag = TopologicalSorter() self._chain_by_id: Dict[int, DependencyChain] = {} # feature name -> serializable feature self._custom_extracted_features: Dict[str, Any] = {} + self._rule_audit_entries: List[WhenRulesAuditEntry] = [] self.enqueue_source(execution_graph.get_entry_point()) @@ -206,6 +272,12 @@ class ExecutionContext: def get_effects(self) -> Mapping[Type[EffectBase], Sequence[EffectBase]]: return dict(self._effects) + def add_rule_audit_entry(self, entry: WhenRulesAuditEntry) -> None: + self._rule_audit_entries.append(entry) + + def get_rule_audit_entries(self) -> List[WhenRulesAuditEntry]: + return list(self._rule_audit_entries) + def add_custom_extracted_feature( self, custom_extracted_feature: CustomExtractedFeature[Any], error_on_duplicate_key: bool = True ) -> None: @@ -282,15 +354,37 @@ class ExecutionContext: def get_external_service_accessor( self, external_service: ExternalService[KeyT, ValueT] - ) -> ExternalServiceAccessor[KeyT, ValueT]: + ) -> 'ExternalServiceAccessor[KeyT, ValueT]': """Given an external service, wraps that service in an accessor that ensures that requests to the service are - cached and debounced by key within this execution.""" + cached and debounced by key within this execution. + + Returns PlainExternalServiceAccessor (no gevent dependency) when gevent is not monkey-patched, + allowing sync ExternalService calls to work in the async worker's thread pool. + """ # No need to lock since not doing any IO accessor = self._external_service_accessors_by_getter_id.get(id(external_service)) if accessor is None: - accessor = ExternalServiceAccessor(external_service) + if _is_gevent_patched(): + from osprey.engine.executor.external_service_utils import ExternalServiceAccessor + + accessor = ExternalServiceAccessor(external_service) + else: + accessor = PlainExternalServiceAccessor(external_service) self._external_service_accessors_by_getter_id[id(external_service)] = accessor + # The cache dict is typed Dict[int, Any] since it holds both the gevent + # ExternalServiceAccessor and PlainExternalServiceAccessor; both satisfy the + # ExternalServiceAccessor[KeyT, ValueT] accessor interface promised here. + return cast('ExternalServiceAccessor[KeyT, ValueT]', accessor) + + def get_async_external_service_accessor( + self, external_service: 'AsyncExternalService[KeyT, ValueT]' + ) -> 'AsyncExternalServiceAccessor[KeyT, ValueT]': + """Async version of get_external_service_accessor, using asyncio.Future for caching.""" + accessor = self._async_external_service_accessors_by_getter_id.get(id(external_service)) + if accessor is None: + accessor = AsyncExternalServiceAccessor(external_service) + self._async_external_service_accessors_by_getter_id[id(external_service)] = accessor return accessor @@ -359,6 +453,8 @@ class ExecutionResult: # some kind of typing when outputing `validator_results` validator_results: Dict[Any, Any] = field(default_factory=dict) sample_rate: int = 100 + trace_id: Optional[str] = None + rule_audit_entries: Sequence[WhenRulesAuditEntry] = field(default_factory=list) def add_custom_extracted_feature(self, custom_extracted_feature: CustomExtractedFeature[Any]) -> None: name = custom_extracted_feature.feature_name() @@ -393,11 +489,20 @@ class ExecutionResult: """ returns a pb2 protobuf of the verdicts declared by the action, along with some extra metadata~ ╰(*°▽°*)╯ + + Also surfaces effective LabelAdd effects as entity verdicts of the form + "//", so sync callers reading + entity_has_label() see labels applied during the action. """ return Verdicts( action_id=self.action.action_id, action_name=self.action.action_name, - verdicts=[v.verdict for v in self.verdicts], + verdicts=[v.verdict for v in self.verdicts] + + [ + f'{e.entity.type}/{e.entity.id}/{e.name}' + for e in self.effects.get(LabelEffect, []) + if isinstance(e, LabelEffect) and _label_effect_takes_effect(e) + ], timestamp=self._get_timestamp_pb2_proto(), ) diff --git a/osprey_worker/src/osprey/engine/executor/execution_visualizer.py b/osprey_worker/src/osprey/engine/executor/execution_visualizer.py index c2523e5..435931d 100644 --- a/osprey_worker/src/osprey/engine/executor/execution_visualizer.py +++ b/osprey_worker/src/osprey/engine/executor/execution_visualizer.py @@ -2,7 +2,7 @@ import copy import enum import pathlib import tempfile -from typing import Any, Dict, Optional, Set, Type, Union +from typing import Any, Dict, Optional, Protocol, Set, Type, Union from graphviz import Digraph from osprey.engine.ast.grammar import ( @@ -19,6 +19,7 @@ from osprey.engine.ast.grammar import ( Number, String, ) +from osprey.engine.config.config_subkey_handler import ModelT from osprey.engine.executor.graph_data import GraphData, LabelType, Node, NodeType from osprey.engine.stdlib.configs.labels_config import LabelsConfig from osprey.worker.lib.singletons import ENGINE @@ -26,6 +27,19 @@ from osprey.worker.lib.singletons import ENGINE from .dependency_chain import DependencyChain from .execution_graph import ExecutionGraph + +class EngineLike(Protocol): + """Subset of the osprey engine surface that ``render_graph`` consumes. + + Implemented structurally by both the sync ``OspreyEngine`` and the async + ``AsyncOspreyEngine``; declaring it here lets non-sync callers pass their + own engine without dragging in the sync singleton. + """ + + def get_known_action_names(self) -> Set[str]: ... + def get_config_subkey(self, model_class: Type[ModelT]) -> ModelT: ... + + debug: bool = False @@ -213,18 +227,26 @@ def render_graph( label_names: Optional[Set[str]] = None, # Label names to show show_label_upstream: bool = False, # Upstream for label view show_label_downstream: bool = True, # Downstream for label view + engine: Optional[ + EngineLike + ] = None, # Engine instance to resolve action/label names from; falls back to the sync ENGINE singleton ) -> 'RenderedDigraph': """ Generate a rules vizualization graph based on the provided parameters. - This method makes a call to the OspreyEngine Singleton to grab all valid label and action names. + + Callers running outside the sync osprey worker (e.g. the async ui_api) must pass ``engine`` + explicitly so this function does not trigger the sync ``ENGINE`` singleton bootstrap. Any + object exposing ``get_known_action_names()`` and ``get_config_subkey(LabelsConfig)`` works. """ if action_names is None: action_names = set() if label_names is None: label_names = set() - all_action_names = ENGINE.instance().get_known_action_names() - all_label_names = {key for key in ENGINE.instance().get_config_subkey(LabelsConfig).labels} + if engine is None: + engine = ENGINE.instance() + all_action_names = engine.get_known_action_names() + all_label_names = {key for key in engine.get_config_subkey(LabelsConfig).labels} return _render_graph( all_action_names, diff --git a/osprey_worker/src/osprey/engine/executor/executor.py b/osprey_worker/src/osprey/engine/executor/executor.py index 4af75fc..9a697f7 100644 --- a/osprey_worker/src/osprey/engine/executor/executor.py +++ b/osprey_worker/src/osprey/engine/executor/executor.py @@ -419,6 +419,8 @@ def execute( 'osprey.action_error_count', len(actionable_error_infos), tags=[f'action:{action.action_name}'] ) + trace_id = str(parent_tracer_span.trace_id) if parent_tracer_span else None + result = ExecutionResult( extracted_features=context.get_extracted_features(), action=action, @@ -426,5 +428,7 @@ def execute( validator_results=validator_results, error_infos=unexpected_error_infos, sample_rate=sample_rate, + trace_id=trace_id, + rule_audit_entries=context.get_rule_audit_entries(), ) return result diff --git a/osprey_worker/src/osprey/engine/executor/external_service_utils.py b/osprey_worker/src/osprey/engine/executor/external_service_utils.py index 36c1de7..fb0918f 100644 --- a/osprey_worker/src/osprey/engine/executor/external_service_utils.py +++ b/osprey_worker/src/osprey/engine/executor/external_service_utils.py @@ -1,46 +1,26 @@ -from abc import ABC, abstractmethod -from datetime import datetime, timedelta -from typing import Dict, Generic, Hashable, Optional, Sequence, Tuple, TypeVar, cast +from datetime import datetime +from typing import Dict, Generic, Optional, Sequence, Tuple, cast from gevent.event import AsyncResult -from result import Err, Ok, Result - -KeyT = TypeVar('KeyT', bound=Hashable) -ValueT = TypeVar('ValueT') - - -class ExternalService(ABC, Generic[KeyT, ValueT]): - @abstractmethod - def get_from_service(self, key: KeyT) -> ValueT: - raise NotImplementedError - - # Not abstract because not all services support batching multiple keys - def batch_get_from_service(self, keys: Sequence[KeyT]) -> Sequence[Result[ValueT, Exception]]: - raise NotImplementedError - - def cache_ttl(self) -> Optional[timedelta]: - """ - Returns a time to live for items in the cache. By default, KVs are cached indefinitely. - To have cache entries auto-expire, override this method in your external service definition. - - Note that timedeltas can accept negative values to represent the past, but only on the days field. - You *can* use timedelta(seconds=0) to disable caching, but a negative time delta *ensures* that even - if a time shift occurs (such as daylight savings), the cache_ttl will still be immediate. - - Therefore, to disable the read cache, it is recommended to set this to `timedelta(days=-1)` - """ - return None - - def count_error_once(self) -> bool: - """ - When True, only the caller that initiated the external service call - receives the exception. Subsequent callers that would hit the cached - error receive None instead. +# Re-export base classes for backward compatibility +from osprey.engine.executor.external_service_utils_base import ( # noqa: F811 + ExternalService, + KeyT, + PlainExternalServiceAccessor, + ValueT, + _CacheEntry, +) +from result import Err, Ok, Result - Only enable this when ValueT is Optional and None is a safe fallback. - """ - return False +__all__ = [ + 'ExternalService', + 'ExternalServiceAccessor', + 'PlainExternalServiceAccessor', + '_CacheEntry', + 'KeyT', + 'ValueT', +] class ExternalServiceAccessor(Generic[KeyT, ValueT]): diff --git a/osprey_worker/src/osprey/engine/executor/external_service_utils_base.py b/osprey_worker/src/osprey/engine/executor/external_service_utils_base.py new file mode 100644 index 0000000..41c45f4 --- /dev/null +++ b/osprey_worker/src/osprey/engine/executor/external_service_utils_base.py @@ -0,0 +1,130 @@ +"""Base external service utilities — no gevent dependency. + +Contains ExternalService ABC, PlainExternalServiceAccessor, and cache helpers. +The gevent-dependent ExternalServiceAccessor remains in external_service_utils.py. +""" + +from abc import ABC, abstractmethod +from datetime import datetime, timedelta +from typing import Any, Dict, Generic, Hashable, Optional, Sequence, Tuple, TypeVar, cast + +from result import Err, Ok, Result + +KeyT = TypeVar('KeyT', bound=Hashable) +ValueT = TypeVar('ValueT') + + +class ExternalService(ABC, Generic[KeyT, ValueT]): + @abstractmethod + def get_from_service(self, key: KeyT) -> ValueT: + raise NotImplementedError + + # Not abstract because not all services support batching multiple keys + def batch_get_from_service(self, keys: Sequence[KeyT]) -> Sequence[Result[ValueT, Exception]]: + raise NotImplementedError + + def cache_ttl(self) -> Optional[timedelta]: + """ + Returns a time to live for items in the cache. By default, KVs are cached indefinitely. + + To have cache entries auto-expire, override this method in your external service definition. + + Note that timedeltas can accept negative values to represent the past, but only on the days field. + You *can* use timedelta(seconds=0) to disable caching, but a negative time delta *ensures* that even + if a time shift occurs (such as daylight savings), the cache_ttl will still be immediate. + + Therefore, to disable the read cache, it is recommended to set this to `timedelta(days=-1)` + """ + return None + + def count_error_once(self) -> bool: + """ + When True, only the caller that initiated the external service call + receives the exception. Subsequent callers that would hit the cached + error receive None instead. + + Only enable this when ValueT is Optional and None is a safe fallback. + """ + return False + + +class _CacheEntry(Generic[ValueT]): + """Value-or-exception container. No gevent, no asyncio.""" + + __slots__ = ('value', 'exception') + + def __init__(self) -> None: + self.value: Any = None + self.exception: Optional[BaseException] = None + + def set_value(self, value: ValueT) -> None: + self.value = value + + def set_exception(self, exc: BaseException) -> None: + self.exception = exc + + def get(self) -> ValueT: + if self.exception is not None: + raise self.exception + return self.value + + +class PlainExternalServiceAccessor(Generic[KeyT, ValueT]): + """ExternalServiceAccessor without gevent. + + Identical caching and count_error_once semantics, but uses a plain + _CacheEntry instead of gevent.event.AsyncResult. Intended for sync + ExternalService calls that run in a thread pool where gevent is not + monkey-patched (e.g. the async worker's legacy batch execution path). + """ + + def __init__(self, service: ExternalService[KeyT, ValueT]): + self._service = service + self._cache: Dict[KeyT, Tuple[_CacheEntry[ValueT], Optional[datetime]]] = {} + + def _is_expired(self, expiration: Optional[datetime]) -> bool: + return expiration is not None and datetime.now() > expiration + + def _expiration(self) -> Optional[datetime]: + ttl = self._service.cache_ttl() + return datetime.now() + ttl if ttl is not None else None + + def get(self, key: KeyT) -> ValueT: + entry = self._cache.get(key) + if entry is not None and not self._is_expired(entry[1]): + return entry[0].get() + + result: _CacheEntry[ValueT] = _CacheEntry() + self._cache[key] = (result, self._expiration()) + try: + result.set_value(self._service.get_from_service(key)) + except Exception as e: + if self._service.count_error_once(): + result.set_value(cast(ValueT, None)) + else: + result.set_exception(e) + raise + return result.get() + + def batch_get(self, keys: Sequence[KeyT]) -> Sequence[Result[ValueT, Exception]]: + non_cached = [k for k in keys if self._cache.get(k) is None or self._is_expired(self._cache[k][1])] + if non_cached: + for k in non_cached: + self._cache[k] = (_CacheEntry(), self._expiration()) + try: + results = self._service.batch_get_from_service(non_cached) + for i, k in enumerate(non_cached): + if results[i].is_ok(): + self._cache[k][0].set_value(results[i].unwrap()) + else: + self._cache[k][0].set_exception(cast(BaseException, results[i].value)) + except Exception as e: + for k in non_cached: + self._cache[k][0].set_exception(e) + + return [ + Ok(self._cache[k][0].get()) + if self._cache[k][0].exception is None + else Err(cast(Exception, self._cache[k][0].exception)) + for k in keys + ] diff --git a/osprey_worker/src/osprey/engine/executor/node_executor/binary_operation_executor.py b/osprey_worker/src/osprey/engine/executor/node_executor/binary_operation_executor.py index 03339fd..3994c05 100644 --- a/osprey_worker/src/osprey/engine/executor/node_executor/binary_operation_executor.py +++ b/osprey_worker/src/osprey/engine/executor/node_executor/binary_operation_executor.py @@ -44,13 +44,31 @@ class BinaryOperationExecutor(BaseNodeExecutor[BinaryOperation, Any]): return [self._node.left, self._node.right] +def _safe_truediv(left: Any, right: Any) -> Any: + if right == 0: + return 0 + return operator.truediv(left, right) + + +def _safe_floordiv(left: Any, right: Any) -> Any: + if right == 0: + return 0 + return operator.floordiv(left, right) + + +def _safe_mod(left: Any, right: Any) -> Any: + if right == 0: + return 0 + return operator.mod(left, right) + + _BINARY_OPERATORS = { Add: operator.add, Subtract: operator.sub, Multiply: operator.mul, - Divide: operator.truediv, - FloorDivide: operator.floordiv, - Modulo: operator.mod, + Divide: _safe_truediv, + FloorDivide: _safe_floordiv, + Modulo: _safe_mod, Pow: operator.pow, LeftShift: operator.lshift, RightShift: operator.rshift, diff --git a/osprey_worker/src/osprey/engine/executor/tests/test_get_verdicts_pb2_proto.py b/osprey_worker/src/osprey/engine/executor/tests/test_get_verdicts_pb2_proto.py new file mode 100644 index 0000000..1ae1b19 --- /dev/null +++ b/osprey_worker/src/osprey/engine/executor/tests/test_get_verdicts_pb2_proto.py @@ -0,0 +1,128 @@ +"""Tests for ExecutionResult.get_verdicts_pb2_proto(). + +Verifies that LabelAdd effects are synthesised into entity verdict strings so +that callers reading entity_has_label() on a sync ProcessAction response see +labels applied during this action. +""" +from datetime import datetime +from typing import Any, Dict, Mapping, Sequence, Type + +import pytest +from osprey.engine.executor.execution_context import Action, ExecutionResult +from osprey.engine.language_types.effects import EffectBase +from osprey.engine.language_types.entities import EntityT +from osprey.engine.language_types.labels import LabelEffect, LabelStatus +from osprey.engine.language_types.rules import RuleT +from osprey.engine.language_types.verdicts import VerdictEffect + + +_ENTITY = EntityT(type='User', id=12345) +_LABEL_NAME = 'require_verified_phone_then_email' + + +def _make_action() -> Action: + return Action( + action_id=1, + action_name='test_action', + data={}, + timestamp=datetime(2026, 1, 1, 12, 0, 0), + ) + + +def _make_result(effects: Mapping[Type[EffectBase], Sequence[EffectBase]]) -> ExecutionResult: + return ExecutionResult( + extracted_features={}, + action=_make_action(), + effects=effects, + error_infos=[], + ) + + +def _make_label_effect( + status: LabelStatus = LabelStatus.ADDED, + suppressed: bool = False, + dependent_rule: RuleT | None = None, +) -> LabelEffect: + return LabelEffect( + entity=_ENTITY, + status=status, + name=_LABEL_NAME, + suppressed=suppressed, + dependent_rule=dependent_rule, + ) + + +def _make_rule(value: bool) -> RuleT: + return RuleT(name='test_rule', value=value, description='', features={}) + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +def test_no_effects_produces_empty_verdicts(): + """Regression: no effects → no verdicts in output.""" + result = _make_result({}) + pb = result.get_verdicts_pb2_proto() + assert list(pb.verdicts) == [] + + +def test_verdict_effects_only_unchanged(): + """Existing VerdictEffects still appear in output when no LabelEffects are present.""" + ve = VerdictEffect(verdict='User/999/some_verdict') + result = _make_result({VerdictEffect: [ve]}) + pb = result.get_verdicts_pb2_proto() + assert list(pb.verdicts) == ['User/999/some_verdict'] + + +def test_label_effect_added_not_suppressed_synthesised(): + """ADDED, not suppressed, no dependent_rule → verdict string synthesised.""" + le = _make_label_effect(status=LabelStatus.ADDED, suppressed=False) + result = _make_result({LabelEffect: [le]}) + pb = result.get_verdicts_pb2_proto() + assert f'User/12345/{_LABEL_NAME}' in list(pb.verdicts) + + +def test_label_effect_suppressed_not_synthesised(): + """Suppressed LabelEffect → NOT synthesised.""" + le = _make_label_effect(suppressed=True) + result = _make_result({LabelEffect: [le]}) + pb = result.get_verdicts_pb2_proto() + assert list(pb.verdicts) == [] + + +def test_label_effect_dependent_rule_false_not_synthesised(): + """dependent_rule.value=False → NOT synthesised.""" + le = _make_label_effect(dependent_rule=_make_rule(value=False)) + result = _make_result({LabelEffect: [le]}) + pb = result.get_verdicts_pb2_proto() + assert list(pb.verdicts) == [] + + +def test_label_effect_dependent_rule_true_synthesised(): + """dependent_rule.value=True → verdict synthesised.""" + le = _make_label_effect(dependent_rule=_make_rule(value=True)) + result = _make_result({LabelEffect: [le]}) + pb = result.get_verdicts_pb2_proto() + assert f'User/12345/{_LABEL_NAME}' in list(pb.verdicts) + + +def test_label_effect_removed_not_synthesised(): + """REMOVED LabelEffect → NOT synthesised (only ADDED emits positive-signal verdicts).""" + le = _make_label_effect(status=LabelStatus.REMOVED) + result = _make_result({LabelEffect: [le]}) + pb = result.get_verdicts_pb2_proto() + assert list(pb.verdicts) == [] + + +def test_mix_of_verdict_and_label_effects(): + """Both VerdictEffects and synthesised LabelEffect verdicts appear together.""" + ve = VerdictEffect(verdict='User/999/some_verdict') + le = _make_label_effect(status=LabelStatus.ADDED, suppressed=False) + result = _make_result({VerdictEffect: [ve], LabelEffect: [le]}) + pb = result.get_verdicts_pb2_proto() + verdicts = list(pb.verdicts) + assert 'User/999/some_verdict' in verdicts + assert f'User/12345/{_LABEL_NAME}' in verdicts + assert len(verdicts) == 2 diff --git a/osprey_worker/src/osprey/engine/executor/udf_execution_helpers.py b/osprey_worker/src/osprey/engine/executor/udf_execution_helpers.py index 2978fd1..13f52ba 100644 --- a/osprey_worker/src/osprey/engine/executor/udf_execution_helpers.py +++ b/osprey_worker/src/osprey/engine/executor/udf_execution_helpers.py @@ -3,10 +3,7 @@ from __future__ import annotations from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, Dict, Generic, Hashable, Type, TypeVar, cast -from osprey.engine.executor.external_service_utils import ( - ExternalService, - ExternalServiceAccessor, -) +from osprey.engine.executor.external_service_utils_base import ExternalService if TYPE_CHECKING: from osprey.engine.executor.execution_context import ExecutionContext @@ -23,7 +20,7 @@ class HasHelperInternal(Generic[HelperT]): def accessor_get(self, execution_context: ExecutionContext, key: KeyT, lock: bool = False) -> ValueT: # type: ignore[type-var] udf = cast(HasHelperInternal[ExternalService[KeyT, ValueT]], self) provider: ExternalService[KeyT, ValueT] = execution_context.get_udf_helper(udf) - accessor: ExternalServiceAccessor[KeyT, ValueT] = execution_context.get_external_service_accessor(provider) + accessor = execution_context.get_external_service_accessor(provider) return accessor.get(key) diff --git a/osprey_worker/src/osprey/engine/stdlib/configs/feature_flags_config.py b/osprey_worker/src/osprey/engine/stdlib/configs/feature_flags_config.py index af77065..b74ba28 100644 --- a/osprey_worker/src/osprey/engine/stdlib/configs/feature_flags_config.py +++ b/osprey_worker/src/osprey/engine/stdlib/configs/feature_flags_config.py @@ -8,10 +8,6 @@ from .._registry import register_config_subkey FEATURE_FLAGS_CONFIG_SUBKEY = 'feature_flags' -### Feature Flags ### -WEBHOOKS_USE_PUBSUB = 'WEBHOOKS_USE_PUBSUB' - - class PercentageFlagInfo(BaseModel): value: float description: str = '' diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/categories.py b/osprey_worker/src/osprey/engine/stdlib/udfs/categories.py index a244cb5..6f12771 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/categories.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/categories.py @@ -3,6 +3,7 @@ from enum import Enum # this needs to be combed thru class UdfCategories(str, Enum): + CAST = 'Cast' DATETIME = 'Datetime' DNS = 'DNS' EMAIL = 'Email' @@ -13,6 +14,5 @@ class UdfCategories(str, Enum): HTTP = 'HTTP' IP = 'IP' PHONE = 'Phone' - CAST = 'Cast' RANDOM = 'Random' STRING = 'String' diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/count_regex_matches.py b/osprey_worker/src/osprey/engine/stdlib/udfs/count_regex_matches.py new file mode 100644 index 0000000..b50ae38 --- /dev/null +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/count_regex_matches.py @@ -0,0 +1,52 @@ +import re +from typing import List, Pattern + +from osprey.engine.stdlib.udfs._prelude import ( + ArgumentsBase, + ConstExpr, + ExecutionContext, + UDFBase, + ValidationContext, +) +from osprey.engine.stdlib.udfs.categories import UdfCategories + + +class CountRegexMatchesArguments(ArgumentsBase): + patterns: ConstExpr[List[str]] + """List of regex patterns to evaluate. The UDF returns the count of distinct + patterns that find at least one match in the target.""" + + target: str + """The string to evaluate the patterns against.""" + + case_insensitive: ConstExpr[bool] = ConstExpr.for_default('case_insensitive', False) + """If `True`, all patterns are matched case-insensitively. Default `False`.""" + + +class CountRegexMatches(UDFBase[CountRegexMatchesArguments, int]): + """Returns the number of distinct regex patterns that match the target string. + + Each pattern in `patterns` is tested independently; the result is the count of + patterns with at least one match. Useful for "N-of-M" rule semantics where a rule + should fire only when several distinct keywords/categories appear in the same text. + """ + + category = UdfCategories.STRING + + def __init__(self, validation_context: 'ValidationContext', arguments: CountRegexMatchesArguments): + super().__init__(validation_context, arguments) + + flags = re.IGNORECASE if arguments.case_insensitive.value else 0 + + self._compiled: List[Pattern[str]] = [] + for pattern in arguments.patterns.value: + try: + self._compiled.append(re.compile(pattern, flags)) + except re.error as exc: + validation_context.add_error( + message=f'invalid regex pattern {pattern!r}: {exc}', + span=arguments.patterns.argument_span, + ) + + def execute(self, execution_context: ExecutionContext, arguments: CountRegexMatchesArguments) -> int: + return sum(1 for compiled in self._compiled if compiled.search(arguments.target) is not None) diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/experiments.py b/osprey_worker/src/osprey/engine/stdlib/udfs/experiments.py index 4db8bf2..14f4ab5 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/experiments.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/experiments.py @@ -299,20 +299,5 @@ class ExperimentWhen(UDFBase[ExperimentWhenArguments, List[bool]]): return arguments.extra_arguments[bucket] -class InExperimentArguments(ArgumentsBase): - experiment: ExperimentT - - -class InExperiment(UDFBase[InExperimentArguments, bool]): - """ - Returns True if the entity was assigned to any bucket in the experiment, - False if the entity fell outside all bucket ranges - """ - - category = UdfCategories.ENGINE - - def execute(self, execution_context: ExecutionContext, arguments: InExperimentArguments) -> bool: - return arguments.experiment.resolved_bucket is not NOT_IN_EXPERIMENT_BUCKET - - ExperimentsBucketAssignment = Experiment.build_cls('experiment_bucket_assignment') +ExperimentsBucketAssignment.__doc__ = 'Returns the experiment bucket assignment for the entity, used for A/B testing.' diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/json_data.py b/osprey_worker/src/osprey/engine/stdlib/udfs/json_data.py index cffc5c0..6716a8d 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/json_data.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/json_data.py @@ -20,11 +20,11 @@ class Arguments(ArgumentsBase): Defaults to `True`. If `False`, will gracefully handle both missing and present-but-null values. """ - coerce_type: bool = False + coerce_type: bool = True """Whether to attempt to convert the value to the expected type. - By default `JsonData` just asserts that the value already is the right type. Setting this to `True` can be useful - to, eg, parse a number from a string if the number was too big to represent in JSON. + By default `JsonData` will attempt to coerce the value to the declared type (e.g., parse a number from a string). + If coercion fails, it still raises `InvalidJsonType`. Set to `False` to require exact type matches. """ diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/rules.py b/osprey_worker/src/osprey/engine/stdlib/udfs/rules.py index 82c0f76..93f750d 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/rules.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/rules.py @@ -3,6 +3,7 @@ from typing import List, Optional, cast from osprey.engine.ast import grammar from osprey.engine.ast_validator.validation_utils import add_must_assign_to_variable_error +from osprey.engine.executor.execution_context import WhenRulesAuditEntry from osprey.engine.executor.node_executor.call_executor import CallExecutor from osprey.engine.language_types.effects import EffectBase from osprey.engine.language_types.rules import RuleT @@ -22,6 +23,8 @@ class RuleArguments(ArgumentsBase): class Rule(UDFBase[RuleArguments, RuleT]): + """Defines a named rule with conditions. Evaluates to true when all conditions in when_all are met.""" + name: Optional[str] = None category = UdfCategories.ENGINE @@ -118,6 +121,8 @@ class WhenRulesArguments(ArgumentsBase): class WhenRules(UDFBase[WhenRulesArguments, None]): + """Binds rules to effects. When any of the referenced rules fire, the then= effects are applied.""" + category = UdfCategories.ENGINE def resolve_arguments(self, execution_context: ExecutionContext, call_executor: CallExecutor) -> WhenRulesArguments: @@ -132,29 +137,81 @@ class WhenRules(UDFBase[WhenRulesArguments, None]): # This is important because if we have a list of LabelAdd() calls, we can end up invalidating # them all if any of the calls fails. then_value = [] + then_failed = 0 assert isinstance(then_node, grammar.List), 'BUG: `then` node is not a list!' for item in then_node.items: resolved = execution_context.resolved(item, return_none_for_failed_values=True) if resolved is not None: then_value.append(resolved) + else: + then_failed += 1 # 3. Perform special resolution of the `rules_any` list. Essentially, we are semi-circumventing the # `ListExecutor`, in order to peek inside, and grab each non-failed item within the list. rules_any_value = [] + rules_any_failed = 0 + failed_rule_names: List[str] = [] assert isinstance(rules_any_node, grammar.List), 'BUG: `rules_any node is not a List!' for item in rules_any_node.items: resolved = execution_context.resolved(item, return_none_for_failed_values=True) if resolved is not None: rules_any_value.append(resolved) + else: + rules_any_failed += 1 + if isinstance(item, grammar.Name): + failed_rule_names.append(item.identifier) + else: + failed_rule_names.append('') + + # 4. Store audit state for execute() to consume. + # Safe because WhenRules is synchronous (execute_async=False) and CallExecutor.execute() + # calls resolve_arguments() then execute() without yielding to the event loop. + self._failed_rule_names = failed_rule_names + self._then_failed = then_failed + + # 5. Emit completeness metrics for this WhenRules block. + action_name = execution_context.get_action_name() + is_degraded = rules_any_failed > 0 or then_failed > 0 + metrics.increment( + 'osprey.whenrules_completeness', + tags=[f'action:{action_name}', f'degraded:{is_degraded}'], + ) - # 4. Construct the resolved arguments, based on our custom argument resolution. + # 6. Construct the resolved arguments, based on our custom argument resolution. return cast( WhenRulesArguments, call_executor.unresolved_arguments.update_with_resolved({'then': then_value, 'rules_any': rules_any_value}), ) def execute(self, execution_context: ExecutionContext, arguments: WhenRulesArguments) -> None: + all_rule_names = [rule.name for rule in arguments.rules_any] passing_rules = [rule for rule in arguments.rules_any if rule.value] + passing_names = [rule.name for rule in passing_rules] + assert hasattr(self, '_failed_rule_names'), 'BUG: resolve_arguments() must be called before execute()' + failed_rule_names: List[str] = self._failed_rule_names + then_failed: int = self._then_failed + + effects_emitted: List[str] = [] + if passing_rules: + effects_emitted = [type(o).__name__ for o in arguments.then if isinstance(o, EffectBase)] + + entry = WhenRulesAuditEntry( + rules_evaluated=all_rule_names, + rules_matched=passing_names, + rules_failed=failed_rule_names, + effects_emitted=effects_emitted, + effects_failed=then_failed, + is_degraded=len(failed_rule_names) > 0 or then_failed > 0, + ) + execution_context.add_rule_audit_entry(entry) + + if entry.effects_failed > 0: + has_gap = not entry.is_degraded + metrics.increment( + 'osprey.enforcement_gap', + tags=[f'action:{execution_context.get_action_name()}', f'gap:{has_gap}'], + ) + if not passing_rules: return diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/string.py b/osprey_worker/src/osprey/engine/stdlib/udfs/string.py index 34b872a..2e133a1 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/string.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/string.py @@ -24,6 +24,8 @@ class StringArguments(ArgumentsBase): class StringLength(UDFBase[StringArguments, int]): + """Returns the length of the string.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringArguments) -> int: @@ -31,6 +33,8 @@ class StringLength(UDFBase[StringArguments, int]): class StringToLower(UDFBase[StringArguments, str]): + """Converts the string to lowercase.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringArguments) -> str: @@ -38,6 +42,8 @@ class StringToLower(UDFBase[StringArguments, str]): class StringToUpper(UDFBase[StringArguments, str]): + """Converts the string to uppercase.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringArguments) -> str: @@ -50,6 +56,8 @@ class StringStartsWithArgument(StringArguments): class StringStartsWith(UDFBase[StringStartsWithArgument, bool]): + """Returns true if the string starts with the given prefix.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringStartsWithArgument) -> bool: @@ -61,6 +69,8 @@ class StringEndsWithArgument(StringArguments): class StringEndsWith(UDFBase[StringEndsWithArgument, bool]): + """Returns true if the string ends with the given suffix.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringEndsWithArgument) -> bool: @@ -72,6 +82,8 @@ class StringStripArguments(StringArguments): class StringStrip(UDFBase[StringStripArguments, str]): + """Strips whitespace (or specified characters) from both ends of the string.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringStripArguments) -> str: @@ -79,6 +91,8 @@ class StringStrip(UDFBase[StringStripArguments, str]): class StringRStrip(UDFBase[StringStripArguments, str]): + """Strips whitespace (or specified characters) from the right side of the string.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringStripArguments) -> str: @@ -86,6 +100,8 @@ class StringRStrip(UDFBase[StringStripArguments, str]): class StringLStrip(UDFBase[StringStripArguments, str]): + """Strips whitespace (or specified characters) from the left side of the string.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringStripArguments) -> str: @@ -98,6 +114,8 @@ class StringReplaceArguments(StringArguments): class StringReplace(UDFBase[StringReplaceArguments, str]): + """Replaces all occurrences of a substring with another string.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringReplaceArguments) -> str: @@ -109,6 +127,8 @@ class StringJoinArguments(StringArguments): class StringJoin(UDFBase[StringJoinArguments, str]): + """Joins a list of strings using the given separator.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringJoinArguments) -> str: @@ -121,6 +141,8 @@ class StringSplitArguments(StringArguments): class StringSplit(UDFBase[StringSplitArguments, List[str]]): + """Splits the string by a delimiter into a list of strings.""" + category = UdfCategories.STRING def execute(self, execution_context: ExecutionContext, arguments: StringSplitArguments) -> List[str]: @@ -133,6 +155,8 @@ class StringSliceArguments(StringArguments): class StringSlice(UDFBase[StringSliceArguments, str]): + """Returns a substring from start index to end index.""" + category = UdfCategories.STRING def __init__(self, validation_context: ValidationContext, arguments: StringSliceArguments): diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_count_regex_matches.py b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_count_regex_matches.py new file mode 100644 index 0000000..36e68ef --- /dev/null +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_count_regex_matches.py @@ -0,0 +1,77 @@ +from typing import Any, Callable, List, Optional + +import pytest +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.conftest import CheckFailureFunction, ExecuteFunction, RunValidationFunction +from osprey.engine.stdlib.udfs.count_regex_matches import CountRegexMatches +from osprey.engine.udf.registry import UDFRegistry + +pytestmark: List[Callable[[Any], Any]] = [ + pytest.mark.use_validators([ValidateCallKwargs, UniqueStoredNames]), + pytest.mark.use_udf_registry(UDFRegistry.with_udfs(CountRegexMatches)), +] + + +@pytest.mark.parametrize( + 'patterns, target, expected', + ( + # No patterns match + (['foo', 'bar'], 'hello world', 0), + # One of two matches + (['foo', 'bar'], 'foo world', 1), + # Both match + (['foo', 'bar'], 'foo bar baz', 2), + # Multiple hits within one pattern still count as 1 + (['foo'], 'foo foo foo', 1), + # Distinct patterns that share a hit each contribute 1 + (['fo+', 'f.o'], 'foo', 2), + # Anchored patterns + (['^abc', 'xyz$'], 'abc and xyz', 2), + (['^abc', 'xyz$'], 'zabc xyzx', 0), + # Empty target with non-trivial patterns + (['foo', 'bar'], '', 0), + # Empty pattern matches anything (re.search with '' returns a match at pos 0) + ([''], 'abc', 1), + ), +) +def test_counts_matching_patterns(execute: ExecuteFunction, patterns: List[str], target: str, expected: int) -> None: + result = execute( + f""" + Count = CountRegexMatches(patterns={patterns!r}, target="{target}") + """ + ) + assert result == {'Count': expected} + + +@pytest.mark.parametrize( + 'patterns, target, case_insensitive, expected', + ( + # Default (case-sensitive) + (['foo', 'BAR'], 'foo bar', None, 1), + # Explicit case-sensitive + (['foo', 'BAR'], 'foo bar', False, 1), + # Case-insensitive applies to all patterns + (['foo', 'BAR'], 'FOO bar', True, 2), + (['foo', 'BAR'], 'FoO BaR', True, 2), + # Insensitive but still no match + (['quux'], 'FoO BaR', True, 0), + ), +) +def test_can_be_case_insensitive( + execute: ExecuteFunction, + patterns: List[str], + target: str, + case_insensitive: Optional[bool], + expected: int, +) -> None: + extra_args = '' + if case_insensitive is not None: + extra_args = f', case_insensitive={case_insensitive}' + result = execute(f'Count = CountRegexMatches(patterns={patterns!r}, target="{target}"{extra_args})') + assert result == {'Count': expected} + + +def test_rejects_invalid_regex(run_validation: RunValidationFunction, check_failure: CheckFailureFunction) -> None: + with check_failure(): + run_validation('Foo = CountRegexMatches(patterns=["valid", "("], target="")') diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_count_regex_matches/test_rejects_invalid_regex.txt b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_count_regex_matches/test_rejects_invalid_regex.txt new file mode 100644 index 0000000..30bf0c2 --- /dev/null +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_count_regex_matches/test_rejects_invalid_regex.txt @@ -0,0 +1,5 @@ +error: invalid regex pattern '(': missing ), unterminated subpattern at position 0 +--> main.sml:1:33 + | + 1 | Foo = CountRegexMatches(patterns=["valid", "("], target="") + | ^ \ No newline at end of file diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_entity.py b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_entity.py index 3dd9b4c..103a4be 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_entity.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_entity.py @@ -105,8 +105,8 @@ def test_entity_literal_arguments_can_be_names_from_other_source( def test_entity_checks_json_value_type(execute_with_result: ExecuteWithResultFunction) -> None: result = execute_with_result( """ - A: Entity[str] = EntityJson(type='A', path='$.my_int') - B: Entity[int] = EntityJson(type='B', path='$.my_str') + A: Entity[str] = EntityJson(type='A', path='$.my_int', coerce_type=False) + B: Entity[int] = EntityJson(type='B', path='$.my_str', coerce_type=False) """, data={'my_int': 123, 'my_str': 'abc'}, ) diff --git a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_experiments.py b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_experiments.py index 2e2bd27..c1ce689 100644 --- a/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_experiments.py +++ b/osprey_worker/src/osprey/engine/stdlib/udfs/tests/test_experiments.py @@ -7,18 +7,12 @@ from osprey.engine.ast_validator.validators.validate_call_kwargs import Validate from osprey.engine.conftest import CheckFailureFunction, ExecuteFunction, RunValidationFunction from osprey.engine.language_types.experiments import NOT_IN_EXPERIMENT_BUCKET, NOT_IN_EXPERIMENT_BUCKET_INDEX from osprey.engine.stdlib.udfs.entity import Entity -from osprey.engine.stdlib.udfs.experiments import ( - CONTROL_BUCKET, - EXPERIMENT_GRANULARITY, - Experiment, - ExperimentWhen, - InExperiment, -) +from osprey.engine.stdlib.udfs.experiments import CONTROL_BUCKET, EXPERIMENT_GRANULARITY, Experiment, ExperimentWhen from osprey.engine.stdlib.udfs.rules import Rule from osprey.engine.udf.registry import UDFRegistry pytestmark: List[Callable[[Any], Any]] = [ - pytest.mark.use_udf_registry(UDFRegistry.with_udfs(Entity, Rule, Experiment, ExperimentWhen, InExperiment)), + pytest.mark.use_udf_registry(UDFRegistry.with_udfs(Entity, Rule, Experiment, ExperimentWhen)), pytest.mark.use_validators([ValidateCallKwargs, UniqueStoredNames]), ] @@ -375,32 +369,6 @@ def test_experimentwhen_results( assert data['B'] == expected_value -@mock.patch.object(Experiment, 'hash_mod') -def test_inexperiment_returns_true_when_in_experiment(hash_mod_mock: mock.MagicMock, execute: ExecuteFunction) -> None: - hash_mod_mock.return_value = 0 - experiment = f""" - E1 = Entity(type='MyEntity', id='entity 1') - A = Experiment(entity=E1, buckets=['{CONTROL_BUCKET}', 'b'], bucket_sizes=[50.0, 50.0], version=1, revision=1) - B = InExperiment(experiment=A) - """ - data = execute(experiment) - assert data['B'] is True - - -@mock.patch.object(Experiment, 'hash_mod') -def test_inexperiment_returns_false_when_not_in_experiment( - hash_mod_mock: mock.MagicMock, execute: ExecuteFunction -) -> None: - hash_mod_mock.return_value = 9999 - experiment = f""" - E1 = Entity(type='MyEntity', id='entity 1') - A = Experiment(entity=E1, buckets=['{CONTROL_BUCKET}', 'b'], bucket_sizes=[1.0, 1.0], version=1, revision=1) - B = InExperiment(experiment=A) - """ - data = execute(experiment) - assert data['B'] is False - - @mock.patch.object(Experiment, 'hash_mod') def test_inline_experimentwhen( hash_mod_mock: mock.MagicMock, diff --git a/osprey_worker/src/osprey/engine/utils/types.py b/osprey_worker/src/osprey/engine/utils/types.py index 73c2449..9a1cb1e 100644 --- a/osprey_worker/src/osprey/engine/utils/types.py +++ b/osprey_worker/src/osprey/engine/utils/types.py @@ -101,7 +101,7 @@ def cached_property(func: Callable[[SelfT], ReturnT]) -> Property[SelfT, ReturnT self_id = id(self) if self_id not in value_cache: value_cache[self_id] = func(self) - weakref.finalize(self, lambda: value_cache.pop(self_id)) + weakref.finalize(self, lambda: value_cache.pop(self_id, None)) return value_cache[self_id] diff --git a/osprey_worker/src/osprey/worker/_stdlibplugin/udf_register.py b/osprey_worker/src/osprey/worker/_stdlibplugin/udf_register.py index b647bd5..2d85007 100644 --- a/osprey_worker/src/osprey/worker/_stdlibplugin/udf_register.py +++ b/osprey_worker/src/osprey/worker/_stdlibplugin/udf_register.py @@ -1,5 +1,6 @@ from typing import Any, Sequence, Type +from osprey.engine.stdlib.udfs.count_regex_matches import CountRegexMatches from osprey.engine.stdlib.udfs.domain_chopper import DomainChopper from osprey.engine.stdlib.udfs.domain_tld import DomainTld from osprey.engine.stdlib.udfs.email_domain import EmailDomain, EmailSubdomain @@ -9,7 +10,6 @@ from osprey.engine.stdlib.udfs.experiments import ( Experiment, ExperimentsBucketAssignment, ExperimentWhen, - InExperiment, ) from osprey.engine.stdlib.udfs.extract_cookie import ExtractCookie from osprey.engine.stdlib.udfs.get_action_id import GetActionId @@ -22,7 +22,6 @@ 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 from osprey.engine.stdlib.udfs.mx_lookup import MXLookup -from osprey.engine.stdlib.udfs.parse_int import ParseInt from osprey.engine.stdlib.udfs.phone_country import PhoneCountry from osprey.engine.stdlib.udfs.phone_prefix import PhonePrefix from osprey.engine.stdlib.udfs.random_bool import RandomBool @@ -85,7 +84,6 @@ def register_udfs() -> Sequence[Type[UDFBase[Any, Any]]]: Experiment, ExperimentWhen, ExperimentsBucketAssignment, - InExperiment, ExtractCookie, GetActionId, GetActionName, @@ -99,7 +97,6 @@ def register_udfs() -> Sequence[Type[UDFBase[Any, Any]]]: ListRead, ListSort, MXLookup, - ParseInt, PhoneCountry, PhonePrefix, RandomBool, @@ -136,4 +133,5 @@ def register_udfs() -> Sequence[Type[UDFBase[Any, Any]]]: TimeDelta, TimeSince, WhenRules, + CountRegexMatches, ] diff --git a/osprey_worker/src/osprey/worker/adaptor/hookspecs/osprey_hooks.py b/osprey_worker/src/osprey/worker/adaptor/hookspecs/osprey_hooks.py index 4f83966..27c58a5 100644 --- a/osprey_worker/src/osprey/worker/adaptor/hookspecs/osprey_hooks.py +++ b/osprey_worker/src/osprey/worker/adaptor/hookspecs/osprey_hooks.py @@ -13,6 +13,7 @@ from osprey.worker.sinks.utils.acking_contexts import BaseAckingContext if TYPE_CHECKING: from osprey.worker.lib.config import Config + from osprey.worker.lib.data_exporters.validation_result_exporter import BaseValidationResultExporter from osprey.worker.lib.storage.stored_execution_result import ExecutionResultStore from osprey.worker.sinks.sink.input_stream import BaseInputStream from osprey.worker.sinks.sink.output_sink import BaseOutputSink @@ -63,6 +64,17 @@ def register_labels_service_or_provider(config: Config) -> LabelsServiceBase | L raise NotImplementedError('register_labels_service_or_provider must be implemented by the plugin') +@hookspec(firstresult=True) +def register_validation_exporter(config: Config) -> 'BaseValidationResultExporter | None': + """ + Optional: Register a custom validation result exporter. + + Called after sources are validated to publish experiment metadata (e.g., bucket definitions). + If None is returned or this hook is not implemented, NullValidationResultExporter is used. + """ + pass + + @hookspec(firstresult=True) def register_label_output_sink(config: Config, labels_provider: LabelsProvider) -> BaseOutputSink | None: """ diff --git a/osprey_worker/src/osprey/worker/adaptor/plugin_manager.py b/osprey_worker/src/osprey/worker/adaptor/plugin_manager.py index ea1be29..243fa9d 100644 --- a/osprey_worker/src/osprey/worker/adaptor/plugin_manager.py +++ b/osprey_worker/src/osprey/worker/adaptor/plugin_manager.py @@ -20,6 +20,7 @@ from osprey.worker.sinks.utils.acking_contexts import BaseAckingContext if TYPE_CHECKING: from osprey.worker.lib.config import Config + from osprey.worker.lib.data_exporters.validation_result_exporter import BaseValidationResultExporter hookimpl_osprey: pluggy.HookimplMarker = pluggy.HookimplMarker(OSPREY_ADAPTOR) @@ -88,6 +89,17 @@ def bootstrap_output_sinks(config: Config) -> BaseOutputSink: return MultiOutputSink(sinks) +def bootstrap_validation_exporter(config: Config) -> 'BaseValidationResultExporter': + from osprey.worker.lib.data_exporters.validation_result_exporter import ( + BaseValidationResultExporter, + NullValidationResultExporter, + ) + + load_all_osprey_plugins() + exporter = plugin_manager.hook.register_validation_exporter(config=config) + return exporter if isinstance(exporter, BaseValidationResultExporter) else NullValidationResultExporter() + + def bootstrap_labels_provider(config: Config) -> LabelsProvider: """ NOTE: If you are looking to get a labels provider to use within Osprey, diff --git a/osprey_worker/src/osprey/worker/cli/sinks.py b/osprey_worker/src/osprey/worker/cli/sinks.py index 13f1af9..ee20c90 100644 --- a/osprey_worker/src/osprey/worker/cli/sinks.py +++ b/osprey_worker/src/osprey/worker/cli/sinks.py @@ -61,6 +61,7 @@ CONFIG_SENTRY_OTHER_SINKS_DSN = 'SENTRY_OTHER_SINKS_DSN' def init_config() -> Config: config = CONFIG.instance() config.configure_from_env() + instruments.set_worker_type_tag('gevent') return config diff --git a/osprey_worker/src/osprey/worker/lib/etcd/tests/test_watcher_mux.py b/osprey_worker/src/osprey/worker/lib/etcd/tests/test_watcher_mux.py index d62cf83..226e108 100644 --- a/osprey_worker/src/osprey/worker/lib/etcd/tests/test_watcher_mux.py +++ b/osprey_worker/src/osprey/worker/lib/etcd/tests/test_watcher_mux.py @@ -70,6 +70,31 @@ def test_recursive_mux_full_syncs_are_muxed(): run_events_thru_mux(events, RecursiveWatchMux()) +def test_recursive_mux_full_sync_does_not_re_emit_when_key_differs_from_value(): + # Regression: the dedup directory must store value (not key). In production keys are etcd + # paths and values are payloads — when they differ, a bug storing the key would cause every + # subsequent full-sync to spuriously re-emit every previously-changed entry as an upsert. + events = [ + # Initial state. + ( + FullSyncRecursive(items=[FullSyncOne(key='/k1', value='v1'), FullSyncOne(key='/k2', value='v2')]), + [FullSyncRecursive(items=[FullSyncOne(key='/k1', value='v1'), FullSyncOne(key='/k2', value='v2')])], + ), + # Full-sync with /k1 actually changing: exercises _synthesize_incremental_events_from_full_sync. + ( + FullSyncRecursive(items=[FullSyncOne(key='/k1', value='v1-new'), FullSyncOne(key='/k2', value='v2')]), + [IncrementalSyncUpsert(key='/k1', value='v1-new')], + ), + # Full-sync again with no changes: must emit nothing. Previously this spuriously re-emitted /k1. + ( + FullSyncRecursive(items=[FullSyncOne(key='/k1', value='v1-new'), FullSyncOne(key='/k2', value='v2')]), + [], + ), + ] + + run_events_thru_mux(events, RecursiveWatchMux()) + + def test_recursive_mux_incrementals(): events = [ ( diff --git a/osprey_worker/src/osprey/worker/lib/etcd/watcher/_mux.py b/osprey_worker/src/osprey/worker/lib/etcd/watcher/_mux.py index 97dde2b..7256bca 100644 --- a/osprey_worker/src/osprey/worker/lib/etcd/watcher/_mux.py +++ b/osprey_worker/src/osprey/worker/lib/etcd/watcher/_mux.py @@ -82,7 +82,7 @@ class RecursiveWatchMux(object): unvisited_keys.discard(sync_one.key) existing = self._directory.get(sync_one.key, self._DOES_NOT_EXIST) if sync_one.value != existing: - self._directory[sync_one.key] = sync_one.key + self._directory[sync_one.key] = sync_one.value yield IncrementalSyncUpsert(key=sync_one.key, value=sync_one.value) for key in sorted(unvisited_keys): diff --git a/osprey_worker/src/osprey/worker/lib/etcd/watcher/watcherd_impl.py b/osprey_worker/src/osprey/worker/lib/etcd/watcher/watcherd_impl.py index f168f47..6dbfff8 100644 --- a/osprey_worker/src/osprey/worker/lib/etcd/watcher/watcherd_impl.py +++ b/osprey_worker/src/osprey/worker/lib/etcd/watcher/watcherd_impl.py @@ -138,14 +138,17 @@ class EtcdWatcherdWatcher(BaseWatcher): except MemoryError as e: raise e - except Exception: + except Exception as e: if self._stream: self._stream.cancel() self._stream = None self._set_state(EtcdWatcherState.RESET) delay = self._backoff.fail() - log.exception( - '%r: etcd-watcherd stream raised an error. sleeping for %.2f sec before retrying', self, delay + log.warning( + '%r: etcd-watcherd stream raised an error (%r). sleeping for %.2f sec before retrying', + self, + e, + delay, ) time.sleep(delay) @@ -161,9 +164,9 @@ class EtcdWatcherdWatcher(BaseWatcher): except MemoryError as e: raise e - except Exception: + except Exception as e: delay = self._backoff.fail() - log.exception('%r: etcd raised an error. sleeping for %.2f sec before retrying', self, delay) + log.warning('%r: etcd raised an error (%r). sleeping for %.2f sec before retrying', self, e, delay) time.sleep(delay) def _translate_event(self, event): diff --git a/osprey_worker/src/osprey/worker/lib/instruments/__init__.py b/osprey_worker/src/osprey/worker/lib/instruments/__init__.py index 0529485..85f09ef 100644 --- a/osprey_worker/src/osprey/worker/lib/instruments/__init__.py +++ b/osprey_worker/src/osprey/worker/lib/instruments/__init__.py @@ -80,6 +80,18 @@ class _DogStatsd(DogStatsd): max_buffer_size=max_buffer_size, # type:ignore constant_tags=constant_tags, use_ms=use_ms, + # Pack multiple metrics per datagram. Default (disable_buffering=True) sends one + # UDP packet per metric, which overruns the DD agent's socket recv buffer at + # high emission rates and causes silent kernel drops. max_buffer_len auto-selects + # 1432 bytes for UDP / 8192 for UDS. + disable_buffering=False, + # Client-side aggregation of counters/gauges/sets: collapses repeated + # increment('foo', tags=T) calls on the same (metric, tags) within a flush + # window into a single send. At ~1.5k action/s/process emitting ~100 distinct + # (metric, action-tag) combos, this cuts counter-packet rate by ~30x without + # changing the aggregated total the agent sees. Histograms and timings are + # untouched — they still emit every sample so percentiles stay accurate. + disable_aggregation=False, ) self.prefix = None self.debug = False @@ -206,6 +218,19 @@ class _DogStatsd(DogStatsd): metrics = _DogStatsd() +def set_worker_type_tag(worker_type: str) -> None: + """Add worker_type as a constant tag on the metrics singleton. + + Must be called once at startup. All subsequent metrics will include + the worker_type tag automatically — no per-call-site changes needed. + """ + tag = f'worker_type:{worker_type}' + if metrics.constant_tags is None: + metrics.constant_tags = [tag] + elif tag not in metrics.constant_tags: + metrics.constant_tags.append(tag) + + class concurrency(contextlib.ContextDecorator): """A decorator for tracking concurrent calls of a function as a Gauge diff --git a/osprey_worker/src/osprey/worker/lib/osprey_engine.py b/osprey_worker/src/osprey/worker/lib/osprey_engine.py index 4bb4501..624a059 100644 --- a/osprey_worker/src/osprey/worker/lib/osprey_engine.py +++ b/osprey_worker/src/osprey/worker/lib/osprey_engine.py @@ -136,6 +136,10 @@ class OspreyEngine: else: # Only do this if no exception occurred above self._config_subkey_handler.dispatch_config(self._execution_graph.validated_sources) + # Confirm to the provider which sources are now live so it dedups no-op + # re-deliveries against what we applied; the except branch above leaves + # it unmarked, so a failed compile retries on the next re-delivery. + self._sources_provider.mark_sources_applied(self._execution_graph.validated_sources.sources.hash()) # noinspection PyBroadException # try to send validation results, should not block osprey_engine if this fails @@ -305,23 +309,31 @@ def bootstrap_engine_with_helpers( sources_provider: Optional[BaseSourcesProvider] = None, ) -> Tuple[OspreyEngine, UDFHelpers]: # Avoid circular imports - from osprey.worker.adaptor.plugin_manager import bootstrap_ast_validators, bootstrap_udfs + from osprey.worker.adaptor.plugin_manager import ( + bootstrap_ast_validators, + bootstrap_udfs, + bootstrap_validation_exporter, + ) udf_registry, udf_helpers = bootstrap_udfs() bootstrap_ast_validators() + config = CONFIG.instance() + if not sources_provider: # Use static rules path if configured, otherwise use etcd - config = CONFIG.instance() rules_path_str = config.get_optional_str('OSPREY_RULES_PATH') rules_path = Path(rules_path_str) if rules_path_str else None sources_provider = get_sources_provider(rules_path=rules_path) + validation_exporter = bootstrap_validation_exporter(config) + return ( OspreyEngine( sources_provider=sources_provider, udf_registry=udf_registry, should_yield_during_compilation=should_yield_during_compilation(), + validation_exporter=validation_exporter, ), udf_helpers, ) diff --git a/osprey_worker/src/osprey/worker/lib/sources_provider.py b/osprey_worker/src/osprey/worker/lib/sources_provider.py index 65f78a8..4d13c34 100644 --- a/osprey_worker/src/osprey/worker/lib/sources_provider.py +++ b/osprey_worker/src/osprey/worker/lib/sources_provider.py @@ -1,39 +1,23 @@ -import abc import logging -from typing import Callable, Dict, Optional +from typing import Dict, Optional from osprey.engine.ast.sources import Sources from osprey.worker.lib.etcd import EtcdClient from osprey.worker.lib.etcd.dict import ReadOnlyEtcdDict +from osprey.worker.lib.sources_provider_base import ( + BaseSourcesProvider, + SourcesWatcherCallback, + StaticSourcesProvider, +) from osprey.worker.lib.utils.input_stream_ready_signaler import InputStreamReadySignaler -SourcesWatcherCallback = Callable[[], None] - - -class BaseSourcesProvider(abc.ABC): - """Provides an interface to get the and be informed of current sources of rules which the rules engine should - evaluate""" - - @abc.abstractmethod - def get_current_sources(self) -> Sources: - raise NotImplementedError - - @abc.abstractmethod - def set_sources_watcher(self, callback: SourcesWatcherCallback) -> None: - raise NotImplementedError - - -class StaticSourcesProvider(BaseSourcesProvider): - """Provides a static sources that won't change for the lifetime of the provider.""" - - def __init__(self, sources: Sources): - self._sources = sources - - def get_current_sources(self) -> Sources: - return self._sources - - def set_sources_watcher(self, callback: SourcesWatcherCallback) -> None: - return None +# Re-export base classes for backward compatibility +__all__ = [ + 'BaseSourcesProvider', + 'StaticSourcesProvider', + 'SourcesWatcherCallback', + 'EtcdSourcesProvider', +] class EtcdSourcesProvider(BaseSourcesProvider): @@ -49,17 +33,36 @@ class EtcdSourcesProvider(BaseSourcesProvider): self._current_sources = Sources.from_dict(self._sources_dict.copy()) self._sources_watcher_callback: Optional[SourcesWatcherCallback] = None self._input_stream_ready_signaler = input_stream_ready_signaler + # The consumer (engine) compiles this initial snapshot at construction, so + # seed the applied-hash to it — a later re-delivery of the same value is a + # genuine no-op and should be skipped. See _notify_watcher for why dedup + # keys off the *applied* hash rather than the last *received* one. + self._applied_sources_hash: Optional[str] = self._current_sources.hash() self._sources_dict.add_watcher(self._notify_watcher) self._sources_dict.watch() def _notify_watcher(self, sources_dict: Dict[str, str]) -> None: + new_sources = Sources.from_dict(sources_dict) + + # Etcd watcher reconnects and session refreshes re-deliver the current + # value as a full snapshot, so we see many notifications where the content + # is unchanged. Skip the (peak-memory-doubling) recompile only when the + # consumer has ALREADY applied this exact hash. Comparing against the last + # *applied* hash — not merely the last *received* one — keeps this + # self-healing: if a recompile fails (engine keeps its old graph) or is + # dropped (the async fire-and-forget bridge), _applied_sources_hash stays + # behind, so the next re-delivery re-fires the recompile instead of being + # suppressed and leaving the engine wedged on stale rules until restart. + if self._applied_sources_hash is not None and new_sources.hash() == self._applied_sources_hash: + return + if self._input_stream_ready_signaler is not None: logging.info('Pausing input streams') self._input_stream_ready_signaler.pause_input_stream() self._input_stream_ready_signaler.wait_for_input_stream_to_pause() - self._current_sources = Sources.from_dict(sources_dict) + self._current_sources = new_sources if self._sources_watcher_callback: self._sources_watcher_callback() @@ -72,3 +75,8 @@ class EtcdSourcesProvider(BaseSourcesProvider): def set_sources_watcher(self, callback: SourcesWatcherCallback) -> None: self._sources_watcher_callback = callback + + def mark_sources_applied(self, sources_hash: str) -> None: + # Advances the dedup baseline only once the consumer confirms it applied + # these sources, so failed/dropped applies retry on the next re-delivery. + self._applied_sources_hash = sources_hash diff --git a/osprey_worker/src/osprey/worker/lib/sources_provider_base.py b/osprey_worker/src/osprey/worker/lib/sources_provider_base.py new file mode 100644 index 0000000..51d99a6 --- /dev/null +++ b/osprey_worker/src/osprey/worker/lib/sources_provider_base.py @@ -0,0 +1,48 @@ +"""Base sources provider classes — no gevent/etcd dependency. + +These classes are used by both the sync (gevent) and async (asyncio) workers. +The gevent-dependent EtcdSourcesProvider remains in sources_provider.py. +""" + +import abc +from typing import Callable + +from osprey.engine.ast.sources import Sources + +SourcesWatcherCallback = Callable[[], None] + + +class BaseSourcesProvider(abc.ABC): + """Provides an interface to get the and be informed of current sources of rules which the rules engine should + evaluate""" + + @abc.abstractmethod + def get_current_sources(self) -> Sources: + raise NotImplementedError + + @abc.abstractmethod + def set_sources_watcher(self, callback: SourcesWatcherCallback) -> None: + raise NotImplementedError + + def mark_sources_applied(self, sources_hash: str) -> None: + """Notify the provider that the consumer (engine) has successfully applied + the sources identified by ``sources_hash``. + + Providers that dedup repeated etcd deliveries use this to track what was + actually *applied* rather than merely *received*, so a recompile that + fails or is dropped self-heals on the next re-delivery instead of being + suppressed. Default: no-op (providers without dedup need not track it).""" + return None + + +class StaticSourcesProvider(BaseSourcesProvider): + """Provides a static sources that won't change for the lifetime of the provider.""" + + def __init__(self, sources: Sources): + self._sources = sources + + def get_current_sources(self) -> Sources: + return self._sources + + def set_sources_watcher(self, callback: SourcesWatcherCallback) -> None: + return None diff --git a/osprey_worker/src/osprey/worker/lib/tests/test_sources_provider.py b/osprey_worker/src/osprey/worker/lib/tests/test_sources_provider.py new file mode 100644 index 0000000..763303a --- /dev/null +++ b/osprey_worker/src/osprey/worker/lib/tests/test_sources_provider.py @@ -0,0 +1,162 @@ +"""Tests for ``EtcdSourcesProvider``'s self-healing dedup of etcd source updates. + +Context: etcd re-delivers the current value as a full snapshot on every watcher +reconnect / session refresh, so the provider dedups to avoid the memory-doubling +recompile on no-op events. The dedup must key off what the engine has *applied*, +not merely what the provider last *received* — otherwise a recompile that fails +(the engine swallows the error and keeps the old graph) or is dropped (the async +fire-and-forget bridge in the gevent config consumers) advances the dedup +baseline while the engine stays on the old graph. Every subsequent identical +re-delivery is then suppressed and the engine is wedged on stale rules until the +pod restarts. +""" + +from typing import Dict +from unittest.mock import patch + +from osprey.engine.ast.sources import Sources +from osprey.engine.udf.registry import UDFRegistry +from osprey.worker.lib import sources_provider as sp_module +from osprey.worker.lib.osprey_engine import OspreyEngine +from osprey.worker.lib.sources_provider import EtcdSourcesProvider + +# Trivial but distinct, validly-compiling rule sources (no UDFs needed). +S0: Dict[str, str] = {'main.sml': ''} +S1: Dict[str, str] = {'main.sml': '# updated\n'} + + +def _make_fake_etcd_dict(initial: Dict[str, str]) -> type: + """Build a ``ReadOnlyEtcdDict`` stand-in seeded with ``initial`` so the + provider can be constructed without a real etcd.""" + + class _FakeEtcdDict: + def __init__(self, etcd_key: str, etcd_client: object = None, **kwargs: object) -> None: + self._initial = dict(initial) + + def copy(self) -> Dict[str, str]: + return dict(self._initial) + + def add_watcher(self, callback: object) -> None: + pass + + def watch(self) -> None: + pass + + return _FakeEtcdDict + + +def _make_provider(initial: Dict[str, str] = S0) -> EtcdSourcesProvider: + with patch.object(sp_module, 'ReadOnlyEtcdDict', _make_fake_etcd_dict(initial)): + # The engine compiles this initial snapshot at construction. + return EtcdSourcesProvider(etcd_key='/test/key') + + +def test_redelivery_after_failed_apply_self_heals(): + """A recompile that fails once must be retried on the next identical + re-delivery — the dedup baseline must not advance past what was applied.""" + provider = _make_provider() + + fired = {'n': 0} + compile_succeeds = {'v': False} + + def fake_engine_callback() -> None: + # Mirrors the engine contract: "compile" the current sources and only + # confirm-applied on success. On failure (swallowed) it does not mark, + # exactly like the engine keeping its old graph after a compile error. + fired['n'] += 1 + if compile_succeeds['v']: + provider.mark_sources_applied(provider.get_current_sources().hash()) + + provider.set_sources_watcher(fake_engine_callback) + + # 1) S1 arrives; the apply FAILS (engine keeps old graph, does not mark). + provider._notify_watcher(dict(S1)) + assert fired['n'] == 1 + + # 2) etcd re-delivers S1 (FullSyncOne on reconnect). It MUST re-fire so the + # transient failure self-heals — with the bug this is suppressed and the + # engine stays wedged on S0. + compile_succeeds['v'] = True + provider._notify_watcher(dict(S1)) + assert fired['n'] == 2, 're-delivery after a failed apply was suppressed — engine wedged on stale rules' + + # 3) Once the engine has applied S1, further S1 re-deliveries are deduped. + provider._notify_watcher(dict(S1)) + assert fired['n'] == 2, 'no-op re-delivery should be skipped once applied' + + +def test_noop_redelivery_of_applied_sources_is_skipped(): + """OOM mitigation preserved: re-delivery of already-applied sources (incl. the + initial snapshot the engine compiled at construction) does not recompile.""" + provider = _make_provider() + + fired = {'n': 0} + + def fake_engine_callback() -> None: + fired['n'] += 1 + provider.mark_sources_applied(provider.get_current_sources().hash()) + + provider.set_sources_watcher(fake_engine_callback) + + # Re-delivery of the initial applied snapshot (S0) must be skipped. + provider._notify_watcher(dict(S0)) + assert fired['n'] == 0, 're-delivery of the initial applied snapshot should be skipped' + + # A genuine change fires once; its subsequent re-deliveries are skipped. + provider._notify_watcher(dict(S1)) + assert fired['n'] == 1 + provider._notify_watcher(dict(S1)) + assert fired['n'] == 1 + + +def test_real_engine_recovers_from_transient_apply_failure_on_redelivery(): + """End-to-end with a REAL OspreyEngine wired to EtcdSourcesProvider. + + A valid rules update whose first apply fails transiently (engine swallows the + error and keeps the old graph — the same shape as a dropped async-bridge + dispatch) must be picked up when etcd re-delivers it. Proves the engine's + live execution graph actually advances, not just the provider's bookkeeping. + """ + s0_hash = Sources.from_dict(dict(S0)).hash() + s1_hash = Sources.from_dict(dict(S1)).hash() + assert s0_hash != s1_hash + + provider = _make_provider(S0) + # Construct the engine directly with an empty UDF registry (trivial rules use + # no UDFs) to avoid loading external plugin entrypoints. + engine = OspreyEngine(sources_provider=provider, udf_registry=UDFRegistry()) + + # The engine compiled the initial snapshot at construction. + assert engine.execution_graph.validated_sources.sources.hash() == s0_hash + + # Make the next recompile fail exactly once — a transient apply failure on a + # valid update (resource blip / dropped dispatch). + real_compile = engine._compile_execution_graph + calls = {'n': 0} + + def flaky_compile(*args: object, **kwargs: object) -> object: + calls['n'] += 1 + if calls['n'] == 1: + raise RuntimeError('transient compile failure') + return real_compile(*args, **kwargs) + + engine._compile_execution_graph = flaky_compile # type: ignore[method-assign] + + # 1) Valid S1 arrives; the apply fails and the engine keeps S0. + provider._notify_watcher(dict(S1)) + assert engine.execution_graph.validated_sources.sources.hash() == s0_hash, ( + 'engine should still serve S0 after a failed apply' + ) + # The OLD dedup (against _current_sources) WOULD now wedge — it advanced to S1 + # so every re-delivery is suppressed... + assert provider._current_sources.hash() == s1_hash + # ...but the fix dedups against the APPLIED hash, still S0, so re-delivery retries. + assert provider._applied_sources_hash == s0_hash + + # 2) etcd re-delivers S1 (FullSyncOne). The engine recompiles and recovers — + # its live execution graph is now S1, with no restart. + provider._notify_watcher(dict(S1)) + assert engine.execution_graph.validated_sources.sources.hash() == s1_hash, ( + 'engine must serve S1 after re-delivery (self-healed)' + ) + assert provider._applied_sources_hash == s1_hash diff --git a/osprey_worker/src/osprey/worker/sinks/sink/monitored_rules_metrics.py b/osprey_worker/src/osprey/worker/sinks/sink/monitored_rules_metrics.py new file mode 100644 index 0000000..3269bb9 --- /dev/null +++ b/osprey_worker/src/osprey/worker/sinks/sink/monitored_rules_metrics.py @@ -0,0 +1,33 @@ +"""Utility for emitting monitored_rules metrics after label mutations.""" + +from typing import Sequence, Set, Tuple + +from osprey.worker.lib.instruments import metrics + + +def emit_monitored_rules_metrics( + mutations: Sequence[Tuple[str, str, str]], + monitored_labels: Set[str], + action_name: str, +) -> None: + """Emit 'monitored_rules' metric for mutations on monitored labels. + + Should be called after label mutations are applied and at least one + label changed (added or removed). Each tuple is (label_name, reason_name, status_tag). + + Args: + mutations: Sequence of (label_name, reason_name, status_tag) tuples. + monitored_labels: Label names to emit metrics for (from AnalyticsConfig). + action_name: The action name for the metric tag. + """ + for label_name, reason_name, status_tag in mutations: + if label_name in monitored_labels: + metrics.increment( + 'monitored_rules', + tags=[ + f'rule:{reason_name}', + f'label:{label_name}', + f'status:{status_tag}', + f'action:{action_name}', + ], + ) diff --git a/osprey_worker/src/osprey/worker/sinks/sink/output_sink.py b/osprey_worker/src/osprey/worker/sinks/sink/output_sink.py index a9cb4e8..bddc0ed 100644 --- a/osprey_worker/src/osprey/worker/sinks/sink/output_sink.py +++ b/osprey_worker/src/osprey/worker/sinks/sink/output_sink.py @@ -16,6 +16,7 @@ from osprey.worker.lib.instruments import metrics from osprey.worker.lib.osprey_shared.labels import EntityLabelMutation from osprey.worker.lib.osprey_shared.logging import DynamicLogSampler, get_logger from osprey.worker.lib.storage.labels import LabelsProvider +from osprey.worker.sinks.sink.monitored_rules_metrics import emit_monitored_rules_metrics from tenacity import RetryCallState, retry, stop_after_attempt, wait_exponential logger = get_logger() @@ -186,18 +187,29 @@ def _get_label_effects_from_result(result: ExecutionResult) -> Mapping[EntityT[A class LabelOutputSink(BaseOutputSink): """An output sink that will send event effects to the label service.""" - def __init__(self, labels_provider: LabelsProvider) -> None: + def __init__(self, labels_provider: LabelsProvider, monitored_labels: set[str] | None = None) -> None: self._labels_provider = labels_provider + self._monitored_labels: set[str] = monitored_labels or set() + + def set_monitored_labels(self, monitored_labels: set[str]) -> None: + """Set monitored labels for metrics emission. Can be called post-engine compilation.""" + self._monitored_labels = monitored_labels def will_do_work(self, result: ExecutionResult) -> bool: return len(_get_label_effects_from_result(result)) > 0 def push(self, result: ExecutionResult) -> None: for entity, mutations in _get_label_effects_from_result(result).items(): - _ = self._labels_provider.apply_entity_label_mutations( + mutation_result = self._labels_provider.apply_entity_label_mutations( entity, mutations, ) + if self._monitored_labels and (mutation_result.labels_added or mutation_result.labels_removed): + emit_monitored_rules_metrics( + mutations=[(m.label_name, m.reason_name, str(m.status)) for m in mutations], + monitored_labels=self._monitored_labels, + action_name=result.action.action_name, + ) def stop(self) -> None: self._labels_provider.stop() diff --git a/osprey_worker/src/osprey/worker/sinks/sink/rules_sink.py b/osprey_worker/src/osprey/worker/sinks/sink/rules_sink.py index 353268b..bf5ca89 100644 --- a/osprey_worker/src/osprey/worker/sinks/sink/rules_sink.py +++ b/osprey_worker/src/osprey/worker/sinks/sink/rules_sink.py @@ -140,9 +140,11 @@ class RulesSink(BaseSink): for message_context in self._input_stream: try: with message_context as action: - if action.data.get('osprey_v2_skip_async_classification', False) or action.data.get( - 'osprey_skip_async', False - ): + action_tags = [f'action:{action.action_name}'] + metrics.increment('rules_sink.input_action_received', tags=action_tags) + + if action.data.get('osprey_skip_async', False): + metrics.increment('rules_sink.skipped', tags=action_tags) continue with tracer.start_span('osprey.classify_one', child_of=None) as span: diff --git a/osprey_worker/src/osprey/worker/sinks/utils/acking_contexts.py b/osprey_worker/src/osprey/worker/sinks/utils/acking_contexts.py index a0248af..80d68df 100644 --- a/osprey_worker/src/osprey/worker/sinks/utils/acking_contexts.py +++ b/osprey_worker/src/osprey/worker/sinks/utils/acking_contexts.py @@ -1,98 +1,33 @@ -import abc from datetime import datetime from types import TracebackType -from typing import Dict, Generic, List, Optional, Type, TypeVar, Union +from typing import Dict, List, Optional, Type, TypeVar, Union import gevent from google.api_core.exceptions import DeadlineExceeded from google.cloud.pubsub_v1 import SubscriberClient from google.cloud.pubsub_v1.subscriber.message import Message -from osprey.rpc.common.v1.verdicts_pb2 import Verdicts from osprey.worker.lib.instruments import metrics from osprey.worker.lib.osprey_shared.logging import get_logger +from osprey.worker.sinks.utils.acking_contexts_base import ( + BaseAckingContext, + NoopAckingContext, + VerdictsAckingContext, +) + +# Re-export base classes for backward compatibility +__all__ = [ + 'BaseAckingContext', + 'NoopAckingContext', + 'VerdictsAckingContext', + 'PubSubMessageAckingContext', + 'PullPubSubMessageContext', +] logger = get_logger() _T = TypeVar('_T') -# TODO: support NACK -class BaseAckingContext(abc.ABC, Generic[_T]): - """An acking context for handling single actions from input streams.""" - - def __init__(self, item: _T) -> None: - super().__init__() - self._item: _T = item - self._should_nack = False - self._publish_time = datetime.now() - self._attributes: Optional[Dict[str, str]] = None - - @abc.abstractmethod - def _ack(self) -> None: - """Acknowledges the message or item that this Acking Context holds.""" - - raise NotImplementedError - - @abc.abstractmethod - def _nack(self) -> None: - """NACKs the message or item that this Acking Context holds.""" - - raise NotImplementedError - - @property - def attributes(self) -> Optional[Dict[str, str]]: - return self._attributes - - def mark_as_nack(self) -> None: - self._should_nack = True - - def __enter__(self) -> _T: - return self._item - - def __exit__( - self, - exc_type: Union[Type[BaseException], None], - exc_value: Union[BaseException, None], - exc_traceback: Union[TracebackType, None], - ) -> None: - if self._should_nack: - self._nack() - else: - self._ack() - - @property - def publish_time(self) -> datetime: - return self._publish_time - - -class NoopAckingContext(BaseAckingContext[_T]): - """A context manager for handling single actions require no acking operations from input streams.""" - - def _ack(self) -> None: - return - - def _nack(self) -> None: - return - - -class VerdictsAckingContext(NoopAckingContext[_T]): - """ - A context manager for storing verdicts from the rules sink inside of a NoopAckingContext :3 - - This is used to send verdicts back to the Osprey Coordinator, if any were captured~ - """ - - def __init__(self, item: _T) -> None: - super().__init__(item) - self._verdicts: Optional[Verdicts] = None - - def set_verdicts(self, verdicts: Verdicts) -> None: - self._verdicts = verdicts - - def get_verdicts(self) -> Optional[Verdicts]: - return self._verdicts - - class PubSubMessageAckingContext(BaseAckingContext[_T]): """A context manager for handling single pubsub messages using the push method. Ennsures that the handling and acking of a specific message will be handled by the same thread.""" diff --git a/osprey_worker/src/osprey/worker/sinks/utils/acking_contexts_base.py b/osprey_worker/src/osprey/worker/sinks/utils/acking_contexts_base.py new file mode 100644 index 0000000..6103dfd --- /dev/null +++ b/osprey_worker/src/osprey/worker/sinks/utils/acking_contexts_base.py @@ -0,0 +1,98 @@ +"""Base acking context classes — no gevent dependency. + +These classes are used by both the sync (gevent) and async (asyncio) workers. +The gevent-dependent PubSub acking contexts remain in acking_contexts.py. +""" + +import abc +from datetime import datetime +from types import TracebackType +from typing import Dict, Generic, Optional, Type, TypeVar, Union + +from osprey.rpc.common.v1.verdicts_pb2 import Verdicts + +_T = TypeVar('_T') + + +# TODO: support NACK +class BaseAckingContext(abc.ABC, Generic[_T]): + """An acking context for handling single actions from input streams.""" + + def __init__(self, item: _T) -> None: + super().__init__() + self._item: _T = item + self._should_nack = False + self._publish_time = datetime.now() + self._attributes: Optional[Dict[str, str]] = None + + @abc.abstractmethod + def _ack(self) -> None: + """Acknowledges the message or item that this Acking Context holds.""" + + raise NotImplementedError + + @abc.abstractmethod + def _nack(self) -> None: + """NACKs the message or item that this Acking Context holds.""" + + raise NotImplementedError + + @property + def attributes(self) -> Optional[Dict[str, str]]: + return self._attributes + + def mark_as_nack(self) -> None: + self._should_nack = True + + @property + def should_nack(self) -> bool: + """Whether this item was marked for nack. Consumers that send ack/nack + out-of-band (e.g. the async coordinator stream) must consult this rather + than always acking.""" + return self._should_nack + + def __enter__(self) -> _T: + return self._item + + def __exit__( + self, + exc_type: Union[Type[BaseException], None], + exc_value: Union[BaseException, None], + exc_traceback: Union[TracebackType, None], + ) -> None: + if self._should_nack: + self._nack() + else: + self._ack() + + @property + def publish_time(self) -> datetime: + return self._publish_time + + +class NoopAckingContext(BaseAckingContext[_T]): + """A context manager for handling single actions require no acking operations from input streams.""" + + def _ack(self) -> None: + return + + def _nack(self) -> None: + return + + +class VerdictsAckingContext(NoopAckingContext[_T]): + """ + A context manager for storing verdicts from the rules sink inside of a NoopAckingContext :3 + + This is used to send verdicts back to the Osprey Coordinator, if any were captured~ + """ + + def __init__(self, item: _T) -> None: + super().__init__(item) + self._verdicts: Optional[Verdicts] = None + + def set_verdicts(self, verdicts: Verdicts) -> None: + self._verdicts = verdicts + + def get_verdicts(self) -> Optional[Verdicts]: + return self._verdicts diff --git a/pyproject.toml b/pyproject.toml index e73c0f7..0ef3c55 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,7 +43,6 @@ common = [ "grpcio-tools==1.53.*; platform_machine != 'x86_64'", "gunicorn", "intervals==0.9.2", - "jslog4kube==1.0.6", "jsonpath-rw", "kafka-python==1.4.7", # Flask 1.1.4 / Jinja2 2.11 import `markupsafe.soft_unicode`, which was removed in @@ -79,7 +78,7 @@ common = [ "tink==1.9.0", "tld==0.12.7", "traitlets==5.14.3", - "typing-extensions==4.6.3", + "typing-extensions==4.12.2", "typing-inspect==0.9.0", "unidecode==1.3.8", "werkzeug==1.0.1", @@ -104,6 +103,7 @@ common = [ ] dev = [ "ipdb==0.13.13", + "pytest-asyncio", "ipython==9.4.0", "memray==1.19.2", "mypy>=1.13.0", @@ -133,12 +133,14 @@ jsonpath-rw = { url = "https://github.com/kennknowles/python-jsonpath-rw/archive nostril = { url = "https://github.com/casics/nostril/archive/v1.2.0.tar.gz" } osprey_rpc = { workspace = true } osprey_worker = { workspace = true } +osprey_async_worker = { workspace = true } example_plugins = { workspace = true } [tool.uv.workspace] members = [ "osprey_rpc", "osprey_worker", + "osprey_async_worker", "example_plugins", ] @@ -172,10 +174,10 @@ ignore = [ ] [tool.ruff.lint.isort] -known-first-party = ["osprey_worker", "osprey_rpc", "example_plugins"] +known-first-party = ["osprey_worker", "osprey_async_worker", "osprey_rpc", "example_plugins"] [tool.fawltydeps] -code = ["osprey_worker/src", "osprey_rpc/src", "example_plugins/src"] +code = ["osprey_worker/src", "osprey_async_worker/src", "osprey_rpc/src", "example_plugins/src"] deps = ["pyproject.toml"] ignore_unused = [ # Type stubs: used by mypy, never imported directly @@ -207,6 +209,7 @@ ignore_unused = [ "mypy-extensions", "mypy-protobuf", "pre-commit", + "pytest-asyncio", "pytest-flask", "pytest-order", "requests-mock", @@ -214,8 +217,6 @@ ignore_unused = [ "setuptools", # Runtime CLI: started via command line, not imported in source "gunicorn", - # Gunicorn logger class: loaded via --logger-class flag, not imported in source - "jslog4kube", # Optional runtime dep of sentry-sdk for FlaskIntegration "blinker", # Transitive of Flask 1.x / Jinja2 2.11; pinned <2.1 here so soft_unicode @@ -233,7 +234,8 @@ ignore_unused = [ ] [tool.pytest.ini_options] -testpaths = ["osprey_worker"] +testpaths = ["osprey_worker", "example_plugins"] +asyncio_mode = "auto" [tool.mypy] plugins = ["pydantic.mypy", "sqlalchemy.ext.mypy.plugin"] @@ -245,6 +247,7 @@ disable_error_code = ["annotation-unchecked"] mypy_path = [ "osprey_rpc/src", "osprey_worker/src", + "osprey_async_worker/src", "example_plugins/src", ] diff --git a/uv.lock b/uv.lock index 618247e..9117960 100644 --- a/uv.lock +++ b/uv.lock @@ -17,6 +17,7 @@ resolution-markers = [ [manifest] members = [ "example-plugins", + "osprey-async-worker", "osprey-rpc", "osprey-worker", ] @@ -63,7 +64,6 @@ common = [ { name = "grpcio-tools", marker = "platform_machine == 'x86_64'", specifier = "==1.49.1" }, { name = "gunicorn", git = "https://github.com/discord/gunicorn.git?rev=979efdcb918daa536d8923668241c6e6bf1edb58" }, { name = "intervals", specifier = "==0.9.2" }, - { name = "jslog4kube", specifier = "==1.0.6" }, { name = "jsonpath-rw", url = "https://github.com/kennknowles/python-jsonpath-rw/archive/6f5647bb3ad2395c20f0191fef07a1df51c9fed8.tar.gz" }, { name = "kafka-python", specifier = "==1.4.7" }, { name = "markupsafe", specifier = "<2.1" }, @@ -114,7 +114,7 @@ common = [ { name = "types-six", specifier = "==1.17.0.20250515" }, { name = "types-urllib3", specifier = "==1.26.25.14" }, { name = "types-werkzeug", specifier = "==1.0.9" }, - { name = "typing-extensions", specifier = "==4.6.3" }, + { name = "typing-extensions", specifier = "==4.12.2" }, { name = "typing-inspect", specifier = "==0.9.0" }, { name = "unidecode", specifier = "==1.3.8" }, { name = "werkzeug", specifier = "==1.0.1" }, @@ -127,6 +127,7 @@ dev = [ { name = "mypy-extensions", specifier = "==1.0.0" }, { name = "mypy-protobuf", specifier = "==3.6.0" }, { name = "pre-commit", specifier = ">=4.3.0" }, + { name = "pytest-asyncio" }, { name = "pytest-flask", specifier = "==1.3.0" }, { name = "pytest-order", specifier = "==1.3.0" }, { name = "requests-mock", specifier = "==1.12.1" }, @@ -146,6 +147,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/18/a6/907a406bb7d359e6a63f99c313846d9eec4f7e6f7437809e03aa00fa3074/absl_py-2.4.0-py3-none-any.whl", hash = "sha256:88476fd881ca8aab94ffa78b7b6c632a782ab3ba1cd19c9bd423abc4fb4cd28d", size = 135750, upload-time = "2026-01-28T10:17:04.19Z" }, ] +[[package]] +name = "aiodns" +version = "4.0.4" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pycares" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/9b/22/a2d928e0e42baad0471d12ec44c71152ac870486e8298dddb2893b888c29/aiodns-4.0.4.tar.gz", hash = "sha256:cb10e0c0d2591636716ad2fe402e977c16d71bdaf76bb8cb49e8a6633596f736", size = 29918, upload-time = "2026-05-20T01:54:15.557Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7f/70/72e4ab117425ccdc4d10bd523a94c1baa051a15586057d64a4c6888f9e3f/aiodns-4.0.4-py3-none-any.whl", hash = "sha256:c24dd605bac70a1676ce503f967a98483ff163507198557d8e9db16267e6cfd2", size = 12696, upload-time = "2026-05-20T01:54:14.134Z" }, +] + [[package]] name = "argon2-cffi" version = "25.1.0" @@ -1271,18 +1284,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/7e/c2/1eece8c95ddbc9b1aeb64f5783a9e07a286de42191b7204d67b7496ddf35/Jinja2-2.11.3-py2.py3-none-any.whl", hash = "sha256:03e47ad063331dd6a3f04a43eddca8a966a26ba0c5b7207a9a9e4e08f1b29419", size = 125699, upload-time = "2021-01-31T16:33:07.289Z" }, ] -[[package]] -name = "jslog4kube" -version = "1.0.6" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "python-json-logger" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/58/19/a1d395d5de8998e889304b76ba2239659611bc43a2bae23cc0c9c792cbb2/jslog4kube-1.0.6.tar.gz", hash = "sha256:4b2fa3a9f9b920b74dae9d6707359b2bf62730ba502b8e92a918e8a5def036ff", size = 10971, upload-time = "2019-06-25T16:04:07.676Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/cc/7d/484b541fdd060ac9a145e061dcd3c45ad29292db21b828add95667728550/jslog4kube-1.0.6-py2.py3-none-any.whl", hash = "sha256:07cbf494e5265198726b8203460401105e099a2816ba4ec9dae3ed5fe21322fb", size = 10881, upload-time = "2019-06-25T16:04:06.37Z" }, -] - [[package]] name = "jsonpath-rw" version = "1.4.0" @@ -1701,6 +1702,25 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a3/ca/9520cc1f3dfbbd03ac5903bbf55833e257bc64b1cf30fa8b0d6df374d821/opentelemetry_api-1.42.1-py3-none-any.whl", hash = "sha256:51a69edacadbc03a8950ace1c4c21099cacc538820ac2c9e36277e78cebba714", size = 61311, upload-time = "2026-05-21T16:32:28.822Z" }, ] +[[package]] +name = "osprey-async-worker" +version = "0.1.0" +source = { editable = "osprey_async_worker" } +dependencies = [ + { name = "aiodns" }, + { name = "osprey-rpc" }, + { name = "osprey-worker" }, + { name = "pycares" }, +] + +[package.metadata] +requires-dist = [ + { name = "aiodns" }, + { name = "osprey-rpc", editable = "osprey_rpc" }, + { name = "osprey-worker", editable = "osprey_worker" }, + { name = "pycares" }, +] + [[package]] name = "osprey-rpc" version = "0.1.0" @@ -1944,6 +1964,77 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/47/8d/d529b5d697919ba8c11ad626e835d4039be708a35b0d22de83a269a6682c/pyasn1_modules-0.4.2-py3-none-any.whl", hash = "sha256:29253a9207ce32b64c3ac6600edc75368f98473906e8fd1043bd6b5b1de2c14a", size = 181259, upload-time = "2025-03-28T02:41:19.028Z" }, ] +[[package]] +name = "pycares" +version = "5.0.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "cffi" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/df/a0/9c823651872e6a0face3f0311de2a40c8bbcb9c8dcb15680bd019ac56ac7/pycares-5.0.1.tar.gz", hash = "sha256:5a3c249c830432631439815f9a818463416f2a8cbdb1e988e78757de9ae75081", size = 652222, upload-time = "2026-01-01T12:37:00.604Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/87/78/43b09f4b8e5fb8a6024661b458b48987abdb39304c78117b106b10a029f1/pycares-5.0.1-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:c29ca77ff9712e20787201ca8e76ad89384771c0e058a0a4f3dc05afbc4b32de", size = 136177, upload-time = "2026-01-01T12:35:11.567Z" }, + { url = "https://files.pythonhosted.org/packages/19/05/194c0e039ff52b166b50e79ff166c61f931fbca2bf94fc0dbaaf39041518/pycares-5.0.1-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:f11424bf5cf6226d0b136ed47daa58434e377c61b62d0100d1de7793f8e34a72", size = 130960, upload-time = "2026-01-01T12:35:12.828Z" }, + { url = "https://files.pythonhosted.org/packages/0d/84/5fce65cc058c5ab619c0dd1370d539667235a5565da72ca77f3f741cdc70/pycares-5.0.1-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:d765afb52d579879f5c4f005763827d3b1eb86b23139e9614e6089c9f98db017", size = 220584, upload-time = "2026-01-01T12:35:14.005Z" }, + { url = "https://files.pythonhosted.org/packages/f6/74/d82304297308f6c24a17961bf589b53eefa5f7f2724158c842c67fa0b302/pycares-5.0.1-cp311-cp311-manylinux_2_26_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:ea0d57ba5add4bfbcc40cbdfa92bbb8a5ef0c4c21881e26c7229d9bdc92a4533", size = 252166, upload-time = "2026-01-01T12:35:15.293Z" }, + { url = "https://files.pythonhosted.org/packages/39/a2/0ead3ba4228a490b52eb44d43514dae172c90421bb30a3659516e5b251a2/pycares-5.0.1-cp311-cp311-manylinux_2_26_s390x.manylinux_2_28_s390x.whl", hash = "sha256:ae9ec2aa3553d33e6220aeb1a05f4853fb83fce4cec3e0dea2dc970338ea47dc", size = 239085, upload-time = "2026-01-01T12:35:16.594Z" }, + { url = "https://files.pythonhosted.org/packages/26/ad/e59f173933f0e696a6afbbd63935114d1400524a72da4f2cbafc6002a398/pycares-5.0.1-cp311-cp311-manylinux_2_26_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:5c63fb2498b05e9f5670a1bf3b900c5d09343b3b6d5001a9714d593f9eb54de1", size = 222936, upload-time = "2026-01-01T12:35:17.521Z" }, + { url = "https://files.pythonhosted.org/packages/98/fa/d85bfe663a9c292efd8e699779027612c0c65ff50dc4cc9eb7a143613460/pycares-5.0.1-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:71316f7a87c15a8d32127ff01374dc2c969c37410693cc0cf6532590b7f18e7a", size = 223506, upload-time = "2026-01-01T12:35:18.535Z" }, + { url = "https://files.pythonhosted.org/packages/2a/6b/4c225a5b10a4c9f88891a20bfe363eca1b1ce7d5244b396e5683c6070998/pycares-5.0.1-cp311-cp311-musllinux_1_2_ppc64le.whl", hash = "sha256:a2117dffbb78615bfdb41ad77b17038689e4e01c66f153649e80d268c6228b4f", size = 251633, upload-time = "2026-01-01T12:35:19.819Z" }, + { url = "https://files.pythonhosted.org/packages/26/ce/ba2349413b5197b72ec19c46e07f6be3a324f80a7b1579c7cbb1b82d6dc2/pycares-5.0.1-cp311-cp311-musllinux_1_2_s390x.whl", hash = "sha256:7d7c4f5d8b88b586ef2288142b806250020e6490b9f2bd8fd5f634a78fd20fcf", size = 237703, upload-time = "2026-01-01T12:35:20.827Z" }, + { url = "https://files.pythonhosted.org/packages/84/2f/1fd794e6fca10d9e20569113d10a4f92cc2b4242d3eb45524419a37cca6b/pycares-5.0.1-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:433b9a4b5a7e10ef8aef0b957e6cd0bfc1bb5bc730d2729f04e93c91c25979c0", size = 222622, upload-time = "2026-01-01T12:35:22.518Z" }, + { url = "https://files.pythonhosted.org/packages/c9/07/7db7977649b210092a7e02d550fcebdfa69bc995c684a3b960c88a5dc4ce/pycares-5.0.1-cp311-cp311-win_amd64.whl", hash = "sha256:cf2699883b88713670d3f9c0a1e44ac24c70aeace9f8c6aa7f0b9f222d5b08a5", size = 117438, upload-time = "2026-01-01T12:35:23.402Z" }, + { url = "https://files.pythonhosted.org/packages/fc/ca/f322ddaa8b3414667de8faeea944ce9d3ddfaf1455839f499a21fcea4cec/pycares-5.0.1-cp311-cp311-win_arm64.whl", hash = "sha256:9528dc11749e5e098c996475b60f879e1db5a6cb3dd0cdc747530620bb1a8941", size = 108920, upload-time = "2026-01-01T12:35:24.599Z" }, + { url = "https://files.pythonhosted.org/packages/75/67/e84ba11d3fec3bf1322c3b302c4df13c85e0a1bc48f16d65cd0f59ad9853/pycares-5.0.1-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:2ee551be4f3f3ac814ac8547586c464c9035e914f5122a534d25de147fa745e1", size = 136241, upload-time = "2026-01-01T12:35:25.439Z" }, + { url = "https://files.pythonhosted.org/packages/ce/ae/50fbb3b4e52b9f1d16a36ffabd051ef8b2106b3f0a0d1c1113904d187a9d/pycares-5.0.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:252d4e5a52a68f825eaa90e16b595f9baee22c760f51e286ab612c6829b96de3", size = 131069, upload-time = "2026-01-01T12:35:26.293Z" }, + { url = "https://files.pythonhosted.org/packages/0e/ea/f431599f1ac42149ea4768e516db7cdae3a503a6646319ae63ab66da1486/pycares-5.0.1-cp312-cp312-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8c1aa549b8c2f2e224215c793d660270778dcba9abc3b85abbc7c41eabe4f1e5", size = 221120, upload-time = "2026-01-01T12:35:27.143Z" }, + { url = "https://files.pythonhosted.org/packages/6e/4f/0a7a6c8b3a64ee5149e935c167cd8ba5d1fdd766ec03e273dbc7502f7bea/pycares-5.0.1-cp312-cp312-manylinux_2_26_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:db7c9c9f16e8311998667a7488e817f8cbeedec2447bac827c71804663f1437e", size = 252228, upload-time = "2026-01-01T12:35:28.443Z" }, + { url = "https://files.pythonhosted.org/packages/49/3d/7f9fd20e97ee30c4b959f87ab26e47ddcef666e5e7717e45f2245fe9d70a/pycares-5.0.1-cp312-cp312-manylinux_2_26_s390x.manylinux_2_28_s390x.whl", hash = "sha256:4b9c4c8bb69bab863f677fa166653bb872bfa5d5a742f1f30bebc2d53b6e71db", size = 239473, upload-time = "2026-01-01T12:35:29.794Z" }, + { url = "https://files.pythonhosted.org/packages/a4/d0/c67967a10abd89529cb9aded9d73f43e5de00cf21243638ef529f6757262/pycares-5.0.1-cp312-cp312-manylinux_2_26_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:09ef90da8da3026fcba4ed223bd71e8057608d5b3fec4f5990b52ae1e8c855cc", size = 223831, upload-time = "2026-01-01T12:35:30.781Z" }, + { url = "https://files.pythonhosted.org/packages/4f/9a/94aacaf22a20b7d342c8f18bf006be57967beef6319adc668d4d86b627be/pycares-5.0.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:ce193ebd54f4c74538b751ebb0923a9208c234ff180589d4d3cec134c001840e", size = 223963, upload-time = "2026-01-01T12:35:31.691Z" }, + { url = "https://files.pythonhosted.org/packages/e6/e1/3666aab6fc5e7d0c669b981fe0407e6a4b67e4e6a37ac429d440274663d5/pycares-5.0.1-cp312-cp312-musllinux_1_2_ppc64le.whl", hash = "sha256:36b9ff18ef231277f99a846feade50b417187a96f742689a9d08b9594e386de4", size = 251813, upload-time = "2026-01-01T12:35:32.918Z" }, + { url = "https://files.pythonhosted.org/packages/94/44/ddab5fbc16ad0084a827167ae8628f54c7a55ce6b743585e6f47a5dd527e/pycares-5.0.1-cp312-cp312-musllinux_1_2_s390x.whl", hash = "sha256:5e40ea4a0ef0c01a02ef7f7390a58c62d237d5ad48d36bc3245e9c2ac181cc22", size = 238181, upload-time = "2026-01-01T12:35:34.078Z" }, + { url = "https://files.pythonhosted.org/packages/66/27/05467933e0e5c4e712302a2d7499797bc3029bf4d0d8ffbfe737254482b7/pycares-5.0.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:3f323b0ddfd2c7896af6fba4f8851d34d3d13387566aa573d93330fb01cb1038", size = 223552, upload-time = "2026-01-01T12:35:35.076Z" }, + { url = "https://files.pythonhosted.org/packages/3e/e2/14f3837e943d46ee12441fe6aaa418fdb2f698d42e179f368eaa9829744b/pycares-5.0.1-cp312-cp312-win_amd64.whl", hash = "sha256:bdc6bcafb72a97b3cdd529fc87210e59e67feb647a7e138110656023599b84da", size = 117478, upload-time = "2026-01-01T12:35:36.133Z" }, + { url = "https://files.pythonhosted.org/packages/d3/c3/3284061f18188d5085338e1f1fd4f03d9c135657acf16f8020b9dd3be5fc/pycares-5.0.1-cp312-cp312-win_arm64.whl", hash = "sha256:f8ef4c70c1edaf022875a8f9ff6c0c064f82831225acc91aa1b4f4d389e2e03a", size = 108889, upload-time = "2026-01-01T12:35:37.135Z" }, + { url = "https://files.pythonhosted.org/packages/92/0a/6bd9bdc2d0ee23ff3aabab7747212e2c5323a081b9b745624d62df88f7e9/pycares-5.0.1-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:7d1b2c6b152c65f14d0e12d741fabb78a487f0f0d22773eede8d8cfc97af612b", size = 136242, upload-time = "2026-01-01T12:35:38.372Z" }, + { url = "https://files.pythonhosted.org/packages/18/2a/2e9f888fc076cfe7a3493a3c4113e787cc4b4533f531dfb562ac9b04898f/pycares-5.0.1-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:8c8ffcc9a48cfc296fe1aefc07d2c8e29a7f97e4bb366ce17effea6a38825f70", size = 131070, upload-time = "2026-01-01T12:35:39.262Z" }, + { url = "https://files.pythonhosted.org/packages/ec/5b/83b5aaf7b6ed102f63cd768a747b6cb5d4624f2eaecd84868d103b9dbf39/pycares-5.0.1-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:b8efc38c2703e3530b823a4165a7b28d7ce0fdcf41960fb7a4ca834a0f8cfe79", size = 221137, upload-time = "2026-01-01T12:35:40.155Z" }, + { url = "https://files.pythonhosted.org/packages/33/d3/d77ab0b33fb805d02896c385176c462e3386d94457a5e508245c39f41829/pycares-5.0.1-cp313-cp313-manylinux_2_26_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:e380bf6eff42c260f829a0a14547e13375e949053a966c23ca204a13647ef265", size = 252252, upload-time = "2026-01-01T12:35:41.287Z" }, + { url = "https://files.pythonhosted.org/packages/14/32/8afbc798bce26dfcc5bc1f6bf1560d31cdd0af837ff52cbede657bf9262e/pycares-5.0.1-cp313-cp313-manylinux_2_26_s390x.manylinux_2_28_s390x.whl", hash = "sha256:35dd5858ee1246bd092a212b5e85a8ef70853f7cfaf16b99569bf4af3ae4695d", size = 239447, upload-time = "2026-01-01T12:35:42.614Z" }, + { url = "https://files.pythonhosted.org/packages/61/1b/a056393fda383b2eda5dab20bd0dd034fd631bf5ae754aabb20da815bdfe/pycares-5.0.1-cp313-cp313-manylinux_2_26_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:c257c6e7bf310cdb5823aa9d9a28f1e370fed8c653a968d38a954a8f8e0375ce", size = 223822, upload-time = "2026-01-01T12:35:43.594Z" }, + { url = "https://files.pythonhosted.org/packages/ca/c7/9817f0fb954ab9926f88403f2b91a3e4984a277e2b7a4563e0118e4e1ffa/pycares-5.0.1-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:07711acb0ef75758f081fb7436acaccc91e8afd5ae34fd35d4edc44297e81f27", size = 223986, upload-time = "2026-01-01T12:35:44.893Z" }, + { url = "https://files.pythonhosted.org/packages/e1/a9/c0ea15c871c77e8c20bcaab18f56ae83988ea4c302155d106cc6a1bd83a9/pycares-5.0.1-cp313-cp313-musllinux_1_2_ppc64le.whl", hash = "sha256:30e5db1ae85cffb031dd8bc1b37903cd74c6d37eb737643bbca3ff2cd4bc6ae2", size = 251838, upload-time = "2026-01-01T12:35:46.271Z" }, + { url = "https://files.pythonhosted.org/packages/be/a4/fe4068abfadf3e06cc22333e87e4730de3c170075572041d5545926062a3/pycares-5.0.1-cp313-cp313-musllinux_1_2_s390x.whl", hash = "sha256:efbe7f89425a14edbc94787042309be77cb3674415eb6079b356e1f9552ba747", size = 238238, upload-time = "2026-01-01T12:35:47.196Z" }, + { url = "https://files.pythonhosted.org/packages/a7/25/4f140518768d974af4221cfd574a30d99d40b3d5c54c479da2c1553be59e/pycares-5.0.1-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:5de9e7ce52d638d78723c24704eb032e60b96fbb6fe90c6b3110882987251377", size = 223574, upload-time = "2026-01-01T12:35:48.191Z" }, + { url = "https://files.pythonhosted.org/packages/1e/0a/6e4afa4a2baffd1eba6c18a90cda17681d4838d3cab5a485e471386e04dc/pycares-5.0.1-cp313-cp313-win_amd64.whl", hash = "sha256:0e99af0a1ce015ab6cc6bd85ce158d95ed89fb3b654515f1d0989d1afcf11026", size = 117472, upload-time = "2026-01-01T12:35:50.674Z" }, + { url = "https://files.pythonhosted.org/packages/57/d0/a99f97e9aa8c8404fc899540cf30be63cda0df5150e3c0837423917c7e4c/pycares-5.0.1-cp313-cp313-win_arm64.whl", hash = "sha256:2a511c9f3b11b7ce9f159c956ea1b8f2de7f419d7ca9fa24528d582cb015dbf9", size = 108889, upload-time = "2026-01-01T12:35:51.902Z" }, + { url = "https://files.pythonhosted.org/packages/38/b2/4af99ff17acb81377c971831520540d1859bf401dc85712eb4abc2e6751f/pycares-5.0.1-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:e330e3561be259ad7a1b7b0ce282c872938625f76587fae7ac8d6bc5af1d0c3d", size = 136635, upload-time = "2026-01-01T12:35:53.365Z" }, + { url = "https://files.pythonhosted.org/packages/42/da/e2e1683811c427492ee0e86e8fae8d55eb5cca032220438599991fdad866/pycares-5.0.1-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:82bd37fec2a3fa62add30d4a3854720f7b051386e2f18e6e8f4ee94b89b5a7b0", size = 131093, upload-time = "2026-01-01T12:35:54.28Z" }, + { url = "https://files.pythonhosted.org/packages/cd/2a/9cf2120cafc19e5c589d5252a9ddd3108cc87e9db09938d16317807de03b/pycares-5.0.1-cp314-cp314-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:258c38aaa82ad1d565b4591cdb93d2c191be8e0a2c70926999c8e0b717a01f2a", size = 221096, upload-time = "2026-01-01T12:35:57.096Z" }, + { url = "https://files.pythonhosted.org/packages/2c/cc/c5fbf6377e2d6b1f1618f147ad898e5d8ae1585fc726d6301f07aeda6cac/pycares-5.0.1-cp314-cp314-manylinux_2_26_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:ccc1b2df8a09ca20eefbe20b9f7a484d376525c0fb173cfadd692320013c6bc5", size = 252330, upload-time = "2026-01-01T12:35:58.182Z" }, + { url = "https://files.pythonhosted.org/packages/3b/df/17a7c518c45bb994f76d9064d2519674e2a3950f895abbe6af123ead04ac/pycares-5.0.1-cp314-cp314-manylinux_2_26_s390x.manylinux_2_28_s390x.whl", hash = "sha256:3c4dfc80cc8b43dc79e02a15486c58eead5cae0a40906d6be64e2522285b5b39", size = 239799, upload-time = "2026-01-01T12:36:00.378Z" }, + { url = "https://files.pythonhosted.org/packages/3f/6c/d79c94809742b56b9180a9a9ec2937607db0b8eb34b8ca75d86d3114d6dd/pycares-5.0.1-cp314-cp314-manylinux_2_26_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f498a6606247bfe896c2a4d837db711eb7b0ba23e409e16e4b23def4bada4b9d", size = 223501, upload-time = "2026-01-01T12:36:02.695Z" }, + { url = "https://files.pythonhosted.org/packages/69/08/83084b67cbce08f44fd803b88816fc80d2fe2fb3d483d5432925df44371b/pycares-5.0.1-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:a7d197835cdb4b202a3b12562b32799e27bb132262d4aa1ac3ee9d440e8ec22c", size = 223708, upload-time = "2026-01-01T12:36:04.357Z" }, + { url = "https://files.pythonhosted.org/packages/15/57/63a6e9ef356c5149b8ec72a694e02207fd8ae643895aeb78a9f0c07f1502/pycares-5.0.1-cp314-cp314-musllinux_1_2_ppc64le.whl", hash = "sha256:f78ab823732b050d658eb735d553726663c9bccdeeee0653247533a23eb2e255", size = 251816, upload-time = "2026-01-01T12:36:05.618Z" }, + { url = "https://files.pythonhosted.org/packages/43/1c/1c85c6355cf7bc3ae86a1024d60f9cabdc12af63306a5f59370ac8718a41/pycares-5.0.1-cp314-cp314-musllinux_1_2_s390x.whl", hash = "sha256:f444ab7f318e9b2c209b45496fb07bff5e7ada606e15d5253a162964aa078527", size = 238259, upload-time = "2026-01-01T12:36:07.609Z" }, + { url = "https://files.pythonhosted.org/packages/5d/7f/bd5ff5a460e50433f993560e4e5d229559a8bf271dbdf6be832faf1973b5/pycares-5.0.1-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:9de80997de7538619b7dd28ec4371e5172e3f9480e4fc648726d3d5ba661ca05", size = 223732, upload-time = "2026-01-01T12:36:09.893Z" }, + { url = "https://files.pythonhosted.org/packages/b5/fe/e77738366e00dc0918bbeb0c8fc63579e5d9cec748a2b838e207e548b5d9/pycares-5.0.1-cp314-cp314-win_amd64.whl", hash = "sha256:206ce9f3cb9d51f5065c81b23c22996230fbc2cf58ae22834c623631b2b473aa", size = 120847, upload-time = "2026-01-01T12:36:11.494Z" }, + { url = "https://files.pythonhosted.org/packages/81/17/758e9af7ee8589ac6deddf7ea56d75b982f155bc2052ef61c45d5f371389/pycares-5.0.1-cp314-cp314-win_arm64.whl", hash = "sha256:45fb3b07231120e8cb5b75be7f15f16115003e9251991dc37a3e5c63733d63b5", size = 112595, upload-time = "2026-01-01T12:36:12.973Z" }, + { url = "https://files.pythonhosted.org/packages/56/12/4f1d418fed957fc96089c69d9ec82314b3b91c48c7f9463385842acad9c4/pycares-5.0.1-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:602f3eac4b880a2527d21f52b2319cb10fde9225d103d338c4d0b2b07f136849", size = 137061, upload-time = "2026-01-01T12:36:15.027Z" }, + { url = "https://files.pythonhosted.org/packages/29/8c/559cea98a8a5d0f38b50b4b812a07fdbcdb1a961bed9e2e9d5d343e53c6f/pycares-5.0.1-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:a1c3736deef003f0c57bc4e7f94d54270d0824350a8f5ceaba3a20b2ce8fb427", size = 131551, upload-time = "2026-01-01T12:36:16.74Z" }, + { url = "https://files.pythonhosted.org/packages/34/cd/aee5d8070888d7be509d4f32a348e2821309ec67980498e5a974cd9e4990/pycares-5.0.1-cp314-cp314t-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:e63328df86d37150ce697fb5d9313d1d468dd4dddee1d09342cb2ed241ce6ad9", size = 230409, upload-time = "2026-01-01T12:36:18.909Z" }, + { url = "https://files.pythonhosted.org/packages/5e/94/15d5cf7d8e7af4b4ce3e19ea117dfe565c08d60d82f043ad23843703a135/pycares-5.0.1-cp314-cp314t-manylinux_2_26_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:57f6fd696213329d9a69b9664a68b1ff2a71ccbdc1fc928a42c9a92858c1ec5d", size = 261297, upload-time = "2026-01-01T12:36:20.771Z" }, + { url = "https://files.pythonhosted.org/packages/af/46/24f6ddc7a37ec6eaa1c38f617f39624211d8e7cdca49b644bfc5f467f275/pycares-5.0.1-cp314-cp314t-manylinux_2_26_s390x.manylinux_2_28_s390x.whl", hash = "sha256:9d0878edabfbecb48a29e8769284003d8dbc05936122fe361849cd5fa52722e0", size = 248071, upload-time = "2026-01-01T12:36:22.925Z" }, + { url = "https://files.pythonhosted.org/packages/fa/f0/7eb7fe44f0db55b9083725ab7a084874c2dc02806d9613e07e719838c2ab/pycares-5.0.1-cp314-cp314t-manylinux_2_26_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:50e21f27a91be122e066ddd78c2d0d2769e547561481d8342a9d652a345b89f7", size = 232073, upload-time = "2026-01-01T12:36:25.773Z" }, + { url = "https://files.pythonhosted.org/packages/1d/cd/993b17e0c049a56b5af4df3fd053acc57b37e17e0dcd709b2d337c22d57d/pycares-5.0.1-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:97ceda969f5a5d5c6b15558b658c29e4301b3a2c4615523797b5f9d4ac74772e", size = 232815, upload-time = "2026-01-01T12:36:27.798Z" }, + { url = "https://files.pythonhosted.org/packages/7a/ff/170177bcc5dff31e735f209f5de63362f513ac18846c83d50e4e68f57866/pycares-5.0.1-cp314-cp314t-musllinux_1_2_ppc64le.whl", hash = "sha256:4d1713e602ab09882c3e65499b2cc763bff0371117327cad704cf524268c2604", size = 261111, upload-time = "2026-01-01T12:36:29.94Z" }, + { url = "https://files.pythonhosted.org/packages/4d/4a/4c6497b8ca9279b4038ee8c7e2c49504008d594d06a044e00678b30c10fe/pycares-5.0.1-cp314-cp314t-musllinux_1_2_s390x.whl", hash = "sha256:954a379055d6c66b2e878b52235b382168d1a3230793ff44454019394aecac5e", size = 246311, upload-time = "2026-01-01T12:36:31.352Z" }, + { url = "https://files.pythonhosted.org/packages/06/19/1603f51f0d73bf34017a9e6967540c2bc138f9541aa7cc1ef38990b3ce9d/pycares-5.0.1-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:145d8a20f7fd1d58a2e49b7ef4309ec9bdcab479ac65c2e49480e20d3f890c23", size = 232027, upload-time = "2026-01-01T12:36:34.374Z" }, + { url = "https://files.pythonhosted.org/packages/7a/de/c000a682757b84688722ac232a24a86b6f195f1f4732432ecf35d0a768a5/pycares-5.0.1-cp314-cp314t-win_amd64.whl", hash = "sha256:ebc9daba03c7ff3f62616c84c6cb37517445d15df00e1754852d6006039eb4a4", size = 121267, upload-time = "2026-01-01T12:36:35.741Z" }, + { url = "https://files.pythonhosted.org/packages/b2/c4/8bfffecd08b9b198113fcff5f0ab84bbe696f07dec46dd1ccae0e7b28c23/pycares-5.0.1-cp314-cp314t-win_arm64.whl", hash = "sha256:e0a86eff6bf9e91d5dd8876b1b82ee45704f46b1104c24291d3dea2c1fc8ebcb", size = 113043, upload-time = "2026-01-01T12:36:37.895Z" }, +] + [[package]] name = "pycparser" version = "3.0" @@ -2046,6 +2137,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d4/24/a372aaf5c9b7208e7112038812994107bc65a84cd00e0354a88c2c77a617/pytest-9.0.3-py3-none-any.whl", hash = "sha256:2c5efc453d45394fdd706ade797c0a81091eccd1d6e4bccfcd476e2b8e0ab5d9", size = 375249, upload-time = "2026-04-07T17:16:16.13Z" }, ] +[[package]] +name = "pytest-asyncio" +version = "1.4.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pytest" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/43/7c/d36d04db312ecf4298932ef77e6e4a9e8ad017906e24e34f0b0c361a2473/pytest_asyncio-1.4.0.tar.gz", hash = "sha256:c6c0d2259945122819f171a32ecea2c349ead889ee28176caaf492143424be42", size = 58514, upload-time = "2026-05-26T09:56:04.083Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/03/e2/08a497ef684b88559c9cc5f4ad53a37e7b99e727094a86d6ea32536d5d3c/pytest_asyncio-1.4.0-py3-none-any.whl", hash = "sha256:933ca923a23075a87fb7070c0ec272a6848489824d887c85c812670932835aa1", size = 16930, upload-time = "2026-05-26T09:56:02.576Z" }, +] + [[package]] name = "pytest-flask" version = "1.3.0" @@ -2720,11 +2824,11 @@ wheels = [ [[package]] name = "typing-extensions" -version = "4.6.3" +version = "4.12.2" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/42/56/cfaa7a5281734dadc842f3a22e50447c675a1c5a5b9f6ad8a07b467bffe7/typing_extensions-4.6.3.tar.gz", hash = "sha256:d91d5919357fe7f681a9f2b5b4cb2a5f1ef0a1e9f59c4d8ff0d3491e05c0ffd5", size = 65757, upload-time = "2023-06-01T23:55:36.332Z" } +sdist = { url = "https://files.pythonhosted.org/packages/df/db/f35a00659bc03fec321ba8bce9420de607a1d37f8342eee1863174c69557/typing_extensions-4.12.2.tar.gz", hash = "sha256:1a7ead55c7e559dd4dee8856e3a88b41225abfe1ce8df57b7c13915fe121ffb8", size = 85321, upload-time = "2024-06-07T18:52:15.995Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5f/86/d9b1518d8e75b346a33eb59fa31bdbbee11459a7e2cc5be502fa779e96c5/typing_extensions-4.6.3-py3-none-any.whl", hash = "sha256:88a4153d8505aabbb4e13aacb7c486c2b4a33ca3b3f807914a9b4c844c471c26", size = 31329, upload-time = "2023-06-01T23:55:34.451Z" }, + { url = "https://files.pythonhosted.org/packages/26/9f/ad63fc0248c5379346306f8668cda6e2e2e9c95e01216d2b8ffd9ff037d0/typing_extensions-4.12.2-py3-none-any.whl", hash = "sha256:04e5ca0351e0f3f85c6853954072df659d0d13fac324d0072316b67d7794700d", size = 37438, upload-time = "2024-06-07T18:52:13.582Z" }, ] [[package]]