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
61 changes: 57 additions & 4 deletions backend/python/common/model_identity.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@
the model itself.
"""

import asyncio
import inspect
import threading

import grpc
Expand Down Expand Up @@ -181,6 +183,40 @@ def guard(request, context):
return _rebuild(handler, guard)


_STREAM_DONE = object()


def _next_or_done(iterator):
"""next(iterator), returning the _STREAM_DONE sentinel at exhaustion.

StopIteration must not propagate out of a function run via run_in_executor:
it cannot travel through a Future and would surface as an opaque error.
"""
try:
return next(iterator)
except StopIteration:
return _STREAM_DONE


async def _call_behavior(behavior, request, context):
"""Invoke a unary servicer behavior without blocking the event loop.

Native async behavior is awaited directly. A sync behavior -- many backends
define `def LoadModel` / `def Embedding`, not `async def` -- is dispatched to
a worker thread so a slow load/inference cannot freeze all aio RPC handling,
mirroring grpc.aio's own sync-handler adaptation. A callable wrapper that
returns an awaitable is supported too: the (cheap) call runs in the thread,
then the awaitable is awaited back on the loop.
"""
if inspect.iscoroutinefunction(behavior):
return await behavior(request, context)
loop = asyncio.get_running_loop()
result = await loop.run_in_executor(None, behavior, request, context)
if inspect.isawaitable(result):
result = await result
return result


class AsyncModelIdentityInterceptor(grpc.aio.ServerInterceptor):
"""Async counterpart for backends running grpc.aio servers."""

Expand All @@ -200,7 +236,11 @@ async def intercept_service(self, continuation, handler_call_details):
original = handler.unary_unary

async def record(request, context):
result = await original(request, context)
# A backend's LoadModel may be a plain sync method (many define
# `def LoadModel`, not `async def`). Dispatch it so it neither
# crashes with "object <T> can't be used in 'await'" nor runs its
# (potentially slow) body on the event loop thread.
result = await _call_behavior(original, request, context)
if getattr(result, "success", True):
self.state.record(getattr(request, "Model", ""))
return result
Expand All @@ -214,8 +254,21 @@ async def guard_stream(request, context):
message = self.state.mismatch(getattr(request, "ModelIdentity", ""))
if message is not None:
await context.abort(grpc.StatusCode.NOT_FOUND, message)
async for response in original_stream(request, context):
yield response
# A sync backend yields a plain generator, an async one an async
# generator. Async: iterate directly. Sync: pull each item via a
# worker thread so a slow producer doesn't block the event loop
# (and so StopIteration can't escape through a Future).
stream = original_stream(request, context)
if hasattr(stream, "__aiter__"):
async for response in stream:
yield response
else:
loop = asyncio.get_running_loop()
while True:
item = await loop.run_in_executor(None, _next_or_done, stream)
if item is _STREAM_DONE:
break
yield item

return _rebuild(handler, guard_stream)

Expand All @@ -225,6 +278,6 @@ async def guard(request, context):
message = self.state.mismatch(getattr(request, "ModelIdentity", ""))
if message is not None:
await context.abort(grpc.StatusCode.NOT_FOUND, message)
return await original_unary(request, context)
return await _call_behavior(original_unary, request, context)

return _rebuild(handler, guard)
180 changes: 180 additions & 0 deletions backend/python/common/model_identity_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@
enforcement that rejects requests it should serve.
"""

import asyncio
import os
import threading
import unittest

import grpc
Expand Down Expand Up @@ -64,6 +66,13 @@ def _handler(behavior, response_streaming=False):
return grpc.unary_unary_rpc_method_handler(behavior)


def _const_continuation(handler):
async def continuation(_):
return handler

return continuation


class TestInterceptorInstalled(unittest.TestCase):
"""The wiring, which is where this can silently do nothing.

Expand Down Expand Up @@ -304,5 +313,176 @@ def test_codec_rpcs_stay_unguarded(self):
self.assertNotIn(method, model_identity._GUARDED_METHODS)


class TestAsyncInterceptorBehavior(unittest.TestCase):
"""The grpc.aio counterpart, which had no behavioral coverage.

