diff --git a/src/nexusrpc/_service.py b/src/nexusrpc/_service.py index 717430a..468ac83 100644 --- a/src/nexusrpc/_service.py +++ b/src/nexusrpc/_service.py @@ -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 @@ -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 diff --git a/src/nexusrpc/handler/_decorators.py b/src/nexusrpc/handler/_decorators.py index 8db2586..1ec9320 100644 --- a/src/nexusrpc/handler/_decorators.py +++ b/src/nexusrpc/handler/_decorators.py @@ -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, @@ -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: diff --git a/src/nexusrpc/handler/_operation_handler.py b/src/nexusrpc/handler/_operation_handler.py index 53ec8c3..b232120 100644 --- a/src/nexusrpc/handler/_operation_handler.py +++ b/src/nexusrpc/handler/_operation_handler.py @@ -1,5 +1,6 @@ from __future__ import annotations +import functools import inspect from abc import ABC, abstractmethod from collections.abc import Awaitable @@ -12,6 +13,8 @@ get_operation_factory, is_async_callable, is_callable, + set_operation, + set_operation_factory, ) from ._common import ( @@ -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]]], diff --git a/tests/test_operation_service.py b/tests/test_operation_service.py new file mode 100644 index 0000000..637d641 --- /dev/null +++ b/tests/test_operation_service.py @@ -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"