Compare commits

..
1 Commits
Author SHA1 Message Date
Evan 4430b0daf9 api streams 2026-05-07 17:07:40 +01:00
29 changed files with 312 additions and 510 deletions

No files matched your search

Generated
+2 -21
View File
@@ -319,26 +319,6 @@ version = "0.6.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "175812e0be2bccb6abe50bb8d566126198344f707e304f45c648fd8f2cc0365e"
[[package]]
name = "bytemuck"
version = "1.25.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c8efb64bd706a16a1bdde310ae86b351e4d21550d98d056f22f8a7f7a2183fec"
dependencies = [
"bytemuck_derive",
]
[[package]]
name = "bytemuck_derive"
version = "1.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f9abbd1bc6865053c427f7198e6af43bfdedc55ab791faed4fbd361d789575ff"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]]
name = "byteorder"
version = "1.5.0"
@@ -818,6 +798,7 @@ dependencies = [
"rand 0.10.1",
"serde_json",
"tokio",
"tracing",
"zenoh",
]
@@ -1840,7 +1821,6 @@ name = "networking"
version = "0.0.1"
dependencies = [
"async-stream",
"bytemuck",
"futures-lite",
"log",
"netwatcher",
@@ -3691,6 +3671,7 @@ dependencies = [
"signal-hook-registry",
"socket2 0.6.1",
"tokio-macros",
"tracing",
"windows-sys 0.61.2",
]
+12 -14
View File
@@ -21,21 +21,10 @@ opt-level = 3
## Crate members as common dependencies
networking = { path = "rust/networking" }
# pyo3
pyo3 = "0.27.2"
pyo3-async-runtimes = "0.27.0"
pyo3-log = "0.13.2"
pyo3-stub-gen = "0.22.2"
# util
# Macro dependecies
extend = "1.2"
tokio = "1.46"
futures-lite = "2.6.1"
async-stream = "0.3.6"
pin-project = "1.1.10"
serde_json = "1.0.149"
rand = "0.10.1"
parking_lot = "0.12.5"
# Tracing/logging
log = "0.4"
@@ -43,10 +32,19 @@ env_logger = "0.11.10"
# networking
zenoh = "=1.9.0"
async-stream = "0.3.6"
netwatcher = "0.6.0"
parking_lot = "0.12.5"
pin-project = "1.1.10"
pyo3 = "0.27.2"
pyo3-async-runtimes = "0.27.0"
pyo3-log = "0.13.2"
pyo3-stub-gen = "0.22.2"
rand = "0.10.1"
serde_json = "1.0.149"
tracing = "0.1.44"
zenoh-plugin-storage-manager = { version = "=1.9.0", default-features = false }
zenoh-plugin-trait = "=1.9.0"
netwatcher = "0.6.0"
bytemuck = "1.25.0"
[workspace.lints.rust]
static_mut_refs = "warn" # Or use "warn" instead of deny
+2 -1
View File
@@ -36,7 +36,7 @@ pyo3-async-runtimes = { workspace = true, features = [
pyo3-log.workspace = true
# async runtime
tokio = { workspace = true, features = ["full"] }
tokio = { workspace = true, features = ["full", "tracing"] }
futures-lite.workspace = true
pin-project.workspace = true
@@ -49,3 +49,4 @@ zenoh.workspace = true
rand.workspace = true
serde_json.workspace = true
parking_lot.workspace = true
tracing.workspace = true
+14 -9
View File
@@ -16,7 +16,7 @@ __all__ = [
@typing.final
class NetReceiver:
def recv(self) -> collections.abc.Awaitable[bytes | None]: ...
def recv(self) -> collections.abc.Awaitable[bytes]: ...
@typing.final
class NetSender:
@@ -25,23 +25,27 @@ class NetSender:
@typing.final
class NetworkingHandle:
@staticmethod
def new(identity: bytes, bootstrap_peers: typing.Sequence[builtins.str], listen_port: builtins.int) -> tuple[NetworkingHandle, PySession]: ...
def new(
identity: bytes,
bootstrap_peers: typing.Sequence[builtins.str],
listen_port: builtins.int,
) -> tuple[NetworkingHandle, PySession]: ...
async def gossipsub_subscribe(self, topic: builtins.str) -> builtins.bool:
r"""
Subscribe to a `GossipSub` topic.
Returns `True` if the subscription worked. Returns `False` if we were already subscribed.
"""
async def gossipsub_unsubscribe(self, topic: builtins.str) -> builtins.bool:
r"""
Unsubscribes from a `GossipSub` topic.
Returns `True` if we were subscribed to this topic. Returns `False` if we were not subscribed.
"""
async def gossipsub_publish(self, topic: builtins.str, data: bytes) -> None:
r"""
Publishes a message with multiple topics to the `GossipSub` network.
If no peers are found that subscribe to this topic, throws `NoPeersSubscribedToTopicError` exception.
"""
async def recv(self) -> PyFromSwarm: ...
@@ -53,16 +57,18 @@ class PyFromSwarm:
@property
def connected(self) -> builtins.bool: ...
def __new__(cls, connected: builtins.bool) -> PyFromSwarm.Connection: ...
@typing.final
class Message(PyFromSwarm):
__match_args__ = ("topic", "data",)
__match_args__ = (
"topic",
"data",
)
@property
def topic(self) -> builtins.str: ...
@property
def data(self) -> bytes: ...
def __new__(cls, topic: builtins.str, data: bytes) -> PyFromSwarm.Message: ...
@typing.final
class PySession:
@@ -73,4 +79,3 @@ class PySession:
@typing.final
class StateProxy:
def snapshot(self) -> collections.abc.Awaitable[str]: ...
+3 -3
View File
@@ -9,12 +9,12 @@ mod allow_threading;
mod networking;
mod point_to_point;
mod session;
mod state;
// mod state;
use crate::networking::networking_submodule;
use crate::point_to_point::{NetReceiver, NetSender};
use crate::session::PySession;
use crate::state::StateProxy;
//use crate::state::StateProxy;
use pyo3::prelude::PyModule;
use pyo3::types::PyModuleMethods;
use pyo3::{Bound, PyResult, pymodule};
@@ -164,7 +164,7 @@ fn main_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
// too many importing issues...
// m.add_class::<PyKeypair>()?;
// networking_submodule(m)?;
m.add_class::<StateProxy>()?;
// m.add_class::<StateProxy>()?;
m.add_class::<PySession>()?;
m.add_class::<NetReceiver>()?;
m.add_class::<NetSender>()?;
+10 -6
View File
@@ -1,9 +1,8 @@
use std::sync::Arc;
use pyo3::exceptions::PyConnectionError;
use pyo3::types::PyBytes;
use pyo3::types::PyNone;
use pyo3::{BoundObject, prelude::*};
use pyo3::prelude::*;
use pyo3::{exceptions::PyStopAsyncIteration, types::PyBytes};
use pyo3_stub_gen::derive::{gen_stub_pyclass, gen_stub_pymethods};
use zenoh::Result;
use zenoh::{
@@ -23,7 +22,7 @@ pub struct NetReceiver {
#[pymethods]
impl NetReceiver {
#[gen_stub(override_return_type(
type_repr="collections.abc.Awaitable[bytes | None]",
type_repr="collections.abc.Awaitable[bytes]",
imports=("collections.abc")
))]
pub fn recv<'py>(&'py self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
@@ -37,9 +36,9 @@ impl NetReceiver {
match subscriber.recv_async().await {
Err(_) => {
// stream closed;
Ok(Python::attach(|py| PyNone::get(py).unbind()).into_any())
Err(PyStopAsyncIteration::new_err(()))
}
Ok(sample) => Ok(sample.payload().to_bytes().to_vec().pybytes().into_any()),
Ok(sample) => Ok(sample.payload().to_bytes().to_vec().pybytes()),
}
}
})
@@ -72,18 +71,23 @@ impl NetSender {
let bytes = data.as_bytes().to_vec();
async move {
if is_first {
log::warn!("sender waiting for listener");
wait_for_listener(&*publisher)
.await
.map_err(|e| PyConnectionError::new_err(e.to_string()))?;
log::warn!("listener found");
}
log::warn!("checking for matcher");
if !publisher
.matching_status()
.await
.map_err(|e| PyConnectionError::new_err(e.to_string()))?
.matching()
{
log::warn!("no matcher found");
return Ok(false);
}
log::warn!("publishing");
publisher
.put(&bytes)
.await
+6 -8
View File
@@ -5,11 +5,10 @@ use pyo3_stub_gen::derive::{gen_stub_pyclass, gen_stub_pymethods};
use zenoh::Session;
use zenoh::Wait;
use zenoh::qos::CongestionControl;
use crate::{
point_to_point::{NetReceiver, NetSender},
state::StateProxy,
//state::StateProxy,
};
#[gen_stub_pyclass]
@@ -55,7 +54,6 @@ impl PySession {
publisher: Arc::new(
self.session
.declare_publisher(key)
.congestion_control(CongestionControl::Block)
.wait()
// C5: key format error, could be declaration error
.map_err(|e| PyValueError::new_err(e.to_string()))?,
@@ -64,9 +62,9 @@ impl PySession {
})
}
pub fn state_proxy(&self) -> StateProxy {
StateProxy {
session: self.session.clone(),
}
}
//pub fn state_proxy(&self) -> StateProxy {
//StateProxy {
//session: self.session.clone(),
//}
//}
}
+62
View File
@@ -0,0 +1,62 @@
[package]
name = "exo_pyo3_bindings"
version = { workspace = true }
edition = { workspace = true }
publish = false
[lib]
doctest = false
path = "src/lib.rs"
name = "exo_pyo3_bindings"
# "cdylib" needed to produce shared library for Python to import
# "rlib" needed for stub-gen to run
crate-type = ["cdylib", "rlib"]
[[bin]]
path = "src/bin/stub_gen.rs"
name = "stub_gen"
doc = false
[lints]
workspace = true
[dependencies]
networking.workspace = true
extend.workspace = true
# interop
pyo3 = { version = "0.27.2", features = [
# "abi3-py313", # tells pyo3 (and maturin) to build using the stable ABI with minimum Python version 3.13
# "nightly", # enables better-supported GIL integration
"experimental-async", # async support in #[pyfunction] & #[pymethods]
#"experimental-inspect", # inspection of generated binary => easier to automate type-hint generation
#"py-clone", # adding Clone-ing of `Py<T>` without GIL (may cause panics - remove if panics happen)
# "multiple-pymethods", # allows multiple #[pymethods] sections per class
# integrations with other libraries
# "arc_lock", "bigdecimal", "either", "hashbrown", "indexmap", "num-bigint", "num-complex", "num-rational",
# "ordered-float", "rust_decimal", "smallvec",
# "anyhow", "chrono", "chrono-local", "chrono-tz", "eyre", "jiff-02", "lock_api", "parking-lot", "time", "serde",
] }
pyo3-stub-gen = { version = "0.17.2" }
pyo3-async-runtimes = { version = "0.27.0", features = [
"attributes",
"tokio-runtime",
"testing",
] }
pyo3-log = "0.13.2"
# async runtime
tokio = { workspace = true, features = ["full", "tracing"] }
futures-lite = { workspace = true }
pin-project = "1.1.10"
# Tracing
log.workspace = true
env_logger.workspace = true
# Networking
zenoh.workspace = true
zerompk.workspace = true
rand = "0.10.1"
@@ -0,0 +1,72 @@
# This file is automatically generated by pyo3_stub_gen
# ruff: noqa: E501, F401
import builtins
import typing
@typing.final
class Keypair:
r"""
Identity keypair of a node.
"""
@staticmethod
def generate() -> Keypair:
r"""
Generate a new Ed25519 keypair.
"""
@staticmethod
def from_bytes(bytes: bytes) -> Keypair:
r"""
Construct an Ed25519 keypair from secret key bytes
"""
def to_bytes(self) -> bytes:
r"""
Get the secret key bytes underlying the keypair
"""
def to_node_id(self) -> builtins.str:
r"""
Convert the `Keypair` into the corresponding `PeerId` string, which we use as our `NodeId`.
"""
@typing.final
class NetworkingHandle:
def __new__(cls, identity: Keypair, bootstrap_peers: typing.Sequence[builtins.str], listen_port: builtins.int) -> NetworkingHandle: ...
async def gossipsub_subscribe(self, topic: builtins.str) -> builtins.bool:
r"""
Subscribe to a `GossipSub` topic.
Returns `True` if the subscription worked. Returns `False` if we were already subscribed.
"""
async def gossipsub_unsubscribe(self, topic: builtins.str) -> builtins.bool:
r"""
Unsubscribes from a `GossipSub` topic.
Returns `True` if we were subscribed to this topic. Returns `False` if we were not subscribed.
"""
async def gossipsub_publish(self, topic: builtins.str, data: bytes) -> None:
r"""
Publishes a message with multiple topics to the `GossipSub` network.
If no peers are found that subscribe to this topic, throws `NoPeersSubscribedToTopicError` exception.
"""
async def recv(self) -> PyFromSwarm: ...
class PyFromSwarm:
@typing.final
class Connection(PyFromSwarm):
__match_args__ = ("connected",)
@property
def connected(self) -> builtins.bool: ...
def __new__(cls, connected: builtins.bool) -> PyFromSwarm.Connection: ...
@typing.final
class Message(PyFromSwarm):
__match_args__ = ("topic", "data",)
@property
def topic(self) -> builtins.str: ...
@property
def data(self) -> bytes: ...
def __new__(cls, topic: builtins.str, data: bytes) -> PyFromSwarm.Message: ...
...
-1
View File
@@ -14,7 +14,6 @@ zenoh-plugin-storage-manager.workspace = true
zenoh-plugin-trait.workspace = true
rand.workspace = true
log.workspace = true
bytemuck = { workspace = true, features = ["derive"] }
[lints]
workspace = true
+1 -1
View File
@@ -7,7 +7,7 @@ async fn main() -> Result<()> {
zenoh::init_log_from_env_or("info");
info!("Opening session...");
let cfg = networking::cfg(rand::random(), 0)?;
let session = networking::open(cfg, 52414).await?;
let session = networking::open(cfg).await?;
let _tok = session
.z
.liveliness()
+2 -2
View File
@@ -6,8 +6,8 @@ use zenoh::Result;
async fn main() -> Result<()> {
zenoh::init_log_from_env_or("info");
info!("Opening session...");
let cfg = networking::cfg(rand::random(), 52414)?;
let session = networking::open(cfg, 52414).await?;
let cfg = networking::cfg(rand::random(), 0)?;
let session = networking::open(cfg).await?;
let _tok = session
.z
.liveliness()
-311
View File
@@ -1,311 +0,0 @@
use std::{
io::{self, ErrorKind},
net::{Ipv6Addr, SocketAddr, SocketAddrV6},
sync::Arc,
time::Duration,
};
use bytemuck::{Pod, Zeroable};
use log::{debug, trace, warn};
use netwatcher::WatchHandle;
use parking_lot::Mutex;
use tokio::{
net::UdpSocket,
time::{Interval, interval},
};
use zenoh::config::ZenohId;
const GROUP: Ipv6Addr = Ipv6Addr::new(0xff12, 0, 0, 0, 0, 0, 0xe0a1, 0xde89);
pub struct Discovery {
sock: Arc<UdpSocket>,
ifaces: Arc<Mutex<Vec<SocketAddr>>>,
last_nonce: Mutex<[u8; 8]>,
/// the port of the service we are doing discovery for - transmitted to peers
listen_port: u16,
zid: ZenohId,
tick: Interval,
_sync: Mutex<WatchHandle>,
}
impl Discovery {
pub async fn new(zid: ZenohId, listen_port: u16) -> io::Result<Self> {
let discovery_port = 52413;
let sock = Arc::new(UdpSocket::bind(format!("[::]:{discovery_port}")).await?);
//sock.set_multicast_loop_v6(false)?;
let ifaces: Arc<Mutex<Vec<SocketAddr>>> = Default::default();
let _sync = Mutex::new(
netwatcher::watch_interfaces_with_callback({
let sock = sock.clone();
let ifaces = ifaces.clone();
move |update| {
for (iface_idx, iface) in update.interfaces.iter() {
if iface
.ipv6_ips()
.all(|addr| addr.is_loopback() || addr.is_unspecified())
{
continue;
}
if let Err(e) = sock.join_multicast_v6(&GROUP, *iface_idx).inspect(|_| {
ifaces.lock().push(SocketAddr::V6(SocketAddrV6::new(
GROUP, 52413, 0, *iface_idx,
)))
}) {
if let Some(iface) = update.interfaces.get(&iface_idx) {
warn!(
"failed to join multicast v6 for interface {}: {e}",
iface.name
)
}
}
}
for iface_idx in update.diff.removed {
ifaces.lock().retain(|addr| {
if let SocketAddr::V6(v6) = addr {
v6.scope_id() != iface_idx
} else {
true
}
});
if let Err(e) = sock.leave_multicast_v6(&GROUP, iface_idx) {
if let Some(iface) = update.interfaces.get(&iface_idx) {
warn!(
"failed to leave multicast v6 for interface {}: {e}",
iface.name
)
}
}
}
}
})
// todo: better error handling here
.expect("failed to bind discovery watcher"),
);
Ok(Self {
sock,
ifaces,
last_nonce: Default::default(),
listen_port,
zid,
tick: interval(Duration::from_secs(1)),
_sync,
})
}
pub async fn next(&mut self) -> io::Result<Discovered> {
let mut buf = [0u8; Hello::buf_size() + WhatsUp::buf_size() + 1];
loop {
tokio::select! {
_ = self.tick.tick() => {
self.announce().await?;
}
res = self.sock.recv_from(&mut buf) => {
let Ok((bytes_read, addr)) = res else { continue; };
if let Some(discovered) = self.respond(bytes_read, addr, &mut buf).await? {
return Ok(discovered)
}
}
}
}
}
async fn respond(
&self,
bytes_read: usize,
addr: SocketAddr,
buf: &mut [u8],
) -> io::Result<Option<Discovered>> {
trace!(
"raw recv: {bytes_read} bytes from {addr}: {:02x?}",
&buf[..bytes_read]
);
if bytes_read < size_of::<Header>() {
trace!("dropped: early EOF");
return Ok(None);
};
let header: &Header = bytemuck::from_bytes(&buf[0..size_of::<Header>()]);
if header.magic != *b"EXO" {
trace!("dropped: wrong magic");
return Ok(None);
};
let Ok(kind) = header.kind.try_into() else {
trace!("dropped: unknown message kind {}", header.kind);
return Ok(None);
};
match kind {
Kind::Hello => {
let total = Hello::buf_size();
if bytes_read != total {
trace!("dropped: hello wrong size");
return Ok(None);
}
let hello: &Hello = bytemuck::from_bytes(&buf[size_of::<Header>()..total]);
if hello.nonce == *self.last_nonce.lock() {
trace!("dropped: local hello nonce");
return Ok(None);
}
// reply
let mut reply_buf = [0u8; WhatsUp::buf_size()];
let reply = WhatsUp {
nonce: hello.nonce,
zid: self.zid.to_le_bytes(),
port_le: self.listen_port.to_le_bytes(),
};
reply.write_into(&mut reply_buf);
for i in 0..4 {
if self
.sock
.send_to(&reply_buf, addr)
.await
.inspect_err(|e| debug!("send to {addr} failed: {e}"))
.is_ok_and(|sent| sent == WhatsUp::buf_size())
{
trace!(
"sent {} bytes to {addr} after {} attempt(s)",
WhatsUp::buf_size(),
i + 1
);
break;
};
tokio::time::sleep(Duration::from_millis(300)).await;
}
Ok(None)
}
Kind::WhatsUp => {
let total = WhatsUp::buf_size();
if bytes_read != total {
trace!("dropped: whatsup wrong size");
return Ok(None);
}
let whats_up: &WhatsUp = bytemuck::from_bytes(&buf[size_of::<Header>()..total]);
if whats_up.nonce == [0u8; 8] || whats_up.nonce != *self.last_nonce.lock() {
trace!("dropped: stale nonce");
return Ok(None);
}
let SocketAddr::V6(v6) = addr else {
trace!("dropped: v4 addr used");
return Ok(None);
};
let Ok(zid) = ZenohId::try_from(&whats_up.zid[..]) else {
trace!("dropped: zenoh conversion failed");
return Ok(None);
};
if zid == self.zid {
trace!("dropped: self zenoh id");
return Ok(None);
}
// discovered
let addr = {
let mut x = v6.clone();
x.set_port(u16::from_le_bytes(whats_up.port_le));
x
};
Ok(Some(Discovered { addr, zid }))
}
}
}
async fn announce(&self) -> io::Result<()> {
let nonce = rand::random();
*self.last_nonce.lock() = nonce;
let hello = Hello { nonce };
let mut buf = [0u8; Hello::buf_size()];
hello.write_into(&mut buf);
let addrs = self.ifaces.lock().clone();
debug!("announcing {hello:?} to {addrs:?}");
// rev so .remove() doesn't break things
for (i, addr) in addrs.into_iter().enumerate().rev() {
match self.sock.send_to(&buf, addr).await {
Ok(bytes) => trace!("sent {bytes} to {addr}"),
Err(e) if e.kind() == ErrorKind::HostUnreachable => {
debug!("disabling discovery address {addr}: {e}");
_ = self.ifaces.lock().swap_remove(i);
}
Err(e) => debug!("failed to reach {addr}: {e}"),
}
}
Ok(())
}
}
pub trait Message: Pod {
const KIND: Kind;
fn header() -> Header {
Header {
magic: *b"EXO",
kind: Self::KIND as u8,
}
}
fn write_into(&self, buf: &mut [u8]) {
let total = size_of::<Header>() + size_of::<Self>();
assert!(total <= buf.len());
buf[0..size_of::<Header>()].copy_from_slice(bytemuck::bytes_of(&Self::header()));
buf[size_of::<Header>()..total].copy_from_slice(bytemuck::bytes_of(self));
}
}
#[repr(u8)]
#[derive(Debug, Clone, Copy)]
// packet & version
pub enum Kind {
Hello = 0,
WhatsUp = 1,
}
#[derive(Debug, Clone, Copy)]
pub struct Discovered {
pub zid: ZenohId,
pub addr: SocketAddrV6,
}
pub struct UnknownKind;
impl TryFrom<u8> for Kind {
type Error = UnknownKind;
fn try_from(value: u8) -> Result<Self, Self::Error> {
match value {
0 => Ok(Kind::Hello),
1 => Ok(Kind::WhatsUp),
_ => Err(UnknownKind),
}
}
}
#[repr(C)]
#[derive(Debug, Clone, Copy, Pod, Zeroable)]
pub struct Header {
magic: [u8; 3],
kind: u8,
}
#[repr(C)]
#[derive(Debug, Clone, Copy, Pod, Zeroable)]
pub struct Hello {
pub nonce: [u8; 8],
}
impl Hello {
const fn buf_size() -> usize {
size_of::<Header>() + size_of::<Self>()
}
}
impl Message for Hello {
const KIND: Kind = Kind::Hello;
}
#[repr(C)]
#[derive(Debug, Clone, Copy, Pod, Zeroable)]
pub struct WhatsUp {
pub nonce: [u8; 8],
pub zid: [u8; 16],
pub port_le: [u8; 2],
}
impl WhatsUp {
const fn buf_size() -> usize {
size_of::<Header>() + size_of::<Self>()
}
}
impl Message for WhatsUp {
const KIND: Kind = Kind::WhatsUp;
}
+51 -38
View File
@@ -1,26 +1,24 @@
use std::env;
use std::{env, panic, sync::Arc};
use tokio::task::JoinHandle;
use zenoh::{Result, Session as ZSession, config::Locator};
use netwatcher::WatchHandle;
use parking_lot::Mutex;
use tokio::{sync::mpsc, task::JoinHandle};
use zenoh::{Result, Session as ZSession, config::WhatAmI, internal::runtime::Runtime};
use zenoh_plugin_storage_manager::StoragesPlugin;
use zenoh_plugin_trait::PluginsManager;
pub use zenoh::{Config, config::ZenohId};
use crate::discovery::Discovery;
pub mod discovery;
pub mod swarm;
pub fn cfg(identity: u128, listen_port: u16) -> Result<zenoh::Config> {
assert!(listen_port != 0, "must used defined listen port port");
let namespace = env::var("EXO_ZENOH_NAMESPACE").unwrap_or_else(|_| "exo".to_string());
let mut cfg = zenoh::Config::default();
// todo: cleanup
cfg.insert_json5("id", &format!("\"{identity:x}\""))?;
cfg.insert_json5("mode", "\"router\"")?;
cfg.insert_json5("mode", "\"peer\"")?;
cfg.insert_json5("listen/endpoints", &format!("[\"tcp/[::]:{listen_port}\"]"))?;
cfg.insert_json5("scouting/multicast/enabled", "false")?;
cfg.insert_json5("scouting/multicast/enabled", "true")?;
cfg.insert_json5("scouting/multicast/autoconnect", "[]")?;
cfg.insert_json5("scouting/gossip/multihop", "true")?;
cfg.insert_json5("namespace", &format!("{namespace:?}"))?;
@@ -41,8 +39,7 @@ pub fn cfg(identity: u128, listen_port: u16) -> Result<zenoh::Config> {
Ok(cfg)
}
pub async fn open(cfg: zenoh::Config, listen_port: u16) -> Result<Session> {
assert!(listen_port != 0, "must used defined listen port");
pub async fn open(cfg: zenoh::Config) -> Result<Session> {
let mut plugins = PluginsManager::static_plugins_only();
plugins.declare_static_plugin::<StoragesPlugin, _>("storage_manager", true);
let mut runtime = zenoh::internal::runtime::RuntimeBuilder::new(cfg)
@@ -51,42 +48,58 @@ pub async fn open(cfg: zenoh::Config, listen_port: u16) -> Result<Session> {
.await?;
let z = zenoh::session::init(runtime.clone().into()).await?;
runtime.start().await?;
let mut discovery = Discovery::new(z.zid(), listen_port).await?;
let _jh = tokio::task::spawn(async move {
let _watch_all_handle = watch_all(runtime).await?;
Ok(Session {
z,
_watch_all_handle,
})
}
async fn watch_all(runtime: Runtime) -> Result<WatchAllHandle> {
log::info!("spawning scout");
let mut cfg = Config::default();
cfg.insert_json5("scouting/multicast/ttl", "3")?;
cfg.insert_json5("scouting/multicast/interface", "\"auto\"")?;
let mut scout = zenoh::scout(WhatAmI::Peer, cfg.clone()).await?;
let (send, mut recv) = mpsc::unbounded_channel();
let _sync = Arc::new(Mutex::new(netwatcher::watch_interfaces_with_callback(
move |u| _ = send.send(u),
)?));
let _async = tokio::task::spawn(async move {
loop {
let Ok(discovered) = discovery.next().await.inspect_err(|e| {
log::warn!("discovery error {e}");
}) else {
continue;
};
if discovered.zid > runtime.zid() {
log::debug!("not connecting to peer with greater zid");
continue;
tokio::select! {
u = recv.recv() => {
if u.is_none() {
return Ok(());
}
log::info!("reloading scout");
scout = zenoh::scout(WhatAmI::Peer, cfg.clone()).await?;
}
hello = scout.recv_async() => {
if let Ok(hello) = hello {
// todo: auth
runtime
.connect_peer(&hello.zid().into(), hello.locators())
.await;
}
}
}
let Ok(locator) =
Locator::new("tcp", discovered.addr.to_string(), "").inspect_err(|e| {
log::warn!("failed to pass locator from addr: {e}");
})
else {
continue;
};
runtime
.connect_peer(&discovered.zid.into(), &[locator])
.await;
}
});
Ok(Session { z, _jh })
Ok(WatchAllHandle { _sync, _async })
}
pub struct Session {
pub z: ZSession,
_jh: JoinHandle<()>,
_watch_all_handle: WatchAllHandle,
}
impl Drop for Session {
impl Drop for WatchAllHandle {
fn drop(&mut self) {
self._jh.abort();
log::error!("aborting iface watcher");
self._async.abort();
}
}
pub struct WatchAllHandle {
_sync: Arc<Mutex<WatchHandle>>,
_async: JoinHandle<Result<()>>,
}
+3 -2
View File
@@ -50,6 +50,7 @@ impl Swarm {
mut from_client,
} = self;
let stream = async_stream::stream! {
// very important!
let mut session = session;
let (mut to_topics, mut from_topics) = mpsc::channel(1024);
let mut topics = Topics::new();
@@ -179,8 +180,8 @@ pub async fn create_swarm(
if !bootstrap_peers.is_empty() || listen_port != 0 {
todo!();
}
let cfg = crate::cfg(identity, 52414)?;
let session = crate::open(cfg, 52414).await?;
let cfg = crate::cfg(identity, listen_port)?;
let session = crate::open(cfg).await?;
Ok(Swarm {
session,
from_client,
+12 -29
View File
@@ -199,7 +199,7 @@ from exo.shared.types.worker.downloads import DownloadCompleted
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
from exo.shared.types.worker.shards import Sharding
from exo.utils.banner import print_startup_banner
from exo.utils.channels import Receiver, Sender, channel
from exo.utils.channels import Receiver, Sender
from exo.utils.disk_event_log import DiskEventLog
from exo.utils.power_sampler import PowerSampler
from exo.utils.task_group import TaskGroup
@@ -234,22 +234,17 @@ def _require_disaggregation_enabled() -> None:
@dataclass
class Transport:
session: PySession
z: PySession
cancel_scopes: dict[CommandId, anyio.CancelScope] = field(
init=False, default_factory=dict
)
command_sender: NetSender = field(init=False)
paused: bool = field(init=False, default=False)
paused_ev: anyio.Event = field(init=False, default_factory=anyio.Event)
tg: TaskGroup = field(init=False, default_factory=TaskGroup)
def __post_init__(self):
# TODO: retire root keyspace
self.command_sender = self.session.net_sender("orchestrator")
async def run(self):
async with self.tg:
await anyio.sleep_forever()
self.command_sender = self.z.net_sender("orchestrator")
async def send_command(self, command: Command) -> bool:
while self.paused:
@@ -275,33 +270,22 @@ class Transport:
async def stream(
self,
command_id: CommandId,
) -> AsyncGenerator[Chunk]:
send, recv = channel[Chunk]()
self.tg.start_soon(self._run_stream, command_id, send)
async with recv:
async for item in recv:
yield item
async def _run_stream(self, command_id: CommandId, send: Sender[Chunk]):
) -> AsyncGenerator[Chunk,]:
try:
with anyio.CancelScope() as cs:
self.cancel_scopes[command_id] = cs
# recv from any node
receiver = self.session.net_receiver(
receiver = self.z.net_receiver(
f"runners/*/active_tasks/{command_id}/chunks"
)
while True:
data = await receiver.recv()
if data is None:
logger.warning(
"stream terminated early without finish reason EOF"
)
break
await send.send(
chunk := (
TypeAdapter[Chunk](Chunk).validate_json(
yield (
chunk := cast(
Chunk,
TypeAdapter(Chunk).validate_json(
data, strict=True, extra="forbid"
)
),
)
)
if (
@@ -309,7 +293,7 @@ class Transport:
and chunk.finish_reason is not None
):
break
except (anyio.get_cancelled_exc_class(), anyio.BrokenResourceError, anyio.ClosedResourceError):
except anyio.get_cancelled_exc_class():
with anyio.CancelScope(shield=True):
await self.command_sender.send(
TaskCancelled(cancelled_command_id=command_id)
@@ -326,7 +310,7 @@ class Transport:
)
def cancel(self, command_id: CommandId) -> bool:
if (cs := self.cancel_scopes.pop(command_id, None)) is not None:
if (cs := self.cancel_scopes.get(command_id, None)) is not None:
cs.cancel()
return True
return False
@@ -1850,7 +1834,6 @@ class API:
try:
async with self._tg as tg:
logger.info("Starting API")
tg.start_soon(self.transport.run)
tg.start_soon(self._apply_state)
tg.start_soon(self._pause_on_new_election)
tg.start_soon(self._cleanup_expired_images)
+14 -10
View File
@@ -1,11 +1,11 @@
# pyright: reportUnusedFunction=false, reportAny=false
from typing import Any
from unittest.mock import AsyncMock
from unittest.mock import AsyncMock, MagicMock
from fastapi import FastAPI
from fastapi.testclient import TestClient
from exo.api.main import API, Transport
from exo.api.main import API
from exo.shared.types.common import CommandId
@@ -15,9 +15,9 @@ def _make_api() -> Any:
app = FastAPI()
api = object.__new__(API)
api.app = app
api.transport = object.__new__(Transport)
api.transport.cancel = AsyncMock()
api.transport.send_command = AsyncMock()
api._text_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api._image_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api._send = AsyncMock() # pyright: ignore[reportPrivateUsage]
api._setup_exception_handlers() # pyright: ignore[reportPrivateUsage]
app.post("/v1/cancel/{command_id}")(api.cancel_command)
return api
@@ -43,14 +43,16 @@ def test_cancel_active_text_generation() -> None:
client = TestClient(api.app)
cid = CommandId("text-cmd-123")
sender = MagicMock()
api._text_generation_queues[cid] = sender
response = client.post(f"/v1/cancel/{cid}")
assert response.status_code == 200
data: dict[str, Any] = response.json()
assert data["message"] == "Command cancelled."
assert data["command_id"] == str(cid)
api.transport.cancel.assert_called_once()
api.transport.send_command.assert_called_once()
sender.close.assert_called_once()
api._send.assert_called_once()
task_cancelled = api._send.call_args[0][0]
assert task_cancelled.cancelled_command_id == cid
@@ -61,13 +63,15 @@ def test_cancel_active_image_generation() -> None:
client = TestClient(api.app)
cid = CommandId("img-cmd-456")
sender = MagicMock()
api._image_generation_queues[cid] = sender
response = client.post(f"/v1/cancel/{cid}")
assert response.status_code == 200
data: dict[str, Any] = response.json()
assert data["message"] == "Command cancelled."
assert data["command_id"] == str(cid)
api.transport.cancel.assert_called_once()
api.transport.send_command.assert_called_once()
task_cancelled = api.transport.send_command.call_args[0][0]
sender.close.assert_called_once()
api._send.assert_called_once()
task_cancelled = api._send.call_args[0][0]
assert task_cancelled.cancelled_command_id == cid
@@ -1,10 +1,9 @@
# pyright: reportUnusedFunction=false, reportAny=false
"""Tests that InstanceDeleted events close active generation streams."""
from typing import Any
from unittest.mock import MagicMock
from exo.api.main import API, Transport
from exo.api.main import API
from exo.api.types import ImageGenerationTaskParams
from exo.shared.types.common import CommandId, ModelId
from exo.shared.types.state import State
@@ -17,11 +16,12 @@ from exo.shared.types.text_generation import (
from exo.shared.types.worker.instances import InstanceId
def _make_api_with_state(state: State) -> Any:
def _make_api_with_state(state: State) -> API:
"""Create a minimal API instance with pre-set state."""
api = object.__new__(API)
api.state = state
api.transport = object.__new__(Transport)
api._text_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api._image_generation_queues = {} # pyright: ignore[reportPrivateUsage]
return api
@@ -47,10 +47,13 @@ def test_close_streams_for_deleted_instance() -> None:
state = State(tasks={task.task_id: task})
api = _make_api_with_state(state)
api._close_streams_for_instance(instance_id)
sender = MagicMock()
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
api.transport.cancel.assert_called_once()
assert api.transport.cancel.call_args[0][0] == command_id
api._close_streams_for_instance(instance_id) # pyright: ignore[reportPrivateUsage]
sender.close.assert_called_once()
assert command_id not in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
def test_close_streams_ignores_unrelated_instances() -> None:
@@ -69,6 +72,7 @@ def test_close_streams_ignores_unrelated_instances() -> None:
api._close_streams_for_instance(target_id) # pyright: ignore[reportPrivateUsage]
sender.close.assert_not_called()
assert other_cmd in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
def test_close_streams_for_deleted_instance_image_generation() -> None:
+2 -2
View File
@@ -114,7 +114,7 @@ class Node:
global_event_sender=router.sender(topics.GLOBAL_EVENTS),
local_event_receiver=router.receiver(topics.LOCAL_EVENTS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
command_receiver=session.net_receiver("orchestrator"),
session=session,
)
er_send, er_recv = channel[ElectionResult]()
@@ -214,7 +214,7 @@ class Node:
download_command_sender=self.router.sender(
topics.DOWNLOAD_COMMANDS
),
command_receiver=self.session.net_receiver("orchestrator"),
session=self.session,
)
self._tg.start_soon(self.master.run)
elif (
+14 -12
View File
@@ -1,7 +1,8 @@
from datetime import datetime, timedelta, timezone
from typing import cast
import anyio
from exo_net import NetReceiver
from exo_net import PySession
from loguru import logger
from pydantic import TypeAdapter
@@ -123,18 +124,18 @@ class Master:
node_id: NodeId,
session_id: SessionId,
*,
command_receiver: NetReceiver, # todo: not this type
event_sender: Sender[Event],
local_event_receiver: Receiver[LocalForwarderEvent],
global_event_sender: Sender[GlobalForwarderEvent],
download_command_sender: Sender[ForwarderDownloadCommand],
session: PySession,
):
self.node_id = node_id
self.session_id = session_id
self.state = State()
self._tg: TaskGroup = TaskGroup()
self.command_task_mapping: dict[CommandId, TaskId] = {}
self.command_receiver = command_receiver
self.session = session
self.local_event_receiver = local_event_receiver
self.global_event_sender = global_event_sender
self.download_command_sender = download_command_sender
@@ -163,12 +164,12 @@ class Master:
self._tg.cancel_tasks()
async def _command_processor(self) -> None:
receiver = self.session.net_receiver("orchestrator")
while True:
data = await self.command_receiver.recv()
if not data:
break
command = cast(
Command, TypeAdapter(Command).validate_json(await receiver.recv())
)
try:
command = TypeAdapter[Command](Command).validate_json(data)
logger.info(f"Executing command: {command}")
generated_events: list[Event] = []
@@ -277,11 +278,10 @@ class Master:
selected_instance_id
)
if selected_instance:
ranks = set(
self._expected_ranks[task_id] = set(
shard.device_rank
for shard in selected_instance.shard_assignments.runner_to_shard.values()
)
self._expected_ranks[task_id] = ranks
case ImageEdits():
for instance in self.state.instances.values():
if (
@@ -329,11 +329,11 @@ class Master:
selected_instance_id
)
if selected_instance:
ranks = set(
self._expected_ranks[task_id] = set(
shard.device_rank
for shard in selected_instance.shard_assignments.runner_to_shard.values()
)
self._expected_ranks[task_id] = ranks
case DeleteInstance():
placement = delete_instance(command, self.state.instances)
transition_events = get_transition_events(
@@ -431,7 +431,9 @@ class Master:
InstanceLinkDeleted(link_id=command.link_id)
)
case RequestEventLog():
end = len(self._event_log)
# We should just be able to send everything, since other buffers will ignore old messages
# rate limit to 1000 at a time
end = min(command.since_idx + 1000, len(self._event_log))
for i, event in enumerate(
self._event_log.read_range(command.since_idx, end),
start=command.since_idx,
+3 -1
View File
@@ -6,6 +6,7 @@ import pytest
from loguru import logger
from exo.master.main import Master
from exo.routing.router import get_node_id_keypair
from exo.shared.models.model_cards import ModelCard, ModelTask
from exo.shared.types.commands import (
CommandId,
@@ -46,7 +47,8 @@ from exo.utils.channels import channel
@pytest.mark.asyncio
async def test_master():
node_id = NodeId("master test")
keypair = get_node_id_keypair()
node_id = NodeId(keypair.to_node_id())
session_id = SessionId(master_node_id=node_id, election_clock=0)
ge_sender, global_event_receiver = channel[GlobalForwarderEvent]()
+2 -1
View File
@@ -4,7 +4,6 @@ from random import random
import anyio
from anyio import BrokenResourceError, ClosedResourceError
from anyio.abc import CancelScope
from exo_net import NetSender
from loguru import logger
from exo.shared.types.commands import RequestEventLog
@@ -20,6 +19,8 @@ from exo.utils.channels import Receiver, Sender, channel
from exo.utils.event_buffer import OrderedBuffer
from exo.utils.task_group import TaskGroup
from exo_net import NetSender
@dataclass
class EventRouter:
+2 -12
View File
@@ -32,21 +32,15 @@ class TokenChunk(BaseChunk):
class ErrorChunk(BaseChunk):
error_message: str
@property
def finish_reason(self) -> Literal["error"]:
return "error"
finish_reason: Literal["error"] = "error"
class ToolCallChunk(BaseChunk):
tool_calls: list[ToolCallItem]
usage: Usage | None
finish_reason: Literal["tool_calls"] = "tool_calls"
stats: GenerationStats | None = None
@property
def finish_reason(self) -> Literal["tool_calls"]:
return "tool_calls"
class ImageChunk(BaseChunk):
data: str
@@ -90,10 +84,6 @@ class PrefillProgressChunk(BaseChunk):
processed_tokens: int
total_tokens: int
@property
def finish_reason(self) -> FinishReason | None:
return None
StatusChunk = PrefillProgressChunk
GenerationChunk = TokenChunk | ImageChunk | ToolCallChunk | ErrorChunk
-1
View File
@@ -161,7 +161,6 @@ Event = (
| NodeDownloadProgress
| TopologyEdgeCreated
| TopologyEdgeDeleted
| ChunkGenerated
| InputChunkReceived
| TracesCollected
| TracesMerged
+2 -2
View File
@@ -7,7 +7,7 @@ import mlx.core as mx
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.shared.types.common import ModelId
from exo.shared.types.events import Event
from exo.shared.types.events import ChunkGenerated, Event
from exo.shared.types.tasks import TaskId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import ModelLoadingResponse
@@ -32,7 +32,7 @@ from .vision import VisionProcessor
@dataclass
class MlxBuilder(Builder):
model_id: ModelId
event_sender: MpSender[Event]
event_sender: MpSender[Event | ChunkGenerated]
cancel_receiver: MpReceiver[TaskId]
inference_model: Model | None = None
tokenizer: TokenizerWrapper | None = None
+1
View File
@@ -15,6 +15,7 @@ from exo.shared.models.model_cards import ModelId, card_cache
from exo.shared.types.chunks import InputImageChunk
from exo.shared.types.commands import (
DeleteInstance,
ForwarderCommand,
ForwarderDownloadCommand,
StartDownload,
)
@@ -96,7 +96,7 @@ class SequentialGenerator(Engine):
model_id: ModelId
device_rank: int
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
event_sender: MpSender[Event | ChunkGenerated]
vision_processor: VisionProcessor | None = None
check_for_cancel_every: int = 50
@@ -327,7 +327,7 @@ class BatchGenerator(Engine):
model_id: ModelId
device_rank: int
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
event_sender: MpSender[Event | ChunkGenerated]
check_for_cancel_every: int = 50
vision_processor: VisionProcessor | None = None
+1 -1
View File
@@ -86,7 +86,7 @@ class Runner:
self,
bound_instance: BoundInstance,
builder: Builder,
event_sender: MpSender[Event],
event_sender: MpSender[Event | ChunkGenerated],
task_receiver: MpReceiver[Task],
):
self.event_sender = event_sender
+6 -13
View File
@@ -53,7 +53,7 @@ class RunnerSupervisor:
shard_metadata: ShardMetadata
bound_instance: BoundInstance
runner_process: mp.Process
_ev_recv: MpReceiver[Event]
_ev_recv: MpReceiver[Event | ChunkGenerated]
_task_sender: MpSender[Task]
_event_sender: Sender[Event]
_cancel_sender: MpSender[TaskId]
@@ -77,7 +77,7 @@ class RunnerSupervisor:
event_sender: Sender[Event],
session: PySession,
) -> Self:
ev_send, ev_recv = mp_channel[Event]()
ev_send, ev_recv = mp_channel[Event | ChunkGenerated]()
task_sender, task_recv = mp_channel[Task]()
cancel_sender, cancel_recv = mp_channel[TaskId]()
@@ -215,19 +215,12 @@ class RunnerSupervisor:
with self._ev_recv as events:
async for event in events:
if isinstance(event, ChunkGenerated):
if (pub := pubs.get(event.command_id, None)) is None:
pub = pubs[event.command_id] = self.session.net_sender(
if event.command_id not in pubs:
pubs[event.command_id] = self.session.net_sender(
f"runners/{self.bound_instance.bound_runner_id}/active_tasks/{event.command_id}/chunks"
)
sent = await pub.send(
event.chunk.model_dump_json().encode("utf-8")
)
if not sent:
logger.warning(
"api node closed communication, dropping chunk"
)
if event.chunk.finish_reason is not None:
pubs.pop(event.command_id, None)
pub = pubs[event.command_id]
await pub.send(event.chunk.model_dump_json().encode("utf-8"))
continue
if isinstance(event, RunnerStatusUpdated):