feat(ohos): 接入 Pro 运行时并自动化 HAR 交付 (#2462)

This commit is contained in:
韩嘉乐
2026-08-01 16:35:14 +08:00
committed by GitHub
parent afbba5d928
commit 40c857748f
4 changed files with 982 additions and 269 deletions
+157 -198
View File
@@ -1,246 +1,205 @@
name: EasyTier OHOS name: ohos
on: on:
push: push:
branches: ["develop", "main", "releases/**"] branches: [develop, main, "releases/**", "ohos/**"]
tags: tags:
- 'v*' - "v*"
- '!*-pre' - "!*-pre"
pull_request: pull_request:
branches: ["develop", "main"] branches: [develop, main, "ohos/**"]
types: [opened, synchronize, reopened, ready_for_review] types: [opened, synchronize, reopened, ready_for_review]
workflow_dispatch: workflow_dispatch:
inputs:
publish:
description: Publish this non-main branch and dispatch downstream builds
required: false
default: false
type: boolean
concurrency: permissions:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} contents: read
cancel-in-progress: true pull-requests: read
env: env:
CARGO_TERM_COLOR: always CARGO_TERM_COLOR: always
defaults: defaults:
run: run:
# necessary for windows
shell: bash shell: bash
jobs: jobs:
cargo_fmt_check: ohos:
name: ohos
if: github.event_name != 'pull_request' || !github.event.pull_request.draft if: github.event_name != 'pull_request' || !github.event.pull_request.draft
runs-on: ubuntu-latest 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 uses: ./.github/actions/prepare-build
with: with:
target: aarch64-unknown-linux-ohos
gui: false gui: false
pnpm: false pnpm: false
token: ${{ secrets.GITHUB_TOKEN }} token: ${{ secrets.GITHUB_TOKEN }}
- uses: actions-rust-lang/setup-rust-toolchain@v1 - name: Set up HarmonyOS
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
uses: ErBWs/setup-ohos@v1 uses: ErBWs/setup-ohos@v1
- name: Download and Extract Custom SDK - name: Install ohrs
run: | uses: taiki-e/install-action@v2
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
with: with:
tool: ohrs tool: ohrs
- name: Create clang wrapper script - name: Build HAR
id: package
env:
DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
run: | run: |
sudo mkdir -p $OHOS_NDK_HOME/native/llvm set -euo pipefail
sudo tee $OHOS_NDK_HOME/native/llvm/aarch64-unknown-linux-ohos-clang.sh > /dev/null <<'EOF' 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 #!/bin/sh
exec $OHOS_NDK_HOME/native/llvm/bin/clang \ exec "$OHOS_NDK_HOME/native/llvm/bin/clang" \
-target aarch64-linux-ohos \ -target aarch64-linux-ohos \
--sysroot=$OHOS_NDK_HOME/native/sysroot \ --sysroot="$OHOS_NDK_HOME/native/sysroot" \
-D__MUSL__ \ -D__MUSL__ "$@"
"$@"
EOF 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 cd easytier-contrib/easytier-ohrs
working-directory: ./easytier-contrib/easytier-ohrs
run: |
sudo apt-get install -y llvm clang lldb lld
sudo apt-get install -y protobuf-compiler
source env.sh source env.sh
ohrs doctor
ohrs build --release --arch aarch ohrs build --release --arch aarch
ohrs artifact ohrs artifact
mv package.har easytier-ohrs.har mv package.har "$package_name.har"
- name: Build Release Package - name: Upload HAR
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
uses: actions/upload-artifact@v5 uses: actions/upload-artifact@v5
with: with:
name: easytier-ohos name: ${{ steps.package.outputs.name }}
path: | path: easytier-contrib/easytier-ohrs/${{ steps.package.outputs.name }}.har
./easytier-contrib/easytier-ohrs/easytier-ohrs.har
retention-days: 5 retention-days: 5
if-no-files-found: error if-no-files-found: error
- name: Publish To Center Ohpm - name: Publish and dispatch
working-directory: ./easytier-contrib/easytier-ohrs 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: env:
OHPM_PRIVATE_KEY: ${{ secrets.OHPM_PRIVATE_KEY }} GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
OHPM_KEY_PASSPHRASE: ${{ secrets.OHPM_KEY_PASSPHRASE }} CODEARTS_PRIVATE_OHPM: ${{ secrets.CODEARTS_PRIVATE_OHPM }}
if: ${{ env.OHPM_PUBLISH_CODE != '' && github.event_name == 'push' }} DOWNSTREAM_DISPATCH_TOKEN: ${{ secrets.DOWNSTREAM_DISPATCH_TOKEN }}
run: | run: |
ohpm config set publish_id "$OHPM_PUBLISH_CODE" set -euo pipefail
ohpm config set publish_registry https://ohpm.openharmony.cn/ohpm if [ "$GITHUB_EVENT_NAME" = push ]; then
TMP_DIR=$(mktemp -d) pull_requests=$(gh api \
PRIVATE_KEY_FILE="$TMP_DIR/private_key" -H "Accept: application/vnd.github+json" \
printf '%s' "$OHPM_PRIVATE_KEY" > "$PRIVATE_KEY_FILE" "/repos/$GITHUB_REPOSITORY/commits/$GITHUB_SHA/pulls")
chmod 600 "$PRIVATE_KEY_FILE" if ! jq -e \
ohpm config set key_path $PRIVATE_KEY_FILE --arg repository "$GITHUB_REPOSITORY" \
unzip ohpm_crypto.zip -d /home/runner/work/ --arg branch "$GITHUB_REF_NAME" \
ohpm config set crypto_path /home/runner/work/ohpm_crypto --arg sha "$GITHUB_SHA" \
chmod 755 /home/runner/work/ohpm_crypto/* 'any(.[];
PASSPHRASE="$(printf '%s' "$OHPM_KEY_PASSPHRASE" | tr -d '\r\n')" .merged_at != null and
ohpm config set key_passphrase "$PASSPHRASE" .base.repo.full_name == $repository and
ohpm publish easytier-ohrs.har .base.ref == $branch and
.merge_commit_sha == $sha)' \
- name: Publish To Private Ohpm <<< "$pull_requests" >/dev/null; then
working-directory: ./easytier-contrib/easytier-ohrs echo "Direct push: HAR built without publishing."
if: ${{ env.OHPM_PUBLISH_CODE != '' && github.event_name == 'push' }} exit 0
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
fi fi
curl --header "Content-Type: application/json" --request POST --data "{}" ${{ secrets.CODEARTS_WEBHOOKS }} fi
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
+65
View File
@@ -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
<core-version>-<branch-id>-<commits-since-tag>-<run-number>-<run-attempt>-g<short-sha>
```
`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/<package_name>/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.
+581 -1
View File
@@ -68,7 +68,7 @@ use kernel_bridge::{
stop_local_socket_server as stop_local_socket_server_inner, stop_local_socket_server as stop_local_socket_server_inner,
}; };
use napi_derive_ohos::napi; 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::collections::{HashMap, HashSet};
use std::format; use std::format;
use std::sync::{Arc, Mutex}; use std::sync::{Arc, Mutex};
@@ -89,10 +89,13 @@ pub(crate) static INSTANCE_MANAGER: once_cell::sync::Lazy<Arc<NativeInstanceMana
}); });
static WEB_CLIENTS: once_cell::sync::Lazy<Mutex<HashMap<String, ManagedWebClient>>> = static WEB_CLIENTS: once_cell::sync::Lazy<Mutex<HashMap<String, ManagedWebClient>>> =
once_cell::sync::Lazy::new(|| Mutex::new(HashMap::new())); once_cell::sync::Lazy::new(|| Mutex::new(HashMap::new()));
const PRO_CONFIG_SERVER_CLIENT_ID: &str = "__easytier_pro_config_server_client__";
#[derive(Default)] #[derive(Default)]
struct TrackedWebClientHooks { struct TrackedWebClientHooks {
instance_ids: Mutex<HashSet<Uuid>>, instance_ids: Mutex<HashSet<Uuid>>,
network_names_by_instance_id: Mutex<HashMap<Uuid, String>>,
events: Mutex<Vec<serde_json::Value>>,
} }
struct ManagedWebClient { struct ManagedWebClient {
@@ -100,6 +103,13 @@ struct ManagedWebClient {
hooks: Arc<TrackedWebClientHooks>, hooks: Arc<TrackedWebClientHooks>,
} }
fn network_name_for_instance(id: &Uuid) -> Option<String> {
INSTANCE_MANAGER
.config(*id)
.map(|config| config.get_network_identity().network_name)
.filter(|name| !name.trim().is_empty())
}
#[async_trait::async_trait] #[async_trait::async_trait]
impl WebClientHooks for TrackedWebClientHooks { impl WebClientHooks for TrackedWebClientHooks {
async fn post_run_network_instance(&self, id: &Uuid) -> Result<(), String> { async fn post_run_network_instance(&self, id: &Uuid) -> Result<(), String> {
@@ -107,13 +117,43 @@ impl WebClientHooks for TrackedWebClientHooks {
.lock() .lock()
.map_err(|err| err.to_string())? .map_err(|err| err.to_string())?
.insert(*id); .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(()) Ok(())
} }
async fn post_remove_network_instances(&self, ids: &[Uuid]) -> Result<(), String> { 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 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 { for id in ids {
guard.remove(id); 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(()) 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<String>,
machine_id: Option<String>,
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<serde_json::Value> {
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<String> {
let value = serde_json::from_str::<serde_json::Value>(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, String>,
) -> 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<String, String>,
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<String, String>,
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<String, String> {
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<String> {
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<String> {
resolve_instance_id_from_state(&collect_runtime_state_inner(), instance_name)
}
pub(crate) fn build_default_network_config_json() -> Result<String, String> { pub(crate) fn build_default_network_config_json() -> Result<String, String> {
let config = NetworkConfig::new_from_config(TomlConfigLoader::default()) let config = NetworkConfig::new_from_config(TomlConfigLoader::default())
.map_err(|e| format!("default_network_config failed {}", e))?; .map_err(|e| format!("default_network_config failed {}", e))?;
@@ -440,6 +862,78 @@ pub fn stop_network_instance(config_ids: Vec<String>) -> bool {
exports::runtime_api::stop_network_instance(config_ids, stop_kernel) exports::runtime_api::stop_network_instance(config_ids, stop_kernel)
} }
#[napi]
pub fn start_config_server_client(
url: String,
hostname: Option<String>,
machine_id: Option<String>,
secure_mode: Option<bool>,
) -> 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(), &registry)
}
#[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(),
&registry,
&service_name,
&method_name,
&payload_json,
)
}
#[napi]
pub fn resolve_instance_id(instance_name: String) -> Option<String> {
resolve_instance_id_inner(&instance_name)
}
#[napi] #[napi]
pub fn easytier_version() -> String { pub fn easytier_version() -> String {
EASYTIER_VERSION.to_string() EASYTIER_VERSION.to_string()
@@ -512,6 +1006,92 @@ mod tests {
.any(|field| field.name == "enabled") .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<String, String> {
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 { pub(crate) fn collect_runtime_state_inner() -> RuntimeAggregateState {
+177 -68
View File
@@ -314,14 +314,19 @@ impl ManualConnectorOptions {
} }
} }
struct ManualConnectorData<H> #[derive(Default)]
where struct ManualConnectorState {
H: ManualConnectorHost,
{
connectors: DashSet<Url>, connectors: DashSet<Url>,
reconnecting: DashSet<Url>, reconnecting: DashSet<Url>,
removed: DashSet<Url>, removed: DashSet<Url>,
state_lock: Mutex<()>, state_lock: Mutex<()>,
}
struct ManualConnectorData<H>
where
H: ManualConnectorHost,
{
state: Arc<ManualConnectorState>,
peer_manager: Weak<PeerManagerCore>, peer_manager: Weak<PeerManagerCore>,
host: Arc<H>, host: Arc<H>,
dns: Arc<dyn DnsResolver>, dns: Arc<dyn DnsResolver>,
@@ -345,6 +350,30 @@ struct ManualConnectorTask {
handle: AbortOnDropHandle<()>, handle: AbortOnDropHandle<()>,
} }
struct ReconnectReservation {
state: Arc<ManualConnectorState>,
url: Url,
}
impl ReconnectReservation {
fn new(state: Arc<ManualConnectorState>, 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<H> ManualConnectorManager<H> impl<H> ManualConnectorManager<H>
where where
H: ManualConnectorHost, H: ManualConnectorHost,
@@ -361,10 +390,7 @@ where
events: Arc<dyn CoreEventSink>, events: Arc<dyn CoreEventSink>,
) -> Self { ) -> Self {
let data = Arc::new(ManualConnectorData { let data = Arc::new(ManualConnectorData {
connectors: DashSet::new(), state: Arc::new(ManualConnectorState::default()),
reconnecting: DashSet::new(),
removed: DashSet::new(),
state_lock: Mutex::new(()),
peer_manager: Arc::downgrade(&peer_manager), peer_manager: Arc::downgrade(&peer_manager),
host, host,
dns, dns,
@@ -404,50 +430,47 @@ where
task.cancel.cancel(); task.cancel.cancel();
let _ = task.handle.await; let _ = task.handle.await;
} }
let _state_guard = self.data.state_lock.lock().unwrap(); let _state_guard = self.data.state.state_lock.lock().unwrap();
restore_interrupted_connectors( restore_interrupted_connectors(&self.data.state);
&self.data.connectors,
&self.data.reconnecting,
&self.data.removed,
);
} }
pub fn add_connector(&self, url: Url) -> anyhow::Result<()> { pub fn add_connector(&self, url: Url) -> anyhow::Result<()> {
validate_manual_url(&url)?; validate_manual_url(&url)?;
let _state_guard = self.data.state_lock.lock().unwrap(); let _state_guard = self.data.state.state_lock.lock().unwrap();
self.data.removed.remove(&url); self.data.state.removed.remove(&url);
if !self.data.reconnecting.contains(&url) { if !self.data.state.reconnecting.contains(&url) {
self.data.connectors.insert(url); self.data.state.connectors.insert(url);
} }
Ok(()) Ok(())
} }
pub fn remove_connector(&self, url: &Url) -> bool { pub fn remove_connector(&self, url: &Url) -> bool {
let _state_guard = self.data.state_lock.lock().unwrap(); let _state_guard = self.data.state.state_lock.lock().unwrap();
if self.data.connectors.remove(url).is_some() { if self.data.state.connectors.remove(url).is_some() {
tracing::warn!(%url, "manual connector removed"); tracing::warn!(%url, "manual connector removed");
return true; return true;
} }
if self.data.reconnecting.contains(url) { if self.data.state.reconnecting.contains(url) {
self.data.removed.insert(url.clone()); self.data.state.removed.insert(url.clone());
return true; return true;
} }
false false
} }
pub fn clear_connectors(&self) { pub fn clear_connectors(&self) {
let _state_guard = self.data.state_lock.lock().unwrap(); let _state_guard = self.data.state.state_lock.lock().unwrap();
self.data.connectors.clear(); self.data.state.connectors.clear();
for url in self.data.reconnecting.iter() { for url in self.data.state.reconnecting.iter() {
self.data.removed.insert(url.key().clone()); self.data.state.removed.insert(url.key().clone());
} }
} }
pub fn list_connectors(&self) -> Vec<ManualConnectorSnapshot> { pub fn list_connectors(&self) -> Vec<ManualConnectorSnapshot> {
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 peer_manager = self.data.peer_manager.upgrade();
let mut snapshots = self let mut snapshots = self
.data .data
.state
.connectors .connectors
.iter() .iter()
.map(|entry| { .map(|entry| {
@@ -465,15 +488,12 @@ where
} }
}) })
.collect::<Vec<_>>(); .collect::<Vec<_>>();
snapshots.extend( snapshots.extend(self.data.state.reconnecting.iter().map(|entry| {
self.data ManualConnectorSnapshot {
.reconnecting
.iter()
.map(|entry| ManualConnectorSnapshot {
url: entry.key().clone(), url: entry.key().clone(),
status: ManualConnectorStatus::Connecting, status: ManualConnectorStatus::Connecting,
}), }
); }));
snapshots snapshots
} }
@@ -487,8 +507,11 @@ where
_ = interval.tick() => { _ = interval.tick() => {
for url in take_dead_connectors_for_reconnect(&data) { for url in take_dead_connectors_for_reconnect(&data) {
let task_data = data.clone(); let task_data = data.clone();
let reservation =
ReconnectReservation::new(data.state.clone(), url.clone());
reconnect_tasks.spawn(async move { reconnect_tasks.spawn(async move {
let result = reconnect(task_data, url.clone()).await; let result = reconnect(task_data, url.clone()).await;
drop(reservation);
(url, result) (url, result)
}); });
} }
@@ -500,13 +523,6 @@ where
match result { match result {
Ok((url, reconnect_result)) => { Ok((url, reconnect_result)) => {
tracing::warn!(?url, ?reconnect_result, "manual reconnect task done"); 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) => { Err(error) => {
tracing::error!(?error, "manual reconnect task failed"); tracing::error!(?error, "manual reconnect task failed");
@@ -519,25 +535,30 @@ where
reconnect_tasks.abort_all(); reconnect_tasks.abort_all();
while reconnect_tasks.join_next().await.is_some() {} while reconnect_tasks.join_next().await.is_some() {}
let _state_guard = data.state_lock.lock().unwrap(); let _state_guard = data.state.state_lock.lock().unwrap();
restore_interrupted_connectors(&data.connectors, &data.reconnecting, &data.removed); restore_interrupted_connectors(&data.state);
} }
} }
fn restore_interrupted_connectors( fn finish_reconnect_attempt(state: &ManualConnectorState, url: &Url) -> bool {
connectors: &DashSet<Url>, if state.reconnecting.remove(url).is_none() {
reconnecting: &DashSet<Url>, return false;
removed: &DashSet<Url>, }
) { if state.removed.remove(url).is_some() {
let interrupted = reconnecting return true;
}
state.connectors.insert(url.clone());
false
}
fn restore_interrupted_connectors(state: &ManualConnectorState) {
let interrupted = state
.reconnecting
.iter() .iter()
.map(|entry| entry.key().clone()) .map(|entry| entry.key().clone())
.collect::<Vec<_>>(); .collect::<Vec<_>>();
for url in interrupted { for url in interrupted {
reconnecting.remove(&url); finish_reconnect_attempt(state, &url);
if removed.remove(&url).is_none() {
connectors.insert(url);
}
} }
} }
@@ -545,12 +566,13 @@ fn take_dead_connectors_for_reconnect<H>(data: &ManualConnectorData<H>) -> BTree
where where
H: ManualConnectorHost, 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 { let Some(peer_manager) = data.peer_manager.upgrade() else {
tracing::warn!("peer manager is gone, skip manual reconnect"); tracing::warn!("peer manager is gone, skip manual reconnect");
return BTreeSet::new(); return BTreeSet::new();
}; };
let dead_connectors = data let dead_connectors = data
.state
.connectors .connectors
.iter() .iter()
.filter_map(|entry| { .filter_map(|entry| {
@@ -559,9 +581,9 @@ where
}) })
.collect::<BTreeSet<_>>(); .collect::<BTreeSet<_>>();
for url in &dead_connectors { for url in &dead_connectors {
let removed = data.connectors.remove(url); let removed = data.state.connectors.remove(url);
debug_assert!(removed.is_some()); debug_assert!(removed.is_some());
let inserted = data.reconnecting.insert(url.clone()); let inserted = data.state.reconnecting.insert(url.clone());
debug_assert!(inserted); debug_assert!(inserted);
} }
dead_connectors dead_connectors
@@ -1040,6 +1062,22 @@ mod tests {
use super::*; use super::*;
fn reserve_pending_connector(
state: &Arc<ManualConnectorState>,
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] #[test]
fn idn_normalization_covers_connector_schemes_and_url_round_trips() { fn idn_normalization_covers_connector_schemes_and_url_round_trips() {
let cases = [ let cases = [
@@ -1316,21 +1354,92 @@ mod tests {
#[test] #[test]
fn interrupted_reconnects_return_to_the_pending_set_unless_removed() { fn interrupted_reconnects_return_to_the_pending_set_unless_removed() {
let connectors = DashSet::new(); let state = ManualConnectorState::default();
let reconnecting = DashSet::new();
let removed = DashSet::new();
let retained: Url = "tcp://127.0.0.1:11010".parse().unwrap(); let retained: Url = "tcp://127.0.0.1:11010".parse().unwrap();
let deleted: Url = "udp://127.0.0.1:11010".parse().unwrap(); let deleted: Url = "udp://127.0.0.1:11010".parse().unwrap();
reconnecting.insert(retained.clone()); state.reconnecting.insert(retained.clone());
reconnecting.insert(deleted.clone()); state.reconnecting.insert(deleted.clone());
removed.insert(deleted.clone()); state.removed.insert(deleted.clone());
restore_interrupted_connectors(&connectors, &reconnecting, &removed); restore_interrupted_connectors(&state);
restore_interrupted_connectors(&connectors, &reconnecting, &removed); restore_interrupted_connectors(&state);
assert!(connectors.contains(&retained)); assert!(state.connectors.contains(&retained));
assert!(!connectors.contains(&deleted)); assert!(!state.connectors.contains(&deleted));
assert!(reconnecting.is_empty()); assert!(state.reconnecting.is_empty());
assert!(removed.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));
} }
} }