AsyncModelIdentityInterceptor wraps a backend's own servicer behavior. That
behavior may be sync or async: several backends define `def LoadModel` and
`def Embedding` (not `async def`), and grpc.aio's dispatch adapts both. The
interceptor invokes the behavior itself, so if it awaits unconditionally it
breaks every sync method it guards with "object <T> can't be used in
'await'". These tests exercise both shapes; the sync ones are the regression.
"""

def setUp(self):
self.interceptor = model_identity.AsyncModelIdentityInterceptor()

def _wrap(self, method, handler):
async def continuation(_):
return handler

return asyncio.run(
self.interceptor.intercept_service(continuation, _FakeCallDetails(method))
)

def _load(self, behavior, model):
wrapped = self._wrap("/backend.Backend/LoadModel", _handler(behavior))
return asyncio.run(wrapped.unary_unary(_Request(Model=model), _FakeContext()))

def _call_unary(self, behavior, identity):
wrapped = self._wrap("/backend.Backend/Predict", _handler(behavior))
return asyncio.run(
wrapped.unary_unary(_Request(ModelIdentity=identity), _FakeContext())
)

def _drain_stream(self, behavior, identity):
wrapped = self._wrap(
"/backend.Backend/PredictStream", _handler(behavior, response_streaming=True)
)

async def drain():
out = []
async for item in wrapped.unary_stream(
_Request(ModelIdentity=identity), _FakeContext()
):
out.append(item)
return out

return asyncio.run(drain())

# --- LoadModel: sync behavior is the regression, async must still work ---

def test_load_records_with_sync_behavior(self):
self._load(lambda request, context: _Result(), "a.gguf")
self.assertEqual(self.interceptor.state.loaded, "a.gguf")

def test_load_records_with_async_behavior(self):
async def behavior(request, context):
return _Result()

self._load(behavior, "a.gguf")
self.assertEqual(self.interceptor.state.loaded, "a.gguf")

def test_failed_sync_load_records_nothing(self):
self._load(lambda request, context: _Result(success=False), "a.gguf")
self.assertEqual(self.interceptor.state.loaded, "")

# --- guarded unary: sync and async behaviors both served / rejected ---

def test_guard_serves_sync_behavior(self):
self._load(lambda request, context: _Result(), "a.gguf")
result = self._call_unary(lambda request, context: "served", "a.gguf")
self.assertEqual(result, "served")

def test_guard_serves_async_behavior(self):
self._load(lambda request, context: _Result(), "a.gguf")

async def behavior(request, context):
return "served"

self.assertEqual(self._call_unary(behavior, "a.gguf"), "served")

def test_guard_rejects_mismatch(self):
self._load(lambda request, context: _Result(), "a.gguf")
with self.assertRaises(_Aborted):
self._call_unary(lambda request, context: "served", "b.gguf")

# --- guarded stream: sync generator and async generator both work ---

def test_guard_stream_serves_sync_generator(self):
self._load(lambda request, context: _Result(), "a.gguf")

def behavior(request, context):
yield "a"
yield "b"

self.assertEqual(self._drain_stream(behavior, "a.gguf"), ["a", "b"])

def test_guard_stream_serves_async_generator(self):
self._load(lambda request, context: _Result(), "a.gguf")

async def behavior(request, context):
yield "a"
yield "b"

self.assertEqual(self._drain_stream(behavior, "a.gguf"), ["a", "b"])

# --- sync behavior must not run on the event-loop thread ---
#
# Awaiting a sync method's return fixed the TypeError, but calling the
# (possibly slow) sync behavior on the event loop still froze all aio RPC
# handling. These record the thread each behavior runs on and assert it is a
# worker thread, not the loop thread.

def _run_capturing_loop_thread(self, method, handler, request):
captured = {}

async def run():
captured["loop"] = threading.get_ident()
wrapped = await self.interceptor.intercept_service(
_const_continuation(handler), _FakeCallDetails(method)
)
behavior = wrapped.unary_stream if handler.response_streaming else wrapped.unary_unary
if handler.response_streaming:
async for _ in behavior(request, _FakeContext()):
pass
else:
await behavior(request, _FakeContext())

asyncio.run(run())
return captured["loop"]

def test_sync_load_runs_off_the_event_loop(self):
ran = {}

def behavior(request, context):
ran["thread"] = threading.get_ident()
return _Result()

loop_thread = self._run_capturing_loop_thread(
"/backend.Backend/LoadModel", _handler(behavior), _Request(Model="a.gguf")
)
self.assertIn("thread", ran)
self.assertNotEqual(ran["thread"], loop_thread)

def test_sync_guarded_unary_runs_off_the_event_loop(self):
self._load(lambda request, context: _Result(), "a.gguf")
ran = {}

def behavior(request, context):
ran["thread"] = threading.get_ident()
return "served"

loop_thread = self._run_capturing_loop_thread(
"/backend.Backend/Predict", _handler(behavior), _Request(ModelIdentity="a.gguf")
)
self.assertNotEqual(ran["thread"], loop_thread)

def test_sync_stream_next_runs_off_the_event_loop(self):
self._load(lambda request, context: _Result(), "a.gguf")
ran = {}

def behavior(request, context):
ran["thread"] = threading.get_ident()
yield "a"

loop_thread = self._run_capturing_loop_thread(
"/backend.Backend/PredictStream",
_handler(behavior, response_streaming=True),
_Request(ModelIdentity="a.gguf"),
)
self.assertNotEqual(ran["thread"], loop_thread)


if __name__ == "__main__":
unittest.main()
Loading