import socket
import struct
import random
import coloredlogs
import logging
import sys
import io
import os
import asyncio
import cbor2
import pprintpp
import toml
import hashlib
from typing import Tuple, Any, Dict, List, Callable
from .transport import DialConfig, TcpDialConfig, UnixDialConfig
MAX_MESSAGE_SIZE = 40 * 1024 * 1024
REPLICA_SUCCESS = 0
REPLICA_ERROR_BOX_ID_NOT_FOUND = 1
REPLICA_ERROR_INVALID_BOX_ID = 2
REPLICA_ERROR_INVALID_SIGNATURE = 3
REPLICA_ERROR_DATABASE_FAILURE = 4
REPLICA_ERROR_INVALID_PAYLOAD = 5
REPLICA_ERROR_STORAGE_FULL = 6
REPLICA_ERROR_INTERNAL_ERROR = 7
REPLICA_ERROR_INVALID_EPOCH = 8
REPLICA_ERROR_REPLICATION_FAILED = 9
REPLICA_ERROR_BOX_ALREADY_EXISTS = 10
REPLICA_ERROR_TOMBSTONE = 11
THIN_CLIENT_SUCCESS = 0
THIN_CLIENT_ERROR_CONNECTION_LOST = 1
THIN_CLIENT_ERROR_TIMEOUT = 2
THIN_CLIENT_ERROR_INVALID_REQUEST = 3
THIN_CLIENT_ERROR_INTERNAL_ERROR = 4
THIN_CLIENT_ERROR_MAX_RETRIES = 5
THIN_CLIENT_ERROR_INVALID_CHANNEL = 6
THIN_CLIENT_ERROR_CHANNEL_NOT_FOUND = 7
THIN_CLIENT_ERROR_PERMISSION_DENIED = 8
THIN_CLIENT_ERROR_INVALID_PAYLOAD = 9
THIN_CLIENT_ERROR_SERVICE_UNAVAILABLE = 10
THIN_CLIENT_ERROR_DUPLICATE_CAPABILITY = 11
THIN_CLIENT_ERROR_COURIER_CACHE_CORRUPTION = 12
THIN_CLIENT_PROPAGATION_ERROR = 13
THIN_CLIENT_ERROR_INVALID_WRITE_CAPABILITY = 14
THIN_CLIENT_ERROR_INVALID_READ_CAPABILITY = 15
THIN_CLIENT_ERROR_INVALID_RESUME_WRITE_CHANNEL_REQUEST = 16
THIN_CLIENT_ERROR_INVALID_RESUME_READ_CHANNEL_REQUEST = 17
THIN_CLIENT_IMPOSSIBLE_HASH_ERROR = 18
THIN_CLIENT_IMPOSSIBLE_NEW_WRITE_CAP_ERROR = 19
THIN_CLIENT_IMPOSSIBLE_NEW_STATEFUL_WRITER_ERROR = 20
THIN_CLIENT_CAPABILITY_ALREADY_IN_USE = 21
THIN_CLIENT_ERROR_MKEM_DECRYPTION_FAILED = 22
THIN_CLIENT_ERROR_BACAP_DECRYPTION_FAILED = 23
THIN_CLIENT_ERROR_START_RESENDING_CANCELLED = 24
THIN_CLIENT_ERROR_INVALID_TOMBSTONE_SIG = 25
THIN_CLIENT_ERROR_COPY_COMMAND_FAILED = 26
THIN_CLIENT_ERROR_PAYLOAD_TOO_LARGE = 27
def thin_client_error_to_string(error_code: int) -> str:
error_messages = {
THIN_CLIENT_SUCCESS: "Success",
THIN_CLIENT_ERROR_CONNECTION_LOST: "Connection lost",
THIN_CLIENT_ERROR_TIMEOUT: "Timeout",
THIN_CLIENT_ERROR_INVALID_REQUEST: "Invalid request",
THIN_CLIENT_ERROR_INTERNAL_ERROR: "Internal error",
THIN_CLIENT_ERROR_MAX_RETRIES: "Maximum retries exceeded",
THIN_CLIENT_ERROR_INVALID_CHANNEL: "Invalid channel",
THIN_CLIENT_ERROR_CHANNEL_NOT_FOUND: "Channel not found",
THIN_CLIENT_ERROR_PERMISSION_DENIED: "Permission denied",
THIN_CLIENT_ERROR_INVALID_PAYLOAD: "Invalid payload",
THIN_CLIENT_ERROR_SERVICE_UNAVAILABLE: "Service unavailable",
THIN_CLIENT_ERROR_DUPLICATE_CAPABILITY: "Duplicate capability",
THIN_CLIENT_ERROR_COURIER_CACHE_CORRUPTION: "Courier cache corruption",
THIN_CLIENT_PROPAGATION_ERROR: "Propagation error",
THIN_CLIENT_ERROR_INVALID_WRITE_CAPABILITY: "Invalid write capability",
THIN_CLIENT_ERROR_INVALID_READ_CAPABILITY: "Invalid read capability",
THIN_CLIENT_ERROR_INVALID_RESUME_WRITE_CHANNEL_REQUEST: "Invalid resume write channel request",
THIN_CLIENT_ERROR_INVALID_RESUME_READ_CHANNEL_REQUEST: "Invalid resume read channel request",
THIN_CLIENT_IMPOSSIBLE_HASH_ERROR: "Impossible hash error",
THIN_CLIENT_IMPOSSIBLE_NEW_WRITE_CAP_ERROR: "Failed to create new write capability",
THIN_CLIENT_IMPOSSIBLE_NEW_STATEFUL_WRITER_ERROR: "Failed to create new stateful writer",
THIN_CLIENT_CAPABILITY_ALREADY_IN_USE: "Capability already in use",
THIN_CLIENT_ERROR_MKEM_DECRYPTION_FAILED: "MKEM decryption failed",
THIN_CLIENT_ERROR_BACAP_DECRYPTION_FAILED: "BACAP decryption failed",
THIN_CLIENT_ERROR_START_RESENDING_CANCELLED: "Start resending cancelled",
THIN_CLIENT_ERROR_INVALID_TOMBSTONE_SIG: "Invalid tombstone signature",
THIN_CLIENT_ERROR_COPY_COMMAND_FAILED: "Copy command failed",
THIN_CLIENT_ERROR_PAYLOAD_TOO_LARGE: "Payload too large",
}
return error_messages.get(error_code, f"Unknown thin client error code: {error_code}")
class ConfigError(Exception):
pass
class ReplicaError(Exception):
pass
class BoxIDNotFoundError(ReplicaError):
pass
class InvalidBoxIDError(ReplicaError):
pass
class InvalidSignatureError(ReplicaError):
pass
class DatabaseFailureError(ReplicaError):
pass
class InvalidPayloadError(ReplicaError):
pass
class StorageFullError(ReplicaError):
pass
class ReplicaInternalError(ReplicaError):
pass
class InvalidEpochError(ReplicaError):
pass
class ReplicationFailedError(ReplicaError):
pass
class BoxAlreadyExistsError(ReplicaError):
pass
class TombstoneError(ReplicaError):
pass
class InvalidTombstoneSignatureError(Exception):
pass
class MKEMDecryptionFailedError(Exception):
pass
class BACAPDecryptionFailedError(Exception):
pass
class StartResendingCancelledError(Exception):
pass
class CopyCommandFailedError(Exception):
def __init__(self, replica_error_code: int = 0, failed_envelope_index: int = 0) -> None:
self.replica_error_code = replica_error_code
self.failed_envelope_index = failed_envelope_index
super().__init__(
f"copy command failed: replica_error_code={replica_error_code}, "
f"failed_envelope_index={failed_envelope_index}"
)
class PayloadTooLargeError(Exception):
pass
def error_code_to_exception(error_code: int) -> Exception:
if error_code == REPLICA_SUCCESS:
return None
if error_code == REPLICA_ERROR_BOX_ID_NOT_FOUND: return BoxIDNotFoundError("box ID not found")
elif error_code == REPLICA_ERROR_INVALID_BOX_ID: return InvalidBoxIDError("invalid box ID")
elif error_code == REPLICA_ERROR_INVALID_SIGNATURE: return InvalidSignatureError("invalid signature")
elif error_code == REPLICA_ERROR_DATABASE_FAILURE: return DatabaseFailureError("database failure")
elif error_code == REPLICA_ERROR_INVALID_PAYLOAD: return InvalidPayloadError("invalid payload")
elif error_code == REPLICA_ERROR_STORAGE_FULL: return StorageFullError("storage full")
elif error_code == REPLICA_ERROR_INTERNAL_ERROR: return ReplicaInternalError("replica internal error")
elif error_code == REPLICA_ERROR_INVALID_EPOCH: return InvalidEpochError("invalid epoch")
elif error_code == REPLICA_ERROR_REPLICATION_FAILED: return ReplicationFailedError("replication failed")
elif error_code == REPLICA_ERROR_BOX_ALREADY_EXISTS: return BoxAlreadyExistsError("box already exists")
elif error_code == REPLICA_ERROR_TOMBSTONE: return TombstoneError("tombstone")
elif error_code == THIN_CLIENT_ERROR_MKEM_DECRYPTION_FAILED: return MKEMDecryptionFailedError("MKEM decryption failed")
elif error_code == THIN_CLIENT_ERROR_BACAP_DECRYPTION_FAILED: return BACAPDecryptionFailedError("BACAP decryption failed")
elif error_code == THIN_CLIENT_ERROR_START_RESENDING_CANCELLED: return StartResendingCancelledError("start resending cancelled")
elif error_code == THIN_CLIENT_ERROR_INVALID_TOMBSTONE_SIG: return InvalidTombstoneSignatureError("invalid tombstone signature")
elif error_code == THIN_CLIENT_ERROR_PAYLOAD_TOO_LARGE: return PayloadTooLargeError("payload too large")
else:
return Exception(thin_client_error_to_string(error_code))
def copy_reply_to_exception(reply: "Dict[str, Any]") -> "Exception | None":
error_code = reply.get("error_code", 0)
if error_code == THIN_CLIENT_SUCCESS:
return None
if error_code == THIN_CLIENT_ERROR_COPY_COMMAND_FAILED:
return CopyCommandFailedError(
replica_error_code=reply.get("replica_error_code", 0),
failed_envelope_index=reply.get("failed_envelope_index", 0),
)
return error_code_to_exception(error_code)
def is_expected_outcome(exc: Exception) -> bool:
return isinstance(exc, (TombstoneError, BoxIDNotFoundError, BoxAlreadyExistsError))
class ThinClientOfflineError(Exception):
pass
SURB_ID_SIZE = 16
MESSAGE_ID_SIZE = 16
class Geometry:
def __init__(self, *, PacketLength:int, NrHops:int, HeaderLength:int, RoutingInfoLength:int, PerHopRoutingInfoLength:int, SURBLength:int, SphinxPlaintextHeaderLength:int, PayloadTagLength:int, ForwardPayloadLength:int, UserForwardPayloadLength:int, NextNodeHopLength:int, SPRPKeyMaterialLength:int, NIKEName:str='', KEMName:str='') -> None:
self.PacketLength = PacketLength
self.NrHops = NrHops
self.HeaderLength = HeaderLength
self.RoutingInfoLength = RoutingInfoLength
self.PerHopRoutingInfoLength = PerHopRoutingInfoLength
self.SURBLength = SURBLength
self.SphinxPlaintextHeaderLength = SphinxPlaintextHeaderLength
self.PayloadTagLength = PayloadTagLength
self.ForwardPayloadLength = ForwardPayloadLength
self.UserForwardPayloadLength = UserForwardPayloadLength
self.NextNodeHopLength = NextNodeHopLength
self.SPRPKeyMaterialLength = SPRPKeyMaterialLength
self.NIKEName = NIKEName
self.KEMName = KEMName
def __str__(self) -> str:
return (
f"PacketLength: {self.PacketLength}\n"
f"NrHops: {self.NrHops}\n"
f"HeaderLength: {self.HeaderLength}\n"
f"RoutingInfoLength: {self.RoutingInfoLength}\n"
f"PerHopRoutingInfoLength: {self.PerHopRoutingInfoLength}\n"
f"SURBLength: {self.SURBLength}\n"
f"SphinxPlaintextHeaderLength: {self.SphinxPlaintextHeaderLength}\n"
f"PayloadTagLength: {self.PayloadTagLength}\n"
f"ForwardPayloadLength: {self.ForwardPayloadLength}\n"
f"UserForwardPayloadLength: {self.UserForwardPayloadLength}\n"
f"NextNodeHopLength: {self.NextNodeHopLength}\n"
f"SPRPKeyMaterialLength: {self.SPRPKeyMaterialLength}\n"
f"NIKEName: {self.NIKEName}\n"
f"KEMName: {self.KEMName}"
)
class PigeonholeGeometry:
LENGTH_PREFIX_SIZE = 4
def __init__(
self,
*,
max_plaintext_payload_length: int,
courier_query_read_length: int = 0,
courier_query_write_length: int = 0,
courier_query_reply_read_length: int = 0,
courier_query_reply_write_length: int = 0,
nike_name: str = "",
signature_scheme_name: str = "Ed25519"
) -> None:
self.max_plaintext_payload_length = max_plaintext_payload_length
self.courier_query_read_length = courier_query_read_length
self.courier_query_write_length = courier_query_write_length
self.courier_query_reply_read_length = courier_query_reply_read_length
self.courier_query_reply_write_length = courier_query_reply_write_length
self.nike_name = nike_name
self.signature_scheme_name = signature_scheme_name
def validate(self) -> None:
if self.max_plaintext_payload_length <= 0:
raise ValueError("max_plaintext_payload_length must be positive")
if not self.nike_name:
raise ValueError("nike_name must be set")
if self.signature_scheme_name != "Ed25519":
raise ValueError("signature_scheme_name must be 'Ed25519'")
def padded_payload_length(self) -> int:
return self.max_plaintext_payload_length + self.LENGTH_PREFIX_SIZE
def __str__(self) -> str:
return (
f"PigeonholeGeometry:\n"
f" max_plaintext_payload_length: {self.max_plaintext_payload_length} bytes\n"
f" courier_query_read_length: {self.courier_query_read_length} bytes\n"
f" courier_query_write_length: {self.courier_query_write_length} bytes\n"
f" courier_query_reply_read_length: {self.courier_query_reply_read_length} bytes\n"
f" courier_query_reply_write_length: {self.courier_query_reply_write_length} bytes\n"
f" nike_name: {self.nike_name}\n"
f" signature_scheme_name: {self.signature_scheme_name}"
)
_EXPECTED_TOP_LEVEL_KEYS = frozenset({"Dial"})
_PIGEONHOLE_TOML_TO_KWARG = {
"MaxPlaintextPayloadLength": "max_plaintext_payload_length",
"CourierQueryReadLength": "courier_query_read_length",
"CourierQueryWriteLength": "courier_query_write_length",
"CourierQueryReplyReadLength": "courier_query_reply_read_length",
"CourierQueryReplyWriteLength": "courier_query_reply_write_length",
"NIKEName": "nike_name",
"SignatureSchemeName": "signature_scheme_name",
}
class ConfigFile:
def __init__(
self,
dial: "DialConfig",
) -> None:
self.dial : "DialConfig" = dial
@classmethod
def load(cls, toml_path:str) -> "ConfigFile":
try:
with open(toml_path, 'r') as f:
data = toml.load(f)
except FileNotFoundError as e:
raise ConfigError(f"config: {toml_path}: file not found") from e
except toml.TomlDecodeError as e:
raise ConfigError(f"config: {toml_path}: TOML parse error: {e}") from e
if not isinstance(data, dict):
raise ConfigError(f"config: {toml_path}: top-level must be a table")
unknown = set(data.keys()) - _EXPECTED_TOP_LEVEL_KEYS
if unknown:
raise ConfigError(
f"config: {toml_path}: unknown top-level key(s) {sorted(unknown)}; "
f"expected exactly {sorted(_EXPECTED_TOP_LEVEL_KEYS)}"
)
missing = _EXPECTED_TOP_LEVEL_KEYS - set(data.keys())
if missing:
raise ConfigError(
f"config: {toml_path}: missing required top-level key(s) {sorted(missing)}"
)
dial = _load_dial(data["Dial"], toml_path)
return cls(dial)
def __str__(self) -> str:
return f"Dial: {self.dial}"
def _load_dial(dial_data: "Any", toml_path: str) -> "DialConfig":
if not isinstance(dial_data, dict):
raise ConfigError(
f"config: {toml_path}: [Dial] must be a table containing "
f"exactly one of [Dial.Unix] or [Dial.Tcp]"
)
try:
return DialConfig.from_toml_dict(dial_data)
except ValueError as e:
raise ConfigError(f"config: {toml_path}: [Dial]: {e}") from e
def _sphinx_geometry_from_event(geometry_data: "Any") -> Geometry:
if not isinstance(geometry_data, dict):
raise ConfigError("daemon sent a malformed sphinx_geometry (not a map)")
try:
return Geometry(**geometry_data)
except TypeError as e:
raise ConfigError(
f"daemon sent a sphinx_geometry with unknown or missing keys: {e}"
) from e
def _pigeonhole_geometry_from_event(geometry_data: "Any") -> PigeonholeGeometry:
if not isinstance(geometry_data, dict):
raise ConfigError("daemon sent a malformed pigeonhole_geometry (not a map)")
unknown = set(geometry_data.keys()) - set(_PIGEONHOLE_TOML_TO_KWARG.keys())
if unknown:
raise ConfigError(
f"daemon sent a pigeonhole_geometry with unknown key(s) {sorted(unknown)}"
)
kwargs = {_PIGEONHOLE_TOML_TO_KWARG[k]: v for k, v in geometry_data.items()}
try:
return PigeonholeGeometry(**kwargs)
except TypeError as e:
raise ConfigError(
f"daemon sent a pigeonhole_geometry with unknown or missing keys: {e}"
) from e
def pretty_print_obj(obj: "Any") -> str:
pp = pprintpp.PrettyPrinter(indent=4)
return pp.pformat(obj)
def blake2_256_sum(data:bytes) -> bytes:
return hashlib.blake2b(data, digest_size=32).digest()
class ServiceDescriptor:
def __init__(self, recipient_queue_id:bytes, mix_descriptor: "Dict[Any,Any]") -> None:
self.recipient_queue_id = recipient_queue_id
self.mix_descriptor = mix_descriptor
def to_destination(self) -> "Tuple[bytes,bytes]":
"provider identity key hash and queue id"
provider_id_hash = blake2_256_sum(self.mix_descriptor['IdentityKey'])
return (provider_id_hash, self.recipient_queue_id)
def find_services(capability:str, doc:"Dict[str,Any]") -> "List[ServiceDescriptor]":
services = []
for node in doc['ServiceNodes']:
mynode = cbor2.loads(node)
if 'Kaetzchen' in mynode:
for cap, details in mynode['Kaetzchen'].items():
if cap == capability:
service_desc = ServiceDescriptor(
recipient_queue_id=bytes(details['endpoint'], 'utf-8'), mix_descriptor=mynode
)
services.append(service_desc)
return services
class Config:
def __init__(self, filepath:str,
on_connection_status:"Callable|None"=None,
on_new_pki_document:"Callable|None"=None,
on_message_sent:"Callable|None"=None,
on_message_reply:"Callable|None"=None,
on_daemon_disconnected:"Callable|None"=None) -> None:
cfgfile = ConfigFile.load(filepath)
self.dial = cfgfile.dial
self.on_connection_status = on_connection_status
self.on_new_pki_document = on_new_pki_document
self.on_message_sent = on_message_sent
self.on_message_reply = on_message_reply
self.on_daemon_disconnected = on_daemon_disconnected
async def handle_connection_status_event(self, event: asyncio.Event) -> None:
if self.on_connection_status:
return await self.on_connection_status(event)
async def handle_new_pki_document_event(self, event: asyncio.Event) -> None:
if self.on_new_pki_document:
await self.on_new_pki_document(event)
async def handle_message_sent_event(self, event: asyncio.Event) -> None:
if self.on_message_sent:
await self.on_message_sent(event)
async def handle_message_reply_event(self, event: asyncio.Event) -> None:
if self.on_message_reply:
await self.on_message_reply(event)
async def handle_daemon_disconnected_event(self, event: dict) -> None:
if self.on_daemon_disconnected:
await self.on_daemon_disconnected(event)
class ThinClient:
def __init__(self, config:Config) -> None:
self.pki_doc : Dict[Any,Any] | None = None
self._pki_doc_cache : Dict[int, Dict[Any,Any]] = {} self.config = config
self.geometry : "Geometry | None" = None
self.pigeonhole_geometry : "PigeonholeGeometry | None" = None
self.reply_received_event = asyncio.Event()
self._is_connected : bool = False self._stopping : bool = False self._received_shutdown : bool = False self._daemon_instance_token : "bytes|None" = None self._in_flight_resends : Dict[bytes, Dict[str, Any]] = {} self.instance_token : bytes = os.urandom(16)
self._send_lock = asyncio.Lock()
self._recv_lock = asyncio.Lock()
self.response_queues : Dict[bytes, asyncio.Queue[Dict[str,Any]]] = {} self.ack_queues : Dict[bytes, asyncio.Queue[Dict[str,Any]]] = {}
self.logger = logging.getLogger('thinclient')
self.logger.setLevel(logging.DEBUG)
if self.config.dial is None:
raise RuntimeError("config.dial is None")
dialer = self.config.dial.resolve()
self.socket, self.server_addr = dialer.setup_socket()
async def start(self, loop:asyncio.AbstractEventLoop) -> None:
self.logger.debug("connecting to daemon")
await loop.sock_connect(self.socket, self.server_addr)
response = await self.recv(loop)
assert response is not None
assert response["connection_status_event"] is not None
await self.handle_response(response)
response = await self.recv(loop)
assert response is not None
assert response["new_pki_document_event"] is not None, response
await self.handle_response(response)
session_token_req = cbor2.dumps({
"session_token": {
"client_instance_token": self.instance_token,
}
})
length_prefix = struct.pack('>I', len(session_token_req))
await self._send_all(length_prefix + session_token_req)
session_reply = await self.recv(loop)
assert session_reply.get("session_token_reply") is not None, f"expected session_token_reply, got {session_reply}"
self.logger.debug(f"Session token reply: resumed={session_reply['session_token_reply'].get('resumed')}")
self.logger.debug("starting read loop")
self.task = loop.create_task(self.worker_loop(loop))
def handle_loop_err(task):
if self._stopping:
return
try:
result = task.result()
except asyncio.CancelledError:
pass
except (BrokenPipeError, ConnectionResetError, OSError) as e:
if not self._stopping:
self.logger.error(f"Unexpected connection error in worker loop: {e}")
except Exception:
import traceback
traceback.print_exc()
raise
self.task.add_done_callback(handle_loop_err)
def get_config(self) -> Config:
return self.config
def is_connected(self) -> bool:
return self._is_connected
def _create_socket(self) -> socket.socket:
dialer = self.config.dial.resolve()
sock, server_addr = dialer.setup_socket()
self.server_addr = server_addr
return sock
def stop(self) -> None:
self.logger.debug("closing connection to daemon")
self._stopping = True try:
close_msg = cbor2.dumps({"thin_close": {}})
length_prefix = struct.pack('>I', len(close_msg))
self.socket.sendall(length_prefix + close_msg)
except Exception:
pass self.socket.close()
self.task.cancel()
def disconnect(self) -> None:
self.logger.debug("disconnecting from daemon (preserving state)")
self._stopping = True
self.socket.close()
self.task.cancel()
async def _send_all(self, data: bytes) -> None:
async with self._send_lock:
loop = asyncio.get_running_loop()
await loop.sock_sendall(self.socket, data)
async def __recv_exactly(self, total:int, loop:asyncio.AbstractEventLoop) -> bytes:
"receive exactly (total) bytes or die trying raising BrokenPipeError"
buf = bytearray(total)
remain = memoryview(buf)
while len(remain):
if not (nread := await loop.sock_recv_into(self.socket, remain)):
raise BrokenPipeError
remain = remain[nread:]
return buf
async def recv(self, loop:asyncio.AbstractEventLoop) -> "Dict[Any,Any]":
async with self._recv_lock:
length_prefix = await self.__recv_exactly(4, loop)
message_length = struct.unpack('>I', length_prefix)[0]
if message_length > MAX_MESSAGE_SIZE:
raise ValueError(
f"daemon response frame too large: {message_length} bytes "
f"(max {MAX_MESSAGE_SIZE})")
raw_data = await self.__recv_exactly(message_length, loop)
try:
response = cbor2.loads(raw_data)
except cbor2.CBORDecodeValueError as e:
self.logger.error(f"{e}")
raise ValueError(f"{e}")
if not (set(response.keys()) & {'new_pki_document_event'}):
self.logger.debug(f"Received daemon response: [{len(raw_data)}] {type(response)} {response}")
return response
async def _reconnect(self, loop: asyncio.AbstractEventLoop) -> None:
backoff = 1.0
max_backoff = 60.0
while not self._stopping:
await asyncio.sleep(backoff)
if self._stopping:
return
try:
self.logger.debug(f"Attempting to reconnect to daemon via {self.config.dial}")
self.socket = self._create_socket()
await loop.sock_connect(self.socket, self.server_addr)
response1 = await self.recv(loop)
if response1.get("connection_status_event") is None:
self.logger.error("Reconnect handshake failed: expected connection_status_event")
self.socket.close()
continue
self.parse_status(response1["connection_status_event"])
await self.config.handle_connection_status_event(response1["connection_status_event"])
response2 = await self.recv(loop)
if response2.get("new_pki_document_event") is not None:
if response2["new_pki_document_event"].get("payload"):
self.parse_pki_doc(response2["new_pki_document_event"])
await self.config.handle_new_pki_document_event(response2["new_pki_document_event"])
session_token_req = cbor2.dumps({
"session_token": {
"client_instance_token": self.instance_token,
}
})
length_prefix = struct.pack('>I', len(session_token_req))
await self._send_all(length_prefix + session_token_req)
response3 = await self.recv(loop)
if response3.get("session_token_reply") is None:
self.logger.error("Reconnect handshake failed: expected session_token_reply")
self.socket.close()
continue
resumed = response3["session_token_reply"].get("resumed", False)
self.logger.info(f"Reconnected to daemon (connected={self._is_connected}, resumed={resumed})")
return
except (BrokenPipeError, ConnectionResetError, OSError, asyncio.CancelledError) as e:
if self._stopping:
return
self.logger.debug(f"Reconnect failed: {e} (backoff {backoff}s)")
backoff = min(backoff * 2, max_backoff)
try:
self.socket.close()
except Exception:
pass
async def _replay_in_flight_resends(self) -> None:
for key, request in list(self._in_flight_resends.items()):
try:
cbor_request = cbor2.dumps(request)
length_prefix = struct.pack('>I', len(cbor_request))
await self._send_all(length_prefix + cbor_request)
self.logger.debug(f"Replayed in-flight request: {key.hex()[:16]}...")
except Exception as e:
self.logger.error(f"Failed to replay in-flight request: {e}")
async def _read_until_disconnect(self, loop: asyncio.AbstractEventLoop) -> "Exception|None":
while not self._stopping:
try:
response = await self.recv(loop)
except asyncio.CancelledError:
return None
except (BrokenPipeError, ConnectionResetError, OSError) as e:
if self._stopping:
return None
return e
except Exception as e:
if self._stopping:
return None
self.logger.error(f"Error reading from socket: {e}")
return e
else:
def handle_response_err(task):
try:
task.result()
except Exception:
import traceback
traceback.print_exc()
resp = asyncio.create_task(self.handle_response(response))
resp.add_done_callback(handle_response_err)
return None
async def worker_loop(self, loop: asyncio.events.AbstractEventLoop) -> None:
while not self._stopping:
disconnect_err = await self._read_until_disconnect(loop)
if disconnect_err is None:
return
self.logger.info(f"Daemon disconnected (graceful={self._received_shutdown}, err={disconnect_err})")
try:
self.socket.close()
except Exception:
pass
self._is_connected = False
previous_token = self._daemon_instance_token
await self.config.handle_daemon_disconnected_event({
"is_graceful": self._received_shutdown,
"error": str(disconnect_err) if disconnect_err else None,
})
self._received_shutdown = False
await self._reconnect(loop)
if self._stopping:
return
if self._daemon_instance_token != previous_token:
self.logger.info("New daemon instance detected, replaying in-flight requests")
await self._replay_in_flight_resends()
else:
self.logger.info("Same daemon instance, skipping replay")
def parse_status(self, event: "Dict[str,Any]") -> None:
self.logger.debug("parse status")
assert event is not None
self._is_connected = event.get("is_connected", False)
token = event.get("instance_token")
if token is not None:
self._daemon_instance_token = bytes(token) if not isinstance(token, bytes) else token
sphinx_geo = event.get("sphinx_geometry")
if sphinx_geo is not None:
self.geometry = _sphinx_geometry_from_event(sphinx_geo)
else:
self.logger.error("Daemon did not supply sphinx_geometry in its ConnectionStatusEvent (incompatible daemon)")
pigeonhole_geo = event.get("pigeonhole_geometry")
if pigeonhole_geo is not None:
self.pigeonhole_geometry = _pigeonhole_geometry_from_event(pigeonhole_geo)
else:
self.logger.error("Daemon did not supply pigeonhole_geometry in its ConnectionStatusEvent (incompatible daemon)")
if self._is_connected:
self.logger.debug("Daemon is connected to mixnet - full functionality available")
else:
self.logger.info("Daemon is not connected to mixnet - entering offline mode (channel operations will work)")
self.logger.debug("parse status success")
def pki_document(self) -> "Dict[str,Any] | None":
return self.pki_doc
def pki_document_for_epoch(self, epoch:int) -> "Dict[str,Any]":
doc = self._pki_doc_cache.get(epoch)
if doc is not None:
return doc
if self.pki_doc is not None:
return self.pki_doc
raise Exception("no PKI document available for the requested epoch")
async def get_pki_document_raw(self, epoch:int = 0) -> "Tuple[bytes,int]":
query_id = self.new_query_id()
request = {
"get_pki_document": {
"query_id": query_id,
"epoch": epoch,
}
}
reply = await self._send_and_wait(query_id=query_id, request=request)
returned_epoch = reply.get("epoch", 0)
error_code = reply.get("error_code", 0)
if error_code != THIN_CLIENT_SUCCESS:
error_msg = thin_client_error_to_string(error_code)
raise Exception(
f"get_pki_document_raw failed for epoch {epoch}: {error_msg}"
)
return reply.get("payload"), returned_epoch
def parse_pki_doc(self, event: "Dict[str,Any]") -> None:
self.logger.debug("parse pki doc")
assert event is not None
assert event["payload"] is not None
raw_pki_doc = cbor2.loads(event["payload"])
self.pki_doc = raw_pki_doc
epoch = raw_pki_doc.get("Epoch")
if epoch is not None:
self._pki_doc_cache[epoch] = raw_pki_doc
self.logger.debug("Cached PKI document for epoch %d", epoch)
max_cached_epochs = 5
if len(self._pki_doc_cache) > max_cached_epochs:
oldest_epoch = epoch - max_cached_epochs
stale = [e for e in self._pki_doc_cache if e < oldest_epoch]
for e in stale:
del self._pki_doc_cache[e]
self.logger.debug("parse pki doc success")
def get_services(self, capability:str) -> "List[ServiceDescriptor]":
doc = self.pki_document()
if doc == None:
raise Exception("pki doc is nil")
descriptors = find_services(capability, doc)
if not descriptors:
raise Exception("service not found in pki doc")
return descriptors
def get_service(self, service_name:str) -> ServiceDescriptor:
service_descriptors = self.get_services(service_name)
return random.choice(service_descriptors)
def get_all_couriers(self) -> "List[Tuple[bytes, bytes]]":
services = self.get_services("courier")
couriers = []
for svc in services:
identity_hash = blake2_256_sum(svc.mix_descriptor['IdentityKey'])
couriers.append((identity_hash, svc.recipient_queue_id))
return couriers
def get_distinct_couriers(self, n:int) -> "List[Tuple[bytes, bytes]]":
couriers = self.get_all_couriers()
if len(couriers) < n:
raise Exception("not enough couriers available")
return random.sample(couriers, n)
async def blocking_send_message(self, payload:bytes|str, dest_node:bytes, dest_queue:bytes, timeout_seconds:float=30.0) -> bytes:
if not self._is_connected:
raise ThinClientOfflineError("cannot send message in offline mode - daemon not connected to mixnet")
surb_id = self.new_surb_id()
reply_future = asyncio.get_event_loop().create_future()
original_handler = self.config.on_message_reply
async def capture_reply(event):
if event.get("surbid") == surb_id and not reply_future.done():
reply_future.set_result(event.get("payload"))
if original_handler:
await original_handler(event)
self.config.on_message_reply = capture_reply
try:
await self.send_message(surb_id, payload, dest_node, dest_queue)
return await asyncio.wait_for(reply_future, timeout=timeout_seconds)
finally:
self.config.on_message_reply = original_handler
@staticmethod
def new_message_id() -> bytes:
return os.urandom(MESSAGE_ID_SIZE)
def new_surb_id(self) -> bytes:
return os.urandom(SURB_ID_SIZE)
def new_query_id(self) -> bytes:
return os.urandom(16)
async def _send_and_wait(self, *, query_id:bytes, request: Dict[str, Any]) -> Dict[str, Any]:
cbor_request = cbor2.dumps(request)
length_prefix = struct.pack('>I', len(cbor_request))
length_prefixed_request = length_prefix + cbor_request
assert query_id not in self.response_queues
self.response_queues[query_id] = asyncio.Queue(maxsize=1)
request_type = list(request.keys())[0]
try:
await self._send_all(length_prefixed_request)
self.logger.info(f"{request_type} request sent.")
reply = await self.response_queues[query_id].get()
self.logger.info(f"{request_type} response received.")
return reply
except asyncio.CancelledError:
self.logger.info("{request_type} task cancelled.")
raise
finally:
del self.response_queues[query_id]
async def handle_response(self, response: "Dict[str,Any]") -> None:
assert response is not None
if response.get("shutdown_event") is not None:
self.logger.info("Received ShutdownEvent from daemon")
self._received_shutdown = True
return
if response.get("connection_status_event") is not None:
self.logger.debug("connection status event")
self.parse_status(response["connection_status_event"])
await self.config.handle_connection_status_event(response["connection_status_event"])
return
if response.get("new_pki_document_event") is not None:
self.logger.debug("new pki doc event")
event = response["new_pki_document_event"]
if event.get("payload") is not None:
self.parse_pki_doc(event)
await self.config.handle_new_pki_document_event(event)
return
if response.get("message_sent_event") is not None:
self.logger.debug("message sent event")
await self.config.handle_message_sent_event(response["message_sent_event"])
return
if response.get("message_reply_event") is not None:
self.logger.debug("message reply event")
reply = response["message_reply_event"]
self.reply_received_event.set()
await self.config.handle_message_reply_event(reply)
return
for reply_type, reply in response.items():
if not reply:
continue
self.logger.debug(f"channel {reply_type} event")
if not reply_type.endswith("_reply") or not (query_id := reply.get("query_id", None)):
self.logger.debug(f"{reply_type} is not a reply, or can't get query_id")
continue
if not (queue := self.response_queues.get(query_id, None)):
self.logger.debug(f"query_id for {reply_type} has no listener")
continue
asyncio.create_task(queue.put(reply))
async def send_message_without_reply(self, payload:bytes|str, dest_node:bytes, dest_queue:bytes) -> None:
if not self._is_connected:
raise ThinClientOfflineError("cannot send_message_without_reply in offline mode - daemon not connected to mixnet")
if not isinstance(payload, bytes):
payload = payload.encode('utf-8')
send_message = {
"id": None, "with_surb": False,
"surbid": None, "destination_id_hash": dest_node,
"recipient_queue_id": dest_queue,
"payload": payload,
}
request = {
"send_message": send_message
}
cbor_request = cbor2.dumps(request)
length_prefix = struct.pack('>I', len(cbor_request))
length_prefixed_request = length_prefix + cbor_request
try:
await self._send_all(length_prefixed_request)
self.logger.info("Message sent successfully.")
except Exception as e:
self.logger.error(f"Error sending message: {e}")
async def send_message(self, surb_id:bytes, payload:bytes|str, dest_node:bytes, dest_queue:bytes) -> None:
if not self._is_connected:
raise ThinClientOfflineError("cannot send message in offline mode - daemon not connected to mixnet")
if not isinstance(payload, bytes):
payload = payload.encode('utf-8')
send_message = {
"id": None, "with_surb": True,
"surbid": surb_id,
"destination_id_hash": dest_node,
"recipient_queue_id": dest_queue,
"payload": payload,
}
request = {
"send_message": send_message
}
cbor_request = cbor2.dumps(request)
length_prefix = struct.pack('>I', len(cbor_request))
length_prefixed_request = length_prefix + cbor_request
try:
await self._send_all(length_prefixed_request)
self.logger.info("Message sent successfully.")
except Exception as e:
self.logger.error(f"Error sending message: {e}")
def pretty_print_pki_doc(self, doc: "Dict[str,Any]") -> None:
assert doc is not None
assert doc['GatewayNodes'] is not None
assert doc['ServiceNodes'] is not None
assert doc['Topology'] is not None
new_doc = doc
gateway_nodes = []
service_nodes = []
topology = []
for gateway_cert_blob in doc['GatewayNodes']:
gateway_cert = cbor2.loads(gateway_cert_blob)
gateway_nodes.append(gateway_cert)
for service_cert_blob in doc['ServiceNodes']:
service_cert = cbor2.loads(service_cert_blob)
service_nodes.append(service_cert)
for layer in doc['Topology']:
for mix_desc_blob in layer:
mix_cert = cbor2.loads(mix_desc_blob)
topology.append(mix_cert)
new_doc['GatewayNodes'] = gateway_nodes
new_doc['ServiceNodes'] = service_nodes
new_doc['Topology'] = topology
pretty_print_obj(new_doc)