From 66293f6c437a334130d56380625386af10c745b9 Mon Sep 17 00:00:00 2001 From: Sandipan Kundu Date: Wed, 26 Aug 2026 17:04:53 +0530 Subject: [PATCH 1/2] feat: add gRPC serving for @bentoml.service() --- docs/source/build-with-bentoml/services.rst | 8 + noxfile.py | 8 +- src/_bentoml_impl/grpc/__init__.py | 6 + src/_bentoml_impl/grpc/codec.py | 226 +++++++++++++ src/_bentoml_impl/server/__init__.py | 3 +- src/_bentoml_impl/server/grpc/__init__.py | 0 .../server/grpc/servicer/__init__.py | 0 .../server/grpc/servicer/v1/__init__.py | 187 +++++++++++ src/_bentoml_impl/server/grpc_app.py | 305 ++++++++++++++++++ src/_bentoml_impl/server/serving.py | 277 ++++++++++++++++ src/_bentoml_impl/worker/grpc_service.py | 157 +++++++++ src/bentoml/bentos.py | 82 +++-- .../grpc/interceptors/opentelemetry.py | 5 + src/bentoml/grpc/v1/service_pb2.py | 2 +- src/bentoml/grpc/v1/service_pb2_grpc.py | 2 +- src/bentoml/grpc/v1alpha1/service_pb2.py | 2 +- src/bentoml/grpc/v1alpha1/service_pb2_grpc.py | 2 +- src/bentoml_cli/serve.py | 73 +++-- tests/e2e/bento_new_sdk_grpc/bentofile.yaml | 3 + tests/e2e/bento_new_sdk_grpc/conftest.py | 10 + tests/e2e/bento_new_sdk_grpc/requirements.txt | 1 + tests/e2e/bento_new_sdk_grpc/service.py | 26 ++ tests/e2e/bento_new_sdk_grpc/test_grpc.py | 92 ++++++ tests/unit/_bentoml_impl/__init__.py | 0 tests/unit/_bentoml_impl/grpc/__init__.py | 0 tests/unit/_bentoml_impl/grpc/test_codec.py | 150 +++++++++ .../unit/_bentoml_impl/grpc/test_servicer.py | 145 +++++++++ 27 files changed, 1718 insertions(+), 54 deletions(-) create mode 100644 src/_bentoml_impl/grpc/__init__.py create mode 100644 src/_bentoml_impl/grpc/codec.py create mode 100644 src/_bentoml_impl/server/grpc/__init__.py create mode 100644 src/_bentoml_impl/server/grpc/servicer/__init__.py create mode 100644 src/_bentoml_impl/server/grpc/servicer/v1/__init__.py create mode 100644 src/_bentoml_impl/server/grpc_app.py create mode 100644 src/_bentoml_impl/worker/grpc_service.py create mode 100644 tests/e2e/bento_new_sdk_grpc/bentofile.yaml create mode 100644 tests/e2e/bento_new_sdk_grpc/conftest.py create mode 100644 tests/e2e/bento_new_sdk_grpc/requirements.txt create mode 100644 tests/e2e/bento_new_sdk_grpc/service.py create mode 100644 tests/e2e/bento_new_sdk_grpc/test_grpc.py create mode 100644 tests/unit/_bentoml_impl/__init__.py create mode 100644 tests/unit/_bentoml_impl/grpc/__init__.py create mode 100644 tests/unit/_bentoml_impl/grpc/test_codec.py create mode 100644 tests/unit/_bentoml_impl/grpc/test_servicer.py diff --git a/docs/source/build-with-bentoml/services.rst b/docs/source/build-with-bentoml/services.rst index c39f941b42a..5838072371c 100644 --- a/docs/source/build-with-bentoml/services.rst +++ b/docs/source/build-with-bentoml/services.rst @@ -60,6 +60,14 @@ By default, the server is accessible at `http://localhost:3000/ dict[str, t.Any]: + ref = schema.get("$ref") + if not ref: + return schema + if not ref.startswith("#/$defs/"): + return schema + name = ref.rsplit("/", 1)[-1] + resolved = (root.get("$defs") or {}).get(name) + return resolved if isinstance(resolved, dict) else schema + + +def _field_from_schema(schema: dict[str, t.Any]) -> str: + type_ = schema.get("type") + fmt = schema.get("format") + if type_ == "tensor": + return "ndarray" + if type_ == "file": + return "file" + if type_ == "string" and fmt in {"binary", "byte"}: + return "file" + if type_ == "string": + return "text" + if type_ == "dataframe": + raise InvalidArgument( + "pandas DataFrame IO is not supported over gRPC for @bentoml.service() yet" + ) + return "json" + + +def proto_binding(spec: type[IODescriptor]) -> ProtoBinding: + """Infer the v1 proto oneof field for an IO spec.""" + schema = spec.model_json_schema() + if issubclass(spec, IORootModel): + return ProtoBinding(_field_from_schema(schema), None) + + props = schema.get("properties") or {} + if schema.get("type") == "object" and len(props) == 1: + key, child = next(iter(props.items())) + child = _resolve_ref(child, schema) + field = _field_from_schema(child) + if field != "json": + return ProtoBinding(field, key) + return ProtoBinding(_field_from_schema(schema), None) + + +def _unwrap_value(spec: type[IODescriptor], binding: ProtoBinding, obj: t.Any) -> t.Any: + if issubclass(spec, IORootModel) and isinstance(obj, IORootModel): + return obj.root + if binding.unwrap_key is None: + if hasattr(obj, "model_dump") and not isinstance( + obj, (bytes, str, Path, dict, list) + ): + dump = obj.model_dump() + if isinstance(dump, dict): + return dump + return obj + key = binding.unwrap_key + if isinstance(obj, dict) and key in obj: + return obj[key] + if hasattr(obj, key): + return getattr(obj, key) + return obj + + +def _wrap_value( + spec: type[IODescriptor], binding: ProtoBinding, value: t.Any +) -> IODescriptor: + if issubclass(spec, IORootModel): + return spec.from_inputs(value) + if binding.unwrap_key is not None: + return spec.model_validate({binding.unwrap_key: value}) + if isinstance(value, spec): + return value + if isinstance(value, dict): + return spec.model_validate(value) + return spec.from_inputs(value) + + +async def _ndarray_from_proto(field: t.Any) -> t.Any: + from bentoml._internal.io_descriptors.numpy import NumpyNdarray + + return await NumpyNdarray().from_proto(field) + + +async def _ndarray_to_proto(obj: t.Any) -> Message: + from bentoml._internal.io_descriptors.numpy import NumpyNdarray + + return await NumpyNdarray().to_proto(obj) + + +def _file_from_proto(field: t.Any) -> bytes: + if isinstance(field, bytes): + return field + content = getattr(field, "content", None) + if content is None: + raise InvalidArgument("File proto is missing content") + return bytes(content) + + +def _bytes_to_path(body: bytes) -> Path: + # FileSchema.decode() looks up the HTTP request temp dir; gRPC has none, + # so persist bytes to a local tempfile the validator can treat as a Path. + fd, name = tempfile.mkstemp(prefix="bentoml-grpc-") + with os.fdopen(fd, "wb") as handle: + handle.write(body) + return Path(name) + + +def _file_to_proto(obj: t.Any) -> pb.File: + if isinstance(obj, bytes): + body = obj + kind = "application/octet-stream" + elif isinstance(obj, (str, Path)): + path = Path(obj) + body = path.read_bytes() + kind = "application/octet-stream" + elif hasattr(obj, "read"): + body = obj.read() + if isinstance(body, str): + body = body.encode("utf-8") + kind = "application/octet-stream" + else: + raise InvalidArgument(f"Cannot encode {type(obj)!r} as a gRPC file") + return pb.File(kind=kind, content=body) + + +def _json_from_proto(field: t.Any) -> t.Any: + if isinstance(field, bytes): + import json + + return json.loads(field) + return json_format.MessageToDict(field, preserving_proto_field_name=True) + + +def _json_to_proto(obj: t.Any) -> struct_pb2.Value: + msg = struct_pb2.Value() + if obj is None: + return msg + if hasattr(obj, "model_dump"): + obj = obj.model_dump(mode="json") + json_format.ParseDict(obj, msg) + return msg + + +async def decode_proto( + spec: type[IODescriptor], field_name: str | None, value: t.Any +) -> IODescriptor: + """Decode a Request/Response oneof value into an IODescriptor instance.""" + if field_name == "serialized_bytes": + raise InvalidArgument( + "serialized_bytes / pickle payloads are not supported over gRPC " + "for @bentoml.service()" + ) + binding = proto_binding(spec) + if field_name is None: + raise InvalidArgument("gRPC request is missing a content field") + if field_name not in _SUPPORTED_FIELDS: + raise InvalidArgument( + f"Unsupported gRPC content field {field_name!r}; " + f"accepted fields: {', '.join(_SUPPORTED_FIELDS)}" + ) + if field_name != binding.field: + raise InvalidArgument( + f"{spec.__name__} expects gRPC field {binding.field!r}, got {field_name!r}" + ) + + if binding.field == "text": + if isinstance(value, bytes): + decoded: t.Any = value.decode("utf-8") + else: + decoded = value.value if hasattr(value, "value") else str(value) + elif binding.field == "json": + decoded = _json_from_proto(value) + elif binding.field == "ndarray": + decoded = await _ndarray_from_proto(value) + else: + decoded = _bytes_to_path(_file_from_proto(value)) + return _wrap_value(spec, binding, decoded) + + +async def encode_proto(spec: type[IODescriptor], obj: t.Any) -> tuple[str, t.Any]: + """Encode a Python value to (oneof field name, proto message).""" + binding = proto_binding(spec) + value = _unwrap_value(spec, binding, obj) + if binding.field == "text": + return binding.field, wrappers_pb2.StringValue(value=str(value)) + if binding.field == "json": + return binding.field, _json_to_proto(value) + if binding.field == "ndarray": + return binding.field, await _ndarray_to_proto(value) + return binding.field, _file_to_proto(value) diff --git a/src/_bentoml_impl/server/__init__.py b/src/_bentoml_impl/server/__init__.py index d81385cbb5b..1b59d339321 100644 --- a/src/_bentoml_impl/server/__init__.py +++ b/src/_bentoml_impl/server/__init__.py @@ -6,6 +6,7 @@ This will be eventually migrated to Rust. """ +from .serving import serve_grpc from .serving import serve_http -__all__ = ["serve_http"] +__all__ = ["serve_http", "serve_grpc"] diff --git a/src/_bentoml_impl/server/grpc/__init__.py b/src/_bentoml_impl/server/grpc/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/src/_bentoml_impl/server/grpc/servicer/__init__.py b/src/_bentoml_impl/server/grpc/servicer/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/src/_bentoml_impl/server/grpc/servicer/v1/__init__.py b/src/_bentoml_impl/server/grpc/servicer/v1/__init__.py new file mode 100644 index 00000000000..2d00ceedfa3 --- /dev/null +++ b/src/_bentoml_impl/server/grpc/servicer/v1/__init__.py @@ -0,0 +1,187 @@ +from __future__ import annotations + +import functools +import inspect +import logging +import sys +import typing as t +from typing import TYPE_CHECKING + +import anyio + +from _bentoml_impl.grpc.codec import decode_proto +from _bentoml_impl.grpc.codec import encode_proto +from _bentoml_sdk.io_models import ARGS +from _bentoml_sdk.io_models import KWARGS +from _bentoml_sdk.io_models import IORootModel +from bentoml._internal.utils import get_original_func +from bentoml._internal.utils import is_async_callable +from bentoml._internal.utils.lazy_loader import LazyLoader +from bentoml.exceptions import BentoMLException +from bentoml.exceptions import InvalidArgument +from bentoml.grpc.utils import grpc_status_code +from bentoml.grpc.utils import import_generated_stubs +from bentoml.grpc.utils import import_grpc + +logger = logging.getLogger(__name__) + +if TYPE_CHECKING: + from logging import _ExcInfoType as ExcInfoType # type: ignore (private warning) + + import grpc + from google.protobuf import struct_pb2 + from grpc import aio + + from _bentoml_sdk import Service + from bentoml.grpc.types import BentoServicerContext + from bentoml.grpc.v1 import service_pb2 as pb + from bentoml.grpc.v1 import service_pb2_grpc as services +else: + grpc, aio = import_grpc() + pb, services = import_generated_stubs(version="v1") + struct_pb2 = LazyLoader("struct_pb2", globals(), "google.protobuf.struct_pb2") + + +def log_exception(request: pb.Request, exc_info: ExcInfoType) -> None: + logger.error("Exception on /%s [POST]", request.api_name, exc_info=exc_info) + + +def _call_args_from_input(method: t.Any, input_data: t.Any, ctx: t.Any) -> tuple[ + tuple[t.Any, ...], dict[str, t.Any] +]: + call_args: tuple[t.Any, ...] = () + call_kwargs: dict[str, t.Any] = {} + if getattr(method.input_spec, "__root_input__", False): + if isinstance(input_data, IORootModel): + call_args = (input_data.root,) + else: + call_args = (input_data,) + else: + call_kwargs = {k: getattr(input_data, k) for k in input_data.model_fields} + if method.ctx_param is not None: + call_kwargs[method.ctx_param] = ctx + if ARGS in call_kwargs: + call_args = (*call_args, call_kwargs.pop(ARGS)) + if KWARGS in call_kwargs: + call_kwargs.update(call_kwargs.pop(KWARGS)) + return call_args, call_kwargs + + +def create_bento_servicer(service: Service[t.Any]) -> services.BentoServiceServicer: + """Create a v1 BentoService servicer for a new-style ``@bentoml.service()``.""" + + class BentoServiceImpl(services.BentoServiceServicer): + instance: t.Any = None + + def set_instance(self, instance: t.Any) -> None: + self.instance = instance + + def _get_instance(self) -> t.Any: + if self.instance is None: + self.instance = service() + return self.instance + + async def Call( # type: ignore (no async types) + self, + request: pb.Request, + context: BentoServicerContext, + ) -> pb.Response | None: + response = pb.Response() + try: + if request.api_name not in service.apis: + raise InvalidArgument( + f"given 'api_name' is not defined in {service.name}", + ) from None + + method = service.apis[request.api_name] + if method.is_stream: + await context.abort( + code=grpc.StatusCode.UNIMPLEMENTED, + details=( + f"API {method.name!r} is a streaming endpoint; " + "gRPC streaming is not supported for @bentoml.service() yet" + ), + ) + return None + if method.batchable: + await context.abort( + code=grpc.StatusCode.UNIMPLEMENTED, + details=( + f"API {method.name!r} is batchable; " + "adaptive batching is not supported over gRPC for @bentoml.service() yet" + ), + ) + return None + if method.is_task: + await context.abort( + code=grpc.StatusCode.UNIMPLEMENTED, + details=( + f"API {method.name!r} is a task endpoint; " + "tasks are not supported over gRPC for @bentoml.service() yet" + ), + ) + return None + + field = request.WhichOneof("content") + input_data = await decode_proto( + method.input_spec, field, getattr(request, field) if field else None + ) + call_args, call_kwargs = _call_args_from_input( + method, input_data, service.context + ) + func = getattr(self._get_instance(), method.name).local + original_func = get_original_func(func) + if is_async_callable(original_func) or inspect.iscoroutinefunction( + original_func + ): + output = await func(*call_args, **call_kwargs) + else: + output = await anyio.to_thread.run_sync( + functools.partial(func, *call_args, **call_kwargs) + ) + + field_name, encoded = await encode_proto(method.output_spec, output) + response = pb.Response(**{field_name: encoded}) + except BentoMLException as e: + log_exception(request, sys.exc_info()) + await context.abort(code=grpc_status_code(e), details=e.message) + except aio.AbortError: + raise + except Exception: # pylint: disable=broad-except + log_exception(request, sys.exc_info()) + await context.abort( + code=grpc.StatusCode.INTERNAL, + details="An error has occurred in BentoML user code when handling this request, find the error details in server logs.", + ) + return response + + async def ServiceMetadata( # type: ignore (no async types) + self, + request: pb.ServiceMetadataRequest, # pylint: disable=unused-argument + context: BentoServicerContext, # pylint: disable=unused-argument + ) -> pb.ServiceMetadataResponse: + from google.protobuf.json_format import ParseDict + + def _schema_metadata(schema: dict[str, t.Any]) -> t.Any: + attributes = struct_pb2.Struct() + ParseDict(schema, attributes) + return pb.ServiceMetadataResponse.DescriptorMetadata( + descriptor_id="bentoml.sdk.IODescriptor", + attributes=attributes, + ) + + return pb.ServiceMetadataResponse( + name=service.name, + docs=service.description or "", + apis=[ + pb.ServiceMetadataResponse.InferenceAPI( + name=api.name, + docs=api.doc or "", + input=_schema_metadata(api.schema().get("input") or {}), + output=_schema_metadata(api.schema().get("output") or {}), + ) + for api in service.apis.values() + ], + ) + + return BentoServiceImpl() diff --git a/src/_bentoml_impl/server/grpc_app.py b/src/_bentoml_impl/server/grpc_app.py new file mode 100644 index 00000000000..9553622f908 --- /dev/null +++ b/src/_bentoml_impl/server/grpc_app.py @@ -0,0 +1,305 @@ +from __future__ import annotations + +import asyncio +import inspect +import logging +import os +import sys +import typing as t +from concurrent.futures import ThreadPoolExecutor +from functools import cached_property +from typing import TYPE_CHECKING + +from simple_di import Provide +from simple_di import inject + +from _bentoml_impl.server.grpc.servicer.v1 import create_bento_servicer +from _bentoml_sdk import Service +from _bentoml_sdk.service import set_current_service +from bentoml._internal.configuration.containers import BentoMLContainer +from bentoml._internal.context import ServiceContext as Context +from bentoml._internal.utils.lazy_loader import LazyLoader +from bentoml.exceptions import BentoMLException +from bentoml.exceptions import MissingDependencyException +from bentoml.grpc.utils import LATEST_PROTOCOL_VERSION +from bentoml.grpc.utils import import_generated_stubs +from bentoml.grpc.utils import import_grpc +from bentoml.grpc.utils import load_from_file + +logger = logging.getLogger(__name__) + +if TYPE_CHECKING: + import grpc + from grpc import aio + from grpc_health.v1 import health + from grpc_health.v1 import health_pb2 as pb_health + from grpc_health.v1 import health_pb2_grpc as services_health + + from bentoml.grpc.types import Interceptors +else: + grpc, aio = import_grpc() + health_exception_msg = "'grpcio-health-checking' is required for using health checking endpoints. Install with 'pip install grpcio-health-checking'." + pb_health = LazyLoader( + "pb_health", + globals(), + "grpc_health.v1.health_pb2", + exc_msg=health_exception_msg, + ) + services_health = LazyLoader( + "services_health", + globals(), + "grpc_health.v1.health_pb2_grpc", + exc_msg=health_exception_msg, + ) + health = LazyLoader( + "health", + globals(), + "grpc_health.v1.health", + exc_msg=health_exception_msg, + ) + + +class Server(aio._server.Server): + """Async gRPC server for ``@bentoml.service()`` services.""" + + @inject + def __init__( + self, + bento_service: Service[t.Any], + bind_address: str, + max_message_length: int | None = Provide[ + BentoMLContainer.grpc.max_message_length + ], + maximum_concurrent_rpcs: int | None = Provide[ + BentoMLContainer.grpc.maximum_concurrent_rpcs + ], + enable_reflection: bool = False, + enable_channelz: bool = False, + max_concurrent_streams: int | None = None, + migration_thread_pool_workers: int = 1, + ssl_certfile: str | None = None, + ssl_keyfile: str | None = None, + ssl_ca_certs: str | None = None, + graceful_shutdown_timeout: float | None = None, + compression: grpc.Compression | None = None, + protocol_version: str = LATEST_PROTOCOL_VERSION, + ): + if protocol_version != "v1": + raise BentoMLException( + f"@bentoml.service() gRPC serving only supports protocol v1, got {protocol_version!r}" + ) + + try: + asyncio.get_running_loop() + except RuntimeError: + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + + pb, _ = import_generated_stubs("v1") + + self.bento_service = bento_service + self.servicer = create_bento_servicer(bento_service) + self._service_instance: t.Any | None = None + + self.max_message_length = max_message_length + self.max_concurrent_streams = max_concurrent_streams + self.bind_address = bind_address + self.enable_reflection = enable_reflection + self.enable_channelz = enable_channelz + self.graceful_shutdown_timeout = graceful_shutdown_timeout + self.ssl_certfile = ssl_certfile + self.ssl_keyfile = ssl_keyfile + self.ssl_ca_certs = ssl_ca_certs + self.protocol_version = protocol_version + + self.health_servicer = health.aio.HealthServicer() + self.service_names = tuple( + service.full_name for service in pb.DESCRIPTOR.services_by_name.values() + ) + (health.SERVICE_NAME,) + + super().__init__( + thread_pool=ThreadPoolExecutor(max_workers=migration_thread_pool_workers), + generic_handlers=() if self.handlers is None else self.handlers, + interceptors=list(map(lambda x: x(), self.interceptors)), + options=self.options, + maximum_concurrent_rpcs=maximum_concurrent_rpcs, + compression=compression, + ) + + @property + def options(self) -> grpc.aio.ChannelArgumentType: + options: grpc.aio.ChannelArgumentType = [] + if sys.platform != "win32": + options.append(("grpc.so_reuseport", 1)) + if self.max_concurrent_streams: + options.append(("grpc.max_concurrent_streams", self.max_concurrent_streams)) + if self.enable_channelz: + options.append(("grpc.enable_channelz", 1)) + if self.max_message_length: + options.extend( + ( + ("grpc.max_message_length", self.max_message_length), + ("grpc.max_receive_message_length", self.max_message_length), + ("grpc.max_send_message_length", self.max_message_length), + ) + ) + return tuple(options) + + @property + def interceptors(self) -> Interceptors: + from bentoml.grpc.interceptors.opentelemetry import ( + AsyncOpenTelemetryServerInterceptor, + ) + + interceptors: Interceptors = [AsyncOpenTelemetryServerInterceptor] + if BentoMLContainer.api_server_config.metrics.enabled.get(): + from bentoml.grpc.interceptors.prometheus import PrometheusServerInterceptor + + interceptors.append(PrometheusServerInterceptor) + if BentoMLContainer.api_server_config.logging.access.enabled.get(): + from bentoml.grpc.interceptors.access import AccessLogServerInterceptor + + access_logger = logging.getLogger("bentoml.access") + if access_logger.getEffectiveLevel() <= logging.INFO: + interceptors.append(AccessLogServerInterceptor) + return interceptors + + @property + def handlers(self) -> t.Sequence[grpc.GenericRpcHandler] | None: + return None + + @cached_property + def loop(self) -> asyncio.AbstractEventLoop: + try: + return asyncio.get_running_loop() + except RuntimeError: + return asyncio.get_event_loop() + + def run(self) -> None: + try: + self.loop.run_until_complete(self.serve()) + finally: + try: + self.loop.call_soon_threadsafe( + lambda: asyncio.ensure_future(self.shutdown()) + ) + except Exception as e: # pylint: disable=broad-except + raise RuntimeError(f"Server failed unexpectedly: {e}") from None + + @cached_property + def context(self) -> Context: + return self.bento_service.context + + def configure_port(self, addr: str) -> None: + if self.ssl_certfile: + client_auth = False + ca_cert = None + assert self.ssl_keyfile, ( + "'ssl_keyfile' is required when 'ssl_certfile' is provided." + ) + if self.ssl_ca_certs is not None: + client_auth = True + ca_cert = load_from_file(self.ssl_ca_certs) + server_credentials = grpc.ssl_server_credentials( + ( + ( + load_from_file(self.ssl_keyfile), + load_from_file(self.ssl_certfile), + ), + ), + root_certificates=ca_cert, + require_client_auth=client_auth, + ) + self.add_secure_port(addr, server_credentials) + else: + self.add_insecure_port(addr) + + async def serve(self) -> None: + self.configure_port(self.bind_address) + await self.startup() + await self.wait_for_termination() + + async def _create_instance(self) -> None: + self._service_instance = self.bento_service() + logger.info("Service %s initialized", self.bento_service.name) + for name in dir(self.bento_service.inner): + member = getattr(self.bento_service.inner, name) + if ( + not name.startswith("__") + and callable(member) + and getattr(member, "__bentoml_startup_hook__", False) + ): + logger.info("Running startup hook: %s", name) + result = getattr(self._service_instance, name)() + if inspect.isawaitable(result): + await result + set_current_service(self._service_instance) + self.servicer.set_instance(self._service_instance) + await asyncio.gather( + *( + real.__aenter__() + for dep_name in self.bento_service.dependencies + if hasattr( + (real := getattr(self._service_instance, dep_name)), "__aenter__" + ) + ) + ) + + async def startup(self) -> None: + _, services = import_generated_stubs("v1") + + await self._create_instance() + + services.add_BentoServiceServicer_to_server(self.servicer, self) + services_health.add_HealthServicer_to_server(self.health_servicer, self) + + service_names = self.service_names + if self.enable_channelz: + try: + from grpc_channelz.v1 import channelz + except ImportError: + raise MissingDependencyException( + "'--debug' is passed, which requires 'grpcio-channelz' to be installed. Install with 'pip install bentoml[grpc-channelz]'." + ) from None + if "GRPC_TRACE" not in os.environ: + logger.debug( + "channelz is enabled, while GRPC_TRACE is not set. No channel tracing will be recorded." + ) + channelz.add_channelz_servicer(self) + if self.enable_reflection: + try: + from grpc_reflection.v1alpha import reflection + except ImportError: + raise MissingDependencyException( + "reflection is enabled, which requires 'grpcio-reflection' to be installed. Install with 'pip install bentoml[grpc-reflection]'." + ) from None + service_names += (reflection.SERVICE_NAME,) + reflection.enable_server_reflection(service_names, self) + for service in service_names: + await self.health_servicer.set( + service, + pb_health.HealthCheckResponse.SERVING, # type: ignore (no types available) + ) + await self.start() + + async def shutdown(self) -> None: + from _bentoml_sdk.service.dependency import cleanup + + if self._service_instance is not None: + for name in dir(self.bento_service.inner): + member = getattr(self.bento_service.inner, name) + if ( + not name.startswith("__") + and callable(member) + and getattr(member, "__bentoml_shutdown_hook__", False) + ): + result = getattr(self._service_instance, name)() + if inspect.isawaitable(result): + await result + await cleanup() + self._service_instance = None + set_current_service(None) + + await self.stop(grace=self.graceful_shutdown_timeout) + await self.health_servicer.enter_graceful_shutdown() + self.loop.stop() diff --git a/src/_bentoml_impl/server/serving.py b/src/_bentoml_impl/server/serving.py index 480e86f28d9..208afb85363 100644 --- a/src/_bentoml_impl/server/serving.py +++ b/src/_bentoml_impl/server/serving.py @@ -404,3 +404,280 @@ def serve_http( except Exception: shutil.rmtree(uds_path, ignore_errors=True) raise + + +_GRPC_SERVICE_WORKER_SCRIPT = "_bentoml_impl.worker.grpc_service" +_PROMETHEUS_SERVER_NAME = "_prometheus_server" + + +@inject(squeeze_none=True) +def serve_grpc( + bento_identifier: str | AnyService, + working_dir: str | None = None, + host: str = Provide[BentoMLContainer.grpc.host], + port: int = Provide[BentoMLContainer.grpc.port], + backlog: int = Provide[BentoMLContainer.api_server_config.backlog], + ssl_certfile: str | None = Provide[BentoMLContainer.ssl.certfile], + ssl_keyfile: str | None = Provide[BentoMLContainer.ssl.keyfile], + ssl_ca_certs: str | None = Provide[BentoMLContainer.ssl.ca_certs], + bentoml_home: str = Provide[BentoMLContainer.bentoml_home], + development_mode: bool = False, + reload: bool = False, + dependency_map: dict[str, str] | None = None, + service_name: str = "", + threaded: bool = False, + max_concurrent_streams: int | None = Provide[ + BentoMLContainer.grpc.max_concurrent_streams + ], + reflection: bool = Provide[BentoMLContainer.grpc.reflection.enabled], + channelz: bool = Provide[BentoMLContainer.grpc.channelz.enabled], + protocol_version: str = "v1", +) -> Server: + import psutil + from circus.sockets import CircusSocket + + from bentoml._internal.log import SERVER_LOGGING_CONFIG + from bentoml._internal.utils import reserve_free_port + from bentoml._internal.utils.analytics.usage_stats import track_serve + from bentoml._internal.utils.circus import create_standalone_arbiter + from bentoml.exceptions import BentoMLException + from bentoml.serving import PROMETHEUS_MESSAGE + from bentoml.serving import SCRIPT_GRPC_PROMETHEUS_SERVER + from bentoml.serving import construct_ssl_args + from bentoml.serving import create_watcher + from bentoml.serving import ensure_prometheus_dir + from bentoml.serving import make_reload_plugin + + from ..loader import load + from .allocator import ResourceAllocator + + if protocol_version != "v1": + raise BentoMLException( + f"@bentoml.service() gRPC serving only supports protocol v1, got {protocol_version!r}" + ) + + if WINDOWS and not development_mode: + raise BentoMLException( + "'grpc' is not supported on Windows without '--development'. The reason being SO_REUSEPORT socket option is only available on UNIX system, and gRPC implementation depends on this behaviour." + ) + if psutil.MACOS or psutil.FREEBSD: + logger.warning( + "Due to gRPC implementation on exposing SO_REUSEPORT, BentoML production server's behaviour on %s is not correct. We recommend to containerize BentoServer as a Linux container instead. For testing locally, use `bentoml serve-grpc --development`", + "MacOS" if psutil.MACOS else "FreeBSD", + ) + + env = {"PROMETHEUS_MULTIPROC_DIR": ensure_prometheus_dir()} + if isinstance(bento_identifier, Service): + svc = bento_identifier + assert working_dir is None, ( + "working_dir should not be set when passing a service in process" + ) + bento_identifier = svc.import_string + bento_path = pathlib.Path(svc.working_dir) + else: + svc = load(bento_identifier, working_dir) + bento_path = pathlib.Path(working_dir or ".") + + for env_var in svc.envs: + if env_var.stage == "build": + continue + if env_var.name in env: + continue + if env_var.name in os.environ: + env[env_var.name] = os.environ[env_var.name] + elif env_var.value: + env[env_var.name] = env_var.value + else: + raise BentoMLException( + f"Environment variable '{env_var.name}' is required but not set. " + f"Either set it in the environment or provide a default value in the service definition." + ) + + watchers: list[Watcher] = [] + sockets: list[CircusSocket] = [] + allocator = ResourceAllocator() + if dependency_map is None: + dependency_map = {} + if service_name and service_name != svc.name: + svc = svc.find_dependent_by_name(service_name) + num_workers, worker_env = allocator.get_worker_env(svc) + if development_mode: + num_workers = 1 + server_on_deployment(svc) + uds_path = tempfile.mkdtemp(prefix="bentoml-uds-") + try: + if not service_name and not development_mode: + with contextlib.ExitStack() as port_stack: + for name, dep_svc in svc.all_services(exclude_urls=True).items(): + if name == svc.name or name in dependency_map: + continue + + dependency_env = env.copy() + for env_var in dep_svc.envs: + if env_var.stage == "build": + continue + if env_var.name in dependency_env: + continue + if env_var.value: + dependency_env[env_var.name] = env_var.value + elif env_var.name in os.environ: + dependency_env[env_var.name] = os.environ[env_var.name] + else: + raise BentoMLException( + f"Environment variable '{env_var.name}' is required for service '{name}' but not set. " + f"Either set it in the environment or provide a default value in the service definition." + ) + + new_watcher, new_socket, uri = create_dependency_watcher( + bento_identifier, + dep_svc, + uds_path, + port_stack, + backlog, + allocator, + str(bento_path.absolute()), + env={k: str(v) for k, v in dependency_env.items()}, + ) + watchers.append(new_watcher) + if new_socket: + sockets.append(new_socket) + dependency_map[name] = uri + server_on_deployment(dep_svc) + port_stack.enter_context(reserve_free_port()) + + env.update(worker_env) + ssl_args = construct_ssl_args( + ssl_certfile=ssl_certfile, + ssl_keyfile=ssl_keyfile, + ssl_ca_certs=ssl_ca_certs, + ) + scheme = "https" if BentoMLContainer.ssl.enabled.get() else "http" + close_child_stdin = not development_mode + + with contextlib.ExitStack() as port_stack: + api_port = port_stack.enter_context( + reserve_free_port(host, port=port, enable_so_reuseport=True) + ) + server_args = [ + "-m", + _GRPC_SERVICE_WORKER_SCRIPT, + bento_identifier, + "--host", + host, + "--port", + str(api_port), + "--working-dir", + str(bento_path.absolute()), + "--worker-id", + "$(CIRCUS.WID)", + "--protocol-version", + protocol_version, + *ssl_args, + ] + if reflection: + server_args.append("--enable-reflection") + if channelz: + server_args.append("--enable-channelz") + if max_concurrent_streams: + server_args.extend( + ["--max-concurrent-streams", str(max_concurrent_streams)] + ) + if development_mode: + server_args.append("--development-mode") + + watchers.append( + create_watcher( + name="grpc_api_server", + args=server_args, + use_sockets=False, + working_dir=str(bento_path.absolute()), + numprocesses=num_workers, + close_child_stdin=close_child_stdin, + env={k: str(v) for k, v in env.items()}, + ) + ) + + if BentoMLContainer.api_server_config.metrics.enabled.get(): + metrics_host = BentoMLContainer.grpc.metrics.host.get() + metrics_port = BentoMLContainer.grpc.metrics.port.get() + sockets.append( + CircusSocket( + name=_PROMETHEUS_SERVER_NAME, + host=metrics_host, + port=metrics_port, + backlog=backlog, + ) + ) + watchers.append( + create_watcher( + name="prom_server", + args=[ + "-m", + SCRIPT_GRPC_PROMETHEUS_SERVER, + "--fd", + f"$(circus.sockets.{_PROMETHEUS_SERVER_NAME})", + "--backlog", + str(backlog), + ], + working_dir=str(bento_path.absolute()), + numprocesses=1, + singleton=True, + close_child_stdin=close_child_stdin, + ) + ) + log_metrics_host = ( + "127.0.0.1" if metrics_host == "0.0.0.0" else metrics_host + ) + logger.info( + PROMETHEUS_MESSAGE, + "gRPC", + bento_identifier, + f"http://{log_metrics_host}:{metrics_port}", + ) + + log_host = "localhost" if host in ["0.0.0.0", "::"] else host + dependency_map[svc.name] = f"{scheme}://{log_host}:{port}" + inject_env = {"BENTOML_RUNNER_MAP": json.dumps(dependency_map)} + for watcher in watchers: + if watcher.env is None: + watcher.env = inject_env + else: + watcher.env.update(inject_env) + + arbiter_kwargs: dict[str, t.Any] = { + "watchers": watchers, + "sockets": sockets, + "threaded": threaded, + } + if reload: + arbiter_kwargs["plugins"] = [ + make_reload_plugin(str(bento_path.absolute()), bentoml_home) + ] + if development_mode: + arbiter_kwargs["debug"] = True if sys.platform != "win32" else False + arbiter_kwargs["loggerconfig"] = SERVER_LOGGING_CONFIG + arbiter_kwargs["loglevel"] = "WARNING" + + arbiter = create_standalone_arbiter(**arbiter_kwargs) + arbiter.exit_stack.enter_context( + track_serve(svc, production=not development_mode, serve_kind="grpc") + ) + arbiter.exit_stack.callback(shutil.rmtree, uds_path, ignore_errors=True) + arbiter.start( + cb=lambda _: ( + logger.info( # type: ignore + 'Starting production %s BentoServer from "%s" listening on %s://%s:%d (Press CTRL+C to quit)', + "gRPC", + bento_identifier, + scheme, + log_host, + port, + ) + if not svc.has_custom_command() + else None + ), + ) + return Server(url=f"{scheme}://{log_host}:{port}", arbiter=arbiter) + except Exception: + shutil.rmtree(uds_path, ignore_errors=True) + raise diff --git a/src/_bentoml_impl/worker/grpc_service.py b/src/_bentoml_impl/worker/grpc_service.py new file mode 100644 index 00000000000..b5b6faca810 --- /dev/null +++ b/src/_bentoml_impl/worker/grpc_service.py @@ -0,0 +1,157 @@ +from __future__ import annotations + +import json +import typing as t + +import click + + +@click.command() +@click.argument("bento_identifier", type=click.STRING, required=False, default=".") +@click.option("--host", type=click.STRING, required=False, default=None) +@click.option("--port", type=click.INT, required=False, default=None) +@click.option( + "--runner-map", + type=click.STRING, + envvar="BENTOML_RUNNER_MAP", + help="JSON string of runners map, default sets to envars `BENTOML_RUNNER_MAP`", +) +@click.option( + "--working-dir", + type=click.Path(exists=True), + help="Working directory for the API server", +) +@click.option( + "--worker-id", + required=False, + type=click.INT, + default=None, + help="If set, start the server as a bare worker with the given worker ID. Otherwise start a standalone server with a supervisor process.", +) +@click.option( + "--enable-reflection", + type=click.BOOL, + is_flag=True, + help="Enable reflection.", +) +@click.option( + "--enable-channelz", + type=click.BOOL, + is_flag=True, + help="Enable channelz.", + default=False, +) +@click.option( + "--max-concurrent-streams", + type=click.INT, + help="Maximum number of concurrent incoming streams to allow on a HTTP2 connection.", + default=None, +) +@click.option( + "--ssl-certfile", + type=str, + default=None, + help="SSL certificate file", +) +@click.option( + "--ssl-keyfile", + type=str, + default=None, + help="SSL key file", +) +@click.option( + "--ssl-ca-certs", + type=str, + default=None, + help="CA certificates file", +) +@click.option( + "--protocol-version", + type=click.Choice(["v1"]), + help="Determine the version of generated gRPC stubs to use.", + default="v1", + show_default=True, +) +@click.option( + "--development-mode", + type=click.BOOL, + help="Run the API server in development mode", + is_flag=True, + default=False, + show_default=True, +) +def main( + bento_identifier: str, + host: str, + port: int, + runner_map: str | None, + working_dir: str | None, + worker_id: int | None, + enable_reflection: bool, + enable_channelz: bool, + max_concurrent_streams: int | None, + ssl_certfile: str | None, + ssl_keyfile: str | None, + ssl_ca_certs: str | None, + protocol_version: str, + development_mode: bool, +): + """ + Start a gRPC BentoML API server for a new-style Service. + \b + This is an internal API, users should not use this directly. Instead use `bentoml serve-grpc [--options]` + """ + + from _bentoml_impl.loader import load + from bentoml._internal.configuration.containers import BentoMLContainer + from bentoml._internal.context import server_context + from bentoml._internal.log import configure_server_logging + + server_context.service_type = "grpc_api_server" + server_context.worker_index = worker_id + configure_server_logging() + + if worker_id is None: + BentoMLContainer.config.runner_probe.enabled.set(False) + + BentoMLContainer.development_mode.set(development_mode) + if runner_map is not None: + BentoMLContainer.remote_runner_mapping.set(json.loads(runner_map)) + + svc = load(bento_identifier, working_dir=working_dir) + svc.inject_config() + if not port: + port = BentoMLContainer.grpc.port.get() + if not host: + host = BentoMLContainer.grpc.host.get() + + server_context.service_name = svc.name + if svc.bento is None or svc.bento.tag is None: + server_context.bento_name = svc.name + server_context.bento_version = "not available" + else: + server_context.bento_name = svc.bento.tag.name + server_context.bento_version = svc.bento.tag.version or "not available" + + from _bentoml_impl.server import grpc_app as grpc + + grpc_options: dict[str, t.Any] = { + "bind_address": f"{host}:{port}", + "enable_reflection": enable_reflection, + "enable_channelz": enable_channelz, + "protocol_version": protocol_version, + } + if max_concurrent_streams: + grpc_options["max_concurrent_streams"] = int(max_concurrent_streams) + if ssl_certfile: + grpc_options["ssl_certfile"] = ssl_certfile + if ssl_keyfile: + grpc_options["ssl_keyfile"] = ssl_keyfile + if ssl_ca_certs: + grpc_options["ssl_ca_certs"] = ssl_ca_certs + + grpc.Server(svc, **grpc_options).run() + + +if __name__ == "__main__": + main() # pylint: disable=no-value-for-parameter diff --git a/src/bentoml/bentos.py b/src/bentoml/bentos.py index fa46a72d9ab..7e301ab5f1f 100644 --- a/src/bentoml/bentos.py +++ b/src/bentoml/bentos.py @@ -575,29 +575,63 @@ def serve( ssl_ciphers=ssl_ciphers, ) elif server_type == "grpc": - from .serving import serve_grpc_production - - if not isinstance(bento, str): - assert isinstance(bento, Service) - bento, working_dir = bento.get_service_import_origin() - - return serve_grpc_production( - bento_identifier=bento, - reload=reload, - host=host, - port=port, - working_dir=working_dir, - api_workers=api_workers, - backlog=backlog, - threaded=not blocking, - development_mode=not production, - ssl_certfile=ssl_certfile, - ssl_keyfile=ssl_keyfile, - ssl_ca_certs=ssl_ca_certs, - max_concurrent_streams=max_concurrent_streams, - reflection=enable_reflection, - channelz=enable_channelz, - protocol_version=grpc_protocol_version, - ) + from _bentoml_sdk import Service as NewService + + from ._internal.service import load + + if not isinstance(bento, (Service, NewService)): + svc = load(bento, working_dir=working_dir) + else: + svc = bento + + if isinstance(svc, Service): # < 1.2 bento + from .serving import serve_grpc_production + + if not isinstance(bento, str): + bento, working_dir = svc.get_service_import_origin() + + return serve_grpc_production( + bento_identifier=bento, + reload=reload, + host=host, + port=port, + working_dir=working_dir, + api_workers=api_workers, + backlog=backlog, + threaded=not blocking, + development_mode=not production, + ssl_certfile=ssl_certfile, + ssl_keyfile=ssl_keyfile, + ssl_ca_certs=ssl_ca_certs, + max_concurrent_streams=max_concurrent_streams, + reflection=enable_reflection, + channelz=enable_channelz, + protocol_version=grpc_protocol_version, + ) + else: # >= 1.2 bento + from _bentoml_impl.server.serving import serve_grpc + + if not isinstance(bento, str): + bento = svc.import_string + working_dir = svc.working_dir + + svc.inject_config() + return serve_grpc( + bento_identifier=bento, + working_dir=working_dir, + reload=reload, + host=host, + port=port, + backlog=backlog, + development_mode=not production, + threaded=not blocking, + ssl_certfile=ssl_certfile, + ssl_keyfile=ssl_keyfile, + ssl_ca_certs=ssl_ca_certs, + max_concurrent_streams=max_concurrent_streams, + reflection=enable_reflection, + channelz=enable_channelz, + protocol_version=grpc_protocol_version, + ) else: raise BadInput(f"Unknown server type: '{server_type}'") diff --git a/src/bentoml/grpc/interceptors/opentelemetry.py b/src/bentoml/grpc/interceptors/opentelemetry.py index 17c3f9e4d08..a35a2c1bb5d 100644 --- a/src/bentoml/grpc/interceptors/opentelemetry.py +++ b/src/bentoml/grpc/interceptors/opentelemetry.py @@ -109,6 +109,11 @@ async def abort( code, details=details, trailing_metadata=trailing_metadata ) + async def abort_with_status(self, status: grpc.Status) -> None: + # grpcio>=1.66 requires this abstract method on aio.ServicerContext. + trailing = getattr(status, "trailing_metadata", None) or tuple() + await self.abort(status.code, details=status.details or "", trailing_metadata=trailing) + def set_code(self, code: grpc.StatusCode) -> None: self._code = code details = self._details or code.value[1] diff --git a/src/bentoml/grpc/v1/service_pb2.py b/src/bentoml/grpc/v1/service_pb2.py index f4372938fb3..e9aceed3963 100644 --- a/src/bentoml/grpc/v1/service_pb2.py +++ b/src/bentoml/grpc/v1/service_pb2.py @@ -1,6 +1,6 @@ from google.protobuf import __version__ -if __version__.startswith("4"): +if int(__version__.split(".", 1)[0]) >= 4: from ._generated_pb4.service_pb2 import * else: from ._generated_pb3.service_pb2 import * diff --git a/src/bentoml/grpc/v1/service_pb2_grpc.py b/src/bentoml/grpc/v1/service_pb2_grpc.py index c256ef99f4d..61866d71820 100644 --- a/src/bentoml/grpc/v1/service_pb2_grpc.py +++ b/src/bentoml/grpc/v1/service_pb2_grpc.py @@ -1,6 +1,6 @@ from google.protobuf import __version__ -if __version__.startswith("4"): +if int(__version__.split(".", 1)[0]) >= 4: from ._generated_pb4.service_pb2_grpc import * else: from ._generated_pb3.service_pb2_grpc import * diff --git a/src/bentoml/grpc/v1alpha1/service_pb2.py b/src/bentoml/grpc/v1alpha1/service_pb2.py index f4372938fb3..e9aceed3963 100644 --- a/src/bentoml/grpc/v1alpha1/service_pb2.py +++ b/src/bentoml/grpc/v1alpha1/service_pb2.py @@ -1,6 +1,6 @@ from google.protobuf import __version__ -if __version__.startswith("4"): +if int(__version__.split(".", 1)[0]) >= 4: from ._generated_pb4.service_pb2 import * else: from ._generated_pb3.service_pb2 import * diff --git a/src/bentoml/grpc/v1alpha1/service_pb2_grpc.py b/src/bentoml/grpc/v1alpha1/service_pb2_grpc.py index c256ef99f4d..61866d71820 100644 --- a/src/bentoml/grpc/v1alpha1/service_pb2_grpc.py +++ b/src/bentoml/grpc/v1alpha1/service_pb2_grpc.py @@ -1,6 +1,6 @@ from google.protobuf import __version__ -if __version__.startswith("4"): +if int(__version__.split(".", 1)[0]) >= 4: from ._generated_pb4.service_pb2_grpc import * else: from ._generated_pb3.service_pb2_grpc import * diff --git a/src/bentoml_cli/serve.py b/src/bentoml_cli/serve.py index a5b71b8d768..0890bb8bdee 100644 --- a/src/bentoml_cli/serve.py +++ b/src/bentoml_cli/serve.py @@ -505,35 +505,62 @@ def serve_grpc( # type: ignore (unused warning) working_dir = os.path.expanduser(bento) else: working_dir = "." + if sys.path[0] != working_dir: + sys.path.insert(0, working_dir) - from bentoml.serving import serve_grpc_production + from bentoml._internal.service.loader import load + from bentoml.legacy import Service - if development: - serve_grpc_production( - bento, - working_dir=working_dir, - port=port, - host=DEFAULT_DEV_SERVER_HOST if not host else host, - backlog=backlog, - api_workers=1, - ssl_keyfile=ssl_keyfile, - ssl_certfile=ssl_certfile, - ssl_ca_certs=ssl_ca_certs, - max_concurrent_streams=max_concurrent_streams, - reflection=enable_reflection, - channelz=enable_channelz, - protocol_version=protocol_version, - reload=reload, - development_mode=True, - ) + svc = load(bento_identifier=bento, working_dir=working_dir) + if isinstance(svc, Service): + from bentoml.serving import serve_grpc_production + + if development: + serve_grpc_production( + bento, + working_dir=working_dir, + port=port, + host=DEFAULT_DEV_SERVER_HOST if not host else host, + backlog=backlog, + api_workers=1, + ssl_keyfile=ssl_keyfile, + ssl_certfile=ssl_certfile, + ssl_ca_certs=ssl_ca_certs, + max_concurrent_streams=max_concurrent_streams, + reflection=enable_reflection, + channelz=enable_channelz, + protocol_version=protocol_version, + reload=reload, + development_mode=True, + ) + else: + serve_grpc_production( + bento, + working_dir=working_dir, + port=port, + backlog=backlog, + host=host, + api_workers=api_workers, + ssl_keyfile=ssl_keyfile, + ssl_certfile=ssl_certfile, + ssl_ca_certs=ssl_ca_certs, + max_concurrent_streams=max_concurrent_streams, + reflection=enable_reflection, + channelz=enable_channelz, + protocol_version=protocol_version, + reload=reload, + development_mode=False, + ) else: - serve_grpc_production( + from _bentoml_impl.server import serve_grpc + + svc.inject_config() + serve_grpc( bento, working_dir=working_dir, + host=DEFAULT_DEV_SERVER_HOST if development and not host else host, port=port, backlog=backlog, - host=host, - api_workers=api_workers, ssl_keyfile=ssl_keyfile, ssl_certfile=ssl_certfile, ssl_ca_certs=ssl_ca_certs, @@ -542,7 +569,7 @@ def serve_grpc( # type: ignore (unused warning) channelz=enable_channelz, protocol_version=protocol_version, reload=reload, - development_mode=False, + development_mode=development, ) return cli diff --git a/tests/e2e/bento_new_sdk_grpc/bentofile.yaml b/tests/e2e/bento_new_sdk_grpc/bentofile.yaml new file mode 100644 index 00000000000..6796837f58a --- /dev/null +++ b/tests/e2e/bento_new_sdk_grpc/bentofile.yaml @@ -0,0 +1,3 @@ +service: "service.py:MyService" +include: + - "*.py" diff --git a/tests/e2e/bento_new_sdk_grpc/conftest.py b/tests/e2e/bento_new_sdk_grpc/conftest.py new file mode 100644 index 00000000000..0de84ec6906 --- /dev/null +++ b/tests/e2e/bento_new_sdk_grpc/conftest.py @@ -0,0 +1,10 @@ +from __future__ import annotations + +import sys + +import pytest + + +@pytest.fixture(autouse=True) +def clear_import_cache() -> None: + sys.modules.pop("service", None) diff --git a/tests/e2e/bento_new_sdk_grpc/requirements.txt b/tests/e2e/bento_new_sdk_grpc/requirements.txt new file mode 100644 index 00000000000..24ce15ab7ea --- /dev/null +++ b/tests/e2e/bento_new_sdk_grpc/requirements.txt @@ -0,0 +1 @@ +numpy diff --git a/tests/e2e/bento_new_sdk_grpc/service.py b/tests/e2e/bento_new_sdk_grpc/service.py new file mode 100644 index 00000000000..6a3edc6b433 --- /dev/null +++ b/tests/e2e/bento_new_sdk_grpc/service.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +import typing as t + +import numpy as np +import numpy.typing as npt + +import bentoml +from bentoml.validators import TensorSchema + +Array = t.Annotated[npt.NDArray[np.float32], TensorSchema("numpy-array")] + + +@bentoml.service +class MyService: + @bentoml.api + def greet(self, name: str) -> str: + return f"hello {name}" + + @bentoml.api + def predict(self, x: Array) -> Array: + return x * 2 + + @bentoml.api + async def stream_greet(self, name: str) -> t.AsyncGenerator[str, None]: + yield f"hello {name}" diff --git a/tests/e2e/bento_new_sdk_grpc/test_grpc.py b/tests/e2e/bento_new_sdk_grpc/test_grpc.py new file mode 100644 index 00000000000..928bab9585a --- /dev/null +++ b/tests/e2e/bento_new_sdk_grpc/test_grpc.py @@ -0,0 +1,92 @@ +from __future__ import annotations + +import asyncio +import time +from pathlib import Path +from urllib.parse import urlparse + +import numpy as np +import psutil +import pytest +from google.protobuf import wrappers_pb2 + +import bentoml +from bentoml.grpc.utils import import_grpc +from bentoml.testing.grpc import async_client_call +from bentoml.testing.grpc import create_channel +from bentoml.testing.grpc import make_pb_ndarray + +pytest.importorskip("grpc") + +grpc, aio = import_grpc() + +PROJECT_DIR = Path(__file__).parent +PORT = 38765 + + +pytestmark = pytest.mark.skipif( + psutil.WINDOWS, reason="gRPC is not supported on Windows." +) + + +def _host_url(url: str) -> str: + parsed = urlparse(url) + return f"{parsed.hostname}:{parsed.port}" + + +async def _wait_until_ready(host_url: str, timeout: float = 100) -> None: + deadline = time.time() + timeout + last_err: Exception | None = None + while time.time() < deadline: + try: + async with aio.insecure_channel(host_url) as channel: + await asyncio.wait_for(channel.channel_ready(), timeout=2) + return + except Exception as exc: + last_err = exc + await asyncio.sleep(0.5) + raise TimeoutError(f"gRPC server at {host_url} was not ready") from last_err + + +@pytest.mark.asyncio +async def test_unary_greet_and_predict(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("BENTOML_CONFIG_OPTIONS", "services.metrics.enabled=false") + with bentoml.serve( + "service.py:MyService", + working_dir=str(PROJECT_DIR), + port=PORT, + server_type="grpc", + production=False, + ) as server: + host_url = _host_url(server.url) + await _wait_until_ready(host_url, timeout=100) + async with create_channel(host_url) as channel: + await async_client_call( + "greet", + channel=channel, + data={"text": wrappers_pb2.StringValue(value="world")}, + assert_code=grpc.StatusCode.OK, + assert_data=lambda resp: resp.text.value == "hello world", + ) + arr = np.array([1.0, 2.0, 3.0], dtype=np.float32) + await async_client_call( + "predict", + channel=channel, + data={"ndarray": make_pb_ndarray(arr)}, + assert_code=grpc.StatusCode.OK, + assert_data=lambda resp: np.allclose( + resp.ndarray.float_values, [2.0, 4.0, 6.0] + ), + ) + await async_client_call( + "missing", + channel=channel, + data={"text": wrappers_pb2.StringValue(value="x")}, + assert_code=grpc.StatusCode.INVALID_ARGUMENT, + ) + await async_client_call( + "stream_greet", + channel=channel, + data={"text": wrappers_pb2.StringValue(value="world")}, + assert_code=grpc.StatusCode.UNIMPLEMENTED, + ) diff --git a/tests/unit/_bentoml_impl/__init__.py b/tests/unit/_bentoml_impl/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/_bentoml_impl/grpc/__init__.py b/tests/unit/_bentoml_impl/grpc/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/_bentoml_impl/grpc/test_codec.py b/tests/unit/_bentoml_impl/grpc/test_codec.py new file mode 100644 index 00000000000..9f5c12b5e1b --- /dev/null +++ b/tests/unit/_bentoml_impl/grpc/test_codec.py @@ -0,0 +1,150 @@ +from __future__ import annotations + +import typing as t +from pathlib import Path + +import numpy as np +import pytest +from google.protobuf import json_format +from google.protobuf import struct_pb2 +from google.protobuf import wrappers_pb2 + +from _bentoml_impl.grpc.codec import decode_proto +from _bentoml_impl.grpc.codec import encode_proto +from _bentoml_impl.grpc.codec import proto_binding +from _bentoml_sdk.io_models import IODescriptor +from _bentoml_sdk.validators import TensorSchema +from bentoml.exceptions import InvalidArgument +from bentoml.grpc.utils import import_generated_stubs + +pb, _ = import_generated_stubs("v1") + +Array = t.Annotated[np.ndarray, TensorSchema("numpy-array")] + + +def _input_of(fn: t.Callable[..., t.Any]) -> type[IODescriptor]: + return IODescriptor.from_input(fn, skip_self=True) + + +def _output_of(fn: t.Callable[..., t.Any]) -> type[IODescriptor]: + return IODescriptor.from_output(fn) + + +def greet(self, name: str) -> str: + return f"hello {name}" + + +def add(self, a: int, b: int) -> dict[str, int]: + return {"sum": a + b} + + +def predict(self, x: Array) -> Array: + return x * 2 + + +def echo_file(self, data: Path) -> Path: + return data + + +def test_proto_binding_maps_single_string_field_to_text(): + binding = proto_binding(_input_of(greet)) + assert binding.field == "text" + assert binding.unwrap_key == "name" + + +def test_proto_binding_maps_string_return_to_text(): + binding = proto_binding(_output_of(greet)) + assert binding.field == "text" + assert binding.unwrap_key is None + + +def test_proto_binding_maps_multi_field_input_to_json(): + binding = proto_binding(_input_of(add)) + assert binding.field == "json" + assert binding.unwrap_key is None + + +def test_proto_binding_maps_tensor_field_to_ndarray(): + binding = proto_binding(_input_of(predict)) + assert binding.field == "ndarray" + assert binding.unwrap_key == "x" + + +def test_proto_binding_maps_path_field_to_file(): + binding = proto_binding(_input_of(echo_file)) + assert binding.field == "file" + assert binding.unwrap_key == "data" + + +@pytest.mark.asyncio +async def test_text_round_trip(): + spec = _input_of(greet) + encoded_field, encoded = await encode_proto(spec, {"name": "willow"}) + assert encoded_field == "text" + assert encoded.value == "willow" + + decoded = await decode_proto(spec, "text", encoded) + assert decoded.name == "willow" + + +@pytest.mark.asyncio +async def test_json_round_trip(): + spec = _input_of(add) + encoded_field, encoded = await encode_proto(spec, {"a": 2, "b": 3}) + assert encoded_field == "json" + parsed = json_format.MessageToDict(encoded) + assert parsed == {"a": 2, "b": 3} + + decoded = await decode_proto(spec, "json", encoded) + assert decoded.a == 2 + assert decoded.b == 3 + + +@pytest.mark.asyncio +async def test_ndarray_round_trip(): + spec = _input_of(predict) + array = np.array([[1.0, 2.0], [3.0, 4.0]], dtype=np.float64) + encoded_field, encoded = await encode_proto(spec, array) + assert encoded_field == "ndarray" + assert list(encoded.shape) == [2, 2] + + decoded = await decode_proto(spec, "ndarray", encoded) + np.testing.assert_array_equal(decoded.x, array) + + +@pytest.mark.asyncio +async def test_file_round_trip(tmp_path: Path): + spec = _input_of(echo_file) + payload = b"grpc-file-bytes" + source = tmp_path / "input.bin" + source.write_bytes(payload) + + encoded_field, encoded = await encode_proto(spec, source) + assert encoded_field == "file" + assert encoded.content == payload + + decoded = await decode_proto(spec, "file", encoded) + assert Path(decoded.data).read_bytes() == payload + + +@pytest.mark.asyncio +async def test_output_string_encodes_as_text(): + spec = _output_of(greet) + field, encoded = await encode_proto(spec, "hello willow") + assert field == "text" + assert isinstance(encoded, wrappers_pb2.StringValue) + assert encoded.value == "hello willow" + + +@pytest.mark.asyncio +async def test_rejects_serialized_bytes(): + spec = _input_of(greet) + with pytest.raises(InvalidArgument, match="serialized_bytes"): + await decode_proto(spec, "serialized_bytes", b"nope") + + +@pytest.mark.asyncio +async def test_rejects_mismatched_proto_field(): + spec = _input_of(greet) + with pytest.raises(InvalidArgument, match="text"): + await decode_proto(spec, "json", struct_pb2.Value()) diff --git a/tests/unit/_bentoml_impl/grpc/test_servicer.py b/tests/unit/_bentoml_impl/grpc/test_servicer.py new file mode 100644 index 00000000000..dc0d31edfdd --- /dev/null +++ b/tests/unit/_bentoml_impl/grpc/test_servicer.py @@ -0,0 +1,145 @@ +from __future__ import annotations + +import typing as t + +import pytest +from google.protobuf import wrappers_pb2 + +import bentoml +from _bentoml_impl.server.grpc.servicer.v1 import create_bento_servicer +from bentoml.exceptions import BentoMLException +from bentoml.grpc.utils import import_generated_stubs +from bentoml.grpc.utils import import_grpc + +pb, _ = import_generated_stubs("v1") +grpc, aio = import_grpc() + + +@bentoml.service +class Greeter: + @bentoml.api + def greet(self, name: str) -> str: + return f"hello {name}" + + @bentoml.api + async def agreet(self, name: str) -> str: + return f"hello {name}" + + @bentoml.api + async def stream_greet(self, name: str) -> t.AsyncGenerator[str, None]: + yield f"hello {name}" + + @bentoml.api(batchable=True) + def batch_greet(self, name: list[str]) -> list[str]: + return [f"hello {n}" for n in name] + + @bentoml.task + def long_job(self, name: str) -> str: + return name + + +class FakeContext: + def __init__(self) -> None: + self.code: grpc.StatusCode | None = None + self.details: str | None = None + + async def abort(self, code: grpc.StatusCode, details: str = "") -> t.NoReturn: + self.code = code + self.details = details + raise aio.AbortError() + + +@pytest.fixture +def servicer(): + return create_bento_servicer(Greeter) + + +@pytest.mark.asyncio +async def test_call_greet(servicer): + ctx = FakeContext() + request = pb.Request( + api_name="greet", text=wrappers_pb2.StringValue(value="world") + ) + response = await servicer.Call(request, ctx) + assert response is not None + assert response.text.value == "hello world" + + +@pytest.mark.asyncio +async def test_call_async_greet(servicer): + ctx = FakeContext() + request = pb.Request( + api_name="agreet", text=wrappers_pb2.StringValue(value="world") + ) + response = await servicer.Call(request, ctx) + assert response is not None + assert response.text.value == "hello world" + + +@pytest.mark.asyncio +async def test_unknown_api_name_aborts(servicer): + ctx = FakeContext() + request = pb.Request( + api_name="missing", text=wrappers_pb2.StringValue(value="x") + ) + with pytest.raises(aio.AbortError): + await servicer.Call(request, ctx) + assert ctx.code == grpc.StatusCode.INVALID_ARGUMENT + assert ctx.details is not None + assert "api_name" in ctx.details + + +@pytest.mark.asyncio +async def test_streaming_api_unimplemented(servicer): + ctx = FakeContext() + request = pb.Request( + api_name="stream_greet", text=wrappers_pb2.StringValue(value="world") + ) + with pytest.raises(aio.AbortError): + await servicer.Call(request, ctx) + assert ctx.code == grpc.StatusCode.UNIMPLEMENTED + assert ctx.details is not None + assert "streaming" in ctx.details + + +@pytest.mark.asyncio +async def test_batchable_api_unimplemented(servicer): + ctx = FakeContext() + request = pb.Request(api_name="batch_greet") + with pytest.raises(aio.AbortError): + await servicer.Call(request, ctx) + assert ctx.code == grpc.StatusCode.UNIMPLEMENTED + assert ctx.details is not None + assert "batchable" in ctx.details + + +@pytest.mark.asyncio +async def test_task_api_unimplemented(servicer): + ctx = FakeContext() + request = pb.Request( + api_name="long_job", text=wrappers_pb2.StringValue(value="world") + ) + with pytest.raises(aio.AbortError): + await servicer.Call(request, ctx) + assert ctx.code == grpc.StatusCode.UNIMPLEMENTED + assert ctx.details is not None + assert "task" in ctx.details + + +@pytest.mark.asyncio +async def test_service_metadata(servicer): + ctx = FakeContext() + meta = await servicer.ServiceMetadata(pb.ServiceMetadataRequest(), ctx) + assert meta.name == "Greeter" + names = {api.name for api in meta.apis} + assert {"greet", "agreet", "stream_greet", "batch_greet", "long_job"} <= names + greet = next(api for api in meta.apis if api.name == "greet") + assert greet.input.descriptor_id == "bentoml.sdk.IODescriptor" + assert greet.output.descriptor_id == "bentoml.sdk.IODescriptor" + + +def test_serve_grpc_rejects_v1alpha1(): + from _bentoml_impl.server.serving import serve_grpc + + with pytest.raises(BentoMLException, match="protocol v1"): + serve_grpc("service.py:Greeter", protocol_version="v1alpha1") From aefbfb5f635d5d9963a0adfda888dd9c94cfa01e Mon Sep 17 00:00:00 2001 From: Sandipan Kundu Date: Thu, 27 Aug 2026 16:18:14 +0530 Subject: [PATCH 2/2] fix: address gRPC worker args, request context, and file cleanup --- docs/source/build-with-bentoml/services.rst | 2 + src/_bentoml_impl/grpc/codec.py | 16 +- .../server/grpc/servicer/v1/__init__.py | 170 ++++++++++++------ src/_bentoml_impl/server/serving.py | 3 + src/_bentoml_impl/worker/grpc_service.py | 5 + tests/e2e/bento_new_sdk_grpc/service.py | 15 +- tests/e2e/bento_new_sdk_grpc/test_grpc.py | 20 +++ tests/unit/_bentoml_impl/grpc/test_codec.py | 13 +- .../unit/_bentoml_impl/grpc/test_servicer.py | 90 +++++++++- tests/unit/_bentoml_impl/grpc/test_worker.py | 30 ++++ 10 files changed, 296 insertions(+), 68 deletions(-) create mode 100644 tests/unit/_bentoml_impl/grpc/test_worker.py diff --git a/docs/source/build-with-bentoml/services.rst b/docs/source/build-with-bentoml/services.rst index 5838072371c..696fab49ea8 100644 --- a/docs/source/build-with-bentoml/services.rst +++ b/docs/source/build-with-bentoml/services.rst @@ -68,6 +68,8 @@ New-style ``@bentoml.service()`` Services can also be served over gRPC: The gRPC server speaks protocol **v1** and currently supports ``str`` (text), JSON (Pydantic models, dicts, and lists), NumPy arrays (``ndarray``), and files. Streaming, batch, and task APIs return ``UNIMPLEMENTED``. Pandas DataFrames, images, multipart payloads, and protocol ``v1alpha1`` are not supported yet. +For APIs using :ref:`bentoml.Context `, incoming gRPC metadata is available through ``ctx.request.headers`` and ``ctx.response.headers`` is returned as trailing gRPC metadata. HTTP-specific response status codes and cookies do not have a gRPC equivalent. + Service configurations ---------------------- diff --git a/src/_bentoml_impl/grpc/codec.py b/src/_bentoml_impl/grpc/codec.py index c69f0f43d5c..d150dae9419 100644 --- a/src/_bentoml_impl/grpc/codec.py +++ b/src/_bentoml_impl/grpc/codec.py @@ -1,7 +1,5 @@ from __future__ import annotations -import os -import tempfile import typing as t from dataclasses import dataclass from pathlib import Path @@ -132,15 +130,6 @@ def _file_from_proto(field: t.Any) -> bytes: return bytes(content) -def _bytes_to_path(body: bytes) -> Path: - # FileSchema.decode() looks up the HTTP request temp dir; gRPC has none, - # so persist bytes to a local tempfile the validator can treat as a Path. - fd, name = tempfile.mkstemp(prefix="bentoml-grpc-") - with os.fdopen(fd, "wb") as handle: - handle.write(body) - return Path(name) - - def _file_to_proto(obj: t.Any) -> pb.File: if isinstance(obj, bytes): body = obj @@ -209,7 +198,10 @@ async def decode_proto( elif binding.field == "ndarray": decoded = await _ndarray_from_proto(value) else: - decoded = _bytes_to_path(_file_from_proto(value)) + # FileSchema stores byte inputs in the active request temp directory. + # The gRPC servicer owns that request lifecycle and cleans it after the + # response has been encoded. + decoded = _file_from_proto(value) return _wrap_value(spec, binding, decoded) diff --git a/src/_bentoml_impl/server/grpc/servicer/v1/__init__.py b/src/_bentoml_impl/server/grpc/servicer/v1/__init__.py index 2d00ceedfa3..12657675e05 100644 --- a/src/_bentoml_impl/server/grpc/servicer/v1/__init__.py +++ b/src/_bentoml_impl/server/grpc/servicer/v1/__init__.py @@ -31,6 +31,7 @@ import grpc from google.protobuf import struct_pb2 from grpc import aio + from starlette.requests import Request from _bentoml_sdk import Service from bentoml.grpc.types import BentoServicerContext @@ -46,6 +47,66 @@ def log_exception(request: pb.Request, exc_info: ExcInfoType) -> None: logger.error("Exception on /%s [POST]", request.api_name, exc_info=exc_info) +def _metadata_item(item: t.Any) -> tuple[t.Any, t.Any]: + if hasattr(item, "key") and hasattr(item, "value"): + return item.key, item.value + return item + + +def _metadata_bytes(value: t.Any) -> bytes: + if isinstance(value, bytes): + return value + return str(value).encode("latin-1") + + +def _request_from_grpc(request: pb.Request, context: BentoServicerContext) -> Request: + from starlette.requests import Request + + invocation_metadata = getattr(context, "invocation_metadata", None) + metadata = invocation_metadata() if invocation_metadata is not None else None + headers = [ + (_metadata_bytes(key), _metadata_bytes(value)) + for key, value in map(_metadata_item, metadata or ()) + ] + + async def receive() -> dict[str, t.Any]: + return {"type": "http.request", "body": b"", "more_body": False} + + return Request( + { + "type": "http", + "asgi": {"version": "3.0", "spec_version": "2.3"}, + "http_version": "2", + "method": "POST", + "scheme": "http", + "path": f"/{request.api_name}", + "raw_path": f"/{request.api_name}".encode(), + "query_string": b"", + "headers": headers, + "client": None, + "server": None, + "root_path": "", + "state": {}, + }, + receive, + ) + + +def _propagate_response_metadata(ctx: t.Any, context: BentoServicerContext) -> None: + raw_metadata = getattr(ctx.response.metadata, "raw", ()) + if not raw_metadata: + return + + outgoing = tuple( + (key.decode("latin-1"), value.decode("latin-1")) + for key, value in raw_metadata + ) + trailing_metadata = getattr(context, "trailing_metadata", None) + existing = trailing_metadata() if trailing_metadata is not None else None + existing_items = tuple(map(_metadata_item, existing or ())) + context.set_trailing_metadata((*existing_items, *outgoing)) + + def _call_args_from_input(method: t.Any, input_data: t.Any, ctx: t.Any) -> tuple[ tuple[t.Any, ...], dict[str, t.Any] ]: @@ -88,60 +149,65 @@ async def Call( # type: ignore (no async types) ) -> pb.Response | None: response = pb.Response() try: - if request.api_name not in service.apis: - raise InvalidArgument( - f"given 'api_name' is not defined in {service.name}", - ) from None - - method = service.apis[request.api_name] - if method.is_stream: - await context.abort( - code=grpc.StatusCode.UNIMPLEMENTED, - details=( - f"API {method.name!r} is a streaming endpoint; " - "gRPC streaming is not supported for @bentoml.service() yet" - ), - ) - return None - if method.batchable: - await context.abort( - code=grpc.StatusCode.UNIMPLEMENTED, - details=( - f"API {method.name!r} is batchable; " - "adaptive batching is not supported over gRPC for @bentoml.service() yet" - ), - ) - return None - if method.is_task: - await context.abort( - code=grpc.StatusCode.UNIMPLEMENTED, - details=( - f"API {method.name!r} is a task endpoint; " - "tasks are not supported over gRPC for @bentoml.service() yet" - ), + grpc_request = _request_from_grpc(request, context) + with service.context.in_request(grpc_request) as ctx: + if request.api_name not in service.apis: + raise InvalidArgument( + f"given 'api_name' is not defined in {service.name}", + ) from None + + method = service.apis[request.api_name] + if method.is_stream: + await context.abort( + code=grpc.StatusCode.UNIMPLEMENTED, + details=( + f"API {method.name!r} is a streaming endpoint; " + "gRPC streaming is not supported for @bentoml.service() yet" + ), + ) + return None + if method.batchable: + await context.abort( + code=grpc.StatusCode.UNIMPLEMENTED, + details=( + f"API {method.name!r} is batchable; " + "adaptive batching is not supported over gRPC for @bentoml.service() yet" + ), + ) + return None + if method.is_task: + await context.abort( + code=grpc.StatusCode.UNIMPLEMENTED, + details=( + f"API {method.name!r} is a task endpoint; " + "tasks are not supported over gRPC for @bentoml.service() yet" + ), + ) + return None + + field = request.WhichOneof("content") + input_data = await decode_proto( + method.input_spec, + field, + getattr(request, field) if field else None, ) - return None - - field = request.WhichOneof("content") - input_data = await decode_proto( - method.input_spec, field, getattr(request, field) if field else None - ) - call_args, call_kwargs = _call_args_from_input( - method, input_data, service.context - ) - func = getattr(self._get_instance(), method.name).local - original_func = get_original_func(func) - if is_async_callable(original_func) or inspect.iscoroutinefunction( - original_func - ): - output = await func(*call_args, **call_kwargs) - else: - output = await anyio.to_thread.run_sync( - functools.partial(func, *call_args, **call_kwargs) + call_args, call_kwargs = _call_args_from_input( + method, input_data, ctx ) - - field_name, encoded = await encode_proto(method.output_spec, output) - response = pb.Response(**{field_name: encoded}) + func = getattr(self._get_instance(), method.name).local + original_func = get_original_func(func) + if is_async_callable(original_func) or inspect.iscoroutinefunction( + original_func + ): + output = await func(*call_args, **call_kwargs) + else: + output = await anyio.to_thread.run_sync( + functools.partial(func, *call_args, **call_kwargs) + ) + + field_name, encoded = await encode_proto(method.output_spec, output) + response = pb.Response(**{field_name: encoded}) + _propagate_response_metadata(ctx, context) except BentoMLException as e: log_exception(request, sys.exc_info()) await context.abort(code=grpc_status_code(e), details=e.message) diff --git a/src/_bentoml_impl/server/serving.py b/src/_bentoml_impl/server/serving.py index 208afb85363..1ff4e3f844b 100644 --- a/src/_bentoml_impl/server/serving.py +++ b/src/_bentoml_impl/server/serving.py @@ -553,6 +553,7 @@ def serve_grpc( ) scheme = "https" if BentoMLContainer.ssl.enabled.get() else "http" close_child_stdin = not development_mode + bento_args = BentoMLContainer.bento_arguments.get() with contextlib.ExitStack() as port_stack: api_port = port_stack.enter_context( @@ -570,6 +571,8 @@ def serve_grpc( str(bento_path.absolute()), "--worker-id", "$(CIRCUS.WID)", + "--args", + json.dumps(bento_args), "--protocol-version", protocol_version, *ssl_args, diff --git a/src/_bentoml_impl/worker/grpc_service.py b/src/_bentoml_impl/worker/grpc_service.py index b5b6faca810..10e81343937 100644 --- a/src/_bentoml_impl/worker/grpc_service.py +++ b/src/_bentoml_impl/worker/grpc_service.py @@ -21,6 +21,7 @@ type=click.Path(exists=True), help="Working directory for the API server", ) +@click.option("--args", type=click.STRING, help="Bento arguments dict for the service") @click.option( "--worker-id", required=False, @@ -86,6 +87,7 @@ def main( port: int, runner_map: str | None, working_dir: str | None, + args: str | None, worker_id: int | None, enable_reflection: bool, enable_channelz: bool, @@ -106,6 +108,7 @@ def main( from bentoml._internal.configuration.containers import BentoMLContainer from bentoml._internal.context import server_context from bentoml._internal.log import configure_server_logging + from bentoml._internal.utils.args import set_arguments server_context.service_type = "grpc_api_server" server_context.worker_index = worker_id @@ -117,6 +120,8 @@ def main( BentoMLContainer.development_mode.set(development_mode) if runner_map is not None: BentoMLContainer.remote_runner_mapping.set(json.loads(runner_map)) + if args: + set_arguments(json.loads(args)) svc = load(bento_identifier, working_dir=working_dir) svc.inject_config() diff --git a/tests/e2e/bento_new_sdk_grpc/service.py b/tests/e2e/bento_new_sdk_grpc/service.py index 6a3edc6b433..766be0cb58e 100644 --- a/tests/e2e/bento_new_sdk_grpc/service.py +++ b/tests/e2e/bento_new_sdk_grpc/service.py @@ -4,6 +4,7 @@ import numpy as np import numpy.typing as npt +from pydantic import BaseModel import bentoml from bentoml.validators import TensorSchema @@ -11,16 +12,28 @@ Array = t.Annotated[npt.NDArray[np.float32], TensorSchema("numpy-array")] +class ServiceArgs(BaseModel): + greeting: str + + +args = bentoml.use_arguments(ServiceArgs) + + @bentoml.service class MyService: @bentoml.api def greet(self, name: str) -> str: - return f"hello {name}" + return f"{args.greeting} {name}" @bentoml.api def predict(self, x: Array) -> Array: return x * 2 + @bentoml.api + def context_greet(self, name: str, ctx: bentoml.Context) -> str: + ctx.response.headers["x-response-source"] = "bentoml-context" + return f"{ctx.request.headers['x-request-source']} {name}" + @bentoml.api async def stream_greet(self, name: str) -> t.AsyncGenerator[str, None]: yield f"hello {name}" diff --git a/tests/e2e/bento_new_sdk_grpc/test_grpc.py b/tests/e2e/bento_new_sdk_grpc/test_grpc.py index 928bab9585a..cc12c2f3bbd 100644 --- a/tests/e2e/bento_new_sdk_grpc/test_grpc.py +++ b/tests/e2e/bento_new_sdk_grpc/test_grpc.py @@ -11,6 +11,7 @@ from google.protobuf import wrappers_pb2 import bentoml +from bentoml.grpc.utils import import_generated_stubs from bentoml.grpc.utils import import_grpc from bentoml.testing.grpc import async_client_call from bentoml.testing.grpc import create_channel @@ -19,6 +20,7 @@ pytest.importorskip("grpc") grpc, aio = import_grpc() +pb, _ = import_generated_stubs("v1") PROJECT_DIR = Path(__file__).parent PORT = 38765 @@ -57,6 +59,7 @@ async def test_unary_greet_and_predict(monkeypatch: pytest.MonkeyPatch) -> None: port=PORT, server_type="grpc", production=False, + args={"greeting": "hello"}, ) as server: host_url = _host_url(server.url) await _wait_until_ready(host_url, timeout=100) @@ -78,6 +81,23 @@ async def test_unary_greet_and_predict(monkeypatch: pytest.MonkeyPatch) -> None: resp.ndarray.float_values, [2.0, 4.0, 6.0] ), ) + call_rpc = channel.unary_unary( + "/bentoml.grpc.v1.BentoService/Call", + request_serializer=pb.Request.SerializeToString, + response_deserializer=pb.Response.FromString, + ) + context_call = call_rpc( + pb.Request( + api_name="context_greet", + text=wrappers_pb2.StringValue(value="world"), + ), + metadata=(("x-request-source", "grpc-client"),), + ) + context_response = await context_call + assert context_response.text.value == "grpc-client world" + assert dict(await context_call.trailing_metadata())[ + "x-response-source" + ] == "bentoml-context" await async_client_call( "missing", channel=channel, diff --git a/tests/unit/_bentoml_impl/grpc/test_codec.py b/tests/unit/_bentoml_impl/grpc/test_codec.py index 9f5c12b5e1b..989996b27b7 100644 --- a/tests/unit/_bentoml_impl/grpc/test_codec.py +++ b/tests/unit/_bentoml_impl/grpc/test_codec.py @@ -16,6 +16,7 @@ from _bentoml_sdk.validators import TensorSchema from bentoml.exceptions import InvalidArgument from bentoml.grpc.utils import import_generated_stubs +from starlette.requests import Request pb, _ = import_generated_stubs("v1") @@ -114,6 +115,8 @@ async def test_ndarray_round_trip(): @pytest.mark.asyncio async def test_file_round_trip(tmp_path: Path): + from bentoml._internal.context import ServiceContext + spec = _input_of(echo_file) payload = b"grpc-file-bytes" source = tmp_path / "input.bin" @@ -123,8 +126,14 @@ async def test_file_round_trip(tmp_path: Path): assert encoded_field == "file" assert encoded.content == payload - decoded = await decode_proto(spec, "file", encoded) - assert Path(decoded.data).read_bytes() == payload + context = ServiceContext() + request = Request({"type": "http", "headers": [], "state": {}}) + with context.in_request(request): + decoded = await decode_proto(spec, "file", encoded) + decoded_path = Path(decoded.data) + assert decoded_path.read_bytes() == payload + + assert not decoded_path.exists() @pytest.mark.asyncio diff --git a/tests/unit/_bentoml_impl/grpc/test_servicer.py b/tests/unit/_bentoml_impl/grpc/test_servicer.py index dc0d31edfdd..196885e5172 100644 --- a/tests/unit/_bentoml_impl/grpc/test_servicer.py +++ b/tests/unit/_bentoml_impl/grpc/test_servicer.py @@ -1,6 +1,7 @@ from __future__ import annotations import typing as t +from pathlib import Path import pytest from google.protobuf import wrappers_pb2 @@ -14,6 +15,8 @@ pb, _ = import_generated_stubs("v1") grpc, aio = import_grpc() +seen_request_files: list[Path] = [] + @bentoml.service class Greeter: @@ -37,11 +40,44 @@ def batch_greet(self, name: list[str]) -> list[str]: def long_job(self, name: str) -> str: return name + @bentoml.api + def context_greet(self, name: str, ctx: bentoml.Context) -> str: + temp_file = Path(ctx.temp_dir) / "context.txt" + temp_file.write_text(name) + seen_request_files.append(temp_file) + ctx.response.headers["x-response-source"] = "bentoml-context" + return f"{ctx.request.headers['x-request-source']} {name}" + + @bentoml.api + def echo_file(self, data: Path) -> Path: + seen_request_files.append(data) + return data + + @bentoml.api + def fail_file(self, data: Path) -> Path: + seen_request_files.append(data) + raise RuntimeError("file handler failed") + class FakeContext: - def __init__(self) -> None: + def __init__( + self, invocation_metadata: tuple[tuple[str, str], ...] = () + ) -> None: self.code: grpc.StatusCode | None = None self.details: str | None = None + self._invocation_metadata = invocation_metadata + self._trailing_metadata: tuple[tuple[str, str], ...] = () + + def invocation_metadata(self) -> tuple[tuple[str, str], ...]: + return self._invocation_metadata + + def trailing_metadata(self) -> tuple[tuple[str, str], ...]: + return self._trailing_metadata + + def set_trailing_metadata( + self, metadata: tuple[tuple[str, str], ...] + ) -> None: + self._trailing_metadata = metadata async def abort(self, code: grpc.StatusCode, details: str = "") -> t.NoReturn: self.code = code @@ -76,6 +112,58 @@ async def test_call_async_greet(servicer): assert response.text.value == "hello world" +@pytest.mark.asyncio +async def test_call_establishes_request_context_and_propagates_metadata(servicer): + seen_request_files.clear() + ctx = FakeContext((("x-request-source", "grpc-client"),)) + request = pb.Request( + api_name="context_greet", + text=wrappers_pb2.StringValue(value="world"), + ) + + response = await servicer.Call(request, ctx) + + assert response is not None + assert response.text.value == "grpc-client world" + assert ("x-response-source", "bentoml-context") in ctx.trailing_metadata() + assert len(seen_request_files) == 1 + assert not seen_request_files[0].exists() + + +@pytest.mark.asyncio +async def test_call_removes_decoded_file_after_response_encoding(servicer): + seen_request_files.clear() + ctx = FakeContext() + request = pb.Request( + api_name="echo_file", + file=pb.File(kind="application/octet-stream", content=b"grpc-file-bytes"), + ) + + response = await servicer.Call(request, ctx) + + assert response is not None + assert response.file.content == b"grpc-file-bytes" + assert len(seen_request_files) == 1 + assert not seen_request_files[0].exists() + + +@pytest.mark.asyncio +async def test_call_removes_decoded_file_after_handler_error(servicer): + seen_request_files.clear() + ctx = FakeContext() + request = pb.Request( + api_name="fail_file", + file=pb.File(kind="application/octet-stream", content=b"grpc-file-bytes"), + ) + + with pytest.raises(aio.AbortError): + await servicer.Call(request, ctx) + + assert ctx.code == grpc.StatusCode.INTERNAL + assert len(seen_request_files) == 1 + assert not seen_request_files[0].exists() + + @pytest.mark.asyncio async def test_unknown_api_name_aborts(servicer): ctx = FakeContext() diff --git a/tests/unit/_bentoml_impl/grpc/test_worker.py b/tests/unit/_bentoml_impl/grpc/test_worker.py new file mode 100644 index 00000000000..0182bf68963 --- /dev/null +++ b/tests/unit/_bentoml_impl/grpc/test_worker.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +import pytest + + +def test_worker_sets_arguments_before_loading_service( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from _bentoml_impl.worker.grpc_service import main + from bentoml._internal.configuration.containers import BentoMLContainer + + class LoadStopped(Exception): + pass + + def fake_load(*args: object, **kwargs: object) -> None: + assert BentoMLContainer.bento_arguments.get() == {"greeting": "hello"} + raise LoadStopped + + monkeypatch.setattr("_bentoml_impl.loader.load", fake_load) + monkeypatch.setattr( + "bentoml._internal.log.configure_server_logging", lambda: None + ) + try: + with pytest.raises(LoadStopped): + main.main( + ["service.py:MyService", "--args", '{"greeting": "hello"}'], + standalone_mode=False, + ) + finally: + BentoMLContainer.bento_arguments.reset()