Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 7 additions & 13 deletions python/restate/ext/pydantic/_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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)
96 changes: 96 additions & 0 deletions tests/ext_pydantic_model.py
Original file line number Diff line number Diff line change
@@ -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
Loading