← back to Exo
implement a health check for peers and discovery should only return healthy peers
de2f6d2e6e1d4e8993707fa4307daa285e2ba133 · 2024-09-23 23:10:32 +0100 · Alex Cheema
Files touched
M exo/networking/grpc/grpc_peer_handle.pyM exo/networking/grpc/node_service.protoM exo/networking/grpc/node_service_pb2.pyM exo/networking/grpc/node_service_pb2_grpc.pyM exo/networking/peer_handle.pyM exo/networking/udp_discovery.py
Diff
commit de2f6d2e6e1d4e8993707fa4307daa285e2ba133
Author: Alex Cheema <alexcheema123@gmail.com>
Date: Mon Sep 23 23:10:32 2024 +0100
implement a health check for peers and discovery should only return healthy peers
---
exo/networking/grpc/grpc_peer_handle.py | 11 +
exo/networking/grpc/node_service.proto | 7 +
exo/networking/grpc/node_service_pb2.py | 79 ++--
exo/networking/grpc/node_service_pb2_grpc.py | 544 ++++++++++++++++-----------
exo/networking/peer_handle.py | 4 +
exo/networking/udp_discovery.py | 40 +-
6 files changed, 405 insertions(+), 280 deletions(-)
diff --git a/exo/networking/grpc/grpc_peer_handle.py b/exo/networking/grpc/grpc_peer_handle.py
index a4ad092a..c9f42ba6 100644
--- a/exo/networking/grpc/grpc_peer_handle.py
+++ b/exo/networking/grpc/grpc_peer_handle.py
@@ -1,5 +1,6 @@
import grpc
import numpy as np
+import asyncio
from typing import Optional, Tuple, List
# These would be generated from the .proto file
@@ -43,6 +44,16 @@ class GRPCPeerHandle(PeerHandle):
self.channel = None
self.stub = None
+ async def health_check(self) -> bool:
+ try:
+ request = node_service_pb2.HealthCheckRequest()
+ response = await asyncio.wait_for(self.stub.HealthCheck(request), timeout=5)
+ return response.is_healthy
+ except asyncio.TimeoutError:
+ return False
+ except:
+ return False
+
async def send_prompt(self, shard: Shard, prompt: str, image_str: Optional[str] = None, request_id: Optional[str] = None, inference_state: Optional[str] = None) -> Optional[np.array]:
request = node_service_pb2.PromptRequest(
prompt=prompt,
diff --git a/exo/networking/grpc/node_service.proto b/exo/networking/grpc/node_service.proto
index 6fcee351..ad10a5c1 100644
--- a/exo/networking/grpc/node_service.proto
+++ b/exo/networking/grpc/node_service.proto
@@ -9,6 +9,7 @@ service NodeService {
rpc CollectTopology (CollectTopologyRequest) returns (Topology) {}
rpc SendResult (SendResultRequest) returns (Empty) {}
rpc SendOpaqueStatus (SendOpaqueStatusRequest) returns (Empty) {}
+ rpc HealthCheck (HealthCheckRequest) returns (HealthCheckResponse) {}
}
message Shard {
@@ -86,4 +87,10 @@ message SendOpaqueStatusRequest {
string status = 2;
}
+message HealthCheckRequest {}
+
+message HealthCheckResponse {
+ bool is_healthy = 1;
+}
+
message Empty {}
\ No newline at end of file
diff --git a/exo/networking/grpc/node_service_pb2.py b/exo/networking/grpc/node_service_pb2.py
index cae2d080..47dd0b5c 100644
--- a/exo/networking/grpc/node_service_pb2.py
+++ b/exo/networking/grpc/node_service_pb2.py
@@ -11,9 +11,10 @@ from google.protobuf.internal import builder as _builder
_sym_db = _symbol_database.Default()
-DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(
- b'\n\x12node_service.proto\x12\x0cnode_service\"S\n\x05Shard\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x13\n\x0bstart_layer\x18\x02 \x01(\x05\x12\x11\n\tend_layer\x18\x03 \x01(\x05\x12\x10\n\x08n_layers\x18\x04 \x01(\x05\"\xc3\x01\n\rPromptRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12\x0e\n\x06prompt\x18\x02 \x01(\t\x12\x16\n\timage_str\x18\x03 \x01(\tH\x00\x88\x01\x01\x12\x17\n\nrequest_id\x18\x04 \x01(\tH\x01\x88\x01\x01\x12\x1c\n\x0finference_state\x18\x05 \x01(\tH\x02\x88\x01\x01\x42\x0c\n\n_image_strB\r\n\x0b_request_idB\x12\n\x10_inference_state\"\xb3\x01\n\rTensorRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12$\n\x06tensor\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x12\x1c\n\x0finference_state\x18\x04 \x01(\tH\x01\x88\x01\x01\x42\r\n\x0b_request_idB\x12\n\x10_inference_state\"/\n\x19GetInferenceResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\"\\\n\x0fInferenceResult\x12)\n\x06tensor\x18\x01 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x12\x13\n\x0bis_finished\x18\x02 \x01(\x08\x42\t\n\x07_tensor\";\n\x06Tensor\x12\x13\n\x0btensor_data\x18\x01 \x01(\x0c\x12\r\n\x05shape\x18\x02 \x03(\x05\x12\r\n\x05\x64type\x18\x03 \x01(\t\"<\n\x16\x43ollectTopologyRequest\x12\x0f\n\x07visited\x18\x01 \x03(\t\x12\x11\n\tmax_depth\x18\x02 \x01(\x05\"\x8e\x02\n\x08Topology\x12\x30\n\x05nodes\x18\x01 \x03(\x0b\x32!.node_service.Topology.NodesEntry\x12\x39\n\npeer_graph\x18\x02 \x03(\x0b\x32%.node_service.Topology.PeerGraphEntry\x1aN\n\nNodesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12/\n\x05value\x18\x02 \x01(\x0b\x32 .node_service.DeviceCapabilities:\x02\x38\x01\x1a\x45\n\x0ePeerGraphEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\"\n\x05value\x18\x02 \x01(\x0b\x32\x13.node_service.Peers:\x02\x38\x01\"\x19\n\x05Peers\x12\x10\n\x08peer_ids\x18\x01 \x03(\t\"7\n\x0b\x44\x65viceFlops\x12\x0c\n\x04\x66p32\x18\x01 \x01(\x02\x12\x0c\n\x04\x66p16\x18\x02 \x01(\x02\x12\x0c\n\x04int8\x18\x03 \x01(\x02\"k\n\x12\x44\x65viceCapabilities\x12\r\n\x05model\x18\x01 \x01(\t\x12\x0c\n\x04\x63hip\x18\x02 \x01(\t\x12\x0e\n\x06memory\x18\x03 \x01(\x05\x12(\n\x05\x66lops\x18\x04 \x01(\x0b\x32\x19.node_service.DeviceFlops\"L\n\x11SendResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06result\x18\x02 \x03(\x05\x12\x13\n\x0bis_finished\x18\x03 \x01(\x08\"=\n\x17SendOpaqueStatusRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t\"\x07\n\x05\x45mpty2\xde\x03\n\x0bNodeService\x12\x41\n\nSendPrompt\x12\x1b.node_service.PromptRequest\x1a\x14.node_service.Tensor\"\x00\x12\x41\n\nSendTensor\x12\x1b.node_service.TensorRequest\x1a\x14.node_service.Tensor\"\x00\x12^\n\x12GetInferenceResult\x12\'.node_service.GetInferenceResultRequest\x1a\x1d.node_service.InferenceResult\"\x00\x12Q\n\x0f\x43ollectTopology\x12$.node_service.CollectTopologyRequest\x1a\x16.node_service.Topology\"\x00\x12\x44\n\nSendResult\x12\x1f.node_service.SendResultRequest\x1a\x13.node_service.Empty\"\x00\x12P\n\x10SendOpaqueStatus\x12%.node_service.SendOpaqueStatusRequest\x1a\x13.node_service.Empty\"\x00\x62\x06proto3'
-)
+
+
+
+DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x12node_service.proto\x12\x0cnode_service\"S\n\x05Shard\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x13\n\x0bstart_layer\x18\x02 \x01(\x05\x12\x11\n\tend_layer\x18\x03 \x01(\x05\x12\x10\n\x08n_layers\x18\x04 \x01(\x05\"\xc3\x01\n\rPromptRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12\x0e\n\x06prompt\x18\x02 \x01(\t\x12\x16\n\timage_str\x18\x03 \x01(\tH\x00\x88\x01\x01\x12\x17\n\nrequest_id\x18\x04 \x01(\tH\x01\x88\x01\x01\x12\x1c\n\x0finference_state\x18\x05 \x01(\tH\x02\x88\x01\x01\x42\x0c\n\n_image_strB\r\n\x0b_request_idB\x12\n\x10_inference_state\"\xb3\x01\n\rTensorRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12$\n\x06tensor\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x12\x1c\n\x0finference_state\x18\x04 \x01(\tH\x01\x88\x01\x01\x42\r\n\x0b_request_idB\x12\n\x10_inference_state\"/\n\x19GetInferenceResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\"\\\n\x0fInferenceResult\x12)\n\x06tensor\x18\x01 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x12\x13\n\x0bis_finished\x18\x02 \x01(\x08\x42\t\n\x07_tensor\";\n\x06Tensor\x12\x13\n\x0btensor_data\x18\x01 \x01(\x0c\x12\r\n\x05shape\x18\x02 \x03(\x05\x12\r\n\x05\x64type\x18\x03 \x01(\t\"<\n\x16\x43ollectTopologyRequest\x12\x0f\n\x07visited\x18\x01 \x03(\t\x12\x11\n\tmax_depth\x18\x02 \x01(\x05\"\x8e\x02\n\x08Topology\x12\x30\n\x05nodes\x18\x01 \x03(\x0b\x32!.node_service.Topology.NodesEntry\x12\x39\n\npeer_graph\x18\x02 \x03(\x0b\x32%.node_service.Topology.PeerGraphEntry\x1aN\n\nNodesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12/\n\x05value\x18\x02 \x01(\x0b\x32 .node_service.DeviceCapabilities:\x02\x38\x01\x1a\x45\n\x0ePeerGraphEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\"\n\x05value\x18\x02 \x01(\x0b\x32\x13.node_service.Peers:\x02\x38\x01\"\x19\n\x05Peers\x12\x10\n\x08peer_ids\x18\x01 \x03(\t\"7\n\x0b\x44\x65viceFlops\x12\x0c\n\x04\x66p32\x18\x01 \x01(\x02\x12\x0c\n\x04\x66p16\x18\x02 \x01(\x02\x12\x0c\n\x04int8\x18\x03 \x01(\x02\"k\n\x12\x44\x65viceCapabilities\x12\r\n\x05model\x18\x01 \x01(\t\x12\x0c\n\x04\x63hip\x18\x02 \x01(\t\x12\x0e\n\x06memory\x18\x03 \x01(\x05\x12(\n\x05\x66lops\x18\x04 \x01(\x0b\x32\x19.node_service.DeviceFlops\"L\n\x11SendResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06result\x18\x02 \x03(\x05\x12\x13\n\x0bis_finished\x18\x03 \x01(\x08\"=\n\x17SendOpaqueStatusRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t\"\x14\n\x12HealthCheckRequest\")\n\x13HealthCheckResponse\x12\x12\n\nis_healthy\x18\x01 \x01(\x08\"\x07\n\x05\x45mpty2\xb4\x04\n\x0bNodeService\x12\x41\n\nSendPrompt\x12\x1b.node_service.PromptRequest\x1a\x14.node_service.Tensor\"\x00\x12\x41\n\nSendTensor\x12\x1b.node_service.TensorRequest\x1a\x14.node_service.Tensor\"\x00\x12^\n\x12GetInferenceResult\x12\'.node_service.GetInferenceResultRequest\x1a\x1d.node_service.InferenceResult\"\x00\x12Q\n\x0f\x43ollectTopology\x12$.node_service.CollectTopologyRequest\x1a\x16.node_service.Topology\"\x00\x12\x44\n\nSendResult\x12\x1f.node_service.SendResultRequest\x1a\x13.node_service.Empty\"\x00\x12P\n\x10SendOpaqueStatus\x12%.node_service.SendOpaqueStatusRequest\x1a\x13.node_service.Empty\"\x00\x12T\n\x0bHealthCheck\x12 .node_service.HealthCheckRequest\x1a!.node_service.HealthCheckResponse\"\x00\x62\x06proto3')
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
@@ -24,38 +25,42 @@ if not _descriptor._USE_C_DESCRIPTORS:
_globals['_TOPOLOGY_NODESENTRY']._serialized_options = b'8\001'
_globals['_TOPOLOGY_PEERGRAPHENTRY']._loaded_options = None
_globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_options = b'8\001'
- _globals['_SHARD']._serialized_start = 36
- _globals['_SHARD']._serialized_end = 119
- _globals['_PROMPTREQUEST']._serialized_start = 122
- _globals['_PROMPTREQUEST']._serialized_end = 317
- _globals['_TENSORREQUEST']._serialized_start = 320
- _globals['_TENSORREQUEST']._serialized_end = 499
- _globals['_GETINFERENCERESULTREQUEST']._serialized_start = 501
- _globals['_GETINFERENCERESULTREQUEST']._serialized_end = 548
- _globals['_INFERENCERESULT']._serialized_start = 550
- _globals['_INFERENCERESULT']._serialized_end = 642
- _globals['_TENSOR']._serialized_start = 644
- _globals['_TENSOR']._serialized_end = 703
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start = 705
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end = 765
- _globals['_TOPOLOGY']._serialized_start = 768
- _globals['_TOPOLOGY']._serialized_end = 1038
- _globals['_TOPOLOGY_NODESENTRY']._serialized_start = 889
- _globals['_TOPOLOGY_NODESENTRY']._serialized_end = 967
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start = 969
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end = 1038
- _globals['_PEERS']._serialized_start = 1040
- _globals['_PEERS']._serialized_end = 1065
- _globals['_DEVICEFLOPS']._serialized_start = 1067
- _globals['_DEVICEFLOPS']._serialized_end = 1122
- _globals['_DEVICECAPABILITIES']._serialized_start = 1124
- _globals['_DEVICECAPABILITIES']._serialized_end = 1231
- _globals['_SENDRESULTREQUEST']._serialized_start = 1233
- _globals['_SENDRESULTREQUEST']._serialized_end = 1309
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start = 1311
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end = 1372
- _globals['_EMPTY']._serialized_start = 1374
- _globals['_EMPTY']._serialized_end = 1381
- _globals['_NODESERVICE']._serialized_start = 1384
- _globals['_NODESERVICE']._serialized_end = 1862
+ _globals['_SHARD']._serialized_start=36
+ _globals['_SHARD']._serialized_end=119
+ _globals['_PROMPTREQUEST']._serialized_start=122
+ _globals['_PROMPTREQUEST']._serialized_end=317
+ _globals['_TENSORREQUEST']._serialized_start=320
+ _globals['_TENSORREQUEST']._serialized_end=499
+ _globals['_GETINFERENCERESULTREQUEST']._serialized_start=501
+ _globals['_GETINFERENCERESULTREQUEST']._serialized_end=548
+ _globals['_INFERENCERESULT']._serialized_start=550
+ _globals['_INFERENCERESULT']._serialized_end=642
+ _globals['_TENSOR']._serialized_start=644
+ _globals['_TENSOR']._serialized_end=703
+ _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=705
+ _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=765
+ _globals['_TOPOLOGY']._serialized_start=768
+ _globals['_TOPOLOGY']._serialized_end=1038
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_start=889
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_end=967
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=969
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1038
+ _globals['_PEERS']._serialized_start=1040
+ _globals['_PEERS']._serialized_end=1065
+ _globals['_DEVICEFLOPS']._serialized_start=1067
+ _globals['_DEVICEFLOPS']._serialized_end=1122
+ _globals['_DEVICECAPABILITIES']._serialized_start=1124
+ _globals['_DEVICECAPABILITIES']._serialized_end=1231
+ _globals['_SENDRESULTREQUEST']._serialized_start=1233
+ _globals['_SENDRESULTREQUEST']._serialized_end=1309
+ _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1311
+ _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1372
+ _globals['_HEALTHCHECKREQUEST']._serialized_start=1374
+ _globals['_HEALTHCHECKREQUEST']._serialized_end=1394
+ _globals['_HEALTHCHECKRESPONSE']._serialized_start=1396
+ _globals['_HEALTHCHECKRESPONSE']._serialized_end=1437
+ _globals['_EMPTY']._serialized_start=1439
+ _globals['_EMPTY']._serialized_end=1446
+ _globals['_NODESERVICE']._serialized_start=1449
+ _globals['_NODESERVICE']._serialized_end=2013
# @@protoc_insertion_point(module_scope)
diff --git a/exo/networking/grpc/node_service_pb2_grpc.py b/exo/networking/grpc/node_service_pb2_grpc.py
index ea1d3c98..dd166ca9 100644
--- a/exo/networking/grpc/node_service_pb2_grpc.py
+++ b/exo/networking/grpc/node_service_pb2_grpc.py
@@ -12,261 +12,349 @@ SCHEDULED_RELEASE_DATE = 'June 25, 2024'
_version_not_supported = False
try:
- from grpc._utilities import first_version_is_lower
- _version_not_supported = first_version_is_lower(GRPC_VERSION, GRPC_GENERATED_VERSION)
+ from grpc._utilities import first_version_is_lower
+ _version_not_supported = first_version_is_lower(GRPC_VERSION, GRPC_GENERATED_VERSION)
except ImportError:
- _version_not_supported = True
+ _version_not_supported = True
if _version_not_supported:
- warnings.warn(
- f'The grpc package installed is at version {GRPC_VERSION},' + f' but the generated code in node_service_pb2_grpc.py depends on' + f' grpcio>={GRPC_GENERATED_VERSION}.' +
- f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' +
- f' This warning will become an error in {EXPECTED_ERROR_RELEASE},' + f' scheduled for release on {SCHEDULED_RELEASE_DATE}.', RuntimeWarning
- )
+ warnings.warn(
+ f'The grpc package installed is at version {GRPC_VERSION},'
+ + f' but the generated code in node_service_pb2_grpc.py depends on'
+ + f' grpcio>={GRPC_GENERATED_VERSION}.'
+ + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}'
+ + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.'
+ + f' This warning will become an error in {EXPECTED_ERROR_RELEASE},'
+ + f' scheduled for release on {SCHEDULED_RELEASE_DATE}.',
+ RuntimeWarning
+ )
class NodeServiceStub(object):
- """Missing associated documentation comment in .proto file."""
- def __init__(self, channel):
- """Constructor.
+ """Missing associated documentation comment in .proto file."""
+
+ def __init__(self, channel):
+ """Constructor.
Args:
channel: A grpc.Channel.
"""
- self.SendPrompt = channel.unary_unary(
- '/node_service.NodeService/SendPrompt',
- request_serializer=node__service__pb2.PromptRequest.SerializeToString,
- response_deserializer=node__service__pb2.Tensor.FromString,
- _registered_method=True
- )
- self.SendTensor = channel.unary_unary(
- '/node_service.NodeService/SendTensor',
- request_serializer=node__service__pb2.TensorRequest.SerializeToString,
- response_deserializer=node__service__pb2.Tensor.FromString,
- _registered_method=True
- )
- self.GetInferenceResult = channel.unary_unary(
- '/node_service.NodeService/GetInferenceResult',
- request_serializer=node__service__pb2.GetInferenceResultRequest.SerializeToString,
- response_deserializer=node__service__pb2.InferenceResult.FromString,
- _registered_method=True
- )
- self.CollectTopology = channel.unary_unary(
- '/node_service.NodeService/CollectTopology',
- request_serializer=node__service__pb2.CollectTopologyRequest.SerializeToString,
- response_deserializer=node__service__pb2.Topology.FromString,
- _registered_method=True
- )
- self.SendResult = channel.unary_unary(
- '/node_service.NodeService/SendResult',
- request_serializer=node__service__pb2.SendResultRequest.SerializeToString,
- response_deserializer=node__service__pb2.Empty.FromString,
- _registered_method=True
- )
- self.SendOpaqueStatus = channel.unary_unary(
- '/node_service.NodeService/SendOpaqueStatus',
- request_serializer=node__service__pb2.SendOpaqueStatusRequest.SerializeToString,
- response_deserializer=node__service__pb2.Empty.FromString,
- _registered_method=True
- )
+ self.SendPrompt = channel.unary_unary(
+ '/node_service.NodeService/SendPrompt',
+ request_serializer=node__service__pb2.PromptRequest.SerializeToString,
+ response_deserializer=node__service__pb2.Tensor.FromString,
+ _registered_method=True)
+ self.SendTensor = channel.unary_unary(
+ '/node_service.NodeService/SendTensor',
+ request_serializer=node__service__pb2.TensorRequest.SerializeToString,
+ response_deserializer=node__service__pb2.Tensor.FromString,
+ _registered_method=True)
+ self.GetInferenceResult = channel.unary_unary(
+ '/node_service.NodeService/GetInferenceResult',
+ request_serializer=node__service__pb2.GetInferenceResultRequest.SerializeToString,
+ response_deserializer=node__service__pb2.InferenceResult.FromString,
+ _registered_method=True)
+ self.CollectTopology = channel.unary_unary(
+ '/node_service.NodeService/CollectTopology',
+ request_serializer=node__service__pb2.CollectTopologyRequest.SerializeToString,
+ response_deserializer=node__service__pb2.Topology.FromString,
+ _registered_method=True)
+ self.SendResult = channel.unary_unary(
+ '/node_service.NodeService/SendResult',
+ request_serializer=node__service__pb2.SendResultRequest.SerializeToString,
+ response_deserializer=node__service__pb2.Empty.FromString,
+ _registered_method=True)
+ self.SendOpaqueStatus = channel.unary_unary(
+ '/node_service.NodeService/SendOpaqueStatus',
+ request_serializer=node__service__pb2.SendOpaqueStatusRequest.SerializeToString,
+ response_deserializer=node__service__pb2.Empty.FromString,
+ _registered_method=True)
+ self.HealthCheck = channel.unary_unary(
+ '/node_service.NodeService/HealthCheck',
+ request_serializer=node__service__pb2.HealthCheckRequest.SerializeToString,
+ response_deserializer=node__service__pb2.HealthCheckResponse.FromString,
+ _registered_method=True)
class NodeServiceServicer(object):
- """Missing associated documentation comment in .proto file."""
- def SendPrompt(self, request, context):
"""Missing associated documentation comment in .proto file."""
- context.set_code(grpc.StatusCode.UNIMPLEMENTED)
- context.set_details('Method not implemented!')
- raise NotImplementedError('Method not implemented!')
- def SendTensor(self, request, context):
- """Missing associated documentation comment in .proto file."""
- context.set_code(grpc.StatusCode.UNIMPLEMENTED)
- context.set_details('Method not implemented!')
- raise NotImplementedError('Method not implemented!')
+ def SendPrompt(self, request, context):
+ """Missing associated documentation comment in .proto file."""
+ context.set_code(grpc.StatusCode.UNIMPLEMENTED)
+ context.set_details('Method not implemented!')
+ raise NotImplementedError('Method not implemented!')
- def GetInferenceResult(self, request, context):
- """Missing associated documentation comment in .proto file."""
- context.set_code(grpc.StatusCode.UNIMPLEMENTED)
- context.set_details('Method not implemented!')
- raise NotImplementedError('Method not implemented!')
+ def SendTensor(self, request, context):
+ """Missing associated documentation comment in .proto file."""
+ context.set_code(grpc.StatusCode.UNIMPLEMENTED)
+ context.set_details('Method not implemented!')
+ raise NotImplementedError('Method not implemented!')
- def CollectTopology(self, request, context):
- """Missing associated documentation comment in .proto file."""
- context.set_code(grpc.StatusCode.UNIMPLEMENTED)
- context.set_details('Method not implemented!')
- raise NotImplementedError('Method not implemented!')
+ def GetInferenceResult(self, request, context):
+ """Missing associated documentation comment in .proto file."""
+ context.set_code(grpc.StatusCode.UNIMPLEMENTED)
+ context.set_details('Method not implemented!')
+ raise NotImplementedError('Method not implemented!')
- def SendResult(self, request, context):
- """Missing associated documentation comment in .proto file."""
- context.set_code(grpc.StatusCode.UNIMPLEMENTED)
- context.set_details('Method not implemented!')
- raise NotImplementedError('Method not implemented!')
+ def CollectTopology(self, request, context):
+ """Missing associated documentation comment in .proto file."""
+ context.set_code(grpc.StatusCode.UNIMPLEMENTED)
+ context.set_details('Method not implemented!')
+ raise NotImplementedError('Method not implemented!')
- def SendOpaqueStatus(self, request, context):
- """Missing associated documentation comment in .proto file."""
- context.set_code(grpc.StatusCode.UNIMPLEMENTED)
- context.set_details('Method not implemented!')
- raise NotImplementedError('Method not implemented!')
+ def SendResult(self, request, context):
+ """Missing associated documentation comment in .proto file."""
+ context.set_code(grpc.StatusCode.UNIMPLEMENTED)
+ context.set_details('Method not implemented!')
+ raise NotImplementedError('Method not implemented!')
+
+ def SendOpaqueStatus(self, request, context):
+ """Missing associated documentation comment in .proto file."""
+ context.set_code(grpc.StatusCode.UNIMPLEMENTED)
+ context.set_details('Method not implemented!')
+ raise NotImplementedError('Method not implemented!')
+
+ def HealthCheck(self, request, context):
+ """Missing associated documentation comment in .proto file."""
+ context.set_code(grpc.StatusCode.UNIMPLEMENTED)
+ context.set_details('Method not implemented!')
+ raise NotImplementedError('Method not implemented!')
def add_NodeServiceServicer_to_server(servicer, server):
- rpc_method_handlers = {
- 'SendPrompt':
- grpc.unary_unary_rpc_method_handler(
- servicer.SendPrompt,
- request_deserializer=node__service__pb2.PromptRequest.FromString,
- response_serializer=node__service__pb2.Tensor.SerializeToString,
- ),
- 'SendTensor':
- grpc.unary_unary_rpc_method_handler(
- servicer.SendTensor,
- request_deserializer=node__service__pb2.TensorRequest.FromString,
- response_serializer=node__service__pb2.Tensor.SerializeToString,
- ),
- 'GetInferenceResult':
- grpc.unary_unary_rpc_method_handler(
- servicer.GetInferenceResult,
- request_deserializer=node__service__pb2.GetInferenceResultRequest.FromString,
- response_serializer=node__service__pb2.InferenceResult.SerializeToString,
- ),
- 'CollectTopology':
- grpc.unary_unary_rpc_method_handler(
- servicer.CollectTopology,
- request_deserializer=node__service__pb2.CollectTopologyRequest.FromString,
- response_serializer=node__service__pb2.Topology.SerializeToString,
- ),
- 'SendResult':
- grpc.unary_unary_rpc_method_handler(
- servicer.SendResult,
- request_deserializer=node__service__pb2.SendResultRequest.FromString,
- response_serializer=node__service__pb2.Empty.SerializeToString,
- ),
- 'SendOpaqueStatus':
- grpc.unary_unary_rpc_method_handler(
- servicer.SendOpaqueStatus,
- request_deserializer=node__service__pb2.SendOpaqueStatusRequest.FromString,
- response_serializer=node__service__pb2.Empty.SerializeToString,
- ),
- }
- generic_handler = grpc.method_handlers_generic_handler('node_service.NodeService', rpc_method_handlers)
- server.add_generic_rpc_handlers((generic_handler,))
- server.add_registered_method_handlers('node_service.NodeService', rpc_method_handlers)
+ rpc_method_handlers = {
+ 'SendPrompt': grpc.unary_unary_rpc_method_handler(
+ servicer.SendPrompt,
+ request_deserializer=node__service__pb2.PromptRequest.FromString,
+ response_serializer=node__service__pb2.Tensor.SerializeToString,
+ ),
+ 'SendTensor': grpc.unary_unary_rpc_method_handler(
+ servicer.SendTensor,
+ request_deserializer=node__service__pb2.TensorRequest.FromString,
+ response_serializer=node__service__pb2.Tensor.SerializeToString,
+ ),
+ 'GetInferenceResult': grpc.unary_unary_rpc_method_handler(
+ servicer.GetInferenceResult,
+ request_deserializer=node__service__pb2.GetInferenceResultRequest.FromString,
+ response_serializer=node__service__pb2.InferenceResult.SerializeToString,
+ ),
+ 'CollectTopology': grpc.unary_unary_rpc_method_handler(
+ servicer.CollectTopology,
+ request_deserializer=node__service__pb2.CollectTopologyRequest.FromString,
+ response_serializer=node__service__pb2.Topology.SerializeToString,
+ ),
+ 'SendResult': grpc.unary_unary_rpc_method_handler(
+ servicer.SendResult,
+ request_deserializer=node__service__pb2.SendResultRequest.FromString,
+ response_serializer=node__service__pb2.Empty.SerializeToString,
+ ),
+ 'SendOpaqueStatus': grpc.unary_unary_rpc_method_handler(
+ servicer.SendOpaqueStatus,
+ request_deserializer=node__service__pb2.SendOpaqueStatusRequest.FromString,
+ response_serializer=node__service__pb2.Empty.SerializeToString,
+ ),
+ 'HealthCheck': grpc.unary_unary_rpc_method_handler(
+ servicer.HealthCheck,
+ request_deserializer=node__service__pb2.HealthCheckRequest.FromString,
+ response_serializer=node__service__pb2.HealthCheckResponse.SerializeToString,
+ ),
+ }
+ generic_handler = grpc.method_handlers_generic_handler(
+ 'node_service.NodeService', rpc_method_handlers)
+ server.add_generic_rpc_handlers((generic_handler,))
+ server.add_registered_method_handlers('node_service.NodeService', rpc_method_handlers)
-# This class is part of an EXPERIMENTAL API.
+ # This class is part of an EXPERIMENTAL API.
class NodeService(object):
- """Missing associated documentation comment in .proto file."""
- @staticmethod
- def SendPrompt(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None):
- return grpc.experimental.unary_unary(
- request,
- target,
- '/node_service.NodeService/SendPrompt',
- node__service__pb2.PromptRequest.SerializeToString,
- node__service__pb2.Tensor.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True
- )
+ """Missing associated documentation comment in .proto file."""
- @staticmethod
- def SendTensor(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None):
- return grpc.experimental.unary_unary(
- request,
- target,
- '/node_service.NodeService/SendTensor',
- node__service__pb2.TensorRequest.SerializeToString,
- node__service__pb2.Tensor.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True
- )
+ @staticmethod
+ def SendPrompt(request,
+ target,
+ options=(),
+ channel_credentials=None,
+ call_credentials=None,
+ insecure=False,
+ compression=None,
+ wait_for_ready=None,
+ timeout=None,
+ metadata=None):
+ return grpc.experimental.unary_unary(
+ request,
+ target,
+ '/node_service.NodeService/SendPrompt',
+ node__service__pb2.PromptRequest.SerializeToString,
+ node__service__pb2.Tensor.FromString,
+ options,
+ channel_credentials,
+ insecure,
+ call_credentials,
+ compression,
+ wait_for_ready,
+ timeout,
+ metadata,
+ _registered_method=True)
- @staticmethod
- def GetInferenceResult(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None):
- return grpc.experimental.unary_unary(
- request,
- target,
- '/node_service.NodeService/GetInferenceResult',
- node__service__pb2.GetInferenceResultRequest.SerializeToString,
- node__service__pb2.InferenceResult.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True
- )
+ @staticmethod
+ def SendTensor(request,
+ target,
+ options=(),
+ channel_credentials=None,
+ call_credentials=None,
+ insecure=False,
+ compression=None,
+ wait_for_ready=None,
+ timeout=None,
+ metadata=None):
+ return grpc.experimental.unary_unary(
+ request,
+ target,
+ '/node_service.NodeService/SendTensor',
+ node__service__pb2.TensorRequest.SerializeToString,
+ node__service__pb2.Tensor.FromString,
+ options,
+ channel_credentials,
+ insecure,
+ call_credentials,
+ compression,
+ wait_for_ready,
+ timeout,
+ metadata,
+ _registered_method=True)
- @staticmethod
- def CollectTopology(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None):
- return grpc.experimental.unary_unary(
- request,
- target,
- '/node_service.NodeService/CollectTopology',
- node__service__pb2.CollectTopologyRequest.SerializeToString,
- node__service__pb2.Topology.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True
- )
+ @staticmethod
+ def GetInferenceResult(request,
+ target,
+ options=(),
+ channel_credentials=None,
+ call_credentials=None,
+ insecure=False,
+ compression=None,
+ wait_for_ready=None,
+ timeout=None,
+ metadata=None):
+ return grpc.experimental.unary_unary(
+ request,
+ target,
+ '/node_service.NodeService/GetInferenceResult',
+ node__service__pb2.GetInferenceResultRequest.SerializeToString,
+ node__service__pb2.InferenceResult.FromString,
+ options,
+ channel_credentials,
+ insecure,
+ call_credentials,
+ compression,
+ wait_for_ready,
+ timeout,
+ metadata,
+ _registered_method=True)
- @staticmethod
- def SendResult(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None):
- return grpc.experimental.unary_unary(
- request,
- target,
- '/node_service.NodeService/SendResult',
- node__service__pb2.SendResultRequest.SerializeToString,
- node__service__pb2.Empty.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True
- )
+ @staticmethod
+ def CollectTopology(request,
+ target,
+ options=(),
+ channel_credentials=None,
+ call_credentials=None,
+ insecure=False,
+ compression=None,
+ wait_for_ready=None,
+ timeout=None,
+ metadata=None):
+ return grpc.experimental.unary_unary(
+ request,
+ target,
+ '/node_service.NodeService/CollectTopology',
+ node__service__pb2.CollectTopologyRequest.SerializeToString,
+ node__service__pb2.Topology.FromString,
+ options,
+ channel_credentials,
+ insecure,
+ call_credentials,
+ compression,
+ wait_for_ready,
+ timeout,
+ metadata,
+ _registered_method=True)
- @staticmethod
- def SendOpaqueStatus(request, target, options=(), channel_credentials=None, call_credentials=None, insecure=False, compression=None, wait_for_ready=None, timeout=None, metadata=None):
- return grpc.experimental.unary_unary(
- request,
- target,
- '/node_service.NodeService/SendOpaqueStatus',
- node__service__pb2.SendOpaqueStatusRequest.SerializeToString,
- node__service__pb2.Empty.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True
- )
+ @staticmethod
+ def SendResult(request,
+ target,
+ options=(),
+ channel_credentials=None,
+ call_credentials=None,
+ insecure=False,
+ compression=None,
+ wait_for_ready=None,
+ timeout=None,
+ metadata=None):
+ return grpc.experimental.unary_unary(
+ request,
+ target,
+ '/node_service.NodeService/SendResult',
+ node__service__pb2.SendResultRequest.SerializeToString,
+ node__service__pb2.Empty.FromString,
+ options,
+ channel_credentials,
+ insecure,
+ call_credentials,
+ compression,
+ wait_for_ready,
+ timeout,
+ metadata,
+ _registered_method=True)
+
+ @staticmethod
+ def SendOpaqueStatus(request,
+ target,
+ options=(),
+ channel_credentials=None,
+ call_credentials=None,
+ insecure=False,
+ compression=None,
+ wait_for_ready=None,
+ timeout=None,
+ metadata=None):
+ return grpc.experimental.unary_unary(
+ request,
+ target,
+ '/node_service.NodeService/SendOpaqueStatus',
+ node__service__pb2.SendOpaqueStatusRequest.SerializeToString,
+ node__service__pb2.Empty.FromString,
+ options,
+ channel_credentials,
+ insecure,
+ call_credentials,
+ compression,
+ wait_for_ready,
+ timeout,
+ metadata,
+ _registered_method=True)
+
+ @staticmethod
+ def HealthCheck(request,
+ target,
+ options=(),
+ channel_credentials=None,
+ call_credentials=None,
+ insecure=False,
+ compression=None,
+ wait_for_ready=None,
+ timeout=None,
+ metadata=None):
+ return grpc.experimental.unary_unary(
+ request,
+ target,
+ '/node_service.NodeService/HealthCheck',
+ node__service__pb2.HealthCheckRequest.SerializeToString,
+ node__service__pb2.HealthCheckResponse.FromString,
+ options,
+ channel_credentials,
+ insecure,
+ call_credentials,
+ compression,
+ wait_for_ready,
+ timeout,
+ metadata,
+ _registered_method=True)
diff --git a/exo/networking/peer_handle.py b/exo/networking/peer_handle.py
index 390966bd..515cb1d1 100644
--- a/exo/networking/peer_handle.py
+++ b/exo/networking/peer_handle.py
@@ -30,6 +30,10 @@ class PeerHandle(ABC):
async def disconnect(self) -> None:
pass
+ @abstractmethod
+ async def health_check(self) -> bool:
+ pass
+
@abstractmethod
async def send_prompt(self, shard: Shard, prompt: str, image_str: Optional[str] = None, request_id: Optional[str] = None, inference_state: Optional[str] = None) -> Optional[np.array]:
pass
diff --git a/exo/networking/udp_discovery.py b/exo/networking/udp_discovery.py
index 733cd1b9b..30859ec7 100644
--- a/exo/networking/udp_discovery.py
+++ b/exo/networking/udp_discovery.py
@@ -139,15 +139,21 @@ class UDPDiscovery(Discovery):
peer_host = addr[0]
peer_port = message["grpc_port"]
device_capabilities = DeviceCapabilities(**message["device_capabilities"])
- if peer_id not in self.known_peers or self.known_peers[peer_id][0].addr() != f"{peer_host}:{peer_port}":
- if DEBUG >= 1: print(
- f"Adding {peer_id=} at {peer_host}:{peer_port}. Replace existing peer_id: {peer_id in self.known_peers}")
- self.known_peers[peer_id] = (
- self.create_peer_handle(peer_id, f"{peer_host}:{peer_port}", device_capabilities),
- time.time(),
- time.time(),
- )
- self.known_peers[peer_id] = (self.known_peers[peer_id][0], self.known_peers[peer_id][1], time.time())
+
+ # Create a new peer handle
+ new_peer_handle = self.create_peer_handle(peer_id, f"{peer_host}:{peer_port}", device_capabilities)
+
+ # Check if the new peer is healthy before adding
+ if await new_peer_handle.health_check():
+ if peer_id not in self.known_peers or self.known_peers[peer_id][0].addr() != f"{peer_host}:{peer_port}":
+ if DEBUG >= 1: print(
+ f"Adding {peer_id=} at {peer_host}:{peer_port}. Replace existing peer_id: {peer_id in self.known_peers}")
+ self.known_peers[peer_id] = (new_peer_handle, time.time(), time.time())
+ else:
+ # Update last seen time for existing peer
+ self.known_peers[peer_id] = (self.known_peers[peer_id][0], self.known_peers[peer_id][1], time.time())
+ else:
+ if DEBUG >= 1: print(f"Peer {peer_id} at {peer_host}:{peer_port} failed health check. Not adding.")
async def task_listen_for_peers(self):
await asyncio.get_event_loop().create_datagram_endpoint(lambda: ListenProtocol(self.on_listen_message),
@@ -158,14 +164,18 @@ class UDPDiscovery(Discovery):
while True:
try:
current_time = time.time()
- peers_to_remove = [
- peer_handle.id() for peer_handle, connected_at, last_seen in self.known_peers.values()
- if (not await peer_handle.is_connected() and current_time - connected_at > self.discovery_timeout) or current_time - last_seen > self.discovery_timeout
- ]
- if DEBUG_DISCOVERY >= 2: print("Peer statuses:", {peer_handle.id(): f"is_connected={await peer_handle.is_connected()}, {connected_at=}, {last_seen=}" for peer_handle, connected_at, last_seen in self.known_peers.values()})
+ peers_to_remove = []
+ for peer_id, (peer_handle, connected_at, last_seen) in self.known_peers.items():
+ if (not await peer_handle.is_connected() and current_time - connected_at > self.discovery_timeout) or \
+ current_time - last_seen > self.discovery_timeout or \
+ not await peer_handle.health_check():
+ peers_to_remove.append(peer_id)
+
+ if DEBUG_DISCOVERY >= 2: print("Peer statuses:", {peer_handle.id(): f"is_connected={await peer_handle.is_connected()}, health_check={await peer_handle.health_check()}, {connected_at=}, {last_seen=}" for peer_handle, connected_at, last_seen in self.known_peers.values()})
+
for peer_id in peers_to_remove:
if peer_id in self.known_peers: del self.known_peers[peer_id]
- if DEBUG_DISCOVERY >= 2: print(f"Removed peer {peer_id} due to inactivity.")
+ if DEBUG_DISCOVERY >= 2: print(f"Removed peer {peer_id} due to inactivity or failed health check.")
except Exception as e:
print(f"Error in cleanup peers: {e}")
print(traceback.format_exc())
← 6e323949 comparing timestamps across distributed systems of consumer
·
back to Exo
·
implement grpc health check 26dc1989 →