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
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:
@@ -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:
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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."""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user