Skip to content
Merged
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
2 changes: 1 addition & 1 deletion .github/workflows/CI.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion skywalking/plugins/sw_aiohttp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
87 changes: 86 additions & 1 deletion tests/unit/test_http_reporter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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()
Loading