diff --git a/.github/workflows/CI.yaml b/.github/workflows/CI.yaml index 5bcd070c0..bfc10e6a9 100644 --- a/.github/workflows/CI.yaml +++ b/.github/workflows/CI.yaml @@ -84,7 +84,7 @@ jobs: with: persist-credentials: false - name: Check for file changes - uses: dorny/paths-filter@fbd0ab8f3e69293af611ebaee6363fc25e6d187d # v4.0.1 + uses: dorny/paths-filter@ceb8a2b8f2d89434be7ff52d3de7ec3738c5cc9d # v4.0.3 id: filter with: # The following filters indicate a category along with diff --git a/skywalking/plugins/sw_aiohttp.py b/skywalking/plugins/sw_aiohttp.py index 65bf42319..b7f25c41c 100644 --- a/skywalking/plugins/sw_aiohttp.py +++ b/skywalking/plugins/sw_aiohttp.py @@ -85,7 +85,7 @@ async def _sw_handle_request(self, request: BaseRequest, start_time: float, *arg if config.agent_protocol == 'http' and config.agent_collector_backend_services.rstrip('/') \ .endswith(f'{request.url.host}:{request.url.port}'): - return _handle_request + return await _handle_request(self, request, start_time, *args, **kwargs) carrier = Carrier() method = request.method diff --git a/tests/unit/test_http_reporter.py b/tests/unit/test_http_reporter.py index 7c5f13fce..bffb9f3b3 100644 --- a/tests/unit/test_http_reporter.py +++ b/tests/unit/test_http_reporter.py @@ -19,7 +19,7 @@ import threading import unittest from http.server import BaseHTTPRequestHandler, HTTPServer -from unittest.mock import patch +from unittest.mock import AsyncMock, MagicMock, patch from skywalking import config @@ -150,5 +150,90 @@ async def run(): get_context.assert_not_called() +class TestAiohttpServerCollectorSkip(unittest.TestCase): + """RequestHandler skip must await the original handler, not return the function.""" + + def setUp(self): + self._saved = ( + config.agent_collector_backend_services, + config.agent_protocol, + ) + config.agent_protocol = 'http' + config.agent_collector_backend_services = '127.0.0.1:12800' + + def tearDown(self): + ( + config.agent_collector_backend_services, + config.agent_protocol, + ) = self._saved + + def test_collector_host_skip_awaits_original_handler(self): + from aiohttp import ClientSession + from aiohttp.web_protocol import RequestHandler + + from skywalking.plugins import sw_aiohttp + + self.addCleanup(setattr, ClientSession, '_request', ClientSession._request) + self.addCleanup(setattr, RequestHandler, '_handle_request', RequestHandler._handle_request) + + stub = AsyncMock(return_value=('resp', False)) + RequestHandler._handle_request = stub + sw_aiohttp.install() + + handler = MagicMock() + request = MagicMock() + request.url.host = '127.0.0.1' + request.url.port = 12800 + + async def run(): + with patch.object(sw_aiohttp, 'get_context') as get_context: + result = await RequestHandler._handle_request(handler, request, 1.5, extra=True) + get_context.assert_not_called() + return result + + result = asyncio.run(run()) + self.assertEqual(('resp', False), result) + stub.assert_awaited_once_with(handler, request, 1.5, extra=True) + + def test_non_collector_host_still_creates_entry_span(self): + from aiohttp import ClientSession + from aiohttp.web_protocol import RequestHandler + + from skywalking.plugins import sw_aiohttp + + self.addCleanup(setattr, ClientSession, '_request', ClientSession._request) + self.addCleanup(setattr, RequestHandler, '_handle_request', RequestHandler._handle_request) + + resp = MagicMock() + resp.status = 200 + stub = AsyncMock(return_value=(resp, False)) + RequestHandler._handle_request = stub + sw_aiohttp.install() + + handler = MagicMock() + request = MagicMock() + request.url.host = '127.0.0.1' + request.url.port = 9999 + request.method = 'GET' + request.path = '/users' + request.headers = {} + request._transport_peername = ('10.0.0.1', 1234) + request.scheme = 'http' + request.host = '127.0.0.1:9999' + + span = MagicMock() + context = MagicMock() + context.new_entry_span.return_value = span + + async def run(): + with patch.object(sw_aiohttp, 'get_context', return_value=context): + return await RequestHandler._handle_request(handler, request, 0.0) + + result = asyncio.run(run()) + self.assertEqual((resp, False), result) + context.new_entry_span.assert_called_once() + stub.assert_awaited_once() + + if __name__ == '__main__': unittest.main()