This is an automated email from the ASF dual-hosted git repository. haonan pushed a commit to branch ssl_between_nodes in repository https://gitbox.apache.org/repos/asf/iotdb.git
commit 0965037890087e0d7654ff2fc723610c18b267d7 Author: HTHou <[email protected]> AuthorDate: Tue Jul 22 09:33:15 2025 +0800 fix ainode --- iotdb-core/ainode/ainode/core/config.py | 14 ++++++++++++++ iotdb-core/ainode/ainode/core/rpc/client.py | 2 ++ iotdb-core/ainode/ainode/core/rpc/service.py | 27 +++++++++++++++++++++++---- 3 files changed, 39 insertions(+), 4 deletions(-) diff --git a/iotdb-core/ainode/ainode/core/config.py b/iotdb-core/ainode/ainode/core/config.py index 3fe7a3d921f..832e45fed8b 100644 --- a/iotdb-core/ainode/ainode/core/config.py +++ b/iotdb-core/ainode/ainode/core/config.py @@ -77,6 +77,7 @@ class AINodeConfig(object): # use for ssl self._ain_thrift_ssl_enabled = False self._ain_thrift_ssl_ca_file = None + self._ain_thrift_ssl_cert_file = None # Cache number of model storage to avoid repeated loading self._ain_model_storage_cache_size = 30 @@ -196,6 +197,14 @@ class AINodeConfig(object): ) -> None: self._ain_thrift_ssl_ca_file = ain_thrift_ssl_ca_file + def get_ain_thrift_ssl_cert_file(self) -> str: + return self._ain_thrift_ssl_cert_file + + def set_ain_thrift_ssl_cert_file( + self, ain_thrift_ssl_cert_file: str + ) -> None: + self._ain_thrift_ssl_cert_file = ain_thrift_ssl_cert_file + def get_ain_model_storage_cache_size(self) -> int: return self._ain_model_storage_cache_size @@ -339,6 +348,11 @@ class AINodeDescriptor(object): file_configs["ain_thrift_ssl_ca_file"] ) + if "ain_thrift_ssl_cert_file" in config_keys: + self._config.set_ain_thrift_ssl_cert_file( + file_configs["ain_thrift_ssl_cert_file"] + ) + if "ain_logs_dir" in config_keys: log_dir = file_configs["ain_logs_dir"] self._config.set_ain_logs_dir(log_dir) diff --git a/iotdb-core/ainode/ainode/core/rpc/client.py b/iotdb-core/ainode/ainode/core/rpc/client.py index abb08e29238..954524cae19 100644 --- a/iotdb-core/ainode/ainode/core/rpc/client.py +++ b/iotdb-core/ainode/ainode/core/rpc/client.py @@ -120,6 +120,8 @@ class ConfigNodeClient(object): context.verify_mode = ssl.CERT_REQUIRED context.check_hostname = True context.load_verify_locations(cafile=AINodeDescriptor().get_config().get_ain_thrift_ssl_ca_file()) + context.load_cert_chain( + certfile=AINodeDescriptor().get_config().get_ain_thrift_ssl_cert_file()) socket = TSSLSocket.TSSLSocket( host=target_config_node.ip, port=target_config_node.port, ssl_context=context ) diff --git a/iotdb-core/ainode/ainode/core/rpc/service.py b/iotdb-core/ainode/ainode/core/rpc/service.py index 72bd38c4ade..e4558db5d5c 100644 --- a/iotdb-core/ainode/ainode/core/rpc/service.py +++ b/iotdb-core/ainode/ainode/core/rpc/service.py @@ -20,6 +20,7 @@ import threading from thrift.protocol import TBinaryProtocol, TCompactProtocol from thrift.server import TServer from thrift.transport import TSocket, TTransport +from thrift.transport.TSSLSocket import TSSLSocket from ainode.core.config import AINodeDescriptor from ainode.core.log import Logger @@ -70,10 +71,28 @@ class AINodeRPCService(threading.Thread): self._stop_event = threading.Event() self._handler = handler processor = IAINodeRPCService.Processor(handler=self._handler) - transport = TSocket.TServerSocket( - host=AINodeDescriptor().get_config().get_ain_inference_rpc_address(), - port=AINodeDescriptor().get_config().get_ain_inference_rpc_port(), - ) + if AINodeDescriptor().get_config().get_ain_thrift_ssl_enabled(): + import ssl,sys + from thrift.transport import TSSLSocket + + if sys.version_info >= (3, 10): + context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH) + else: + context = ssl.SSLContext(ssl.PROTOCOL_TLS) + context.verify_mode = ssl.CERT_REQUIRED + context.check_hostname = True + context.load_verify_locations(cafile=AINodeDescriptor().get_config().get_ain_thrift_ssl_ca_file()) + context.load_cert_chain(certfile=AINodeDescriptor().get_config().get_ain_thrift_ssl_cert_file()) + transport = TSSLSocket.TSSLServerSocket( + host=AINodeDescriptor().get_config().get_ain_inference_rpc_address(), + port=AINodeDescriptor().get_config().get_ain_inference_rpc_port(), + ssl_context=context + ) + else: + transport = TSocket.TServerSocket( + host=AINodeDescriptor().get_config().get_ain_inference_rpc_address(), + port=AINodeDescriptor().get_config().get_ain_inference_rpc_port(), + ) transport_factory = TTransport.TFramedTransportFactory() if AINodeDescriptor().get_config().get_ain_thrift_compression_enabled(): protocol_factory = TCompactProtocol.TCompactProtocolFactory()
