Skip to content
Merged
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
55 changes: 33 additions & 22 deletions sdk/python/src/agent_mesh/mesh.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,10 @@ def __init__(self, identity: Identity, control_plane: ControlPlaneClient, creden
self.control_plane = control_plane
self._credential = credential
self._state_dir = state_dir
# refresh() and sync_control_plane() run in worker threads of the
# session's loops; the control plane redeems only the last biscuit it
# issued, and both write the same state files.
self._lock = threading.RLock()

@property
def peer_id(self) -> str:
Expand Down Expand Up @@ -173,29 +177,31 @@ def refresh(self) -> MeshCredential:
"""Trades the current biscuit for a fresh one and persists it. The control
plane redeems only the last biscuit it issued, so a lost refresh result
means re-enrolling; persisting before returning keeps that rare."""
result = self.control_plane.refresh(self.identity, self._credential.biscuit)
control_plane_keys = self._credential.control_plane_keys
try:
control_plane_keys = self.control_plane.keys(control_plane_keys)
except Exception: # noqa: BLE001 - a failed /keys sync must not cost the new biscuit
pass
self._credential = replace(
self._credential,
biscuit=result.biscuit,
expiration=result.expiration,
control_plane_keys=control_plane_keys,
issued_under_keys=list(control_plane_keys),
)
self.save()
return self._credential
with self._lock:
result = self.control_plane.refresh(self.identity, self._credential.biscuit)
control_plane_keys = self._credential.control_plane_keys
try:
control_plane_keys = self.control_plane.keys(control_plane_keys)
except Exception: # noqa: BLE001 - a failed /keys sync must not cost the new biscuit
pass
self._credential = replace(
self._credential,
biscuit=result.biscuit,
expiration=result.expiration,
control_plane_keys=control_plane_keys,
issued_under_keys=list(control_plane_keys),
)
self.save()
return self._credential

def add_trusted_key(self, key: bytes) -> bool:
"""Adopts a signing key announced by a KEY_ROTATION event, so peers whose
credentials the new key signs verify before the next pull confirms it."""
if any(bytes(k) == bytes(key) for k in self._credential.control_plane_keys):
return False
self._credential = replace(self._credential, control_plane_keys=[*self._credential.control_plane_keys, bytes(key)])
return True
with self._lock:
if any(bytes(k) == bytes(key) for k in self._credential.control_plane_keys):
return False
self._credential = replace(self._credential, control_plane_keys=[*self._credential.control_plane_keys, bytes(key)])
return True

def sync_control_plane(self) -> ControlPlaneSync:
"""The member's pull from the control plane, as sam-node's SyncControlPlane:
Expand All @@ -205,6 +211,10 @@ def sync_control_plane(self) -> ControlPlaneSync:
addresses and the ban set. Each part is attempted even when another
fails; the errors are reported together. Gossip events only bring this
forward; they are never the only way state arrives."""
with self._lock:
return self._sync_control_plane()

def _sync_control_plane(self) -> ControlPlaneSync:
errors: list[str] = []
keys_changed = False
refreshed = False
Expand Down Expand Up @@ -255,9 +265,10 @@ def save(self) -> None:
"""Writes identity and credential to the state directory, if one is configured."""
if self._state_dir is None:
return
self._state_dir.mkdir(parents=True, exist_ok=True, mode=0o700)
_write_atomic(self._state_dir / _IDENTITY_FILE, self.identity.to_libp2p_private_key())
_write_atomic(self._state_dir / _CREDENTIAL_FILE, self._credential.to_json().encode())
with self._lock:
self._state_dir.mkdir(parents=True, exist_ok=True, mode=0o700)
_write_atomic(self._state_dir / _IDENTITY_FILE, self.identity.to_libp2p_private_key())
_write_atomic(self._state_dir / _CREDENTIAL_FILE, self._credential.to_json().encode())


def _load_identity(state: Optional[Path]) -> Optional[Identity]:
Expand Down
23 changes: 20 additions & 3 deletions sdk/python/src/agent_mesh/relay.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@
from libp2p.utils.varint import encode_varint_prefixed, read_varint_prefixed_bytes

from ._proto import circuit_pb2 as circuit
from .host import open_stream
from .host import DIAL_TIMEOUT, HANGUP_GRACE, open_stream
from .identity import canonical_peer_id

logger = logging.getLogger("agent_mesh")
Expand Down Expand Up @@ -82,7 +82,11 @@ async def reserve_relay(host: IHost, relay_peer_id: ID) -> circuit.Reservation:

async def dial_through_relay(host: IHost, relay_peer_id: ID, target: ID) -> INetConn:
"""Opens a connection to `target` through a relay we are connected to, and
registers it with the host so streams can be opened on it."""
registers it with the host so streams can be opened on it. The relay's
answer and the handshake with the target that follows are each bounded:
py-libp2p's own limits on that upgrade add up to about a minute when the
relay accepted the circuit and the far end never speaks, and a relayed
dial is held to DIAL_TIMEOUT as a direct one is."""
stream = await open_stream(host, relay_peer_id, HOP_PROTOCOL, RELAY_MESSAGE_TIMEOUT)
try:
with trio.fail_after(RELAY_MESSAGE_TIMEOUT):
Expand All @@ -99,7 +103,20 @@ async def dial_through_relay(host: IHost, relay_peer_id: ID, target: ID) -> INet
# From here the stream is the wire; the usual TLS + yamux upgrade runs on it.
circuit_addr = multiaddr.Multiaddr(f"/p2p/{relay_peer_id}/p2p-circuit/p2p/{target}")
raw = RawConnection(stream=stream, initiator=True, connection_type=ConnectionType.RELAYED, addresses=[circuit_addr])
return await host.upgrade_outbound_connection(raw, target)
try:
with trio.fail_after(DIAL_TIMEOUT):
return await host.upgrade_outbound_connection(raw, target)
except BaseException as err:
# A circuit given up on, timed out, failed or cancelled for another
# path that won, is reset so the relay drops it too.
with trio.CancelScope(shield=True), trio.move_on_after(HANGUP_GRACE):
try:
await stream.reset()
except Exception: # noqa: BLE001 - already gone
pass
if isinstance(err, trio.TooSlowError):
raise ConnectionError(f"no secure connection to {target} through relay {relay_peer_id} within {DIAL_TIMEOUT:g}s") from None
raise


def stop_stream_handler(host: IHost) -> Callable[[INetStream], object]:
Expand Down
62 changes: 39 additions & 23 deletions sdk/python/src/agent_mesh/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -204,7 +204,10 @@ async def connect(self, peer: Peer) -> ID:
router the control plane lists that this member has not joined
through; a relay opens a circuit only for a source it authenticated,
so each such router is admitted first. Which router each side joined
through does not decide whether they can talk."""
through does not decide whether they can talk. The routers of each
step are dialed at once: a relay whose destination is gone answers
only after its own timeout, and a peer a rollout replaced must cost
the caller one such wait, not one per router."""
if isinstance(peer, multiaddr.Multiaddr) or (isinstance(peer, str) and peer.startswith("/")):
return await self._connect_addr(multiaddr.Multiaddr(str(peer)))
if isinstance(peer, str):
Expand All @@ -230,21 +233,33 @@ async def connect(self, peer: Peer) -> ID:
return target
except Exception as err: # noqa: BLE001 - the relayed path is tried next
failures.append(f"direct {[str(a) for a in direct]}: {err}")
for r in self.routers:
try:
await self._connect_addr(multiaddr.Multiaddr(f"{r.addr}/p2p-circuit{suffix}"))
return target
except Exception as err: # noqa: BLE001 - the next router is tried
failures.append(f"via router {r.peer_id}: {err}")
if await self._connect_through(target, [multiaddr.Multiaddr(f"{r.addr}/p2p-circuit{suffix}") for r in self.routers], failures):
return target
for more in (self._routed_addresses, self._unjoined_router_addresses):
for ma in await more(target):
try:
await self._connect_addr(ma)
return target
except Exception as err: # noqa: BLE001 - the next address is tried
failures.append(f"{ma}: {err}")
if await self._connect_through(target, await more(target), failures):
return target
raise ConnectionError(f"cannot reach {target}:\n " + "\n ".join(failures))

async def _connect_through(self, target: ID, addrs: list[multiaddr.Multiaddr], failures: list[str]) -> bool:
"""Dials the addresses at once; the first that reaches target ends the
others, and each that failed adds its reason to failures."""
reached = False

async def attempt(ma: multiaddr.Multiaddr, nursery: trio.Nursery) -> None:
nonlocal reached
try:
await self._connect_addr(ma)
except Exception as err: # noqa: BLE001 - the other attempts go on
failures.append(f"{ma}: {err}")
return
reached = True
nursery.cancel_scope.cancel()

async with trio.open_nursery() as nursery:
for ma in addrs:
nursery.start_soon(attempt, ma, nursery)
return reached or target in self.host.get_connected_peers()

async def _unjoined_router_addresses(self, target: ID) -> list[multiaddr.Multiaddr]:
"""The relayed paths to a peer through the routers the control plane
lists that this member has not joined through, each admitted first.
Expand All @@ -269,9 +284,11 @@ async def _unjoined_router_addresses(self, target: ID) -> list[multiaddr.Multiad

async def _routed_addresses(self, target: ID) -> list[multiaddr.Multiaddr]:
"""The addresses the routers' DHT knows for a peer, relayed ones
through routers this member has admitted by then: a relay opens a
circuit only for a source it authenticated, so a router met this way
is dialed and passed the handshake first, and joins the admitted set."""
through routers this member has not dialed for it yet: a relay that
is an admitted router was tried by connect() already, and a relay
opens a circuit only for a source it authenticated, so a router met
this way is dialed and passed the handshake first, and joins the
admitted set."""
out: list[multiaddr.Multiaddr] = []
seeds = [ID.from_base58(r.peer_id) for r in self.routers]
for ma in await find_peer(self.host, target, seeds):
Expand All @@ -284,14 +301,13 @@ async def _routed_addresses(self, target: ID) -> list[multiaddr.Multiaddr]:
relay = info_from_p2p_addr(relay_addr).peer_id
except Exception: # noqa: BLE001 - a circuit address naming no relay is useless
continue
if str(relay) in self.banned:
if str(relay) in self.banned or any(r.peer_id == str(relay) for r in self.routers):
continue
try:
await self._admit_router(relay_addr)
except Exception as err: # noqa: BLE001 - a relay that is not a router of this mesh is not used
logger.debug("router %s named by the DHT did not admit us: %s", relay, err)
continue
if not any(r.peer_id == str(relay) for r in self.routers):
try:
await self._admit_router(relay_addr)
except Exception as err: # noqa: BLE001 - a relay that is not a router of this mesh is not used
logger.debug("router %s named by the DHT did not admit us: %s", relay, err)
continue
out.append(multiaddr.Multiaddr(f"{relay_addr}/p2p-circuit/p2p/{target}"))
return out

Expand Down
60 changes: 52 additions & 8 deletions sdk/python/tests/test_mesh.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,10 @@
# See the License for the specific language governing permissions and
# limitations under the License.

import base64
import json
import stat
import threading
import time
import urllib.parse
from datetime import datetime, timezone
Expand Down Expand Up @@ -42,11 +44,22 @@ def _ts_s(seconds: int) -> Timestamp:


class FakeControlPlane:
"""Approves everything and hands out numbered biscuits."""
"""Approves everything and hands out numbered biscuits. With strict_refresh
it redeems only the last biscuit it issued, as the real one does, and
refresh_delay is how long a /refresh takes."""

def __init__(self, keys_ok=True):
def __init__(self, keys_ok=True, strict_refresh=False, refresh_delay=0.0):
self.issued = 0
self.keys_ok = keys_ok
self.strict_refresh = strict_refresh
self.refresh_delay = refresh_delay
self.last_biscuit = b""
self._lock = threading.Lock()

def _issue(self, biscuit):
self.issued += 1
self.last_biscuit = biscuit
return biscuit

def _signed_keys(self):
ts = int(time.time() * 1000)
Expand All @@ -58,25 +71,28 @@ def transport(self, method, url, headers, body):
if (method, path) == ("POST", "/register"):
# The biscuit names the JWT that was presented, so a test can see which.
jwt = pb.EnrollRequest.FromString(body).jwt
self.issued += 1
return 200, pb.EnrollResponse(
biscuit_token=f"biscuit-for-{jwt}".encode(),
biscuit_token=self._issue(f"biscuit-for-{jwt}".encode()),
control_plane_public_key=CP_KEY.public_key_raw,
router_addresses=["/dns4/router.example/tcp/4001/p2p/12D3KooWP8iKhDf3iCMo2H3butNVfdTUtYwYWYQ75jTGnynXPFMp"],
expire_time=_ts_s(int(time.time()) + 3600),
).SerializeToString()
if (method, path) == ("POST", "/enroll"):
self.issued += 1
return 200, pb.BootstrapEnrollResponse(
status=pb.ENROLLMENT_STATUS_APPROVED,
biscuit_token=f"biscuit-{self.issued}".encode(),
biscuit_token=self._issue(f"biscuit-{self.issued + 1}".encode()),
control_plane_public_key=CP_KEY.public_key_raw,
router_addresses=["/dns4/router.example/tcp/4001/p2p/12D3KooWP8iKhDf3iCMo2H3butNVfdTUtYwYWYQ75jTGnynXPFMp"],
expire_time=_ts_s(int(time.time()) + 3600),
).SerializeToString()
if (method, path) == ("POST", "/refresh"):
self.issued += 1
return 200, pb.TokenRefreshResponse(biscuit_token=f"biscuit-{self.issued}".encode(), expire_time=_ts_s(int(time.time()) + 7200)).SerializeToString()
time.sleep(self.refresh_delay)
with self._lock:
presented = base64.b64decode(headers["Authorization"].removeprefix("Bearer "))
if self.strict_refresh and presented != self.last_biscuit:
return 200, pb.TokenRefreshResponse(error_message="biscuit already redeemed").SerializeToString()
biscuit = self._issue(f"biscuit-{self.issued + 1}".encode())
return 200, pb.TokenRefreshResponse(biscuit_token=biscuit, expire_time=_ts_s(int(time.time()) + 7200)).SerializeToString()
if (method, path) == ("GET", "/keys"):
return (200, self._signed_keys().SerializeToString()) if self.keys_ok else (500, b"boom")
return 404, f"no route for {method} {path}".encode()
Expand Down Expand Up @@ -153,6 +169,34 @@ def test_enroll_without_state_dir_keeps_enrollment_key_when_keys_fails():
assert mesh.credential.biscuit == b"biscuit-2"


def test_concurrent_refreshes_run_one_after_the_other(tmp_path):
"""A session refreshes the credential from one worker thread and pulls
from the control plane, which may refresh too, from another. The control
plane redeems only the last biscuit it issued and both write the same
state files, so two refreshes at once leave the loser with a spent
biscuit, or renaming a temp file the winner already renamed."""
cp = FakeControlPlane(strict_refresh=True, refresh_delay=0.05)
state = tmp_path / "state"
mesh = AgentMesh.enroll("http://127.0.0.1:1", bootstrap_token="sbt_secret", state_dir=state, transport=cp.transport)
errors = []

def refresh():
try:
mesh.refresh()
except Exception as err: # noqa: BLE001 - collected for the assertion
errors.append(err)

threads = [threading.Thread(target=refresh) for _ in range(8)]
for t in threads:
t.start()
for t in threads:
t.join()
assert errors == []
assert mesh.credential.biscuit == b"biscuit-9"
assert json.loads((state / "credential.json").read_text())["biscuit"] == base64.b64encode(b"biscuit-9").decode()
assert AgentMesh.load(state, transport=cp.transport).credential.biscuit == b"biscuit-9"


def test_enroll_refuses_ambiguous_credentials():
cp = FakeControlPlane()
with pytest.raises(ValueError, match="exactly one of"):
Expand Down
Loading
Loading