Skip to content
Draft
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
15 changes: 11 additions & 4 deletions src/litserve/loops/streaming_loops.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,19 +82,21 @@ def run_streaming_loop(
context = {}
if hasattr(lit_spec, "populate_context"):
lit_spec.populate_context(context, x_enc)

callback_runner.trigger_event(EventTypes.BEFORE_DECODE_REQUEST.value, lit_api=lit_api)
x = _inject_context(
context,
lit_api.decode_request,
x_enc,
)
callback_runner.trigger_event(EventTypes.AFTER_DECODE_REQUEST.value, lit_api=lit_api)

callback_runner.trigger_event(EventTypes.BEFORE_PREDICT.value, lit_api=lit_api)
y_gen = _inject_context(
context,
lit_api.predict,
x,
)
callback_runner.trigger_event(EventTypes.AFTER_PREDICT.value, lit_api=lit_api)

callback_runner.trigger_event(EventTypes.BEFORE_ENCODE_RESPONSE.value, lit_api=lit_api)
y_enc_gen = _inject_context(
Expand All @@ -111,6 +113,8 @@ def run_streaming_loop(
transport, response_queue_id, uid, "", LitAPIStatus.FINISH_STREAMING, LoopResponseType.STREAMING
)

# predict/encode_response are lazy generators: the work happens while the stream is
# consumed above, so both "after" events only become meaningful once it is exhausted.
callback_runner.trigger_event(EventTypes.AFTER_PREDICT.value, lit_api=lit_api)
callback_runner.trigger_event(EventTypes.AFTER_ENCODE_RESPONSE.value, lit_api=lit_api)

Expand Down Expand Up @@ -163,7 +167,6 @@ async def _process_streaming_request(
lit_api.predict,
x,
)
callback_runner.trigger_event(EventTypes.AFTER_PREDICT.value, lit_api=lit_api)

callback_runner.trigger_event(EventTypes.BEFORE_ENCODE_RESPONSE.value, lit_api=lit_api)

Expand All @@ -184,6 +187,8 @@ async def _process_streaming_request(
self.put_response(
transport, response_queue_id, uid, "", LitAPIStatus.FINISH_STREAMING, LoopResponseType.STREAMING
)
# See run_streaming_loop: the async generators only do their work as the stream is consumed.
callback_runner.trigger_event(EventTypes.AFTER_PREDICT.value, lit_api=lit_api)
callback_runner.trigger_event(EventTypes.AFTER_ENCODE_RESPONSE.value, lit_api=lit_api)

except HTTPException as e:
Expand Down Expand Up @@ -351,13 +356,11 @@ def run_batched_streaming_loop(

callback_runner.trigger_event(EventTypes.BEFORE_PREDICT.value, lit_api=lit_api)
y_iter = _inject_context(contexts, lit_api.predict, x)
callback_runner.trigger_event(EventTypes.AFTER_PREDICT.value, lit_api=lit_api)

unbatched_iter = _inject_context(contexts, lit_api.unbatch, y_iter)

callback_runner.trigger_event(EventTypes.BEFORE_ENCODE_RESPONSE.value, lit_api=lit_api)
y_enc_iter = _inject_context(contexts, lit_api.encode_response, unbatched_iter)
callback_runner.trigger_event(EventTypes.AFTER_ENCODE_RESPONSE.value, lit_api=lit_api)

# y_enc_iter -> [[response-1, response-2], [response-1, response-2]]
for y_batch in y_enc_iter:
Expand All @@ -371,6 +374,10 @@ def run_batched_streaming_loop(
self.put_response(
transport, response_queue_id, uid, "", LitAPIStatus.FINISH_STREAMING, LoopResponseType.STREAMING
)

# See run_streaming_loop: predict/unbatch/encode_response are lazy here too.
callback_runner.trigger_event(EventTypes.AFTER_PREDICT.value, lit_api=lit_api)
callback_runner.trigger_event(EventTypes.AFTER_ENCODE_RESPONSE.value, lit_api=lit_api)
except KeyboardInterrupt: # pragma: no cover
self.kill()
return
Expand Down
133 changes: 132 additions & 1 deletion tests/unit/test_loops.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@

import litserve as ls
from litserve import LitAPI
from litserve.callbacks import CallbackRunner
from litserve.callbacks import CallbackRunner, EventTypes
from litserve.loops import BatchedStreamingLoop, LitLoop, Output, StreamingLoop, inference_worker
from litserve.loops.base import (
_SENTINEL_VALUE,
Expand Down Expand Up @@ -600,6 +600,137 @@ async def test_run_streaming_loop(mock_transport):
assert response == {"output": f"{i}: Hello"}


# predict/encode_response return generators, so no work happens until the stream is consumed. Both
# "after" events therefore land together at the end, once the generators are exhausted.
_EXPECTED_STREAMING_EVENTS = [
EventTypes.BEFORE_DECODE_REQUEST.value,
EventTypes.AFTER_DECODE_REQUEST.value,
EventTypes.BEFORE_PREDICT.value,
EventTypes.BEFORE_ENCODE_RESPONSE.value,
EventTypes.AFTER_PREDICT.value,
EventTypes.AFTER_ENCODE_RESPONSE.value,
]


class _EventRecorder(ls.Callback):
"""Records every lifecycle event fired by the loop, in order."""

def __init__(self):
self.events = []

def on_before_decode_request(self, *args, **kwargs):
self.events.append(EventTypes.BEFORE_DECODE_REQUEST.value)

def on_after_decode_request(self, *args, **kwargs):
self.events.append(EventTypes.AFTER_DECODE_REQUEST.value)

def on_before_predict(self, *args, **kwargs):
self.events.append(EventTypes.BEFORE_PREDICT.value)

def on_after_predict(self, *args, **kwargs):
self.events.append(EventTypes.AFTER_PREDICT.value)

def on_before_encode_response(self, *args, **kwargs):
self.events.append(EventTypes.BEFORE_ENCODE_RESPONSE.value)

def on_after_encode_response(self, *args, **kwargs):
self.events.append(EventTypes.AFTER_ENCODE_RESPONSE.value)


@pytest.mark.asyncio
async def test_run_streaming_loop_triggers_each_callback_event_once(mock_transport):
"""The sync streaming loop must fire the same event sequence as the other loops, exactly once each."""
lit_api = ls.test_examples.SimpleStreamAPI()
lit_api.setup(None)
lit_api.request_timeout = 1

recorder = _EventRecorder()
cb_runner = CallbackRunner()
cb_runner._add_callbacks(recorder)

request_queue = Queue()
request_queue.put((0, "UUID-001", time.monotonic(), {"input": "Hello"}))

lit_loop = StreamingLoop()
lit_loop._restart_workers = True
loop_thread = threading.Thread(
target=lit_loop.run_streaming_loop, args=(lit_api, request_queue, mock_transport, cb_runner)
)
loop_thread.start()
time.sleep(1)
request_queue.put(_SENTINEL_VALUE)
loop_thread.join()

assert recorder.events == _EXPECTED_STREAMING_EVENTS


def test_run_streaming_loop_async_triggers_each_callback_event_once(mock_transport, monkeypatch):
"""The async streaming path must fire the same sequence as the sync one."""
recorder = _EventRecorder()
cb_runner = CallbackRunner()
cb_runner._add_callbacks(recorder)

requests_queue = TestQueue()
requests_queue.put((0, "uuid-123", time.monotonic(), {"input": 5}))
requests_queue.put(_SENTINEL_VALUE)

loop = StreamingLoop()
loop._restart_workers = True
monkeypatch.setattr(loop, "kill", lambda: None)

with contextlib.suppress(KeyboardInterrupt):
loop.run_streaming_loop_async(AsyncTestStreamLitAPI(), requests_queue, mock_transport, cb_runner)

assert recorder.events == _EXPECTED_STREAMING_EVENTS


@pytest.mark.asyncio
async def test_run_streaming_loop_after_events_fire_only_once_stream_is_consumed(mock_transport):
"""AFTER_PREDICT must reflect real prediction time, not the instant the generator was created."""

class SlowStreamAPI(ls.test_examples.SimpleStreamAPI):
def predict(self, x):
for i in range(3):
time.sleep(0.1)
yield f"{i}: {x}"

lit_api = SlowStreamAPI()
lit_api.setup(None)
lit_api.request_timeout = 10

before_predict = []
after_predict = []

class TimingRecorder(ls.Callback):
def on_before_predict(self, *args, **kwargs):
before_predict.append(time.perf_counter())

def on_after_predict(self, *args, **kwargs):
after_predict.append(time.perf_counter())

cb_runner = CallbackRunner()
cb_runner._add_callbacks(TimingRecorder())

request_queue = Queue()
request_queue.put((0, "UUID-001", time.monotonic(), {"input": "Hello"}))

lit_loop = StreamingLoop()
lit_loop._restart_workers = True
loop_thread = threading.Thread(
target=lit_loop.run_streaming_loop, args=(lit_api, request_queue, mock_transport, cb_runner)
)
loop_thread.start()
time.sleep(1)
request_queue.put(_SENTINEL_VALUE)
loop_thread.join()

assert len(before_predict) == 1, f"BEFORE_PREDICT fired {len(before_predict)} times, expected 1"
assert len(after_predict) == 1, f"AFTER_PREDICT fired {len(after_predict)} times, expected 1"

elapsed = after_predict[0] - before_predict[0]
assert elapsed >= 0.3, f"AFTER_PREDICT fired before the stream was consumed (measured {elapsed:.3f}s)"


@pytest.mark.asyncio
async def test_run_streaming_loop_timeout(mock_transport):
stream = io.StringIO()
Expand Down
Loading