Compare commits

..
10 Commits
Author SHA1 Message Date
Evan 9c25164744 rebase fix 2026-05-09 09:53:40 +01:00
Evan ece77c738a custom discovery take 2 2026-05-09 00:58:12 +01:00
Evan 62343fb090 api streams 2026-05-09 00:58:12 +01:00
Evan 6efb77ca2d json state proxy 2026-05-09 00:56:30 +01:00
Evan 46cd9c0af3 uncap 2026-05-09 00:55:38 +01:00
Evan 0a3e5fdc3c libp2p -> zenoh 2026-05-09 00:55:38 +01:00
Evan c58563bef3 rename 2026-05-09 00:55:09 +01:00
Andrei Cravtov e5a1e5dadb Create PID file locking for EXO (#2072)
## Motivation

EXO should be PID file locked, to prevent duplicate processes from
clobbering the log, right now this isn't the case.

## Changes

I added a wrapper around a Rust PID file lock library, and used it to
implement PID locking for EXO, with the PID file being in exo cache
directory.

## Test Plan

### Manual Testing
Tested on e11, trying to spawn duplicate EXO processes prevented.
2026-05-08 18:50:18 +01:00
ciaranbor fa57131374 Integration tests infra (#1995)
## Motivation

No automated integration tests exist for exo. Manual testing against
real hardware clusters is slow and error-prone. We need a pytest
framework that deploys clusters via `eco`, runs inference scenarios, and
tears down cleanly.

## Changes

- **`tools/src/exo_tools/`** — New workspace member shared by bench,
eval, and tests:
- `client.py` — `ExoClient` HTTP client (extracted from
`bench/harness.py`)
- `harness.py` — instance lifecycle helpers (placement, wait-for-ready,
etc.)
- `cluster.py` — `EcoSession` for eco cluster lifecycle
(deploy/stop/start/release/logs/exec) with unique `USER=<prefix>-<uuid>`
per session and atexit/signal cleanup
- **`tests/integration/`** — 17 pytest tests across 5 files:
- `test_1node.py` — place, chat, multi-turn, delete, state/models
endpoints, cluster snapshot, download-from-scratch
- `test_2node.py` — parametrized tensor/jaccl + pipeline/ring inference
and multi-turn
- `test_4node.py` — parametrized 4-node pipeline/ring inference, cluster
state
- `test_resilience.py` — full disconnect/reconnect cycle (2-node →
disconnect → 1-node → reconnect → 2-node)
- `test_dashboard.py` — Playwright: dashboard loads, shows node info,
chat flow
- `helpers.py` — placement/inference helpers, re-exports from
`exo_tools`
- `conftest.py` — session-scoped cluster fixtures with constraint-based
eco reservations; `--hosts` override; `EXO_REF` env var for CI
deployments from a GitHub branch
- **`bench/`** — Updated imports from `exo_tools.client` /
`exo_tools.harness`
- **`pyproject.toml`** — Added `tools` workspace member, `playwright`
dev dep, `--ignore=tests/integration`

## Why It Works

Tests use `eco` for cluster lifecycle and `ExoClient` for API
interactions — same tools humans use. Session-scoped fixtures deploy
once per file. Unique eco users prevent test runs from interfering with
each other or manual usage.

## Test Plan

### Automated Testing

- `uv run pytest tests/integration/ -v -s` — full suite (~4-5 min, 17/17
passing)
- `uv run pytest tests/integration/ -v -s --hosts s4,s9,s10,s22` — pin
specific hosts
- `EXO_REF=main uv run pytest tests/integration/ -v` — deploy from a
GitHub branch (CI)
- `uv run pytest` — confirms integration tests are excluded from default
runs
2026-05-08 17:15:08 +01:00
Alex Cheema 414132ae9c Use time-weighted power sampling (#2038)
## Why

The power sampler currently averages sampled wattage values
arithmetically. That can be materially wrong when sample intervals are
uneven: a short high-power spike gets the same weight as a long steady
interval. Energy should be computed by integrating power over time, and
average power should be derived from energy / elapsed time.

## How

- Store each power sample with its relative timestamp.
- Anchor the first sample at `t=0` and take a final sample at `elapsed`
when producing results.
- Integrate per-node power using the trapezoidal rule.
- Sum node energy for total cluster energy, then derive total average
system power from total energy / elapsed.
- Add focused unit tests for uneven sample intervals and the
single-sample fallback.

## Tests

- `uv run pytest src/exo/utils/tests/test_power_sampler.py`
- `uv run basedpyright`
- `uv run ruff check src/exo/utils/power_sampler.py
src/exo/utils/tests/test_power_sampler.py`
- `nix fmt`
2026-05-07 10:42:14 +00:00
77 changed files with 3373 additions and 1631 deletions

No files matched your search

+1
View File
@@ -40,3 +40,4 @@ bench/**/*.json
tmp/models
/build/exo
/.claude/skills
/.claude
Generated
+177 -116
View File
@@ -142,9 +142,9 @@ dependencies = [
[[package]]
name = "anyhow"
version = "1.0.100"
version = "1.0.102"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a23eb6b1614318a8071c9b2521f36b424b2c83db5eb3a0fead4a6c0809af6e61"
checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c"
[[package]]
name = "arc-swap"
@@ -182,7 +182,7 @@ checksum = "3109e49b1e4909e9db6515a30c633684d68cdeaa252f215214cb4fa1a5bfee2c"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
"synstructure",
]
@@ -194,7 +194,7 @@ checksum = "7b18050c2cd6fe86c3a76584ef5e0baf286d038cda203eb6223df2cc413565f7"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -216,7 +216,7 @@ checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -227,7 +227,7 @@ checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -319,6 +319,26 @@ 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"
@@ -372,9 +392,9 @@ dependencies = [
[[package]]
name = "chrono"
version = "0.4.42"
version = "0.4.44"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "145052bdd345b87320e369255277e3fb5152762ad123a901ef5c262dd38fe8d2"
checksum = "c673075a2e0e5f4a1dde27ce9dee1ea4558c7ffe648f576438a20ca1d2acc4b0"
dependencies = [
"iana-time-zone",
"js-sys",
@@ -608,7 +628,7 @@ dependencies = [
"proc-macro2",
"quote",
"strsim",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -619,7 +639,7 @@ checksum = "ac3984ec7bd6cfa798e62b4a642426a5be0e68f9401cfc2a01e3fa9ea2fcdb8d"
dependencies = [
"darling_core",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -681,7 +701,7 @@ dependencies = [
"proc-macro2",
"quote",
"rustc_version",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -725,7 +745,7 @@ checksum = "97369cbbc041bc366949bc74d34658d6cda5621039731c6310521892a3a20ae0"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -781,7 +801,7 @@ dependencies = [
]
[[package]]
name = "exo_pyo3_bindings"
name = "exo_net"
version = "0.0.1"
dependencies = [
"env_logger",
@@ -789,15 +809,17 @@ dependencies = [
"futures-lite",
"log",
"networking",
"parking_lot",
"pidfile-rs",
"pin-project",
"pyo3",
"pyo3-async-runtimes",
"pyo3-log",
"pyo3-stub-gen",
"rand 0.10.1",
"serde_json",
"tokio",
"zenoh",
"zerompk",
]
[[package]]
@@ -808,7 +830,7 @@ checksum = "311a6d2f1f9d60bff73d2c78a0af97ed27f79672f15c238192a5bbb64db56d00"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -862,6 +884,16 @@ dependencies = [
"miniz_oxide",
]
[[package]]
name = "flopen"
version = "0.1.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fbfb8b5fbd1f27929f216650081a07b6ceb0741f0542c8c43ff7ef8e93a35a5d"
dependencies = [
"libc",
"nix 0.31.2",
]
[[package]]
name = "flume"
version = "0.11.1"
@@ -974,7 +1006,7 @@ checksum = "162ee34ebcb7c64a8abebc059ce0fee27c2262618d7b60ed8faf72fef13c3650"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -1084,7 +1116,7 @@ checksum = "53010ccb100b96a67bc32c0175f0ed1426b31b655d562898e57325f81c023ac0"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -1123,6 +1155,12 @@ dependencies = [
"foldhash 0.2.0",
]
[[package]]
name = "hashbrown"
version = "0.17.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4f467dd6dccf739c208452f8014c75c18bb8301b050ad1cfb27153803edb0f51"
[[package]]
name = "heck"
version = "0.5.0"
@@ -1332,12 +1370,12 @@ dependencies = [
[[package]]
name = "indexmap"
version = "2.12.1"
version = "2.14.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0ad4bb2b565bca0645f4d68c5c9af97fba094e9791da685bf83cb5f3ce74acf2"
checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9"
dependencies = [
"equivalent",
"hashbrown 0.16.1",
"hashbrown 0.17.0",
"serde",
"serde_core",
]
@@ -1362,9 +1400,9 @@ dependencies = [
[[package]]
name = "inventory"
version = "0.3.21"
version = "0.3.24"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bc61209c082fbeb19919bee74b176221b27223e27b65d781eb91af24eb1fb46e"
checksum = "a4f0c30c76f2f4ccee3fe55a2435f691ca00c0e4bd87abe4f4a851b1d4dac39b"
dependencies = [
"rustversion",
]
@@ -1387,7 +1425,7 @@ dependencies = [
"heck",
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -1450,7 +1488,7 @@ checksum = "e000de030ff8022ea1da3f466fbb0f3a809f5e51ed31f6dd931c35181ad8e6d7"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -1496,7 +1534,7 @@ dependencies = [
"quote",
"rustc_version",
"simd_cesu8",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -1524,7 +1562,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264"
dependencies = [
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -1813,17 +1851,16 @@ name = "networking"
version = "0.0.1"
dependencies = [
"async-stream",
"bytemuck",
"futures-lite",
"log",
"netwatcher",
"parking_lot",
"rand 0.10.1",
"tokio",
"tracing",
"zenoh",
"zenoh-plugin-storage-manager",
"zenoh-plugin-trait",
"zerompk",
]
[[package]]
@@ -2057,9 +2094,9 @@ checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d"
[[package]]
name = "ordered-float"
version = "5.1.0"
version = "5.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f4779c6901a562440c3786d08192c6fbda7c1c2060edd10006b05ee35d10f2d"
checksum = "b7d950ca161dc355eaf28f82b11345ed76c6e1f6eb1f4f4479e0323b9e2fbd0e"
dependencies = [
"num-traits",
]
@@ -2160,7 +2197,7 @@ dependencies = [
"pest_meta",
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -2181,7 +2218,7 @@ checksum = "8701b58ea97060d5e5b155d383a69952a60943f0e6dfe30b04c287beb0b27455"
dependencies = [
"fixedbitset",
"hashbrown 0.15.5",
"indexmap 2.12.1",
"indexmap 2.14.0",
"serde",
]
@@ -2245,7 +2282,7 @@ dependencies = [
"phf_shared 0.13.1",
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -2266,6 +2303,18 @@ dependencies = [
"siphasher",
]
[[package]]
name = "pidfile-rs"
version = "0.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d1a8aa9a30b1b65ef48b333931b80f2324a14e00208eb2b8f5788f1180791bcc"
dependencies = [
"flopen",
"libc",
"log",
"thiserror 1.0.69",
]
[[package]]
name = "pin-project"
version = "1.1.10"
@@ -2283,7 +2332,7 @@ checksum = "6e918e4ff8c4549eb882f14b3a4bc8c8bc93de829416eacf579f1207a8fbf861"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -2397,7 +2446,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b"
dependencies = [
"proc-macro2",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -2411,9 +2460,9 @@ dependencies = [
[[package]]
name = "proc-macro2"
version = "1.0.103"
version = "1.0.106"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5ee95bc4ef87b8d5ba32e8b7714ccc834865276eab0aed5c9958d00ec45f49e8"
checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934"
dependencies = [
"unicode-ident",
]
@@ -2459,7 +2508,7 @@ checksum = "bcd7d70ee0ca1661c40407e6f84e4463ef2658c90a9e2fbbd4515b2bcdfcaeca"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -2501,7 +2550,7 @@ dependencies = [
"proc-macro2",
"pyo3-macros-backend",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -2514,19 +2563,19 @@ dependencies = [
"proc-macro2",
"pyo3-build-config",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
name = "pyo3-stub-gen"
version = "0.17.2"
version = "0.22.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "398b833826a83ca72c1e26d1b2c7c71f9ca7c3bfc74eacc663901895c362ae33"
checksum = "b695173a33bec37b6acd288efe74f902f123f5891e5ff5ad53ff429e18833c14"
dependencies = [
"anyhow",
"chrono",
"either",
"indexmap 2.12.1",
"indexmap 2.14.0",
"inventory",
"itertools 0.14.0",
"log",
@@ -2536,22 +2585,25 @@ dependencies = [
"ordered-float",
"pyo3",
"pyo3-stub-gen-derive",
"rustpython-parser",
"serde",
"toml",
"serde_json",
"time",
"toml 1.1.2+spec-1.1.0",
]
[[package]]
name = "pyo3-stub-gen-derive"
version = "0.17.2"
version = "0.22.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2426ba759d848787239d80f9fdb1f223786976f87fb6c3da8188ca7c17744b28"
checksum = "c2f901834d55c74f36be3353994062e3c3b3d47fda25572f5535e99434612ac8"
dependencies = [
"heck",
"indexmap 2.12.1",
"indexmap 2.14.0",
"proc-macro2",
"quote",
"rustpython-parser",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -2613,9 +2665,9 @@ dependencies = [
[[package]]
name = "quote"
version = "1.0.42"
version = "1.0.45"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a338cc41d27e6cc6dce6cefc13a0729dfbb81c262b1f519331575dd80ef3067f"
checksum = "41f2619966050689382d2b44f664f4bc593e129785a36d6ee376ddf37259b924"
dependencies = [
"proc-macro2",
]
@@ -2765,7 +2817,7 @@ checksum = "b7186006dcb21920990093f30e3dea63b7d6e977bf1256be20c3563a5db070da"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3091,7 +3143,7 @@ dependencies = [
"proc-macro2",
"quote",
"serde_derive_internals",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3166,7 +3218,7 @@ checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3177,7 +3229,7 @@ checksum = "18d26a20a969b9e3fdf2fc2d9f21eda6c40e2de84c9408bb5d3b05d499aae711"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3195,9 +3247,9 @@ dependencies = [
[[package]]
name = "serde_spanned"
version = "1.0.3"
version = "1.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e24345aa0fe688594e73770a5f6d1b216508b4f93484c0026d521acd30134392"
checksum = "6662b5879511e06e8999a8a235d848113e942c9124f211511b16466ee2995f26"
dependencies = [
"serde_core",
]
@@ -3212,7 +3264,7 @@ dependencies = [
"chrono",
"hex",
"indexmap 1.9.3",
"indexmap 2.12.1",
"indexmap 2.14.0",
"schemars 0.9.0",
"schemars 1.2.1",
"serde_core",
@@ -3230,7 +3282,7 @@ dependencies = [
"darling",
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3239,7 +3291,7 @@ version = "0.9.34+deprecated"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6a8b1a1a2ebf674015cc02edccce75287f1a0130d394307b36743c2f5d504b47"
dependencies = [
"indexmap 2.12.1",
"indexmap 2.14.0",
"itoa",
"ryu",
"serde",
@@ -3484,9 +3536,9 @@ dependencies = [
[[package]]
name = "syn"
version = "2.0.111"
version = "2.0.117"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "390cc9a294ab71bdb1aa2e99d13be9c753cd2d7bd6560c77118597410c4d2e87"
checksum = "e665b8803e7b1d2a727f4023456bbbbe74da67099c585258af0ad9c5013b9b99"
dependencies = [
"proc-macro2",
"quote",
@@ -3501,7 +3553,7 @@ checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3536,7 +3588,7 @@ checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3547,7 +3599,7 @@ checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3662,7 +3714,6 @@ dependencies = [
"signal-hook-registry",
"socket2 0.6.1",
"tokio-macros",
"tracing",
"windows-sys 0.61.2",
]
@@ -3674,7 +3725,7 @@ checksum = "af407857209536a95c8e56f8231ef2c2e2aff839b22e07a1ffcbc617e9db9fa5"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -3719,13 +3770,28 @@ version = "0.9.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f0dc8b1fb61449e27716ec0e1bdf0f6b8f3e8f6b05391e8497b8b6d7804ea6d8"
dependencies = [
"indexmap 2.12.1",
"indexmap 2.14.0",
"serde_core",
"serde_spanned",
"toml_datetime",
"toml_datetime 0.7.3",
"toml_parser",
"toml_writer",
"winnow",
"winnow 0.7.14",
]
[[package]]
name = "toml"
version = "1.1.2+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "81f3d15e84cbcd896376e6730314d59fb5a87f31e4b038454184435cd57defee"
dependencies = [
"indexmap 2.14.0",
"serde_core",
"serde_spanned",
"toml_datetime 1.1.1+spec-1.1.0",
"toml_parser",
"toml_writer",
"winnow 1.0.2",
]
[[package]]
@@ -3737,32 +3803,41 @@ dependencies = [
"serde_core",
]
[[package]]
name = "toml_datetime"
version = "1.1.1+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3165f65f62e28e0115a00b2ebdd37eb6f3b641855f9d636d3cd4103767159ad7"
dependencies = [
"serde_core",
]
[[package]]
name = "toml_edit"
version = "0.23.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5d7cbc3b4b49633d57a0509303158ca50de80ae32c265093b24c414705807832"
dependencies = [
"indexmap 2.12.1",
"toml_datetime",
"indexmap 2.14.0",
"toml_datetime 0.7.3",
"toml_parser",
"winnow",
"winnow 0.7.14",
]
[[package]]
name = "toml_parser"
version = "1.0.4"
version = "1.1.2+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c0cbe268d35bdb4bb5a56a2de88d0ad0eb70af5384a99d648cd4b3d04039800e"
checksum = "a2abe9b86193656635d2411dc43050282ca48aa31c2451210f4202550afb7526"
dependencies = [
"winnow",
"winnow 1.0.2",
]
[[package]]
name = "toml_writer"
version = "1.0.4"
version = "1.1.1+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df8b2b54733674ad286d16267dcfc7a71ed5c776e4ac7aa3c3e2561f7c637bf2"
checksum = "756daf9b1013ebe47a8776667b466417e2d4c5679d441c26230efd9ef78692db"
[[package]]
name = "tracing"
@@ -3784,7 +3859,7 @@ checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -4017,7 +4092,7 @@ checksum = "3b5bb2756c16fb66f80cfbf5fb0e0c09a7001e739f453c9ec241b9c8b1556fda"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -4092,7 +4167,7 @@ checksum = "8c44ce98e7227a04eeb4cf9c784109a5c9710e54849ceb4f09f8597247897f1e"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
"unzip-n",
]
@@ -4186,7 +4261,7 @@ dependencies = [
"bumpalo",
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
"wasm-bindgen-shared",
]
@@ -4216,7 +4291,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bb0e353e6a2fbdc176932bbaab493762eb1255a7900fe0fea1a2f96c296cc909"
dependencies = [
"anyhow",
"indexmap 2.12.1",
"indexmap 2.14.0",
"wasm-encoder",
"wasmparser",
]
@@ -4229,7 +4304,7 @@ checksum = "47b807c72e1bac69382b3a6fb3dbe8ea4c0ed87ff5629b8685ae6b9a611028fe"
dependencies = [
"bitflags",
"hashbrown 0.15.5",
"indexmap 2.12.1",
"indexmap 2.14.0",
"semver",
]
@@ -4345,7 +4420,7 @@ checksum = "053e2e040ab57b9dc951b72c264860db7eb3b0200ba345b4e4c3b14f67855ddf"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -4356,7 +4431,7 @@ checksum = "3f316c4a2570ba26bbec722032c4099d8c8bc095efccdc15688708623367e358"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -4633,6 +4708,12 @@ dependencies = [
"memchr",
]
[[package]]
name = "winnow"
version = "1.0.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2ee1708bef14716a11bae175f579062d4554d95be2c6829f518df847b7b3fdd0"
[[package]]
name = "wit-bindgen"
version = "0.51.0"
@@ -4667,9 +4748,9 @@ checksum = "b7c566e0f4b284dd6561c786d9cb0142da491f46a9fbed79ea69cdad5db17f21"
dependencies = [
"anyhow",
"heck",
"indexmap 2.12.1",
"indexmap 2.14.0",
"prettyplease",
"syn 2.0.111",
"syn 2.0.117",
"wasm-metadata",
"wit-bindgen-core",
"wit-component",
@@ -4685,7 +4766,7 @@ dependencies = [
"prettyplease",
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
"wit-bindgen-core",
"wit-bindgen-rust",
]
@@ -4698,7 +4779,7 @@ checksum = "9d66ea20e9553b30172b5e831994e35fbde2d165325bec84fc43dbf6f4eb9cb2"
dependencies = [
"anyhow",
"bitflags",
"indexmap 2.12.1",
"indexmap 2.14.0",
"log",
"serde",
"serde_derive",
@@ -4717,7 +4798,7 @@ checksum = "ecc8ac4bc1dc3381b7f59c34f00b67e18f910c2c0f50015669dde7def656a736"
dependencies = [
"anyhow",
"id-arena",
"indexmap 2.12.1",
"indexmap 2.14.0",
"log",
"semver",
"serde",
@@ -4785,7 +4866,7 @@ checksum = "b659052874eb698efe5b9e8cf382204678a0086ebf46982b79d6ca3182927e5d"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
"synstructure",
]
@@ -4884,7 +4965,7 @@ dependencies = [
"serde_json",
"serde_with",
"serde_yaml",
"toml",
"toml 0.9.8",
"tracing",
"uhlc",
"validated_struct",
@@ -5147,7 +5228,7 @@ checksum = "9310b02a8f6dc4bd04d9ce6b318b9d00182aeeeeca60410003307d63a2569a3f"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
"zenoh-keyexpr",
]
@@ -5363,7 +5444,7 @@ checksum = "d8a8d209fdf45cf5138cbb5a506f6b52522a25afccc534d1475dad8e31105c6a"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
@@ -5383,7 +5464,7 @@ checksum = "d71e5d6e06ab090c67b5e44993ec16b72dcbaabc526db883a360057678b48502"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
"synstructure",
]
@@ -5393,26 +5474,6 @@ version = "1.8.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0"
[[package]]
name = "zerompk"
version = "0.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b4bfaccde2671513aa564585b2e9ccf0a0ccaf9689477abfd60e53d3806afa51"
dependencies = [
"zerompk_derive",
]
[[package]]
name = "zerompk_derive"
version = "0.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2197ce2b0c4f95ebfd8b09134ecfe0dcc152481b21bacd7a4b1976df1ccfdfc8"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
]
[[package]]
name = "zerotrie"
version = "0.2.3"
@@ -5443,7 +5504,7 @@ checksum = "eadce39539ca5cb3985590102671f2567e659fca9666581ad3411d59207951f3"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.111",
"syn 2.0.117",
]
[[package]]
+18 -3
View File
@@ -1,6 +1,6 @@
[workspace]
resolver = "3"
members = ["rust/exo_pyo3_bindings", "rust/networking"]
members = ["rust/exo_net", "rust/networking"]
[workspace.package]
version = "0.0.1"
@@ -21,10 +21,22 @@ opt-level = 3
## Crate members as common dependencies
networking = { path = "rust/networking" }
# Macro dependecies
# pyo3
pyo3 = "0.27.2"
pyo3-async-runtimes = "0.27.0"
pyo3-log = "0.13.2"
pyo3-stub-gen = "0.22.2"
# util
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"
pidfile-rs = "0.3.1"
# Tracing/logging
log = "0.4"
@@ -32,7 +44,10 @@ env_logger = "0.11.10"
# networking
zenoh = "=1.9.0"
zerompk = "0.4.2"
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 -3
View File
@@ -15,9 +15,8 @@ from pathlib import Path
from typing import Any, Literal
import httpx
from harness import (
ExoClient,
ExoHttpError,
from exo_tools.client import ExoClient, ExoHttpError
from exo_tools.harness import (
add_common_instance_args,
capture_cluster_snapshot,
instance_id_from_instance,
+2 -3
View File
@@ -30,9 +30,8 @@ from pathlib import Path
from statistics import mean
from typing import Any
from harness import (
ExoClient,
ExoHttpError,
from exo_tools.client import ExoClient, ExoHttpError
from exo_tools.harness import (
add_common_instance_args,
capture_cluster_snapshot,
find_existing_instance,
+2 -3
View File
@@ -42,9 +42,8 @@ from pathlib import Path
from typing import Any
import httpx
from harness import (
ExoClient,
ExoHttpError,
from exo_tools.client import ExoClient, ExoHttpError
from exo_tools.harness import (
add_common_instance_args,
capture_cluster_snapshot,
find_existing_instance,
+2 -3
View File
@@ -35,9 +35,8 @@ from exo_bench import (
load_tokenizer_for_bench,
parse_int_list,
)
from harness import (
ExoClient,
ExoHttpError,
from exo_tools.client import ExoClient, ExoHttpError
from exo_tools.harness import (
add_common_instance_args,
instance_id_from_instance,
node_ids_from_instance,
+1 -1
View File
@@ -110,7 +110,7 @@
nixpkgs-fmt.enable = true;
ruff-format = {
enable = true;
excludes = [ "rust/exo_pyo3_bindings/exo_pyo3_bindings.pyi" ];
excludes = [ "rust/exo_net/exo_net.pyi" ];
};
rustfmt = {
enable = true;
+1 -1
View File
@@ -23,7 +23,7 @@ sync-clean:
rust-rebuild:
PYO3_PYTHON="$(uv run python -c 'import sys; print(sys.executable)')" cargo run --bin stub_gen
uv sync --reinstall-package exo_pyo3_bindings
uv sync --reinstall-package exo_net
build-dashboard:
#!/usr/bin/env bash
+14 -10
View File
@@ -15,7 +15,7 @@ dependencies = [
"huggingface-hub>=1.8.0",
"psutil>=7.0.0",
"loguru>=0.7.3",
"exo-pyo3-bindings", # rust bindings
"exo-net", # rust bindings
"anyio==4.11.0",
"mlx==0.31.2; sys_platform == 'darwin'",
"mlx-lm; sys_platform=='darwin'",
@@ -40,6 +40,7 @@ exo = "exo.main:main"
dev = [
"basedpyright>=1.29.0",
"pyinstaller>=6.17.0",
"playwright>=1.52.0",
"pytest>=8.4.0",
"pytest-asyncio>=1.0.0",
"pytest-env",
@@ -75,10 +76,10 @@ cuda13 = [
###
[tool.uv.workspace]
members = ["rust/exo_pyo3_bindings", "bench"]
members = ["rust/exo_net", "bench", "tools"]
[tool.uv.sources]
exo-pyo3-bindings = { workspace = true }
exo-net = { workspace = true }
mlx = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git", branch = "address-rdma-gpu-locks", marker = "sys_platform == 'darwin'" }
mlx-lm = { git = "https://github.com/rltakashige/mlx-lm", branch = "leo/deepseek-v4" }
torch = [
@@ -112,7 +113,7 @@ build-backend = "uv_build"
###
[tool.basedpyright]
include = ["src", "bench"]
include = ["src", "bench", "tools"]
typeCheckingMode = "strict"
failOnWarnings = true
@@ -146,6 +147,13 @@ reportMissingModuleSource = false
[[tool.basedpyright.executionEnvironments]]
root = "src"
[[tool.basedpyright.executionEnvironments]]
root = "bench"
extraPaths = ["tools/src"]
[[tool.basedpyright.executionEnvironments]]
root = "tools/src"
###
# uv configuration
@@ -206,11 +214,7 @@ torchaudio = ["torch"]
###
[tool.ruff]
extend-exclude = [
"*mlx_typings/**",
"rust/exo_pyo3_bindings/**",
"bench/vendor/**",
]
extend-exclude = ["*mlx_typings/**", "rust/exo_net/**", "bench/vendor/**"]
[tool.ruff.lint]
extend-select = ["I", "N", "B", "A", "PIE", "SIM"]
@@ -220,5 +224,5 @@ pythonpath = "."
asyncio_mode = "auto"
markers = ["slow: marks tests as slow (deselected by default)"]
env = ["EXO_TESTS=1"]
addopts = "-m 'not slow' --ignore=tests/start_distributed_test.py"
addopts = "-m 'not slow' --ignore=tests"
filterwarnings = ["ignore:builtin type Swig:DeprecationWarning"]
+6 -5
View File
@@ -35,17 +35,18 @@ let
# Replace workspace exo_pyo3_bindings with Nix-built wheel.
# Preserve passthru so mkVirtualEnv can resolve dependency groups.
# Copy .pyi stub + py.typed marker so basedpyright can find the types.
exo-pyo3-bindings = pkgs.stdenv.mkDerivation {
pname = "exo-pyo3-bindings";
exo-net = pkgs.stdenv.mkDerivation {
pname = "exo-net";
version = "0.1.0";
src = self'.packages.exo_pyo3_bindings;
src = self'.packages.exo-net;
# Install from pre-built wheel
nativeBuildInputs = [ final.pyprojectWheelHook ];
dontStrip = true;
passthru = prev.exo-pyo3-bindings.passthru or { };
postInstall = ''
local siteDir=$out/${final.python.sitePackages}/exo_pyo3_bindings
cp ${inputs.self}/rust/exo_pyo3_bindings/exo_pyo3_bindings.pyi $siteDir/
local siteDir=$out/${final.python.sitePackages}/exo_net
cp ${inputs.self}/rust/exo_net/exo_net.pyi $siteDir/
touch $siteDir/py.typed
'';
};
+52
View File
@@ -0,0 +1,52 @@
[package]
name = "exo_net"
version = { workspace = true }
edition = { workspace = true }
publish = false
[lib]
doctest = false
path = "src/lib.rs"
name = "exo_net"
# "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 = { workspace = true, features = ["experimental-async"] }
pyo3-stub-gen.workspace = true
pyo3-async-runtimes = { workspace = true, features = [
"attributes",
"tokio-runtime",
"testing",
] }
pyo3-log.workspace = true
# async runtime
tokio = { workspace = true, features = ["full"] }
futures-lite.workspace = true
pin-project.workspace = true
# Tracing
log.workspace = true
env_logger.workspace = true
# Networking
zenoh.workspace = true
rand.workspace = true
serde_json.workspace = true
parking_lot.workspace = true
pidfile-rs.workspace = true
File renamed without changes.
+122
View File
@@ -0,0 +1,122 @@
# This file is automatically generated by pyo3_stub_gen
# ruff: noqa: E501, F401, F403, F405
import builtins
import collections.abc
import os
import pathlib
import typing
__all__ = [
"NetReceiver",
"NetSender",
"NetworkingHandle",
"Pidfile",
"PidfileError",
"PyFromSwarm",
"PySession",
"StateProxy",
]
@typing.final
class NetReceiver:
def recv(self) -> collections.abc.Awaitable[bytes | None]: ...
@typing.final
class NetSender:
def send(self, data: bytes) -> collections.abc.Awaitable[bool]: ...
@typing.final
class NetworkingHandle:
@staticmethod
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: ...
@typing.final
class Pidfile:
r"""
A PID file protected with a lock.
An instance of `Pidfile` can be used to manage a PID file: create it,
lock it, detect already running daemons. It is backed by [`pidfile`][]
functions of `libbsd`/`libutil` which use `flopen` to lock the PID
file.
When a PID file is created, the process ID of the current process is
*not* written there, making it possible to lock the PID file before
forking and only write the ID of the forked process when it is ready.
The PID file is deleted automatically when the `Pidfile` comes out of
the scope. To close the PID file without deleting it, for example, in
the parent process of a forked daemon, call `close()`.
[`exit`]: https://doc.rust-lang.org/std/process/fn.exit.html
[`pidfile`]: https://linux.die.net/man/3/pidfile
[`daemon`(3)]: https://linux.die.net/man/3/daemon
"""
def __new__(cls, path: builtins.str | os.PathLike | pathlib.Path, mode: builtins.int) -> Pidfile:
r"""
Creates a new PID file and locks it.
If the PID file cannot be locked, returns `PidfileError::AlreadyRunning` with
a PID of the already running process, or `None` if no PID has been written to
the PID file yet.
"""
def write(self) -> None:
r"""
Writes the current process ID to the PID file.
The file is truncated before writing.
"""
@typing.final
class PidfileError(builtins.Exception):
def __repr__(self) -> builtins.str: ...
def __str__(self) -> builtins.str: ...
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: ...
...
@typing.final
class PySession:
def net_receiver(self, key: builtins.str) -> NetReceiver: ...
def net_sender(self, key: builtins.str) -> NetSender: ...
def state_proxy(self) -> StateProxy: ...
@typing.final
class StateProxy:
def snapshot(self) -> collections.abc.Awaitable[str]: ...
@@ -3,24 +3,22 @@ requires = ["maturin>=1.0,<2.0"]
build-backend = "maturin"
[project]
name = "exo_pyo3_bindings"
version = "0.2.1"
name = "exo_net"
version = "0.3.0"
description = "Add your description here"
readme = "README.md"
authors = [
{ name = "Andrei Cravtov", email = "the.andrei.cravtov@gmail.com" },
{ name = "Evan Quiney", email = "evanev7@gmail.com" },
{ name = "Andrei Cravtov", email = "the.andrei.cravtov@gmail.com" },
]
requires-python = ">=3.13"
dependencies = []
[dependency-groups]
dev = ["exo_pyo3_bindings", "pytest>=8.4.0", "pytest-asyncio>=1.0.0"]
dev = ["exo-net", "pytest>=8.4.0", "pytest-asyncio>=1.0.0"]
[tool.maturin]
#purelib = true
#python-source = "python"
module-name = "exo_pyo3_bindings"
module-name = "exo_net"
features = ["pyo3/extension-module", "pyo3/experimental-async"]
[tool.pytest.ini_options]
@@ -2,7 +2,7 @@ use pyo3_stub_gen::Result;
fn main() -> Result<()> {
env_logger::Builder::from_env(env_logger::Env::default().filter_or("RUST_LOG", "info")).init();
let stub = exo_pyo3_bindings::stub_info()?;
let stub = exo_net::stub_info()?;
stub.generate()?;
Ok(())
}
File renamed without changes.
@@ -5,21 +5,23 @@
//!
mod allow_threading;
mod ident;
mod pidfile;
// mod ident;
mod networking;
mod point_to_point;
mod session;
mod state;
use crate::ident::PyKeypair;
use crate::networking::networking_submodule;
use crate::pidfile::pidfile_submodule;
use crate::point_to_point::{NetReceiver, NetSender};
use crate::session::PySession;
use crate::state::StateProxy;
use pyo3::prelude::PyModule;
use pyo3::types::PyModuleMethods;
use pyo3::{Bound, PyResult, pyclass, pymodule};
use pyo3::{Bound, PyResult, pymodule};
use pyo3_stub_gen::define_stub_info_gatherer;
/// Namespace for all the constants used by this crate.
pub(crate) mod r#const {
pub const MPSC_CHANNEL_SIZE: usize = 1024;
}
/// Namespace for crate-wide extension traits/methods
pub(crate) mod ext {
use crate::allow_threading::AllowThreads;
@@ -151,7 +153,7 @@ pub(crate) mod ext {
/// A Python module implemented in Rust. The name of this function must match
/// the `lib.name` setting in the `Cargo.toml`, else Python will not be able to
/// import the module.
#[pymodule(name = "exo_pyo3_bindings")]
#[pymodule(name = "exo_net")]
fn main_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
// install logger
pyo3_log::init();
@@ -162,7 +164,13 @@ fn main_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
// TODO: for now this is all NOT a submodule, but figure out how to make the submodule system
// work with maturin, where the types generate correctly, in the right folder, without
// too many importing issues...
m.add_class::<PyKeypair>()?;
pidfile_submodule(m)?;
// m.add_class::<PyKeypair>()?;
// networking_submodule(m)?;
m.add_class::<StateProxy>()?;
m.add_class::<PySession>()?;
m.add_class::<NetReceiver>()?;
m.add_class::<NetSender>()?;
networking_submodule(m)?;
// top-level constructs
@@ -1,15 +1,13 @@
use std::pin::Pin;
use std::sync::Arc;
use crate::r#const::MPSC_CHANNEL_SIZE;
use crate::ext::{ByteArrayExt as _, FutureExt, PyErrExt as _};
use crate::ext::{ResultExt as _, TokioMpscSenderExt as _};
use crate::ident::PyKeypair;
use crate::pyclass;
use crate::session::PySession;
use futures_lite::{Stream, StreamExt as _};
use networking::swarm::{FromSwarm, ToSwarm, create_swarm};
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::{PyModule, PyModuleMethods as _};
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::PyBytes;
use pyo3::{Bound, Py, PyAny, PyErr, PyResult, Python, pymethods};
use pyo3_stub_gen::derive::{
@@ -53,29 +51,42 @@ impl PyNetworkingHandle {
// ---- Lifecycle management methods ----
#[new]
#[pyo3(signature = (identity, bootstrap_peers, listen_port))]
fn py_new(
identity: Bound<'_, PyKeypair>,
#[staticmethod]
fn new<'py>(
identity: Bound<'py, PyBytes>,
bootstrap_peers: Vec<String>,
listen_port: u16,
) -> PyResult<Self> {
) -> PyResult<(PyNetworkingHandle, PySession)> {
// create communication channels
let (to_swarm, from_client) = mpsc::channel(MPSC_CHANNEL_SIZE);
let (to_swarm, from_client) = mpsc::channel(1024);
// get identity
let identity = identity.borrow().0.clone();
let identity = u128::from_le_bytes(
identity
.extract::<'_, Vec<u8>>()?
.try_into()
.map_err(|_| PyValueError::new_err("invalid identity bytes"))?,
);
// create networking swarm (within tokio context!! or it crashes)
let _guard = pyo3_async_runtimes::tokio::get_runtime().enter();
let swarm = create_swarm(identity, from_client, bootstrap_peers, listen_port)
.map(|it| it.into_stream())
let swarm = pyo3_async_runtimes::tokio::get_runtime()
.block_on(create_swarm(
identity,
from_client,
bootstrap_peers,
listen_port,
))
.pyerr()?;
Ok(Self {
swarm: Arc::new(Mutex::new(swarm)),
to_swarm,
})
let session = swarm.session.z.clone();
Ok((
PyNetworkingHandle {
swarm: Arc::new(Mutex::new(swarm.into_stream())),
to_swarm,
},
PySession { session },
))
}
#[gen_stub(skip)]
+87
View File
@@ -0,0 +1,87 @@
use pidfile_rs::{Pidfile, PidfileError};
use pyo3::exceptions::PyException;
use pyo3::prelude::{PyModule, PyModuleMethods};
use pyo3::{Bound, PyErr, PyResult, Python, pyclass, pymethods};
use pyo3_stub_gen::derive::{gen_stub_pyclass, gen_stub_pymethods};
use std::fs::Permissions;
use std::os::unix::prelude::PermissionsExt;
use std::path::PathBuf;
#[gen_stub_pyclass]
#[pyclass(frozen, extends=PyException, name="PidfileError")]
pub struct PyPidfileError(PidfileError);
impl PyPidfileError {
// TODO: I actually like this pattern a LOT more but how to abstract??
fn into_pyerr(self, py: Python) -> PyErr {
match Bound::new(py, self) {
Ok(err) => PyErr::from_value(err.into_any()),
Err(err) => err,
}
}
}
#[gen_stub_pymethods]
#[pymethods]
impl PyPidfileError {
fn __repr__(&self) -> String {
format!("PidfileError(\"{}\")", self.0)
}
fn __str__(&self) -> String {
self.0.to_string()
}
}
/// A PID file protected with a lock.
///
/// An instance of `Pidfile` can be used to manage a PID file: create it,
/// lock it, detect already running daemons. It is backed by [`pidfile`][]
/// functions of `libbsd`/`libutil` which use `flopen` to lock the PID
/// file.
///
/// When a PID file is created, the process ID of the current process is
/// *not* written there, making it possible to lock the PID file before
/// forking and only write the ID of the forked process when it is ready.
///
/// The PID file is deleted automatically when the `Pidfile` comes out of
/// the scope. To close the PID file without deleting it, for example, in
/// the parent process of a forked daemon, call `close()`.
///
/// [`exit`]: https://doc.rust-lang.org/std/process/fn.exit.html
/// [`pidfile`]: https://linux.die.net/man/3/pidfile
/// [`daemon`(3)]: https://linux.die.net/man/3/daemon
#[gen_stub_pyclass]
#[pyclass(name = "Pidfile")]
pub struct PyPidfile(Pidfile);
#[gen_stub_pymethods]
#[pymethods]
impl PyPidfile {
/// Creates a new PID file and locks it.
///
/// If the PID file cannot be locked, returns `PidfileError::AlreadyRunning` with
/// a PID of the already running process, or `None` if no PID has been written to
/// the PID file yet.
#[new]
fn py_new(py: Python, path: PathBuf, mode: u32) -> PyResult<Self> {
Ok(Self(
Pidfile::new(&path, Permissions::from_mode(mode))
.map_err(|e| PyPidfileError(e).into_pyerr(py))?,
))
}
/// Writes the current process ID to the PID file.
///
/// The file is truncated before writing.
fn write<'py>(&mut self, py: Python<'py>) -> PyResult<()> {
self.0.write().map_err(|e| PyPidfileError(e).into_pyerr(py))
}
}
pub fn pidfile_submodule(m: &Bound<PyModule>) -> PyResult<()> {
m.add_class::<PyPidfileError>()?;
m.add_class::<PyPidfile>()?;
Ok(())
}
+108
View File
@@ -0,0 +1,108 @@
use std::sync::Arc;
use pyo3::exceptions::PyConnectionError;
use pyo3::types::PyBytes;
use pyo3::types::PyNone;
use pyo3::{BoundObject, prelude::*};
use pyo3_stub_gen::derive::{gen_stub_pyclass, gen_stub_pymethods};
use zenoh::Result;
use zenoh::{
handlers::FifoChannelHandler,
pubsub::{Publisher, Subscriber},
sample::Sample,
};
use crate::ext::ByteArrayExt;
#[gen_stub_pyclass]
#[pyclass]
pub struct NetReceiver {
pub subscriber: Subscriber<FifoChannelHandler<Sample>>,
}
#[gen_stub_pymethods]
#[pymethods]
impl NetReceiver {
#[gen_stub(override_return_type(
type_repr="collections.abc.Awaitable[bytes | None]",
imports=("collections.abc")
))]
pub fn recv<'py>(&'py self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
pyo3_async_runtimes::tokio::future_into_py(py, {
assert!(
self.subscriber.receiver_count() == 1,
"tried to receive twice on the same receiver"
);
let subscriber = self.subscriber.clone();
async move {
match subscriber.recv_async().await {
Err(_) => {
// stream closed;
Ok(Python::attach(|py| PyNone::get(py).unbind()).into_any())
}
Ok(sample) => Ok(sample.payload().to_bytes().to_vec().pybytes().into_any()),
}
}
})
}
}
#[gen_stub_pyclass]
#[pyclass]
pub struct NetSender {
pub publisher: Arc<Publisher<'static>>,
pub first: bool,
}
#[gen_stub_pymethods]
#[pymethods]
impl NetSender {
#[gen_stub(override_return_type(
type_repr="collections.abc.Awaitable[bool]",
imports=("collections.abc")
))]
pub fn send<'py>(
&'py mut self,
py: Python<'py>,
data: Bound<'py, PyBytes>,
) -> PyResult<Bound<'py, PyAny>> {
let is_first = self.first;
self.first = false;
pyo3_async_runtimes::tokio::future_into_py(py, {
let publisher = Arc::clone(&self.publisher);
// clone the data so py can have it back
let bytes = data.as_bytes().to_vec();
async move {
if is_first {
wait_for_listener(&*publisher)
.await
.map_err(|e| PyConnectionError::new_err(e.to_string()))?;
}
if !publisher
.matching_status()
.await
.map_err(|e| PyConnectionError::new_err(e.to_string()))?
.matching()
{
return Ok(false);
}
publisher
.put(&bytes)
.await
.map_err(|e| PyConnectionError::new_err(e.to_string()))?;
Ok(true)
}
})
}
}
async fn wait_for_listener<'a>(publisher: &Publisher<'a>) -> Result<()> {
let matcher = publisher.matching_listener().await?;
if publisher.matching_status().await?.matching() {
return Ok(());
}
while let Ok(status) = matcher.recv_async().await {
if status.matching() {
break;
}
}
Ok(())
}
+72
View File
@@ -0,0 +1,72 @@
use std::sync::Arc;
use pyo3::{exceptions::PyValueError, prelude::*};
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,
};
#[gen_stub_pyclass]
#[pyclass]
pub struct PySession {
pub session: Session,
}
#[gen_stub_pymethods]
#[pymethods]
impl PySession {
/* for now construct with NetworkingHandle
#[staticmethod]
#[gen_stub(override_return_type(
type_repr="collections.abc.Awaitable[PySession]",
imports=("collections.abc")
))]
pub fn init<'py>(py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
pyo3_async_runtimes::tokio::future_into_py(py, async move {
Ok(Self {
session: networking::open(
networking::cfg(rand::random(), 0).expect("default cfg is valid"),
)
.await
.map_err(|e| PyRuntimeError::new_err(e.to_string()))?,
})
})
}
*/
pub fn net_receiver<'py>(&self, key: String) -> PyResult<NetReceiver> {
Ok(NetReceiver {
subscriber: self
.session
.declare_subscriber(key)
.wait()
// C5: key format error
.map_err(|e| PyValueError::new_err(e.to_string()))?,
})
}
pub fn net_sender<'py>(&self, key: String) -> PyResult<NetSender> {
Ok(NetSender {
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()))?,
),
first: true,
})
}
pub fn state_proxy(&self) -> StateProxy {
StateProxy {
session: self.session.clone(),
}
}
}
+78
View File
@@ -0,0 +1,78 @@
use pyo3::{exceptions::PyValueError, prelude::*};
use pyo3_stub_gen::derive::{gen_stub_pyclass, gen_stub_pymethods};
use serde_json::{Map, Value};
use zenoh::{Result, Session, sample::SampleFields};
#[gen_stub_pyclass]
#[pyclass]
pub struct StateProxy {
pub session: Session,
}
#[gen_stub_pymethods]
#[pymethods]
impl StateProxy {
#[gen_stub(override_return_type(
type_repr="collections.abc.Awaitable[str]",
imports=("collections.abc")
))]
pub fn snapshot<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
pyo3_async_runtimes::tokio::future_into_py(py, {
let session = self.session.clone();
async move {
Self::_snapshot(session)
.await
.map_err(|e| PyValueError::new_err(e.to_string()))
.map(|v| v.to_string())
}
})
}
}
impl StateProxy {
async fn _snapshot(session: Session) -> Result<Value> {
let q = session.get("storage/mem1/**").await?;
let mut v = Value::Object(Map::default());
while let Ok(sample) = q.recv_async().await {
let mut cur_v = &mut v;
let Ok(sample) = sample.into_result() else {
continue;
};
// skip storage/mem1
let SampleFields {
payload, key_expr, ..
} = sample.into();
let mut iter = key_expr.split('/').skip(2).peekable();
loop {
let Some(p) = iter.next() else {
break;
};
if iter.peek().is_none() {
// terminal; write value into json
let existing = cur_v
.as_object_mut()
.expect("path terminated unexpectedly - value stored at some/path and some/path/two")
.insert(p.to_owned(), Value::String(payload.try_to_string()?.to_string()));
if let Some(value) = existing {
assert!(value.is_string())
// could log, but string overwrites are fine
}
} else {
// non-terminal; ensure key exists in v, then replace cur with that object
cur_v = cur_v
.as_object_mut()
.expect("path terminated unexpectedly - value stored at some/path and some/path/two")
.entry(p)
.or_insert(Value::Object(Map::default()));
assert!(
cur_v.is_object(),
"path terminated unexpectedly - value stored at some/path and some/path/two"
)
}
}
}
Ok(v)
}
}
+54
View File
@@ -0,0 +1,54 @@
#[cfg(test)]
mod tests {
use core::mem::drop;
use core::option::Option::Some;
use core::time::Duration;
use tokio;
use tokio::sync::mpsc;
#[tokio::test]
async fn test_drop_channel() {
struct Ping;
let (tx, mut rx) = mpsc::channel::<Ping>(10);
let _ = tokio::spawn(async move {
println!("TASK: entered");
loop {
tokio::select! {
result = rx.recv() => {
match result {
Some(_) => {
println!("TASK: pinged");
}
None => {
println!("TASK: closing channel");
break;
}
}
}
_ = tokio::time::sleep(Duration::from_secs_f32(0.1)) => {
println!("TASK: heartbeat");
}
}
}
println!("TASK: exited");
});
let tx2 = tx.clone();
tokio::time::sleep(Duration::from_secs_f32(0.11)).await;
tx.send(Ping).await.expect("Should not fail");
drop(tx);
tokio::time::sleep(Duration::from_secs_f32(0.11)).await;
tx2.send(Ping).await.expect("Should not fail");
drop(tx2);
tokio::time::sleep(Duration::from_secs_f32(0.11)).await;
}
}
@@ -1,9 +1,11 @@
import asyncio
import pytest
from _pytest.capture import CaptureFixture
from exo_pyo3_bindings import (
Keypair,
NetworkingHandle,
Pidfile,
PyFromSwarm,
)
@@ -23,6 +25,13 @@ async def test_sleep_on_multiple_items() -> None:
await h.gossipsub_publish("topic", b"somehting or other")
def test_pidfile(capsys: CaptureFixture[str]):
with capsys.disabled():
print("\nbefore python")
scoped_lock_file()
print("after python")
async def _await_recv(h: NetworkingHandle):
while True:
event = await h.recv()
@@ -32,5 +41,10 @@ async def _await_recv(h: NetworkingHandle):
case PyFromSwarm.Message() as m:
print(f"PYTHON: message: {m}")
def scoped_lock_file():
a = Pidfile("/tmp/lock.pid", 0o0600)
if __name__ == "__main__":
asyncio.run(test_sleep_on_multiple_items())
-62
View File
@@ -1,62 +0,0 @@
[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"
@@ -1,72 +0,0 @@
# 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: ...
...
+8 -10
View File
@@ -4,19 +4,17 @@ version.workspace = true
edition.workspace = true
[dependencies]
async-stream = "0.3.6"
async-stream.workspace = true
futures-lite.workspace = true
netwatcher = { version = "0.6.0", features = ["tokio"] }
parking_lot = "0.12.5"
netwatcher = { workspace = true, features = ["tokio"] }
parking_lot.workspace = true
tokio = { workspace = true, features = ["full"] }
zenoh = { version = "=1.9.0", features = ["internal", "plugins", "unstable"] }
zenoh-plugin-storage-manager = { version = "=1.9.0", default-features = false }
zenoh-plugin-trait = "=1.9.0"
zerompk = { version = "0.4.2", features = ["derive"] }
rand = "0.10.1"
tracing = "0.1.44"
zenoh = { workspace = true, features = ["internal", "plugins", "unstable"] }
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
+5 -4
View File
@@ -1,5 +1,5 @@
use log::info;
use networking;
use tracing::info;
use zenoh::Result;
#[tokio::main]
@@ -7,16 +7,17 @@ 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).await?;
let session = networking::open(cfg, 52414).await?;
let _tok = session
.z
.liveliness()
.declare_token(format!("nodes/{}/live", session.zid()))
.declare_token(format!("nodes/{}/live", session.z.zid()))
.await?;
let key_expr = "storage/mem1/name";
let payload = "me";
info!("Putting Data ('{key_expr}': '{payload}')...");
session.put(key_expr, payload).await?;
session.z.put(key_expr, payload).await?;
tokio::signal::ctrl_c().await?;
Ok(())
}
+6 -4
View File
@@ -1,18 +1,20 @@
use log::info;
use networking;
use tracing::info;
use zenoh::Result;
#[tokio::main]
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).await?;
let cfg = networking::cfg(rand::random(), 52414)?;
let session = networking::open(cfg, 52414).await?;
let _tok = session
.z
.liveliness()
.declare_token(format!("nodes/{}/live", session.zid()))
.declare_token(format!("nodes/{}/live", session.z.zid()))
.await?;
let _sub = session
.z
.liveliness()
.declare_subscriber("nodes/*/live")
.history(true)
+311
View File
@@ -0,0 +1,311 @@
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;
}
+40 -64
View File
@@ -1,27 +1,26 @@
use std::{
env,
ops::{Deref, DerefMut},
panic,
};
use std::env;
use netwatcher::WatchHandle;
use tokio::{sync::mpsc, task::JoinHandle};
use zenoh::{Result, Session as ZSession, config::WhatAmI, internal::runtime::Runtime};
use tokio::task::JoinHandle;
use zenoh::{Result, Session as ZSession, config::Locator};
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", "\"peer\"")?;
cfg.insert_json5("mode", "\"router\"")?;
cfg.insert_json5("listen/endpoints", &format!("[\"tcp/[::]:{listen_port}\"]"))?;
cfg.insert_json5("scouting/multicast/enabled", "true")?;
cfg.insert_json5("scouting/multicast/enabled", "false")?;
cfg.insert_json5("scouting/multicast/autoconnect", "[]")?;
cfg.insert_json5("scouting/gossip/multihop", "true")?;
cfg.insert_json5("namespace", &format!("{namespace:?}"))?;
@@ -42,75 +41,52 @@ pub fn cfg(identity: u128, listen_port: u16) -> Result<zenoh::Config> {
Ok(cfg)
}
pub async fn open(cfg: zenoh::Config) -> Result<Session> {
pub async fn open(cfg: zenoh::Config, listen_port: u16) -> Result<Session> {
assert!(listen_port != 0, "must used defined listen port");
let mut plugins = PluginsManager::static_plugins_only();
plugins.declare_static_plugin::<StoragesPlugin, _>("storage_manager", true);
let mut runtime = zenoh::internal::runtime::RuntimeBuilder::new(cfg)
.plugins_manager(plugins)
.build()
.await?;
let session = zenoh::session::init(runtime.clone().into()).await?;
let z = zenoh::session::init(runtime.clone().into()).await?;
runtime.start().await?;
let _watch_all_handle = watch_all(runtime).await?;
Ok(Session {
session,
_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 = netwatcher::watch_interfaces_with_callback(move |u| _ = send.send(u))?;
let _async = tokio::task::spawn(async move {
let mut discovery = Discovery::new(z.zid(), listen_port).await?;
let _jh = tokio::task::spawn(async move {
loop {
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(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;
}
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(WatchAllHandle { _sync, _async })
Ok(Session { z, _jh })
}
pub struct Session {
pub session: ZSession,
_watch_all_handle: WatchAllHandle,
pub z: ZSession,
_jh: JoinHandle<()>,
}
impl Deref for Session {
type Target = ZSession;
fn deref(&self) -> &Self::Target {
&self.session
}
}
impl DerefMut for Session {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.session
}
}
impl Drop for WatchAllHandle {
impl Drop for Session {
fn drop(&mut self) {
self._async.abort();
self._jh.abort();
}
}
pub struct WatchAllHandle {
_sync: WatchHandle,
_async: JoinHandle<Result<()>>,
}
+15 -13
View File
@@ -6,7 +6,6 @@ use std::pin::Pin;
use futures_lite::Stream;
use tokio::sync::mpsc;
use tokio::sync::oneshot;
use tracing::info;
use zenoh::Result;
use zenoh::Session;
use zenoh::handlers::FifoChannelHandler;
@@ -14,7 +13,6 @@ use zenoh::liveliness::LivelinessToken;
use zenoh::pubsub::Subscriber;
use zenoh::sample::Sample;
use zenoh::sample::SampleKind;
use zerompk::{FromMessagePack, ToMessagePack};
#[derive(Debug)]
pub enum ToSwarm {
@@ -32,7 +30,7 @@ pub enum ToSwarm {
result_sender: oneshot::Sender<Result<()>>,
},
}
#[derive(Debug, ToMessagePack, FromMessagePack)]
#[derive(Debug)]
pub enum FromSwarm {
Message { topic: String, data: Vec<u8> },
Discovered {},
@@ -41,26 +39,26 @@ pub enum FromSwarm {
pub type Topics = HashMap<String, Subscriber<()>>;
pub struct Swarm {
cfg: zenoh::Config,
pub session: crate::Session,
from_client: mpsc::Receiver<ToSwarm>,
}
impl Swarm {
pub fn into_stream(self) -> Pin<Box<dyn Stream<Item = FromSwarm> + Send>> {
let Swarm {
cfg,
session,
mut from_client,
} = self;
let stream = async_stream::stream! {
let mut session = session;
let (mut to_topics, mut from_topics) = mpsc::channel(1024);
let mut topics = Topics::new();
let Ok(mut session) = crate::open(cfg).await else { return; };
let Ok((_token, discovery)) = register_liveness(&mut session).await else { return; };
let Ok((_token, discovery)) = register_liveness(&mut session.z).await else { return; };
loop {
tokio::select! {
msg = from_client.recv() => {
let Some(msg) = msg else { break };
on_message(&mut session, &mut topics, &mut to_topics, msg).await;
on_message(&mut session.z, &mut topics, &mut to_topics, msg).await;
}
event = from_topics.recv() => {
if let Some(event) = event {
@@ -73,11 +71,11 @@ impl Swarm {
let nid = key_expr.strip_prefix("nodes/").and_then(|s| s.strip_suffix("/live"));
yield match token.kind() {
SampleKind::Put => {
info!("discovered: {nid:?}");
log::info!("discovered: {nid:?}");
FromSwarm::Discovered {}
}
SampleKind::Delete => {
info!("expired: {nid:?}");
log::info!("expired: {nid:?}");
FromSwarm::Expired {}
}
}
@@ -171,7 +169,7 @@ async fn on_message(
}
}
pub fn create_swarm(
pub async fn create_swarm(
identity: u128,
from_client: mpsc::Receiver<ToSwarm>,
bootstrap_peers: Vec<String>,
@@ -181,6 +179,10 @@ pub fn create_swarm(
if !bootstrap_peers.is_empty() || listen_port != 0 {
todo!();
}
let cfg = crate::cfg(identity, listen_port)?;
Ok(Swarm { cfg, from_client })
let cfg = crate::cfg(identity, 52414)?;
let session = crate::open(cfg, 52414).await?;
Ok(Swarm {
session,
from_client,
})
}
+4 -3
View File
@@ -55,6 +55,7 @@
];
OPENSSL_NO_VENDOR = "1";
MATURIN_NO_INSTALL_RUST = "1";
# Required for pyo3 tests to find libpython
LD_LIBRARY_PATH = lib.makeLibraryPath [ pkgs.python313 ];
@@ -81,11 +82,11 @@
config = {
packages = {
# Python bindings wheel via maturin
exo_pyo3_bindings = craneLib.buildPackage (
exo-net = craneLib.buildPackage (
commonArgs
// {
inherit cargoArtifacts;
pname = "exo_pyo3_bindings";
pname = "exo-net";
nativeBuildInputs = commonArgs.nativeBuildInputs ++ [
pkgs.maturin
@@ -95,7 +96,7 @@
maturin build \
--release \
--manylinux off \
--manifest-path rust/exo_pyo3_bindings/Cargo.toml \
--manifest-path rust/exo_net/Cargo.toml \
--features "pyo3/extension-module,pyo3/experimental-async" \
--interpreter ${pkgs.python313}/bin/python \
--out dist
-5
View File
@@ -20,7 +20,6 @@ from exo.shared.types.chunks import (
TokenChunk,
ToolCallChunk,
)
from exo.shared.types.common import CommandId
from exo.shared.types.text_generation import (
Base64Image,
InputMessage,
@@ -181,7 +180,6 @@ def ollama_request_to_text_generation(
async def generate_ollama_chat_stream(
_command_id: CommandId,
chunk_stream: AsyncGenerator[
ErrorChunk | ToolCallChunk | TokenChunk | PrefillProgressChunk, None
],
@@ -264,7 +262,6 @@ async def generate_ollama_chat_stream(
async def collect_ollama_chat_response(
_command_id: CommandId,
chunk_stream: AsyncGenerator[
ErrorChunk | ToolCallChunk | TokenChunk | PrefillProgressChunk, None
],
@@ -369,7 +366,6 @@ def ollama_generate_request_to_text_generation(
async def generate_ollama_generate_stream(
_command_id: CommandId,
chunk_stream: AsyncGenerator[
ErrorChunk | ToolCallChunk | TokenChunk | PrefillProgressChunk, None
],
@@ -442,7 +438,6 @@ async def generate_ollama_generate_stream(
async def collect_ollama_generate_response(
_command_id: CommandId,
chunk_stream: AsyncGenerator[
ErrorChunk | ToolCallChunk | TokenChunk | PrefillProgressChunk, None
],
+278 -298
View File
@@ -5,6 +5,7 @@ import json
import random
import time
from collections.abc import AsyncGenerator, Awaitable, Callable, Iterable
from dataclasses import dataclass, field
from datetime import datetime, timezone
from http import HTTPStatus
from pathlib import Path
@@ -12,7 +13,7 @@ from typing import Annotated, Any, Literal, cast
from uuid import uuid4
import anyio
from anyio import BrokenResourceError, ClosedResourceError
from exo_net import NetSender, PySession
from fastapi import FastAPI, File, Form, HTTPException, Query, Request, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, JSONResponse, StreamingResponse
@@ -21,6 +22,7 @@ from hypercorn.asyncio import serve # pyright: ignore[reportUnknownVariableType
from hypercorn.config import Config
from hypercorn.typing import ASGIFramework
from loguru import logger
from pydantic import TypeAdapter
from exo.api.adapters.chat_completions import (
chat_request_to_text_generation,
@@ -140,10 +142,12 @@ from exo.shared.models.model_cards import (
)
from exo.shared.tracing import TraceEvent, compute_stats, export_trace, load_trace_file
from exo.shared.types.chunks import (
ErrorChunk,
ImageChunk,
Chunk,
ImageGenerationChunk,
InputImageChunk,
PrefillProgressChunk,
StatusChunk,
TextGenerationChunk,
TokenChunk,
ToolCallChunk,
)
@@ -157,7 +161,6 @@ from exo.shared.types.commands import (
DeleteInstance,
DeleteInstanceLink,
DownloadCommand,
ForwarderCommand,
ForwarderDownloadCommand,
ImageEdits,
ImageGeneration,
@@ -171,7 +174,6 @@ from exo.shared.types.commands import (
)
from exo.shared.types.common import CommandId, Id, NodeId, SystemId
from exo.shared.types.events import (
ChunkGenerated,
Event,
IndexedEvent,
InstanceDeleted,
@@ -230,6 +232,110 @@ def _require_disaggregation_enabled() -> None:
)
@dataclass
class Transport:
session: 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()
async def send_command(self, command: Command) -> bool:
while self.paused:
await self.paused_ev.wait()
return await self.command_sender.send(command.model_dump_json().encode("utf-8"))
async def stream_text(
self,
command_id: CommandId,
) -> AsyncGenerator[TextGenerationChunk | StatusChunk]:
async for chunk in self.stream(command_id):
if isinstance(chunk, (TextGenerationChunk | StatusChunk)):
yield chunk
async def stream_images(
self,
command_id: CommandId,
) -> AsyncGenerator[ImageGenerationChunk]:
async for chunk in self.stream(command_id):
if isinstance(chunk, (ImageGenerationChunk)):
yield chunk
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]):
try:
with anyio.CancelScope() as cs:
self.cancel_scopes[command_id] = cs
# recv from any node
receiver = self.session.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(
data, strict=True, extra="forbid"
)
)
)
if (
not isinstance(chunk, StatusChunk)
and chunk.finish_reason is not None
):
break
except (
anyio.get_cancelled_exc_class(),
anyio.BrokenResourceError,
anyio.ClosedResourceError,
):
with anyio.CancelScope(shield=True):
await self.command_sender.send(
TaskCancelled(cancelled_command_id=command_id)
.model_dump_json()
.encode("utf-8")
)
finally:
self.cancel_scopes.pop(command_id, None)
with anyio.CancelScope(shield=True):
await self.command_sender.send(
TaskFinished(finished_command_id=command_id)
.model_dump_json()
.encode("utf-8")
)
def cancel(self, command_id: CommandId) -> bool:
if (cs := self.cancel_scopes.pop(command_id, None)) is not None:
cs.cancel()
return True
return False
class API:
def __init__(
self,
@@ -237,15 +343,14 @@ class API:
*,
port: int,
event_receiver: Receiver[IndexedEvent],
command_sender: Sender[ForwarderCommand],
download_command_sender: Sender[ForwarderDownloadCommand],
# This lets us pause the API if an election is running
election_receiver: Receiver[ElectionMessage],
session: PySession,
) -> None:
self.state = State()
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
self._system_id = SystemId()
self.command_sender = command_sender
self.download_command_sender = download_command_sender
self.event_receiver = event_receiver
self.election_receiver = election_receiver
@@ -254,9 +359,6 @@ class API:
self.port = port
self._sent_image_hashes: set[str] = set()
self.paused: bool = False
self.paused_ev: anyio.Event = anyio.Event()
self.app = FastAPI()
@self.app.middleware("http")
@@ -280,13 +382,7 @@ class API:
name="dashboard",
)
self._text_generation_queues: dict[
CommandId,
Sender[TokenChunk | ErrorChunk | ToolCallChunk | PrefillProgressChunk],
] = {}
self._image_generation_queues: dict[
CommandId, Sender[ImageChunk | ErrorChunk]
] = {}
self.transport = Transport(session)
self._image_store = ImageStore(EXO_IMAGE_CACHE_DIR)
self._tg: TaskGroup = TaskGroup()
@@ -307,9 +403,9 @@ class API:
def unpause(self, result_clock: int):
logger.info("Unpausing API")
self.last_completed_election = result_clock
self.paused = False
self.paused_ev.set()
self.paused_ev = anyio.Event()
self.transport.paused = False
self.transport.paused_ev.set()
self.transport.paused_ev = anyio.Event()
def _setup_exception_handlers(self) -> None:
self.app.exception_handler(HTTPException)(self.http_exception_handler)
@@ -424,7 +520,7 @@ class API:
instance_meta=payload.instance_meta,
min_nodes=payload.min_nodes,
)
await self._send(command)
await self.transport.send_command(command)
return CreateInstanceResponse(
message="Command received.",
@@ -449,7 +545,7 @@ class API:
command = CreateInstance(
instance=instance,
)
await self._send(command)
await self.transport.send_command(command)
return CreateInstanceResponse(
message="Command received.",
@@ -632,7 +728,7 @@ class API:
command = DeleteInstance(
instance_id=instance_id,
)
await self._send(command)
await self.transport.send_command(command)
return DeleteInstanceResponse(
message="Command received.",
command_id=command.command_id,
@@ -667,7 +763,7 @@ class API:
prefill_instances=list(body.prefill_instances),
decode_instances=list(body.decode_instances),
)
await self._send(command)
await self.transport.send_command(command)
return InstanceLinkResponse(
message="Command received.", command_id=command.command_id
)
@@ -677,64 +773,24 @@ class API:
) -> InstanceLinkResponse:
_require_disaggregation_enabled()
command = DeleteInstanceLink(link_id=link_id)
await self._send(command)
await self.transport.send_command(command)
return InstanceLinkResponse(
message="Command received.", command_id=command.command_id
)
async def cancel_command(self, command_id: CommandId) -> CancelCommandResponse:
"""Cancel an active command by closing its stream and notifying workers."""
sender = self._text_generation_queues.get(
command_id
) or self._image_generation_queues.get(command_id)
if sender is None:
if self.transport.cancel(command_id):
return CancelCommandResponse(
message="Command cancelled.",
command_id=command_id,
)
else:
raise HTTPException(
status_code=404,
detail="Command not found or already completed",
)
await self._send(TaskCancelled(cancelled_command_id=command_id))
sender.close()
return CancelCommandResponse(
message="Command cancelled.",
command_id=command_id,
)
async def _token_chunk_stream(
self, command_id: CommandId
) -> AsyncGenerator[
TokenChunk | ErrorChunk | ToolCallChunk | PrefillProgressChunk, None
]:
"""Yield chunks for a given command until completion.
This is the internal low-level stream used by all API adapters.
"""
try:
self._text_generation_queues[command_id], recv = channel[
TokenChunk | ErrorChunk | ToolCallChunk | PrefillProgressChunk
]()
with recv as token_chunks:
async for chunk in token_chunks:
yield chunk
if isinstance(chunk, PrefillProgressChunk):
continue
if chunk.finish_reason is not None:
break
except anyio.get_cancelled_exc_class():
command = TaskCancelled(cancelled_command_id=command_id)
with anyio.CancelScope(shield=True):
await self.command_sender.send(
ForwarderCommand(origin=self._system_id, command=command)
)
raise
finally:
await self._send(TaskFinished(finished_command_id=command_id))
if command_id in self._text_generation_queues:
del self._text_generation_queues[command_id]
async def _collect_text_generation_with_stats(
self, command_id: CommandId
) -> BenchChatCompletionResponse:
@@ -749,7 +805,7 @@ class API:
async with anyio.create_task_group() as tg:
tg.start_soon(sampler.run)
async for chunk in self._token_chunk_stream(command_id):
async for chunk in self.transport.stream_text(command_id):
if isinstance(chunk, PrefillProgressChunk):
continue
@@ -816,7 +872,7 @@ class API:
images = task_params.images
if not images:
command = TextGeneration(task_params=task_params)
await self._send(command)
await self.transport.send_command(command)
return command
hashes = [hashlib.sha256(img.encode("ascii")).hexdigest() for img in images]
@@ -833,7 +889,7 @@ class API:
new_images.append((idx, img))
if not new_images:
await self._send(command)
await self.transport.send_command(command)
return command
all_chunks: list[tuple[int, str]] = []
@@ -842,7 +898,7 @@ class API:
all_chunks.append((img_idx, img_data[i : i + EXO_MAX_CHUNK_SIZE]))
for global_idx, (img_idx, chunk_data) in enumerate(all_chunks):
await self._send(
await self.transport.send_command(
SendInputChunk(
chunk=InputImageChunk(
model=task_params.model,
@@ -855,7 +911,7 @@ class API:
)
)
await self._send(command)
await self.transport.send_command(command)
return command
async def chat_completions(
@@ -875,7 +931,7 @@ class API:
with_sse_keepalive(
generate_chat_stream(
command.command_id,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
),
media_type="text/event-stream",
@@ -889,7 +945,7 @@ class API:
return StreamingResponse(
collect_chat_response(
command.command_id,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
media_type="application/json",
)
@@ -918,7 +974,7 @@ class API:
with_sse_keepalive(
generate_chat_stream(
command.command_id,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
),
media_type="text/event-stream",
@@ -1024,7 +1080,7 @@ class API:
command = ImageGeneration(
task_params=payload,
)
await self._send(command)
await self.transport.send_command(command)
# Check if streaming is requested
if payload.stream and payload.partial_images and payload.partial_images > 0:
@@ -1060,105 +1116,85 @@ class API:
image_metadata: dict[tuple[int, bool], tuple[int | None, int | None]] = {}
images_complete = 0
try:
self._image_generation_queues[command_id], recv = channel[
ImageChunk | ErrorChunk
]()
with recv as chunks:
async for chunk in chunks:
if chunk.finish_reason == "error":
error_response = ErrorResponse(
error=ErrorInfo(
message=chunk.error_message or "Internal server error",
type="InternalServerError",
code=500,
)
)
yield f"data: {error_response.model_dump_json()}\n\n"
yield "data: [DONE]\n\n"
return
key = (chunk.image_index, chunk.is_partial)
if key not in image_chunks:
image_chunks[key] = {}
image_total_chunks[key] = chunk.total_chunks
image_metadata[key] = (
chunk.partial_index,
chunk.total_partials,
)
image_chunks[key][chunk.chunk_index] = chunk.data
# Check if this image is complete
if len(image_chunks[key]) == image_total_chunks[key]:
full_data = "".join(
image_chunks[key][i] for i in range(len(image_chunks[key]))
)
partial_idx, total_partials = image_metadata[key]
if chunk.is_partial:
# Yield partial image event (always use b64_json for partials)
event_data = {
"type": "partial",
"image_index": chunk.image_index,
"partial_index": partial_idx,
"total_partials": total_partials,
"format": str(chunk.format),
"data": {
"b64_json": full_data
if response_format == "b64_json"
else None,
},
}
yield f"data: {json.dumps(event_data)}\n\n"
else:
# Final image
if response_format == "url":
image_bytes = base64.b64decode(full_data)
content_type = _format_to_content_type(chunk.format)
stored = self._image_store.store(
image_bytes, content_type
)
url = self._build_image_url(request, stored.image_id)
event_data = {
"type": "final",
"image_index": chunk.image_index,
"format": str(chunk.format),
"data": {"url": url},
}
else:
event_data = {
"type": "final",
"image_index": chunk.image_index,
"format": str(chunk.format),
"data": {"b64_json": full_data},
}
yield f"data: {json.dumps(event_data)}\n\n"
images_complete += 1
if images_complete >= num_images:
yield "data: [DONE]\n\n"
break
# Clean up completed image chunks
del image_chunks[key]
del image_total_chunks[key]
del image_metadata[key]
except anyio.get_cancelled_exc_class():
command = TaskCancelled(cancelled_command_id=command_id)
with anyio.CancelScope(shield=True):
await self.command_sender.send(
ForwarderCommand(origin=self._system_id, command=command)
async for chunk in self.transport.stream_images(command_id):
if chunk.finish_reason == "error":
error_response = ErrorResponse(
error=ErrorInfo(
message=chunk.error_message or "Internal server error",
type="InternalServerError",
code=500,
)
)
raise
finally:
await self._send(TaskFinished(finished_command_id=command_id))
if command_id in self._image_generation_queues:
del self._image_generation_queues[command_id]
yield f"data: {error_response.model_dump_json()}\n\n"
yield "data: [DONE]\n\n"
return
key = (chunk.image_index, chunk.is_partial)
if key not in image_chunks:
image_chunks[key] = {}
image_total_chunks[key] = chunk.total_chunks
image_metadata[key] = (
chunk.partial_index,
chunk.total_partials,
)
image_chunks[key][chunk.chunk_index] = chunk.data
# Check if this image is complete
if len(image_chunks[key]) == image_total_chunks[key]:
full_data = "".join(
image_chunks[key][i] for i in range(len(image_chunks[key]))
)
partial_idx, total_partials = image_metadata[key]
if chunk.is_partial:
# Yield partial image event (always use b64_json for partials)
event_data = {
"type": "partial",
"image_index": chunk.image_index,
"partial_index": partial_idx,
"total_partials": total_partials,
"format": str(chunk.format),
"data": {
"b64_json": full_data
if response_format == "b64_json"
else None,
},
}
yield f"data: {json.dumps(event_data)}\n\n"
else:
# Final image
if response_format == "url":
image_bytes = base64.b64decode(full_data)
content_type = _format_to_content_type(chunk.format)
stored = self._image_store.store(image_bytes, content_type)
url = self._build_image_url(request, stored.image_id)
event_data = {
"type": "final",
"image_index": chunk.image_index,
"format": str(chunk.format),
"data": {"url": url},
}
else:
event_data = {
"type": "final",
"image_index": chunk.image_index,
"format": str(chunk.format),
"data": {"b64_json": full_data},
}
yield f"data: {json.dumps(event_data)}\n\n"
images_complete += 1
if images_complete >= num_images:
yield "data: [DONE]\n\n"
break
# Clean up completed image chunks
del image_chunks[key]
del image_total_chunks[key]
del image_metadata[key]
async def _collect_image_chunks(
self,
@@ -1177,74 +1213,55 @@ class API:
images_complete = 0
stats: ImageGenerationStats | None = None
try:
self._image_generation_queues[command_id], recv = channel[
ImageChunk | ErrorChunk
]()
while images_complete < num_images:
with recv as chunks:
async for chunk in chunks:
if chunk.finish_reason == "error":
raise HTTPException(
status_code=500,
detail=chunk.error_message or "Internal server error",
)
if chunk.is_partial:
continue
if chunk.image_index not in image_chunks:
image_chunks[chunk.image_index] = {}
image_total_chunks[chunk.image_index] = chunk.total_chunks
image_formats[chunk.image_index] = chunk.format
image_chunks[chunk.image_index][chunk.chunk_index] = chunk.data
if capture_stats and chunk.stats is not None:
stats = chunk.stats
if (
len(image_chunks[chunk.image_index])
== image_total_chunks[chunk.image_index]
):
images_complete += 1
if images_complete >= num_images:
break
images: list[ImageData] = []
for image_idx in range(num_images):
chunks_dict = image_chunks[image_idx]
full_data = "".join(chunks_dict[i] for i in range(len(chunks_dict)))
if response_format == "url" and request is not None:
image_bytes = base64.b64decode(full_data)
content_type = _format_to_content_type(image_formats.get(image_idx))
stored = self._image_store.store(image_bytes, content_type)
url = self._build_image_url(request, stored.image_id)
images.append(ImageData(b64_json=None, url=url))
else:
images.append(
ImageData(
b64_json=full_data
if response_format == "b64_json"
else None,
url=None,
)
while images_complete < num_images:
async for chunk in self.transport.stream_images(command_id):
if chunk.finish_reason == "error":
raise HTTPException(
status_code=500,
detail=chunk.error_message or "Internal server error",
)
return (images, stats if capture_stats else None)
except anyio.get_cancelled_exc_class():
command = TaskCancelled(cancelled_command_id=command_id)
with anyio.CancelScope(shield=True):
await self.command_sender.send(
ForwarderCommand(origin=self._system_id, command=command)
if chunk.is_partial:
continue
if chunk.image_index not in image_chunks:
image_chunks[chunk.image_index] = {}
image_total_chunks[chunk.image_index] = chunk.total_chunks
image_formats[chunk.image_index] = chunk.format
image_chunks[chunk.image_index][chunk.chunk_index] = chunk.data
if capture_stats and chunk.stats is not None:
stats = chunk.stats
if (
len(image_chunks[chunk.image_index])
== image_total_chunks[chunk.image_index]
):
images_complete += 1
if images_complete >= num_images:
break
images: list[ImageData] = []
for image_idx in range(num_images):
chunks_dict = image_chunks[image_idx]
full_data = "".join(chunks_dict[i] for i in range(len(chunks_dict)))
if response_format == "url" and request is not None:
image_bytes = base64.b64decode(full_data)
content_type = _format_to_content_type(image_formats.get(image_idx))
stored = self._image_store.store(image_bytes, content_type)
url = self._build_image_url(request, stored.image_id)
images.append(ImageData(b64_json=None, url=url))
else:
images.append(
ImageData(
b64_json=full_data if response_format == "b64_json" else None,
url=None,
)
)
raise
finally:
await self._send(TaskFinished(finished_command_id=command_id))
if command_id in self._image_generation_queues:
del self._image_generation_queues[command_id]
return (images, stats if capture_stats else None)
async def _collect_image_generation(
self,
@@ -1294,7 +1311,7 @@ class API:
command = ImageGeneration(
task_params=payload,
)
await self._send(command)
await self.transport.send_command(command)
return await self._collect_image_generation_with_stats(
request=request,
@@ -1357,7 +1374,7 @@ class API:
f"Sending input image: {len(image_data)} bytes in {total_chunks} chunks"
)
for chunk_index, chunk_data in enumerate(data_chunks):
await self._send(
await self.transport.send_command(
SendInputChunk(
chunk=InputImageChunk(
model=resolved_model,
@@ -1369,7 +1386,7 @@ class API:
)
)
await self._send(command)
await self.transport.send_command(command)
return command
async def image_edits(
@@ -1497,7 +1514,7 @@ class API:
generate_claude_stream(
command.command_id,
payload.model,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
),
media_type="text/event-stream",
@@ -1512,7 +1529,7 @@ class API:
collect_claude_response(
command.command_id,
payload.model,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
media_type="application/json",
)
@@ -1533,7 +1550,7 @@ class API:
generate_responses_stream(
command.command_id,
payload.model,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
),
media_type="text/event-stream",
@@ -1549,7 +1566,7 @@ class API:
collect_responses_response(
command.command_id,
payload.model,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
media_type="application/json",
)
@@ -1575,8 +1592,7 @@ class API:
if payload.stream:
return StreamingResponse(
generate_ollama_chat_stream(
command.command_id,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
media_type="application/x-ndjson",
headers={
@@ -1588,8 +1604,7 @@ class API:
else:
return StreamingResponse(
collect_ollama_chat_response(
command.command_id,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
media_type="application/json",
)
@@ -1611,8 +1626,7 @@ class API:
if payload.stream:
return StreamingResponse(
generate_ollama_generate_stream(
command.command_id,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
media_type="application/x-ndjson",
headers={
@@ -1624,8 +1638,7 @@ class API:
else:
return StreamingResponse(
collect_ollama_generate_response(
command.command_id,
self._token_chunk_stream(command.command_id),
self.transport.stream_text(command.command_id),
),
media_type="application/json",
)
@@ -1761,11 +1774,8 @@ class API:
status_code=400, detail=f"Failed to fetch model: {exc}"
) from exc
await self.command_sender.send(
ForwarderCommand(
origin=self._system_id,
command=AddCustomModelCard(model_card=card),
)
await self.transport.command_sender.send(
AddCustomModelCard(model_card=card).model_dump_json().encode("utf-8")
)
# Immediately update the local cache so the subsequent GET /models
@@ -1790,11 +1800,8 @@ class API:
if card is None or not card.is_custom:
raise HTTPException(status_code=404, detail="Custom model card not found")
await self.command_sender.send(
ForwarderCommand(
origin=self._system_id,
command=DeleteCustomModelCard(model_id=model_id),
)
await self.transport.command_sender.send(
DeleteCustomModelCard(model_id=model_id).model_dump_json().encode("utf-8")
)
return JSONResponse(
@@ -1847,6 +1854,7 @@ 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)
@@ -1859,7 +1867,6 @@ class API:
shutdown_ev.set()
finally:
self._event_log.close()
self.command_sender.close()
self.event_receiver.close()
async def run_api(self, ev: anyio.Event):
@@ -1883,23 +1890,6 @@ class API:
self.state = apply(self.state, i_event)
event = i_event.event
if isinstance(event, ChunkGenerated):
if queue := self._image_generation_queues.get(
event.command_id, None
):
assert isinstance(event.chunk, ImageChunk)
try:
await queue.send(event.chunk)
except (BrokenResourceError, ClosedResourceError):
self._image_generation_queues.pop(event.command_id, None)
if queue := self._text_generation_queues.get(
event.command_id, None
):
assert not isinstance(event.chunk, ImageChunk)
try:
await queue.send(event.chunk)
except (BrokenResourceError, ClosedResourceError):
self._text_generation_queues.pop(event.command_id, None)
if isinstance(event, InstanceDeleted):
self._close_streams_for_instance(event.instance_id)
if isinstance(event, TracesMerged):
@@ -1914,10 +1904,7 @@ class API:
task, (TextGenerationTask, ImageGenerationTask, ImageEditsTask)
):
continue
if sender := self._text_generation_queues.pop(task.command_id, None):
sender.close()
if sender := self._image_generation_queues.pop(task.command_id, None):
sender.close()
self.transport.cancel(task.command_id)
def _save_merged_trace(self, event: TracesMerged) -> None:
traces = [
@@ -1938,7 +1925,7 @@ class API:
with self.election_receiver as ems:
async for message in ems:
if message.clock > self.last_completed_election:
self.paused = True
self.transport.paused = True
async def _cleanup_expired_images(self):
"""Periodically clean up expired images from the store."""
@@ -1949,13 +1936,6 @@ class API:
if removed > 0:
logger.debug(f"Cleaned up {removed} expired images")
async def _send(self, command: Command):
while self.paused:
await self.paused_ev.wait()
await self.command_sender.send(
ForwarderCommand(origin=self._system_id, command=command)
)
async def _send_download(self, command: DownloadCommand):
await self.download_command_sender.send(
ForwarderDownloadCommand(origin=self._system_id, command=command)
+10 -14
View File
@@ -1,11 +1,11 @@
# pyright: reportUnusedFunction=false, reportAny=false
from typing import Any
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock
from fastapi import FastAPI
from fastapi.testclient import TestClient
from exo.api.main import API
from exo.api.main import API, Transport
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._text_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api._image_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api._send = AsyncMock() # pyright: ignore[reportPrivateUsage]
api.transport = object.__new__(Transport)
api.transport.cancel = AsyncMock()
api.transport.send_command = AsyncMock()
api._setup_exception_handlers() # pyright: ignore[reportPrivateUsage]
app.post("/v1/cancel/{command_id}")(api.cancel_command)
return api
@@ -43,16 +43,14 @@ 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)
sender.close.assert_called_once()
api._send.assert_called_once()
api.transport.cancel.assert_called_once()
api.transport.send_command.assert_called_once()
task_cancelled = api._send.call_args[0][0]
assert task_cancelled.cancelled_command_id == cid
@@ -63,15 +61,13 @@ 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)
sender.close.assert_called_once()
api._send.assert_called_once()
task_cancelled = api._send.call_args[0][0]
api.transport.cancel.assert_called_once()
api.transport.send_command.assert_called_once()
task_cancelled = api.transport.send_command.call_args[0][0]
assert task_cancelled.cancelled_command_id == cid
@@ -1,9 +1,10 @@
# 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
from exo.api.main import API, Transport
from exo.api.types import ImageGenerationTaskParams
from exo.shared.types.common import CommandId, ModelId
from exo.shared.types.state import State
@@ -16,12 +17,11 @@ from exo.shared.types.text_generation import (
from exo.shared.types.worker.instances import InstanceId
def _make_api_with_state(state: State) -> API:
def _make_api_with_state(state: State) -> Any:
"""Create a minimal API instance with pre-set state."""
api = object.__new__(API)
api.state = state
api._text_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api._image_generation_queues = {} # pyright: ignore[reportPrivateUsage]
api.transport = object.__new__(Transport)
return api
@@ -47,13 +47,10 @@ def test_close_streams_for_deleted_instance() -> None:
state = State(tasks={task.task_id: task})
api = _make_api_with_state(state)
sender = MagicMock()
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
api._close_streams_for_instance(instance_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]
api.transport.cancel.assert_called_once()
assert api.transport.cancel.call_args[0][0] == command_id
def test_close_streams_ignores_unrelated_instances() -> None:
@@ -72,7 +69,6 @@ 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:
+28 -15
View File
@@ -3,10 +3,13 @@ import multiprocessing as mp
import os
import resource
import signal
import sys
from dataclasses import dataclass, field
from typing import Self
from uuid import uuid4
import anyio
from exo_net import Pidfile, PidfileError, PySession
from loguru import logger
from pydantic import PositiveInt
@@ -16,8 +19,8 @@ from exo.download.coordinator import DownloadCoordinator
from exo.download.impl_shard_downloader import exo_shard_downloader
from exo.master.main import Master
from exo.routing.event_router import EventRouter
from exo.routing.router import Router, get_node_id_keypair
from exo.shared.constants import EXO_LOG
from exo.routing.router import Router
from exo.shared.constants import EXO_LOG, EXO_PID_FILE
from exo.shared.election import Election, ElectionResult
from exo.shared.logging import logger_cleanup, logger_setup
from exo.shared.types.common import NodeId, SessionId
@@ -39,31 +42,31 @@ class Node:
api: API | None
node_id: NodeId
session: PySession
offline: bool
_api_port: int
_tg: TaskGroup = field(init=False, default_factory=TaskGroup)
@classmethod
async def create(cls, args: "Args") -> Self:
keypair = get_node_id_keypair()
node_id = NodeId(keypair.to_node_id())
node_id_bytes = uuid4()
node_id = NodeId(str(node_id_bytes))
session_id = SessionId(master_node_id=node_id, election_clock=0)
router = Router.create(
keypair,
router, session = Router.create(
node_id_bytes.bytes,
bootstrap_peers=args.bootstrap_peers,
listen_port=args.libp2p_port,
)
await router.register_topic(topics.GLOBAL_EVENTS)
await router.register_topic(topics.LOCAL_EVENTS)
await router.register_topic(topics.COMMANDS)
await router.register_topic(topics.ELECTION_MESSAGES)
await router.register_topic(topics.CONNECTION_MESSAGES)
await router.register_topic(topics.DOWNLOAD_COMMANDS)
event_router = EventRouter(
session_id,
command_sender=router.sender(topics.COMMANDS),
external_outbound=router.sender(topics.LOCAL_EVENTS),
external_inbound=router.receiver(topics.GLOBAL_EVENTS),
command_sender=session.net_sender("orchestrator"),
)
logger.info(f"Starting node {node_id}")
@@ -85,9 +88,9 @@ class Node:
node_id,
port=args.api_port,
event_receiver=event_router.receiver(),
command_sender=router.sender(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
election_receiver=router.receiver(topics.ELECTION_MESSAGES),
session=session,
)
else:
api = None
@@ -97,9 +100,9 @@ class Node:
node_id,
event_receiver=event_router.receiver(),
event_sender=event_router.sender(),
command_sender=router.sender(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
api_port=args.api_port,
session=session,
)
else:
worker = None
@@ -111,8 +114,8 @@ class Node:
event_sender=event_router.sender(),
global_event_sender=router.sender(topics.GLOBAL_EVENTS),
local_event_receiver=router.receiver(topics.LOCAL_EVENTS),
command_receiver=router.receiver(topics.COMMANDS),
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
command_receiver=session.net_receiver("orchestrator"),
)
er_send, er_recv = channel[ElectionResult]()
@@ -125,7 +128,6 @@ class Node:
election_message_sender=router.sender(topics.ELECTION_MESSAGES),
election_message_receiver=router.receiver(topics.ELECTION_MESSAGES),
connection_message_receiver=router.receiver(topics.CONNECTION_MESSAGES),
command_receiver=router.receiver(topics.COMMANDS),
election_result_sender=er_send,
)
@@ -139,6 +141,7 @@ class Node:
master,
api,
node_id,
session,
args.offline,
args.api_port,
)
@@ -188,7 +191,7 @@ class Node:
self.event_router.shutdown()
self.event_router = EventRouter(
result.session_id,
self.router.sender(topics.COMMANDS),
self.session.net_sender("orchestrator"),
self.router.receiver(topics.GLOBAL_EVENTS),
self.router.sender(topics.LOCAL_EVENTS),
)
@@ -209,10 +212,10 @@ class Node:
event_sender=self.event_router.sender(),
global_event_sender=self.router.sender(topics.GLOBAL_EVENTS),
local_event_receiver=self.router.receiver(topics.LOCAL_EVENTS),
command_receiver=self.router.receiver(topics.COMMANDS),
download_command_sender=self.router.sender(
topics.DOWNLOAD_COMMANDS
),
command_receiver=self.session.net_receiver("orchestrator"),
)
self._tg.start_soon(self.master.run)
elif (
@@ -248,11 +251,11 @@ class Node:
self.node_id,
event_receiver=self.event_router.receiver(),
event_sender=self.event_router.sender(),
command_sender=self.router.sender(topics.COMMANDS),
download_command_sender=self.router.sender(
topics.DOWNLOAD_COMMANDS
),
api_port=self._api_port,
session=self.session,
)
self._tg.start_soon(self.worker.run)
if self.api:
@@ -264,12 +267,21 @@ class Node:
def main():
# Exit early if no PID file (not compatible with double-for daemonization yet)
try:
pidfile = Pidfile(EXO_PID_FILE, 0o0600)
pidfile.write()
except (PidfileError, OSError) as exception:
print(exception, file=sys.stderr)
raise SystemExit(1) from exception
args = Args.parse()
soft, hard = resource.getrlimit(resource.RLIMIT_NOFILE)
target = min(max(soft, 65535), hard)
resource.setrlimit(resource.RLIMIT_NOFILE, (target, hard))
mp.set_start_method("spawn", force=True)
# TODO: Refactor the current verbosity system
logger_setup(EXO_LOG, args.verbosity)
logger.info(f"{'=' * 40}")
@@ -306,6 +318,7 @@ def main():
finally:
logger.info("EXO Shutdown complete")
logger_cleanup()
del pidfile
class Args(FrozenModel):
+272 -281
View File
@@ -1,7 +1,9 @@
from datetime import datetime, timedelta, timezone
import anyio
from exo_net import NetReceiver
from loguru import logger
from pydantic import TypeAdapter
from exo.master.placement import (
add_instance_to_placements,
@@ -15,11 +17,11 @@ from exo.shared.apply import apply
from exo.shared.constants import EXO_EVENT_LOG_DIR, EXO_TRACING_ENABLED
from exo.shared.types.commands import (
AddCustomModelCard,
Command,
CreateInstance,
DeleteCustomModelCard,
DeleteInstance,
DeleteInstanceLink,
ForwarderCommand,
ForwarderDownloadCommand,
ImageEdits,
ImageGeneration,
@@ -121,7 +123,7 @@ class Master:
node_id: NodeId,
session_id: SessionId,
*,
command_receiver: Receiver[ForwarderCommand],
command_receiver: NetReceiver, # todo: not this type
event_sender: Sender[Event],
local_event_receiver: Receiver[LocalForwarderEvent],
global_event_sender: Sender[GlobalForwarderEvent],
@@ -155,301 +157,290 @@ class Master:
self._event_log.close()
self.global_event_sender.close()
self.local_event_receiver.close()
self.command_receiver.close()
async def shutdown(self):
logger.info("Stopping Master")
self._tg.cancel_tasks()
async def _command_processor(self) -> None:
with self.command_receiver as commands:
async for forwarder_command in commands:
try:
logger.info(f"Executing command: {forwarder_command.command}")
while True:
data = await self.command_receiver.recv()
if not data:
break
try:
command = TypeAdapter[Command](Command).validate_json(data)
logger.info(f"Executing command: {command}")
generated_events: list[Event] = []
command = forwarder_command.command
instance_task_counts: dict[InstanceId, int] = {}
match command:
case TestCommand():
pass
case TextGeneration():
prefill_only: set[InstanceId] = set()
for link in self.state.instance_links.values():
prefill_only.update(link.prefill_instances)
for link in self.state.instance_links.values():
prefill_only.difference_update(link.decode_instances)
generated_events: list[Event] = []
instance_task_counts: dict[InstanceId, int] = {}
match command:
case TestCommand():
pass
case TextGeneration():
prefill_only: set[InstanceId] = set()
for link in self.state.instance_links.values():
prefill_only.update(link.prefill_instances)
for link in self.state.instance_links.values():
prefill_only.difference_update(link.decode_instances)
for instance in self.state.instances.values():
if (
instance.shard_assignments.model_id
== command.task_params.model
and instance.instance_id not in prefill_only
):
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = (
task_count
)
if not instance_task_counts:
raise ValueError(
f"No instance found for model {command.task_params.model}"
)
available_instance_ids = sorted(
instance_task_counts.keys(),
key=lambda instance_id: instance_task_counts[
instance_id
],
)
decode_instance_id = available_instance_ids[0]
task_id = TaskId()
params = command.task_params.model_copy(
update={
"prefill_endpoint": _prefill_endpoint_for(
self.state, decode_instance_id
),
}
)
generated_events.append(
TaskCreated(
task_id=task_id,
task=TextGenerationTask(
task_id=task_id,
command_id=command.command_id,
instance_id=decode_instance_id,
task_status=TaskStatus.Pending,
task_params=params,
),
)
)
self.command_task_mapping[command.command_id] = task_id
case ImageGeneration():
for instance in self.state.instances.values():
if (
instance.shard_assignments.model_id
== command.task_params.model
):
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = (
task_count
)
if not instance_task_counts:
raise ValueError(
f"No instance found for model {command.task_params.model}"
)
available_instance_ids = sorted(
instance_task_counts.keys(),
key=lambda instance_id: instance_task_counts[
instance_id
],
)
task_id = TaskId()
selected_instance_id = available_instance_ids[0]
generated_events.append(
TaskCreated(
task_id=task_id,
task=ImageGenerationTask(
task_id=task_id,
command_id=command.command_id,
instance_id=selected_instance_id,
task_status=TaskStatus.Pending,
task_params=command.task_params,
),
)
)
self.command_task_mapping[command.command_id] = task_id
if EXO_TRACING_ENABLED:
selected_instance = self.state.instances.get(
selected_instance_id
)
if selected_instance:
ranks = 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 (
instance.shard_assignments.model_id
== command.task_params.model
):
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = (
task_count
)
if not instance_task_counts:
raise ValueError(
f"No instance found for model {command.task_params.model}"
)
available_instance_ids = sorted(
instance_task_counts.keys(),
key=lambda instance_id: instance_task_counts[
instance_id
],
)
task_id = TaskId()
selected_instance_id = available_instance_ids[0]
generated_events.append(
TaskCreated(
task_id=task_id,
task=ImageEditsTask(
task_id=task_id,
command_id=command.command_id,
instance_id=selected_instance_id,
task_status=TaskStatus.Pending,
task_params=command.task_params,
),
)
)
self.command_task_mapping[command.command_id] = task_id
if EXO_TRACING_ENABLED:
selected_instance = self.state.instances.get(
selected_instance_id
)
if selected_instance:
ranks = 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(
self.state.instances, placement, self.state.tasks
)
for cmd in cancel_unnecessary_downloads(
placement, self.state.downloads
for instance in self.state.instances.values():
if (
instance.shard_assignments.model_id
== command.task_params.model
and instance.instance_id not in prefill_only
):
await self.download_command_sender.send(
ForwarderDownloadCommand(
origin=self._system_id, command=cmd
)
)
generated_events.extend(transition_events)
case PlaceInstance():
placement = place_instance(
command,
self.state.topology,
self.state.instances,
self.state.node_memory,
self.state.node_network,
download_status=self.state.downloads,
node_rdma_ctl=self.state.node_rdma_ctl,
)
transition_events = get_transition_events(
self.state.instances, placement, self.state.tasks
)
generated_events.extend(transition_events)
case CreateInstance():
placement = add_instance_to_placements(
command,
self.state.topology,
self.state.instances,
)
transition_events = get_transition_events(
self.state.instances, placement, self.state.tasks
)
generated_events.extend(transition_events)
case SendInputChunk(chunk=chunk):
generated_events.append(
InputChunkReceived(
command_id=chunk.command_id,
chunk=chunk,
)
)
case TaskCancelled():
if (
task_id := self.command_task_mapping.get(
command.cancelled_command_id
)
) is not None:
generated_events.append(
TaskStatusUpdated(
task_status=TaskStatus.Cancelled,
task_id=task_id,
)
)
else:
logger.warning(
f"Nonexistent command {command.cancelled_command_id} cancelled"
)
case TaskFinished():
if (
task_id := self.command_task_mapping.pop(
command.finished_command_id, None
)
) is not None:
generated_events.append(TaskDeleted(task_id=task_id))
else:
logger.warning(
f"Finished command {command.finished_command_id} finished"
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = task_count
case AddCustomModelCard():
generated_events.append(
CustomModelCardAdded(model_card=command.model_card)
if not instance_task_counts:
raise ValueError(
f"No instance found for model {command.task_params.model}"
)
case DeleteCustomModelCard():
generated_events.append(
CustomModelCardDeleted(model_id=command.model_id)
)
case SetInstanceLink():
link = InstanceLink(
link_id=command.link_id,
prefill_instances=list(
dict.fromkeys(command.prefill_instances)
available_instance_ids = sorted(
instance_task_counts.keys(),
key=lambda instance_id: instance_task_counts[instance_id],
)
decode_instance_id = available_instance_ids[0]
task_id = TaskId()
params = command.task_params.model_copy(
update={
"prefill_endpoint": _prefill_endpoint_for(
self.state, decode_instance_id
),
decode_instances=list(
dict.fromkeys(command.decode_instances)
}
)
generated_events.append(
TaskCreated(
task_id=task_id,
task=TextGenerationTask(
task_id=task_id,
command_id=command.command_id,
instance_id=decode_instance_id,
task_status=TaskStatus.Pending,
task_params=params,
),
)
generated_events.append(InstanceLinkCreated(link=link))
case DeleteInstanceLink():
generated_events.append(
InstanceLinkDeleted(link_id=command.link_id)
)
case RequestEventLog():
end = len(self._event_log)
for i, event in enumerate(
self._event_log.read_range(command.since_idx, end),
start=command.since_idx,
)
self.command_task_mapping[command.command_id] = task_id
case ImageGeneration():
for instance in self.state.instances.values():
if (
instance.shard_assignments.model_id
== command.task_params.model
):
await self._send_event(IndexedEvent(idx=i, event=event))
for event in generated_events:
await self.event_sender.send(event)
except ValueError as e:
logger.opt(exception=e).warning("Error in command processor")
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = task_count
if not instance_task_counts:
raise ValueError(
f"No instance found for model {command.task_params.model}"
)
available_instance_ids = sorted(
instance_task_counts.keys(),
key=lambda instance_id: instance_task_counts[instance_id],
)
task_id = TaskId()
selected_instance_id = available_instance_ids[0]
generated_events.append(
TaskCreated(
task_id=task_id,
task=ImageGenerationTask(
task_id=task_id,
command_id=command.command_id,
instance_id=selected_instance_id,
task_status=TaskStatus.Pending,
task_params=command.task_params,
),
)
)
self.command_task_mapping[command.command_id] = task_id
if EXO_TRACING_ENABLED:
selected_instance = self.state.instances.get(
selected_instance_id
)
if selected_instance:
ranks = 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 (
instance.shard_assignments.model_id
== command.task_params.model
):
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = task_count
if not instance_task_counts:
raise ValueError(
f"No instance found for model {command.task_params.model}"
)
available_instance_ids = sorted(
instance_task_counts.keys(),
key=lambda instance_id: instance_task_counts[instance_id],
)
task_id = TaskId()
selected_instance_id = available_instance_ids[0]
generated_events.append(
TaskCreated(
task_id=task_id,
task=ImageEditsTask(
task_id=task_id,
command_id=command.command_id,
instance_id=selected_instance_id,
task_status=TaskStatus.Pending,
task_params=command.task_params,
),
)
)
self.command_task_mapping[command.command_id] = task_id
if EXO_TRACING_ENABLED:
selected_instance = self.state.instances.get(
selected_instance_id
)
if selected_instance:
ranks = 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(
self.state.instances, placement, self.state.tasks
)
for cmd in cancel_unnecessary_downloads(
placement, self.state.downloads
):
await self.download_command_sender.send(
ForwarderDownloadCommand(
origin=self._system_id, command=cmd
)
)
generated_events.extend(transition_events)
case PlaceInstance():
placement = place_instance(
command,
self.state.topology,
self.state.instances,
self.state.node_memory,
self.state.node_network,
download_status=self.state.downloads,
node_rdma_ctl=self.state.node_rdma_ctl,
)
transition_events = get_transition_events(
self.state.instances, placement, self.state.tasks
)
generated_events.extend(transition_events)
case CreateInstance():
placement = add_instance_to_placements(
command,
self.state.topology,
self.state.instances,
)
transition_events = get_transition_events(
self.state.instances, placement, self.state.tasks
)
generated_events.extend(transition_events)
case SendInputChunk(chunk=chunk):
generated_events.append(
InputChunkReceived(
command_id=chunk.command_id,
chunk=chunk,
)
)
case TaskCancelled():
if (
task_id := self.command_task_mapping.get(
command.cancelled_command_id
)
) is not None:
generated_events.append(
TaskStatusUpdated(
task_status=TaskStatus.Cancelled,
task_id=task_id,
)
)
else:
logger.warning(
f"Nonexistent command {command.cancelled_command_id} cancelled"
)
case TaskFinished():
if (
task_id := self.command_task_mapping.pop(
command.finished_command_id, None
)
) is not None:
generated_events.append(TaskDeleted(task_id=task_id))
else:
logger.warning(
f"Finished command {command.finished_command_id} finished"
)
case AddCustomModelCard():
generated_events.append(
CustomModelCardAdded(model_card=command.model_card)
)
case DeleteCustomModelCard():
generated_events.append(
CustomModelCardDeleted(model_id=command.model_id)
)
case SetInstanceLink():
link = InstanceLink(
link_id=command.link_id,
prefill_instances=list(
dict.fromkeys(command.prefill_instances)
),
decode_instances=list(
dict.fromkeys(command.decode_instances)
),
)
generated_events.append(InstanceLinkCreated(link=link))
case DeleteInstanceLink():
generated_events.append(
InstanceLinkDeleted(link_id=command.link_id)
)
case RequestEventLog():
end = len(self._event_log)
for i, event in enumerate(
self._event_log.read_range(command.since_idx, end),
start=command.since_idx,
):
await self._send_event(IndexedEvent(idx=i, event=event))
for event in generated_events:
await self.event_sender.send(event)
except ValueError as e:
logger.opt(exception=e).warning("Error in command processor")
# These plan loops are the cracks showing in our event sourcing architecture - more things could be commands
async def _plan(self) -> None:
+1 -3
View File
@@ -6,7 +6,6 @@ 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,
@@ -47,8 +46,7 @@ from exo.utils.channels import channel
@pytest.mark.asyncio
async def test_master():
keypair = get_node_id_keypair()
node_id = NodeId(keypair.to_node_id())
node_id = NodeId("master test")
session_id = SessionId(master_node_id=node_id, election_clock=0)
ge_sender, global_event_receiver = channel[GlobalForwarderEvent]()
+1 -1
View File
@@ -1,4 +1,4 @@
from exo_pyo3_bindings import PyFromSwarm
from exo_net import PyFromSwarm
from exo.utils.pydantic_ext import FrozenModel
+6 -6
View File
@@ -4,9 +4,10 @@ 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 ForwarderCommand, RequestEventLog
from exo.shared.types.commands import RequestEventLog
from exo.shared.types.common import SessionId, SystemId
from exo.shared.types.events import (
Event,
@@ -23,7 +24,7 @@ from exo.utils.task_group import TaskGroup
@dataclass
class EventRouter:
session_id: SessionId
command_sender: Sender[ForwarderCommand]
command_sender: NetSender
external_inbound: Receiver[GlobalForwarderEvent]
external_outbound: Sender[LocalForwarderEvent]
_system_id: SystemId = field(init=False, default_factory=SystemId)
@@ -152,10 +153,9 @@ class EventRouter:
f"Nack attempt {self._nack_attempts}: Requesting Event Log from {since_idx}"
)
await self.command_sender.send(
ForwarderCommand(
origin=self._system_id,
command=RequestEventLog(since_idx=since_idx),
)
RequestEventLog(since_idx=since_idx)
.model_dump_json()
.encode("utf-8")
)
finally:
if self._nack_cancel_scope is scope:
+7 -49
View File
@@ -2,8 +2,6 @@ from collections.abc import Sequence
from copy import copy
from itertools import count
from math import inf
from os import PathLike
from pathlib import Path
from typing import cast
from anyio import (
@@ -12,15 +10,9 @@ from anyio import (
move_on_after,
sleep_forever,
)
from exo_pyo3_bindings import (
Keypair,
NetworkingHandle,
PyFromSwarm,
)
from filelock import FileLock
from exo_net import NetworkingHandle, PyFromSwarm, PySession
from loguru import logger
from exo.shared.constants import EXO_NODE_ID_KEYPAIR
from exo.utils.channels import Receiver, Sender, channel
from exo.utils.pydantic_ext import FrozenModel
from exo.utils.task_group import TaskGroup
@@ -102,13 +94,14 @@ class Router:
@classmethod
def create(
cls,
identity: Keypair,
identity: bytes,
bootstrap_peers: Sequence[str] = (),
listen_port: int = 0,
) -> "Router":
return cls(
handle=NetworkingHandle(identity, list(bootstrap_peers), listen_port)
) -> "tuple[Router, PySession]":
handle, session = NetworkingHandle.new(
identity, list(bootstrap_peers), listen_port
)
return cls(handle=handle), session
def __init__(self, handle: NetworkingHandle):
self.topic_routers: dict[str, TopicRouter[FrozenModel]] = {}
@@ -189,9 +182,7 @@ class Router:
logger.debug(from_swarm)
match from_swarm:
case PyFromSwarm.Message(topic, data):
logger.trace(
f"Received message on {topic} with payload {data}"
)
logger.trace(f"Received message on {topic} with payload {data}")
if topic not in self.topic_routers:
logger.warning(
f"Received message on unknown or inactive topic {topic}"
@@ -228,36 +219,3 @@ class Router:
"Sending overlarge payload, network performance may be temporarily degraded"
)
await self._net.gossipsub_publish(topic, data)
def get_node_id_keypair(
path: str | bytes | PathLike[str] | PathLike[bytes] = EXO_NODE_ID_KEYPAIR,
) -> Keypair:
"""
Obtains the :class:`Keypair` associated with this node-ID.
Obtain the :class:`PeerId` by from it.
"""
# TODO(evan): bring back node id persistence once we figure out how to deal with duplicates
return Keypair.generate()
def lock_path(path: str | bytes | PathLike[str] | PathLike[bytes]) -> Path:
return Path(str(path) + ".lock")
# operate with cross-process lock to avoid race conditions
with FileLock(lock_path(path)):
with open(path, "a+b") as f: # opens in append-mode => starts at EOF
# if non-zero EOF, then file exists => use to get node-ID
if f.tell() != 0:
f.seek(0) # go to start & read protobuf-encoded bytes
protobuf_encoded = f.read()
try: # if decoded successfully, save & return
return Keypair.from_bytes(protobuf_encoded)
except ValueError as e: # on runtime error, assume corrupt file
logger.warning(f"Encountered error when trying to get keypair: {e}")
# if no valid credentials, create new ones and persist
with open(path, "w+b") as f:
keypair = Keypair.generate()
f.write(keypair.to_bytes())
return keypair
+1
View File
@@ -69,6 +69,7 @@ DASHBOARD_DIR = (
EXO_LOG_DIR = EXO_CACHE_HOME / "exo_log"
EXO_LOG = EXO_LOG_DIR / "exo.log"
EXO_TEST_LOG = EXO_CACHE_HOME / "exo_test.log"
EXO_PID_FILE = EXO_CACHE_HOME / "exo.pid"
# Identity (config)
EXO_NODE_ID_KEYPAIR = EXO_CONFIG_HOME / "node_id.keypair"
-14
View File
@@ -9,7 +9,6 @@ from anyio import (
from loguru import logger
from exo.routing.connection_message import ConnectionMessage
from exo.shared.types.commands import ForwarderCommand
from exo.shared.types.common import NodeId, SessionId
from exo.utils.channels import Receiver, Sender
from exo.utils.pydantic_ext import FrozenModel
@@ -22,7 +21,6 @@ class ElectionMessage(FrozenModel):
clock: int
seniority: int
proposed_session: SessionId
commands_seen: int
# Could eventually include a list of neighbour nodes for centrality
def __lt__(self, other: Self) -> bool:
@@ -30,8 +28,6 @@ class ElectionMessage(FrozenModel):
return self.clock < other.clock
if self.seniority != other.seniority:
return self.seniority < other.seniority
elif self.commands_seen != other.commands_seen:
return self.commands_seen < other.commands_seen
else:
return (
self.proposed_session.master_node_id
@@ -54,7 +50,6 @@ class Election:
election_message_sender: Sender[ElectionMessage],
election_result_sender: Sender[ElectionResult],
connection_message_receiver: Receiver[ConnectionMessage],
command_receiver: Receiver[ForwarderCommand],
is_candidate: bool = True,
seniority: int = 0,
):
@@ -64,7 +59,6 @@ class Election:
self.seniority = seniority if is_candidate else -1
self.clock = 0
self.node_id = node_id
self.commands_seen = 0
# Every node spawns as master
self.current_session: SessionId = SessionId(
master_node_id=node_id, election_clock=0
@@ -75,7 +69,6 @@ class Election:
self._em_receiver = election_message_receiver
self._er_sender = election_result_sender
self._cm_receiver = connection_message_receiver
self._co_receiver = command_receiver
# Campaign state
self._candidates: list[ElectionMessage] = []
@@ -89,7 +82,6 @@ class Election:
async with self._tg as tg:
tg.start_soon(self._election_receiver)
tg.start_soon(self._connection_receiver)
tg.start_soon(self._command_counter)
# And start an election immediately, that instantly resolves
candidates: list[ElectionMessage] = []
@@ -179,11 +171,6 @@ class Election:
logger.debug("Campaign started")
logger.debug("Connection message added")
async def _command_counter(self) -> None:
with self._co_receiver as commands:
async for _command in commands:
self.commands_seen += 1
async def _campaign(
self, candidates: list[ElectionMessage], campaign_timeout: float
) -> None:
@@ -261,5 +248,4 @@ class Election:
),
clock=c,
seniority=self.seniority,
commands_seen=self.commands_seen,
)
+14 -2
View File
@@ -32,15 +32,21 @@ class TokenChunk(BaseChunk):
class ErrorChunk(BaseChunk):
error_message: str
finish_reason: Literal["error"] = "error"
@property
def finish_reason(self) -> Literal["error"]:
return "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
@@ -84,7 +90,13 @@ 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
TextGenerationChunk = TokenChunk | ToolCallChunk | ErrorChunk
ImageGenerationChunk = ImageChunk | ErrorChunk
Chunk = StatusChunk | GenerationChunk
+2 -2
View File
@@ -159,10 +159,10 @@ Event = (
| NodeTimedOut
| NodeGatheredInfo
| NodeDownloadProgress
| ChunkGenerated
| InputChunkReceived
| TopologyEdgeCreated
| TopologyEdgeDeleted
| ChunkGenerated
| InputChunkReceived
| TracesCollected
| TracesMerged
| CustomModelCardAdded
+40 -13
View File
@@ -19,19 +19,21 @@ class PowerSampler:
):
self._get_node_system = get_node_system
self._interval = interval
self._samples: defaultdict[NodeId, list[SystemPerformanceProfile]] = (
defaultdict(list)
)
self._samples: defaultdict[
NodeId, list[tuple[float, SystemPerformanceProfile]]
] = defaultdict(list)
self._start_time: float | None = None
self._stopped = False
def _take_sample(self) -> None:
def _take_sample(self, t_rel: float | None = None) -> None:
assert self._start_time is not None
ts = t_rel if t_rel is not None else time.perf_counter() - self._start_time
for node_id, profile in self._get_node_system().items():
self._samples[node_id].append(profile)
self._samples[node_id].append((ts, profile))
async def run(self) -> None:
self._start_time = time.perf_counter()
self._take_sample()
self._take_sample(t_rel=0.0)
while not self._stopped:
await anyio.sleep(self._interval)
self._take_sample()
@@ -39,26 +41,51 @@ class PowerSampler:
def result(self) -> PowerUsage:
self._stopped = True
assert self._start_time is not None, "result() called before run()"
self._take_sample()
elapsed = time.perf_counter() - self._start_time
self._take_sample(t_rel=elapsed)
node_stats: list[NodePowerStats] = []
for node_id, profiles in self._samples.items():
n = len(profiles)
total_energy_j = 0.0
for node_id, ts_profiles in self._samples.items():
n = len(ts_profiles)
if n == 0:
continue
node_energy_j = trapezoidal_energy(ts_profiles, elapsed)
avg_power_w = node_energy_j / elapsed if elapsed > 0 else 0.0
total_energy_j += node_energy_j
node_stats.append(
NodePowerStats(
node_id=node_id,
samples=n,
avg_sys_power=sum(p.sys_power for p in profiles) / n,
avg_sys_power=avg_power_w,
)
)
total_avg_sys = sum(ns.avg_sys_power for ns in node_stats)
total_avg_sys_w = total_energy_j / elapsed if elapsed > 0 else 0.0
return PowerUsage(
elapsed_seconds=elapsed,
nodes=node_stats,
total_avg_sys_power_watts=total_avg_sys,
total_energy_joules=total_avg_sys * elapsed,
total_avg_sys_power_watts=total_avg_sys_w,
total_energy_joules=total_energy_j,
)
def trapezoidal_energy(
ts_profiles: list[tuple[float, SystemPerformanceProfile]],
elapsed: float,
) -> float:
"""Integrate sys_power(t) over the sample window using the trapezoidal rule.
First sample is anchored at t=0 and last at t=elapsed (set by `run` /
`result`), so the integral spans the full request interval. Falls back to
power * elapsed when only one sample exists (constant-power assumption)."""
if len(ts_profiles) == 1:
return ts_profiles[0][1].sys_power * elapsed
energy_j = 0.0
for i in range(1, len(ts_profiles)):
t_prev, p_prev = ts_profiles[i - 1]
t_cur, p_cur = ts_profiles[i]
dt = t_cur - t_prev
if dt <= 0:
continue
energy_j += (p_prev.sys_power + p_cur.sys_power) / 2.0 * dt
return energy_j
+8
View File
@@ -0,0 +1,8 @@
import multiprocessing as mp
import pytest
@pytest.fixture(scope="session", autouse=True)
def mp_force_spawn():
mp.set_start_method("spawn", force=True)
+83
View File
@@ -0,0 +1,83 @@
from __future__ import annotations
import gc
import os
import subprocess
import sys
import textwrap
from pathlib import Path
from typing import Final
import exo.utils.pidfile as pidfile
import pytest
from exo.utils.pidfile import acquire_exo_pidfile
_CHILD_ACQUIRE_PIDFILE_SCRIPT: Final = textwrap.dedent(
"""
import sys
from pathlib import Path
from unittest.mock import patch
import exo.utils.pidfile as pidfile
from exo.utils.pidfile import PidfileLockError, acquire_exo_pidfile
with patch.object(pidfile, "EXO_PID_FILE", Path(sys.argv[1])):
try:
handle = acquire_exo_pidfile()
except PidfileLockError as exception:
print(str(exception))
raise SystemExit(73) from exception
del handle
"""
)
def _use_pidfile_path(monkeypatch: pytest.MonkeyPatch, path: Path) -> None:
monkeypatch.setattr(pidfile, "EXO_PID_FILE", path)
def _run_child_acquire_pidfile(path: Path) -> subprocess.CompletedProcess[str]:
return subprocess.run(
[sys.executable, "-c", _CHILD_ACQUIRE_PIDFILE_SCRIPT, str(path)],
check=False,
capture_output=True,
text=True,
)
def test_acquire_exo_pidfile_writes_current_pid_and_removes_on_drop(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
path = tmp_path / "exo.pid"
_use_pidfile_path(monkeypatch, path)
handle = acquire_exo_pidfile()
assert path.read_text() == str(os.getpid())
del handle
gc.collect()
assert not path.exists()
def test_acquire_exo_pidfile_rejects_second_process(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
path = tmp_path / "exo.pid"
_use_pidfile_path(monkeypatch, path)
handle = acquire_exo_pidfile()
try:
blocked_child = _run_child_acquire_pidfile(path)
assert blocked_child.returncode == 73
assert "Failed to acquire EXO pidfile" in blocked_child.stdout
finally:
del handle
gc.collect()
unblocked_child = _run_child_acquire_pidfile(path)
assert unblocked_child.returncode == 0
assert unblocked_child.stdout == ""
+30
View File
@@ -111,6 +111,36 @@ async def test_empty_state() -> None:
assert result.total_energy_joules == 0.0
def test_trapezoidal_unit_dt_weighting() -> None:
"""Pure unit test on the integration helper. Crafted samples where the
arithmetic mean is wildly wrong vs the time-weighted result."""
from exo.utils.power_sampler import trapezoidal_energy
# 5 s window. Power = 10 W for the first 4.9 s, then 100 W for the last 0.1 s.
# Three samples: t=0 W=10, t=4.9 W=10, t=5.0 W=100.
samples = [
(0.0, _make_profile(10.0)),
(4.9, _make_profile(10.0)),
(5.0, _make_profile(100.0)),
]
energy = trapezoidal_energy(samples, elapsed=5.0)
# (10+10)/2 * 4.9 + (10+100)/2 * 0.1 = 49 + 5.5 = 54.5 J
assert abs(energy - 54.5) < 1e-9
avg = energy / 5.0 # 10.9 W
# Arithmetic mean of the three samples would be (10+10+100)/3 ≈ 40 W.
# Trapezoidal correctly weights each segment by its dt.
assert abs(avg - 10.9) < 1e-9
def test_trapezoidal_unit_single_sample() -> None:
"""One sample: no window to integrate over, so fall back to constant power
over the elapsed duration."""
from exo.utils.power_sampler import trapezoidal_energy
samples = [(0.0, _make_profile(42.0))]
assert trapezoidal_energy(samples, elapsed=3.0) == 42.0 * 3.0
async def test_result_stops_sampling() -> None:
"""Calling result() should stop the sampler's run loop."""
state: dict[NodeId, SystemPerformanceProfile] = {
+8 -8
View File
@@ -4,6 +4,7 @@ from datetime import datetime, timezone
import anyio
from anyio import fail_after, to_thread
from exo_net import PySession
from loguru import logger
from exo.api.types import ImageEditsTaskParams
@@ -14,7 +15,6 @@ 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,
)
@@ -65,16 +65,17 @@ class Worker:
*,
event_receiver: Receiver[IndexedEvent],
event_sender: Sender[Event],
session: PySession,
# This is for requesting updates. It doesn't need to be a general command sender right now,
# but I think it's the correct way to be thinking about commands
command_sender: Sender[ForwarderCommand],
download_command_sender: Sender[ForwarderDownloadCommand],
api_port: int,
):
self.node_id: NodeId = node_id
self.event_receiver = event_receiver
self.event_sender = event_sender
self.command_sender = command_sender
self.session = session
self.command_sender = session.net_sender("orchestrator")
self.download_command_sender = download_command_sender
self.api_port = api_port
@@ -114,7 +115,6 @@ class Worker:
# Actual shutdown code - waits for all tasks to complete before executing.
logger.info("Stopping Worker")
self.event_sender.close()
self.command_sender.close()
self.download_command_sender.close()
for runner in self.runners.values():
runner.shutdown()
@@ -209,10 +209,9 @@ class Worker:
f"Instance {iid} exceeded {EXO_MAX_INSTANCE_RETRIES} retries, requesting deletion"
)
await self.command_sender.send(
ForwarderCommand(
origin=self._system_id,
command=DeleteInstance(instance_id=iid),
)
DeleteInstance(instance_id=iid)
.model_dump_json()
.encode("utf-8")
)
continue
@@ -375,6 +374,7 @@ class Worker:
runner = RunnerSupervisor.create(
bound_instance=task.bound_instance,
event_sender=self.event_sender.clone(),
session=self.session,
)
self.runners[task.bound_instance.bound_runner_id] = runner
self._tg.start_soon(runner.run)
+33 -15
View File
@@ -10,9 +10,11 @@ from anyio import (
ClosedResourceError,
to_thread,
)
from exo_net import NetSender, PySession
from loguru import logger
from exo.shared.types.chunks import ErrorChunk
from exo.shared.types.commands import CommandId
from exo.shared.types.events import (
ChunkGenerated,
Event,
@@ -45,20 +47,18 @@ from exo.utils.channels import MpReceiver, MpSender, Sender, mp_channel
from exo.utils.task_group import TaskGroup
from exo.worker.runner.bootstrap import entrypoint
PREFILL_TIMEOUT_SECONDS = 60
DECODE_TIMEOUT_SECONDS = 5
@dataclass(eq=False)
class RunnerSupervisor:
shard_metadata: ShardMetadata
bound_instance: BoundInstance
runner_process: mp.Process
initialize_timeout: float
_ev_recv: MpReceiver[Event]
_task_sender: MpSender[Task]
_event_sender: Sender[Event]
_cancel_sender: MpSender[TaskId]
session: PySession
_tg: TaskGroup = field(default_factory=TaskGroup, init=False)
status: RunnerStatus = field(default_factory=RunnerIdle, init=False)
pending: dict[TaskId, anyio.Event] = field(default_factory=dict, init=False)
@@ -75,7 +75,7 @@ class RunnerSupervisor:
*,
bound_instance: BoundInstance,
event_sender: Sender[Event],
initialize_timeout: float = 400,
session: PySession,
) -> Self:
ev_send, ev_recv = mp_channel[Event]()
task_sender, task_recv = mp_channel[Task]()
@@ -99,11 +99,11 @@ class RunnerSupervisor:
bound_instance=bound_instance,
shard_metadata=shard_metadata,
runner_process=runner_process,
initialize_timeout=initialize_timeout,
_ev_recv=ev_recv,
_task_sender=task_sender,
_cancel_sender=cancel_sender,
_event_sender=event_sender,
session=session,
)
return self
@@ -210,9 +210,26 @@ class RunnerSupervisor:
await self._check_runner(TimeoutError("cancel pipe blocked"))
async def _forward_events(self):
pubs: dict[CommandId, NetSender] = {}
try:
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(
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)
continue
if isinstance(event, RunnerStatusUpdated):
self.status = event.runner_status
if isinstance(event, TaskAcknowledged):
@@ -275,17 +292,18 @@ class RunnerSupervisor:
for task in self.in_progress.values():
if isinstance(task, (TextGeneration, ImageGeneration, ImageEdits)):
with anyio.CancelScope(shield=True):
await self._event_sender.send(
ChunkGenerated(
command_id=task.command_id,
chunk=ErrorChunk(
model=self.shard_metadata.model_card.model_id,
error_message=(
"Runner shutdown before completing command "
f"({cause})"
),
send = self.session.net_sender(
f"runners/{self.bound_instance.bound_runner_id}/active_tasks/{task.command_id}/chunks"
)
await send.send(
ErrorChunk(
model=self.shard_metadata.model_card.model_id,
error_message=(
f"Runner shutdown before completing command ({cause})"
),
)
.model_dump_json()
.encode("utf-8")
)
try:
File renamed without changes.
+181
View File
@@ -0,0 +1,181 @@
# type: ignore
"""Pytest configuration for marker-driven exo integration tests.
Test authors declare requirements via markers:
@pytest.mark.cluster(count=2, thunderbolt='a2a')
@pytest.mark.instance('mlx-community/Llama-3.2-1B-Instruct-4bit',
sharding='tensor', comm='jaccl')
def test_jaccl_inference(session):
resp = session.chat('What is 2+2?')
assert '4' in resp
Clusters are cached by `ClusterSpec`; tests with the same cluster_spec
share a deployment. Each test places its own instance (matching its
`@pytest.mark.instance`), and instances are cleaned up after the test.
Run with:
uv run pytest tests/ -v
uv run pytest tests/ -v --hosts s2,s4,s9,s10
"""
from __future__ import annotations
import contextlib
import json
import pytest
from exo_tools.cluster import ClusterInfo, EcoSession
from exo_tools.harness import cleanup_all_instances, place_instance
from .framework import (
ClusterSpec,
Session,
parse_cluster_marker,
parse_instance_marker,
)
# Single eco session for the entire test process.
eco = EcoSession(user_prefix="test")
# Cluster cache keyed by ClusterSpec — tests with the same spec share a deployment.
# Cleared at session teardown.
_cluster_cache: dict[ClusterSpec, ClusterInfo] = {}
def pytest_addoption(parser):
parser.addoption(
"--hosts",
default=None,
help="Comma-separated list of hosts (e.g. s2,s4,s9,s10). "
"Overrides constraint-based reservation.",
)
def pytest_configure(config):
"""Register custom markers."""
config.addinivalue_line(
"markers",
"cluster(count=N, thunderbolt=Thunderbolt|None, min_memory=GB, chip=PATTERN): "
"declare cluster requirements for a test",
)
config.addinivalue_line(
"markers",
"instance(model_id, sharding=Sharding, comm=Comm, min_nodes=N): "
"declare instance placement for a test",
)
def pytest_report_header(config):
"""Show the eco user and hosts for this test session."""
hosts = config.getoption("--hosts")
lines = [f"eco user: {eco.user}"]
if hosts:
lines.append(f"hosts override: {hosts}")
return lines
@pytest.fixture(scope="session")
def _host_pool(request) -> list[str] | None:
raw = request.config.getoption("--hosts")
if raw:
return [h.strip() for h in raw.split(",") if h.strip()]
return None
@pytest.fixture
def session(request, _host_pool) -> Session:
"""Per-test fixture providing a Session matching the test's markers.
Reads @pytest.mark.cluster and @pytest.mark.instance from the test, deploys
a matching cluster (cached across tests with the same spec), places the
model, and yields a Session for the test to interact with. Cleans up the
instance after the test, and invalidates the cluster cache if the test
left nodes disconnected.
"""
cluster_marker = request.node.get_closest_marker("cluster")
instance_marker = request.node.get_closest_marker("instance")
cluster_spec = parse_cluster_marker(cluster_marker)
instance_spec = parse_instance_marker(instance_marker)
# Deploy or reuse a cluster matching the spec
cluster = _cluster_cache.get(cluster_spec)
if cluster is None:
if _host_pool:
cluster = eco.start_deploy(
hosts=_host_pool[: cluster_spec.count], wait=True
)
else:
cluster = eco.start_deploy(
count=cluster_spec.count,
thunderbolt=cluster_spec.thunderbolt,
chip=cluster_spec.chip,
min_memory_gb=cluster_spec.min_memory_gb,
wait=True,
)
_cluster_cache[cluster_spec] = cluster
# Place an instance for this test if the test specified one
instance_id = None
if instance_spec is not None:
client = cluster.make_client()
instance_id = place_instance(
client,
instance_spec.model_id,
sharding=instance_spec.sharding,
comm=instance_spec.comm,
min_nodes=instance_spec.min_nodes,
)
sess = Session(
cluster=cluster,
eco=eco,
instance_spec=instance_spec,
instance_id=instance_id,
)
yield sess
# ---- Teardown ----
# If the test left nodes disconnected, invalidate the cluster cache and
# stop the cluster so the next test deploys fresh.
if sess._stopped_hosts:
_cluster_cache.pop(cluster_spec, None)
with contextlib.suppress(Exception):
eco.stop(sess.cluster.hosts)
return
# Otherwise, clean up any instances created during the test
with contextlib.suppress(Exception):
cleanup_all_instances(sess.client)
# ---------------------------------------------------------------------------
# Session-level teardown — stop all cached clusters
# ---------------------------------------------------------------------------
@pytest.fixture(scope="session", autouse=True)
def _teardown_clusters():
yield
for cluster in _cluster_cache.values():
with contextlib.suppress(Exception):
eco.stop(cluster.hosts)
_cluster_cache.clear()
def pytest_runtest_makereport(item, call):
"""Attach cluster logs to the test report when a test fails."""
if call.when != "call" or call.excinfo is None:
return
sess = item.funcargs.get("session")
if sess is None:
return
try:
logs = eco.logs(sess.cluster.hosts, lines=200)
item.add_report_section("call", "Cluster Logs", json.dumps(logs, indent=2))
except Exception:
pass
+199
View File
@@ -0,0 +1,199 @@
"""Marker-driven test framework for exo integration tests.
Test authors declare requirements via markers:
@pytest.mark.cluster(count=2, thunderbolt='a2a')
@pytest.mark.instance('mlx-community/Llama-3.2-1B-Instruct-4bit',
sharding='tensor', comm='jaccl')
def test_jaccl_inference(session):
resp = session.chat('What is 2+2?')
assert '4' in resp
The `session` fixture reads the markers, deploys the cluster, places the
instance, and provides a `Session` object. All cluster/instance orchestration
lives in `exo_tools.harness`; this module is purely the pytest-facing layer.
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
from typing import Any
from exo_tools.client import ExoClient
from exo_tools.cluster import (
Chip,
ClusterInfo,
EcoSession,
Thunderbolt,
make_client_from_url,
)
from exo_tools.harness import Comm, Sharding
from exo.api.types.api import (
ChatCompletionChoice,
ChatCompletionRequest,
ChatCompletionResponse,
)
DEFAULT_MODEL = "mlx-community/Llama-3.2-1B-Instruct-4bit"
def _extract_content(resp: ChatCompletionResponse) -> str:
"""Extract plain-text content from a non-streaming chat completion."""
choice = resp.choices[0]
if not isinstance(choice, ChatCompletionChoice):
raise RuntimeError(
f"Expected non-streaming choice, got {type(choice).__name__}"
)
content = choice.message.content
if not isinstance(content, str):
raise RuntimeError(f"Expected string content, got {type(content).__name__}")
return content
@dataclass(frozen=True)
class ClusterSpec:
count: int = 1
thunderbolt: Thunderbolt | None = None
min_memory_gb: float | None = None
chip: Chip | None = None
@dataclass(frozen=True)
class InstanceSpec:
model_id: str
sharding: Sharding = Sharding.PIPELINE
comm: Comm = Comm.RING
min_nodes: int = 1
def parse_cluster_marker(marker) -> ClusterSpec:
if marker is None:
return ClusterSpec()
return ClusterSpec(
count=marker.kwargs.get("count", 1),
thunderbolt=marker.kwargs.get("thunderbolt"),
min_memory_gb=marker.kwargs.get("min_memory"),
chip=marker.kwargs.get("chip"),
)
def parse_instance_marker(marker) -> InstanceSpec | None:
if marker is None:
return None
if not marker.args:
raise ValueError(
"@pytest.mark.instance requires a positional model_id argument"
)
return InstanceSpec(
model_id=marker.args[0],
sharding=marker.kwargs.get("sharding", Sharding.PIPELINE),
comm=marker.kwargs.get("comm", Comm.RING),
min_nodes=marker.kwargs.get("min_nodes", 1),
)
@dataclass
class Session:
cluster: ClusterInfo
eco: EcoSession
instance_spec: InstanceSpec | None = None
instance_id: str | None = None
_stopped_hosts: set[str] = field(default_factory=set)
@property
def client(self) -> ExoClient:
for host in self.cluster.hosts:
if host not in self._stopped_hosts:
return make_client_from_url(self.cluster.api_endpoints[host])
return self.cluster.make_client()
@property
def state(self) -> dict[str, Any]:
return self.client.request_json("GET", "/state") or {}
@property
def instances(self) -> dict[str, Any]:
return self.state.get("instances", {})
# ---- Inference ----
def chat(self, prompt: str, max_tokens: int = 100) -> str:
resp = self.chat_raw(prompt, max_tokens=max_tokens)
return _extract_content(resp)
def chat_raw(self, prompt: str, **kwargs: Any) -> ChatCompletionResponse:
if not self.instance_spec:
raise RuntimeError(
"No instance placed; add @pytest.mark.instance to the test"
)
max_tokens = kwargs.pop("max_tokens", 100)
request = ChatCompletionRequest.model_validate(
{
"model": self.instance_spec.model_id,
"messages": [{"role": "user", "content": prompt}],
"max_tokens": max_tokens,
**kwargs,
}
)
return self._post_chat(request)
def multi_turn(self, messages: list[dict[str, str]], max_tokens: int = 100) -> str:
if not self.instance_spec:
raise RuntimeError(
"No instance placed; add @pytest.mark.instance to the test"
)
request = ChatCompletionRequest.model_validate(
{
"model": self.instance_spec.model_id,
"messages": messages,
"max_tokens": max_tokens,
}
)
return _extract_content(self._post_chat(request))
def _post_chat(self, request: ChatCompletionRequest) -> ChatCompletionResponse:
raw = self.client.request_json(
"POST",
"/v1/chat/completions",
body=request.model_dump(exclude_none=True),
)
return ChatCompletionResponse.model_validate(raw)
def disconnect_node(self, index: int) -> None:
"""Stop exo on a node and wait for the cluster to observe the disconnect."""
host = self.cluster.hosts[index]
self.eco.stop([host], keep=True)
self._stopped_hosts.add(host)
def reconnect_node(self, index: int) -> None:
"""Restart a previously disconnected node into the existing namespace."""
host = self.cluster.hosts[index]
self.eco.start_hosts([host], namespace=self.cluster.namespace)
self._stopped_hosts.discard(host)
def wait_ready(
self, expected_nodes: int | None = None, timeout: float = 60
) -> None:
"""Wait until the cluster has exactly `expected_nodes` visible and reporting memory.
Defaults to the count of non-stopped hosts. Use this after
`disconnect_node` / `reconnect_node` to wait for the cluster to settle.
"""
if expected_nodes is None:
expected_nodes = len(self.cluster.hosts) - len(self._stopped_hosts)
start = time.time()
while time.time() - start < timeout:
try:
state = self.state
identities = len(state.get("nodeIdentities", {}))
memory = len(state.get("nodeMemory", {}))
if identities == expected_nodes and memory == expected_nodes:
return
except Exception:
pass
time.sleep(2.0)
raise TimeoutError(
f"Cluster did not reach exactly {expected_nodes} ready nodes within {timeout}s"
)
-264
View File
@@ -1,264 +0,0 @@
import socket
from typing import Literal
import anyio
from fastapi import FastAPI
from fastapi.responses import Response, StreamingResponse
from hypercorn import Config
from hypercorn.asyncio import serve # pyright: ignore[reportUnknownVariableType]
from loguru import logger
from pydantic import BaseModel
from exo.shared.constants import EXO_DEFAULT_MODELS_DIR
from exo.shared.models.model_cards import ModelCard, ModelId
from exo.shared.types.chunks import TokenChunk
from exo.shared.types.commands import CommandId
from exo.shared.types.common import Host, NodeId
from exo.shared.types.events import ChunkGenerated, Event, RunnerStatusUpdated
from exo.shared.types.tasks import (
ConnectToGroup,
LoadModel,
Shutdown,
StartWarmup,
Task,
TextGeneration,
)
from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams
from exo.shared.types.worker.instances import (
BoundInstance,
Instance,
InstanceId,
MlxJacclInstance,
MlxRingInstance,
)
from exo.shared.types.worker.runners import (
RunnerFailed,
RunnerId,
RunnerShutdown,
ShardAssignments,
)
from exo.shared.types.worker.shards import PipelineShardMetadata, TensorShardMetadata
from exo.utils.channels import channel, mp_channel
from exo.utils.info_gatherer.info_gatherer import GatheredInfo, InfoGatherer
from exo.worker.runner.bootstrap import entrypoint
class Tests(BaseModel):
# list[hostname, ip addr]
devs: list[list[str]]
ibv_devs: list[list[str | None]] | None
model_id: ModelId
kind: Literal["ring", "jaccl", "both"]
iid = InstanceId("im testing here")
async def main():
logger.info("starting cool server majig")
cfg = Config()
cfg.bind = "0.0.0.0:52414"
# nb: shared.logging needs updating if any of this changes
cfg.accesslog = "-"
cfg.errorlog = "-"
ev = anyio.Event()
app = FastAPI()
app.post("/run_test")(run_test)
app.post("/kill")(lambda: kill(ev))
app.get("/tb_detection")(tb_detection)
app.get("/models")(list_models)
await serve(
app, # type: ignore
cfg,
shutdown_trigger=lambda: ev.wait(),
)
def kill(ev: anyio.Event):
ev.set()
return Response(status_code=204)
async def tb_detection():
send, recv = channel[GatheredInfo]()
ig = InfoGatherer(send)
with anyio.move_on_after(1):
await ig._monitor_system_profiler_thunderbolt_data() # pyright: ignore[reportPrivateUsage]
with recv:
return recv.collect()
def list_models():
sent = set[str]()
for path in EXO_DEFAULT_MODELS_DIR.rglob("model-*.safetensors"):
if "--" not in path.parent.name:
continue
name = path.parent.name.replace("--", "/")
if name in sent:
continue
sent.add(name)
yield ModelId(path.parent.name.replace("--", "/"))
async def run_test(test: Tests):
weird_hn = socket.gethostname()
for dev in test.devs:
if weird_hn.startswith(dev[0]) or dev[0].startswith(weird_hn):
hn = dev[0]
break
else:
raise ValueError(f"{weird_hn} not in {test.devs}")
async def run():
logger.info(f"testing {test.model_id}")
instances: list[Instance] = []
if test.kind in ["ring", "both"]:
i = await ring_instance(test, hn)
if i is None:
yield "no model found"
return
instances.append(i)
if test.kind in ["jaccl", "both"]:
i = await jaccl_instance(test)
if i is None:
yield "no model found"
return
instances.append(i)
for instance in instances:
recv = await execute_test(test, instance, hn)
str_out = ""
for item in recv:
if isinstance(item, ChunkGenerated):
assert isinstance(item.chunk, TokenChunk)
str_out += item.chunk.text
if isinstance(item, RunnerStatusUpdated) and isinstance(
item.runner_status, (RunnerFailed, RunnerShutdown)
):
yield str_out + "\n"
yield item.model_dump_json() + "\n"
return StreamingResponse(run())
async def ring_instance(test: Tests, hn: str) -> Instance | None:
hbn = [Host(ip="198.51.100.0", port=52417) for _ in test.devs]
world_size = len(test.devs)
for i in range(world_size):
if test.devs[i][0] == hn:
hn = test.devs[i][0]
hbn[(i - 1) % world_size] = Host(ip=test.devs[i - 1][1], port=52417)
hbn[(i + 1) % world_size] = Host(ip=test.devs[i + 1][1], port=52417)
hbn[i] = Host(ip="0.0.0.0", port=52417)
break
else:
raise ValueError(f"{hn} not in {test.devs}")
card = await ModelCard.load(test.model_id)
instance = MlxRingInstance(
instance_id=iid,
ephemeral_port=52417,
hosts_by_node={NodeId(hn): hbn},
shard_assignments=ShardAssignments(
model_id=test.model_id,
node_to_runner={NodeId(host[0]): RunnerId(host[0]) for host in test.devs},
runner_to_shard={
RunnerId(test.devs[i][0]): PipelineShardMetadata(
model_card=card,
device_rank=i,
world_size=world_size,
start_layer=(card.n_layers // world_size) * i,
end_layer=min(
card.n_layers, (card.n_layers // world_size) * (i + 1)
),
n_layers=min(card.n_layers, (card.n_layers // world_size) * (i + 1))
- (card.n_layers // world_size) * i,
)
for i in range(world_size)
},
),
)
return instance
async def execute_test(test: Tests, instance: Instance, hn: str) -> list[Event]:
world_size = len(test.devs)
commands: list[Task] = [
(LoadModel(instance_id=iid)),
(StartWarmup(instance_id=iid)),
(
TextGeneration(
task_params=TextGenerationTaskParams(
model=test.model_id,
instructions="You are a helpful assistant",
input=[
InputMessage(
role="user", content="What is the capital of France?"
)
],
),
command_id=CommandId("yo"),
instance_id=iid,
)
),
(Shutdown(runner_id=RunnerId(hn), instance_id=iid)),
]
if world_size > 1:
commands.insert(0, ConnectToGroup(instance_id=iid))
bound_instance = BoundInstance(
instance=instance, bound_runner_id=RunnerId(hn), bound_node_id=NodeId(hn)
)
ev_send, _ev_recv = mp_channel[Event]()
task_send, task_recv = mp_channel[Task]()
for command in commands:
task_send.send(command)
entrypoint(
bound_instance,
ev_send,
task_recv,
logger,
)
# TODO(evan): return ev_recv.collect()
return []
async def jaccl_instance(test: Tests) -> MlxJacclInstance | None:
card = await ModelCard.load(test.model_id)
world_size = len(test.devs)
assert test.ibv_devs
return MlxJacclInstance(
instance_id=iid,
jaccl_devices=test.ibv_devs,
# rank 0 is always coordinator
jaccl_coordinators={
NodeId(host[0]): test.devs[0][1] + ":52417" for host in test.devs
},
shard_assignments=ShardAssignments(
model_id=test.model_id,
node_to_runner={NodeId(host[0]): RunnerId(host[0]) for host in test.devs},
runner_to_shard={
RunnerId(host[0]): TensorShardMetadata(
model_card=card,
device_rank=i,
world_size=world_size,
start_layer=0,
end_layer=card.n_layers,
n_layers=card.n_layers,
)
for i, host in enumerate(test.devs)
},
),
)
if __name__ == "__main__":
anyio.run(main)
-85
View File
@@ -1,85 +0,0 @@
#!/usr/bin/env python3
import itertools
import json
import subprocess
import sys
from concurrent.futures import ThreadPoolExecutor
from typing import Any, cast
from urllib.request import Request, urlopen
if not (args := sys.argv[1:]):
sys.exit(
f"USAGE: {sys.argv[0]} <kind> [host1] [host2] ...\nkind is optional, and should be jaccl or ring"
)
kind = args[0] if args[0] in ("jaccl", "ring") else "both"
hosts = args[1:] if kind != "both" else args
ts = subprocess.run(
["tailscale", "status"], check=True, text=True, capture_output=True
).stdout.splitlines()
ip = {sl[1]: sl[0] for line in ts if len(sl := line.split()) >= 2}
ips = [ip[h] for h in hosts]
devs = [[h, ip[h]] for h in hosts]
n = len(hosts)
def get_tb(a: str) -> list[dict[str, Any]]:
with urlopen(f"http://{a}:52414/tb_detection", timeout=5) as r: # pyright: ignore[reportAny]
return json.loads(r.read()) # pyright: ignore[reportAny]
def get_models(a: str) -> set[str]:
with urlopen(f"http://{a}:52414/models", timeout=5) as r: # pyright: ignore[reportAny]
return set(json.loads(r.read())) # pyright: ignore[reportAny]
def run(h: str, a: str, body: bytes) -> None:
with urlopen(
Request(
f"http://{a}:52414/run_test",
data=body,
method="POST",
headers={"Content-Type": "application/json"},
),
timeout=300,
) as r: # pyright: ignore[reportAny]
for line in r.read().decode(errors="replace").splitlines(): # pyright: ignore[reportAny]
print(f"\n{h}@{a}: {line}", flush=True)
with ThreadPoolExecutor(n) as exctr:
if kind in ("jaccl", "both"):
payloads = list(exctr.map(get_tb, ips))
u2e = {
ident["domainUuid"]: (i, ident["rdmaInterface"])
for i, p in enumerate(payloads)
for d in p
for ident in cast(
list[dict[str, str]],
d.get("MacThunderboltIdentifiers", {}).get("idents", []), # pyright: ignore[reportAny]
)
}
edges = {
(u2e[s][0], u2e[t][0]): u2e[t][1]
for p in payloads
for d in p
for c in d.get("MacThunderboltConnections", {}).get("conns", []) # pyright: ignore[reportAny]
if (s := c["sourceUuid"]) in u2e and (t := c["sinkUuid"]) in u2e # pyright: ignore[reportAny]
}
ibv_devs = [[edges.get((i, j)) for j in range(n)] for i in range(n)]
else:
ibv_devs = None
models = set[str].intersection(*exctr.map(get_models, ips))
print("\n")
print("=" * 70)
print(f"Starting test with {models}")
print("=" * 70)
print("\n")
for model in models:
body = json.dumps(
{"devs": devs, "model_id": model, "ibv_devs": ibv_devs, "kind": kind}
).encode()
list(exctr.map(run, hosts, ips, itertools.repeat(body)))
+75
View File
@@ -0,0 +1,75 @@
# type: ignore
"""Single-node integration tests.
Run with:
uv run pytest tests/test_1node.py -v
"""
from __future__ import annotations
import time
import pytest
from exo_tools.harness import is_model_downloaded, place_instance
from .framework import DEFAULT_MODEL, InstanceSpec
@pytest.mark.cluster(count=1)
@pytest.mark.instance(DEFAULT_MODEL)
def test_place_instance_and_chat(session):
resp = session.chat("Say hello in one sentence.")
assert len(resp) > 0
@pytest.mark.cluster(count=1)
@pytest.mark.instance(DEFAULT_MODEL)
def test_chat_multiple_turns(session):
first_reply = session.chat("What is 2 + 2?")
assert len(first_reply) > 0
second_reply = session.multi_turn(
[
{"role": "user", "content": "What is 2 + 2?"},
{"role": "assistant", "content": first_reply},
{"role": "user", "content": "Now multiply that by 3."},
]
)
assert len(second_reply) > 0
@pytest.mark.cluster(count=1)
@pytest.mark.instance(DEFAULT_MODEL)
def test_delete_instance(session):
from exo_tools.harness import wait_for_instance_gone
session.client.request_json("DELETE", f"/instance/{session.instance_id}")
wait_for_instance_gone(session.client, session.instance_id, timeout=30.0)
assert len(session.instances) == 0, (
f"Expected no instances, found {len(session.instances)}"
)
@pytest.mark.cluster(count=1)
def test_download_from_scratch(session):
"""Ensure the model is not on the cluster, then place an instance to
trigger a fresh download and verify inference.
"""
node_id = next(iter(session.state.get("nodeIdentities", {})))
# Delete any existing download — the API call is idempotent
session.client.request_json("DELETE", f"/download/{node_id}/{DEFAULT_MODEL}")
# Poll until the model is gone (it may already be gone)
deadline = time.time() + 60.0
while time.time() < deadline:
if not is_model_downloaded(session.client, DEFAULT_MODEL):
break
time.sleep(2.0)
else:
raise AssertionError(f"Expected {DEFAULT_MODEL} to be deleted from cluster")
place_instance(session.client, DEFAULT_MODEL, timeout=900.0)
session.instance_spec = InstanceSpec(model_id=DEFAULT_MODEL)
resp = session.chat("Say hello in one sentence.")
assert len(resp) > 0
+49
View File
@@ -0,0 +1,49 @@
# type: ignore
"""Two-node integration tests (ring + jaccl parallelism).
Run with:
uv run pytest tests/test_2node.py -v
"""
from __future__ import annotations
import pytest
from exo_tools.cluster import Thunderbolt
from exo_tools.harness import Comm, Sharding
from .framework import DEFAULT_MODEL
@pytest.mark.cluster(count=2, thunderbolt=Thunderbolt.A2A)
@pytest.mark.instance(
DEFAULT_MODEL, sharding=Sharding.TENSOR, comm=Comm.JACCL, min_nodes=2
)
def test_2node_jaccl(session):
resp = session.chat("Say hello in one sentence.")
assert len(resp) > 0
@pytest.mark.cluster(count=2, thunderbolt=Thunderbolt.A2A)
@pytest.mark.instance(
DEFAULT_MODEL, sharding=Sharding.PIPELINE, comm=Comm.RING, min_nodes=2
)
def test_2node_ring(session):
resp = session.chat("Say hello in one sentence.")
assert len(resp) > 0
@pytest.mark.cluster(count=2, thunderbolt=Thunderbolt.A2A)
@pytest.mark.instance(
DEFAULT_MODEL, sharding=Sharding.TENSOR, comm=Comm.JACCL, min_nodes=2
)
def test_2node_jaccl_multi_turn(session):
first = session.chat("What is the capital of France?")
assert len(first) > 0
second = session.multi_turn(
[
{"role": "user", "content": "What is the capital of France?"},
{"role": "assistant", "content": first},
{"role": "user", "content": "What country is it in?"},
]
)
assert len(second) > 0
+32
View File
@@ -0,0 +1,32 @@
# type: ignore
"""Four-node integration tests.
Run with:
uv run pytest tests/test_4node.py -v
"""
from __future__ import annotations
import pytest
from exo_tools.cluster import Thunderbolt
from exo_tools.harness import Comm, Sharding
from .framework import DEFAULT_MODEL
@pytest.mark.cluster(count=4, thunderbolt=Thunderbolt.A2A)
@pytest.mark.instance(
DEFAULT_MODEL, sharding=Sharding.PIPELINE, comm=Comm.RING, min_nodes=4
)
def test_4node_pipeline_ring(session):
resp = session.chat("Say hello in one sentence.")
assert len(resp) > 0
@pytest.mark.cluster(count=4, thunderbolt=Thunderbolt.A2A)
@pytest.mark.instance(
DEFAULT_MODEL, sharding=Sharding.TENSOR, comm=Comm.JACCL, min_nodes=4
)
def test_4node_tensor_jaccl(session):
resp = session.chat("Say hello in one sentence.")
assert len(resp) > 0
+102
View File
@@ -0,0 +1,102 @@
# type: ignore
"""Dashboard end-to-end tests using Playwright (headless Chromium).
Prerequisites:
uv run playwright install chromium
Run with:
uv run pytest tests/test_dashboard.py -v
"""
from __future__ import annotations
import contextlib
import pytest
try:
from playwright.sync_api import sync_playwright
_HAS_PLAYWRIGHT = True
except ImportError:
_HAS_PLAYWRIGHT = False
# Check if Chromium is installed by attempting a quick launch
_HAS_CHROMIUM = False
if _HAS_PLAYWRIGHT:
try:
with sync_playwright() as p:
browser = p.chromium.launch(headless=True)
browser.close()
_HAS_CHROMIUM = True
except Exception:
pass
pytestmark = pytest.mark.skipif(
not _HAS_PLAYWRIGHT or not _HAS_CHROMIUM,
reason="playwright or chromium not installed (run: uv run playwright install chromium)",
)
def _mark_onboarding_complete(session) -> None:
"""Mark onboarding complete on the server so the wizard doesn't auto-launch a model."""
with contextlib.suppress(Exception):
session.client.request_json("POST", "/onboarding")
@pytest.mark.cluster(count=1)
def test_dashboard_chat_inference(session):
"""Full UI flow: open dashboard, pick a model, send a chat, verify response.
The instance is created via the dashboard UI (model picker chat send
triggers the dashboard's auto-launch flow), not via @pytest.mark.instance.
"""
_mark_onboarding_complete(session)
with sync_playwright() as p:
browser = p.chromium.launch(headless=True)
page = browser.new_page(viewport={"width": 1280, "height": 800})
page.goto(session.cluster.api_url, wait_until="networkidle")
page.wait_for_timeout(3000)
page.screenshot(path="/tmp/dashboard_initial.png")
# Open the model picker by clicking the "SELECT MODEL" button
page.get_by_text("SELECT MODEL", exact=False).first.click()
page.wait_for_timeout(1000)
page.screenshot(path="/tmp/dashboard_picker_open.png")
# Search for the model — uses the model id substring; the picker
# matches against name/id so "Llama-3.2-1B" filters to the small Llama.
search_input = page.locator('input[placeholder*="Search models"]').first
search_input.fill("Llama-3.2-1B")
page.wait_for_timeout(1500)
page.screenshot(path="/tmp/dashboard_picker_search.png")
# Click the only matching result. The picker shows the model's
# display name (e.g. "Llama 3.2 1B") which differs from the model_id.
# We click the first visible button-like row in the result list.
page.get_by_text("Llama 3.2 1B", exact=False).first.click()
page.wait_for_timeout(1500)
page.screenshot(path="/tmp/dashboard_model_selected.png")
# Type a chat message — sending triggers the dashboard's auto-launch
# flow: it picks an optimal placement for the selected model and POSTs
# to /instance, then sends the chat once the runner is ready.
chat_input = page.locator("textarea").first
chat_input.fill("Say hello")
chat_input.press("Enter")
page.screenshot(path="/tmp/dashboard_chat_sent.png")
# Wait for the instance to launch and respond. Generous timeout
# because this includes model placement + load + generation.
page.wait_for_timeout(60000)
page.screenshot(path="/tmp/dashboard_after_chat.png")
# Verify an instance was created and the chat got a response
instances = session.client.request_json("GET", "/state").get("instances", {})
assert len(instances) > 0, "Expected the dashboard to have created an instance"
body_text = page.text_content("body") or ""
assert len(body_text) > 0
browser.close()
+56
View File
@@ -0,0 +1,56 @@
# type: ignore
"""Resilience tests: disconnect/reconnect nodes and verify cluster recovery.
Run with:
uv run pytest tests/test_resilience.py -v
"""
from __future__ import annotations
import pytest
from exo_tools.cluster import Thunderbolt
from exo_tools.harness import Comm, Sharding, cleanup_all_instances, place_instance
from .framework import DEFAULT_MODEL, InstanceSpec
@pytest.mark.cluster(count=2, thunderbolt=Thunderbolt.A2A)
@pytest.mark.instance(
DEFAULT_MODEL, sharding=Sharding.PIPELINE, comm=Comm.RING, min_nodes=2
)
def test_node_recovery(session):
"""Full disconnect/reconnect cycle.
1. Place a 2-node instance, verify inference
2. Disconnect one node
3. Place a 1-node instance on remaining node, verify inference
4. Reconnect the stopped node, wait for the cluster to reform
5. Place a 2-node instance again, verify inference
"""
# --- Phase 1: 2-node inference ---
resp = session.chat("Hello")
assert len(resp) > 0
# --- Phase 2: disconnect one node ---
session.disconnect_node(1)
session.wait_ready(60)
# Clean up the now-broken 2-node instance
cleanup_all_instances(session.client)
# --- Phase 3: 1-node inference on the remaining node ---
place_instance(session.client, DEFAULT_MODEL, min_nodes=1)
session.instance_spec = InstanceSpec(model_id=DEFAULT_MODEL, min_nodes=1)
resp = session.chat("Hello")
assert len(resp) > 0
# --- Phase 4: reconnect and restore 2-node cluster ---
cleanup_all_instances(session.client)
session.reconnect_node(1)
session.wait_ready(60)
# --- Phase 5: 2-node inference again ---
place_instance(session.client, DEFAULT_MODEL, min_nodes=2)
session.instance_spec = InstanceSpec(model_id=DEFAULT_MODEL, min_nodes=2)
resp = session.chat("Hello again")
assert len(resp) > 0
File renamed without changes.
File renamed without changes.
File renamed without changes.
File renamed without changes.
+15
View File
@@ -0,0 +1,15 @@
import anyio
from exo_net import StateProxy
async def main():
sp = await StateProxy.init()
while True:
data = await sp.snapshot()
if data != "{}":
print(data)
await anyio.sleep(1)
if __name__ == "__main__":
anyio.run(main)
+10
View File
@@ -0,0 +1,10 @@
[project]
name = "exo-tools"
version = "0.1.0"
description = "Shared tooling for interacting with exo clusters"
requires-python = ">=3.13"
dependencies = ["loguru>=0.7.3"]
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
View File
Whitespace-only changes.
+117
View File
@@ -0,0 +1,117 @@
# type: ignore
"""HTTP client for the exo API."""
from __future__ import annotations
import http.client
import json
from collections.abc import Iterator
from typing import Any
from urllib.parse import urlencode
class ExoHttpError(RuntimeError):
def __init__(self, status: int, reason: str, body_preview: str):
super().__init__(f"HTTP {status} {reason}: {body_preview}")
self.status = status
class ExoClient:
def __init__(self, host: str, port: int, timeout_s: float = 7200.0):
self.host = host
self.port = port
self.timeout_s = timeout_s
def request_json(
self,
method: str,
path: str,
params: dict[str, Any] | None = None,
body: dict[str, Any] | None = None,
headers: dict[str, str] | None = None,
) -> Any:
if not path.startswith("/"):
path = "/" + path
if params:
path = path + "?" + urlencode(params)
conn = http.client.HTTPConnection(self.host, self.port, timeout=self.timeout_s)
try:
payload: bytes | None = None
hdrs: dict[str, str] = {"Accept": "application/json"}
if body is not None:
payload = json.dumps(body).encode("utf-8")
hdrs["Content-Type"] = "application/json"
if headers:
hdrs.update(headers)
conn.request(method.upper(), path, body=payload, headers=hdrs)
resp = conn.getresponse()
raw = resp.read()
text = raw.decode("utf-8", errors="replace") if raw else ""
if resp.status >= 400:
raise ExoHttpError(resp.status, resp.reason, text[:300])
if not text:
return None
return json.loads(text)
finally:
conn.close()
def post_bench_chat_completions(self, payload: dict[str, Any]) -> dict[str, Any]:
return self.request_json("POST", "/bench/chat/completions", body=payload)
def stream_bench_chat_completions(self, payload: dict[str, Any]) -> Iterator[str]:
"""POST /bench/chat/completions with stream=True, yielding raw SSE lines."""
payload = {**payload, "stream": True}
data = json.dumps(payload).encode("utf-8")
conn = http.client.HTTPConnection(self.host, self.port, timeout=self.timeout_s)
try:
conn.request(
"POST",
"/bench/chat/completions",
body=data,
headers={
"Content-Type": "application/json",
"Accept": "text/event-stream",
},
)
resp = conn.getresponse()
if resp.status >= 400:
raw = resp.read().decode("utf-8", errors="replace")
raise ExoHttpError(resp.status, resp.reason, raw[:300])
for line in resp:
yield line.decode("utf-8", errors="replace")
finally:
conn.close()
def get_state_path(self, path: str) -> Any:
try:
return self.request_json("GET", f"/state/{path}")
except ExoHttpError as e:
if e.status == 404:
return None
raise
def get_instance(self, instance_id: str) -> dict[str, Any] | None:
return self.get_state_path(f"instances/{instance_id}")
def get_runner(self, runner_id: str) -> dict[str, Any] | None:
return self.get_state_path(f"runners/{runner_id}")
def get_node_downloads(self, node_id: str) -> list[dict[str, Any]] | None:
return self.get_state_path(f"downloads/{node_id}")
def get_node_disk(self, node_id: str) -> dict[str, Any] | None:
return self.get_state_path(f"nodeDisk/{node_id}")
def get_node_system(self, node_id: str) -> dict[str, Any] | None:
return self.get_state_path(f"nodeSystem/{node_id}")
def get_node_identities(self) -> dict[str, Any] | None:
return self.get_state_path("nodeIdentities")
def get_topology(self) -> dict[str, Any] | None:
return self.get_state_path("topology")
+243
View File
@@ -0,0 +1,243 @@
# type: ignore
"""Cluster lifecycle management via eco.
Provides subprocess wrappers for eco commands (deploy, stop, start, release,
logs, exec) and a ClusterInfo dataclass. Reusable by integration tests,
bench, eval, and CI workflows.
"""
from __future__ import annotations
import atexit
import contextlib
import json
import logging
import os
import signal
import subprocess
import uuid
from dataclasses import dataclass, field
from enum import Enum
from .client import ExoClient
class Thunderbolt(str, Enum):
A2A = "a2a" # all-to-all (eco --tb-a2a)
RING = "ring" # ring topology (eco --tb-ring)
class Chip(str, Enum):
M1 = "M1"
M1_PRO = "M1 Pro"
M1_MAX = "M1 Max"
M1_ULTRA = "M1 Ultra"
M2 = "M2"
M2_PRO = "M2 Pro"
M2_MAX = "M2 Max"
M2_ULTRA = "M2 Ultra"
M3 = "M3"
M3_PRO = "M3 Pro"
M3_MAX = "M3 Max"
M3_ULTRA = "M3 Ultra"
M4 = "M4"
M4_PRO = "M4 Pro"
M4_MAX = "M4 Max"
M4_ULTRA = "M4 Ultra"
logger = logging.getLogger("exo_tools.cluster")
# When set, deploy from a GitHub branch/tag instead of local source (rsync).
_EXO_REF = os.environ.get("EXO_REF")
@dataclass
class ClusterInfo:
"""Holds the result of an `eco start --deploy` invocation."""
hosts: list[str]
namespace: str
api_endpoints: dict[str, str] # host -> url
api_url: str # primary endpoint for ExoClient
primary_host: str = ""
_host: str = field(init=False, repr=False, default="")
_port: int = field(init=False, repr=False, default=52415)
def __post_init__(self) -> None:
if not self.primary_host:
self.primary_host = self.hosts[0]
url = self.api_url.replace("http://", "").replace("https://", "")
parts = url.split(":")
self._host = parts[0]
self._port = int(parts[1]) if len(parts) > 1 else 52415
def make_client(self, timeout_s: float = 7200.0) -> ExoClient:
return ExoClient(self._host, self._port, timeout_s=timeout_s)
class EcoSession:
"""Manages an eco session with a unique user and automatic cleanup.
Usage:
session = EcoSession(user_prefix="test")
cluster = session.start_deploy(count=2, thunderbolt=True)
...
session.stop_all() # or let atexit handle it
The session registers atexit and signal handlers to ensure cleanup
on normal exit, uncaught exceptions, SIGTERM, and SIGHUP. SIGINT
is left unhandled so KeyboardInterrupt propagates normally.
"""
def __init__(self, user_prefix: str = "test") -> None:
self._session_id = uuid.uuid4().hex[:8]
self.user = f"{user_prefix}-{self._session_id}"
self._env = {**os.environ, "USER": self.user}
# Register cleanup handlers
atexit.register(self.stop_all)
for sig in (signal.SIGTERM, signal.SIGHUP):
signal.signal(sig, self._signal_handler)
def _signal_handler(self, signum: int, _frame: object) -> None:
self.stop_all()
raise SystemExit(128 + signum)
def stop_all(self) -> None:
"""Stop all clusters and release all reservations for this session."""
with contextlib.suppress(Exception):
subprocess.run(
["eco", "stop"],
capture_output=True,
text=True,
timeout=30,
env=self._env,
)
def _run(
self, args: list[str], *, check: bool = True, timeout: int = 120
) -> subprocess.CompletedProcess[str]:
"""Run an eco command as this session's user.
stdout is captured (JSON output), stderr is passed through to the
console so eco's progress messages are visible.
"""
logger.info(f"eco: {' '.join(args)}")
return subprocess.run(
args,
stdout=subprocess.PIPE,
stderr=None,
text=True,
check=check,
timeout=timeout,
env=self._env,
)
def start_deploy(
self,
hosts: list[str] | None = None,
*,
count: int | None = None,
thunderbolt: Thunderbolt | None = None,
chip: Chip | None = None,
min_memory_gb: float | None = None,
wait: bool = True,
ref: str | None = _EXO_REF,
timeout: int = 600,
) -> ClusterInfo:
"""Start and deploy exo on a set of hosts via eco.
By default, deploys from local source via rsync. Set EXO_REF
or pass ref= to deploy from a GitHub branch/tag instead (for CI).
"""
cmd: list[str] = ["eco", "--json", "start", "--deploy"]
if hosts:
cmd.extend(hosts)
if count is not None:
cmd.extend(["--count", str(count)])
if thunderbolt is not None:
cmd.append(f"--tb-{thunderbolt.value}")
if chip is not None:
cmd.extend(["--chip", chip.value])
if min_memory_gb is not None:
cmd.extend(["--min-memory", str(min_memory_gb)])
if wait:
cmd.append("--wait")
if ref:
cmd.extend(["--ref", ref])
result = self._run(cmd, timeout=timeout)
data = json.loads(result.stdout)["data"]
endpoints: dict[str, str] = data["api_endpoints"]
primary_host = data["hosts"][0]
return ClusterInfo(
hosts=data["hosts"],
namespace=data["namespace"],
api_endpoints=endpoints,
api_url=endpoints[primary_host],
primary_host=primary_host,
)
def stop(self, hosts: list[str], *, keep: bool = False, timeout: int = 120) -> None:
"""Stop exo on the given hosts. If keep=True, keep the reservation."""
cmd: list[str] = ["eco", "stop"]
cmd.extend(hosts)
if keep:
cmd.append("--keep")
self._run(cmd, timeout=timeout)
def start_hosts(
self, hosts: list[str], *, namespace: str, timeout: int = 300
) -> None:
"""Start (previously stopped) hosts back into an existing namespace."""
cmd: list[str] = ["eco", "--json", "start"]
cmd.extend(hosts)
cmd.extend(["--namespace", namespace])
self._run(cmd, timeout=timeout)
def release(self, hosts: list[str], timeout: int = 120) -> None:
"""Release hosts from the reservation."""
cmd: list[str] = ["eco", "release"]
cmd.extend(hosts)
self._run(cmd, timeout=timeout)
def logs(
self, hosts: list[str], lines: int = 500, timeout: int = 60
) -> dict[str, list[str]]:
"""Fetch recent logs from cluster hosts."""
cmd: list[str] = ["eco", "--json", "logs"]
cmd.extend(hosts)
cmd.extend(["-n", str(lines), "--raw"])
result = self._run(cmd, check=False, timeout=timeout)
if result.returncode != 0:
return {"_error": [result.stderr]}
try:
return json.loads(result.stdout)
except json.JSONDecodeError:
return {"_raw": result.stdout.splitlines()}
def exec(self, hosts: list[str], command: str, timeout: int = 120) -> str:
"""Run an arbitrary command on the given hosts via eco."""
cmd: list[str] = ["eco", "exec"]
cmd.extend(hosts)
cmd.append("--")
cmd.extend(command.split())
result = self._run(cmd, check=False, timeout=timeout)
return result.stdout
def make_client(cluster: ClusterInfo, timeout_s: float = 7200.0) -> ExoClient:
"""Create an ExoClient from a ClusterInfo."""
return cluster.make_client(timeout_s=timeout_s)
def make_client_from_url(url: str, timeout_s: float = 7200.0) -> ExoClient:
"""Create an ExoClient from a URL string like 'http://host:port'."""
url_clean = url.replace("http://", "").replace("https://", "")
parts = url_clean.split(":")
host = parts[0]
port = int(parts[1]) if len(parts) > 1 else 52415
return ExoClient(host, port, timeout_s=timeout_s)
@@ -1,129 +1,39 @@
# type: ignore
"""Instance lifecycle helpers for exo clusters.
Provides utilities for placing instances, waiting for readiness,
managing downloads, filtering placements, and common CLI arguments.
"""
from __future__ import annotations
import argparse
import http.client
import json
import contextlib
import os
import time
from collections.abc import Iterator
from enum import Enum
from typing import Any
from urllib.parse import urlencode
from loguru import logger
from .client import ExoClient, ExoHttpError
class Sharding(str, Enum):
PIPELINE = "Pipeline" # layers split across nodes
TENSOR = "Tensor" # layers split within (across nodes)
class Comm(str, Enum):
RING = "MlxRing" # ring all-reduce over network
JACCL = "MlxJaccl" # RDMA over Thunderbolt
_SETTLE_INITIAL_BACKOFF_S = 1.0
_SETTLE_MAX_BACKOFF_S = 60.0
_SETTLE_BACKOFF_MULTIPLIER = 2.0
class ExoHttpError(RuntimeError):
def __init__(self, status: int, reason: str, body_preview: str):
super().__init__(f"HTTP {status} {reason}: {body_preview}")
self.status = status
class ExoClient:
def __init__(self, host: str, port: int, timeout_s: float = 7200.0):
self.host = host
self.port = port
self.timeout_s = timeout_s
def request_json(
self,
method: str,
path: str,
params: dict[str, Any] | None = None,
body: dict[str, Any] | None = None,
headers: dict[str, str] | None = None,
) -> Any:
if not path.startswith("/"):
path = "/" + path
if params:
path = path + "?" + urlencode(params)
conn = http.client.HTTPConnection(self.host, self.port, timeout=self.timeout_s)
try:
payload: bytes | None = None
hdrs: dict[str, str] = {"Accept": "application/json"}
if body is not None:
payload = json.dumps(body).encode("utf-8")
hdrs["Content-Type"] = "application/json"
if headers:
hdrs.update(headers)
conn.request(method.upper(), path, body=payload, headers=hdrs)
resp = conn.getresponse()
raw = resp.read()
text = raw.decode("utf-8", errors="replace") if raw else ""
if resp.status >= 400:
raise ExoHttpError(resp.status, resp.reason, text[:300])
if not text:
return None
return json.loads(text)
finally:
conn.close()
def post_bench_chat_completions(self, payload: dict[str, Any]) -> dict[str, Any]:
return self.request_json("POST", "/bench/chat/completions", body=payload)
def stream_bench_chat_completions(self, payload: dict[str, Any]) -> Iterator[str]:
"""POST /bench/chat/completions with stream=True, yielding raw SSE lines."""
payload = {**payload, "stream": True}
data = json.dumps(payload).encode("utf-8")
conn = http.client.HTTPConnection(self.host, self.port, timeout=self.timeout_s)
try:
conn.request(
"POST",
"/bench/chat/completions",
body=data,
headers={
"Content-Type": "application/json",
"Accept": "text/event-stream",
},
)
resp = conn.getresponse()
if resp.status >= 400:
raw = resp.read().decode("utf-8", errors="replace")
raise ExoHttpError(resp.status, resp.reason, raw[:300])
for line in resp:
yield line.decode("utf-8", errors="replace")
finally:
conn.close()
def get_state_path(self, path: str) -> Any:
try:
return self.request_json("GET", f"/state/{path}")
except ExoHttpError as e:
if e.status == 404:
return None
raise
def get_instance(self, instance_id: str) -> dict[str, Any] | None:
return self.get_state_path(f"instances/{instance_id}")
def get_runner(self, runner_id: str) -> dict[str, Any] | None:
return self.get_state_path(f"runners/{runner_id}")
def get_node_downloads(self, node_id: str) -> list[dict[str, Any]] | None:
return self.get_state_path(f"downloads/{node_id}")
def get_node_disk(self, node_id: str) -> dict[str, Any] | None:
return self.get_state_path(f"nodeDisk/{node_id}")
def get_node_system(self, node_id: str) -> dict[str, Any] | None:
return self.get_state_path(f"nodeSystem/{node_id}")
def get_node_identities(self) -> dict[str, Any] | None:
return self.get_state_path("nodeIdentities")
def get_topology(self) -> dict[str, Any] | None:
return self.get_state_path("topology")
def unwrap_instance(instance: dict[str, Any]) -> dict[str, Any]:
if len(instance) != 1:
raise KeyError(f"Expected 1 key, got keys={list(instance.keys())}")
@@ -555,7 +465,6 @@ def find_existing_instance(client: ExoClient, model_id: str) -> str | None:
except Exception:
return None
for inst_id, inst in state.get("instances", {}).items():
# Instance structure is nested: {"MlxJacclInstance": {"shardAssignments": {"modelId": ...}}}
for _inst_type, inner in inst.items():
if not isinstance(inner, dict):
continue
@@ -623,3 +532,112 @@ def add_common_instance_args(ap: argparse.ArgumentParser) -> None:
action="store_true",
help="Reuse an existing running instance for this model instead of creating a new one.",
)
# ---------------------------------------------------------------------------
# Cluster/instance orchestration helpers (used by tests, bench, eval)
# ---------------------------------------------------------------------------
def get_instance_ids(client: ExoClient) -> set[str]:
"""Return the set of current instance IDs from cluster state."""
state = client.request_json("GET", "/state") or {}
result: set[str] = set()
for instance in state.get("instances", {}).values():
with contextlib.suppress(Exception):
result.add(instance_id_from_instance(instance))
return result
def wait_for_cluster_ready(
client: ExoClient, expected_nodes: int = 1, timeout: float = 120.0
) -> None:
"""Wait until the cluster has all expected nodes visible and reporting memory.
Placement requires nodeMemory for all nodes in a cycle. This polls until
both nodeIdentities and nodeMemory have at least `expected_nodes` entries.
"""
start = time.time()
while time.time() - start < timeout:
try:
state = client.request_json("GET", "/state") or {}
if (
len(state.get("nodeIdentities", {})) >= expected_nodes
and len(state.get("nodeMemory", {})) >= expected_nodes
):
return
except Exception:
pass
time.sleep(1.0)
raise TimeoutError(f"Cluster not ready: expected {expected_nodes} nodes")
def place_instance(
client: ExoClient,
model_id: str,
*,
sharding: Sharding = Sharding.PIPELINE,
comm: Comm = Comm.RING,
min_nodes: int = 1,
timeout: float = 600.0,
placement_retries: int = 10,
placement_retry_delay: float = 10.0,
) -> str:
"""Place an instance and wait for it to be ready. Returns the instance_id.
The /place_instance API returns a command_id, but instances are stored
under a separately-generated instance_id. This polls cluster state for the
new instance, retrying placement if the cluster is still settling.
"""
wait_for_cluster_ready(client, expected_nodes=min_nodes)
body = {
"model_id": model_id,
"sharding": sharding.value,
"instance_meta": comm.value,
"min_nodes": min_nodes,
}
instance_id: str | None = None
for attempt in range(placement_retries):
before_ids = get_instance_ids(client)
client.request_json("POST", "/place_instance", body=body)
poll_deadline = time.time() + 30.0
while time.time() < poll_deadline:
new_ids = get_instance_ids(client) - before_ids
if new_ids:
instance_id = next(iter(new_ids))
break
time.sleep(1.0)
if instance_id is not None:
break
if attempt < placement_retries - 1:
time.sleep(placement_retry_delay)
if instance_id is None:
raise TimeoutError(
f"Placement failed after {placement_retries} attempts "
f"({sharding.value}/{comm.value} for {model_id})"
)
wait_for_instance_ready(client, instance_id, timeout=timeout)
return instance_id
def cleanup_all_instances(client: ExoClient) -> None:
"""Remove all running instances from the cluster."""
state = client.request_json("GET", "/state") or {}
for instance in state.get("instances", {}).values():
with contextlib.suppress(Exception):
iid = instance_id_from_instance(instance)
client.request_json("DELETE", f"/instance/{iid}")
wait_for_instance_gone(client, iid, timeout=30.0)
def is_model_downloaded(client: ExoClient, model_id: str) -> bool:
response = client.request_json("GET", "/models", params={"status": "downloaded"})
data = (response or {}).get("data", [])
return all(model.get("id") == model_id for model in data)
Generated
+79 -16
View File
@@ -22,7 +22,8 @@ prerelease-mode = "allow"
members = [
"exo",
"exo-bench",
"exo-pyo3-bindings",
"exo-net",
"exo-tools",
]
constraints = [{ name = "transformers", specifier = ">=5.6.2" }]
overrides = [
@@ -385,7 +386,7 @@ dependencies = [
{ name = "aiofiles", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "aiohttp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "anyio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "exo-pyo3-bindings", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "exo-net", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "fastapi", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "filelock", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "httpx", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -394,7 +395,7 @@ dependencies = [
{ name = "loguru", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mflux", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.1", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260427+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260429+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-lm", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-vlm", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "msgspec", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -416,21 +417,21 @@ build = [
]
cpu = [
{ name = "mlx", version = "0.31.1", source = { registry = "https://pypi.org/simple" }, marker = "(sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260427+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "(sys_platform == 'darwin' and extra == 'extra-3-exo-cpu') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260429+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "(sys_platform == 'darwin' and extra == 'extra-3-exo-cpu') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-cpu", marker = "sys_platform == 'linux'" },
{ name = "mlx-lm", marker = "sys_platform == 'linux'" },
{ name = "mlx-vlm", marker = "sys_platform == 'linux'" },
]
cuda12 = [
{ name = "mlx", version = "0.31.1", source = { registry = "https://pypi.org/simple" }, marker = "(sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260427+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "(sys_platform == 'darwin' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260429+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "(sys_platform == 'darwin' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-cuda-12", marker = "sys_platform == 'linux'" },
{ name = "mlx-lm", marker = "sys_platform == 'linux'" },
{ name = "mlx-vlm", marker = "sys_platform == 'linux'" },
]
cuda13 = [
{ name = "mlx", version = "0.31.1", source = { registry = "https://pypi.org/simple" }, marker = "(sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260427+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "(sys_platform == 'darwin' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260429+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "(sys_platform == 'darwin' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-cuda-13", marker = "sys_platform == 'linux'" },
{ name = "mlx-lm", marker = "sys_platform == 'linux'" },
{ name = "mlx-vlm", marker = "sys_platform == 'linux'" },
@@ -439,6 +440,7 @@ cuda13 = [
[package.dev-dependencies]
dev = [
{ name = "basedpyright", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "playwright", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "pyinstaller", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "pytest", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "pytest-asyncio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -451,7 +453,7 @@ requires-dist = [
{ name = "aiofiles", specifier = ">=24.1.0" },
{ name = "aiohttp", specifier = ">=3.12.14" },
{ name = "anyio", specifier = "==4.11.0" },
{ name = "exo-pyo3-bindings", editable = "rust/exo_pyo3_bindings" },
{ name = "exo-net", editable = "rust/exo_net" },
{ name = "fastapi", specifier = ">=0.116.1" },
{ name = "filelock", specifier = ">=3.18.0" },
{ name = "httpx", specifier = ">=0.28.1" },
@@ -498,6 +500,7 @@ provides-extras = ["build", "cpu", "cuda12", "cuda13"]
[package.metadata.requires-dev]
dev = [
{ name = "basedpyright", specifier = ">=1.29.0" },
{ name = "playwright", specifier = ">=1.52.0" },
{ name = "pyinstaller", specifier = ">=6.17.0" },
{ name = "pytest", specifier = ">=8.4.0" },
{ name = "pytest-asyncio", specifier = ">=1.0.0" },
@@ -541,13 +544,13 @@ requires-dist = [
]
[[package]]
name = "exo-pyo3-bindings"
version = "0.2.1"
source = { editable = "rust/exo_pyo3_bindings" }
name = "exo-net"
version = "0.3.0"
source = { editable = "rust/exo_net" }
[package.dev-dependencies]
dev = [
{ name = "exo-pyo3-bindings", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "exo-net", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "pytest", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "pytest-asyncio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
]
@@ -556,11 +559,22 @@ dev = [
[package.metadata.requires-dev]
dev = [
{ name = "exo-pyo3-bindings", editable = "rust/exo_pyo3_bindings" },
{ name = "exo-net", editable = "rust/exo_net" },
{ name = "pytest", specifier = ">=8.4.0" },
{ name = "pytest-asyncio", specifier = ">=1.0.0" },
]
[[package]]
name = "exo-tools"
version = "0.1.0"
source = { editable = "tools" }
dependencies = [
{ name = "loguru", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
]
[package.metadata]
requires-dist = [{ name = "loguru", specifier = ">=0.7.3" }]
[[package]]
name = "fastapi"
version = "0.128.0"
@@ -669,6 +683,24 @@ http = [
{ name = "aiohttp", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
]
[[package]]
name = "greenlet"
version = "3.5.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/3c/3f/dbf99fb14bfeb88c28f16729215478c0e265cacd6dc22270c8f31bb6892f/greenlet-3.5.0.tar.gz", hash = "sha256:d419647372241bc68e957bf38d5c1f98852155e4146bd1e4121adea81f4f01e4", size = 196995, upload-time = "2026-04-27T13:37:15.544Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/0c/58/fc576f99037ce19c5aa16628e4c3226b6d1419f72a62c79f5f40576e6eb3/greenlet-3.5.0-cp313-cp313-macosx_11_0_universal2.whl", hash = "sha256:5a5ed18de6a0f6cc7087f1563f6bd93fc7df1c19165ca01e9bde5a5dc281d106", size = 285066, upload-time = "2026-04-27T12:23:05.033Z" },
{ url = "https://files.pythonhosted.org/packages/4a/ba/b28ddbe6bfad6a8ac196ef0e8cff37bc65b79735995b9e410923fffeeb70/greenlet-3.5.0-cp313-cp313-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:3a717fbc46d8a354fa675f7c1e813485b6ba3885f9bef0cd56e5ba27d758ff5b", size = 604414, upload-time = "2026-04-27T12:52:42.358Z" },
{ url = "https://files.pythonhosted.org/packages/09/06/4b69f8f0b67603a8be2790e55107a190b376f2627fe0eaf5695d85ffb3cd/greenlet-3.5.0-cp313-cp313-manylinux_2_24_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:ddc090c5c1792b10246a78e8c2163ebbe04cf877f9d785c230a7b27b39ad038e", size = 617349, upload-time = "2026-04-27T12:59:43.32Z" },
{ url = "https://files.pythonhosted.org/packages/6a/15/a643b4ecd09969e30b8a150d5919960caae0abe4f5af75ab040b1ab85e78/greenlet-3.5.0-cp313-cp313-manylinux_2_24_s390x.manylinux_2_28_s390x.whl", hash = "sha256:4964101b8585c144cbda5532b1aa644255126c08a265dae90c16e7a0e63aaa9d", size = 623234, upload-time = "2026-04-27T13:02:40.611Z" },
{ url = "https://files.pythonhosted.org/packages/8a/17/a3918541fd0ddefe024a69de6d16aa7b46d36ac19562adaa63c7fa180eff/greenlet-3.5.0-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:2094acd54b272cb6eae8c03dd87b3fa1820a4cef18d6889c378d503500a1dc13", size = 613927, upload-time = "2026-04-27T12:25:30.28Z" },
{ url = "https://files.pythonhosted.org/packages/77/18/3b13d5ef1275b0ffaf933b05efa21408ac4ca95823c7411d79682e4fdcff/greenlet-3.5.0-cp313-cp313-manylinux_2_39_riscv64.whl", hash = "sha256:7022615368890680e67b9965d33f5773aade330d5343bbe25560135aaa849eae", size = 425243, upload-time = "2026-04-27T13:05:15.689Z" },
{ url = "https://files.pythonhosted.org/packages/ee/e1/bd0af6213c7dd33175d8a462d4c1fe1175124ebed4855bc1475a5b5242c2/greenlet-3.5.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:5e05ba267789ea87b5a155cf0e810b1ab88bf18e9e8740813945ceb8ee4350ba", size = 1570893, upload-time = "2026-04-27T12:53:29.483Z" },
{ url = "https://files.pythonhosted.org/packages/9b/2a/0789702f864f5382cb476b93d7a9c823c10472658102ccd65f415747d2e2/greenlet-3.5.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:0ecec963079cd58cbd14723582384f11f166fd58883c15dcbfb342e0bc9b5846", size = 1636060, upload-time = "2026-04-27T12:25:28.845Z" },
{ url = "https://files.pythonhosted.org/packages/b2/8f/22bf9df92bbff0eb07842b60f7e63bf7675a9742df628437a9f02d09137f/greenlet-3.5.0-cp313-cp313-win_amd64.whl", hash = "sha256:728d9667d8f2f586644b748dbd9bb67e50d6a9381767d1357714ea6825bb3bf5", size = 238740, upload-time = "2026-04-27T12:24:01.341Z" },
{ url = "https://files.pythonhosted.org/packages/b6/b7/9c5c3d653bd4ff614277c049ac676422e2c557db47b4fe43e6313fc005dc/greenlet-3.5.0-cp313-cp313-win_arm64.whl", hash = "sha256:47422135b1d308c14b2c6e758beedb1acd33bb91679f5670edf77bf46244722b", size = 235525, upload-time = "2026-04-27T12:23:12.308Z" },
]
[[package]]
name = "h11"
version = "0.16.0"
@@ -1213,7 +1245,7 @@ dependencies = [
{ name = "hf-transfer", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "huggingface-hub", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "matplotlib", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260427+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260429+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "numpy", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "opencv-python", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "piexif", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -1263,7 +1295,7 @@ wheels = [
[[package]]
name = "mlx"
version = "0.32.0.dev20260427+cc3f3e60"
version = "0.32.0.dev20260429+cc3f3e60"
source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }
resolution-markers = [
"sys_platform == 'darwin'",
@@ -1315,7 +1347,7 @@ source = { git = "https://github.com/rltakashige/mlx-lm?branch=leo%2Fdeepseek-v4
dependencies = [
{ name = "jinja2", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.1", source = { registry = "https://pypi.org/simple" }, marker = "(sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260427+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260429+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "numpy", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "protobuf", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "pyyaml", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -1332,7 +1364,7 @@ dependencies = [
{ name = "fastapi", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "miniaudio", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.31.1", source = { registry = "https://pypi.org/simple" }, marker = "(sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260427+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx", version = "0.32.0.dev20260429+cc3f3e60", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#cc3f3e60be1289506125f2fa19b73b05aa770df8" }, marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "mlx-lm", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "numpy", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "opencv-python", marker = "sys_platform == 'darwin' or (sys_platform == 'linux' and extra == 'extra-3-exo-cpu') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda12') or (sys_platform == 'linux' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
@@ -1768,6 +1800,25 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/cb/28/3bfe2fa5a7b9c46fe7e13c97bda14c895fb10fa2ebf1d0abb90e0cea7ee1/platformdirs-4.5.1-py3-none-any.whl", hash = "sha256:d03afa3963c806a9bed9d5125c8f4cb2fdaf74a55ab60e5d59b3fde758104d31", size = 18731, upload-time = "2025-12-05T13:52:56.823Z" },
]
[[package]]
name = "playwright"
version = "1.58.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "greenlet", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
{ name = "pyee", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/f8/c9/9c6061d5703267f1baae6a4647bfd1862e386fbfdb97d889f6f6ae9e3f64/playwright-1.58.0-py3-none-macosx_10_13_x86_64.whl", hash = "sha256:96e3204aac292ee639edbfdef6298b4be2ea0a55a16b7068df91adac077cc606", size = 42251098, upload-time = "2026-01-30T15:09:24.028Z" },
{ url = "https://files.pythonhosted.org/packages/e0/40/59d34a756e02f8c670f0fee987d46f7ee53d05447d43cd114ca015cb168c/playwright-1.58.0-py3-none-macosx_11_0_arm64.whl", hash = "sha256:70c763694739d28df71ed578b9c8202bb83e8fe8fb9268c04dd13afe36301f71", size = 41039625, upload-time = "2026-01-30T15:09:27.558Z" },
{ url = "https://files.pythonhosted.org/packages/e1/ee/3ce6209c9c74a650aac9028c621f357a34ea5cd4d950700f8e2c4b7fe2c4/playwright-1.58.0-py3-none-macosx_11_0_universal2.whl", hash = "sha256:185e0132578733d02802dfddfbbc35f42be23a45ff49ccae5081f25952238117", size = 42251098, upload-time = "2026-01-30T15:09:30.461Z" },
{ url = "https://files.pythonhosted.org/packages/f1/af/009958cbf23fac551a940d34e3206e6c7eed2b8c940d0c3afd1feb0b0589/playwright-1.58.0-py3-none-manylinux1_x86_64.whl", hash = "sha256:c95568ba1eda83812598c1dc9be60b4406dffd60b149bc1536180ad108723d6b", size = 46235268, upload-time = "2026-01-30T15:09:33.787Z" },
{ url = "https://files.pythonhosted.org/packages/d9/a6/0e66ad04b6d3440dae73efb39540c5685c5fc95b17c8b29340b62abbd952/playwright-1.58.0-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:8f9999948f1ab541d98812de25e3a8c410776aa516d948807140aff797b4bffa", size = 45964214, upload-time = "2026-01-30T15:09:36.751Z" },
{ url = "https://files.pythonhosted.org/packages/0e/4b/236e60ab9f6d62ed0fd32150d61f1f494cefbf02304c0061e78ed80c1c32/playwright-1.58.0-py3-none-win32.whl", hash = "sha256:1e03be090e75a0fabbdaeab65ce17c308c425d879fa48bb1d7986f96bfad0b99", size = 36815998, upload-time = "2026-01-30T15:09:39.627Z" },
{ url = "https://files.pythonhosted.org/packages/41/f8/5ec599c5e59d2f2f336a05b4f318e733077cd5044f24adb6f86900c3e6a7/playwright-1.58.0-py3-none-win_amd64.whl", hash = "sha256:a2bf639d0ce33b3ba38de777e08697b0d8f3dc07ab6802e4ac53fb65e3907af8", size = 36816005, upload-time = "2026-01-30T15:09:42.449Z" },
{ url = "https://files.pythonhosted.org/packages/c8/c4/cc0229fea55c87d6c9c67fe44a21e2cd28d1d558a5478ed4d617e9fb0c93/playwright-1.58.0-py3-none-win_arm64.whl", hash = "sha256:32ffe5c303901a13a0ecab91d1c3f74baf73b84f4bedbb6b935f5bc11cc98e1b", size = 33085919, upload-time = "2026-01-30T15:09:45.71Z" },
]
[[package]]
name = "pluggy"
version = "1.6.0"
@@ -1941,6 +1992,18 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/73/7d/f2f9db34af103bea3e09735bb40b021788a5e834c81eedb541991badf8f5/pydantic_core-2.41.5-cp313-cp313-win_arm64.whl", hash = "sha256:3f84d5c1b4ab906093bdc1ff10484838aca54ef08de4afa9de0f5f14d69639cd", size = 1981005, upload-time = "2025-11-04T13:40:54.734Z" },
]
[[package]]
name = "pyee"
version = "13.0.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
]
sdist = { url = "https://files.pythonhosted.org/packages/8b/04/e7c1fe4dc78a6fdbfd6c337b1c3732ff543b8a397683ab38378447baa331/pyee-13.0.1.tar.gz", hash = "sha256:0b931f7c14535667ed4c7e0d531716368715e860b988770fc7eb8578d1f67fc8", size = 31655, upload-time = "2026-02-14T21:12:28.044Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/a0/c4/b4d4827c93ef43c01f599ef31453ccc1c132b353284fc6c87d535c233129/pyee-13.0.1-py3-none-any.whl", hash = "sha256:af2f8fede4171ef667dfded53f96e2ed0d6e6bd7ee3bb46437f77e3b57689228", size = 15659, upload-time = "2026-02-14T21:12:26.263Z" },
]
[[package]]
name = "pygments"
version = "2.19.2"