mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-02 05:46:39 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b26a91de4f | ||
|
|
1f304916bd | ||
|
|
0086362663 | ||
|
|
e51b3c3621 | ||
|
|
6243d2a26e | ||
|
|
35e616d3b9 | ||
|
|
08cb6e6bca | ||
|
|
48ad0300ea | ||
|
|
0fc379203f | ||
|
|
fc8f8a451d | ||
|
|
e5e85ca9da | ||
|
|
7780db9bbe | ||
|
|
2953d44734 | ||
|
|
7b8ade3ec5 | ||
|
|
5dda894e29 | ||
|
|
47a2c2ffdc | ||
|
|
7a018833ec | ||
|
|
65e853ed84 | ||
|
|
3519dfecbd | ||
|
|
df261e4479 | ||
|
|
61cad5ec8b | ||
|
|
a642a190ed | ||
|
|
60e2a0c502 | ||
|
|
7d3e44fee2 | ||
|
|
9927942aaa | ||
|
|
a308ded2e6 | ||
|
|
7741e9e77e | ||
|
|
d562d8947d | ||
|
|
dbb1ea30ba | ||
|
|
efc9e6da62 | ||
|
|
8267cf953a | ||
|
|
24e6f6d551 | ||
|
|
dcdfc57ccd | ||
|
|
3461c511aa | ||
|
|
c412e77a9b | ||
|
|
ccb69ea5e2 | ||
|
|
52a412d9e2 | ||
|
|
18a1b5042a | ||
|
|
c26d2eda24 | ||
|
|
a1bf968be9 | ||
|
|
c037ccd98d | ||
|
|
37ceb8b4b6 | ||
|
|
fd2ca74822 | ||
|
|
47cfe9994a | ||
|
|
3e2f040cd8 | ||
|
|
c7245c0336 | ||
|
|
eef6e63bc1 | ||
|
|
6ce8dc53e7 | ||
|
|
01a034be53 | ||
|
|
de2caf3cef | ||
|
|
cecc88f43c | ||
|
|
cd4ce973e9 | ||
|
|
fc7b980636 | ||
|
|
8ee131cbbb | ||
|
|
2776ea6d74 | ||
|
|
5e245b082e | ||
|
|
d9c54026c5 | ||
|
|
c1958dba04 | ||
|
|
540b9070f5 | ||
|
|
ada99a4eb0 | ||
|
|
65458e919f | ||
|
|
aa3d6589da | ||
|
|
25c11e2d2b | ||
|
|
dffc7ada5e | ||
|
|
77f98eba09 | ||
|
|
f124daf5a3 | ||
|
|
598bde7412 | ||
|
|
9b373e39ca | ||
|
|
c7e569b037 | ||
|
|
f02a357861 | ||
|
|
5fe6d6217a | ||
|
|
0604ffa957 | ||
|
|
d3f1a24285 | ||
|
|
2323273e37 | ||
|
|
09107b71dc | ||
|
|
7021606ad3 | ||
|
|
7d214f8b09 | ||
|
|
8b419d833d | ||
|
|
a12801c13b | ||
|
|
a000371b2a | ||
|
|
bc6e966af8 | ||
|
|
fc5620de98 | ||
|
|
b02bdcf4cc | ||
|
|
2b329b3675 | ||
|
|
5ca6f4b7d4 | ||
|
|
18e283909c | ||
|
|
6ab123bf8f | ||
|
|
4aba687dd3 | ||
|
|
5b1b41058e | ||
|
|
6e3322d219 | ||
|
|
1d8eb81d70 | ||
|
|
e78d8ef184 | ||
|
|
6ce924ad56 | ||
|
|
035d438979 | ||
|
|
50231eaff9 | ||
|
|
1f74c480d6 | ||
|
|
af7eb68028 | ||
|
|
35387572e0 | ||
|
|
8f15190c23 | ||
|
|
64fada32b5 | ||
|
|
0bafca9486 | ||
|
|
d5bc58dc6b | ||
|
|
c18b39ed80 | ||
|
|
e2ad0acf60 | ||
|
|
c320e89108 | ||
|
|
412898fed7 | ||
|
|
5c62d50d43 | ||
|
|
1aabe7ea78 | ||
|
|
e4e7614c62 | ||
|
|
987290ba48 | ||
|
|
d7fa2076c3 | ||
|
|
fb548f54d2 | ||
|
|
0495b17650 | ||
|
|
65f6f0a43b | ||
|
|
3263ae9255 | ||
|
|
3dc8bf3d8b | ||
|
|
695e68ef9e | ||
|
|
dfdbcf86cc | ||
|
|
2b828b7bc2 | ||
|
|
45cf2898ab | ||
|
|
18b85adb4e | ||
|
|
452b719504 | ||
|
|
345c76f9a8 | ||
|
|
f496437b84 | ||
|
|
b12bc504c8 | ||
|
|
dda2b10c9d | ||
|
|
241aa38ac0 | ||
|
|
7e7e820763 | ||
|
|
f9eb1597ad | ||
|
|
ac04c445bd | ||
|
|
e7e9254630 | ||
|
|
fab4bcc1ed | ||
|
|
b99c3e5657 | ||
|
|
583bb4a63f | ||
|
|
9cd9382e3d | ||
|
|
567500c4af | ||
|
|
5aefcb41fb | ||
|
|
be8009c625 | ||
|
|
8734774e4a | ||
|
|
1e036ce1c5 | ||
|
|
c815c2f2df | ||
|
|
986c512e0f | ||
|
|
711aea4e34 | ||
|
|
6412738486 | ||
|
|
ad2e4cb0e1 | ||
|
|
829d54d7be | ||
|
|
862631172d | ||
|
|
d27b3e46e2 | ||
|
|
da21a8f77f | ||
|
|
e10347bf01 | ||
|
|
26a022c905 | ||
|
|
95e9816223 | ||
|
|
3239d21168 | ||
|
|
06b4931743 | ||
|
|
6189d2bfd5 | ||
|
|
a0e9347f1b | ||
|
|
83cf229909 | ||
|
|
2249f8b5c6 | ||
|
|
fdb9b616fc | ||
|
|
d792fba59c | ||
|
|
55956f8d70 |
@@ -0,0 +1 @@
|
|||||||
|
powershell.exe -ExecutionPolicy Bypass -File ".\xray_no_window.ps1"
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
Start-Process -FilePath ".\xray.exe" -ArgumentList "-config .\config.json" -WindowStyle Hidden
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
CreateObject("Wscript.Shell").Run "xray.exe -config config.json",0
|
||||||
@@ -37,6 +37,7 @@ 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,6 +37,7 @@ 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,8 +64,16 @@ 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@v6
|
uses: actions/checkout@v7
|
||||||
|
|
||||||
- name: Set up QEMU
|
- name: Set up QEMU
|
||||||
uses: docker/setup-qemu-action@v4
|
uses: docker/setup-qemu-action@v4
|
||||||
@@ -74,7 +82,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
|
uses: docker/login-action@v4.6.0
|
||||||
with:
|
with:
|
||||||
registry: ghcr.io
|
registry: ghcr.io
|
||||||
username: ${{ github.repository_owner }}
|
username: ${{ github.repository_owner }}
|
||||||
@@ -124,6 +132,13 @@ 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 }}
|
||||||
@@ -131,3 +146,7 @@ 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,15 +11,16 @@ 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@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|
||||||
- name: Restore Wintun Cache
|
- name: Restore Wintun Cache
|
||||||
uses: actions/cache/restore@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -75,13 +76,14 @@ 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@v6
|
uses: actions/checkout@v7
|
||||||
|
|
||||||
- name: Show workflow information
|
- name: Show workflow information
|
||||||
run: |
|
run: |
|
||||||
@@ -90,7 +92,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@v6
|
uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
check-latest: true
|
check-latest: true
|
||||||
@@ -117,13 +119,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@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|
||||||
- name: Restore Wintun Cache
|
- name: Restore Wintun Cache
|
||||||
uses: actions/cache/restore@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -132,15 +134,17 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
mv -f resources/geo* build_assets/
|
mv -f resources/geo* build_assets/
|
||||||
if [[ ${GOOS} == 'windows' ]]; then
|
if [[ ${GOOS} == 'windows' ]]; then
|
||||||
echo 'CreateObject("Wscript.Shell").Run "xray.exe -config config.json",0' > build_assets/xray_no_window.vbs
|
cp .github/build/windows/* build_assets/
|
||||||
echo 'Start-Process -FilePath ".\xray.exe" -ArgumentList "-config .\config.json" -WindowStyle Hidden' > build_assets/xray_no_window.ps1
|
fi
|
||||||
|
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.txt
|
mv resources/wintun/LICENSE.txt build_assets/LICENSE-Wintun
|
||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Copy README.md & LICENSE
|
- name: Copy README.md & LICENSE
|
||||||
|
|||||||
@@ -11,15 +11,16 @@ 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@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|
||||||
- name: Restore Wintun Cache
|
- name: Restore Wintun Cache
|
||||||
uses: actions/cache/restore@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -161,6 +162,7 @@ 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 }}
|
||||||
@@ -168,7 +170,7 @@ jobs:
|
|||||||
CGO_ENABLED: 0
|
CGO_ENABLED: 0
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout codebase
|
- name: Checkout codebase
|
||||||
uses: actions/checkout@v6
|
uses: actions/checkout@v7
|
||||||
|
|
||||||
- name: Set up NDK
|
- name: Set up NDK
|
||||||
if: matrix.goos == 'android'
|
if: matrix.goos == 'android'
|
||||||
@@ -191,7 +193,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@v6
|
uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
check-latest: true
|
check-latest: true
|
||||||
@@ -223,14 +225,14 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Restore Geodat Cache
|
- name: Restore Geodat Cache
|
||||||
uses: actions/cache/restore@v5
|
uses: actions/cache/restore@v6
|
||||||
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@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -239,8 +241,10 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
mv -f resources/geo* build_assets/
|
mv -f resources/geo* build_assets/
|
||||||
if [[ ${GOOS} == 'windows' ]]; then
|
if [[ ${GOOS} == 'windows' ]]; then
|
||||||
echo 'CreateObject("Wscript.Shell").Run "xray.exe -config config.json",0' > build_assets/xray_no_window.vbs
|
cp .github/build/windows/* build_assets/
|
||||||
echo 'Start-Process -FilePath ".\xray.exe" -ArgumentList "-config .\config.json" -WindowStyle Hidden' > build_assets/xray_no_window.ps1
|
fi
|
||||||
|
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
|
||||||
@@ -250,7 +254,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.txt
|
mv resources/wintun/LICENSE.txt build_assets/LICENSE-Wintun
|
||||||
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@v5
|
uses: actions/cache/restore@v6
|
||||||
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@v5
|
uses: actions/cache/save@v6
|
||||||
if: ${{ steps.update.outputs.unhit }}
|
if: ${{ steps.update.outputs.unhit }}
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
@@ -68,9 +68,12 @@ 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@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-wintun-
|
key: xray-wintun-
|
||||||
@@ -96,7 +99,6 @@ 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
|
||||||
@@ -113,16 +115,21 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
if [[ "$missing" == true ]]; then
|
if [[ "$missing" == true ]]; then
|
||||||
FILENAME=wintun.zip
|
FILENAME=wintun.zip
|
||||||
DOWNLOAD_FILE=wintun-0.14.1.zip
|
DOWNLOAD_FILE=wintun-${ASSETVER}.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}"
|
||||||
echo -e "Unpacking wintun..."
|
if [[ "$(sha256sum "./${FILENAME}" | awk -F ' ' '{print $1}')" == "${ASSETHASH}" ]]; then
|
||||||
unzip -u ${FILENAME} -d resources/
|
echo -e "Unpacking wintun..."
|
||||||
echo "unhit=true" >> $GITHUB_OUTPUT
|
unzip -u ${FILENAME} -d resources/
|
||||||
|
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@v5
|
uses: actions/cache/save@v6
|
||||||
if: ${{ steps.update.outputs.unhit }}
|
if: ${{ steps.update.outputs.unhit }}
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ jobs:
|
|||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
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@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
@@ -40,7 +40,7 @@ jobs:
|
|||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
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@v6
|
uses: actions/checkout@v7
|
||||||
- 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
|
||||||
@@ -59,17 +59,15 @@ jobs:
|
|||||||
contents: read
|
contents: read
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout codebase
|
- name: Checkout codebase
|
||||||
uses: actions/checkout@v6
|
uses: actions/checkout@v7
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v6
|
uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version-file: go.mod
|
go-version-file: go.mod
|
||||||
check-latest: true
|
check-latest: true
|
||||||
cache: false
|
cache: false
|
||||||
- name: Check Format
|
- name: Check Format
|
||||||
run: |
|
run: go run ./infra/vformat/main.go -mode check -pwd ./
|
||||||
go install -v mvdan.cc/gofumpt@latest
|
|
||||||
go run ./infra/vformat/main.go -mode check -pwd ./
|
|
||||||
|
|
||||||
test:
|
test:
|
||||||
needs: check-assets
|
needs: check-assets
|
||||||
@@ -83,14 +81,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@v6
|
uses: actions/checkout@v7
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v6
|
uses: actions/setup-go@v7
|
||||||
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@v5
|
uses: actions/cache/restore@v6
|
||||||
with:
|
with:
|
||||||
path: resources
|
path: resources
|
||||||
key: xray-geodat-
|
key: xray-geodat-
|
||||||
|
|||||||
@@ -73,6 +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)
|
||||||
- [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`
|
||||||
@@ -145,6 +146,8 @@
|
|||||||
- [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...
|
||||||
|
|
||||||
@@ -184,6 +187,27 @@
|
|||||||
- [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)
|
||||||
|
|||||||
@@ -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, _ := stats.GetOrRegisterCounter(d.stats, name); c != nil {
|
if c, _ := d.stats.GetOrRegisterCounter(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, _ := stats.GetOrRegisterCounter(d.stats, name); c != nil {
|
if c, _ := d.stats.GetOrRegisterCounter(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, _ := stats.GetOrRegisterCounter(statsManager, name); c != nil {
|
if c, _ := statsManager.GetOrRegisterCounter(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, _ := stats.GetOrRegisterCounter(statsManager, name); c != nil {
|
if c, _ := statsManager.GetOrRegisterCounter(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, _ := stats.GetOrRegisterOnlineMap(sm, name); om != nil {
|
if om, _ := sm.GetOrRegisterOnlineMap(name); om != nil {
|
||||||
om.AddIP(ip)
|
om.AddIP(ip)
|
||||||
context.AfterFunc(ctx, func() { om.RemoveIP(ip) })
|
context.AfterFunc(ctx, func() { om.RemoveIP(ip) })
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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").AtError()
|
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used")
|
||||||
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) {
|
||||||
|
|||||||
+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.").AtWarning()
|
return nil, errors.New("Failed to convert address", addr, "to Net IP.")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return ips, nil
|
return ips, nil
|
||||||
|
|||||||
@@ -212,6 +212,28 @@ 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
|
||||||
|
|||||||
@@ -0,0 +1,59 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -188,19 +188,24 @@ 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).AtWarning()
|
return nil, errors.New("failed to parse DNS response").Base(err)
|
||||||
}
|
}
|
||||||
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).AtWarning()
|
return nil, errors.New("failed to skip questions in DNS response").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
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 {
|
||||||
@@ -217,7 +222,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.After(expire) {
|
if ipRecord.Expire.IsZero() || ipRecord.Expire.After(expire) {
|
||||||
ipRecord.Expire = expire
|
ipRecord.Expire = expire
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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).AtError()
|
return nil, errors.New("Unable to create Fake Dns Engine").Base(err)
|
||||||
}
|
}
|
||||||
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).AtError()
|
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
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").AtError()
|
return errors.New("LRU size is bigger than subnet size")
|
||||||
}
|
}
|
||||||
fkdns.domainToIP = cache.NewLru(lruSize)
|
fkdns.domainToIP = cache.NewLru(lruSize)
|
||||||
fkdns.ipRange = ipRange
|
fkdns.ipRange = ipRange
|
||||||
|
|||||||
@@ -84,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).AtWarning()
|
return nil, errors.New("No available name server could be created from ", dest)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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.
|
||||||
@@ -102,7 +102,7 @@ func NewClient(
|
|||||||
// 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).AtWarning()
|
return errors.New("failed to create nameserver").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, isLocalDNS := server.(*LocalNameServer)
|
_, isLocalDNS := server.(*LocalNameServer)
|
||||||
@@ -113,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).AtWarning()
|
return errors.New("failed to create expected ip matcher").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -122,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).AtWarning()
|
return errors.New("failed to create unexpected ip matcher").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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").AtError()
|
return nil, 0, errors.New("Unable to locate a fake DNS Engine")
|
||||||
}
|
}
|
||||||
|
|
||||||
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).AtError()
|
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
||||||
|
|||||||
+85
-43
@@ -2,6 +2,7 @@ package geodata
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/tls"
|
||||||
go_errors "errors"
|
go_errors "errors"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -9,6 +10,7 @@ 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"
|
||||||
@@ -16,6 +18,7 @@ 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
|
||||||
@@ -26,8 +29,9 @@ type stage struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type downloader struct {
|
type downloader struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
client *http.Client
|
httpClient *http.Client
|
||||||
|
httpsClient *http.Client
|
||||||
}
|
}
|
||||||
|
|
||||||
type idleConn struct {
|
type idleConn struct {
|
||||||
@@ -53,52 +57,84 @@ 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,
|
||||||
client: newClient(ctx, dispatcher, outbound),
|
httpClient: newClient(ctx, dispatcher, outbound, false),
|
||||||
|
httpsClient: newClient(ctx, dispatcher, outbound, true),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func newClient(baseCtx context.Context, dispatcher routing.Dispatcher, outbound string) *http.Client {
|
func newClient(baseCtx context.Context, dispatcher routing.Dispatcher, outbound string, isHTTPS bool) *http.Client {
|
||||||
return &http.Client{
|
dial := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||||
Transport: &http.Transport{
|
var conn net.Conn
|
||||||
Proxy: nil,
|
err := task.Run(ctx, func() error {
|
||||||
DisableKeepAlives: true,
|
if tagged.Dialer == nil {
|
||||||
DialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
return errors.New("tagged dialer is not initialized")
|
||||||
var conn net.Conn
|
|
||||||
err := task.Run(ctx, func() error {
|
|
||||||
if tagged.Dialer == nil {
|
|
||||||
return errors.New("tagged dialer is not initialized")
|
|
||||||
}
|
|
||||||
dest, err := net.ParseDestination(network + ":" + address)
|
|
||||||
if err != nil {
|
|
||||||
return errors.New("cannot understand address").Base(err)
|
|
||||||
}
|
|
||||||
c, err := tagged.Dialer(baseCtx, dispatcher, dest, outbound)
|
|
||||||
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
|
|
||||||
},
|
|
||||||
TLSHandshakeTimeout: idleTimeout,
|
|
||||||
ResponseHeaderTimeout: idleTimeout,
|
|
||||||
},
|
|
||||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
|
||||||
if req.URL.Scheme != "https" {
|
|
||||||
return errors.New("redirected to non-https URL: ", req.URL.String())
|
|
||||||
}
|
}
|
||||||
if len(via) >= 10 {
|
dest, err := net.ParseDestination(network + ":" + address)
|
||||||
return errors.New("stopped after 10 redirects")
|
if err != nil {
|
||||||
|
return errors.New("cannot understand address").Base(err)
|
||||||
}
|
}
|
||||||
|
c, err := tagged.Dialer(baseCtx, dispatcher, dest, outbound)
|
||||||
|
if err != nil {
|
||||||
|
return errors.New("cannot dial remote address ", dest).Base(err)
|
||||||
|
}
|
||||||
|
conn = c
|
||||||
return nil
|
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 {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
host, _, _ := net.SplitHostPort(address)
|
||||||
|
tlsConn := utls.UClient(conn, &utls.Config{ServerName: host}, utls.HelloChrome_Auto)
|
||||||
|
handshakeCtx, cancel := context.WithTimeout(ctx, idleTimeout)
|
||||||
|
defer cancel()
|
||||||
|
if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
|
||||||
|
conn.Close()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return tlsConn, nil
|
||||||
|
},
|
||||||
|
},
|
||||||
|
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||||
|
if req.URL.Scheme != "https" {
|
||||||
|
return errors.New("redirected to non-https URL: ", req.URL.String())
|
||||||
|
}
|
||||||
|
if len(via) >= 10 {
|
||||||
|
return errors.New("stopped after 10 redirects")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
return &http.Client{
|
||||||
|
Transport: &http.Transport{
|
||||||
|
Proxy: nil,
|
||||||
|
DisableKeepAlives: true,
|
||||||
|
DialContext: dial,
|
||||||
|
ResponseHeaderTimeout: idleTimeout,
|
||||||
|
},
|
||||||
|
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||||
|
if req.URL.Scheme != "https" {
|
||||||
|
return errors.New("redirected to non-https URL: ", req.URL.String())
|
||||||
|
}
|
||||||
|
if len(via) >= 10 {
|
||||||
|
return errors.New("stopped after 10 redirects")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -160,7 +196,13 @@ func (d *downloader) fetch(rawURL string, writer io.Writer) error {
|
|||||||
}
|
}
|
||||||
utils.TryDefaultHeadersWith(req.Header, "nav")
|
utils.TryDefaultHeadersWith(req.Header, "nav")
|
||||||
|
|
||||||
resp, err := d.client.Do(req)
|
var client *http.Client
|
||||||
|
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
|
||||||
}
|
}
|
||||||
|
|||||||
+6
-2
@@ -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).AtWarning()
|
return errors.New("failed to initialize access logger").Base(err)
|
||||||
}
|
}
|
||||||
if err := g.initErrorLogger(); err != nil {
|
if err := g.initErrorLogger(); err != nil {
|
||||||
return errors.New("failed to initialize error logger").Base(err).AtWarning()
|
return errors.New("failed to initialize error logger").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -141,6 +141,10 @@ 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")
|
||||||
|
|||||||
+151
-59
@@ -2,15 +2,18 @@ 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"
|
||||||
"github.com/xtls/xray-core/common/net"
|
xnet "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"
|
||||||
@@ -21,15 +24,17 @@ import (
|
|||||||
type MetricsHandler struct {
|
type MetricsHandler struct {
|
||||||
ohm outbound.Manager
|
ohm outbound.Manager
|
||||||
statsManager feature_stats.Manager
|
statsManager feature_stats.Manager
|
||||||
observatory extension.Observatory
|
ctx context.Context
|
||||||
tag string
|
tag string
|
||||||
listen string
|
listen string
|
||||||
tcpListener net.Listener
|
tcpListener xnet.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,
|
||||||
}
|
}
|
||||||
@@ -37,46 +42,6 @@ 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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -85,45 +50,172 @@ 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 := net.Listen("tcp", p.listen)
|
TCPlistener, err := xnet.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 func() {
|
go p.serve(TCPlistener, handler)
|
||||||
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 net.Conn, 4),
|
buffer: make(chan xnet.Conn, 4),
|
||||||
done: done.New(),
|
done: done.New(),
|
||||||
}
|
}
|
||||||
|
p.listener = listener
|
||||||
|
|
||||||
go func() {
|
go p.serve(listener, handler)
|
||||||
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")
|
||||||
}
|
}
|
||||||
|
|
||||||
return p.ohm.AddHandler(context.Background(), &Outbound{
|
if err := 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 {
|
||||||
return nil
|
var errs []error
|
||||||
|
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() {
|
||||||
|
|||||||
@@ -0,0 +1,161 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -78,6 +78,12 @@ 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 {
|
||||||
|
|||||||
+11
-22
@@ -330,7 +330,6 @@ 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"`
|
||||||
@@ -382,13 +381,6 @@ 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
|
||||||
@@ -506,14 +498,13 @@ 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\"\x9d\x03\n" +
|
"\x0eOutboundConfig\"\xd6\x02\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\x12K\n" +
|
"\x0fstream_settings\x18\x02 \x01(\v2%.xray.transport.internet.StreamConfigR\x0estreamSettings\x12T\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\x0etargetStrategy\"\xa4\x01\n" +
|
"\x0ftarget_strategy\x18\x06 \x01(\x0e2'.xray.transport.internet.DomainStrategyR\x0etargetStrategyJ\x04\b\x03\x10\x04\"\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" +
|
||||||
@@ -548,8 +539,7 @@ 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.ProxyConfig)(nil), // 13: xray.transport.internet.ProxyConfig
|
(internet.DomainStrategy)(0), // 13: xray.transport.internet.DomainStrategy
|
||||||
(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
|
||||||
@@ -562,14 +552,13 @@ 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
|
||||||
13, // 10: xray.app.proxyman.SenderConfig.proxy_settings:type_name -> xray.transport.internet.ProxyConfig
|
6, // 10: xray.app.proxyman.SenderConfig.multiplex_settings:type_name -> xray.app.proxyman.MultiplexingConfig
|
||||||
6, // 11: xray.app.proxyman.SenderConfig.multiplex_settings:type_name -> xray.app.proxyman.MultiplexingConfig
|
13, // 11: xray.app.proxyman.SenderConfig.target_strategy:type_name -> xray.transport.internet.DomainStrategy
|
||||||
14, // 12: xray.app.proxyman.SenderConfig.target_strategy:type_name -> xray.transport.internet.DomainStrategy
|
12, // [12:12] is the sub-list for method output_type
|
||||||
13, // [13:13] is the sub-list for method output_type
|
12, // [12:12] is the sub-list for method input_type
|
||||||
13, // [13:13] is the sub-list for method input_type
|
12, // [12:12] is the sub-list for extension type_name
|
||||||
13, // [13:13] is the sub-list for extension type_name
|
12, // [12:12] is the sub-list for extension extendee
|
||||||
13, // [13:13] is the sub-list for extension extendee
|
0, // [0:12] is the sub-list for field type_name
|
||||||
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;
|
||||||
xray.transport.internet.ProxyConfig proxy_settings = 3;
|
reserved 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, _ := stats.GetOrRegisterCounter(statsManager, name)
|
c, _ := statsManager.GetOrRegisterCounter(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, _ := stats.GetOrRegisterCounter(statsManager, name)
|
c, _ := statsManager.GetOrRegisterCounter(name)
|
||||||
if c != nil {
|
if c != nil {
|
||||||
downlinkCounter = c
|
downlinkCounter = c
|
||||||
}
|
}
|
||||||
@@ -57,16 +57,23 @@ 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)
|
||||||
// Set tag and sniffing config in context before creating proxy
|
if receiverConfig.Listen != nil {
|
||||||
// This allows proxies like TUN to access these settings
|
src.Address = receiverConfig.Listen.AsAddress()
|
||||||
ctx = session.ContextWithInbound(ctx, &session.Inbound{Tag: tag})
|
|
||||||
if receiverConfig.SniffingSettings != nil {
|
|
||||||
ctx = session.ContextWithContent(ctx, &session.Content{
|
|
||||||
SniffingRequest: sniffingRequest,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
rawProxy, err := common.CreateObject(ctx, proxyConfig)
|
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})
|
||||||
|
newCtx = session.ContextWithContent(newCtx, &session.Content{SniffingRequest: sniffingRequest})
|
||||||
|
newCtx = session.ContextWithStreamSettings(newCtx, mss)
|
||||||
|
|
||||||
|
rawProxy, err := common.CreateObject(newCtx, proxyConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -92,11 +99,6 @@ 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{}
|
||||||
@@ -170,6 +172,12 @@ 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
|
||||||
|
|||||||
@@ -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").AtError()
|
return nil, errors.New("not a ReceiverConfig")
|
||||||
}
|
}
|
||||||
|
|
||||||
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).AtWarning().Base(err)
|
return errors.New("failed to listen TCP on ", w.port).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).AtWarning().Base(err)
|
return errors.New("failed to listen Unix Domain Socket on ", w.address).Base(err)
|
||||||
}
|
}
|
||||||
w.hub = hub
|
w.hub = hub
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ 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"
|
||||||
|
|
||||||
@@ -16,7 +15,6 @@ 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"
|
||||||
@@ -27,8 +25,6 @@ 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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -40,7 +36,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, _ := stats.GetOrRegisterCounter(statsManager, name)
|
c, _ := statsManager.GetOrRegisterCounter(name)
|
||||||
if c != nil {
|
if c != nil {
|
||||||
uplinkCounter = c
|
uplinkCounter = c
|
||||||
}
|
}
|
||||||
@@ -48,7 +44,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, _ := stats.GetOrRegisterCounter(statsManager, name)
|
c, _ := statsManager.GetOrRegisterCounter(name)
|
||||||
if c != nil {
|
if c != nil {
|
||||||
downlinkCounter = c
|
downlinkCounter = c
|
||||||
}
|
}
|
||||||
@@ -64,7 +60,6 @@ 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
|
||||||
@@ -78,7 +73,6 @@ 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,
|
||||||
}
|
}
|
||||||
@@ -93,7 +87,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).AtWarning()
|
return nil, errors.New("failed to parse stream settings").Base(err)
|
||||||
}
|
}
|
||||||
h.streamSettings = mss
|
h.streamSettings = mss
|
||||||
default:
|
default:
|
||||||
@@ -109,6 +103,10 @@ 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
|
||||||
@@ -196,7 +194,6 @@ 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))])
|
||||||
@@ -220,7 +217,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").AtInfo())
|
test(errors.New("XUDP rejected UDP/443 traffic"))
|
||||||
return
|
return
|
||||||
case "skip":
|
case "skip":
|
||||||
goto out
|
goto out
|
||||||
@@ -269,71 +266,26 @@ 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 {
|
if h.senderSettings != nil && h.senderSettings.Via != nil {
|
||||||
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
if h.senderSettings.ProxySettings.HasTag() {
|
ob := outbounds[len(outbounds)-1]
|
||||||
|
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 && !h.senderSettings.ProxySettings.HasTag() && (h.streamSettings.SocketSettings == nil || len(h.streamSettings.SocketSettings.DialerProxy) == 0) {
|
if ob.Gateway == nil && h.senderSettings != nil && h.senderSettings.Via != nil &&
|
||||||
|
(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() {
|
||||||
@@ -348,12 +300,9 @@ 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())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// case addr.Family().IsDomain():
|
default: // case addr.Family().IsDomain():
|
||||||
default:
|
|
||||||
ob.Gateway = addr
|
ob.Gateway = addr
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,35 +0,0 @@
|
|||||||
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").AtError()
|
return errors.New("outbound metadata not found")
|
||||||
}
|
}
|
||||||
|
|
||||||
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).AtWarning()
|
return errors.New("failed to create mux client worker").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
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[tag]; ok {
|
if b, ok := (*r.balancers.Load())[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[tag]; ok {
|
if b, ok := (*r.balancers.Load())[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[tag]; ok {
|
if b, ok := (*r.balancers.Load())[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,25 +2,8 @@ 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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"runtime"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -393,3 +394,22 @@ 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,7 +2,9 @@ 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"
|
||||||
@@ -343,6 +345,31 @@ 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,6 +33,10 @@ 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))
|
||||||
}
|
}
|
||||||
@@ -111,7 +115,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").AtWarning()
|
return nil, errors.New("this rule has no effective fields")
|
||||||
}
|
}
|
||||||
|
|
||||||
return conds, nil
|
return conds, nil
|
||||||
@@ -141,7 +145,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").AtError()
|
return nil, errors.New("not a StrategyLeastLoadConfig")
|
||||||
}
|
}
|
||||||
leastLoadStrategy := NewLeastLoadStrategy(s)
|
leastLoadStrategy := NewLeastLoadStrategy(s)
|
||||||
return &Balancer{
|
return &Balancer{
|
||||||
|
|||||||
+14
-4
@@ -107,8 +107,10 @@ 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"`
|
||||||
unknownFields protoimpl.UnknownFields
|
// List of operating systems for matching the one Xray itself is running on.
|
||||||
sizeCache protoimpl.SizeCache
|
LocalOs []string `protobuf:"bytes,23,rep,name=local_os,json=localOs,proto3" json:"local_os,omitempty"`
|
||||||
|
unknownFields protoimpl.UnknownFields
|
||||||
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *RoutingRule) Reset() {
|
func (x *RoutingRule) Reset() {
|
||||||
@@ -278,6 +280,13 @@ 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()
|
||||||
}
|
}
|
||||||
@@ -637,7 +646,7 @@ 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\"\xc1\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\"\xdc\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" +
|
||||||
@@ -661,7 +670,8 @@ 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\x1a=\n" +
|
"\awebhook\x18\x16 \x01(\v2\x1e.xray.app.router.WebhookConfigR\awebhook\x12\x19\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" +
|
||||||
|
|||||||
@@ -56,6 +56,9 @@ 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 {
|
||||||
|
|||||||
+59
-114
@@ -2,7 +2,9 @@ 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"
|
||||||
@@ -17,8 +19,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 []*Rule
|
rules atomic.Pointer[[]*Rule]
|
||||||
balancers map[string]*Balancer
|
balancers atomic.Pointer[map[string]*Balancer]
|
||||||
dns dns.Client
|
dns dns.Client
|
||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
@@ -43,52 +45,9 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out
|
|||||||
r.ohm = ohm
|
r.ohm = ohm
|
||||||
r.dispatcher = dispatcher
|
r.dispatcher = dispatcher
|
||||||
|
|
||||||
r.balancers = make(map[string]*Balancer, len(config.BalancingRule))
|
r.rules.Store(new([]*Rule))
|
||||||
for _, rule := range config.BalancingRule {
|
r.balancers.Store(&map[string]*Balancer{})
|
||||||
balancer, err := rule.Build(ohm, dispatcher)
|
return r.ReloadRules(config, false)
|
||||||
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.
|
||||||
@@ -124,18 +83,22 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
|
|||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
if !shouldAppend {
|
oldRules := *r.rules.Load()
|
||||||
for _, rule := range r.rules {
|
oldBalancers := *r.balancers.Load()
|
||||||
if rule.Webhook != nil {
|
|
||||||
rule.Webhook.Close()
|
var newRules []*Rule
|
||||||
}
|
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 {
|
||||||
_, found := r.balancers[rule.Tag]
|
if _, found := newBalancers[rule.Tag]; found {
|
||||||
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)
|
||||||
@@ -143,27 +106,12 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
balancer.InjectContext(r.ctx)
|
balancer.InjectContext(r.ctx)
|
||||||
r.balancers[rule.Tag] = balancer
|
newBalancers[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{
|
||||||
@@ -171,69 +119,64 @@ 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
|
||||||
}
|
}
|
||||||
btag := rule.GetBalancingTag()
|
if btag := rule.GetBalancingTag(); len(btag) > 0 {
|
||||||
if len(btag) > 0 {
|
brule, found := newBalancers[btag]
|
||||||
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
|
||||||
}
|
}
|
||||||
r.rules = append(r.rules, rr)
|
newRules = append(newRules, 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()
|
||||||
|
|
||||||
newRules := []*Rule{}
|
oldRules := *r.rules.Load()
|
||||||
if tag != "" {
|
newRules := make([]*Rule, 0, len(oldRules))
|
||||||
for _, rule := range r.rules {
|
var removed []*Rule
|
||||||
if rule.RuleTag != tag {
|
for _, rule := range oldRules {
|
||||||
newRules = append(newRules, rule)
|
if rule.RuleTag != tag {
|
||||||
} else if rule.Webhook != nil {
|
newRules = append(newRules, rule)
|
||||||
rule.Webhook.Close()
|
} else {
|
||||||
}
|
removed = append(removed, rule)
|
||||||
}
|
}
|
||||||
r.rules = newRules
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
return errors.New("empty tag name!")
|
r.rules.Store(&newRules)
|
||||||
|
closeWebhooks(removed)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListRule implements routing.Router
|
// ListRule implements routing.Router
|
||||||
func (r *Router) ListRule() []routing.Route {
|
func (r *Router) ListRule() []routing.Route {
|
||||||
r.mu.Lock()
|
rules := *r.rules.Load()
|
||||||
defer r.mu.Unlock()
|
ruleList := make([]routing.Route, 0, len(rules))
|
||||||
ruleList := make([]routing.Route, 0)
|
for _, rule := range rules {
|
||||||
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,
|
||||||
@@ -252,7 +195,9 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
|||||||
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
|
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, rule := range r.rules {
|
rules := *r.rules.Load()
|
||||||
|
|
||||||
|
for _, rule := range rules {
|
||||||
if rule.Apply(ctx) {
|
if rule.Apply(ctx) {
|
||||||
return rule, ctx, nil
|
return rule, ctx, nil
|
||||||
}
|
}
|
||||||
@@ -265,7 +210,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 r.rules {
|
for _, rule := range rules {
|
||||||
if rule.Apply(ctx) {
|
if rule.Apply(ctx) {
|
||||||
return rule, ctx, nil
|
return rule, ctx, nil
|
||||||
}
|
}
|
||||||
@@ -279,9 +224,9 @@ func (r *Router) Start() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// closeWebhooks closes all webhook notifiers in the current rule set.
|
// closeWebhooks closes all webhook notifiers in the given rule set.
|
||||||
func (r *Router) closeWebhooks() {
|
func closeWebhooks(rules []*Rule) {
|
||||||
for _, rule := range r.rules {
|
for _, rule := range rules {
|
||||||
if rule.Webhook != nil {
|
if rule.Webhook != nil {
|
||||||
rule.Webhook.Close()
|
rule.Webhook.Close()
|
||||||
}
|
}
|
||||||
@@ -292,7 +237,7 @@ func (r *Router) closeWebhooks() {
|
|||||||
func (r *Router) Close() error {
|
func (r *Router) Close() error {
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
r.closeWebhooks()
|
closeWebhooks(*r.rules.Load())
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+17
-23
@@ -8,6 +8,7 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
@@ -40,6 +41,7 @@ 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
|
||||||
@@ -77,11 +79,6 @@ func NewWebhookNotifier(cfg *WebhookConfig) (*WebhookNotifier, error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if h.deduplication > 0 {
|
|
||||||
h.wg.Add(1)
|
|
||||||
go h.cleanupLoop()
|
|
||||||
}
|
|
||||||
|
|
||||||
return h, nil
|
return h, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -201,6 +198,7 @@ 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
|
||||||
@@ -210,27 +208,23 @@ func (h *WebhookNotifier) isDuplicate(email string) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *WebhookNotifier) cleanupLoop() {
|
func (h *WebhookNotifier) maybeSweep(now time.Time, ttl time.Duration) {
|
||||||
defer h.wg.Done()
|
last := h.lastSweep.Load()
|
||||||
ttl := time.Duration(h.deduplication) * time.Second
|
if now.UnixNano()-last < int64(ttl) {
|
||||||
ticker := time.NewTicker(ttl)
|
return
|
||||||
defer ticker.Stop()
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-h.done:
|
|
||||||
return
|
|
||||||
case <-ticker.C:
|
|
||||||
now := time.Now()
|
|
||||||
h.seen.Range(func(key, value any) bool {
|
|
||||||
if now.Sub(value.(time.Time)) >= ttl {
|
|
||||||
h.seen.Delete(key)
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
if !h.lastSweep.CompareAndSwap(last, now.UnixNano()) {
|
||||||
|
return // another goroutine did the sweep
|
||||||
|
}
|
||||||
|
h.seen.Range(func(key, value any) bool {
|
||||||
|
if now.Sub(value.(time.Time)) >= ttl {
|
||||||
|
h.seen.Delete(key)
|
||||||
|
}
|
||||||
|
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,6 +48,20 @@ 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()
|
||||||
@@ -97,6 +111,20 @@ 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()
|
||||||
@@ -149,6 +177,26 @@ 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()
|
||||||
|
|||||||
+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 {
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package buf
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"io"
|
"io"
|
||||||
|
"sync/atomic"
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/platform"
|
"github.com/xtls/xray-core/common/platform"
|
||||||
@@ -143,13 +144,24 @@ func (r *ReadVReader) ReadMultiBuffer() (MultiBuffer, error) {
|
|||||||
return mb, nil
|
return mb, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var useReadv bool
|
var useReadv atomic.Bool
|
||||||
|
|
||||||
func init() {
|
func useReadV() bool {
|
||||||
|
return useReadv.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
func reloadEnvSettings() error {
|
||||||
const defaultFlagValue = "NOT_DEFINED_AT_ALL"
|
const defaultFlagValue = "NOT_DEFINED_AT_ALL"
|
||||||
value := platform.NewEnvFlag(platform.UseReadV).GetValue(func() string { return defaultFlagValue })
|
value := platform.NewEnvFlag(platform.UseReadV).GetValue(func() string { return defaultFlagValue })
|
||||||
|
enabled := false
|
||||||
switch value {
|
switch value {
|
||||||
case defaultFlagValue, "auto", "enable":
|
case defaultFlagValue, "auto", "enable":
|
||||||
useReadv = true
|
enabled = true
|
||||||
}
|
}
|
||||||
|
useReadv.Store(enabled)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
platform.RegisterEnvReload(reloadEnvSettings)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,7 +10,9 @@ import (
|
|||||||
"github.com/xtls/xray-core/features/stats"
|
"github.com/xtls/xray-core/features/stats"
|
||||||
)
|
)
|
||||||
|
|
||||||
const useReadv = false
|
func useReadV() bool {
|
||||||
|
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,7 +5,8 @@ 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) {
|
||||||
@@ -15,6 +16,7 @@ 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() {
|
||||||
@@ -25,6 +27,14 @@ 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,7 +118,9 @@ 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()
|
||||||
}
|
}
|
||||||
|
|||||||
+13
-65
@@ -18,17 +18,12 @@ 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().
|
||||||
@@ -69,46 +64,6 @@ 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()
|
||||||
@@ -132,9 +87,8 @@ func New(msg ...interface{}) *Error {
|
|||||||
details = details[:i]
|
details = details[:i]
|
||||||
}
|
}
|
||||||
return &Error{
|
return &Error{
|
||||||
message: msg,
|
message: msg,
|
||||||
severity: log.Severity_Info,
|
caller: details,
|
||||||
caller: details,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -171,6 +125,9 @@ 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 {
|
||||||
@@ -181,10 +138,9 @@ 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,
|
||||||
severity: severity,
|
caller: details,
|
||||||
caller: details,
|
inner: inner,
|
||||||
inner: inner,
|
|
||||||
}
|
}
|
||||||
if ctx != nil && ctx != context.Background() {
|
if ctx != nil && ctx != context.Background() {
|
||||||
id := uint32(c.IDFromContext(ctx))
|
id := uint32(c.IDFromContext(ctx))
|
||||||
@@ -193,7 +149,7 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
log.Record(&log.GeneralMessage{
|
log.Record(&log.GeneralMessage{
|
||||||
Severity: GetSeverity(err),
|
Severity: severity,
|
||||||
Content: err,
|
Content: err,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -217,11 +173,3 @@ 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,30 +7,21 @@ 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 := GetSeverity(err); v != log.Severity_Info {
|
if v := err.Error(); !strings.Contains(v, "TestError") {
|
||||||
t.Error("severity: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = New("TestError2").Base(io.EOF)
|
err = New("TestError2").Base(io.EOF)
|
||||||
if v := GetSeverity(err); v != log.Severity_Info {
|
if v := err.Error(); !strings.Contains(v, "EOF") {
|
||||||
t.Error("severity: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = New("TestError3").Base(io.EOF).AtWarning()
|
err = New("TestError3").Base(io.EOF)
|
||||||
if v := GetSeverity(err); v != log.Severity_Warning {
|
err = New("TestError4").Base(err)
|
||||||
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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,49 @@
|
|||||||
|
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() }
|
||||||
@@ -82,19 +82,10 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
}
|
}
|
||||||
g.Add(m, uint32(i))
|
g.Add(m, uint32(i))
|
||||||
case *DomainRule_Geosite:
|
case *DomainRule_Geosite:
|
||||||
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
|
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
|
||||||
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")
|
||||||
}
|
}
|
||||||
@@ -108,12 +99,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
return g, nil
|
return g, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type CompactDomainMatcherFactory struct {
|
type CompactMphDomainMatcherFactory struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
|
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
|
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
|
||||||
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
||||||
|
|
||||||
f.Lock()
|
f.Lock()
|
||||||
@@ -125,33 +116,23 @@ func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmat
|
|||||||
}
|
}
|
||||||
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
||||||
|
|
||||||
s := strmatcher.NewLinearAnyMatcher()
|
s := strmatcher.NewMphValueMatcher()
|
||||||
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
|
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
for i, d := range domains {
|
if err := s.Build(); err != nil {
|
||||||
domains[i] = nil // peak mem
|
return nil, err
|
||||||
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.Store(key, s)
|
||||||
return s, err
|
return s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// BuildMatcher implements DomainMatcherFactory.
|
// BuildMatcher implements DomainMatcherFactory.
|
||||||
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
func (f *CompactMphDomainMatcherFactory) 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 := &CompactDomainMatcher{
|
compact := new(CompactMphDomainMatcher)
|
||||||
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:
|
||||||
@@ -168,8 +149,7 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
compact.matchers = append(compact.matchers, m)
|
compact.combiner.Add(m, uint32(i))
|
||||||
compact.values = append(compact.values, uint32(i))
|
|
||||||
default:
|
default:
|
||||||
panic("unknown domain rule type")
|
panic("unknown domain rule type")
|
||||||
}
|
}
|
||||||
@@ -177,37 +157,40 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
|
|||||||
return compact, nil
|
return compact, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type CompactDomainMatcher struct {
|
type CompactMphDomainMatcher struct {
|
||||||
custom strmatcher.ValueMatcher
|
custom strmatcher.ValueMatcher
|
||||||
matchers []strmatcher.MatcherSet
|
combiner strmatcher.MphValueMatcherCombiner
|
||||||
values []uint32
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Match implements DomainMatcher.
|
// Match implements DomainMatcher.
|
||||||
func (c *CompactDomainMatcher) Match(input string) []uint32 {
|
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
|
||||||
var result []uint32
|
result := c.combiner.Match(input)
|
||||||
if c.custom != nil {
|
if c.custom != nil {
|
||||||
result = append(result, c.custom.Match(input)...)
|
result = append(c.custom.Match(input), result...)
|
||||||
}
|
|
||||||
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 *CompactDomainMatcher) MatchAny(input string) bool {
|
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
|
||||||
if c.custom != nil && c.custom.MatchAny(input) {
|
if c.custom != nil && c.custom.MatchAny(input) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
for _, m := range c.matchers {
|
return c.combiner.MatchAny(input)
|
||||||
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) {
|
||||||
@@ -220,7 +203,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(d.Value)
|
return strmatcher.Domain.New(strings.ToLower(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:
|
||||||
@@ -231,7 +214,7 @@ 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 &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
default:
|
default:
|
||||||
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ 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"
|
||||||
@@ -11,7 +12,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
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"}}},
|
||||||
@@ -32,7 +33,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 := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
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"}}},
|
||||||
@@ -72,3 +73,76 @@ 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,12 +6,14 @@ 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 []*DynamicDomainMatcher
|
matchers *utils.WeakCacheMap[uuid.UUID, DynamicDomainMatcher]
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *DomainRegistry) BuildDomainMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
func (r *DomainRegistry) BuildDomainMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||||
@@ -24,7 +26,7 @@ func (r *DomainRegistry) BuildDomainMatcher(rules []*DomainRule) (DomainMatcher,
|
|||||||
}
|
}
|
||||||
|
|
||||||
d := NewDynamicDomainMatcher(rules, m)
|
d := NewDynamicDomainMatcher(rules, m)
|
||||||
r.matchers = append(r.matchers, d)
|
r.matchers.Store(uuid.New(), d)
|
||||||
return d, nil
|
return d, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -32,15 +34,20 @@ func (r *DomainRegistry) Reload() error {
|
|||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
errors.LogInfo(context.Background(), "reloading GeoSite data for ", len(r.matchers), " domain matcher(s)")
|
var matchers []*DynamicDomainMatcher
|
||||||
|
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(r.matchers))
|
reloaded := make([]reloadEntry, len(matchers))
|
||||||
for i, d := range r.matchers {
|
for i, d := range 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)
|
||||||
@@ -52,13 +59,14 @@ 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(r.matchers), " domain matcher(s)")
|
errors.LogInfo(context.Background(), "reloaded GeoSite data for ", len(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](),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+213
-62
@@ -5,11 +5,14 @@ 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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -52,17 +55,56 @@ func loadIP(file, code string) ([]*CIDR, error) {
|
|||||||
return geoip.Cidr, nil
|
return geoip.Cidr, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadSite(file, code string) ([]*Domain, error) {
|
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
|
||||||
bs, err := loadFile(file, code)
|
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
|
||||||
|
// 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 nil, err
|
return errors.New("failed to open ", file).Base(err)
|
||||||
}
|
}
|
||||||
defer runtime.GC() // peak mem
|
defer r.Close()
|
||||||
var geosite GeoSite
|
br := bufio.NewReaderSize(r, 64*1024)
|
||||||
if err := proto.Unmarshal(bs, &geosite); err != nil {
|
n, err := seek(br, []byte(code))
|
||||||
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
if err != nil {
|
||||||
|
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||||
}
|
}
|
||||||
return geosite.Domain, nil
|
loadErr := func(err error) error {
|
||||||
|
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) {
|
||||||
@@ -82,68 +124,63 @@ 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 nil, errors.New("empty code")
|
return 0, 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 nil, err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
x, err := decodeVarint(br)
|
x, err := decodeVarint(br)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return 0, err
|
||||||
}
|
}
|
||||||
bodyL := int(x)
|
bodyL := int(x)
|
||||||
if bodyL <= 0 {
|
if bodyL <= 0 {
|
||||||
return nil, errors.New("invalid body length: ", bodyL)
|
return 0, errors.New("invalid body length: ", bodyL)
|
||||||
}
|
}
|
||||||
|
|
||||||
prefixL := bodyL
|
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
|
||||||
if prefixL > need {
|
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
|
||||||
prefixL = need
|
prefix, err := br.Peek(min(bodyL, need, br.Size()))
|
||||||
}
|
if err != nil {
|
||||||
prefix := prefixBuf[:prefixL]
|
if err == io.EOF && len(prefix) > 0 {
|
||||||
if _, err := io.ReadFull(br, prefix); err != nil {
|
err = io.ErrUnexpectedEOF // as io.ReadFull
|
||||||
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) {
|
||||||
remain := bodyL - prefixL
|
return bodyL, nil
|
||||||
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 {
|
||||||
if remain > 0 {
|
return 0, err
|
||||||
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
|
||||||
}
|
}
|
||||||
@@ -185,23 +222,137 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
|
|||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
|
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
|
||||||
domains, err := loadSite(file, code)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
matcher := NewAllAttrsMatcher(attrs)
|
type siteDecoder struct {
|
||||||
if matcher == nil {
|
want []string
|
||||||
return domains, nil
|
has []bool
|
||||||
}
|
fn func(Domain_Type, []byte)
|
||||||
|
}
|
||||||
|
|
||||||
filtered := make([]*Domain, 0, len(domains))
|
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
|
||||||
for _, d := range domains {
|
d := &siteDecoder{fn: fn}
|
||||||
if matcher.Match(d) {
|
if attrs != "" {
|
||||||
filtered = append(filtered, d)
|
d.want = strings.Split(attrs, "@")
|
||||||
|
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),
|
||||||
|
// 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
|
||||||
return filtered, 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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,283 @@
|
|||||||
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -7,25 +7,27 @@ 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
|
||||||
ipsetFactory *IPSetFactory
|
factory *IPSetFactory
|
||||||
matchers []*DynamicIPMatcher
|
matchers *utils.WeakCacheMap[uuid.UUID, 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.ipsetFactory, rules)
|
m, err := buildOptimizedIPMatcher(r.factory, rules)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
d := NewDynamicIPMatcher(rules, m)
|
d := NewDynamicIPMatcher(rules, m)
|
||||||
r.matchers = append(r.matchers, d)
|
r.matchers.Store(uuid.New(), d)
|
||||||
return d, nil
|
return d, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -33,15 +35,20 @@ func (r *IPRegistry) Reload() error {
|
|||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|
||||||
errors.LogInfo(context.Background(), "reloading GeoIP data for ", len(r.matchers), " IP matcher(s)")
|
var matchers []*DynamicIPMatcher
|
||||||
|
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(r.matchers))
|
reloaded := make([]reloadEntry, len(matchers))
|
||||||
for i, d := range r.matchers {
|
for i, d := range 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)
|
||||||
@@ -52,14 +59,15 @@ func (r *IPRegistry) Reload() error {
|
|||||||
for _, entry := range reloaded {
|
for _, entry := range reloaded {
|
||||||
entry.dynamic.Reload(entry.matcher)
|
entry.dynamic.Reload(entry.matcher)
|
||||||
}
|
}
|
||||||
r.ipsetFactory = factory
|
r.factory = factory
|
||||||
errors.LogInfo(context.Background(), "reloaded GeoIP data for ", len(r.matchers), " IP matcher(s)")
|
errors.LogInfo(context.Background(), "reloaded GeoIP data for ", len(matchers), " IP matcher(s)")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newIPRegistry() *IPRegistry {
|
func newIPRegistry() *IPRegistry {
|
||||||
return &IPRegistry{
|
return &IPRegistry{
|
||||||
ipsetFactory: newIPSetFactory(),
|
factory: newIPSetFactory(),
|
||||||
|
matchers: utils.NewWeakCacheMap[uuid.UUID, DynamicIPMatcher](),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -138,7 +138,7 @@ func ParseDomainRule(r string, defaultType Domain_Type) (*DomainRule, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
prefix := 0
|
prefix := 0
|
||||||
for _, ext := range [...]string{"ext:", "ext-domain:"} {
|
for _, ext := range [...]string{"ext:", "ext-domain:", "ext-site:"} {
|
||||||
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:"} {
|
for _, ext := range [...]string{"ext:", "ext-domain:", "ext-site:"} {
|
||||||
if strings.HasPrefix(r, ext) {
|
if strings.HasPrefix(r, ext) {
|
||||||
prefix = len(ext)
|
prefix = len(ext)
|
||||||
break
|
break
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"regexp"
|
||||||
"strconv"
|
"strconv"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -72,6 +73,64 @@ 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,7 +52,9 @@ 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
|
||||||
g.mph.Build()
|
if err := g.mph.Build(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
@@ -64,23 +66,17 @@ 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 {
|
||||||
result := make([][]uint32, 0, 5)
|
var result []uint32
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
result = append(result, g.ac.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.regex != nil {
|
if g.regex != nil {
|
||||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
result = append(result, g.regex.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return CompositeMatches(result)
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements IndexMatcher.MatchAny.
|
// MatchAny implements IndexMatcher.MatchAny.
|
||||||
|
|||||||
@@ -78,6 +78,10 @@ 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 {
|
||||||
@@ -87,8 +91,13 @@ func TestMphIndexMatcher(t *testing.T) {
|
|||||||
}
|
}
|
||||||
matcherGroup.Build()
|
matcherGroup.Build()
|
||||||
for _, test := range cases {
|
for _, test := range cases {
|
||||||
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
m := matcherGroup.Match(test.Input)
|
||||||
|
if !reflect.DeepEqual(m, test.Output) {
|
||||||
t.Error("unexpected output: ", m, " for test case ", test)
|
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) {
|
||||||
|
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,198 +1,440 @@
|
|||||||
package strmatcher
|
package strmatcher
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"math/bits"
|
"bytes"
|
||||||
"runtime"
|
"cmp"
|
||||||
"sort"
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"math"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"unsafe"
|
"unsafe"
|
||||||
)
|
)
|
||||||
|
|
||||||
// PrimeRK is the prime base used in Rabin-Karp algorithm.
|
// Flags of a level1 slot, stored above the record offset.
|
||||||
const PrimeRK = 16777619
|
|
||||||
|
|
||||||
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
|
|
||||||
func RollingHash(hash uint32, input string) uint32 {
|
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
|
||||||
hash = hash*PrimeRK + uint32(input[i])
|
|
||||||
}
|
|
||||||
return hash
|
|
||||||
}
|
|
||||||
|
|
||||||
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
|
|
||||||
// as aeshash if aes instruction is available).
|
|
||||||
// With different seed, each MemHash<seed> performs as distinct hash functions.
|
|
||||||
func MemHash(seed uint32, input string) uint32 {
|
|
||||||
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
|
|
||||||
}
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
mphMatchTypeCount = 2 // Full and Domain
|
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
|
||||||
)
|
)
|
||||||
|
|
||||||
type mphRuleInfo struct {
|
// Kinds of an added pattern, indexes of mphKinds.
|
||||||
rollingHash uint32
|
const (
|
||||||
matchers [mphMatchTypeCount][]uint32
|
mphKindFull = iota
|
||||||
|
mphKindParent
|
||||||
|
mphKindDomain
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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.
|
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
|
||||||
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
|
||||||
|
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
|
||||||
|
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
|
||||||
|
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
|
||||||
type MphMatcherGroup struct {
|
type MphMatcherGroup struct {
|
||||||
rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup
|
arena string
|
||||||
values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence)
|
level0 []uint16 // bucket -> seed
|
||||||
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
level1 []uint32 // slot -> flags | record offset
|
||||||
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
|
||||||
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
n0, n1 uint32
|
||||||
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
mul uint64 // multiplier of the suffix hash
|
||||||
ruleInfos *map[string]mphRuleInfo
|
single uint32 // the only value if !multi
|
||||||
|
multi bool
|
||||||
|
|
||||||
|
buf []byte // build only, patterns in Add order
|
||||||
|
entries []mphEntry
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||||
return &MphMatcherGroup{
|
return new(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) {
|
||||||
pattern := strings.ToLower(matcher.Pattern())
|
g.add(matcher.Pattern(), mphKindFull, value)
|
||||||
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) {
|
||||||
pattern := strings.ToLower(matcher.Pattern())
|
g.add(matcher.Pattern(), mphKindDomain, value)
|
||||||
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) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
|
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
|
||||||
fullPattern := pattern + suffixPattern
|
if g.arena != "" {
|
||||||
info, found := (*g.ruleInfos)[fullPattern]
|
panic(errMphBuilt)
|
||||||
if !found {
|
}
|
||||||
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
pattern = strings.ToLower(pattern)
|
||||||
g.rules = append(g.rules, fullPattern)
|
off := uint32(len(g.buf))
|
||||||
g.values = append(g.values, nil)
|
g.buf = append(g.buf, pattern...)
|
||||||
|
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build builds a minimal perfect hash table for insert rules.
|
func (g *MphMatcherGroup) key(i uint32) []byte {
|
||||||
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
|
e := &g.entries[i]
|
||||||
|
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 {
|
||||||
ruleCount := len(*g.ruleInfos)
|
if g.arena != "" {
|
||||||
g.level0 = make([]uint32, nextPow2(ruleCount/4))
|
return errMphBuilt
|
||||||
g.level0Mask = uint32(len(g.level0) - 1)
|
|
||||||
g.level1 = make([]uint32, nextPow2(ruleCount))
|
|
||||||
g.level1Mask = uint32(len(g.level1) - 1)
|
|
||||||
|
|
||||||
// Create buckets based on all rule's rolling hash
|
|
||||||
buckets := make([][]uint32, len(g.level0))
|
|
||||||
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
|
||||||
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
|
||||||
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
|
||||||
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
|
||||||
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
|
|
||||||
}
|
}
|
||||||
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
if uint64(len(g.buf)) > math.MaxUint32 {
|
||||||
runtime.GC() // peak mem
|
return errors.New("too many rules for MphMatcherGroup")
|
||||||
|
|
||||||
// Sort buckets in descending order with respect to each bucket's size
|
|
||||||
bucketIdxs := make([]int, len(buckets))
|
|
||||||
for bucketIdx := range buckets {
|
|
||||||
bucketIdxs[bucketIdx] = bucketIdx
|
|
||||||
}
|
}
|
||||||
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
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
|
||||||
|
}
|
||||||
|
|
||||||
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
|
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
|
||||||
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
|
func (g *MphMatcherGroup) writeRecords() []uint32 {
|
||||||
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
|
g.multi = false
|
||||||
for _, bucketIdx := range bucketIdxs {
|
if len(g.entries) > 0 {
|
||||||
bucket := buckets[bucketIdx]
|
g.single = g.entries[0].value
|
||||||
hashedBucket = hashedBucket[:0]
|
for _, e := range g.entries {
|
||||||
seed := uint32(0)
|
if e.value != g.single {
|
||||||
for len(hashedBucket) != len(bucket) {
|
g.multi = true
|
||||||
for _, ruleIdx := range bucket {
|
break
|
||||||
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
|
|
||||||
occupied[hash] = false
|
|
||||||
g.level1[hash] = 0
|
|
||||||
}
|
|
||||||
hashedBucket = hashedBucket[:0]
|
|
||||||
seed++ // Try next seed
|
|
||||||
break
|
|
||||||
}
|
|
||||||
occupied[memHash] = true
|
|
||||||
g.level1[memHash] = ruleIdx // The final value in the hash table
|
|
||||||
hashedBucket = append(hashedBucket, memHash)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
g.level0[bucketIdx] = seed // Displacement value for this bucket
|
}
|
||||||
|
// Equal patterns become neighbours in Add order, so their values keep their priority
|
||||||
|
order := make([]uint32, len(g.entries))
|
||||||
|
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
|
||||||
|
if g.multi {
|
||||||
|
size += 3 * len(g.entries)
|
||||||
|
}
|
||||||
|
arena := make([]byte, 0, size)
|
||||||
|
recs := make([]uint32, 0, len(order))
|
||||||
|
var vals [len(mphKinds)][]uint32
|
||||||
|
for i := 0; i < len(order); {
|
||||||
|
k := g.key(order[i])
|
||||||
|
for t := range vals {
|
||||||
|
vals[t] = vals[t][:0]
|
||||||
|
}
|
||||||
|
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
|
||||||
|
e := &g.entries[order[i]]
|
||||||
|
if !slices.Contains(vals[e.kind], e.value) {
|
||||||
|
vals[e.kind] = append(vals[e.kind], e.value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rec := uint32(len(arena))
|
||||||
|
if len(k) < 255 {
|
||||||
|
arena = append(arena, byte(len(k)))
|
||||||
|
} else {
|
||||||
|
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
|
||||||
|
}
|
||||||
|
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))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
recs = append(recs, rec)
|
||||||
|
}
|
||||||
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
|
||||||
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
func mphHash(mul uint64, s string) uint64 {
|
||||||
i0 := rollingHash & g.level0Mask
|
h := uint64(0)
|
||||||
seed := g.level0[i0]
|
for i := len(s) - 1; i >= 0; i-- {
|
||||||
i1 := MemHash(seed, input) & g.level1Mask
|
h = h*mul + uint64(s[i])
|
||||||
if n := g.level1[i1]; g.rules[n] == input {
|
}
|
||||||
return n
|
return h
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
// Match implements MatcherGroup.Match.
|
// appendValues appends the values of record e for the flags in want, in mphKinds order.
|
||||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
|
||||||
matches := make([][]uint32, 0, 5)
|
if !g.multi {
|
||||||
hash := uint32(0)
|
for _, flag := range mphKinds {
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
if e&want&flag != 0 {
|
||||||
hash = hash*PrimeRK + uint32(input[i])
|
dst = append(dst, g.single)
|
||||||
if input[i] == '.' {
|
}
|
||||||
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
}
|
||||||
matches = append(matches, g.values[mphIdx])
|
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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
return dst
|
||||||
matches = append(matches, g.values[mphIdx])
|
}
|
||||||
|
|
||||||
|
// 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 {
|
||||||
|
var stack [8]uint32
|
||||||
|
parents := stack[:0] // TLD side first
|
||||||
|
h, mul := uint64(0), g.mul
|
||||||
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
|
if input[i] == '.' {
|
||||||
|
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
|
||||||
|
parents = append(parents, e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
h = h*mul + uint64(input[i])
|
||||||
}
|
}
|
||||||
return CompositeMatchesReverse(matches)
|
exact := g.lookup(h, input)
|
||||||
|
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
|
||||||
|
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 {
|
||||||
hash := uint32(0)
|
h, mul := uint64(0), g.mul
|
||||||
|
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] == '.' {
|
||||||
if g.Lookup(hash, input[i:]) != 0 {
|
dst = append(dst, mphSuffix{h, i + 1})
|
||||||
return true
|
}
|
||||||
}
|
h = h*mul + uint64(input[i])
|
||||||
|
}
|
||||||
|
return dst, h
|
||||||
|
}
|
||||||
|
|
||||||
|
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
|
||||||
|
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||||
|
if g.mul != mul {
|
||||||
|
return g.MatchAny(input) // built with a later multiplier after a collision
|
||||||
|
}
|
||||||
|
for _, p := range parents {
|
||||||
|
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
|
||||||
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return g.Lookup(hash, input) != 0
|
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func nextPow2(v int) int {
|
|
||||||
if v <= 1 {
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
const MaxUInt = ^uint(0)
|
|
||||||
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
|
|
||||||
return int(n)
|
|
||||||
}
|
|
||||||
|
|
||||||
//go:noescape
|
|
||||||
//go:linkname strhash runtime.strhash
|
|
||||||
func strhash(p unsafe.Pointer, h uintptr) uintptr
|
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
package strmatcher
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMphMatcherGroupHashCollision(t *testing.T) {
|
||||||
|
saved := mphMultipliers
|
||||||
|
defer func() { mphMultipliers = saved }()
|
||||||
|
|
||||||
|
mphMultipliers[0] = 1 // anagrams collide
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("ab.com"), 1)
|
||||||
|
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
|
||||||
|
g.AddDomainMatcher(DomainMatcher("com"), 3)
|
||||||
|
if err := g.Build(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if g.mul != saved[1] {
|
||||||
|
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
|
||||||
|
}
|
||||||
|
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
|
||||||
|
if m := g.Match(input); !slices.Equal(m, want) {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", input, m, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
|
||||||
|
mphMultipliers = saved
|
||||||
|
a, b := make([]byte, 2048), make([]byte, 2048)
|
||||||
|
for i := range a {
|
||||||
|
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
|
||||||
|
}
|
||||||
|
g = NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher(a), 1)
|
||||||
|
g.AddFullMatcher(FullMatcher(b), 1)
|
||||||
|
if err := g.Build(); err != errMphCollision {
|
||||||
|
t.Errorf("Build() = %v, want %v", err, errMphCollision)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func bitsOnes(i int) int {
|
||||||
|
n := 0
|
||||||
|
for ; i > 0; i &= i - 1 {
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphValueMatcherCombiner(t *testing.T) {
|
||||||
|
build := func(matchers ...Matcher) *MphValueMatcher {
|
||||||
|
m := NewMphValueMatcher()
|
||||||
|
for _, x := range matchers {
|
||||||
|
m.Add(x, 0)
|
||||||
|
}
|
||||||
|
if err := m.Build(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
regex, err := Regex.New(`^a\d+\.net$`)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
saved := mphMultipliers
|
||||||
|
t.Cleanup(func() { mphMultipliers = saved })
|
||||||
|
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
|
||||||
|
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
|
||||||
|
mphMultipliers = saved
|
||||||
|
if collided.mph.mul == mphMultipliers[0] {
|
||||||
|
t.Fatal("collided matcher uses the first multiplier")
|
||||||
|
}
|
||||||
|
matchers := []*MphValueMatcher{
|
||||||
|
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
|
||||||
|
collided,
|
||||||
|
build(regex, SubstrMatcher("keyword")),
|
||||||
|
build(),
|
||||||
|
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
|
||||||
|
}
|
||||||
|
var s MphValueMatcherCombiner
|
||||||
|
for i, m := range matchers {
|
||||||
|
s.Add(m, uint32(10+i))
|
||||||
|
}
|
||||||
|
inputs := []string{
|
||||||
|
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
|
||||||
|
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
|
||||||
|
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
|
||||||
|
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
|
||||||
|
}
|
||||||
|
for _, input := range inputs {
|
||||||
|
var want []uint32
|
||||||
|
for i, m := range matchers {
|
||||||
|
if m.MatchAny(input) {
|
||||||
|
want = append(want, uint32(10+i))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := s.Match(input); !slices.Equal(got, want) {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", input, got, want)
|
||||||
|
}
|
||||||
|
if got := s.MatchAny(input); got != (len(want) > 0) {
|
||||||
|
t.Errorf("MatchAny(%q) = %v", input, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
|
||||||
|
t.Errorf("MatchAny allocates %v times", n)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,7 +1,10 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"math/rand"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
@@ -276,3 +279,142 @@ func TestEmptyMphMatcherGroup(t *testing.T) {
|
|||||||
t.Error("Expect [], but ", r)
|
t.Error("Expect [], but ", r)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupRandom(t *testing.T) {
|
||||||
|
inputs := []string{""} // All strings over "ab." up to 7 bytes
|
||||||
|
for i := 0; len(inputs[i]) < 7; i++ {
|
||||||
|
for _, c := range []string{"a", "b", "."} {
|
||||||
|
inputs = append(inputs, inputs[i]+c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for seed := int64(0); seed < 300; seed++ {
|
||||||
|
r := rand.New(rand.NewSource(seed))
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
full, domain := map[string][]uint32{}, map[string][]uint32{} // Stored pattern -> values
|
||||||
|
for value := uint32(r.Intn(200)); value > 0; value-- {
|
||||||
|
pattern := make([]byte, r.Intn(8))
|
||||||
|
for i := range pattern {
|
||||||
|
pattern[i] = "ab."[r.Intn(3)]
|
||||||
|
}
|
||||||
|
if p := string(pattern); r.Intn(2) == 0 {
|
||||||
|
g.AddFullMatcher(FullMatcher(p), value)
|
||||||
|
full[p] = append(full[p], value)
|
||||||
|
} else {
|
||||||
|
g.AddDomainMatcher(DomainMatcher(p), value)
|
||||||
|
domain[p] = append(domain[p], value)
|
||||||
|
domain["."+p] = append(domain["."+p], value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
common.Must(g.Build())
|
||||||
|
for _, input := range inputs {
|
||||||
|
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
||||||
|
for i := range len(input) {
|
||||||
|
if input[i] == '.' {
|
||||||
|
keys = append(keys, input[i:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var want []uint32
|
||||||
|
for _, k := range keys {
|
||||||
|
want = append(append(want, full[k]...), domain[k]...)
|
||||||
|
}
|
||||||
|
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
|
||||||
|
// from want for patterns and inputs with a leading dot
|
||||||
|
m := g.Match(input)
|
||||||
|
if !slices.Equal(sortedSet(m), sortedSet(want)) {
|
||||||
|
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
||||||
|
}
|
||||||
|
if m := g.MatchAny(input); m != (len(want) > 0) {
|
||||||
|
t.Fatalf("seed %d: MatchAny(%q) = %v", seed, input, m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupAppend(t *testing.T) {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||||
|
g.AddFullMatcher(FullMatcher("b.com"), 2)
|
||||||
|
g.Build()
|
||||||
|
if m := append(g.Match("a.com"), 3); !slices.Equal(m, []uint32{1, 3}) {
|
||||||
|
t.Error("expect [1 3], but ", m)
|
||||||
|
}
|
||||||
|
if m := g.Match("b.com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Error("expect [2], but ", m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortedSet(v []uint32) []uint32 {
|
||||||
|
v = slices.Clone(v)
|
||||||
|
slices.Sort(v)
|
||||||
|
return slices.Compact(v)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupLongPattern(t *testing.T) {
|
||||||
|
long := strings.Repeat("a", 300) + ".com"
|
||||||
|
for _, values := range [][4]uint32{{1, 2, 3, 4}, {7, 7, 7, 7}} {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddDomainMatcher(DomainMatcher(long), values[0])
|
||||||
|
g.AddFullMatcher(FullMatcher("x."+long), values[1])
|
||||||
|
g.AddFullMatcher(FullMatcher(long[:255]), values[2]) // the shortest pattern stored with a long length
|
||||||
|
g.AddFullMatcher(FullMatcher(long[:254]), values[3])
|
||||||
|
common.Must(g.Build())
|
||||||
|
cases := []struct {
|
||||||
|
input string
|
||||||
|
want []uint32
|
||||||
|
}{
|
||||||
|
{long, []uint32{values[0]}},
|
||||||
|
{"www." + long, []uint32{values[0]}},
|
||||||
|
{"x." + long, []uint32{values[1], values[0]}},
|
||||||
|
{long[1:], nil},
|
||||||
|
{"a" + long, nil},
|
||||||
|
{long[:255], []uint32{values[2]}},
|
||||||
|
{long[:254], []uint32{values[3]}},
|
||||||
|
{long[:256], nil},
|
||||||
|
{long[:253], nil},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if m := g.Match(c.input); !slices.Equal(m, c.want) {
|
||||||
|
t.Errorf("Match(%d bytes) = %v, want %v", len(c.input), m, c.want)
|
||||||
|
}
|
||||||
|
if m := g.MatchAny(c.input); m != (c.want != nil) {
|
||||||
|
t.Errorf("MatchAny(%d bytes) = %v", len(c.input), m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A pattern longer than 65535 bytes builds and matches: a record's length is a uvarint,
|
||||||
|
// so the only cap was the build-time length field, now widened to uint32.
|
||||||
|
huge := strings.Repeat("a", 70000)
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher(strings.Repeat("a", 65535)), 1)
|
||||||
|
g.AddDomainMatcher(DomainMatcher(huge+".com"), 2)
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 3)
|
||||||
|
common.Must(g.Build())
|
||||||
|
if !g.MatchAny(strings.Repeat("a", 65535)) || g.MatchAny(strings.Repeat("a", 65534)) {
|
||||||
|
t.Error("wrong answer for a 65535-byte pattern")
|
||||||
|
}
|
||||||
|
if m := g.Match(huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Errorf("Match(%d-byte input) = %v, want [2]", len(huge)+4, m)
|
||||||
|
}
|
||||||
|
if m := g.Match("x." + huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Errorf("Match(subdomain of a %d-byte pattern) = %v, want [2]", len(huge)+4, m)
|
||||||
|
}
|
||||||
|
if g.MatchAny(huge) { // the 70000-byte label on its own is not a rule
|
||||||
|
t.Error("unexpected match for the bare 70000-byte label")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupBuildOnce(t *testing.T) {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||||
|
common.Must(g.Build())
|
||||||
|
if err := g.Build(); err == nil || !g.MatchAny("a.com") {
|
||||||
|
t.Errorf("second Build() = %v, MatchAny(a.com) = %v", err, g.MatchAny("a.com"))
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if recover() == nil {
|
||||||
|
t.Error("Add after Build did not panic")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
g.AddDomainMatcher(DomainMatcher("b.com"), 2)
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,9 +2,12 @@ package strmatcher
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
|
"math/bits"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"regexp/syntax"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
"unicode"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
|
|
||||||
"golang.org/x/net/idna"
|
"golang.org/x/net/idna"
|
||||||
@@ -73,7 +76,274 @@ func (m SubstrMatcher) Match(s string) bool {
|
|||||||
|
|
||||||
// RegexMatcher is an implementation of Matcher.
|
// RegexMatcher is an implementation of Matcher.
|
||||||
type RegexMatcher struct {
|
type RegexMatcher struct {
|
||||||
pattern *regexp.Regexp
|
pattern *regexp.Regexp
|
||||||
|
literals []string // every match contains all of them, longest first
|
||||||
|
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
|
||||||
|
rest *byteSet // the bytes it can have further before, nil if any
|
||||||
|
}
|
||||||
|
|
||||||
|
func newRegexMatcher(pattern string) (Matcher, error) {
|
||||||
|
regex, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
m := &RegexMatcher{pattern: regex}
|
||||||
|
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
||||||
|
m.literals = requiredLiterals(re, nil)
|
||||||
|
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
||||||
|
m.tail, m.rest = tailGuard(re)
|
||||||
|
}
|
||||||
|
return m, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
|
||||||
|
type byteSet [4]uint32
|
||||||
|
|
||||||
|
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
|
||||||
|
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
|
||||||
|
func (s *byteSet) or(t *byteSet) {
|
||||||
|
for i := range s {
|
||||||
|
s[i] |= t[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
|
||||||
|
|
||||||
|
// tailLen is how many positions before the end of the input tailGuard tells apart.
|
||||||
|
const tailLen = 8
|
||||||
|
|
||||||
|
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
|
||||||
|
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
|
||||||
|
// its guard.
|
||||||
|
const tailBudget = 100000
|
||||||
|
|
||||||
|
// tailWalk is a set of positions in the input, counted in bytes before its end.
|
||||||
|
type tailWalk struct {
|
||||||
|
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
|
||||||
|
far bool // tailLen or more bytes before the end
|
||||||
|
free bool // not tied to the end of the input yet
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w tailWalk) union(v tailWalk) tailWalk {
|
||||||
|
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
|
||||||
|
}
|
||||||
|
|
||||||
|
type tailBuilder struct {
|
||||||
|
tail [tailLen]byteSet
|
||||||
|
rest byteSet
|
||||||
|
void bool
|
||||||
|
work int
|
||||||
|
}
|
||||||
|
|
||||||
|
// tailGuard walks re backwards from the end of the input and collects the bytes an input
|
||||||
|
// matching re can have at each position before its end. It returns nil, nil when a branch
|
||||||
|
// of re does not end with $ or when nested repeats push the walk past tailBudget.
|
||||||
|
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
|
||||||
|
var b tailBuilder
|
||||||
|
w := b.walk(re, tailWalk{free: true})
|
||||||
|
b.stop(w)
|
||||||
|
if b.void {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
if w.at != 0 { // a match can start here, so any bytes can come before
|
||||||
|
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
|
||||||
|
b.tail[i] = allBytes
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if w.at != 0 || w.far {
|
||||||
|
b.rest = allBytes
|
||||||
|
}
|
||||||
|
n := tailLen
|
||||||
|
for n > 0 && b.tail[n-1] == b.rest {
|
||||||
|
n--
|
||||||
|
}
|
||||||
|
var tail []byteSet
|
||||||
|
if n > 0 {
|
||||||
|
tail = slices.Clone(b.tail[:n])
|
||||||
|
}
|
||||||
|
if b.rest != allBytes {
|
||||||
|
rest := b.rest
|
||||||
|
return tail, &rest
|
||||||
|
}
|
||||||
|
return tail, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
|
||||||
|
func (b *tailBuilder) stop(w tailWalk) {
|
||||||
|
if w.free {
|
||||||
|
b.void = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
|
||||||
|
if w == (tailWalk{}) || b.void {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpNoMatch:
|
||||||
|
return tailWalk{}
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
for i := len(re.Rune) - 1; i >= 0; i-- {
|
||||||
|
var set byteSet
|
||||||
|
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
|
||||||
|
if re.Flags&syntax.FoldCase != 0 {
|
||||||
|
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
|
||||||
|
set.add(byte(min(f, utf8.RuneSelf)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
w = b.step(w, &set)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
case syntax.OpCharClass:
|
||||||
|
var set byteSet
|
||||||
|
for i := 0; i+1 < len(re.Rune); i += 2 {
|
||||||
|
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
|
||||||
|
set.add(byte(r))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return b.step(w, &set)
|
||||||
|
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
|
||||||
|
return b.step(w, &allBytes)
|
||||||
|
case syntax.OpBeginText: // nothing comes before
|
||||||
|
b.stop(w)
|
||||||
|
return tailWalk{}
|
||||||
|
case syntax.OpEndText:
|
||||||
|
out := tailWalk{at: w.at & 1}
|
||||||
|
if w.free {
|
||||||
|
out.at = 1
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case syntax.OpCapture:
|
||||||
|
return b.walk(re.Sub[0], w)
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for i := len(re.Sub) - 1; i >= 0; i-- {
|
||||||
|
w = b.walk(re.Sub[i], w)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
case syntax.OpAlternate:
|
||||||
|
var out tailWalk
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
out = out.union(b.walk(sub, w))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case syntax.OpQuest:
|
||||||
|
return b.repeat(re.Sub[0], w, 1)
|
||||||
|
case syntax.OpStar:
|
||||||
|
return b.repeat(re.Sub[0], w, -1)
|
||||||
|
case syntax.OpPlus:
|
||||||
|
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
for i := 0; i < re.Min; i++ {
|
||||||
|
if b.charge() {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
w = b.walk(re.Sub[0], w)
|
||||||
|
}
|
||||||
|
if re.Max < 0 {
|
||||||
|
return b.repeat(re.Sub[0], w, -1)
|
||||||
|
}
|
||||||
|
return b.repeat(re.Sub[0], w, re.Max-re.Min)
|
||||||
|
}
|
||||||
|
return w // empty match, line and word boundaries: no constraint
|
||||||
|
}
|
||||||
|
|
||||||
|
// charge counts one repetition step and reports whether the walk has run out of budget. Only
|
||||||
|
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
|
||||||
|
// leaving a single linear pass, of any length, free.
|
||||||
|
func (b *tailBuilder) charge() bool {
|
||||||
|
b.work++
|
||||||
|
if b.work > tailBudget {
|
||||||
|
b.void = true
|
||||||
|
}
|
||||||
|
return b.void
|
||||||
|
}
|
||||||
|
|
||||||
|
// repeat walks back over up to n more repetitions of re, any number if n < 0.
|
||||||
|
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
|
||||||
|
for ; n != 0; n-- {
|
||||||
|
if b.charge() {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
next := w.union(b.walk(re, w))
|
||||||
|
if next == w {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
w = next
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
// step walks back over one character whose last byte is in set. A character that can be
|
||||||
|
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
|
||||||
|
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
|
||||||
|
out := tailWalk{far: w.far, free: w.free}
|
||||||
|
if w.far {
|
||||||
|
b.rest.or(set)
|
||||||
|
}
|
||||||
|
width := 1
|
||||||
|
if set.has(0x80) {
|
||||||
|
width = utf8.UTFMax
|
||||||
|
}
|
||||||
|
for i := 0; i < tailLen; i++ {
|
||||||
|
if w.at&(1<<i) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
b.tail[i].or(set)
|
||||||
|
for n := 1; n <= width; n++ {
|
||||||
|
if j := i + n; j < tailLen {
|
||||||
|
out.at |= 1 << j
|
||||||
|
if n < width {
|
||||||
|
b.tail[j].add(0x80)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
out.far = true
|
||||||
|
if n < width {
|
||||||
|
b.rest.add(0x80)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// mayMatch reports whether s passes the tail guard.
|
||||||
|
func (m *RegexMatcher) mayMatch(s string) bool {
|
||||||
|
n := len(s)
|
||||||
|
if m.rest == nil {
|
||||||
|
n = min(n, len(m.tail))
|
||||||
|
}
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
set := m.rest
|
||||||
|
if i < len(m.tail) {
|
||||||
|
set = &m.tail[i]
|
||||||
|
}
|
||||||
|
if !set.has(s[len(s)-1-i]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
||||||
|
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
// regexp matches U+FFFD against invalid UTF-8 bytes, strings.Contains does not
|
||||||
|
if re.Flags&syntax.FoldCase == 0 && !slices.Contains(re.Rune, utf8.RuneError) {
|
||||||
|
dst = append(dst, string(re.Rune))
|
||||||
|
}
|
||||||
|
case syntax.OpCapture, syntax.OpPlus:
|
||||||
|
dst = requiredLiterals(re.Sub[0], dst)
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
if re.Min > 0 {
|
||||||
|
dst = requiredLiterals(re.Sub[0], dst)
|
||||||
|
}
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
dst = requiredLiterals(sub, dst)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return dst
|
||||||
}
|
}
|
||||||
|
|
||||||
func (*RegexMatcher) Type() Type {
|
func (*RegexMatcher) Type() Type {
|
||||||
@@ -89,6 +359,14 @@ func (m *RegexMatcher) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (m *RegexMatcher) Match(s string) bool {
|
func (m *RegexMatcher) Match(s string) bool {
|
||||||
|
if !m.mayMatch(s) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, l := range m.literals {
|
||||||
|
if !strings.Contains(s, l) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
return m.pattern.MatchString(s)
|
return m.pattern.MatchString(s)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -102,11 +380,7 @@ func (t Type) New(pattern string) (Matcher, error) {
|
|||||||
case Domain:
|
case Domain:
|
||||||
return DomainMatcher(pattern), nil
|
return DomainMatcher(pattern), nil
|
||||||
case Regex: // 1. regex matching is case-sensitive
|
case Regex: // 1. regex matching is case-sensitive
|
||||||
regex, err := regexp.Compile(pattern)
|
return newRegexMatcher(pattern)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &RegexMatcher{pattern: regex}, nil
|
|
||||||
default:
|
default:
|
||||||
return nil, errors.New("unknown matcher type")
|
return nil, errors.New("unknown matcher type")
|
||||||
}
|
}
|
||||||
@@ -135,11 +409,7 @@ func (t Type) NewDomainPattern(pattern string) (Matcher, error) {
|
|||||||
}
|
}
|
||||||
return DomainMatcher(pattern), nil
|
return DomainMatcher(pattern), nil
|
||||||
case Regex: // Regex's charset not in LDH subset
|
case Regex: // Regex's charset not in LDH subset
|
||||||
regex, err := regexp.Compile(pattern)
|
return newRegexMatcher(pattern)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &RegexMatcher{pattern: regex}, nil
|
|
||||||
default:
|
default:
|
||||||
return nil, errors.New("unknown matcher type")
|
return nil, errors.New("unknown matcher type")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
package strmatcher
|
||||||
|
|
||||||
|
import (
|
||||||
|
"hash/fnv"
|
||||||
|
"math/rand/v2"
|
||||||
|
"regexp"
|
||||||
|
"regexp/syntax"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"unicode"
|
||||||
|
"unicode/utf8"
|
||||||
|
)
|
||||||
|
|
||||||
|
var regexLiteralCases = []struct {
|
||||||
|
pattern string
|
||||||
|
literals []string
|
||||||
|
}{
|
||||||
|
{`(^|\.)91porn\.(best|com)$`, []string{"91porn."}},
|
||||||
|
{`.+\.awsdns-cn-[0-9][0-9]\.(biz|com|net|top)$`, []string{".awsdns-cn-", "."}},
|
||||||
|
{`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`, []string{".googlevideo.com", "sn-", "r"}},
|
||||||
|
{`(?i)abc`, nil},
|
||||||
|
{`ab(?i:CD)ef`, []string{"ab", "ef"}},
|
||||||
|
{`(abc)?x`, []string{"x"}},
|
||||||
|
{`(abc)*x`, []string{"x"}},
|
||||||
|
{`x{0,3}yy`, []string{"yy"}},
|
||||||
|
{`(ab)+c{2}`, []string{"ab", "c"}},
|
||||||
|
{`abc|abd`, []string{"ab"}},
|
||||||
|
{`\Qa.b\E`, []string{"a.b"}},
|
||||||
|
{`a\x{FFFD}b`, nil},
|
||||||
|
{`^[^.]+$`, nil},
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexRequiredLiterals(t *testing.T) {
|
||||||
|
for _, test := range regexLiteralCases {
|
||||||
|
m, err := newRegexMatcher(test.pattern)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := m.(*RegexMatcher).literals; !slices.Equal(got, test.literals) {
|
||||||
|
t.Errorf("%s: got %q, want %q", test.pattern, got, test.literals)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var regexTailCases = []struct {
|
||||||
|
pattern string
|
||||||
|
guard bool
|
||||||
|
match []string // inputs the pattern matches
|
||||||
|
reject []string // inputs the tail guard alone rejects
|
||||||
|
}{
|
||||||
|
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
|
||||||
|
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
|
||||||
|
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
|
||||||
|
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
|
||||||
|
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
|
||||||
|
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
|
||||||
|
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
|
||||||
|
{`^$`, true, []string{""}, []string{"a"}},
|
||||||
|
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
|
||||||
|
{`abc`, false, []string{"abc", "xabcx"}, nil},
|
||||||
|
{`^ab`, false, []string{"ab", "abc"}, nil},
|
||||||
|
{`a$|b`, false, []string{"a", "bx"}, nil},
|
||||||
|
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
|
||||||
|
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexTailGuard(t *testing.T) {
|
||||||
|
for _, test := range regexTailCases {
|
||||||
|
m, err := newRegexMatcher(test.pattern)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rm := m.(*RegexMatcher)
|
||||||
|
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
|
||||||
|
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
|
||||||
|
}
|
||||||
|
for _, s := range test.match {
|
||||||
|
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||||
|
t.Errorf("%s: %q does not match", test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range test.reject {
|
||||||
|
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
|
||||||
|
t.Errorf("%s: %q passes the guard", test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
|
||||||
|
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
|
||||||
|
// names, however large, is walked once and guarded; its guard is checked against regexp.
|
||||||
|
func TestRegexTailGuardFlatAlternation(t *testing.T) {
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("(?:")
|
||||||
|
for i := 0; i < 20000; i++ {
|
||||||
|
if i > 0 {
|
||||||
|
sb.WriteByte('|')
|
||||||
|
}
|
||||||
|
sb.WriteString("name")
|
||||||
|
sb.WriteString(strconv.Itoa(i))
|
||||||
|
}
|
||||||
|
sb.WriteString(`)\.example\.com$`)
|
||||||
|
m, err := newRegexMatcher(sb.String())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rm := m.(*RegexMatcher)
|
||||||
|
if rm.tail == nil && rm.rest == nil {
|
||||||
|
t.Fatal("flat alternation of 20000 names lost its guard")
|
||||||
|
}
|
||||||
|
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
|
||||||
|
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||||
|
t.Errorf("%q should match", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
|
||||||
|
if rm.pattern.MatchString(s) {
|
||||||
|
t.Fatalf("test bug: %q matches the pattern", s)
|
||||||
|
}
|
||||||
|
if rm.mayMatch(s) {
|
||||||
|
t.Errorf("%q should be rejected by the guard", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
|
||||||
|
// budget, which it spends one per call so that nested repeats stay cheap.
|
||||||
|
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
|
||||||
|
if *budget <= 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
*budget--
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
for _, r := range re.Rune {
|
||||||
|
if re.Flags&syntax.FoldCase != 0 {
|
||||||
|
for n := rnd.IntN(4); n > 0; n-- {
|
||||||
|
r = unicode.SimpleFold(r)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sampleRune(sb, r, rnd)
|
||||||
|
}
|
||||||
|
case syntax.OpCharClass:
|
||||||
|
if len(re.Rune) > 0 {
|
||||||
|
i := rnd.IntN(len(re.Rune)/2) * 2
|
||||||
|
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
|
||||||
|
}
|
||||||
|
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
|
||||||
|
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
|
||||||
|
case syntax.OpCapture:
|
||||||
|
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
sampleMatch(sb, sub, rnd, budget)
|
||||||
|
}
|
||||||
|
case syntax.OpAlternate:
|
||||||
|
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
|
||||||
|
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
|
||||||
|
lo, hi := 0, 3
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpQuest:
|
||||||
|
hi = 1
|
||||||
|
case syntax.OpPlus:
|
||||||
|
lo = 1
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
lo, hi = re.Min, re.Min+3
|
||||||
|
if re.Max >= 0 {
|
||||||
|
hi = min(hi, re.Max)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
|
||||||
|
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
|
||||||
|
if r == utf8.RuneError && rnd.IntN(2) == 0 {
|
||||||
|
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sb.WriteRune(r)
|
||||||
|
}
|
||||||
|
|
||||||
|
func FuzzRegexMatcher(f *testing.F) {
|
||||||
|
inputs := []string{
|
||||||
|
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
||||||
|
"www.91porn.com", "ns1.awsdns-cn-01.top", "r1---sn-2x3abcde.googlevideo.com",
|
||||||
|
}
|
||||||
|
for _, test := range regexLiteralCases {
|
||||||
|
for _, s := range inputs {
|
||||||
|
f.Add(test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, test := range regexTailCases {
|
||||||
|
for _, s := range append(test.match, test.reject...) {
|
||||||
|
f.Add(test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
f.Fuzz(func(t *testing.T, pattern, s string) {
|
||||||
|
re, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
m, _ := newRegexMatcher(pattern)
|
||||||
|
check := func(s string) {
|
||||||
|
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||||
|
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
check(s)
|
||||||
|
// random inputs seldom match, so also try strings built from the pattern
|
||||||
|
parsed, _ := syntax.Parse(pattern, syntax.Perl)
|
||||||
|
h := fnv.New64a()
|
||||||
|
h.Write([]byte(s))
|
||||||
|
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
|
||||||
|
for range 8 {
|
||||||
|
var sb strings.Builder
|
||||||
|
budget := 256
|
||||||
|
sampleMatch(&sb, parsed, rnd, &budget)
|
||||||
|
sample := sb.String()
|
||||||
|
check(sample)
|
||||||
|
check(s + sample)
|
||||||
|
if len(sample) > 0 && len(s) > 0 {
|
||||||
|
i := rnd.IntN(len(sample))
|
||||||
|
check(sample[:i] + s[:1] + sample[i+1:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -46,7 +46,9 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
|
|||||||
func (g *MphValueMatcher) Build() error {
|
func (g *MphValueMatcher) Build() error {
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
g.mph.Build()
|
if err := g.mph.Build(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
@@ -58,23 +60,17 @@ func (g *MphValueMatcher) Build() error {
|
|||||||
|
|
||||||
// Match implements ValueMatcher.Match.
|
// Match implements ValueMatcher.Match.
|
||||||
func (g *MphValueMatcher) Match(input string) []uint32 {
|
func (g *MphValueMatcher) Match(input string) []uint32 {
|
||||||
result := make([][]uint32, 0, 5)
|
var result []uint32
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
result = append(result, g.ac.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.regex != nil {
|
if g.regex != nil {
|
||||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
result = append(result, g.regex.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return CompositeMatches(result)
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements ValueMatcher.MatchAny.
|
// MatchAny implements ValueMatcher.MatchAny.
|
||||||
@@ -87,3 +83,62 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
|
|||||||
}
|
}
|
||||||
return g.regex != nil && g.regex.MatchAny(input)
|
return g.regex != nil && g.regex.MatchAny(input)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||||
|
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if g.ac != nil && g.ac.MatchAny(input) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return g.regex != nil && g.regex.MatchAny(input)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
|
||||||
|
// against them as their MatchAny would, hashing the input once for all of them.
|
||||||
|
type MphValueMatcherCombiner struct {
|
||||||
|
matchers []*MphValueMatcher
|
||||||
|
values []uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add adds a built matcher that stands for value.
|
||||||
|
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
|
||||||
|
s.matchers = append(s.matchers, m)
|
||||||
|
s.values = append(s.values, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match returns the values of the matchers that match input, in Add order.
|
||||||
|
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
|
||||||
|
if len(s.matchers) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var stack [16]mphSuffix
|
||||||
|
mul := mphMultipliers[0]
|
||||||
|
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||||
|
var result []uint32
|
||||||
|
for i, m := range s.matchers {
|
||||||
|
if m.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
result = append(result, s.values[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// MatchAny returns true as soon as one matcher matches input.
|
||||||
|
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
|
||||||
|
switch len(s.matchers) {
|
||||||
|
case 0:
|
||||||
|
return false
|
||||||
|
case 1:
|
||||||
|
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
|
||||||
|
}
|
||||||
|
var stack [16]mphSuffix
|
||||||
|
mul := mphMultipliers[0]
|
||||||
|
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||||
|
for _, m := range s.matchers {
|
||||||
|
if m.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|||||||
+21
-25
@@ -1,7 +1,7 @@
|
|||||||
package log // import "github.com/xtls/xray-core/common/log"
|
package log // import "github.com/xtls/xray-core/common/log"
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"sync"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
)
|
)
|
||||||
@@ -29,36 +29,32 @@ func (m *GeneralMessage) String() string {
|
|||||||
|
|
||||||
// Record writes a message into log stream.
|
// Record writes a message into log stream.
|
||||||
func Record(msg Message) {
|
func Record(msg Message) {
|
||||||
logHandler.Handle(msg)
|
if h := logHandler.Load(); h != nil {
|
||||||
|
(*h).Handle(msg)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var logHandler syncHandler
|
type SeverityLogger interface {
|
||||||
|
Handler
|
||||||
|
Severity() Severity
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetSeverity() Severity {
|
||||||
|
if h := logHandler.Load(); h != nil {
|
||||||
|
if sh, ok := (*h).(SeverityLogger); ok {
|
||||||
|
return sh.Severity()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// log everything by default
|
||||||
|
return Severity_Debug
|
||||||
|
}
|
||||||
|
|
||||||
|
var logHandler atomic.Pointer[Handler]
|
||||||
|
|
||||||
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
||||||
func RegisterHandler(handler Handler) {
|
func RegisterHandler(handler Handler) {
|
||||||
if handler == nil {
|
if handler == nil {
|
||||||
panic("Log handler is nil")
|
panic("Log handler is nil")
|
||||||
}
|
}
|
||||||
logHandler.Set(handler)
|
logHandler.Store(&handler)
|
||||||
}
|
|
||||||
|
|
||||||
type syncHandler struct {
|
|
||||||
sync.RWMutex
|
|
||||||
Handler
|
|
||||||
}
|
|
||||||
|
|
||||||
func (h *syncHandler) Handle(msg Message) {
|
|
||||||
h.RLock()
|
|
||||||
defer h.RUnlock()
|
|
||||||
|
|
||||||
if h.Handler != nil {
|
|
||||||
h.Handler.Handle(msg)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (h *syncHandler) Set(handler Handler) {
|
|
||||||
h.Lock()
|
|
||||||
defer h.Unlock()
|
|
||||||
|
|
||||||
h.Handler = handler
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -68,6 +68,10 @@ func (l *serverityLogger) Handle(msg Message) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (l *serverityLogger) Severity() Severity {
|
||||||
|
return l.logLevel
|
||||||
|
}
|
||||||
|
|
||||||
func (l *generalLogger) run() {
|
func (l *generalLogger) run() {
|
||||||
defer l.access.Signal()
|
defer l.access.Signal()
|
||||||
|
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ func (m *ClientManager) Dispatch(ctx context.Context, link *transport.Link) erro
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return errors.New("unable to find an available mux client").AtWarning()
|
return errors.New("unable to find an available mux client")
|
||||||
}
|
}
|
||||||
|
|
||||||
type WorkerPicker interface {
|
type WorkerPicker interface {
|
||||||
|
|||||||
+1
-1
@@ -117,7 +117,7 @@ func (f *FrameMetadata) Unmarshal(reader io.Reader, readSourceAndLocal bool) err
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if metaLen > 512 {
|
if metaLen > 512 {
|
||||||
return errors.New("invalid metalen ", metaLen).AtError()
|
return errors.New("invalid metalen ", metaLen)
|
||||||
}
|
}
|
||||||
|
|
||||||
b := buf.New()
|
b := buf.New()
|
||||||
|
|||||||
@@ -351,7 +351,7 @@ func (w *ServerWorker) handleFrame(ctx context.Context, reader *buf.BufferedRead
|
|||||||
err = w.handleStatusKeep(&meta, reader)
|
err = w.handleStatusKeep(&meta, reader)
|
||||||
default:
|
default:
|
||||||
status := meta.SessionStatus
|
status := meta.SessionStatus
|
||||||
return errors.New("unknown status: ", status).AtError()
|
return errors.New("unknown status: ", status)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,350 @@
|
|||||||
|
//go:build darwin && !ios
|
||||||
|
|
||||||
|
package net
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
darwinProcPIDListFDs = 1
|
||||||
|
darwinProcPIDFDSocketInfo = 3
|
||||||
|
darwinProcFDTypeSocket = 2
|
||||||
|
darwinProcFDInfoSize = 8
|
||||||
|
darwinSocketFDInfoSize = 792
|
||||||
|
darwinSocketFDInfoPSIOff = 24
|
||||||
|
darwinSocketInfoProtoOff = darwinSocketFDInfoPSIOff + 156
|
||||||
|
darwinSocketInfoFamilyOff = darwinSocketFDInfoPSIOff + 160
|
||||||
|
darwinSocketInfoKindOff = darwinSocketFDInfoPSIOff + 232
|
||||||
|
darwinSocketInfoInSockOff = darwinSocketFDInfoPSIOff + 240
|
||||||
|
darwinInSockInfoFPortOff = darwinSocketInfoInSockOff
|
||||||
|
darwinInSockInfoLPortOff = darwinSocketInfoInSockOff + 4
|
||||||
|
darwinInSockInfoVFlagOff = darwinSocketInfoInSockOff + 24
|
||||||
|
darwinInSockInfoFAddrOff = darwinSocketInfoInSockOff + 32
|
||||||
|
darwinInSockInfoLAddrOff = darwinSocketInfoInSockOff + 48
|
||||||
|
darwinInSockInfoSize = 80
|
||||||
|
darwinInSockInfoIPv4 = 0x1
|
||||||
|
darwinInSockInfoIPv6 = 0x2
|
||||||
|
darwinSockInfoIN = 1
|
||||||
|
darwinSockInfoTCP = 2
|
||||||
|
)
|
||||||
|
|
||||||
|
type darwinSocketMatchLevel int
|
||||||
|
|
||||||
|
const (
|
||||||
|
darwinSocketNoMatch darwinSocketMatchLevel = iota
|
||||||
|
darwinSocketPortMatch
|
||||||
|
darwinSocketRemoteMatch
|
||||||
|
darwinSocketLocalMatch
|
||||||
|
darwinSocketExactMatch
|
||||||
|
)
|
||||||
|
|
||||||
|
func FindProcess(network, srcIP string, srcPort uint16, destIP string, destPort uint16) (PID int, Name string, AbsolutePath string, err error) {
|
||||||
|
isLocal, err := IsLocal(net.ParseIP(srcIP))
|
||||||
|
if err != nil {
|
||||||
|
return 0, "", "", errors.New("failed to determine if address is local: ", err)
|
||||||
|
}
|
||||||
|
if !isLocal {
|
||||||
|
return 0, "", "", ErrNotLocal
|
||||||
|
}
|
||||||
|
if network != "tcp" && network != "udp" {
|
||||||
|
panic("Unsupported network type for process lookup.")
|
||||||
|
}
|
||||||
|
|
||||||
|
srcAddr, err := netip.ParseAddr(srcIP)
|
||||||
|
if err != nil {
|
||||||
|
return 0, "", "", errors.New("invalid source IP address: ", srcIP)
|
||||||
|
}
|
||||||
|
srcAddr = srcAddr.Unmap()
|
||||||
|
|
||||||
|
var dstAddr netip.Addr
|
||||||
|
hasDstAddr := false
|
||||||
|
if destIP != "" && destPort != 0 {
|
||||||
|
dstAddr, err = netip.ParseAddr(destIP)
|
||||||
|
if err != nil {
|
||||||
|
return 0, "", "", errors.New("invalid destination IP address: ", destIP)
|
||||||
|
}
|
||||||
|
dstAddr = dstAddr.Unmap()
|
||||||
|
hasDstAddr = true
|
||||||
|
}
|
||||||
|
|
||||||
|
processes, err := unix.SysctlKinfoProcSlice("kern.proc.all")
|
||||||
|
if err != nil {
|
||||||
|
return 0, "", "", errors.New("failed to list processes").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var bestPID int32
|
||||||
|
bestLevel := darwinSocketNoMatch
|
||||||
|
ambiguousBest := false
|
||||||
|
|
||||||
|
for _, process := range processes {
|
||||||
|
pid := process.Proc.P_pid
|
||||||
|
if pid <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
matchLevel, err := darwinProcessSocketMatchLevel(pid, network, srcAddr, srcPort, dstAddr, destPort, hasDstAddr)
|
||||||
|
if err != nil || matchLevel == darwinSocketNoMatch {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if matchLevel == darwinSocketExactMatch {
|
||||||
|
bestPID = pid
|
||||||
|
bestLevel = matchLevel
|
||||||
|
ambiguousBest = false
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if matchLevel > bestLevel {
|
||||||
|
bestPID = pid
|
||||||
|
bestLevel = matchLevel
|
||||||
|
ambiguousBest = false
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if matchLevel == bestLevel {
|
||||||
|
ambiguousBest = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if bestLevel == darwinSocketNoMatch {
|
||||||
|
return 0, "", "", errors.New("process not found for ", network, " connection from ", srcIP, ":", srcPort, " to ", destIP, ":", destPort)
|
||||||
|
}
|
||||||
|
if ambiguousBest {
|
||||||
|
return 0, "", "", errors.New("ambiguous process match for ", network, " connection from ", srcIP, ":", srcPort, " to ", destIP, ":", destPort)
|
||||||
|
}
|
||||||
|
|
||||||
|
absPath, err := darwinProcessPath(bestPID)
|
||||||
|
if err != nil {
|
||||||
|
return 0, "", "", errors.New("could not get process path for PID ", bestPID, ": ", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
absPath = filepath.ToSlash(absPath)
|
||||||
|
return int(bestPID), filepath.Base(absPath), absPath, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinProcessSocketMatchLevel(pid int32, network string, srcAddr netip.Addr, srcPort uint16, dstAddr netip.Addr, dstPort uint16, hasDstAddr bool) (darwinSocketMatchLevel, error) {
|
||||||
|
fds, err := darwinProcessFDs(pid)
|
||||||
|
if err != nil {
|
||||||
|
return darwinSocketNoMatch, err
|
||||||
|
}
|
||||||
|
|
||||||
|
bestLevel := darwinSocketNoMatch
|
||||||
|
info := make([]byte, darwinSocketFDInfoSize)
|
||||||
|
for fd := 0; fd+darwinProcFDInfoSize <= len(fds); fd += darwinProcFDInfoSize {
|
||||||
|
fdNumber := int32(darwinReadNativeUint32(fds[fd : fd+4]))
|
||||||
|
fdType := darwinReadNativeUint32(fds[fd+4 : fd+8])
|
||||||
|
if fdType != darwinProcFDTypeSocket {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
n, err := darwinProcPIDFDInfo(pid, fdNumber, darwinProcPIDFDSocketInfo, info)
|
||||||
|
if err != nil || n < darwinSocketInfoInSockOff+darwinInSockInfoSize {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
level := darwinSocketInfoMatchLevel(info[:n], network, srcAddr, srcPort, dstAddr, dstPort, hasDstAddr)
|
||||||
|
if level == darwinSocketExactMatch {
|
||||||
|
return level, nil
|
||||||
|
}
|
||||||
|
if level > bestLevel {
|
||||||
|
bestLevel = level
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return bestLevel, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinProcessFDs(pid int32) ([]byte, error) {
|
||||||
|
n, err := darwinProcPIDInfo(pid, darwinProcPIDListFDs, 0, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if n <= 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
buf := make([]byte, n)
|
||||||
|
n, err = darwinProcPIDInfo(pid, darwinProcPIDListFDs, 0, buf)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return buf[:n], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinSocketInfoMatchLevel(info []byte, network string, srcAddr netip.Addr, srcPort uint16, dstAddr netip.Addr, dstPort uint16, hasDstAddr bool) darwinSocketMatchLevel {
|
||||||
|
protocol := int(darwinReadNativeUint32(info[darwinSocketInfoProtoOff : darwinSocketInfoProtoOff+4]))
|
||||||
|
family := int(darwinReadNativeUint32(info[darwinSocketInfoFamilyOff : darwinSocketInfoFamilyOff+4]))
|
||||||
|
kind := int(darwinReadNativeUint32(info[darwinSocketInfoKindOff : darwinSocketInfoKindOff+4]))
|
||||||
|
|
||||||
|
switch network {
|
||||||
|
case "tcp":
|
||||||
|
if protocol != unix.IPPROTO_TCP || kind != darwinSockInfoTCP {
|
||||||
|
return darwinSocketNoMatch
|
||||||
|
}
|
||||||
|
case "udp":
|
||||||
|
if protocol != unix.IPPROTO_UDP || kind != darwinSockInfoIN {
|
||||||
|
return darwinSocketNoMatch
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return darwinSocketNoMatch
|
||||||
|
}
|
||||||
|
|
||||||
|
vflag := info[darwinInSockInfoVFlagOff]
|
||||||
|
if srcAddr.Is4() {
|
||||||
|
// Dual-stack sockets expose IPv4-mapped connections as AF_INET6
|
||||||
|
// while marking the endpoint as IPv4 in ini_vflag.
|
||||||
|
if (family != unix.AF_INET && family != unix.AF_INET6) || vflag&darwinInSockInfoIPv4 == 0 {
|
||||||
|
return darwinSocketNoMatch
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if family != unix.AF_INET6 || vflag&darwinInSockInfoIPv6 == 0 {
|
||||||
|
return darwinSocketNoMatch
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
localPort := int32(darwinReadNativeUint32(info[darwinInSockInfoLPortOff : darwinInSockInfoLPortOff+4]))
|
||||||
|
if !darwinPortMatches(localPort, srcPort) {
|
||||||
|
return darwinSocketNoMatch
|
||||||
|
}
|
||||||
|
localAddrMatches := darwinAddrMatchesOrUnspecified(info[darwinInSockInfoLAddrOff:darwinInSockInfoLAddrOff+16], srcAddr)
|
||||||
|
|
||||||
|
foreignAddrRaw := info[darwinInSockInfoFAddrOff : darwinInSockInfoFAddrOff+16]
|
||||||
|
foreignPort := int32(darwinReadNativeUint32(info[darwinInSockInfoFPortOff : darwinInSockInfoFPortOff+4]))
|
||||||
|
|
||||||
|
if !hasDstAddr {
|
||||||
|
if localAddrMatches {
|
||||||
|
return darwinSocketExactMatch
|
||||||
|
}
|
||||||
|
return darwinSocketNoMatch
|
||||||
|
}
|
||||||
|
|
||||||
|
remoteMatches := darwinPortMatches(foreignPort, dstPort) && darwinAddrMatches(foreignAddrRaw, dstAddr)
|
||||||
|
if network == "udp" && darwinEndpointIsZero(foreignAddrRaw, foreignPort) && localAddrMatches {
|
||||||
|
return darwinSocketExactMatch
|
||||||
|
}
|
||||||
|
switch {
|
||||||
|
case localAddrMatches && remoteMatches:
|
||||||
|
return darwinSocketExactMatch
|
||||||
|
case localAddrMatches:
|
||||||
|
return darwinSocketLocalMatch
|
||||||
|
case remoteMatches:
|
||||||
|
return darwinSocketRemoteMatch
|
||||||
|
default:
|
||||||
|
return darwinSocketPortMatch
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinPortMatches(value int32, port uint16) bool {
|
||||||
|
raw := uint16(value)
|
||||||
|
return raw == port || darwinNtohs(raw) == port
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinNtohs(value uint16) uint16 {
|
||||||
|
return value<<8 | value>>8
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinAddrMatches(raw []byte, addr netip.Addr) bool {
|
||||||
|
if addr.Is4() {
|
||||||
|
ip := addr.As4()
|
||||||
|
return bytes.Equal(raw[12:16], ip[:])
|
||||||
|
}
|
||||||
|
ip := addr.As16()
|
||||||
|
return bytes.Equal(raw, ip[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinAddrMatchesOrUnspecified(raw []byte, addr netip.Addr) bool {
|
||||||
|
if darwinAddrMatches(raw, addr) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if addr.Is4() {
|
||||||
|
return darwinBytesAreZero(raw[12:16])
|
||||||
|
}
|
||||||
|
return darwinBytesAreZero(raw)
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinEndpointIsZero(rawAddr []byte, port int32) bool {
|
||||||
|
return uint32(port) == 0 && darwinBytesAreZero(rawAddr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinBytesAreZero(raw []byte) bool {
|
||||||
|
for _, value := range raw {
|
||||||
|
if value != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinReadNativeUint32(b []byte) uint32 {
|
||||||
|
return *(*uint32)(unsafe.Pointer(&b[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinProcessPath(pid int32) (string, error) {
|
||||||
|
buf := make([]byte, unix.PathMax)
|
||||||
|
n, err := darwinProcPIDPath(pid, buf)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
if n <= 0 {
|
||||||
|
return "", errors.New("empty process path")
|
||||||
|
}
|
||||||
|
return strings.TrimRight(string(buf[:n]), "\x00"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinProcPIDInfo(pid int32, flavor int, arg uint64, buf []byte) (int, error) {
|
||||||
|
var ptr unsafe.Pointer
|
||||||
|
if len(buf) > 0 {
|
||||||
|
ptr = unsafe.Pointer(&buf[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
r0, _, errno := syscall_syscall6(libc_proc_pidinfo_trampoline_addr, uintptr(pid), uintptr(flavor), uintptr(arg), uintptr(ptr), uintptr(len(buf)), 0)
|
||||||
|
if errno != 0 {
|
||||||
|
return 0, errno
|
||||||
|
}
|
||||||
|
return int(r0), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinProcPIDFDInfo(pid int32, fd int32, flavor int, buf []byte) (int, error) {
|
||||||
|
var ptr unsafe.Pointer
|
||||||
|
if len(buf) > 0 {
|
||||||
|
ptr = unsafe.Pointer(&buf[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
r0, _, errno := syscall_syscall6(libc_proc_pidfdinfo_trampoline_addr, uintptr(pid), uintptr(fd), uintptr(flavor), uintptr(ptr), uintptr(len(buf)), 0)
|
||||||
|
if errno != 0 {
|
||||||
|
return 0, errno
|
||||||
|
}
|
||||||
|
return int(r0), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func darwinProcPIDPath(pid int32, buf []byte) (int, error) {
|
||||||
|
r0, _, errno := syscall_syscall6(libc_proc_pidpath_trampoline_addr, uintptr(pid), uintptr(unsafe.Pointer(&buf[0])), uintptr(len(buf)), 0, 0, 0)
|
||||||
|
if errno != 0 {
|
||||||
|
return 0, errno
|
||||||
|
}
|
||||||
|
return int(r0), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var libc_proc_pidinfo_trampoline_addr uintptr
|
||||||
|
|
||||||
|
//go:cgo_import_dynamic libc_proc_pidinfo proc_pidinfo "/usr/lib/libproc.dylib"
|
||||||
|
|
||||||
|
var libc_proc_pidfdinfo_trampoline_addr uintptr
|
||||||
|
|
||||||
|
//go:cgo_import_dynamic libc_proc_pidfdinfo proc_pidfdinfo "/usr/lib/libproc.dylib"
|
||||||
|
|
||||||
|
var libc_proc_pidpath_trampoline_addr uintptr
|
||||||
|
|
||||||
|
//go:cgo_import_dynamic libc_proc_pidpath proc_pidpath "/usr/lib/libproc.dylib"
|
||||||
|
|
||||||
|
// Implemented in the runtime package (runtime/sys_darwin.go).
|
||||||
|
func syscall_syscall6(fn, a1, a2, a3, a4, a5, a6 uintptr) (r1, r2 uintptr, err syscall.Errno)
|
||||||
|
|
||||||
|
//go:linkname syscall_syscall6 syscall.syscall6
|
||||||
@@ -0,0 +1,18 @@
|
|||||||
|
//go:build darwin && !ios
|
||||||
|
|
||||||
|
#include "textflag.h"
|
||||||
|
|
||||||
|
TEXT libc_proc_pidinfo_trampoline<>(SB),NOSPLIT,$0-0
|
||||||
|
JMP libc_proc_pidinfo(SB)
|
||||||
|
GLOBL ·libc_proc_pidinfo_trampoline_addr(SB), RODATA, $8
|
||||||
|
DATA ·libc_proc_pidinfo_trampoline_addr(SB)/8, $libc_proc_pidinfo_trampoline<>(SB)
|
||||||
|
|
||||||
|
TEXT libc_proc_pidfdinfo_trampoline<>(SB),NOSPLIT,$0-0
|
||||||
|
JMP libc_proc_pidfdinfo(SB)
|
||||||
|
GLOBL ·libc_proc_pidfdinfo_trampoline_addr(SB), RODATA, $8
|
||||||
|
DATA ·libc_proc_pidfdinfo_trampoline_addr(SB)/8, $libc_proc_pidfdinfo_trampoline<>(SB)
|
||||||
|
|
||||||
|
TEXT libc_proc_pidpath_trampoline<>(SB),NOSPLIT,$0-0
|
||||||
|
JMP libc_proc_pidpath(SB)
|
||||||
|
GLOBL ·libc_proc_pidpath_trampoline_addr(SB), RODATA, $8
|
||||||
|
DATA ·libc_proc_pidpath_trampoline_addr(SB)/8, $libc_proc_pidpath_trampoline<>(SB)
|
||||||
@@ -0,0 +1,356 @@
|
|||||||
|
//go:build darwin && !ios
|
||||||
|
|
||||||
|
package net
|
||||||
|
|
||||||
|
import (
|
||||||
|
stdnet "net"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestFindProcessDarwinTCP(t *testing.T) {
|
||||||
|
listener, err := stdnet.Listen("tcp", "127.0.0.1:0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer listener.Close()
|
||||||
|
|
||||||
|
accepted := make(chan stdnet.Conn, 1)
|
||||||
|
go func() {
|
||||||
|
conn, err := listener.Accept()
|
||||||
|
if err == nil {
|
||||||
|
accepted <- conn
|
||||||
|
return
|
||||||
|
}
|
||||||
|
close(accepted)
|
||||||
|
}()
|
||||||
|
|
||||||
|
conn, err := stdnet.Dial("tcp", listener.Addr().String())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
serverConn := <-accepted
|
||||||
|
if serverConn == nil {
|
||||||
|
t.Fatal("server did not accept tcp connection")
|
||||||
|
}
|
||||||
|
defer serverConn.Close()
|
||||||
|
|
||||||
|
local := conn.LocalAddr().(*stdnet.TCPAddr)
|
||||||
|
remote := conn.RemoteAddr().(*stdnet.TCPAddr)
|
||||||
|
|
||||||
|
pid, name, path, err := FindProcess("tcp", local.IP.String(), uint16(local.Port), remote.IP.String(), uint16(remote.Port))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
assertCurrentProcess(t, pid, name, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindProcessDarwinTCPIPv4Mapped(t *testing.T) {
|
||||||
|
listener, err := stdnet.Listen("tcp4", "127.0.0.1:0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer listener.Close()
|
||||||
|
|
||||||
|
accepted := make(chan stdnet.Conn, 1)
|
||||||
|
go func() {
|
||||||
|
conn, err := listener.Accept()
|
||||||
|
if err == nil {
|
||||||
|
accepted <- conn
|
||||||
|
return
|
||||||
|
}
|
||||||
|
close(accepted)
|
||||||
|
}()
|
||||||
|
|
||||||
|
listenerAddr := listener.Addr().(*stdnet.TCPAddr)
|
||||||
|
fd, err := unix.Socket(unix.AF_INET6, unix.SOCK_STREAM, unix.IPPROTO_TCP)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer unix.Close(fd)
|
||||||
|
|
||||||
|
mappedAddr := [16]byte{10: 0xff, 11: 0xff, 12: 127, 15: 1}
|
||||||
|
if err := unix.Connect(fd, &unix.SockaddrInet6{
|
||||||
|
Port: listenerAddr.Port,
|
||||||
|
Addr: mappedAddr,
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
serverConn := <-accepted
|
||||||
|
if serverConn == nil {
|
||||||
|
t.Fatal("server did not accept tcp connection")
|
||||||
|
}
|
||||||
|
defer serverConn.Close()
|
||||||
|
|
||||||
|
local, err := unix.Getsockname(fd)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
localPort := local.(*unix.SockaddrInet6).Port
|
||||||
|
|
||||||
|
pid, name, path, err := FindProcess("tcp", "127.0.0.1", uint16(localPort), "127.0.0.1", uint16(listenerAddr.Port))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
assertCurrentProcess(t, pid, name, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindProcessDarwinUDP(t *testing.T) {
|
||||||
|
conn, err := stdnet.ListenUDP("udp", &stdnet.UDPAddr{IP: stdnet.ParseIP("127.0.0.1")})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
local := conn.LocalAddr().(*stdnet.UDPAddr)
|
||||||
|
|
||||||
|
pid, name, path, err := FindProcess("udp", local.IP.String(), uint16(local.Port), "", 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
assertCurrentProcess(t, pid, name, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindProcessDarwinNonLocal(t *testing.T) {
|
||||||
|
_, _, _, err := FindProcess("tcp", "203.0.113.1", 80, "", 0)
|
||||||
|
if err != ErrNotLocal {
|
||||||
|
t.Fatalf("expected ErrNotLocal, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindProcessDarwinUnsupportedNetwork(t *testing.T) {
|
||||||
|
defer func() {
|
||||||
|
if recover() == nil {
|
||||||
|
t.Fatal("expected panic")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
_, _, _, _ = FindProcess("icmp", "127.0.0.1", 0, "", 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertCurrentProcess(t *testing.T, pid int, name string, path string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if pid != os.Getpid() {
|
||||||
|
t.Fatalf("expected pid %d, got %d (%s, %s)", os.Getpid(), pid, name, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
executable, err := os.Executable()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if path == "" || name == "" {
|
||||||
|
t.Fatalf("expected process path and name, got name=%q path=%q", name, path)
|
||||||
|
}
|
||||||
|
if sameFile(executable, path) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
t.Fatalf("expected executable %q, got %q", executable, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
func sameFile(left string, right string) bool {
|
||||||
|
leftInfo, leftErr := os.Stat(left)
|
||||||
|
rightInfo, rightErr := os.Stat(right)
|
||||||
|
if leftErr != nil || rightErr != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return os.SameFile(leftInfo, rightInfo)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindProcessDarwinTCPWithoutDestination(t *testing.T) {
|
||||||
|
listener, err := stdnet.Listen("tcp", "127.0.0.1:0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer listener.Close()
|
||||||
|
|
||||||
|
accepted := make(chan stdnet.Conn, 1)
|
||||||
|
go func() {
|
||||||
|
conn, err := listener.Accept()
|
||||||
|
if err == nil {
|
||||||
|
accepted <- conn
|
||||||
|
return
|
||||||
|
}
|
||||||
|
close(accepted)
|
||||||
|
}()
|
||||||
|
|
||||||
|
conn, err := stdnet.DialTimeout("tcp", listener.Addr().String(), time.Second)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
serverConn := <-accepted
|
||||||
|
if serverConn == nil {
|
||||||
|
t.Fatal("server did not accept tcp connection")
|
||||||
|
}
|
||||||
|
defer serverConn.Close()
|
||||||
|
|
||||||
|
local := conn.LocalAddr().(*stdnet.TCPAddr)
|
||||||
|
pid, name, path, err := FindProcess("tcp", local.IP.String(), uint16(local.Port), "", 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
assertCurrentProcess(t, pid, name, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindProcessDarwinTCPWithDifferentDestination(t *testing.T) {
|
||||||
|
listener, err := stdnet.Listen("tcp", "127.0.0.1:0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer listener.Close()
|
||||||
|
|
||||||
|
accepted := make(chan stdnet.Conn, 1)
|
||||||
|
go func() {
|
||||||
|
conn, err := listener.Accept()
|
||||||
|
if err == nil {
|
||||||
|
accepted <- conn
|
||||||
|
return
|
||||||
|
}
|
||||||
|
close(accepted)
|
||||||
|
}()
|
||||||
|
|
||||||
|
conn, err := stdnet.DialTimeout("tcp", listener.Addr().String(), time.Second)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
serverConn := <-accepted
|
||||||
|
if serverConn == nil {
|
||||||
|
t.Fatal("server did not accept tcp connection")
|
||||||
|
}
|
||||||
|
defer serverConn.Close()
|
||||||
|
|
||||||
|
local := conn.LocalAddr().(*stdnet.TCPAddr)
|
||||||
|
pid, name, path, err := FindProcess("tcp", local.IP.String(), uint16(local.Port), "203.0.113.10", 443)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
assertCurrentProcess(t, pid, name, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDarwinSocketInfoMatchLevelFallbacks(t *testing.T) {
|
||||||
|
src := netip.MustParseAddr("198.18.0.2")
|
||||||
|
dst := netip.MustParseAddr("203.0.113.10")
|
||||||
|
otherLocal := netip.MustParseAddr("192.168.1.10")
|
||||||
|
otherRemote := netip.MustParseAddr("198.51.100.10")
|
||||||
|
unspecifiedLocal := netip.MustParseAddr("0.0.0.0")
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
local netip.Addr
|
||||||
|
remote netip.Addr
|
||||||
|
hasDst bool
|
||||||
|
wantLevel darwinSocketMatchLevel
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "exact",
|
||||||
|
local: src,
|
||||||
|
remote: dst,
|
||||||
|
hasDst: true,
|
||||||
|
wantLevel: darwinSocketExactMatch,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unspecified local with matching remote",
|
||||||
|
local: unspecifiedLocal,
|
||||||
|
remote: dst,
|
||||||
|
hasDst: true,
|
||||||
|
wantLevel: darwinSocketExactMatch,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unspecified local without destination",
|
||||||
|
local: unspecifiedLocal,
|
||||||
|
remote: otherRemote,
|
||||||
|
hasDst: false,
|
||||||
|
wantLevel: darwinSocketExactMatch,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "local match with different remote",
|
||||||
|
local: src,
|
||||||
|
remote: otherRemote,
|
||||||
|
hasDst: true,
|
||||||
|
wantLevel: darwinSocketLocalMatch,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "remote match with different local",
|
||||||
|
local: otherLocal,
|
||||||
|
remote: dst,
|
||||||
|
hasDst: true,
|
||||||
|
wantLevel: darwinSocketRemoteMatch,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "port only with destination",
|
||||||
|
local: otherLocal,
|
||||||
|
remote: otherRemote,
|
||||||
|
hasDst: true,
|
||||||
|
wantLevel: darwinSocketPortMatch,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "different local without destination",
|
||||||
|
local: otherLocal,
|
||||||
|
remote: otherRemote,
|
||||||
|
hasDst: false,
|
||||||
|
wantLevel: darwinSocketNoMatch,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, test := range tests {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
info := newDarwinSocketInfo("tcp", test.local, 12345, test.remote, 443)
|
||||||
|
level := darwinSocketInfoMatchLevel(info, "tcp", src, 12345, dst, 443, test.hasDst)
|
||||||
|
if level != test.wantLevel {
|
||||||
|
t.Fatalf("unexpected match level: got %d, want %d", level, test.wantLevel)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDarwinSocketInfoMatchLevelIPv4Mapped(t *testing.T) {
|
||||||
|
src := netip.MustParseAddr("127.0.0.1")
|
||||||
|
dst := netip.MustParseAddr("203.0.113.10")
|
||||||
|
info := newDarwinSocketInfo("tcp", src, 12345, dst, 443)
|
||||||
|
writeDarwinNativeUint32(info, darwinSocketInfoFamilyOff, uint32(unix.AF_INET6))
|
||||||
|
|
||||||
|
level := darwinSocketInfoMatchLevel(info, "tcp", src, 12345, dst, 443, true)
|
||||||
|
if level != darwinSocketExactMatch {
|
||||||
|
t.Fatalf("unexpected match level: got %d, want %d", level, darwinSocketExactMatch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newDarwinSocketInfo(network string, local netip.Addr, localPort uint16, remote netip.Addr, remotePort uint16) []byte {
|
||||||
|
info := make([]byte, darwinSocketFDInfoSize)
|
||||||
|
switch network {
|
||||||
|
case "tcp":
|
||||||
|
writeDarwinNativeUint32(info, darwinSocketInfoProtoOff, uint32(unix.IPPROTO_TCP))
|
||||||
|
writeDarwinNativeUint32(info, darwinSocketInfoKindOff, uint32(darwinSockInfoTCP))
|
||||||
|
case "udp":
|
||||||
|
writeDarwinNativeUint32(info, darwinSocketInfoProtoOff, uint32(unix.IPPROTO_UDP))
|
||||||
|
writeDarwinNativeUint32(info, darwinSocketInfoKindOff, uint32(darwinSockInfoIN))
|
||||||
|
}
|
||||||
|
writeDarwinNativeUint32(info, darwinSocketInfoFamilyOff, uint32(unix.AF_INET))
|
||||||
|
info[darwinInSockInfoVFlagOff] = darwinInSockInfoIPv4
|
||||||
|
writeDarwinNativeUint32(info, darwinInSockInfoLPortOff, uint32(localPort))
|
||||||
|
writeDarwinNativeUint32(info, darwinInSockInfoFPortOff, uint32(remotePort))
|
||||||
|
copyDarwinIPv4(info[darwinInSockInfoLAddrOff:darwinInSockInfoLAddrOff+16], local)
|
||||||
|
copyDarwinIPv4(info[darwinInSockInfoFAddrOff:darwinInSockInfoFAddrOff+16], remote)
|
||||||
|
return info
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeDarwinNativeUint32(b []byte, offset int, value uint32) {
|
||||||
|
*(*uint32)(unsafe.Pointer(&b[offset])) = value
|
||||||
|
}
|
||||||
|
|
||||||
|
func copyDarwinIPv4(dst []byte, addr netip.Addr) {
|
||||||
|
ip := addr.As4()
|
||||||
|
copy(dst[12:16], ip[:])
|
||||||
|
}
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
//go:build ios
|
||||||
|
|
||||||
|
package net
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
)
|
||||||
|
|
||||||
|
func FindProcess(network, srcIP string, srcPort uint16, destIP string, destPort uint16) (int, string, string, error) {
|
||||||
|
return 0, "", "", errors.New("process lookup is not supported on this platform")
|
||||||
|
}
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
//go:build !windows && !linux && !android
|
//go:build !windows && !linux && !android && !darwin
|
||||||
|
|
||||||
package net
|
package net
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,20 @@
|
|||||||
|
package net
|
||||||
|
|
||||||
|
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
|
||||||
|
type PacketConnWrapper struct {
|
||||||
|
PacketConn
|
||||||
|
Dest Addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
|
||||||
|
n, _, err := c.PacketConn.ReadFrom(p)
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
|
||||||
|
return c.PacketConn.WriteTo(p, c.Dest)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) RemoteAddr() Addr {
|
||||||
|
return c.Dest
|
||||||
|
}
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
package platform
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
var envReloadRegistry = struct {
|
||||||
|
sync.RWMutex
|
||||||
|
handlers []func() error
|
||||||
|
}{}
|
||||||
|
|
||||||
|
// RegisterEnvReload registers an environment reload handler and runs it once
|
||||||
|
// immediately so package defaults keep the same behavior as init-time reads.
|
||||||
|
func RegisterEnvReload(handler func() error) {
|
||||||
|
if handler == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
envReloadRegistry.Lock()
|
||||||
|
envReloadRegistry.handlers = append(envReloadRegistry.handlers, handler)
|
||||||
|
envReloadRegistry.Unlock()
|
||||||
|
if err := handler(); err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReloadEnvSettings refreshes all registered environment-backed package state.
|
||||||
|
func ReloadEnvSettings() error {
|
||||||
|
envReloadRegistry.RLock()
|
||||||
|
handlers := append([]func() error{}, envReloadRegistry.handlers...)
|
||||||
|
envReloadRegistry.RUnlock()
|
||||||
|
|
||||||
|
var errs []error
|
||||||
|
for _, handler := range handlers {
|
||||||
|
if err := handler(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
@@ -3,11 +3,8 @@ package bittorrent
|
|||||||
import (
|
import (
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"errors"
|
"errors"
|
||||||
"math"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type SniffHeader struct{}
|
type SniffHeader struct{}
|
||||||
@@ -39,50 +36,44 @@ func SniffUTP(b []byte) (*SniffHeader, error) {
|
|||||||
return nil, common.ErrNoClue
|
return nil, common.ErrNoClue
|
||||||
}
|
}
|
||||||
|
|
||||||
buffer := buf.FromBytes(b)
|
// type 4 (ST_SYN), version 1
|
||||||
|
if b[0] != 0x41 {
|
||||||
var typeAndVersion uint8
|
|
||||||
|
|
||||||
if binary.Read(buffer, binary.BigEndian, &typeAndVersion) != nil {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
} else if b[0]>>4&0xF > 4 || b[0]&0xF != 1 {
|
|
||||||
return nil, errNotBittorrent
|
return nil, errNotBittorrent
|
||||||
}
|
}
|
||||||
|
|
||||||
var extension uint8
|
// timestamp_difference is always 0 in new connections
|
||||||
|
if binary.BigEndian.Uint32(b[8:12]) != 0 {
|
||||||
if binary.Read(buffer, binary.BigEndian, &extension) != nil {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
} else if extension != 0 && extension != 1 {
|
|
||||||
return nil, errNotBittorrent
|
return nil, errNotBittorrent
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Walk the extension chain. Selective ack (1) and extension bits (2)
|
||||||
|
extension, offset := b[1], 20
|
||||||
for extension != 0 {
|
for extension != 0 {
|
||||||
if extension != 1 {
|
if len(b) < offset+2 {
|
||||||
return nil, errNotBittorrent
|
return nil, errNotBittorrent
|
||||||
}
|
}
|
||||||
if binary.Read(buffer, binary.BigEndian, &extension) != nil {
|
length := int(b[offset+1])
|
||||||
return nil, common.ErrNoClue
|
switch extension {
|
||||||
|
case 1: // selective ack
|
||||||
|
if length < 4 || length%4 != 0 {
|
||||||
|
return nil, errNotBittorrent
|
||||||
|
}
|
||||||
|
case 2: // extension bits: fixed 8 bytes, sent in ST_SYN by µTorrent
|
||||||
|
if length != 8 {
|
||||||
|
return nil, errNotBittorrent
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return nil, errNotBittorrent
|
||||||
}
|
}
|
||||||
|
if len(b) < offset+2+length {
|
||||||
var length uint8
|
return nil, errNotBittorrent
|
||||||
if err := binary.Read(buffer, binary.BigEndian, &length); err != nil {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
}
|
|
||||||
if common.Error2(buffer.ReadBytes(int32(length))) != nil {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
}
|
}
|
||||||
|
extension = b[offset]
|
||||||
|
offset += 2 + length
|
||||||
}
|
}
|
||||||
|
|
||||||
if common.Error2(buffer.ReadBytes(2)) != nil {
|
// extensions should consume all ST_SYN payload
|
||||||
return nil, common.ErrNoClue
|
if len(b) != offset {
|
||||||
}
|
|
||||||
|
|
||||||
var timestamp uint32
|
|
||||||
if err := binary.Read(buffer, binary.BigEndian, ×tamp); err != nil {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
}
|
|
||||||
if math.Abs(float64(time.Now().UnixMicro()-int64(timestamp))) > float64(24*time.Hour) {
|
|
||||||
return nil, errNotBittorrent
|
return nil, errNotBittorrent
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,67 @@
|
|||||||
|
package bittorrent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common"
|
||||||
|
)
|
||||||
|
|
||||||
|
// utpPacket builds the fixed 20-byte header defined by BEP 29.
|
||||||
|
func utpPacket(packetType, extension byte, tsDiff uint32, payload ...byte) []byte {
|
||||||
|
b := make([]byte, 20)
|
||||||
|
b[0] = packetType<<4 | 1
|
||||||
|
b[1] = extension
|
||||||
|
binary.BigEndian.PutUint16(b[2:4], 0x4a3f) // connection_id, random
|
||||||
|
binary.BigEndian.PutUint32(b[4:8], 0x8c3a91d2) // timestamp_microseconds, sender's clock
|
||||||
|
binary.BigEndian.PutUint32(b[8:12], tsDiff)
|
||||||
|
binary.BigEndian.PutUint32(b[12:16], 0x00100000) // wnd_size
|
||||||
|
binary.BigEndian.PutUint16(b[16:18], 0x71ee) // seq_nr, random in libutp/libtorrent
|
||||||
|
binary.BigEndian.PutUint16(b[18:20], 0x0000) // ack_nr
|
||||||
|
return append(b, payload...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSniffUTP(t *testing.T) {
|
||||||
|
selectiveAck := []byte{0, 4, 0xff, 0x00, 0xff, 0x00}
|
||||||
|
wrongVersion := utpPacket(4, 0, 0)
|
||||||
|
wrongVersion[0] = 4<<4 | 2
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
payload []byte
|
||||||
|
err error
|
||||||
|
}{
|
||||||
|
{"syn", utpPacket(4, 0, 0), nil},
|
||||||
|
{"syn with selective ack", append(utpPacket(4, 1, 0), selectiveAck...), nil},
|
||||||
|
{"syn with extension bits", append(utpPacket(4, 2, 0), 0, 8, 1, 2, 3, 4, 5, 6, 7, 8), nil},
|
||||||
|
{"extension bits with wrong length", append(utpPacket(4, 2, 0), 0, 4, 1, 2, 3, 4), errNotBittorrent},
|
||||||
|
{"syn with nonzero timestamp_difference", utpPacket(4, 0, 0x1234), errNotBittorrent},
|
||||||
|
{"syn with trailing payload", utpPacket(4, 0, 0, 'x'), errNotBittorrent},
|
||||||
|
// txid 0x4100, no EDNS0: the worst case colliding with the uTP header
|
||||||
|
{"dns query", []byte{
|
||||||
|
0x41, 0x00, 0x01, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
|
||||||
|
0x01, 'a', 0x02, 'c', 'o', 0x00, 0x00, 0x01, 0x00, 0x01,
|
||||||
|
}, errNotBittorrent},
|
||||||
|
{"established connection packets", utpPacket(0, 0, 0x5678, 'x', 'y', 'z'), errNotBittorrent},
|
||||||
|
{"state", utpPacket(2, 0, 0x5678), errNotBittorrent},
|
||||||
|
{"fin", utpPacket(1, 0, 0x5678), errNotBittorrent},
|
||||||
|
{"wrong version", wrongVersion, errNotBittorrent},
|
||||||
|
{"unknown packet type", utpPacket(5, 0, 0), errNotBittorrent},
|
||||||
|
{"unknown extension", utpPacket(4, 2, 0), errNotBittorrent},
|
||||||
|
{"extension chain past the datagram", utpPacket(4, 1, 0, 0, 8, 0xff), errNotBittorrent},
|
||||||
|
{"selective ack not in multiples of 4", append(utpPacket(4, 1, 0), 0, 3, 0xff, 0x00, 0xff), errNotBittorrent},
|
||||||
|
{"shorter than the header", utpPacket(4, 0, 0)[:19], common.ErrNoClue},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
h, err := SniffUTP(c.payload)
|
||||||
|
if err != c.err {
|
||||||
|
t.Fatalf("expected error %v, got %v", c.err, err)
|
||||||
|
}
|
||||||
|
if err == nil && h == nil {
|
||||||
|
t.Fatal("expected a sniff header, got nil")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -28,8 +28,6 @@ const (
|
|||||||
SecurityType_AUTO SecurityType = 2
|
SecurityType_AUTO SecurityType = 2
|
||||||
SecurityType_AES128_GCM SecurityType = 3
|
SecurityType_AES128_GCM SecurityType = 3
|
||||||
SecurityType_CHACHA20_POLY1305 SecurityType = 4
|
SecurityType_CHACHA20_POLY1305 SecurityType = 4
|
||||||
SecurityType_NONE SecurityType = 5 // [DEPRECATED 2023-06]
|
|
||||||
SecurityType_ZERO SecurityType = 6
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Enum value maps for SecurityType.
|
// Enum value maps for SecurityType.
|
||||||
@@ -39,16 +37,12 @@ var (
|
|||||||
2: "AUTO",
|
2: "AUTO",
|
||||||
3: "AES128_GCM",
|
3: "AES128_GCM",
|
||||||
4: "CHACHA20_POLY1305",
|
4: "CHACHA20_POLY1305",
|
||||||
5: "NONE",
|
|
||||||
6: "ZERO",
|
|
||||||
}
|
}
|
||||||
SecurityType_value = map[string]int32{
|
SecurityType_value = map[string]int32{
|
||||||
"UNKNOWN": 0,
|
"UNKNOWN": 0,
|
||||||
"AUTO": 2,
|
"AUTO": 2,
|
||||||
"AES128_GCM": 3,
|
"AES128_GCM": 3,
|
||||||
"CHACHA20_POLY1305": 4,
|
"CHACHA20_POLY1305": 4,
|
||||||
"NONE": 5,
|
|
||||||
"ZERO": 6,
|
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -129,15 +123,13 @@ const file_common_protocol_headers_proto_rawDesc = "" +
|
|||||||
"\n" +
|
"\n" +
|
||||||
"\x1dcommon/protocol/headers.proto\x12\x14xray.common.protocol\"H\n" +
|
"\x1dcommon/protocol/headers.proto\x12\x14xray.common.protocol\"H\n" +
|
||||||
"\x0eSecurityConfig\x126\n" +
|
"\x0eSecurityConfig\x126\n" +
|
||||||
"\x04type\x18\x01 \x01(\x0e2\".xray.common.protocol.SecurityTypeR\x04type*`\n" +
|
"\x04type\x18\x01 \x01(\x0e2\".xray.common.protocol.SecurityTypeR\x04type*L\n" +
|
||||||
"\fSecurityType\x12\v\n" +
|
"\fSecurityType\x12\v\n" +
|
||||||
"\aUNKNOWN\x10\x00\x12\b\n" +
|
"\aUNKNOWN\x10\x00\x12\b\n" +
|
||||||
"\x04AUTO\x10\x02\x12\x0e\n" +
|
"\x04AUTO\x10\x02\x12\x0e\n" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"AES128_GCM\x10\x03\x12\x15\n" +
|
"AES128_GCM\x10\x03\x12\x15\n" +
|
||||||
"\x11CHACHA20_POLY1305\x10\x04\x12\b\n" +
|
"\x11CHACHA20_POLY1305\x10\x04B^\n" +
|
||||||
"\x04NONE\x10\x05\x12\b\n" +
|
|
||||||
"\x04ZERO\x10\x06B^\n" +
|
|
||||||
"\x18com.xray.common.protocolP\x01Z)github.com/xtls/xray-core/common/protocol\xaa\x02\x14Xray.Common.Protocolb\x06proto3"
|
"\x18com.xray.common.protocolP\x01Z)github.com/xtls/xray-core/common/protocol\xaa\x02\x14Xray.Common.Protocolb\x06proto3"
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|||||||
@@ -11,8 +11,6 @@ enum SecurityType {
|
|||||||
AUTO = 2;
|
AUTO = 2;
|
||||||
AES128_GCM = 3;
|
AES128_GCM = 3;
|
||||||
CHACHA20_POLY1305 = 4;
|
CHACHA20_POLY1305 = 4;
|
||||||
NONE = 5; // [DEPRECATED 2023-06]
|
|
||||||
ZERO = 6;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
message SecurityConfig {
|
message SecurityConfig {
|
||||||
|
|||||||
@@ -1,25 +1,41 @@
|
|||||||
package http
|
package http
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ParseXForwardedFor parses X-Forwarded-For header in http headers, and return the IP list in it.
|
// ApplyTrustedXForwardedFor returns remoteAddr overridden by X-Forwarded-For only when a configured trusted header is present.
|
||||||
func ParseXForwardedFor(header http.Header) []net.Address {
|
func ApplyTrustedXForwardedFor(header http.Header, trusted []string, remoteAddr net.Addr) net.Addr {
|
||||||
xff := header.Get("X-Forwarded-For")
|
value := header.Get("X-Forwarded-For")
|
||||||
if xff == "" {
|
if value == "" {
|
||||||
return nil
|
return remoteAddr
|
||||||
}
|
}
|
||||||
list := strings.Split(xff, ",")
|
for _, t := range trusted {
|
||||||
addrs := make([]net.Address, 0, len(list))
|
if len(header.Values(t)) > 0 {
|
||||||
for _, proxy := range list {
|
if idx := strings.IndexByte(value, ','); idx >= 0 {
|
||||||
addrs = append(addrs, net.ParseAddress(proxy))
|
value = value[:idx]
|
||||||
|
}
|
||||||
|
if addr := net.ParseAddress(value); addr.Family().IsIP() {
|
||||||
|
return &net.TCPAddr{
|
||||||
|
IP: addr.IP(),
|
||||||
|
Port: 0,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return remoteAddr
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return addrs
|
if len(trusted) == 0 {
|
||||||
|
errors.LogWarning(context.Background(), `received "X-Forwarded-For" from `, remoteAddr, ` but "sockopt.trustedXForwardedFor" is not configured; ignoring it and using the real remote address`)
|
||||||
|
} else {
|
||||||
|
errors.LogError(context.Background(), `ignored potentially forged "X-Forwarded-For" from `, remoteAddr, `: `, value)
|
||||||
|
}
|
||||||
|
return remoteAddr
|
||||||
}
|
}
|
||||||
|
|
||||||
// RemoveHopByHopHeaders removes hop by hop headers in http header list.
|
// RemoveHopByHopHeaders removes hop by hop headers in http header list.
|
||||||
|
|||||||
@@ -2,23 +2,48 @@ package http_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
|
gonet "net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/google/go-cmp/cmp"
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
. "github.com/xtls/xray-core/common/protocol/http"
|
. "github.com/xtls/xray-core/common/protocol/http"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestParseXForwardedFor(t *testing.T) {
|
func TestApplyTrustedXForwardedFor(t *testing.T) {
|
||||||
header := http.Header{}
|
remoteAddr := &gonet.TCPAddr{IP: gonet.ParseIP("127.0.0.1"), Port: 12345}
|
||||||
header.Add("X-Forwarded-For", "129.78.138.66, 129.78.64.103")
|
|
||||||
addrs := ParseXForwardedFor(header)
|
t.Run("ignore X-Forwarded-For without trusted header", func(t *testing.T) {
|
||||||
if r := cmp.Diff(addrs, []net.Address{net.ParseAddress("129.78.138.66"), net.ParseAddress("129.78.64.103")}); r != "" {
|
header := http.Header{}
|
||||||
t.Error(r)
|
header.Add("X-Forwarded-For", "129.78.138.66, 129.78.64.103")
|
||||||
}
|
|
||||||
|
if addr := ApplyTrustedXForwardedFor(header, nil, remoteAddr); addr != remoteAddr {
|
||||||
|
t.Fatalf("unexpected remote address: %v", addr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("trust X-Forwarded-For", func(t *testing.T) {
|
||||||
|
header := http.Header{}
|
||||||
|
header.Add("X-Forwarded-For", "129.78.138.66, 129.78.64.103")
|
||||||
|
header.Add("X-Trusted-CDN", "")
|
||||||
|
|
||||||
|
addr := ApplyTrustedXForwardedFor(header, []string{"X-Trusted-CDN"}, remoteAddr)
|
||||||
|
if addr.String() != "129.78.138.66:0" {
|
||||||
|
t.Fatalf("unexpected remote address: %v", addr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("ignore non-IP X-Forwarded-For", func(t *testing.T) {
|
||||||
|
header := http.Header{}
|
||||||
|
header.Add("X-Forwarded-For", "example.com")
|
||||||
|
header.Add("X-Trusted-CDN", "")
|
||||||
|
|
||||||
|
if addr := ApplyTrustedXForwardedFor(header, []string{"X-Trusted-CDN"}, remoteAddr); addr != remoteAddr {
|
||||||
|
t.Fatalf("unexpected remote address: %v", addr)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHopByHopHeadersRemoving(t *testing.T) {
|
func TestHopByHopHeadersRemoving(t *testing.T) {
|
||||||
|
|||||||
@@ -1,18 +1,10 @@
|
|||||||
package quic
|
package quic
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto"
|
|
||||||
"crypto/cipher"
|
"crypto/cipher"
|
||||||
_ "crypto/tls"
|
_ "crypto/tls"
|
||||||
_ "unsafe"
|
_ "unsafe"
|
||||||
)
|
)
|
||||||
|
|
||||||
type CipherSuiteTLS13 struct {
|
|
||||||
ID uint16
|
|
||||||
KeyLen int
|
|
||||||
AEAD func(key, fixedNonce []byte) cipher.AEAD
|
|
||||||
Hash crypto.Hash
|
|
||||||
}
|
|
||||||
|
|
||||||
//go:linkname AEADAESGCMTLS13 crypto/tls.aeadAESGCMTLS13
|
//go:linkname AEADAESGCMTLS13 crypto/tls.aeadAESGCMTLS13
|
||||||
func AEADAESGCMTLS13(key, nonceMask []byte) cipher.AEAD
|
func AEADAESGCMTLS13(key, nonceMask []byte) cipher.AEAD
|
||||||
|
|||||||
@@ -3,7 +3,6 @@ package quic
|
|||||||
import (
|
import (
|
||||||
"crypto"
|
"crypto"
|
||||||
"crypto/aes"
|
"crypto/aes"
|
||||||
"crypto/tls"
|
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
"io"
|
||||||
|
|
||||||
@@ -28,22 +27,43 @@ func (s SniffHeader) Domain() string {
|
|||||||
return s.domain
|
return s.domain
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
var (
|
||||||
versionDraft29 uint32 = 0xff00001d
|
errNotQUIC = errors.New("not quic")
|
||||||
version1 uint32 = 0x1
|
errNotQUICInitial = errors.New("not initial packet")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type quicVersionSpec struct {
|
||||||
|
ver uint32
|
||||||
|
typeInitial byte
|
||||||
|
initialSalt []byte
|
||||||
|
labelPrefix string
|
||||||
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
quicSaltOld = []byte{0xaf, 0xbf, 0xec, 0x28, 0x99, 0x93, 0xd2, 0x4c, 0x9e, 0x97, 0x86, 0xf1, 0x9c, 0x61, 0x11, 0xe0, 0x43, 0x90, 0xa8, 0x99}
|
quicDraft29 = quicVersionSpec{
|
||||||
quicSalt = []byte{0x38, 0x76, 0x2c, 0xf7, 0xf5, 0x59, 0x34, 0xb3, 0x4d, 0x17, 0x9a, 0xe6, 0xa4, 0xc8, 0x0c, 0xad, 0xcc, 0xbb, 0x7f, 0x0a}
|
ver: 0xff00001d,
|
||||||
initialSuite = &CipherSuiteTLS13{
|
typeInitial: 0b00,
|
||||||
ID: tls.TLS_AES_128_GCM_SHA256,
|
initialSalt: []byte{0xaf, 0xbf, 0xec, 0x28, 0x99, 0x93, 0xd2, 0x4c, 0x9e, 0x97, 0x86, 0xf1, 0x9c, 0x61, 0x11, 0xe0, 0x43, 0x90, 0xa8, 0x99},
|
||||||
KeyLen: 16,
|
labelPrefix: "quic",
|
||||||
AEAD: AEADAESGCMTLS13,
|
}
|
||||||
Hash: crypto.SHA256,
|
quicV1 = quicVersionSpec{
|
||||||
|
ver: 0x1,
|
||||||
|
typeInitial: 0b00,
|
||||||
|
initialSalt: []byte{0x38, 0x76, 0x2c, 0xf7, 0xf5, 0x59, 0x34, 0xb3, 0x4d, 0x17, 0x9a, 0xe6, 0xa4, 0xc8, 0x0c, 0xad, 0xcc, 0xbb, 0x7f, 0x0a},
|
||||||
|
labelPrefix: "quic",
|
||||||
|
}
|
||||||
|
quicV2 = quicVersionSpec{
|
||||||
|
ver: 0x6b3343cf,
|
||||||
|
typeInitial: 0b01,
|
||||||
|
initialSalt: []byte{0x0d, 0xed, 0xe3, 0xde, 0xf7, 0x00, 0xa6, 0xdb, 0x81, 0x93, 0x81, 0xbe, 0x6e, 0x26, 0x9d, 0xcb, 0xf9, 0xbd, 0x2e, 0xd9},
|
||||||
|
labelPrefix: "quicv2",
|
||||||
|
}
|
||||||
|
|
||||||
|
quicVersionSpecMap = map[uint32]*quicVersionSpec{
|
||||||
|
quicDraft29.ver: &quicDraft29,
|
||||||
|
quicV1.ver: &quicV1,
|
||||||
|
quicV2.ver: &quicV2,
|
||||||
}
|
}
|
||||||
errNotQuic = errors.New("not quic")
|
|
||||||
errNotQuicInitial = errors.New("not initial packet")
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func SniffQUIC(b []byte) (*SniffHeader, error) {
|
func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||||
@@ -63,60 +83,61 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
|||||||
buffer := buf.FromBytes(b)
|
buffer := buf.FromBytes(b)
|
||||||
typeByte, err := buffer.ReadByte()
|
typeByte, err := buffer.ReadByte()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
|
|
||||||
isLongHeader := typeByte&0x80 > 0
|
isLongHeader := typeByte&0x80 > 0
|
||||||
if !isLongHeader || typeByte&0x40 == 0 {
|
if !isLongHeader || typeByte&0x40 == 0 {
|
||||||
return nil, errNotQuicInitial
|
return nil, errNotQUICInitial
|
||||||
}
|
}
|
||||||
|
|
||||||
vb, err := buffer.ReadBytes(4)
|
vb, err := buffer.ReadBytes(4)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
|
|
||||||
versionNumber := binary.BigEndian.Uint32(vb)
|
versionNumber := binary.BigEndian.Uint32(vb)
|
||||||
if versionNumber != 0 && typeByte&0x40 == 0 {
|
var s *quicVersionSpec
|
||||||
return nil, errNotQuic
|
if v, ok := quicVersionSpecMap[versionNumber]; ok {
|
||||||
} else if versionNumber != versionDraft29 && versionNumber != version1 {
|
s = v
|
||||||
return nil, errNotQuic
|
} else {
|
||||||
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
|
|
||||||
packetType := (typeByte & 0x30) >> 4
|
|
||||||
isQuicInitial := packetType == 0x0
|
|
||||||
|
|
||||||
var destConnID []byte
|
var destConnID []byte
|
||||||
if l, err := buffer.ReadByte(); err != nil {
|
if l, err := buffer.ReadByte(); err != nil {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
} else if destConnID, err = buffer.ReadBytes(int32(l)); err != nil {
|
} else if destConnID, err = buffer.ReadBytes(int32(l)); err != nil {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
|
|
||||||
if l, err := buffer.ReadByte(); err != nil {
|
if l, err := buffer.ReadByte(); err != nil {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
} else if common.Error2(buffer.ReadBytes(int32(l))) != nil {
|
} else if common.Error2(buffer.ReadBytes(int32(l))) != nil {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
|
|
||||||
if isQuicInitial { // Only initial packets have token, see https://datatracker.ietf.org/doc/html/rfc9000#section-17.2.2
|
packetType := (typeByte & 0x30) >> 4
|
||||||
tokenLen, err := readShortQuicVarint(buffer)
|
isQUICInitial := packetType == s.typeInitial
|
||||||
|
|
||||||
|
if isQUICInitial { // Only initial packets have token, see https://datatracker.ietf.org/doc/html/rfc9000#section-17.2.2
|
||||||
|
tokenLen, err := readShortQUICVarint(buffer)
|
||||||
if err != nil || tokenLen > int32(len(b)) {
|
if err != nil || tokenLen > int32(len(b)) {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err = buffer.ReadBytes(tokenLen); err != nil {
|
if _, err = buffer.ReadBytes(tokenLen); err != nil {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
packetLen, err := readShortQuicVarint(buffer)
|
packetLen, err := readShortQUICVarint(buffer)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
// packetLen is impossible to be shorter than this
|
// packetLen is impossible to be shorter than this
|
||||||
if packetLen < 4 {
|
if packetLen < 4 {
|
||||||
return nil, errNotQuic
|
return nil, errNotQUIC
|
||||||
}
|
}
|
||||||
|
|
||||||
hdrLen := len(b) - int(buffer.Len())
|
hdrLen := len(b) - int(buffer.Len())
|
||||||
@@ -125,25 +146,23 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
restPayload := b[hdrLen+int(packetLen):]
|
restPayload := b[hdrLen+int(packetLen):]
|
||||||
if !isQuicInitial { // Skip this packet if it's not initial packet
|
if !isQUICInitial { // Skip this packet if it's not initial packet
|
||||||
b = restPayload
|
b = restPayload
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
var salt []byte
|
salt := s.initialSalt
|
||||||
if versionNumber == version1 {
|
label := s.labelPrefix
|
||||||
salt = quicSalt
|
|
||||||
} else {
|
|
||||||
salt = quicSaltOld
|
|
||||||
}
|
|
||||||
initialSecret := hkdf.Extract(crypto.SHA256.New, destConnID, salt)
|
initialSecret := hkdf.Extract(crypto.SHA256.New, destConnID, salt)
|
||||||
secret := hkdfExpandLabel(crypto.SHA256, initialSecret, []byte{}, "client in", crypto.SHA256.Size())
|
secret := hkdfExpandLabel(initialSecret, "client in", crypto.SHA256.Size())
|
||||||
hpKey := hkdfExpandLabel(initialSuite.Hash, secret, []byte{}, "quic hp", initialSuite.KeyLen)
|
hpKey := hkdfExpandLabel(secret, label+" hp", 16)
|
||||||
block, err := aes.NewCipher(hpKey)
|
block, err := aes.NewCipher(hpKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if len(b) < hdrLen+4+block.BlockSize() {
|
||||||
|
return nil, errNotQUIC
|
||||||
|
}
|
||||||
cache.Clear()
|
cache.Clear()
|
||||||
mask := cache.Extend(int32(block.BlockSize()))
|
mask := cache.Extend(int32(block.BlockSize()))
|
||||||
block.Encrypt(mask, b[hdrLen+4:hdrLen+4+len(mask)])
|
block.Encrypt(mask, b[hdrLen+4:hdrLen+4+len(mask)])
|
||||||
@@ -153,8 +172,8 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
|||||||
b[hdrLen+i] ^= mask[i+1]
|
b[hdrLen+i] ^= mask[i+1]
|
||||||
}
|
}
|
||||||
|
|
||||||
key := hkdfExpandLabel(crypto.SHA256, secret, []byte{}, "quic key", 16)
|
key := hkdfExpandLabel(secret, label+" key", 16)
|
||||||
iv := hkdfExpandLabel(crypto.SHA256, secret, []byte{}, "quic iv", 12)
|
iv := hkdfExpandLabel(secret, label+" iv", 12)
|
||||||
cipher := AEADAESGCMTLS13(key, iv)
|
cipher := AEADAESGCMTLS13(key, iv)
|
||||||
|
|
||||||
nonce := cache.Extend(int32(cipher.NonceSize()))
|
nonce := cache.Extend(int32(cipher.NonceSize()))
|
||||||
@@ -179,44 +198,44 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
|||||||
case 0x00: // PADDING frame
|
case 0x00: // PADDING frame
|
||||||
case 0x01: // PING frame
|
case 0x01: // PING frame
|
||||||
case 0x02, 0x03: // ACK frame
|
case 0x02, 0x03: // ACK frame
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: Largest Acknowledged
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: Largest Acknowledged
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ACK Delay
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ACK Delay
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
ackRangeCount, err := readShortQuicVarint(buffer) // Field: ACK Range Count
|
ackRangeCount, err := readShortQUICVarint(buffer) // Field: ACK Range Count
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: First ACK Range
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: First ACK Range
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
for i := 0; i < int(ackRangeCount); i++ { // Field: ACK Range
|
for i := 0; i < int(ackRangeCount); i++ { // Field: ACK Range
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ACK Range -> Gap
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ACK Range -> Gap
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ACK Range -> ACK Range Length
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ACK Range -> ACK Range Length
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if frameType == 0x03 {
|
if frameType == 0x03 {
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ECN Counts -> ECT0 Count
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ECN Counts -> ECT0 Count
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ECN Counts -> ECT1 Count
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ECN Counts -> ECT1 Count
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { //nolint:misspell // Field: ECN Counts -> ECT-CE Count
|
if _, err = readShortQUICVarint(buffer); err != nil { //nolint:misspell // Field: ECN Counts -> ECT-CE Count
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
case 0x06: // CRYPTO frame, we will use this frame
|
case 0x06: // CRYPTO frame, we will use this frame
|
||||||
offset, err := readShortQuicVarint(buffer) // Field: Offset
|
offset, err := readShortQUICVarint(buffer) // Field: Offset
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
length, err := readShortQuicVarint(buffer) // Field: Length
|
length, err := readShortQUICVarint(buffer) // Field: Length
|
||||||
if err != nil || length > buffer.Len() {
|
if err != nil || length > buffer.Len() {
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
@@ -232,13 +251,13 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
|||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
case 0x1c: // CONNECTION_CLOSE frame, only 0x1c is permitted in initial packet
|
case 0x1c: // CONNECTION_CLOSE frame, only 0x1c is permitted in initial packet
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: Error Code
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: Error Code
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: Frame Type
|
if _, err = readShortQUICVarint(buffer); err != nil { // Field: Frame Type
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
length, err := readShortQuicVarint(buffer) // Field: Reason Phrase Length
|
length, err := readShortQUICVarint(buffer) // Field: Reason Phrase Length
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, io.ErrUnexpectedEOF
|
return nil, io.ErrUnexpectedEOF
|
||||||
}
|
}
|
||||||
@@ -248,7 +267,7 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
|||||||
default:
|
default:
|
||||||
// Only above frame types are permitted in initial packet.
|
// Only above frame types are permitted in initial packet.
|
||||||
// See https://www.rfc-editor.org/rfc/rfc9000.html#section-17.2.2-8
|
// See https://www.rfc-editor.org/rfc/rfc9000.html#section-17.2.2-8
|
||||||
return nil, errNotQuicInitial
|
return nil, errNotQUICInitial
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -266,35 +285,33 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
|||||||
return nil, protocol.ErrProtoNeedMoreData
|
return nil, protocol.ErrProtoNeedMoreData
|
||||||
}
|
}
|
||||||
|
|
||||||
func hkdfExpandLabel(hash crypto.Hash, secret, context []byte, label string, length int) []byte {
|
func hkdfExpandLabel(secret []byte, label string, length int) []byte {
|
||||||
b := make([]byte, 3, 3+6+len(label)+1+len(context))
|
b := make([]byte, 0, 2+1+6+len(label)+1)
|
||||||
binary.BigEndian.PutUint16(b, uint16(length))
|
b = binary.BigEndian.AppendUint16(b, uint16(length))
|
||||||
b[2] = uint8(6 + len(label))
|
b = append(b, byte(6+len(label)))
|
||||||
b = append(b, []byte("tls13 ")...)
|
b = append(b, "tls13 "...)
|
||||||
b = append(b, []byte(label)...)
|
b = append(b, label...)
|
||||||
b = b[:3+6+len(label)+1]
|
b = append(b, 0) // context
|
||||||
b[3+6+len(label)] = uint8(len(context))
|
|
||||||
b = append(b, context...)
|
|
||||||
|
|
||||||
out := make([]byte, length)
|
out := make([]byte, length)
|
||||||
n, err := hkdf.Expand(hash.New, secret, b).Read(out)
|
n, err := hkdf.Expand(crypto.SHA256.New, secret, b).Read(out)
|
||||||
if err != nil || n != length {
|
if err != nil || n != length {
|
||||||
panic("quic: HKDF-Expand-Label invocation failed unexpectedly")
|
panic("quic: HKDF-Expand-Label invocation failed unexpectedly")
|
||||||
}
|
}
|
||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
// readShortQuicVarint wraps quicvarint.Read with a max limit for length related fields.
|
// readShortQUICVarint wraps quicvarint.Read with a max limit for length related fields.
|
||||||
// we only handle QUIC Initial so these numbers should not exceed 65535
|
// we only handle QUIC Initial so these numbers should not exceed 65535
|
||||||
// returns int32 to reduce type conversion
|
// returns int32 to reduce type conversion
|
||||||
func readShortQuicVarint(reader io.ByteReader) (int32, error) {
|
func readShortQUICVarint(reader io.ByteReader) (int32, error) {
|
||||||
v, err := quicvarint.Read(reader)
|
v, err := quicvarint.Read(reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
if v > 65535 {
|
if v > 65535 {
|
||||||
// not used(
|
// not used(
|
||||||
return 0, errNotQuicInitial
|
return 0, errNotQUICInitial
|
||||||
}
|
}
|
||||||
return int32(v), nil
|
return int32(v), nil
|
||||||
}
|
}
|
||||||
|
|||||||
File diff suppressed because one or more lines are too long
@@ -7,7 +7,7 @@ import (
|
|||||||
|
|
||||||
func (u *User) GetTypedAccount() (Account, error) {
|
func (u *User) GetTypedAccount() (Account, error) {
|
||||||
if u.GetAccount() == nil {
|
if u.GetAccount() == nil {
|
||||||
return nil, errors.New("Account is missing").AtWarning()
|
return nil, errors.New("Account is missing")
|
||||||
}
|
}
|
||||||
|
|
||||||
rawAccount, err := u.Account.GetInstance()
|
rawAccount, err := u.Account.GetInstance()
|
||||||
|
|||||||
@@ -207,6 +207,7 @@ func getConfig() string {
|
|||||||
"tag": "XHTTP_IN",
|
"tag": "XHTTP_IN",
|
||||||
"streamSettings": {
|
"streamSettings": {
|
||||||
"network": "xhttp",
|
"network": "xhttp",
|
||||||
|
"security": "tls",
|
||||||
"xhttpSettings": {
|
"xhttpSettings": {
|
||||||
"host": "bing.com",
|
"host": "bing.com",
|
||||||
"path": "/xhttp_client_upload",
|
"path": "/xhttp_client_upload",
|
||||||
|
|||||||
@@ -26,6 +26,8 @@ const (
|
|||||||
fullHandlerKey ctx.SessionKey = 10 // outbound gets full handler
|
fullHandlerKey ctx.SessionKey = 10 // outbound gets full handler
|
||||||
mitmAlpn11Key ctx.SessionKey = 11 // used by TLS dialer
|
mitmAlpn11Key ctx.SessionKey = 11 // used by TLS dialer
|
||||||
mitmServerNameKey ctx.SessionKey = 12 // used by TLS dialer
|
mitmServerNameKey ctx.SessionKey = 12 // used by TLS dialer
|
||||||
|
|
||||||
|
streamSettingsKey ctx.SessionKey = 13
|
||||||
)
|
)
|
||||||
|
|
||||||
func ContextWithInbound(ctx context.Context, inbound *Inbound) context.Context {
|
func ContextWithInbound(ctx context.Context, inbound *Inbound) context.Context {
|
||||||
@@ -192,3 +194,11 @@ func MitmServerNameFromContext(ctx context.Context) string {
|
|||||||
}
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func ContextWithStreamSettings(ctx context.Context, streamSettings any) context.Context {
|
||||||
|
return context.WithValue(ctx, streamSettingsKey, streamSettings)
|
||||||
|
}
|
||||||
|
|
||||||
|
func StreamSettingsFromContext(ctx context.Context) any {
|
||||||
|
return ctx.Value(streamSettingsKey)
|
||||||
|
}
|
||||||
|
|||||||
@@ -70,8 +70,6 @@ type Outbound struct {
|
|||||||
Tag string
|
Tag string
|
||||||
// Name of the outbound proxy that handles the connection.
|
// Name of the outbound proxy that handles the connection.
|
||||||
Name string
|
Name string
|
||||||
// Unused. Conn is actually internet.Connection. May be nil. It is currently nil for outbound with proxySettings
|
|
||||||
Conn net.Conn
|
|
||||||
// CanSpliceCopy is a property for this connection
|
// CanSpliceCopy is a property for this connection
|
||||||
// 1 = can, 2 = after processing protocol info should be able to, 3 = cannot
|
// 1 = can, 2 = after processing protocol info should be able to, 3 = cannot
|
||||||
CanSpliceCopy int
|
CanSpliceCopy int
|
||||||
|
|||||||
@@ -1,53 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
)
|
|
||||||
|
|
||||||
func ToNetwork(network string) net.Network {
|
|
||||||
switch N.NetworkName(network) {
|
|
||||||
case N.NetworkTCP:
|
|
||||||
return net.Network_TCP
|
|
||||||
case N.NetworkUDP:
|
|
||||||
return net.Network_UDP
|
|
||||||
default:
|
|
||||||
return net.Network_Unknown
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func ToDestination(socksaddr M.Socksaddr, network net.Network) (net.Destination, error) {
|
|
||||||
// IsFqdn() implicitly checks if the domain name is valid
|
|
||||||
if socksaddr.IsFqdn() {
|
|
||||||
return net.Destination{
|
|
||||||
Network: network,
|
|
||||||
Address: net.DomainAddress(socksaddr.Fqdn),
|
|
||||||
Port: net.Port(socksaddr.Port),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// IsIP() implicitly checks if the IP address is valid
|
|
||||||
if socksaddr.IsIP() {
|
|
||||||
return net.Destination{
|
|
||||||
Network: network,
|
|
||||||
Address: net.IPAddress(socksaddr.Addr.AsSlice()),
|
|
||||||
Port: net.Port(socksaddr.Port),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return net.Destination{}, errors.New("invalid socks address: ", socksaddr)
|
|
||||||
}
|
|
||||||
|
|
||||||
func ToSocksaddr(destination net.Destination) M.Socksaddr {
|
|
||||||
var addr M.Socksaddr
|
|
||||||
switch destination.Address.Family() {
|
|
||||||
case net.AddressFamilyDomain:
|
|
||||||
addr.Fqdn = destination.Address.Domain()
|
|
||||||
default:
|
|
||||||
addr.Addr = M.AddrFromIP(destination.Address.IP())
|
|
||||||
}
|
|
||||||
addr.Port = uint16(destination.Port)
|
|
||||||
return addr
|
|
||||||
}
|
|
||||||
@@ -1,72 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/common/net/cnc"
|
|
||||||
"github.com/xtls/xray-core/common/session"
|
|
||||||
"github.com/xtls/xray-core/proxy"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
"github.com/xtls/xray-core/transport/pipe"
|
|
||||||
)
|
|
||||||
|
|
||||||
var _ N.Dialer = (*XrayDialer)(nil)
|
|
||||||
|
|
||||||
type XrayDialer struct {
|
|
||||||
internet.Dialer
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewDialer(dialer internet.Dialer) *XrayDialer {
|
|
||||||
return &XrayDialer{dialer}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
|
||||||
dest, err := ToDestination(destination, ToNetwork(network))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return d.Dialer.Dial(ctx, dest)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
|
||||||
return nil, os.ErrInvalid
|
|
||||||
}
|
|
||||||
|
|
||||||
type XrayOutboundDialer struct {
|
|
||||||
outbound proxy.Outbound
|
|
||||||
dialer internet.Dialer
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewOutboundDialer(outbound proxy.Outbound, dialer internet.Dialer) *XrayOutboundDialer {
|
|
||||||
return &XrayOutboundDialer{outbound, dialer}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayOutboundDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
|
||||||
dest, err := ToDestination(destination, ToNetwork(network))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
outbounds := session.OutboundsFromContext(ctx)
|
|
||||||
if len(outbounds) == 0 {
|
|
||||||
outbounds = []*session.Outbound{{}}
|
|
||||||
ctx = session.ContextWithOutbounds(ctx, outbounds)
|
|
||||||
}
|
|
||||||
ob := outbounds[len(outbounds)-1]
|
|
||||||
ob.Target = dest
|
|
||||||
|
|
||||||
opts := []pipe.Option{pipe.WithSizeLimit(64 * 1024)}
|
|
||||||
uplinkReader, uplinkWriter := pipe.New(opts...)
|
|
||||||
downlinkReader, downlinkWriter := pipe.New(opts...)
|
|
||||||
conn := cnc.NewConnection(cnc.ConnectionInputMulti(downlinkWriter), cnc.ConnectionOutputMulti(uplinkReader))
|
|
||||||
go d.outbound.Process(ctx, &transport.Link{Reader: downlinkReader, Writer: uplinkWriter}, d.dialer)
|
|
||||||
return conn, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayOutboundDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
|
||||||
return nil, os.ErrInvalid
|
|
||||||
}
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import E "github.com/sagernet/sing/common/exceptions"
|
|
||||||
|
|
||||||
func ReturnError(err error) error {
|
|
||||||
if E.IsClosedOrCanceled(err) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
@@ -1,58 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/features/routing"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
_ N.TCPConnectionHandler = (*Dispatcher)(nil)
|
|
||||||
_ N.UDPConnectionHandler = (*Dispatcher)(nil)
|
|
||||||
)
|
|
||||||
|
|
||||||
type Dispatcher struct {
|
|
||||||
upstream routing.Dispatcher
|
|
||||||
newErrorFunc func(values ...any) *errors.Error
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewDispatcher(dispatcher routing.Dispatcher, newErrorFunc func(values ...any) *errors.Error) *Dispatcher {
|
|
||||||
return &Dispatcher{
|
|
||||||
upstream: dispatcher,
|
|
||||||
newErrorFunc: newErrorFunc,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
|
||||||
dest, err := ToDestination(metadata.Destination, net.Network_TCP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
xConn := NewConn(conn)
|
|
||||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
|
||||||
Reader: xConn,
|
|
||||||
Writer: xConn,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
|
||||||
dest, err := ToDestination(metadata.Destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
|
||||||
Reader: buf.NewPacketReader(conn.(io.Reader)),
|
|
||||||
Writer: buf.NewWriter(conn.(io.Writer)),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewError(ctx context.Context, err error) {
|
|
||||||
errors.LogInfo(ctx, err.Error())
|
|
||||||
}
|
|
||||||
@@ -1,70 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
|
|
||||||
"github.com/sagernet/sing/common/logger"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
)
|
|
||||||
|
|
||||||
var _ logger.ContextLogger = (*XrayLogger)(nil)
|
|
||||||
|
|
||||||
type XrayLogger struct {
|
|
||||||
newError func(values ...any) *errors.Error
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewLogger(newErrorFunc func(values ...any) *errors.Error) *XrayLogger {
|
|
||||||
return &XrayLogger{
|
|
||||||
newErrorFunc,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Trace(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Debug(args ...any) {
|
|
||||||
errors.LogDebug(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Info(args ...any) {
|
|
||||||
errors.LogInfo(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Warn(args ...any) {
|
|
||||||
errors.LogWarning(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Error(args ...any) {
|
|
||||||
errors.LogError(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Fatal(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Panic(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) TraceContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) DebugContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogDebug(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) InfoContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogInfo(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) WarnContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogWarning(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) ErrorContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogError(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) FatalContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) PanicContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
@@ -1,107 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
B "github.com/sagernet/sing/common/buf"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
func CopyPacketConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, destination net.Destination, serverConn net.PacketConn) error {
|
|
||||||
cancel := func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(serverConn)
|
|
||||||
}
|
|
||||||
conn := &PacketConnWrapper{
|
|
||||||
Reader: link.Reader,
|
|
||||||
Writer: link.Writer,
|
|
||||||
Dest: destination,
|
|
||||||
Conn: inboundConn,
|
|
||||||
T: signal.CancelAfterInactivity(ctx, cancel, 300*time.Second),
|
|
||||||
}
|
|
||||||
return ReturnError(bufio.CopyPacketConn(ctx, conn, bufio.NewPacketConn(serverConn)))
|
|
||||||
}
|
|
||||||
|
|
||||||
type PacketConnWrapper struct {
|
|
||||||
buf.Reader
|
|
||||||
buf.Writer
|
|
||||||
net.Conn
|
|
||||||
Dest net.Destination
|
|
||||||
cached buf.MultiBuffer
|
|
||||||
|
|
||||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
|
||||||
T *signal.ActivityTimer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
defer func() {
|
|
||||||
if err != nil {
|
|
||||||
// uplinkonly
|
|
||||||
w.T.SetTimeout(2 * time.Second)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
if w.cached != nil {
|
|
||||||
mb, bb := buf.SplitFirst(w.cached)
|
|
||||||
if bb == nil {
|
|
||||||
w.cached = nil
|
|
||||||
} else {
|
|
||||||
buffer.Write(bb.Bytes())
|
|
||||||
w.cached = mb
|
|
||||||
var destination net.Destination
|
|
||||||
if bb.UDP != nil {
|
|
||||||
destination = *bb.UDP
|
|
||||||
} else {
|
|
||||||
destination = w.Dest
|
|
||||||
}
|
|
||||||
bb.Release()
|
|
||||||
return ToSocksaddr(destination), nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
mb, err := w.ReadMultiBuffer()
|
|
||||||
nb, bb := buf.SplitFirst(mb)
|
|
||||||
if bb == nil {
|
|
||||||
return M.Socksaddr{}, nil
|
|
||||||
} else {
|
|
||||||
buffer.Write(bb.Bytes())
|
|
||||||
w.cached = nb
|
|
||||||
var destination net.Destination
|
|
||||||
if bb.UDP != nil {
|
|
||||||
destination = *bb.UDP
|
|
||||||
} else {
|
|
||||||
destination = w.Dest
|
|
||||||
}
|
|
||||||
bb.Release()
|
|
||||||
return ToSocksaddr(destination), nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) WritePacket(buffer *B.Buffer, destination M.Socksaddr) (err error) {
|
|
||||||
w.T.Update()
|
|
||||||
defer func() {
|
|
||||||
if err != nil {
|
|
||||||
// downlinkonly
|
|
||||||
w.T.SetTimeout(5 * time.Second)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
endpoint, err := ToDestination(destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
vBuf := buf.New()
|
|
||||||
vBuf.Write(buffer.Bytes())
|
|
||||||
vBuf.UDP = &endpoint
|
|
||||||
return w.WriteMultiBuffer(buf.MultiBuffer{vBuf})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) Close() error {
|
|
||||||
buf.ReleaseMulti(w.cached)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,81 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
func CopyConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, serverConn net.Conn) error {
|
|
||||||
conn := &PipeConnWrapper{
|
|
||||||
W: link.Writer,
|
|
||||||
Conn: inboundConn,
|
|
||||||
}
|
|
||||||
if ir, ok := link.Reader.(io.Reader); ok {
|
|
||||||
conn.R = ir
|
|
||||||
} else {
|
|
||||||
conn.R = &buf.BufferedReader{Reader: link.Reader}
|
|
||||||
}
|
|
||||||
cancel := func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(serverConn)
|
|
||||||
}
|
|
||||||
conn.T = signal.CancelAfterInactivity(ctx, cancel, 300*time.Second)
|
|
||||||
return ReturnError(bufio.CopyConn(ctx, conn, serverConn))
|
|
||||||
}
|
|
||||||
|
|
||||||
type PipeConnWrapper struct {
|
|
||||||
R io.Reader
|
|
||||||
W buf.Writer
|
|
||||||
net.Conn
|
|
||||||
|
|
||||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
|
||||||
T *signal.ActivityTimer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Close() error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Read(b []byte) (n int, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
n, err = w.R.Read(b)
|
|
||||||
if err != nil {
|
|
||||||
// uplinkonly
|
|
||||||
w.T.SetTimeout(2 * time.Second)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Write(p []byte) (n int, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
n = len(p)
|
|
||||||
var mb buf.MultiBuffer
|
|
||||||
pLen := len(p)
|
|
||||||
for pLen > 0 {
|
|
||||||
buffer := buf.New()
|
|
||||||
if pLen > buf.Size {
|
|
||||||
_, err = buffer.Write(p[:buf.Size])
|
|
||||||
p = p[buf.Size:]
|
|
||||||
} else {
|
|
||||||
buffer.Write(p)
|
|
||||||
}
|
|
||||||
pLen -= int(buffer.Len())
|
|
||||||
mb = append(mb, buffer)
|
|
||||||
}
|
|
||||||
err = w.W.WriteMultiBuffer(mb)
|
|
||||||
if err != nil {
|
|
||||||
n = 0
|
|
||||||
buf.ReleaseMulti(mb)
|
|
||||||
// downlinkonly
|
|
||||||
w.T.SetTimeout(5 * time.Second)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user