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
26 changes: 18 additions & 8 deletions src/nexusrpc/_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,17 @@ class MyNexusService:
method_name: Optional[str] = dataclasses.field(default=None, init=False)
input_type: Optional[type[InputT]] = dataclasses.field(default=None)
output_type: Optional[type[OutputT]] = dataclasses.field(default=None)
service: Optional[ServiceDefinition] = dataclasses.field(
default=None, init=False, compare=False, repr=False
)
"""The service that contains this operation.

:py:func:`nexusrpc.service` and :py:func:`nexusrpc.handler.service_handler`
set this value. Each decorated class has its own operations. Thus, when a child
service inherits an operation from a parent service, the operation reports the
child service. The value is None when the operation is not part of a decorated
class.
"""


@dataclass
Expand Down Expand Up @@ -133,20 +144,19 @@ def decorator(cls: type[ServiceT]) -> type[ServiceT]:
# In order for callers to refer to operation definitions at run-time, a decorated user
# service class must itself have a class attribute for every operation, even if
# declared only via a type annotation, and whether inherited from a parent class
# or not.
#
# TODO(preview): it is sufficient to do this setattr only for the subset of
# operations that were declared on *this* class. Currently however we are
# setting all inherited operations.
for op_name, op_defn in defn.operation_definitions.items():
if not hasattr(cls, op_name):
# or not. Each inherited operation gets a new Operation instance. Thus, each
# operation reports the service of the class that holds it.
for op_defn in defn.operation_definitions.values():
op = cls.__dict__.get(op_defn.method_name)
if not isinstance(op, Operation):
op = Operation(
name=op_defn.name,
input_type=op_defn.input_type,
output_type=op_defn.output_type,
)
op.method_name = op_defn.method_name
setattr(cls, op_name, op)
setattr(cls, op_defn.method_name, op)
op.service = defn

return cls

Expand Down
2 changes: 2 additions & 0 deletions src/nexusrpc/handler/_decorators.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
from ._operation_handler import (
OperationHandler,
SyncOperationHandler,
bind_operation_handler_methods_to_service,
collect_operation_handler_factories_by_method_name,
service_definition_from_operation_handler_methods,
validate_operation_handler_methods,
Expand Down Expand Up @@ -112,6 +113,7 @@ def decorator(cls: type[ServiceHandlerT]) -> type[ServiceHandlerT]:
)
validate_operation_handler_methods(cls, factories_by_method_name, service)
set_service_definition(cls, service)
bind_operation_handler_methods_to_service(cls, service)
return cls

if cls is None:
Expand Down
70 changes: 70 additions & 0 deletions src/nexusrpc/handler/_operation_handler.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import functools
import inspect
from abc import ABC, abstractmethod
from collections.abc import Awaitable
Expand All @@ -12,6 +13,8 @@
get_operation_factory,
is_async_callable,
is_callable,
set_operation,
set_operation_factory,
)

