mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-04 21:08:11 +03:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b92aa0358a | ||
|
|
316bcd6343 |
@@ -1 +0,0 @@
|
|||||||
powershell.exe -ExecutionPolicy Bypass -File ".\xray_no_window.ps1"
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
Start-Process -FilePath ".\xray.exe" -ArgumentList "-config .\config.json" -WindowStyle Hidden
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
CreateObject("Wscript.Shell").Run "xray.exe -config config.json",0
|
|
||||||
@@ -37,7 +37,6 @@ RUN echo '{}' >/tmp/usr/local/etc/xray/08_fakedns.json
|
|||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/09_metrics.json
|
RUN echo '{}' >/tmp/usr/local/etc/xray/09_metrics.json
|
||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/10_observatory.json
|
RUN echo '{}' >/tmp/usr/local/etc/xray/10_observatory.json
|
||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/11_geodata.json
|
RUN echo '{}' >/tmp/usr/local/etc/xray/11_geodata.json
|
||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/12_env.json
|
|
||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/99_version.json
|
RUN echo '{}' >/tmp/usr/local/etc/xray/99_version.json
|
||||||
|
|
||||||
# Create log files
|
# Create log files
|
||||||
|
|||||||
@@ -37,7 +37,6 @@ RUN echo '{}' >/tmp/usr/local/etc/xray/08_fakedns.json
|
|||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/09_metrics.json
|
RUN echo '{}' >/tmp/usr/local/etc/xray/09_metrics.json
|
||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/10_observatory.json
|
RUN echo '{}' >/tmp/usr/local/etc/xray/10_observatory.json
|
||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/11_geodata.json
|
RUN echo '{}' >/tmp/usr/local/etc/xray/11_geodata.json
|
||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/12_env.json
|
|
||||||
RUN echo '{}' >/tmp/usr/local/etc/xray/99_version.json
|
RUN echo '{}' >/tmp/usr/local/etc/xray/99_version.json
|
||||||
|
|
||||||
# Create log files
|
# Create log files
|
||||||
|
|||||||
@@ -64,16 +64,8 @@ jobs:
|
|||||||
echo "Latest: '$LATEST'."
|
echo "Latest: '$LATEST'."
|
||||||
echo "LATEST=$LATEST" >>${GITHUB_ENV}
|
echo "LATEST=$LATEST" >>${GITHUB_ENV}
|
||||||
|
|
||||||
NEWEST=false
|
|
||||||
if [[ "${{ github.event_name }}" == "release" ]]; then
|
|
||||||
NEWEST=true
|
|
||||||
fi
|
|
||||||
|
|
||||||
echo "Newest: '$NEWEST'."
|
|
||||||
echo "NEWEST=$NEWEST" >>${GITHUB_ENV}
|
|
||||||
|
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
uses: actions/checkout@v7
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Set up QEMU
|
- name: Set up QEMU
|
||||||
uses: docker/setup-qemu-action@v4
|
uses: docker/setup-qemu-action@v4
|
||||||
@@ -82,7 +74,7 @@ jobs:
|
|||||||
uses: docker/setup-buildx-action@v4
|
uses: docker/setup-buildx-action@v4
|
||||||
|
|
||||||
- name: Login to GitHub Container Registry
|
- name: Login to GitHub Container Registry
|
||||||
uses: docker/login-action@v4.6.0
|
uses: docker/login-action@v4
|
||||||
with:
|
with:
|
||||||
registry: ghcr.io
|
registry: ghcr.io
|
||||||
username: ${{ github.repository_owner }}
|
username: ${{ github.repository_owner }}
|
||||||
@@ -132,13 +124,6 @@ jobs:
|
|||||||
${{ env.FULL_IMAGE_NAME }}:${{ env.IMAGE_TAG }}
|
${{ env.FULL_IMAGE_NAME }}:${{ env.IMAGE_TAG }}
|
||||||
fi
|
fi
|
||||||
|
|
||||||
if [[ "${{ env.NEWEST }}" == "true" ]]; then
|
|
||||||
echo "Adding 'pre-release' tag to manifest: '${{ env.FULL_IMAGE_NAME }}:pre-release'."
|
|
||||||
docker buildx imagetools create \
|
|
||||||
--tag ${{ env.FULL_IMAGE_NAME }}:pre-release \
|
|
||||||
${{ env.FULL_IMAGE_NAME }}:${{ env.IMAGE_TAG }}
|
|
||||||
fi
|
|
||||||
|
|
||||||
- name: Inspect image
|
- name: Inspect image
|
||||||
run: |
|
run: |
|
||||||
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:${{ env.IMAGE_TAG }}
|
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:${{ env.IMAGE_TAG }}
|
||||||
@@ -146,7 +131,3 @@ jobs:
|
|||||||
if [[ "${{ env.LATEST }}" == "true" ]]; then
|
if [[ "${{ env.LATEST }}" == "true" ]]; then
|
||||||
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:latest
|
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:latest
|
||||||
fi
|
fi
|
||||||
|
|
||||||
if [[ "${{ env.NEWEST }}" == "true" ]]; then
|
|
||||||
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:pre-release
|
|
||||||
fi
|
|
||||||
|
|||||||
@@ -11,16 +11,15 @@ on:
|
|||||||
jobs:
|
jobs:
|
||||||
check-assets:
|
check-assets:
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
|
||||||
steps:
|
steps:
|
||||||
- name: Restore Geodat Cache
|
- name: Restore Geodat Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|
||||||
- name: Restore Wintun Cache
|
- name: Restore Wintun Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -76,14 +75,13 @@ jobs:
|
|||||||
fail-fast: false
|
fail-fast: false
|
||||||
|
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
|
||||||
env:
|
env:
|
||||||
GOOS: ${{ matrix.goos}}
|
GOOS: ${{ matrix.goos}}
|
||||||
GOARCH: ${{ matrix.goarch }}
|
GOARCH: ${{ matrix.goarch }}
|
||||||
CGO_ENABLED: 0
|
CGO_ENABLED: 0
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout codebase
|
- name: Checkout codebase
|
||||||
uses: actions/checkout@v7
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Show workflow information
|
- name: Show workflow information
|
||||||
run: |
|
run: |
|
||||||
@@ -92,7 +90,7 @@ jobs:
|
|||||||
echo "ASSET_NAME=$_NAME" >> $GITHUB_ENV
|
echo "ASSET_NAME=$_NAME" >> $GITHUB_ENV
|
||||||
|
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v7
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
check-latest: true
|
check-latest: true
|
||||||
@@ -119,13 +117,13 @@ jobs:
|
|||||||
# go build -o build_assets/wxray.exe -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-H windowsgui -X github.com/xtls/xray-core/core.build=${COMMID} -s -w -buildid=" -v ./main
|
# go build -o build_assets/wxray.exe -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-H windowsgui -X github.com/xtls/xray-core/core.build=${COMMID} -s -w -buildid=" -v ./main
|
||||||
|
|
||||||
- name: Restore Geodat Cache
|
- name: Restore Geodat Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|
||||||
- name: Restore Wintun Cache
|
- name: Restore Wintun Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -134,17 +132,15 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
mv -f resources/geo* build_assets/
|
mv -f resources/geo* build_assets/
|
||||||
if [[ ${GOOS} == 'windows' ]]; then
|
if [[ ${GOOS} == 'windows' ]]; then
|
||||||
cp .github/build/windows/* build_assets/
|
echo 'CreateObject("Wscript.Shell").Run "xray.exe -config config.json",0' > build_assets/xray_no_window.vbs
|
||||||
fi
|
echo 'Start-Process -FilePath ".\xray.exe" -ArgumentList "-config .\config.json" -WindowStyle Hidden' > build_assets/xray_no_window.ps1
|
||||||
if [[ ${GOOS} == 'windows' ]]; then
|
|
||||||
echo 'Adding Wintun into packages'
|
|
||||||
if [[ ${GOARCH} == 'amd64' ]]; then
|
if [[ ${GOARCH} == 'amd64' ]]; then
|
||||||
mv resources/wintun/bin/amd64/wintun.dll build_assets/
|
mv resources/wintun/bin/amd64/wintun.dll build_assets/
|
||||||
fi
|
fi
|
||||||
if [[ ${GOARCH} == '386' ]]; then
|
if [[ ${GOARCH} == '386' ]]; then
|
||||||
mv resources/wintun/bin/x86/wintun.dll build_assets/
|
mv resources/wintun/bin/x86/wintun.dll build_assets/
|
||||||
fi
|
fi
|
||||||
mv resources/wintun/LICENSE.txt build_assets/LICENSE-Wintun
|
mv resources/wintun/LICENSE.txt build_assets/LICENSE-wintun.txt
|
||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Copy README.md & LICENSE
|
- name: Copy README.md & LICENSE
|
||||||
|
|||||||
@@ -11,16 +11,15 @@ on:
|
|||||||
jobs:
|
jobs:
|
||||||
check-assets:
|
check-assets:
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
|
||||||
steps:
|
steps:
|
||||||
- name: Restore Geodat Cache
|
- name: Restore Geodat Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|
||||||
- name: Restore Wintun Cache
|
- name: Restore Wintun Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -162,7 +161,6 @@ jobs:
|
|||||||
fail-fast: false
|
fail-fast: false
|
||||||
|
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
|
||||||
env:
|
env:
|
||||||
GOOS: ${{ matrix.goos }}
|
GOOS: ${{ matrix.goos }}
|
||||||
GOARCH: ${{ matrix.goarch }}
|
GOARCH: ${{ matrix.goarch }}
|
||||||
@@ -170,7 +168,7 @@ jobs:
|
|||||||
CGO_ENABLED: 0
|
CGO_ENABLED: 0
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout codebase
|
- name: Checkout codebase
|
||||||
uses: actions/checkout@v7
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Set up NDK
|
- name: Set up NDK
|
||||||
if: matrix.goos == 'android'
|
if: matrix.goos == 'android'
|
||||||
@@ -193,7 +191,7 @@ jobs:
|
|||||||
echo "ASSET_NAME=$_NAME" >> $GITHUB_ENV
|
echo "ASSET_NAME=$_NAME" >> $GITHUB_ENV
|
||||||
|
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v7
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
check-latest: true
|
check-latest: true
|
||||||
@@ -210,9 +208,6 @@ jobs:
|
|||||||
go build -o build_assets/xray.exe -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-X github.com/xtls/xray-core/core.build=${COMMID} -s -w -buildid=" -v ./main
|
go build -o build_assets/xray.exe -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-X github.com/xtls/xray-core/core.build=${COMMID} -s -w -buildid=" -v ./main
|
||||||
# The line below is for without running conhost.exe version. Commented for not being used. Provided for reference.
|
# The line below is for without running conhost.exe version. Commented for not being used. Provided for reference.
|
||||||
# go build -o build_assets/wxray.exe -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-H windowsgui -X github.com/xtls/xray-core/core.build=${COMMID} -s -w -buildid=" -v ./main
|
# go build -o build_assets/wxray.exe -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-H windowsgui -X github.com/xtls/xray-core/core.build=${COMMID} -s -w -buildid=" -v ./main
|
||||||
elif [[ ${GOOS} == 'android' ]]; then
|
|
||||||
echo 'Building Xray for Android...'
|
|
||||||
go build -o build_assets/xray -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-X github.com/xtls/xray-core/core.build=${COMMID} -s -w -buildid= -checklinkname=0" -v ./main
|
|
||||||
else
|
else
|
||||||
echo 'Building Xray...'
|
echo 'Building Xray...'
|
||||||
if [[ ${GOARCH} == 'mips' || ${GOARCH} == 'mipsle' ]]; then
|
if [[ ${GOARCH} == 'mips' || ${GOARCH} == 'mipsle' ]]; then
|
||||||
@@ -225,14 +220,14 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Restore Geodat Cache
|
- name: Restore Geodat Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|
||||||
- name: Restore Wintun Cache
|
- name: Restore Wintun Cache
|
||||||
if: matrix.goos == 'windows'
|
if: matrix.goos == 'windows'
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -241,10 +236,8 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
mv -f resources/geo* build_assets/
|
mv -f resources/geo* build_assets/
|
||||||
if [[ ${GOOS} == 'windows' ]]; then
|
if [[ ${GOOS} == 'windows' ]]; then
|
||||||
cp .github/build/windows/* build_assets/
|
echo 'CreateObject("Wscript.Shell").Run "xray.exe -config config.json",0' > build_assets/xray_no_window.vbs
|
||||||
fi
|
echo 'Start-Process -FilePath ".\xray.exe" -ArgumentList "-config .\config.json" -WindowStyle Hidden' > build_assets/xray_no_window.ps1
|
||||||
if [[ ${GOOS} == 'windows' ]]; then
|
|
||||||
echo 'Adding Wintun into packages'
|
|
||||||
if [[ ${GOARCH} == 'amd64' ]]; then
|
if [[ ${GOARCH} == 'amd64' ]]; then
|
||||||
mv resources/wintun/bin/amd64/wintun.dll build_assets/
|
mv resources/wintun/bin/amd64/wintun.dll build_assets/
|
||||||
fi
|
fi
|
||||||
@@ -254,7 +247,7 @@ jobs:
|
|||||||
if [[ ${GOARCH} == 'arm64' ]]; then
|
if [[ ${GOARCH} == 'arm64' ]]; then
|
||||||
mv resources/wintun/bin/arm64/wintun.dll build_assets/
|
mv resources/wintun/bin/arm64/wintun.dll build_assets/
|
||||||
fi
|
fi
|
||||||
mv resources/wintun/LICENSE.txt build_assets/LICENSE-Wintun
|
mv resources/wintun/LICENSE.txt build_assets/LICENSE-wintun.txt
|
||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Copy README.md & LICENSE
|
- name: Copy README.md & LICENSE
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- name: Restore Geodat Cache
|
- name: Restore Geodat Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
@@ -59,7 +59,7 @@ jobs:
|
|||||||
done
|
done
|
||||||
|
|
||||||
- name: Save Geodat Cache
|
- name: Save Geodat Cache
|
||||||
uses: actions/cache/save@v6
|
uses: actions/cache/save@v5
|
||||||
if: ${{ steps.update.outputs.unhit }}
|
if: ${{ steps.update.outputs.unhit }}
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
@@ -68,12 +68,9 @@ jobs:
|
|||||||
wintun:
|
wintun:
|
||||||
if: github.event.schedule == '30 22 * * *' || github.event_name == 'push' || github.event_name == 'pull_request' || github.event_name == 'workflow_dispatch'
|
if: github.event.schedule == '30 22 * * *' || github.event_name == 'push' || github.event_name == 'pull_request' || github.event_name == 'workflow_dispatch'
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
env:
|
|
||||||
ASSETVER: 0.14.1
|
|
||||||
ASSETHASH: 07c256185d6ee3652e09fa55c0b673e2624b565e02c4b9091c79ca7d2f24ef51
|
|
||||||
steps:
|
steps:
|
||||||
- name: Restore Wintun Cache
|
- name: Restore Wintun Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -99,6 +96,7 @@ jobs:
|
|||||||
echo -e "Checking if wintun.dll for ${ARCHITECTURE} exists..."
|
echo -e "Checking if wintun.dll for ${ARCHITECTURE} exists..."
|
||||||
if [ -s "./resources/wintun/bin/${ARCHITECTURE}/wintun.dll" ]; then
|
if [ -s "./resources/wintun/bin/${ARCHITECTURE}/wintun.dll" ]; then
|
||||||
echo -e "wintun.dll for ${ARCHITECTURE} exists"
|
echo -e "wintun.dll for ${ARCHITECTURE} exists"
|
||||||
|
continue
|
||||||
else
|
else
|
||||||
echo -e "wintun.dll for ${ARCHITECTURE} is missing"
|
echo -e "wintun.dll for ${ARCHITECTURE} is missing"
|
||||||
missing=true
|
missing=true
|
||||||
@@ -115,21 +113,16 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
if [[ "$missing" == true ]]; then
|
if [[ "$missing" == true ]]; then
|
||||||
FILENAME=wintun.zip
|
FILENAME=wintun.zip
|
||||||
DOWNLOAD_FILE=wintun-${ASSETVER}.zip
|
DOWNLOAD_FILE=wintun-0.14.1.zip
|
||||||
echo -e "Downloading https://www.wintun.net/builds/${DOWNLOAD_FILE}..."
|
echo -e "Downloading https://www.wintun.net/builds/${DOWNLOAD_FILE}..."
|
||||||
curl -L "https://www.wintun.net/builds/${DOWNLOAD_FILE}" -o "${FILENAME}"
|
curl -L "https://www.wintun.net/builds/${DOWNLOAD_FILE}" -o "${FILENAME}"
|
||||||
if [[ "$(sha256sum "./${FILENAME}" | awk -F ' ' '{print $1}')" == "${ASSETHASH}" ]]; then
|
echo -e "Unpacking wintun..."
|
||||||
echo -e "Unpacking wintun..."
|
unzip -u ${FILENAME} -d resources/
|
||||||
unzip -u ${FILENAME} -d resources/
|
echo "unhit=true" >> $GITHUB_OUTPUT
|
||||||
echo "unhit=true" >> $GITHUB_OUTPUT
|
|
||||||
else
|
|
||||||
echo -e "Digest of ${FILENAME} mismatch."
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Save Wintun Cache
|
- name: Save Wintun Cache
|
||||||
uses: actions/cache/save@v6
|
uses: actions/cache/save@v5
|
||||||
if: ${{ steps.update.outputs.unhit }}
|
if: ${{ steps.update.outputs.unhit }}
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
name: Tests and Checkings
|
name: Test
|
||||||
|
|
||||||
on:
|
on:
|
||||||
push:
|
push:
|
||||||
@@ -8,10 +8,9 @@ on:
|
|||||||
jobs:
|
jobs:
|
||||||
check-assets:
|
check-assets:
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
|
||||||
steps:
|
steps:
|
||||||
- name: Restore Geodat Cache
|
- name: Restore Geodat Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
@@ -37,10 +36,9 @@ jobs:
|
|||||||
|
|
||||||
check-proto:
|
check-proto:
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout codebase
|
- name: Checkout codebase
|
||||||
uses: actions/checkout@v7
|
uses: actions/checkout@v6
|
||||||
- name: Check Proto Version Header
|
- name: Check Proto Version Header
|
||||||
run: |
|
run: |
|
||||||
head -n 4 core/config.pb.go > ref.txt
|
head -n 4 core/config.pb.go > ref.txt
|
||||||
@@ -52,26 +50,8 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
done
|
done
|
||||||
|
|
||||||
check-format:
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
|
||||||
permissions:
|
|
||||||
contents: read
|
|
||||||
steps:
|
|
||||||
- name: Checkout codebase
|
|
||||||
uses: actions/checkout@v7
|
|
||||||
- name: Set up Go
|
|
||||||
uses: actions/setup-go@v7
|
|
||||||
with:
|
|
||||||
go-version-file: go.mod
|
|
||||||
check-latest: true
|
|
||||||
cache: false
|
|
||||||
- name: Check Format
|
|
||||||
run: go run ./infra/vformat/main.go -mode check -pwd ./
|
|
||||||
|
|
||||||
test:
|
test:
|
||||||
needs: check-assets
|
needs: check-assets
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
|
||||||
permissions:
|
permissions:
|
||||||
contents: read
|
contents: read
|
||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
@@ -81,14 +61,14 @@ jobs:
|
|||||||
os: [windows-latest, ubuntu-latest, macos-latest]
|
os: [windows-latest, ubuntu-latest, macos-latest]
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout codebase
|
- name: Checkout codebase
|
||||||
uses: actions/checkout@v7
|
uses: actions/checkout@v6
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v7
|
uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
check-latest: true
|
check-latest: true
|
||||||
- name: Restore Geodat Cache
|
- name: Restore Geodat Cache
|
||||||
uses: actions/cache/restore@v6
|
uses: actions/cache/restore@v5
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|||||||
@@ -73,7 +73,7 @@
|
|||||||
- [Xray_bash_onekey](https://github.com/hello-yunshu/Xray_bash_onekey), [XTool](https://github.com/LordPenguin666/XTool), [VPainLess](https://github.com/vpainless/vpainless)
|
- [Xray_bash_onekey](https://github.com/hello-yunshu/Xray_bash_onekey), [XTool](https://github.com/LordPenguin666/XTool), [VPainLess](https://github.com/vpainless/vpainless)
|
||||||
- [v2ray-agent](https://github.com/mack-a/v2ray-agent), [Xray_onekey](https://github.com/wulabing/Xray_onekey), [ProxySU](https://github.com/proxysu/ProxySU)
|
- [v2ray-agent](https://github.com/mack-a/v2ray-agent), [Xray_onekey](https://github.com/wulabing/Xray_onekey), [ProxySU](https://github.com/proxysu/ProxySU)
|
||||||
- Magisk
|
- Magisk
|
||||||
- [Magic_V2Ray](https://github.com/vincentng295/Magic_V2Ray)
|
- [Xray4Magisk](https://github.com/Asterisk4Magisk/Xray4Magisk)
|
||||||
- [Xray_For_Magisk](https://github.com/E7KMbb/Xray_For_Magisk)
|
- [Xray_For_Magisk](https://github.com/E7KMbb/Xray_For_Magisk)
|
||||||
- Homebrew
|
- Homebrew
|
||||||
- `brew install xray`
|
- `brew install xray`
|
||||||
@@ -120,7 +120,6 @@
|
|||||||
- [XrayFA](https://github.com/Q7DF1/XrayFA)
|
- [XrayFA](https://github.com/Q7DF1/XrayFA)
|
||||||
- [AnyPortal](https://github.com/AnyPortal/AnyPortal)
|
- [AnyPortal](https://github.com/AnyPortal/AnyPortal)
|
||||||
- [OneXray](https://github.com/OneXray/OneXray)
|
- [OneXray](https://github.com/OneXray/OneXray)
|
||||||
- [AsteriskNG](https://github.com/Asterisk4Magisk/AsteriskNG)
|
|
||||||
- iOS & macOS arm64 & tvOS
|
- iOS & macOS arm64 & tvOS
|
||||||
- [Happ](https://apps.apple.com/app/happ-proxy-utility/id6504287215) | [Happ RU](https://apps.apple.com/ru/app/happ-proxy-utility-plus/id6746188973) | [Happ tvOS](https://apps.apple.com/us/app/happ-proxy-utility-for-tv/id6748297274)
|
- [Happ](https://apps.apple.com/app/happ-proxy-utility/id6504287215) | [Happ RU](https://apps.apple.com/ru/app/happ-proxy-utility-plus/id6746188973) | [Happ tvOS](https://apps.apple.com/us/app/happ-proxy-utility-for-tv/id6748297274)
|
||||||
- [Streisand](https://apps.apple.com/app/streisand/id6450534064)
|
- [Streisand](https://apps.apple.com/app/streisand/id6450534064)
|
||||||
@@ -146,8 +145,6 @@
|
|||||||
- [v2rayN](https://github.com/2dust/v2rayN)
|
- [v2rayN](https://github.com/2dust/v2rayN)
|
||||||
- [GenyConnect](https://github.com/genyleap/GenyConnect)
|
- [GenyConnect](https://github.com/genyleap/GenyConnect)
|
||||||
- [OneXray](https://github.com/OneXray/OneXray)
|
- [OneXray](https://github.com/OneXray/OneXray)
|
||||||
- HarmonyOS
|
|
||||||
- [Hey](https://github.com/popsiclelmlm/Hey)
|
|
||||||
|
|
||||||
## Others that support VLESS, XTLS, REALITY, XUDP, PLUX...
|
## Others that support VLESS, XTLS, REALITY, XUDP, PLUX...
|
||||||
|
|
||||||
@@ -165,7 +162,6 @@
|
|||||||
- [xtls-sdk](https://github.com/remnawave/xtls-sdk)
|
- [xtls-sdk](https://github.com/remnawave/xtls-sdk)
|
||||||
- [xtlsapi](https://github.com/hiddify/xtlsapi)
|
- [xtlsapi](https://github.com/hiddify/xtlsapi)
|
||||||
- [AndroidLibXrayLite](https://github.com/2dust/AndroidLibXrayLite)
|
- [AndroidLibXrayLite](https://github.com/2dust/AndroidLibXrayLite)
|
||||||
- [flutter_vless](https://github.com/XIIIFOX/flutter_vless)
|
|
||||||
- [Xray-core-python](https://github.com/LorenEteval/Xray-core-python)
|
- [Xray-core-python](https://github.com/LorenEteval/Xray-core-python)
|
||||||
- [xray-api](https://github.com/XVGuardian/xray-api)
|
- [xray-api](https://github.com/XVGuardian/xray-api)
|
||||||
- [XrayR](https://github.com/XrayR-project/XrayR)
|
- [XrayR](https://github.com/XrayR-project/XrayR)
|
||||||
@@ -187,27 +183,6 @@
|
|||||||
- [Xray-core v1.0.0](https://github.com/XTLS/Xray-core/releases/tag/v1.0.0) was forked from [v2fly-core 9a03cc5](https://github.com/v2fly/v2ray-core/commit/9a03cc5c98d04cc28320fcee26dbc236b3291256), and we have made & accumulated a huge number of enhancements over time, check [the release notes for each version](https://github.com/XTLS/Xray-core/releases).
|
- [Xray-core v1.0.0](https://github.com/XTLS/Xray-core/releases/tag/v1.0.0) was forked from [v2fly-core 9a03cc5](https://github.com/v2fly/v2ray-core/commit/9a03cc5c98d04cc28320fcee26dbc236b3291256), and we have made & accumulated a huge number of enhancements over time, check [the release notes for each version](https://github.com/XTLS/Xray-core/releases).
|
||||||
- For third-party projects used in [Xray-core](https://github.com/XTLS/Xray-core), check your local or [the latest go.mod](https://github.com/XTLS/Xray-core/blob/main/go.mod).
|
- For third-party projects used in [Xray-core](https://github.com/XTLS/Xray-core), check your local or [the latest go.mod](https://github.com/XTLS/Xray-core/blob/main/go.mod).
|
||||||
|
|
||||||
### Bundled Third-Party Components Redistribution
|
|
||||||
|
|
||||||
**Certain optional features dynamically load third-party components. These optional components are separate works distributed under their own licenses, and are bundled into the ZIP package for ease of use. Users may replace these components under the licenses from these components.**
|
|
||||||
|
|
||||||
These components include:
|
|
||||||
|
|
||||||
#### Wintun
|
|
||||||
|
|
||||||
This distribution contains unmodified official precompiled and pre-signed Wintun binaries.
|
|
||||||
|
|
||||||
- Project: Wintun
|
|
||||||
- Copyright: Copyright (C) 2018-2021 WireGuard LLC. All Rights Reserved.
|
|
||||||
- Redistribution License: Prebuilt Binaries License (PBL) bundled with official precompiled and pre-signed binaries from wintun.net
|
|
||||||
- Component(s): wintun.dll
|
|
||||||
- Source: https://www.wintun.net/
|
|
||||||
- Included in:
|
|
||||||
- Windows x86 (windows-32, win7-32)
|
|
||||||
- Windows x86-64 (windows-64, win7-64)
|
|
||||||
- Windows AArch64 (windows-arm64)
|
|
||||||
- Notes: Wintun is an optional runtime-loaded component only used for TUN inbound functionality on supported Windows platforms.
|
|
||||||
|
|
||||||
## One-line Compilation
|
## One-line Compilation
|
||||||
|
|
||||||
### Windows (PowerShell)
|
### Windows (PowerShell)
|
||||||
@@ -231,13 +206,6 @@ Make sure that you are using the same Go version, and remember to set the git co
|
|||||||
CGO_ENABLED=0 go build -o xray -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-X github.com/xtls/xray-core/core.build=REPLACE -s -w -buildid=" -v ./main
|
CGO_ENABLED=0 go build -o xray -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-X github.com/xtls/xray-core/core.build=REPLACE -s -w -buildid=" -v ./main
|
||||||
```
|
```
|
||||||
|
|
||||||
For Android:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
GOOS=android GOARCH=arm64 CGO_ENABLED=1 CC=/path/to/aarch64-linux-android24-clang go build -o xray -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-X github.com/xtls/xray-core/core.build=REPLACE -s -w -buildid= -checklinkname=0" -v ./main
|
|
||||||
GOOS=android GOARCH=amd64 CGO_ENABLED=1 CC=/path/to/x86_64-linux-android24-clang go build -o xray -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-X github.com/xtls/xray-core/core.build=REPLACE -s -w -buildid= -checklinkname=0" -v ./main
|
|
||||||
```
|
|
||||||
|
|
||||||
If you are compiling a 32-bit MIPS/MIPSLE target, use this command instead:
|
If you are compiling a 32-bit MIPS/MIPSLE target, use this command instead:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|||||||
@@ -3,15 +3,15 @@ package commander
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"net"
|
"net"
|
||||||
"strings"
|
|
||||||
"sync"
|
"sync"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/signal/done"
|
"github.com/xtls/xray-core/common/signal/done"
|
||||||
core "github.com/xtls/xray-core/core"
|
core "github.com/xtls/xray-core/core"
|
||||||
"github.com/xtls/xray-core/features/outbound"
|
"github.com/xtls/xray-core/features/outbound"
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -69,12 +69,13 @@ func (c *Commander) Start() error {
|
|||||||
}
|
}
|
||||||
c.Unlock()
|
c.Unlock()
|
||||||
|
|
||||||
listen := func(listener net.Listener) {
|
var listen = func(listener net.Listener) {
|
||||||
if err := c.server.Serve(listener); err != nil {
|
if err := c.server.Serve(listener); err != nil {
|
||||||
errors.LogErrorInner(context.Background(), err, "failed to start grpc server")
|
errors.LogErrorInner(context.Background(), err, "failed to start grpc server")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
if len(c.listen) > 0 {
|
if len(c.listen) > 0 {
|
||||||
var addr net.Addr
|
var addr net.Addr
|
||||||
|
|
||||||
@@ -88,7 +89,7 @@ func (c *Commander) Start() error {
|
|||||||
}
|
}
|
||||||
addr = tcpAddr
|
addr = tcpAddr
|
||||||
}
|
}
|
||||||
l, err := internet.ListenSystem(context.Background(), addr, nil)
|
l, err := internet.ListenSystem(context.Background(), addr, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
errors.LogErrorInner(context.Background(), err, "API server failed to listen on ", c.listen)
|
errors.LogErrorInner(context.Background(), err, "API server failed to listen on ", c.listen)
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -162,7 +162,7 @@ func (d *DefaultDispatcher) getLink(ctx context.Context) (*transport.Link, *tran
|
|||||||
p := d.policy.ForLevel(user.Level)
|
p := d.policy.ForLevel(user.Level)
|
||||||
if p.Stats.UserUplink {
|
if p.Stats.UserUplink {
|
||||||
name := "user>>>" + user.Email + ">>>traffic>>>uplink"
|
name := "user>>>" + user.Email + ">>>traffic>>>uplink"
|
||||||
if c, _ := d.stats.GetOrRegisterCounter(name); c != nil {
|
if c, _ := stats.GetOrRegisterCounter(d.stats, name); c != nil {
|
||||||
inboundLink.Writer = &SizeStatWriter{
|
inboundLink.Writer = &SizeStatWriter{
|
||||||
Counter: c,
|
Counter: c,
|
||||||
Writer: inboundLink.Writer,
|
Writer: inboundLink.Writer,
|
||||||
@@ -171,7 +171,7 @@ func (d *DefaultDispatcher) getLink(ctx context.Context) (*transport.Link, *tran
|
|||||||
}
|
}
|
||||||
if p.Stats.UserDownlink {
|
if p.Stats.UserDownlink {
|
||||||
name := "user>>>" + user.Email + ">>>traffic>>>downlink"
|
name := "user>>>" + user.Email + ">>>traffic>>>downlink"
|
||||||
if c, _ := d.stats.GetOrRegisterCounter(name); c != nil {
|
if c, _ := stats.GetOrRegisterCounter(d.stats, name); c != nil {
|
||||||
outboundLink.Writer = &SizeStatWriter{
|
outboundLink.Writer = &SizeStatWriter{
|
||||||
Counter: c,
|
Counter: c,
|
||||||
Writer: outboundLink.Writer,
|
Writer: outboundLink.Writer,
|
||||||
@@ -200,13 +200,13 @@ func WrapLink(ctx context.Context, policyManager policy.Manager, statsManager st
|
|||||||
p := policyManager.ForLevel(user.Level)
|
p := policyManager.ForLevel(user.Level)
|
||||||
if p.Stats.UserUplink {
|
if p.Stats.UserUplink {
|
||||||
name := "user>>>" + user.Email + ">>>traffic>>>uplink"
|
name := "user>>>" + user.Email + ">>>traffic>>>uplink"
|
||||||
if c, _ := statsManager.GetOrRegisterCounter(name); c != nil {
|
if c, _ := stats.GetOrRegisterCounter(statsManager, name); c != nil {
|
||||||
link.Reader.(*buf.TimeoutWrapperReader).Counter = c
|
link.Reader.(*buf.TimeoutWrapperReader).Counter = c
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if p.Stats.UserDownlink {
|
if p.Stats.UserDownlink {
|
||||||
name := "user>>>" + user.Email + ">>>traffic>>>downlink"
|
name := "user>>>" + user.Email + ">>>traffic>>>downlink"
|
||||||
if c, _ := statsManager.GetOrRegisterCounter(name); c != nil {
|
if c, _ := stats.GetOrRegisterCounter(statsManager, name); c != nil {
|
||||||
link.Writer = &SizeStatWriter{
|
link.Writer = &SizeStatWriter{
|
||||||
Counter: c,
|
Counter: c,
|
||||||
Writer: link.Writer,
|
Writer: link.Writer,
|
||||||
@@ -223,7 +223,7 @@ func WrapLink(ctx context.Context, policyManager policy.Manager, statsManager st
|
|||||||
|
|
||||||
func trackOnlineIP(ctx context.Context, sm stats.Manager, email, ip string) {
|
func trackOnlineIP(ctx context.Context, sm stats.Manager, email, ip string) {
|
||||||
name := "user>>>" + email + ">>>online"
|
name := "user>>>" + email + ">>>online"
|
||||||
if om, _ := sm.GetOrRegisterOnlineMap(name); om != nil {
|
if om, _ := stats.GetOrRegisterOnlineMap(sm, name); om != nil {
|
||||||
om.AddIP(ip)
|
om.AddIP(ip)
|
||||||
context.AfterFunc(ctx, func() { om.RemoveIP(ip) })
|
context.AfterFunc(ctx, func() { om.RemoveIP(ip) })
|
||||||
}
|
}
|
||||||
@@ -470,9 +470,6 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport.
|
|||||||
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
|
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
if err != common.ErrNoClue {
|
|
||||||
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
|
|
||||||
}
|
|
||||||
errors.LogInfo(ctx, "default route for ", destination)
|
errors.LogInfo(ctx, "default route for ", destination)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ func newFakeDNSSniffer(ctx context.Context) (protocolSnifferWithMetadata, error)
|
|||||||
}
|
}
|
||||||
|
|
||||||
if fakeDNSEngine == nil {
|
if fakeDNSEngine == nil {
|
||||||
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used")
|
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used").AtError()
|
||||||
return protocolSnifferWithMetadata{}, errNotInit
|
return protocolSnifferWithMetadata{}, errNotInit
|
||||||
}
|
}
|
||||||
return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) {
|
return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) {
|
||||||
|
|||||||
@@ -139,8 +139,7 @@ func (c *CacheController) writeAndShrink(expiredKeys []string) {
|
|||||||
|
|
||||||
if lenAfter == 0 {
|
if lenAfter == 0 {
|
||||||
if c.highWatermark >= minSizeForEmptyRebuild {
|
if c.highWatermark >= minSizeForEmptyRebuild {
|
||||||
errors.LogDebug(
|
errors.LogDebug(context.Background(), c.name,
|
||||||
context.Background(), c.name,
|
|
||||||
" rebuilding empty cache map to reclaim memory.",
|
" rebuilding empty cache map to reclaim memory.",
|
||||||
" size_before_cleanup=", lenBefore,
|
" size_before_cleanup=", lenBefore,
|
||||||
" peak_size_before_rebuild=", c.highWatermark,
|
" peak_size_before_rebuild=", c.highWatermark,
|
||||||
@@ -154,8 +153,7 @@ func (c *CacheController) writeAndShrink(expiredKeys []string) {
|
|||||||
|
|
||||||
if reductionFromPeak := c.highWatermark - lenAfter; reductionFromPeak > shrinkAbsoluteThreshold &&
|
if reductionFromPeak := c.highWatermark - lenAfter; reductionFromPeak > shrinkAbsoluteThreshold &&
|
||||||
float64(reductionFromPeak) > float64(c.highWatermark)*shrinkRatioThreshold {
|
float64(reductionFromPeak) > float64(c.highWatermark)*shrinkRatioThreshold {
|
||||||
errors.LogDebug(
|
errors.LogDebug(context.Background(), c.name,
|
||||||
context.Background(), c.name,
|
|
||||||
" shrinking cache map to reclaim memory.",
|
" shrinking cache map to reclaim memory.",
|
||||||
" new_size=", lenAfter,
|
" new_size=", lenAfter,
|
||||||
" peak_size_before_shrink=", c.highWatermark,
|
" peak_size_before_shrink=", c.highWatermark,
|
||||||
@@ -167,6 +165,7 @@ func (c *CacheController) writeAndShrink(expiredKeys []string) {
|
|||||||
c.highWatermark = lenAfter
|
c.highWatermark = lenAfter
|
||||||
go c.migrate()
|
go c.migrate()
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type migrationEntry struct {
|
type migrationEntry struct {
|
||||||
|
|||||||
+1
-1
@@ -28,7 +28,7 @@ func toNetIP(addrs []net.Address) ([]net.IP, error) {
|
|||||||
if addr.Family().IsIP() {
|
if addr.Family().IsIP() {
|
||||||
ips = append(ips, addr.IP())
|
ips = append(ips, addr.IP())
|
||||||
} else {
|
} else {
|
||||||
return nil, errors.New("Failed to convert address", addr, "to Net IP.")
|
return nil, errors.New("Failed to convert address", addr, "to Net IP.").AtWarning()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return ips, nil
|
return ips, nil
|
||||||
|
|||||||
+6
-25
@@ -93,7 +93,6 @@ type NameServer struct {
|
|||||||
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
|
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
|
||||||
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
|
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
|
||||||
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
|
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
|
||||||
Id string `protobuf:"bytes,18,opt,name=id,proto3" json:"id,omitempty"`
|
|
||||||
unknownFields protoimpl.UnknownFields
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -240,13 +239,6 @@ func (x *NameServer) GetPolicyID() uint32 {
|
|||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *NameServer) GetId() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Id
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
// NameServer list used by this DNS client.
|
// NameServer list used by this DNS client.
|
||||||
@@ -266,10 +258,8 @@ type Config struct {
|
|||||||
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
||||||
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
|
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
|
||||||
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
|
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
|
||||||
// Absolute path to the Lua DNS query script.
|
unknownFields protoimpl.UnknownFields
|
||||||
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
|
sizeCache protoimpl.SizeCache
|
||||||
unknownFields protoimpl.UnknownFields
|
|
||||||
sizeCache protoimpl.SizeCache
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) Reset() {
|
func (x *Config) Reset() {
|
||||||
@@ -379,13 +369,6 @@ func (x *Config) GetEnableParallelQuery() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) GetScript() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Script
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
type Config_HostMapping struct {
|
type Config_HostMapping struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
|
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
|
||||||
@@ -452,7 +435,7 @@ var File_app_dns_config_proto protoreflect.FileDescriptor
|
|||||||
|
|
||||||
const file_app_dns_config_proto_rawDesc = "" +
|
const file_app_dns_config_proto_rawDesc = "" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xee\x05\n" +
|
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xde\x05\n" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"NameServer\x123\n" +
|
"NameServer\x123\n" +
|
||||||
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
|
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
|
||||||
@@ -478,11 +461,10 @@ const file_app_dns_config_proto_rawDesc = "" +
|
|||||||
"\n" +
|
"\n" +
|
||||||
"actUnprior\x18\x0e \x01(\bR\n" +
|
"actUnprior\x18\x0e \x01(\bR\n" +
|
||||||
"actUnprior\x12\x1a\n" +
|
"actUnprior\x12\x1a\n" +
|
||||||
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
|
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
|
||||||
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
|
|
||||||
"\r_disableCacheB\r\n" +
|
"\r_disableCacheB\r\n" +
|
||||||
"\v_serveStaleB\x12\n" +
|
"\v_serveStaleB\x12\n" +
|
||||||
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x9a\x05\n" +
|
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x82\x05\n" +
|
||||||
"\x06Config\x129\n" +
|
"\x06Config\x129\n" +
|
||||||
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
|
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
|
||||||
"nameServer\x12\x1b\n" +
|
"nameServer\x12\x1b\n" +
|
||||||
@@ -498,8 +480,7 @@ const file_app_dns_config_proto_rawDesc = "" +
|
|||||||
"\x0fdisableFallback\x18\n" +
|
"\x0fdisableFallback\x18\n" +
|
||||||
" \x01(\bR\x0fdisableFallback\x126\n" +
|
" \x01(\bR\x0fdisableFallback\x126\n" +
|
||||||
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
|
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
|
||||||
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
|
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x1a}\n" +
|
||||||
"\x06script\x18\x0f \x01(\tR\x06script\x1a}\n" +
|
|
||||||
"\vHostMapping\x127\n" +
|
"\vHostMapping\x127\n" +
|
||||||
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
|
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
|
||||||
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
|
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ message NameServer {
|
|||||||
repeated xray.common.geodata.IPRule unexpected_ip = 13;
|
repeated xray.common.geodata.IPRule unexpected_ip = 13;
|
||||||
bool actUnprior = 14;
|
bool actUnprior = 14;
|
||||||
uint32 policyID = 17;
|
uint32 policyID = 17;
|
||||||
string id = 18;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
enum QueryStrategy {
|
enum QueryStrategy {
|
||||||
@@ -74,7 +73,4 @@ message Config {
|
|||||||
bool disableFallbackIfMatch = 11;
|
bool disableFallbackIfMatch = 11;
|
||||||
|
|
||||||
bool enableParallelQuery = 14;
|
bool enableParallelQuery = 14;
|
||||||
|
|
||||||
// Absolute path to the Lua DNS query script.
|
|
||||||
string script = 15;
|
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-40
@@ -31,8 +31,6 @@ type DNS struct {
|
|||||||
domainMatcher geodata.DomainMatcher
|
domainMatcher geodata.DomainMatcher
|
||||||
matcherInfos []*DomainMatcherInfo
|
matcherInfos []*DomainMatcherInfo
|
||||||
checkSystem bool
|
checkSystem bool
|
||||||
script *scriptEngine
|
|
||||||
scriptPath string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
|
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
|
||||||
@@ -88,7 +86,7 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
|
|||||||
return nil, errors.New("failed to create hosts").Base(err)
|
return nil, errors.New("failed to create hosts").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
defaultTag := config.Tag
|
var defaultTag = config.Tag
|
||||||
if len(config.Tag) == 0 {
|
if len(config.Tag) == 0 {
|
||||||
defaultTag = generateRandomTag()
|
defaultTag = generateRandomTag()
|
||||||
}
|
}
|
||||||
@@ -141,7 +139,7 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
|
|||||||
serveExpiredTTL = *ns.ServeExpiredTTL
|
serveExpiredTTL = *ns.ServeExpiredTTL
|
||||||
}
|
}
|
||||||
|
|
||||||
tag := defaultTag
|
var tag = defaultTag
|
||||||
if len(ns.Tag) > 0 {
|
if len(ns.Tag) > 0 {
|
||||||
tag = ns.Tag
|
tag = ns.Tag
|
||||||
}
|
}
|
||||||
@@ -182,7 +180,6 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
|
|||||||
disableFallbackIfMatch: config.DisableFallbackIfMatch,
|
disableFallbackIfMatch: config.DisableFallbackIfMatch,
|
||||||
enableParallelQuery: config.EnableParallelQuery,
|
enableParallelQuery: config.EnableParallelQuery,
|
||||||
checkSystem: checkSystem,
|
checkSystem: checkSystem,
|
||||||
scriptPath: config.Script,
|
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -193,21 +190,11 @@ func (*DNS) Type() interface{} {
|
|||||||
|
|
||||||
// Start implements common.Runnable.
|
// Start implements common.Runnable.
|
||||||
func (s *DNS) Start() error {
|
func (s *DNS) Start() error {
|
||||||
if s.scriptPath != "" {
|
|
||||||
engine, err := newScriptEngine(s.scriptPath, s)
|
|
||||||
if err != nil {
|
|
||||||
return errors.New("failed to initialize DNS script").Base(err)
|
|
||||||
}
|
|
||||||
s.script = engine
|
|
||||||
}
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close implements common.Closable.
|
// Close implements common.Closable.
|
||||||
func (s *DNS) Close() error {
|
func (s *DNS) Close() error {
|
||||||
if s.script != nil {
|
|
||||||
s.script.close()
|
|
||||||
}
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -225,28 +212,6 @@ func (s *DNS) IsOwnLink(ctx context.Context) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// MayUseSystemResolver reports whether any name server configured here could
|
|
||||||
// still resolve through the system resolver. That is what happens when no name
|
|
||||||
// server is configured at all, and it is also what a name server pointed at
|
|
||||||
// "localhost" does. Callers that are about to redirect the system resolver need
|
|
||||||
// to know, because a resolution path that reaches it would then loop back to
|
|
||||||
// them.
|
|
||||||
//
|
|
||||||
// Any such server is enough: name servers can be selected per domain, so a
|
|
||||||
// single local one makes some query reach the system resolver even when
|
|
||||||
// independent upstreams are configured alongside it.
|
|
||||||
func (s *DNS) MayUseSystemResolver() bool {
|
|
||||||
if len(s.clients) == 0 {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
for _, client := range s.clients {
|
|
||||||
if _, isLocal := client.server.(*LocalNameServer); isLocal {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// LookupIP implements dns.Client.
|
// LookupIP implements dns.Client.
|
||||||
func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
||||||
// Normalize the FQDN form query
|
// Normalize the FQDN form query
|
||||||
@@ -292,9 +257,6 @@ func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, er
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Name servers lookup
|
// Name servers lookup
|
||||||
if s.script != nil {
|
|
||||||
return s.script.query(domain, option)
|
|
||||||
}
|
|
||||||
if s.enableParallelQuery {
|
if s.enableParallelQuery {
|
||||||
return s.parallelQuery(domain, option)
|
return s.parallelQuery(domain, option)
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -1,59 +0,0 @@
|
|||||||
package dns
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
feature_dns "github.com/xtls/xray-core/features/dns"
|
|
||||||
)
|
|
||||||
|
|
||||||
// fakeServer stands in for any name server that is not the system resolver.
|
|
||||||
type fakeServer struct{}
|
|
||||||
|
|
||||||
func (fakeServer) Name() string { return "fake" }
|
|
||||||
func (fakeServer) IsDisableCache() bool { return false }
|
|
||||||
func (fakeServer) QueryIP(context.Context, string, feature_dns.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
return nil, 0, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Callers that are about to redirect the system resolver rely on this to tell
|
|
||||||
// whether any resolution path could still reach the system resolver, so the
|
|
||||||
// mixed shape has to be reported as reachable: a domain-specific rule can
|
|
||||||
// select the system resolver even when an independent upstream also exists.
|
|
||||||
func TestMayUseSystemResolver(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
clients []*Client
|
|
||||||
want bool
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "no clients at all",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "only the system resolver",
|
|
||||||
clients: []*Client{{server: NewLocalNameServer()}},
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "the system resolver alongside an independent name server",
|
|
||||||
clients: []*Client{{server: fakeServer{}}, {server: NewLocalNameServer()}},
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "only independent name servers",
|
|
||||||
clients: []*Client{{server: fakeServer{}}},
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
server := &DNS{clients: tt.clients}
|
|
||||||
if got := server.MayUseSystemResolver(); got != tt.want {
|
|
||||||
t.Errorf("MayUseSystemResolver() = %v, want %v", got, tt.want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+8
-18
@@ -127,20 +127,15 @@ func genEDNS0Options(clientIP net.IP, padding int) *dnsmessage.Resource {
|
|||||||
return opt
|
return opt
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildReqMsgs(domain string, option dns_feature.IPOption, reqIDGen func() uint16, reqOpts *dnsmessage.Resource) ([]*dnsRequest, error) {
|
func buildReqMsgs(domain string, option dns_feature.IPOption, reqIDGen func() uint16, reqOpts *dnsmessage.Resource) []*dnsRequest {
|
||||||
name, err := dnsmessage.NewName(domain)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
qA := dnsmessage.Question{
|
qA := dnsmessage.Question{
|
||||||
Name: name,
|
Name: dnsmessage.MustNewName(domain),
|
||||||
Type: dnsmessage.TypeA,
|
Type: dnsmessage.TypeA,
|
||||||
Class: dnsmessage.ClassINET,
|
Class: dnsmessage.ClassINET,
|
||||||
}
|
}
|
||||||
|
|
||||||
qAAAA := dnsmessage.Question{
|
qAAAA := dnsmessage.Question{
|
||||||
Name: name,
|
Name: dnsmessage.MustNewName(domain),
|
||||||
Type: dnsmessage.TypeAAAA,
|
Type: dnsmessage.TypeAAAA,
|
||||||
Class: dnsmessage.ClassINET,
|
Class: dnsmessage.ClassINET,
|
||||||
}
|
}
|
||||||
@@ -180,7 +175,7 @@ func buildReqMsgs(domain string, option dns_feature.IPOption, reqIDGen func() ui
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
return reqs, nil
|
return reqs
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseResponse parses DNS answers from the returned payload
|
// parseResponse parses DNS answers from the returned payload
|
||||||
@@ -188,24 +183,19 @@ func parseResponse(payload []byte) (*IPRecord, error) {
|
|||||||
var parser dnsmessage.Parser
|
var parser dnsmessage.Parser
|
||||||
h, err := parser.Start(payload)
|
h, err := parser.Start(payload)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to parse DNS response").Base(err)
|
return nil, errors.New("failed to parse DNS response").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
if err := parser.SkipAllQuestions(); err != nil {
|
if err := parser.SkipAllQuestions(); err != nil {
|
||||||
return nil, errors.New("failed to skip questions in DNS response").Base(err)
|
return nil, errors.New("failed to skip questions in DNS response").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
ipRecord := &IPRecord{
|
ipRecord := &IPRecord{
|
||||||
ReqID: h.ID,
|
ReqID: h.ID,
|
||||||
RCode: h.RCode,
|
RCode: h.RCode,
|
||||||
|
Expire: now.Add(time.Second * dns_feature.DefaultTTL),
|
||||||
RawHeader: &h,
|
RawHeader: &h,
|
||||||
}
|
}
|
||||||
defer func() {
|
|
||||||
// set to default TTL if no valid TTL is found
|
|
||||||
if ipRecord.Expire.IsZero() {
|
|
||||||
ipRecord.Expire = now.Add(time.Second * dns_feature.DefaultTTL)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
L:
|
L:
|
||||||
for {
|
for {
|
||||||
@@ -222,7 +212,7 @@ L:
|
|||||||
ttl = 1
|
ttl = 1
|
||||||
}
|
}
|
||||||
expire := now.Add(time.Duration(ttl) * time.Second)
|
expire := now.Add(time.Duration(ttl) * time.Second)
|
||||||
if ipRecord.Expire.IsZero() || ipRecord.Expire.After(expire) {
|
if ipRecord.Expire.After(expire) {
|
||||||
ipRecord.Expire = expire
|
ipRecord.Expire = expire
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ package dns
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"math/rand"
|
"math/rand"
|
||||||
"strings"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -25,8 +24,7 @@ func Test_parseResponse(t *testing.T) {
|
|||||||
|
|
||||||
ans = new(dns.Msg)
|
ans = new(dns.Msg)
|
||||||
ans.Id = 1
|
ans.Id = 1
|
||||||
ans.Answer = append(
|
ans.Answer = append(ans.Answer,
|
||||||
ans.Answer,
|
|
||||||
common.Must2(dns.NewRR("google.com. IN CNAME m.test.google.com")),
|
common.Must2(dns.NewRR("google.com. IN CNAME m.test.google.com")),
|
||||||
common.Must2(dns.NewRR("google.com. IN CNAME fake.google.com")),
|
common.Must2(dns.NewRR("google.com. IN CNAME fake.google.com")),
|
||||||
common.Must2(dns.NewRR("google.com. IN A 8.8.8.8")),
|
common.Must2(dns.NewRR("google.com. IN A 8.8.8.8")),
|
||||||
@@ -36,8 +34,7 @@ func Test_parseResponse(t *testing.T) {
|
|||||||
|
|
||||||
ans = new(dns.Msg)
|
ans = new(dns.Msg)
|
||||||
ans.Id = 2
|
ans.Id = 2
|
||||||
ans.Answer = append(
|
ans.Answer = append(ans.Answer,
|
||||||
ans.Answer,
|
|
||||||
common.Must2(dns.NewRR("google.com. IN CNAME m.test.google.com")),
|
common.Must2(dns.NewRR("google.com. IN CNAME m.test.google.com")),
|
||||||
common.Must2(dns.NewRR("google.com. IN CNAME fake.google.com")),
|
common.Must2(dns.NewRR("google.com. IN CNAME fake.google.com")),
|
||||||
common.Must2(dns.NewRR("google.com. IN CNAME m.test.google.com")),
|
common.Must2(dns.NewRR("google.com. IN CNAME m.test.google.com")),
|
||||||
@@ -134,15 +131,10 @@ func Test_buildReqMsgs(t *testing.T) {
|
|||||||
IPv6Enable: false,
|
IPv6Enable: false,
|
||||||
FakeEnable: false,
|
FakeEnable: false,
|
||||||
}, nil}, 0},
|
}, nil}, 0},
|
||||||
{"name too long", args{strings.Repeat("a", 256), dns_feature.IPOption{
|
|
||||||
IPv4Enable: true,
|
|
||||||
IPv6Enable: true,
|
|
||||||
FakeEnable: false,
|
|
||||||
}, nil}, 0},
|
|
||||||
}
|
}
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
if got, _ := buildReqMsgs(tt.args.domain, tt.args.option, stubID, tt.args.reqOpts); !(len(got) == tt.want) {
|
if got := buildReqMsgs(tt.args.domain, tt.args.option, stubID, tt.args.reqOpts); !(len(got) == tt.want) {
|
||||||
t.Errorf("buildReqMsgs() = %v, want %v", got, tt.want)
|
t.Errorf("buildReqMsgs() = %v, want %v", got, tt.want)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -58,7 +58,7 @@ func NewFakeDNSHolder() (*Holder, error) {
|
|||||||
var err error
|
var err error
|
||||||
|
|
||||||
if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil {
|
if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil {
|
||||||
return nil, errors.New("Unable to create Fake Dns Engine").Base(err)
|
return nil, errors.New("Unable to create Fake Dns Engine").Base(err).AtError()
|
||||||
}
|
}
|
||||||
err = fkdns.initialize(dns.FakeIPv4Pool, 65535)
|
err = fkdns.initialize(dns.FakeIPv4Pool, 65535)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -80,13 +80,13 @@ func (fkdns *Holder) initialize(ipPoolCidr string, lruSize int) error {
|
|||||||
var err error
|
var err error
|
||||||
|
|
||||||
if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil {
|
if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil {
|
||||||
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err)
|
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err).AtError()
|
||||||
}
|
}
|
||||||
|
|
||||||
ones, bits := ipRange.Mask.Size()
|
ones, bits := ipRange.Mask.Size()
|
||||||
rooms := bits - ones
|
rooms := bits - ones
|
||||||
if math.Log2(float64(lruSize)) >= float64(rooms) {
|
if math.Log2(float64(lruSize)) >= float64(rooms) {
|
||||||
return errors.New("LRU size is bigger than subnet size")
|
return errors.New("LRU size is bigger than subnet size").AtError()
|
||||||
}
|
}
|
||||||
fkdns.domainToIP = cache.NewLru(lruSize)
|
fkdns.domainToIP = cache.NewLru(lruSize)
|
||||||
fkdns.ipRange = ipRange
|
fkdns.ipRange = ipRange
|
||||||
|
|||||||
@@ -129,16 +129,15 @@ func TestFakeDnsHolderCreateMappingAndRollOver(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestFakeDNSMulti(t *testing.T) {
|
func TestFakeDNSMulti(t *testing.T) {
|
||||||
fakeMulti, err := NewFakeDNSHolderMulti(
|
fakeMulti, err := NewFakeDNSHolderMulti(&FakeDnsPoolMulti{
|
||||||
&FakeDnsPoolMulti{
|
Pools: []*FakeDnsPool{{
|
||||||
Pools: []*FakeDnsPool{{
|
IpPool: "240.0.0.0/12",
|
||||||
IpPool: "240.0.0.0/12",
|
LruSize: 256,
|
||||||
LruSize: 256,
|
}, {
|
||||||
}, {
|
IpPool: "fddd:c5b4:ff5f:f4f0::/64",
|
||||||
IpPool: "fddd:c5b4:ff5f:f4f0::/64",
|
LruSize: 256,
|
||||||
LruSize: 256,
|
}},
|
||||||
}},
|
},
|
||||||
},
|
|
||||||
)
|
)
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
|
|
||||||
|
|||||||
-165
@@ -1,165 +0,0 @@
|
|||||||
package dns
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
xlua "github.com/xtls/xray-core/common/lua"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
|
||||||
"github.com/xtls/xray-core/features/dns/localdns"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
// luaDNSServer adapts configured and local DNS to the same Lua API.
|
|
||||||
type luaDNSServer struct {
|
|
||||||
id string
|
|
||||||
name string
|
|
||||||
query func(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// RegisterLua makes xray.dns available to scripts backed by client.
|
|
||||||
func RegisterLua(L *lua.LState, client featureDNS.Client) {
|
|
||||||
var servers []luaDNSServer
|
|
||||||
switch client := client.(type) {
|
|
||||||
case *DNS:
|
|
||||||
servers = luaServers(client)
|
|
||||||
case *localdns.Client:
|
|
||||||
servers = []luaDNSServer{{
|
|
||||||
id: "localhost",
|
|
||||||
name: "localhost",
|
|
||||||
query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
return client.LookupIP(domain, option)
|
|
||||||
},
|
|
||||||
}}
|
|
||||||
}
|
|
||||||
registerLua(L, servers, client)
|
|
||||||
}
|
|
||||||
|
|
||||||
// registerLua makes xray.dns available to DNS scripts.
|
|
||||||
func (s *DNS) registerLua(L *lua.LState) {
|
|
||||||
registerLua(L, luaServers(s), nil)
|
|
||||||
}
|
|
||||||
|
|
||||||
func luaServers(s *DNS) []luaDNSServer {
|
|
||||||
servers := make([]luaDNSServer, len(s.clients))
|
|
||||||
for i, client := range s.clients {
|
|
||||||
servers[i] = luaDNSServer{id: client.id, name: client.Name(), query: client.QueryIP}
|
|
||||||
}
|
|
||||||
return servers
|
|
||||||
}
|
|
||||||
|
|
||||||
func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client) {
|
|
||||||
L.PreloadModule("xray.dns", func(L *lua.LState) int {
|
|
||||||
serverList := L.NewTable()
|
|
||||||
for i, client := range servers {
|
|
||||||
server := L.NewTable()
|
|
||||||
|
|
||||||
server.RawSetString("ID", lua.LString(client.id))
|
|
||||||
|
|
||||||
server.RawSetString("Query", L.NewFunction(func(L *lua.LState) int {
|
|
||||||
domain, ok := L.Get(2).(lua.LString)
|
|
||||||
if !ok {
|
|
||||||
L.RaiseError("server:Query requires a domain")
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
option := featureDNS.IPOption{
|
|
||||||
IPv4Enable: L.CheckBool(3),
|
|
||||||
IPv6Enable: L.CheckBool(4),
|
|
||||||
FakeEnable: L.CheckBool(5),
|
|
||||||
}
|
|
||||||
ctx := L.Context()
|
|
||||||
if ctx == nil {
|
|
||||||
L.RaiseError("server:Query requires an active DNS query")
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
var ips []net.IP
|
|
||||||
var ttl uint32
|
|
||||||
var err error
|
|
||||||
if !option.FakeEnable && strings.EqualFold(client.name, "FakeDNS") {
|
|
||||||
err = featureDNS.ErrEmptyResponse
|
|
||||||
} else {
|
|
||||||
ips, ttl, err = client.query(ctx, string(domain), option)
|
|
||||||
}
|
|
||||||
xlua.PushUserData(L, ips)
|
|
||||||
xlua.PushNumber(L, ttl)
|
|
||||||
xlua.PushError(L, err)
|
|
||||||
return 3
|
|
||||||
}))
|
|
||||||
serverList.RawSetInt(i+1, server)
|
|
||||||
}
|
|
||||||
|
|
||||||
module := L.NewTable()
|
|
||||||
if servers != nil {
|
|
||||||
module.RawSetString("Servers", serverList)
|
|
||||||
}
|
|
||||||
if client != nil {
|
|
||||||
module.RawSetString("Query", newLuaClientQuery(L, client))
|
|
||||||
}
|
|
||||||
L.Push(module)
|
|
||||||
return 1
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
|
|
||||||
return L.NewFunction(func(L *lua.LState) int {
|
|
||||||
domain, ok := L.Get(1).(lua.LString)
|
|
||||||
if !ok {
|
|
||||||
L.RaiseError("dns.Query requires a domain")
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
option := featureDNS.IPOption{
|
|
||||||
IPv4Enable: L.CheckBool(2),
|
|
||||||
IPv6Enable: L.CheckBool(3),
|
|
||||||
FakeEnable: L.CheckBool(4),
|
|
||||||
}
|
|
||||||
if L.Context() == nil {
|
|
||||||
L.RaiseError("dns.Query requires an active DNS query")
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
ips, ttl, err := client.LookupIP(string(domain), option)
|
|
||||||
xlua.PushUserData(L, ips)
|
|
||||||
xlua.PushNumber(L, ttl)
|
|
||||||
xlua.PushError(L, err)
|
|
||||||
return 3
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// callLuaHook invokes HandleDNSQuery in the supplied state.
|
|
||||||
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
|
|
||||||
func (s *DNS) callLuaHook(L *lua.LState, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
top := L.GetTop()
|
|
||||||
defer L.SetTop(top)
|
|
||||||
fn := L.GetGlobal("HandleDNSQuery")
|
|
||||||
if fn.Type() != lua.LTFunction {
|
|
||||||
return nil, 0, errors.New("DNS script must define HandleDNSQuery(...)")
|
|
||||||
}
|
|
||||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
|
||||||
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
|
|
||||||
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
|
||||||
}
|
|
||||||
|
|
||||||
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
|
|
||||||
if err := xlua.ReadError(errorValue, "DNS script error must be an error or string"); err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
ttl, err := xlua.ReadUint32(ttlValue, "DNS script returned invalid TTL")
|
|
||||||
if err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
if addresses == lua.LNil {
|
|
||||||
return nil, 0, featureDNS.ErrEmptyResponse
|
|
||||||
}
|
|
||||||
ips, err := xlua.ReadUserData[[]net.IP](addresses, "DNS script IPs must be native IP slice userdata")
|
|
||||||
if err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
if len(ips) == 0 {
|
|
||||||
return nil, 0, featureDNS.ErrEmptyResponse
|
|
||||||
}
|
|
||||||
return ips, ttl, nil
|
|
||||||
}
|
|
||||||
@@ -1,294 +0,0 @@
|
|||||||
package dns
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
go_errors "errors"
|
|
||||||
"math"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
|
||||||
"github.com/xtls/xray-core/features/dns/localdns"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestReadLuaDNSResult(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
|
|
||||||
addresses := L.NewUserData()
|
|
||||||
addresses.Value = want
|
|
||||||
ips, ttl, err := readLuaDNSResult(addresses, lua.LNumber(45), lua.LNil)
|
|
||||||
if err != nil || ttl != 45 || len(ips) != len(want) {
|
|
||||||
t.Fatalf("readLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
|
|
||||||
}
|
|
||||||
for i := range want {
|
|
||||||
if !ips[i].Equal(want[i]) {
|
|
||||||
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReadLuaDNSResultValidation(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
change func(*[3]lua.LValue)
|
|
||||||
want string
|
|
||||||
}{
|
|
||||||
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
|
|
||||||
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
|
|
||||||
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
|
|
||||||
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
|
|
||||||
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
|
|
||||||
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
|
|
||||||
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
|
|
||||||
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
|
|
||||||
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
addresses := L.NewUserData()
|
|
||||||
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
|
||||||
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
|
|
||||||
tc.change(&values)
|
|
||||||
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
|
|
||||||
if err == nil || !strings.Contains(err.Error(), tc.want) {
|
|
||||||
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
addresses := L.NewUserData()
|
|
||||||
addresses.Value = []net.IP(nil)
|
|
||||||
for _, empty := range []lua.LValue{addresses, lua.LNil} {
|
|
||||||
if _, _, err := readLuaDNSResult(empty, lua.LNumber(0), lua.LNil); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
|
|
||||||
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
wantErr := go_errors.New("upstream failed")
|
|
||||||
errorValue := L.NewUserData()
|
|
||||||
errorValue.Value = wantErr
|
|
||||||
if _, _, err := readLuaDNSResult(lua.LNil, lua.LNil, errorValue); err != wantErr {
|
|
||||||
t.Fatalf("upstream error = %v, want original error %v", err, wantErr)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCallLuaHookCancellation(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
cancel()
|
|
||||||
L.SetContext(ctx)
|
|
||||||
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("CallLuaHook did not stop after context cancellation")
|
|
||||||
}
|
|
||||||
if L.Context() != ctx {
|
|
||||||
t.Fatal("CallLuaHook changed the Lua state's context")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCallLuaHookNormalizesDomain(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
addresses := L.NewUserData()
|
|
||||||
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
|
||||||
L.SetGlobal("ips", addresses)
|
|
||||||
if err := L.DoString(`
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
assert(domain == "example.com")
|
|
||||||
assert(ipv4 and not ipv6 and not fake)
|
|
||||||
return ips, 60, nil
|
|
||||||
end
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
s := &DNS{}
|
|
||||||
if _, _, err := s.callLuaHook(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCallLuaHookRestoresStack(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
body string
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{"success", `return ips, 60`, false},
|
|
||||||
{"error", `error("failed")`, true},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
addresses := L.NewUserData()
|
|
||||||
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
|
||||||
L.SetGlobal("ips", addresses)
|
|
||||||
if err := L.DoString("function HandleDNSQuery() " + tc.body + " end"); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
L.Push(lua.LTrue)
|
|
||||||
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
|
||||||
if (err != nil) != tc.wantErr {
|
|
||||||
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
|
|
||||||
}
|
|
||||||
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
|
||||||
t.Fatal("hook did not restore the stack")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaDNSServerQuery(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
geodata.RegisterLua(L)
|
|
||||||
option := featureDNS.IPOption{IPv4Enable: true}
|
|
||||||
ips := []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("8.8.8.8")}
|
|
||||||
server := &DNS{clients: []*Client{{server: &benchmarkLuaNameServer{ips: ips}, ipOption: &option, timeoutMs: time.Second}}}
|
|
||||||
server.registerLua(L)
|
|
||||||
if err := L.DoString(`
|
|
||||||
local server = require("xray.dns").Servers[1]
|
|
||||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
|
|
||||||
assert(type(ips) == "userdata" and not err)
|
|
||||||
assert(matcher:AnyMatch(ips))
|
|
||||||
local matched = matcher:FilterIPs(ips)
|
|
||||||
return matched, ttl, err
|
|
||||||
end
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
L.SetContext(context.Background())
|
|
||||||
got, ttl, err := server.callLuaHook(L, "example.com", option)
|
|
||||||
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
|
|
||||||
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type luaDNSClient struct {
|
|
||||||
featureDNS.Client
|
|
||||||
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *luaDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
return c.lookup(domain, option)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaDNSClientQuery(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
L.SetContext(context.Background())
|
|
||||||
geodata.RegisterLua(L)
|
|
||||||
want := []net.IP{{127, 0, 0, 1}}
|
|
||||||
client := &luaDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
if domain != "MiXeD.Example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
|
|
||||||
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
|
|
||||||
}
|
|
||||||
return want, 42, nil
|
|
||||||
}}
|
|
||||||
RegisterLua(L, client)
|
|
||||||
if err := L.DoString(`
|
|
||||||
local dns = require("xray.dns")
|
|
||||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.1")
|
|
||||||
assert(dns.Servers == nil)
|
|
||||||
ips, ttl, err = dns.Query("MiXeD.Example.", true, false, true)
|
|
||||||
assert(not err and ttl == 42 and matcher:AnyMatch(ips))
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
got := L.GetGlobal("ips").(*lua.LUserData).Value.([]net.IP)
|
|
||||||
if &got[0] != &want[0] {
|
|
||||||
t.Fatal("dns.Query copied the IP slice")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaDNSLocalClient(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
L.SetContext(context.Background())
|
|
||||||
RegisterLua(L, localdns.New())
|
|
||||||
if err := L.DoString(`
|
|
||||||
local dns = require("xray.dns")
|
|
||||||
assert(dns.Servers[1].ID == "localhost")
|
|
||||||
serverIPs, _, serverErr = dns.Servers[1]:Query("127.0.0.1", true, false, false)
|
|
||||||
clientIPs, _, clientErr = dns.Query("127.0.0.1", true, false, false)
|
|
||||||
assert(not serverErr and not clientErr)
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
for _, name := range []string{"serverIPs", "clientIPs"} {
|
|
||||||
ips := L.GetGlobal(name).(*lua.LUserData).Value.([]net.IP)
|
|
||||||
if len(ips) != 1 || !ips[0].Equal(net.ParseIP("127.0.0.1")) {
|
|
||||||
t.Fatalf("%s = %v", name, ips)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type benchmarkLuaNameServer struct {
|
|
||||||
ips []net.IP
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*benchmarkLuaNameServer) Name() string { return "benchmark" }
|
|
||||||
func (*benchmarkLuaNameServer) IsDisableCache() bool { return true }
|
|
||||||
func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
return s.ips, 60, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
|
|
||||||
// The direct case measures the same DNS client without Lua.
|
|
||||||
func BenchmarkLuaDNSHookCall(b *testing.B) {
|
|
||||||
option := featureDNS.IPOption{IPv4Enable: true}
|
|
||||||
ip := net.ParseIP("127.0.0.1")
|
|
||||||
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
|
|
||||||
client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second}
|
|
||||||
server := &DNS{clients: []*Client{client}}
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
server.registerLua(L)
|
|
||||||
if err := L.DoString(`
|
|
||||||
local server = require("xray.dns").Servers[1]
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
return server:Query(domain, ipv4, ipv6, fake)
|
|
||||||
end
|
|
||||||
`); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
L.SetContext(ctx)
|
|
||||||
for _, bench := range []struct {
|
|
||||||
name string
|
|
||||||
query func() ([]net.IP, uint32, error)
|
|
||||||
}{
|
|
||||||
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
|
|
||||||
{"lua_hook", func() ([]net.IP, uint32, error) { return server.callLuaHook(L, "example.com", option) }},
|
|
||||||
} {
|
|
||||||
b.Run(bench.name, func(b *testing.B) {
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
var ips []net.IP
|
|
||||||
var ttl uint32
|
|
||||||
var err error
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
ips, ttl, err = bench.query()
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
b.StopTimer()
|
|
||||||
if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
|
|
||||||
b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -29,7 +29,6 @@ type Server interface {
|
|||||||
|
|
||||||
// Client is the interface for DNS client.
|
// Client is the interface for DNS client.
|
||||||
type Client struct {
|
type Client struct {
|
||||||
id string
|
|
||||||
server Server
|
server Server
|
||||||
skipFallback bool
|
skipFallback bool
|
||||||
expectedIPs geodata.IPMatcher
|
expectedIPs geodata.IPMatcher
|
||||||
@@ -85,7 +84,7 @@ func NewServer(ctx context.Context, dest net.Destination, dispatcher routing.Dis
|
|||||||
if dest.Network == net.Network_UDP { // UDP classic DNS mode
|
if dest.Network == net.Network_UDP { // UDP classic DNS mode
|
||||||
return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil
|
return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil
|
||||||
}
|
}
|
||||||
return nil, errors.New("No available name server could be created from ", dest)
|
return nil, errors.New("No available name server could be created from ", dest).AtWarning()
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs.
|
// NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs.
|
||||||
@@ -98,12 +97,12 @@ func NewClient(
|
|||||||
ipOption dns.IPOption,
|
ipOption dns.IPOption,
|
||||||
updateRules func(bool),
|
updateRules func(bool),
|
||||||
) (*Client, error) {
|
) (*Client, error) {
|
||||||
client := &Client{id: ns.Id}
|
client := &Client{}
|
||||||
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
|
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
|
||||||
// Create a new server for each client for now
|
// Create a new server for each client for now
|
||||||
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create nameserver").Base(err)
|
return errors.New("failed to create nameserver").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
|
|
||||||
_, isLocalDNS := server.(*LocalNameServer)
|
_, isLocalDNS := server.(*LocalNameServer)
|
||||||
@@ -114,7 +113,7 @@ func NewClient(
|
|||||||
if len(ns.ExpectedIp) > 0 {
|
if len(ns.ExpectedIp) > 0 {
|
||||||
expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp)
|
expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create expected ip matcher").Base(err)
|
return errors.New("failed to create expected ip matcher").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -123,7 +122,7 @@ func NewClient(
|
|||||||
if len(ns.UnexpectedIp) > 0 {
|
if len(ns.UnexpectedIp) > 0 {
|
||||||
unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp)
|
unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create unexpected ip matcher").Base(err)
|
return errors.New("failed to create unexpected ip matcher").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -136,7 +135,7 @@ func NewClient(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
timeoutMs := 4000 * time.Millisecond
|
var timeoutMs = 4000 * time.Millisecond
|
||||||
if ns.TimeoutMs > 0 {
|
if ns.TimeoutMs > 0 {
|
||||||
timeoutMs = time.Duration(ns.TimeoutMs) * time.Millisecond
|
timeoutMs = time.Duration(ns.TimeoutMs) * time.Millisecond
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -137,32 +137,14 @@ func (s *DoHNameServer) sendQuery(ctx context.Context, noResponseErrCh chan<- er
|
|||||||
if s.Name()+"." == "DOH//"+fqdn {
|
if s.Name()+"." == "DOH//"+fqdn {
|
||||||
errors.LogError(ctx, s.Name(), " tries to resolve itself! Use IP or set \"hosts\" instead")
|
errors.LogError(ctx, s.Name(), " tries to resolve itself! Use IP or set \"hosts\" instead")
|
||||||
if noResponseErrCh != nil {
|
if noResponseErrCh != nil {
|
||||||
err := errors.New("tries to resolve itself!", s.Name())
|
noResponseErrCh <- errors.New("tries to resolve itself!", s.Name())
|
||||||
if option.IPv4Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
if option.IPv6Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// As we don't want our traffic pattern looks like DoH, we use Random-Length Padding instead of Block-Length Padding recommended in RFC 8467
|
// As we don't want our traffic pattern looks like DoH, we use Random-Length Padding instead of Block-Length Padding recommended in RFC 8467
|
||||||
// Although DoH server like 1.1.1.1 will pad the response to Block-Length 468, at least it is better than no padding for response at all
|
// Although DoH server like 1.1.1.1 will pad the response to Block-Length 468, at least it is better than no padding for response at all
|
||||||
reqs, err := buildReqMsgs(fqdn, option, s.newReqID, genEDNS0Options(s.clientIP, int(crypto.RandBetween(100, 300))))
|
reqs := buildReqMsgs(fqdn, option, s.newReqID, genEDNS0Options(s.clientIP, int(crypto.RandBetween(100, 300))))
|
||||||
if err != nil {
|
|
||||||
errors.LogErrorInner(ctx, err, "failed to build dns query for ", fqdn)
|
|
||||||
if noResponseErrCh != nil {
|
|
||||||
if option.IPv4Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
if option.IPv6Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var deadline time.Time
|
var deadline time.Time
|
||||||
if d, ok := ctx.Deadline(); ok {
|
if d, ok := ctx.Deadline(); ok {
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ func (s *FakeDNSServer) IsDisableCache() bool {
|
|||||||
|
|
||||||
func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) {
|
func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) {
|
||||||
if f.fakeDNSEngine == nil {
|
if f.fakeDNSEngine == nil {
|
||||||
return nil, 0, errors.New("Unable to locate a fake DNS Engine")
|
return nil, 0, errors.New("Unable to locate a fake DNS Engine").AtError()
|
||||||
}
|
}
|
||||||
|
|
||||||
var ips []net.Address
|
var ips []net.Address
|
||||||
@@ -39,7 +39,7 @@ func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOp
|
|||||||
|
|
||||||
netIP, err := toNetIP(ips)
|
netIP, err := toNetIP(ips)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err)
|
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err).AtError()
|
||||||
}
|
}
|
||||||
|
|
||||||
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ type LocalNameServer struct {
|
|||||||
|
|
||||||
// QueryIP implements Server.
|
// QueryIP implements Server.
|
||||||
func (s *LocalNameServer) QueryIP(ctx context.Context, domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) {
|
func (s *LocalNameServer) QueryIP(ctx context.Context, domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) {
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
ips, ttl, err = s.client.LookupIP(domain, option)
|
ips, ttl, err = s.client.LookupIP(domain, option)
|
||||||
|
|
||||||
@@ -49,5 +50,5 @@ func NewLocalNameServer() *LocalNameServer {
|
|||||||
|
|
||||||
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
|
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
|
||||||
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
|
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
|
||||||
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
|
return &Client{server: NewLocalNameServer(), ipOption: &ipOption}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -78,19 +78,7 @@ func (s *QUICNameServer) getCacheController() *CacheController { return s.cacheC
|
|||||||
func (s *QUICNameServer) sendQuery(ctx context.Context, noResponseErrCh chan<- error, fqdn string, option dns_feature.IPOption) {
|
func (s *QUICNameServer) sendQuery(ctx context.Context, noResponseErrCh chan<- error, fqdn string, option dns_feature.IPOption) {
|
||||||
errors.LogInfo(ctx, s.Name(), " querying: ", fqdn)
|
errors.LogInfo(ctx, s.Name(), " querying: ", fqdn)
|
||||||
|
|
||||||
reqs, err := buildReqMsgs(fqdn, option, s.newReqID, genEDNS0Options(s.clientIP, 0))
|
reqs := buildReqMsgs(fqdn, option, s.newReqID, genEDNS0Options(s.clientIP, 0))
|
||||||
if err != nil {
|
|
||||||
errors.LogErrorInner(ctx, err, "failed to build dns query for ", fqdn)
|
|
||||||
if noResponseErrCh != nil {
|
|
||||||
if option.IPv4Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
if option.IPv6Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var deadline time.Time
|
var deadline time.Time
|
||||||
if d, ok := ctx.Deadline(); ok {
|
if d, ok := ctx.Deadline(); ok {
|
||||||
|
|||||||
@@ -113,19 +113,7 @@ func (s *TCPNameServer) getCacheController() *CacheController {
|
|||||||
func (s *TCPNameServer) sendQuery(ctx context.Context, noResponseErrCh chan<- error, fqdn string, option dns_feature.IPOption) {
|
func (s *TCPNameServer) sendQuery(ctx context.Context, noResponseErrCh chan<- error, fqdn string, option dns_feature.IPOption) {
|
||||||
errors.LogInfo(ctx, s.Name(), " querying DNS for: ", fqdn)
|
errors.LogInfo(ctx, s.Name(), " querying DNS for: ", fqdn)
|
||||||
|
|
||||||
reqs, err := buildReqMsgs(fqdn, option, s.newReqID, genEDNS0Options(s.clientIP, 0))
|
reqs := buildReqMsgs(fqdn, option, s.newReqID, genEDNS0Options(s.clientIP, 0))
|
||||||
if err != nil {
|
|
||||||
errors.LogErrorInner(ctx, err, "failed to build dns query for ", fqdn)
|
|
||||||
if noResponseErrCh != nil {
|
|
||||||
if option.IPv4Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
if option.IPv6Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var deadline time.Time
|
var deadline time.Time
|
||||||
if d, ok := ctx.Deadline(); ok {
|
if d, ok := ctx.Deadline(); ok {
|
||||||
|
|||||||
@@ -131,6 +131,8 @@ func (s *ClassicNameServer) HandleResponse(ctx context.Context, packet *udp_prot
|
|||||||
newReq.msg = &newMsg
|
newReq.msg = &newMsg
|
||||||
s.addPendingRequest(&newReq)
|
s.addPendingRequest(&newReq)
|
||||||
b, _ := dns.PackMessage(newReq.msg)
|
b, _ := dns.PackMessage(newReq.msg)
|
||||||
|
copyDest := net.UDPDestination(s.address.Address, s.address.Port)
|
||||||
|
b.UDP = ©Dest
|
||||||
s.udpServer.Dispatch(toDnsContext(newReq.ctx, s.address.String()), *s.address, b)
|
s.udpServer.Dispatch(toDnsContext(newReq.ctx, s.address.String()), *s.address, b)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -161,19 +163,7 @@ func (s *ClassicNameServer) getCacheController() *CacheController {
|
|||||||
func (s *ClassicNameServer) sendQuery(ctx context.Context, noResponseErrCh chan<- error, fqdn string, option dns_feature.IPOption) {
|
func (s *ClassicNameServer) sendQuery(ctx context.Context, noResponseErrCh chan<- error, fqdn string, option dns_feature.IPOption) {
|
||||||
errors.LogInfo(ctx, s.Name(), " querying DNS for: ", fqdn)
|
errors.LogInfo(ctx, s.Name(), " querying DNS for: ", fqdn)
|
||||||
|
|
||||||
reqs, err := buildReqMsgs(fqdn, option, s.newReqID, genEDNS0Options(s.clientIP, 0))
|
reqs := buildReqMsgs(fqdn, option, s.newReqID, genEDNS0Options(s.clientIP, 0))
|
||||||
if err != nil {
|
|
||||||
errors.LogErrorInner(ctx, err, "failed to build dns query for ", fqdn)
|
|
||||||
if noResponseErrCh != nil {
|
|
||||||
if option.IPv4Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
if option.IPv6Enable {
|
|
||||||
noResponseErrCh <- err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, req := range reqs {
|
for _, req := range reqs {
|
||||||
udpReq := &udpDnsRequest{
|
udpReq := &udpDnsRequest{
|
||||||
@@ -189,6 +179,8 @@ func (s *ClassicNameServer) sendQuery(ctx context.Context, noResponseErrCh chan<
|
|||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
copyDest := net.UDPDestination(s.address.Address, s.address.Port)
|
||||||
|
b.UDP = ©Dest
|
||||||
s.udpServer.Dispatch(toDnsContext(ctx, s.address.String()), *s.address, b)
|
s.udpServer.Dispatch(toDnsContext(ctx, s.address.String()), *s.address, b)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,59 +0,0 @@
|
|||||||
package dns
|
|
||||||
|
|
||||||
import (
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
|
||||||
"github.com/xtls/xray-core/common/log"
|
|
||||||
xlua "github.com/xtls/xray-core/common/lua"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/features/dns"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
const scriptExecutionTimeout = 6 * time.Second
|
|
||||||
|
|
||||||
type scriptEngine struct {
|
|
||||||
dns *DNS
|
|
||||||
pool *xlua.Pool
|
|
||||||
}
|
|
||||||
|
|
||||||
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
|
||||||
program, err := xlua.CompileFile(path)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
e := &scriptEngine{dns: server}
|
|
||||||
e.pool, err = xlua.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory(
|
|
||||||
scriptExecutionTimeout*20,
|
|
||||||
func(L *lua.LState) {
|
|
||||||
geodata.RegisterLua(L)
|
|
||||||
log.RegisterLua(L)
|
|
||||||
server.registerLua(L)
|
|
||||||
},
|
|
||||||
func(L *lua.LState) error {
|
|
||||||
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
|
|
||||||
return errors.New("DNS script must define HandleDNSQuery(...)")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
|
|
||||||
return e, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *scriptEngine) close() {
|
|
||||||
e.pool.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) {
|
|
||||||
err = e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
|
||||||
var hookErr error
|
|
||||||
ips, ttl, hookErr = e.dns.callLuaHook(L, domain, option)
|
|
||||||
return hookErr
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
@@ -1,197 +0,0 @@
|
|||||||
package dns
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
|
||||||
)
|
|
||||||
|
|
||||||
type geoIPScriptNameServer struct {
|
|
||||||
name string
|
|
||||||
answers map[string]net.IP
|
|
||||||
ttl uint32
|
|
||||||
calls int
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *geoIPScriptNameServer) Name() string { return s.name }
|
|
||||||
func (s *geoIPScriptNameServer) IsDisableCache() bool { return true }
|
|
||||||
|
|
||||||
func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
if err := ctx.Err(); err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
s.calls++
|
|
||||||
ip, ok := s.answers[domain]
|
|
||||||
if !ok {
|
|
||||||
return nil, 0, featureDNS.ErrEmptyResponse
|
|
||||||
}
|
|
||||||
return []net.IP{ip}, s.ttl, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDNSScriptGeoIPFallback(t *testing.T) {
|
|
||||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
|
||||||
script := `
|
|
||||||
local servers = require("xray.dns").Servers
|
|
||||||
local us_ips = require("xray.geodata").BuildIPMatcher("geoip:us")
|
|
||||||
|
|
||||||
local by_id = {}
|
|
||||||
for _, server in ipairs(servers) do
|
|
||||||
by_id[server.ID] = server
|
|
||||||
end
|
|
||||||
assert(by_id.primary and by_id.fallback, "primary and fallback DNS servers are required")
|
|
||||||
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
local ips, ttl, err = by_id.primary:Query(domain, ipv4, ipv6, fake)
|
|
||||||
if not err and us_ips:AnyMatch(ips) then
|
|
||||||
return ips, ttl, nil
|
|
||||||
end
|
|
||||||
return by_id.fallback:Query(domain, ipv4, ipv6, fake)
|
|
||||||
end
|
|
||||||
`
|
|
||||||
scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua")
|
|
||||||
if err := os.WriteFile(scriptPath, []byte(script), 0o600); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
primary := &geoIPScriptNameServer{
|
|
||||||
name: "primary",
|
|
||||||
answers: map[string]net.IP{
|
|
||||||
"us.example": net.ParseIP("2001:4860:4860::8888"),
|
|
||||||
"other.example": net.ParseIP("127.0.0.1"),
|
|
||||||
},
|
|
||||||
ttl: 30,
|
|
||||||
}
|
|
||||||
fallback := &geoIPScriptNameServer{
|
|
||||||
name: "fallback",
|
|
||||||
answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")},
|
|
||||||
ttl: 60,
|
|
||||||
}
|
|
||||||
option := featureDNS.IPOption{IPv4Enable: true, IPv6Enable: true}
|
|
||||||
hosts, err := NewStaticHosts(nil)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
server := &DNS{
|
|
||||||
ctx: context.Background(),
|
|
||||||
hosts: hosts,
|
|
||||||
ipOption: &option,
|
|
||||||
scriptPath: scriptPath,
|
|
||||||
clients: []*Client{
|
|
||||||
{id: "primary", server: primary, ipOption: &option, timeoutMs: 2 * time.Second},
|
|
||||||
{id: "fallback", server: fallback, ipOption: &option, timeoutMs: 2 * time.Second},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
if err := server.Start(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
defer server.Close()
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
domain string
|
|
||||||
ip net.IP
|
|
||||||
ttl uint32
|
|
||||||
}{
|
|
||||||
{"Us.Example.", net.ParseIP("2001:4860:4860::8888"), 30},
|
|
||||||
{"other.example", net.ParseIP("9.9.9.9"), 60},
|
|
||||||
} {
|
|
||||||
ips, ttl, err := server.LookupIP(tc.domain, option)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("LookupIP(%q): %v", tc.domain, err)
|
|
||||||
}
|
|
||||||
if ttl != tc.ttl || len(ips) != 1 || !ips[0].Equal(tc.ip) {
|
|
||||||
t.Fatalf("LookupIP(%q) = %v, TTL %d; want %v, TTL %d", tc.domain, ips, ttl, tc.ip, tc.ttl)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if primary.calls != 2 || fallback.calls != 1 {
|
|
||||||
t.Fatalf("upstream calls: primary %d, fallback %d; want 2 and 1", primary.calls, fallback.calls)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDNSScriptRejectsInvalidStartup(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
script string
|
|
||||||
}{
|
|
||||||
{"syntax", "function HandleDNSQuery("},
|
|
||||||
{"missing hook", "value = 1"},
|
|
||||||
{"top-level error", `error("setup failed")`},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
path := filepath.Join(t.TempDir(), "script.lua")
|
|
||||||
if err := os.WriteFile(path, []byte(tc.script), 0o600); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
server := &DNS{ctx: context.Background(), scriptPath: path}
|
|
||||||
if err := server.Start(); err == nil {
|
|
||||||
t.Fatal("Start accepted an invalid DNS script")
|
|
||||||
}
|
|
||||||
if server.script != nil {
|
|
||||||
t.Fatal("Start retained a script engine after failure")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) {
|
|
||||||
path := filepath.Join(t.TempDir(), "script.lua")
|
|
||||||
script := `
|
|
||||||
local server = require("xray.dns").Servers[1]
|
|
||||||
local log = require("xray.log")
|
|
||||||
log.Info("DNS script loaded")
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
log.Debug("DNS query: ", domain)
|
|
||||||
if domain == "bad.example" then error("script failure") end
|
|
||||||
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
|
|
||||||
if err then log.Error("DNS failed: ", err) end
|
|
||||||
return ips, ttl, err
|
|
||||||
end
|
|
||||||
`
|
|
||||||
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
option := featureDNS.IPOption{IPv4Enable: true}
|
|
||||||
hosts, err := NewStaticHosts(nil)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
upstream := &geoIPScriptNameServer{
|
|
||||||
name: "FakeDNS",
|
|
||||||
answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")},
|
|
||||||
ttl: 30,
|
|
||||||
}
|
|
||||||
server := &DNS{
|
|
||||||
ctx: context.Background(),
|
|
||||||
hosts: hosts,
|
|
||||||
ipOption: &option,
|
|
||||||
scriptPath: path,
|
|
||||||
clients: []*Client{{id: "fake", server: upstream, ipOption: &option, timeoutMs: time.Second}},
|
|
||||||
}
|
|
||||||
if err := server.Start(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
defer server.Close()
|
|
||||||
|
|
||||||
if _, _, err := server.LookupIP("bad.example", option); err == nil || !strings.Contains(err.Error(), "script failure") {
|
|
||||||
t.Fatalf("hook failure = %v, want script failure", err)
|
|
||||||
}
|
|
||||||
if _, _, err := server.LookupIP("good.example", option); err != featureDNS.ErrEmptyResponse {
|
|
||||||
t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err)
|
|
||||||
}
|
|
||||||
if upstream.calls != 0 {
|
|
||||||
t.Fatalf("FakeDNS was queried without FakeEnable: %d calls", upstream.calls)
|
|
||||||
}
|
|
||||||
withFake := featureDNS.IPOption{IPv4Enable: true, FakeEnable: true}
|
|
||||||
ips, ttl, err := server.LookupIP("good.example", withFake)
|
|
||||||
if err != nil || ttl != 30 || len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.18.0.1")) {
|
|
||||||
t.Fatalf("FakeDNS with FakeEnable = %v, TTL %d, %v", ips, ttl, err)
|
|
||||||
}
|
|
||||||
if upstream.calls != 1 {
|
|
||||||
t.Fatalf("FakeDNS query count = %d, want 1", upstream.calls)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+41
-83
@@ -2,7 +2,6 @@ package geodata
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/tls"
|
|
||||||
go_errors "errors"
|
go_errors "errors"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -10,7 +9,6 @@ import (
|
|||||||
"path/filepath"
|
"path/filepath"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
utls "github.com/refraction-networking/utls"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/platform/filesystem"
|
"github.com/xtls/xray-core/common/platform/filesystem"
|
||||||
@@ -18,7 +16,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/utils"
|
"github.com/xtls/xray-core/common/utils"
|
||||||
"github.com/xtls/xray-core/features/routing"
|
"github.com/xtls/xray-core/features/routing"
|
||||||
"github.com/xtls/xray-core/transport/internet/tagged"
|
"github.com/xtls/xray-core/transport/internet/tagged"
|
||||||
"golang.org/x/net/http2"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const idleTimeout = 30 * time.Second
|
const idleTimeout = 30 * time.Second
|
||||||
@@ -29,9 +26,8 @@ type stage struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type downloader struct {
|
type downloader struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
httpClient *http.Client
|
client *http.Client
|
||||||
httpsClient *http.Client
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type idleConn struct {
|
type idleConn struct {
|
||||||
@@ -57,84 +53,52 @@ func (c *idleConn) Write(b []byte) (int, error) {
|
|||||||
|
|
||||||
func newDownloader(ctx context.Context, dispatcher routing.Dispatcher, outbound string) *downloader {
|
func newDownloader(ctx context.Context, dispatcher routing.Dispatcher, outbound string) *downloader {
|
||||||
return &downloader{
|
return &downloader{
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
httpClient: newClient(ctx, dispatcher, outbound, false),
|
client: newClient(ctx, dispatcher, outbound),
|
||||||
httpsClient: newClient(ctx, dispatcher, outbound, true),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func newClient(baseCtx context.Context, dispatcher routing.Dispatcher, outbound string, isHTTPS bool) *http.Client {
|
func newClient(baseCtx context.Context, dispatcher routing.Dispatcher, outbound string) *http.Client {
|
||||||
dial := func(ctx context.Context, network, address string) (net.Conn, error) {
|
return &http.Client{
|
||||||
var conn net.Conn
|
Transport: &http.Transport{
|
||||||
err := task.Run(ctx, func() error {
|
Proxy: nil,
|
||||||
if tagged.Dialer == nil {
|
DisableKeepAlives: true,
|
||||||
return errors.New("tagged dialer is not initialized")
|
DialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||||
}
|
var conn net.Conn
|
||||||
dest, err := net.ParseDestination(network + ":" + address)
|
err := task.Run(ctx, func() error {
|
||||||
if err != nil {
|
if tagged.Dialer == nil {
|
||||||
return errors.New("cannot understand address").Base(err)
|
return errors.New("tagged dialer is not initialized")
|
||||||
}
|
}
|
||||||
c, err := tagged.Dialer(baseCtx, dispatcher, dest, outbound)
|
dest, err := net.ParseDestination(network + ":" + address)
|
||||||
if err != nil {
|
|
||||||
return errors.New("cannot dial remote address ", dest).Base(err)
|
|
||||||
}
|
|
||||||
conn = c
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("cannot finish connection").Base(err)
|
|
||||||
}
|
|
||||||
return &idleConn{
|
|
||||||
Conn: conn,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
if isHTTPS {
|
|
||||||
return &http.Client{
|
|
||||||
Transport: &http2.Transport{
|
|
||||||
DialTLSContext: func(ctx context.Context, network string, address string, cfg *tls.Config) (net.Conn, error) {
|
|
||||||
conn, err := dial(ctx, network, address)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return errors.New("cannot understand address").Base(err)
|
||||||
}
|
}
|
||||||
host, _, _ := net.SplitHostPort(address)
|
c, err := tagged.Dialer(baseCtx, dispatcher, dest, outbound)
|
||||||
tlsConn := utls.UClient(conn, &utls.Config{ServerName: host}, utls.HelloChrome_Auto)
|
if err != nil {
|
||||||
handshakeCtx, cancel := context.WithTimeout(ctx, idleTimeout)
|
return errors.New("cannot dial remote address ", dest).Base(err)
|
||||||
defer cancel()
|
|
||||||
if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
|
|
||||||
conn.Close()
|
|
||||||
return nil, err
|
|
||||||
}
|
}
|
||||||
return tlsConn, nil
|
conn = c
|
||||||
},
|
return nil
|
||||||
},
|
})
|
||||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
if err != nil {
|
||||||
if req.URL.Scheme != "https" {
|
return nil, errors.New("cannot finish connection").Base(err)
|
||||||
return errors.New("redirected to non-https URL: ", req.URL.String())
|
|
||||||
}
|
}
|
||||||
if len(via) >= 10 {
|
return &idleConn{
|
||||||
return errors.New("stopped after 10 redirects")
|
Conn: conn,
|
||||||
}
|
}, nil
|
||||||
return nil
|
|
||||||
},
|
},
|
||||||
}
|
TLSHandshakeTimeout: idleTimeout,
|
||||||
} else {
|
ResponseHeaderTimeout: idleTimeout,
|
||||||
return &http.Client{
|
},
|
||||||
Transport: &http.Transport{
|
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||||
Proxy: nil,
|
if req.URL.Scheme != "https" {
|
||||||
DisableKeepAlives: true,
|
return errors.New("redirected to non-https URL: ", req.URL.String())
|
||||||
DialContext: dial,
|
}
|
||||||
ResponseHeaderTimeout: idleTimeout,
|
if len(via) >= 10 {
|
||||||
},
|
return errors.New("stopped after 10 redirects")
|
||||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
}
|
||||||
if req.URL.Scheme != "https" {
|
return nil
|
||||||
return errors.New("redirected to non-https URL: ", req.URL.String())
|
},
|
||||||
}
|
|
||||||
if len(via) >= 10 {
|
|
||||||
return errors.New("stopped after 10 redirects")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -196,13 +160,7 @@ func (d *downloader) fetch(rawURL string, writer io.Writer) error {
|
|||||||
}
|
}
|
||||||
utils.TryDefaultHeadersWith(req.Header, "nav")
|
utils.TryDefaultHeadersWith(req.Header, "nav")
|
||||||
|
|
||||||
var client *http.Client
|
resp, err := d.client.Do(req)
|
||||||
if req.URL.Scheme == "https" {
|
|
||||||
client = d.httpsClient
|
|
||||||
} else {
|
|
||||||
client = d.httpClient
|
|
||||||
}
|
|
||||||
resp, err := client.Do(req)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-6
@@ -89,10 +89,10 @@ func (g *Instance) startInternal() error {
|
|||||||
g.active = true
|
g.active = true
|
||||||
|
|
||||||
if err := g.initAccessLogger(); err != nil {
|
if err := g.initAccessLogger(); err != nil {
|
||||||
return errors.New("failed to initialize access logger").Base(err)
|
return errors.New("failed to initialize access logger").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
if err := g.initErrorLogger(); err != nil {
|
if err := g.initErrorLogger(); err != nil {
|
||||||
return errors.New("failed to initialize error logger").Base(err)
|
return errors.New("failed to initialize error logger").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -141,10 +141,6 @@ func (g *Instance) Handle(msg log.Message) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *Instance) Severity() log.Severity {
|
|
||||||
return g.config.ErrorLogLevel
|
|
||||||
}
|
|
||||||
|
|
||||||
// Close implements common.Closable.Close().
|
// Close implements common.Closable.Close().
|
||||||
func (g *Instance) Close() error {
|
func (g *Instance) Close() error {
|
||||||
errors.LogDebug(context.Background(), "Logger closing")
|
errors.LogDebug(context.Background(), "Logger closing")
|
||||||
|
|||||||
+59
-150
@@ -2,18 +2,15 @@ package metrics
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
|
||||||
stderrors "errors"
|
|
||||||
"expvar"
|
"expvar"
|
||||||
stdnet "net"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/pprof"
|
_ "net/http/pprof"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/app/observatory"
|
"github.com/xtls/xray-core/app/observatory"
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
xnet "github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/signal/done"
|
"github.com/xtls/xray-core/common/signal/done"
|
||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
"github.com/xtls/xray-core/features/extension"
|
"github.com/xtls/xray-core/features/extension"
|
||||||
@@ -24,17 +21,15 @@ import (
|
|||||||
type MetricsHandler struct {
|
type MetricsHandler struct {
|
||||||
ohm outbound.Manager
|
ohm outbound.Manager
|
||||||
statsManager feature_stats.Manager
|
statsManager feature_stats.Manager
|
||||||
ctx context.Context
|
observatory extension.Observatory
|
||||||
tag string
|
tag string
|
||||||
listen string
|
listen string
|
||||||
tcpListener xnet.Listener
|
tcpListener net.Listener
|
||||||
listener *OutboundListener
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewMetricsHandler creates a new MetricsHandler based on the given config.
|
// NewMetricsHandler creates a new MetricsHandler based on the given config.
|
||||||
func NewMetricsHandler(ctx context.Context, config *Config) (*MetricsHandler, error) {
|
func NewMetricsHandler(ctx context.Context, config *Config) (*MetricsHandler, error) {
|
||||||
c := &MetricsHandler{
|
c := &MetricsHandler{
|
||||||
ctx: ctx,
|
|
||||||
tag: config.Tag,
|
tag: config.Tag,
|
||||||
listen: config.Listen,
|
listen: config.Listen,
|
||||||
}
|
}
|
||||||
@@ -42,6 +37,46 @@ func NewMetricsHandler(ctx context.Context, config *Config) (*MetricsHandler, er
|
|||||||
c.statsManager = sm
|
c.statsManager = sm
|
||||||
c.ohm = om
|
c.ohm = om
|
||||||
}))
|
}))
|
||||||
|
expvar.Publish("stats", expvar.Func(func() interface{} {
|
||||||
|
resp := map[string]map[string]map[string]int64{
|
||||||
|
"inbound": {},
|
||||||
|
"outbound": {},
|
||||||
|
"user": {},
|
||||||
|
}
|
||||||
|
c.statsManager.VisitCounters(func(name string, counter feature_stats.Counter) bool {
|
||||||
|
nameSplit := strings.Split(name, ">>>")
|
||||||
|
typeName, tagOrUser, direction := nameSplit[0], nameSplit[1], nameSplit[3]
|
||||||
|
if item, found := resp[typeName][tagOrUser]; found {
|
||||||
|
item[direction] = counter.Value()
|
||||||
|
} else {
|
||||||
|
resp[typeName][tagOrUser] = map[string]int64{
|
||||||
|
direction: counter.Value(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
return resp
|
||||||
|
}))
|
||||||
|
expvar.Publish("observatory", expvar.Func(func() interface{} {
|
||||||
|
if c.observatory == nil {
|
||||||
|
common.Must(core.RequireFeatures(ctx, func(observatory extension.Observatory) error {
|
||||||
|
c.observatory = observatory
|
||||||
|
return nil
|
||||||
|
}))
|
||||||
|
if c.observatory == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
resp := map[string]*observatory.OutboundStatus{}
|
||||||
|
if o, err := c.observatory.GetObservation(context.Background()); err != nil {
|
||||||
|
return err
|
||||||
|
} else {
|
||||||
|
for _, x := range o.(*observatory.ObservationResult).GetStatus() {
|
||||||
|
resp[x.OutboundTag] = x
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return resp
|
||||||
|
}))
|
||||||
return c, nil
|
return c, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -50,172 +85,46 @@ func (p *MetricsHandler) Type() interface{} {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (p *MetricsHandler) Start() error {
|
func (p *MetricsHandler) Start() error {
|
||||||
handler := p.httpHandler()
|
|
||||||
|
|
||||||
// direct listen a port if listen is set
|
// direct listen a port if listen is set
|
||||||
if p.listen != "" {
|
if p.listen != "" {
|
||||||
TCPlistener, err := xnet.Listen("tcp", p.listen)
|
TCPlistener, err := net.Listen("tcp", p.listen)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
p.tcpListener = TCPlistener
|
p.tcpListener = TCPlistener
|
||||||
errors.LogInfo(context.Background(), "Metrics server listening on ", p.listen)
|
errors.LogInfo(context.Background(), "Metrics server listening on ", p.listen)
|
||||||
|
|
||||||
go p.serve(TCPlistener, handler)
|
go func() {
|
||||||
}
|
if err := http.Serve(TCPlistener, http.DefaultServeMux); err != nil {
|
||||||
|
errors.LogErrorInner(context.Background(), err, "failed to start metrics server")
|
||||||
if p.tag == "" {
|
}
|
||||||
if p.tcpListener == nil {
|
}()
|
||||||
return errors.New("metrics must have a tag or listen address")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
listener := &OutboundListener{
|
listener := &OutboundListener{
|
||||||
buffer: make(chan xnet.Conn, 4),
|
buffer: make(chan net.Conn, 4),
|
||||||
done: done.New(),
|
done: done.New(),
|
||||||
}
|
}
|
||||||
p.listener = listener
|
|
||||||
|
|
||||||
go p.serve(listener, handler)
|
go func() {
|
||||||
|
if err := http.Serve(listener, http.DefaultServeMux); err != nil {
|
||||||
|
errors.LogErrorInner(context.Background(), err, "failed to start metrics server")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
if err := p.ohm.RemoveHandler(context.Background(), p.tag); err != nil {
|
if err := p.ohm.RemoveHandler(context.Background(), p.tag); err != nil {
|
||||||
errors.LogInfo(context.Background(), "failed to remove existing handler")
|
errors.LogInfo(context.Background(), "failed to remove existing handler")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := p.ohm.AddHandler(context.Background(), &Outbound{
|
return p.ohm.AddHandler(context.Background(), &Outbound{
|
||||||
tag: p.tag,
|
tag: p.tag,
|
||||||
listener: listener,
|
listener: listener,
|
||||||
}); err != nil {
|
})
|
||||||
if closeErr := p.Close(); closeErr != nil {
|
|
||||||
errors.LogErrorInner(context.Background(), closeErr, "failed to close metrics server after start failure")
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *MetricsHandler) Close() error {
|
func (p *MetricsHandler) Close() error {
|
||||||
var errs []error
|
return nil
|
||||||
if p.tcpListener != nil {
|
|
||||||
errs = append(errs, p.tcpListener.Close())
|
|
||||||
p.tcpListener = nil
|
|
||||||
}
|
|
||||||
if p.listener != nil {
|
|
||||||
errs = append(errs, p.listener.Close())
|
|
||||||
p.listener = nil
|
|
||||||
}
|
|
||||||
if p.ohm != nil && p.tag != "" {
|
|
||||||
if err := p.ohm.RemoveHandler(context.Background(), p.tag); err != nil {
|
|
||||||
errors.LogInfo(context.Background(), "failed to remove metrics handler")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return errors.Combine(errs...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *MetricsHandler) serve(listener xnet.Listener, handler http.Handler) {
|
|
||||||
if err := http.Serve(listener, handler); err != nil && !isClosedListenerError(err) {
|
|
||||||
errors.LogErrorInner(context.Background(), err, "failed to start metrics server")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func isClosedListenerError(err error) bool {
|
|
||||||
if err == nil {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
if stderrors.Is(err, stdnet.ErrClosed) || stderrors.Is(err, http.ErrServerClosed) {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
errText := err.Error()
|
|
||||||
return strings.Contains(errText, "listen closed") ||
|
|
||||||
strings.Contains(errText, "use of closed network connection")
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *MetricsHandler) httpHandler() http.Handler {
|
|
||||||
mux := http.NewServeMux()
|
|
||||||
mux.HandleFunc("/debug/vars", p.handleDebugVars)
|
|
||||||
mux.HandleFunc("/debug/pprof/", pprof.Index)
|
|
||||||
mux.HandleFunc("/debug/pprof/cmdline", pprof.Cmdline)
|
|
||||||
mux.HandleFunc("/debug/pprof/profile", pprof.Profile)
|
|
||||||
mux.HandleFunc("/debug/pprof/symbol", pprof.Symbol)
|
|
||||||
mux.HandleFunc("/debug/pprof/trace", pprof.Trace)
|
|
||||||
return mux
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *MetricsHandler) handleDebugVars(w http.ResponseWriter, r *http.Request) {
|
|
||||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
|
||||||
vars := map[string]json.RawMessage{}
|
|
||||||
expvar.Do(func(kv expvar.KeyValue) {
|
|
||||||
value := json.RawMessage(kv.Value.String())
|
|
||||||
if !json.Valid(value) {
|
|
||||||
value = json.RawMessage("null")
|
|
||||||
}
|
|
||||||
vars[kv.Key] = value
|
|
||||||
})
|
|
||||||
vars["stats"] = marshalJSON(p.stats())
|
|
||||||
vars["observatory"] = marshalJSON(p.observatoryStatus())
|
|
||||||
|
|
||||||
payload, err := json.Marshal(vars)
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
w.Write(payload)
|
|
||||||
}
|
|
||||||
|
|
||||||
func marshalJSON(value interface{}) json.RawMessage {
|
|
||||||
data, err := json.Marshal(value)
|
|
||||||
if err != nil {
|
|
||||||
return json.RawMessage("null")
|
|
||||||
}
|
|
||||||
return data
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *MetricsHandler) stats() map[string]map[string]map[string]int64 {
|
|
||||||
resp := map[string]map[string]map[string]int64{
|
|
||||||
"inbound": {},
|
|
||||||
"outbound": {},
|
|
||||||
"user": {},
|
|
||||||
}
|
|
||||||
p.statsManager.VisitCounters(func(name string, counter feature_stats.Counter) bool {
|
|
||||||
nameSplit := strings.Split(name, ">>>")
|
|
||||||
if len(nameSplit) < 4 {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
typeName, tagOrUser, direction := nameSplit[0], nameSplit[1], nameSplit[3]
|
|
||||||
items, found := resp[typeName]
|
|
||||||
if !found {
|
|
||||||
items = map[string]map[string]int64{}
|
|
||||||
resp[typeName] = items
|
|
||||||
}
|
|
||||||
if item, found := items[tagOrUser]; found {
|
|
||||||
item[direction] = counter.Value()
|
|
||||||
} else {
|
|
||||||
items[tagOrUser] = map[string]int64{
|
|
||||||
direction: counter.Value(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
return resp
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *MetricsHandler) observatoryStatus() interface{} {
|
|
||||||
feature := core.MustFromContext(p.ctx).GetFeature(extension.ObservatoryType())
|
|
||||||
if feature == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
observatoryFeature := feature.(extension.Observatory)
|
|
||||||
resp := map[string]*observatory.OutboundStatus{}
|
|
||||||
if o, err := observatoryFeature.GetObservation(context.Background()); err != nil {
|
|
||||||
return err
|
|
||||||
} else {
|
|
||||||
for _, x := range o.(*observatory.ObservationResult).GetStatus() {
|
|
||||||
resp[x.OutboundTag] = x
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return resp
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
|
|||||||
@@ -1,161 +0,0 @@
|
|||||||
package metrics
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
stdnet "net"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/app/dispatcher"
|
|
||||||
"github.com/xtls/xray-core/app/proxyman"
|
|
||||||
_ "github.com/xtls/xray-core/app/proxyman/inbound"
|
|
||||||
_ "github.com/xtls/xray-core/app/proxyman/outbound"
|
|
||||||
appstats "github.com/xtls/xray-core/app/stats"
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
|
||||||
"github.com/xtls/xray-core/core"
|
|
||||||
feature_outbound "github.com/xtls/xray-core/features/outbound"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestMetricsCanRestartInSameProcess(t *testing.T) {
|
|
||||||
for i := 0; i < 2; i++ {
|
|
||||||
server := startMetricsTestServer(t)
|
|
||||||
readMetricsVars(t, server)
|
|
||||||
readMetricsPprof(t, server)
|
|
||||||
if err := server.Close(); err != nil {
|
|
||||||
t.Fatalf("failed to close metrics server: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMetricsCanRunMultipleInstancesInSameProcess(t *testing.T) {
|
|
||||||
server1 := startMetricsTestServer(t)
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = server1.Close()
|
|
||||||
})
|
|
||||||
server2 := startMetricsTestServer(t)
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = server2.Close()
|
|
||||||
})
|
|
||||||
|
|
||||||
readMetricsVars(t, server1)
|
|
||||||
readMetricsVars(t, server2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMetricsListenOnlyWithoutTagDoesNotRegisterOutbound(t *testing.T) {
|
|
||||||
listen := pickMetricsListenAddress(t)
|
|
||||||
server := startMetricsTestServerWithMetricsConfig(t, &Config{
|
|
||||||
Listen: listen,
|
|
||||||
})
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = server.Close()
|
|
||||||
})
|
|
||||||
|
|
||||||
response, err := http.Get("http://" + listen + "/debug/vars")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to read listen-only metrics: %v", err)
|
|
||||||
}
|
|
||||||
defer response.Body.Close()
|
|
||||||
if response.StatusCode != http.StatusOK {
|
|
||||||
t.Fatalf("unexpected listen-only metrics status: %d", response.StatusCode)
|
|
||||||
}
|
|
||||||
|
|
||||||
outboundManager := server.GetFeature(feature_outbound.ManagerType()).(feature_outbound.Manager)
|
|
||||||
if handlers := outboundManager.ListHandlers(context.Background()); len(handlers) != 0 {
|
|
||||||
t.Fatalf("listen-only metrics registered outbound handlers: got %d, want 0", len(handlers))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func startMetricsTestServer(t *testing.T) *core.Instance {
|
|
||||||
return startMetricsTestServerWithMetricsConfig(t, &Config{
|
|
||||||
Tag: "metrics_out",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func startMetricsTestServerWithMetricsConfig(t *testing.T, metricsConfig *Config) *core.Instance {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
server, err := core.New(metricsTestConfig(metricsConfig))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to create metrics server: %v", err)
|
|
||||||
}
|
|
||||||
if err := server.Start(); err != nil {
|
|
||||||
_ = server.Close()
|
|
||||||
t.Fatalf("failed to start metrics server: %v", err)
|
|
||||||
}
|
|
||||||
return server
|
|
||||||
}
|
|
||||||
|
|
||||||
func metricsTestConfig(metricsConfig *Config) *core.Config {
|
|
||||||
return &core.Config{
|
|
||||||
App: []*serial.TypedMessage{
|
|
||||||
serial.ToTypedMessage(&dispatcher.Config{}),
|
|
||||||
serial.ToTypedMessage(&proxyman.InboundConfig{}),
|
|
||||||
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
|
|
||||||
serial.ToTypedMessage(&appstats.Config{}),
|
|
||||||
serial.ToTypedMessage(metricsConfig),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func pickMetricsListenAddress(t *testing.T) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
listener, err := stdnet.Listen("tcp", "127.0.0.1:0")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to pick metrics listen address: %v", err)
|
|
||||||
}
|
|
||||||
defer listener.Close()
|
|
||||||
return listener.Addr().String()
|
|
||||||
}
|
|
||||||
|
|
||||||
func readMetricsVars(t *testing.T, server *core.Instance) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
recorder := httptest.NewRecorder()
|
|
||||||
metricsHandler(t, server).httpHandler().ServeHTTP(
|
|
||||||
recorder,
|
|
||||||
httptest.NewRequest(http.MethodGet, "/debug/vars", nil),
|
|
||||||
)
|
|
||||||
|
|
||||||
if recorder.Code != http.StatusOK {
|
|
||||||
t.Fatalf("unexpected metrics vars status: %d", recorder.Code)
|
|
||||||
}
|
|
||||||
|
|
||||||
var payload map[string]interface{}
|
|
||||||
if err := json.NewDecoder(recorder.Body).Decode(&payload); err != nil {
|
|
||||||
t.Fatalf("failed to decode metrics vars: %v", err)
|
|
||||||
}
|
|
||||||
if _, found := payload["stats"]; !found {
|
|
||||||
t.Fatal("metrics vars missing stats")
|
|
||||||
}
|
|
||||||
if _, found := payload["observatory"]; !found {
|
|
||||||
t.Fatal("metrics vars missing observatory")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func readMetricsPprof(t *testing.T, server *core.Instance) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
recorder := httptest.NewRecorder()
|
|
||||||
metricsHandler(t, server).httpHandler().ServeHTTP(
|
|
||||||
recorder,
|
|
||||||
httptest.NewRequest(http.MethodGet, "/debug/pprof/goroutine?debug=1", nil),
|
|
||||||
)
|
|
||||||
|
|
||||||
if recorder.Code != http.StatusOK {
|
|
||||||
t.Fatalf("unexpected metrics pprof status: %d", recorder.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func metricsHandler(t *testing.T, server *core.Instance) *MetricsHandler {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
feature := server.GetFeature((*MetricsHandler)(nil))
|
|
||||||
handler, ok := feature.(*MetricsHandler)
|
|
||||||
if !ok || handler == nil {
|
|
||||||
t.Fatal("metrics handler not registered")
|
|
||||||
}
|
|
||||||
return handler
|
|
||||||
}
|
|
||||||
@@ -2,6 +2,7 @@ package burst
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/app/observatory"
|
"github.com/xtls/xray-core/app/observatory"
|
||||||
@@ -71,6 +72,7 @@ func (o *Observer) Start() error {
|
|||||||
o.hp.StartScheduler(func() ([]string, error) {
|
o.hp.StartScheduler(func() ([]string, error) {
|
||||||
hs, ok := o.ohm.(outbound.HandlerSelector)
|
hs, ok := o.ohm.(outbound.HandlerSelector)
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|
||||||
return nil, errors.New("outbound.Manager is not a HandlerSelector")
|
return nil, errors.New("outbound.Manager is not a HandlerSelector")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/dice"
|
"github.com/xtls/xray-core/common/dice"
|
||||||
@@ -25,12 +24,11 @@ type HealthPingSettings struct {
|
|||||||
|
|
||||||
// HealthPing is the health checker for balancers
|
// HealthPing is the health checker for balancers
|
||||||
type HealthPing struct {
|
type HealthPing struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancelCtx context.CancelFunc
|
dispatcher routing.Dispatcher
|
||||||
cancelPending atomic.Pointer[context.CancelFunc]
|
access sync.Mutex
|
||||||
dispatcher routing.Dispatcher
|
ticker *time.Ticker
|
||||||
access sync.Mutex
|
tickerClose chan struct{}
|
||||||
ticker *time.Ticker
|
|
||||||
|
|
||||||
Settings *HealthPingSettings
|
Settings *HealthPingSettings
|
||||||
Results map[string]*HealthPingRTTS
|
Results map[string]*HealthPingRTTS
|
||||||
@@ -64,10 +62,10 @@ func NewHealthPing(ctx context.Context, dispatcher routing.Dispatcher, config *H
|
|||||||
settings.Destination = "https://connectivitycheck.gstatic.com/generate_204"
|
settings.Destination = "https://connectivitycheck.gstatic.com/generate_204"
|
||||||
}
|
}
|
||||||
if settings.Interval == 0 {
|
if settings.Interval == 0 {
|
||||||
settings.Interval = 1 * time.Minute
|
settings.Interval = time.Duration(1) * time.Minute
|
||||||
} else if settings.Interval < 10*time.Second {
|
} else if settings.Interval < 10 {
|
||||||
errors.LogWarning(ctx, "health check interval is too small, 10s is applied")
|
errors.LogWarning(ctx, "health check interval is too small, 10s is applied")
|
||||||
settings.Interval = 10 * time.Second
|
settings.Interval = time.Duration(10) * time.Second
|
||||||
}
|
}
|
||||||
if settings.SamplingCount <= 0 {
|
if settings.SamplingCount <= 0 {
|
||||||
settings.SamplingCount = 10
|
settings.SamplingCount = 10
|
||||||
@@ -75,12 +73,10 @@ func NewHealthPing(ctx context.Context, dispatcher routing.Dispatcher, config *H
|
|||||||
if settings.Timeout <= 0 {
|
if settings.Timeout <= 0 {
|
||||||
// results are saved after all health pings finish,
|
// results are saved after all health pings finish,
|
||||||
// a larger timeout could possibly makes checks run longer
|
// a larger timeout could possibly makes checks run longer
|
||||||
settings.Timeout = 5 * time.Second
|
settings.Timeout = time.Duration(5) * time.Second
|
||||||
}
|
}
|
||||||
ctx, cancel := context.WithCancel(ctx)
|
|
||||||
return &HealthPing{
|
return &HealthPing{
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
cancelCtx: cancel,
|
|
||||||
dispatcher: dispatcher,
|
dispatcher: dispatcher,
|
||||||
Settings: settings,
|
Settings: settings,
|
||||||
Results: nil,
|
Results: nil,
|
||||||
@@ -94,9 +90,9 @@ func (h *HealthPing) StartScheduler(selector func() ([]string, error)) {
|
|||||||
}
|
}
|
||||||
interval := h.Settings.Interval * time.Duration(h.Settings.SamplingCount)
|
interval := h.Settings.Interval * time.Duration(h.Settings.SamplingCount)
|
||||||
ticker := time.NewTicker(interval)
|
ticker := time.NewTicker(interval)
|
||||||
|
tickerClose := make(chan struct{})
|
||||||
h.ticker = ticker
|
h.ticker = ticker
|
||||||
|
h.tickerClose = tickerClose
|
||||||
// init run to get a fast check result
|
|
||||||
go func() {
|
go func() {
|
||||||
tags, err := selector()
|
tags, err := selector()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -114,20 +110,13 @@ func (h *HealthPing) StartScheduler(selector func() ([]string, error)) {
|
|||||||
errors.LogWarning(h.ctx, "error select outbounds for scheduled health check: ", err)
|
errors.LogWarning(h.ctx, "error select outbounds for scheduled health check: ", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
subCtx, cancel := context.WithCancel(h.ctx)
|
h.doCheck(tags, interval, h.Settings.SamplingCount)
|
||||||
old := h.cancelPending.Swap(&cancel)
|
|
||||||
if old != nil {
|
|
||||||
errors.LogDebug(h.ctx, "scheduled health check not finished before next round, canceling previous one")
|
|
||||||
(*old)()
|
|
||||||
}
|
|
||||||
h.doCheck(subCtx, tags, interval, h.Settings.SamplingCount)
|
|
||||||
h.cancelPending.CompareAndSwap(&cancel, nil)
|
|
||||||
h.Cleanup(tags)
|
h.Cleanup(tags)
|
||||||
}()
|
}()
|
||||||
select {
|
select {
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
continue
|
continue
|
||||||
case <-h.ctx.Done():
|
case <-tickerClose:
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -141,7 +130,8 @@ func (h *HealthPing) StopScheduler() {
|
|||||||
}
|
}
|
||||||
h.ticker.Stop()
|
h.ticker.Stop()
|
||||||
h.ticker = nil
|
h.ticker = nil
|
||||||
h.cancelCtx()
|
close(h.tickerClose)
|
||||||
|
h.tickerClose = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check implements the HealthChecker
|
// Check implements the HealthChecker
|
||||||
@@ -150,7 +140,7 @@ func (h *HealthPing) Check(tags []string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
errors.LogInfo(h.ctx, "perform one-time health check for tags ", tags)
|
errors.LogInfo(h.ctx, "perform one-time health check for tags ", tags)
|
||||||
h.doCheck(h.ctx, tags, 0, 1)
|
h.doCheck(tags, 0, 1)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -161,14 +151,13 @@ type rtt struct {
|
|||||||
|
|
||||||
// doCheck performs the 'rounds' amount checks in given 'duration'. You should make
|
// doCheck performs the 'rounds' amount checks in given 'duration'. You should make
|
||||||
// sure all tags are valid for current balancer
|
// sure all tags are valid for current balancer
|
||||||
// cancel ctx will stop all pending checks
|
func (h *HealthPing) doCheck(tags []string, duration time.Duration, rounds int) {
|
||||||
func (h *HealthPing) doCheck(ctx context.Context, tags []string, duration time.Duration, rounds int) {
|
|
||||||
count := len(tags) * rounds
|
count := len(tags) * rounds
|
||||||
if count == 0 {
|
if count == 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
ch := make(chan *rtt, count)
|
ch := make(chan *rtt, count)
|
||||||
timers := make([]*time.Timer, 0, count)
|
|
||||||
for _, tag := range tags {
|
for _, tag := range tags {
|
||||||
handler := tag
|
handler := tag
|
||||||
client := newPingClient(
|
client := newPingClient(
|
||||||
@@ -183,7 +172,7 @@ func (h *HealthPing) doCheck(ctx context.Context, tags []string, duration time.D
|
|||||||
if duration > 0 {
|
if duration > 0 {
|
||||||
delay = time.Duration(dice.RollInt63n(int64(duration)))
|
delay = time.Duration(dice.RollInt63n(int64(duration)))
|
||||||
}
|
}
|
||||||
timers = append(timers, time.AfterFunc(delay, func() {
|
time.AfterFunc(delay, func() {
|
||||||
errors.LogDebug(h.ctx, "checking ", handler)
|
errors.LogDebug(h.ctx, "checking ", handler)
|
||||||
delay, err := client.MeasureDelay(h.Settings.HttpMethod)
|
delay, err := client.MeasureDelay(h.Settings.HttpMethod)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
@@ -211,21 +200,14 @@ func (h *HealthPing) doCheck(ctx context.Context, tags []string, duration time.D
|
|||||||
handler: handler,
|
handler: handler,
|
||||||
value: rttFailed,
|
value: rttFailed,
|
||||||
}
|
}
|
||||||
}))
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for i := 0; i < count; i++ {
|
for i := 0; i < count; i++ {
|
||||||
select {
|
rtt := <-ch
|
||||||
case rtt := <-ch:
|
if rtt.value > 0 {
|
||||||
if rtt.value > 0 {
|
// should not put results when network is down
|
||||||
// should not put results when network is down
|
h.PutResult(rtt.handler, rtt.value)
|
||||||
h.PutResult(rtt.handler, rtt.value)
|
|
||||||
}
|
|
||||||
case <-ctx.Done():
|
|
||||||
for _, timer := range timers {
|
|
||||||
timer.Stop()
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ func (h *HealthPingRTTS) Put(d time.Duration) {
|
|||||||
if h.rtts == nil {
|
if h.rtts == nil {
|
||||||
h.rtts = make([]*pingRTT, h.cap)
|
h.rtts = make([]*pingRTT, h.cap)
|
||||||
for i := 0; i < h.cap; i++ {
|
for i := 0; i < h.cap; i++ {
|
||||||
h.rtts[i] = &pingRTT{value: rttUntested}
|
h.rtts[i] = &pingRTT{}
|
||||||
}
|
}
|
||||||
h.idx = -1
|
h.idx = -1
|
||||||
}
|
}
|
||||||
@@ -88,7 +88,7 @@ func (h *HealthPingRTTS) getStatistics() *HealthPingStats {
|
|||||||
validRTTs := make([]time.Duration, 0)
|
validRTTs := make([]time.Duration, 0)
|
||||||
for _, rtt := range h.rtts {
|
for _, rtt := range h.rtts {
|
||||||
switch {
|
switch {
|
||||||
case rtt.value == rttUntested || time.Since(rtt.time) > h.validity:
|
case rtt.value == 0 || time.Since(rtt.time) > h.validity:
|
||||||
continue
|
continue
|
||||||
case rtt.value == rttFailed:
|
case rtt.value == rttFailed:
|
||||||
stats.Fail++
|
stats.Fail++
|
||||||
|
|||||||
@@ -78,12 +78,6 @@ func (o *Observer) background() {
|
|||||||
sleepTime = time.Duration(o.config.ProbeInterval)
|
sleepTime = time.Duration(o.config.ProbeInterval)
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(outbounds) == 0 {
|
|
||||||
errors.LogWarning(o.ctx, "no outbound matches subjectSelector ", o.config.SubjectSelector)
|
|
||||||
time.Sleep(sleepTime)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if !o.config.EnableConcurrency {
|
if !o.config.EnableConcurrency {
|
||||||
sort.Strings(outbounds)
|
sort.Strings(outbounds)
|
||||||
for _, v := range outbounds {
|
for _, v := range outbounds {
|
||||||
@@ -192,7 +186,7 @@ func (o *Observer) probe(outbound string) ProbeResult {
|
|||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
errorMessage := "the outbound " + outbound + " is dead: GET request failed:" + err.Error() + "with outbound handler report underlying connection failed"
|
var errorMessage = "the outbound " + outbound + " is dead: GET request failed:" + err.Error() + "with outbound handler report underlying connection failed"
|
||||||
errors.LogInfoInner(o.ctx, errorCollectorForRequest.UnderlyingError(), errorMessage)
|
errors.LogInfoInner(o.ctx, errorCollectorForRequest.UnderlyingError(), errorMessage)
|
||||||
return ProbeResult{Alive: false, LastErrorReason: errorMessage}
|
return ProbeResult{Alive: false, LastErrorReason: errorMessage}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -138,7 +138,7 @@ func (s *handlerServer) GetInboundUsers(ctx context.Context, request *GetInbound
|
|||||||
if len(request.Email) > 0 {
|
if len(request.Email) > 0 {
|
||||||
return &GetInboundUserResponse{Users: []*protocol.User{protocol.ToProtoUser(um.GetUser(ctx, request.Email))}}, nil
|
return &GetInboundUserResponse{Users: []*protocol.User{protocol.ToProtoUser(um.GetUser(ctx, request.Email))}}, nil
|
||||||
}
|
}
|
||||||
result := make([]*protocol.User, 0, 100)
|
var result = make([]*protocol.User, 0, 100)
|
||||||
users := um.GetUsers(ctx)
|
users := um.GetUsers(ctx)
|
||||||
for _, u := range users {
|
for _, u := range users {
|
||||||
result = append(result, protocol.ToProtoUser(u))
|
result = append(result, protocol.ToProtoUser(u))
|
||||||
|
|||||||
+22
-11
@@ -330,6 +330,7 @@ type SenderConfig struct {
|
|||||||
// Send traffic through the given IP. Only IP is allowed.
|
// Send traffic through the given IP. Only IP is allowed.
|
||||||
Via *net.IPOrDomain `protobuf:"bytes,1,opt,name=via,proto3" json:"via,omitempty"`
|
Via *net.IPOrDomain `protobuf:"bytes,1,opt,name=via,proto3" json:"via,omitempty"`
|
||||||
StreamSettings *internet.StreamConfig `protobuf:"bytes,2,opt,name=stream_settings,json=streamSettings,proto3" json:"stream_settings,omitempty"`
|
StreamSettings *internet.StreamConfig `protobuf:"bytes,2,opt,name=stream_settings,json=streamSettings,proto3" json:"stream_settings,omitempty"`
|
||||||
|
ProxySettings *internet.ProxyConfig `protobuf:"bytes,3,opt,name=proxy_settings,json=proxySettings,proto3" json:"proxy_settings,omitempty"`
|
||||||
MultiplexSettings *MultiplexingConfig `protobuf:"bytes,4,opt,name=multiplex_settings,json=multiplexSettings,proto3" json:"multiplex_settings,omitempty"`
|
MultiplexSettings *MultiplexingConfig `protobuf:"bytes,4,opt,name=multiplex_settings,json=multiplexSettings,proto3" json:"multiplex_settings,omitempty"`
|
||||||
ViaCidr string `protobuf:"bytes,5,opt,name=via_cidr,json=viaCidr,proto3" json:"via_cidr,omitempty"`
|
ViaCidr string `protobuf:"bytes,5,opt,name=via_cidr,json=viaCidr,proto3" json:"via_cidr,omitempty"`
|
||||||
TargetStrategy internet.DomainStrategy `protobuf:"varint,6,opt,name=target_strategy,json=targetStrategy,proto3,enum=xray.transport.internet.DomainStrategy" json:"target_strategy,omitempty"`
|
TargetStrategy internet.DomainStrategy `protobuf:"varint,6,opt,name=target_strategy,json=targetStrategy,proto3,enum=xray.transport.internet.DomainStrategy" json:"target_strategy,omitempty"`
|
||||||
@@ -381,6 +382,13 @@ func (x *SenderConfig) GetStreamSettings() *internet.StreamConfig {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (x *SenderConfig) GetProxySettings() *internet.ProxyConfig {
|
||||||
|
if x != nil {
|
||||||
|
return x.ProxySettings
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (x *SenderConfig) GetMultiplexSettings() *MultiplexingConfig {
|
func (x *SenderConfig) GetMultiplexSettings() *MultiplexingConfig {
|
||||||
if x != nil {
|
if x != nil {
|
||||||
return x.MultiplexSettings
|
return x.MultiplexSettings
|
||||||
@@ -498,13 +506,14 @@ const file_app_proxyman_config_proto_rawDesc = "" +
|
|||||||
"\x03tag\x18\x01 \x01(\tR\x03tag\x12M\n" +
|
"\x03tag\x18\x01 \x01(\tR\x03tag\x12M\n" +
|
||||||
"\x11receiver_settings\x18\x02 \x01(\v2 .xray.common.serial.TypedMessageR\x10receiverSettings\x12G\n" +
|
"\x11receiver_settings\x18\x02 \x01(\v2 .xray.common.serial.TypedMessageR\x10receiverSettings\x12G\n" +
|
||||||
"\x0eproxy_settings\x18\x03 \x01(\v2 .xray.common.serial.TypedMessageR\rproxySettings\"\x10\n" +
|
"\x0eproxy_settings\x18\x03 \x01(\v2 .xray.common.serial.TypedMessageR\rproxySettings\"\x10\n" +
|
||||||
"\x0eOutboundConfig\"\xd6\x02\n" +
|
"\x0eOutboundConfig\"\x9d\x03\n" +
|
||||||
"\fSenderConfig\x12-\n" +
|
"\fSenderConfig\x12-\n" +
|
||||||
"\x03via\x18\x01 \x01(\v2\x1b.xray.common.net.IPOrDomainR\x03via\x12N\n" +
|
"\x03via\x18\x01 \x01(\v2\x1b.xray.common.net.IPOrDomainR\x03via\x12N\n" +
|
||||||
"\x0fstream_settings\x18\x02 \x01(\v2%.xray.transport.internet.StreamConfigR\x0estreamSettings\x12T\n" +
|
"\x0fstream_settings\x18\x02 \x01(\v2%.xray.transport.internet.StreamConfigR\x0estreamSettings\x12K\n" +
|
||||||
|
"\x0eproxy_settings\x18\x03 \x01(\v2$.xray.transport.internet.ProxyConfigR\rproxySettings\x12T\n" +
|
||||||
"\x12multiplex_settings\x18\x04 \x01(\v2%.xray.app.proxyman.MultiplexingConfigR\x11multiplexSettings\x12\x19\n" +
|
"\x12multiplex_settings\x18\x04 \x01(\v2%.xray.app.proxyman.MultiplexingConfigR\x11multiplexSettings\x12\x19\n" +
|
||||||
"\bvia_cidr\x18\x05 \x01(\tR\aviaCidr\x12P\n" +
|
"\bvia_cidr\x18\x05 \x01(\tR\aviaCidr\x12P\n" +
|
||||||
"\x0ftarget_strategy\x18\x06 \x01(\x0e2'.xray.transport.internet.DomainStrategyR\x0etargetStrategyJ\x04\b\x03\x10\x04\"\xa4\x01\n" +
|
"\x0ftarget_strategy\x18\x06 \x01(\x0e2'.xray.transport.internet.DomainStrategyR\x0etargetStrategy\"\xa4\x01\n" +
|
||||||
"\x12MultiplexingConfig\x12\x18\n" +
|
"\x12MultiplexingConfig\x12\x18\n" +
|
||||||
"\aenabled\x18\x01 \x01(\bR\aenabled\x12 \n" +
|
"\aenabled\x18\x01 \x01(\bR\aenabled\x12 \n" +
|
||||||
"\vconcurrency\x18\x02 \x01(\x05R\vconcurrency\x12(\n" +
|
"\vconcurrency\x18\x02 \x01(\x05R\vconcurrency\x12(\n" +
|
||||||
@@ -539,7 +548,8 @@ var file_app_proxyman_config_proto_goTypes = []any{
|
|||||||
(*net.IPOrDomain)(nil), // 10: xray.common.net.IPOrDomain
|
(*net.IPOrDomain)(nil), // 10: xray.common.net.IPOrDomain
|
||||||
(*internet.StreamConfig)(nil), // 11: xray.transport.internet.StreamConfig
|
(*internet.StreamConfig)(nil), // 11: xray.transport.internet.StreamConfig
|
||||||
(*serial.TypedMessage)(nil), // 12: xray.common.serial.TypedMessage
|
(*serial.TypedMessage)(nil), // 12: xray.common.serial.TypedMessage
|
||||||
(internet.DomainStrategy)(0), // 13: xray.transport.internet.DomainStrategy
|
(*internet.ProxyConfig)(nil), // 13: xray.transport.internet.ProxyConfig
|
||||||
|
(internet.DomainStrategy)(0), // 14: xray.transport.internet.DomainStrategy
|
||||||
}
|
}
|
||||||
var file_app_proxyman_config_proto_depIdxs = []int32{
|
var file_app_proxyman_config_proto_depIdxs = []int32{
|
||||||
7, // 0: xray.app.proxyman.SniffingConfig.domains_excluded:type_name -> xray.common.geodata.DomainRule
|
7, // 0: xray.app.proxyman.SniffingConfig.domains_excluded:type_name -> xray.common.geodata.DomainRule
|
||||||
@@ -552,13 +562,14 @@ var file_app_proxyman_config_proto_depIdxs = []int32{
|
|||||||
12, // 7: xray.app.proxyman.InboundHandlerConfig.proxy_settings:type_name -> xray.common.serial.TypedMessage
|
12, // 7: xray.app.proxyman.InboundHandlerConfig.proxy_settings:type_name -> xray.common.serial.TypedMessage
|
||||||
10, // 8: xray.app.proxyman.SenderConfig.via:type_name -> xray.common.net.IPOrDomain
|
10, // 8: xray.app.proxyman.SenderConfig.via:type_name -> xray.common.net.IPOrDomain
|
||||||
11, // 9: xray.app.proxyman.SenderConfig.stream_settings:type_name -> xray.transport.internet.StreamConfig
|
11, // 9: xray.app.proxyman.SenderConfig.stream_settings:type_name -> xray.transport.internet.StreamConfig
|
||||||
6, // 10: xray.app.proxyman.SenderConfig.multiplex_settings:type_name -> xray.app.proxyman.MultiplexingConfig
|
13, // 10: xray.app.proxyman.SenderConfig.proxy_settings:type_name -> xray.transport.internet.ProxyConfig
|
||||||
13, // 11: xray.app.proxyman.SenderConfig.target_strategy:type_name -> xray.transport.internet.DomainStrategy
|
6, // 11: xray.app.proxyman.SenderConfig.multiplex_settings:type_name -> xray.app.proxyman.MultiplexingConfig
|
||||||
12, // [12:12] is the sub-list for method output_type
|
14, // 12: xray.app.proxyman.SenderConfig.target_strategy:type_name -> xray.transport.internet.DomainStrategy
|
||||||
12, // [12:12] is the sub-list for method input_type
|
13, // [13:13] is the sub-list for method output_type
|
||||||
12, // [12:12] is the sub-list for extension type_name
|
13, // [13:13] is the sub-list for method input_type
|
||||||
12, // [12:12] is the sub-list for extension extendee
|
13, // [13:13] is the sub-list for extension type_name
|
||||||
0, // [0:12] is the sub-list for field type_name
|
13, // [13:13] is the sub-list for extension extendee
|
||||||
|
0, // [0:13] is the sub-list for field type_name
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() { file_app_proxyman_config_proto_init() }
|
func init() { file_app_proxyman_config_proto_init() }
|
||||||
|
|||||||
@@ -57,7 +57,7 @@ message SenderConfig {
|
|||||||
// Send traffic through the given IP. Only IP is allowed.
|
// Send traffic through the given IP. Only IP is allowed.
|
||||||
xray.common.net.IPOrDomain via = 1;
|
xray.common.net.IPOrDomain via = 1;
|
||||||
xray.transport.internet.StreamConfig stream_settings = 2;
|
xray.transport.internet.StreamConfig stream_settings = 2;
|
||||||
reserved 3;
|
xray.transport.internet.ProxyConfig proxy_settings = 3;
|
||||||
MultiplexingConfig multiplex_settings = 4;
|
MultiplexingConfig multiplex_settings = 4;
|
||||||
string via_cidr = 5;
|
string via_cidr = 5;
|
||||||
xray.transport.internet.DomainStrategy target_strategy = 6;
|
xray.transport.internet.DomainStrategy target_strategy = 6;
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ func getStatCounter(v *core.Instance, tag string) (stats.Counter, stats.Counter)
|
|||||||
if len(tag) > 0 && policy.ForSystem().Stats.InboundUplink {
|
if len(tag) > 0 && policy.ForSystem().Stats.InboundUplink {
|
||||||
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
||||||
name := "inbound>>>" + tag + ">>>traffic>>>uplink"
|
name := "inbound>>>" + tag + ">>>traffic>>>uplink"
|
||||||
c, _ := statsManager.GetOrRegisterCounter(name)
|
c, _ := stats.GetOrRegisterCounter(statsManager, name)
|
||||||
if c != nil {
|
if c != nil {
|
||||||
uplinkCounter = c
|
uplinkCounter = c
|
||||||
}
|
}
|
||||||
@@ -34,7 +34,7 @@ func getStatCounter(v *core.Instance, tag string) (stats.Counter, stats.Counter)
|
|||||||
if len(tag) > 0 && policy.ForSystem().Stats.InboundDownlink {
|
if len(tag) > 0 && policy.ForSystem().Stats.InboundDownlink {
|
||||||
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
||||||
name := "inbound>>>" + tag + ">>>traffic>>>downlink"
|
name := "inbound>>>" + tag + ">>>traffic>>>downlink"
|
||||||
c, _ := statsManager.GetOrRegisterCounter(name)
|
c, _ := stats.GetOrRegisterCounter(statsManager, name)
|
||||||
if c != nil {
|
if c != nil {
|
||||||
downlinkCounter = c
|
downlinkCounter = c
|
||||||
}
|
}
|
||||||
@@ -57,23 +57,16 @@ func NewAlwaysOnInboundHandler(ctx context.Context, tag string, receiverConfig *
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
src := net.TCPDestination(net.AnyIP, 0)
|
|
||||||
if receiverConfig.Listen != nil {
|
|
||||||
src.Address = receiverConfig.Listen.AsAddress()
|
|
||||||
}
|
|
||||||
if receiverConfig.PortList != nil && len(receiverConfig.PortList.Range) > 0 {
|
|
||||||
src.Port = net.Port(receiverConfig.PortList.Range[0].From)
|
|
||||||
}
|
|
||||||
mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("failed to parse stream config").Base(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
newCtx := session.ContextWithInbound(ctx, &session.Inbound{Tag: tag, Source: src})
|
// Set tag and sniffing config in context before creating proxy
|
||||||
newCtx = session.ContextWithContent(newCtx, &session.Content{SniffingRequest: sniffingRequest})
|
// This allows proxies like TUN to access these settings
|
||||||
newCtx = session.ContextWithStreamSettings(newCtx, mss)
|
ctx = session.ContextWithInbound(ctx, &session.Inbound{Tag: tag})
|
||||||
|
if receiverConfig.SniffingSettings != nil {
|
||||||
rawProxy, err := common.CreateObject(newCtx, proxyConfig)
|
ctx = session.ContextWithContent(ctx, &session.Content{
|
||||||
|
SniffingRequest: sniffingRequest,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
rawProxy, err := common.CreateObject(ctx, proxyConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -99,6 +92,11 @@ func NewAlwaysOnInboundHandler(ctx context.Context, tag string, receiverConfig *
|
|||||||
address = net.AnyIP
|
address = net.AnyIP
|
||||||
}
|
}
|
||||||
|
|
||||||
|
mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to parse stream config").Base(err).AtWarning()
|
||||||
|
}
|
||||||
|
|
||||||
if receiverConfig.ReceiveOriginalDestination {
|
if receiverConfig.ReceiveOriginalDestination {
|
||||||
if mss.SocketSettings == nil {
|
if mss.SocketSettings == nil {
|
||||||
mss.SocketSettings = &internet.SocketConfig{}
|
mss.SocketSettings = &internet.SocketConfig{}
|
||||||
@@ -172,12 +170,6 @@ func NewAlwaysOnInboundHandler(ctx context.Context, tag string, receiverConfig *
|
|||||||
|
|
||||||
// Start implements common.Runnable.
|
// Start implements common.Runnable.
|
||||||
func (h *AlwaysOnInboundHandler) Start() error {
|
func (h *AlwaysOnInboundHandler) Start() error {
|
||||||
// for inbound without worker (TUN)
|
|
||||||
if run, ok := h.proxy.(common.Runnable); ok {
|
|
||||||
if err := run.Start(); err != nil {
|
|
||||||
return errors.New("failed to start proxy").Base(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for _, worker := range h.workers {
|
for _, worker := range h.workers {
|
||||||
if err := worker.Start(); err != nil {
|
if err := worker.Start(); err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -16,10 +16,10 @@ import (
|
|||||||
|
|
||||||
// Manager manages all inbound handlers.
|
// Manager manages all inbound handlers.
|
||||||
type Manager struct {
|
type Manager struct {
|
||||||
access sync.RWMutex
|
access sync.RWMutex
|
||||||
untaggedHandlers []inbound.Handler
|
untaggedHandlers []inbound.Handler
|
||||||
taggedHandlers map[string]inbound.Handler
|
taggedHandlers map[string]inbound.Handler
|
||||||
running bool
|
running bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// New returns a new Manager for inbound handlers.
|
// New returns a new Manager for inbound handlers.
|
||||||
@@ -165,7 +165,7 @@ func NewHandler(ctx context.Context, config *core.InboundHandlerConfig) (inbound
|
|||||||
|
|
||||||
receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig)
|
receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, errors.New("not a ReceiverConfig")
|
return nil, errors.New("not a ReceiverConfig").AtError()
|
||||||
}
|
}
|
||||||
|
|
||||||
streamSettings := receiverSettings.StreamSettings
|
streamSettings := receiverSettings.StreamSettings
|
||||||
|
|||||||
@@ -142,7 +142,7 @@ func (w *tcpWorker) Start() error {
|
|||||||
go w.callback(conn)
|
go w.callback(conn)
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to listen TCP on ", w.port).Base(err)
|
return errors.New("failed to listen TCP on ", w.port).AtWarning().Base(err)
|
||||||
}
|
}
|
||||||
w.hub = hub
|
w.hub = hub
|
||||||
return nil
|
return nil
|
||||||
@@ -528,7 +528,7 @@ func (w *dsWorker) Start() error {
|
|||||||
go w.callback(conn)
|
go w.callback(conn)
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to listen Unix Domain Socket on ", w.address).Base(err)
|
return errors.New("failed to listen Unix Domain Socket on ", w.address).AtWarning().Base(err)
|
||||||
}
|
}
|
||||||
w.hub = hub
|
w.hub = hub
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
goerrors "errors"
|
goerrors "errors"
|
||||||
"io"
|
"io"
|
||||||
"math/big"
|
"math/big"
|
||||||
|
"os"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/dice"
|
"github.com/xtls/xray-core/common/dice"
|
||||||
|
|
||||||
@@ -15,6 +16,7 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/mux"
|
"github.com/xtls/xray-core/common/mux"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/net/cnc"
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
@@ -25,6 +27,8 @@ import (
|
|||||||
"github.com/xtls/xray-core/transport"
|
"github.com/xtls/xray-core/transport"
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
"github.com/xtls/xray-core/transport/internet/stat"
|
"github.com/xtls/xray-core/transport/internet/stat"
|
||||||
|
"github.com/xtls/xray-core/transport/internet/tls"
|
||||||
|
"github.com/xtls/xray-core/transport/pipe"
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -36,7 +40,7 @@ func getStatCounter(v *core.Instance, tag string) (stats.Counter, stats.Counter)
|
|||||||
if len(tag) > 0 && policy.ForSystem().Stats.OutboundUplink {
|
if len(tag) > 0 && policy.ForSystem().Stats.OutboundUplink {
|
||||||
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
||||||
name := "outbound>>>" + tag + ">>>traffic>>>uplink"
|
name := "outbound>>>" + tag + ">>>traffic>>>uplink"
|
||||||
c, _ := statsManager.GetOrRegisterCounter(name)
|
c, _ := stats.GetOrRegisterCounter(statsManager, name)
|
||||||
if c != nil {
|
if c != nil {
|
||||||
uplinkCounter = c
|
uplinkCounter = c
|
||||||
}
|
}
|
||||||
@@ -44,7 +48,7 @@ func getStatCounter(v *core.Instance, tag string) (stats.Counter, stats.Counter)
|
|||||||
if len(tag) > 0 && policy.ForSystem().Stats.OutboundDownlink {
|
if len(tag) > 0 && policy.ForSystem().Stats.OutboundDownlink {
|
||||||
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
||||||
name := "outbound>>>" + tag + ">>>traffic>>>downlink"
|
name := "outbound>>>" + tag + ">>>traffic>>>downlink"
|
||||||
c, _ := statsManager.GetOrRegisterCounter(name)
|
c, _ := stats.GetOrRegisterCounter(statsManager, name)
|
||||||
if c != nil {
|
if c != nil {
|
||||||
downlinkCounter = c
|
downlinkCounter = c
|
||||||
}
|
}
|
||||||
@@ -60,6 +64,7 @@ type Handler struct {
|
|||||||
streamSettings *internet.MemoryStreamConfig
|
streamSettings *internet.MemoryStreamConfig
|
||||||
proxyConfig proto.Message
|
proxyConfig proto.Message
|
||||||
proxy proxy.Outbound
|
proxy proxy.Outbound
|
||||||
|
outboundManager outbound.Manager
|
||||||
mux *mux.ClientManager
|
mux *mux.ClientManager
|
||||||
xudp *mux.ClientManager
|
xudp *mux.ClientManager
|
||||||
udp443 string
|
udp443 string
|
||||||
@@ -73,6 +78,7 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
|||||||
uplinkCounter, downlinkCounter := getStatCounter(v, config.Tag)
|
uplinkCounter, downlinkCounter := getStatCounter(v, config.Tag)
|
||||||
h := &Handler{
|
h := &Handler{
|
||||||
tag: config.Tag,
|
tag: config.Tag,
|
||||||
|
outboundManager: v.GetFeature(outbound.ManagerType()).(outbound.Manager),
|
||||||
uplinkCounter: uplinkCounter,
|
uplinkCounter: uplinkCounter,
|
||||||
downlinkCounter: downlinkCounter,
|
downlinkCounter: downlinkCounter,
|
||||||
}
|
}
|
||||||
@@ -87,7 +93,7 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
|||||||
h.senderSettings = s
|
h.senderSettings = s
|
||||||
mss, err := internet.ToMemoryStreamConfig(s.StreamSettings)
|
mss, err := internet.ToMemoryStreamConfig(s.StreamSettings)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to parse stream settings").Base(err)
|
return nil, errors.New("failed to parse stream settings").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
h.streamSettings = mss
|
h.streamSettings = mss
|
||||||
default:
|
default:
|
||||||
@@ -103,10 +109,6 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
|||||||
|
|
||||||
ctx = session.ContextWithFullHandler(ctx, h)
|
ctx = session.ContextWithFullHandler(ctx, h)
|
||||||
|
|
||||||
if h.streamSettings != nil {
|
|
||||||
ctx = session.ContextWithStreamSettings(ctx, h.streamSettings)
|
|
||||||
}
|
|
||||||
|
|
||||||
rawProxyHandler, err := common.CreateObject(ctx, proxyConfig)
|
rawProxyHandler, err := common.CreateObject(ctx, proxyConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -194,6 +196,7 @@ func (h *Handler) Dispatch(ctx context.Context, link *transport.Link) {
|
|||||||
common.Interrupt(link.Reader)
|
common.Interrupt(link.Reader)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
unchangedDomain := ob.Target.Address.Domain()
|
unchangedDomain := ob.Target.Address.Domain()
|
||||||
ob.Target.Address = net.IPAddress(ips[dice.Roll(len(ips))])
|
ob.Target.Address = net.IPAddress(ips[dice.Roll(len(ips))])
|
||||||
@@ -217,7 +220,7 @@ func (h *Handler) Dispatch(ctx context.Context, link *transport.Link) {
|
|||||||
if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 {
|
if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 {
|
||||||
switch h.udp443 {
|
switch h.udp443 {
|
||||||
case "reject":
|
case "reject":
|
||||||
test(errors.New("XUDP rejected UDP/443 traffic"))
|
test(errors.New("XUDP rejected UDP/443 traffic").AtInfo())
|
||||||
return
|
return
|
||||||
case "skip":
|
case "skip":
|
||||||
goto out
|
goto out
|
||||||
@@ -266,26 +269,71 @@ func (h *Handler) DestIpAddress() net.IP {
|
|||||||
|
|
||||||
// Dial implements internet.Dialer.
|
// Dial implements internet.Dialer.
|
||||||
func (h *Handler) Dial(ctx context.Context, dest net.Destination) (stat.Connection, error) {
|
func (h *Handler) Dial(ctx context.Context, dest net.Destination) (stat.Connection, error) {
|
||||||
if h.senderSettings != nil && h.senderSettings.Via != nil {
|
if h.senderSettings != nil {
|
||||||
outbounds := session.OutboundsFromContext(ctx)
|
|
||||||
ob := outbounds[len(outbounds)-1]
|
if h.senderSettings.ProxySettings.HasTag() {
|
||||||
h.SetOutboundGateway(ctx, ob)
|
|
||||||
|
tag := h.senderSettings.ProxySettings.Tag
|
||||||
|
handler := h.outboundManager.GetHandler(tag)
|
||||||
|
if handler != nil {
|
||||||
|
errors.LogDebug(ctx, "proxying to ", tag, " for dest ", dest)
|
||||||
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
|
ctx = session.ContextWithOutbounds(ctx, append(outbounds, &session.Outbound{
|
||||||
|
Target: dest,
|
||||||
|
Tag: tag,
|
||||||
|
})) // add another outbound in session ctx
|
||||||
|
opts := pipe.OptionsFromContext(ctx)
|
||||||
|
uplinkReader, uplinkWriter := pipe.New(opts...)
|
||||||
|
downlinkReader, downlinkWriter := pipe.New(opts...)
|
||||||
|
|
||||||
|
go handler.Dispatch(ctx, &transport.Link{Reader: uplinkReader, Writer: downlinkWriter})
|
||||||
|
conn := cnc.NewConnection(cnc.ConnectionInputMulti(uplinkWriter), cnc.ConnectionOutputMulti(downlinkReader))
|
||||||
|
|
||||||
|
if config := tls.ConfigFromStreamSettings(h.streamSettings); config != nil {
|
||||||
|
tlsConfig := config.GetTLSConfig(tls.WithDestination(dest))
|
||||||
|
conn = tls.Client(conn, tlsConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
return h.getStatCouterConnection(conn), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
errors.LogError(ctx, "failed to get outbound handler with tag: ", tag)
|
||||||
|
return nil, errors.New("failed to get outbound handler with tag: " + tag)
|
||||||
|
}
|
||||||
|
|
||||||
|
if h.senderSettings.Via != nil {
|
||||||
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
|
ob := outbounds[len(outbounds)-1]
|
||||||
|
h.SetOutboundGateway(ctx, ob)
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
if conn, err := h.getUoTConnection(ctx, dest); err != os.ErrInvalid {
|
||||||
|
return conn, err
|
||||||
}
|
}
|
||||||
|
|
||||||
conn, err := internet.Dial(ctx, dest, h.streamSettings)
|
conn, err := internet.Dial(ctx, dest, h.streamSettings)
|
||||||
conn = h.getStatCouterConnection(conn)
|
conn = h.getStatCouterConnection(conn)
|
||||||
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
|
if outbounds != nil {
|
||||||
|
ob := outbounds[len(outbounds)-1]
|
||||||
|
ob.Conn = conn
|
||||||
|
} else {
|
||||||
|
// for Vision's pre-connect
|
||||||
|
}
|
||||||
return conn, err
|
return conn, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) SetOutboundGateway(ctx context.Context, ob *session.Outbound) {
|
func (h *Handler) SetOutboundGateway(ctx context.Context, ob *session.Outbound) {
|
||||||
if ob.Gateway == nil && h.senderSettings != nil && h.senderSettings.Via != nil &&
|
if ob.Gateway == nil && h.senderSettings != nil && h.senderSettings.Via != nil && !h.senderSettings.ProxySettings.HasTag() && (h.streamSettings.SocketSettings == nil || len(h.streamSettings.SocketSettings.DialerProxy) == 0) {
|
||||||
(h.streamSettings.SocketSettings == nil || len(h.streamSettings.SocketSettings.DialerProxy) == 0) {
|
|
||||||
var domain string
|
var domain string
|
||||||
addr := h.senderSettings.Via.AsAddress()
|
addr := h.senderSettings.Via.AsAddress()
|
||||||
domain = h.senderSettings.Via.GetDomain()
|
domain = h.senderSettings.Via.GetDomain()
|
||||||
switch {
|
switch {
|
||||||
case h.senderSettings.ViaCidr != "":
|
case h.senderSettings.ViaCidr != "":
|
||||||
ob.Gateway = ParseRandomIP(addr, h.senderSettings.ViaCidr)
|
ob.Gateway = ParseRandomIP(addr, h.senderSettings.ViaCidr)
|
||||||
|
|
||||||
case domain == "origin":
|
case domain == "origin":
|
||||||
if inbound := session.InboundFromContext(ctx); inbound != nil {
|
if inbound := session.InboundFromContext(ctx); inbound != nil {
|
||||||
if inbound.Local.IsValid() && inbound.Local.Address.Family().IsIP() {
|
if inbound.Local.IsValid() && inbound.Local.Address.Family().IsIP() {
|
||||||
@@ -300,9 +348,12 @@ func (h *Handler) SetOutboundGateway(ctx context.Context, ob *session.Outbound)
|
|||||||
errors.LogDebug(ctx, "use inbound source ip as sendthrough: ", inbound.Source.Address.String())
|
errors.LogDebug(ctx, "use inbound source ip as sendthrough: ", inbound.Source.Address.String())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
default: // case addr.Family().IsDomain():
|
//case addr.Family().IsDomain():
|
||||||
|
default:
|
||||||
ob.Gateway = addr
|
ob.Gateway = addr
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -345,6 +396,7 @@ func (h *Handler) ProxySettings() *serial.TypedMessage {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func ParseRandomIP(addr net.Address, prefix string) net.Address {
|
func ParseRandomIP(addr net.Address, prefix string) net.Address {
|
||||||
|
|
||||||
_, ipnet, _ := net.ParseCIDR(addr.IP().String() + "/" + prefix)
|
_, ipnet, _ := net.ParseCIDR(addr.IP().String() + "/" + prefix)
|
||||||
|
|
||||||
ones, bits := ipnet.Mask.Size()
|
ones, bits := ipnet.Mask.Size()
|
||||||
|
|||||||
@@ -22,8 +22,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestInterfaces(t *testing.T) {
|
func TestInterfaces(t *testing.T) {
|
||||||
_ = outbound.Handler(new(Handler))
|
_ = (outbound.Handler)(new(Handler))
|
||||||
_ = outbound.Manager(new(Manager))
|
_ = (outbound.Manager)(new(Manager))
|
||||||
}
|
}
|
||||||
|
|
||||||
const xrayKey core.XrayKey = 1
|
const xrayKey core.XrayKey = 1
|
||||||
@@ -43,7 +43,7 @@ func TestOutboundWithoutStatCounter(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
v, _ := core.New(config)
|
v, _ := core.New(config)
|
||||||
v.AddFeature(outbound.Manager(new(Manager)))
|
v.AddFeature((outbound.Manager)(new(Manager)))
|
||||||
ctx := context.WithValue(context.Background(), xrayKey, v)
|
ctx := context.WithValue(context.Background(), xrayKey, v)
|
||||||
ctx = session.ContextWithOutbounds(ctx, []*session.Outbound{{}})
|
ctx = session.ContextWithOutbounds(ctx, []*session.Outbound{{}})
|
||||||
h, _ := NewHandler(ctx, &core.OutboundHandlerConfig{
|
h, _ := NewHandler(ctx, &core.OutboundHandlerConfig{
|
||||||
@@ -73,7 +73,7 @@ func TestOutboundWithStatCounter(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
v, _ := core.New(config)
|
v, _ := core.New(config)
|
||||||
v.AddFeature(outbound.Manager(new(Manager)))
|
v.AddFeature((outbound.Manager)(new(Manager)))
|
||||||
ctx := context.WithValue(context.Background(), xrayKey, v)
|
ctx := context.WithValue(context.Background(), xrayKey, v)
|
||||||
ctx = session.ContextWithOutbounds(ctx, []*session.Outbound{{}})
|
ctx = session.ContextWithOutbounds(ctx, []*session.Outbound{{}})
|
||||||
h, _ := NewHandler(ctx, &core.OutboundHandlerConfig{
|
h, _ := NewHandler(ctx, &core.OutboundHandlerConfig{
|
||||||
@@ -88,6 +88,7 @@ func TestOutboundWithStatCounter(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestTagsCache(t *testing.T) {
|
func TestTagsCache(t *testing.T) {
|
||||||
|
|
||||||
test_duration := 10 * time.Second
|
test_duration := 10 * time.Second
|
||||||
threads_num := 50
|
threads_num := 50
|
||||||
delay := 10 * time.Millisecond
|
delay := 10 * time.Millisecond
|
||||||
|
|||||||
@@ -162,6 +162,7 @@ func (m *Manager) ListHandlers(ctx context.Context) []outbound.Handler {
|
|||||||
|
|
||||||
// Select implements outbound.HandlerSelector.
|
// Select implements outbound.HandlerSelector.
|
||||||
func (m *Manager) Select(selectors []string) []string {
|
func (m *Manager) Select(selectors []string) []string {
|
||||||
|
|
||||||
key := strings.Join(selectors, ",")
|
key := strings.Join(selectors, ",")
|
||||||
if cache, ok := m.tagsCache.Load(key); ok {
|
if cache, ok := m.tagsCache.Load(key); ok {
|
||||||
return cache.([]string)
|
return cache.([]string)
|
||||||
|
|||||||
@@ -0,0 +1,35 @@
|
|||||||
|
package outbound
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/sagernet/sing/common/uot"
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
|
"github.com/xtls/xray-core/transport/internet/stat"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (h *Handler) getUoTConnection(ctx context.Context, dest net.Destination) (stat.Connection, error) {
|
||||||
|
if dest.Address == nil {
|
||||||
|
return nil, errors.New("nil destination address")
|
||||||
|
}
|
||||||
|
if !dest.Address.Family().IsDomain() {
|
||||||
|
return nil, os.ErrInvalid
|
||||||
|
}
|
||||||
|
var uotVersion int
|
||||||
|
if dest.Address.Domain() == uot.MagicAddress {
|
||||||
|
uotVersion = uot.Version
|
||||||
|
} else if dest.Address.Domain() == uot.LegacyMagicAddress {
|
||||||
|
uotVersion = uot.LegacyVersion
|
||||||
|
} else {
|
||||||
|
return nil, os.ErrInvalid
|
||||||
|
}
|
||||||
|
packetConn, err := internet.ListenSystemPacket(ctx, &net.UDPAddr{IP: net.AnyIP.IP(), Port: 0}, h.streamSettings.SocketSettings)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("unable to listen socket").Base(err)
|
||||||
|
}
|
||||||
|
conn := uot.NewServerConn(packetConn, uotVersion)
|
||||||
|
return h.getStatCouterConnection(conn), nil
|
||||||
|
}
|
||||||
@@ -68,13 +68,13 @@ func (p *Portal) HandleConnection(ctx context.Context, link *transport.Link) err
|
|||||||
outbounds := session.OutboundsFromContext(ctx)
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
ob := outbounds[len(outbounds)-1]
|
ob := outbounds[len(outbounds)-1]
|
||||||
if ob == nil {
|
if ob == nil {
|
||||||
return errors.New("outbound metadata not found")
|
return errors.New("outbound metadata not found").AtError()
|
||||||
}
|
}
|
||||||
|
|
||||||
if isDomain(ob.Target, p.domain) {
|
if isDomain(ob.Target, p.domain) {
|
||||||
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
|
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create mux client worker").Base(err)
|
return errors.New("failed to create mux client worker").Base(err).AtWarning()
|
||||||
}
|
}
|
||||||
|
|
||||||
worker, err := NewPortalWorker(muxClient)
|
worker, err := NewPortalWorker(muxClient)
|
||||||
|
|||||||
@@ -136,7 +136,7 @@ func (b *Balancer) SelectOutbounds() ([]string, error) {
|
|||||||
|
|
||||||
// GetPrincipleTarget implements routing.BalancerPrincipleTarget
|
// GetPrincipleTarget implements routing.BalancerPrincipleTarget
|
||||||
func (r *Router) GetPrincipleTarget(tag string) ([]string, error) {
|
func (r *Router) GetPrincipleTarget(tag string) ([]string, error) {
|
||||||
if b, ok := (*r.balancers.Load())[tag]; ok {
|
if b, ok := r.balancers[tag]; ok {
|
||||||
if s, ok := b.strategy.(BalancingPrincipleTarget); ok {
|
if s, ok := b.strategy.(BalancingPrincipleTarget); ok {
|
||||||
candidates, err := b.SelectOutbounds()
|
candidates, err := b.SelectOutbounds()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -151,7 +151,7 @@ func (r *Router) GetPrincipleTarget(tag string) ([]string, error) {
|
|||||||
|
|
||||||
// SetOverrideTarget implements routing.BalancerOverrider
|
// SetOverrideTarget implements routing.BalancerOverrider
|
||||||
func (r *Router) SetOverrideTarget(tag, target string) error {
|
func (r *Router) SetOverrideTarget(tag, target string) error {
|
||||||
if b, ok := (*r.balancers.Load())[tag]; ok {
|
if b, ok := r.balancers[tag]; ok {
|
||||||
b.override.Put(target)
|
b.override.Put(target)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -160,7 +160,7 @@ func (r *Router) SetOverrideTarget(tag, target string) error {
|
|||||||
|
|
||||||
// GetOverrideTarget implements routing.BalancerOverrider
|
// GetOverrideTarget implements routing.BalancerOverrider
|
||||||
func (r *Router) GetOverrideTarget(tag string) (string, error) {
|
func (r *Router) GetOverrideTarget(tag string) (string, error) {
|
||||||
if b, ok := (*r.balancers.Load())[tag]; ok {
|
if b, ok := r.balancers[tag]; ok {
|
||||||
return b.override.Get(), nil
|
return b.override.Get(), nil
|
||||||
}
|
}
|
||||||
return "", errors.New("cannot find tag")
|
return "", errors.New("cannot find tag")
|
||||||
|
|||||||
@@ -2,8 +2,25 @@ package router
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
sync "sync"
|
sync "sync"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func (r *Router) OverrideBalancer(balancer string, target string) error {
|
||||||
|
var b *Balancer
|
||||||
|
for tag, bl := range r.balancers {
|
||||||
|
if tag == balancer {
|
||||||
|
b = bl
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if b == nil {
|
||||||
|
return errors.New("balancer '", balancer, "' not found")
|
||||||
|
}
|
||||||
|
b.override.Put(target)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
type overrideSettings struct {
|
type overrideSettings struct {
|
||||||
target string
|
target string
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ func (s *routingServer) AddRule(ctx context.Context, request *AddRuleRequest) (*
|
|||||||
return &AddRuleResponse{}, bo.AddRule(request.Config, request.ShouldAppend)
|
return &AddRuleResponse{}, bo.AddRule(request.Config, request.ShouldAppend)
|
||||||
}
|
}
|
||||||
return nil, errors.New("unsupported router implementation")
|
return nil, errors.New("unsupported router implementation")
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *routingServer) RemoveRule(ctx context.Context, request *RemoveRuleRequest) (*RemoveRuleResponse, error) {
|
func (s *routingServer) RemoveRule(ctx context.Context, request *RemoveRuleRequest) (*RemoveRuleResponse, error) {
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"regexp"
|
"regexp"
|
||||||
"runtime"
|
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -394,22 +393,3 @@ func (m *ProcessNameMatcher) Apply(ctx routing.Context) bool {
|
|||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// LocalOSMatcher matches the operating system Xray itself is running on. That never
|
|
||||||
// changes while Xray is running, so the result is resolved when the rule is built.
|
|
||||||
type LocalOSMatcher struct {
|
|
||||||
matched bool
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewLocalOSMatcher(names []string) *LocalOSMatcher {
|
|
||||||
return &LocalOSMatcher{
|
|
||||||
matched: slices.ContainsFunc(names, func(name string) bool {
|
|
||||||
return strings.EqualFold(name, runtime.GOOS)
|
|
||||||
}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Apply implements Condition.
|
|
||||||
func (m *LocalOSMatcher) Apply(_ routing.Context) bool {
|
|
||||||
return m.matched
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -2,9 +2,7 @@ package router_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"runtime"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
. "github.com/xtls/xray-core/app/router"
|
. "github.com/xtls/xray-core/app/router"
|
||||||
@@ -345,31 +343,6 @@ func TestChinaSites(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestLocalOSRule(t *testing.T) {
|
|
||||||
otherOS := "plan9"
|
|
||||||
if runtime.GOOS == otherOS {
|
|
||||||
otherOS = "linux"
|
|
||||||
}
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
localOS []string
|
|
||||||
output bool
|
|
||||||
}{
|
|
||||||
{localOS: []string{runtime.GOOS}, output: true},
|
|
||||||
{localOS: []string{otherOS}, output: false},
|
|
||||||
{localOS: []string{otherOS, runtime.GOOS}, output: true},
|
|
||||||
{localOS: []string{strings.ToUpper(runtime.GOOS)}, output: true},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, test := range cases {
|
|
||||||
cond, err := (&RoutingRule{LocalOs: test.localOS}).BuildCondition()
|
|
||||||
common.Must(err)
|
|
||||||
if got := cond.Apply(withBackground()); got != test.output {
|
|
||||||
t.Errorf("for localOS %v on %s: expected %v, got %v", test.localOS, runtime.GOOS, test.output, got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkMphDomainMatcher(b *testing.B) {
|
func BenchmarkMphDomainMatcher(b *testing.B) {
|
||||||
b.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
b.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||||
rules, err := geodata.ParseDomainRules([]string{"geosite:cn"}, geodata.Domain_Substr)
|
rules, err := geodata.ParseDomainRules([]string{"geosite:cn"}, geodata.Domain_Substr)
|
||||||
|
|||||||
@@ -33,10 +33,6 @@ func (r *Rule) Apply(ctx routing.Context) bool {
|
|||||||
func (rr *RoutingRule) BuildCondition() (Condition, error) {
|
func (rr *RoutingRule) BuildCondition() (Condition, error) {
|
||||||
conds := NewConditionChan()
|
conds := NewConditionChan()
|
||||||
|
|
||||||
if len(rr.LocalOs) > 0 {
|
|
||||||
conds.Add(NewLocalOSMatcher(rr.LocalOs))
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(rr.InboundTag) > 0 {
|
if len(rr.InboundTag) > 0 {
|
||||||
conds.Add(NewInboundTagMatcher(rr.InboundTag))
|
conds.Add(NewInboundTagMatcher(rr.InboundTag))
|
||||||
}
|
}
|
||||||
@@ -115,7 +111,7 @@ func (rr *RoutingRule) BuildCondition() (Condition, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if conds.Len() == 0 {
|
if conds.Len() == 0 {
|
||||||
return nil, errors.New("this rule has no effective fields")
|
return nil, errors.New("this rule has no effective fields").AtWarning()
|
||||||
}
|
}
|
||||||
|
|
||||||
return conds, nil
|
return conds, nil
|
||||||
@@ -145,7 +141,7 @@ func (br *BalancingRule) Build(ohm outbound.Manager, dispatcher routing.Dispatch
|
|||||||
}
|
}
|
||||||
s, ok := i.(*StrategyLeastLoadConfig)
|
s, ok := i.(*StrategyLeastLoadConfig)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, errors.New("not a StrategyLeastLoadConfig")
|
return nil, errors.New("not a StrategyLeastLoadConfig").AtError()
|
||||||
}
|
}
|
||||||
leastLoadStrategy := NewLeastLoadStrategy(s)
|
leastLoadStrategy := NewLeastLoadStrategy(s)
|
||||||
return &Balancer{
|
return &Balancer{
|
||||||
|
|||||||
+8
-28
@@ -107,10 +107,8 @@ type RoutingRule struct {
|
|||||||
VlessRouteList *net.PortList `protobuf:"bytes,20,opt,name=vless_route_list,json=vlessRouteList,proto3" json:"vless_route_list,omitempty"`
|
VlessRouteList *net.PortList `protobuf:"bytes,20,opt,name=vless_route_list,json=vlessRouteList,proto3" json:"vless_route_list,omitempty"`
|
||||||
Process []string `protobuf:"bytes,21,rep,name=process,proto3" json:"process,omitempty"`
|
Process []string `protobuf:"bytes,21,rep,name=process,proto3" json:"process,omitempty"`
|
||||||
Webhook *WebhookConfig `protobuf:"bytes,22,opt,name=webhook,proto3" json:"webhook,omitempty"`
|
Webhook *WebhookConfig `protobuf:"bytes,22,opt,name=webhook,proto3" json:"webhook,omitempty"`
|
||||||
// List of operating systems for matching the one Xray itself is running on.
|
unknownFields protoimpl.UnknownFields
|
||||||
LocalOs []string `protobuf:"bytes,23,rep,name=local_os,json=localOs,proto3" json:"local_os,omitempty"`
|
sizeCache protoimpl.SizeCache
|
||||||
unknownFields protoimpl.UnknownFields
|
|
||||||
sizeCache protoimpl.SizeCache
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *RoutingRule) Reset() {
|
func (x *RoutingRule) Reset() {
|
||||||
@@ -280,13 +278,6 @@ func (x *RoutingRule) GetWebhook() *WebhookConfig {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *RoutingRule) GetLocalOs() []string {
|
|
||||||
if x != nil {
|
|
||||||
return x.LocalOs
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type isRoutingRule_TargetTag interface {
|
type isRoutingRule_TargetTag interface {
|
||||||
isRoutingRule_TargetTag()
|
isRoutingRule_TargetTag()
|
||||||
}
|
}
|
||||||
@@ -587,10 +578,8 @@ type Config struct {
|
|||||||
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
||||||
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
||||||
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
||||||
// Absolute path to the Lua routing script.
|
unknownFields protoimpl.UnknownFields
|
||||||
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
|
sizeCache protoimpl.SizeCache
|
||||||
unknownFields protoimpl.UnknownFields
|
|
||||||
sizeCache protoimpl.SizeCache
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) Reset() {
|
func (x *Config) Reset() {
|
||||||
@@ -644,18 +633,11 @@ func (x *Config) GetBalancingRule() []*BalancingRule {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) GetScript() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Script
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
var File_app_router_config_proto protoreflect.FileDescriptor
|
var File_app_router_config_proto protoreflect.FileDescriptor
|
||||||
|
|
||||||
const file_app_router_config_proto_rawDesc = "" +
|
const file_app_router_config_proto_rawDesc = "" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"\x17app/router/config.proto\x12\x0fxray.app.router\x1a!common/serial/typed_message.proto\x1a\x15common/net/port.proto\x1a\x18common/net/network.proto\x1a\x1bcommon/geodata/geodat.proto\"\xdc\a\n" +
|
"\x17app/router/config.proto\x12\x0fxray.app.router\x1a!common/serial/typed_message.proto\x1a\x15common/net/port.proto\x1a\x18common/net/network.proto\x1a\x1bcommon/geodata/geodat.proto\"\xc1\a\n" +
|
||||||
"\vRoutingRule\x12\x12\n" +
|
"\vRoutingRule\x12\x12\n" +
|
||||||
"\x03tag\x18\x01 \x01(\tH\x00R\x03tag\x12%\n" +
|
"\x03tag\x18\x01 \x01(\tH\x00R\x03tag\x12%\n" +
|
||||||
"\rbalancing_tag\x18\f \x01(\tH\x00R\fbalancingTag\x12\x19\n" +
|
"\rbalancing_tag\x18\f \x01(\tH\x00R\fbalancingTag\x12\x19\n" +
|
||||||
@@ -679,8 +661,7 @@ const file_app_router_config_proto_rawDesc = "" +
|
|||||||
"\x0flocal_port_list\x18\x12 \x01(\v2\x19.xray.common.net.PortListR\rlocalPortList\x12C\n" +
|
"\x0flocal_port_list\x18\x12 \x01(\v2\x19.xray.common.net.PortListR\rlocalPortList\x12C\n" +
|
||||||
"\x10vless_route_list\x18\x14 \x01(\v2\x19.xray.common.net.PortListR\x0evlessRouteList\x12\x18\n" +
|
"\x10vless_route_list\x18\x14 \x01(\v2\x19.xray.common.net.PortListR\x0evlessRouteList\x12\x18\n" +
|
||||||
"\aprocess\x18\x15 \x03(\tR\aprocess\x128\n" +
|
"\aprocess\x18\x15 \x03(\tR\aprocess\x128\n" +
|
||||||
"\awebhook\x18\x16 \x01(\v2\x1e.xray.app.router.WebhookConfigR\awebhook\x12\x19\n" +
|
"\awebhook\x18\x16 \x01(\v2\x1e.xray.app.router.WebhookConfigR\awebhook\x1a=\n" +
|
||||||
"\blocal_os\x18\x17 \x03(\tR\alocalOs\x1a=\n" +
|
|
||||||
"\x0fAttributesEntry\x12\x10\n" +
|
"\x0fAttributesEntry\x12\x10\n" +
|
||||||
"\x03key\x18\x01 \x01(\tR\x03key\x12\x14\n" +
|
"\x03key\x18\x01 \x01(\tR\x03key\x12\x14\n" +
|
||||||
"\x05value\x18\x02 \x01(\tR\x05value:\x028\x01B\f\n" +
|
"\x05value\x18\x02 \x01(\tR\x05value:\x028\x01B\f\n" +
|
||||||
@@ -708,12 +689,11 @@ const file_app_router_config_proto_rawDesc = "" +
|
|||||||
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
|
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
|
||||||
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
|
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
|
||||||
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
|
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
|
||||||
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\x02\n" +
|
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\x96\x02\n" +
|
||||||
"\x06Config\x12O\n" +
|
"\x06Config\x12O\n" +
|
||||||
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
|
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
|
||||||
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
|
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
|
||||||
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" +
|
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" +
|
||||||
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
|
|
||||||
"\x0eDomainStrategy\x12\b\n" +
|
"\x0eDomainStrategy\x12\b\n" +
|
||||||
"\x04AsIs\x10\x00\x12\x10\n" +
|
"\x04AsIs\x10\x00\x12\x10\n" +
|
||||||
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
|
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
|
||||||
|
|||||||
@@ -56,9 +56,6 @@ message RoutingRule {
|
|||||||
|
|
||||||
repeated string process = 21;
|
repeated string process = 21;
|
||||||
WebhookConfig webhook = 22;
|
WebhookConfig webhook = 22;
|
||||||
|
|
||||||
// List of operating systems for matching the one Xray itself is running on.
|
|
||||||
repeated string local_os = 23;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
message WebhookConfig {
|
message WebhookConfig {
|
||||||
@@ -110,6 +107,4 @@ message Config {
|
|||||||
DomainStrategy domain_strategy = 1;
|
DomainStrategy domain_strategy = 1;
|
||||||
repeated RoutingRule rule = 2;
|
repeated RoutingRule rule = 2;
|
||||||
repeated BalancingRule balancing_rule = 3;
|
repeated BalancingRule balancing_rule = 3;
|
||||||
// Absolute path to the Lua routing script.
|
|
||||||
string script = 4;
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,167 +0,0 @@
|
|||||||
package router
|
|
||||||
|
|
||||||
import (
|
|
||||||
"runtime"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
xlua "github.com/xtls/xray-core/common/lua"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/features/routing"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
luaContextType = "xray.router.Context"
|
|
||||||
luaAttributesType = "xray.router.Attributes"
|
|
||||||
)
|
|
||||||
|
|
||||||
// RegisterLua makes xray.router available to routing scripts.
|
|
||||||
func (r *Router) RegisterLua(L *lua.LState) {
|
|
||||||
registerLuaContext(L)
|
|
||||||
|
|
||||||
L.PreloadModule("xray.router", func(L *lua.LState) int {
|
|
||||||
module := L.NewTable()
|
|
||||||
|
|
||||||
module.RawSetString("NetworkUnknown", lua.LNumber(net.Network_Unknown))
|
|
||||||
module.RawSetString("NetworkTCP", lua.LNumber(net.Network_TCP))
|
|
||||||
module.RawSetString("NetworkUDP", lua.LNumber(net.Network_UDP))
|
|
||||||
module.RawSetString("NetworkUNIX", lua.LNumber(net.Network_UNIX))
|
|
||||||
module.RawSetString("LocalOS", lua.LString(runtime.GOOS))
|
|
||||||
|
|
||||||
module.RawSetString("PickOutbound", L.NewFunction(func(L *lua.LState) int {
|
|
||||||
tag, ok := L.Get(2).(lua.LString)
|
|
||||||
if !ok {
|
|
||||||
L.ArgError(2, "balancer tag must be a string")
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
balancer, found := (*r.balancers.Load())[string(tag)]
|
|
||||||
if !found {
|
|
||||||
xlua.PushNil(L)
|
|
||||||
xlua.PushError(L, errors.New("balancer ", tag, " not found"))
|
|
||||||
return 2
|
|
||||||
}
|
|
||||||
outboundTag, err := balancer.PickOutbound()
|
|
||||||
xlua.PushString(L, outboundTag)
|
|
||||||
xlua.PushError(L, err)
|
|
||||||
return 2
|
|
||||||
}))
|
|
||||||
|
|
||||||
module.RawSetString("FindProcess", L.NewFunction(func(L *lua.LState) int {
|
|
||||||
pid, name, path, err := findProcess(checkLuaContext(L), net.FindProcess)
|
|
||||||
xlua.PushNumber(L, pid)
|
|
||||||
xlua.PushString(L, name)
|
|
||||||
xlua.PushString(L, path)
|
|
||||||
xlua.PushError(L, err)
|
|
||||||
return 4
|
|
||||||
}))
|
|
||||||
|
|
||||||
L.Push(module)
|
|
||||||
return 1
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func registerLuaContext(L *lua.LState) {
|
|
||||||
attributes := L.NewTypeMetatable(luaAttributesType)
|
|
||||||
L.SetField(attributes, "__index", L.NewFunction(func(L *lua.LState) int {
|
|
||||||
values := L.CheckUserData(1).Value.(map[string]string)
|
|
||||||
key := L.CheckString(2)
|
|
||||||
if value, found := values[key]; found {
|
|
||||||
xlua.PushString(L, value)
|
|
||||||
} else {
|
|
||||||
xlua.PushNil(L)
|
|
||||||
}
|
|
||||||
return 1
|
|
||||||
}))
|
|
||||||
methods := L.NewTable()
|
|
||||||
L.SetFuncs(methods, map[string]lua.LGFunction{
|
|
||||||
"GetSourceIPs": func(L *lua.LState) int {
|
|
||||||
xlua.PushUserData(L, checkLuaContext(L).GetSourceIPs())
|
|
||||||
return 1
|
|
||||||
},
|
|
||||||
"GetTargetIPs": func(L *lua.LState) int {
|
|
||||||
xlua.PushUserData(L, checkLuaContext(L).GetTargetIPs())
|
|
||||||
return 1
|
|
||||||
},
|
|
||||||
"GetLocalIPs": func(L *lua.LState) int {
|
|
||||||
xlua.PushUserData(L, checkLuaContext(L).GetLocalIPs())
|
|
||||||
return 1
|
|
||||||
},
|
|
||||||
"GetAttributes": func(L *lua.LState) int {
|
|
||||||
values := L.NewUserData()
|
|
||||||
values.Value = checkLuaContext(L).GetAttributes()
|
|
||||||
L.SetMetatable(values, attributes)
|
|
||||||
L.Push(values)
|
|
||||||
return 1
|
|
||||||
},
|
|
||||||
})
|
|
||||||
L.SetField(L.NewTypeMetatable(luaContextType), "__index", methods)
|
|
||||||
}
|
|
||||||
|
|
||||||
func checkLuaContext(L *lua.LState) routing.Context {
|
|
||||||
ctx, ok := L.CheckUserData(1).Value.(routing.Context)
|
|
||||||
if !ok {
|
|
||||||
L.ArgError(1, "routing context expected")
|
|
||||||
}
|
|
||||||
return ctx
|
|
||||||
}
|
|
||||||
|
|
||||||
// callLuaHook invokes HandleRoute in the supplied state.
|
|
||||||
func (r *Router) callLuaHook(L *lua.LState, routeCtx routing.Context) (string, string, error) {
|
|
||||||
top := L.GetTop()
|
|
||||||
defer L.SetTop(top)
|
|
||||||
fn := L.GetGlobal("HandleRoute")
|
|
||||||
if fn.Type() != lua.LTFunction {
|
|
||||||
return "", "", errors.New("routing script must define HandleRoute(...)")
|
|
||||||
}
|
|
||||||
value := L.NewUserData()
|
|
||||||
value.Value = routeCtx
|
|
||||||
L.SetMetatable(value, L.GetTypeMetatable(luaContextType))
|
|
||||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
|
||||||
value, lua.LString(routeCtx.GetInboundTag()), lua.LNumber(routeCtx.GetSourcePort()),
|
|
||||||
lua.LNumber(routeCtx.GetTargetPort()), lua.LNumber(routeCtx.GetLocalPort()),
|
|
||||||
lua.LString(strings.ToLower(routeCtx.GetTargetDomain())), lua.LNumber(routeCtx.GetNetwork()),
|
|
||||||
lua.LString(routeCtx.GetProtocol()), lua.LString(routeCtx.GetUser()),
|
|
||||||
lua.LNumber(routeCtx.GetVlessRoute()), lua.LBool(routeCtx.GetSkipDNSResolve())); err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
return readLuaRouteResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
|
||||||
}
|
|
||||||
|
|
||||||
func readLuaRouteResult(tagValue, ruleValue, errorValue lua.LValue) (string, string, error) {
|
|
||||||
if err := xlua.ReadError(errorValue, "routing script error must be an error or string"); err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
tag, err := xlua.ReadOptionalString(tagValue, "routing script outboundTag must be a string or nil")
|
|
||||||
if err != nil || tag == "" {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
ruleTag, err := xlua.ReadOptionalString(ruleValue, "routing script ruleTag must be a string")
|
|
||||||
if err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
return tag, ruleTag, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type processFinder func(string, string, uint16, string, uint16) (int, string, string, error)
|
|
||||||
|
|
||||||
func findProcess(ctx routing.Context, finder processFinder) (int, string, string, error) {
|
|
||||||
sources := ctx.GetSourceIPs()
|
|
||||||
if len(sources) == 0 {
|
|
||||||
return 0, "", "", errors.New("process lookup requires a source IP")
|
|
||||||
}
|
|
||||||
var network string
|
|
||||||
switch ctx.GetNetwork() {
|
|
||||||
case net.Network_TCP:
|
|
||||||
network = "tcp"
|
|
||||||
case net.Network_UDP:
|
|
||||||
network = "udp"
|
|
||||||
default:
|
|
||||||
return 0, "", "", errors.New("process lookup requires TCP or UDP")
|
|
||||||
}
|
|
||||||
targetIP, targetPort := "", uint16(0)
|
|
||||||
if targets := ctx.GetTargetIPs(); len(targets) > 0 {
|
|
||||||
targetIP, targetPort = targets[0].String(), uint16(ctx.GetTargetPort())
|
|
||||||
}
|
|
||||||
return finder(network, sources[0].String(), uint16(ctx.GetSourcePort()), targetIP, targetPort)
|
|
||||||
}
|
|
||||||
@@ -1,300 +0,0 @@
|
|||||||
package router
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
go_errors "errors"
|
|
||||||
"runtime"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
|
||||||
"github.com/xtls/xray-core/common/session"
|
|
||||||
"github.com/xtls/xray-core/features/routing"
|
|
||||||
routing_session "github.com/xtls/xray-core/features/routing/session"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
type luaRouteTestContext struct {
|
|
||||||
*routing_session.Context
|
|
||||||
sourceIPs, targetIPs, localIPs []net.IP
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *luaRouteTestContext) GetSourceIPs() []net.IP { return c.sourceIPs }
|
|
||||||
func (c *luaRouteTestContext) GetTargetIPs() []net.IP { return c.targetIPs }
|
|
||||||
func (c *luaRouteTestContext) GetLocalIPs() []net.IP { return c.localIPs }
|
|
||||||
|
|
||||||
func newLuaRouteTestContext() *luaRouteTestContext {
|
|
||||||
return &luaRouteTestContext{
|
|
||||||
Context: &routing_session.Context{
|
|
||||||
Inbound: &session.Inbound{
|
|
||||||
Tag: "in", VlessRoute: 4321,
|
|
||||||
Source: net.TCPDestination(net.LocalHostIP, 1234),
|
|
||||||
Local: net.TCPDestination(net.LocalHostIP, 5678),
|
|
||||||
User: &protocol.MemoryUser{Email: "user@example.com"},
|
|
||||||
},
|
|
||||||
Outbound: &session.Outbound{
|
|
||||||
Target: net.TCPDestination(net.LocalHostIP, 443),
|
|
||||||
RouteTarget: net.TCPDestination(net.DomainAddress("MiXeD.Example."), 443),
|
|
||||||
},
|
|
||||||
Content: &session.Content{
|
|
||||||
Protocol: "tls", Attributes: map[string]string{"key": "value"}, SkipDNSResolve: true,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
sourceIPs: []net.IP{{127, 0, 0, 2}},
|
|
||||||
targetIPs: []net.IP{{127, 0, 0, 3}},
|
|
||||||
localIPs: []net.IP{{127, 0, 0, 1}},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLuaRouterState(t *testing.T, script string) (*Router, *lua.LState) {
|
|
||||||
t.Helper()
|
|
||||||
r := new(Router)
|
|
||||||
if err := r.Init(context.Background(), &Config{}, nil, nil, nil); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
L := lua.NewState()
|
|
||||||
t.Cleanup(L.Close)
|
|
||||||
r.RegisterLua(L)
|
|
||||||
geodata.RegisterLua(L)
|
|
||||||
if err := L.DoString(script); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
return r, L
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaRouteBinding(t *testing.T) {
|
|
||||||
r, L := newLuaRouterState(t, `
|
|
||||||
local router = require("xray.router")
|
|
||||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
|
||||||
assert(router.NetworkUnknown == 0 and router.NetworkTCP == 2)
|
|
||||||
assert(router.NetworkUDP == 3 and router.NetworkUNIX == 4)
|
|
||||||
assert(router.BuildIPMatcher == nil and router.BuildDomainMatcher == nil)
|
|
||||||
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
|
||||||
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve, ...)
|
|
||||||
assert(select("#", ...) == 0)
|
|
||||||
assert(inboundTag == "in" and sourcePort == 1234 and targetPort == 443 and localPort == 5678)
|
|
||||||
assert(targetDomain == "mixed.example." and network == router.NetworkTCP)
|
|
||||||
assert(protocol == "tls" and user == "user@example.com" and vlessRoute == 4321 and skipDNSResolve)
|
|
||||||
assert(ctx.GetNetwork == nil and ctx.Context == nil)
|
|
||||||
savedContext = ctx
|
|
||||||
sourceIPs, targetIPs, localIPs = ctx:GetSourceIPs(), ctx:GetTargetIPs(), ctx:GetLocalIPs()
|
|
||||||
attributes = ctx:GetAttributes()
|
|
||||||
assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
|
|
||||||
assert(attributes.key == "value" and attributes.missing == nil)
|
|
||||||
assert(not pcall(function() attributes.key = "changed" end))
|
|
||||||
return "out", "rule"
|
|
||||||
end`)
|
|
||||||
|
|
||||||
ctx := newLuaRouteTestContext()
|
|
||||||
tag, rule, err := r.callLuaHook(L, ctx)
|
|
||||||
if err != nil || tag != "out" || rule != "rule" {
|
|
||||||
t.Fatalf("hook = %q, %q, %v", tag, rule, err)
|
|
||||||
}
|
|
||||||
if L.GetGlobal("savedContext").(*lua.LUserData).Value != ctx {
|
|
||||||
t.Fatal("routing context was copied")
|
|
||||||
}
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
want []net.IP
|
|
||||||
}{
|
|
||||||
{"sourceIPs", ctx.sourceIPs},
|
|
||||||
{"targetIPs", ctx.targetIPs},
|
|
||||||
{"localIPs", ctx.localIPs},
|
|
||||||
} {
|
|
||||||
got := L.GetGlobal(tc.name).(*lua.LUserData).Value.([]net.IP)
|
|
||||||
if &got[0] != &tc.want[0] {
|
|
||||||
t.Fatalf("%s storage was copied", tc.name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
ctx.Content.Attributes["key"] = "updated"
|
|
||||||
L.SetGlobal("expectedOS", lua.LString(runtime.GOOS))
|
|
||||||
if err := L.DoString(`
|
|
||||||
assert(attributes.key == "updated")
|
|
||||||
assert(require("xray.router").LocalOS == expectedOS)`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaRouteResult(t *testing.T) {
|
|
||||||
nativeErr := go_errors.New("native failure")
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name, body, tag, rule, wantErr string
|
|
||||||
native bool
|
|
||||||
}{
|
|
||||||
{name: "route", body: `return "out", "rule"`, tag: "out", rule: "rule"},
|
|
||||||
{name: "no match", body: `return nil`},
|
|
||||||
{name: "empty tag", body: `return ""`},
|
|
||||||
{name: "no match ignores rule", body: `return nil, false`},
|
|
||||||
{name: "empty tag ignores rule", body: `return "", false`},
|
|
||||||
{name: "missing rule", body: `return "out"`, tag: "out"},
|
|
||||||
{name: "invalid tag", body: `return 1`, wantErr: "outboundTag"},
|
|
||||||
{name: "invalid rule", body: `return "out", false`, wantErr: "ruleTag"},
|
|
||||||
{name: "string error", body: `return nil, nil, "script failure"`, wantErr: "script failure"},
|
|
||||||
{name: "native error", body: `return nil, nil, nativeError`, native: true},
|
|
||||||
{name: "error overrides invalid tags", body: `return false, false, nativeError`, native: true},
|
|
||||||
{name: "invalid error", body: `return "out", "rule", false`, wantErr: "error or string"},
|
|
||||||
{name: "wrong error userdata", body: `return "out", "rule", wrongError`, wantErr: "error or string"},
|
|
||||||
{name: "runtime error", body: `error("runtime failure")`, wantErr: "runtime failure"},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
r, L := newLuaRouterState(t, "function HandleRoute() "+tc.body+" end")
|
|
||||||
value := L.NewUserData()
|
|
||||||
value.Value = nativeErr
|
|
||||||
L.SetGlobal("nativeError", value)
|
|
||||||
wrong := L.NewUserData()
|
|
||||||
wrong.Value = "not a native error"
|
|
||||||
L.SetGlobal("wrongError", wrong)
|
|
||||||
L.Push(lua.LTrue)
|
|
||||||
|
|
||||||
tag, rule, err := r.callLuaHook(L, &routing_session.Context{})
|
|
||||||
if tag != tc.tag || rule != tc.rule {
|
|
||||||
t.Fatalf("result = %q, %q, %v", tag, rule, err)
|
|
||||||
}
|
|
||||||
switch {
|
|
||||||
case tc.native:
|
|
||||||
if err != nativeErr {
|
|
||||||
t.Fatalf("error = %v, want original error", err)
|
|
||||||
}
|
|
||||||
case tc.wantErr != "":
|
|
||||||
if err == nil || !strings.Contains(err.Error(), tc.wantErr) {
|
|
||||||
t.Fatalf("error = %v, want %q", err, tc.wantErr)
|
|
||||||
}
|
|
||||||
case err != nil:
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
|
||||||
t.Fatal("hook did not restore the stack")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaRouteCancellation(t *testing.T) {
|
|
||||||
r, L := newLuaRouterState(t, `function HandleRoute() while true do end end`)
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
cancel()
|
|
||||||
L.SetContext(ctx)
|
|
||||||
if _, _, err := r.callLuaHook(L, &routing_session.Context{}); err == nil {
|
|
||||||
t.Fatal("CallLuaHook did not stop after context cancellation")
|
|
||||||
}
|
|
||||||
if L.Context() != ctx || L.GetTop() != 0 {
|
|
||||||
t.Fatal("CallLuaHook did not restore the Lua state")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFindProcess(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name, network, target string
|
|
||||||
targetPort uint16
|
|
||||||
modify func(*luaRouteTestContext)
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{name: "TCP", network: "tcp", target: "127.0.0.3", targetPort: 443},
|
|
||||||
{name: "UDP", network: "udp", target: "127.0.0.3", targetPort: 443, modify: func(c *luaRouteTestContext) {
|
|
||||||
c.Outbound.Target.Network = net.Network_UDP
|
|
||||||
}},
|
|
||||||
{name: "domain target", network: "tcp", modify: func(c *luaRouteTestContext) { c.targetIPs = nil }},
|
|
||||||
{name: "missing source", modify: func(c *luaRouteTestContext) { c.sourceIPs = nil }, wantErr: true},
|
|
||||||
{name: "unsupported network", modify: func(c *luaRouteTestContext) {
|
|
||||||
c.Outbound.Target.Network = net.Network_UNIX
|
|
||||||
}, wantErr: true},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
ctx := newLuaRouteTestContext()
|
|
||||||
if tc.modify != nil {
|
|
||||||
tc.modify(ctx)
|
|
||||||
}
|
|
||||||
called := false
|
|
||||||
pid, name, path, err := findProcess(ctx, func(network, source string, sourcePort uint16, target string, targetPort uint16) (int, string, string, error) {
|
|
||||||
called = true
|
|
||||||
if network != tc.network || source != "127.0.0.2" || sourcePort != 1234 || target != tc.target || targetPort != tc.targetPort {
|
|
||||||
t.Fatalf("endpoints = %s %s:%d -> %s:%d", network, source, sourcePort, target, targetPort)
|
|
||||||
}
|
|
||||||
return 42, "process", "/path/process", nil
|
|
||||||
})
|
|
||||||
if tc.wantErr {
|
|
||||||
if err == nil || called {
|
|
||||||
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if err != nil || !called || pid != 42 || name != "process" || path != "/path/process" {
|
|
||||||
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkLuaRouteHookCall isolates a preloaded Lua hook and its routing context bridge.
|
|
||||||
// The direct case runs an equivalent native routing rule.
|
|
||||||
func BenchmarkLuaRouteHookCall(b *testing.B) {
|
|
||||||
r := new(Router)
|
|
||||||
if err := r.Init(context.Background(), &Config{Rule: []*RoutingRule{{
|
|
||||||
TargetTag: &RoutingRule_Tag{Tag: "out"},
|
|
||||||
RuleTag: "rule",
|
|
||||||
InboundTag: []string{"in"},
|
|
||||||
Networks: []net.Network{net.Network_TCP},
|
|
||||||
Ip: []*geodata.IPRule{{
|
|
||||||
Value: &geodata.IPRule_Custom{Custom: &geodata.CIDRRule{
|
|
||||||
Cidr: &geodata.CIDR{Ip: []byte{127, 0, 0, 0}, Prefix: 8},
|
|
||||||
}},
|
|
||||||
}},
|
|
||||||
}}}, nil, nil, nil); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
r.RegisterLua(L)
|
|
||||||
geodata.RegisterLua(L)
|
|
||||||
if err := L.DoString(`
|
|
||||||
local router = require("xray.router")
|
|
||||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
|
||||||
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
|
||||||
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
|
|
||||||
if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then
|
|
||||||
return "out", "rule"
|
|
||||||
end
|
|
||||||
end
|
|
||||||
`); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
L.SetContext(context.Background())
|
|
||||||
routeCtx := newLuaRouteTestContext()
|
|
||||||
for _, benchmark := range []struct {
|
|
||||||
name string
|
|
||||||
route func() (string, string, error)
|
|
||||||
}{
|
|
||||||
{"direct", func() (string, string, error) {
|
|
||||||
route, err := r.PickRoute(routeCtx)
|
|
||||||
if err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
return route.GetOutboundTag(), route.GetRuleTag(), nil
|
|
||||||
}},
|
|
||||||
{"lua_hook", func() (string, string, error) {
|
|
||||||
return r.callLuaHook(L, routeCtx)
|
|
||||||
}},
|
|
||||||
} {
|
|
||||||
b.Run(benchmark.name, func(b *testing.B) {
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
var tag, rule string
|
|
||||||
var err error
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
tag, rule, err = benchmark.route()
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
b.StopTimer()
|
|
||||||
if tag != "out" || rule != "rule" {
|
|
||||||
b.Fatalf("route() = %q, %q; want out, rule", tag, rule)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var _ routing.Context = (*luaRouteTestContext)(nil)
|
|
||||||
+116
-76
@@ -2,9 +2,7 @@ package router
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"maps"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
@@ -19,10 +17,8 @@ import (
|
|||||||
// Router is an implementation of routing.Router.
|
// Router is an implementation of routing.Router.
|
||||||
type Router struct {
|
type Router struct {
|
||||||
domainStrategy Config_DomainStrategy
|
domainStrategy Config_DomainStrategy
|
||||||
rules atomic.Pointer[[]*Rule]
|
rules []*Rule
|
||||||
scriptPath string
|
balancers map[string]*Balancer
|
||||||
script *scriptEngine
|
|
||||||
balancers atomic.Pointer[map[string]*Balancer]
|
|
||||||
dns dns.Client
|
dns dns.Client
|
||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
@@ -42,23 +38,61 @@ type Route struct {
|
|||||||
// Init initializes the Router.
|
// Init initializes the Router.
|
||||||
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
|
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
|
||||||
r.domainStrategy = config.DomainStrategy
|
r.domainStrategy = config.DomainStrategy
|
||||||
r.scriptPath = config.Script
|
|
||||||
r.dns = d
|
r.dns = d
|
||||||
r.ctx = ctx
|
r.ctx = ctx
|
||||||
r.ohm = ohm
|
r.ohm = ohm
|
||||||
r.dispatcher = dispatcher
|
r.dispatcher = dispatcher
|
||||||
|
|
||||||
r.rules.Store(new([]*Rule))
|
r.balancers = make(map[string]*Balancer, len(config.BalancingRule))
|
||||||
r.balancers.Store(&map[string]*Balancer{})
|
for _, rule := range config.BalancingRule {
|
||||||
return r.ReloadRules(config, false)
|
balancer, err := rule.Build(ohm, dispatcher)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
balancer.InjectContext(ctx)
|
||||||
|
r.balancers[rule.Tag] = balancer
|
||||||
|
}
|
||||||
|
|
||||||
|
r.rules = make([]*Rule, 0, len(config.Rule))
|
||||||
|
for _, rule := range config.Rule {
|
||||||
|
cond, err := rule.BuildCondition()
|
||||||
|
if err != nil {
|
||||||
|
r.closeWebhooks()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
rr := &Rule{
|
||||||
|
Condition: cond,
|
||||||
|
Tag: rule.GetTag(),
|
||||||
|
RuleTag: rule.GetRuleTag(),
|
||||||
|
}
|
||||||
|
if wh := rule.GetWebhook(); wh != nil {
|
||||||
|
notifier, err := NewWebhookNotifier(wh)
|
||||||
|
if err != nil {
|
||||||
|
r.closeWebhooks()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
rr.Webhook = notifier
|
||||||
|
}
|
||||||
|
btag := rule.GetBalancingTag()
|
||||||
|
if len(btag) > 0 {
|
||||||
|
brule, found := r.balancers[btag]
|
||||||
|
if !found {
|
||||||
|
if rr.Webhook != nil {
|
||||||
|
rr.Webhook.Close()
|
||||||
|
}
|
||||||
|
r.closeWebhooks()
|
||||||
|
return errors.New("balancer ", btag, " not found")
|
||||||
|
}
|
||||||
|
rr.Balancer = brule
|
||||||
|
}
|
||||||
|
r.rules = append(r.rules, rr)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// PickRoute implements routing.Router.
|
// PickRoute implements routing.Router.
|
||||||
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
|
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
|
||||||
if r.script != nil {
|
|
||||||
return r.script.pickRoute(ctx)
|
|
||||||
}
|
|
||||||
|
|
||||||
originalCtx := ctx
|
originalCtx := ctx
|
||||||
rule, ctx, err := r.pickRouteInternal(ctx)
|
rule, ctx, err := r.pickRouteInternal(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -76,6 +110,7 @@ func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
|
|||||||
|
|
||||||
// AddRule implements routing.Router.
|
// AddRule implements routing.Router.
|
||||||
func (r *Router) AddRule(config *serial.TypedMessage, shouldAppend bool) error {
|
func (r *Router) AddRule(config *serial.TypedMessage, shouldAppend bool) error {
|
||||||
|
|
||||||
inst, err := config.GetInstance()
|
inst, err := config.GetInstance()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -90,22 +125,18 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
|
|||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
oldRules := *r.rules.Load()
|
if !shouldAppend {
|
||||||
oldBalancers := *r.balancers.Load()
|
for _, rule := range r.rules {
|
||||||
|
if rule.Webhook != nil {
|
||||||
var newRules []*Rule
|
rule.Webhook.Close()
|
||||||
newBalancers := make(map[string]*Balancer)
|
}
|
||||||
existTags := make(map[string]bool, len(oldRules)+len(config.Rule))
|
|
||||||
if shouldAppend {
|
|
||||||
newRules = append(newRules, oldRules...)
|
|
||||||
maps.Copy(newBalancers, oldBalancers)
|
|
||||||
for _, rule := range oldRules {
|
|
||||||
existTags[rule.RuleTag] = true
|
|
||||||
}
|
}
|
||||||
|
r.balancers = make(map[string]*Balancer, len(config.BalancingRule))
|
||||||
|
r.rules = make([]*Rule, 0, len(config.Rule))
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, rule := range config.BalancingRule {
|
for _, rule := range config.BalancingRule {
|
||||||
if _, found := newBalancers[rule.Tag]; found {
|
_, found := r.balancers[rule.Tag]
|
||||||
|
if found {
|
||||||
return errors.New("duplicate balancer tag")
|
return errors.New("duplicate balancer tag")
|
||||||
}
|
}
|
||||||
balancer, err := rule.Build(r.ohm, r.dispatcher)
|
balancer, err := rule.Build(r.ohm, r.dispatcher)
|
||||||
@@ -113,12 +144,27 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
balancer.InjectContext(r.ctx)
|
balancer.InjectContext(r.ctx)
|
||||||
newBalancers[rule.Tag] = balancer
|
r.balancers[rule.Tag] = balancer
|
||||||
|
}
|
||||||
|
|
||||||
|
startIdx := len(r.rules)
|
||||||
|
closeNewWebhooks := func() {
|
||||||
|
for i := startIdx; i < len(r.rules); i++ {
|
||||||
|
if r.rules[i].Webhook != nil {
|
||||||
|
r.rules[i].Webhook.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
r.rules = r.rules[:startIdx]
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, rule := range config.Rule {
|
for _, rule := range config.Rule {
|
||||||
|
if r.RuleExists(rule.GetRuleTag()) {
|
||||||
|
closeNewWebhooks()
|
||||||
|
return errors.New("duplicate ruleTag ", rule.GetRuleTag())
|
||||||
|
}
|
||||||
cond, err := rule.BuildCondition()
|
cond, err := rule.BuildCondition()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
closeNewWebhooks()
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
rr := &Rule{
|
rr := &Rule{
|
||||||
@@ -126,64 +172,70 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
|
|||||||
Tag: rule.GetTag(),
|
Tag: rule.GetTag(),
|
||||||
RuleTag: rule.GetRuleTag(),
|
RuleTag: rule.GetRuleTag(),
|
||||||
}
|
}
|
||||||
if rr.RuleTag != "" && existTags[rr.RuleTag] {
|
|
||||||
return errors.New("duplicate ruleTag ", rr.RuleTag)
|
|
||||||
}
|
|
||||||
existTags[rr.RuleTag] = true
|
|
||||||
if wh := rule.GetWebhook(); wh != nil {
|
if wh := rule.GetWebhook(); wh != nil {
|
||||||
notifier, err := NewWebhookNotifier(wh)
|
notifier, err := NewWebhookNotifier(wh)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
closeNewWebhooks()
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
rr.Webhook = notifier
|
rr.Webhook = notifier
|
||||||
}
|
}
|
||||||
if btag := rule.GetBalancingTag(); len(btag) > 0 {
|
btag := rule.GetBalancingTag()
|
||||||
brule, found := newBalancers[btag]
|
if len(btag) > 0 {
|
||||||
|
brule, found := r.balancers[btag]
|
||||||
if !found {
|
if !found {
|
||||||
|
if rr.Webhook != nil {
|
||||||
|
rr.Webhook.Close()
|
||||||
|
}
|
||||||
|
closeNewWebhooks()
|
||||||
return errors.New("balancer ", btag, " not found")
|
return errors.New("balancer ", btag, " not found")
|
||||||
}
|
}
|
||||||
rr.Balancer = brule
|
rr.Balancer = brule
|
||||||
}
|
}
|
||||||
newRules = append(newRules, rr)
|
r.rules = append(r.rules, rr)
|
||||||
}
|
}
|
||||||
|
|
||||||
r.balancers.Store(&newBalancers)
|
|
||||||
r.rules.Store(&newRules)
|
|
||||||
if !shouldAppend {
|
|
||||||
closeWebhooks(oldRules)
|
|
||||||
}
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *Router) RuleExists(tag string) bool {
|
||||||
|
if tag != "" {
|
||||||
|
for _, rule := range r.rules {
|
||||||
|
if rule.RuleTag == tag {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
// RemoveRule implements routing.Router.
|
// RemoveRule implements routing.Router.
|
||||||
func (r *Router) RemoveRule(tag string) error {
|
func (r *Router) RemoveRule(tag string) error {
|
||||||
if tag == "" {
|
|
||||||
return errors.New("empty tag name!")
|
|
||||||
}
|
|
||||||
|
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
oldRules := *r.rules.Load()
|
newRules := []*Rule{}
|
||||||
newRules := make([]*Rule, 0, len(oldRules))
|
if tag != "" {
|
||||||
var removed []*Rule
|
for _, rule := range r.rules {
|
||||||
for _, rule := range oldRules {
|
if rule.RuleTag != tag {
|
||||||
if rule.RuleTag != tag {
|
newRules = append(newRules, rule)
|
||||||
newRules = append(newRules, rule)
|
} else if rule.Webhook != nil {
|
||||||
} else {
|
rule.Webhook.Close()
|
||||||
removed = append(removed, rule)
|
}
|
||||||
}
|
}
|
||||||
|
r.rules = newRules
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
r.rules.Store(&newRules)
|
return errors.New("empty tag name!")
|
||||||
closeWebhooks(removed)
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListRule implements routing.Router
|
// ListRule implements routing.Router
|
||||||
func (r *Router) ListRule() []routing.Route {
|
func (r *Router) ListRule() []routing.Route {
|
||||||
rules := *r.rules.Load()
|
r.mu.Lock()
|
||||||
ruleList := make([]routing.Route, 0, len(rules))
|
defer r.mu.Unlock()
|
||||||
for _, rule := range rules {
|
ruleList := make([]routing.Route, 0)
|
||||||
|
for _, rule := range r.rules {
|
||||||
ruleList = append(ruleList, &Route{
|
ruleList = append(ruleList, &Route{
|
||||||
outboundTag: rule.Tag,
|
outboundTag: rule.Tag,
|
||||||
ruleTag: rule.RuleTag,
|
ruleTag: rule.RuleTag,
|
||||||
@@ -202,9 +254,7 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
|||||||
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
|
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
|
||||||
}
|
}
|
||||||
|
|
||||||
rules := *r.rules.Load()
|
for _, rule := range r.rules {
|
||||||
|
|
||||||
for _, rule := range rules {
|
|
||||||
if rule.Apply(ctx) {
|
if rule.Apply(ctx) {
|
||||||
return rule, ctx, nil
|
return rule, ctx, nil
|
||||||
}
|
}
|
||||||
@@ -217,7 +267,7 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
|||||||
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
|
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
|
||||||
|
|
||||||
// Try applying rules again if we have IPs.
|
// Try applying rules again if we have IPs.
|
||||||
for _, rule := range rules {
|
for _, rule := range r.rules {
|
||||||
if rule.Apply(ctx) {
|
if rule.Apply(ctx) {
|
||||||
return rule, ctx, nil
|
return rule, ctx, nil
|
||||||
}
|
}
|
||||||
@@ -228,19 +278,12 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
|||||||
|
|
||||||
// Start implements common.Runnable.
|
// Start implements common.Runnable.
|
||||||
func (r *Router) Start() error {
|
func (r *Router) Start() error {
|
||||||
if r.scriptPath != "" {
|
|
||||||
engine, err := newScriptEngine(r.scriptPath, r)
|
|
||||||
if err != nil {
|
|
||||||
return errors.New("failed to initialize routing script").Base(err)
|
|
||||||
}
|
|
||||||
r.script = engine
|
|
||||||
}
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// closeWebhooks closes all webhook notifiers in the given rule set.
|
// closeWebhooks closes all webhook notifiers in the current rule set.
|
||||||
func closeWebhooks(rules []*Rule) {
|
func (r *Router) closeWebhooks() {
|
||||||
for _, rule := range rules {
|
for _, rule := range r.rules {
|
||||||
if rule.Webhook != nil {
|
if rule.Webhook != nil {
|
||||||
rule.Webhook.Close()
|
rule.Webhook.Close()
|
||||||
}
|
}
|
||||||
@@ -249,12 +292,9 @@ func closeWebhooks(rules []*Rule) {
|
|||||||
|
|
||||||
// Close implements common.Closable.
|
// Close implements common.Closable.
|
||||||
func (r *Router) Close() error {
|
func (r *Router) Close() error {
|
||||||
if r.script != nil {
|
|
||||||
r.script.close()
|
|
||||||
}
|
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
closeWebhooks(*r.rules.Load())
|
r.closeWebhooks()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,68 +0,0 @@
|
|||||||
package router
|
|
||||||
|
|
||||||
import (
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/app/dns"
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
|
||||||
"github.com/xtls/xray-core/common/log"
|
|
||||||
xlua "github.com/xtls/xray-core/common/lua"
|
|
||||||
"github.com/xtls/xray-core/features/routing"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
const scriptExecutionTimeout = 6 * time.Second
|
|
||||||
|
|
||||||
type scriptEngine struct {
|
|
||||||
router *Router
|
|
||||||
pool *xlua.Pool
|
|
||||||
}
|
|
||||||
|
|
||||||
func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
|
|
||||||
program, err := xlua.CompileFile(path)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
e := &scriptEngine{router: router}
|
|
||||||
e.pool, err = xlua.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory(
|
|
||||||
scriptExecutionTimeout*20,
|
|
||||||
func(L *lua.LState) {
|
|
||||||
geodata.RegisterLua(L)
|
|
||||||
log.RegisterLua(L)
|
|
||||||
router.RegisterLua(L)
|
|
||||||
dns.RegisterLua(L, router.dns)
|
|
||||||
},
|
|
||||||
func(L *lua.LState) error {
|
|
||||||
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
|
|
||||||
return errors.New("routing script must define HandleRoute(...)")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
errors.LogInfo(router.ctx, "routing script initialized from ", path)
|
|
||||||
return e, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *scriptEngine) close() {
|
|
||||||
e.pool.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
|
|
||||||
var tag, ruleTag string
|
|
||||||
err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
|
||||||
var hookErr error
|
|
||||||
tag, ruleTag, hookErr = e.router.callLuaHook(L, ctx)
|
|
||||||
return hookErr
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if tag == "" {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
}
|
|
||||||
return &Route{Context: ctx, outboundTag: tag, ruleTag: ruleTag}, nil
|
|
||||||
}
|
|
||||||
@@ -1,372 +0,0 @@
|
|||||||
package router
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
stdnet "net"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
wireDNS "github.com/miekg/dns"
|
|
||||||
"github.com/xtls/xray-core/app/dispatcher"
|
|
||||||
appdns "github.com/xtls/xray-core/app/dns"
|
|
||||||
"github.com/xtls/xray-core/app/proxyman"
|
|
||||||
_ "github.com/xtls/xray-core/app/proxyman/outbound"
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
|
||||||
"github.com/xtls/xray-core/core"
|
|
||||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
|
||||||
"github.com/xtls/xray-core/features/outbound"
|
|
||||||
"github.com/xtls/xray-core/features/routing"
|
|
||||||
routing_session "github.com/xtls/xray-core/features/routing/session"
|
|
||||||
"github.com/xtls/xray-core/proxy/blackhole"
|
|
||||||
"github.com/xtls/xray-core/proxy/freedom"
|
|
||||||
)
|
|
||||||
|
|
||||||
type luaRouteDNSClient struct {
|
|
||||||
featureDNS.Client
|
|
||||||
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *luaRouteDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
return d.lookup(domain, option)
|
|
||||||
}
|
|
||||||
|
|
||||||
type luaRouteOutboundManager struct{ outbound.Manager }
|
|
||||||
|
|
||||||
func (*luaRouteOutboundManager) Select(selectors []string) []string { return selectors }
|
|
||||||
|
|
||||||
func writeRouteScript(t *testing.T, script string) string {
|
|
||||||
t.Helper()
|
|
||||||
path := filepath.Join(t.TempDir(), "route.lua")
|
|
||||||
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
return path
|
|
||||||
}
|
|
||||||
|
|
||||||
func startLuaRouter(t *testing.T, script string, d featureDNS.Client, config *Config) *Router {
|
|
||||||
t.Helper()
|
|
||||||
if config == nil {
|
|
||||||
config = &Config{}
|
|
||||||
}
|
|
||||||
config.Script = writeRouteScript(t, script)
|
|
||||||
r := new(Router)
|
|
||||||
if err := r.Init(context.Background(), config, d, &luaRouteOutboundManager{}, nil); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := r.Start(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
if err := r.Close(); err != nil {
|
|
||||||
t.Error(err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
return r
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRouterScriptStartup(t *testing.T) {
|
|
||||||
for _, tc := range []struct{ name, script string }{
|
|
||||||
{"syntax error", "function HandleRoute("},
|
|
||||||
{"missing hook", "value = 1"},
|
|
||||||
{"initialization error", `error("setup failed")`},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
r := new(Router)
|
|
||||||
if err := r.Init(context.Background(), &Config{Script: writeRouteScript(t, tc.script)}, nil, nil, nil); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
defer r.Close()
|
|
||||||
if err := r.Start(); err == nil {
|
|
||||||
t.Fatal("Start accepted an invalid routing script")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRouterScriptRouting(t *testing.T) {
|
|
||||||
var dnsCalls atomic.Int32
|
|
||||||
d := &luaRouteDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
dnsCalls.Add(1)
|
|
||||||
return []net.IP{{1, 2, 3, 4}}, 60, nil
|
|
||||||
}}
|
|
||||||
r := startLuaRouter(t, `
|
|
||||||
function HandleRoute(ctx, inbound)
|
|
||||||
if inbound == "miss" then return nil end
|
|
||||||
return "lua-out", "lua-rule"
|
|
||||||
end`, d, &Config{
|
|
||||||
DomainStrategy: Config_IpOnDemand,
|
|
||||||
Rule: []*RoutingRule{{
|
|
||||||
TargetTag: &RoutingRule_Tag{Tag: "json-out"},
|
|
||||||
Networks: []net.Network{net.Network_TCP},
|
|
||||||
}},
|
|
||||||
})
|
|
||||||
ctx := newLuaRouteTestContext()
|
|
||||||
ctx.Content.SkipDNSResolve = false
|
|
||||||
route, err := r.PickRoute(ctx)
|
|
||||||
if err != nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != "lua-rule" || route.(*Route).Context != ctx {
|
|
||||||
t.Fatalf("route = %v, %v", route, err)
|
|
||||||
}
|
|
||||||
ctx.Inbound.Tag = "miss"
|
|
||||||
if route, err := r.PickRoute(ctx); route != nil || err != common.ErrNoClue {
|
|
||||||
t.Fatalf("miss = %v, %v", route, err)
|
|
||||||
}
|
|
||||||
if dnsCalls.Load() != 0 {
|
|
||||||
t.Fatal("script routing implicitly resolved DNS")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRouterScriptModules(t *testing.T) {
|
|
||||||
ips := []net.IP{{127, 0, 0, 7}}
|
|
||||||
calls := 0
|
|
||||||
d := &luaRouteDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
calls++
|
|
||||||
if domain != "mixed.example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
|
|
||||||
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
|
|
||||||
}
|
|
||||||
return ips, 17, nil
|
|
||||||
}}
|
|
||||||
r := startLuaRouter(t, `
|
|
||||||
local dns = require("xray.dns")
|
|
||||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
|
||||||
assert(dns.Servers == nil and type(dns.Query) == "function")
|
|
||||||
assert(type(require("xray.log").Info) == "function")
|
|
||||||
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain)
|
|
||||||
local ips, ttl, err = dns.Query(domain, true, false, true)
|
|
||||||
assert(not err and ttl == 17)
|
|
||||||
assert(matcher:AnyMatch(ips) and matcher:AnyMatch(ctx:GetTargetIPs()))
|
|
||||||
return "out"
|
|
||||||
end`, d, nil)
|
|
||||||
if _, err := r.PickRoute(newLuaRouteTestContext()); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if calls != 1 {
|
|
||||||
t.Fatalf("DNS calls = %d, want 1", calls)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRouterScriptBalancerReload(t *testing.T) {
|
|
||||||
config := func(tag string) *Config {
|
|
||||||
return &Config{BalancingRule: []*BalancingRule{{
|
|
||||||
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
|
|
||||||
}}}
|
|
||||||
}
|
|
||||||
r := startLuaRouter(t, `
|
|
||||||
local router = require("xray.router")
|
|
||||||
function HandleRoute()
|
|
||||||
local tag, err = router:PickOutbound("balance")
|
|
||||||
return tag, "balanced", err
|
|
||||||
end`, nil, config("old"))
|
|
||||||
pick := func(want string) {
|
|
||||||
t.Helper()
|
|
||||||
route, err := r.PickRoute(&routing_session.Context{})
|
|
||||||
if err != nil || route.GetOutboundTag() != want || route.GetRuleTag() != "balanced" {
|
|
||||||
t.Fatalf("route = %v, %v, want %q", route, err, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
pick("old")
|
|
||||||
if err := r.SetOverrideTarget("balance", "override"); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
pick("override")
|
|
||||||
if err := r.SetOverrideTarget("balance", ""); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := r.ReloadRules(config("new"), false); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
pick("new")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRouterScriptConcurrentBalancerReload(t *testing.T) {
|
|
||||||
config := func(tag string) *Config {
|
|
||||||
return &Config{BalancingRule: []*BalancingRule{{
|
|
||||||
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
|
|
||||||
}}}
|
|
||||||
}
|
|
||||||
r := startLuaRouter(t, `
|
|
||||||
local router = require("xray.router")
|
|
||||||
function HandleRoute()
|
|
||||||
local tag, err = router:PickOutbound("balance")
|
|
||||||
return tag, nil, err
|
|
||||||
end`, nil, config("a"))
|
|
||||||
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
for range 4 {
|
|
||||||
wg.Go(func() {
|
|
||||||
for range 20 {
|
|
||||||
route, err := r.PickRoute(&routing_session.Context{})
|
|
||||||
if err != nil {
|
|
||||||
t.Errorf("PickRoute: %v", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if tag := route.GetOutboundTag(); tag != "a" && tag != "b" {
|
|
||||||
t.Errorf("unexpected tag %q", tag)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
wg.Go(func() {
|
|
||||||
for range 20 {
|
|
||||||
for _, tag := range []string{"a", "b"} {
|
|
||||||
if err := r.ReloadRules(config(tag), false); err != nil {
|
|
||||||
t.Error(err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
wg.Wait()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRouterScriptStateReuse(t *testing.T) {
|
|
||||||
r := startLuaRouter(t, `
|
|
||||||
local calls = 0
|
|
||||||
function HandleRoute(ctx, inbound)
|
|
||||||
calls = calls + 1
|
|
||||||
if inbound == "miss" then return nil end
|
|
||||||
if inbound == "fail" then error("failed") end
|
|
||||||
return tostring(calls)
|
|
||||||
end`, nil, nil)
|
|
||||||
ctx := newLuaRouteTestContext()
|
|
||||||
pick := func(want string) {
|
|
||||||
t.Helper()
|
|
||||||
route, err := r.PickRoute(ctx)
|
|
||||||
if err != nil || route.GetOutboundTag() != want {
|
|
||||||
t.Fatalf("route = %v, %v, want %q", route, err, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
pick("1")
|
|
||||||
ctx.Inbound.Tag = "miss"
|
|
||||||
if _, err := r.PickRoute(ctx); err != common.ErrNoClue {
|
|
||||||
t.Fatalf("miss = %v", err)
|
|
||||||
}
|
|
||||||
ctx.Inbound.Tag = "in"
|
|
||||||
pick("3")
|
|
||||||
ctx.Inbound.Tag = "fail"
|
|
||||||
if _, err := r.PickRoute(ctx); err == nil {
|
|
||||||
t.Fatal("script error was ignored")
|
|
||||||
}
|
|
||||||
ctx.Inbound.Tag = "in"
|
|
||||||
pick("1")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRouterScriptDNSDispatcherReentry(t *testing.T) {
|
|
||||||
conn, err := stdnet.ListenPacket("udp4", "127.0.0.1:0")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
port := conn.LocalAddr().(*stdnet.UDPAddr).Port
|
|
||||||
ready, stopped := make(chan struct{}), make(chan error, 1)
|
|
||||||
var queries atomic.Int32
|
|
||||||
server := &wireDNS.Server{
|
|
||||||
PacketConn: conn,
|
|
||||||
NotifyStartedFunc: func() {
|
|
||||||
close(ready)
|
|
||||||
},
|
|
||||||
Handler: wireDNS.HandlerFunc(func(w wireDNS.ResponseWriter, query *wireDNS.Msg) {
|
|
||||||
queries.Add(1)
|
|
||||||
response := new(wireDNS.Msg).SetReply(query)
|
|
||||||
for _, question := range query.Question {
|
|
||||||
if question.Name == "nested.example." && question.Qtype == wireDNS.TypeA {
|
|
||||||
response.Answer = append(response.Answer, &wireDNS.A{
|
|
||||||
Hdr: wireDNS.RR_Header{Name: question.Name, Rrtype: wireDNS.TypeA, Class: wireDNS.ClassINET, Ttl: 60},
|
|
||||||
A: stdnet.IP{127, 0, 0, 7},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := w.WriteMsg(response); err != nil {
|
|
||||||
t.Error(err)
|
|
||||||
}
|
|
||||||
}),
|
|
||||||
}
|
|
||||||
go func() { stopped <- server.ActivateAndServe() }()
|
|
||||||
defer func() {
|
|
||||||
server.Shutdown()
|
|
||||||
select {
|
|
||||||
case err := <-stopped:
|
|
||||||
if err != nil {
|
|
||||||
t.Error(err)
|
|
||||||
}
|
|
||||||
case <-time.After(3 * time.Second):
|
|
||||||
t.Error("DNS server did not stop")
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
select {
|
|
||||||
case <-ready:
|
|
||||||
case err := <-stopped:
|
|
||||||
t.Fatalf("DNS server startup: %v", err)
|
|
||||||
case <-time.After(3 * time.Second):
|
|
||||||
t.Fatal("DNS server did not start")
|
|
||||||
}
|
|
||||||
|
|
||||||
dnsScript := writeRouteScript(t, `
|
|
||||||
local server = require("xray.dns").Servers[1]
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
return server:Query(domain, ipv4, ipv6, fake)
|
|
||||||
end`)
|
|
||||||
routerScript := writeRouteScript(t, `
|
|
||||||
local router = require("xray.router")
|
|
||||||
local dns = require("xray.dns")
|
|
||||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.7")
|
|
||||||
local active = false
|
|
||||||
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain, network,
|
|
||||||
protocol, user, vlessRoute, skipDNSResolve)
|
|
||||||
assert(not active, "borrowed Router VM reentered")
|
|
||||||
if inbound == "dns" then
|
|
||||||
assert(network == router.NetworkUDP and skipDNSResolve == false)
|
|
||||||
return "direct", "dns-route"
|
|
||||||
end
|
|
||||||
active = true
|
|
||||||
local ips, ttl, err = dns.Query("nested.example", true, false, false)
|
|
||||||
assert(not err and matcher:AnyMatch(ips) and active)
|
|
||||||
active = false
|
|
||||||
return "direct", "outer-route"
|
|
||||||
end`)
|
|
||||||
instance, err := core.New(&core.Config{
|
|
||||||
App: []*serial.TypedMessage{
|
|
||||||
serial.ToTypedMessage(&appdns.Config{
|
|
||||||
Tag: "dns", Script: dnsScript, DisableCache: true,
|
|
||||||
NameServer: []*appdns.NameServer{{
|
|
||||||
Id: "upstream", TimeoutMs: 1000,
|
|
||||||
Address: &net.Endpoint{
|
|
||||||
Network: net.Network_UDP,
|
|
||||||
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
|
|
||||||
Port: uint32(port),
|
|
||||||
},
|
|
||||||
}},
|
|
||||||
}),
|
|
||||||
serial.ToTypedMessage(&Config{Script: routerScript}),
|
|
||||||
serial.ToTypedMessage(&dispatcher.Config{}),
|
|
||||||
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
|
|
||||||
},
|
|
||||||
Outbound: []*core.OutboundHandlerConfig{
|
|
||||||
{Tag: "default", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})},
|
|
||||||
{Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{
|
|
||||||
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
|
|
||||||
})},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
defer instance.Close()
|
|
||||||
if err := instance.Start(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
r := instance.GetFeature(routing.RouterType()).(*Router)
|
|
||||||
route, err := r.PickRoute(newLuaRouteTestContext())
|
|
||||||
if err != nil || route.GetOutboundTag() != "direct" || route.GetRuleTag() != "outer-route" {
|
|
||||||
t.Fatalf("nested DNS routing = %v, %v", route, err)
|
|
||||||
}
|
|
||||||
if queries.Load() == 0 {
|
|
||||||
t.Fatal("DNS query did not pass through the dispatcher")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -3,7 +3,6 @@ package router
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"math"
|
"math"
|
||||||
"slices"
|
|
||||||
"sort"
|
"sort"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -78,7 +77,7 @@ func (s *LeastLoadStrategy) PickOutbound(candidates []string) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *LeastLoadStrategy) pickOutbounds(candidates []string) []*node {
|
func (s *LeastLoadStrategy) pickOutbounds(candidates []string) []*node {
|
||||||
qualified := s.getNodes(candidates)
|
qualified := s.getNodes(candidates, time.Duration(s.settings.MaxRTT))
|
||||||
selects := s.selectLeastLoad(qualified)
|
selects := s.selectLeastLoad(qualified)
|
||||||
return selects
|
return selects
|
||||||
}
|
}
|
||||||
@@ -139,7 +138,7 @@ func (s *LeastLoadStrategy) selectLeastLoad(nodes []*node) []*node {
|
|||||||
return nodes[:count]
|
return nodes[:count]
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *LeastLoadStrategy) getNodes(candidates []string) []*node {
|
func (s *LeastLoadStrategy) getNodes(candidates []string, maxRTT time.Duration) []*node {
|
||||||
if s.observer == nil {
|
if s.observer == nil {
|
||||||
errors.LogError(s.ctx, "observer is nil")
|
errors.LogError(s.ctx, "observer is nil")
|
||||||
return make([]*node, 0)
|
return make([]*node, 0)
|
||||||
@@ -152,10 +151,12 @@ func (s *LeastLoadStrategy) getNodes(candidates []string) []*node {
|
|||||||
|
|
||||||
results := observeResult.(*observatory.ObservationResult)
|
results := observeResult.(*observatory.ObservationResult)
|
||||||
|
|
||||||
|
outboundlist := outboundList(candidates)
|
||||||
|
|
||||||
var ret []*node
|
var ret []*node
|
||||||
|
|
||||||
for _, v := range results.Status {
|
for _, v := range results.Status {
|
||||||
if s.shouldSelectNode(v, candidates) {
|
if v.Alive && (v.Delay < maxRTT.Milliseconds() || maxRTT == 0) && outboundlist.contains(v.OutboundTag) {
|
||||||
record := &node{
|
record := &node{
|
||||||
Tag: v.OutboundTag,
|
Tag: v.OutboundTag,
|
||||||
CountAll: 1,
|
CountAll: 1,
|
||||||
@@ -171,8 +172,8 @@ func (s *LeastLoadStrategy) getNodes(candidates []string) []*node {
|
|||||||
record.RTTDeviationCost = time.Duration(s.costs.Apply(v.OutboundTag, float64(v.HealthPing.Deviation)))
|
record.RTTDeviationCost = time.Duration(s.costs.Apply(v.OutboundTag, float64(v.HealthPing.Deviation)))
|
||||||
record.CountAll = int(v.HealthPing.All)
|
record.CountAll = int(v.HealthPing.All)
|
||||||
record.CountFail = int(v.HealthPing.Fail)
|
record.CountFail = int(v.HealthPing.Fail)
|
||||||
}
|
|
||||||
|
|
||||||
|
}
|
||||||
ret = append(ret, record)
|
ret = append(ret, record)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -181,23 +182,6 @@ func (s *LeastLoadStrategy) getNodes(candidates []string) []*node {
|
|||||||
return ret
|
return ret
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *LeastLoadStrategy) shouldSelectNode(v *observatory.OutboundStatus, candidates []string) bool {
|
|
||||||
maxRTT := time.Duration(s.settings.MaxRTT)
|
|
||||||
if !v.Alive {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if maxRTT != 0 && v.Delay >= maxRTT.Milliseconds() {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !slices.Contains(candidates, v.OutboundTag) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if v.HealthPing != nil && v.HealthPing.All > 0 && s.settings.Tolerance > 0 && float64(v.HealthPing.Fail)/float64(v.HealthPing.All) > float64(s.settings.Tolerance) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func leastloadSort(nodes []*node) {
|
func leastloadSort(nodes []*node) {
|
||||||
sort.Slice(nodes, func(i, j int) bool {
|
sort.Slice(nodes, func(i, j int) bool {
|
||||||
left := nodes[i]
|
left := nodes[i]
|
||||||
|
|||||||
@@ -85,7 +85,6 @@ func TestSelectLeastExpected(t *testing.T) {
|
|||||||
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSelectLeastExpected2(t *testing.T) {
|
func TestSelectLeastExpected2(t *testing.T) {
|
||||||
strategy := &LeastLoadStrategy{
|
strategy := &LeastLoadStrategy{
|
||||||
settings: &StrategyLeastLoadConfig{
|
settings: &StrategyLeastLoadConfig{
|
||||||
@@ -103,7 +102,6 @@ func TestSelectLeastExpected2(t *testing.T) {
|
|||||||
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSelectLeastExpectedAndBaselines(t *testing.T) {
|
func TestSelectLeastExpectedAndBaselines(t *testing.T) {
|
||||||
strategy := &LeastLoadStrategy{
|
strategy := &LeastLoadStrategy{
|
||||||
settings: &StrategyLeastLoadConfig{
|
settings: &StrategyLeastLoadConfig{
|
||||||
@@ -124,7 +122,6 @@ func TestSelectLeastExpectedAndBaselines(t *testing.T) {
|
|||||||
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSelectLeastExpectedAndBaselines2(t *testing.T) {
|
func TestSelectLeastExpectedAndBaselines2(t *testing.T) {
|
||||||
strategy := &LeastLoadStrategy{
|
strategy := &LeastLoadStrategy{
|
||||||
settings: &StrategyLeastLoadConfig{
|
settings: &StrategyLeastLoadConfig{
|
||||||
@@ -145,7 +142,6 @@ func TestSelectLeastExpectedAndBaselines2(t *testing.T) {
|
|||||||
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSelectLeastLoadBaselines(t *testing.T) {
|
func TestSelectLeastLoadBaselines(t *testing.T) {
|
||||||
strategy := &LeastLoadStrategy{
|
strategy := &LeastLoadStrategy{
|
||||||
settings: &StrategyLeastLoadConfig{
|
settings: &StrategyLeastLoadConfig{
|
||||||
@@ -164,7 +160,6 @@ func TestSelectLeastLoadBaselines(t *testing.T) {
|
|||||||
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
t.Errorf("expected: %v, actual: %v", expected, len(ns))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSelectLeastLoadBaselinesNoQualified(t *testing.T) {
|
func TestSelectLeastLoadBaselinesNoQualified(t *testing.T) {
|
||||||
strategy := &LeastLoadStrategy{
|
strategy := &LeastLoadStrategy{
|
||||||
settings: &StrategyLeastLoadConfig{
|
settings: &StrategyLeastLoadConfig{
|
||||||
|
|||||||
+72
-20
@@ -7,16 +7,61 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/utils"
|
|
||||||
"github.com/xtls/xray-core/features/routing"
|
"github.com/xtls/xray-core/features/routing"
|
||||||
routing_session "github.com/xtls/xray-core/features/routing/session"
|
routing_session "github.com/xtls/xray-core/features/routing/session"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// parseURL splits a webhook URL into an HTTP URL and an optional Unix socket
|
||||||
|
// path. For regular http/https URLs the input is returned unchanged with an
|
||||||
|
// empty socketPath. For Unix sockets the format is:
|
||||||
|
//
|
||||||
|
// /path/to/socket.sock:/http/path
|
||||||
|
// @abstract:/http/path
|
||||||
|
// @@padded:/http/path
|
||||||
|
//
|
||||||
|
// The :/ separator after the socket path delimits the HTTP request path.
|
||||||
|
// If omitted, "/" is used.
|
||||||
|
func parseURL(raw string) (httpURL, socketPath string) {
|
||||||
|
if len(raw) == 0 || (!filepath.IsAbs(raw) && raw[0] != '@') {
|
||||||
|
return raw, ""
|
||||||
|
}
|
||||||
|
if idx := strings.Index(raw, ":/"); idx >= 0 {
|
||||||
|
return "http://localhost" + raw[idx+1:], raw[:idx]
|
||||||
|
}
|
||||||
|
return "http://localhost/", raw
|
||||||
|
}
|
||||||
|
|
||||||
|
// resolveSocketPath applies platform-specific transformations to a Unix
|
||||||
|
// socket path, matching the behaviour of the listen side in
|
||||||
|
// transport/internet/system_listener.go.
|
||||||
|
//
|
||||||
|
// For abstract sockets (prefix @) on Linux/Android:
|
||||||
|
// - single @ — used as-is (lock-free abstract socket)
|
||||||
|
// - double @@ — stripped to single @ and padded to
|
||||||
|
// syscall.RawSockaddrUnix{}.Path length (HAProxy compat)
|
||||||
|
func resolveSocketPath(path string) string {
|
||||||
|
if len(path) == 0 || path[0] != '@' {
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
if runtime.GOOS != "linux" && runtime.GOOS != "android" {
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
if len(path) > 1 && path[1] == '@' {
|
||||||
|
fullAddr := make([]byte, len(syscall.RawSockaddrUnix{}.Path))
|
||||||
|
copy(fullAddr, path[1:])
|
||||||
|
return string(fullAddr)
|
||||||
|
}
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
|
||||||
func ptr[T any](v T) *T { return &v }
|
func ptr[T any](v T) *T { return &v }
|
||||||
|
|
||||||
type event struct {
|
type event struct {
|
||||||
@@ -41,7 +86,6 @@ type WebhookNotifier struct {
|
|||||||
deduplication uint32
|
deduplication uint32
|
||||||
client *http.Client
|
client *http.Client
|
||||||
seen sync.Map
|
seen sync.Map
|
||||||
lastSweep atomic.Int64
|
|
||||||
done chan struct{}
|
done chan struct{}
|
||||||
wg sync.WaitGroup
|
wg sync.WaitGroup
|
||||||
closeOnce sync.Once
|
closeOnce sync.Once
|
||||||
@@ -52,7 +96,7 @@ func NewWebhookNotifier(cfg *WebhookConfig) (*WebhookNotifier, error) {
|
|||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
httpURL, socketPath := utils.SplitHTTPUnixURL(cfg.Url)
|
httpURL, socketPath := parseURL(cfg.Url)
|
||||||
h := &WebhookNotifier{
|
h := &WebhookNotifier{
|
||||||
url: httpURL,
|
url: httpURL,
|
||||||
deduplication: cfg.Deduplication,
|
deduplication: cfg.Deduplication,
|
||||||
@@ -63,7 +107,7 @@ func NewWebhookNotifier(cfg *WebhookConfig) (*WebhookNotifier, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if socketPath != "" {
|
if socketPath != "" {
|
||||||
dialAddr := utils.ResolveSocketPath(socketPath)
|
dialAddr := resolveSocketPath(socketPath)
|
||||||
h.client.Transport = &http.Transport{
|
h.client.Transport = &http.Transport{
|
||||||
DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) {
|
DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) {
|
||||||
var d net.Dialer
|
var d net.Dialer
|
||||||
@@ -79,6 +123,11 @@ func NewWebhookNotifier(cfg *WebhookConfig) (*WebhookNotifier, error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if h.deduplication > 0 {
|
||||||
|
h.wg.Add(1)
|
||||||
|
go h.cleanupLoop()
|
||||||
|
}
|
||||||
|
|
||||||
return h, nil
|
return h, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -198,7 +247,6 @@ func (h *WebhookNotifier) isDuplicate(email string) bool {
|
|||||||
}
|
}
|
||||||
ttl := time.Duration(h.deduplication) * time.Second
|
ttl := time.Duration(h.deduplication) * time.Second
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
h.maybeSweep(now, ttl)
|
|
||||||
if v, loaded := h.seen.LoadOrStore(email, now); loaded {
|
if v, loaded := h.seen.LoadOrStore(email, now); loaded {
|
||||||
if now.Sub(v.(time.Time)) < ttl {
|
if now.Sub(v.(time.Time)) < ttl {
|
||||||
return true
|
return true
|
||||||
@@ -208,23 +256,27 @@ func (h *WebhookNotifier) isDuplicate(email string) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *WebhookNotifier) maybeSweep(now time.Time, ttl time.Duration) {
|
func (h *WebhookNotifier) cleanupLoop() {
|
||||||
last := h.lastSweep.Load()
|
defer h.wg.Done()
|
||||||
if now.UnixNano()-last < int64(ttl) {
|
ttl := time.Duration(h.deduplication) * time.Second
|
||||||
return
|
ticker := time.NewTicker(ttl)
|
||||||
}
|
defer ticker.Stop()
|
||||||
if !h.lastSweep.CompareAndSwap(last, now.UnixNano()) {
|
for {
|
||||||
return // another goroutine did the sweep
|
select {
|
||||||
}
|
case <-h.done:
|
||||||
h.seen.Range(func(key, value any) bool {
|
return
|
||||||
if now.Sub(value.(time.Time)) >= ttl {
|
case <-ticker.C:
|
||||||
h.seen.Delete(key)
|
now := time.Now()
|
||||||
|
h.seen.Range(func(key, value any) bool {
|
||||||
|
if now.Sub(value.(time.Time)) >= ttl {
|
||||||
|
h.seen.Delete(key)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
}
|
}
|
||||||
return true
|
}
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Only need to call if the Notifier is really used, otherwise GC can clean it
|
|
||||||
func (h *WebhookNotifier) Close() error {
|
func (h *WebhookNotifier) Close() error {
|
||||||
h.closeOnce.Do(func() {
|
h.closeOnce.Do(func() {
|
||||||
close(h.done)
|
close(h.done)
|
||||||
|
|||||||
@@ -48,20 +48,6 @@ func (m *Manager) RegisterCounter(name string) (stats.Counter, error) {
|
|||||||
return c, nil
|
return c, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetOrRegisterCounter implements stats.Manager.
|
|
||||||
func (m *Manager) GetOrRegisterCounter(name string) (stats.Counter, error) {
|
|
||||||
m.access.Lock()
|
|
||||||
defer m.access.Unlock()
|
|
||||||
|
|
||||||
if c, found := m.counters[name]; found {
|
|
||||||
return c, nil
|
|
||||||
}
|
|
||||||
errors.LogDebug(context.Background(), "create new counter ", name)
|
|
||||||
c := new(Counter)
|
|
||||||
m.counters[name] = c
|
|
||||||
return c, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// UnregisterCounter implements stats.Manager.
|
// UnregisterCounter implements stats.Manager.
|
||||||
func (m *Manager) UnregisterCounter(name string) error {
|
func (m *Manager) UnregisterCounter(name string) error {
|
||||||
m.access.Lock()
|
m.access.Lock()
|
||||||
@@ -111,20 +97,6 @@ func (m *Manager) RegisterOnlineMap(name string) (stats.OnlineMap, error) {
|
|||||||
return om, nil
|
return om, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetOrRegisterOnlineMap implements stats.Manager.
|
|
||||||
func (m *Manager) GetOrRegisterOnlineMap(name string) (stats.OnlineMap, error) {
|
|
||||||
m.access.Lock()
|
|
||||||
defer m.access.Unlock()
|
|
||||||
|
|
||||||
if om, found := m.onlineMaps[name]; found {
|
|
||||||
return om, nil
|
|
||||||
}
|
|
||||||
errors.LogDebug(context.Background(), "create new OnlineMap ", name)
|
|
||||||
om := NewOnlineMap()
|
|
||||||
m.onlineMaps[name] = om
|
|
||||||
return om, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// UnregisterOnlineMap implements stats.Manager.
|
// UnregisterOnlineMap implements stats.Manager.
|
||||||
func (m *Manager) UnregisterOnlineMap(name string) error {
|
func (m *Manager) UnregisterOnlineMap(name string) error {
|
||||||
m.access.Lock()
|
m.access.Lock()
|
||||||
@@ -177,26 +149,6 @@ func (m *Manager) RegisterChannel(name string) (stats.Channel, error) {
|
|||||||
return c, nil
|
return c, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetOrRegisterChannel implements stats.Manager.
|
|
||||||
func (m *Manager) GetOrRegisterChannel(name string) (stats.Channel, error) {
|
|
||||||
m.access.Lock()
|
|
||||||
defer m.access.Unlock()
|
|
||||||
|
|
||||||
if c, found := m.channels[name]; found {
|
|
||||||
return c, nil
|
|
||||||
}
|
|
||||||
errors.LogDebug(context.Background(), "create new channel ", name)
|
|
||||||
c := NewChannel(&ChannelConfig{BufferSize: 64, Blocking: false})
|
|
||||||
if m.running {
|
|
||||||
// Start before publishing so no goroutine can observe an unstarted channel.
|
|
||||||
if err := c.Start(); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
m.channels[name] = c
|
|
||||||
return c, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// UnregisterChannel implements stats.Manager.
|
// UnregisterChannel implements stats.Manager.
|
||||||
func (m *Manager) UnregisterChannel(name string) error {
|
func (m *Manager) UnregisterChannel(name string) error {
|
||||||
m.access.Lock()
|
m.access.Lock()
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestInterface(t *testing.T) {
|
func TestInterface(t *testing.T) {
|
||||||
_ = stats.Manager(new(Manager))
|
_ = (stats.Manager)(new(Manager))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestStatsChannelRunnable(t *testing.T) {
|
func TestStatsChannelRunnable(t *testing.T) {
|
||||||
|
|||||||
@@ -2,11 +2,10 @@ package version
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Version struct {
|
type Version struct {
|
||||||
|
|||||||
+1
-1
@@ -122,7 +122,7 @@ func NewReader(reader io.Reader) Reader {
|
|||||||
}
|
}
|
||||||
|
|
||||||
_, isFile := reader.(*os.File)
|
_, isFile := reader.(*os.File)
|
||||||
if !isFile && useReadV() {
|
if !isFile && useReadv {
|
||||||
if sc, ok := reader.(syscall.Conn); ok {
|
if sc, ok := reader.(syscall.Conn); ok {
|
||||||
rawConn, err := sc.SyscallConn()
|
rawConn, err := sc.SyscallConn()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -38,8 +38,8 @@ func MergeMulti(dest MultiBuffer, src MultiBuffer) (MultiBuffer, MultiBuffer) {
|
|||||||
// MergeBytes merges the given bytes into MultiBuffer and return the new address of the merged MultiBuffer.
|
// MergeBytes merges the given bytes into MultiBuffer and return the new address of the merged MultiBuffer.
|
||||||
func MergeBytes(dest MultiBuffer, src []byte) MultiBuffer {
|
func MergeBytes(dest MultiBuffer, src []byte) MultiBuffer {
|
||||||
n := len(dest)
|
n := len(dest)
|
||||||
if n > 0 && !dest[n-1].IsFull() {
|
if n > 0 && !(dest)[n-1].IsFull() {
|
||||||
nBytes, _ := dest[n-1].Write(src)
|
nBytes, _ := (dest)[n-1].Write(src)
|
||||||
src = src[nBytes:]
|
src = src[nBytes:]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -121,11 +121,11 @@ func TestPacketReader_ReadMultiBuffer(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestReaderInterface(t *testing.T) {
|
func TestReaderInterface(t *testing.T) {
|
||||||
_ = io.Reader(new(ReadVReader))
|
_ = (io.Reader)(new(ReadVReader))
|
||||||
_ = Reader(new(ReadVReader))
|
_ = (Reader)(new(ReadVReader))
|
||||||
|
|
||||||
_ = Reader(new(BufferedReader))
|
_ = (Reader)(new(BufferedReader))
|
||||||
_ = io.Reader(new(BufferedReader))
|
_ = (io.Reader)(new(BufferedReader))
|
||||||
_ = io.ByteReader(new(BufferedReader))
|
_ = (io.ByteReader)(new(BufferedReader))
|
||||||
_ = io.WriterTo(new(BufferedReader))
|
_ = (io.WriterTo)(new(BufferedReader))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ func (r *posixReader) Init(bs []*Buffer) {
|
|||||||
}
|
}
|
||||||
for idx, b := range bs {
|
for idx, b := range bs {
|
||||||
iovecs = append(iovecs, syscall.Iovec{
|
iovecs = append(iovecs, syscall.Iovec{
|
||||||
Base: &b.v[0],
|
Base: &(b.v[0]),
|
||||||
})
|
})
|
||||||
iovecs[idx].SetLen(int(Size))
|
iovecs[idx].SetLen(int(Size))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ package buf
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"io"
|
"io"
|
||||||
"sync/atomic"
|
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/platform"
|
"github.com/xtls/xray-core/common/platform"
|
||||||
@@ -144,24 +143,13 @@ func (r *ReadVReader) ReadMultiBuffer() (MultiBuffer, error) {
|
|||||||
return mb, nil
|
return mb, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var useReadv atomic.Bool
|
var useReadv bool
|
||||||
|
|
||||||
func useReadV() bool {
|
|
||||||
return useReadv.Load()
|
|
||||||
}
|
|
||||||
|
|
||||||
func reloadEnvSettings() error {
|
|
||||||
const defaultFlagValue = "NOT_DEFINED_AT_ALL"
|
|
||||||
value := platform.NewEnvFlag(platform.UseReadV).GetValue(func() string { return defaultFlagValue })
|
|
||||||
enabled := false
|
|
||||||
switch value {
|
|
||||||
case defaultFlagValue, "auto", "enable":
|
|
||||||
enabled = true
|
|
||||||
}
|
|
||||||
useReadv.Store(enabled)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
platform.RegisterEnvReload(reloadEnvSettings)
|
const defaultFlagValue = "NOT_DEFINED_AT_ALL"
|
||||||
|
value := platform.NewEnvFlag(platform.UseReadV).GetValue(func() string { return defaultFlagValue })
|
||||||
|
switch value {
|
||||||
|
case defaultFlagValue, "auto", "enable":
|
||||||
|
useReadv = true
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,9 +10,7 @@ import (
|
|||||||
"github.com/xtls/xray-core/features/stats"
|
"github.com/xtls/xray-core/features/stats"
|
||||||
)
|
)
|
||||||
|
|
||||||
func useReadV() bool {
|
const useReadv = false
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewReadVReader(reader io.Reader, rawConn syscall.RawConn, counter stats.Counter) Reader {
|
func NewReadVReader(reader io.Reader, rawConn syscall.RawConn, counter stats.Counter) Reader {
|
||||||
panic("not implemented")
|
panic("not implemented")
|
||||||
|
|||||||
@@ -5,8 +5,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type windowsReader struct {
|
type windowsReader struct {
|
||||||
bufs []syscall.WSABuf
|
bufs []syscall.WSABuf
|
||||||
ready bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Init(bs []*Buffer) {
|
func (r *windowsReader) Init(bs []*Buffer) {
|
||||||
@@ -16,7 +15,6 @@ func (r *windowsReader) Init(bs []*Buffer) {
|
|||||||
for _, b := range bs {
|
for _, b := range bs {
|
||||||
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
|
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
|
||||||
}
|
}
|
||||||
r.ready = false
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Clear() {
|
func (r *windowsReader) Clear() {
|
||||||
@@ -27,14 +25,6 @@ func (r *windowsReader) Clear() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Read(fd uintptr) int32 {
|
func (r *windowsReader) Read(fd uintptr) int32 {
|
||||||
// On the first invocation, we return -1 to indicate "not ready"
|
|
||||||
// to make rawConn.Read wait for readability using the runtime's own mechanism
|
|
||||||
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
|
|
||||||
if !r.ready {
|
|
||||||
r.ready = true
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
var nBytes uint32
|
var nBytes uint32
|
||||||
var flags uint32
|
var flags uint32
|
||||||
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
|
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
|
||||||
|
|||||||
@@ -118,9 +118,7 @@ func (w *BufferedWriter) Write(b []byte) (int, error) {
|
|||||||
|
|
||||||
nBytes, err := w.buffer.Write(b)
|
nBytes, err := w.buffer.Write(b)
|
||||||
totalBytes += nBytes
|
totalBytes += nBytes
|
||||||
|
if err != nil {
|
||||||
// ErrBufferFull means a partial write, so flush below and continue
|
|
||||||
if err != nil && err != ErrBufferFull {
|
|
||||||
return totalBytes, err
|
return totalBytes, err
|
||||||
}
|
}
|
||||||
if !w.buffered || w.buffer.IsFull() {
|
if !w.buffered || w.buffer.IsFull() {
|
||||||
|
|||||||
@@ -10,12 +10,12 @@ import (
|
|||||||
|
|
||||||
// [,)
|
// [,)
|
||||||
func RandBetween(from int64, to int64) int64 {
|
func RandBetween(from int64, to int64) int64 {
|
||||||
|
if from == to {
|
||||||
|
return from
|
||||||
|
}
|
||||||
if from > to {
|
if from > to {
|
||||||
from, to = to, from
|
from, to = to, from
|
||||||
}
|
}
|
||||||
if d := to - from; d == 0 || d == 1 {
|
|
||||||
return from
|
|
||||||
}
|
|
||||||
bigInt, _ := rand.Int(rand.Reader, big.NewInt(to-from))
|
bigInt, _ := rand.Int(rand.Reader, big.NewInt(to-from))
|
||||||
return from + bigInt.Int64()
|
return from + bigInt.Int64()
|
||||||
}
|
}
|
||||||
|
|||||||
+65
-13
@@ -18,12 +18,17 @@ type hasInnerError interface {
|
|||||||
Unwrap() error
|
Unwrap() error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type hasSeverity interface {
|
||||||
|
Severity() log.Severity
|
||||||
|
}
|
||||||
|
|
||||||
// Error is an error object with underlying error.
|
// Error is an error object with underlying error.
|
||||||
type Error struct {
|
type Error struct {
|
||||||
prefix []interface{}
|
prefix []interface{}
|
||||||
message []interface{}
|
message []interface{}
|
||||||
caller string
|
caller string
|
||||||
inner error
|
inner error
|
||||||
|
severity log.Severity
|
||||||
}
|
}
|
||||||
|
|
||||||
// Error implements error.Error().
|
// Error implements error.Error().
|
||||||
@@ -64,6 +69,46 @@ func (err *Error) Base(e error) *Error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (err *Error) atSeverity(s log.Severity) *Error {
|
||||||
|
err.severity = s
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (err *Error) Severity() log.Severity {
|
||||||
|
if err.inner == nil {
|
||||||
|
return err.severity
|
||||||
|
}
|
||||||
|
|
||||||
|
if s, ok := err.inner.(hasSeverity); ok {
|
||||||
|
as := s.Severity()
|
||||||
|
if as < err.severity {
|
||||||
|
return as
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return err.severity
|
||||||
|
}
|
||||||
|
|
||||||
|
// AtDebug sets the severity to debug.
|
||||||
|
func (err *Error) AtDebug() *Error {
|
||||||
|
return err.atSeverity(log.Severity_Debug)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AtInfo sets the severity to info.
|
||||||
|
func (err *Error) AtInfo() *Error {
|
||||||
|
return err.atSeverity(log.Severity_Info)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AtWarning sets the severity to warning.
|
||||||
|
func (err *Error) AtWarning() *Error {
|
||||||
|
return err.atSeverity(log.Severity_Warning)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AtError sets the severity to error.
|
||||||
|
func (err *Error) AtError() *Error {
|
||||||
|
return err.atSeverity(log.Severity_Error)
|
||||||
|
}
|
||||||
|
|
||||||
// String returns the string representation of this error.
|
// String returns the string representation of this error.
|
||||||
func (err *Error) String() string {
|
func (err *Error) String() string {
|
||||||
return err.Error()
|
return err.Error()
|
||||||
@@ -87,8 +132,9 @@ func New(msg ...interface{}) *Error {
|
|||||||
details = details[:i]
|
details = details[:i]
|
||||||
}
|
}
|
||||||
return &Error{
|
return &Error{
|
||||||
message: msg,
|
message: msg,
|
||||||
caller: details,
|
severity: log.Severity_Info,
|
||||||
|
caller: details,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -125,9 +171,6 @@ func LogErrorInner(ctx context.Context, inner error, msg ...interface{}) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) {
|
func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) {
|
||||||
if log.GetSeverity() < severity {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
pc, _, _, _ := runtime.Caller(2)
|
pc, _, _, _ := runtime.Caller(2)
|
||||||
details := runtime.FuncForPC(pc).Name()
|
details := runtime.FuncForPC(pc).Name()
|
||||||
if len(details) >= trim {
|
if len(details) >= trim {
|
||||||
@@ -138,9 +181,10 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
|||||||
details = details[:i]
|
details = details[:i]
|
||||||
}
|
}
|
||||||
err := &Error{
|
err := &Error{
|
||||||
message: msg,
|
message: msg,
|
||||||
caller: details,
|
severity: severity,
|
||||||
inner: inner,
|
caller: details,
|
||||||
|
inner: inner,
|
||||||
}
|
}
|
||||||
if ctx != nil && ctx != context.Background() {
|
if ctx != nil && ctx != context.Background() {
|
||||||
id := uint32(c.IDFromContext(ctx))
|
id := uint32(c.IDFromContext(ctx))
|
||||||
@@ -149,7 +193,7 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
log.Record(&log.GeneralMessage{
|
log.Record(&log.GeneralMessage{
|
||||||
Severity: severity,
|
Severity: GetSeverity(err),
|
||||||
Content: err,
|
Content: err,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -173,3 +217,11 @@ L:
|
|||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetSeverity returns the actual severity of the error, including inner errors.
|
||||||
|
func GetSeverity(err error) log.Severity {
|
||||||
|
if s, ok := err.(hasSeverity); ok {
|
||||||
|
return s.Severity()
|
||||||
|
}
|
||||||
|
return log.Severity_Info
|
||||||
|
}
|
||||||
|
|||||||
@@ -7,21 +7,30 @@ import (
|
|||||||
|
|
||||||
"github.com/google/go-cmp/cmp"
|
"github.com/google/go-cmp/cmp"
|
||||||
. "github.com/xtls/xray-core/common/errors"
|
. "github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/log"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestError(t *testing.T) {
|
func TestError(t *testing.T) {
|
||||||
err := New("TestError")
|
err := New("TestError")
|
||||||
if v := err.Error(); !strings.Contains(v, "TestError") {
|
if v := GetSeverity(err); v != log.Severity_Info {
|
||||||
t.Error("error: ", v)
|
t.Error("severity: ", v)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = New("TestError2").Base(io.EOF)
|
err = New("TestError2").Base(io.EOF)
|
||||||
if v := err.Error(); !strings.Contains(v, "EOF") {
|
if v := GetSeverity(err); v != log.Severity_Info {
|
||||||
t.Error("error: ", v)
|
t.Error("severity: ", v)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = New("TestError3").Base(io.EOF)
|
err = New("TestError3").Base(io.EOF).AtWarning()
|
||||||
err = New("TestError4").Base(err)
|
if v := GetSeverity(err); v != log.Severity_Warning {
|
||||||
|
t.Error("severity: ", v)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = New("TestError4").Base(io.EOF).AtWarning()
|
||||||
|
err = New("TestError5").Base(err)
|
||||||
|
if v := GetSeverity(err); v != log.Severity_Warning {
|
||||||
|
t.Error("severity: ", v)
|
||||||
|
}
|
||||||
if v := err.Error(); !strings.Contains(v, "EOF") {
|
if v := err.Error(); !strings.Contains(v, "EOF") {
|
||||||
t.Error("error: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,49 +0,0 @@
|
|||||||
package geodata
|
|
||||||
|
|
||||||
import (
|
|
||||||
"sync"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
)
|
|
||||||
|
|
||||||
var privateIPMatcher = sync.OnceValue(func() IPMatcher {
|
|
||||||
return common.Must2(IPReg.BuildIPMatcher(common.Must2(ParseIPRules([]string{
|
|
||||||
"0.0.0.0/8",
|
|
||||||
"10.0.0.0/8",
|
|
||||||
"100.64.0.0/10",
|
|
||||||
"127.0.0.0/8",
|
|
||||||
"169.254.0.0/16",
|
|
||||||
"172.16.0.0/12",
|
|
||||||
"192.0.0.0/24",
|
|
||||||
"192.0.2.0/24",
|
|
||||||
"192.88.99.0/24",
|
|
||||||
"192.168.0.0/16",
|
|
||||||
"198.18.0.0/15",
|
|
||||||
"198.51.100.0/24",
|
|
||||||
"203.0.113.0/24",
|
|
||||||
"224.0.0.0/3",
|
|
||||||
"::/127",
|
|
||||||
"fc00::/7",
|
|
||||||
"fe80::/10",
|
|
||||||
"ff00::/8",
|
|
||||||
}))))
|
|
||||||
})
|
|
||||||
|
|
||||||
func GetPrivateIPMatcher() IPMatcher { return privateIPMatcher() }
|
|
||||||
|
|
||||||
var privateDomainMatcher = sync.OnceValue(func() DomainMatcher {
|
|
||||||
return common.Must2(DomainReg.BuildDomainMatcher(common.Must2(ParseDomainRules([]string{
|
|
||||||
"lan",
|
|
||||||
"localdomain",
|
|
||||||
"example",
|
|
||||||
"invalid",
|
|
||||||
"localhost",
|
|
||||||
"test",
|
|
||||||
"local",
|
|
||||||
"home.arpa",
|
|
||||||
"internal",
|
|
||||||
"regexp:^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$", // Dotless domains
|
|
||||||
}, Domain_Domain))))
|
|
||||||
})
|
|
||||||
|
|
||||||
func GetPrivateDomainMatcher() DomainMatcher { return privateDomainMatcher() }
|
|
||||||
@@ -8,7 +8,6 @@ import (
|
|||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
||||||
"github.com/xtls/xray-core/common/utils"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type DomainMatcher interface {
|
type DomainMatcher interface {
|
||||||
@@ -26,7 +25,7 @@ type DomainMatcherFactory interface {
|
|||||||
|
|
||||||
type MphDomainMatcherFactory struct {
|
type MphDomainMatcherFactory struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
|
shared map[string]strmatcher.MatcherGroup // TODO: cleanup
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildDomainRulesKey(rules []*DomainRule) string {
|
func buildDomainRulesKey(rules []*DomainRule) string {
|
||||||
@@ -66,7 +65,7 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
if key != "" {
|
if key != "" {
|
||||||
f.Lock()
|
f.Lock()
|
||||||
defer f.Unlock()
|
defer f.Unlock()
|
||||||
if g, ok := f.shared.Load(key); ok {
|
if g := f.shared[key]; g != nil {
|
||||||
errors.LogDebug(context.Background(), "geodata mph domain matcher cache HIT for ", len(rules), " rules")
|
errors.LogDebug(context.Background(), "geodata mph domain matcher cache HIT for ", len(rules), " rules")
|
||||||
return g, nil
|
return g, nil
|
||||||
}
|
}
|
||||||
@@ -82,10 +81,19 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
}
|
}
|
||||||
g.Add(m, uint32(i))
|
g.Add(m, uint32(i))
|
||||||
case *DomainRule_Geosite:
|
case *DomainRule_Geosite:
|
||||||
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
|
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
for j, d := range domains {
|
||||||
|
domains[j] = nil // peak mem
|
||||||
|
m, err := parseDomain(d)
|
||||||
|
if err != nil {
|
||||||
|
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
g.Add(m, uint32(i))
|
||||||
|
}
|
||||||
default:
|
default:
|
||||||
panic("unknown domain rule type")
|
panic("unknown domain rule type")
|
||||||
}
|
}
|
||||||
@@ -94,45 +102,55 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if key != "" {
|
if key != "" {
|
||||||
f.shared.Store(key, g)
|
f.shared[key] = g
|
||||||
}
|
}
|
||||||
return g, nil
|
return g, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type CompactMphDomainMatcherFactory struct {
|
type CompactDomainMatcherFactory struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
|
shared map[string]strmatcher.MatcherSet // TODO: cleanup
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
|
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
|
||||||
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
||||||
|
|
||||||
f.Lock()
|
f.Lock()
|
||||||
defer f.Unlock()
|
defer f.Unlock()
|
||||||
|
|
||||||
if s, ok := f.shared.Load(key); ok {
|
if s := f.shared[key]; s != nil {
|
||||||
errors.LogDebug(context.Background(), "geodata geosite matcher cache HIT ", key)
|
errors.LogDebug(context.Background(), "geodata geosite matcher cache HIT ", key)
|
||||||
return s, nil
|
return s, nil
|
||||||
}
|
}
|
||||||
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
||||||
|
|
||||||
s := strmatcher.NewMphValueMatcher()
|
s := strmatcher.NewLinearAnyMatcher()
|
||||||
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
|
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
|
||||||
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := s.Build(); err != nil {
|
for i, d := range domains {
|
||||||
return nil, err
|
domains[i] = nil // peak mem
|
||||||
|
m, err := parseDomain(d)
|
||||||
|
if err != nil {
|
||||||
|
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
s.Add(m)
|
||||||
}
|
}
|
||||||
f.shared.Store(key, s)
|
f.shared[key] = s
|
||||||
return s, nil
|
return s, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// BuildMatcher implements DomainMatcherFactory.
|
// BuildMatcher implements DomainMatcherFactory.
|
||||||
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||||
if len(rules) == 0 {
|
if len(rules) == 0 {
|
||||||
return nil, errors.New("empty domain rule list")
|
return nil, errors.New("empty domain rule list")
|
||||||
}
|
}
|
||||||
compact := new(CompactMphDomainMatcher)
|
compact := &CompactDomainMatcher{
|
||||||
|
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
|
||||||
|
values: make([]uint32, 0, len(rules)),
|
||||||
|
}
|
||||||
for i, r := range rules {
|
for i, r := range rules {
|
||||||
switch v := r.Value.(type) {
|
switch v := r.Value.(type) {
|
||||||
case *DomainRule_Custom:
|
case *DomainRule_Custom:
|
||||||
@@ -149,7 +167,8 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
compact.combiner.Add(m, uint32(i))
|
compact.matchers = append(compact.matchers, m)
|
||||||
|
compact.values = append(compact.values, uint32(i))
|
||||||
default:
|
default:
|
||||||
panic("unknown domain rule type")
|
panic("unknown domain rule type")
|
||||||
}
|
}
|
||||||
@@ -157,40 +176,37 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
|
|||||||
return compact, nil
|
return compact, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type CompactMphDomainMatcher struct {
|
type CompactDomainMatcher struct {
|
||||||
custom strmatcher.ValueMatcher
|
custom strmatcher.ValueMatcher
|
||||||
combiner strmatcher.MphValueMatcherCombiner
|
matchers []strmatcher.MatcherSet
|
||||||
|
values []uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
// Match implements DomainMatcher.
|
// Match implements DomainMatcher.
|
||||||
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
|
func (c *CompactDomainMatcher) Match(input string) []uint32 {
|
||||||
result := c.combiner.Match(input)
|
var result []uint32
|
||||||
if c.custom != nil {
|
if c.custom != nil {
|
||||||
result = append(c.custom.Match(input), result...)
|
result = append(result, c.custom.Match(input)...)
|
||||||
|
}
|
||||||
|
for i, m := range c.matchers {
|
||||||
|
if m.MatchAny(input) {
|
||||||
|
result = append(result, c.values[i])
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements DomainMatcher.
|
// MatchAny implements DomainMatcher.
|
||||||
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
|
func (c *CompactDomainMatcher) MatchAny(input string) bool {
|
||||||
if c.custom != nil && c.custom.MatchAny(input) {
|
if c.custom != nil && c.custom.MatchAny(input) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
return c.combiner.MatchAny(input)
|
for _, m := range c.matchers {
|
||||||
}
|
if m.MatchAny(input) {
|
||||||
|
return true
|
||||||
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
|
|
||||||
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
|
|
||||||
i := 0
|
|
||||||
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
|
|
||||||
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
|
|
||||||
if err != nil {
|
|
||||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
|
||||||
} else {
|
|
||||||
add(m)
|
|
||||||
}
|
}
|
||||||
i++
|
}
|
||||||
})
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||||
@@ -203,7 +219,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
|||||||
case Domain_Regex:
|
case Domain_Regex:
|
||||||
return strmatcher.Regex.New(d.Value)
|
return strmatcher.Regex.New(d.Value)
|
||||||
case Domain_Domain:
|
case Domain_Domain:
|
||||||
return strmatcher.Domain.New(strings.ToLower(d.Value))
|
return strmatcher.Domain.New(d.Value)
|
||||||
case Domain_Full:
|
case Domain_Full:
|
||||||
return strmatcher.Full.New(strings.ToLower(d.Value))
|
return strmatcher.Full.New(strings.ToLower(d.Value))
|
||||||
default:
|
default:
|
||||||
@@ -214,8 +230,8 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
|||||||
func newDomainMatcherFactory() DomainMatcherFactory {
|
func newDomainMatcherFactory() DomainMatcherFactory {
|
||||||
switch runtime.GOOS {
|
switch runtime.GOOS {
|
||||||
case "ios", "android":
|
case "ios", "android":
|
||||||
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
return &CompactDomainMatcherFactory{shared: make(map[string]strmatcher.MatcherSet)}
|
||||||
default:
|
default:
|
||||||
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
return &MphDomainMatcherFactory{shared: make(map[string]strmatcher.MatcherGroup)}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,15 +4,13 @@ import (
|
|||||||
"path/filepath"
|
"path/filepath"
|
||||||
"reflect"
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
"sync"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
||||||
"github.com/xtls/xray-core/common/utils"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||||
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
factory := &CompactDomainMatcherFactory{shared: make(map[string]strmatcher.MatcherSet)}
|
||||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||||
@@ -33,7 +31,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
|||||||
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
||||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||||
|
|
||||||
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
factory := &CompactDomainMatcherFactory{shared: make(map[string]strmatcher.MatcherSet)}
|
||||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
||||||
@@ -52,11 +50,10 @@ func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
|
func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
|
||||||
matcher, err := (&MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}).
|
matcher, err := (&MphDomainMatcherFactory{shared: make(map[string]strmatcher.MatcherGroup)}).BuildMatcher([]*DomainRule{
|
||||||
BuildMatcher([]*DomainRule{
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
})
|
||||||
})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("BuildMatcher() failed: %v", err)
|
t.Fatalf("BuildMatcher() failed: %v", err)
|
||||||
}
|
}
|
||||||
@@ -73,76 +70,3 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
|
|||||||
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// DNS sorts every Match result in place, so a matcher must never hand out a
|
|
||||||
// slice it keeps, also when only its keyword or regex part matches.
|
|
||||||
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
|
|
||||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
|
||||||
|
|
||||||
rules := []*DomainRule{
|
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
|
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
|
||||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
|
|
||||||
}
|
|
||||||
cases := []struct {
|
|
||||||
input string
|
|
||||||
want []uint32
|
|
||||||
}{
|
|
||||||
{"example.com", []uint32{0, 1, 2, 4}},
|
|
||||||
{"www.example.com", []uint32{1, 2, 4}},
|
|
||||||
{"exam.net", []uint32{2, 4}}, // keyword part only
|
|
||||||
{"example.org", []uint32{2, 3, 4}},
|
|
||||||
{"163.com", []uint32{5}},
|
|
||||||
{"www.163.com", []uint32{5}},
|
|
||||||
{"only.full.test", []uint32{6}}, // full part only
|
|
||||||
{"nomatch.test", nil},
|
|
||||||
}
|
|
||||||
factories := map[string]DomainMatcherFactory{
|
|
||||||
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
|
||||||
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
|
||||||
}
|
|
||||||
for name, factory := range factories {
|
|
||||||
t.Run(name, func(t *testing.T) {
|
|
||||||
matcher, err := factory.BuildMatcher(rules)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("BuildMatcher() failed: %v", err)
|
|
||||||
}
|
|
||||||
for _, c := range cases {
|
|
||||||
got := matcher.Match(c.input)
|
|
||||||
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
|
|
||||||
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
|
|
||||||
}
|
|
||||||
got = got[:cap(got)]
|
|
||||||
for j := range got {
|
|
||||||
got[j] = ^uint32(0)
|
|
||||||
}
|
|
||||||
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
|
|
||||||
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
for range 8 {
|
|
||||||
wg.Add(1)
|
|
||||||
go func() {
|
|
||||||
defer wg.Done()
|
|
||||||
for range 500 {
|
|
||||||
for _, c := range cases {
|
|
||||||
got := matcher.Match(c.input)
|
|
||||||
slices.Sort(got)
|
|
||||||
if !slices.Equal(got, c.want) {
|
|
||||||
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
}
|
|
||||||
wg.Wait()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -6,14 +6,12 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/utils"
|
|
||||||
"github.com/xtls/xray-core/common/uuid"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type DomainRegistry struct {
|
type DomainRegistry struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
factory DomainMatcherFactory
|
factory DomainMatcherFactory
|
||||||
matchers *utils.WeakCacheMap[uuid.UUID, DynamicDomainMatcher]
|
matchers []*DynamicDomainMatcher
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *DomainRegistry) BuildDomainMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
func (r *DomainRegistry) BuildDomainMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||||
@@ -26,7 +24,7 @@ func (r *DomainRegistry) BuildDomainMatcher(rules []*DomainRule) (DomainMatcher,
|
|||||||
}
|
}
|
||||||
|
|
||||||
d := NewDynamicDomainMatcher(rules, m)
|
d := NewDynamicDomainMatcher(rules, m)
|
||||||
r.matchers.Store(uuid.New(), d)
|
r.matchers = append(r.matchers, d)
|
||||||
return d, nil
|
return d, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -34,20 +32,15 @@ func (r *DomainRegistry) Reload() error {
|
|||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
var matchers []*DynamicDomainMatcher
|
errors.LogInfo(context.Background(), "reloading GeoSite data for ", len(r.matchers), " domain matcher(s)")
|
||||||
r.matchers.Range(func(_ uuid.UUID, matcher *DynamicDomainMatcher) bool {
|
|
||||||
matchers = append(matchers, matcher)
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
errors.LogInfo(context.Background(), "reloading GeoSite data for ", len(matchers), " domain matcher(s)")
|
|
||||||
|
|
||||||
factory := newDomainMatcherFactory()
|
factory := newDomainMatcherFactory()
|
||||||
type reloadEntry struct {
|
type reloadEntry struct {
|
||||||
dynamic *DynamicDomainMatcher
|
dynamic *DynamicDomainMatcher
|
||||||
matcher DomainMatcher
|
matcher DomainMatcher
|
||||||
}
|
}
|
||||||
reloaded := make([]reloadEntry, len(matchers))
|
reloaded := make([]reloadEntry, len(r.matchers))
|
||||||
for i, d := range matchers {
|
for i, d := range r.matchers {
|
||||||
m, err := factory.BuildMatcher(d.rules)
|
m, err := factory.BuildMatcher(d.rules)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
errors.LogErrorInner(context.Background(), err, "failed to reload GeoSite data for domain matcher ", i)
|
errors.LogErrorInner(context.Background(), err, "failed to reload GeoSite data for domain matcher ", i)
|
||||||
@@ -59,14 +52,13 @@ func (r *DomainRegistry) Reload() error {
|
|||||||
entry.dynamic.Reload(entry.matcher)
|
entry.dynamic.Reload(entry.matcher)
|
||||||
}
|
}
|
||||||
r.factory = factory
|
r.factory = factory
|
||||||
errors.LogInfo(context.Background(), "reloaded GeoSite data for ", len(matchers), " domain matcher(s)")
|
errors.LogInfo(context.Background(), "reloaded GeoSite data for ", len(r.matchers), " domain matcher(s)")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newDomainRegistry() *DomainRegistry {
|
func newDomainRegistry() *DomainRegistry {
|
||||||
return &DomainRegistry{
|
return &DomainRegistry{
|
||||||
factory: newDomainMatcherFactory(),
|
factory: newDomainMatcherFactory(),
|
||||||
matchers: utils.NewWeakCacheMap[uuid.UUID, DynamicDomainMatcher](),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+62
-213
@@ -5,14 +5,11 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"io"
|
"io"
|
||||||
"runtime"
|
"runtime"
|
||||||
"slices"
|
|
||||||
"strings"
|
"strings"
|
||||||
"unicode/utf8"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/platform/filesystem"
|
"github.com/xtls/xray-core/common/platform/filesystem"
|
||||||
|
|
||||||
"google.golang.org/protobuf/encoding/protowire"
|
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -55,56 +52,17 @@ func loadIP(file, code string) ([]*CIDR, error) {
|
|||||||
return geoip.Cidr, nil
|
return geoip.Cidr, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
|
func loadSite(file, code string) ([]*Domain, error) {
|
||||||
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
|
bs, err := loadFile(file, code)
|
||||||
// unmarshalling it into a []*Domain, so value is only valid during fn.
|
|
||||||
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
|
|
||||||
runtime.GC() // peak mem
|
|
||||||
r, err := filesystem.OpenAsset(file)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to open ", file).Base(err)
|
return nil, err
|
||||||
}
|
}
|
||||||
defer r.Close()
|
defer runtime.GC() // peak mem
|
||||||
br := bufio.NewReaderSize(r, 64*1024)
|
var geosite GeoSite
|
||||||
n, err := seek(br, []byte(code))
|
if err := proto.Unmarshal(bs, &geosite); err != nil {
|
||||||
if err != nil {
|
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||||
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
|
||||||
}
|
}
|
||||||
loadErr := func(err error) error {
|
return geosite.Domain, nil
|
||||||
if err == io.EOF {
|
|
||||||
err = io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
|
||||||
}
|
|
||||||
unmarshalErr := func(err error) error {
|
|
||||||
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
|
||||||
}
|
|
||||||
d := newSiteDecoder(attrs, fn)
|
|
||||||
for n > 0 {
|
|
||||||
w, err := br.Peek(min(n, br.Size()))
|
|
||||||
if err != nil {
|
|
||||||
return loadErr(err)
|
|
||||||
}
|
|
||||||
used, err := d.decode(w, len(w) < n)
|
|
||||||
if err != nil {
|
|
||||||
return unmarshalErr(err)
|
|
||||||
}
|
|
||||||
if used == 0 {
|
|
||||||
break // a field longer than the buffer
|
|
||||||
}
|
|
||||||
br.Discard(used)
|
|
||||||
n -= used
|
|
||||||
}
|
|
||||||
if n > 0 {
|
|
||||||
w := make([]byte, n)
|
|
||||||
if _, err := io.ReadFull(br, w); err != nil {
|
|
||||||
return loadErr(err)
|
|
||||||
}
|
|
||||||
if _, err := d.decode(w, false); err != nil {
|
|
||||||
return unmarshalErr(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||||
@@ -124,63 +82,68 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
||||||
br := bufio.NewReaderSize(r, 64*1024)
|
|
||||||
bodyL, err := seek(br, code)
|
|
||||||
if err != nil || !readBody {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
out := make([]byte, bodyL)
|
|
||||||
if _, err := io.ReadFull(br, out); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// seek advances br to the body of the entry for code and returns the body length.
|
|
||||||
func seek(br *bufio.Reader, code []byte) (int, error) {
|
|
||||||
codeL := len(code)
|
codeL := len(code)
|
||||||
if codeL == 0 {
|
if codeL == 0 {
|
||||||
return 0, errors.New("empty code")
|
return nil, errors.New("empty code")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
br := bufio.NewReaderSize(r, 64*1024)
|
||||||
need := 2 + codeL // TODO: if code too long
|
need := 2 + codeL // TODO: if code too long
|
||||||
|
prefixBuf := make([]byte, need)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
if _, err := br.ReadByte(); err != nil {
|
if _, err := br.ReadByte(); err != nil {
|
||||||
return 0, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
x, err := decodeVarint(br)
|
x, err := decodeVarint(br)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
return nil, err
|
||||||
}
|
}
|
||||||
bodyL := int(x)
|
bodyL := int(x)
|
||||||
if bodyL <= 0 {
|
if bodyL <= 0 {
|
||||||
return 0, errors.New("invalid body length: ", bodyL)
|
return nil, errors.New("invalid body length: ", bodyL)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
|
prefixL := bodyL
|
||||||
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
|
if prefixL > need {
|
||||||
prefix, err := br.Peek(min(bodyL, need, br.Size()))
|
prefixL = need
|
||||||
if err != nil {
|
}
|
||||||
if err == io.EOF && len(prefix) > 0 {
|
prefix := prefixBuf[:prefixL]
|
||||||
err = io.ErrUnexpectedEOF // as io.ReadFull
|
if _, err := io.ReadFull(br, prefix); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
match := false
|
||||||
|
if bodyL >= need {
|
||||||
|
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
|
||||||
|
if !readBody {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
match = true
|
||||||
}
|
}
|
||||||
return 0, err
|
|
||||||
}
|
}
|
||||||
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
|
|
||||||
return bodyL, nil
|
remain := bodyL - prefixL
|
||||||
|
if match {
|
||||||
|
out := make([]byte, bodyL)
|
||||||
|
copy(out, prefix)
|
||||||
|
if remain > 0 {
|
||||||
|
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
}
|
}
|
||||||
if _, err := br.Discard(bodyL); err != nil {
|
|
||||||
return 0, err
|
if remain > 0 {
|
||||||
|
if _, err := br.Discard(remain); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
|
|
||||||
// attribute helpers that have been part of this package's API since #5814. The streaming loader
|
|
||||||
// above filters attributes itself without building a *Domain, so it does not use them, but they
|
|
||||||
// are kept for external callers. Their behaviour is unchanged.
|
|
||||||
|
|
||||||
type AttributeMatcher interface {
|
type AttributeMatcher interface {
|
||||||
Match(*Domain) bool
|
Match(*Domain) bool
|
||||||
}
|
}
|
||||||
@@ -222,137 +185,23 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
|
|||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
|
|
||||||
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
|
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
|
||||||
|
domains, err := loadSite(file, code)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
type siteDecoder struct {
|
matcher := NewAllAttrsMatcher(attrs)
|
||||||
want []string
|
if matcher == nil {
|
||||||
has []bool
|
return domains, nil
|
||||||
fn func(Domain_Type, []byte)
|
}
|
||||||
}
|
|
||||||
|
|
||||||
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
|
filtered := make([]*Domain, 0, len(domains))
|
||||||
d := &siteDecoder{fn: fn}
|
for _, d := range domains {
|
||||||
if attrs != "" {
|
if matcher.Match(d) {
|
||||||
d.want = strings.Split(attrs, "@")
|
filtered = append(filtered, d)
|
||||||
d.has = make([]bool, len(d.want))
|
}
|
||||||
}
|
}
|
||||||
return d
|
|
||||||
}
|
|
||||||
|
|
||||||
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
|
return filtered, nil
|
||||||
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
|
|
||||||
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
|
|
||||||
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
|
|
||||||
used := 0
|
|
||||||
for used < len(b) {
|
|
||||||
f, n, err := consumeField(b[used:])
|
|
||||||
if err == io.ErrUnexpectedEOF && more {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return used, err
|
|
||||||
}
|
|
||||||
used += n
|
|
||||||
if f.typ != protowire.BytesType {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
switch f.num {
|
|
||||||
case 1: // code
|
|
||||||
if !utf8.Valid(f.v) {
|
|
||||||
return used, errInvalidUTF8
|
|
||||||
}
|
|
||||||
case 2: // domain
|
|
||||||
t, value, err := decodeDomain(f.v, d.want, d.has)
|
|
||||||
if err != nil {
|
|
||||||
return used, err
|
|
||||||
}
|
|
||||||
if !slices.Contains(d.has, false) {
|
|
||||||
d.fn(t, value)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return used, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
|
|
||||||
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
|
|
||||||
clear(has)
|
|
||||||
for len(b) > 0 {
|
|
||||||
f, n, err := consumeField(b)
|
|
||||||
if err != nil {
|
|
||||||
return 0, nil, err
|
|
||||||
}
|
|
||||||
b = b[n:]
|
|
||||||
switch {
|
|
||||||
case f.num == 1 && f.typ == protowire.VarintType: // type
|
|
||||||
t = Domain_Type(f.x)
|
|
||||||
case f.num == 2 && f.typ == protowire.BytesType: // value
|
|
||||||
if !utf8.Valid(f.v) {
|
|
||||||
return 0, nil, errInvalidUTF8
|
|
||||||
}
|
|
||||||
value = f.v
|
|
||||||
case f.num == 3 && f.typ == protowire.BytesType: // attribute
|
|
||||||
key, err := decodeAttributeKey(f.v)
|
|
||||||
if err != nil {
|
|
||||||
return 0, nil, err
|
|
||||||
}
|
|
||||||
for i, w := range want {
|
|
||||||
if string(key) == w {
|
|
||||||
has[i] = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return t, value, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
|
|
||||||
func decodeAttributeKey(b []byte) ([]byte, error) {
|
|
||||||
var key []byte
|
|
||||||
for len(b) > 0 {
|
|
||||||
f, n, err := consumeField(b)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
b = b[n:]
|
|
||||||
if f.num == 1 && f.typ == protowire.BytesType {
|
|
||||||
if !utf8.Valid(f.v) {
|
|
||||||
return nil, errInvalidUTF8
|
|
||||||
}
|
|
||||||
key = f.v
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return key, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type protoField struct {
|
|
||||||
num protowire.Number
|
|
||||||
typ protowire.Type
|
|
||||||
v []byte // payload of a length-delimited field
|
|
||||||
x uint64 // value of a varint field
|
|
||||||
}
|
|
||||||
|
|
||||||
// consumeField parses the first field of an encoded message and returns it with its length.
|
|
||||||
func consumeField(b []byte) (protoField, int, error) {
|
|
||||||
num, typ, n := protowire.ConsumeTag(b)
|
|
||||||
if n < 0 {
|
|
||||||
return protoField{}, 0, protowire.ParseError(n)
|
|
||||||
}
|
|
||||||
if num > protowire.MaxValidNumber {
|
|
||||||
return protoField{}, 0, errors.New("invalid field number ", num)
|
|
||||||
}
|
|
||||||
f := protoField{num: num, typ: typ}
|
|
||||||
var m int
|
|
||||||
switch typ {
|
|
||||||
case protowire.BytesType:
|
|
||||||
f.v, m = protowire.ConsumeBytes(b[n:])
|
|
||||||
case protowire.VarintType:
|
|
||||||
f.x, m = protowire.ConsumeVarint(b[n:])
|
|
||||||
default:
|
|
||||||
m = protowire.ConsumeFieldValue(num, typ, b[n:])
|
|
||||||
}
|
|
||||||
if m < 0 {
|
|
||||||
return protoField{}, 0, protowire.ParseError(m)
|
|
||||||
}
|
|
||||||
return f, n + m, nil
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,283 +0,0 @@
|
|||||||
package geodata
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"slices"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"google.golang.org/protobuf/encoding/protowire"
|
|
||||||
"google.golang.org/protobuf/proto"
|
|
||||||
)
|
|
||||||
|
|
||||||
type siteEntry struct {
|
|
||||||
Type Domain_Type
|
|
||||||
Value string
|
|
||||||
}
|
|
||||||
|
|
||||||
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
|
|
||||||
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
|
|
||||||
var site GeoSite
|
|
||||||
if err := proto.Unmarshal(b, &site); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
var entries []siteEntry
|
|
||||||
for _, d := range site.Domain {
|
|
||||||
ok := true
|
|
||||||
for _, key := range strings.Split(attrs, "@") {
|
|
||||||
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
|
|
||||||
}
|
|
||||||
if ok {
|
|
||||||
entries = append(entries, siteEntry{d.Type, d.Value})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return entries, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
|
|
||||||
t.Helper()
|
|
||||||
want, wantErr := unmarshalSite(b, attrs)
|
|
||||||
var got []siteEntry
|
|
||||||
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
|
||||||
got = append(got, siteEntry{typ, string(value)})
|
|
||||||
}).decode(b, false)
|
|
||||||
if (err == nil) != (wantErr == nil) {
|
|
||||||
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
|
|
||||||
}
|
|
||||||
if err == nil && !slices.Equal(got, want) {
|
|
||||||
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
|
|
||||||
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
for len(bs) > 0 {
|
|
||||||
num, typ, n := protowire.ConsumeTag(bs)
|
|
||||||
if n < 0 || num != 1 || typ != protowire.BytesType {
|
|
||||||
t.Fatal("unexpected GeoSiteList field")
|
|
||||||
}
|
|
||||||
entry, m := protowire.ConsumeBytes(bs[n:])
|
|
||||||
if m < 0 {
|
|
||||||
t.Fatal(protowire.ParseError(m))
|
|
||||||
}
|
|
||||||
bs = bs[n+m:]
|
|
||||||
|
|
||||||
var site GeoSite
|
|
||||||
if err := proto.Unmarshal(entry, &site); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
queries := []string{"", "none"}
|
|
||||||
for _, d := range site.Domain {
|
|
||||||
for _, a := range d.Attribute {
|
|
||||||
if !slices.Contains(queries, a.Key) {
|
|
||||||
queries = append(queries, a.Key, a.Key+"@none")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for _, attrs := range queries {
|
|
||||||
checkDecodeSite(t, site.Code, entry, attrs)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDecodeSiteUnusualEncodings(t *testing.T) {
|
|
||||||
field := func(num protowire.Number, v []byte) []byte {
|
|
||||||
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
|
|
||||||
}
|
|
||||||
typ := func(v Domain_Type) []byte {
|
|
||||||
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
|
|
||||||
}
|
|
||||||
value := func(s string) []byte { return field(2, []byte(s)) }
|
|
||||||
attr := func(keys ...string) []byte {
|
|
||||||
var b []byte
|
|
||||||
for _, k := range keys {
|
|
||||||
b = append(b, field(1, []byte(k))...)
|
|
||||||
}
|
|
||||||
return field(3, b)
|
|
||||||
}
|
|
||||||
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
|
|
||||||
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
|
|
||||||
|
|
||||||
for name, b := range map[string][]byte{
|
|
||||||
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
|
|
||||||
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
|
|
||||||
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
|
|
||||||
"repeated key": domain(value("a.com"), attr("cn", "ads")),
|
|
||||||
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
|
|
||||||
"no value": domain(typ(Domain_Domain), attr("cn")),
|
|
||||||
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
|
|
||||||
"invalid utf8": domain(value("example.\xff")),
|
|
||||||
"invalid key": domain(value("a.com"), attr("\xff")),
|
|
||||||
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
|
|
||||||
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
|
|
||||||
} {
|
|
||||||
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
|
|
||||||
checkDecodeSite(t, name, b, attrs)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
|
|
||||||
// buffer, with a field longer than the buffer in the middle, and a file cut short.
|
|
||||||
func TestLoadSiteReadsInPieces(t *testing.T) {
|
|
||||||
site := &GeoSite{Code: "BIG"}
|
|
||||||
for i := range 5000 {
|
|
||||||
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
|
|
||||||
if i%3 == 0 {
|
|
||||||
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
|
||||||
}
|
|
||||||
if i == 2500 {
|
|
||||||
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
|
|
||||||
}
|
|
||||||
site.Domain = append(site.Domain, d)
|
|
||||||
}
|
|
||||||
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
|
|
||||||
bs, err := proto.Marshal(list)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
entry, err := proto.Marshal(site)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
dir := t.TempDir()
|
|
||||||
t.Setenv("xray.location.asset", dir)
|
|
||||||
write := func(b []byte) {
|
|
||||||
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for _, attrs := range []string{"", "cn"} {
|
|
||||||
want, _ := unmarshalSite(entry, attrs)
|
|
||||||
var got []siteEntry
|
|
||||||
write(bs)
|
|
||||||
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
|
||||||
got = append(got, siteEntry{typ, string(value)})
|
|
||||||
})
|
|
||||||
if err != nil || !slices.Equal(got, want) {
|
|
||||||
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
|
|
||||||
}
|
|
||||||
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
|
|
||||||
write(bs[:cut])
|
|
||||||
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
|
|
||||||
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
|
|
||||||
func oneEntryGeoSiteFile(entry []byte) []byte {
|
|
||||||
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
|
|
||||||
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
|
|
||||||
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
|
|
||||||
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
|
|
||||||
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
|
|
||||||
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
|
|
||||||
const window = 64 * 1024
|
|
||||||
site := &GeoSite{Code: "BIG"}
|
|
||||||
for i := range 12000 { // ~250 KiB, four windows
|
|
||||||
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
|
|
||||||
if i%3 == 0 {
|
|
||||||
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
|
||||||
}
|
|
||||||
site.Domain = append(site.Domain, d)
|
|
||||||
}
|
|
||||||
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
|
|
||||||
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
|
|
||||||
entry, err := proto.Marshal(site)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
dir := t.TempDir()
|
|
||||||
t.Setenv("xray.location.asset", dir)
|
|
||||||
|
|
||||||
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
|
|
||||||
// either side of a window edge), and truncations at the same places.
|
|
||||||
type mut struct {
|
|
||||||
name string
|
|
||||||
make func([]byte) []byte
|
|
||||||
}
|
|
||||||
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
|
|
||||||
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
|
|
||||||
if off < len(entry) {
|
|
||||||
off := off
|
|
||||||
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
|
|
||||||
c := slices.Clone(b)
|
|
||||||
c[off] ^= 0xff
|
|
||||||
return c
|
|
||||||
}})
|
|
||||||
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, attrs := range []string{"", "cn"} {
|
|
||||||
for _, m := range muts {
|
|
||||||
e := m.make(entry)
|
|
||||||
// single-shot reference: decode the whole entry in one call
|
|
||||||
var want []siteEntry
|
|
||||||
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
|
||||||
want = append(want, siteEntry{typ, string(value)})
|
|
||||||
}).decode(e, false)
|
|
||||||
// windowed: loadSite reads the file 64 KiB at a time
|
|
||||||
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
var got []siteEntry
|
|
||||||
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
|
||||||
got = append(got, siteEntry{typ, string(value)})
|
|
||||||
})
|
|
||||||
if (gotErr == nil) != (wantErr == nil) {
|
|
||||||
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
|
|
||||||
}
|
|
||||||
if gotErr == nil && !slices.Equal(got, want) {
|
|
||||||
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
|
|
||||||
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
|
|
||||||
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
|
|
||||||
func TestLoadSiteLongCode(t *testing.T) {
|
|
||||||
longCode := strings.Repeat("Z", 70000)
|
|
||||||
list := &GeoSiteList{Entry: []*GeoSite{
|
|
||||||
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
|
|
||||||
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
|
|
||||||
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
|
|
||||||
}}
|
|
||||||
bs, err := proto.Marshal(list)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
dir := t.TempDir()
|
|
||||||
t.Setenv("xray.location.asset", dir)
|
|
||||||
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
collect := func(code string) ([]siteEntry, error) {
|
|
||||||
var got []siteEntry
|
|
||||||
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
|
|
||||||
got = append(got, siteEntry{typ, string(value)})
|
|
||||||
})
|
|
||||||
return got, err
|
|
||||||
}
|
|
||||||
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
|
|
||||||
t.Fatalf("FIRST: %v %v", got, err)
|
|
||||||
}
|
|
||||||
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
|
|
||||||
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
|
|
||||||
}
|
|
||||||
if _, err := collect(longCode); err == nil {
|
|
||||||
t.Fatal("oversized code: expected a not-found error, got nil")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -11,7 +11,6 @@ import (
|
|||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/utils"
|
|
||||||
|
|
||||||
"go4.org/netipx"
|
"go4.org/netipx"
|
||||||
)
|
)
|
||||||
@@ -807,7 +806,7 @@ func (mm *HeuristicMultiIPMatcher) SetReverse(reverse bool) {
|
|||||||
|
|
||||||
type IPSetFactory struct {
|
type IPSetFactory struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
shared *utils.WeakCacheMap[string, IPSet]
|
shared map[string]*IPSet // TODO: cleanup
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *IPSetFactory) GetOrCreateFromGeoIPRules(rules []*GeoIPRule) (*IPSet, error) {
|
func (f *IPSetFactory) GetOrCreateFromGeoIPRules(rules []*GeoIPRule) (*IPSet, error) {
|
||||||
@@ -816,7 +815,7 @@ func (f *IPSetFactory) GetOrCreateFromGeoIPRules(rules []*GeoIPRule) (*IPSet, er
|
|||||||
f.Lock()
|
f.Lock()
|
||||||
defer f.Unlock()
|
defer f.Unlock()
|
||||||
|
|
||||||
if ipset, ok := f.shared.Load(key); ok {
|
if ipset := f.shared[key]; ipset != nil {
|
||||||
errors.LogDebug(context.Background(), "geodata geoip matcher cache HIT ", key)
|
errors.LogDebug(context.Background(), "geodata geoip matcher cache HIT ", key)
|
||||||
return ipset, nil
|
return ipset, nil
|
||||||
}
|
}
|
||||||
@@ -836,7 +835,7 @@ func (f *IPSetFactory) GetOrCreateFromGeoIPRules(rules []*GeoIPRule) (*IPSet, er
|
|||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
if err == nil {
|
if err == nil {
|
||||||
f.shared.Store(key, ipset)
|
f.shared[key] = ipset
|
||||||
}
|
}
|
||||||
return ipset, err
|
return ipset, err
|
||||||
}
|
}
|
||||||
@@ -1019,5 +1018,5 @@ func buildOptimizedIPMatcher(f *IPSetFactory, rules []*IPRule) (IPMatcher, error
|
|||||||
}
|
}
|
||||||
|
|
||||||
func newIPSetFactory() *IPSetFactory {
|
func newIPSetFactory() *IPSetFactory {
|
||||||
return &IPSetFactory{shared: utils.NewWeakCacheMap[string, IPSet]()}
|
return &IPSetFactory{shared: make(map[string]*IPSet)}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -315,7 +315,7 @@ func TestIPMatcherAnyMatchAndMatches(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !matcher.AnyMatch([]net.IP{
|
if !matcher.AnyMatch([]net.IP{
|
||||||
{},
|
net.IP{},
|
||||||
ip("1.1.1.1"),
|
ip("1.1.1.1"),
|
||||||
ip("8.8.8.8"),
|
ip("8.8.8.8"),
|
||||||
}) {
|
}) {
|
||||||
@@ -345,7 +345,7 @@ func TestIPMatcherAnyMatchAndMatches(t *testing.T) {
|
|||||||
|
|
||||||
if matcher.Matches([]net.IP{
|
if matcher.Matches([]net.IP{
|
||||||
ip("8.8.8.8"),
|
ip("8.8.8.8"),
|
||||||
{},
|
net.IP{},
|
||||||
}) {
|
}) {
|
||||||
t.Fatal("expect Matches to be false when any IP is invalid")
|
t.Fatal("expect Matches to be false when any IP is invalid")
|
||||||
}
|
}
|
||||||
@@ -362,7 +362,7 @@ func TestIPMatcherFilterIPs(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
matched, unmatched := matcher.FilterIPs([]net.IP{
|
matched, unmatched := matcher.FilterIPs([]net.IP{
|
||||||
{},
|
net.IP{},
|
||||||
ip("8.8.8.8"),
|
ip("8.8.8.8"),
|
||||||
ip("91.108.255.254"),
|
ip("91.108.255.254"),
|
||||||
ip("1.1.1.1"),
|
ip("1.1.1.1"),
|
||||||
|
|||||||
@@ -7,27 +7,25 @@ import (
|
|||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/utils"
|
|
||||||
"github.com/xtls/xray-core/common/uuid"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type IPRegistry struct {
|
type IPRegistry struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
factory *IPSetFactory
|
ipsetFactory *IPSetFactory
|
||||||
matchers *utils.WeakCacheMap[uuid.UUID, DynamicIPMatcher]
|
matchers []*DynamicIPMatcher
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *IPRegistry) BuildIPMatcher(rules []*IPRule) (IPMatcher, error) {
|
func (r *IPRegistry) BuildIPMatcher(rules []*IPRule) (IPMatcher, error) {
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
m, err := buildOptimizedIPMatcher(r.factory, rules)
|
m, err := buildOptimizedIPMatcher(r.ipsetFactory, rules)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
d := NewDynamicIPMatcher(rules, m)
|
d := NewDynamicIPMatcher(rules, m)
|
||||||
r.matchers.Store(uuid.New(), d)
|
r.matchers = append(r.matchers, d)
|
||||||
return d, nil
|
return d, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -35,20 +33,15 @@ func (r *IPRegistry) Reload() error {
|
|||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
var matchers []*DynamicIPMatcher
|
errors.LogInfo(context.Background(), "reloading GeoIP data for ", len(r.matchers), " IP matcher(s)")
|
||||||
r.matchers.Range(func(_ uuid.UUID, matcher *DynamicIPMatcher) bool {
|
|
||||||
matchers = append(matchers, matcher)
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
errors.LogInfo(context.Background(), "reloading GeoIP data for ", len(matchers), " IP matcher(s)")
|
|
||||||
|
|
||||||
factory := newIPSetFactory()
|
factory := newIPSetFactory()
|
||||||
type reloadEntry struct {
|
type reloadEntry struct {
|
||||||
dynamic *DynamicIPMatcher
|
dynamic *DynamicIPMatcher
|
||||||
matcher IPMatcher
|
matcher IPMatcher
|
||||||
}
|
}
|
||||||
reloaded := make([]reloadEntry, len(matchers))
|
reloaded := make([]reloadEntry, len(r.matchers))
|
||||||
for i, d := range matchers {
|
for i, d := range r.matchers {
|
||||||
m, err := buildOptimizedIPMatcher(factory, d.rules)
|
m, err := buildOptimizedIPMatcher(factory, d.rules)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
errors.LogErrorInner(context.Background(), err, "failed to reload GeoIP data for IP matcher ", i)
|
errors.LogErrorInner(context.Background(), err, "failed to reload GeoIP data for IP matcher ", i)
|
||||||
@@ -59,15 +52,14 @@ func (r *IPRegistry) Reload() error {
|
|||||||
for _, entry := range reloaded {
|
for _, entry := range reloaded {
|
||||||
entry.dynamic.Reload(entry.matcher)
|
entry.dynamic.Reload(entry.matcher)
|
||||||
}
|
}
|
||||||
r.factory = factory
|
r.ipsetFactory = factory
|
||||||
errors.LogInfo(context.Background(), "reloaded GeoIP data for ", len(matchers), " IP matcher(s)")
|
errors.LogInfo(context.Background(), "reloaded GeoIP data for ", len(r.matchers), " IP matcher(s)")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newIPRegistry() *IPRegistry {
|
func newIPRegistry() *IPRegistry {
|
||||||
return &IPRegistry{
|
return &IPRegistry{
|
||||||
factory: newIPSetFactory(),
|
ipsetFactory: newIPSetFactory(),
|
||||||
matchers: utils.NewWeakCacheMap[uuid.UUID, DynamicIPMatcher](),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,58 +0,0 @@
|
|||||||
package geodata
|
|
||||||
|
|
||||||
import (
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
luar "layeh.com/gopher-luar"
|
|
||||||
)
|
|
||||||
|
|
||||||
// RegisterLua makes xray.geodata available to require in an LState.
|
|
||||||
func RegisterLua(L *lua.LState) {
|
|
||||||
L.PreloadModule("xray.geodata", func(L *lua.LState) int {
|
|
||||||
module := L.NewTable()
|
|
||||||
|
|
||||||
module.RawSetString("BuildDomainMatcher", L.NewFunction(func(L *lua.LState) int {
|
|
||||||
parsed, err := ParseDomainRules(luaRules(L), Domain_Domain)
|
|
||||||
if err != nil {
|
|
||||||
L.RaiseError("%v", err)
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
matcher, err := DomainReg.BuildDomainMatcher(parsed)
|
|
||||||
if err != nil {
|
|
||||||
L.RaiseError("%v", err)
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
L.Push(luar.New(L, matcher))
|
|
||||||
return 1
|
|
||||||
}))
|
|
||||||
|
|
||||||
module.RawSetString("BuildIPMatcher", L.NewFunction(func(L *lua.LState) int {
|
|
||||||
parsed, err := ParseIPRules(luaRules(L))
|
|
||||||
if err != nil {
|
|
||||||
L.RaiseError("%v", err)
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
matcher, err := IPReg.BuildIPMatcher(parsed)
|
|
||||||
if err != nil {
|
|
||||||
L.RaiseError("%v", err)
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
L.Push(luar.New(L, matcher))
|
|
||||||
return 1
|
|
||||||
}))
|
|
||||||
L.Push(module)
|
|
||||||
return 1
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func luaRules(L *lua.LState) []string {
|
|
||||||
rules := make([]string, L.GetTop())
|
|
||||||
for i := range rules {
|
|
||||||
value, ok := L.Get(i + 1).(lua.LString)
|
|
||||||
if !ok {
|
|
||||||
L.RaiseError("geodata rules must be strings")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
rules[i] = string(value)
|
|
||||||
}
|
|
||||||
return rules
|
|
||||||
}
|
|
||||||
@@ -1,66 +0,0 @@
|
|||||||
package geodata
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestLuaIPMatcher(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
RegisterLua(L)
|
|
||||||
ip := L.NewUserData()
|
|
||||||
ip.Value = net.ParseIP("127.0.0.1")
|
|
||||||
L.SetGlobal("ip", ip)
|
|
||||||
ips := L.NewUserData()
|
|
||||||
ips.Value = []net.IP{ip.Value.(net.IP), net.ParseIP("8.8.8.8")}
|
|
||||||
L.SetGlobal("ips", ips)
|
|
||||||
if err := L.DoString(`
|
|
||||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8", "::1")
|
|
||||||
assert(matcher:Match(ip))
|
|
||||||
assert(matcher:AnyMatch(ips))
|
|
||||||
assert(not matcher:Matches(ips))
|
|
||||||
local matched, unmatched = matcher:FilterIPs(ips)
|
|
||||||
assert(type(matched) == "userdata" and type(unmatched) == "userdata")
|
|
||||||
assert(#matched == 1 and #unmatched == 1)
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaDomainMatcher(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
RegisterLua(L)
|
|
||||||
if err := L.DoString(`
|
|
||||||
local matcher = require("xray.geodata").BuildDomainMatcher("example.com", "full:other.com")
|
|
||||||
assert(matcher:MatchAny("example.com"))
|
|
||||||
assert(matcher:MatchAny("www.example.com"))
|
|
||||||
assert(matcher:MatchAny("other.com"))
|
|
||||||
assert(not matcher:MatchAny("www.other.com"))
|
|
||||||
assert(#(matcher:Match("www.example.com")) == 1)
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaMatchersRejectInvalidRules(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
script string
|
|
||||||
}{
|
|
||||||
{"IP rule", `require("xray.geodata").BuildIPMatcher("not-an-ip")`},
|
|
||||||
{"non-string domain rule", `require("xray.geodata").BuildDomainMatcher("example.com", true)`},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
RegisterLua(L)
|
|
||||||
if err := L.DoString(tc.script); err == nil {
|
|
||||||
t.Fatal("invalid geodata rule was accepted")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -138,7 +138,7 @@ func ParseDomainRule(r string, defaultType Domain_Type) (*DomainRule, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
prefix := 0
|
prefix := 0
|
||||||
for _, ext := range [...]string{"ext:", "ext-domain:", "ext-site:"} {
|
for _, ext := range [...]string{"ext:", "ext-domain:"} {
|
||||||
if strings.HasPrefix(r, ext) {
|
if strings.HasPrefix(r, ext) {
|
||||||
prefix = len(ext)
|
prefix = len(ext)
|
||||||
break
|
break
|
||||||
@@ -167,7 +167,7 @@ func ParseDomainRules(rules []string, defaultType Domain_Type) ([]*DomainRule, e
|
|||||||
}
|
}
|
||||||
|
|
||||||
prefix := 0
|
prefix := 0
|
||||||
for _, ext := range [...]string{"ext:", "ext-domain:", "ext-site:"} {
|
for _, ext := range [...]string{"ext:", "ext-domain:"} {
|
||||||
if strings.HasPrefix(r, ext) {
|
if strings.HasPrefix(r, ext) {
|
||||||
prefix = len(ext)
|
prefix = len(ext)
|
||||||
break
|
break
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"regexp"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -73,64 +72,6 @@ func BenchmarkSubstrMatcher(b *testing.B) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func BenchmarkRegexMatcher(b *testing.B) {
|
|
||||||
patterns := []string{ // taken from geosite
|
|
||||||
`(^|\.)91porn\.(best|com|cool|fun|group|party|plus|site|tw|work)$`,
|
|
||||||
`(^|\.)91porn[0-9]{3}\.me$`,
|
|
||||||
`(^|\.)apiproxy-device-prod-nlb-.+\.amazonaws\.com$`,
|
|
||||||
`(^|\.)dualstack\.apiproxy-.+\.amazonaws\.com$`,
|
|
||||||
`(^|\.)aqdk[0-9]{3}\.com$`,
|
|
||||||
`(^|\.)bilibili3(0[1-9]|1[0-2])\.xyz$`,
|
|
||||||
`(^|\.)byyum([3589]|2[235689]|3[34]|4[1-9]|5[1-79]|6[0134679])?\.com$`,
|
|
||||||
`(^|\.)fiftymvapi\..+$`,
|
|
||||||
`(^|\.)gossipfuli[0-9]{3,4}\.xyz$`,
|
|
||||||
`(^|\.)kpkuang\.(bond|fun|info|one|us)$`,
|
|
||||||
`(^|\.)rule34\.(asia|us|world|xxx|xyz)$`,
|
|
||||||
`(^|\.)[a-z][1-9][0-9][a-z]\.com$`,
|
|
||||||
`.+\.awsdns-[0-9][0-9]\.(co\.uk|com|net|org)$`,
|
|
||||||
`.+\.dkr\.ecr\.[^\.]+\.amazonaws\.com$`,
|
|
||||||
`^(.+\.)*zh\.okaapps\.com$`,
|
|
||||||
`^cdn\d-epicgames-\d+\.file\.myqcloud\.com$`,
|
|
||||||
`^chatgpt-async-webps-prod-\S+-\d+\.webpubsub\.azure\.com$`,
|
|
||||||
`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`,
|
|
||||||
`^speed\.(coe|open)\.ad\.[a-z]{2,6}\.prod\.hosts\.ooklaserver\.net$`,
|
|
||||||
`javdb\d+\.com$`,
|
|
||||||
}
|
|
||||||
domains := []string{
|
|
||||||
"www.google.com", "rr3---sn-4g5edndy.googlevideo.com", "r1---sn-2x3abcde.googlevideo.com", "i.ytimg.com",
|
|
||||||
"graph.facebook.com", "api.twitter.com", "www.baidu.com", "github.com", "objects.githubusercontent.com",
|
|
||||||
"login.microsoftonline.com", "e1234.dscb.akamaiedge.net", "d1a2b3c4d5e6f7.cloudfront.net",
|
|
||||||
"s3.us-east-1.amazonaws.com", "123456789012.dkr.ecr.us-east-1.amazonaws.com", "www.wikipedia.org",
|
|
||||||
"discord.com", "telegram.org", "store.steampowered.com", "www.91porn.com", "ns-1234.awsdns-12.org",
|
|
||||||
}
|
|
||||||
bench := func(b *testing.B, ctor func(pattern string) func(string) bool) {
|
|
||||||
var matchers []func(string) bool
|
|
||||||
for _, p := range patterns {
|
|
||||||
matchers = append(matchers, ctor(p))
|
|
||||||
}
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
for _, d := range domains {
|
|
||||||
for _, match := range matchers {
|
|
||||||
_ = match(d)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
b.Run("regexp", func(b *testing.B) {
|
|
||||||
bench(b, func(pattern string) func(string) bool {
|
|
||||||
return regexp.MustCompile(pattern).MatchString
|
|
||||||
})
|
|
||||||
})
|
|
||||||
b.Run("prefilter", func(b *testing.B) {
|
|
||||||
bench(b, func(pattern string) func(string) bool {
|
|
||||||
m, err := Regex.New(pattern)
|
|
||||||
common.Must(err)
|
|
||||||
return m.Match
|
|
||||||
})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Utility functions for benchmark
|
// Utility functions for benchmark
|
||||||
|
|
||||||
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
||||||
|
|||||||
@@ -52,9 +52,7 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
|
|||||||
func (g *MphIndexMatcher) Build() error {
|
func (g *MphIndexMatcher) Build() error {
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
if err := g.mph.Build(); err != nil {
|
g.mph.Build()
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
@@ -66,17 +64,23 @@ func (g *MphIndexMatcher) Build() error {
|
|||||||
|
|
||||||
// Match implements IndexMatcher.Match.
|
// Match implements IndexMatcher.Match.
|
||||||
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
||||||
var result []uint32
|
result := make([][]uint32, 0, 5)
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
result = g.mph.Match(input) // a new slice, returned without another copy
|
if matches := g.mph.Match(input); len(matches) > 0 {
|
||||||
|
result = append(result, matches)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
result = append(result, g.ac.Match(input)...)
|
if matches := g.ac.Match(input); len(matches) > 0 {
|
||||||
|
result = append(result, matches)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if g.regex != nil {
|
if g.regex != nil {
|
||||||
result = append(result, g.regex.Match(input)...)
|
if matches := g.regex.Match(input); len(matches) > 0 {
|
||||||
|
result = append(result, matches)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return result
|
return CompositeMatches(result)
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements IndexMatcher.MatchAny.
|
// MatchAny implements IndexMatcher.MatchAny.
|
||||||
|
|||||||
@@ -78,10 +78,6 @@ func TestMphIndexMatcher(t *testing.T) {
|
|||||||
Input: "example.com",
|
Input: "example.com",
|
||||||
Output: []uint32{10, 4},
|
Output: []uint32{10, 4},
|
||||||
},
|
},
|
||||||
{
|
|
||||||
Input: "apis.org",
|
|
||||||
Output: []uint32{2, 6},
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
matcherGroup := NewMphIndexMatcher()
|
matcherGroup := NewMphIndexMatcher()
|
||||||
for _, rule := range rules {
|
for _, rule := range rules {
|
||||||
@@ -91,13 +87,8 @@ func TestMphIndexMatcher(t *testing.T) {
|
|||||||
}
|
}
|
||||||
matcherGroup.Build()
|
matcherGroup.Build()
|
||||||
for _, test := range cases {
|
for _, test := range cases {
|
||||||
m := matcherGroup.Match(test.Input)
|
|
||||||
if !reflect.DeepEqual(m, test.Output) {
|
|
||||||
t.Error("unexpected output: ", m, " for test case ", test)
|
|
||||||
}
|
|
||||||
clear(m) // the caller owns the result, so this must not change the next one
|
|
||||||
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
||||||
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
|
t.Error("unexpected output: ", m, " for test case ", test)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,440 +1,198 @@
|
|||||||
package strmatcher
|
package strmatcher
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"math/bits"
|
||||||
"cmp"
|
"runtime"
|
||||||
"encoding/binary"
|
"sort"
|
||||||
"errors"
|
|
||||||
"math"
|
|
||||||
"slices"
|
|
||||||
"strings"
|
"strings"
|
||||||
"unsafe"
|
"unsafe"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Flags of a level1 slot, stored above the record offset.
|
// PrimeRK is the prime base used in Rabin-Karp algorithm.
|
||||||
const (
|
const PrimeRK = 16777619
|
||||||
mphDomain = 1 << 31 // matches the pattern and its subdomains
|
|
||||||
mphFull = 1 << 30 // matches the pattern only
|
|
||||||
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
|
|
||||||
mphOffMask = mphParent - 1
|
|
||||||
)
|
|
||||||
|
|
||||||
// Kinds of an added pattern, indexes of mphKinds.
|
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
|
||||||
const (
|
func RollingHash(hash uint32, input string) uint32 {
|
||||||
mphKindFull = iota
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
mphKindParent
|
hash = hash*PrimeRK + uint32(input[i])
|
||||||
mphKindDomain
|
}
|
||||||
)
|
return hash
|
||||||
|
|
||||||
// mphKinds are the slot flags in the order Match reports their values.
|
|
||||||
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
|
|
||||||
|
|
||||||
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
|
|
||||||
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
|
|
||||||
|
|
||||||
var (
|
|
||||||
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
|
|
||||||
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
|
|
||||||
)
|
|
||||||
|
|
||||||
type mphEntry struct {
|
|
||||||
off uint32 // pattern start in buf
|
|
||||||
value uint32
|
|
||||||
n uint32 // pattern length
|
|
||||||
kind uint8
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
|
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
|
||||||
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
|
// as aeshash if aes instruction is available).
|
||||||
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
|
// With different seed, each MemHash<seed> performs as distinct hash functions.
|
||||||
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
|
func MemHash(seed uint32, input string) uint32 {
|
||||||
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
|
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
|
||||||
type MphMatcherGroup struct {
|
}
|
||||||
arena string
|
|
||||||
level0 []uint16 // bucket -> seed
|
|
||||||
level1 []uint32 // slot -> flags | record offset
|
|
||||||
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
|
|
||||||
n0, n1 uint32
|
|
||||||
mul uint64 // multiplier of the suffix hash
|
|
||||||
single uint32 // the only value if !multi
|
|
||||||
multi bool
|
|
||||||
|
|
||||||
buf []byte // build only, patterns in Add order
|
const (
|
||||||
entries []mphEntry
|
mphMatchTypeCount = 2 // Full and Domain
|
||||||
|
)
|
||||||
|
|
||||||
|
type mphRuleInfo struct {
|
||||||
|
rollingHash uint32
|
||||||
|
matchers [mphMatchTypeCount][]uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
// MphMatcherGroup is an implementation of MatcherGroup.
|
||||||
|
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
||||||
|
type MphMatcherGroup struct {
|
||||||
|
rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup
|
||||||
|
values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence)
|
||||||
|
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
||||||
|
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
||||||
|
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
||||||
|
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
||||||
|
ruleInfos *map[string]mphRuleInfo
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||||
return new(MphMatcherGroup)
|
return &MphMatcherGroup{
|
||||||
|
rules: []string{""},
|
||||||
|
values: [][]uint32{nil},
|
||||||
|
level0: nil,
|
||||||
|
level0Mask: 0,
|
||||||
|
level1: nil,
|
||||||
|
level1Mask: 0,
|
||||||
|
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddFullMatcher implements MatcherGroupForFull.
|
// AddFullMatcher implements MatcherGroupForFull.
|
||||||
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
||||||
g.add(matcher.Pattern(), mphKindFull, value)
|
pattern := strings.ToLower(matcher.Pattern())
|
||||||
|
g.addPattern(0, "", pattern, matcher.Type(), value)
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddDomainMatcher implements MatcherGroupForDomain.
|
// AddDomainMatcher implements MatcherGroupForDomain.
|
||||||
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
||||||
g.add(matcher.Pattern(), mphKindDomain, value)
|
pattern := strings.ToLower(matcher.Pattern())
|
||||||
|
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
|
||||||
|
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
|
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
|
||||||
if g.arena != "" {
|
fullPattern := pattern + suffixPattern
|
||||||
panic(errMphBuilt)
|
info, found := (*g.ruleInfos)[fullPattern]
|
||||||
}
|
if !found {
|
||||||
pattern = strings.ToLower(pattern)
|
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
||||||
off := uint32(len(g.buf))
|
g.rules = append(g.rules, fullPattern)
|
||||||
g.buf = append(g.buf, pattern...)
|
g.values = append(g.values, nil)
|
||||||
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
|
|
||||||
if len(pattern) > 0 && pattern[0] == '.' {
|
|
||||||
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
|
|
||||||
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
|
|
||||||
}
|
}
|
||||||
|
info.matchers[matcherType] = append(info.matchers[matcherType], value)
|
||||||
|
(*g.ruleInfos)[fullPattern] = info
|
||||||
|
return info.rollingHash
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *MphMatcherGroup) key(i uint32) []byte {
|
// Build builds a minimal perfect hash table for insert rules.
|
||||||
e := &g.entries[i]
|
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
|
||||||
return g.buf[e.off : e.off+e.n]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Build builds the hash table. It must be called once, after the last Add.
|
|
||||||
func (g *MphMatcherGroup) Build() error {
|
func (g *MphMatcherGroup) Build() error {
|
||||||
if g.arena != "" {
|
ruleCount := len(*g.ruleInfos)
|
||||||
return errMphBuilt
|
g.level0 = make([]uint32, nextPow2(ruleCount/4))
|
||||||
}
|
g.level0Mask = uint32(len(g.level0) - 1)
|
||||||
if uint64(len(g.buf)) > math.MaxUint32 {
|
g.level1 = make([]uint32, nextPow2(ruleCount))
|
||||||
return errors.New("too many rules for MphMatcherGroup")
|
g.level1Mask = uint32(len(g.level1) - 1)
|
||||||
}
|
|
||||||
recs := g.writeRecords()
|
|
||||||
if len(g.arena) > mphOffMask {
|
|
||||||
return errors.New("too many rules for MphMatcherGroup")
|
|
||||||
}
|
|
||||||
hashes := make([]uint64, len(recs))
|
|
||||||
for _, mul := range mphMultipliers {
|
|
||||||
for i, rec := range recs {
|
|
||||||
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
|
|
||||||
}
|
|
||||||
g.mul = mul
|
|
||||||
if err := g.place(recs, hashes); err != errMphCollision {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return errMphCollision
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
|
// Create buckets based on all rule's rolling hash
|
||||||
func (g *MphMatcherGroup) writeRecords() []uint32 {
|
buckets := make([][]uint32, len(g.level0))
|
||||||
g.multi = false
|
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
||||||
if len(g.entries) > 0 {
|
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
||||||
g.single = g.entries[0].value
|
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
||||||
for _, e := range g.entries {
|
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
||||||
if e.value != g.single {
|
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
|
||||||
g.multi = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
// Equal patterns become neighbours in Add order, so their values keep their priority
|
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
||||||
order := make([]uint32, len(g.entries))
|
runtime.GC() // peak mem
|
||||||
for i := range order {
|
|
||||||
order[i] = uint32(i)
|
|
||||||
}
|
|
||||||
slices.SortFunc(order, func(a, b uint32) int {
|
|
||||||
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
|
|
||||||
})
|
|
||||||
|
|
||||||
size := len(g.buf) + len(g.entries) + 2
|
// Sort buckets in descending order with respect to each bucket's size
|
||||||
if g.multi {
|
bucketIdxs := make([]int, len(buckets))
|
||||||
size += 3 * len(g.entries)
|
for bucketIdx := range buckets {
|
||||||
|
bucketIdxs[bucketIdx] = bucketIdx
|
||||||
}
|
}
|
||||||
arena := make([]byte, 0, size)
|
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
||||||
recs := make([]uint32, 0, len(order))
|
|
||||||
var vals [len(mphKinds)][]uint32
|
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
|
||||||
for i := 0; i < len(order); {
|
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
|
||||||
k := g.key(order[i])
|
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
|
||||||
for t := range vals {
|
for _, bucketIdx := range bucketIdxs {
|
||||||
vals[t] = vals[t][:0]
|
bucket := buckets[bucketIdx]
|
||||||
}
|
hashedBucket = hashedBucket[:0]
|
||||||
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
|
seed := uint32(0)
|
||||||
e := &g.entries[order[i]]
|
for len(hashedBucket) != len(bucket) {
|
||||||
if !slices.Contains(vals[e.kind], e.value) {
|
for _, ruleIdx := range bucket {
|
||||||
vals[e.kind] = append(vals[e.kind], e.value)
|
memHash := MemHash(seed, g.rules[ruleIdx]) & g.level1Mask
|
||||||
}
|
if occupied[memHash] { // Collision occurred with this seed
|
||||||
}
|
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
|
||||||
rec := uint32(len(arena))
|
occupied[hash] = false
|
||||||
if len(k) < 255 {
|
g.level1[hash] = 0
|
||||||
arena = append(arena, byte(len(k)))
|
}
|
||||||
} else {
|
hashedBucket = hashedBucket[:0]
|
||||||
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
|
seed++ // Try next seed
|
||||||
}
|
break
|
||||||
arena = append(arena, k...)
|
|
||||||
for t, v := range vals {
|
|
||||||
if len(v) == 0 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
rec |= mphKinds[t]
|
|
||||||
if g.multi {
|
|
||||||
arena = binary.AppendUvarint(arena, uint64(len(v)))
|
|
||||||
for _, x := range v {
|
|
||||||
arena = binary.AppendUvarint(arena, uint64(x))
|
|
||||||
}
|
}
|
||||||
|
occupied[memHash] = true
|
||||||
|
g.level1[memHash] = ruleIdx // The final value in the hash table
|
||||||
|
hashedBucket = append(hashedBucket, memHash)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
recs = append(recs, rec)
|
g.level0[bucketIdx] = seed // Displacement value for this bucket
|
||||||
}
|
|
||||||
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
|
|
||||||
arena = append(arena, 0)
|
|
||||||
if len(recs) == 0 {
|
|
||||||
arena = append(arena, 0)
|
|
||||||
}
|
|
||||||
g.buf, g.entries = nil, nil
|
|
||||||
if cap(arena)-len(arena) > len(arena)/32 {
|
|
||||||
arena = slices.Clone(arena)
|
|
||||||
}
|
|
||||||
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
|
|
||||||
return recs
|
|
||||||
}
|
|
||||||
|
|
||||||
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
|
|
||||||
// the first seed that puts all its records in free slots.
|
|
||||||
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
|
|
||||||
r := len(recs)
|
|
||||||
n0, n1 := max(1, r/3), max(1, r+r/99)
|
|
||||||
g.n0, g.n1 = uint32(n0), uint32(n1)
|
|
||||||
g.level0 = make([]uint16, n0)
|
|
||||||
g.level1 = make([]uint32, n1)
|
|
||||||
g.fp = make([]uint8, n1)
|
|
||||||
|
|
||||||
start := make([]uint32, n0+1)
|
|
||||||
for _, h := range hashes {
|
|
||||||
start[g.bucket(h)+1]++
|
|
||||||
}
|
|
||||||
for b := range n0 {
|
|
||||||
start[b+1] += start[b]
|
|
||||||
}
|
|
||||||
members := make([]uint32, r)
|
|
||||||
fill := slices.Clone(start[:n0])
|
|
||||||
for i, h := range hashes {
|
|
||||||
b := g.bucket(h)
|
|
||||||
members[fill[b]] = uint32(i)
|
|
||||||
fill[b]++
|
|
||||||
}
|
|
||||||
fill = nil
|
|
||||||
buckets := make([]uint32, n0)
|
|
||||||
for b := range buckets {
|
|
||||||
buckets[b] = uint32(b)
|
|
||||||
}
|
|
||||||
slices.SortStableFunc(buckets, func(a, b uint32) int {
|
|
||||||
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
|
|
||||||
})
|
|
||||||
|
|
||||||
occupied := make([]uint64, (n1+63)/64)
|
|
||||||
var slots []uint32
|
|
||||||
next:
|
|
||||||
for _, b := range buckets {
|
|
||||||
m := members[start[b]:start[b+1]]
|
|
||||||
if len(m) == 0 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
for i := range m {
|
|
||||||
for j := range i {
|
|
||||||
if hashes[m[i]] == hashes[m[j]] {
|
|
||||||
return errMphCollision // no seed can separate them
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
search:
|
|
||||||
for seed := range math.MaxUint16 + 1 {
|
|
||||||
slots = slots[:0]
|
|
||||||
for _, ri := range m {
|
|
||||||
s := g.slot(hashes[ri], uint16(seed))
|
|
||||||
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
|
|
||||||
continue search
|
|
||||||
}
|
|
||||||
slots = append(slots, s)
|
|
||||||
}
|
|
||||||
for k, ri := range m {
|
|
||||||
s := slots[k]
|
|
||||||
occupied[s/64] |= 1 << (s % 64)
|
|
||||||
g.level1[s] = recs[ri]
|
|
||||||
g.fp[s] = uint8(hashes[ri])
|
|
||||||
}
|
|
||||||
g.level0[b] = uint16(seed)
|
|
||||||
continue next
|
|
||||||
}
|
|
||||||
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
|
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
|
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
||||||
func mphHash(mul uint64, s string) uint64 {
|
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
||||||
h := uint64(0)
|
i0 := rollingHash & g.level0Mask
|
||||||
for i := len(s) - 1; i >= 0; i-- {
|
seed := g.level0[i0]
|
||||||
h = h*mul + uint64(s[i])
|
i1 := MemHash(seed, input) & g.level1Mask
|
||||||
}
|
if n := g.level1[i1]; g.rules[n] == input {
|
||||||
return h
|
return n
|
||||||
}
|
|
||||||
|
|
||||||
// mphMix spreads the weak low bits of a suffix hash.
|
|
||||||
func mphMix(h uint64) uint64 {
|
|
||||||
h ^= h >> 32
|
|
||||||
h *= 0xd6e8feb86659fd93
|
|
||||||
return h ^ h>>32
|
|
||||||
}
|
|
||||||
|
|
||||||
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
|
|
||||||
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
|
|
||||||
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
|
|
||||||
return uint32((x * uint64(g.n1)) >> 32)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
|
|
||||||
for shift := 0; ; shift += 7 {
|
|
||||||
c := g.arena[p]
|
|
||||||
p++
|
|
||||||
x |= uint32(c&0x7f) << shift
|
|
||||||
if c < 0x80 {
|
|
||||||
return x, p
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// recSpan returns where the pattern of the record at off starts and how long it is.
|
|
||||||
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
|
|
||||||
n, p = uint32(g.arena[off]), off+1
|
|
||||||
if n == 255 {
|
|
||||||
n, p = g.uvarint(p)
|
|
||||||
}
|
|
||||||
return p, n
|
|
||||||
}
|
|
||||||
|
|
||||||
func (g *MphMatcherGroup) recKey(rec uint32) string {
|
|
||||||
p, n := g.recSpan(rec & mphOffMask)
|
|
||||||
return g.arena[p : p+n]
|
|
||||||
}
|
|
||||||
|
|
||||||
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
|
|
||||||
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
|
|
||||||
f := mphMix(h)
|
|
||||||
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
|
|
||||||
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
|
|
||||||
slot := uintptr(g.slot(f, seed))
|
|
||||||
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
|
|
||||||
if len(s) < 255 {
|
|
||||||
// A record whose length byte is len(s) has len(s) pattern bytes after it
|
|
||||||
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
|
|
||||||
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
|
|
||||||
return e
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
if g.recKey(e) == s {
|
|
||||||
return e
|
|
||||||
}
|
}
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
// appendValues appends the values of record e for the flags in want, in mphKinds order.
|
// Match implements MatcherGroup.Match.
|
||||||
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
|
|
||||||
if !g.multi {
|
|
||||||
for _, flag := range mphKinds {
|
|
||||||
if e&want&flag != 0 {
|
|
||||||
dst = append(dst, g.single)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return dst
|
|
||||||
}
|
|
||||||
if e&want == 0 {
|
|
||||||
return dst
|
|
||||||
}
|
|
||||||
p, n := g.recSpan(e & mphOffMask)
|
|
||||||
p += n
|
|
||||||
for _, flag := range mphKinds {
|
|
||||||
if e&flag == 0 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
var count, v uint32
|
|
||||||
for count, p = g.uvarint(p); count > 0; count-- {
|
|
||||||
v, p = g.uvarint(p)
|
|
||||||
if want&flag != 0 {
|
|
||||||
dst = append(dst, v)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return dst
|
|
||||||
}
|
|
||||||
|
|
||||||
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
|
|
||||||
// the parent domains, nearest first.
|
|
||||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
||||||
var stack [8]uint32
|
matches := make([][]uint32, 0, 5)
|
||||||
parents := stack[:0] // TLD side first
|
hash := uint32(0)
|
||||||
h, mul := uint64(0), g.mul
|
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
|
hash = hash*PrimeRK + uint32(input[i])
|
||||||
if input[i] == '.' {
|
if input[i] == '.' {
|
||||||
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
|
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
||||||
parents = append(parents, e)
|
matches = append(matches, g.values[mphIdx])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
h = h*mul + uint64(input[i])
|
|
||||||
}
|
}
|
||||||
exact := g.lookup(h, input)
|
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
||||||
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
|
matches = append(matches, g.values[mphIdx])
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
|
return CompositeMatchesReverse(matches)
|
||||||
for k := len(parents) - 1; k >= 0; k-- {
|
|
||||||
result = g.appendValues(result, parents[k], mphParent|mphDomain)
|
|
||||||
}
|
|
||||||
return result
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements MatcherGroup.MatchAny.
|
// MatchAny implements MatcherGroup.MatchAny.
|
||||||
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
||||||
h, mul := uint64(0), g.mul
|
hash := uint32(0)
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
|
||||||
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
h = h*mul + uint64(input[i])
|
|
||||||
}
|
|
||||||
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
|
|
||||||
type mphSuffix struct {
|
|
||||||
h uint64
|
|
||||||
off int
|
|
||||||
}
|
|
||||||
|
|
||||||
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
|
|
||||||
// with the hash of input itself: what MatchAny computes, computed once for several groups.
|
|
||||||
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
|
|
||||||
h := uint64(0)
|
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
|
hash = hash*PrimeRK + uint32(input[i])
|
||||||
if input[i] == '.' {
|
if input[i] == '.' {
|
||||||
dst = append(dst, mphSuffix{h, i + 1})
|
if g.Lookup(hash, input[i:]) != 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
}
|
}
|
||||||
h = h*mul + uint64(input[i])
|
|
||||||
}
|
}
|
||||||
return dst, h
|
return g.Lookup(hash, input) != 0
|
||||||
}
|
}
|
||||||
|
|
||||||
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
|
func nextPow2(v int) int {
|
||||||
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
if v <= 1 {
|
||||||
if g.mul != mul {
|
return 1
|
||||||
return g.MatchAny(input) // built with a later multiplier after a collision
|
|
||||||
}
|
}
|
||||||
for _, p := range parents {
|
const MaxUInt = ^uint(0)
|
||||||
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
|
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
|
||||||
return true
|
return int(n)
|
||||||
}
|
|
||||||
}
|
|
||||||
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
//go:noescape
|
||||||
|
//go:linkname strhash runtime.strhash
|
||||||
|
func strhash(p unsafe.Pointer, h uintptr) uintptr
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user