mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-11 04:49:25 -04:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4430b0daf9 |
No files matched your search
Generated
+2
-21
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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]: ...
|
||||
|
||||
@@ -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>()?;
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(),
|
||||
//}
|
||||
//}
|
||||
}
|
||||
@@ -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: ...
|
||||
|
||||
...
|
||||
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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<()>>,
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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]()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -161,7 +161,6 @@ Event = (
|
||||
| NodeDownloadProgress
|
||||
| TopologyEdgeCreated
|
||||
| TopologyEdgeDeleted
|
||||
| ChunkGenerated
|
||||
| InputChunkReceived
|
||||
| TracesCollected
|
||||
| TracesMerged
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in new issue
Block a user