From 0cd9b15c7be24289cafe2bc1a9c3afb6320b2460 Mon Sep 17 00:00:00 2001 From: songzhendong12315 Date: Wed, 16 Sep 2026 09:43:37 +0800 Subject: [PATCH 1/2] fix: keep long-lived aiohttp ClientSession across HTTP reports Replace async-with on the shared ClientSession (closes it after the first request) with async-with on the response. Add aclose() for session cleanup. --- skywalking/agent/protocol/http_aio.py | 7 ++ skywalking/client/http_aio.py | 145 ++++++++++++++------------ tests/unit/test_http_reporter.py | 117 +++++++++++++++++++++ 3 files changed, 203 insertions(+), 66 deletions(-) create mode 100644 tests/unit/test_http_reporter.py diff --git a/skywalking/agent/protocol/http_aio.py b/skywalking/agent/protocol/http_aio.py index 3316d10ca..37c18ce38 100644 --- a/skywalking/agent/protocol/http_aio.py +++ b/skywalking/agent/protocol/http_aio.py @@ -32,6 +32,13 @@ def __init__(self): self.traces_reporter = HttpTraceSegmentReportServiceAsync() self.log_reporter = HttpLogDataReportServiceAsync() + async def aclose(self) -> None: + # Close long-lived aiohttp ClientSessions owned by the HTTP clients. + for part in (self.service_management, self.traces_reporter, self.log_reporter): + aclose = getattr(part, 'aclose', None) + if callable(aclose): + await aclose() + async def heartbeat(self): if not self.properties_sent.is_set(): logger.debug('Sending instance properties') diff --git a/skywalking/client/http_aio.py b/skywalking/client/http_aio.py index 248797ae9..72b7d490b 100644 --- a/skywalking/client/http_aio.py +++ b/skywalking/client/http_aio.py @@ -32,6 +32,16 @@ def _aiohttp_session(material=tls_mod._MATERIAL_UNSET): return aiohttp.ClientSession(connector=aiohttp.TCPConnector(ssl=ssl_ctx)) +async def _aclose_session(session) -> None: + """Close a long-lived ClientSession; never raise into agent shutdown.""" + if session is None or getattr(session, 'closed', True): + return + try: + await session.close() + except Exception: # noqa: BLE001 + pass + + class HttpServiceManagementClientAsync(ServiceManagementClientAsync): def __init__(self): super().__init__() @@ -42,19 +52,21 @@ def __init__(self): proto = collector_http_scheme(material) self.url_instance_props = f"{proto}{config.agent_collector_backend_services.rstrip('/')}/v3/management/reportProperties" self.url_heart_beat = f"{proto}{config.agent_collector_backend_services.rstrip('/')}/v3/management/keepAlive" - # self.client = httpx.AsyncClient() + # Long-lived session: never `async with self.client` (that closes the session). self.client = _aiohttp_session(material) - async def send_instance_props(self): + async def aclose(self) -> None: + await _aclose_session(self.client) - async with self.client as client: - res = await client.post(self.url_instance_props, json={ - 'service': config.agent_name, - 'serviceInstance': config.agent_instance_name, - 'properties': self.instance_properties, - }) - if logger_debug_enabled: - logger.debug('heartbeat response: %s', res) + async def send_instance_props(self): + # `async with session.post(...)` closes the response, not the session. + async with self.client.post(self.url_instance_props, json={ + 'service': config.agent_name, + 'serviceInstance': config.agent_instance_name, + 'properties': self.instance_properties, + }) as res: + if logger_debug_enabled: + logger.debug('heartbeat response: %s', res.status) async def send_heart_beat(self): await self.refresh_instance_props() @@ -65,13 +77,12 @@ async def send_heart_beat(self): config.agent_name, config.agent_instance_name, ) - async with self.client as client: - res = await client.post(self.url_heart_beat, json={ - 'service': config.agent_name, - 'serviceInstance': config.agent_instance_name, - }) - if logger_debug_enabled: - logger.debug('heartbeat response: %s', res) + async with self.client.post(self.url_heart_beat, json={ + 'service': config.agent_name, + 'serviceInstance': config.agent_instance_name, + }) as res: + if logger_debug_enabled: + logger.debug('heartbeat response: %s', res.status) class HttpTraceSegmentReportServiceAsync(TraceSegmentReportServiceAsync): @@ -79,54 +90,55 @@ def __init__(self): material = safe_tls_pem_material() proto = collector_http_scheme(material) self.url_report = f"{proto}{config.agent_collector_backend_services.rstrip('/')}/v3/segment" - # self.client = httpx.AsyncClient() self.client = _aiohttp_session(material) + async def aclose(self) -> None: + await _aclose_session(self.client) + async def report(self, generator): async for segment in generator: - async with self.client as client: - res = await client.post(self.url_report, json={ - 'traceId': str(segment.related_traces[0]), - 'traceSegmentId': str(segment.segment_id), - 'service': config.agent_name, - 'serviceInstance': config.agent_instance_name, - 'isSizeLimited': segment.is_size_limited, - 'spans': [{ - 'spanId': span.sid, - 'parentSpanId': span.pid, - 'startTime': span.start_time, - 'endTime': span.end_time, - 'operationName': span.op, - 'peer': span.peer, - 'spanType': span.kind.name, - 'spanLayer': span.layer.name, - 'componentId': span.component.value, - 'isError': span.error_occurred, - 'logs': [{ - 'time': int(log.timestamp * 1000), - 'data': [{ - 'key': item.key, - 'value': item.val, - } for item in log.items], - } for log in span.logs], - 'tags': [{ - 'key': tag.key, - 'value': tag.val, - } for tag in span.iter_tags()], - 'refs': [{ - 'refType': 0, - 'traceId': ref.trace_id, - 'parentTraceSegmentId': ref.segment_id, - 'parentSpanId': ref.span_id, - 'parentService': ref.service, - 'parentServiceInstance': ref.service_instance, - 'parentEndpoint': ref.endpoint, - 'networkAddressUsedAtPeer': ref.client_address, - } for ref in span.refs if ref.trace_id] - } for span in segment.spans] - }) - if logger_debug_enabled: - logger.debug('report traces response: %s', res) + async with self.client.post(self.url_report, json={ + 'traceId': str(segment.related_traces[0]), + 'traceSegmentId': str(segment.segment_id), + 'service': config.agent_name, + 'serviceInstance': config.agent_instance_name, + 'isSizeLimited': segment.is_size_limited, + 'spans': [{ + 'spanId': span.sid, + 'parentSpanId': span.pid, + 'startTime': span.start_time, + 'endTime': span.end_time, + 'operationName': span.op, + 'peer': span.peer, + 'spanType': span.kind.name, + 'spanLayer': span.layer.name, + 'componentId': span.component.value, + 'isError': span.error_occurred, + 'logs': [{ + 'time': int(log.timestamp * 1000), + 'data': [{ + 'key': item.key, + 'value': item.val, + } for item in log.items], + } for log in span.logs], + 'tags': [{ + 'key': tag.key, + 'value': tag.val, + } for tag in span.iter_tags()], + 'refs': [{ + 'refType': 0, + 'traceId': ref.trace_id, + 'parentTraceSegmentId': ref.segment_id, + 'parentSpanId': ref.span_id, + 'parentService': ref.service, + 'parentServiceInstance': ref.service_instance, + 'parentEndpoint': ref.endpoint, + 'networkAddressUsedAtPeer': ref.client_address, + } for ref in span.refs if ref.trace_id] + } for span in segment.spans] + }) as res: + if logger_debug_enabled: + logger.debug('report traces response: %s', res.status) class HttpLogDataReportServiceAsync(LogDataReportServiceAsync): @@ -134,13 +146,14 @@ def __init__(self): material = safe_tls_pem_material() proto = collector_http_scheme(material) self.url_report = f"{proto}{config.agent_collector_backend_services.rstrip('/')}/v3/logs" - # self.client = httpx.AsyncClient() self.client = _aiohttp_session(material) + async def aclose(self) -> None: + await _aclose_session(self.client) + async def report(self, generator): log_batch = [json.loads(json_format.MessageToJson(log_data)) async for log_data in generator] if log_batch: # prevent empty batches - async with self.client as client: - res = await client.post(self.url_report, json=log_batch) - if logger_debug_enabled: - logger.debug('report batch log response: %s', res) + async with self.client.post(self.url_report, json=log_batch) as res: + if logger_debug_enabled: + logger.debug('report batch log response: %s', res.status) diff --git a/tests/unit/test_http_reporter.py b/tests/unit/test_http_reporter.py new file mode 100644 index 000000000..28e31a672 --- /dev/null +++ b/tests/unit/test_http_reporter.py @@ -0,0 +1,117 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +import asyncio +import threading +import unittest +from http.server import BaseHTTPRequestHandler, HTTPServer + +from skywalking import config + + +class _OkHandler(BaseHTTPRequestHandler): + def do_POST(self): # noqa + length = int(self.headers.get('Content-Length') or 0) + self.rfile.read(length) + self.send_response(200) + self.send_header('Content-Type', 'application/json') + self.end_headers() + self.wfile.write(b'{}') + + def log_message(self, *_args): + pass + + +def _start_http_server(): + server = HTTPServer(('127.0.0.1', 0), _OkHandler) + threading.Thread(target=server.serve_forever, daemon=True).start() + return server, server.server_address[1] + + +class TestAsyncHttpClientSession(unittest.TestCase): + def setUp(self): + self._saved = ( + config.agent_collector_backend_services, + config.agent_protocol, + config.agent_name, + config.agent_instance_name, + ) + config.agent_protocol = 'http' + config.agent_name = 'test-service' + config.agent_instance_name = 'test-instance' + + def tearDown(self): + ( + config.agent_collector_backend_services, + config.agent_protocol, + config.agent_name, + config.agent_instance_name, + ) = self._saved + + def test_async_http_session_survives_repeated_posts(self): + """async with session.post closes the response, not the long-lived session.""" + server, port = _start_http_server() + config.agent_collector_backend_services = f'127.0.0.1:{port}' + try: + from skywalking.client.http_aio import HttpServiceManagementClientAsync + + async def run(): + client = HttpServiceManagementClientAsync() + self.assertFalse(client.client.closed) + await client.send_instance_props() + self.assertFalse( + client.client.closed, + 'ClientSession must stay open after the first request', + ) + await client.send_heart_beat() + self.assertFalse(client.client.closed) + await client.aclose() + self.assertTrue(client.client.closed) + + asyncio.run(run()) + finally: + server.shutdown() + + def test_async_http_segment_reporter_reuses_session(self): + server, port = _start_http_server() + config.agent_collector_backend_services = f'127.0.0.1:{port}' + try: + from skywalking.client.http_aio import HttpTraceSegmentReportServiceAsync + + class _Seg: + related_traces = ['t1'] + segment_id = 's1' + is_size_limited = False + spans = [] + + async def gen(): + yield _Seg() + yield _Seg() + + async def run(): + reporter = HttpTraceSegmentReportServiceAsync() + await reporter.report(gen()) + self.assertFalse(reporter.client.closed) + await reporter.aclose() + + asyncio.run(run()) + finally: + server.shutdown() + + +if __name__ == '__main__': + unittest.main() From 8f83983cf63538089057b0ba6930b618224bd060 Mon Sep 17 00:00:00 2001 From: Wu Sheng Date: Tue, 22 Sep 2026 16:45:26 +0800 Subject: [PATCH 2/2] fix: await aiohttp collector requests with instrumentation enabled --- skywalking/plugins/sw_aiohttp.py | 2 +- tests/unit/test_http_reporter.py | 37 ++++++++++++++++++++++++++++++++ 2 files changed, 38 insertions(+), 1 deletion(-) diff --git a/skywalking/plugins/sw_aiohttp.py b/skywalking/plugins/sw_aiohttp.py index 34c3eeaca..65bf42319 100644 --- a/skywalking/plugins/sw_aiohttp.py +++ b/skywalking/plugins/sw_aiohttp.py @@ -46,7 +46,7 @@ async def _sw_request(self: ClientSession, method: str, str_or_url, **kwargs): if config.agent_protocol == 'http' and config.agent_collector_backend_services.rstrip('/') \ .endswith(f'{url.host}:{url.port}'): - return _request + return await _request(self, method, str_or_url, **kwargs) span = NoopSpan(NoopContext()) if config.ignore_http_method_check(method) \ else get_context().new_exit_span(op=url.path or '/', peer=peer, component=Component.AioHttp) diff --git a/tests/unit/test_http_reporter.py b/tests/unit/test_http_reporter.py index 28e31a672..7c5f13fce 100644 --- a/tests/unit/test_http_reporter.py +++ b/tests/unit/test_http_reporter.py @@ -19,6 +19,7 @@ import threading import unittest from http.server import BaseHTTPRequestHandler, HTTPServer +from unittest.mock import patch from skywalking import config @@ -112,6 +113,42 @@ async def run(): finally: server.shutdown() + def test_async_http_heartbeat_with_aiohttp_plugin_installed(self): + from aiohttp import ClientSession + from aiohttp.web_protocol import RequestHandler + + from skywalking.agent.protocol.http_aio import HttpProtocolAsync + from skywalking.plugins import sw_aiohttp + + self.addCleanup(setattr, ClientSession, '_request', ClientSession._request) + self.addCleanup(setattr, RequestHandler, '_handle_request', RequestHandler._handle_request) + sw_aiohttp.install() + + server, port = _start_http_server() + self.addCleanup(server.server_close) + self.addCleanup(server.shutdown) + config.agent_collector_backend_services = f'127.0.0.1:{port}' + + async def run(): + protocol = HttpProtocolAsync() + sessions = [part.client for part in ( + protocol.service_management, protocol.traces_reporter, protocol.log_reporter, + )] + try: + await asyncio.wait_for(protocol.heartbeat(), timeout=5) + await asyncio.wait_for(protocol.heartbeat(), timeout=5) + self.assertTrue(all(not session.closed for session in sessions)) + finally: + await protocol.aclose() + self.assertTrue(all(session.closed for session in sessions)) + + with patch.object(_OkHandler, 'do_POST', autospec=True, side_effect=_OkHandler.do_POST) as post, \ + patch.object(sw_aiohttp, 'get_context') as get_context, \ + patch.object(config, 'agent_collector_properties_report_period_factor', 10): + asyncio.run(run()) + self.assertEqual(post.call_count, 3) + get_context.assert_not_called() + if __name__ == '__main__': unittest.main()