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

Reply via email to