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
7 changes: 7 additions & 0 deletions skywalking/agent/protocol/http_aio.py
Original file line number Diff line number Diff line change
Expand Up @@ -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')
Expand Down
145 changes: 79 additions & 66 deletions skywalking/client/http_aio.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__()
Expand All @@ -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()
Expand All @@ -65,82 +77,83 @@ 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):
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):
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)
2 changes: 1 addition & 1 deletion skywalking/plugins/sw_aiohttp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
154 changes: 154 additions & 0 deletions tests/unit/test_http_reporter.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,154 @@
#
# 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 unittest.mock import patch

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()

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()
Loading