From 40c857748fc6ad5b07e2dafee10b516dc9df21cd Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9F=A9=E5=98=89=E4=B9=90?= Date: Sat, 1 Aug 2026 16:35:14 +0800 Subject: [PATCH] =?UTF-8?q?feat(ohos):=20=E6=8E=A5=E5=85=A5=20Pro=20?= =?UTF-8?q?=E8=BF=90=E8=A1=8C=E6=97=B6=E5=B9=B6=E8=87=AA=E5=8A=A8=E5=8C=96?= =?UTF-8?q?=20HAR=20=E4=BA=A4=E4=BB=98=20(#2462)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .github/workflows/ohos.yml | 355 +++++------ docs/ohos-downstream-builds.md | 65 +++ easytier-contrib/easytier-ohrs/src/lib.rs | 582 ++++++++++++++++++- easytier-core/src/connectivity/manual/mod.rs | 249 +++++--- 4 files changed, 982 insertions(+), 269 deletions(-) create mode 100644 docs/ohos-downstream-builds.md diff --git a/.github/workflows/ohos.yml b/.github/workflows/ohos.yml index 4308e0fb..a393b7bf 100644 --- a/.github/workflows/ohos.yml +++ b/.github/workflows/ohos.yml @@ -1,246 +1,205 @@ -name: EasyTier OHOS +name: ohos on: push: - branches: ["develop", "main", "releases/**"] + branches: [develop, main, "releases/**", "ohos/**"] tags: - - 'v*' - - '!*-pre' + - "v*" + - "!*-pre" pull_request: - branches: ["develop", "main"] + branches: [develop, main, "ohos/**"] types: [opened, synchronize, reopened, ready_for_review] workflow_dispatch: + inputs: + publish: + description: Publish this non-main branch and dispatch downstream builds + required: false + default: false + type: boolean -concurrency: - group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} - cancel-in-progress: true +permissions: + contents: read + pull-requests: read env: CARGO_TERM_COLOR: always defaults: run: - # necessary for windows shell: bash jobs: - cargo_fmt_check: + ohos: + name: ohos if: github.event_name != 'pull_request' || !github.event.pull_request.draft runs-on: ubuntu-latest - steps: - - uses: actions/checkout@v5 - - name: Prepare build environment + steps: + - name: Checkout + uses: actions/checkout@v5 + with: + fetch-depth: 0 + + - name: Set up Rust uses: ./.github/actions/prepare-build with: + target: aarch64-unknown-linux-ohos gui: false pnpm: false token: ${{ secrets.GITHUB_TOKEN }} - - uses: actions-rust-lang/setup-rust-toolchain@v1 - with: - components: rustfmt - - - name: Check formatting - working-directory: ./easytier-contrib/easytier-ohrs - run: cargo fmt --all -- --check - - pre_job: - # continue-on-error: true # Uncomment once integration is finished - runs-on: ubuntu-latest - if: github.event_name != 'pull_request' || !github.event.pull_request.draft - # Map a step output to a job output - outputs: - # do not skip push on branch starts with releases/ - should_skip: ${{ steps.skip_check.outputs.should_skip == 'true' && !startsWith(github.ref_name, 'releases/') }} - steps: - - id: skip_check - uses: fkirc/skip-duplicate-actions@v5 - with: - # All of these options are optional, so you can remove them if you are happy with the defaults - concurrent_skipping: "same_content_newer" - skip_after_successful_duplicate: "true" - cancel_others: "true" - paths: '["Cargo.toml", "Cargo.lock", "easytier/**", "easytier-core/**", "easytier-contrib/easytier-ohrs/**", ".github/workflows/ohos.yml", ".github/actions/**"]' - - build-ohos: - runs-on: ubuntu-latest - needs: pre_job - env: - OHPM_PUBLISH_CODE: ${{ secrets.OHPM_PUBLISH_CODE }} - if: needs.pre_job.outputs.should_skip != 'true' - steps: - - uses: actions/checkout@v5 - - name: Install dependencies - run: | - sudo apt-get update - sudo apt-get install -qq \ - build-essential \ - wget \ - unzip \ - git \ - pkg-config curl libgl1-mesa-dev expect - - - name: Resolve easytier version - run: | - set -e - - UPSTREAM_REPO="https://github.com/EasyTier/EasyTier.git" - - git remote add upstream "$UPSTREAM_REPO" 2>/dev/null || true - git fetch --unshallow upstream main || git fetch upstream main - git fetch --tags upstream --force - - # 读取 cargo 版本 - CARGO_VERSION=$(cargo metadata --format-version 1 --no-deps --manifest-path easytier/Cargo.toml \ - | jq -r '.packages[0].version') - - # 获取 upstream/main 最新 tag - LAST_TAG=$(git describe --tags --abbrev=0 upstream/main 2>/dev/null || echo "") - LAST_TAG_VERSION="${LAST_TAG#v}" - - # 语义版本比较 - version_gt() { - [ "$(printf '%s\n' "$1" "$2" | sort -V | tail -n1)" = "$1" ] && [ "$1" != "$2" ] - } - - if [ -z "$LAST_TAG_VERSION" ]; then - BASE_VERSION="$CARGO_VERSION" - DIFF_COUNT=$(git rev-list --count upstream/main) - elif version_gt "$CARGO_VERSION" "$LAST_TAG_VERSION"; then - BASE_VERSION="$CARGO_VERSION" - DIFF_COUNT=0 - else - BASE_VERSION="$LAST_TAG_VERSION" - DIFF_COUNT=$(git rev-list --count "${LAST_TAG}..upstream/main") - fi - - COMMIT_HASH=$(git rev-parse --short upstream/main) - EASYTIER_VERSION="${BASE_VERSION}-${DIFF_COUNT}-${COMMIT_HASH}" - - echo "EASYTIER_VERSION=$EASYTIER_VERSION" - echo "EASYTIER_VERSION=$EASYTIER_VERSION" >> $GITHUB_ENV - - cd ./easytier-contrib/easytier-ohrs/package - jq --arg v "$EASYTIER_VERSION" '.version = $v' oh-package.json5 > oh-package.tmp.json5 - mv oh-package.tmp.json5 oh-package.json5 - - - - name: Generate CHANGELOG.md for current commit - working-directory: ./easytier-contrib/easytier-ohrs/package - run: | - { - echo "## easytier-ohrs ${EASYTIER_VERSION}" - echo - git log -1 --pretty=format:"- %s" - echo - } > CHANGELOG.md - - - name: Setup HarmonyOS CLI tools + - name: Set up HarmonyOS uses: ErBWs/setup-ohos@v1 - - name: Download and Extract Custom SDK - run: | - wget https://github.com/FrankHan052176/Easytier-OHOS-sdk/releases/download/v1/ohos-sdk.zip -O /tmp/ohos-sdk.zip - sudo unzip -o /tmp/ohos-sdk.zip -d /tmp/custom-sdk - sudo cp -rf /tmp/custom-sdk/linux/native/* $OHOS_NDK_HOME/native - echo "Custom SDK files deployed to $OHOS_NDK_HOME/native" - ls -a $OHOS_NDK_HOME/native - - - name: Setup build environment - run: | - echo "TARGET_ARCH=aarch64-linux-ohos" >> $GITHUB_ENV - - rustup install stable - rustup default stable - - rustup target add aarch64-unknown-linux-ohos - - - uses: taiki-e/install-action@v2 + - name: Install ohrs + uses: taiki-e/install-action@v2 with: tool: ohrs - - name: Create clang wrapper script + - name: Build HAR + id: package + env: + DEFAULT_BRANCH: ${{ github.event.repository.default_branch }} run: | - sudo mkdir -p $OHOS_NDK_HOME/native/llvm - sudo tee $OHOS_NDK_HOME/native/llvm/aarch64-unknown-linux-ohos-clang.sh > /dev/null <<'EOF' + set -euo pipefail + sudo apt-get install -qqy \ + pkg-config curl libgl1-mesa-dev expect llvm clang lldb lld + rustup component add rustfmt + cargo fmt --all --manifest-path \ + easytier-contrib/easytier-ohrs/Cargo.toml -- --check + + cargo_version=$(cargo metadata --format-version 1 --no-deps \ + --manifest-path easytier/Cargo.toml | jq -r '.packages[0].version') + last_tag=$(git describe --tags --abbrev=0 HEAD 2>/dev/null || true) + if [ -n "$last_tag" ]; then + base_version=$(printf '%s\n' "$cargo_version" "${last_tag#v}" \ + | sort -V | tail -n 1) + commit_count=$(git rev-list --count "$last_tag..HEAD") + else + base_version=$cargo_version + commit_count=0 + fi + + source_branch=${GITHUB_HEAD_REF:-} + if [ -z "$source_branch" ]; then + if [ "$GITHUB_REF_TYPE" = branch ]; then + source_branch=$GITHUB_REF_NAME + else + source_branch=${DEFAULT_BRANCH:-main} + fi + fi + branch_id=$(printf '%s' "$source_branch" \ + | tr '[:upper:]' '[:lower:]' \ + | sed -E 's/[^a-z0-9-]+/-/g; s/^-+//; s/-+$//' \ + | cut -c1-64) + branch_id=${branch_id:-main} + + package_name=easytier-ohrs + package_version="${base_version}-${branch_id}-${commit_count}-${GITHUB_RUN_NUMBER}-${GITHUB_RUN_ATTEMPT}-g$(git rev-parse --short=8 HEAD)" + echo "name=$package_name" >> "$GITHUB_OUTPUT" + echo "EASYTIER_PACKAGE_NAME=$package_name" >> "$GITHUB_ENV" + echo "EASYTIER_VERSION=$package_version" >> "$GITHUB_ENV" + + package_dir=easytier-contrib/easytier-ohrs/package + jq --arg name "$package_name" --arg version "$package_version" \ + '.name = $name | .version = $version' \ + "$package_dir/oh-package.json5" > "$package_dir/oh-package.tmp.json5" + mv "$package_dir/oh-package.tmp.json5" "$package_dir/oh-package.json5" + { + echo "## $package_name $package_version" + echo + echo "- Core version: $base_version" + echo "- Core commit: $GITHUB_SHA" + git log -1 --pretty=format:'- %s' + echo + } > "$package_dir/CHANGELOG.md" + + sudo mkdir -p "$OHOS_NDK_HOME/native/llvm" + sudo tee "$OHOS_NDK_HOME/native/llvm/aarch64-unknown-linux-ohos-clang.sh" >/dev/null <<'EOF' #!/bin/sh - exec $OHOS_NDK_HOME/native/llvm/bin/clang \ + exec "$OHOS_NDK_HOME/native/llvm/bin/clang" \ -target aarch64-linux-ohos \ - --sysroot=$OHOS_NDK_HOME/native/sysroot \ - -D__MUSL__ \ - "$@" + --sysroot="$OHOS_NDK_HOME/native/sysroot" \ + -D__MUSL__ "$@" EOF - sudo chmod +x $OHOS_NDK_HOME/native/llvm/aarch64-unknown-linux-ohos-clang.sh + sudo chmod +x \ + "$OHOS_NDK_HOME/native/llvm/aarch64-unknown-linux-ohos-clang.sh" - - name: Build latest Har - working-directory: ./easytier-contrib/easytier-ohrs - run: | - sudo apt-get install -y llvm clang lldb lld - sudo apt-get install -y protobuf-compiler + cd easytier-contrib/easytier-ohrs source env.sh - ohrs doctor ohrs build --release --arch aarch ohrs artifact - mv package.har easytier-ohrs.har + mv package.har "$package_name.har" - - name: Build Release Package - if: startsWith(github.ref, 'refs/tags/') - working-directory: ./easytier-contrib/easytier-ohrs - run: | - echo "🎉 Official Release detected. Building easytier-release..." - TAG_NAME="${{ github.ref_name }}" - TAG_VERSION="${TAG_NAME#v}" - echo "Release Version: $TAG_VERSION" - cd package - jq --arg v "$TAG_VERSION" '.name = "easytier-release" | .version = $v' oh-package.json5 > oh-package.tmp.json5 && mv oh-package.tmp.json5 oh-package.json5 - cd .. - ohrs build --release --arch aarch - cd dist/arm64-v8a - mv libeasytier_ohrs.so libeasytier_release.so - cd ../.. - ohrs artifact - mv package.har easytier-release.har - - - name: Upload artifact + - name: Upload HAR uses: actions/upload-artifact@v5 with: - name: easytier-ohos - path: | - ./easytier-contrib/easytier-ohrs/easytier-ohrs.har + name: ${{ steps.package.outputs.name }} + path: easytier-contrib/easytier-ohrs/${{ steps.package.outputs.name }}.har retention-days: 5 if-no-files-found: error - - name: Publish To Center Ohpm - working-directory: ./easytier-contrib/easytier-ohrs + - name: Publish and dispatch + if: >- + (github.event_name == 'push' && + github.ref_type == 'branch' && + github.ref_name == 'main' && + github.event.forced != true) || + (github.event_name == 'workflow_dispatch' && + github.ref_type == 'branch' && + (github.ref_name == 'main' || inputs.publish)) + working-directory: easytier-contrib/easytier-ohrs env: - OHPM_PRIVATE_KEY: ${{ secrets.OHPM_PRIVATE_KEY }} - OHPM_KEY_PASSPHRASE: ${{ secrets.OHPM_KEY_PASSPHRASE }} - if: ${{ env.OHPM_PUBLISH_CODE != '' && github.event_name == 'push' }} + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + CODEARTS_PRIVATE_OHPM: ${{ secrets.CODEARTS_PRIVATE_OHPM }} + DOWNSTREAM_DISPATCH_TOKEN: ${{ secrets.DOWNSTREAM_DISPATCH_TOKEN }} run: | - ohpm config set publish_id "$OHPM_PUBLISH_CODE" - ohpm config set publish_registry https://ohpm.openharmony.cn/ohpm - TMP_DIR=$(mktemp -d) - PRIVATE_KEY_FILE="$TMP_DIR/private_key" - printf '%s' "$OHPM_PRIVATE_KEY" > "$PRIVATE_KEY_FILE" - chmod 600 "$PRIVATE_KEY_FILE" - ohpm config set key_path $PRIVATE_KEY_FILE - unzip ohpm_crypto.zip -d /home/runner/work/ - ohpm config set crypto_path /home/runner/work/ohpm_crypto - chmod 755 /home/runner/work/ohpm_crypto/* - PASSPHRASE="$(printf '%s' "$OHPM_KEY_PASSPHRASE" | tr -d '\r\n')" - ohpm config set key_passphrase "$PASSPHRASE" - ohpm publish easytier-ohrs.har - - - name: Publish To Private Ohpm - working-directory: ./easytier-contrib/easytier-ohrs - if: ${{ env.OHPM_PUBLISH_CODE != '' && github.event_name == 'push' }} - run: | - printf '%s' "${{ secrets.CODEARTS_PRIVATE_OHPM }}" > ~/.ohpm/.ohpmrc - ohpm config set strict_ssl false - ohpm publish easytier-ohrs.har - if [ -f "easytier-release.har" ]; then - echo "🚀 Publishing Release package..." - ohpm publish easytier-release.har + set -euo pipefail + if [ "$GITHUB_EVENT_NAME" = push ]; then + pull_requests=$(gh api \ + -H "Accept: application/vnd.github+json" \ + "/repos/$GITHUB_REPOSITORY/commits/$GITHUB_SHA/pulls") + if ! jq -e \ + --arg repository "$GITHUB_REPOSITORY" \ + --arg branch "$GITHUB_REF_NAME" \ + --arg sha "$GITHUB_SHA" \ + 'any(.[]; + .merged_at != null and + .base.repo.full_name == $repository and + .base.ref == $branch and + .merge_commit_sha == $sha)' \ + <<< "$pull_requests" >/dev/null; then + echo "Direct push: HAR built without publishing." + exit 0 + fi fi - curl --header "Content-Type: application/json" --request POST --data "{}" ${{ secrets.CODEARTS_WEBHOOKS }} + + mkdir -p "$HOME/.ohpm" + umask 077 + printf '%s' "$CODEARTS_PRIVATE_OHPM" > "$HOME/.ohpm/.ohpmrc" + trap 'rm -f "$HOME/.ohpm/.ohpmrc"' EXIT + ohpm publish "$EASYTIER_PACKAGE_NAME.har" + + payload=$(jq -nc \ + --arg repository "$GITHUB_REPOSITORY" \ + --arg ref "refs/heads/$GITHUB_REF_NAME" \ + --arg package "$EASYTIER_PACKAGE_NAME" \ + '{ + event_type: "core-har-published", + client_payload: { + core_repository: $repository, + core_ref: $ref, + package_name: $package + } + }') + for repository in \ + FrankHan052176/EasyTier-ArkTS \ + FrankHan052176/easytier-pro-app; do + curl --fail-with-body --silent --show-error \ + -X POST \ + -H "Accept: application/vnd.github+json" \ + -H "Authorization: Bearer $DOWNSTREAM_DISPATCH_TOKEN" \ + -H "X-GitHub-Api-Version: 2022-11-28" \ + "$GITHUB_API_URL/repos/$repository/dispatches" \ + --data "$payload" + done diff --git a/docs/ohos-downstream-builds.md b/docs/ohos-downstream-builds.md new file mode 100644 index 00000000..b08a4933 --- /dev/null +++ b/docs/ohos-downstream-builds.md @@ -0,0 +1,65 @@ +# HarmonyOS HAR delivery + +The `ohos` workflow builds the Core HAR on pushes, pull requests, tags, and +manual runs. Every successful run retains a short-lived HAR artifact, while +publication to the private OHPM registry is deliberately restricted: + +- A push to `main` publishes only when the pushed SHA is the merge commit of a + pull request targeting `main` and the push is not forced. +- A manual run on `main` publishes by default. +- A manual run on another branch publishes only when its `publish` input is + enabled. +- Direct pushes, pull requests, tags, and ordinary non-main branch builds do + not publish. + +## Package identity + +All branches publish the same private package name, `easytier-ohrs`. The +source branch is encoded in the package version instead of the package name: + +```text +-----g +``` + +`branch-id` is a lowercase, OHPM-safe form of the source branch. Publishing a +new version advances the registry's `latest` version. After publication, Core +sends the `core-har-published` repository dispatch to the ArkTS and Pro +repositories. The payload contains only `core_repository`, `core_ref`, and +`package_name`. + +## App install sequence + +ArkTS and Pro use the same three OHPM commands: + +```bash +ohpm uninstall "$CORE_HAR_PACKAGE" +ohpm install "$CORE_HAR_PACKAGE@latest" \ + --registry "$CORE_HAR_REGISTRY" +ohpm install +``` + +The App workflow then reads the installed version from: + +```text +oh_modules//oh-package.json5 +``` + +The existing `oh-package-lock.json5` and `oh_modules` directory are not +manually deleted. Because the package name remains `easytier-ohrs`, downstream +source imports do not need to be rewritten. + +## Secrets + +Core requires: + +- `CODEARTS_PRIVATE_OHPM`: publish-capable OHPM configuration. +- `DOWNSTREAM_DISPATCH_TOKEN`: permission to dispatch both App repositories. + +ArkTS and Pro require: + +- `CODEARTS_PRIVATE_OHPM_READ`: read-only private OHPM authentication. +- `SIGNING_REPOSITORY_TOKEN`: read access to the corresponding private signing + repository. + +Signing and AppGallery Connect credentials remain downstream application +concerns and are not passed through the Core dispatch payload. diff --git a/easytier-contrib/easytier-ohrs/src/lib.rs b/easytier-contrib/easytier-ohrs/src/lib.rs index 47ea077f..0b224d05 100644 --- a/easytier-contrib/easytier-ohrs/src/lib.rs +++ b/easytier-contrib/easytier-ohrs/src/lib.rs @@ -68,7 +68,7 @@ use kernel_bridge::{ stop_local_socket_server as stop_local_socket_server_inner, }; use napi_derive_ohos::napi; -use runtime::state::runtime_state::RuntimeAggregateState; +use runtime::state::runtime_state::{RuntimeAggregateState, RuntimeInstanceState}; use std::collections::{HashMap, HashSet}; use std::format; use std::sync::{Arc, Mutex}; @@ -89,10 +89,13 @@ pub(crate) static INSTANCE_MANAGER: once_cell::sync::Lazy>> = once_cell::sync::Lazy::new(|| Mutex::new(HashMap::new())); +const PRO_CONFIG_SERVER_CLIENT_ID: &str = "__easytier_pro_config_server_client__"; #[derive(Default)] struct TrackedWebClientHooks { instance_ids: Mutex>, + network_names_by_instance_id: Mutex>, + events: Mutex>, } struct ManagedWebClient { @@ -100,6 +103,13 @@ struct ManagedWebClient { hooks: Arc, } +fn network_name_for_instance(id: &Uuid) -> Option { + INSTANCE_MANAGER + .config(*id) + .map(|config| config.get_network_identity().network_name) + .filter(|name| !name.trim().is_empty()) +} + #[async_trait::async_trait] impl WebClientHooks for TrackedWebClientHooks { async fn post_run_network_instance(&self, id: &Uuid) -> Result<(), String> { @@ -107,13 +117,43 @@ impl WebClientHooks for TrackedWebClientHooks { .lock() .map_err(|err| err.to_string())? .insert(*id); + let network_name = network_name_for_instance(id); + if let Some(network_name) = &network_name { + self.network_names_by_instance_id + .lock() + .map_err(|err| err.to_string())? + .insert(*id, network_name.clone()); + } + self.events + .lock() + .map_err(|err| err.to_string())? + .push(serde_json::json!({ + "event": "run_network_instance", + "success": true, + "instance_id": id.to_string(), + "instance_name": id.to_string(), + "network_name": network_name, + })); Ok(()) } async fn post_remove_network_instances(&self, ids: &[Uuid]) -> Result<(), String> { let mut guard = self.instance_ids.lock().map_err(|err| err.to_string())?; + let mut events = self.events.lock().map_err(|err| err.to_string())?; + let mut network_names_by_instance_id = self + .network_names_by_instance_id + .lock() + .map_err(|err| err.to_string())?; for id in ids { guard.remove(id); + let network_name = network_names_by_instance_id.remove(id); + events.push(serde_json::json!({ + "event": "delete_network_instance", + "success": true, + "instance_id": id.to_string(), + "instance_name": id.to_string(), + "network_name": network_name, + })); } Ok(()) } @@ -243,6 +283,388 @@ fn run_config_server_instance(config_id: &str, config: &NetworkConfig) -> bool { } } +fn run_config_server_client( + url: &str, + hostname: Option, + machine_id: Option, + secure_mode: bool, +) -> bool { + let trimmed_url = url.trim(); + if trimmed_url.is_empty() { + ohrs_log_error!("[Rust] config server url missing"); + return false; + } + + let _ = stop_web_client(PRO_CONFIG_SERVER_CLIENT_ID); + let hooks = Arc::new(TrackedWebClientHooks::default()); + + if !ensure_local_socket_server_started() { + return false; + } + + let machine_id_opts = MachineIdOptions { + explicit_machine_id: machine_id.filter(|value| !value.trim().is_empty()), + state_dir: None, + }; + let client = ASYNC_RUNTIME.block_on(run_web_client( + trimmed_url, + machine_id_opts, + hostname.filter(|value| !value.trim().is_empty()), + secure_mode, + INSTANCE_MANAGER.clone(), + Some(hooks.clone()), + )); + + let client = match client { + Ok(client) => client, + Err(err) => { + ohrs_log_error!("[Rust] start pro config server client failed {}", err); + return false; + } + }; + + match WEB_CLIENTS.lock() { + Ok(mut guard) => { + guard.insert( + PRO_CONFIG_SERVER_CLIENT_ID.to_string(), + ManagedWebClient { + _client: client, + hooks, + }, + ); + true + } + Err(err) => { + ohrs_log_error!("[Rust] store pro config server client failed {}", err); + false + } + } +} + +fn pro_config_server_client_connected() -> bool { + WEB_CLIENTS + .lock() + .ok() + .and_then(|guard| { + guard + .get(PRO_CONFIG_SERVER_CLIENT_ID) + .map(|managed| managed._client.is_connected()) + }) + .unwrap_or(false) +} + +fn drain_config_server_events_inner() -> Vec { + let Ok(guard) = WEB_CLIENTS.lock() else { + return Vec::new(); + }; + let Some(managed) = guard.get(PRO_CONFIG_SERVER_CLIENT_ID) else { + return Vec::new(); + }; + let Ok(mut events) = managed.hooks.events.lock() else { + return Vec::new(); + }; + events.drain(..).collect() +} + +fn stop_runtime_inner() -> bool { + let mut ok = stop_web_client(PRO_CONFIG_SERVER_CLIENT_ID); + let ids = INSTANCE_MANAGER.instance_ids(); + if !ids.is_empty() { + ok = ASYNC_RUNTIME + .block_on(INSTANCE_MANAGER.delete_network_instances(ids)) + .map(|_| true) + .unwrap_or_else(|err| { + ohrs_log_error!("[Rust] stop runtime instances failed {}", err); + false + }) + && ok; + } + maybe_stop_local_socket_server(); + ok +} + +fn is_pro_internal_instance(instance: &RuntimeInstanceState) -> bool { + instance.instance_id == PRO_CONFIG_SERVER_CLIENT_ID + || instance.config_id == PRO_CONFIG_SERVER_CLIENT_ID + || instance.display_name == PRO_CONFIG_SERVER_CLIENT_ID +} + +fn runtime_instance_label(instance: &RuntimeInstanceState) -> String { + let display_name = instance.display_name.trim(); + if !display_name.is_empty() && display_name != PRO_CONFIG_SERVER_CLIENT_ID { + return display_name.to_string(); + } + let instance_id = instance.instance_id.trim(); + if !instance_id.is_empty() { + return instance_id.to_string(); + } + instance.config_id.clone() +} + +fn runtime_instance_matches(instance: &RuntimeInstanceState, selector: &str) -> bool { + let target = selector.trim(); + if target.is_empty() { + return false; + } + instance.instance_id == target + || instance.config_id == target + || instance.display_name == target + || runtime_instance_label(instance) == target +} + +fn read_json_string_path<'a>(value: &'a serde_json::Value, path: &[&str]) -> Option<&'a str> { + let mut cursor = value; + for key in path { + cursor = cursor.get(*key)?; + } + cursor.as_str().filter(|value| !value.trim().is_empty()) +} + +fn selected_instance_from_payload(payload_json: &str) -> Option { + let value = serde_json::from_str::(payload_json).ok()?; + for path in [ + &["instance", "instance_selector", "name"][..], + &["instance", "instanceSelector", "name"][..], + &["instance", "instance_selector", "id"][..], + &["instance", "instanceSelector", "id"][..], + &["instance", "name"][..], + &["instance", "id"][..], + &["instance_name"][..], + &["instanceName"][..], + &["instance_id"][..], + &["instanceId"][..], + &["id"][..], + ] { + if let Some(value) = read_json_string_path(&value, path) { + return Some(value.to_string()); + } + } + None +} + +fn find_runtime_instance<'a>( + state: &'a RuntimeAggregateState, + selector: Option<&str>, +) -> Option<&'a RuntimeInstanceState> { + if let Some(selector) = selector + && let Some(instance) = state.instances.iter().find(|instance| { + !is_pro_internal_instance(instance) && runtime_instance_matches(instance, selector) + }) + { + return Some(instance); + } + state + .instances + .iter() + .find(|instance| !is_pro_internal_instance(instance) && instance.running) +} + +fn list_instances_json_inner(state: &RuntimeAggregateState) -> String { + let mut instances = serde_json::Map::new(); + for instance in state + .instances + .iter() + .filter(|instance| !is_pro_internal_instance(instance) && instance.running) + { + let label = runtime_instance_label(instance); + if !label.trim().is_empty() { + instances.insert( + label, + serde_json::Value::String(instance.instance_id.clone()), + ); + } + } + serde_json::Value::Object(instances).to_string() +} + +fn list_pro_instances_json_inner( + state: &RuntimeAggregateState, + network_names_by_instance_id: &HashMap, +) -> String { + let mut instances = serde_json::Map::new(); + for instance in state.instances.iter().filter(|instance| instance.running) { + let Some(network_name) = network_names_by_instance_id.get(&instance.instance_id) else { + continue; + }; + let label = if network_name.trim().is_empty() { + instance.instance_id.clone() + } else { + network_name.clone() + }; + instances.insert( + label, + serde_json::Value::String(instance.instance_id.clone()), + ); + } + serde_json::Value::Object(instances).to_string() +} + +fn find_pro_runtime_instance<'a>( + state: &'a RuntimeAggregateState, + network_names_by_instance_id: &HashMap, + selector: Option<&str>, +) -> Option<(&'a RuntimeInstanceState, String)> { + let mut tracked_instances = state.instances.iter().filter_map(|instance| { + let network_name = network_names_by_instance_id.get(&instance.instance_id)?; + Some((instance, network_name)) + }); + if let Some(selector) = selector { + return tracked_instances + .filter(|(instance, _)| instance.running) + .find(|(instance, network_name)| { + selector == network_name.as_str() || runtime_instance_matches(instance, selector) + }) + .map(|(instance, network_name)| (instance, network_name.clone())); + } + tracked_instances + .find(|(instance, _)| instance.running) + .map(|(instance, network_name)| (instance, network_name.clone())) +} + +fn call_pro_json_rpc_inner( + state: &RuntimeAggregateState, + network_names_by_instance_id: &HashMap, + service_name: &str, + method_name: &str, + payload_json: &str, +) -> String { + let selector = selected_instance_from_payload(payload_json); + let Some((instance, network_name)) = + find_pro_runtime_instance(state, network_names_by_instance_id, selector.as_deref()) + else { + return "{}".to_string(); + }; + + let method = method_name.trim(); + let service = service_name.trim(); + let response = match (service, method) { + (_, "show_node_info") => serde_json::json!({ + "node_info": instance.my_node_info, + }), + (_, "list_route") => serde_json::json!({ + "routes": instance.routes, + }), + (_, "list_peer") => serde_json::json!({ + "my_info": instance.my_node_info, + "peer_infos": instance.peers, + }), + (_, "get_stats") => { + let mut rx_bytes = 0_i64; + let mut tx_bytes = 0_i64; + for peer in &instance.peers { + for conn in &peer.conns { + if let Some(stats) = &conn.stats { + rx_bytes = rx_bytes.saturating_add(stats.rx_bytes); + tx_bytes = tx_bytes.saturating_add(stats.tx_bytes); + } + } + } + serde_json::json!({ + "metrics": [ + { + "name": "traffic_bytes_self_rx", + "labels": { "network_name": network_name }, + "value": rx_bytes, + }, + { + "name": "traffic_bytes_self_tx", + "labels": { "network_name": network_name }, + "value": tx_bytes, + } + ] + }) + } + _ => serde_json::json!({}), + }; + response.to_string() +} + +fn pro_runtime_registry_snapshot() -> HashMap { + WEB_CLIENTS + .lock() + .ok() + .and_then(|clients| { + clients + .get(PRO_CONFIG_SERVER_CLIENT_ID) + .and_then(|managed| managed.hooks.network_names_by_instance_id.lock().ok()) + .map(|registry| { + registry + .iter() + .map(|(instance_id, network_name)| { + (instance_id.to_string(), network_name.clone()) + }) + .collect() + }) + }) + .unwrap_or_default() +} + +fn call_json_rpc_inner(service_name: &str, method_name: &str, payload_json: &str) -> String { + let state = collect_runtime_state_inner(); + let selector = selected_instance_from_payload(payload_json); + let Some(instance) = find_runtime_instance(&state, selector.as_deref()) else { + return "{}".to_string(); + }; + + let method = method_name.trim(); + let service = service_name.trim(); + let response = match (service, method) { + (_, "show_node_info") => serde_json::json!({ + "node_info": instance.my_node_info, + }), + (_, "list_route") => serde_json::json!({ + "routes": instance.routes, + }), + (_, "list_peer") => serde_json::json!({ + "my_info": instance.my_node_info, + "peer_infos": instance.peers, + }), + (_, "get_stats") => { + let mut rx_bytes = 0_i64; + let mut tx_bytes = 0_i64; + for peer in &instance.peers { + for conn in &peer.conns { + if let Some(stats) = &conn.stats { + rx_bytes = rx_bytes.saturating_add(stats.rx_bytes); + tx_bytes = tx_bytes.saturating_add(stats.tx_bytes); + } + } + } + let network_name = runtime_instance_label(instance); + serde_json::json!({ + "metrics": [ + { + "name": "traffic_bytes_self_rx", + "labels": { "network_name": network_name }, + "value": rx_bytes, + }, + { + "name": "traffic_bytes_self_tx", + "labels": { "network_name": network_name }, + "value": tx_bytes, + } + ] + }) + } + _ => serde_json::json!({}), + }; + response.to_string() +} + +fn resolve_instance_id_from_state( + state: &RuntimeAggregateState, + instance_name: &str, +) -> Option { + let instance = state.instances.iter().find(|instance| { + !is_pro_internal_instance(instance) && runtime_instance_matches(instance, instance_name) + })?; + Some(instance.instance_id.clone()) +} + +fn resolve_instance_id_inner(instance_name: &str) -> Option { + resolve_instance_id_from_state(&collect_runtime_state_inner(), instance_name) +} + pub(crate) fn build_default_network_config_json() -> Result { let config = NetworkConfig::new_from_config(TomlConfigLoader::default()) .map_err(|e| format!("default_network_config failed {}", e))?; @@ -440,6 +862,78 @@ pub fn stop_network_instance(config_ids: Vec) -> bool { exports::runtime_api::stop_network_instance(config_ids, stop_kernel) } +#[napi] +pub fn start_config_server_client( + url: String, + hostname: Option, + machine_id: Option, + secure_mode: Option, +) -> bool { + run_config_server_client(&url, hostname, machine_id, secure_mode.unwrap_or(false)) +} + +#[napi] +pub fn stop_config_server_client() -> bool { + stop_web_client(PRO_CONFIG_SERVER_CLIENT_ID) +} + +#[napi] +pub fn is_config_server_client_connected() -> bool { + pro_config_server_client_connected() +} + +#[napi] +pub fn stop_runtime() -> bool { + stop_runtime_inner() +} + +#[napi] +pub fn drain_config_server_events() -> String { + serde_json::to_string(&drain_config_server_events_inner()).unwrap_or_else(|_| "[]".to_string()) +} + +#[napi] +pub fn collect_runtime_state_json() -> String { + serde_json::to_string(&collect_runtime_state_inner()).unwrap_or_else(|_| "{}".to_string()) +} + +#[napi] +pub fn list_instances_json() -> String { + list_instances_json_inner(&collect_runtime_state_inner()) +} + +#[napi] +pub fn list_pro_instances_json() -> String { + let registry = pro_runtime_registry_snapshot(); + list_pro_instances_json_inner(&collect_runtime_state_inner(), ®istry) +} + +#[napi] +pub fn call_json_rpc(service_name: String, method_name: String, payload_json: String) -> String { + call_json_rpc_inner(&service_name, &method_name, &payload_json) +} + +#[napi] +pub fn call_pro_json_rpc( + service_name: String, + method_name: String, + payload_json: String, +) -> String { + let registry = pro_runtime_registry_snapshot(); + call_pro_json_rpc_inner( + &collect_runtime_state_inner(), + ®istry, + &service_name, + &method_name, + &payload_json, + ) +} + +#[napi] +pub fn resolve_instance_id(instance_name: String) -> Option { + resolve_instance_id_inner(&instance_name) +} + #[napi] pub fn easytier_version() -> String { EASYTIER_VERSION.to_string() @@ -512,6 +1006,92 @@ mod tests { .any(|field| field.name == "enabled") ); } + + fn pro_test_state() -> RuntimeAggregateState { + RuntimeAggregateState { + instances: vec![ + RuntimeInstanceState { + config_id: "0c4b33ba-4ed5-42d8-9095-21b786c66e94".to_string(), + instance_id: "0c4b33ba-4ed5-42d8-9095-21b786c66e94".to_string(), + display_name: "0c4b33ba-4ed5-42d8-9095-21b786c66e94".to_string(), + running: true, + tun_required: false, + tun_attached: false, + magic_dns_enabled: false, + need_exit_node: false, + error_message: None, + my_node_info: None, + events: vec![], + routes: vec![], + peers: vec![], + }, + RuntimeInstanceState { + config_id: "ec7b6a3c-aeae-4c0e-844e-f7ec2dbdc2ce".to_string(), + instance_id: "ec7b6a3c-aeae-4c0e-844e-f7ec2dbdc2ce".to_string(), + display_name: "ec7b6a3c-aeae-4c0e-844e-f7ec2dbdc2ce".to_string(), + running: true, + tun_required: false, + tun_attached: false, + magic_dns_enabled: false, + need_exit_node: false, + error_message: None, + my_node_info: None, + events: vec![], + routes: vec![], + peers: vec![], + }, + ], + tun: runtime::state::runtime_state::TunAggregateState { + active: false, + attached_instance_ids: vec![], + aggregated_routes: vec![], + dns_servers: vec![], + need_rebuild: false, + }, + running_instance_count: 2, + } + } + + fn pro_test_registry() -> HashMap { + HashMap::from([( + "0c4b33ba-4ed5-42d8-9095-21b786c66e94".to_string(), + "office-network".to_string(), + )]) + } + + #[test] + fn pro_instance_list_uses_registry_network_name_and_excludes_untracked_instances() { + assert_eq!( + list_pro_instances_json_inner(&pro_test_state(), &pro_test_registry()), + r#"{"office-network":"0c4b33ba-4ed5-42d8-9095-21b786c66e94"}"#, + ); + } + + #[test] + fn pro_json_rpc_selects_instance_by_network_name_and_labels_traffic() { + let response = call_pro_json_rpc_inner( + &pro_test_state(), + &pro_test_registry(), + "api.instance.StatsRpcService", + "get_stats", + r#"{"instance":{"instance_selector":{"name":"office-network"}}}"#, + ); + let response: serde_json::Value = serde_json::from_str(&response).unwrap(); + assert_eq!( + response["metrics"][0]["labels"]["network_name"], + "office-network", + ); + } + + #[test] + fn resolve_instance_id_does_not_fall_back_for_unknown_selector() { + let state = pro_test_state(); + assert_eq!( + resolve_instance_id_from_state(&state, "0c4b33ba-4ed5-42d8-9095-21b786c66e94"), + Some("0c4b33ba-4ed5-42d8-9095-21b786c66e94".to_string()), + ); + assert_eq!(resolve_instance_id_from_state(&state, "stale-name"), None); + } } pub(crate) fn collect_runtime_state_inner() -> RuntimeAggregateState { diff --git a/easytier-core/src/connectivity/manual/mod.rs b/easytier-core/src/connectivity/manual/mod.rs index cf57feec..20f84691 100644 --- a/easytier-core/src/connectivity/manual/mod.rs +++ b/easytier-core/src/connectivity/manual/mod.rs @@ -314,14 +314,19 @@ impl ManualConnectorOptions { } } -struct ManualConnectorData -where - H: ManualConnectorHost, -{ +#[derive(Default)] +struct ManualConnectorState { connectors: DashSet, reconnecting: DashSet, removed: DashSet, state_lock: Mutex<()>, +} + +struct ManualConnectorData +where + H: ManualConnectorHost, +{ + state: Arc, peer_manager: Weak, host: Arc, dns: Arc, @@ -345,6 +350,30 @@ struct ManualConnectorTask { handle: AbortOnDropHandle<()>, } +struct ReconnectReservation { + state: Arc, + url: Url, +} + +impl ReconnectReservation { + fn new(state: Arc, url: Url) -> Self { + Self { state, url } + } +} + +impl Drop for ReconnectReservation { + fn drop(&mut self) { + let _state_guard = self + .state + .state_lock + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if finish_reconnect_attempt(&self.state, &self.url) { + tracing::warn!(url = %self.url, "manual connector removed after reconnect"); + } + } +} + impl ManualConnectorManager where H: ManualConnectorHost, @@ -361,10 +390,7 @@ where events: Arc, ) -> Self { let data = Arc::new(ManualConnectorData { - connectors: DashSet::new(), - reconnecting: DashSet::new(), - removed: DashSet::new(), - state_lock: Mutex::new(()), + state: Arc::new(ManualConnectorState::default()), peer_manager: Arc::downgrade(&peer_manager), host, dns, @@ -404,50 +430,47 @@ where task.cancel.cancel(); let _ = task.handle.await; } - let _state_guard = self.data.state_lock.lock().unwrap(); - restore_interrupted_connectors( - &self.data.connectors, - &self.data.reconnecting, - &self.data.removed, - ); + let _state_guard = self.data.state.state_lock.lock().unwrap(); + restore_interrupted_connectors(&self.data.state); } pub fn add_connector(&self, url: Url) -> anyhow::Result<()> { validate_manual_url(&url)?; - let _state_guard = self.data.state_lock.lock().unwrap(); - self.data.removed.remove(&url); - if !self.data.reconnecting.contains(&url) { - self.data.connectors.insert(url); + let _state_guard = self.data.state.state_lock.lock().unwrap(); + self.data.state.removed.remove(&url); + if !self.data.state.reconnecting.contains(&url) { + self.data.state.connectors.insert(url); } Ok(()) } pub fn remove_connector(&self, url: &Url) -> bool { - let _state_guard = self.data.state_lock.lock().unwrap(); - if self.data.connectors.remove(url).is_some() { + let _state_guard = self.data.state.state_lock.lock().unwrap(); + if self.data.state.connectors.remove(url).is_some() { tracing::warn!(%url, "manual connector removed"); return true; } - if self.data.reconnecting.contains(url) { - self.data.removed.insert(url.clone()); + if self.data.state.reconnecting.contains(url) { + self.data.state.removed.insert(url.clone()); return true; } false } pub fn clear_connectors(&self) { - let _state_guard = self.data.state_lock.lock().unwrap(); - self.data.connectors.clear(); - for url in self.data.reconnecting.iter() { - self.data.removed.insert(url.key().clone()); + let _state_guard = self.data.state.state_lock.lock().unwrap(); + self.data.state.connectors.clear(); + for url in self.data.state.reconnecting.iter() { + self.data.state.removed.insert(url.key().clone()); } } pub fn list_connectors(&self) -> Vec { - let _state_guard = self.data.state_lock.lock().unwrap(); + let _state_guard = self.data.state.state_lock.lock().unwrap(); let peer_manager = self.data.peer_manager.upgrade(); let mut snapshots = self .data + .state .connectors .iter() .map(|entry| { @@ -465,15 +488,12 @@ where } }) .collect::>(); - snapshots.extend( - self.data - .reconnecting - .iter() - .map(|entry| ManualConnectorSnapshot { - url: entry.key().clone(), - status: ManualConnectorStatus::Connecting, - }), - ); + snapshots.extend(self.data.state.reconnecting.iter().map(|entry| { + ManualConnectorSnapshot { + url: entry.key().clone(), + status: ManualConnectorStatus::Connecting, + } + })); snapshots } @@ -487,8 +507,11 @@ where _ = interval.tick() => { for url in take_dead_connectors_for_reconnect(&data) { let task_data = data.clone(); + let reservation = + ReconnectReservation::new(data.state.clone(), url.clone()); reconnect_tasks.spawn(async move { let result = reconnect(task_data, url.clone()).await; + drop(reservation); (url, result) }); } @@ -500,13 +523,6 @@ where match result { Ok((url, reconnect_result)) => { tracing::warn!(?url, ?reconnect_result, "manual reconnect task done"); - let _state_guard = data.state_lock.lock().unwrap(); - data.reconnecting.remove(&url); - if data.removed.remove(&url).is_some() { - tracing::warn!(%url, "manual connector removed after reconnect"); - } else { - data.connectors.insert(url); - } } Err(error) => { tracing::error!(?error, "manual reconnect task failed"); @@ -519,25 +535,30 @@ where reconnect_tasks.abort_all(); while reconnect_tasks.join_next().await.is_some() {} - let _state_guard = data.state_lock.lock().unwrap(); - restore_interrupted_connectors(&data.connectors, &data.reconnecting, &data.removed); + let _state_guard = data.state.state_lock.lock().unwrap(); + restore_interrupted_connectors(&data.state); } } -fn restore_interrupted_connectors( - connectors: &DashSet, - reconnecting: &DashSet, - removed: &DashSet, -) { - let interrupted = reconnecting +fn finish_reconnect_attempt(state: &ManualConnectorState, url: &Url) -> bool { + if state.reconnecting.remove(url).is_none() { + return false; + } + if state.removed.remove(url).is_some() { + return true; + } + state.connectors.insert(url.clone()); + false +} + +fn restore_interrupted_connectors(state: &ManualConnectorState) { + let interrupted = state + .reconnecting .iter() .map(|entry| entry.key().clone()) .collect::>(); for url in interrupted { - reconnecting.remove(&url); - if removed.remove(&url).is_none() { - connectors.insert(url); - } + finish_reconnect_attempt(state, &url); } } @@ -545,12 +566,13 @@ fn take_dead_connectors_for_reconnect(data: &ManualConnectorData) -> BTree where H: ManualConnectorHost, { - let _state_guard = data.state_lock.lock().unwrap(); + let _state_guard = data.state.state_lock.lock().unwrap(); let Some(peer_manager) = data.peer_manager.upgrade() else { tracing::warn!("peer manager is gone, skip manual reconnect"); return BTreeSet::new(); }; let dead_connectors = data + .state .connectors .iter() .filter_map(|entry| { @@ -559,9 +581,9 @@ where }) .collect::>(); for url in &dead_connectors { - let removed = data.connectors.remove(url); + let removed = data.state.connectors.remove(url); debug_assert!(removed.is_some()); - let inserted = data.reconnecting.insert(url.clone()); + let inserted = data.state.reconnecting.insert(url.clone()); debug_assert!(inserted); } dead_connectors @@ -1040,6 +1062,22 @@ mod tests { use super::*; + fn reserve_pending_connector( + state: &Arc, + url: &Url, + ) -> ReconnectReservation { + let _state_guard = state.state_lock.lock().unwrap(); + assert!(state.connectors.remove(url).is_some()); + assert!(state.reconnecting.insert(url.clone())); + ReconnectReservation::new(state.clone(), url.clone()) + } + + fn assert_connector_is_pending(state: &ManualConnectorState, url: &Url) { + assert!(state.connectors.contains(url)); + assert!(!state.reconnecting.contains(url)); + assert!(!state.removed.contains(url)); + } + #[test] fn idn_normalization_covers_connector_schemes_and_url_round_trips() { let cases = [ @@ -1316,21 +1354,92 @@ mod tests { #[test] fn interrupted_reconnects_return_to_the_pending_set_unless_removed() { - let connectors = DashSet::new(); - let reconnecting = DashSet::new(); - let removed = DashSet::new(); + let state = ManualConnectorState::default(); let retained: Url = "tcp://127.0.0.1:11010".parse().unwrap(); let deleted: Url = "udp://127.0.0.1:11010".parse().unwrap(); - reconnecting.insert(retained.clone()); - reconnecting.insert(deleted.clone()); - removed.insert(deleted.clone()); + state.reconnecting.insert(retained.clone()); + state.reconnecting.insert(deleted.clone()); + state.removed.insert(deleted.clone()); - restore_interrupted_connectors(&connectors, &reconnecting, &removed); - restore_interrupted_connectors(&connectors, &reconnecting, &removed); + restore_interrupted_connectors(&state); + restore_interrupted_connectors(&state); - assert!(connectors.contains(&retained)); - assert!(!connectors.contains(&deleted)); - assert!(reconnecting.is_empty()); - assert!(removed.is_empty()); + assert!(state.connectors.contains(&retained)); + assert!(!state.connectors.contains(&deleted)); + assert!(state.reconnecting.is_empty()); + assert!(state.removed.is_empty()); + } + + #[test] + fn reconnect_reservation_restores_connector_on_normal_drop() { + let state = Arc::new(ManualConnectorState::default()); + let url: Url = "tcp://127.0.0.1:11010".parse().unwrap(); + state.connectors.insert(url.clone()); + + let reservation = reserve_pending_connector(&state, &url); + assert!(!state.connectors.contains(&url)); + assert!(state.reconnecting.contains(&url)); + + drop(reservation); + + assert_connector_is_pending(&state, &url); + } + + #[tokio::test] + async fn reconnect_reservation_restores_connector_after_attempt_panic() { + let state = Arc::new(ManualConnectorState::default()); + let url: Url = "tcp://127.0.0.1:11010".parse().unwrap(); + state.connectors.insert(url.clone()); + let reservation = reserve_pending_connector(&state, &url); + + let task = tokio::spawn(async move { + let _reservation = reservation; + panic!("reconnect attempt panic"); + }); + assert!(task.await.unwrap_err().is_panic()); + assert_connector_is_pending(&state, &url); + + let retry_reservation = reserve_pending_connector(&state, &url); + drop(retry_reservation); + assert_connector_is_pending(&state, &url); + } + + #[tokio::test] + async fn reconnect_reservation_restores_connector_after_attempt_abort() { + let state = Arc::new(ManualConnectorState::default()); + let url: Url = "tcp://127.0.0.1:11010".parse().unwrap(); + state.connectors.insert(url.clone()); + let reservation = reserve_pending_connector(&state, &url); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + + let task = tokio::spawn(async move { + let _reservation = reservation; + let _ = started_tx.send(()); + std::future::pending::<()>().await; + }); + started_rx.await.unwrap(); + task.abort(); + + assert!(task.await.unwrap_err().is_cancelled()); + assert_connector_is_pending(&state, &url); + + let retry_reservation = reserve_pending_connector(&state, &url); + drop(retry_reservation); + assert_connector_is_pending(&state, &url); + } + + #[test] + fn reconnect_reservation_does_not_restore_removed_connector() { + let state = Arc::new(ManualConnectorState::default()); + let url: Url = "tcp://127.0.0.1:11010".parse().unwrap(); + state.connectors.insert(url.clone()); + let reservation = reserve_pending_connector(&state, &url); + state.removed.insert(url.clone()); + + drop(reservation); + + assert!(!state.connectors.contains(&url)); + assert!(!state.reconnecting.contains(&url)); + assert!(!state.removed.contains(&url)); } }