diff --git a/.gitignore b/.gitignore index 602a7a36..c7fcd3a7 100644 --- a/.gitignore +++ b/.gitignore @@ -19,7 +19,6 @@ probing.egg-info/ .idea/ app/dist/ dist/ -python/probing/bundled_web/ .python-version report.json pkg/ @@ -35,4 +34,6 @@ docs/site/ .coverage .pytest_cache/ web/target/ +probing/server/web-assets/ +/web/dist frontend/ diff --git a/Cargo.lock b/Cargo.lock index 36933ad2..a4e6268c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3440,7 +3440,6 @@ dependencies = [ "tempfile", "thiserror 2.0.12", "tokio", - "ureq", "url", "uuid", ] @@ -3628,6 +3627,7 @@ dependencies = [ "http-body-util", "hyper", "hyper-util", + "include_dir", "log", "nix 0.31.3", "nu-ansi-term", @@ -3663,6 +3663,7 @@ dependencies = [ "serde", "serde_json", "serde_yaml", + "sqlparser", "tokio", ] diff --git a/Cargo.toml b/Cargo.toml index 014a3145..9cdeb393 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -130,11 +130,11 @@ pyo3 = { version = "0.29.0", default-features = false, features = [ [build-dependencies] pyo3-build-config = "0.29.0" -# Dev (`DEBUG=1 make develop`, `cargo check`): prioritize compile speed + lean intermediates. -# debug=1 → line tables only (breakpoints/backtraces without full DWARF). +# Dev (`DEBUG=1 make develop`, `cargo check`): favor incremental compile speed +# without paying for split debug symbols. Opt in to symbols with +# CARGO_PROFILE_DEV_DEBUG=1. [profile.dev] -debug = 1 -split-debuginfo = "unpacked" +debug = 0 incremental = true codegen-units = 256 diff --git a/Makefile b/Makefile index 90183ecc..07dc0075 100644 --- a/Makefile +++ b/Makefile @@ -1,7 +1,7 @@ # Probing Makefile # # develop → maturin develop (Rust/Python daily loop) -# frontend → dx bundle → python/probing/bundled_web/public (web/dist symlink) +# frontend → dx bundle → probing/server/web-assets/public (web/dist symlink) # wheel → bundle skills + UI, then maturin build # frontend wheel → full release path # @@ -27,7 +27,7 @@ else endif MATURIN_FLAGS := $(MATURIN_RELEASE) --features $(MATURIN_FEATURES) -BUNDLED_WEB_PUBLIC := python/probing/bundled_web/public +EMBEDDED_WEB_PUBLIC := probing/server/web-assets/public ifdef ZIG ifdef TARGET @@ -85,8 +85,8 @@ help: @echo " develop / dev Bootstrap: _core, CLI, pytest, site hook" @echo " Tip: DEBUG=1 make develop → dev profile (faster link)" @echo " core Rebuild probing._core after Rust edits" - @echo " frontend Build UI into python/probing/bundled_web (dx bundle)" - @echo " wheel Build dist/*.whl (needs bundled_web; bundles skills + UI)" + @echo " frontend Build UI for compile-time embedding into probing-server" + @echo " wheel Build dist/*.whl (UI is embedded in probing._core)" @echo " wheel-ci alias for wheel (native build; PyPI uses maturin-action + zig)" @echo " install-wheel pip install dist/probing-*.whl" @echo " venv create/sync project .venv (from .python-version when set)" @@ -206,23 +206,24 @@ soak-quick: check-dev DURATION_SEC=60 MAX_STEPS=8 PYTHON=$(VENV_PYTHON) PROBING=1 ./examples/imagenet/run_soak.sh frontend: - @test -n "$$SKIP_FRONTEND_CLEAN" || rm -rf python/probing/bundled_web - cd web && dx bundle --release - @test -f $(BUNDLED_WEB_PUBLIC)/index.html - @chmod +x scripts/prune-bundled-web.sh - @./scripts/prune-bundled-web.sh $(BUNDLED_WEB_PUBLIC) - @mkdir -p $(BUNDLED_WEB_PUBLIC)/assets - @cp -f web/assets/logo.svg $(BUNDLED_WEB_PUBLIC)/logo.svg 2>/dev/null || true - @cp -f web/assets/logo.svg $(BUNDLED_WEB_PUBLIC)/assets/logo.svg 2>/dev/null || true - @cp -f web/assets/tailwind.css $(BUNDLED_WEB_PUBLIC)/assets/tailwind.css - @$(PYTHON) scripts/verify_web_assets.py $(BUNDLED_WEB_PUBLIC) + @test -n "$$SKIP_FRONTEND_CLEAN" || rm -rf probing/server/web-assets + cd web && dx bundle --release --debug-symbols false + @test -f $(EMBEDDED_WEB_PUBLIC)/index.html + @chmod +x scripts/prune-web-assets.sh + @./scripts/prune-web-assets.sh $(EMBEDDED_WEB_PUBLIC) + @mkdir -p $(EMBEDDED_WEB_PUBLIC)/assets + @cp -f web/assets/logo.svg $(EMBEDDED_WEB_PUBLIC)/logo.svg 2>/dev/null || true + @cp -f web/assets/logo.svg $(EMBEDDED_WEB_PUBLIC)/assets/logo.svg 2>/dev/null || true + @cp -f web/assets/tailwind.css $(EMBEDDED_WEB_PUBLIC)/assets/tailwind.css + @printf '%s\n' '__PROBING_EMBEDDED_WEB_ASSETS_V1__' > $(EMBEDDED_WEB_PUBLIC)/embedded.manifest + @PROBING_CLI_MODE=1 $(PYTHON) scripts/verify_web_assets.py $(EMBEDDED_WEB_PUBLIC) @rm -rf web/dist - @ln -sfn ../python/probing/bundled_web/public web/dist - @echo "$(BUNDLED_WEB_PUBLIC) ($$(du -sh $(BUNDLED_WEB_PUBLIC) | cut -f1))" + @ln -sfn ../probing/server/web-assets/public web/dist + @echo "$(EMBEDDED_WEB_PUBLIC) ($$(du -sh $(EMBEDDED_WEB_PUBLIC) | cut -f1))" wheel-bundle: - @test -f $(BUNDLED_WEB_PUBLIC)/index.html || { echo "error: run 'make frontend' first"; exit 1; } - @$(PYTHON) scripts/verify_web_assets.py $(BUNDLED_WEB_PUBLIC) + @test -f $(EMBEDDED_WEB_PUBLIC)/index.html || { echo "error: run 'make frontend' first"; exit 1; } + @PROBING_CLI_MODE=1 $(PYTHON) scripts/verify_web_assets.py $(EMBEDDED_WEB_PUBLIC) @test -f python/probing/bundled_skills/catalog.yaml \ || { echo "error: missing python/probing/bundled_skills/catalog.yaml"; exit 1; } @@ -234,7 +235,7 @@ wheel-ci: $(MAKE) wheel verify-wheel-contents: - @$(PYTHON) scripts/verify_wheel_contents.py + @PROBING_CLI_MODE=1 $(PYTHON) scripts/verify_wheel_contents.py install-wheel: verify-wheel-contents @WH=$$(ls -1 dist/probing-*.whl 2>/dev/null | head -1); \ @@ -246,7 +247,7 @@ import importlib.util; \ import probing; from probing import _core; from pathlib import Path; \ root = Path(probing.__file__).resolve().parent; \ assert (root / 'bundled_skills' / 'catalog.yaml').is_file(), f'missing bundled skills under {root}'; \ -assert (root / 'bundled_web' / 'public' / 'index.html').is_file() or (root / 'bundled_web' / 'index.html').is_file(), f'missing bundled web UI under {root}'; \ +assert not (root / 'bundled_web').exists(), f'legacy bundled_web must not be shipped under {root}'; \ assert all(importlib.util.find_spec(m) for m in ('probing.skills.loader', 'probing.ext', 'probing.handlers')), 'missing probing.* modules in installed wheel'; \ print('probing', probing.VERSION)" @@ -393,7 +394,7 @@ docs-clean: @cd docs && $(MAKE) clean clean: - rm -rf dist docs/site python/probing/bundled_web web/dist web/target + rm -rf dist docs/site probing/server/web-assets web/dist web/target rm -rf python/probing/libs python/probing/shim/hccl rm -rf .pytest_cache .coverage coverage.xml coverage.lcov cargo clean diff --git a/docs/src/contributing.md b/docs/src/contributing.md index e6057e78..eb89026a 100644 --- a/docs/src/contributing.md +++ b/docs/src/contributing.md @@ -180,7 +180,7 @@ See [examples/README.md](https://github.com/DeepLink-org/probing/blob/main/examp - **Authoring**: repo root `skills/` (`SKILL.md`, `steps.yaml`, `catalog.yaml`) - **Install to IDE agents**: `./skills/install.sh` or `probing skill install` -- **Bundled in wheel**: `make wheel` copies skills into `python/probing/bundled_skills/` and UI into `python/probing/bundled_web/` +- **Bundled in wheel**: `make wheel` copies skills into `python/probing/bundled_skills/`; `make frontend` builds ignored UI artifacts under `probing/server/web-assets/`, and the server build script embeds them into `probing._core` at compile time (plain Rust builds use a tracked fallback page) - **Docs**: `skills/README.md`, [Extensibility — Diagnostic skill](design/extensibility.md#path-2-diagnostic-skill) ## Development workflow @@ -250,13 +250,11 @@ probing/ # repo root ├── python/ │ ├── probing/ # Python PACKAGE (not Rust) │ │ ├── skills/ # skill loader/install CODE — see python/probing/skills/README.md -│ │ ├── web_assets.py # wheel _web/ + editable web/dist → PROBING_ASSETS_ROOT -│ │ ├── bundled_skills/ # skill DATA bundled in wheel (author in repo-root skills/) -│ │ └── bundled_web/ # UI bundled in wheel (make frontend) +│ │ └── bundled_skills/ # skill DATA bundled in wheel (author in repo-root skills/) │ ├── probing_hook.py # .pth → site hook │ └── probing.pth ├── src/lib.rs # PyO3 entry → probing._core (maturin) -├── probing/ # Rust WORKSPACE (core, server, cli, extensions) +├── probing/ # Rust WORKSPACE (server/web-assets is embedded in probing._core) ├── web/ # Dioxus UI (`make frontend` → web/dist/) ├── tests/ # see tests/README.md ├── examples/ # optional torch/etc. diff --git a/docs/src/contributing.zh.md b/docs/src/contributing.zh.md index 7c0bd0d7..a8899385 100644 --- a/docs/src/contributing.zh.md +++ b/docs/src/contributing.zh.md @@ -176,7 +176,7 @@ PROBING=1 python examples/getting-started/tracing.py - **编写**:仓库根 `skills/`(`SKILL.md`、`steps.yaml`、`catalog.yaml`) - **安装到 IDE**:`./skills/install.sh` 或 `probing skill install` -- **打进 wheel**:`make wheel` 自动复制到 `python/probing/bundled_skills/`、`python/probing/bundled_web/` +- **打进 wheel**:`make wheel` 复制 Skill 到 `python/probing/bundled_skills/`;`make frontend` 将被忽略的 UI 产物构建到 `probing/server/web-assets/`,server build script 在编译时把它嵌入 `probing._core`(普通 Rust 构建使用受版本控制的 fallback 页面) - **说明**:`skills/README.md`、[扩展机制 — 诊断 skill](design/extensibility.zh.md#path-2-diagnostic-skill) ## 开发流程 @@ -245,13 +245,11 @@ probing/ # 仓库根 ├── python/ │ ├── probing/ # Python 包(不是 Rust) │ │ ├── skills/ # skill 加载/安装代码 — 见 python/probing/skills/README.md -│ │ ├── web_assets.py # wheel _web/ + editable web/dist → PROBING_ASSETS_ROOT -│ │ ├── bundled_skills/ # wheel 打包的 skill 数据(编写在 repo-root skills/) -│ │ └── bundled_web/ # wheel 打包的 UI(make frontend) +│ │ └── bundled_skills/ # wheel 打包的 skill 数据(编写在 repo-root skills/) │ ├── probing_hook.py │ └── probing.pth ├── src/lib.rs # PyO3 → probing._core -├── probing/ # Rust workspace +├── probing/ # Rust workspace(server/web-assets 编译进 probing._core) ├── web/ # Dioxus UI(`make frontend` → web/dist/) ├── tests/ # 见 tests/README.md ├── examples/ diff --git a/docs/src/design/data-layer.md b/docs/src/design/data-layer.md index 36c33cbf..33889af8 100644 --- a/docs/src/design/data-layer.md +++ b/docs/src/design/data-layer.md @@ -205,10 +205,12 @@ Each column is encoded independently (`ColEncoding`): ### Crash Recovery -- A **sealed** segment is read via its footer page directory — O(1) location of every page. -- An **unsealed or torn** segment is recovered by **forward scan**: walk blocks from the start, - verifying each block's header and payload checksum, stopping at the first bad block and dropping - the torn tail. Table-definition blocks are always scanned (cheap, and they precede pages). +- A **sealed** segment is read via its footer page directory — O(1) location of every page. Footer, + block, payload, and page integrity failures reject the segment and fail the SQL scan; queries + never return a successful but incomplete cold result. +- An **unsealed or torn** segment is recovered by **forward scan**. Only an incomplete final + header/payload is dropped as a crash tail; corruption of a complete checksummed block is an + error. Table-definition blocks are always scanned (cheap, and they precede pages). There is no heuristic that tries to repair a half-written record. @@ -251,6 +253,12 @@ compactor a single lifecycle home: over time), drains each into the shared `ColdStore`, rolls by age, and enforces the budget; - on startup it calls `prime_from_cold()`; on stop it flushes (seals the open segment). +Discovery, segment enumeration, write/roll, and retention I/O are fallible operations. The worker +logs failures and records an observable `CompactorRuntimeStats` snapshot (`error_count` plus the +operation-qualified `last_error`). Startup priming fails closed, because continuing without recovered +watermarks could duplicate cold rows; retention never reports success after an unreadable directory +or metadata failure. + It is **opt-in** (off by default) to avoid spawning a compaction thread in every forked worker. Configuration is applied via the `MemTableProbeExtension` option surface or environment variables; the server calls `start_cold_compaction_from_env()` at engine init. diff --git a/docs/src/design/data-layer.zh.md b/docs/src/design/data-layer.zh.md index fae490e3..af354370 100644 --- a/docs/src/design/data-layer.zh.md +++ b/docs/src/design/data-layer.zh.md @@ -181,9 +181,10 @@ MEMC v2 reader 会明确拒绝 v1 段,因为 v1 缺少实例身份,无法安 ### 崩溃恢复 -- **已封存**段通过 footer 的 page 目录读取——O(1) 定位每个 page; -- **未封存或撕裂**的段通过**前向扫描**恢复:从头遍历 block,校验每个 block 的头部和 payload - 校验和,在第一个坏 block 处停止并丢弃撕裂的尾部。表定义 block 总会被扫描(开销小,且位于 page 之前)。 +- **已封存**段通过 footer 的 page 目录读取——O(1) 定位每个 page。footer、block、payload 或 page + 的完整性校验一旦失败,整个段及 SQL 扫描都会报错,查询不会再成功返回不完整的冷层结果; +- **未封存或撕裂**的段通过**前向扫描**恢复。只有不完整的最后一个头部或 payload 会被当作崩溃尾部 + 丢弃;完整 block 的校验和损坏仍会报错。表定义 block 总会被扫描(开销小,且位于 page 之前)。 不存在任何试图修复半行记录的启发式逻辑。 @@ -218,6 +219,11 @@ MEMC v2 reader 会明确拒绝 v1 段,因为 v1 缺少实例身份,无法安 `ColdStore`,按时长滚动,并执行预算约束; - 启动时调用 `prime_from_cold()`;停止时 flush(封存打开的段)。 +发现、段枚举、写入/滚动与 retention I/O 都是可失败操作。worker 会记录 warning,并通过 +`CompactorRuntimeStats` 暴露 `error_count` 和带操作上下文的 `last_error`。启动时的 watermark +恢复失败会 fail-closed,因为继续运行可能重复写入冷层;目录或 metadata 读取失败时 retention 也不会 +再被当作成功。 + 它**默认关闭**(opt-in),以避免在每个 fork 出来的 worker 中都启动一个压缩线程。配置通过 `MemTableProbeExtension` 选项面或环境变量下发;server 在引擎初始化时调用 `start_cold_compaction_from_env()`。 diff --git a/docs/src/design/distributed.md b/docs/src/design/distributed.md index 1ccaa202..8935c9a7 100644 --- a/docs/src/design/distributed.md +++ b/docs/src/design/distributed.md @@ -116,6 +116,11 @@ See [Federated query engine](federation.md) for engine paths and acceptance test At wan scale, **`cluster query` defaults to [hierarchical fan-out](hierarchical-fanout.md)** (coordinator → per-machine local0 → on-node leaves). Set `PROBING_CLUSTER_FANOUT_HIERARCHICAL=0` or use CLI `--flat` for legacy flat fan-out. +Raw `global.*` scans use the same topology: the coordinator reads its own partition, queries its +local leaf ranks directly, and sends node-aggregate requests to remote local0 peers. Hierarchical +execution requires `group_rank` and `local_rank` on every live registry entry; partial metadata is +reported as an error instead of falling back to a potentially incomplete flat/partial scan. + ## Synchronized Debugging ### Capture All Stacks @@ -244,6 +249,10 @@ PROBING_AUTH_TOKEN=secret python train.py probing -t host:8080 --token secret query "..." ``` +Use the same token on every peer. Probing applies the configured credential to internal +node discovery, heartbeat, flat federation queries, and hierarchical fan-out requests; +load-balancer health endpoints remain public. + ## Best Practices ### Consistent environment diff --git a/docs/src/design/distributed.zh.md b/docs/src/design/distributed.zh.md index 35910094..f2b7619e 100644 --- a/docs/src/design/distributed.zh.md +++ b/docs/src/design/distributed.zh.md @@ -111,6 +111,10 @@ ORDER BY avg_ms DESC" 万卡场景下 `cluster query` 默认走 **[分层 fan-out](hierarchical-fanout.zh.md)**(coordinator 仅联系各机 local0,local0 再聚合本机 leaf rank),可用 `PROBING_CLUSTER_FANOUT_HIERARCHICAL=0` 或 CLI `--flat` 恢复扁平 fan-out。 +普通 `global.*` 扫描采用同一拓扑:coordinator 读取自身分区、直接查询本机 leaf ranks,并向异机 +local0 发送节点聚合请求。层级执行要求每个存活注册节点都具有 `group_rank` 与 `local_rank`;元数据 +只要部分缺失就会明确报错,不再降级为可能漏数的 flat/partial 扫描。 + ## 同步调试 ### 捕获所有堆栈 @@ -239,6 +243,9 @@ PROBING_AUTH_TOKEN=secret python train.py probing -t host:8080 --token secret query "..." ``` +所有 peer 必须使用相同令牌。Probing 会把已配置的凭据附加到内部节点发现、心跳、 +普通联邦查询和层级 fan-out 请求;负载均衡器使用的健康检查端点仍保持公开。 + ## 最佳实践 ### 一致的环境变量 diff --git a/docs/src/reference/skill-format.md b/docs/src/reference/skill-format.md index 9ec637bb..e1ec70c5 100644 --- a/docs/src/reference/skill-format.md +++ b/docs/src/reference/skill-format.md @@ -110,8 +110,10 @@ parameters: description: "Number of recent steps to analyze" ``` -Types: `integer`, `boolean`, `string`. Parameter values are referenced in SQL as -`{param_name}`. +Types: `integer`, `number`, `boolean`, `string`. Parameter values are referenced in SQL as +`{param_name}`. Overrides are validated and normalized before planning or execution; +unknown parameters are rejected. String parameters represent SQL literal contents and +single quotes are escaped during template expansion. ### Requires @@ -158,6 +160,8 @@ Step fields: - `cluster`: If `true`, uses federation fan-out (`POST /apis/cluster/query`). - `when`: Optional condition. `"always"` or `"{use_global}"` (runs only when the boolean variable is true). +- `platform`: Optional execution platform: `linux`, `macos`, or `windows`. A step for a + different platform is reported as skipped without issuing its SQL/API request. **`api`** — Call an HTTP API on the probing endpoint. @@ -236,6 +240,7 @@ interpretation: | Row count | `rows == 0`, `rows >= 1` | | Column predicate | `column: \| ` — pairs with the tail on the right | | Numeric compare | `value == 0`, `value > 10`, `max > 1e6`, `avg > 5` | +| Text equality | `value = propagated_victim` — exact match on the first row | | Spread | `max/min(ratio) > 1.5` | | Ratio of columns | `ratio(num_col/den_col) > 0.3` | | Text search | `any_contains('dead', 'stale')` — case-insensitive substring | @@ -341,5 +346,8 @@ python -m probing.skills validate my_skill python -m probing.skills validate --all ``` -The validator checks: missing steps, duplicate step IDs, read-only SQL compliance -(all statements must start with SELECT/WITH/SHOW/DESCRIBE), and missing SKILL.md. +The validator compiles every catalog skill through the Rust execution schema. It rejects +unknown YAML fields, mismatched parameter defaults/overrides, unsupported platforms and +interpretation grammar, duplicate IDs, unresolved templates, and missing required step +fields. SQL is parsed as an AST; every statement and nested CTE must be read-only, so a +read followed by a trailing write is rejected. It also checks that `SKILL.md` exists. diff --git a/probing/cli/src/cli/bench/runners/mixed.rs b/probing/cli/src/cli/bench/runners/mixed.rs index 516e429f..b4502bc2 100644 --- a/probing/cli/src/cli/bench/runners/mixed.rs +++ b/probing/cli/src/cli/bench/runners/mixed.rs @@ -89,7 +89,7 @@ pub fn run(args: &MixedArgs, json: bool, seed: u64) -> Result<()> { ttl: args.ttl_secs.map(Duration::from_secs), }; let handle = attach.open()?; - Some(Compactor::new(store, config).spawn(vec![("bench".to_string(), handle)])) + Some(Compactor::new(store, config).spawn(vec![("bench".to_string(), handle)])?) }; let mut threads = Vec::new(); diff --git a/probing/core/Cargo.toml b/probing/core/Cargo.toml index ec77f9d6..a4cd34ba 100644 --- a/probing/core/Cargo.toml +++ b/probing/core/Cargo.toml @@ -49,7 +49,6 @@ thiserror = { workspace = true } async-trait = "0.1.83" datafusion = { workspace = true } futures = "0.3.31" -ureq = { workspace = true } sled = "0.34.7" bincode = "1.3.3" uuid = { version = "1.0", features = ["v4", "serde"] } diff --git a/probing/core/src/core/cluster.rs b/probing/core/src/core/cluster.rs index b7c95354..ec4c5932 100644 --- a/probing/core/src/core/cluster.rs +++ b/probing/core/src/core/cluster.rs @@ -380,7 +380,20 @@ pub fn node_aggregator_peers() -> Vec { /// Leaf ranks on this node (same ``group_rank``, excluding self). pub fn local_leaf_peers() -> Vec { let local_addrs = local_listen_addrs(); - let group_rank = env_i32("GROUP_RANK").or_else(|| env_i32("NODE_RANK")); + let group_rank = env_i32("GROUP_RANK") + .or_else(|| env_i32("NODE_RANK")) + .or_else(|| { + get_nodes() + .into_iter() + .find(|node| local_addrs.iter().any(|local| local == &node.addr)) + .and_then(|node| node.group_rank) + }); + let Some(group_rank) = group_rank else { + log::warn!( + "hierarchical fan-out cannot identify the local group_rank; refusing cross-node leaf selection" + ); + return Vec::new(); + }; let self_rank = env_i32("RANK"); get_nodes() @@ -390,10 +403,8 @@ pub fn local_leaf_peers() -> Vec { if local_addrs.iter().any(|local| local == &node.addr) { return false; } - if let (Some(g), Some(expected)) = (node.group_rank, group_rank) { - if g != expected { - return false; - } + if node.group_rank != Some(group_rank) { + return false; } if let (Some(r), Some(self_r)) = (node.rank, self_rank) { if r == self_r { @@ -411,10 +422,11 @@ fn env_i32(name: &str) -> Option { /// Whether ``cluster.nodes`` has enough metadata for hierarchical fan-out. pub fn hierarchical_metadata_available() -> bool { - get_nodes() - .iter() - .filter(|n| is_node_alive(n)) - .any(|n| n.group_rank.is_some() && n.local_rank.is_some()) + let alive: Vec = get_nodes().into_iter().filter(is_node_alive).collect(); + !alive.is_empty() + && alive + .iter() + .all(|n| n.group_rank.is_some() && n.local_rank.is_some()) } /// Error prefix returned when hierarchical fan-out is requested but metadata is missing. @@ -650,4 +662,24 @@ mod tests { std::env::remove_var("GROUP_RANK"); std::env::remove_var("RANK"); } + + #[test] + fn hierarchical_metadata_requires_every_alive_node() { + let _guard = TEST_LOCK.lock().unwrap(); + reset_cluster_for_tests(); + for (rank, group_rank, local_rank) in [(0, Some(0), Some(0)), (1, None, Some(1))] { + write_cluster().put(Node { + host: "h".into(), + addr: format!("10.0.0.1:{}", 8080 + rank), + rank: Some(rank), + group_rank, + local_rank, + status: Some("running".into()), + ..Default::default() + }); + } + + assert!(!hierarchical_metadata_available()); + reset_cluster_for_tests(); + } } diff --git a/probing/core/src/core/engine.rs b/probing/core/src/core/engine.rs index 959b5685..fbc36be1 100644 --- a/probing/core/src/core/engine.rs +++ b/probing/core/src/core/engine.rs @@ -1,6 +1,6 @@ use std::collections::HashMap; use std::sync::Arc; -use tokio::sync::RwLock; +use std::sync::RwLock; use arrow::compute::concat_batches; use datafusion::catalog::MemoryCatalogProvider; @@ -17,6 +17,7 @@ use super::probe_extension::ProbeExtensionManager; use super::data_source::{ProbeDataSource, ProbeDataSourceKind}; use super::federation; +use super::federation::PeerQueryTransport; use super::metadata_rewrite; use super::semantic_catalog; @@ -53,27 +54,17 @@ pub struct Engine { pub context: SessionContext, /// Registry of enabled plugins, mapped by their fully qualified names data_sources: RwLock>>, + peer_query_transport: Option>, } impl Clone for Engine { fn clone(&self) -> Self { - let plugins_clone = match self.data_sources.try_read() { - Ok(guard) => guard.clone(), - Err(_) => { - log::warn!("Engine::clone: data_sources contended; blocking for read"); - if let Ok(handle) = tokio::runtime::Handle::try_current() { - tokio::task::block_in_place(|| { - handle.block_on(async { self.data_sources.read().await.clone() }) - }) - } else { - crate::runtime::CORE_RUNTIME - .block_on(async { self.data_sources.read().await.clone() }) - } - } - }; + let plugins_clone = + crate::sync::read_rwlock(&self.data_sources, "engine data_sources").clone(); Self { context: self.context.clone(), data_sources: RwLock::new(plugins_clone), + peer_query_transport: self.peer_query_transport.clone(), } } } @@ -92,6 +83,7 @@ impl Default for Engine { Engine { context: SessionContext::new_with_config(config), data_sources: Default::default(), + peer_query_transport: None, } } } @@ -192,7 +184,7 @@ impl Engine { if let Some(wrapper) = data_source.provide_catalog(catalog) { self.context.register_catalog("probe", wrapper); } - let mut maps = self.data_sources.write().await; + let mut maps = crate::sync::write_rwlock(&self.data_sources, "engine data_sources"); maps.insert(format!("probe.{namespace}"), data_source); } else if data_source.kind() == ProbeDataSourceKind::Table { // In DataFusion, schemas are used to implement namespaces @@ -205,7 +197,7 @@ impl Engine { })?; let state: SessionState = self.context.state(); data_source.register_table(schema, &state)?; - let mut maps = self.data_sources.write().await; + let mut maps = crate::sync::write_rwlock(&self.data_sources, "engine data_sources"); maps.insert( format!("probe.{}.{}", namespace, data_source.name()), data_source, @@ -213,6 +205,10 @@ impl Engine { } Ok(()) } + + pub fn peer_query_transport(&self) -> Option> { + self.peer_query_transport.clone() + } } // Define the EngineBuilder struct @@ -221,6 +217,7 @@ pub struct EngineBuilder { default_namespace: Option, data_sources: Vec>, probe_extensions: HashMap>>, + peer_query_transport: Option>, } impl EngineBuilder { @@ -231,6 +228,7 @@ impl EngineBuilder { default_namespace: None, data_sources: Vec::new(), probe_extensions: Default::default(), + peer_query_transport: None, } } @@ -257,9 +255,14 @@ impl EngineBuilder { self } + pub fn with_peer_query_transport(mut self, transport: Arc) -> Self { + self.peer_query_transport = Some(transport); + self + } + // Build the Engine with the specified configurations pub async fn build(mut self) -> Result { - let mut eem = ProbeExtensionManager; + let mut eem = ProbeExtensionManager::default(); for (name, extension) in self.probe_extensions.iter() { eem.register(name.clone(), extension.clone()).await; } @@ -279,12 +282,13 @@ impl EngineBuilder { let engine = Engine { context, data_sources: Default::default(), + peer_query_transport: self.peer_query_transport, }; for data_source in self.data_sources { engine.enable(data_source).await?; } semantic_catalog::install_semantic_catalog(&engine.context)?; - federation::install_global_catalog(&engine.context)?; + federation::install_global_catalog(&engine.context, engine.peer_query_transport())?; Ok(engine) } @@ -314,6 +318,33 @@ mod tests { use probing_proto::prelude::Seq; use std::sync::Arc; + #[derive(Debug)] + struct TestPeerTransport; + + impl PeerQueryTransport for TestPeerTransport { + fn query( + &self, + _addr: &str, + _sql: &str, + _scope: federation::FanoutScope, + ) -> Result { + unreachable!("transport isolation test does not execute queries") + } + } + + #[tokio::test] + async fn peer_transport_is_scoped_to_the_built_engine() { + let configured = Engine::builder() + .with_peer_query_transport(Arc::new(TestPeerTransport)) + .build() + .await + .unwrap(); + let plain = Engine::builder().build().await.unwrap(); + + assert!(configured.peer_query_transport().is_some()); + assert!(plain.peer_query_transport().is_none()); + } + #[derive(Debug, Clone)] struct TestTableProbeDataSource { schema: SchemaRef, diff --git a/probing/core/src/core/federation/aggregate_pushdown.rs b/probing/core/src/core/federation/aggregate_pushdown.rs index 7e7b6faa..2426cc86 100644 --- a/probing/core/src/core/federation/aggregate_pushdown.rs +++ b/probing/core/src/core/federation/aggregate_pushdown.rs @@ -95,12 +95,18 @@ pub async fn try_execute_aggregate_pushdown( } let per_node_sql = plan.per_node_sql.clone(); + let transport = engine.peer_query_transport(); let mut stats = FanoutStats::default(); if scope == FanoutScope::Coordinator && is_local0_from_env() { let leaf_sql = per_node_sql.clone(); + let leaf_transport = transport.clone(); let leaf_outcomes = tokio::task::spawn_blocking(move || { - ProbeClusterExecutor::fanout_query_to_peers_scoped(&leaf_sql, FanoutScope::Node) + ProbeClusterExecutor::fanout_query_to_peers_scoped( + &leaf_sql, + FanoutScope::Node, + leaf_transport, + ) }) .await .map_err(|e| DataFusionError::Execution(format!("local leaf fan-out failed: {e}")))?; @@ -110,7 +116,7 @@ pub async fn try_execute_aggregate_pushdown( let remote_scope = scope; let remote_sql = per_node_sql.clone(); let outcomes = tokio::task::spawn_blocking(move || { - ProbeClusterExecutor::fanout_query_to_peers_scoped(&remote_sql, remote_scope) + ProbeClusterExecutor::fanout_query_to_peers_scoped(&remote_sql, remote_scope, transport) }) .await .map_err(|e| DataFusionError::Execution(format!("aggregate fan-out join failed: {e}")))?; @@ -135,6 +141,7 @@ pub async fn try_execute_aggregate_pushdown( } } stats.peer_batches_dropped += convert_failed; + stats.partial |= convert_failed > 0; merge_on_coordinator(&engine.context, &merge_sql, batches).await? } else if proto_parts.len() == 1 { proto_parts.remove(0) @@ -159,18 +166,20 @@ fn append_fanout_outcomes( plan: &FederatedAggregatePlan, outcomes: Vec, ) { - for outcome in outcomes { - match outcome.result { - Ok(mut df) => { - stats.nodes_succeeded += 1; + for remote in outcomes { + match remote.result { + Ok(outcome) => { + stats.absorb(outcome.stats); + let mut df = outcome.dataframe; if plan.inject_tags { - tag_proto_dataframe(&mut df, &outcome.host, &outcome.addr, outcome.rank); + tag_proto_dataframe(&mut df, &remote.host, &remote.addr, remote.rank); } proto_parts.push(df); } Err(err) => { - log::warn!("aggregate pushdown skipped {}: {err}", outcome.addr); - stats.nodes_failed.push(outcome.addr); + log::warn!("aggregate pushdown skipped {}: {err}", remote.addr); + stats.nodes_failed.push(remote.addr); + stats.partial = true; } } } diff --git a/probing/core/src/core/federation/cluster_executor.rs b/probing/core/src/core/federation/cluster_executor.rs index 21af2f0c..a38705f9 100644 --- a/probing/core/src/core/federation/cluster_executor.rs +++ b/probing/core/src/core/federation/cluster_executor.rs @@ -1,14 +1,15 @@ -use std::sync::LazyLock; +use std::fmt::Debug; +use std::sync::{Arc, LazyLock}; #[cfg(any(test, feature = "test-utils"))] use std::sync::{Mutex, MutexGuard}; use std::time::Duration; use datafusion::error::{DataFusionError, Result}; -use probing_proto::prelude::{DataFrame, Message, Node, Query, QueryDataFormat}; +use probing_proto::prelude::{DataFrame, Node}; use crate::core::cluster::{ - hierarchical_metadata_available, local_leaf_peers, node_aggregator_peers, - remote_peers_excluding_local, + hierarchical_metadata_available, hierarchical_metadata_unavailable_err, local_leaf_peers, + node_aggregator_peers, remote_peers_excluding_local, }; use crate::core::federation::fanout_scope::{ current_fanout_scope, current_fanout_stats_handle, resolve_fanout_scope, FanoutScope, @@ -21,6 +22,35 @@ type RemoteQueryHook = Box Result + Send + Sync static REMOTE_QUERY_HOOK: LazyLock>> = LazyLock::new(|| Mutex::new(None)); +/// L3-provided transport for L1 federation execution. +pub trait PeerQueryTransport: Debug + Send + Sync { + fn query(&self, addr: &str, sql: &str, scope: FanoutScope) -> Result; +} + +/// Data and completeness metadata returned for one remote subtree. +/// +/// `stats` accounts for the addressed endpoint and every descendant it queried. +/// A leaf success therefore reports one successful node, while a node aggregator +/// reports the complete subtree counts received from its L3 response. +#[derive(Debug)] +pub struct PeerQueryOutcome { + pub dataframe: DataFrame, + pub stats: FanoutStats, +} + +impl PeerQueryOutcome { + pub fn complete(dataframe: DataFrame) -> Self { + Self { + dataframe, + stats: FanoutStats::complete_node(), + } + } + + pub fn with_stats(dataframe: DataFrame, stats: FanoutStats) -> Self { + Self { dataframe, stats } + } +} + /// Install an in-process remote query handler for federation integration tests. #[cfg(any(test, feature = "test-utils"))] pub fn set_remote_query_hook(hook: Option) { @@ -35,10 +65,6 @@ const REMOTE_QUERY_TIMEOUT_ENV: &str = "PROBING_REMOTE_QUERY_TIMEOUT_SECS"; const REMOTE_FANOUT_CONCURRENCY_ENV: &str = "PROBING_FANOUT_CONCURRENCY"; const DEFAULT_REMOTE_FANOUT_CONCURRENCY: usize = 128; -fn external(err: E) -> DataFusionError { - DataFusionError::External(Box::new(err)) -} - /// Per-node timeout for remote federated queries. /// /// Defaults to [`DEFAULT_REMOTE_QUERY_TIMEOUT_SECS`]; override via the @@ -72,7 +98,18 @@ pub struct RemoteFanoutResult { pub addr: String, pub host: String, pub rank: Option, - pub result: Result, + pub result: Result, +} + +/// One remote partition in a raw federated scan. +/// +/// Coordinator scans deliberately mix direct on-node leaf queries with +/// node-aggregate queries to remote nodes, so the routing scope belongs to +/// each target rather than to the scan as a whole. +#[derive(Debug, Clone)] +pub(crate) struct FederatedScanTarget { + pub node: Node, + pub scope: FanoutScope, } pub use super::fanout_scope::FanoutStats; @@ -117,7 +154,7 @@ pub fn fanout_strict_enabled() -> bool { } pub fn fanout_stats_partial(stats: &FanoutStats) -> bool { - !stats.nodes_failed.is_empty() || stats.peer_batches_dropped > 0 + stats.partial || !stats.nodes_failed.is_empty() || stats.peer_batches_dropped > 0 } /// Fail the query when strict fan-out is enabled and any peer was dropped. @@ -201,11 +238,18 @@ impl ProbeClusterExecutor { /// Requests run in parallel (one OS thread per peer via [`std::thread::scope`]), /// so total latency is bounded by the slowest peer rather than the sum of all /// peers. Node identity is preserved for row tagging and fan-out accounting. - pub fn fanout_query_to_peers(sql: &str) -> Vec { - Self::fanout_query_to_peers_scoped(sql, current_fanout_scope()) + pub fn fanout_query_to_peers( + sql: &str, + transport: Option>, + ) -> Vec { + Self::fanout_query_to_peers_scoped(sql, current_fanout_scope(), transport) } - pub fn fanout_query_to_peers_scoped(sql: &str, scope: FanoutScope) -> Vec { + pub fn fanout_query_to_peers_scoped( + sql: &str, + scope: FanoutScope, + transport: Option>, + ) -> Vec { let nodes = Self::remote_nodes_for_scope(scope); if nodes.is_empty() { return Vec::new(); @@ -219,13 +263,19 @@ impl ProbeClusterExecutor { .iter() .map(|node| { let node = node.clone(); + let transport = transport.clone(); s.spawn(move || { let host = if node.host.is_empty() { node.addr.clone() } else { node.host.clone() }; - let result = Self::execute_remote_scoped(&node.addr, sql, scope); + let result = Self::execute_remote_scoped( + transport.as_ref(), + &node.addr, + sql, + scope, + ); RemoteFanoutResult { addr: node.addr, host, @@ -250,146 +300,102 @@ impl ProbeClusterExecutor { results } - /// Peer nodes and execution scope for a federated table scan. + /// Remote partitions for a federated table scan. /// - /// When hierarchical coordinator tier finds no ``local_rank == 0`` aggregators, - /// only falls back to flat peers if hierarchical fan-out is disabled. - pub fn federated_scan_targets() -> (Vec, FanoutScope) { + /// A coordinator owns its own local partition, queries sibling ranks on + /// the same node directly, and asks one aggregator on every other node to + /// fan in there. Hierarchical scopes fail closed when metadata is partial; + /// silently switching to a flat scan can duplicate or omit ranks. + pub(crate) fn federated_scan_targets() -> Result> { let resolved = resolve_fanout_scope(current_fanout_scope()); match resolved { FanoutScope::Coordinator => { - let peers = Self::remote_nodes_for_scope(FanoutScope::Coordinator); - if peers.is_empty() { - log::debug!("federated scan: no node aggregators; falling back to flat peers"); - ( - Self::remote_nodes_for_scope(FanoutScope::Flat), - FanoutScope::Flat, - ) - } else { - (peers, FanoutScope::Coordinator) + if remote_peers_excluding_local().is_empty() { + return Ok(Vec::new()); + } + if !hierarchical_metadata_available() { + return Err(hierarchical_metadata_unavailable_err().into()); } + Ok(Self::hierarchical_scan_targets( + local_leaf_peers(), + node_aggregator_peers(), + )) } FanoutScope::Node => { - let peers = Self::remote_nodes_for_scope(FanoutScope::Node); - if peers.is_empty() { - log::debug!("federated scan: no local leaf peers; falling back to flat peers"); - ( - Self::remote_nodes_for_scope(FanoutScope::Flat), - FanoutScope::Flat, - ) - } else { - (peers, FanoutScope::Node) + if remote_peers_excluding_local().is_empty() { + return Ok(Vec::new()); + } + if !hierarchical_metadata_available() { + return Err(hierarchical_metadata_unavailable_err().into()); } + Ok(local_leaf_peers() + .into_iter() + .map(|node| FederatedScanTarget { + node, + scope: FanoutScope::Node, + }) + .collect()) } - scope => (Self::remote_nodes_for_scope(scope), scope), + scope => Ok(Self::remote_nodes_for_scope(scope) + .into_iter() + .map(|node| FederatedScanTarget { node, scope }) + .collect()), } } - pub fn execute_remote_query(addr: &str, sql: &str) -> Result { - Self::execute_remote_for_scope(addr, sql, current_fanout_scope()) + fn hierarchical_scan_targets( + local_leaves: Vec, + remote_aggregators: Vec, + ) -> Vec { + local_leaves + .into_iter() + .map(|node| FederatedScanTarget { + node, + scope: FanoutScope::Node, + }) + .chain( + remote_aggregators + .into_iter() + .map(|node| FederatedScanTarget { + node, + scope: FanoutScope::Coordinator, + }), + ) + .collect() + } + + pub fn execute_remote_query( + transport: Option<&Arc>, + addr: &str, + sql: &str, + ) -> Result { + Self::execute_remote_for_scope(transport, addr, sql, current_fanout_scope()) } pub fn execute_remote_for_scope( + transport: Option<&Arc>, addr: &str, sql: &str, scope: FanoutScope, - ) -> Result { - Self::execute_remote_scoped(addr, sql, scope) + ) -> Result { + Self::execute_remote_scoped(transport, addr, sql, scope) } - fn execute_remote_scoped(addr: &str, sql: &str, scope: FanoutScope) -> Result { + fn execute_remote_scoped( + transport: Option<&Arc>, + addr: &str, + sql: &str, + scope: FanoutScope, + ) -> Result { let scope = resolve_fanout_scope(scope); - if scope == FanoutScope::Coordinator { - return Self::execute_remote_node_aggregate(addr, sql); - } - Self::execute_remote_plain(addr, sql) - } - - /// Ask a node aggregator to fan in on-node ranks (``POST /apis/cluster/query``). - fn execute_remote_node_aggregate(addr: &str, sql: &str) -> Result { #[cfg(any(test, feature = "test-utils"))] if let Some(hook) = lock_remote_query_hook().as_ref() { - return hook(addr, sql); - } - - let url = format!("http://{addr}/apis/cluster/query"); - let body = serde_json::json!({ - "expr": sql, - "cluster": true, - "hierarchical": true, - "scope": "node", - }); - let body = serde_json::to_string(&body).map_err(external)?; - let addr_owned = addr.to_string(); - let response = ureq::post(&url) - .config() - .timeout_global(Some(remote_query_timeout())) - .build() - .send(body) - .map_err(external)?; - - let status = response.status().as_u16(); - let text = response.into_body().read_to_string().map_err(external)?; - if status >= 400 { - return Err(DataFusionError::Execution(format!( - "remote node aggregate {addr_owned} failed: HTTP {status}: {text}" - ))); - } - - let value: serde_json::Value = serde_json::from_str(&text).map_err(external)?; - if let Some(err) = value.get("error").and_then(|v| v.as_str()) { - return Err(DataFusionError::Execution(format!( - "remote node aggregate {addr_owned}: {err}" - ))); + return hook(addr, sql).map(PeerQueryOutcome::complete); } - let df_value = value.get("dataframe").ok_or_else(|| { - DataFusionError::Execution(format!( - "remote node aggregate {addr_owned}: missing dataframe" - )) + let transport = transport.ok_or_else(|| { + DataFusionError::Execution("peer query transport is not configured".to_string()) })?; - serde_json::from_value(df_value.clone()).map_err(external) - } - - fn execute_remote_plain(addr: &str, sql: &str) -> Result { - #[cfg(any(test, feature = "test-utils"))] - if let Some(hook) = lock_remote_query_hook().as_ref() { - return hook(addr, sql); - } - - let url = format!("http://{addr}/query"); - let request = Message::new(Query { - expr: sql.to_string(), - ..Default::default() - }); - let body = serde_json::to_string(&request).map_err(external)?; - let addr_owned = addr.to_string(); - let response = ureq::post(&url) - .config() - .timeout_global(Some(remote_query_timeout())) - .build() - .send(body) - .map_err(external)?; - - let status = response.status().as_u16(); - let text = response.into_body().read_to_string().map_err(external)?; - if status >= 400 { - return Err(DataFusionError::Execution(format!( - "remote query {addr_owned} failed: HTTP {status}: {text}" - ))); - } - - let msg: Message = serde_json::from_str(&text).map_err(external)?; - match msg.payload { - QueryDataFormat::DataFrame(df) => Ok(df), - QueryDataFormat::Nil => Ok(DataFrame::default()), - QueryDataFormat::Error(err) => Err(DataFusionError::Execution(format!( - "remote query {addr_owned}: {}", - err.message - ))), - QueryDataFormat::TimeSeries(_) => Err(DataFusionError::NotImplemented( - "remote timeseries query not supported".into(), - )), - } + transport.query(addr, sql, scope) } } @@ -409,4 +415,28 @@ mod fanout_strict_tests { assert!(enforce_fanout_strict(&stats).is_err()); std::env::remove_var(FANOUT_STRICT_ENV); } + + #[test] + fn coordinator_raw_scan_includes_local_leaves_and_remote_aggregators() { + let leaf = Node { + addr: "10.0.0.1:8081".into(), + rank: Some(1), + ..Default::default() + }; + let aggregator = Node { + addr: "10.0.0.2:8080".into(), + rank: Some(8), + ..Default::default() + }; + + let targets = ProbeClusterExecutor::hierarchical_scan_targets( + vec![leaf.clone()], + vec![aggregator.clone()], + ); + assert_eq!(targets.len(), 2); + assert_eq!(targets[0].node.addr, leaf.addr); + assert_eq!(targets[0].scope, FanoutScope::Node); + assert_eq!(targets[1].node.addr, aggregator.addr); + assert_eq!(targets[1].scope, FanoutScope::Coordinator); + } } diff --git a/probing/core/src/core/federation/convert.rs b/probing/core/src/core/federation/convert.rs index 4cdfd153..b5cc3d65 100644 --- a/probing/core/src/core/federation/convert.rs +++ b/probing/core/src/core/federation/convert.rs @@ -244,33 +244,8 @@ pub fn dataframe_to_record_batch( return Ok(RecordBatch::new_empty(Arc::new(Schema::empty()))); } - let mut tags = federation_tags_for_endpoint(host, addr); - if let Some(rank) = rank { - tags.rank = rank; - } - let mut columns = Vec::with_capacity(df.cols.len() + FEDERATION_TAG_COLUMNS.len()); - let mut fields = Vec::with_capacity(df.names.len() + FEDERATION_TAG_COLUMNS.len()); - - for (name, col) in df.names.iter().zip(df.cols.iter()) { - fields.push(Field::new(name, array_data_type(col), true)); - columns.push(seq_to_array(col)?); - } - - let rows = df.len(); - fields.push(Field::new(PROBE_HOST_COL, DataType::Utf8, false)); - fields.push(Field::new(PROBE_ADDR_COL, DataType::Utf8, false)); - fields.push(Field::new(PROBE_RANK_COL, DataType::Int32, true)); - fields.push(Field::new(PROBE_NODE_RANK_COL, DataType::Int32, true)); - fields.push(Field::new(PROBE_LOCAL_RANK_COL, DataType::Int32, true)); - fields.push(Field::new(PROBE_ROLE_COL, DataType::Utf8, true)); - columns.push(Arc::new(StringArray::from(vec![tags.host; rows]))); - columns.push(Arc::new(StringArray::from(vec![tags.addr; rows]))); - columns.push(Arc::new(Int32Array::from(vec![tags.rank; rows]))); - columns.push(Arc::new(Int32Array::from(vec![tags.node_rank; rows]))); - columns.push(Arc::new(Int32Array::from(vec![tags.local_rank; rows]))); - columns.push(Arc::new(StringArray::from(vec![tags.role; rows]))); - - record_batch(Schema::new(fields), columns, "dataframe conversion failed") + let batch = proto_dataframe_to_record_batch(df)?; + tag_record_batch(batch, host, addr, rank) } pub fn tag_record_batch( @@ -495,6 +470,28 @@ mod tests { } } + #[test] + fn dataframe_conversion_preserves_existing_federation_tags() { + let df = DataFrame { + names: vec!["value".into(), PROBE_ADDR_COL.into(), PROBE_RANK_COL.into()], + cols: vec![ + Seq::SeqI32(vec![7]), + Seq::SeqText(vec!["leaf:8081".into()]), + Seq::SeqI32(vec![1]), + ], + size: 1, + }; + + let batch = dataframe_to_record_batch(&df, "aggregator", "node:8080", Some(0)).unwrap(); + assert_eq!(batch.schema().fields().len(), 7); + let addr = batch + .column(batch.schema().index_of(PROBE_ADDR_COL).unwrap()) + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(addr.value(0), "leaf:8081"); + } + #[test] fn extend_projection_honors_explicit_selection() { let local = Arc::new(Schema::new(vec![Field::new( diff --git a/probing/core/src/core/federation/fanout_scope.rs b/probing/core/src/core/federation/fanout_scope.rs index fc74238b..f9ef42da 100644 --- a/probing/core/src/core/federation/fanout_scope.rs +++ b/probing/core/src/core/federation/fanout_scope.rs @@ -40,6 +40,27 @@ pub struct FanoutStats { pub nodes_failed: Vec, /// Peer partial DataFrames dropped during coordinator merge (conversion failure). pub peer_batches_dropped: usize, + /// True when any descendant reported incomplete data, even if it did not + /// provide a concrete failed-node or dropped-batch count. + pub partial: bool, +} + +impl FanoutStats { + /// A successful transport response from one leaf endpoint. + pub fn complete_node() -> Self { + Self { + nodes_succeeded: 1, + ..Self::default() + } + } + + /// Merge a complete child-subtree outcome into the current request. + pub fn absorb(&mut self, child: Self) { + self.nodes_succeeded += child.nodes_succeeded; + self.nodes_failed.extend(child.nodes_failed); + self.peer_batches_dropped += child.peer_batches_dropped; + self.partial |= child.partial; + } } /// Shareable stats sink captured by physical execution plans. @@ -62,12 +83,25 @@ impl FanoutStatsHandle { *self.lock() = stats; } + #[cfg(test)] pub(crate) fn record_success(&self) { self.lock().nodes_succeeded += 1; } pub(crate) fn record_failure(&self, addr: &str) { - self.lock().nodes_failed.push(addr.to_string()); + let mut stats = self.lock(); + stats.nodes_failed.push(addr.to_string()); + stats.partial = true; + } + + pub(crate) fn record_batch_drop(&self) { + let mut stats = self.lock(); + stats.peer_batches_dropped += 1; + stats.partial = true; + } + + pub(crate) fn absorb(&self, child: FanoutStats) { + self.lock().absorb(child); } pub(crate) fn take(&self) -> FanoutStats { diff --git a/probing/core/src/core/federation/federated_scan_exec.rs b/probing/core/src/core/federation/federated_scan_exec.rs index 5e4e5213..a2f25d9f 100644 --- a/probing/core/src/core/federation/federated_scan_exec.rs +++ b/probing/core/src/core/federation/federated_scan_exec.rs @@ -25,11 +25,13 @@ use datafusion::physical_plan::{ SendableRecordBatchStream, }; use futures::StreamExt; -use probing_proto::prelude::{DataFrame, Node}; +use probing_proto::prelude::DataFrame; -use super::cluster_executor::{fanout_strict_enabled, ProbeClusterExecutor}; +use super::cluster_executor::{ + fanout_strict_enabled, FederatedScanTarget, PeerQueryTransport, ProbeClusterExecutor, +}; use super::convert::{align_batch_to_schema, dataframe_to_record_batch, tag_record_batch}; -use super::fanout_scope::{FanoutScope, FanoutStatsHandle}; +use super::fanout_scope::FanoutStatsHandle; fn log_federated_peer_failure(context: &str, addr_tag: &str, detail: &str) { if fanout_strict_enabled() { @@ -52,15 +54,14 @@ pub struct FederatedScanExec { projection: Vec, /// `probe.*` SQL executed on each peer node. remote_sql: String, - /// Snapshot of peer nodes captured at planning time (one partition each). - remote_nodes: Vec, - /// How [`Self::execute_remote`] talks to each peer (plain SQL vs node aggregate API). - remote_scope: FanoutScope, + /// Snapshot of peer nodes and their routing scope (one partition each). + remote_targets: Vec, local_host: String, local_addr: String, local_rank: Option, /// Request-owned sink; explicit because DataFusion may poll partitions in child tasks. fanout_stats: FanoutStatsHandle, + transport: Option>, properties: Arc, } @@ -71,19 +72,19 @@ impl FederatedScanExec { output_schema: SchemaRef, projection: Vec, remote_sql: String, - remote_nodes: Vec, - remote_scope: FanoutScope, + remote_targets: Vec, local_host: String, local_addr: String, local_rank: Option, fanout_stats: FanoutStatsHandle, + transport: Option>, ) -> Result { let projected_schema = Arc::new( output_schema .project(&projection) .map_err(DataFusionError::from)?, ); - let num_partitions = 1 + remote_nodes.len(); + let num_partitions = 1 + remote_targets.len(); let properties = PlanProperties::new( EquivalenceProperties::new(projected_schema.clone()), Partitioning::UnknownPartitioning(num_partitions), @@ -96,12 +97,12 @@ impl FederatedScanExec { projected_schema, projection, remote_sql, - remote_nodes, - remote_scope, + remote_targets, local_host, local_addr, local_rank, fanout_stats, + transport, properties: Arc::new(properties), }) } @@ -126,37 +127,43 @@ impl FederatedScanExec { } fn execute_remote(&self, node_index: usize) -> Result { - let node = self.remote_nodes.get(node_index).ok_or_else(|| { + let target = self.remote_targets.get(node_index).ok_or_else(|| { DataFusionError::Internal(format!( "FederatedScanExec: no peer node at index {node_index}" )) })?; - let addr_query = node.addr.clone(); - let addr_tag = node.addr.clone(); - let host = if node.host.is_empty() { - node.addr.clone() + let addr_query = target.node.addr.clone(); + let addr_tag = target.node.addr.clone(); + let host = if target.node.host.is_empty() { + target.node.addr.clone() } else { - node.host.clone() + target.node.host.clone() }; - let rank = node.rank; + let rank = target.node.rank; let sql = self.remote_sql.clone(); - let remote_scope = self.remote_scope; + let remote_scope = target.scope; let full = self.output_schema.clone(); let projection = self.projection.clone(); let projected_schema = self.projected_schema.clone(); let fanout_stats = self.fanout_stats.clone(); + let transport = self.transport.clone(); // Best-effort fetch: failures (network, conversion) drop the node from // the result set and are recorded in the fan-out stats rather than // failing the whole query, matching the legacy partial-result behavior. let fut = async move { let joined = tokio::task::spawn_blocking(move || { - ProbeClusterExecutor::execute_remote_for_scope(&addr_query, &sql, remote_scope) + ProbeClusterExecutor::execute_remote_for_scope( + transport.as_ref(), + &addr_query, + &sql, + remote_scope, + ) }) .await; match joined { - Ok(Ok(df)) => match finalize_remote_dataframe( - &df, + Ok(Ok(outcome)) => match finalize_remote_dataframe( + &outcome.dataframe, &host, &addr_tag, rank, @@ -164,16 +171,17 @@ impl FederatedScanExec { &projection, ) { Ok(opt) => { - fanout_stats.record_success(); + fanout_stats.absorb(outcome.stats); opt } Err(err) => { + fanout_stats.absorb(outcome.stats); + fanout_stats.record_batch_drop(); log_federated_peer_failure( "federated scan dropped", &addr_tag, &err.to_string(), ); - fanout_stats.record_failure(&addr_tag); None } }, @@ -246,7 +254,7 @@ impl DisplayAs for FederatedScanExec { write!( f, "FederatedScanExec: peers={}, remote_sql={}", - self.remote_nodes.len(), + self.remote_targets.len(), self.remote_sql ) } @@ -278,12 +286,12 @@ impl ExecutionPlan for FederatedScanExec { projected_schema: self.projected_schema.clone(), projection: self.projection.clone(), remote_sql: self.remote_sql.clone(), - remote_nodes: self.remote_nodes.clone(), - remote_scope: self.remote_scope, + remote_targets: self.remote_targets.clone(), local_host: self.local_host.clone(), local_addr: self.local_addr.clone(), local_rank: self.local_rank, fanout_stats: self.fanout_stats.clone(), + transport: self.transport.clone(), properties: self.properties.clone(), })) } diff --git a/probing/core/src/core/federation/global_catalog.rs b/probing/core/src/core/federation/global_catalog.rs index 09f04e0f..40a0b3df 100644 --- a/probing/core/src/core/federation/global_catalog.rs +++ b/probing/core/src/core/federation/global_catalog.rs @@ -7,6 +7,7 @@ use datafusion::error::Result; use datafusion::execution::context::SessionContext; use super::global_table::GlobalFederatedTable; +use super::PeerQueryTransport; pub const GLOBAL_CATALOG: &str = "global"; const SKIP_SCHEMAS: &[&str] = &["information_schema"]; @@ -16,11 +17,14 @@ const SKIP_SCHEMAS: &[&str] = &["information_schema"]; /// Schemas and tables are discovered on demand at query time, so tables registered /// after engine build (e.g. new Python extensions or memtable files) are visible /// under `global.*` without refreshing or rebuilding the catalog. -pub fn install_global_catalog(ctx: &SessionContext) -> Result<()> { +pub fn install_global_catalog( + ctx: &SessionContext, + transport: Option>, +) -> Result<()> { let shared_ctx = Arc::new(ctx.clone()); ctx.register_catalog( GLOBAL_CATALOG, - Arc::new(DynamicGlobalCatalog::new(shared_ctx)), + Arc::new(DynamicGlobalCatalog::new(shared_ctx, transport)), ); Ok(()) } @@ -28,6 +32,7 @@ pub fn install_global_catalog(ctx: &SessionContext) -> Result<()> { /// Read-only view over `probe` that exposes federated wrappers for every table. struct DynamicGlobalCatalog { ctx: Arc, + transport: Option>, } impl std::fmt::Debug for DynamicGlobalCatalog { @@ -39,8 +44,8 @@ impl std::fmt::Debug for DynamicGlobalCatalog { } impl DynamicGlobalCatalog { - fn new(ctx: Arc) -> Self { - Self { ctx } + fn new(ctx: Arc, transport: Option>) -> Self { + Self { ctx, transport } } fn probe_catalog(&self) -> Option> { @@ -66,7 +71,11 @@ impl CatalogProvider for DynamicGlobalCatalog { } let probe = self.probe_catalog()?; let inner = probe.schema(name)?; - Some(Arc::new(GlobalSchemaProvider::new(name.to_string(), inner))) + Some(Arc::new(GlobalSchemaProvider::new( + name.to_string(), + inner, + self.transport.clone(), + ))) } } @@ -75,11 +84,20 @@ impl CatalogProvider for DynamicGlobalCatalog { struct GlobalSchemaProvider { schema_name: String, inner: Arc, + transport: Option>, } impl GlobalSchemaProvider { - fn new(schema_name: String, inner: Arc) -> Self { - Self { schema_name, inner } + fn new( + schema_name: String, + inner: Arc, + transport: Option>, + ) -> Self { + Self { + schema_name, + inner, + transport, + } } } @@ -97,6 +115,7 @@ impl SchemaProvider for GlobalSchemaProvider { &self.schema_name, name, local, + self.transport.clone(), )))) } diff --git a/probing/core/src/core/federation/global_table.rs b/probing/core/src/core/federation/global_table.rs index e2bc5275..5c145160 100644 --- a/probing/core/src/core/federation/global_table.rs +++ b/probing/core/src/core/federation/global_table.rs @@ -1,6 +1,6 @@ use std::sync::Arc; -use super::cluster_executor::{reset_fanout_stats, ProbeClusterExecutor}; +use super::cluster_executor::{reset_fanout_stats, PeerQueryTransport, ProbeClusterExecutor}; use super::convert::{ cluster_rank_for_endpoint, extend_projection_with_probe_tags, federated_output_schema, }; @@ -59,6 +59,7 @@ pub struct GlobalFederatedTable { schema_name: String, table_name: String, local: Arc, + transport: Option>, } impl GlobalFederatedTable { @@ -66,11 +67,13 @@ impl GlobalFederatedTable { schema_name: impl Into, table_name: impl Into, local: Arc, + transport: Option>, ) -> Self { Self { schema_name: schema_name.into(), table_name: table_name.into(), local, + transport, } } } @@ -101,9 +104,13 @@ impl TableProvider for GlobalFederatedTable { let local_rank = cluster_rank_for_endpoint(&host, &addr); reset_fanout_stats(); - let (remote_nodes, remote_scope) = ProbeClusterExecutor::federated_scan_targets(); + let remote_targets = ProbeClusterExecutor::federated_scan_targets()?; // With peers registered, LIMIT is global top-K at the coordinator only. - let scan_limit = if remote_nodes.is_empty() { limit } else { None }; + let scan_limit = if remote_targets.is_empty() { + limit + } else { + None + }; // Local scan stays lazy; coalesce to a single partition so the federated // plan can expose it as partition 0 without losing rows from sub-partitions. @@ -130,12 +137,12 @@ impl TableProvider for GlobalFederatedTable { output_schema, scan_projection, remote_sql, - remote_nodes, - remote_scope, + remote_targets, host, addr, local_rank, current_fanout_stats_handle(), + self.transport.clone(), )?; Ok(Arc::new(exec)) } diff --git a/probing/core/src/core/federation/mod.rs b/probing/core/src/core/federation/mod.rs index a206f3cd..3d19bd89 100644 --- a/probing/core/src/core/federation/mod.rs +++ b/probing/core/src/core/federation/mod.rs @@ -18,7 +18,8 @@ pub use cluster_executor::set_remote_query_hook; pub use cluster_executor::{ check_fanout_strict, enforce_fanout_strict, fanout_stats_partial, fanout_strict_enabled, remote_fanout_concurrency, remote_query_timeout, reset_fanout_stats, set_fanout_stats, - take_fanout_stats, FanoutStats, ProbeClusterExecutor, RemoteFanoutResult, + take_fanout_stats, FanoutStats, PeerQueryOutcome, PeerQueryTransport, ProbeClusterExecutor, + RemoteFanoutResult, }; pub use convert::{ cluster_local_rank_for_endpoint, cluster_node_rank_for_endpoint, cluster_rank_for_endpoint, diff --git a/probing/core/src/core/memtable_sql.rs b/probing/core/src/core/memtable_sql.rs index d286e812..1a9d7cb5 100644 --- a/probing/core/src/core/memtable_sql.rs +++ b/probing/core/src/core/memtable_sql.rs @@ -28,7 +28,7 @@ use std::collections::{BTreeSet, HashSet}; use std::panic::AssertUnwindSafe; -use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::{Arc, Mutex, MutexGuard}; use std::thread::JoinHandle; use std::time::Duration; @@ -55,7 +55,8 @@ use once_cell::sync::Lazy; use probing_memtable::discover::{default_dir, MappedFile}; use probing_memtable::memc::{ - ColdStats, ColdStore, ColumnData, Compactor, CompactorConfig, SegmentReader, + ColdStats, ColdStore, ColumnData, Compactor, CompactorConfig, CompactorRuntimeStats, + SegmentReader, }; use probing_memtable::{ detect_table, DType, MemTableView, MemhView, MemtableError, TableKind, TypedValue, @@ -626,19 +627,26 @@ impl TableProvider for RingMmapTable { // ── Cold segments (MEMC) → Arrow, with two-level time pruning ───────── +type ColdCoverage = HashSet<(u64, usize, u64)>; +type ColdScanResult = (Vec, ColdCoverage); + /// `.memc` segment paths in `dir`, or empty if the dir does not exist. /// Read-only: never creates the directory (unlike `ColdStore::open`). -fn cold_segment_paths(dir: &std::path::Path) -> Vec { +fn cold_segment_paths(dir: &std::path::Path) -> std::io::Result> { let mut out = Vec::new(); - if let Ok(entries) = std::fs::read_dir(dir) { - for e in entries.flatten() { - let p = e.path(); - if p.extension().and_then(|s| s.to_str()) == Some("memc") { - out.push(p); + match std::fs::read_dir(dir) { + Ok(entries) => { + for entry in entries { + let p = entry?.path(); + if p.extension().and_then(|s| s.to_str()) == Some("memc") { + out.push(p); + } } } + Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} + Err(error) => return Err(error), } - out + Ok(out) } /// One decoded cold column → an Arrow array (schema order is preserved). @@ -670,13 +678,16 @@ fn cold_scan( table: &str, schema: &SchemaRef, bounds: &TsBounds, -) -> (Vec, HashSet<(u64, usize, u64)>) { +) -> DfResult { let mut out = Vec::new(); let mut covered: HashSet<(u64, usize, u64)> = HashSet::new(); - for path in cold_segment_paths(dir) { - let Ok(reader) = SegmentReader::open(&path) else { - continue; // unreadable/foreign file: skip rather than fail the scan - }; + for path in cold_segment_paths(dir).map_err(DataFusionError::IoError)? { + let reader = SegmentReader::open(&path).map_err(|error| { + DataFusionError::External(Box::new(std::io::Error::new( + error.kind(), + format!("failed to read MEMC segment {}: {error}", path.display()), + ))) + })?; if let Some((smin, smax)) = reader.ts_range() { if bounds.lower.is_some_and(|lo| smax < lo) || bounds.upper.is_some_and(|hi| smin > hi) { @@ -688,26 +699,28 @@ fn cold_scan( }; let pages = reader.pages(); for idx in reader.pages_in_range(tid, bounds.lower, bounds.upper) { + let cols = reader.read_page(idx).map_err(|error| { + DataFusionError::External(Box::new(std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!( + "failed to decode MEMC page {idx} in {}: {error}", + path.display() + ), + ))) + })?; + let arrays: Vec = cols.into_iter().map(cold_column_to_array).collect(); + let batch = RecordBatch::try_new(Arc::clone(schema), arrays)?; if let Some(p) = pages.get(idx) { if p.source_chunk != probing_memtable::memc::SOURCE_CHUNK_NONE { covered.insert((p.source_instance, p.source_chunk as usize, p.source_gen)); } } - match reader.read_page(idx) { - Ok(cols) => { - let arrays: Vec = - cols.into_iter().map(cold_column_to_array).collect(); - match RecordBatch::try_new(Arc::clone(schema), arrays) { - Ok(b) if b.num_rows() > 0 => out.push(b), - Ok(_) => {} - Err(e) => log::error!("cold page {idx} → RecordBatch failed: {e}"), - } - } - Err(e) => log::debug!("cold page {idx} decode skipped: {e}"), + if batch.num_rows() > 0 { + out.push(batch); } } } - (out, covered) + Ok((out, covered)) } /// [`TableProvider`] unioning a hot ring with its cold MEMC segments under one @@ -760,7 +773,7 @@ impl TableProvider for HotColdTable { limit: Option, ) -> DfResult> { let bounds = self.hot.bounds_for(filters); - let (cold, covered) = cold_scan(&self.cold_dir, &self.table, &self.schema, &bounds); + let (cold, covered) = cold_scan(&self.cold_dir, &self.table, &self.schema, &bounds)?; // Drop hot chunks already in cold so each row is counted once. let hot = self.hot.pruned_batches_excluding(&bounds, &covered); @@ -1140,21 +1153,24 @@ impl ColdRuntimeConfig { /// returned as `(on-disk basename, path)`. The basename is the cold table /// identity (matching the SQL read path), so names never collide across /// schemas. The `cold/` subdir is skipped (it is a directory, not a file). -fn cold_source_candidates() -> Vec<(String, std::path::PathBuf)> { +fn cold_source_candidates() -> std::io::Result> { let mut out = Vec::new(); - if let Ok(entries) = std::fs::read_dir(self_dir()) { - for e in entries.flatten() { - let p = e.path(); - if !p.is_file() { - continue; - } - let name = e.file_name().to_string_lossy().to_string(); - if classify_mmap_basename(&name).is_some() { - out.push((name, p)); - } + let entries = match std::fs::read_dir(self_dir()) { + Ok(entries) => entries, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(out), + Err(error) => return Err(error), + }; + for entry in entries { + let entry = entry?; + if !entry.metadata()?.is_file() { + continue; + } + let name = entry.file_name().to_string_lossy().to_string(); + if classify_mmap_basename(&name).is_some() { + out.push((name, entry.path())); } } - out + Ok(out) } /// Process-global owner of the background hot→cold compactor thread. @@ -1167,6 +1183,8 @@ fn cold_source_candidates() -> Vec<(String, std::path::PathBuf)> { pub struct ColdCompactor { running: Arc, handle: Mutex>>, + error_count: Arc, + last_error: Arc>>, } fn lock_compactor_handle( @@ -1175,11 +1193,28 @@ fn lock_compactor_handle( crate::sync::lock_mutex(m, "ColdCompactor handle") } +fn lock_compactor_last_error(m: &Mutex>) -> MutexGuard<'_, Option> { + crate::sync::lock_mutex(m, "ColdCompactor last_error") +} + +fn record_cold_compactor_error( + error_count: &AtomicU64, + last_error: &Mutex>, + message: impl Into, +) { + let message = message.into(); + log::warn!("cold compactor: {message}"); + error_count.fetch_add(1, Ordering::Relaxed); + *lock_compactor_last_error(last_error) = Some(message); +} + impl ColdCompactor { pub fn instance() -> &'static Self { static INSTANCE: Lazy = Lazy::new(|| ColdCompactor { running: Arc::new(AtomicBool::new(false)), handle: Mutex::new(None), + error_count: Arc::new(AtomicU64::new(0)), + last_error: Arc::new(Mutex::new(None)), }); &INSTANCE } @@ -1188,6 +1223,19 @@ impl ColdCompactor { self.running.load(Ordering::Acquire) } + /// Background write/roll/retention health for the current run. + pub fn runtime_stats(&self) -> CompactorRuntimeStats { + CompactorRuntimeStats { + error_count: self.error_count.load(Ordering::Relaxed), + last_error: lock_compactor_last_error(&self.last_error).clone(), + } + } + + fn reset_runtime_stats(&self) { + self.error_count.store(0, Ordering::Relaxed); + *lock_compactor_last_error(&self.last_error) = None; + } + /// (Re)apply `cfg`: stop any running thread, then start a fresh one when /// `cfg.enabled`. Idempotent and the single entry point for the config /// surface, so changing a knob simply restarts with the new settings. @@ -1202,11 +1250,16 @@ impl ColdCompactor { if self.running.swap(true, Ordering::SeqCst) { return; // already running } + self.reset_runtime_stats(); let dir = cold_dir(); let store = match ColdStore::open(&dir) { Ok(s) => s, Err(e) => { - log::error!("cold compactor: cannot open {}: {e}", dir.display()); + record_cold_compactor_error( + &self.error_count, + &self.last_error, + format!("cannot open {}: {e}", dir.display()), + ); self.running.store(false, Ordering::SeqCst); return; } @@ -1215,42 +1268,99 @@ impl ColdCompactor { // Exactly-once across restarts: recover per-chunk watermarks from any // segments already on disk before draining. if let Err(e) = compactor.prime_from_cold() { - log::warn!("cold compactor: prime_from_cold failed: {e}"); + record_cold_compactor_error( + &self.error_count, + &self.last_error, + format!("prime_from_cold failed: {e}"), + ); + self.running.store(false, Ordering::SeqCst); + return; } let running = self.running.clone(); + let error_count = self.error_count.clone(); + let last_error = self.last_error.clone(); let poll = cfg.poll; match std::thread::Builder::new() .name("memc-compactor".into()) .spawn(move || { while running.load(Ordering::SeqCst) { - for (name, path) in cold_source_candidates() { - let Ok(mapped) = MappedFile::open(&path) else { - continue; - }; - if !matches!(detect_table(mapped.as_bytes()), Some(TableKind::Ring)) { - continue; // only ring tables tier to cold - } - if let Ok(view) = MemTableView::new(mapped.as_bytes()) { - if let Err(e) = compactor.drain_view(&name, &view) { - log::debug!("cold compactor: drain {name}: {e}"); + match cold_source_candidates() { + Ok(candidates) => { + for (name, path) in candidates { + let mapped = match MappedFile::open(&path) { + Ok(mapped) => mapped, + Err(error) => { + record_cold_compactor_error( + &error_count, + &last_error, + format!("open source {}: {error}", path.display()), + ); + continue; + } + }; + if !matches!(detect_table(mapped.as_bytes()), Some(TableKind::Ring)) + { + continue; // only ring tables tier to cold + } + match MemTableView::new(mapped.as_bytes()) { + Ok(view) => { + if let Err(error) = compactor.drain_view(&name, &view) { + record_cold_compactor_error( + &error_count, + &last_error, + format!("drain {name}: {error}"), + ); + } + } + Err(error) => record_cold_compactor_error( + &error_count, + &last_error, + format!("open table view {name}: {error}"), + ), + } } } + Err(error) => record_cold_compactor_error( + &error_count, + &last_error, + format!("discover hot sources: {error}"), + ), + } + if let Err(error) = compactor.maybe_roll_on_age() { + record_cold_compactor_error( + &error_count, + &last_error, + format!("age-triggered roll: {error}"), + ); + } + if let Err(error) = compactor.enforce_checked() { + record_cold_compactor_error( + &error_count, + &last_error, + format!("retention enforcement: {error}"), + ); } - let _ = compactor.maybe_roll_on_age(); - let _ = compactor.enforce(); sleep_interruptible(&running, poll); } // Final flush so the last open segment is sealed on shutdown. - if let Err(e) = compactor.flush() { - log::debug!("cold compactor: final flush: {e}"); + if let Err(error) = compactor.flush() { + record_cold_compactor_error( + &error_count, + &last_error, + format!("final flush: {error}"), + ); } }) { Ok(handle) => { *lock_compactor_handle(&self.handle) = Some(handle); } Err(e) => { - log::error!("cold compactor: failed to spawn background thread: {e}"); + record_cold_compactor_error( + &self.error_count, + &self.last_error, + format!("failed to spawn background thread: {e}"), + ); self.running.store(false, Ordering::SeqCst); } } @@ -1262,7 +1372,13 @@ impl ColdCompactor { return; } if let Some(h) = lock_compactor_handle(&self.handle).take() { - let _ = h.join(); + if h.join().is_err() { + record_cold_compactor_error( + &self.error_count, + &self.last_error, + "background thread panicked", + ); + } } } @@ -1753,6 +1869,29 @@ mod tests { .unwrap(); assert_eq!(collect_i32(&span), vec![2, 3, 4, 5]); + // A corrupted sealed segment must fail the query instead of returning + // only the still-resident hot rows. + let segment = std::fs::read_dir(&cold) + .unwrap() + .map(|entry| entry.unwrap().path()) + .find(|path| path.extension().and_then(|ext| ext.to_str()) == Some("memc")) + .expect("cold segment"); + let mut bytes = std::fs::read(&segment).unwrap(); + let footer_off = u64::from_le_bytes(bytes[32..40].try_into().unwrap()) as usize; + bytes[footer_off + 12] ^= 0xff; + std::fs::write(&segment, bytes).unwrap(); + let error = ctx + .sql("SELECT v FROM hc_demo ORDER BY v") + .await + .unwrap() + .collect() + .await + .unwrap_err(); + assert!( + error.to_string().contains("MEMC footer checksum mismatch"), + "unexpected error: {error}" + ); + drop(t); } @@ -1874,6 +2013,38 @@ mod tests { } } + #[test] + fn cold_compactor_start_failure_is_observable() { + let _lock = PROBING_DATA_DIR_LOCK.lock().unwrap(); + let tmp = tempfile::tempdir().unwrap(); + let original = std::env::var("PROBING_DATA_DIR").ok(); + std::env::set_var("PROBING_DATA_DIR", tmp.path()); + + ColdCompactor::instance().stop(); + std::fs::create_dir_all(self_dir()).unwrap(); + std::fs::File::create(cold_dir()).unwrap(); + ColdCompactor::instance().apply(ColdRuntimeConfig { + enabled: true, + ..Default::default() + }); + + let stats = ColdCompactor::instance().runtime_stats(); + assert!(!ColdCompactor::instance().is_running()); + assert_eq!(stats.error_count, 1); + assert!( + stats + .last_error + .as_deref() + .is_some_and(|message| message.contains("cannot open")), + "unexpected runtime stats: {stats:?}" + ); + + match original { + Some(value) => std::env::set_var("PROBING_DATA_DIR", value), + None => std::env::remove_var("PROBING_DATA_DIR"), + } + } + #[tokio::test] async fn engine_catalog_query_unions_cold_tier() { let _lock = PROBING_DATA_DIR_LOCK.lock().unwrap(); diff --git a/probing/core/src/core/probe_extension.rs b/probing/core/src/core/probe_extension.rs index c2e38560..fdfc4003 100644 --- a/probing/core/src/core/probe_extension.rs +++ b/probing/core/src/core/probe_extension.rs @@ -8,7 +8,6 @@ use std::sync::Arc; use async_trait::async_trait; use datafusion::config::{ConfigExtension, ExtensionOptions}; -use once_cell::sync::Lazy; use tokio::sync::{Mutex, RwLock}; use super::error::EngineError; @@ -17,12 +16,6 @@ use crate::config; /// Shared probe extension instances keyed by extension name. pub type ProbeExtensionMap = BTreeMap>>; -/// Global probe extension registry. -/// -/// Shared storage for [`ProbeExtension`] instances; [`ProbeExtensionManager`] operates on this map. -pub static PROBE_EXTENSIONS: Lazy> = - Lazy::new(|| RwLock::new(BTreeMap::new())); - #[derive(Clone, Debug, Default)] pub enum Maybe { Just(T), @@ -221,22 +214,31 @@ pub trait ProbeExtension: Debug + Send + Sync + ProbeExtensionCall { /// // Or if used in a #[tokio::test] or #[tokio::main] annotated function: /// // manager_usage_example().await.unwrap(); /// ``` -/// Engine extension manager that operates on the global extensions registry. +/// Engine-scoped extension manager. /// -/// This struct no longer holds extensions directly. Instead, it operates -/// on the global `PROBE_EXTENSIONS` registry, allowing multiple instances to -/// work with the same set of extensions. -#[derive(Clone, Debug, Default)] -pub struct ProbeExtensionManager; +/// Clones share one engine's registry, while independently built engines keep +/// separate extension instances and configuration state. +#[derive(Clone, Debug)] +pub struct ProbeExtensionManager { + extensions: Arc>, +} + +impl Default for ProbeExtensionManager { + fn default() -> Self { + Self { + extensions: Arc::new(RwLock::new(BTreeMap::new())), + } + } +} impl ProbeExtensionManager { - /// Register an extension in the global extensions registry. + /// Register an extension in this manager's engine-scoped registry. pub async fn register( &mut self, name: String, extension: Arc>, ) { - PROBE_EXTENSIONS.write().await.insert(name, extension); + self.extensions.write().await.insert(name, extension); } /// Extract namespace from extension name by removing "extension" suffix and converting to lowercase @@ -254,7 +256,7 @@ impl ProbeExtensionManager { /// ConfigStore is not updated by this method. pub async fn set_option(&mut self, key: &str, value: &str) -> Result<(), EngineError> { let extensions_clone: Vec<_> = { - let extensions = PROBE_EXTENSIONS.read().await; + let extensions = self.extensions.read().await; extensions.values().cloned().collect() }; // Lock is released here @@ -306,7 +308,7 @@ impl ProbeExtensionManager { pub async fn get_option(&self, key: &str) -> Result { let extensions_clone: Vec<_> = { - let extensions = PROBE_EXTENSIONS.read().await; + let extensions = self.extensions.read().await; extensions.values().cloned().collect() }; // Lock is released here @@ -332,7 +334,7 @@ impl ProbeExtensionManager { pub async fn options(&self) -> Vec { let mut all_options = Vec::new(); let extensions_clone: Vec<_> = { - let extensions = PROBE_EXTENSIONS.read().await; + let extensions = self.extensions.read().await; extensions.values().cloned().collect() }; // Lock is released here @@ -350,7 +352,7 @@ impl ProbeExtensionManager { body: &[u8], ) -> Result, EngineError> { let extensions_clone: Vec<_> = { - let extensions = PROBE_EXTENSIONS.read().await; + let extensions = self.extensions.read().await; extensions.values().cloned().collect() }; // Lock is released here @@ -416,25 +418,23 @@ impl ExtensionOptions for ProbeExtensionManager { } fn cloned(&self) -> Box { - // ProbeExtensionManager is now a zero-sized type, so cloning is trivial - Box::new(ProbeExtensionManager) + Box::new(self.clone()) } fn set(&mut self, key: &str, value: &str) -> datafusion::error::Result<()> { - let fut = self.set_option(key, value); - if let Ok(handle) = tokio::runtime::Handle::try_current() { - tokio::task::block_in_place(|| handle.block_on(fut)) - .map_err(datafusion::error::DataFusionError::from) - } else { - crate::runtime::CORE_RUNTIME - .block_on(fut) - .map_err(datafusion::error::DataFusionError::from) - } + let mut manager = self.clone(); + let key = key.to_string(); + let value = value.to_string(); + crate::runtime::block_on(async move { manager.set_option(&key, &value).await }) + .map_err(datafusion::error::DataFusionError::from)? + .map_err(datafusion::error::DataFusionError::from) } fn entries(&self) -> Vec { - let fut = async { - self.options() + let manager = self.clone(); + match crate::runtime::block_on(async move { + manager + .options() .await .iter() .map(|option| datafusion::config::ConfigEntry { @@ -443,11 +443,12 @@ impl ExtensionOptions for ProbeExtensionManager { description: option.help, }) .collect() - }; - if let Ok(handle) = tokio::runtime::Handle::try_current() { - tokio::task::block_in_place(|| handle.block_on(fut)) - } else { - crate::runtime::CORE_RUNTIME.block_on(fut) + }) { + Ok(entries) => entries, + Err(error) => { + log::error!("failed to enumerate probing extension options: {error}"); + Vec::new() + } } } } @@ -461,18 +462,12 @@ mod tests { async fn setup_test() -> tokio::sync::MutexGuard<'static, ()> { let guard = config::TEST_STATE_LOCK.lock().await; config::clear().await; - PROBE_EXTENSIONS.write().await.clear(); guard } // Helper to ensure clean state after each test async fn teardown_test() { config::clear().await; - // 确保在清空之前所有锁都已释放 - let mut extensions = PROBE_EXTENSIONS.write().await; - extensions.clear(); - // 显式释放写锁 - drop(extensions); } #[derive(Debug)] @@ -526,9 +521,11 @@ mod tests { async fn test_set_option_syncs_to_config_store() { let _state_guard = setup_test().await; - let mut manager = ProbeExtensionManager; + let mut manager = ProbeExtensionManager::default(); let extension = Arc::new(Mutex::new(TestExtension::default())); - manager.register("test".to_string(), extension).await; + manager + .register("test".to_string(), extension.clone()) + .await; // Set option through manager using set_option_with_store_update manager @@ -541,12 +538,10 @@ mod tests { assert_eq!(value, Some("new_value".to_string())); // Verify extension was updated - { - let extensions = PROBE_EXTENSIONS.read().await; - let ext_guard = extensions.get("test").unwrap().lock().await; - let value = ext_guard.get("option").unwrap(); - assert_eq!(value, "new_value"); - } // 确保锁在这里释放 + let ext_guard = extension.lock().await; + let value = ext_guard.get("option").unwrap(); + assert_eq!(value, "new_value"); + drop(ext_guard); teardown_test().await; } @@ -558,7 +553,7 @@ mod tests { // Pre-populate ConfigStore config::set("test.option", "old_value").await; - let mut manager = ProbeExtensionManager; + let mut manager = ProbeExtensionManager::default(); let extension = Arc::new(Mutex::new(TestExtension::default())); manager.register("test".to_string(), extension).await; @@ -579,7 +574,7 @@ mod tests { async fn test_set_option_unsupported_key() { let _state_guard = setup_test().await; - let mut manager = ProbeExtensionManager; + let mut manager = ProbeExtensionManager::default(); let extension = Arc::new(Mutex::new(TestExtension::default())); manager.register("test".to_string(), extension).await; @@ -610,4 +605,23 @@ mod tests { teardown_test().await; } + + #[tokio::test(flavor = "multi_thread")] + async fn independently_built_managers_do_not_share_extensions() { + let mut first = ProbeExtensionManager::default(); + let second = ProbeExtensionManager::default(); + first + .register( + "test".to_string(), + Arc::new(Mutex::new(TestExtension::default())), + ) + .await; + + first.set_option("test.option", "first").await.unwrap(); + assert_eq!(first.get_option("test.option").await.unwrap(), "first"); + assert!(matches!( + second.get_option("test.option").await, + Err(EngineError::UnsupportedOption(_)) + )); + } } diff --git a/probing/core/src/runtime.rs b/probing/core/src/runtime.rs index 31426538..575dbf8b 100644 --- a/probing/core/src/runtime.rs +++ b/probing/core/src/runtime.rs @@ -292,28 +292,16 @@ impl CoreRuntime { } } - pub fn spawn(&self, future: F) -> tokio::task::JoinHandle + pub fn spawn(&self, future: F) -> Result, RuntimeError> where F: Future + Send + 'static, F::Output: Send + 'static, { match self.ensure_runtime() { - Some(rt) => rt.spawn(future), + Some(rt) => Ok(rt.spawn(future)), None => { self.mark_degraded(); - log::error!("probing: no tokio runtime for spawn; creating per-call ephemeral"); - match tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - { - Ok(rt) => rt.spawn(future), - Err(e) => { - log::error!("probing: per-call spawn runtime build failed: {e}"); - try_ephemeral_runtime() - .expect("probing: no tokio runtime for spawn") - .spawn(future) - } - } + Err(RuntimeError::Unavailable) } } } @@ -347,30 +335,6 @@ impl CoreRuntime { ); block_on_ephemeral(future) } - - /// Prefer [`try_block_on`] when the caller can surface bridge failures. - pub fn block_on(&self, future: F) -> T - where - F: Future, - { - if let Some(rt) = &self.current_slot().runtime { - return rt.block_on(future); - } - if let Some(rt) = fallback_runtime() { - self.mark_degraded(); - return rt.block_on(future); - } - if let Some(rt) = try_ephemeral_runtime() { - self.mark_degraded(); - return rt.block_on(future); - } - self.mark_degraded(); - log::error!("probing: CoreRuntime::block_on using per-call ephemeral executor"); - block_on_ephemeral(future).unwrap_or_else(|err| { - log::error!("probing: CoreRuntime::block_on failed: {err}"); - panic!("probing: async bridge unavailable: {err}"); - }) - } } /// Shared Tokio runtime for all sync→async bridges (Python bindings, local server, etc.). @@ -695,6 +659,25 @@ where T: Send + 'static, { if is_inside_core_runtime() { + if on_native_bridge() { + let Ok(handle) = tokio::runtime::Handle::try_current() else { + return block_on_failed("nested native bridge has no Tokio runtime handle"); + }; + if !matches!( + handle.runtime_flavor(), + tokio::runtime::RuntimeFlavor::MultiThread + ) { + return block_on_failed( + "nested native bridge requires a multi-thread Tokio runtime", + ); + } + return match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + tokio::task::block_in_place(|| handle.block_on(future)) + })) { + Ok(value) => Ok(value), + Err(_) => block_on_failed("nested native bridge block_on panicked"), + }; + } return spawn_block_on_thread(future); } run_on_native_thread(move || { @@ -741,6 +724,16 @@ mod tests { assert_eq!(value, 42); } + #[test] + fn nested_block_on_preserves_native_bridge_thread() { + let stayed_on_bridge = run_on_native_bridge(|| { + block_on(async { block_on(async { on_native_bridge_thread() }) }) + }) + .expect("outer bridge") + .expect("inner bridge"); + assert!(stayed_on_bridge); + } + #[test] fn native_bridge_serializes_calls() { run_on_native_bridge(|| ()); @@ -808,14 +801,16 @@ mod tests { let native_bridge_ok = block_on(async { 6 * 7 }).is_ok_and(|value| value == 42); let completed = Arc::new(AtomicBool::new(false)); let task_completed = Arc::clone(&completed); - std::mem::drop(CORE_RUNTIME.spawn(async move { - task_completed.store(true, Ordering::Release); - })); + let spawn_ok = CORE_RUNTIME + .spawn(async move { + task_completed.store(true, Ordering::Release); + }) + .is_ok(); let spawn_deadline = std::time::Instant::now() + std::time::Duration::from_secs(3); while !completed.load(Ordering::Acquire) && std::time::Instant::now() < spawn_deadline { std::thread::sleep(std::time::Duration::from_millis(10)); } - let spawn_ok = completed.load(Ordering::Acquire); + let spawn_ok = spawn_ok && completed.load(Ordering::Acquire); // SAFETY: `_exit` avoids running inherited parent-only destructors in the child. unsafe { libc::_exit(if block_on_ok && native_bridge_ok && spawn_ok { diff --git a/probing/core/src/storage/distributed.rs b/probing/core/src/storage/distributed.rs index e9148258..701bcb8c 100644 --- a/probing/core/src/storage/distributed.rs +++ b/probing/core/src/storage/distributed.rs @@ -288,10 +288,10 @@ impl EntityStore for DistributedEntityStore { async fn list_paginated( &self, - limit: usize, offset: usize, + limit: usize, ) -> Result<(Vec, bool)> { - self.local_store().list_paginated(limit, offset).await + self.local_store().list_paginated(offset, limit).await } } @@ -426,4 +426,31 @@ mod tests { "Job should not exist after deletion" ); } + + #[tokio::test] + async fn pagination_uses_offset_then_limit_without_overflow() { + let store = setup_default_store(); + for id in ["job-1", "job-2"] { + store + .put(&ClusterJob { + id: id.to_string(), + name: id.to_string(), + tasks_count: 1, + status: "Pending".to_string(), + }) + .await + .unwrap(); + } + + let (first_page, has_more) = store.list_paginated::(0, 1).await.unwrap(); + assert_eq!(first_page.len(), 1); + assert!(has_more); + + let (past_end, has_more) = store + .list_paginated::(usize::MAX, usize::MAX) + .await + .unwrap(); + assert!(past_end.is_empty()); + assert!(!has_more); + } } diff --git a/probing/core/src/storage/mem_store.rs b/probing/core/src/storage/mem_store.rs index 4335f650..98b52f4e 100644 --- a/probing/core/src/storage/mem_store.rs +++ b/probing/core/src/storage/mem_store.rs @@ -89,7 +89,7 @@ impl EntityStore for MemoryStore { let total = all_entities.len(); let start = offset.min(total); - let end = (offset + limit).min(total); + let end = offset.saturating_add(limit).min(total); let has_more = end < total; Ok((all_entities[start..end].to_vec(), has_more)) diff --git a/probing/crates/skills/Cargo.toml b/probing/crates/skills/Cargo.toml index 2d112b45..135adb77 100644 --- a/probing/crates/skills/Cargo.toml +++ b/probing/crates/skills/Cargo.toml @@ -15,6 +15,7 @@ anyhow = { workspace = true } serde = { workspace = true } serde_json = { workspace = true } serde_yaml = "0.9" +sqlparser = { version = "0.62", default-features = false, features = ["std"] } async-trait = "0.1" pyo3 = { version = "0.29.0", optional = true, default-features = false, features = [ "macros", diff --git a/probing/crates/skills/src/api.rs b/probing/crates/skills/src/api.rs index 85d722c5..55eddf19 100644 --- a/probing/crates/skills/src/api.rs +++ b/probing/crates/skills/src/api.rs @@ -7,7 +7,10 @@ use serde::Deserialize; use serde_json::{json, Value}; use super::catalog::{load_catalog, load_intents, load_pages, CatalogEntry}; -use super::loader::{load_skill, InterpretRule, KeywordsSpec, Skill, SkillParameter, SkillStep}; +use super::loader::{ + load_skill, validate_skill_contract, InterpretRule, KeywordsSpec, RequiresSpec, Skill, + SkillParameter, SkillParameterType, SkillPlatform, SkillStep, +}; use super::routing::match_skills; pub fn skill_to_json(skill: &Skill) -> Value { @@ -21,10 +24,15 @@ pub fn skill_to_json(skill: &Skill) -> Value { "parameters": skill.parameters.iter().map(|p| { json!({ "name": p.name, + "type": p.parameter_type.as_str(), "default": yaml_to_json(&p.default), + "description": p.description, }) }).collect::>(), "steps": skill.steps.iter().map(step_to_json).collect::>(), + "requires": { + "any_tables": skill.requires.any_tables, + }, "interpretation": { "rules": skill.interpretation.iter().map(|r| { json!({ @@ -50,12 +58,15 @@ fn step_to_json(step: &SkillStep) -> Value { "title": step.title, "type": step.step_type, "sql": step.sql, + "method": step.method, "path": step.path, "view": step.view, "on_empty": step.on_empty, "empty_message": step.empty_message, "when": step.when, "cluster": step.cluster, + "platform": step.platform.map(SkillPlatform::as_str), + "action": step.action, }) } @@ -84,7 +95,9 @@ pub fn load_skill_json(id: &str) -> Result { pub fn skill_from_api(value: &Value) -> Result { let payload: SkillApiPayload = serde_json::from_value(value.clone()).context("deserialize skill")?; - Ok(payload.into_skill()) + let skill = payload.into_skill(); + validate_skill_contract(&skill).context("validate skill contract")?; + Ok(skill) } #[derive(Debug, Deserialize)] @@ -105,6 +118,8 @@ struct SkillApiPayload { #[serde(default)] steps: Vec, #[serde(default)] + requires: RequiresSpec, + #[serde(default)] interpretation: InterpretationApi, #[serde(default)] summary_template: String, @@ -115,8 +130,12 @@ struct SkillApiPayload { #[derive(Debug, Deserialize)] struct SkillParameterApi { name: String, + #[serde(rename = "type", default)] + parameter_type: SkillParameterType, #[serde(default)] default: Value, + #[serde(default)] + description: String, } #[derive(Debug, Deserialize)] @@ -128,6 +147,8 @@ struct SkillStepApi { #[serde(default)] sql: Option, #[serde(default)] + method: Option, + #[serde(default)] path: Option, #[serde(default)] view: Option, @@ -139,6 +160,10 @@ struct SkillStepApi { when: Option, #[serde(default)] cluster: Option, + #[serde(default)] + platform: Option, + #[serde(default)] + action: Option, } #[derive(Debug, Default, Deserialize)] @@ -196,7 +221,9 @@ impl SkillApiPayload { .into_iter() .map(|p| SkillParameter { name: p.name, + parameter_type: p.parameter_type, default: json_to_yaml(&p.default), + description: p.description, }) .collect(), steps: self @@ -207,12 +234,15 @@ impl SkillApiPayload { title: s.title, step_type: s.step_type, sql: s.sql, + method: s.method, path: s.path, view: s.view, on_empty: s.on_empty, empty_message: s.empty_message, when: s.when, cluster: s.cluster, + platform: s.platform, + action: s.action, }) .collect(), interpretation: self @@ -229,6 +259,7 @@ impl SkillApiPayload { summary_template: self.summary_template, next_steps: self.next_steps, variables: HashMap::new(), + requires: self.requires, } } } @@ -338,3 +369,29 @@ fn yaml_to_json(value: &serde_yaml::Value) -> Value { serde_yaml::Value::Tagged(tagged) => yaml_to_json(&tagged.value), } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn skill_json_preserves_parameter_and_platform_contracts() { + let skill = load_skill("comm_bottleneck").expect("comm_bottleneck skill"); + let value = skill_to_json(&skill); + assert_eq!(value["parameters"][0]["type"], "integer"); + assert!(value["parameters"][0]["description"].is_string()); + let rdma = value["steps"] + .as_array() + .expect("steps") + .iter() + .find(|step| step["id"] == "rdma_hint") + .expect("rdma_hint step"); + assert_eq!(rdma["platform"], "linux"); + + let roundtrip = skill_from_api(&value).expect("valid roundtrip contract"); + assert_eq!( + roundtrip.parameters[0].parameter_type, + SkillParameterType::Integer + ); + } +} diff --git a/probing/crates/skills/src/interpret.rs b/probing/crates/skills/src/interpret.rs index 129b6dd0..b334f94b 100644 --- a/probing/crates/skills/src/interpret.rs +++ b/probing/crates/skills/src/interpret.rs @@ -6,6 +6,154 @@ use probing_proto::prelude::{DataFrame, Ele}; use super::loader::InterpretRule; +/// Validate the complete interpretation-rule grammar accepted by this module. +pub fn validate_rule_expression(when: &str) -> Result<(), String> { + let parts: Vec<&str> = when + .split('|') + .map(str::trim) + .filter(|part| !part.is_empty()) + .collect(); + if parts.is_empty() { + return Err("empty interpretation expression".to_string()); + } + + let mut index = 0; + let mut has_step = false; + if let Some(step_id) = parts[0].strip_prefix("step:") { + if step_id.trim().is_empty() { + return Err("step binding requires an id".to_string()); + } + has_step = true; + index = 1; + } + + while index < parts.len() { + let clause = parts[index]; + if clause == "always" { + index += 1; + continue; + } + if let Some(predicate) = clause.strip_prefix("rows ") { + if !has_step { + return Err("rows predicate requires step:".to_string()); + } + validate_rows_predicate(predicate)?; + index += 1; + continue; + } + if clause.contains("top(row)") { + if !has_step || !valid_top_median_clause(clause) { + return Err(format!("unsupported top/median predicate `{clause}`")); + } + index += 1; + continue; + } + if let Some(column) = clause.strip_prefix("column:") { + if !has_step || column.trim().is_empty() { + return Err("column predicate requires step: and a column name".to_string()); + } + let tail = parts + .get(index + 1) + .ok_or_else(|| format!("column `{}` is missing a predicate", column.trim()))?; + if !valid_column_predicate(tail) { + return Err(format!("unsupported column predicate `{tail}`")); + } + index += 2; + continue; + } + return Err(format!("unsupported interpretation clause `{clause}`")); + } + Ok(()) +} + +fn validate_rows_predicate(predicate: &str) -> Result<(), String> { + let Some((operator, expression)) = predicate.split_once(' ') else { + return Err(format!("invalid rows predicate `{predicate}`")); + }; + if !matches!(operator, "==" | ">=" | ">" | "<=" | "<") + || !valid_numeric_expression(expression.trim()) + { + return Err(format!("invalid rows predicate `{predicate}`")); + } + Ok(()) +} + +fn valid_numeric_expression(expression: &str) -> bool { + expression.split('*').all(|term| { + let term = term.trim(); + term.parse::().is_ok() + || (term.starts_with('{') + && term.ends_with('}') + && !term[1..term.len() - 1].trim().is_empty()) + }) +} + +fn valid_column_predicate(predicate: &str) -> bool { + let predicate = predicate.trim(); + if let Some(rhs) = predicate.strip_prefix("max/min(ratio) >") { + return rhs.trim().parse::().is_ok(); + } + for prefix in ["max >", "avg >", "top >", "value >", "value <"] { + if let Some(rhs) = predicate.strip_prefix(prefix) { + return rhs.trim().parse::().is_ok(); + } + } + if let Some(rhs) = predicate.strip_prefix("value ==") { + return rhs.trim().parse::().is_ok(); + } + if let Some(rhs) = predicate.strip_prefix("value =") { + return !rhs.trim().trim_matches(['\'', '"']).is_empty(); + } + if let Some(rest) = predicate.strip_prefix("ratio(") { + let Some((columns, comparison)) = rest.split_once(')') else { + return false; + }; + let Some((numerator, denominator)) = columns.split_once('/') else { + return false; + }; + return !numerator.trim().is_empty() + && !denominator.trim().is_empty() + && comparison + .trim() + .strip_prefix('>') + .is_some_and(|rhs| rhs.trim().parse::().is_ok()); + } + if let Some(rhs) = predicate.strip_prefix("last >") { + let Some((factor, column)) = rhs.trim().split_once("* avg(") else { + return false; + }; + return factor.trim().parse::().is_ok() + && column.ends_with(')') + && !column.trim_end_matches(')').trim().is_empty(); + } + if let Some(inner) = predicate + .strip_prefix("any_contains(") + .and_then(|rest| rest.strip_suffix(')')) + { + return !inner.trim().is_empty() + && inner + .split(',') + .all(|item| !item.trim().trim_matches(['\'', '"']).is_empty()); + } + false +} + +fn valid_top_median_clause(clause: &str) -> bool { + let Some(rest) = clause.strip_prefix("top(row).") else { + return false; + }; + let Some((top_column, rhs)) = rest.split_once(" > ") else { + return false; + }; + let Some((factor, median_column)) = rhs.split_once(" * median(") else { + return false; + }; + !top_column.trim().is_empty() + && factor.trim().parse::().is_ok() + && median_column.ends_with(')') + && median_column.trim_end_matches(')').trim() == top_column.trim() +} + #[derive(Debug, Clone)] pub struct StepEvidence { pub step_id: String, @@ -212,6 +360,10 @@ fn eval_column_predicate(col_name: &str, tail: &str, ev: &StepEvidence) -> bool let value = nums.first().copied().unwrap_or(0.0); return (value - threshold).abs() < f64::EPSILON; } + if let Some(rhs) = tail.strip_prefix("value =") { + let expected = rhs.trim().trim_matches(['\'', '"']); + return texts.first().is_some_and(|value| value == expected); + } if let Some(rest) = tail.strip_prefix("ratio(") { if let Some((expr, pred)) = rest.split_once(')') { let threshold = pred @@ -451,6 +603,35 @@ mod tests { assert_eq!(findings.len(), 1); } + #[test] + fn text_value_equality_rule() { + let rules = vec![InterpretRule { + id: "propagated".into(), + when: "step:topology | column:attribution_class | value = propagated_victim".into(), + severity: "warning".into(), + message: "propagated victim".into(), + }]; + let steps = vec![StepEvidence { + step_id: "topology".into(), + row_count: 1, + dataframe: DataFrame::new( + vec!["attribution_class".into()], + vec![Seq::SeqText(vec!["propagated_victim".into()])], + ), + }]; + assert_eq!(evaluate_rules(&rules, &steps, &HashMap::new()).len(), 1); + } + + #[test] + fn validates_supported_grammar_and_rejects_unknown_syntax() { + assert!(validate_rule_expression( + "step:topology | rows >= 1 | column:attribution_class | value = propagated_victim" + ) + .is_ok()); + assert!(validate_rule_expression("step:x | column:y | slope() > 1").is_err()); + assert!(validate_rule_expression("rows > 0").is_err()); + } + #[test] fn ratio_rule() { let rules = vec![InterpretRule { diff --git a/probing/crates/skills/src/lib.rs b/probing/crates/skills/src/lib.rs index 48b99d3e..f365d204 100644 --- a/probing/crates/skills/src/lib.rs +++ b/probing/crates/skills/src/lib.rs @@ -8,6 +8,7 @@ pub mod interpret; pub mod loader; pub mod routing; pub mod runner; +mod sql_guard; #[cfg(feature = "python-bridge")] pub mod pyo3; @@ -21,7 +22,8 @@ pub use catalog::{load_catalog, load_intents, load_pages, CatalogEntry}; pub use interpret::{evaluate_rules, InterpretFinding, StepEvidence}; pub use loader::{ build_context, default_parameters, derive_variables, expand_template, list_skill_ids, - load_skill, InterpretRule, KeywordsSpec, Skill, SkillParameter, SkillStep, + load_skill, normalize_parameter_overrides, InterpretRule, KeywordsSpec, RequiresSpec, Skill, + SkillParameter, SkillParameterType, SkillPlatform, SkillStep, }; pub use routing::{ match_intent_routes, match_routed_skills, match_skills, IntentRoute, SkillRoute, diff --git a/probing/crates/skills/src/loader.rs b/probing/crates/skills/src/loader.rs index 48ae8b33..8b3af372 100644 --- a/probing/crates/skills/src/loader.rs +++ b/probing/crates/skills/src/loader.rs @@ -2,8 +2,8 @@ use std::collections::HashMap; -use anyhow::{anyhow, Result}; -use serde::Deserialize; +use anyhow::{anyhow, bail, Context, Result}; +use serde::{Deserialize, Serialize}; use super::catalog; use super::discovery; @@ -13,12 +13,17 @@ pub use discovery::all_skill_root_paths; pub use catalog::CatalogEntry; #[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] struct SkillFile { + #[serde(rename = "apiVersion")] + api_version: String, + kind: String, metadata: SkillMeta, spec: SkillSpec, } #[derive(Debug, Clone, Default, Deserialize, serde::Serialize)] +#[serde(deny_unknown_fields)] pub struct KeywordsSpec { #[serde(default)] pub zh: Vec, @@ -27,15 +32,19 @@ pub struct KeywordsSpec { } #[derive(Debug, Clone, Default, Deserialize)] +#[serde(deny_unknown_fields)] struct TriggersSpec { #[serde(default)] keywords: KeywordsSpec, } #[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] struct SkillMeta { id: String, title: String, + #[serde(default, rename = "title_en")] + _title_en: String, #[serde(default)] category: String, #[serde(default)] @@ -47,6 +56,7 @@ struct SkillMeta { } #[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] struct SkillSpec { #[serde(default)] parameters: Vec, @@ -60,16 +70,45 @@ struct SkillSpec { next_steps: Vec, #[serde(default)] variables: HashMap, + #[serde(default)] + requires: RequiresSpec, +} + +#[derive(Debug, Clone, Copy, Default, Deserialize, Serialize, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum SkillParameterType { + Integer, + Number, + Boolean, + #[default] + String, +} + +impl SkillParameterType { + pub fn as_str(self) -> &'static str { + match self { + Self::Integer => "integer", + Self::Number => "number", + Self::Boolean => "boolean", + Self::String => "string", + } + } } #[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] pub struct SkillParameter { pub name: String, + #[serde(rename = "type")] + pub parameter_type: SkillParameterType, #[serde(default)] pub default: serde_yaml::Value, + #[serde(default)] + pub description: String, } #[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] struct SkillStepRaw { id: String, title: String, @@ -78,7 +117,7 @@ struct SkillStepRaw { #[serde(default)] sql: Option, #[serde(default, rename = "method")] - _method: Option, + method: Option, #[serde(default)] path: Option, #[serde(default)] @@ -91,15 +130,21 @@ struct SkillStepRaw { when: Option, #[serde(default)] cluster: Option, + #[serde(default)] + platform: Option, + #[serde(default)] + action: Option, } #[derive(Debug, Clone, Default, Deserialize)] +#[serde(deny_unknown_fields)] struct InterpretationSpec { #[serde(default)] rules: Vec, } #[derive(Debug, Clone, Deserialize)] +#[serde(deny_unknown_fields)] pub struct InterpretRule { pub id: String, pub when: String, @@ -108,6 +153,35 @@ pub struct InterpretRule { pub message: String, } +#[derive(Debug, Clone, Default, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct RequiresSpec { + #[serde(default)] + pub any_tables: Vec, +} + +#[derive(Debug, Clone, Copy, Deserialize, Serialize, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum SkillPlatform { + Linux, + Macos, + Windows, +} + +impl SkillPlatform { + pub fn as_str(self) -> &'static str { + match self { + Self::Linux => "linux", + Self::Macos => "macos", + Self::Windows => "windows", + } + } + + pub fn is_current(self) -> bool { + self.as_str() == std::env::consts::OS + } +} + fn default_severity() -> String { "info".to_string() } @@ -126,12 +200,15 @@ pub struct SkillStep { pub title: String, pub step_type: String, pub sql: Option, + pub method: Option, pub path: Option, pub view: Option, pub on_empty: String, pub empty_message: Option, pub when: Option, pub cluster: Option, + pub platform: Option, + pub action: Option, } #[derive(Debug, Clone)] @@ -149,6 +226,7 @@ pub struct Skill { pub summary_template: String, pub next_steps: Vec, pub variables: HashMap, + pub requires: RequiresSpec, } impl Skill { @@ -167,7 +245,9 @@ pub fn list_skill_ids() -> Vec { pub fn load_skill(id: &str) -> Result { let yaml = discovery::load_fs_steps_yaml(id).ok_or_else(|| anyhow!("Unknown skill: {id}"))?; - let file: SkillFile = serde_yaml::from_str(&yaml)?; + let file: SkillFile = + serde_yaml::from_str(&yaml).with_context(|| format!("invalid skill schema for `{id}`"))?; + validate_skill_file(id, &file)?; let steps = file .spec .steps @@ -177,16 +257,19 @@ pub fn load_skill(id: &str) -> Result { title: s.title, step_type: s.step_type, sql: s.sql, + method: s.method, path: s.path, view: s.view, on_empty: s.on_empty, empty_message: s.empty_message, when: s.when, cluster: s.cluster, + platform: s.platform, + action: s.action, }) .collect(); let keywords = collect_keywords(&file.metadata); - Ok(Skill { + let skill = Skill { id: file.metadata.id, title: file.metadata.title, category: file.metadata.category, @@ -200,7 +283,184 @@ pub fn load_skill(id: &str) -> Result { summary_template: file.spec.summary_template.trim().to_string(), next_steps: file.spec.next_steps, variables: file.spec.variables, - }) + requires: file.spec.requires, + }; + validate_skill_contract(&skill)?; + Ok(skill) +} + +fn validate_default_sql(skill: &Skill) -> Result<()> { + let context = build_context(skill, &HashMap::new()); + for step in &skill.steps { + let Some(sql) = &step.sql else { continue }; + let expanded = expand_template(sql, &context); + if contains_template_placeholder(&expanded) { + bail!("step `{}` contains an unresolved SQL template", step.id); + } + crate::sql_guard::ensure_read_only_sql(&expanded) + .map_err(|error| anyhow!("step `{}`: {error}", step.id))?; + } + Ok(()) +} + +fn contains_template_placeholder(value: &str) -> bool { + let mut rest = value; + while let Some(open) = rest.find('{') { + rest = &rest[open + 1..]; + let Some(close) = rest.find('}') else { + return false; + }; + let name = &rest[..close]; + if !name.is_empty() + && name + .chars() + .all(|character| character.is_ascii_alphanumeric() || character == '_') + { + return true; + } + rest = &rest[close + 1..]; + } + false +} + +fn validate_skill_file(requested_id: &str, file: &SkillFile) -> Result<()> { + if file.api_version != "probing.dev/v1" { + bail!("unsupported apiVersion `{}`", file.api_version); + } + if file.kind != "Skill" { + bail!("expected kind `Skill`, got `{}`", file.kind); + } + if file.metadata.id != requested_id { + bail!( + "skill id `{}` does not match requested id `{requested_id}`", + file.metadata.id + ); + } + + Ok(()) +} + +pub(crate) fn validate_skill_contract(skill: &Skill) -> Result<()> { + let mut parameter_names = std::collections::HashSet::new(); + for parameter in &skill.parameters { + if !parameter_names.insert(parameter.name.as_str()) { + bail!("duplicate parameter id `{}`", parameter.name); + } + validate_parameter_default(parameter)?; + } + + let mut step_ids = std::collections::HashSet::new(); + for step in &skill.steps { + if !step_ids.insert(step.id.as_str()) { + bail!("duplicate step id `{}`", step.id); + } + if !matches!(step.step_type.as_str(), "sql" | "api" | "ui" | "config") { + bail!( + "step `{}` has unsupported type `{}`", + step.id, + step.step_type + ); + } + if !matches!(step.on_empty.as_str(), "skip" | "warn" | "abort") { + bail!( + "step `{}` has invalid on_empty `{}`", + step.id, + step.on_empty + ); + } + match step.step_type.as_str() { + "sql" if step.sql.as_ref().is_none_or(|sql| sql.trim().is_empty()) => { + bail!("SQL step `{}` is missing sql", step.id) + } + "api" if step.path.as_ref().is_none_or(|path| path.trim().is_empty()) => { + bail!("API step `{}` is missing path", step.id) + } + "api" + if step + .method + .as_deref() + .is_some_and(|method| !method.eq_ignore_ascii_case("GET")) => + { + bail!("API step `{}` only supports method GET", step.id) + } + "ui" if step.view.as_ref().is_none_or(|view| view.trim().is_empty()) => { + bail!("UI step `{}` is missing view", step.id) + } + _ => {} + } + } + + let mut rule_ids = std::collections::HashSet::new(); + for rule in &skill.interpretation { + if !rule_ids.insert(rule.id.as_str()) { + bail!("duplicate interpretation rule id `{}`", rule.id); + } + if !matches!(rule.severity.as_str(), "error" | "warning" | "info") { + bail!( + "rule `{}` has invalid severity `{}`", + rule.id, + rule.severity + ); + } + crate::interpret::validate_rule_expression(&rule.when) + .map_err(|error| anyhow!("rule `{}`: {error}", rule.id))?; + } + validate_default_sql(skill) +} + +fn validate_parameter_default(parameter: &SkillParameter) -> Result<()> { + let valid = match parameter.parameter_type { + SkillParameterType::Integer => parameter.default.as_i64().is_some(), + SkillParameterType::Number => parameter.default.as_f64().is_some(), + SkillParameterType::Boolean => parameter.default.as_bool().is_some(), + SkillParameterType::String => parameter.default.as_str().is_some(), + }; + if valid { + Ok(()) + } else { + bail!( + "parameter `{}` default does not match type `{}`", + parameter.name, + parameter.parameter_type.as_str() + ) + } +} + +/// Validate and normalize user-supplied parameter values before expansion. +pub fn normalize_parameter_overrides( + skill: &Skill, + overrides: &mut HashMap, +) -> Result<()> { + for (name, value) in overrides.iter_mut() { + let parameter = skill + .parameters + .iter() + .find(|parameter| parameter.name == *name) + .ok_or_else(|| anyhow!("unknown parameter `{name}` for skill `{}`", skill.id))?; + let normalized = match parameter.parameter_type { + SkillParameterType::Integer => value + .parse::() + .map(|value| value.to_string()) + .map_err(|_| anyhow!("parameter `{name}` must be an integer"))?, + SkillParameterType::Number => { + let number = value + .parse::() + .map_err(|_| anyhow!("parameter `{name}` must be a number"))?; + if !number.is_finite() { + bail!("parameter `{name}` must be a finite number"); + } + number.to_string() + } + SkillParameterType::Boolean => match value.to_ascii_lowercase().as_str() { + "true" => "true".to_string(), + "false" => "false".to_string(), + _ => bail!("parameter `{name}` must be true or false"), + }, + SkillParameterType::String => value.clone(), + }; + *value = normalized; + } + Ok(()) } fn collect_keywords(meta: &SkillMeta) -> Vec { @@ -304,6 +564,15 @@ pub fn build_context(pb: &Skill, overrides: &HashMap) -> HashMap ctx.insert(k.clone(), v.clone()); } ctx.extend(derive_variables(&ctx)); + // String parameters are authored as SQL literal contents (for example + // `stage = '{stage_filter}'`). Escape quotes before template expansion. + for parameter in &pb.parameters { + if parameter.parameter_type == SkillParameterType::String { + if let Some(value) = ctx.get_mut(¶meter.name) { + *value = value.replace('\'', "''"); + } + } + } for (key, template) in &pb.variables { let expanded = expand_template(template, &ctx); ctx.insert(key.clone(), expanded); @@ -327,6 +596,64 @@ mod tests { sql.split_whitespace().collect::>().join(" ") } + #[test] + fn every_catalog_skill_compiles() { + for id in list_skill_ids() { + load_skill(&id).unwrap_or_else(|error| panic!("skill {id}: {error:#}")); + } + } + + #[test] + fn parameter_overrides_are_typed_and_normalized() { + let skill = load_skill("slow_rank").expect("slow_rank skill"); + let mut valid = HashMap::from([ + ("step_window".to_string(), "0042".to_string()), + ("use_global".to_string(), "TRUE".to_string()), + ]); + normalize_parameter_overrides(&skill, &mut valid).expect("valid overrides"); + assert_eq!(valid["step_window"], "42"); + assert_eq!(valid["use_global"], "true"); + + let mut invalid = HashMap::from([("step_window".to_string(), "1; SELECT 1".to_string())]); + assert!(normalize_parameter_overrides(&skill, &mut invalid).is_err()); + let mut unknown = HashMap::from([("typo".to_string(), "1".to_string())]); + assert!(normalize_parameter_overrides(&skill, &mut unknown).is_err()); + + let string_skill = load_skill("module_bottleneck").expect("module_bottleneck skill"); + let string_overrides = + HashMap::from([("stage_filter".to_string(), "x' OR '1'='1".to_string())]); + let context = build_context(&string_skill, &string_overrides); + assert_eq!(context["stage_filter"], "x'' OR ''1''=''1"); + } + + #[test] + fn yaml_schema_rejects_unknown_fields() { + let yaml = r#" +apiVersion: probing.dev/v1 +kind: Skill +metadata: + id: strict + title: Strict +spec: + parameters: + - name: limit + type: integer + default: 1 + typo: ignored-before + steps: [] +"#; + let error = serde_yaml::from_str::(yaml).expect_err("unknown field must fail"); + assert!(error.to_string().contains("unknown field `typo`")); + } + + #[test] + fn unresolved_template_detection_ignores_non_template_braces() { + assert!(contains_template_placeholder("SELECT {missing_value}")); + assert!(!contains_template_placeholder( + "SELECT '{\"json\": true}' AS payload" + )); + } + #[test] fn slow_rank_rank_latency_sql_golden() { let skill = load_skill("slow_rank").expect("slow_rank skill"); diff --git a/probing/crates/skills/src/runner.rs b/probing/crates/skills/src/runner.rs index 046726af..798d55ec 100644 --- a/probing/crates/skills/src/runner.rs +++ b/probing/crates/skills/src/runner.rs @@ -8,8 +8,10 @@ use probing_proto::prelude::DataFrame; use crate::backend::{cluster_meta_note, SkillBackend}; use crate::interpret::{evaluate_rules, InterpretFinding, StepEvidence}; use crate::loader::{ - build_context, default_parameters, expand_template, load_skill, Skill, SkillStep, + build_context, default_parameters, expand_template, load_skill, normalize_parameter_overrides, + Skill, SkillStep, }; +use crate::sql_guard::ensure_read_only_sql; #[derive(Debug, Clone)] pub struct SkillRunError(pub String); @@ -82,6 +84,8 @@ pub struct RunResult { pub fn plan_skill(skill_id: &str, overrides: HashMap) -> Result { let pb = load_skill(skill_id).map_err(|e| SkillRunError(e.to_string()))?; + let mut overrides = overrides; + normalize_parameter_overrides(&pb, &mut overrides).map_err(SkillRunError::from)?; let mut params = default_parameters(&pb); params.extend(overrides); let ctx = build_context(&pb, ¶ms); @@ -115,6 +119,9 @@ fn step_plan_json(step: &SkillStep, ctx: &HashMap) -> serde_json if let Some(when) = &step.when { item["when"] = serde_json::Value::String(when.clone()); } + if let Some(platform) = step.platform { + item["platform"] = serde_json::Value::String(platform.as_str().to_string()); + } item } @@ -123,6 +130,7 @@ pub async fn resolve_use_global( pb: &Skill, overrides: &mut HashMap, ) -> Result<()> { + normalize_parameter_overrides(pb, overrides).map_err(SkillRunError::from)?; if overrides.contains_key("use_global") { return Ok(()); } @@ -221,6 +229,19 @@ pub async fn run_step( ctx: &HashMap, options: &RunOptions, ) -> StepOutcome { + if let Some(platform) = step.platform { + if !platform.is_current() { + return StepOutcome::Skipped { + step_id: step.id.clone(), + title: step.title.clone(), + reason: format!( + "step requires platform {}, current platform is {}", + platform.as_str(), + std::env::consts::OS + ), + }; + } + } if let Some(reason) = should_skip_step(step, ctx) { return StepOutcome::Skipped { step_id: step.id.clone(), @@ -288,26 +309,12 @@ fn sql_needs_cluster(sql: &str, step_cluster: bool) -> bool { step_cluster || sql.to_lowercase().contains("global.") } -fn ensure_read_only_sql(sql: &str) -> Result<()> { - let upper = sql.trim().to_uppercase(); - if upper.starts_with("SELECT") - || upper.starts_with("WITH") - || upper.starts_with("SHOW") - || upper.starts_with("DESCRIBE") - { - return Ok(()); - } - Err(SkillRunError( - "Only read-only SQL is allowed in skills".to_string(), - )) -} - async fn run_sql_step(backend: &B, step: &SkillStep, sql: &str) -> StepOutcome { if let Err(e) = ensure_read_only_sql(sql) { return StepOutcome::Error { step_id: step.id.clone(), title: step.title.clone(), - message: e.0, + message: e, }; } let cluster = sql_needs_cluster(sql, step.cluster.unwrap_or(false)); @@ -672,12 +679,15 @@ mod tests { title: format!("title-{id}"), step_type: "sql".into(), sql: Some(sql.into()), + method: None, path: None, view: None, on_empty: "skip".into(), empty_message: None, when: None, cluster: None, + platform: None, + action: None, } } @@ -692,13 +702,16 @@ mod tests { trigger_keywords: Default::default(), parameters: vec![SkillParameter { name: "use_global".into(), + parameter_type: crate::loader::SkillParameterType::Boolean, default: serde_yaml::Value::Bool(true), + description: String::new(), }], steps, interpretation, summary_template: "rows={available_tables.row_count}".into(), next_steps: vec![], variables: HashMap::new(), + requires: Default::default(), } } @@ -760,6 +773,29 @@ mod tests { assert!(matches!(outcome, StepOutcome::Error { .. })); } + #[tokio::test] + async fn run_step_rejects_trailing_mutating_sql() { + let backend = MockBackend::new(0, 1); + let step = sample_step("bad", "SELECT 1; DELETE FROM python.t"); + let outcome = run_step(&backend, &step, &HashMap::new(), &RunOptions::default()).await; + assert!(matches!(outcome, StepOutcome::Error { .. })); + assert_eq!(backend.calls.load(Ordering::Relaxed), 0); + } + + #[tokio::test] + async fn run_step_skips_other_platform() { + let backend = MockBackend::new(0, 1); + let mut step = sample_step("platform", "SELECT 1"); + step.platform = Some(if cfg!(target_os = "linux") { + crate::loader::SkillPlatform::Macos + } else { + crate::loader::SkillPlatform::Linux + }); + let outcome = run_step(&backend, &step, &HashMap::new(), &RunOptions::default()).await; + assert!(matches!(outcome, StepOutcome::Skipped { .. })); + assert_eq!(backend.calls.load(Ordering::Relaxed), 0); + } + #[tokio::test] async fn resolve_use_global_honors_override() { let backend = MockBackend::new(4, 1); diff --git a/probing/crates/skills/src/sql_guard.rs b/probing/crates/skills/src/sql_guard.rs new file mode 100644 index 00000000..f488577b --- /dev/null +++ b/probing/crates/skills/src/sql_guard.rs @@ -0,0 +1,88 @@ +//! AST-based read-only SQL validation shared by skill loading and execution. + +use sqlparser::ast::{Query, SetExpr, Statement}; +use sqlparser::dialect::GenericDialect; +use sqlparser::parser::Parser; + +pub(crate) fn ensure_read_only_sql(sql: &str) -> Result<(), String> { + let trimmed = sql.trim(); + if trimmed.is_empty() { + return Err("SQL must not be empty".to_string()); + } + + let statements = Parser::parse_sql(&GenericDialect {}, trimmed) + .map_err(|error| format!("invalid SQL: {error}"))?; + if statements.is_empty() { + return Err("SQL must not be empty".to_string()); + } + if statements.iter().all(statement_is_read_only) { + Ok(()) + } else { + Err("Only read-only SQL is allowed (SELECT/WITH/SHOW/DESCRIBE/EXPLAIN)".to_string()) + } +} + +fn statement_is_read_only(statement: &Statement) -> bool { + match statement { + Statement::Query(query) => query_is_read_only(query), + Statement::Explain { .. } | Statement::ExplainTable { .. } => true, + Statement::ShowFunctions { .. } + | Statement::ShowVariable { .. } + | Statement::ShowStatus { .. } + | Statement::ShowVariables { .. } + | Statement::ShowCreate { .. } + | Statement::ShowColumns { .. } + | Statement::ShowCatalogs { .. } + | Statement::ShowDatabases { .. } + | Statement::ShowProcessList { .. } + | Statement::ShowSchemas { .. } + | Statement::ShowCharset(_) + | Statement::ShowObjects(_) + | Statement::ShowTables { .. } + | Statement::ShowViews { .. } + | Statement::ShowCollation { .. } => true, + _ => false, + } +} + +fn query_is_read_only(query: &Query) -> bool { + query.with.as_ref().is_none_or(|with| { + with.cte_tables + .iter() + .all(|cte| query_is_read_only(&cte.query)) + }) && set_expr_is_read_only(query.body.as_ref()) +} + +fn set_expr_is_read_only(expression: &SetExpr) -> bool { + match expression { + SetExpr::Select(_) | SetExpr::Values(_) | SetExpr::Table(_) => true, + SetExpr::Query(query) => query_is_read_only(query), + SetExpr::SetOperation { left, right, .. } => { + set_expr_is_read_only(left) && set_expr_is_read_only(right) + } + SetExpr::Insert(_) | SetExpr::Update(_) | SetExpr::Delete(_) | SetExpr::Merge(_) => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn accepts_read_only_statements() { + assert!(ensure_read_only_sql("SELECT 1").is_ok()); + assert!(ensure_read_only_sql("WITH x AS (SELECT 1) SELECT * FROM x").is_ok()); + assert!(ensure_read_only_sql("SHOW TABLES; DESCRIBE python.t").is_ok()); + } + + #[test] + fn rejects_write_hidden_after_read() { + assert!(ensure_read_only_sql("SELECT 1; DELETE FROM python.t").is_err()); + assert!(ensure_read_only_sql("SELECT 1; SET probing.x=1").is_err()); + } + + #[test] + fn rejects_write_cte() { + assert!(ensure_read_only_sql("WITH x AS (DELETE FROM python.t) SELECT 1").is_err()); + } +} diff --git a/probing/memtable/src/memc/compactor.rs b/probing/memtable/src/memc/compactor.rs index 8dc12e42..0abd2c11 100644 --- a/probing/memtable/src/memc/compactor.rs +++ b/probing/memtable/src/memc/compactor.rs @@ -101,6 +101,7 @@ struct SharedRuntimeStats { impl SharedRuntimeStats { fn record(&self, message: String) { + log::warn!("MEMC compactor: {message}"); self.error_count.fetch_add(1, Ordering::Relaxed); match self.last_error.lock() { Ok(mut last) => *last = Some(message), @@ -168,10 +169,13 @@ impl Compactor { /// `(source_instance, source_gen, source_chunk)` it came from; we keep /// the max generation per instance/chunk. pub fn prime_from_cold(&mut self) -> io::Result<()> { - for path in self.store.segment_paths() { - let Ok(reader) = SegmentReader::open(&path) else { - continue; // unreadable/foreign file: skip, never fail priming - }; + for path in self.store.segment_paths_checked()? { + let reader = SegmentReader::open(&path).map_err(|error| { + io::Error::new( + error.kind(), + format!("failed to prime from {}: {error}", path.display()), + ) + })?; for page in reader.pages() { if page.source_chunk == SOURCE_CHUNK_NONE { continue; @@ -336,7 +340,11 @@ impl Compactor { if w.page_count() == 0 { let path = w.path().to_path_buf(); drop(w); - let _ = std::fs::remove_file(&path); + match std::fs::remove_file(&path) { + Ok(()) => {} + Err(error) if error.kind() == io::ErrorKind::NotFound => {} + Err(error) => return Err(error), + } return Ok(None); } Ok(Some(w.seal()?)) @@ -353,7 +361,8 @@ impl Compactor { .enforce_limits(self.config.max_total_bytes, self.config.ttl) } - fn enforce_checked(&self) -> io::Result> { + /// Fallible retention pass for workers that must surface filesystem errors. + pub fn enforce_checked(&self) -> io::Result> { self.store .enforce_limits_checked(self.config.max_total_bytes, self.config.ttl) } @@ -395,7 +404,7 @@ impl Compactor { /// shared/file-backed [`MemTable`] the application is writing elsewhere. /// Dropping (or [`stop`](CompactorHandle::stop)ping) the returned handle /// does a final drain + flush so no sealed chunk is left behind. - pub fn spawn(mut self, sources: Vec<(String, MemTable)>) -> CompactorHandle { + pub fn spawn(mut self, sources: Vec<(String, MemTable)>) -> io::Result { let stop = Arc::new(AtomicBool::new(false)); let stop_thread = stop.clone(); let runtime_stats = Arc::new(SharedRuntimeStats::default()); @@ -431,13 +440,12 @@ impl Compactor { if let Err(error) = self.enforce_checked() { thread_stats.record(format!("final retention enforcement: {error}")); } - }) - .expect("spawn memc-compactor thread"); - CompactorHandle { + })?; + Ok(CompactorHandle { stop, thread: Some(thread), runtime_stats, - } + }) } } diff --git a/probing/memtable/src/memc/reader.rs b/probing/memtable/src/memc/reader.rs index eeb2b665..3ef4bc11 100644 --- a/probing/memtable/src/memc/reader.rs +++ b/probing/memtable/src/memc/reader.rs @@ -1,9 +1,9 @@ //! [`SegmentReader`]: mmap a `.memc` file and read its tables and pages. //! -//! A sealed segment is read via its footer page directory. An unsealed or -//! torn segment (writer crashed before `seal`) falls back to a forward -//! scan of checksummed blocks, stopping at the first damaged/partial block -//! — so a half-written tail is silently dropped rather than surfaced. +//! A sealed segment is read via its footer page directory and any integrity +//! failure is surfaced. An unsealed segment (writer crashed before `seal`) +//! falls back to a forward scan of checksummed blocks; only an incomplete tail +//! is ignored, while corruption of a complete block remains an error. use std::collections::HashMap; use std::io; @@ -56,13 +56,23 @@ impl SegmentReader { let mut tables = HashMap::new(); let mut pages = Vec::new(); - let footer_ok = header.is_sealed() - && header.footer_off != 0 - && Self::load_footer(&mmap, &header, &mut pages); + let footer_ok = if header.is_sealed() { + if header.footer_off == 0 { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "sealed MEMC segment has no footer", + )); + } + Self::load_footer(&mmap, &header, &mut pages).map_err(invalid_data)?; + true + } else { + false + }; // Always scan blocks for table definitions (cheap; MCTB blocks live - // before pages). On footer failure this also recovers page metadata. - Self::scan_blocks(&mmap, &header, &mut tables, footer_ok, &mut pages); + // before pages). Unsealed segments also recover page metadata here. + Self::scan_blocks(&mmap, &header, &mut tables, footer_ok, &mut pages) + .map_err(invalid_data)?; Ok(Self { mmap, @@ -73,56 +83,52 @@ impl SegmentReader { }) } - /// Parse the footer page directory. Returns `false` (and leaves `pages` - /// untouched) if the footer is malformed or fails its checksum. - fn load_footer(mmap: &[u8], header: &SegmentHeader, pages: &mut Vec) -> bool { - let Ok(foff) = usize::try_from(header.footer_off) else { - return false; - }; - let Some(footer_header_end) = foff.checked_add(16) else { - return false; - }; + /// Parse and verify the footer page directory. + fn load_footer( + mmap: &[u8], + header: &SegmentHeader, + pages: &mut Vec, + ) -> Result<(), &'static str> { + let foff = usize::try_from(header.footer_off).map_err(|_| "footer offset overflow")?; + let footer_header_end = foff.checked_add(16).ok_or("footer header overflow")?; if footer_header_end > mmap.len() || get_u32(mmap, foff) != MAGIC_FOOTER { - return false; + return Err("MEMC footer missing or invalid"); } let count = get_u32(mmap, foff + 4) as usize; let entries_len = get_u32(mmap, foff + 8) as usize; let checksum = get_u32(mmap, foff + 12); - let Some(expected_entries_len) = count.checked_mul(PAGE_DIR_ENTRY_SIZE) else { - return false; - }; + let expected_entries_len = count + .checked_mul(PAGE_DIR_ENTRY_SIZE) + .ok_or("footer page count overflow")?; if count != header.page_count as usize || entries_len != expected_entries_len { - return false; + return Err("MEMC footer page count mismatch"); } let entries_start = footer_header_end; - let Some(entries_end) = entries_start.checked_add(entries_len) else { - return false; - }; + let entries_end = entries_start + .checked_add(entries_len) + .ok_or("footer entries length overflow")?; if entries_end > mmap.len() || xxh32(&mmap[entries_start..entries_end]) != checksum { - return false; + return Err("MEMC footer checksum mismatch"); } let mut out = Vec::with_capacity(count); for i in 0..count { - let Some(entry_off) = i + let entry_off = i .checked_mul(PAGE_DIR_ENTRY_SIZE) .and_then(|offset| entries_start.checked_add(offset)) - else { - return false; - }; + .ok_or("footer entry offset overflow")?; let block_off = super::layout::get_u64(mmap, entry_off + 24); let block_len = get_u32(mmap, entry_off + 32); - let Ok(block_start) = usize::try_from(block_off) else { - return false; - }; - let Some(block_end) = block_start.checked_add(block_len as usize) else { - return false; - }; + let block_start = + usize::try_from(block_off).map_err(|_| "footer block offset overflow")?; + let block_end = block_start + .checked_add(block_len as usize) + .ok_or("footer block length overflow")?; if block_start < SEGMENT_HEADER_SIZE || block_len < BLOCK_HEADER_SIZE as u32 || block_end > foff { - return false; + return Err("MEMC footer page block range invalid"); } out.push(PageMeta { table_id: get_u32(mmap, entry_off), @@ -138,23 +144,24 @@ impl SegmentReader { }); } *pages = out; - true + Ok(()) } /// Forward-scan blocks from the first block to `footer_off`/EOF. /// Collects table definitions always; collects page metadata only when - /// `footer_ok` is false (recovery path). Stops at the first block that - /// fails to decode or whose payload checksum mismatches. + /// `footer_ok` is false (recovery path). Complete corrupt blocks are + /// rejected; only an incomplete tail of an unsealed segment is ignored. fn scan_blocks( mmap: &[u8], header: &SegmentHeader, tables: &mut HashMap, footer_ok: bool, pages: &mut Vec, - ) { + ) -> Result<(), &'static str> { + let sealed = header.is_sealed(); let limit = if header.footer_off != 0 { usize::try_from(header.footer_off) - .unwrap_or(usize::MAX) + .map_err(|_| "footer offset overflow")? .min(mmap.len()) } else { mmap.len() @@ -164,42 +171,46 @@ impl SegmentReader { .checked_add(BLOCK_HEADER_SIZE) .is_some_and(|end| end <= limit) { - let Some(bh) = BlockHeader::decode(&mmap[off..]) else { - break; - }; - let Some(payload_start) = off.checked_add(BLOCK_HEADER_SIZE) else { - break; - }; - let Some(payload_end) = payload_start.checked_add(bh.payload_len as usize) else { - break; - }; + let bh = BlockHeader::decode(&mmap[off..]).ok_or("MEMC block header invalid")?; + let payload_start = off + .checked_add(BLOCK_HEADER_SIZE) + .ok_or("MEMC block payload offset overflow")?; + let payload_end = payload_start + .checked_add(bh.payload_len as usize) + .ok_or("MEMC block payload length overflow")?; if payload_end > limit { - break; // torn tail + if sealed { + return Err("sealed MEMC block payload out of bounds"); + } + break; // crash during an unsealed tail write } if xxh32(&mmap[payload_start..payload_end]) != bh.payload_xxh { - break; // corrupt payload — stop here + return Err("MEMC block payload checksum mismatch"); } - let Some(raw_block_len) = BLOCK_HEADER_SIZE.checked_add(bh.payload_len as usize) else { - break; - }; - let Some(block_len) = raw_block_len.checked_add(63).map(|n| n & !63) else { - break; - }; - let Some(next_off) = off.checked_add(block_len) else { - break; - }; + let raw_block_len = BLOCK_HEADER_SIZE + .checked_add(bh.payload_len as usize) + .ok_or("MEMC block length overflow")?; + let block_len = raw_block_len + .checked_add(63) + .map(|n| n & !63) + .ok_or("MEMC aligned block length overflow")?; + let next_off = off + .checked_add(block_len) + .ok_or("MEMC next block offset overflow")?; if next_off > limit { + if sealed { + return Err("sealed MEMC aligned block out of bounds"); + } break; } match bh.magic { MAGIC_TABLE_BLOCK => { - if let Ok(def) = super::layout::decode_table_payload( + let def = super::layout::decode_table_payload( bh.table_id, &mmap[payload_start..payload_end], - ) { - tables.insert(bh.table_id, def); - } + )?; + tables.insert(bh.table_id, def); } MAGIC_PAGE_BLOCK if !footer_ok => { pages.push(PageMeta { @@ -219,6 +230,10 @@ impl SegmentReader { } off = next_off; } + if sealed && off != limit { + return Err("sealed MEMC segment has trailing or truncated block bytes"); + } + Ok(()) } pub fn path(&self) -> &Path { @@ -333,3 +348,7 @@ impl SegmentReader { Ok(cols) } } + +fn invalid_data(message: &'static str) -> io::Error { + io::Error::new(io::ErrorKind::InvalidData, message) +} diff --git a/probing/memtable/src/memc/store.rs b/probing/memtable/src/memc/store.rs index b3714097..96eff97a 100644 --- a/probing/memtable/src/memc/store.rs +++ b/probing/memtable/src/memc/store.rs @@ -72,7 +72,7 @@ impl ColdStore { std::fs::create_dir_all(&dir)?; let pid = std::process::id(); let wid = writer_id(pid, process_start_time(pid)); - let next_seq = Self::max_seq_for(&dir, &wid) + 1; + let next_seq = Self::max_seq_for(&dir, &wid)?.saturating_add(1); Ok(Self { dir, writer_id: wid, @@ -89,19 +89,18 @@ impl ColdStore { } /// Highest existing sequence number for `wid` in `dir` (0 if none). - fn max_seq_for(dir: &Path, wid: &str) -> u32 { + fn max_seq_for(dir: &Path, wid: &str) -> io::Result { let mut max = 0u32; - if let Ok(entries) = std::fs::read_dir(dir) { - for e in entries.flatten() { - let name = e.file_name().to_string_lossy().to_string(); - if let Some((w, seq)) = parse_segment_name(&name) { - if w == wid { - max = max.max(seq); - } + for entry in std::fs::read_dir(dir)? { + let entry = entry?; + let name = entry.file_name().to_string_lossy().to_string(); + if let Some((w, seq)) = parse_segment_name(&name) { + if w == wid { + max = max.max(seq); } } } - max + Ok(max) } /// Path for the next segment (does not create the file). @@ -133,45 +132,67 @@ impl ColdStore { /// All segment files in the directory (any writer), sorted oldest → /// newest by modification time. pub fn segment_paths(&self) -> Vec { + match self.segment_paths_checked() { + Ok(paths) => paths, + Err(error) => { + log::warn!( + "MEMC segment enumeration failed for {}: {error}", + self.dir.display() + ); + Vec::new() + } + } + } + + /// Fallible segment enumeration for correctness-sensitive workers. + pub fn segment_paths_checked(&self) -> io::Result> { let mut segs: Vec<(SystemTime, PathBuf)> = Vec::new(); - if let Ok(entries) = std::fs::read_dir(&self.dir) { - for e in entries.flatten() { - let path = e.path(); - if path.extension().and_then(|s| s.to_str()) != Some(SEGMENT_EXT) { - continue; - } - let mtime = e - .metadata() - .and_then(|m| m.modified()) - .unwrap_or(SystemTime::UNIX_EPOCH); - segs.push((mtime, path)); + for entry in std::fs::read_dir(&self.dir)? { + let entry = entry?; + let path = entry.path(); + if path.extension().and_then(|s| s.to_str()) != Some(SEGMENT_EXT) { + continue; } + let mtime = entry.metadata()?.modified()?; + segs.push((mtime, path)); } segs.sort_by_key(|a| a.0); - segs.into_iter().map(|(_, p)| p).collect() + Ok(segs.into_iter().map(|(_, p)| p).collect()) } pub fn stats(&self) -> ColdStats { - let paths = self.segment_paths(); + match self.stats_checked() { + Ok(stats) => stats, + Err(error) => { + log::warn!( + "MEMC stats collection failed for {}: {error}", + self.dir.display() + ); + ColdStats::default() + } + } + } + + /// Fallible capacity snapshot used by retention accounting. + pub fn stats_checked(&self) -> io::Result { + let paths = self.segment_paths_checked()?; let mut total = 0u64; let mut oldest = u64::MAX; for p in &paths { - if let Ok(meta) = std::fs::metadata(p) { - total += meta.len(); - if let Ok(mtime) = meta.modified() { - let ms = mtime - .duration_since(SystemTime::UNIX_EPOCH) - .map(|d| d.as_millis() as u64) - .unwrap_or(0); - oldest = oldest.min(ms); - } - } + let meta = std::fs::metadata(p)?; + total = total.saturating_add(meta.len()); + let ms = meta + .modified()? + .duration_since(SystemTime::UNIX_EPOCH) + .map(|d| d.as_millis() as u64) + .unwrap_or(0); + oldest = oldest.min(ms); } - ColdStats { + Ok(ColdStats { segment_count: paths.len(), total_bytes: total, oldest_unix_ms: if paths.is_empty() { 0 } else { oldest }, - } + }) } /// Evict oldest segments until under `max_bytes` and within `ttl`. @@ -197,7 +218,7 @@ impl ColdStore { max_bytes: Option, ttl: Option, ) -> io::Result> { - let paths = self.segment_paths(); + let paths = self.segment_paths_checked()?; if paths.len() <= 1 { return Ok(Vec::new()); } @@ -208,9 +229,18 @@ impl ColdStore { // A concurrent writer may be between create/header write or // actively appending. Conservatively retain both cases. Err(_) if legacy_segment_is_sealed(path) => {} - Ok(_) | Err(_) => { + Ok(_) => { protected.insert(path.clone()); } + Err(error) => { + return Err(io::Error::new( + error.kind(), + format!( + "cannot validate MEMC segment {} for retention: {error}", + path.display() + ), + )); + } } } // Preserve the existing newest-segment retention guarantee even @@ -220,28 +250,32 @@ impl ColdStore { } let now = SystemTime::now(); - let mut total: u64 = self.stats().total_bytes; + let mut total = 0u64; + for path in &paths { + total = total.saturating_add(std::fs::metadata(path)?.len()); + } let mut removed = Vec::new(); for path in paths { if protected.contains(&path) { continue; } - let too_old = ttl - .and_then(|ttl| { - let mtime = std::fs::metadata(&path).ok()?.modified().ok()?; - now.duration_since(mtime).ok().map(|age| age > ttl) - }) - .unwrap_or(false); + let metadata = match std::fs::metadata(&path) { + Ok(metadata) => metadata, + Err(error) if error.kind() == io::ErrorKind::NotFound => continue, + Err(error) => return Err(error), + }; + let too_old = match ttl { + Some(ttl) => now + .duration_since(metadata.modified()?) + .is_ok_and(|age| age > ttl), + None => false, + }; let over_budget = max_bytes.is_some_and(|max| total > max); if !(too_old || over_budget) { break; // sorted oldest-first: nothing newer qualifies either } - let sz = match std::fs::metadata(&path) { - Ok(metadata) => metadata.len(), - Err(error) if error.kind() == io::ErrorKind::NotFound => continue, - Err(error) => return Err(error), - }; + let sz = metadata.len(); match std::fs::remove_file(&path) { Ok(()) => { total = total.saturating_sub(sz); diff --git a/probing/memtable/src/memc/tests.rs b/probing/memtable/src/memc/tests.rs index d341a860..630eb63e 100644 --- a/probing/memtable/src/memc/tests.rs +++ b/probing/memtable/src/memc/tests.rs @@ -288,7 +288,7 @@ fn decoded_column_count_must_match_page_row_count() { } #[test] -fn malformed_footer_block_range_falls_back_to_checked_scan() { +fn malformed_footer_block_range_is_rejected() { let dir = tmp_dir("footer-range"); let path = dir.join("seg.memc"); let mut w = SegmentWriter::create(&path).unwrap(); @@ -308,11 +308,12 @@ fn malformed_footer_block_range_falls_back_to_checked_scan() { super::layout::put_u32(&mut bytes, footer_off + 12, checksum); std::fs::write(&path, bytes).unwrap(); - let reader = SegmentReader::open(&path).unwrap(); - assert_eq!(reader.pages().len(), 1); - assert_eq!( - reader.read_page(0).unwrap()[0], - ColumnData::I64(vec![1, 2, 3]) + let error = SegmentReader::open(&path) + .err() + .expect("footer must be rejected"); + assert!( + error.to_string().contains("footer"), + "unexpected error: {error}" ); let _ = std::fs::remove_dir_all(&dir); @@ -702,7 +703,8 @@ fn compactor_background_thread_drains_on_stop() { ..Default::default() }, ) - .spawn(vec![("metrics".to_string(), reader)]); + .spawn(vec![("metrics".to_string(), reader)]) + .expect("spawn compactor"); for i in 0..4 { writer.push_row(&[Value::I64(i), Value::F64(i as f64)]); @@ -739,7 +741,8 @@ fn compactor_background_errors_are_observable() { ..Default::default() }, ) - .spawn(vec![("metrics".to_string(), source)]); + .spawn(vec![("metrics".to_string(), source)]) + .expect("spawn compactor"); std::thread::sleep(Duration::from_millis(30)); let stats = handle.stop(); @@ -748,13 +751,48 @@ fn compactor_background_errors_are_observable() { stats .last_error .as_deref() - .is_some_and(|message| message.contains("drain table metrics")), + .is_some_and(|message| message.contains(':')), "last error should retain operation context: {stats:?}" ); std::fs::remove_file(&dir).unwrap(); } +#[test] +fn retention_directory_errors_are_not_reported_as_success() { + let dir = tmp_dir("retention-directory-error"); + let store = ColdStore::open(&dir).unwrap(); + std::fs::remove_dir_all(&dir).unwrap(); + std::fs::File::create(&dir).unwrap(); + + assert!(store.segment_paths_checked().is_err()); + assert!(store + .enforce_limits_checked(Some(1), Some(Duration::ZERO)) + .is_err()); + + std::fs::remove_file(&dir).unwrap(); +} + +#[test] +fn retention_and_priming_report_corrupt_segments() { + let dir = tmp_dir("retention-corrupt-segment"); + let store = ColdStore::open(&dir).unwrap(); + std::fs::write(dir.join("broken-000001.memc"), b"not a MEMC segment").unwrap(); + std::fs::write(dir.join("broken-000002.memc"), b"not a MEMC segment").unwrap(); + + let prime_error = Compactor::new(ColdStore::open(&dir).unwrap(), CompactorConfig::default()) + .prime_from_cold() + .unwrap_err(); + assert!(prime_error.to_string().contains("failed to prime")); + + let retention_error = store + .enforce_limits_checked(Some(1), Some(Duration::ZERO)) + .unwrap_err(); + assert!(retention_error.to_string().contains("for retention")); + + std::fs::remove_dir_all(&dir).unwrap(); +} + #[test] fn compactor_enforce_evicts_oldest_segments() { let dir = tmp_dir("compact-evict"); diff --git a/probing/server/Cargo.toml b/probing/server/Cargo.toml index c2d01e7e..62ed6adf 100644 --- a/probing/server/Cargo.toml +++ b/probing/server/Cargo.toml @@ -39,6 +39,7 @@ tokio = { workspace = true, features = ["fs"] } async-trait = "0.1.83" bytes = "1" +include_dir = "=0.7.4" nu-ansi-term = "0.50.1" base64 = "0.21.5" ureq = { workspace = true, features = ["json"] } diff --git a/probing/server/build.rs b/probing/server/build.rs new file mode 100644 index 00000000..88def696 --- /dev/null +++ b/probing/server/build.rs @@ -0,0 +1,44 @@ +use std::fs; +use std::io; +use std::path::{Path, PathBuf}; + +fn copy_dir(source: &Path, destination: &Path) -> io::Result<()> { + fs::create_dir_all(destination)?; + for entry in fs::read_dir(source)? { + let entry = entry?; + let source_path = entry.path(); + let destination_path = destination.join(entry.file_name()); + if entry.file_type()?.is_dir() { + copy_dir(&source_path, &destination_path)?; + } else { + fs::copy(source_path, destination_path)?; + } + } + Ok(()) +} + +fn main() -> io::Result<()> { + let manifest_dir = PathBuf::from(std::env::var_os("CARGO_MANIFEST_DIR").ok_or_else(|| { + io::Error::new(io::ErrorKind::NotFound, "CARGO_MANIFEST_DIR is unavailable") + })?); + let generated = manifest_dir.join("web-assets/public"); + let fallback = manifest_dir.join("web-fallback"); + let source = if generated.join("embedded.manifest").is_file() { + generated.as_path() + } else { + fallback.as_path() + }; + + println!("cargo:rerun-if-changed={}", generated.display()); + println!("cargo:rerun-if-changed={}", fallback.display()); + + let out_dir = PathBuf::from( + std::env::var_os("OUT_DIR") + .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "OUT_DIR is unavailable"))?, + ); + let destination = out_dir.join("probing-web-assets"); + if destination.exists() { + fs::remove_dir_all(&destination)?; + } + copy_dir(source, &destination) +} diff --git a/probing/server/src/asset.rs b/probing/server/src/asset.rs index 5e95b93b..73f09765 100644 --- a/probing/server/src/asset.rs +++ b/probing/server/src/asset.rs @@ -5,8 +5,11 @@ use axum::body::Body; use axum::http::{header, HeaderMap, StatusCode, Uri}; use axum::response::{IntoResponse, Response}; use bytes::Bytes; +use include_dir::{include_dir, Dir}; use once_cell::sync::Lazy; +static EMBEDDED_ASSETS: Dir<'_> = include_dir!("$OUT_DIR/probing-web-assets"); + static BASE_PATH: Lazy = Lazy::new(|| { env::var("PROBING_BASE_PATH") .unwrap_or_default() @@ -43,12 +46,18 @@ fn read_from_disk(key: &str) -> Option { Some(Bytes::from(content)) } +fn read_embedded(key: &str) -> Option { + EMBEDDED_ASSETS + .get_file(key) + .map(|file| Bytes::copy_from_slice(file.contents())) +} + pub fn contains(path: &str) -> bool { let key = normalize_asset_path(path); if let Some(root) = assets_root() { return Path::new(&root).join(&key).exists(); } - key == "index.html" + EMBEDDED_ASSETS.get_file(&key).is_some() } pub fn get(path: &str) -> Bytes { @@ -56,6 +65,9 @@ pub fn get(path: &str) -> Bytes { if let Some(data) = read_from_disk(&key) { return data; } + if let Some(data) = read_embedded(&key) { + return data; + } if key == "index.html" { return Bytes::from_static(MISSING_UI_HTML.as_bytes()); } @@ -101,6 +113,11 @@ fn resolve_asset(path: &str, accept_encoding: &str) -> (Bytes, Option<&'static s return (data, Some("br")); } } + if let Some(data) = read_embedded(&br_key) { + if !data.is_empty() { + return (data, Some("br")); + } + } } (get(path), None) @@ -230,7 +247,7 @@ pub async fn static_files(uri: Uri, headers: HeaderMap) -> Result = pub static AUTH_REALM: Lazy = Lazy::new(|| env::var(AUTH_REALM_ENV).unwrap_or_else(|_| "Probe Server".to_string())); +static PEER_AUTH_TOKEN: Lazy>> = Lazy::new(|| RwLock::new(None)); + +fn configure_peer_auth_token(token: &str) { + let token = token.trim(); + *probing_core::sync::write_rwlock(&PEER_AUTH_TOKEN, "PEER_AUTH_TOKEN") = + (!token.is_empty()).then(|| token.to_string()); +} + +pub(crate) fn peer_auth_header_value() -> Option { + probing_core::sync::read_rwlock(&PEER_AUTH_TOKEN, "PEER_AUTH_TOKEN") + .as_ref() + .map(|token| format!("Bearer {token}")) +} + /// Load `PROBING_AUTH_TOKEN` into the config store so middleware can enforce it. pub async fn bootstrap_auth_from_env() { let Ok(token) = env::var(AUTH_TOKEN_ENV) else { @@ -33,14 +47,17 @@ pub async fn bootstrap_auth_from_env() { if token.is_empty() { return; } - if let Err(err) = config::write(AUTH_TOKEN_CONFIG_KEY, token).await { - log::error!("failed to bootstrap auth token from {AUTH_TOKEN_ENV}: {err}"); + match config::write(AUTH_TOKEN_CONFIG_KEY, token).await { + Ok(()) => configure_peer_auth_token(token), + Err(err) => log::error!("failed to bootstrap auth token from {AUTH_TOKEN_ENV}: {err}"), } } /// Persist auth token to the config store (used by SET and extension options). pub async fn persist_auth_token(token: &str) -> Result<(), probing_core::core::EngineError> { - config::write(AUTH_TOKEN_CONFIG_KEY, token).await + config::write(AUTH_TOKEN_CONFIG_KEY, token).await?; + configure_peer_auth_token(token); + Ok(()) } /// Get the auth token from the request diff --git a/probing/server/src/cluster_http.rs b/probing/server/src/cluster_http.rs index 40864b74..e0184a56 100644 --- a/probing/server/src/cluster_http.rs +++ b/probing/server/src/cluster_http.rs @@ -3,6 +3,15 @@ use std::time::Duration; use anyhow::Result; use probing_proto::prelude::{Node, NodeListResponse, NodeReportRequest, NodeReportResponse}; +use crate::auth::peer_auth_header_value; + +fn authenticate_peer_request(request: ureq::RequestBuilder) -> ureq::RequestBuilder { + match peer_auth_header_value() { + Some(value) => request.header("Authorization", value), + None => request, + } +} + pub fn get_i32_env(name: &str) -> Option { std::env::var(name) .ok() @@ -25,11 +34,12 @@ pub fn fetch_nodes_blocking(http_base: &str) -> Result> { let mut all = Vec::new(); loop { let url = format!("{base}/apis/nodes?offset={offset}&limit={page_size}"); - let text = ureq::get(&url) + let request = ureq::get(&url) .config() .no_delay(true) .timeout_global(Some(Duration::from_secs(10))) - .build() + .build(); + let text = authenticate_peer_request(request) .call()? .body_mut() .read_to_string()?; @@ -54,7 +64,7 @@ pub fn put_nodes_blocking( nodes, seen_version, }; - let text = ureq::put(&url) + let request = ureq::put(&url) .config() .no_delay(true) .timeout_global(Some(Duration::from_secs( @@ -63,7 +73,8 @@ pub fn put_nodes_blocking( .and_then(|v| v.parse().ok()) .unwrap_or(5), ))) - .build() + .build(); + let text = authenticate_peer_request(request) .send_json(body)? .body_mut() .read_to_string()?; diff --git a/probing/server/src/engine.rs b/probing/server/src/engine.rs index 033951b6..c1d1685a 100644 --- a/probing/server/src/engine.rs +++ b/probing/server/src/engine.rs @@ -17,7 +17,8 @@ use probing_core::config; use crate::server::error::{ApiError, ApiResult}; use probing_core::core::federation::{ - reset_fanout_stats, take_fanout_stats, with_fanout_scope_async, FanoutScope, + fanout_stats_partial, reset_fanout_stats, take_fanout_stats, with_fanout_scope_async, + FanoutScope, }; use probing_core::core::UnifiedMemtableProbeDataSource; pub use probing_core::ENGINE; @@ -28,8 +29,8 @@ use probing_python::extensions::python::PythonProbeDataSource; pub async fn initialize_engine() -> Result<()> { probing_hccl_shim::register_docs(); probing_nccl_profiler::register_docs(); - let builder = probing_core::create_engine() + .with_peer_query_transport(crate::server::cluster_fanout::core_transport()) .with_data_source(cc::ClusterProbeDataSource::create("cluster", "nodes")) .with_data_source(cc::EnvProbeDataSource::create("process", "envs")) .with_data_source(cc::FilesProbeDataSource::create("files")) @@ -171,7 +172,7 @@ fn is_missing_table_error(err: &impl std::fmt::Display) -> bool { fn fanout_meta_from_stats( stats: probing_core::core::federation::FanoutStats, ) -> Option { - if stats.nodes_failed.is_empty() && stats.peer_batches_dropped == 0 { + if !fanout_stats_partial(&stats) { return None; } Some(serde_json::json!({ @@ -185,7 +186,7 @@ fn fanout_meta_from_stats( } fn query_response_partial(stats: &probing_core::core::federation::FanoutStats) -> bool { - !stats.nodes_failed.is_empty() || stats.peer_batches_dropped > 0 + fanout_stats_partial(stats) } /// Serialized `/query` body plus whether federated fan-out was partial. @@ -251,3 +252,22 @@ async fn query_in_fanout_context(req: String) -> ApiResult { error, }) } + +#[cfg(test)] +mod tests { + use super::{fanout_meta_from_stats, query_response_partial}; + use probing_core::core::federation::FanoutStats; + + #[test] + fn child_partial_flag_surfaces_without_failure_details() { + let stats = FanoutStats { + nodes_succeeded: 2, + partial: true, + ..FanoutStats::default() + }; + + assert!(query_response_partial(&stats)); + let meta = fanout_meta_from_stats(stats).expect("partial response metadata"); + assert_eq!(meta["fanout"]["partial"], true); + } +} diff --git a/probing/server/src/mcp/helpers.rs b/probing/server/src/mcp/helpers.rs index 48ef896e..b0f3e9ec 100644 --- a/probing/server/src/mcp/helpers.rs +++ b/probing/server/src/mcp/helpers.rs @@ -118,7 +118,7 @@ pub async fn engine_query_json(sql: String, limit: usize) -> Result Result<(), ErrorData> { let stats = take_fanout_stats(); - if stats.nodes_failed.is_empty() && stats.peer_batches_dropped == 0 { + if !probing_core::core::federation::fanout_stats_partial(&stats) { return Ok(()); } if fanout_strict_enabled() { diff --git a/probing/server/src/server/cluster_fanout.rs b/probing/server/src/server/cluster_fanout.rs index 7f061f70..3f85dd9e 100644 --- a/probing/server/src/server/cluster_fanout.rs +++ b/probing/server/src/server/cluster_fanout.rs @@ -24,6 +24,7 @@ use planner::{peers_for_scope, plan_fanout}; use transport::HttpPeerQueryClient; use types::{finish_fanout, FanoutOutcome}; +pub(crate) use transport::core_transport; pub use transport::remote_query_df; pub use types::{ClusterFanoutScope, FanoutMeta, FanoutQueryResponse}; @@ -187,7 +188,7 @@ async fn fanout_via_global_catalog(sql: &str, scope: FanoutScope) -> anyhow::Res } else { 0 }, - partial: false, + partial: stats.partial, }, "global-catalog", ) diff --git a/probing/server/src/server/cluster_fanout/transport.rs b/probing/server/src/server/cluster_fanout/transport.rs index 30ba7aac..6d43c9c9 100644 --- a/probing/server/src/server/cluster_fanout/transport.rs +++ b/probing/server/src/server/cluster_fanout/transport.rs @@ -1,8 +1,13 @@ use async_trait::async_trait; -use probing_core::core::federation::{fanout_strict_enabled, remote_query_timeout}; +use datafusion::error::{DataFusionError, Result as DataFusionResult}; +use probing_core::core::federation::{ + fanout_strict_enabled, remote_query_timeout, FanoutScope, FanoutStats, PeerQueryOutcome, + PeerQueryTransport, +}; use probing_proto::prelude::*; use super::types::FanoutOutcome; +use crate::auth::peer_auth_header_value; #[async_trait] pub(super) trait PeerQueryClient: Send + Sync { @@ -13,6 +18,59 @@ pub(super) trait PeerQueryClient: Send + Sync { #[derive(Debug, Default, Clone, Copy)] pub(super) struct HttpPeerQueryClient; +#[derive(Debug, Default)] +struct HttpFederationPeerTransport; + +#[derive(Debug)] +struct PeerTransportError(anyhow::Error); + +impl std::fmt::Display for PeerTransportError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(formatter, "{:#}", self.0) + } +} + +impl std::error::Error for PeerTransportError {} + +fn transport_error(error: anyhow::Error) -> DataFusionError { + DataFusionError::External(Box::new(PeerTransportError(error))) +} + +impl PeerQueryTransport for HttpFederationPeerTransport { + fn query( + &self, + addr: &str, + sql: &str, + scope: FanoutScope, + ) -> DataFusionResult { + if scope == FanoutScope::Coordinator { + return remote_node_aggregate_blocking(addr, sql) + .map(peer_query_outcome) + .map_err(transport_error); + } + remote_query_df_blocking(addr, sql) + .map(PeerQueryOutcome::complete) + .map_err(transport_error) + } +} + +fn peer_query_outcome(outcome: FanoutOutcome) -> PeerQueryOutcome { + let meta = outcome.meta; + PeerQueryOutcome::with_stats( + outcome.dataframe, + FanoutStats { + nodes_succeeded: meta.nodes_queried.saturating_sub(meta.nodes_failed.len()), + nodes_failed: meta.nodes_failed, + peer_batches_dropped: meta.peer_batches_dropped, + partial: meta.partial, + }, + ) +} + +pub(crate) fn core_transport() -> std::sync::Arc { + std::sync::Arc::new(HttpFederationPeerTransport) +} + #[async_trait] impl PeerQueryClient for HttpPeerQueryClient { async fn query_leaf(&self, addr: &str, sql: &str) -> anyhow::Result { @@ -25,17 +83,29 @@ impl PeerQueryClient for HttpPeerQueryClient { } pub async fn remote_query_df(addr: &str, sql: &str) -> anyhow::Result { + let addr = addr.to_string(); + let sql = sql.to_string(); + tokio::task::spawn_blocking(move || remote_query_df_blocking(&addr, &sql)).await? +} + +fn remote_query_df_blocking(addr: &str, sql: &str) -> anyhow::Result { let url = format!("http://{addr}/query"); let request = Message::new(Query { expr: sql.to_string(), ..Default::default() }); let body = serde_json::to_string(&request)?; - let response = send_post(url, None, body).await?; + let response = send_post_blocking(url, None, body)?; parse_leaf_response(response.0, &response.1, addr) } async fn remote_node_aggregate(addr: &str, sql: &str) -> anyhow::Result { + let addr = addr.to_string(); + let sql = sql.to_string(); + tokio::task::spawn_blocking(move || remote_node_aggregate_blocking(&addr, &sql)).await? +} + +fn remote_node_aggregate_blocking(addr: &str, sql: &str) -> anyhow::Result { let url = format!("http://{addr}/apis/cluster/query"); let body = serde_json::to_string(&serde_json::json!({ "expr": sql, @@ -43,30 +113,30 @@ async fn remote_node_aggregate(addr: &str, sql: &str) -> anyhow::Result, body: String, ) -> anyhow::Result<(u16, String)> { let timeout = remote_query_timeout(); - tokio::task::spawn_blocking(move || { - let mut request = ureq::post(&url) - .config() - .timeout_global(Some(timeout)) - .build(); - if let Some(content_type) = content_type { - request = request.header("Content-Type", content_type); - } - let response = request.send(body).map_err(anyhow::Error::new)?; - let status = response.status().as_u16(); - let text = response.into_body().read_to_string()?; - Ok((status, text)) - }) - .await? + let mut request = ureq::post(&url) + .config() + .timeout_global(Some(timeout)) + .build(); + if let Some(content_type) = content_type { + request = request.header("Content-Type", content_type); + } + if let Some(value) = peer_auth_header_value() { + request = request.header("Authorization", value); + } + let response = request.send(body).map_err(anyhow::Error::new)?; + let status = response.status().as_u16(); + let text = response.into_body().read_to_string()?; + Ok((status, text)) } fn parse_leaf_response(status: u16, text: &str, addr: &str) -> anyhow::Result { @@ -116,3 +186,32 @@ fn decode_cluster_query_response(text: &str) -> anyhow::Result { } Ok(serde_json::from_value(value)?) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::server::cluster_fanout::FanoutMeta; + + #[test] + fn core_outcome_preserves_child_partial_metadata() { + let outcome = peer_query_outcome(FanoutOutcome { + dataframe: DataFrame::default(), + meta: FanoutMeta { + cluster: true, + hierarchical: true, + scope: "node".into(), + nodes_queried: 3, + nodes_failed: vec!["rank-3: timeout".into()], + peer_batches_dropped: 1, + node_aggregators_queried: 0, + local_ranks_queried: 2, + partial: true, + }, + }); + + assert_eq!(outcome.stats.nodes_succeeded, 2); + assert_eq!(outcome.stats.nodes_failed, vec!["rank-3: timeout"]); + assert_eq!(outcome.stats.peer_batches_dropped, 1); + assert!(outcome.stats.partial); + } +} diff --git a/probing/server/src/server/cluster_fanout/types.rs b/probing/server/src/server/cluster_fanout/types.rs index c9139e21..165bb787 100644 --- a/probing/server/src/server/cluster_fanout/types.rs +++ b/probing/server/src/server/cluster_fanout/types.rs @@ -87,6 +87,7 @@ impl FanoutMeta { nodes_succeeded: self.nodes_queried.saturating_sub(self.nodes_failed.len()), nodes_failed: self.nodes_failed.clone(), peer_batches_dropped: self.peer_batches_dropped, + partial: self.partial, } } } diff --git a/probing/server/src/server/settings.rs b/probing/server/src/server/settings.rs index 28653c80..17855fb0 100644 --- a/probing/server/src/server/settings.rs +++ b/probing/server/src/server/settings.rs @@ -34,7 +34,7 @@ pub fn sync_env_settings() { }) .collect(); - super::SERVER_RUNTIME.spawn(async move { + if let Err(error) = super::SERVER_RUNTIME.spawn(async move { for (key, value) in env_vars { let key = key.replace('_', ".").to_lowercase(); match config::write(&key, &value).await { @@ -42,7 +42,9 @@ pub fn sync_env_settings() { Err(error) => error!("Failed to sync env setting '{key}': {error}"), }; } - }); + }) { + log::error!("failed to schedule server settings update: {error}"); + } } #[cfg(test)] diff --git a/probing/server/src/supervisor.rs b/probing/server/src/supervisor.rs index 8abb4930..a009ffcd 100644 --- a/probing/server/src/supervisor.rs +++ b/probing/server/src/supervisor.rs @@ -217,13 +217,27 @@ impl ServerSupervisor { let (start_tx, start_rx) = oneshot::channel(); let future = factory(generation); let supervisor = self; - let handle = crate::server::SERVER_RUNTIME.spawn(async move { + let handle = match crate::server::SERVER_RUNTIME.spawn(async move { if start_rx.await.is_err() { return; } let result = future.await; supervisor.remote_listener_finished(generation, result.err()); - }); + }) { + Ok(handle) => handle, + Err(error) => { + transition( + "remote_listener", + &mut slot.state, + ComponentState::Failed { + generation, + key, + error: error.to_string(), + }, + ); + return; + } + }; slot.candidate = Some(ManagedTask { generation, key, @@ -347,14 +361,28 @@ impl ServerSupervisor { let (start_tx, start_rx) = oneshot::channel(); let future = factory(generation); let supervisor = self; - let handle = crate::server::SERVER_RUNTIME.spawn(async move { + let handle = match crate::server::SERVER_RUNTIME.spawn(async move { if start_rx.await.is_err() { return; } supervisor.mark_component_running(component, generation); let result = future.await; supervisor.component_finished(component, generation, result.err()); - }); + }) { + Ok(handle) => handle, + Err(error) => { + transition( + component.name(), + &mut slot.state, + ComponentState::Failed { + generation, + key, + error: error.to_string(), + }, + ); + return; + } + }; slot.active = Some(ManagedTask { generation, key, diff --git a/probing/server/src/torchrun_cluster.rs b/probing/server/src/torchrun_cluster.rs index 3e53ac44..a208e429 100644 --- a/probing/server/src/torchrun_cluster.rs +++ b/probing/server/src/torchrun_cluster.rs @@ -554,13 +554,15 @@ pub fn refresh_torchrun_role() -> bool { return false; } lock_mutex(&REPORT_BACKOFF, "torchrun REPORT_BACKOFF").reset(); - SERVER_RUNTIME.spawn(async { + if let Err(error) = SERVER_RUNTIME.spawn(async { let outcome = tokio::task::spawn_blocking(report_once) .await .unwrap_or(ReportOutcome::Failed); let mut backoff = lock_mutex(&REPORT_BACKOFF, "torchrun REPORT_BACKOFF"); backoff.record(outcome); - }); + }) { + log::error!("failed to schedule torchrun cluster reporter: {error}"); + } true } diff --git a/probing/server/tests/hierarchical_cluster_report.rs b/probing/server/tests/hierarchical_cluster_report.rs index 30a06953..3a78e0f8 100644 --- a/probing/server/tests/hierarchical_cluster_report.rs +++ b/probing/server/tests/hierarchical_cluster_report.rs @@ -4,12 +4,18 @@ use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, LazyLock, Mutex}; use std::time::{SystemTime, UNIX_EPOCH}; -use axum::{extract::State, routing::get, Json, Router}; +use axum::{extract::State, middleware, routing::get, routing::post, Json, Router}; +use probing_core::core::federation::{FanoutScope, ProbeClusterExecutor}; use probing_core::sync::lock_mutex; use probing_proto::prelude::{ - Cluster, Node, NodeListResponse, NodeReportRequest, NodeReportResponse, + Cluster, DataFrame, Message, Node, NodeListResponse, NodeReportRequest, NodeReportResponse, + QueryDataFormat, +}; +use probing_server::auth::{ + bootstrap_auth_from_env, persist_auth_token, selective_auth_middleware, AUTH_TOKEN_ENV, }; use probing_server::cluster_http::{fetch_nodes_blocking, put_nodes_blocking}; +use probing_server::server::cluster_fanout::remote_query_df; use probing_server::server::SERVER_RUNTIME; use tokio::net::TcpListener; @@ -22,18 +28,19 @@ struct AppState { static ENV_LOCK: LazyLock> = LazyLock::new(|| Mutex::new(())); fn local_http_available() -> bool { - match SERVER_RUNTIME.block_on(TcpListener::bind("127.0.0.1:0")) { - Ok(listener) => { + match SERVER_RUNTIME.try_block_on(TcpListener::bind("127.0.0.1:0")) { + Ok(Ok(listener)) => { drop(listener); true } - Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => { + Ok(Err(error)) if error.kind() == std::io::ErrorKind::PermissionDenied => { eprintln!( "skipping hierarchical cluster report test: environment denied TCP bind ({error})" ); false } - Err(error) => panic!("probe local HTTP bind capability: {error}"), + Ok(Err(error)) => panic!("probe local HTTP bind capability: {error}"), + Err(error) => panic!("probe runtime unavailable: {error}"), } } @@ -119,6 +126,48 @@ async fn spawn_cluster_server() -> String { format!("http://{addr}") } +async fn query_handler() -> Json> { + Json(Message::new(QueryDataFormat::DataFrame( + DataFrame::default(), + ))) +} + +async fn cluster_query_handler() -> Json { + Json(serde_json::json!({ + "dataframe": DataFrame::default(), + "meta": { + "cluster": true, + "hierarchical": true, + "scope": "node", + "nodes_queried": 1, + "nodes_failed": [], + "peer_batches_dropped": 0, + "node_aggregators_queried": 0, + "local_ranks_queried": 1, + "partial": false + } + })) +} + +async fn spawn_authenticated_cluster_server() -> String { + let state = AppState { + cluster: Arc::new(Mutex::new(Cluster::default())), + version: Arc::new(AtomicU64::new(0)), + }; + let app = Router::new() + .route("/apis/nodes", get(get_nodes_handler).put(put_nodes_handler)) + .route("/query", post(query_handler)) + .route("/apis/cluster/query", post(cluster_query_handler)) + .with_state(state) + .layer(middleware::from_fn(selective_auth_middleware)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + format!("http://{addr}") +} + /// Mirrors production ``local_leaf_nodes`` for aggregator simulation. fn aggregator_payload(store: &[Node], self_node: Node) -> Vec { let group_rank = self_node.group_rank.unwrap_or(0); @@ -161,44 +210,110 @@ fn hierarchical_two_nodes_times_two_gpus_converges_on_master() { std::env::remove_var(key); } - SERVER_RUNTIME.block_on(async { - let master = spawn_cluster_server().await; - let node1_local0 = spawn_cluster_server().await; - - put_nodes_blocking(&master, vec![test_node(1, 0, "127.0.0.1:9101")], 0) - .expect("leaf rank1 put"); - - put_nodes_blocking(&node1_local0, vec![test_node(3, 1, "127.0.0.1:9103")], 0) - .expect("leaf rank3 put"); - - let node0_store = fetch_nodes_blocking(&master).expect("read node0 local store"); - let rank0 = test_node(0, 0, "127.0.0.1:9100"); - let node0_batch = aggregator_payload(&node0_store, rank0); - assert_eq!( - node0_batch - .iter() - .filter_map(|n| n.rank) - .collect::>(), - vec![0, 1] - ); - put_nodes_blocking(&master, node0_batch, 1).expect("rank0 aggregator put"); - - let node1_store = fetch_nodes_blocking(&node1_local0).expect("read node1 local store"); - let rank2 = test_node(2, 1, "127.0.0.1:9102"); - let node1_batch = aggregator_payload(&node1_store, rank2); - assert_eq!( - node1_batch - .iter() - .filter_map(|n| n.rank) - .collect::>(), - vec![2, 3] - ); - put_nodes_blocking(&master, node1_batch, 2).expect("rank2 aggregator put"); - - let snapshot = fetch_nodes_blocking(&master).expect("master snapshot"); - let ranks: Vec = snapshot.iter().filter_map(|n| n.rank).collect(); - assert_eq!(ranks, vec![0, 1, 2, 3]); - - assert_eq!(local_group_ranks(&snapshot, 0), vec![0, 1]); - }); + SERVER_RUNTIME + .try_block_on(async { + let master = spawn_cluster_server().await; + let node1_local0 = spawn_cluster_server().await; + + put_nodes_blocking(&master, vec![test_node(1, 0, "127.0.0.1:9101")], 0) + .expect("leaf rank1 put"); + + put_nodes_blocking(&node1_local0, vec![test_node(3, 1, "127.0.0.1:9103")], 0) + .expect("leaf rank3 put"); + + let node0_store = fetch_nodes_blocking(&master).expect("read node0 local store"); + let rank0 = test_node(0, 0, "127.0.0.1:9100"); + let node0_batch = aggregator_payload(&node0_store, rank0); + assert_eq!( + node0_batch + .iter() + .filter_map(|n| n.rank) + .collect::>(), + vec![0, 1] + ); + put_nodes_blocking(&master, node0_batch, 1).expect("rank0 aggregator put"); + + let node1_store = fetch_nodes_blocking(&node1_local0).expect("read node1 local store"); + let rank2 = test_node(2, 1, "127.0.0.1:9102"); + let node1_batch = aggregator_payload(&node1_store, rank2); + assert_eq!( + node1_batch + .iter() + .filter_map(|n| n.rank) + .collect::>(), + vec![2, 3] + ); + put_nodes_blocking(&master, node1_batch, 2).expect("rank2 aggregator put"); + + let snapshot = fetch_nodes_blocking(&master).expect("master snapshot"); + let ranks: Vec = snapshot.iter().filter_map(|n| n.rank).collect(); + assert_eq!(ranks, vec![0, 1, 2, 3]); + + assert_eq!(local_group_ranks(&snapshot, 0), vec![0, 1]); + }) + .expect("probing runtime"); +} + +#[test] +fn authenticated_peer_traffic_supports_heartbeat_and_fanout() { + let _guard = lock_mutex(&ENV_LOCK, "hierarchical_cluster_report ENV_LOCK"); + if !local_http_available() { + return; + } + + SERVER_RUNTIME + .try_block_on(async { + probing_server::initialize_engine() + .await + .expect("initialize composition root"); + std::env::set_var(AUTH_TOKEN_ENV, "cluster-secret"); + bootstrap_auth_from_env().await; + std::env::remove_var(AUTH_TOKEN_ENV); + let base = spawn_authenticated_cluster_server().await; + let addr = base.trim_start_matches("http://"); + let unauthenticated_url = format!("{base}/apis/nodes"); + let unauthenticated = tokio::task::spawn_blocking(move || { + ureq::get(&unauthenticated_url) + .call() + .expect_err("token is required") + }) + .await + .expect("join unauthenticated request"); + assert!(matches!(unauthenticated, ureq::Error::StatusCode(401))); + + put_nodes_blocking(&base, vec![test_node(0, 0, "127.0.0.1:9100")], 0) + .expect("authenticated heartbeat"); + assert_eq!( + fetch_nodes_blocking(&base) + .expect("authenticated node discovery") + .len(), + 1 + ); + + remote_query_df(addr, "SELECT 1") + .await + .expect("authenticated server leaf fan-out"); + let transport = probing_core::ENGINE + .read() + .await + .peer_query_transport() + .expect("composition root transport"); + ProbeClusterExecutor::execute_remote_for_scope( + Some(&transport), + addr, + "SELECT 1", + FanoutScope::Flat, + ) + .expect("authenticated core leaf fan-out"); + ProbeClusterExecutor::execute_remote_for_scope( + Some(&transport), + addr, + "SELECT 1", + FanoutScope::Coordinator, + ) + .expect("authenticated hierarchical fan-out"); + + persist_auth_token("").await.expect("clear auth token"); + }) + .expect("probing runtime"); } diff --git a/probing/server/web-fallback/index.html b/probing/server/web-fallback/index.html new file mode 100644 index 00000000..81ee84b7 --- /dev/null +++ b/probing/server/web-fallback/index.html @@ -0,0 +1,10 @@ + + + + + Probing Web Interface + + +

Web UI not available. Run make frontend, then rebuild probing.

+ + diff --git a/pyproject.toml b/pyproject.toml index f668374a..26e9b137 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -120,7 +120,6 @@ include = [ "probing/**/*.py", "probing/libs/*", "probing/bundled_skills/**", - "probing/bundled_web/**", ] [build-system] diff --git a/python/probing/__init__.py b/python/probing/__init__.py index 5f30c4e6..543c7018 100644 --- a/python/probing/__init__.py +++ b/python/probing/__init__.py @@ -30,9 +30,6 @@ cli_main = _core.cli_main __all__ = ["VERSION", "cli_main"] else: - from probing.web_assets import configure_assets_root - - configure_assets_root() import probing.config as config from probing import _core from probing.external_table import ExternalTable diff --git a/python/probing/skills/loader.py b/python/probing/skills/loader.py index 60a9461a..182edd1b 100644 --- a/python/probing/skills/loader.py +++ b/python/probing/skills/loader.py @@ -178,12 +178,27 @@ def validate_skill(skill: Skill) -> List[str]: if not skill.summary_template.strip() and not skill.next_steps: warnings.append(f"{skill.id}: no steps defined") seen_ids: set[str] = set() + seen_parameters: set[str] = set() + for parameter in skill.parameters: + name = str(parameter.get("name", "")) + parameter_type = parameter.get("type") + if name in seen_parameters: + warnings.append(f"{skill.id}: duplicate parameter id {name}") + seen_parameters.add(name) + if parameter_type not in {"integer", "number", "boolean", "string"}: + warnings.append( + f"{skill.id}.{name}: unsupported parameter type {parameter_type!r}" + ) for step in skill.steps: if step.id in seen_ids: warnings.append(f"{skill.id}: duplicate step id {step.id}") seen_ids.add(step.id) if step.type == "sql" and not step.sql: warnings.append(f"{skill.id}.{step.id}: sql step missing sql") + if step.platform not in {None, "linux", "macos", "windows"}: + warnings.append( + f"{skill.id}.{step.id}: unsupported platform {step.platform!r}" + ) skill_md = skill.path.parent / "SKILL.md" if not skill_md.is_file(): warnings.append(f"{skill.id}: missing SKILL.md") @@ -194,6 +209,10 @@ def validate_all() -> List[str]: catalog = load_catalog() all_warnings: List[str] = [] for entry in catalog.skills: - skill = load_skill(entry.id) + try: + skill = load_skill(entry.id) + except Exception as error: + all_warnings.append(f"{entry.id}: invalid skill contract: {error}") + continue all_warnings.extend(validate_skill(skill)) return all_warnings diff --git a/python/probing/web_assets.py b/python/probing/web_assets.py deleted file mode 100644 index 359296bb..00000000 --- a/python/probing/web_assets.py +++ /dev/null @@ -1,107 +0,0 @@ -"""Web UI static assets bundled in the wheel or available in editable installs.""" - -from __future__ import annotations - -import os -from pathlib import Path - -_ENV = "PROBING_ASSETS_ROOT" -_BUNDLED_DIRNAME = "bundled_web" - - -def _package_dir() -> Path: - return Path(__file__).resolve().parent - - -def _bundled_web_candidates() -> list[Path]: - """``dx bundle`` emits ``bundled_web/public/``; legacy wheels may be flat.""" - root = _package_dir() / _BUNDLED_DIRNAME - return [root / "public", root] - - -def bundled_web_dir() -> Path | None: - """Wheel / install tree: ``python/probing/bundled_web[/public]``.""" - for candidate in _bundled_web_candidates(): - if (candidate / "index.html").is_file(): - return candidate - for rel in ("bundled_web/public", "bundled_web"): - root = _resource_dir(rel, "index.html") - if root is not None: - return root - return None - - -def dev_web_dir() -> Path | None: - """Editable checkout: ``web/dist`` symlink → ``bundled_web/public``.""" - if _running_from_installed_wheel(): - return None - root = _repo_root_from_editable() / "web" / "dist" - if (root / "index.html").is_file(): - return root - return bundled_web_dir() - - -def _running_from_installed_wheel() -> bool: - maybe_repo = Path(__file__).resolve().parents[2] - return not (maybe_repo / "pyproject.toml").is_file() - - -def _repo_root_from_editable() -> Path: - return Path(__file__).resolve().parents[2] - - -def _resource_dir(rel: str, marker: str) -> Path | None: - try: - from importlib.resources import as_file, files - - bundle = files("probing") - for part in rel.split("/"): - bundle = bundle / part - if not (bundle / marker).is_file(): - return None - with as_file(bundle) as path: - return Path(path) - except (TypeError, ModuleNotFoundError, FileNotFoundError, OSError): - return None - - -def _looks_like_built_ui(root: Path) -> bool: - """True when ``index.html`` is a Dioxus bundle, not the checkout placeholder.""" - index = root / "index.html" - if not index.is_file(): - return False - try: - body = index.read_text(encoding="utf-8", errors="ignore") - except OSError: - return False - return "web-dxh" in body or '
' in body - - -def resolve_web_assets_root() -> Path | None: - """Return the directory that contains ``index.html``, if any.""" - override = os.environ.get(_ENV) - if override: - root = Path(override) - if (root / "index.html").is_file(): - return root - return None - - for getter in (dev_web_dir, bundled_web_dir): - root = getter() - if root and _looks_like_built_ui(root): - return root - - return dev_web_dir() - - -def configure_assets_root() -> Path | None: - """Set ``PROBING_ASSETS_ROOT`` for ``probing-server`` when UI files are available.""" - if os.environ.get(_ENV): - root = Path(os.environ[_ENV]) - if (root / "index.html").is_file(): - return root - return None - root = resolve_web_assets_root() - if root is not None: - os.environ[_ENV] = str(root) - return root diff --git a/scripts/prune-bundled-web.sh b/scripts/prune-web-assets.sh similarity index 93% rename from scripts/prune-bundled-web.sh rename to scripts/prune-web-assets.sh index a562bae0..5aa62de7 100755 --- a/scripts/prune-bundled-web.sh +++ b/scripts/prune-web-assets.sh @@ -2,12 +2,12 @@ # Drop stale dx hashed assets after `dx bundle` — keep only what index.html loads. set -euo pipefail -PUBLIC="${1:-python/probing/bundled_web/public}" +PUBLIC="${1:-probing/server/web-assets/public}" INDEX="$PUBLIC/index.html" ASSETS="$PUBLIC/assets" if [[ ! -f "$INDEX" ]]; then - echo "error: missing bundled web index: $INDEX" >&2 + echo "error: missing embedded web index: $INDEX" >&2 exit 1 fi diff --git a/scripts/verify_web_assets.py b/scripts/verify_web_assets.py index e8feeafa..3cbcac6c 100644 --- a/scripts/verify_web_assets.py +++ b/scripts/verify_web_assets.py @@ -82,7 +82,7 @@ def main(argv: list[str] | None = None) -> int: "root", nargs="?", type=Path, - default=Path("python/probing/bundled_web/public"), + default=Path("probing/server/web-assets/public"), help="bundle root containing index.html", ) args = parser.parse_args(argv) diff --git a/scripts/verify_wheel_contents.py b/scripts/verify_wheel_contents.py index 0e41a6e3..4e2e0d4e 100644 --- a/scripts/verify_wheel_contents.py +++ b/scripts/verify_wheel_contents.py @@ -8,11 +8,6 @@ import zipfile from pathlib import Path -try: - from verify_web_assets import verify_web_files -except ModuleNotFoundError: # Imported as `scripts.verify_wheel_contents` in tests. - from scripts.verify_web_assets import verify_web_files - # Paths that must exist in every release wheel (wheel archive member names). REQUIRED_PATHS = ( "probing/__init__.py", @@ -21,9 +16,22 @@ "probing/handlers/router.py", "probing/profiling/torch_probe.py", "probing/bundled_skills/catalog.yaml", - "probing/bundled_web/public/index.html", ) +EMBEDDED_WEB_MARKER = b"__PROBING_EMBEDDED_WEB_ASSETS_V1__" + + +def _native_extension(names: set[str]) -> str | None: + return next( + ( + name + for name in names + if name.startswith("probing/_core") + and name.endswith((".so", ".pyd", ".dylib")) + ), + None, + ) + def _pick_wheel(path: Path | None) -> Path: if path is not None: @@ -43,24 +51,18 @@ def verify_wheel(wheel: Path) -> list[str]: names = set(zf.namelist()) for member in REQUIRED_PATHS: if member not in names: - # Legacy dx layout (no public/ subdir). - if member == "probing/bundled_web/public/index.html": - if "probing/bundled_web/index.html" in names: - continue missing.append(member) - web_root = ( - "probing/bundled_web/public/" - if "probing/bundled_web/public/index.html" in names - else "probing/bundled_web/" - ) - index_member = f"{web_root}index.html" - if index_member in names: - errors = verify_web_files( - zf.read(index_member).decode("utf-8"), - lambda path: f"{web_root}{path}" in names, - lambda path: zf.read(f"{web_root}{path}").decode("utf-8"), - ) - missing.extend(f"invalid web bundle: {error}" for error in errors) + legacy_web = sorted(name for name in names if name.startswith("probing/bundled_web/")) + if legacy_web: + missing.append("legacy probing/bundled_web files must not be shipped") + + native = _native_extension(names) + if native is None: + missing.append("probing/_core native extension") + else: + binary = zf.read(native) + if EMBEDDED_WEB_MARKER not in binary: + missing.append("native extension does not contain embedded Web assets") return missing @@ -80,7 +82,7 @@ def main(argv: list[str] | None = None) -> int: for path in missing: print(f" - {path}", file=sys.stderr) print( - "hint: run 'make frontend && make wheel' before install-wheel", + "hint: run 'make frontend && make wheel' so Web assets are embedded in probing._core", file=sys.stderr, ) return 1 diff --git a/tests/regression/rust/probing/core/extension_routing_spec.rs b/tests/regression/rust/probing/core/extension_routing_spec.rs index f1626d77..74f4dcc3 100644 --- a/tests/regression/rust/probing/core/extension_routing_spec.rs +++ b/tests/regression/rust/probing/core/extension_routing_spec.rs @@ -34,7 +34,7 @@ fn load_spec() -> serde_json::Value { } async fn register_pythonext_stub() -> ProbeExtensionManager { - let mut manager = ProbeExtensionManager; + let mut manager = ProbeExtensionManager::default(); manager .register( "pythonext".to_string(), diff --git a/tests/regression/rust/probing/core/federation_explain_tests.rs b/tests/regression/rust/probing/core/federation_explain_tests.rs index 291dc24f..0c5b9284 100644 --- a/tests/regression/rust/probing/core/federation_explain_tests.rs +++ b/tests/regression/rust/probing/core/federation_explain_tests.rs @@ -19,12 +19,16 @@ async fn metrics_engine(values: Vec) -> Engine { host: "explain-coord".into(), addr: "127.0.0.1:19999".into(), rank: Some(0), + group_rank: Some(0), + local_rank: Some(0), ..Default::default() }); update_node(Node { host: "explain-peer".into(), addr: "127.0.0.1:20001".into(), rank: Some(1), + group_rank: Some(1), + local_rank: Some(0), ..Default::default() }); diff --git a/tests/regression/rust/probing/core/federation_tests.rs b/tests/regression/rust/probing/core/federation_tests.rs index 9c45ec07..82997d28 100644 --- a/tests/regression/rust/probing/core/federation_tests.rs +++ b/tests/regression/rust/probing/core/federation_tests.rs @@ -1,13 +1,16 @@ //! Regression tests for the `global` federated catalog path: //! probe catalog (local) vs global catalog (fan-out + `_addr` / `_rank` tagging). +use std::fmt; use std::sync::Arc; -use probing_core::core::cluster::{reset_cluster_for_tests, update_node}; +use probing_core::core::cluster::{ + reset_cluster_for_tests, update_node, HIERARCHICAL_METADATA_UNAVAILABLE, +}; use probing_core::core::federation::{ - set_remote_query_hook, take_fanout_stats, FEDERATION_TAG_COLUMNS, GLOBAL_CATALOG, - PROBE_ADDR_COL, PROBE_HOST_COL, PROBE_LOCAL_RANK_COL, PROBE_NODE_RANK_COL, PROBE_RANK_COL, - PROBE_ROLE_COL, + set_remote_query_hook, take_fanout_stats, FanoutScope, FanoutStats, PeerQueryOutcome, + PeerQueryTransport, FEDERATION_TAG_COLUMNS, GLOBAL_CATALOG, PROBE_ADDR_COL, PROBE_HOST_COL, + PROBE_LOCAL_RANK_COL, PROBE_NODE_RANK_COL, PROBE_RANK_COL, PROBE_ROLE_COL, }; use probing_core::core::{Engine, ProbeDataSource}; use probing_proto::prelude::{Node, Seq}; @@ -84,6 +87,45 @@ struct FederatedTestCluster { peer_addr: String, } +struct PartialSubtreeTransport { + peer_engine: Engine, + peer_addr: String, +} + +impl fmt::Debug for PartialSubtreeTransport { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("PartialSubtreeTransport") + .field("peer_addr", &self.peer_addr) + .finish() + } +} + +impl PeerQueryTransport for PartialSubtreeTransport { + fn query( + &self, + addr: &str, + sql: &str, + scope: FanoutScope, + ) -> datafusion::error::Result { + assert_eq!(addr, self.peer_addr); + assert_eq!(scope, FanoutScope::Coordinator); + let dataframe = futures::executor::block_on(self.peer_engine.async_query(sql))? + .ok_or_else(|| { + datafusion::error::DataFusionError::Execution("missing peer data".into()) + })?; + Ok(PeerQueryOutcome::with_stats( + dataframe, + FanoutStats { + nodes_succeeded: 1, + nodes_failed: vec!["remote-leaf: timeout".into()], + peer_batches_dropped: 0, + partial: true, + }, + )) + } +} + impl FederatedTestCluster { async fn setup(local_values: Vec, peer_values: Vec) -> Self { reset_cluster_for_tests(); @@ -149,6 +191,72 @@ impl FederatedTestCluster { } } +#[tokio::test] +async fn global_query_preserves_nested_partial_transport_stats() { + let _lock = federation_test_lock().await; + reset_cluster_for_tests(); + set_remote_query_hook(None); + std::env::set_var("PROBING_ADDRESS", "127.0.0.1:19999"); + std::env::set_var("HOSTNAME", "coord-host"); + std::env::set_var("RANK", "0"); + std::env::set_var("LOCAL_RANK", "0"); + std::env::set_var("GROUP_RANK", "0"); + std::env::set_var("PROBING_CLUSTER_FANOUT_HIERARCHICAL", "1"); + + register_local_node(0, "127.0.0.1:19999", "coord-host"); + let peer_addr = "127.0.0.1:20002".to_string(); + update_node(Node { + host: "remote-node".into(), + addr: peer_addr.clone(), + rank: Some(2), + group_rank: Some(1), + local_rank: Some(0), + status: Some("running".into()), + ..Default::default() + }); + + let peer_table = + GenericTableProbeDataSource::single_column_table("metrics", "demo", "v", vec![2]); + let peer_engine = Engine::builder() + .with_data_source(Arc::new(peer_table) as Arc) + .build() + .await + .expect("peer engine"); + let local_table = + GenericTableProbeDataSource::single_column_table("metrics", "demo", "v", vec![0]); + let engine = Engine::builder() + .with_peer_query_transport(Arc::new(PartialSubtreeTransport { + peer_engine, + peer_addr, + })) + .with_data_source(Arc::new(local_table) as Arc) + .build() + .await + .expect("coordinator engine"); + + let dataframe = engine + .async_query("SELECT v FROM global.demo.metrics ORDER BY v") + .await + .expect("global query") + .expect("global dataframe"); + assert_eq!(df_col_i32(&dataframe, "v"), vec![0, 2]); + + let stats = take_fanout_stats(); + assert!(stats.partial); + assert_eq!(stats.nodes_succeeded, 1); + assert_eq!(stats.nodes_failed, vec!["remote-leaf: timeout"]); + + for key in [ + "RANK", + "LOCAL_RANK", + "GROUP_RANK", + "PROBING_CLUSTER_FANOUT_HIERARCHICAL", + ] { + std::env::remove_var(key); + } + reset_cluster_for_tests(); +} + #[tokio::test] async fn global_catalog_discovers_probe_schema() { let _lock = federation_test_lock().await; @@ -250,6 +358,36 @@ async fn global_and_probe_return_same_ranks_without_peers() { ); } +#[tokio::test] +async fn global_query_with_remote_peer_missing_hierarchical_metadata_fails_closed() { + let _lock = federation_test_lock().await; + reset_cluster_for_tests(); + set_remote_query_hook(None); + std::env::set_var("PROBING_ADDRESS", "127.0.0.1:19999"); + std::env::set_var("PROBING_CLUSTER_FANOUT_HIERARCHICAL", "1"); + update_node(Node { + host: "peer-host".into(), + addr: "127.0.0.1:20001".into(), + rank: Some(1), + ..Default::default() + }); + + let engine = build_demo_engine().await; + let error = engine + .async_query("SELECT rank FROM global.demo.metrics") + .await + .expect_err("remote fan-out with incomplete metadata must fail closed"); + assert!( + error + .to_string() + .contains(HIERARCHICAL_METADATA_UNAVAILABLE), + "unexpected error: {error}" + ); + + std::env::remove_var("PROBING_CLUSTER_FANOUT_HIERARCHICAL"); + reset_cluster_for_tests(); +} + #[tokio::test] async fn global_select_name_returns_only_name() { let _lock = federation_test_lock().await; diff --git a/tests/regression/rust/probing/server/hierarchical_fanout_query.rs b/tests/regression/rust/probing/server/hierarchical_fanout_query.rs index 72ae7097..ae484b9c 100644 --- a/tests/regression/rust/probing/server/hierarchical_fanout_query.rs +++ b/tests/regression/rust/probing/server/hierarchical_fanout_query.rs @@ -1,5 +1,6 @@ //! Hierarchical cluster query fan-out integration test (mock HTTP peers + real engine). +use std::future::Future; use std::sync::{LazyLock, Mutex}; use axum::{extract::State, http::StatusCode, routing::post, Json, Router}; @@ -22,19 +23,26 @@ fn lock_test_env() -> std::sync::MutexGuard<'static, ()> { } fn local_http_available() -> bool { - match SERVER_RUNTIME.block_on(TcpListener::bind("127.0.0.1:0")) { - Ok(listener) => { + match SERVER_RUNTIME.try_block_on(TcpListener::bind("127.0.0.1:0")) { + Ok(Ok(listener)) => { drop(listener); true } - Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => { + Ok(Err(error)) if error.kind() == std::io::ErrorKind::PermissionDenied => { eprintln!("skipping local HTTP fan-out test: environment denied TCP bind ({error})"); false } - Err(error) => panic!("probe local HTTP bind capability: {error}"), + Ok(Err(error)) => panic!("probe local HTTP bind capability: {error}"), + Err(error) => panic!("probe runtime unavailable: {error}"), } } +fn run_on_server_runtime(future: impl Future) { + SERVER_RUNTIME + .try_block_on(future) + .expect("probing runtime"); +} + #[derive(Clone)] struct QueryState { rank: i32, @@ -182,7 +190,7 @@ fn hierarchical_fanout_contacts_node_aggregators_not_every_rank() { } clear_rank_env(); - SERVER_RUNTIME.block_on(async { + run_on_server_runtime(async { initialize_engine() .await .expect("initialize probing engine"); @@ -237,7 +245,7 @@ fn flat_fanout_contacts_all_remote_peers() { } clear_rank_env(); - SERVER_RUNTIME.block_on(async { + run_on_server_runtime(async { initialize_engine() .await .expect("initialize probing engine"); @@ -275,7 +283,7 @@ fn hierarchical_fanout_rejects_without_metadata() { } clear_rank_env(); - SERVER_RUNTIME.block_on(async { + run_on_server_runtime(async { initialize_engine() .await .expect("initialize probing engine"); @@ -311,7 +319,7 @@ fn hierarchical_fanout_reports_failed_remote_node_aggregator() { } clear_rank_env(); - SERVER_RUNTIME.block_on(async { + run_on_server_runtime(async { initialize_engine() .await .expect("initialize probing engine"); @@ -356,7 +364,7 @@ fn hierarchical_fanout_reports_failed_local_leaf() { } clear_rank_env(); - SERVER_RUNTIME.block_on(async { + run_on_server_runtime(async { initialize_engine() .await .expect("initialize probing engine"); @@ -416,7 +424,7 @@ fn hierarchical_fanout_leaf_rank_stays_local_only() { } clear_rank_env(); - SERVER_RUNTIME.block_on(async { + run_on_server_runtime(async { initialize_engine() .await .expect("initialize probing engine"); diff --git a/tests/unit/probing/skills/test_loader.py b/tests/unit/probing/skills/test_loader.py index b3ac32da..b3759b7f 100644 --- a/tests/unit/probing/skills/test_loader.py +++ b/tests/unit/probing/skills/test_loader.py @@ -52,6 +52,10 @@ def test_catalog_loads_all_skills(): def test_load_slow_rank_global(): skill = load_skill("slow_rank") + assert {parameter["type"] for parameter in skill.parameters} == { + "integer", + "boolean", + } steps = expand_skill(skill, {"use_global": True, "step_window": 10}) assert steps sql = " ".join(s.sql or "" for s in steps if s.sql) diff --git a/tests/unit/probing/test_web_assets.py b/tests/unit/probing/test_web_assets.py index f10ac2bc..0c112a58 100644 --- a/tests/unit/probing/test_web_assets.py +++ b/tests/unit/probing/test_web_assets.py @@ -1,98 +1,16 @@ -"""Tests for wheel / editable web UI asset resolution.""" +"""Tests for Web bundle integrity and native-wheel embedding checks.""" from __future__ import annotations -import os import zipfile from pathlib import Path -import pytest - -from probing import web_assets from scripts.verify_web_assets import verify_web_bundle -from scripts.verify_wheel_contents import REQUIRED_PATHS, verify_wheel - -from tests.conftest import is_wheel_install, repo_root - - -def test_bundled_web_dir_missing_without_sync(): - root = web_assets.bundled_web_dir() - checkout_bundled = ( - repo_root() / "python" / "probing" / "bundled_web" / "public" / "index.html" - ) - legacy_bundled = repo_root() / "python" / "probing" / "bundled_web" / "index.html" - if is_wheel_install(): - assert root is not None, "installed wheel is missing probing/bundled_web" - assert (root / "index.html").is_file() - return - if root is None: - assert not checkout_bundled.is_file() and not legacy_bundled.is_file() - else: - assert (root / "index.html").is_file() - - -def test_dev_web_dir_when_frontend_built(): - root = web_assets.dev_web_dir() - built = repo_root() / "python" / "probing" / "bundled_web" / "public" / "index.html" - if is_wheel_install(): - pytest.skip("dev_web_dir applies to editable checkout layout only") - if built.is_file(): - assert root is not None - assert (root / "index.html").is_file() - assert root.resolve() == built.parent.resolve() - else: - assert root is None - - -def test_configure_assets_root_prefers_dev_in_editable(monkeypatch, tmp_path: Path): - bundled = tmp_path / "_web" - bundled.mkdir() - (bundled / "index.html").write_text("bundled", encoding="utf-8") - - dev = tmp_path / "web" / "dist" - dev.mkdir(parents=True) - (dev / "index.html").write_text( - '
', - encoding="utf-8", - ) - - monkeypatch.setattr(web_assets, "bundled_web_dir", lambda: bundled) - monkeypatch.setattr(web_assets, "dev_web_dir", lambda: dev) - monkeypatch.setattr(web_assets, "_running_from_installed_wheel", lambda: False) - monkeypatch.delenv(web_assets._ENV, raising=False) - - assert web_assets.configure_assets_root() == dev - assert os.environ[web_assets._ENV] == str(dev) - - -def test_configure_assets_root_prefers_bundled_on_wheel(monkeypatch, tmp_path: Path): - bundled = tmp_path / "_web" - bundled.mkdir() - (bundled / "index.html").write_text( - '
', - encoding="utf-8", - ) - - dev = tmp_path / "web" / "dist" - dev.mkdir(parents=True) - (dev / "index.html").write_text("dev", encoding="utf-8") - - monkeypatch.setattr(web_assets, "bundled_web_dir", lambda: bundled) - monkeypatch.setattr(web_assets, "dev_web_dir", lambda: dev) - monkeypatch.setattr(web_assets, "_running_from_installed_wheel", lambda: True) - monkeypatch.delenv(web_assets._ENV, raising=False) - - assert web_assets.configure_assets_root() == bundled - assert os.environ[web_assets._ENV] == str(bundled) - - -def test_configure_assets_root_respects_override(monkeypatch, tmp_path: Path): - override = tmp_path / "custom" - override.mkdir() - (override / "index.html").write_text("custom", encoding="utf-8") - monkeypatch.setenv(web_assets._ENV, str(override)) - - assert web_assets.configure_assets_root() == override +from scripts.verify_wheel_contents import ( + EMBEDDED_WEB_MARKER, + REQUIRED_PATHS, + verify_wheel, +) def _write_valid_web_bundle(root: Path) -> None: @@ -132,15 +50,20 @@ def test_verify_web_bundle_rejects_missing_wasm(tmp_path: Path): assert any("missing WASM module" in error for error in errors) -def test_verify_wheel_rejects_broken_web_reference(tmp_path: Path): +def test_verify_wheel_rejects_native_extension_without_embedded_web(tmp_path: Path): wheel = tmp_path / "probing-test.whl" with zipfile.ZipFile(wheel, "w") as archive: for path in REQUIRED_PATHS: - content = ( - '' - if path.endswith("bundled_web/public/index.html") - else "" - ) - archive.writestr(path, content) + archive.writestr(path, "") + archive.writestr("probing/_core.test.so", b"native-without-assets") errors = verify_wheel(wheel) - assert any("invalid web bundle" in error for error in errors) + assert "native extension does not contain embedded Web assets" in errors + + +def test_verify_wheel_accepts_embedded_web_marker(tmp_path: Path): + wheel = tmp_path / "probing-test.whl" + with zipfile.ZipFile(wheel, "w") as archive: + for path in REQUIRED_PATHS: + archive.writestr(path, "") + archive.writestr("probing/_core.test.so", b"native" + EMBEDDED_WEB_MARKER) + assert verify_wheel(wheel) == [] diff --git a/web/Cargo.toml b/web/Cargo.toml index 0d68df39..393a775a 100644 --- a/web/Cargo.toml +++ b/web/Cargo.toml @@ -77,10 +77,10 @@ dioxus-code = { version = "0.1.2", default-features = false, features = [ "lang-toml", ] } -# Match workspace root: fast dev compiles, lean debug artifacts. +# Match the workspace root's incremental, symbol-free dev profile. Opt in to +# symbols with CARGO_PROFILE_DEV_DEBUG=1. [profile.dev] -debug = 1 -split-debuginfo = "unpacked" +debug = 0 incremental = true codegen-units = 256 diff --git a/web/DESIGN.md b/web/DESIGN.md index bbcddae6..58df9767 100644 --- a/web/DESIGN.md +++ b/web/DESIGN.md @@ -423,7 +423,7 @@ web/src/ ## 九、构建与部署 - 开发 / 构建:`dx serve` / `dx build --release`;仓库根 `make frontend` 复制产物到 `web/dist/`。 -- UI 静态资源由 Python 包提供(wheel:`python/probing/_web/`;editable:`web/dist/`),经 `probing.web_assets` 设置 `PROBING_ASSETS_ROOT`,`probing-server` 只读该目录;未配置时返回占位页。 +- UI 静态资源由 `make frontend` 生成到被 Git 忽略的 `probing/server/web-assets/`,build script 将其复制到 `$OUT_DIR` 后通过 `include_dir` 编译进 `probing._core`;没有前端产物的普通 Rust 构建使用轻量 fallback,`PROBING_ASSETS_ROOT` 仅作为开发期显式磁盘覆盖。 --- diff --git a/web/Dioxus.toml b/web/Dioxus.toml index 4929cbc2..6c50e4fb 100644 --- a/web/Dioxus.toml +++ b/web/Dioxus.toml @@ -1,7 +1,8 @@ [application] name = "probing-web" -# dx bundle copies the web public folder here (see Makefile frontend). -out_dir = "../python/probing/bundled_web" +# dx bundle writes the public folder into the server crate; include_dir embeds +# it into the native library at compile time (see Makefile frontend). +out_dir = "../probing/server/web-assets" # dx autodetects tailwind.config.js (v3) and compiles on build/serve. tailwind_input = "tailwind.css" tailwind_output = "assets/tailwind.css" diff --git a/web/assets/tailwind.css b/web/assets/tailwind.css index 1495cecd..e8003625 100644 --- a/web/assets/tailwind.css +++ b/web/assets/tailwind.css @@ -1190,6 +1190,13 @@ video { margin-top: 1rem; } +.line-clamp-2 { + overflow: hidden; + display: -webkit-box; + -webkit-box-orient: vertical; + -webkit-line-clamp: 2; +} + .block { display: block; } @@ -2001,6 +2008,10 @@ video { grid-template-columns: minmax(0,1fr) auto auto; } +.grid-cols-\[minmax\(180px\2c 0\.7fr\)_minmax\(0\2c 1\.7fr\)_auto\] { + grid-template-columns: minmax(180px,0.7fr) minmax(0,1.7fr) auto; +} + .grid-cols-\[minmax\(180px\2c 22\%\)_130px_1fr\] { grid-template-columns: minmax(180px,22%) 130px 1fr; } @@ -8164,10 +8175,6 @@ video { grid-template-columns: repeat(4, minmax(0, 1fr)); } - .xl\:grid-cols-5 { - grid-template-columns: repeat(5, minmax(0, 1fr)); - } - .xl\:grid-cols-\[320px_minmax\(0\2c 1fr\)\] { grid-template-columns: 320px minmax(0,1fr); } diff --git a/web/dist b/web/dist deleted file mode 120000 index d94bccd8..00000000 --- a/web/dist +++ /dev/null @@ -1 +0,0 @@ -../python/probing/bundled_web/public \ No newline at end of file