Added route handler draining during shutdown to promptly close streaming connections.
CI / Formatting (push) Successful in 6s
CI / Linting (push) Successful in 4s
CI / Tests (Python 3.12) (push) Successful in 20s
CI / Tests (Python 3.13) (push) Successful in 53s
CI / Type Checking (push) Has been cancelled
CI / Spelling (push) Has been cancelled
CI / Tests (Python 3.14) (push) Has been cancelled

This commit is contained in:
2026-04-04 10:53:36 -04:00
parent 5425241f74
commit 7e44496baa
3 changed files with 222 additions and 1 deletions
+5
View File
@@ -80,6 +80,7 @@ class HttpServer:
self._runner: web.AppRunner | None = None self._runner: web.AppRunner | None = None
self.app = web.Application(middlewares=[_server_header_middleware]) self.app = web.Application(middlewares=[_server_header_middleware])
self.app.on_shutdown.append(self._on_shutdown)
self.app.router.add_post(self.config.webhook_path, self._handle_webhook) self.app.router.add_post(self.config.webhook_path, self._handle_webhook)
# Static assets served from the owlbot package. # Static assets served from the owlbot package.
@@ -135,6 +136,10 @@ class HttpServer:
await asyncio.gather(*self._pending_tasks) await asyncio.gather(*self._pending_tasks)
logger.debug("All pending events drained.") logger.debug("All pending events drained.")
async def _on_shutdown(self, _app: web.Application) -> None:
"""Drain active route handlers so connections close promptly."""
await self._route_dispatcher.drain_handlers()
async def _handle_webhook(self, request: web.Request) -> web.Response: async def _handle_webhook(self, request: web.Request) -> web.Response:
"""Handle incoming webhook requests from Owncast.""" """Handle incoming webhook requests from Owncast."""
try: try:
+45 -1
View File
@@ -24,7 +24,7 @@ import logging
import time import time
from collections import defaultdict from collections import defaultdict
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, cast, overload from typing import TYPE_CHECKING, Any, cast, overload
from aiohttp import web from aiohttp import web
from aiohttp.web import DynamicResource from aiohttp.web import DynamicResource
@@ -382,6 +382,8 @@ class RouteDispatcher:
self._route_registry = RouteRegistry() self._route_registry = RouteRegistry()
self._get_module_context = get_module_context self._get_module_context = get_module_context
self._handler_timeout = handler_timeout self._handler_timeout = handler_timeout
self._streaming_tasks: set[asyncio.Task[Any]] = set()
self._handler_tasks: set[asyncio.Task[Any]] = set()
def register( def register(
self, self,
@@ -474,6 +476,37 @@ class RouteDispatcher:
""" """
return self._route_registry.unregister_by_module(module_name) return self._route_registry.unregister_by_module(module_name)
async def drain_handlers(self) -> None:
"""Drain all active route handlers during shutdown.
Waits for in-flight non-streaming handlers to complete, then
cancels long-lived streaming handlers so their connections
close promptly instead of blocking until aiohttp's shutdown
timeout expires.
Intended to be called via ``app.on_shutdown``.
"""
if self._handler_tasks:
logger.info(
"Waiting for %d non-streaming handler(s) to complete...",
len(self._handler_tasks),
)
await asyncio.gather(*list(self._handler_tasks), return_exceptions=True)
logger.debug("All non-streaming handlers completed.")
if not self._streaming_tasks:
return
logger.info(
"Cancelling %d active streaming handler(s)...",
len(self._streaming_tasks),
)
tasks = list(self._streaming_tasks)
for task in tasks:
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
logger.debug("All streaming handlers cancelled.")
async def dispatch(self, request: web.Request) -> web.StreamResponse: async def dispatch(self, request: web.Request) -> web.StreamResponse:
"""Dispatch an HTTP request to the appropriate module route handler. """Dispatch an HTTP request to the appropriate module route handler.
@@ -543,6 +576,13 @@ class RouteDispatcher:
module_name, module_name,
) )
task = asyncio.current_task()
if task is not None:
if route_info.streaming:
self._streaming_tasks.add(task)
else:
self._handler_tasks.add(task)
try: try:
start = time.perf_counter() start = time.perf_counter()
if route_info.streaming: if route_info.streaming:
@@ -587,6 +627,10 @@ class RouteDispatcher:
f"Route handler '{route_info.full_path}' raised exception: {e}" f"Route handler '{route_info.full_path}' raised exception: {e}"
) )
return web.Response(status=500) return web.Response(status=500)
finally:
if task is not None:
self._streaming_tasks.discard(task)
self._handler_tasks.discard(task)
class ModuleRoutes: class ModuleRoutes:
+172
View File
@@ -698,6 +698,178 @@ class TestRouteDispatcherDispatch:
assert response.status == 200 assert response.status == 200
class TestRouteDispatcherDrainHandlers:
"""Tests RouteDispatcher.drain_handlers() from owlbot.registries.routes."""
async def test_drain_handlers_cancels_active_handler(self) -> None:
"""drain_handlers() cancels a running streaming handler."""
entered = asyncio.Event()
async def forever_stream(ctx: RouteContext) -> web.Response:
entered.set()
await asyncio.sleep(3600)
return web.Response(text="done")
dispatcher = _make_dispatcher()
dispatcher.register(
"/stream", forever_stream, module_name="mod", streaming=True
)
request = make_mocked_request("GET", "/owlbot/mod/stream")
request.match_info["module_name"] = "mod"
request.match_info["path"] = "stream"
# Run dispatch in a task so we can cancel it from the outside.
task = asyncio.create_task(dispatcher.dispatch(request))
await entered.wait()
assert len(dispatcher._streaming_tasks) == 1
await dispatcher.drain_handlers()
assert len(dispatcher._streaming_tasks) == 0
# The dispatch task should now be done (cancelled).
with pytest.raises(asyncio.CancelledError):
await task
async def test_drain_handlers_noop_when_empty(self) -> None:
"""drain_handlers() with no active handlers completes immediately."""
dispatcher = _make_dispatcher()
await dispatcher.drain_handlers()
assert len(dispatcher._streaming_tasks) == 0
async def test_streaming_task_cleaned_up_on_normal_return(self) -> None:
"""Streaming task is removed from tracking after normal completion."""
async def quick_stream(ctx: RouteContext) -> dict[str, str]:
return {"ok": "true"}
dispatcher = _make_dispatcher()
dispatcher.register("/quick", quick_stream, module_name="mod", streaming=True)
request = make_mocked_request("GET", "/owlbot/mod/quick")
request.match_info["module_name"] = "mod"
request.match_info["path"] = "quick"
await dispatcher.dispatch(request)
assert len(dispatcher._streaming_tasks) == 0
async def test_streaming_task_cleaned_up_on_exception(self) -> None:
"""Streaming task is removed from tracking after handler exception."""
async def bad_stream(ctx: RouteContext) -> web.Response:
raise RuntimeError("boom")
dispatcher = _make_dispatcher()
dispatcher.register("/bad", bad_stream, module_name="mod", streaming=True)
request = make_mocked_request("GET", "/owlbot/mod/bad")
request.match_info["module_name"] = "mod"
request.match_info["path"] = "bad"
await dispatcher.dispatch(request)
assert len(dispatcher._streaming_tasks) == 0
async def test_drain_handlers_multiple_handlers(self) -> None:
"""drain_handlers() cancels all active streaming handlers."""
entered = asyncio.Event()
count = 0
async def forever_stream(ctx: RouteContext) -> web.Response:
nonlocal count
count += 1
if count == 2:
entered.set()
await asyncio.sleep(3600)
return web.Response(text="done")
dispatcher = _make_dispatcher()
dispatcher.register("/s1", forever_stream, module_name="mod", streaming=True)
dispatcher.register("/s2", forever_stream, module_name="mod", streaming=True)
request1 = make_mocked_request("GET", "/owlbot/mod/s1")
request1.match_info["module_name"] = "mod"
request1.match_info["path"] = "s1"
request2 = make_mocked_request("GET", "/owlbot/mod/s2")
request2.match_info["module_name"] = "mod"
request2.match_info["path"] = "s2"
task1 = asyncio.create_task(dispatcher.dispatch(request1))
task2 = asyncio.create_task(dispatcher.dispatch(request2))
await entered.wait()
assert len(dispatcher._streaming_tasks) == 2
await dispatcher.drain_handlers()
assert len(dispatcher._streaming_tasks) == 0
with pytest.raises(asyncio.CancelledError):
await task1
with pytest.raises(asyncio.CancelledError):
await task2
async def test_nonstreaming_task_tracked_and_cleaned_up(self) -> None:
"""Non-streaming handler tasks are tracked and cleaned up."""
async def handler(ctx: RouteContext) -> dict[str, str]:
return {"ok": "true"}
dispatcher = _make_dispatcher()
dispatcher.register("/api", handler, module_name="mod")
request = make_mocked_request("GET", "/owlbot/mod/api")
request.match_info["module_name"] = "mod"
request.match_info["path"] = "api"
await dispatcher.dispatch(request)
assert len(dispatcher._handler_tasks) == 0
async def test_drain_handlers_waits_for_nonstreaming_first(self) -> None:
"""drain_handlers() waits for non-streaming handlers before cancelling."""
order: list[str] = []
handler_entered = asyncio.Event()
stream_entered = asyncio.Event()
async def slow_handler(ctx: RouteContext) -> dict[str, str]:
handler_entered.set()
await asyncio.sleep(0.15)
order.append("handler_done")
return {"ok": "true"}
async def forever_stream(ctx: RouteContext) -> web.Response:
stream_entered.set()
try:
await asyncio.sleep(3600)
except asyncio.CancelledError:
order.append("stream_cancelled")
raise
return web.Response(text="done")
dispatcher = _make_dispatcher(handler_timeout=5.0)
dispatcher.register("/api", slow_handler, module_name="mod")
dispatcher.register(
"/stream", forever_stream, module_name="mod", streaming=True
)
req_handler = make_mocked_request("GET", "/owlbot/mod/api")
req_handler.match_info["module_name"] = "mod"
req_handler.match_info["path"] = "api"
req_stream = make_mocked_request("GET", "/owlbot/mod/stream")
req_stream.match_info["module_name"] = "mod"
req_stream.match_info["path"] = "stream"
handler_task = asyncio.create_task(dispatcher.dispatch(req_handler))
stream_task = asyncio.create_task(dispatcher.dispatch(req_stream))
await handler_entered.wait()
await stream_entered.wait()
await dispatcher.drain_handlers()
assert order == ["handler_done", "stream_cancelled"]
await handler_task
with pytest.raises(asyncio.CancelledError):
await stream_task
class TestRouteDispatcherDelegation: class TestRouteDispatcherDelegation:
"""Tests RouteDispatcher delegation methods from owlbot.registries.routes.""" """Tests RouteDispatcher delegation methods from owlbot.registries.routes."""