Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,16 @@ await tunnel.forward_to("127.0.0.1", 8000)
It keeps accepting rstream streams and relays them to the local TCP service
until the tunnel or control channel is closed.

The control channel uses a negotiated heartbeat deadline. If that deadline
expires, or the control transport ends unexpectedly, the SDK immediately stops
accepting new streams but lets already accepted streams and active
`forward_to()` relays drain normally. This avoids interrupting an established
application session because of a transient or asymmetric control-path outage.
An explicit tunnel/control close or a protocol violation remains a hard close
and terminates active local forwarding. Applications should use each stream's
EOF/error as its payload-lifecycle signal; `control.done()` only describes the
control plane.

## Private dial

```python
Expand Down
14 changes: 12 additions & 2 deletions proto/rstream.proto
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@ extend google.protobuf.FieldOptions {
string access = 51234;
}

option (protocol_version) = "1.4.4";
option (protocol_version) = "1.4.5";

package rstream.io_rstrm.protobuf;

Expand Down Expand Up @@ -136,6 +136,7 @@ message TunnelProperties {

message OpenControlChannelReq {
ClientDetails client_details = 1;
ControlChannelLiveness liveness = 2;
}

// The server responds to an 'OpenControlChannelReq' message with an
Expand All @@ -148,6 +149,7 @@ message OpenControlChannelRsp {
message Ok {
string client_id = 1;
ServerDetails server_details = 2;
ControlChannelLiveness liveness = 3;
}
oneof payload {
Ok ok = 1;
Expand Down Expand Up @@ -263,7 +265,15 @@ message DatagramChannelClose {

// Sent by client and/or server to maintain the control channel active

message Heartbeat { }
message ControlChannelLiveness {
uint32 heartbeat_interval_ms = 1;
uint32 heartbeat_timeout_ms = 2;
}

message Heartbeat {
uint64 sequence = 1;
uint64 acknowledgement = 2;
}

// Allows the server to send unsolicited messages to the client
message ServerMessage {
Expand Down
92 changes: 47 additions & 45 deletions src/rstream/_proto/rstream_pb2.py

Large diffs are not rendered by default.

28 changes: 22 additions & 6 deletions src/rstream/_proto/rstream_pb2.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -159,20 +159,24 @@ class TunnelProperties(_message.Message):
def __init__(self, id: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., creation_date: _Optional[_Union[datetime.datetime, _timestamp_pb2.Timestamp, _Mapping]] = ..., name: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., type: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., publish: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ..., protocol: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., labels: _Optional[_Mapping[str, str]] = ..., geoip: _Optional[_Iterable[str]] = ..., trusted_ips: _Optional[_Iterable[str]] = ..., host: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., tls_mode: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., tls_alpns: _Optional[_Iterable[str]] = ..., tls_min_version: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., tls_ciphers: _Optional[_Iterable[str]] = ..., mtls_auth: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ..., mtls_cacert_pem: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., http_version: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., http_use_tls: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ..., token_auth: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ..., rstream_auth: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ..., challenge_mode: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ..., hostname: _Optional[_Union[_wrappers_pb2.StringValue, _Mapping]] = ..., port: _Optional[_Union[_wrappers_pb2.UInt32Value, _Mapping]] = ..., upstream_tls: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ..., datagram_guaranteed_delivery: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ..., allow_cross_region_routing: _Optional[_Union[_wrappers_pb2.BoolValue, _Mapping]] = ...) -> None: ...

class OpenControlChannelReq(_message.Message):
__slots__ = ("client_details",)
__slots__ = ("client_details", "liveness")
CLIENT_DETAILS_FIELD_NUMBER: _ClassVar[int]
LIVENESS_FIELD_NUMBER: _ClassVar[int]
client_details: ClientDetails
def __init__(self, client_details: _Optional[_Union[ClientDetails, _Mapping]] = ...) -> None: ...
liveness: ControlChannelLiveness
def __init__(self, client_details: _Optional[_Union[ClientDetails, _Mapping]] = ..., liveness: _Optional[_Union[ControlChannelLiveness, _Mapping]] = ...) -> None: ...

class OpenControlChannelRsp(_message.Message):
__slots__ = ("ok", "error")
class Ok(_message.Message):
__slots__ = ("client_id", "server_details")
__slots__ = ("client_id", "server_details", "liveness")
CLIENT_ID_FIELD_NUMBER: _ClassVar[int]
SERVER_DETAILS_FIELD_NUMBER: _ClassVar[int]
LIVENESS_FIELD_NUMBER: _ClassVar[int]
client_id: str
server_details: ServerDetails
def __init__(self, client_id: _Optional[str] = ..., server_details: _Optional[_Union[ServerDetails, _Mapping]] = ...) -> None: ...
liveness: ControlChannelLiveness
def __init__(self, client_id: _Optional[str] = ..., server_details: _Optional[_Union[ServerDetails, _Mapping]] = ..., liveness: _Optional[_Union[ControlChannelLiveness, _Mapping]] = ...) -> None: ...
OK_FIELD_NUMBER: _ClassVar[int]
ERROR_FIELD_NUMBER: _ClassVar[int]
ok: OpenControlChannelRsp.Ok
Expand Down Expand Up @@ -283,9 +287,21 @@ class DatagramChannelClose(_message.Message):
error: Error
def __init__(self, stream_id: _Optional[str] = ..., error: _Optional[_Union[Error, _Mapping]] = ...) -> None: ...

class ControlChannelLiveness(_message.Message):
__slots__ = ("heartbeat_interval_ms", "heartbeat_timeout_ms")
HEARTBEAT_INTERVAL_MS_FIELD_NUMBER: _ClassVar[int]
HEARTBEAT_TIMEOUT_MS_FIELD_NUMBER: _ClassVar[int]
heartbeat_interval_ms: int
heartbeat_timeout_ms: int
def __init__(self, heartbeat_interval_ms: _Optional[int] = ..., heartbeat_timeout_ms: _Optional[int] = ...) -> None: ...

class Heartbeat(_message.Message):
__slots__ = ()
def __init__(self) -> None: ...
__slots__ = ("sequence", "acknowledgement")
SEQUENCE_FIELD_NUMBER: _ClassVar[int]
ACKNOWLEDGEMENT_FIELD_NUMBER: _ClassVar[int]
sequence: int
acknowledgement: int
def __init__(self, sequence: _Optional[int] = ..., acknowledgement: _Optional[int] = ...) -> None: ...

class ServerMessage(_message.Message):
__slots__ = ("message",)
Expand Down
38 changes: 37 additions & 1 deletion src/rstream/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,28 @@
from rstream.stream import RstreamStream

_T = TypeVar("_T")
_MAX_HEARTBEAT_TIMEOUT_MS = 900_000


def _negotiated_heartbeat_timeout(
heartbeat: bool,
heartbeat_interval_ms: int,
response: pb.OpenControlChannelRsp.Ok,
) -> float:
if not response.HasField("liveness"):
return 0.0
liveness = response.liveness
if (
not heartbeat
or liveness.heartbeat_interval_ms != heartbeat_interval_ms
or liveness.heartbeat_timeout_ms < liveness.heartbeat_interval_ms
or liveness.heartbeat_timeout_ms > _MAX_HEARTBEAT_TIMEOUT_MS
):
raise ProtocolError(
"Engine returned an invalid liveness policy.",
code="ERR_RSTREAM_PROTOCOL",
)
return liveness.heartbeat_timeout_ms / 1_000


class Client:
Expand Down Expand Up @@ -146,8 +168,17 @@ async def connect(self) -> ControlChannel:
engine = await self._resolve_engine(resolved)
token = await self._resolve_token(resolved, engine)
reader, writer = await self._dial_engine(engine, resolved)
heartbeat_interval_ms = (
round(resolved.heartbeat_interval * 1_000) if resolved.heartbeat else 0
)
try:
await write_message(writer, message_with_open_control_channel_req(token))
await write_message(
writer,
message_with_open_control_channel_req(
token,
heartbeat_interval_ms if resolved.heartbeat else None,
),
)
response = await _wait_for_operation(
read_message(reader),
resolved.operation_timeout,
Expand All @@ -172,6 +203,11 @@ async def connect(self) -> ControlChannel:
writer,
heartbeat=resolved.heartbeat,
heartbeat_interval=resolved.heartbeat_interval,
heartbeat_timeout=_negotiated_heartbeat_timeout(
resolved.heartbeat,
heartbeat_interval_ms,
payload.ok,
),
operation_timeout=resolved.operation_timeout,
open_proxy_connection=lambda request: self._open_proxy_connection(
engine,
Expand Down
27 changes: 27 additions & 0 deletions src/rstream/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import base64
import json
import math
import os
import re
import ssl
Expand All @@ -17,6 +18,8 @@
from rstream.errors import ConfigurationError, UnsupportedFeatureError

DEFAULT_API_URL = "https://rstream.io"
MIN_HEARTBEAT_INTERVAL_MS = 1_000
MAX_HEARTBEAT_INTERVAL_MS = 300_000


@dataclass(frozen=True)
Expand Down Expand Up @@ -229,6 +232,30 @@ async def resolve_client_options(options: ClientOptions) -> ResolvedClientOption
"operation_timeout must be positive.",
code="ERR_RSTREAM_INVALID_TIMEOUT",
)
if options.heartbeat and not math.isfinite(options.heartbeat_interval):
raise ConfigurationError(
"heartbeat_interval must be between 1 and 300 seconds with "
"millisecond precision.",
code="ERR_RSTREAM_INVALID_CONFIG",
)
heartbeat_interval_ms = (
round(options.heartbeat_interval * 1_000) if options.heartbeat else 0
)
if options.heartbeat and (
not math.isclose(
options.heartbeat_interval * 1_000,
heartbeat_interval_ms,
abs_tol=1e-9,
)
or not MIN_HEARTBEAT_INTERVAL_MS
<= heartbeat_interval_ms
<= MAX_HEARTBEAT_INTERVAL_MS
):
raise ConfigurationError(
"heartbeat_interval must be between 1 and 300 seconds with "
"millisecond precision.",
code="ERR_RSTREAM_INVALID_CONFIG",
)
region = _normalize_region(
_first_defined(options.region, env.region, config.region)
)
Expand Down
Loading