diff --git a/influxdb_client_3/write_client/client/util/multiprocessing_helper.py b/influxdb_client_3/write_client/client/util/multiprocessing_helper.py index 96ba6469..1bc8245a 100644 --- a/influxdb_client_3/write_client/client/util/multiprocessing_helper.py +++ b/influxdb_client_3/write_client/client/util/multiprocessing_helper.py @@ -7,9 +7,9 @@ import logging import multiprocessing -from influxdb_client_3 import InfluxDBClient3, write_client_options -from influxdb_client_3.write_client import WriteOptions +from influxdb_client_3 import write_client_options from influxdb_client_3.exceptions import InfluxDBError +from influxdb_client_3.write_client import WriteOptions, WriteApi logger = logging.getLogger('influxdb_client.client.util.multiprocessing_helper') @@ -35,9 +35,9 @@ class _PoisonPill: pass -class MultiprocessingWriter(multiprocessing.Process): +class MultiprocessingWriter: """ - The Helper class to write data into InfluxDB in independent OS process. + The Helper class to write data into InfluxDB in an independent OS process. Example: .. code-block:: python @@ -119,27 +119,27 @@ def main(): __started__ = False __disposed__ = False - def __init__(self, **kwargs) -> None: + def __init__(self, start_method='spawn', **kwargs) -> None: """ Initialize defaults. - For more information how to initialize the writer see the examples above. + For more information on how to initialize the writer, see the examples above. - :param kwargs: arguments are passed into ``__init__`` function of ``InfluxDBClient`` and ``write_api``. + :param kwargs: arguments are passed into the `` _ _init__`` function of ``InfluxDBClient`` and ``write_api``. """ - multiprocessing.Process.__init__(self) + self.ctx = multiprocessing.get_context(start_method) + self.process = self.ctx.Process(target=self.run) self.kwargs = kwargs - self.client = None self.write_api = None - self.queue_ = multiprocessing.Manager().Queue() + self.queue_ = self.ctx.JoinableQueue() def write(self, **kwargs) -> None: """ - Append time-series data into underlying queue. + Append time-series data into the underlying queue. - For more information how to pass arguments see the examples above. + For more information on how to pass arguments, see the examples above. - :param kwargs: arguments are passed into ``write`` function of ``WriteApi`` + :param kwargs: arguments are passed into the `` write `` function of ``WriteApi`` :return: None """ assert self.__disposed__ is False, 'Cannot write data: the writer is closed.' @@ -155,16 +155,13 @@ def run(self): retry_callback=self.kwargs.get('retry_callback', _retry_callback) ) - # Still need to create the InfluxDBClient3 because the init logics of InfluxDBClient3 will create the WriteApi. - # it will make WriteApi class created properly. - self.client = InfluxDBClient3(write_client_options=wco, **self.kwargs) - - # Close and set _query_api to None because query_api is not needed in this process. - # We only need write_api. - self.client._query_api.close() - self.client._query_api = None - - self.write_api = self.client._write_api + self.write_api = WriteApi( + bucket=self.kwargs.get('database'), + org=self.kwargs.get('org'), + default_header=self.kwargs.get('default_header'), + rest_client=self.kwargs.get('rest_client'), + **wco + ) # Infinite loop - until poison pill while True: next_record = self.queue_.get() @@ -177,26 +174,26 @@ def run(self): self.queue_.task_done() def start(self) -> None: - """Start independent process for writing data into InfluxDB.""" - super().start() + """Start an independent process for writing data into InfluxDB.""" + self.process.start() self.__started__ = True def terminate(self) -> None: """ - Cleanup resources in independent process. + Cleanup resources in an independent process. This function **cannot be used** to terminate the ``MultiprocessingWriter``. - If you want to finish your writes please call: ``__del__``. + If you want to finish your writes, please call: ``__del__``. """ if self.write_api: logger.info("flushing data...") - self.write_api.__del__() + self.write_api.close() self.write_api = None - if self.client: - self.client.close() - self.client = None logger.info("closed") + def get_start_processing_method(self): + return self.ctx.get_start_method() + def __enter__(self): """Enter the runtime context related to this object.""" self.start() @@ -207,11 +204,11 @@ def __exit__(self, exc_type, exc_value, traceback): self.__del__() def __del__(self): - """Dispose the client and write_api.""" + """Dispose of the client and write_api.""" if self.__started__: self.queue_.put(_PoisonPill()) self.queue_.join() - self.join() + self.process.join() self.queue_ = None self.__started__ = False self.__disposed__ = True diff --git a/tests/test_influxdb_client_3_integration.py b/tests/test_influxdb_client_3_integration.py index b6225bb5..9f22173a 100644 --- a/tests/test_influxdb_client_3_integration.py +++ b/tests/test_influxdb_client_3_integration.py @@ -1,10 +1,10 @@ +import asyncio import json import logging import os import random import string import time -import asyncio import unittest import pandas as pd @@ -21,8 +21,6 @@ from influxdb_client_3.write_client.write_exceptions import ApiException from tests.util import asyncio_run, lp_to_py_object -running_on_posix = os.name == 'posix' - def random_hex(len=6): return ''.join(random.choice(string.hexdigits) for i in range(len)) @@ -343,28 +341,76 @@ def test_batch_write_closed(self): list_results = reader.to_pylist() self.assertEqual(data_size, len(list_results)) - @pytest.mark.skipif(running_on_posix, reason="Skipping this test in POSIX environments") def test_multiprocessing_helper(self): - org = 'my-org' - writer = MultiprocessingWriter( - host=self.host, - database=self.database, - token=self.token, - org=org, - write_options=WriteOptions(batch_size=1)) - writer.start() - measurement = f'test{random_hex(6)}'.lower() - for x in range(1, 10): - time.sleep(0.2) - writer.write( - bucket=self.database, - record=f"{measurement},tag=a value=\"number{x}\" {time.time_ns()}" - ) - writer.__del__() + default_header = { + 'Authorization': f'Token {self.token}' + } + rest = rest_client.RestClient( + base_url=self.host, + default_header=default_header, + ) + + with MultiprocessingWriter( + host=self.host, + database=self.database, + token=self.token, + org='my-org', + default_header=default_header, + rest_client=rest, + write_options=WriteOptions(batch_size=1)) as mp: + self.assertEqual(mp.get_start_processing_method(), 'spawn') + + measurement = f'test{random_hex(6)}'.lower() + for x in range(1, 5): + time.sleep(0.5) + mp.write( + bucket=self.database, + record=f"{measurement},tag=a value=\"number{x}\" {time.time_ns()}" + ) time.sleep(1) df = self.client.query(f'select * from {measurement}', mode="pandas") - self.assertEqual(9, len(df)) + self.assertEqual(4, len(df)) + + def test_multiprocessing_start_method_forkserver(self): + default_header = { + 'Authorization': f'Token {self.token}' + } + + with MultiprocessingWriter( + host=self.host, + database=self.database, + token=self.token, + org='my-org', + default_header=default_header, + rest_client=(rest_client.RestClient( + base_url=self.host, + default_header=default_header, + )), + write_options=WriteOptions(batch_size=1), + start_method='forkserver' + ) as mp: + self.assertEqual(mp.get_start_processing_method(), 'forkserver') + + def test_multiprocessing_start_method_fork(self): + default_header = { + 'Authorization': f'Token {self.token}' + } + + with MultiprocessingWriter( + host=self.host, + database=self.database, + token=self.token, + org='my-org', + default_header=default_header, + rest_client=(rest_client.RestClient( + base_url=self.host, + default_header=default_header, + )), + write_options=WriteOptions(batch_size=1), + start_method='fork' + ) as mp: + self.assertEqual(mp.get_start_processing_method(), 'fork') test_cert = """-----BEGIN CERTIFICATE----- MIIDUzCCAjugAwIBAgIUZB55ULutbc9gy6xLp1BkTQU7siowDQYJKoZIhvcNAQEL