diff --git a/packages/pynumaflow-lite/manifests/mapstream/.dockerignore b/packages/pynumaflow-lite/manifests/mapstream/.dockerignore new file mode 100644 index 00000000..21d0b898 --- /dev/null +++ b/packages/pynumaflow-lite/manifests/mapstream/.dockerignore @@ -0,0 +1 @@ +.venv/ diff --git a/packages/pynumaflow-lite/manifests/mapstream/Dockerfile b/packages/pynumaflow-lite/manifests/mapstream/Dockerfile index 3caa0111..aea0c342 100644 --- a/packages/pynumaflow-lite/manifests/mapstream/Dockerfile +++ b/packages/pynumaflow-lite/manifests/mapstream/Dockerfile @@ -1,38 +1,43 @@ -FROM python:3.11-slim-bullseye AS builder - -ENV PYTHONFAULTHANDLER=1 \ - PYTHONUNBUFFERED=1 \ - PYTHONHASHSEED=random \ - PIP_NO_CACHE_DIR=on \ - PIP_DISABLE_PIP_VERSION_CHECK=on \ - PIP_DEFAULT_TIMEOUT=100 \ - POETRY_HOME="/opt/poetry" \ - POETRY_VIRTUALENVS_IN_PROJECT=true \ - POETRY_NO_INTERACTION=1 \ - PYSETUP_PATH="/opt/pysetup" - - ENV PATH="$POETRY_HOME/bin:$PATH" - -RUN apt-get update \ - && apt-get install --no-install-recommends -y \ - curl \ - wget \ - # deps for building python deps - build-essential \ - && apt-get install -y git \ - && apt-get clean && rm -rf /var/lib/apt/lists/* \ - && curl -sSL https://install.python-poetry.org | python3 - - -FROM builder AS udf - -WORKDIR $PYSETUP_PATH -COPY ./ ./ - -RUN pip install $PYSETUP_PATH/pynumaflow_lite-0.1.0-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl - -RUN poetry lock -RUN poetry install --no-cache --no-root && \ - rm -rf ~/.cache/pypoetry/ +FROM python:3.11-slim-trixie AS builder -CMD ["python", "mapstream_cat.py"] +COPY --from=ghcr.io/astral-sh/uv:0.12.13 /uv /uvx /bin/ + +ENV UV_COMPILE_BYTECODE=1 \ + UV_LINK_MODE=copy \ + UV_PYTHON_DOWNLOADS=never + +WORKDIR /app + +RUN --mount=type=cache,target=/root/.cache/uv \ + --mount=type=bind,source=uv.lock,target=uv.lock \ + --mount=type=bind,source=pyproject.toml,target=pyproject.toml \ + uv sync --locked --no-install-project --no-dev + +COPY . . + +RUN --mount=type=cache,target=/root/.cache/uv \ + uv sync --locked --no-dev + +# Install the local pynumaflow-lite wheel that matches the build platform. +RUN uv pip install --no-index --find-links . pynumaflow-lite + +FROM python:3.11-slim-trixie +# Setup a non-root user +RUN groupadd --system --gid 999 nonroot \ + && useradd --system --gid 999 --uid 999 --create-home nonroot + +COPY --from=builder --chown=nonroot:nonroot /app /app + +ENV PATH="/app/.venv/bin:$PATH" + +# Keeps Python from buffering stdout and stderr to avoid situations where +# the application crashes without emitting any logs due to buffering. +ENV PYTHONUNBUFFERED=1 + +# Use the non-root user to run our application +USER nonroot + +WORKDIR /app + +CMD ["python", "mapstream_cat.py"] diff --git a/packages/pynumaflow-lite/manifests/mapstream/mapstream_cat.py b/packages/pynumaflow-lite/manifests/mapstream/mapstream_cat.py index 9cfe0beb..2ca0ef97 100644 --- a/packages/pynumaflow-lite/manifests/mapstream/mapstream_cat.py +++ b/packages/pynumaflow-lite/manifests/mapstream/mapstream_cat.py @@ -1,41 +1,24 @@ import asyncio -import signal -from collections.abc import AsyncIterator, Callable +from collections.abc import AsyncIterator -from pynumaflow_lite import mapstreamer -from pynumaflow_lite.mapstreamer import Message +from pynumaflow_lite.mapstreamer import Datum, MapStreamAsyncServer, MapStreamer, Message -class SimpleStreamCat(mapstreamer.MapStreamer): - async def handler(self, keys: list[str], datum: mapstreamer.Datum) -> AsyncIterator[Message]: - parts = datum.value.decode("utf-8").split(",") - if not parts: +class SimpleStreamCat(MapStreamer): + async def handler(self, datum: Datum) -> AsyncIterator[Message]: + if not datum.value: yield Message.to_drop() return - for s in parts: - yield Message(s.encode(), keys) + for s in datum.value.decode("utf-8").split(","): + yield Message(s.encode(), keys=datum.keys) -async def start(f: Callable[[list[str], mapstreamer.Datum], AsyncIterator[Message]]): - # Use default socket/info file locations; no explicit sock file passed - server = mapstreamer.MapStreamAsyncServer() - - # Register loop-level signal handlers so we control shutdown and avoid asyncio.run noise. - loop = asyncio.get_running_loop() - try: - loop.add_signal_handler(signal.SIGINT, lambda: server.stop()) - loop.add_signal_handler(signal.SIGTERM, lambda: server.stop()) - except (NotImplementedError, RuntimeError): - pass - - try: - await server.start(f) - print("Shutting down gracefully...") - except asyncio.CancelledError: - server.stop() - return +async def main() -> None: + print("Starting MapStream server") + # `serve` returns when SIGINT or SIGTERM arrives. + await MapStreamAsyncServer(SimpleStreamCat()).serve() + print("MapStream server stopped") if __name__ == "__main__": - async_handler = SimpleStreamCat() - asyncio.run(start(async_handler)) + asyncio.run(main()) diff --git a/packages/pynumaflow-lite/manifests/mapstream/pyproject.toml b/packages/pynumaflow-lite/manifests/mapstream/pyproject.toml index 73b6aba5..c28c807a 100644 --- a/packages/pynumaflow-lite/manifests/mapstream/pyproject.toml +++ b/packages/pynumaflow-lite/manifests/mapstream/pyproject.toml @@ -6,11 +6,6 @@ authors = [ { name = "Vigith Maurice", email = "vigith@gmail.com" } ] readme = "README.md" -requires-python = ">=3.11" +requires-python = "==3.11.*" dependencies = [ ] - - -[build-system] -requires = ["poetry-core>=2.0.0,<3.0.0"] -build-backend = "poetry.core.masonry.api" \ No newline at end of file diff --git a/packages/pynumaflow-lite/manifests/mapstream/uv.lock b/packages/pynumaflow-lite/manifests/mapstream/uv.lock new file mode 100644 index 00000000..396da73f --- /dev/null +++ b/packages/pynumaflow-lite/manifests/mapstream/uv.lock @@ -0,0 +1,8 @@ +version = 1 +revision = 3 +requires-python = "==3.11.*" + +[[package]] +name = "mapstream-cat" +version = "0.1.0" +source = { virtual = "." } diff --git a/packages/pynumaflow-lite/pynumaflow_lite/__init__.py b/packages/pynumaflow-lite/pynumaflow_lite/__init__.py index c346b880..8b23548b 100644 --- a/packages/pynumaflow-lite/pynumaflow_lite/__init__.py +++ b/packages/pynumaflow-lite/pynumaflow_lite/__init__.py @@ -10,6 +10,7 @@ from ._map_dtypes import Mapper from ._map_server import MapAsyncServer from ._mapstream_dtypes import MapStreamer +from ._mapstream_server import MapStreamAsyncServer from ._reduce_dtypes import Reducer from ._reducestreamer_dtypes import ReduceStreamer from ._session_reduce_dtypes import SessionReducer @@ -41,6 +42,7 @@ batchmapper.BatchMapper = BatchMapper batchmapper.BatchMapAsyncServer = BatchMapAsyncServer mapstreamer.MapStreamer = MapStreamer +mapstreamer.MapStreamAsyncServer = MapStreamAsyncServer reducer.Reducer = Reducer session_reducer.SessionReducer = SessionReducer reducestreamer.ReduceStreamer = ReduceStreamer diff --git a/packages/pynumaflow-lite/pynumaflow_lite/_mapstream_dtypes.py b/packages/pynumaflow-lite/pynumaflow_lite/_mapstream_dtypes.py index 32c2b9cb..d563fa74 100644 --- a/packages/pynumaflow-lite/pynumaflow_lite/_mapstream_dtypes.py +++ b/packages/pynumaflow-lite/pynumaflow_lite/_mapstream_dtypes.py @@ -14,9 +14,10 @@ def __call__(self, *args, **kwargs): return self.handler(*args, **kwargs) @abstractmethod - async def handler(self, keys: list[str], datum: Datum) -> AsyncIterator[Message]: + async def handler(self, datum: Datum) -> AsyncIterator[Message]: """ Implement this handler function for streaming mapping. It should be an async generator yielding Message objects. """ - pass + raise NotImplementedError + yield # makes this an async generator, so overrides that yield type-check diff --git a/packages/pynumaflow-lite/pynumaflow_lite/_mapstream_server.py b/packages/pynumaflow-lite/pynumaflow_lite/_mapstream_server.py new file mode 100644 index 00000000..c4d25ace --- /dev/null +++ b/packages/pynumaflow-lite/pynumaflow_lite/_mapstream_server.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +import asyncio +import contextlib +import signal +from collections.abc import AsyncIterator, Callable +from types import TracebackType +from typing import TypeAlias + +from .pynumaflow_lite import mapstreamer as _mapstreamer + +Datum: TypeAlias = _mapstreamer.Datum +Message: TypeAlias = _mapstreamer.Message + +_SHUTDOWN_SIGNALS = (signal.SIGINT, signal.SIGTERM) + + +class MapStreamAsyncServer: + def __init__( + self, + handler: Callable[[Datum], AsyncIterator[Message]], + *, + sock_file: str | None = None, + server_info_file: str | None = None, + install_signal_handlers: bool = True, + ) -> None: + self._core = _mapstreamer._MapStreamAsyncServer(sock_file, server_info_file) + self._handler = handler + self._install_signal_handlers = install_signal_handlers + self._task: asyncio.Task[None] | None = None + self._serving = False + self._installed_signals: list[signal.Signals] = [] + + async def serve(self) -> None: + """Run the mapstream server until it stops. + + This is the entrypoint for an application that already runs an event + loop. It returns when a shutdown signal arrives or when `stop()` runs. + """ + await self._serve(install_signal_handlers=self._install_signal_handlers) + + async def _serve(self, *, install_signal_handlers: bool) -> None: + if self._serving: + raise RuntimeError("mapstream server is already serving") + self._serving = True + try: + if install_signal_handlers: + self._add_signal_handlers() + await self._core.start(self._handler) + finally: + self._remove_signal_handlers() + self._serving = False + + def stop(self) -> None: + self._core.stop() + + async def wait_ready(self, timeout: float = 30.0) -> None: + await self._core.wait_ready(timeout) + + async def wait_for_termination(self) -> None: + """Wait until the background server task ends. + + Use this inside an `async with` block. It raises the handler error if + the server task failed. + """ + if self._task is None: + raise RuntimeError("mapstream server is not serving") + await asyncio.shield(self._task) + + def _add_signal_handlers(self) -> None: + loop = asyncio.get_running_loop() + for sig in _SHUTDOWN_SIGNALS: + try: + loop.add_signal_handler(sig, self.stop) + except (NotImplementedError, RuntimeError, OSError): + continue + self._installed_signals.append(sig) + + def _remove_signal_handlers(self) -> None: + if not self._installed_signals: + return + try: + loop = asyncio.get_running_loop() + except RuntimeError: + self._installed_signals.clear() + return + for sig in self._installed_signals: + with contextlib.suppress(NotImplementedError, OSError): + loop.remove_signal_handler(sig) + self._installed_signals.clear() + + async def __aenter__(self) -> MapStreamAsyncServer: + """Start the server in a background task and wait until it is ready. + + This form is for tests and for code that must run other work next to + the server. It never installs signal handlers. + """ + if self._task is not None and not self._task.done(): + raise RuntimeError("mapstream server is already serving") + + self._task = asyncio.create_task(self._serve(install_signal_handlers=False)) + try: + await self.wait_ready() + except BaseException: + self.stop() + task, self._task = self._task, None + if task is not None: + # Surface the server error, if there is one. It explains the + # failure better than the `wait_ready` error does. + await task + raise + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: TracebackType | None, + ) -> None: + self.stop() + if self._task is not None: + try: + await self._task + finally: + self._task = None + + def run(self) -> None: + """Run the mapstream server in a new event loop until it stops.""" + try: + asyncio.run(self.serve()) + except KeyboardInterrupt: + self.stop() diff --git a/packages/pynumaflow-lite/pynumaflow_lite/mapstreamer.pyi b/packages/pynumaflow-lite/pynumaflow_lite/mapstreamer.pyi index 8ae79e16..06232864 100644 --- a/packages/pynumaflow-lite/pynumaflow_lite/mapstreamer.pyi +++ b/packages/pynumaflow-lite/pynumaflow_lite/mapstreamer.pyi @@ -2,6 +2,9 @@ from __future__ import annotations import datetime as _dt from collections.abc import AsyncIterator, Awaitable, Callable +from types import TracebackType + +from ._mapstream_dtypes import MapStreamer as MapStreamer class NackOptions: """Per-message redelivery options for a nack.""" @@ -32,8 +35,8 @@ class Message: keys: list[str] | None = ..., tags: list[str] | None = ..., ) -> None: ... - @staticmethod - def message_to_drop() -> Message: ... + def __repr__(self) -> str: ... + def __str__(self) -> str: ... @staticmethod def to_drop() -> Message: ... @staticmethod @@ -45,23 +48,52 @@ class Datum: keys: list[str] value: bytes watermark: _dt.datetime - eventtime: _dt.datetime + event_time: _dt.datetime headers: dict[str, str] + def __init__( + self, + *, + keys: list[str] | None = ..., + value: bytes | None = ..., + event_time: _dt.datetime | None = ..., + watermark: _dt.datetime | None = ..., + headers: dict[str, str] | None = ..., + ) -> None: ... def __repr__(self) -> str: ... def __str__(self) -> str: ... -class MapStreamAsyncServer: +class _MapStreamAsyncServer: def __init__( self, sock_file: str | None = ..., - info_file: str | None = ..., + server_info_file: str | None = ..., ) -> None: ... - def start(self, py_func: Callable[[list[str], Datum], AsyncIterator[Message]]) -> Awaitable[None]: ... + def start(self, handler: Callable[[Datum], AsyncIterator[Message]]) -> Awaitable[None]: ... + def wait_ready(self, timeout: float = ...) -> Awaitable[None]: ... def stop(self) -> None: ... -class MapStreamer: - async def handler(self, keys: list[str], datum: Datum) -> AsyncIterator[Message]: ... +class MapStreamAsyncServer: + def __init__( + self, + handler: Callable[[Datum], AsyncIterator[Message]], + *, + sock_file: str | None = ..., + server_info_file: str | None = ..., + install_signal_handlers: bool = ..., + ) -> None: ... + def run(self) -> None: ... + async def serve(self) -> None: ... + def stop(self) -> None: ... + async def wait_ready(self, timeout: float = ...) -> None: ... + async def wait_for_termination(self) -> None: ... + async def __aenter__(self) -> MapStreamAsyncServer: ... + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: TracebackType | None, + ) -> None: ... __all__ = [ "Datum", diff --git a/packages/pynumaflow-lite/src/batchmap/server.rs b/packages/pynumaflow-lite/src/batchmap/server.rs index aca523dd..ebc08a32 100644 --- a/packages/pynumaflow-lite/src/batchmap/server.rs +++ b/packages/pynumaflow-lite/src/batchmap/server.rs @@ -119,8 +119,9 @@ pub(super) async fn start( let errors = Arc::new(Mutex::new(Vec::new())); - let (sig_handle, combined_rx) = crate::pyrs::setup_sig_handler(shutdown_rx); - + // Shutdown has two sources, and neither one needs a channel here. The Python + // side signals stop() through shutdown_rx. An uncaught Python error panics in + // fail(), and numaflow then shuts the server down on its own. let py_runner = PyBatchMapRunner { py_func: Arc::new(py_func), event_loop: event_loop.clone(), @@ -132,7 +133,7 @@ pub(super) async fn start( .with_server_info_file(info_file); let result = server - .start_with_shutdown(combined_rx) + .start_with_shutdown(shutdown_rx) .await .map_err(|e| pyo3::PyErr::new::(e.to_string())); @@ -153,11 +154,5 @@ pub(super) async fn start( return Err(Python::attach(|py| combine_errors(py, errors))); } - // if not finished, abort it - if !sig_handle.is_finished() { - println!("Aborting signal handler"); - sig_handle.abort(); - } - result } diff --git a/packages/pynumaflow-lite/src/mapstream/mod.rs b/packages/pynumaflow-lite/src/mapstream/mod.rs index 08f2b854..c35c5560 100644 --- a/packages/pynumaflow-lite/src/mapstream/mod.rs +++ b/packages/pynumaflow-lite/src/mapstream/mod.rs @@ -1,14 +1,16 @@ -use chrono::{DateTime, Utc}; -use numaflow::mapstream; use std::collections::HashMap; use std::sync::Mutex; +use std::time::Duration; -pub mod server; - -/// Types for streaming handler +use chrono::{DateTime, Utc}; +use numaflow::mapstream; use pyo3::prelude::*; +pub mod server; + +use crate::map::wait_for_ready; use crate::nack::NackOptions; +use crate::pyrs::bytes_literal; /// Streaming Datum mirrors MapStreamRequest for Python #[pyclass(module = "pynumaflow_lite.mapstreamer", from_py_object)] @@ -26,39 +28,64 @@ pub struct Datum { pub watermark: DateTime, /// Time of the element as seen at source or aligned after a reduce operation. #[pyo3(get)] - pub eventtime: DateTime, + pub event_time: DateTime, /// Headers associated with the message. #[pyo3(get)] pub headers: HashMap, } +#[pymethods] impl Datum { - pub(crate) fn new( - keys: Vec, - value: Vec, - watermark: DateTime, - eventtime: DateTime, - headers: HashMap, + #[new] + #[pyo3(signature = ( + *, + keys: "list[str] | None"=None, + value: "bytes | None"=None, + event_time: "datetime.datetime | None"=None, + watermark: "datetime.datetime | None"=None, + headers: "dict[str, str] | None"=None, + ) -> "Datum")] + fn new( + keys: Option>, + value: Option>, + event_time: Option>, + watermark: Option>, + headers: Option>, ) -> Self { Self { - keys, - value, - watermark, - eventtime, - headers, + keys: keys.unwrap_or_default(), + value: value.unwrap_or_default(), + watermark: watermark.unwrap_or(DateTime::::UNIX_EPOCH), + event_time: event_time.unwrap_or(DateTime::::UNIX_EPOCH), + headers: headers.unwrap_or_default(), } } + + fn __repr__(&self) -> String { + format!( + "Datum(keys={:?}, value={}, watermark={}, event_time={}, headers={:?})", + self.keys, + bytes_literal(&self.value), + self.watermark, + self.event_time, + self.headers, + ) + } + + fn __str__(&self) -> String { + self.__repr__() + } } impl From for Datum { fn from(value: numaflow::mapstream::MapStreamRequest) -> Self { - Self::new( - value.keys, - value.value, - value.watermark, - value.eventtime, - value.headers, - ) + Self { + keys: value.keys, + value: value.value, + watermark: value.watermark, + event_time: value.eventtime, + headers: value.headers, + } } } @@ -92,9 +119,9 @@ impl Message { } /// Drop a Message, do not forward to the next vertex. - #[pyo3(signature = ())] #[staticmethod] - fn message_to_drop() -> Self { + #[pyo3(signature = () -> "Message")] + fn to_drop() -> Self { Self { keys: None, value: vec![], @@ -103,13 +130,6 @@ impl Message { } } - /// Convenience alias to match example usage: Message.to_drop() - #[pyo3(signature = ())] - #[staticmethod] - fn to_drop() -> Self { - Self::message_to_drop() - } - /// A Message marked to be negatively acknowledged (retried), with optional nack options. #[staticmethod] #[pyo3(signature = (nack_options: "NackOptions | None"=None) -> "Message")] @@ -133,6 +153,26 @@ impl Message { nack_options: None, } } + + fn __repr__(&self) -> String { + format!( + "Message(value={}, keys={}, tags={}, user_metadata={})", + bytes_literal(&self.value), + self.keys + .as_ref() + .map_or_else(|| "None".to_string(), |keys| format!("{keys:?}")), + self.tags + .as_ref() + .map_or_else(|| "None".to_string(), |tags| format!("{tags:?}")), + self.nack_options + .as_ref() + .map_or_else(|| "None".to_string(), NackOptions::__repr__), + ) + } + + fn __str__(&self) -> String { + self.__repr__() + } } impl From for mapstream::Message { @@ -147,43 +187,62 @@ impl From for mapstream::Message { } /// Async MapStream Server that can be started from Python code which will run the Python UDF async generator. -#[pyclass(module = "pynumaflow_lite.mapstreamer")] +#[pyclass(name = "_MapStreamAsyncServer", module = "pynumaflow_lite.mapstreamer")] pub struct MapStreamAsyncServer { sock_file: String, - info_file: String, + server_info_file: String, shutdown_tx: Mutex>>, } #[pymethods] impl MapStreamAsyncServer { #[new] - #[pyo3(signature = (sock_file: "str | None"=mapstream::SOCK_ADDR.to_string(), info_file: "str | None"=mapstream::SERVER_INFO_FILE.to_string()) -> "MapStreamAsyncServer" - )] - fn new(sock_file: String, info_file: String) -> Self { + #[pyo3(signature = ( + sock_file: "str | None"=None, + server_info_file: "str | None"=None, + ) -> "_MapStreamAsyncServer")] + fn new(sock_file: Option, server_info_file: Option) -> Self { Self { - sock_file, - info_file, + sock_file: sock_file.unwrap_or_else(|| mapstream::SOCK_ADDR.to_string()), + server_info_file: server_info_file + .unwrap_or_else(|| mapstream::SERVER_INFO_FILE.to_string()), shutdown_tx: Mutex::new(None), } } /// Start the server with the given Python async generator function. - #[pyo3(signature = (py_func: "callable") -> "None")] - pub fn start<'a>(&self, py: Python<'a>, py_func: Py) -> PyResult> { + #[pyo3(signature = (handler: "callable") -> "None")] + pub fn start<'a>(&self, py: Python<'a>, handler: Py) -> PyResult> { let sock_file = self.sock_file.clone(); - let info_file = self.info_file.clone(); + let info_file = self.server_info_file.clone(); let (tx, rx) = tokio::sync::oneshot::channel::<()>(); { let mut guard = self.shutdown_tx.lock().unwrap(); *guard = Some(tx); } - pyo3_async_runtimes::tokio::future_into_py(py, async move { - crate::mapstream::server::start(py_func, sock_file, info_file, rx) - .await - .expect("server failed to start"); - Ok(()) - }) + pyo3_async_runtimes::tokio::future_into_py( + py, + crate::mapstream::server::start(handler, sock_file, info_file, rx), + ) + } + + /// Wait until the Numaflow IsReady probe succeeds over the mapstream UDS. + #[pyo3(signature = (timeout: "float"=30.0) -> "None")] + pub fn wait_ready<'a>(&self, py: Python<'a>, timeout: f64) -> PyResult> { + if !timeout.is_finite() || timeout < 0.0 { + return Err(pyo3::PyErr::new::( + "timeout must be a non-negative finite float", + )); + } + + let sock_file = self.sock_file.clone(); + let timeout = Duration::from_secs_f64(timeout); + + pyo3_async_runtimes::tokio::future_into_py( + py, + wait_for_ready(sock_file, timeout, "mapstream"), + ) } /// Trigger server shutdown from Python (idempotent). diff --git a/packages/pynumaflow-lite/src/mapstream/server.rs b/packages/pynumaflow-lite/src/mapstream/server.rs index 7dbcaee1..710277f5 100644 --- a/packages/pynumaflow-lite/src/mapstream/server.rs +++ b/packages/pynumaflow-lite/src/mapstream/server.rs @@ -1,53 +1,86 @@ -use crate::mapstream::Datum; -use crate::mapstream::Message as PyMessage; -use crate::pyiterables::PyAsyncIterStream; +use std::sync::{Arc, Mutex}; use numaflow::mapstream; use numaflow::shared::ServerExtras; - +use pyo3::exceptions::PyTypeError; use pyo3::prelude::*; -use std::sync::Arc; use tokio::sync::mpsc::Sender; use tokio_stream::StreamExt; +use crate::mapstream::Datum; +use crate::mapstream::Message as PyMessage; +use crate::pyiterables::PyAsyncIterStream; +use crate::pyrs::{combine_errors, format_error}; + pub(crate) struct PyMapStreamRunner { pub(crate) event_loop: Arc>, pub(crate) py_func: Arc>, + pub(crate) errors: Arc>>, +} + +impl PyMapStreamRunner { + fn fail(&self, error: PyErr) -> ! { + // numaflow calls mapstreamer concurrently, so each error belongs to a different + // message. Keep all of them. start() raises them together, which lets + // Python format every traceback instead of Rust printing them by hand. + let message = Python::attach(|py| format_error(py, &error)); + self.errors.lock().unwrap().push(error); + + // numaflow catches this panic, sends a gRPC error for this message, and + // starts the server shutdown. An empty result would instead look like a + // message that the handler dropped on purpose. + panic!("{message}"); + } } #[tonic::async_trait] impl mapstream::MapStreamer for PyMapStreamRunner { async fn map_stream(&self, input: mapstream::MapStreamRequest, tx: Sender) { - // Call Python handler: handler(keys, datum) -> AsyncIterator - let agen_obj = Python::attach(|py| { - let keys = input.keys.clone(); + // Call the Python handler: py_func(datum: Datum) -> AsyncIterator[Message] + let agen = match Python::attach(|py| -> PyResult<_> { let datum: Datum = input.into(); - let py_func = self.py_func.clone(); - let agen = py_func - .call1(py, (keys, datum)) - .expect("python handler raised before returning async iterable"); - // Keep as Py - agen.clone_ref(py).extract(py).unwrap_or(agen) - }); - - // Wrap the Python AsyncIterable in a Rust Stream that yields incrementally - let mut stream = PyAsyncIterStream::::new(agen_obj, self.event_loop.clone()) - .expect("failed to construct PyAsyncIterStream"); + self.py_func.call1(py, (datum,)) + }) { + Ok(agen) => agen, + Err(error) => self.fail(error), + }; + + // Wrap the Python AsyncIterable in a Rust Stream that yields incrementally. + // Items stay as Py so that a wrong item type gives a clear error below. + let mut stream = match PyAsyncIterStream::>::new(agen, self.event_loop.clone()) { + Ok(stream) => stream, + Err(_) => self.fail(PyErr::new::( + "mapstream handler must be an async generator (return AsyncIterator[Message])", + )), + }; // Forward each yielded message immediately to the sender while let Some(item) = stream.next().await { - match item { - Ok(py_msg) => { - let out: mapstream::Message = py_msg.into(); - if tx.send(out).await.is_err() { - break; - } - } - Err(e) => { - // Non-stop errors are surfaced per-item; log and stop this stream. - eprintln!("Python async iteration error: {:?}", e); - break; - } + let obj = match item { + Ok(obj) => obj, + Err(error) => self.fail(error), + }; + + let message: PyMessage = match Python::attach(|py| { + obj.extract(py).map_err(|_| { + let type_name = obj + .bind(py) + .get_type() + .name() + .map(|name| name.to_string_lossy().into_owned()) + .unwrap_or_else(|_| "".to_string()); + PyErr::new::(format!( + "mapstream handler must yield Message, got {type_name}" + )) + }) + }) { + Ok(message) => message, + Err(error) => self.fail(error), + }; + + // The receiver is gone, so the client does not want more messages. + if tx.send(message.into()).await.is_err() { + break; } } } @@ -61,14 +94,24 @@ pub(super) async fn start( shutdown_rx: tokio::sync::oneshot::Receiver<()>, ) -> Result<(), pyo3::PyErr> { let (tx, rx) = tokio::sync::oneshot::channel(); - let py_asyncio_loop_handle = tokio::task::spawn_blocking(move || crate::pyrs::run_asyncio(tx)); + let py_asyncio_loop_handle = tokio::task::spawn_blocking({ + println!( + "Starting MapStream UDF. socket={}, server_info={}", + sock_file, info_file + ); + move || crate::pyrs::run_asyncio(tx) + }); let event_loop = rx.await.unwrap(); - let (sig_handle, combined_rx) = crate::pyrs::setup_sig_handler(shutdown_rx); + let errors = Arc::new(Mutex::new(Vec::new())); + // Shutdown has two sources, and neither one needs a channel here. The Python + // side signals stop() through shutdown_rx. An uncaught Python error panics in + // fail(), and numaflow then shuts the server down on its own. let py_runner = PyMapStreamRunner { py_func: Arc::new(py_func), event_loop: event_loop.clone(), + errors: Arc::clone(&errors), }; let server = numaflow::mapstream::Server::new(py_runner) @@ -76,7 +119,7 @@ pub(super) async fn start( .with_server_info_file(info_file); let result = server - .start_with_shutdown(combined_rx) + .start_with_shutdown(shutdown_rx) .await .map_err(|e| pyo3::PyErr::new::(e.to_string())); @@ -87,15 +130,14 @@ pub(super) async fn start( } }); - println!("Numaflow Core (stream) has shutdown..."); + println!("Numaflow MapStream has shutdown..."); // Wait for the blocking asyncio thread to finish. let _ = py_asyncio_loop_handle.await; - // if not finished, abort it - if !sig_handle.is_finished() { - println!("Aborting signal handler"); - sig_handle.abort(); + let errors = std::mem::take(&mut *errors.lock().unwrap()); + if !errors.is_empty() { + return Err(Python::attach(|py| combine_errors(py, errors))); } result diff --git a/packages/pynumaflow-lite/tests/examples/mapstream_cat.py b/packages/pynumaflow-lite/tests/examples/mapstream_cat.py index 0fd0bcbb..299769f2 100644 --- a/packages/pynumaflow-lite/tests/examples/mapstream_cat.py +++ b/packages/pynumaflow-lite/tests/examples/mapstream_cat.py @@ -1,44 +1,28 @@ import asyncio -import signal -from collections.abc import AsyncIterator, Callable +from collections.abc import AsyncIterator -from pynumaflow_lite import mapstreamer -from pynumaflow_lite.mapstreamer import Message +from pynumaflow_lite.mapstreamer import Datum, MapStreamAsyncServer, Message -async def async_handler(keys: list[str], datum: mapstreamer.Datum) -> AsyncIterator[Message]: +async def async_handler(datum: Datum) -> AsyncIterator[Message]: """ A handler that splits the input datum value into multiple strings by `,` separator and emits them as a stream. """ - parts = datum.value.decode("utf-8").split(",") - if not parts: + if not datum.value: yield Message.to_drop() return - for s in parts: - yield Message(s.encode(), keys) + for s in datum.value.decode("utf-8").split(","): + yield Message(s.encode(), keys=datum.keys) -async def start(f: Callable[[list[str], mapstreamer.Datum], AsyncIterator[Message]]): - sock_file = "/tmp/var/run/numaflow/mapstream.sock" - server_info_file = "/tmp/var/run/numaflow/mapper-server-info" - server = mapstreamer.MapStreamAsyncServer(sock_file, server_info_file) - - # Register loop-level signal handlers to request graceful shutdown - loop = asyncio.get_running_loop() - try: - loop.add_signal_handler(signal.SIGINT, lambda: server.stop()) - loop.add_signal_handler(signal.SIGTERM, lambda: server.stop()) - except (NotImplementedError, RuntimeError): - pass - - try: - await server.start(f) - print("Shutting down gracefully...") - except asyncio.CancelledError: - server.stop() - return +async def main(): + await MapStreamAsyncServer( + handler=async_handler, + sock_file="/tmp/var/run/numaflow/mapstream.sock", + server_info_file="/tmp/var/run/numaflow/mapper-server-info", + ).serve() if __name__ == "__main__": - asyncio.run(start(async_handler)) + asyncio.run(main()) diff --git a/packages/pynumaflow-lite/tests/examples/mapstream_cat_class.py b/packages/pynumaflow-lite/tests/examples/mapstream_cat_class.py index c90ffc75..1a62caf4 100644 --- a/packages/pynumaflow-lite/tests/examples/mapstream_cat_class.py +++ b/packages/pynumaflow-lite/tests/examples/mapstream_cat_class.py @@ -1,42 +1,25 @@ import asyncio -import signal -from collections.abc import AsyncIterator, Callable +from collections.abc import AsyncIterator -from pynumaflow_lite import mapstreamer -from pynumaflow_lite.mapstreamer import Message +from pynumaflow_lite.mapstreamer import Datum, MapStreamAsyncServer, MapStreamer, Message -class SimpleStreamCat(mapstreamer.MapStreamer): - async def handler(self, keys: list[str], datum: mapstreamer.Datum) -> AsyncIterator[Message]: - parts = datum.value.decode("utf-8").split(",") - if not parts: +class SimpleStreamCat(MapStreamer): + async def handler(self, datum: Datum) -> AsyncIterator[Message]: + if not datum.value: yield Message.to_drop() return - for s in parts: - yield Message(s.encode(), keys) + for s in datum.value.decode("utf-8").split(","): + yield Message(s.encode(), datum.keys) -async def start(f: Callable[[list[str], mapstreamer.Datum], AsyncIterator[Message]]): - sock_file = "/tmp/var/run/numaflow/mapstream.sock" - server_info_file = "/tmp/var/run/numaflow/mapper-server-info" - server = mapstreamer.MapStreamAsyncServer(sock_file, server_info_file) - - # Register loop-level signal handlers so we control shutdown and avoid asyncio.run noise. - loop = asyncio.get_running_loop() - try: - loop.add_signal_handler(signal.SIGINT, lambda: server.stop()) - loop.add_signal_handler(signal.SIGTERM, lambda: server.stop()) - except (NotImplementedError, RuntimeError): - pass - - try: - await server.start(f) - print("Shutting down gracefully...") - except asyncio.CancelledError: - server.stop() - return +async def main(): + await MapStreamAsyncServer( + SimpleStreamCat(), + sock_file="/tmp/var/run/numaflow/mapstream.sock", + server_info_file="/tmp/var/run/numaflow/mapper-server-info", + ).serve() if __name__ == "__main__": - async_handler = SimpleStreamCat() - asyncio.run(start(async_handler)) + asyncio.run(main())