Skip to content

Commit 553aea0

Browse files
authored
chore: Refactor mock_server setup for network level TLS handling and added thread safety (#2047)
Signed-off-by: Manish Dait <daitmanish88@gmail.com>
1 parent 8fa79db commit 553aea0

2 files changed

Lines changed: 24 additions & 29 deletions

File tree

CHANGELOG.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ This changelog is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.
1010
-
1111

1212
### Tests
13+
- Refactor `mock_server` setup for network level TLS handling and added thread safety
1314

1415

1516
### Examples

tests/unit/mock_server.py

Lines changed: 23 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,12 @@
11
import grpc
2+
import threading
23
from concurrent import futures
34
from contextlib import contextmanager
45
from hiero_sdk_python.client.network import Network
56
from hiero_sdk_python.client.client import Client
67
from hiero_sdk_python.account.account_id import AccountId
78
from hiero_sdk_python.crypto.private_key import PrivateKey
89
from hiero_sdk_python.client.network import _Node
9-
import socket
10-
from contextlib import closing
1110
from hiero_sdk_python.hapi.services import (
1211
crypto_service_pb2_grpc,
1312
token_service_pb2_grpc,
@@ -32,14 +31,15 @@ def __init__(self, responses):
3231
responses (list): List of response objects to return in sequence
3332
"""
3433
self.responses = responses
34+
self._lock = threading.Lock()
3535
self.server = grpc.server(futures.ThreadPoolExecutor(max_workers=10))
36-
self.port = _find_free_port()
36+
37+
self.port = self.server.add_insecure_port('[::]:0')
3738
self.address = f"localhost:{self.port}"
3839

3940
self._register_services()
4041

4142
# Start the server
42-
self.server.add_insecure_port(self.address)
4343
self.server.start()
4444

4545
def _register_services(self):
@@ -95,6 +95,7 @@ def _create_mock_servicer(self, servicer_class):
9595
A mock servicer object
9696
"""
9797
responses = self.responses
98+
lock = self._lock;
9899

99100
class MockServicer(servicer_class):
100101
def __getattribute__(self, name):
@@ -103,12 +104,11 @@ def __getattribute__(self, name):
103104
return super().__getattribute__(name)
104105

105106
def method_wrapper(request, context):
106-
nonlocal responses
107-
if not responses:
108-
# If no more responses are available, return None
109-
return None
107+
with lock:
108+
if not responses:
109+
return None
110110

111-
response = responses.pop(0)
111+
response = responses.pop(0)
112112

113113
if isinstance(response, RealRpcError):
114114
# Abort with custom error
@@ -122,21 +122,11 @@ def method_wrapper(request, context):
122122

123123
def close(self):
124124
"""Stop the server."""
125-
self.server.stop(0)
125+
shutdown_event = self.server.stop(0)
126+
success = shutdown_event.wait(timeout=2.0)
126127

127-
128-
def _find_free_port():
129-
"""Find a free port on localhost."""
130-
with closing(socket.socket(socket.AF_INET, socket.SOCK_STREAM)) as s:
131-
s.bind(("", 0))
132-
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
133-
port = s.getsockname()[1]
134-
135-
# If we get the tls port 50212 port skip it
136-
if port in [50212]:
137-
return port + 1
138-
139-
return port
128+
if not success:
129+
pass
140130

141131

142132
class RealRpcError(grpc.RpcError):
@@ -172,16 +162,20 @@ def mock_hedera_servers(response_sequences):
172162
nodes = []
173163
for i, server in enumerate(servers):
174164
node = _Node(AccountId(0, 0, 3 + i), server.address, None)
175-
176-
# force insecure transport and mock cert even if we get the tls-port
177-
node._set_root_certificates(b"mock-tls-cert-for-unit-tests")
178-
node._apply_transport_security(False)
179-
node._set_verify_certificates(False)
180-
181165
nodes.append(node)
182166

183167
# Create network and client
184168
network = Network(nodes=nodes)
169+
network.set_transport_security(False)
170+
network.set_verify_certificates(False)
171+
client = Client(network)
172+
173+
# Force non-tls for channel
174+
for node in client.network.nodes:
175+
node._address._is_transport_security = lambda: False
176+
node._set_verify_certificates(False)
177+
node._close()
178+
185179
client = Client(network)
186180
client.logger.set_level(LogLevel.DISABLED)
187181
# Set the operator

0 commit comments

Comments
 (0)