from ._common import (
Expand Down Expand Up @@ -236,6 +239,73 @@ def validate_operation_handler_methods(
)


def bind_operation_handler_methods_to_service(
service_cls: type[ServiceHandlerT],
service_definition: ServiceDefinition,
) -> None:
"""Give each operation handler method of a service handler class its own operation.

The operation of each method reports ``service_definition`` as its service. The
operation gets its name and types from the service definition. When
``service_cls`` inherits a method from a parent class, this function replaces the
method on ``service_cls`` with a wrapper. Thus, the method of the parent class
keeps its own operation.

Validate the methods against the service definition before you call this
function.
"""
for op_defn in service_definition.operation_definitions.values():
method = getattr(service_cls, op_defn.method_name)
# A start method with a decorator such as @sync_operation holds a factory. A
# method with the @operation_handler decorator is the factory.
factory = getattr(method, "__nexus_operation_factory__", None)

op: Operation[Any, Any] = Operation(
name=op_defn.name,
input_type=op_defn.input_type,
output_type=op_defn.output_type,
)
op.method_name = op_defn.method_name
op.service = service_definition

if op_defn.method_name not in service_cls.__dict__:
method = _wrap_inherited_method(method)
setattr(service_cls, op_defn.method_name, method)

if factory is None:
set_operation(method, op)
else:
set_operation_factory(method, _wrap_factory(factory, op))


def _wrap_factory(
factory: Callable[[Any], OperationHandler[Any, Any]],
op: Operation[Any, Any],
) -> Callable[[Any], OperationHandler[Any, Any]]:
@functools.wraps(factory)
def wrapper(self: Any) -> OperationHandler[Any, Any]:
return factory(self)

set_operation(wrapper, op)
return wrapper


def _wrap_inherited_method(method: Callable[..., Any]) -> Callable[..., Any]:
if inspect.iscoroutinefunction(method):

@functools.wraps(method)
async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
return await method(*args, **kwargs)

return async_wrapper

@functools.wraps(method)
def wrapper(*args: Any, **kwargs: Any) -> Any:
return method(*args, **kwargs)

return wrapper


def service_definition_from_operation_handler_methods(
service_name: str,
user_methods: dict[str, Callable[[ServiceHandlerT], OperationHandler[Any, Any]]],
Expand Down
188 changes: 188 additions & 0 deletions tests/test_operation_service.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,188 @@
"""Tests that operations report the service that they belong to."""

from __future__ import annotations

import inspect
from typing import Any

import pytest

import nexusrpc
from nexusrpc import LazyValue, Operation
from nexusrpc.handler import (
Handler,
OperationHandler,
StartOperationContext,
StartOperationResultSync,
operation_handler,
service_handler,
sync_operation,
)
from nexusrpc.handler._operation_handler import SyncOperationHandler
from tests.helpers import DummySerializer, TestOperationTaskCancellation


def _service_name(obj: Any) -> str | None:
op = obj if isinstance(obj, Operation) else nexusrpc.get_operation(obj)
assert op is not None
return op.service.name if op.service else None


def _operation(obj: Any) -> Operation[Any, Any]:
op = nexusrpc.get_operation(obj)
assert op is not None
return op


@nexusrpc.service(name="base-definition")
class BaseDefinition:
annotated: Operation[int, int]
renamed: Operation[int, int] = Operation(name="Renamed")


@nexusrpc.service(name="sub-definition")
class SubDefinition(BaseDefinition):
added: Operation[int, int]


class UndecoratedSubDefinition(BaseDefinition):
pass


def test_definition_operations_report_their_service():
assert _service_name(BaseDefinition.annotated) == "base-definition"
assert _service_name(BaseDefinition.renamed) == "base-definition"


def test_inherited_definition_operations_report_the_child_service():
assert _service_name(SubDefinition.annotated) == "sub-definition"
assert _service_name(SubDefinition.renamed) == "sub-definition"
assert _service_name(SubDefinition.added) == "sub-definition"
assert SubDefinition.renamed.name == "Renamed"

# The parent keeps its own operations
assert SubDefinition.annotated is not BaseDefinition.annotated
assert _service_name(BaseDefinition.annotated) == "base-definition"
# Equality ignores the service
assert SubDefinition.annotated == BaseDefinition.annotated


def test_inherited_renamed_operation_is_set_under_its_method_name():
assert "renamed" in SubDefinition.__dict__
assert "Renamed" not in SubDefinition.__dict__


def test_undecorated_subclass_reports_the_parent_service():
assert _service_name(UndecoratedSubDefinition.annotated) == "base-definition"


def test_operation_outside_a_service_has_no_service():
assert Operation[int, int](name="free").service is None


@service_handler(service=BaseDefinition)
class BaseHandler:
@sync_operation
async def annotated(self, _ctx: StartOperationContext, input: int) -> int:
return input + 1

@sync_operation
async def renamed(self, _ctx: StartOperationContext, input: int) -> int:
return input + 2


@service_handler(name="sub-handler")
class SubHandler(BaseHandler):
pass


def test_handler_methods_report_their_service():
assert _service_name(BaseHandler.annotated) == "base-definition"
assert _service_name(BaseHandler.renamed) == "base-definition"


def test_handler_method_operations_use_the_definition_name_and_types():
op = _operation(BaseHandler.renamed)
assert op.name == "Renamed"
assert op.method_name == "renamed"
assert (op.input_type, op.output_type) == (int, int)


def test_inherited_handler_methods_report_the_child_service():
assert _service_name(SubHandler.annotated) == "sub-handler"
assert _service_name(SubHandler.renamed) == "sub-handler"
assert _operation(SubHandler.renamed).name == "Renamed"

# The parent keeps its own methods and operations
assert SubHandler.annotated is not BaseHandler.annotated
assert _service_name(BaseHandler.annotated) == "base-definition"
assert inspect.iscoroutinefunction(SubHandler.annotated)


@pytest.mark.asyncio
async def test_inherited_handler_methods_still_handle_requests():
assert await SubHandler().annotated(None, 1) == 2 # type: ignore[arg-type]

handler = Handler(user_service_handlers=[SubHandler()])
result = await handler.start_operation(
StartOperationContext(
service="sub-handler",
operation="Renamed",
headers={},
request_id="request_id",
task_cancellation=TestOperationTaskCancellation(),
),
LazyValue(serializer=DummySerializer(1), headers={}),
)
assert result == StartOperationResultSync(3)


class UndecoratedSyncBase:
@sync_operation
def sync_op(self, _ctx: StartOperationContext, input: int) -> int:
return input


@service_handler(name="sync-handler")
class SyncHandler(UndecoratedSyncBase):
pass


def test_inherited_sync_handler_method_stays_sync():
assert _service_name(SyncHandler.sync_op) == "sync-handler"
assert not inspect.iscoroutinefunction(SyncHandler.sync_op)
assert SyncHandler().sync_op(None, 5) == 5 # type: ignore[arg-type]
# The undecorated parent is not a service
assert _service_name(UndecoratedSyncBase.sync_op) is None


@service_handler(name="factory-handler")
class FactoryHandler:
@operation_handler(name="Factory")
def factory(self) -> OperationHandler[int, int]:
async def start(_ctx: StartOperationContext, input: int) -> int:
return input

return SyncOperationHandler(start)


@service_handler(name="sub-factory-handler")
class SubFactoryHandler(FactoryHandler):
pass


def test_operation_handler_factories_report_their_service():
assert _service_name(FactoryHandler.factory) == "factory-handler"
assert _service_name(SubFactoryHandler.factory) == "sub-factory-handler"
assert _operation(SubFactoryHandler.factory).name == "Factory"
assert isinstance(SubFactoryHandler().factory(), SyncOperationHandler)


def test_local_handler_class_reports_its_service():
@service_handler(name="local-handler")
class LocalHandler:
@sync_operation
async def op(self, _ctx: StartOperationContext, input: int) -> int:
return input

assert _service_name(LocalHandler.op) == "local-handler"
Loading