From 41e2458a9f81bc0512a733231f6ad66393ebe762 Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Thu, 24 Sep 2026 23:29:16 -0700 Subject: [PATCH] fix(pydantic): propagate SDK internal exceptions Signed-off-by: 1fanwang <1fannnw@gmail.com> --- python/restate/ext/pydantic/_model.py | 20 ++---- tests/ext_pydantic_model.py | 96 +++++++++++++++++++++++++++ 2 files changed, 103 insertions(+), 13 deletions(-) create mode 100644 tests/ext_pydantic_model.py diff --git a/python/restate/ext/pydantic/_model.py b/python/restate/ext/pydantic/_model.py index fd162da..4f16a08 100644 --- a/python/restate/ext/pydantic/_model.py +++ b/python/restate/ext/pydantic/_model.py @@ -4,7 +4,7 @@ from typing import Any import dataclasses -from restate import RunOptions, SdkInternalBaseException +from restate import RunOptions from restate.ext.pydantic._utils import current_state from restate.extensions import current_context from restate.ext.turnstile import Turnstile @@ -76,13 +76,10 @@ async def request(self, *args: Any, **kwargs: Any) -> ModelResponse: raise UserError( "A model cannot be used without a Restate context. Make sure to run it within an agent or a run context." ) - try: - res = await context.run_typed("Model call", self.wrapped.request, self._options, *args, **kwargs) - ids = [c.tool_call_id for c in res.tool_calls] - current_state().turnstile = Turnstile(ids) - return res - except SdkInternalBaseException as e: - raise Exception("Internal error during model call") from e + res = await context.run_typed("Model call", self.wrapped.request, self._options, *args, **kwargs) + ids = [c.tool_call_id for c in res.tool_calls] + current_state().turnstile = Turnstile(ids) + return res @asynccontextmanager async def request_stream( @@ -118,8 +115,5 @@ async def request_stream_run(): raise UserError( "A model cannot be used without a Restate context. Make sure to run it within an agent or a run context." ) - try: - response = await context.run_typed("Model stream call", request_stream_run, self._options) - yield RestateStreamedResponse(model_request_parameters, response) - except SdkInternalBaseException as e: - raise Exception("Internal error during model stream call") from e + response = await context.run_typed("Model stream call", request_stream_run, self._options) + yield RestateStreamedResponse(model_request_parameters, response) diff --git a/tests/ext_pydantic_model.py b/tests/ext_pydantic_model.py new file mode 100644 index 0000000..0b9d987 --- /dev/null +++ b/tests/ext_pydantic_model.py @@ -0,0 +1,96 @@ +import importlib +import sys +import types +from pathlib import Path +from typing import Any, cast + +import pytest +from pydantic_ai.messages import ModelMessage, ModelResponse +from pydantic_ai.models import Model, ModelRequestParameters +from pydantic_ai.settings import ModelSettings +from pydantic_ai.tools import RunContext + +from restate import RunOptions +from restate.context import RunAction +from restate.exceptions import SdkInternalException + + +pytestmark = [ + pytest.mark.anyio, +] + + +@pytest.fixture +def anyio_backend() -> str: + return "asyncio" + + +def _model_module() -> Any: + package_name = "restate.ext.pydantic" + if package_name not in sys.modules: + package = types.ModuleType(package_name) + package.__path__ = [str(Path(__file__).parents[1] / "python" / "restate" / "ext" / "pydantic")] + sys.modules[package_name] = package + return importlib.import_module("restate.ext.pydantic._model") + + +class FailingContext: + async def run_typed( + self, + name: str, + action: RunAction[Any], + opts: RunOptions[Any], + *args: Any, + **kwargs: Any, + ) -> Any: + raise SdkInternalException() + + +class DummyModel(Model): + @property + def system(self) -> str: + return "dummy" + + @property + def model_name(self) -> str: + return "dummy" + + async def request( + self, + messages: list[ModelMessage], + model_settings: ModelSettings | None, + model_request_parameters: ModelRequestParameters, + ) -> ModelResponse: + return ModelResponse(parts=[]) + + +def _model_request_parameters() -> ModelRequestParameters: + return ModelRequestParameters(function_tools=[], allow_text_output=True, output_mode="native", output_object=None) + + +async def test_request_propagates_sdk_internal_exception(monkeypatch: pytest.MonkeyPatch) -> None: + model_module = _model_module() + monkeypatch.setattr(model_module, "current_context", lambda: FailingContext()) + wrapper = model_module.RestateModelWrapper(DummyModel(), RunOptions()) + + with pytest.raises(SdkInternalException): + await wrapper.request([], None, _model_request_parameters()) + + +async def test_request_stream_propagates_sdk_internal_exception(monkeypatch: pytest.MonkeyPatch) -> None: + model_module = _model_module() + monkeypatch.setattr(model_module, "current_context", lambda: FailingContext()) + wrapper = model_module.RestateModelWrapper( + DummyModel(), + RunOptions(), + event_stream_handler=lambda run_context, streamed_response: None, + ) + + with pytest.raises(SdkInternalException): + async with wrapper.request_stream( + [], + None, + _model_request_parameters(), + cast(RunContext[Any], object()), + ): + pass