mirror of
https://github.com/fatedier/frp.git
synced 2026-10-04 23:45:56 +08:00
Compare commits
10
Commits
dev
..
8dd26c6961
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8dd26c6961 | ||
|
|
c8c1e5116c | ||
|
|
4ec8de973f | ||
|
|
5bfcea3d0c | ||
|
|
0a1b4ab21f | ||
|
|
5f575b8442 | ||
|
|
a1348cdf00 | ||
|
|
2f5e1f7945 | ||
|
|
22ae8166d3 | ||
|
|
af6bc6369d |
+10
-12
@@ -1,4 +1,4 @@
|
|||||||
version: 2.1
|
version: 2
|
||||||
jobs:
|
jobs:
|
||||||
go-version-latest:
|
go-version-latest:
|
||||||
docker:
|
docker:
|
||||||
@@ -7,20 +7,18 @@ jobs:
|
|||||||
steps:
|
steps:
|
||||||
- checkout
|
- checkout
|
||||||
- run:
|
- run:
|
||||||
name: Test and build web assets
|
name: Build web assets (frps)
|
||||||
command: make web-ci
|
command: make install build
|
||||||
|
working_directory: web/frps
|
||||||
- run:
|
- run:
|
||||||
name: Check Go formatting and build binaries
|
name: Build web assets (frpc)
|
||||||
command: |
|
command: make install build
|
||||||
set -e
|
working_directory: web/frpc
|
||||||
make env fmt
|
- run: make
|
||||||
git diff --exit-code
|
- run: make alltest
|
||||||
make build
|
|
||||||
- run:
|
|
||||||
name: Run all tests
|
|
||||||
command: make alltest
|
|
||||||
|
|
||||||
workflows:
|
workflows:
|
||||||
|
version: 2
|
||||||
build_and_test:
|
build_and_test:
|
||||||
jobs:
|
jobs:
|
||||||
- go-version-latest
|
- go-version-latest
|
||||||
|
|||||||
@@ -65,7 +65,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./dockerfiles/Dockerfile-for-frpc
|
file: ./dockerfiles/Dockerfile-for-frpc
|
||||||
platforms: linux/amd64,linux/arm/v7,linux/arm64,linux/ppc64le
|
platforms: linux/amd64,linux/arm/v7,linux/arm64,linux/ppc64le,linux/s390x
|
||||||
push: true
|
push: true
|
||||||
tags: |
|
tags: |
|
||||||
${{ env.TAG_FRPC }}
|
${{ env.TAG_FRPC }}
|
||||||
@@ -76,7 +76,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./dockerfiles/Dockerfile-for-frps
|
file: ./dockerfiles/Dockerfile-for-frps
|
||||||
platforms: linux/amd64,linux/arm/v7,linux/arm64,linux/ppc64le
|
platforms: linux/amd64,linux/arm/v7,linux/arm64,linux/ppc64le,linux/s390x
|
||||||
push: true
|
push: true
|
||||||
tags: |
|
tags: |
|
||||||
${{ env.TAG_FRPS }}
|
${{ env.TAG_FRPS }}
|
||||||
|
|||||||
@@ -22,10 +22,14 @@ jobs:
|
|||||||
- uses: actions/setup-node@v6
|
- uses: actions/setup-node@v6
|
||||||
with:
|
with:
|
||||||
node-version: '22'
|
node-version: '22'
|
||||||
- name: Test and build web assets
|
- name: Build web assets (frps)
|
||||||
run: make web-ci
|
run: make build
|
||||||
|
working-directory: web/frps
|
||||||
|
- name: Build web assets (frpc)
|
||||||
|
run: make build
|
||||||
|
working-directory: web/frpc
|
||||||
- name: golangci-lint
|
- name: golangci-lint
|
||||||
uses: golangci/golangci-lint-action@v9
|
uses: golangci/golangci-lint-action@v9
|
||||||
with:
|
with:
|
||||||
# Optional: version of golangci-lint to use in form of v1.2 or v1.2.3 or `latest` to use the latest version
|
# Optional: version of golangci-lint to use in form of v1.2 or v1.2.3 or `latest` to use the latest version
|
||||||
version: v2.12.2
|
version: v2.11
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ NOWEB_TAG = $(shell [ ! -d web/frps/dist ] || [ ! -d web/frpc/dist ] && echo ',n
|
|||||||
FRP_COMPAT_BASELINE_COUNT ?= 8
|
FRP_COMPAT_BASELINE_COUNT ?= 8
|
||||||
FRP_COMPAT_FLOOR_VERSION ?= 0.61.0
|
FRP_COMPAT_FLOOR_VERSION ?= 0.61.0
|
||||||
|
|
||||||
.PHONY: web web-ci frps-web frpc-web frps frpc e2e-compatibility-smoke e2e-compatibility e2e-compatibility-floor
|
.PHONY: web frps-web frpc-web frps frpc e2e-compatibility-smoke e2e-compatibility e2e-compatibility-floor
|
||||||
|
|
||||||
all: env fmt web build
|
all: env fmt web build
|
||||||
|
|
||||||
@@ -16,9 +16,6 @@ env:
|
|||||||
|
|
||||||
web: frps-web frpc-web
|
web: frps-web frpc-web
|
||||||
|
|
||||||
web-ci:
|
|
||||||
cd web && npm ci && npm run lint:check --workspace frps && npm run lint:check --workspace frpc && npm run test:unit && npm run build --workspace frps && npm run build --workspace frpc
|
|
||||||
|
|
||||||
frps-web:
|
frps-web:
|
||||||
$(MAKE) -C web/frps build
|
$(MAKE) -C web/frps build
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ all: build
|
|||||||
build: app
|
build: app
|
||||||
|
|
||||||
app:
|
app:
|
||||||
@set -e; $(foreach n, $(os-archs), \
|
@$(foreach n, $(os-archs), \
|
||||||
os=$(shell echo "$(n)" | cut -d : -f 1); \
|
os=$(shell echo "$(n)" | cut -d : -f 1); \
|
||||||
arch=$(shell echo "$(n)" | cut -d : -f 2); \
|
arch=$(shell echo "$(n)" | cut -d : -f 2); \
|
||||||
extra=$(shell echo "$(n)" | cut -d : -f 3); \
|
extra=$(shell echo "$(n)" | cut -d : -f 3); \
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
[](https://circleci.com/gh/fatedier/frp)
|
[](https://circleci.com/gh/fatedier/frp)
|
||||||
[](https://github.com/fatedier/frp/releases)
|
[](https://github.com/fatedier/frp/releases)
|
||||||
|
[](https://goreportcard.com/report/github.com/fatedier/frp)
|
||||||
[](https://somsubhra.github.io/github-release-stats/?username=fatedier&repository=frp)
|
[](https://somsubhra.github.io/github-release-stats/?username=fatedier&repository=frp)
|
||||||
|
|
||||||
[README](README.md) | [中文文档](README_zh.md)
|
[README](README.md) | [中文文档](README_zh.md)
|
||||||
@@ -12,18 +13,6 @@ frp is an open source project with its ongoing development made possible entirel
|
|||||||
|
|
||||||
<h3 align="center">Gold Sponsors</h3>
|
<h3 align="center">Gold Sponsors</h3>
|
||||||
<!--gold sponsors start-->
|
<!--gold sponsors start-->
|
||||||
<p align="center">
|
|
||||||
<a href="https://www.rapidproxy.io/?ref=frp" target="_blank">
|
|
||||||
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_rapidproxy.png">
|
|
||||||
<br>
|
|
||||||
<b>High-performance residential and ISP proxies for developers</b>
|
|
||||||
</a>
|
|
||||||
<br>
|
|
||||||
<sub>90M+ residential IPs worldwide. Rotating IPs, sticky sessions, and traffic that never expires.</sub>
|
|
||||||
<br>
|
|
||||||
<sub>From $0.55/GB. Use RAPID10 for 10% off. Try it for free.</sub>
|
|
||||||
</p>
|
|
||||||
|
|
||||||
<p align="center">
|
<p align="center">
|
||||||
<a href="https://jb.gg/frp" target="_blank">
|
<a href="https://jb.gg/frp" target="_blank">
|
||||||
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
|
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
|
||||||
@@ -51,7 +40,6 @@ If you're looking for a meeting recording API, consider checking out [Recall.ai]
|
|||||||
an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
|
an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
|
||||||
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!--gold sponsors end-->
|
<!--gold sponsors end-->
|
||||||
|
|
||||||
## What is frp?
|
## What is frp?
|
||||||
|
|||||||
+1
-13
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
[](https://circleci.com/gh/fatedier/frp)
|
[](https://circleci.com/gh/fatedier/frp)
|
||||||
[](https://github.com/fatedier/frp/releases)
|
[](https://github.com/fatedier/frp/releases)
|
||||||
|
[](https://goreportcard.com/report/github.com/fatedier/frp)
|
||||||
[](https://somsubhra.github.io/github-release-stats/?username=fatedier&repository=frp)
|
[](https://somsubhra.github.io/github-release-stats/?username=fatedier&repository=frp)
|
||||||
|
|
||||||
[README](README.md) | [中文文档](README_zh.md)
|
[README](README.md) | [中文文档](README_zh.md)
|
||||||
@@ -14,18 +15,6 @@ frp 是一个完全开源的项目,我们的开发工作完全依靠赞助者
|
|||||||
|
|
||||||
<h3 align="center">Gold Sponsors</h3>
|
<h3 align="center">Gold Sponsors</h3>
|
||||||
<!--gold sponsors start-->
|
<!--gold sponsors start-->
|
||||||
<p align="center">
|
|
||||||
<a href="https://www.rapidproxy.io/?ref=frp" target="_blank">
|
|
||||||
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_rapidproxy.png">
|
|
||||||
<br>
|
|
||||||
<b>High-performance residential and ISP proxies for developers</b>
|
|
||||||
</a>
|
|
||||||
<br>
|
|
||||||
<sub>90M+ residential IPs worldwide. Rotating IPs, sticky sessions, and traffic that never expires.</sub>
|
|
||||||
<br>
|
|
||||||
<sub>From $0.55/GB. Use RAPID10 for 10% off. Try it for free.</sub>
|
|
||||||
</p>
|
|
||||||
|
|
||||||
<p align="center">
|
<p align="center">
|
||||||
<a href="https://jb.gg/frp" target="_blank">
|
<a href="https://jb.gg/frp" target="_blank">
|
||||||
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
|
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
|
||||||
@@ -53,7 +42,6 @@ If you're looking for a meeting recording API, consider checking out [Recall.ai]
|
|||||||
an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
|
an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
|
||||||
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!--gold sponsors end-->
|
<!--gold sponsors end-->
|
||||||
|
|
||||||
## 为什么使用 frp ?
|
## 为什么使用 frp ?
|
||||||
|
|||||||
+8
-3
@@ -1,4 +1,9 @@
|
|||||||
## Fixes
|
## Features
|
||||||
|
|
||||||
* Fixed VirtualNet route lifecycle issues during reconnect and shutdown, including stale route cleanup, shutdown races, and reconnect backoff overflow.
|
* `transport.wireProtocol = "v2"` now also applies to UDP-based proxy payloads, including ordinary UDP and SUDP, so their payload framing is consistent with the selected wire protocol.
|
||||||
* Fixed health check failure counts not resetting after a successful check, ensuring `healthCheck.maxFailed` applies to consecutive failures.
|
* Improved SUDP compatibility during mixed `transport.wireProtocol` deployments, allowing frps to bridge payloads between v1/default and v2 SUDP clients.
|
||||||
|
* XTCP work connection `NatHoleSid` messages now follow the selected `transport.wireProtocol`.
|
||||||
|
|
||||||
|
## Compatibility Notes
|
||||||
|
|
||||||
|
* When enabling `transport.wireProtocol = "v2"` for SUDP, upgrade both the proxy and visitor frpc instances first, or keep them on `v1` until both sides are upgraded.
|
||||||
|
|||||||
@@ -2,16 +2,12 @@ package client
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"os"
|
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/fatedier/frp/client/configmgmt"
|
"github.com/fatedier/frp/client/configmgmt"
|
||||||
"github.com/fatedier/frp/pkg/config/source"
|
"github.com/fatedier/frp/pkg/config/source"
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
v1 "github.com/fatedier/frp/pkg/config/v1"
|
||||||
"github.com/fatedier/frp/pkg/policy/security"
|
|
||||||
"github.com/fatedier/frp/pkg/vnet"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func newTestRawTCPProxyConfig(name string) *v1.TCPProxyConfig {
|
func newTestRawTCPProxyConfig(name string) *v1.TCPProxyConfig {
|
||||||
@@ -26,256 +22,6 @@ func newTestRawTCPProxyConfig(name string) *v1.TCPProxyConfig {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func newTestVirtualNetProxyConfig(name string) *v1.STCPProxyConfig {
|
|
||||||
return &v1.STCPProxyConfig{
|
|
||||||
ProxyBaseConfig: v1.ProxyBaseConfig{
|
|
||||||
Name: name,
|
|
||||||
Type: "stcp",
|
|
||||||
ProxyBackend: v1.ProxyBackend{
|
|
||||||
Plugin: v1.TypedClientPluginOptions{
|
|
||||||
Type: v1.PluginVirtualNet,
|
|
||||||
ClientPluginOptions: &v1.VirtualNetPluginOptions{Type: v1.PluginVirtualNet},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTestVirtualNetVisitorConfig(name string) *v1.STCPVisitorConfig {
|
|
||||||
return &v1.STCPVisitorConfig{
|
|
||||||
VisitorBaseConfig: v1.VisitorBaseConfig{
|
|
||||||
Name: name,
|
|
||||||
Type: "stcp",
|
|
||||||
ServerName: "vnet-server",
|
|
||||||
SecretKey: "secret",
|
|
||||||
BindPort: -1,
|
|
||||||
Plugin: v1.TypedVisitorPluginOptions{
|
|
||||||
Type: v1.VisitorPluginVirtualNet,
|
|
||||||
VisitorPluginOptions: &v1.VirtualNetVisitorPluginOptions{
|
|
||||||
Type: v1.VisitorPluginVirtualNet,
|
|
||||||
DestinationIP: "100.86.0.1",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceConfigManagerReloadVirtualNetRuntimeDependency(t *testing.T) {
|
|
||||||
const runtimeErr = "VirtualNet-dependent configuration requires a VirtualNet runtime enabled at startup"
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
startupVirtualNetAddr string
|
|
||||||
nextConfig string
|
|
||||||
wantRuntimeDependency bool
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "unrelated common config",
|
|
||||||
nextConfig: `serverAddr = "0.0.0.0"`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "VirtualNet address without startup runtime",
|
|
||||||
nextConfig: `featureGates = { VirtualNet = true }
|
|
||||||
virtualNet.address = "100.86.0.4/24"
|
|
||||||
`,
|
|
||||||
wantRuntimeDependency: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "VirtualNet proxy without startup runtime",
|
|
||||||
nextConfig: `[[proxies]]
|
|
||||||
name = "vnet-proxy"
|
|
||||||
type = "stcp"
|
|
||||||
secretKey = "secret"
|
|
||||||
[proxies.plugin]
|
|
||||||
type = "virtual_net"
|
|
||||||
`,
|
|
||||||
wantRuntimeDependency: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "VirtualNet visitor without startup runtime",
|
|
||||||
nextConfig: `[[visitors]]
|
|
||||||
name = "vnet-visitor"
|
|
||||||
type = "stcp"
|
|
||||||
serverName = "vnet-server"
|
|
||||||
secretKey = "secret"
|
|
||||||
bindPort = -1
|
|
||||||
[visitors.plugin]
|
|
||||||
type = "virtual_net"
|
|
||||||
destinationIP = "100.86.0.1"
|
|
||||||
`,
|
|
||||||
wantRuntimeDependency: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "existing VirtualNet startup runtime",
|
|
||||||
startupVirtualNetAddr: "100.86.0.4/24",
|
|
||||||
nextConfig: `featureGates = { VirtualNet = true }
|
|
||||||
virtualNet.address = "100.86.0.5/24"
|
|
||||||
`,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
current := &v1.ClientCommonConfig{}
|
|
||||||
if tc.startupVirtualNetAddr != "" {
|
|
||||||
current.FeatureGates = map[string]bool{"VirtualNet": true}
|
|
||||||
current.VirtualNet.Address = tc.startupVirtualNetAddr
|
|
||||||
}
|
|
||||||
if err := current.Complete(); err != nil {
|
|
||||||
t.Fatalf("complete current config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
configFile := filepath.Join(t.TempDir(), "frpc.toml")
|
|
||||||
if err := os.WriteFile(configFile, []byte(tc.nextConfig), 0o600); err != nil {
|
|
||||||
t.Fatalf("write config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
configSource := source.NewConfigSource()
|
|
||||||
aggregator := source.NewAggregator(configSource)
|
|
||||||
svr := &Service{
|
|
||||||
common: current,
|
|
||||||
reloadCommon: current,
|
|
||||||
configFilePath: configFile,
|
|
||||||
unsafeFeatures: security.NewUnsafeFeatures(nil),
|
|
||||||
aggregator: aggregator,
|
|
||||||
configSource: configSource,
|
|
||||||
}
|
|
||||||
if tc.startupVirtualNetAddr != "" {
|
|
||||||
svr.vnetController = vnet.NewController(current.VirtualNet)
|
|
||||||
}
|
|
||||||
|
|
||||||
err := (&serviceConfigManager{svr: svr}).ReloadFromFile(true)
|
|
||||||
if tc.wantRuntimeDependency {
|
|
||||||
if !errors.Is(err, configmgmt.ErrApplyConfig) || !strings.Contains(err.Error(), runtimeErr) {
|
|
||||||
t.Fatalf("expected VirtualNet runtime dependency error, got: %v", err)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("reload config: %v", err)
|
|
||||||
}
|
|
||||||
if svr.common != current {
|
|
||||||
t.Fatal("reload should not replace startup common config")
|
|
||||||
}
|
|
||||||
if tc.startupVirtualNetAddr == "" && svr.vnetController != nil {
|
|
||||||
t.Fatal("reload should not enable startup-only VirtualNet runtime state")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceConfigManagerReloadVirtualNetRuntimeDependencyUsesMergedSources(t *testing.T) {
|
|
||||||
const runtimeErr = "VirtualNet-dependent configuration requires a VirtualNet runtime enabled at startup"
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
nextConfig string
|
|
||||||
storeProxy v1.ProxyConfigurer
|
|
||||||
storeVisitor v1.VisitorConfigurer
|
|
||||||
wantRuntimeDependency bool
|
|
||||||
wantProxyPlugin string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "Store VirtualNet proxy is rejected",
|
|
||||||
nextConfig: `serverAddr = "0.0.0.0"`,
|
|
||||||
storeProxy: newTestVirtualNetProxyConfig("store-vnet"),
|
|
||||||
wantRuntimeDependency: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "Store VirtualNet visitor is rejected",
|
|
||||||
nextConfig: `serverAddr = "0.0.0.0"`,
|
|
||||||
storeVisitor: newTestVirtualNetVisitorConfig("store-vnet"),
|
|
||||||
wantRuntimeDependency: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "Store VirtualNet proxy overrides file proxy",
|
|
||||||
nextConfig: `[[proxies]]
|
|
||||||
name = "shared"
|
|
||||||
type = "tcp"
|
|
||||||
localPort = 10080
|
|
||||||
remotePort = 10081
|
|
||||||
`,
|
|
||||||
storeProxy: newTestVirtualNetProxyConfig("shared"),
|
|
||||||
wantRuntimeDependency: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "Store non-VirtualNet proxy overrides file VirtualNet proxy",
|
|
||||||
nextConfig: `[[proxies]]
|
|
||||||
name = "shared"
|
|
||||||
type = "stcp"
|
|
||||||
secretKey = "secret"
|
|
||||||
[proxies.plugin]
|
|
||||||
type = "virtual_net"
|
|
||||||
`,
|
|
||||||
storeProxy: newTestRawTCPProxyConfig("shared"),
|
|
||||||
wantProxyPlugin: "",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
current := &v1.ClientCommonConfig{}
|
|
||||||
if err := current.Complete(); err != nil {
|
|
||||||
t.Fatalf("complete current config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
configFile := filepath.Join(t.TempDir(), "frpc.toml")
|
|
||||||
if err := os.WriteFile(configFile, []byte(tc.nextConfig), 0o600); err != nil {
|
|
||||||
t.Fatalf("write config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
storeSource, err := source.NewStoreSource(source.StoreSourceConfig{
|
|
||||||
Path: filepath.Join(t.TempDir(), "store.json"),
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("new store source: %v", err)
|
|
||||||
}
|
|
||||||
if tc.storeProxy != nil {
|
|
||||||
if err := storeSource.AddProxy(tc.storeProxy); err != nil {
|
|
||||||
t.Fatalf("add store proxy: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if tc.storeVisitor != nil {
|
|
||||||
if err := storeSource.AddVisitor(tc.storeVisitor); err != nil {
|
|
||||||
t.Fatalf("add store visitor: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
configSource := source.NewConfigSource()
|
|
||||||
aggregator := source.NewAggregator(configSource)
|
|
||||||
aggregator.SetStoreSource(storeSource)
|
|
||||||
svr := &Service{
|
|
||||||
common: current,
|
|
||||||
reloadCommon: current,
|
|
||||||
configFilePath: configFile,
|
|
||||||
unsafeFeatures: security.NewUnsafeFeatures(nil),
|
|
||||||
aggregator: aggregator,
|
|
||||||
configSource: configSource,
|
|
||||||
storeSource: storeSource,
|
|
||||||
}
|
|
||||||
|
|
||||||
err = (&serviceConfigManager{svr: svr}).ReloadFromFile(true)
|
|
||||||
if tc.wantRuntimeDependency {
|
|
||||||
if !errors.Is(err, configmgmt.ErrApplyConfig) || !strings.Contains(err.Error(), runtimeErr) {
|
|
||||||
t.Fatalf("expected VirtualNet runtime dependency error, got: %v", err)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("reload config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(svr.proxyCfgs) != 1 {
|
|
||||||
t.Fatalf("expected one applied proxy, got %d", len(svr.proxyCfgs))
|
|
||||||
}
|
|
||||||
if got := svr.proxyCfgs[0].GetBaseConfig().Plugin.Type; got != tc.wantProxyPlugin {
|
|
||||||
t.Fatalf("unexpected applied proxy plugin: %q", got)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceConfigManagerCreateStoreProxyConflict(t *testing.T) {
|
func TestServiceConfigManagerCreateStoreProxyConflict(t *testing.T) {
|
||||||
storeSource, err := source.NewStoreSource(source.StoreSourceConfig{
|
storeSource, err := source.NewStoreSource(source.StoreSourceConfig{
|
||||||
Path: filepath.Join(t.TempDir(), "store.json"),
|
Path: filepath.Join(t.TempDir(), "store.json"),
|
||||||
|
|||||||
+1
-1
@@ -24,7 +24,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
libnet "github.com/fatedier/golib/net"
|
libnet "github.com/fatedier/golib/net"
|
||||||
fmux "github.com/fatedier/yamux"
|
fmux "github.com/hashicorp/yamux"
|
||||||
quic "github.com/quic-go/quic-go"
|
quic "github.com/quic-go/quic-go"
|
||||||
"github.com/samber/lo"
|
"github.com/samber/lo"
|
||||||
|
|
||||||
|
|||||||
+2
-11
@@ -47,8 +47,6 @@ type SessionContext struct {
|
|||||||
Connector MessageConnector
|
Connector MessageConnector
|
||||||
// Virtual net controller
|
// Virtual net controller
|
||||||
VnetController *vnet.Controller
|
VnetController *vnet.Controller
|
||||||
// UDPPacketCodec is immutable for the lifetime of this negotiated session.
|
|
||||||
UDPPacketCodec string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type Control struct {
|
type Control struct {
|
||||||
@@ -94,16 +92,9 @@ func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, erro
|
|||||||
ctl.registerMsgHandlers()
|
ctl.registerMsgHandlers()
|
||||||
ctl.msgTransporter = transport.NewMessageTransporter(ctl.msgDispatcher)
|
ctl.msgTransporter = transport.NewMessageTransporter(ctl.msgDispatcher)
|
||||||
|
|
||||||
ctl.pm = proxy.NewManager(
|
ctl.pm = proxy.NewManager(ctl.ctx, sessionCtx.Common, sessionCtx.Auth.EncryptionKey(), ctl.msgTransporter, sessionCtx.VnetController)
|
||||||
ctl.ctx,
|
|
||||||
sessionCtx.Common,
|
|
||||||
sessionCtx.Auth.EncryptionKey(),
|
|
||||||
ctl.msgTransporter,
|
|
||||||
sessionCtx.VnetController,
|
|
||||||
sessionCtx.UDPPacketCodec,
|
|
||||||
)
|
|
||||||
ctl.vm = visitor.NewManager(ctl.ctx, sessionCtx.RunID, sessionCtx.Common,
|
ctl.vm = visitor.NewManager(ctl.ctx, sessionCtx.RunID, sessionCtx.Common,
|
||||||
ctl.connectServer, ctl.msgTransporter, sessionCtx.VnetController, sessionCtx.UDPPacketCodec)
|
ctl.connectServer, ctl.msgTransporter, sessionCtx.VnetController)
|
||||||
return ctl, nil
|
return ctl, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -99,7 +99,6 @@ func (d *controlSessionDialer) Dial(previousRunID string) (*SessionContext, erro
|
|||||||
Auth: d.auth,
|
Auth: d.auth,
|
||||||
Connector: newMessageConnector(connector, d.common.Transport.WireProtocol),
|
Connector: newMessageConnector(connector, d.common.Transport.WireProtocol),
|
||||||
VnetController: d.vnetController,
|
VnetController: d.vnetController,
|
||||||
UDPPacketCodec: loginResult.udpPacketCodec,
|
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -128,9 +127,8 @@ func (d *controlSessionDialer) buildLoginMsg(previousRunID string) (*msg.Login,
|
|||||||
}
|
}
|
||||||
|
|
||||||
type loginExchangeResult struct {
|
type loginExchangeResult struct {
|
||||||
resp *msg.LoginResp
|
resp *msg.LoginResp
|
||||||
crypto *wire.CryptoContext
|
crypto *wire.CryptoContext
|
||||||
udpPacketCodec string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) (*loginExchangeResult, error) {
|
func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) (*loginExchangeResult, error) {
|
||||||
@@ -174,7 +172,6 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login)
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
var cryptoContext *wire.CryptoContext
|
var cryptoContext *wire.CryptoContext
|
||||||
var udpPacketCodec string
|
|
||||||
if wireConn != nil {
|
if wireConn != nil {
|
||||||
serverHelloFrame, err := wireConn.ReadFrame()
|
serverHelloFrame, err := wireConn.ReadFrame()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -194,7 +191,6 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login)
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
udpPacketCodec = serverHello.Selected.Message.UDPPacketCodec
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var loginRespMsg msg.LoginResp
|
var loginRespMsg msg.LoginResp
|
||||||
@@ -202,9 +198,8 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login)
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &loginExchangeResult{
|
return &loginExchangeResult{
|
||||||
resp: &loginRespMsg,
|
resp: &loginRespMsg,
|
||||||
crypto: cryptoContext,
|
crypto: cryptoContext,
|
||||||
udpPacketCodec: udpPacketCodec,
|
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -117,7 +117,6 @@ func TestControlSessionDialerDialV1(t *testing.T) {
|
|||||||
defer sessionCtx.Connector.Close()
|
defer sessionCtx.Connector.Close()
|
||||||
|
|
||||||
require.Equal(t, "run-v1", sessionCtx.RunID)
|
require.Equal(t, "run-v1", sessionCtx.RunID)
|
||||||
require.Empty(t, sessionCtx.UDPPacketCodec)
|
|
||||||
require.NotNil(t, sessionCtx.Conn)
|
require.NotNil(t, sessionCtx.Conn)
|
||||||
require.NotNil(t, sessionCtx.Connector)
|
require.NotNil(t, sessionCtx.Connector)
|
||||||
require.False(t, connector.closed.Load())
|
require.False(t, connector.closed.Load())
|
||||||
@@ -226,7 +225,6 @@ func TestControlSessionDialerDialV2(t *testing.T) {
|
|||||||
defer sessionCtx.Connector.Close()
|
defer sessionCtx.Connector.Close()
|
||||||
|
|
||||||
require.Equal(t, "run-v2", sessionCtx.RunID)
|
require.Equal(t, "run-v2", sessionCtx.RunID)
|
||||||
require.Equal(t, wire.UDPPacketCodecBinary, sessionCtx.UDPPacketCodec)
|
|
||||||
require.NotNil(t, sessionCtx.Conn)
|
require.NotNil(t, sessionCtx.Conn)
|
||||||
require.NotNil(t, sessionCtx.Connector)
|
require.NotNil(t, sessionCtx.Connector)
|
||||||
require.False(t, connector.closed.Load())
|
require.False(t, connector.closed.Load())
|
||||||
|
|||||||
@@ -1,125 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
//go:build !frps
|
|
||||||
|
|
||||||
package client
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/binary"
|
|
||||||
"net"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
clientproxy "github.com/fatedier/frp/client/proxy"
|
|
||||||
"github.com/fatedier/frp/pkg/auth"
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
|
||||||
"github.com/fatedier/frp/pkg/proto/wire"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestControlPropagatesBinaryUDPPacketCodecToWorkConn(t *testing.T) {
|
|
||||||
echoConn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = echoConn.Close() })
|
|
||||||
|
|
||||||
echoDone := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
buf := make([]byte, 64)
|
|
||||||
n, addr, err := echoConn.ReadFromUDP(buf)
|
|
||||||
if err == nil {
|
|
||||||
_, err = echoConn.WriteToUDP(buf[:n], addr)
|
|
||||||
}
|
|
||||||
echoDone <- err
|
|
||||||
}()
|
|
||||||
|
|
||||||
authRuntime, err := auth.BuildClientAuth(&v1.AuthClientConfig{
|
|
||||||
Method: v1.AuthMethodToken,
|
|
||||||
Token: "token",
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
controlConn, controlPeer := net.Pipe()
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = controlConn.Close()
|
|
||||||
_ = controlPeer.Close()
|
|
||||||
})
|
|
||||||
common := &v1.ClientCommonConfig{
|
|
||||||
Transport: v1.ClientTransportConfig{WireProtocol: wire.ProtocolV2},
|
|
||||||
UDPPacketSize: 1500,
|
|
||||||
}
|
|
||||||
ctl, err := NewControl(context.Background(), &SessionContext{
|
|
||||||
Common: common,
|
|
||||||
RunID: "binary-udp-test",
|
|
||||||
Conn: msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn)),
|
|
||||||
Auth: authRuntime,
|
|
||||||
UDPPacketCodec: wire.UDPPacketCodecBinary,
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(ctl.pm.Close)
|
|
||||||
|
|
||||||
echoAddr := echoConn.LocalAddr().(*net.UDPAddr)
|
|
||||||
proxyCfg := &v1.UDPProxyConfig{
|
|
||||||
ProxyBaseConfig: v1.ProxyBaseConfig{
|
|
||||||
Name: "udp",
|
|
||||||
Type: string(v1.ProxyTypeUDP),
|
|
||||||
ProxyBackend: v1.ProxyBackend{
|
|
||||||
LocalIP: "127.0.0.1",
|
|
||||||
LocalPort: echoAddr.Port,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
ctl.pm.UpdateAll([]v1.ProxyConfigurer{proxyCfg})
|
|
||||||
require.Eventually(t, func() bool {
|
|
||||||
status, ok := ctl.pm.GetProxyStatus("udp")
|
|
||||||
return ok && status.Phase == clientproxy.ProxyPhaseWaitStart
|
|
||||||
}, time.Second, 10*time.Millisecond)
|
|
||||||
require.NoError(t, ctl.pm.StartProxy("udp", "", ""))
|
|
||||||
|
|
||||||
workClient, workServer := net.Pipe()
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = workClient.Close()
|
|
||||||
_ = workServer.Close()
|
|
||||||
})
|
|
||||||
deadline := time.Now().Add(3 * time.Second)
|
|
||||||
require.NoError(t, workClient.SetDeadline(deadline))
|
|
||||||
require.NoError(t, workServer.SetDeadline(deadline))
|
|
||||||
ctl.pm.HandleWorkConn("udp", workClient, &msg.StartWorkConn{ProxyName: "udp"})
|
|
||||||
|
|
||||||
serverRW, err := msg.NewUDPPacketReadWriter(workServer, wire.ProtocolV2, wire.UDPPacketCodecBinary)
|
|
||||||
require.NoError(t, err)
|
|
||||||
writeDone := make(chan error, 1)
|
|
||||||
in := &msg.UDPPacket{
|
|
||||||
Content: []byte("binary udp"),
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 12345},
|
|
||||||
}
|
|
||||||
go func() {
|
|
||||||
writeDone <- serverRW.WriteMsg(in)
|
|
||||||
}()
|
|
||||||
|
|
||||||
frame, err := wire.NewConn(workServer).ReadFrame()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, wire.FrameTypeMessage, frame.Type)
|
|
||||||
require.GreaterOrEqual(t, len(frame.Payload), 2)
|
|
||||||
require.Equal(t, msg.V2TypeUDPPacketBinary, binary.BigEndian.Uint16(frame.Payload[:2]))
|
|
||||||
out, err := msg.DecodeUDPPacketBinary(frame.Payload[2:])
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, in.Content, out.Content)
|
|
||||||
require.Equal(t, in.RemoteAddr.String(), out.RemoteAddr.String())
|
|
||||||
require.NoError(t, <-writeDone)
|
|
||||||
require.NoError(t, <-echoDone)
|
|
||||||
}
|
|
||||||
+22
-59
@@ -30,11 +30,6 @@ import (
|
|||||||
|
|
||||||
var ErrHealthCheckType = errors.New("error health check type")
|
var ErrHealthCheckType = errors.New("error health check type")
|
||||||
|
|
||||||
func newHealthTimer(interval time.Duration) (<-chan time.Time, func()) {
|
|
||||||
timer := time.NewTimer(interval)
|
|
||||||
return timer.C, func() { timer.Stop() }
|
|
||||||
}
|
|
||||||
|
|
||||||
type Monitor struct {
|
type Monitor struct {
|
||||||
checkType string
|
checkType string
|
||||||
interval time.Duration
|
interval time.Duration
|
||||||
@@ -54,9 +49,6 @@ type Monitor struct {
|
|||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
doneCh chan struct{}
|
|
||||||
|
|
||||||
timerFactory func(time.Duration) (<-chan time.Time, func())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMonitor(ctx context.Context, cfg v1.HealthCheckConfig, addr string,
|
func NewMonitor(ctx context.Context, cfg v1.HealthCheckConfig, addr string,
|
||||||
@@ -99,8 +91,6 @@ func NewMonitor(ctx context.Context, cfg v1.HealthCheckConfig, addr string,
|
|||||||
statusFailedFn: statusFailedFn,
|
statusFailedFn: statusFailedFn,
|
||||||
ctx: newctx,
|
ctx: newctx,
|
||||||
cancel: cancel,
|
cancel: cancel,
|
||||||
doneCh: make(chan struct{}),
|
|
||||||
timerFactory: newHealthTimer,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -112,66 +102,39 @@ func (monitor *Monitor) Stop() {
|
|||||||
monitor.cancel()
|
monitor.cancel()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Done is closed when the worker launched by Start has exited.
|
|
||||||
func (monitor *Monitor) Done() <-chan struct{} {
|
|
||||||
return monitor.doneCh
|
|
||||||
}
|
|
||||||
|
|
||||||
func (monitor *Monitor) checkWorker() {
|
func (monitor *Monitor) checkWorker() {
|
||||||
defer close(monitor.doneCh)
|
xl := xlog.FromContextSafe(monitor.ctx)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
if monitor.ctx.Err() != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
doCtx, cancel := context.WithDeadline(monitor.ctx, time.Now().Add(monitor.timeout))
|
doCtx, cancel := context.WithDeadline(monitor.ctx, time.Now().Add(monitor.timeout))
|
||||||
err := monitor.doCheck(doCtx)
|
err := monitor.doCheck(doCtx)
|
||||||
cancel()
|
|
||||||
|
|
||||||
// check if this monitor has been closed
|
// check if this monitor has been closed
|
||||||
if monitor.ctx.Err() != nil {
|
select {
|
||||||
|
case <-monitor.ctx.Done():
|
||||||
|
cancel()
|
||||||
return
|
return
|
||||||
|
default:
|
||||||
|
cancel()
|
||||||
}
|
}
|
||||||
monitor.handleCheckResult(err)
|
|
||||||
|
|
||||||
if !monitor.waitForNextCheck() {
|
if err == nil {
|
||||||
return
|
xl.Tracef("do one health check success")
|
||||||
|
if !monitor.statusOK && monitor.statusNormalFn != nil {
|
||||||
|
xl.Infof("health check status change to success")
|
||||||
|
monitor.statusOK = true
|
||||||
|
monitor.statusNormalFn()
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
xl.Warnf("do one health check failed: %v", err)
|
||||||
|
monitor.failedTimes++
|
||||||
|
if monitor.statusOK && int(monitor.failedTimes) >= monitor.maxFailedTimes && monitor.statusFailedFn != nil {
|
||||||
|
xl.Warnf("health check status change to failed")
|
||||||
|
monitor.statusOK = false
|
||||||
|
monitor.statusFailedFn()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (monitor *Monitor) handleCheckResult(err error) {
|
time.Sleep(monitor.interval)
|
||||||
xl := xlog.FromContextSafe(monitor.ctx)
|
|
||||||
if err == nil {
|
|
||||||
xl.Tracef("do one health check success")
|
|
||||||
monitor.failedTimes = 0
|
|
||||||
if !monitor.statusOK && monitor.statusNormalFn != nil {
|
|
||||||
xl.Infof("health check status change to success")
|
|
||||||
monitor.statusOK = true
|
|
||||||
monitor.statusNormalFn()
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
xl.Warnf("do one health check failed: %v", err)
|
|
||||||
monitor.failedTimes++
|
|
||||||
if monitor.statusOK && int(monitor.failedTimes) >= monitor.maxFailedTimes && monitor.statusFailedFn != nil {
|
|
||||||
xl.Warnf("health check status change to failed")
|
|
||||||
monitor.statusOK = false
|
|
||||||
monitor.statusFailedFn()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (monitor *Monitor) waitForNextCheck() bool {
|
|
||||||
timerC, stopTimer := monitor.timerFactory(monitor.interval)
|
|
||||||
defer stopTimer()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-monitor.ctx.Done():
|
|
||||||
return false
|
|
||||||
case <-timerC:
|
|
||||||
return monitor.ctx.Err() == nil
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,575 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package health
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
)
|
|
||||||
|
|
||||||
type tcpHealthBackend struct {
|
|
||||||
listener net.Listener
|
|
||||||
accepted chan struct{}
|
|
||||||
done chan struct{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTCPHealthBackend(t *testing.T, addr string, accepted chan struct{}) *tcpHealthBackend {
|
|
||||||
listener, err := net.Listen("tcp", addr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
backend := &tcpHealthBackend{
|
|
||||||
listener: listener,
|
|
||||||
accepted: accepted,
|
|
||||||
done: make(chan struct{}),
|
|
||||||
}
|
|
||||||
go func() {
|
|
||||||
defer close(backend.done)
|
|
||||||
for {
|
|
||||||
conn, err := backend.listener.Accept()
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
_ = conn.Close()
|
|
||||||
select {
|
|
||||||
case backend.accepted <- struct{}{}:
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
return backend
|
|
||||||
}
|
|
||||||
|
|
||||||
func (backend *tcpHealthBackend) Close() {
|
|
||||||
_ = backend.listener.Close()
|
|
||||||
<-backend.done
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitorConsecutiveFailureWindows(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
setup func(*testing.T, func(), func()) (*Monitor, func(bool))
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "HTTP",
|
|
||||||
setup: func(t *testing.T, normalFn, failedFn func()) (*Monitor, func(bool)) {
|
|
||||||
var healthy atomic.Bool
|
|
||||||
healthy.Store(true)
|
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
if !healthy.Load() {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
}
|
|
||||||
}))
|
|
||||||
t.Cleanup(server.Close)
|
|
||||||
|
|
||||||
monitor := NewMonitor(
|
|
||||||
context.Background(),
|
|
||||||
v1.HealthCheckConfig{
|
|
||||||
Type: "http",
|
|
||||||
Path: "/health",
|
|
||||||
TimeoutSeconds: 1,
|
|
||||||
IntervalSeconds: 1,
|
|
||||||
MaxFailed: 3,
|
|
||||||
},
|
|
||||||
strings.TrimPrefix(server.URL, "http://"),
|
|
||||||
normalFn,
|
|
||||||
failedFn,
|
|
||||||
)
|
|
||||||
return monitor, healthy.Store
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "TCP",
|
|
||||||
setup: func(t *testing.T, normalFn, failedFn func()) (*Monitor, func(bool)) {
|
|
||||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = listener.Close() })
|
|
||||||
go func() {
|
|
||||||
for {
|
|
||||||
conn, err := listener.Accept()
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
_ = conn.Close()
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
monitor := NewMonitor(
|
|
||||||
context.Background(),
|
|
||||||
v1.HealthCheckConfig{
|
|
||||||
Type: "tcp",
|
|
||||||
TimeoutSeconds: 1,
|
|
||||||
IntervalSeconds: 1,
|
|
||||||
MaxFailed: 3,
|
|
||||||
},
|
|
||||||
listener.Addr().String(),
|
|
||||||
normalFn,
|
|
||||||
failedFn,
|
|
||||||
)
|
|
||||||
healthyAddr := monitor.addr
|
|
||||||
return monitor, func(healthy bool) {
|
|
||||||
if healthy {
|
|
||||||
monitor.addr = healthyAddr
|
|
||||||
} else {
|
|
||||||
monitor.addr = "127.0.0.1:0"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, test := range tests {
|
|
||||||
t.Run(test.name, func(t *testing.T) {
|
|
||||||
var events []string
|
|
||||||
monitor, setHealthy := test.setup(
|
|
||||||
t,
|
|
||||||
func() { events = append(events, "normal") },
|
|
||||||
func() { events = append(events, "failed") },
|
|
||||||
)
|
|
||||||
t.Cleanup(monitor.Stop)
|
|
||||||
|
|
||||||
runCheck := func(healthy bool) {
|
|
||||||
t.Helper()
|
|
||||||
setHealthy(healthy)
|
|
||||||
ctx, cancel := context.WithTimeout(monitor.ctx, time.Second)
|
|
||||||
err := monitor.doCheck(ctx)
|
|
||||||
cancel()
|
|
||||||
if healthy {
|
|
||||||
require.NoError(t, err)
|
|
||||||
} else {
|
|
||||||
require.Error(t, err)
|
|
||||||
}
|
|
||||||
monitor.handleCheckResult(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
runCheck(true)
|
|
||||||
require.True(t, monitor.statusOK)
|
|
||||||
require.Zero(t, monitor.failedTimes)
|
|
||||||
require.Equal(t, []string{"normal"}, events)
|
|
||||||
|
|
||||||
runCheck(false)
|
|
||||||
runCheck(false)
|
|
||||||
require.True(t, monitor.statusOK)
|
|
||||||
require.Equal(t, uint64(2), monitor.failedTimes)
|
|
||||||
require.Equal(t, []string{"normal"}, events)
|
|
||||||
|
|
||||||
runCheck(true)
|
|
||||||
require.True(t, monitor.statusOK)
|
|
||||||
require.Zero(t, monitor.failedTimes)
|
|
||||||
require.Equal(t, []string{"normal"}, events)
|
|
||||||
|
|
||||||
runCheck(false)
|
|
||||||
require.True(t, monitor.statusOK)
|
|
||||||
require.Equal(t, uint64(1), monitor.failedTimes)
|
|
||||||
require.Equal(t, []string{"normal"}, events)
|
|
||||||
|
|
||||||
runCheck(false)
|
|
||||||
require.True(t, monitor.statusOK)
|
|
||||||
require.Equal(t, uint64(2), monitor.failedTimes)
|
|
||||||
|
|
||||||
runCheck(false)
|
|
||||||
require.False(t, monitor.statusOK)
|
|
||||||
require.Equal(t, uint64(3), monitor.failedTimes)
|
|
||||||
require.Equal(t, []string{"normal", "failed"}, events)
|
|
||||||
|
|
||||||
runCheck(true)
|
|
||||||
require.True(t, monitor.statusOK)
|
|
||||||
require.Zero(t, monitor.failedTimes)
|
|
||||||
require.Equal(t, []string{"normal", "failed", "normal"}, events)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type workerStatus struct {
|
|
||||||
failedTimes uint64
|
|
||||||
statusOK bool
|
|
||||||
}
|
|
||||||
|
|
||||||
type workerInterval struct {
|
|
||||||
status workerStatus
|
|
||||||
timer chan time.Time
|
|
||||||
}
|
|
||||||
|
|
||||||
func observeWorkerIntervals(monitor *Monitor) <-chan workerInterval {
|
|
||||||
intervals := make(chan workerInterval, 1)
|
|
||||||
monitor.timerFactory = func(time.Duration) (<-chan time.Time, func()) {
|
|
||||||
timer := make(chan time.Time, 1)
|
|
||||||
// The worker takes this snapshot after handleCheckResult. Sending it
|
|
||||||
// publishes the state to the test; no test goroutine reads live fields.
|
|
||||||
interval := workerInterval{
|
|
||||||
status: workerStatus{failedTimes: monitor.failedTimes, statusOK: monitor.statusOK},
|
|
||||||
timer: timer,
|
|
||||||
}
|
|
||||||
select {
|
|
||||||
case intervals <- interval:
|
|
||||||
case <-monitor.ctx.Done():
|
|
||||||
}
|
|
||||||
return timer, func() {}
|
|
||||||
}
|
|
||||||
return intervals
|
|
||||||
}
|
|
||||||
|
|
||||||
func awaitWorkerInterval(t *testing.T, intervals <-chan workerInterval, want workerStatus) chan time.Time {
|
|
||||||
t.Helper()
|
|
||||||
select {
|
|
||||||
case interval := <-intervals:
|
|
||||||
require.Equal(t, want, interval.status)
|
|
||||||
return interval.timer
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("health worker did not reach the interval barrier")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func stopMonitorWorker(t *testing.T, monitor *Monitor) {
|
|
||||||
t.Helper()
|
|
||||||
monitor.Stop()
|
|
||||||
select {
|
|
||||||
case <-monitor.Done():
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Error("health worker did not exit after Stop")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitorWorkerProcessesResults(t *testing.T) {
|
|
||||||
requestReady := make(chan struct{}, 1)
|
|
||||||
responses := make(chan int)
|
|
||||||
requestCanceled := make(chan struct{})
|
|
||||||
var cancelOnce sync.Once
|
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
select {
|
|
||||||
case requestReady <- struct{}{}:
|
|
||||||
case <-r.Context().Done():
|
|
||||||
return
|
|
||||||
}
|
|
||||||
select {
|
|
||||||
case code := <-responses:
|
|
||||||
w.WriteHeader(code)
|
|
||||||
case <-r.Context().Done():
|
|
||||||
cancelOnce.Do(func() { close(requestCanceled) })
|
|
||||||
}
|
|
||||||
}))
|
|
||||||
t.Cleanup(server.Close)
|
|
||||||
|
|
||||||
events := make(chan string, 3)
|
|
||||||
monitor := NewMonitor(
|
|
||||||
context.Background(),
|
|
||||||
v1.HealthCheckConfig{
|
|
||||||
Type: "http",
|
|
||||||
Path: "/health",
|
|
||||||
TimeoutSeconds: 5,
|
|
||||||
MaxFailed: 3,
|
|
||||||
},
|
|
||||||
strings.TrimPrefix(server.URL, "http://"),
|
|
||||||
func() { events <- "normal" },
|
|
||||||
func() { events <- "failed" },
|
|
||||||
)
|
|
||||||
intervals := observeWorkerIntervals(monitor)
|
|
||||||
t.Cleanup(func() { stopMonitorWorker(t, monitor) })
|
|
||||||
|
|
||||||
awaitRequest := func() {
|
|
||||||
t.Helper()
|
|
||||||
select {
|
|
||||||
case <-requestReady:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("health check request did not start")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
awaitEvent := func(want string) {
|
|
||||||
t.Helper()
|
|
||||||
select {
|
|
||||||
case got := <-events:
|
|
||||||
require.Equal(t, want, got)
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatalf("health check callback %q was not called", want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
checkResult := func(code int, want workerStatus, event string) {
|
|
||||||
t.Helper()
|
|
||||||
awaitRequest()
|
|
||||||
select {
|
|
||||||
case responses <- code:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("health check handler did not accept the response")
|
|
||||||
}
|
|
||||||
timer := awaitWorkerInterval(t, intervals, want)
|
|
||||||
if event != "" {
|
|
||||||
awaitEvent(event)
|
|
||||||
}
|
|
||||||
// The worker has finished this result and cannot emit another
|
|
||||||
// callback until the test releases the next check.
|
|
||||||
select {
|
|
||||||
case got := <-events:
|
|
||||||
t.Fatalf("unexpected health check callback %q", got)
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
timer <- time.Now()
|
|
||||||
}
|
|
||||||
|
|
||||||
monitor.Start()
|
|
||||||
checkResult(http.StatusOK, workerStatus{statusOK: true}, "normal")
|
|
||||||
checkResult(http.StatusServiceUnavailable, workerStatus{failedTimes: 1, statusOK: true}, "")
|
|
||||||
checkResult(http.StatusServiceUnavailable, workerStatus{failedTimes: 2, statusOK: true}, "")
|
|
||||||
checkResult(http.StatusOK, workerStatus{statusOK: true}, "")
|
|
||||||
checkResult(http.StatusServiceUnavailable, workerStatus{failedTimes: 1, statusOK: true}, "")
|
|
||||||
checkResult(http.StatusServiceUnavailable, workerStatus{failedTimes: 2, statusOK: true}, "")
|
|
||||||
checkResult(http.StatusServiceUnavailable, workerStatus{failedTimes: 3}, "failed")
|
|
||||||
checkResult(http.StatusOK, workerStatus{statusOK: true}, "normal")
|
|
||||||
|
|
||||||
awaitRequest()
|
|
||||||
stopMonitorWorker(t, monitor)
|
|
||||||
select {
|
|
||||||
case <-requestCanceled:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("health check request context was not canceled")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitorTCPWorkerProcessesResults(t *testing.T) {
|
|
||||||
initialAccepted := make(chan struct{}, 1)
|
|
||||||
initialBackend := newTCPHealthBackend(t, "127.0.0.1:0", initialAccepted)
|
|
||||||
addr := initialBackend.listener.Addr().String()
|
|
||||||
recoveryAccepted := make(chan struct{}, 1)
|
|
||||||
|
|
||||||
normalCallbacks := make(chan workerStatus, 2)
|
|
||||||
failedCallbacks := make(chan workerStatus, 2)
|
|
||||||
var monitor *Monitor
|
|
||||||
monitor = NewMonitor(
|
|
||||||
context.Background(),
|
|
||||||
v1.HealthCheckConfig{
|
|
||||||
Type: "tcp",
|
|
||||||
TimeoutSeconds: 1,
|
|
||||||
MaxFailed: 3,
|
|
||||||
},
|
|
||||||
addr,
|
|
||||||
func() {
|
|
||||||
normalCallbacks <- workerStatus{failedTimes: monitor.failedTimes, statusOK: monitor.statusOK}
|
|
||||||
},
|
|
||||||
func() {
|
|
||||||
failedCallbacks <- workerStatus{failedTimes: monitor.failedTimes, statusOK: monitor.statusOK}
|
|
||||||
},
|
|
||||||
)
|
|
||||||
intervals := observeWorkerIntervals(monitor)
|
|
||||||
recoveryBackend := (*tcpHealthBackend)(nil)
|
|
||||||
t.Cleanup(func() {
|
|
||||||
stopMonitorWorker(t, monitor)
|
|
||||||
initialBackend.Close()
|
|
||||||
if recoveryBackend != nil {
|
|
||||||
recoveryBackend.Close()
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
awaitStatus := func(ch <-chan workerStatus, want workerStatus, message string) {
|
|
||||||
t.Helper()
|
|
||||||
select {
|
|
||||||
case got := <-ch:
|
|
||||||
require.Equal(t, want, got)
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal(message)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
monitor.Start()
|
|
||||||
select {
|
|
||||||
case <-initialAccepted:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("TCP health check did not reach the initial backend")
|
|
||||||
}
|
|
||||||
awaitStatus(normalCallbacks, workerStatus{failedTimes: 0, statusOK: true}, "TCP worker did not report the initial success")
|
|
||||||
|
|
||||||
initialTimer := awaitWorkerInterval(t, intervals, workerStatus{statusOK: true})
|
|
||||||
initialBackend.Close()
|
|
||||||
initialTimer <- time.Now()
|
|
||||||
|
|
||||||
firstFailureTimer := awaitWorkerInterval(t, intervals, workerStatus{failedTimes: 1, statusOK: true})
|
|
||||||
firstFailureTimer <- time.Now()
|
|
||||||
|
|
||||||
secondFailureTimer := awaitWorkerInterval(t, intervals, workerStatus{failedTimes: 2, statusOK: true})
|
|
||||||
secondFailureTimer <- time.Now()
|
|
||||||
|
|
||||||
awaitStatus(failedCallbacks, workerStatus{failedTimes: 3, statusOK: false}, "TCP worker did not report the third failed health check")
|
|
||||||
thirdFailureTimer := awaitWorkerInterval(t, intervals, workerStatus{failedTimes: 3})
|
|
||||||
|
|
||||||
recoveryBackend = newTCPHealthBackend(t, addr, recoveryAccepted)
|
|
||||||
thirdFailureTimer <- time.Now()
|
|
||||||
select {
|
|
||||||
case <-recoveryAccepted:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("TCP health check did not reach the recovery backend")
|
|
||||||
}
|
|
||||||
awaitStatus(normalCallbacks, workerStatus{failedTimes: 0, statusOK: true}, "TCP worker did not report recovery")
|
|
||||||
|
|
||||||
recoveryTimer := awaitWorkerInterval(t, intervals, workerStatus{statusOK: true})
|
|
||||||
recoveryBackend.Close()
|
|
||||||
recoveryTimer <- time.Now()
|
|
||||||
|
|
||||||
firstRecoveryFailureTimer := awaitWorkerInterval(t, intervals, workerStatus{failedTimes: 1, statusOK: true})
|
|
||||||
firstRecoveryFailureTimer <- time.Now()
|
|
||||||
|
|
||||||
secondRecoveryFailureTimer := awaitWorkerInterval(t, intervals, workerStatus{failedTimes: 2, statusOK: true})
|
|
||||||
secondRecoveryFailureTimer <- time.Now()
|
|
||||||
|
|
||||||
awaitStatus(failedCallbacks, workerStatus{failedTimes: 3, statusOK: false}, "TCP worker did not report the third failed health check after recovery")
|
|
||||||
stopMonitorWorker(t, monitor)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitorMaxFailedOne(t *testing.T) {
|
|
||||||
var events []string
|
|
||||||
monitor := NewMonitor(
|
|
||||||
context.Background(),
|
|
||||||
v1.HealthCheckConfig{Type: "tcp", MaxFailed: 1},
|
|
||||||
"",
|
|
||||||
func() { events = append(events, "normal") },
|
|
||||||
func() { events = append(events, "failed") },
|
|
||||||
)
|
|
||||||
t.Cleanup(monitor.Stop)
|
|
||||||
|
|
||||||
checkErr := errors.New("health check failed")
|
|
||||||
monitor.handleCheckResult(nil)
|
|
||||||
monitor.handleCheckResult(checkErr)
|
|
||||||
require.Equal(t, []string{"normal", "failed"}, events)
|
|
||||||
require.False(t, monitor.statusOK)
|
|
||||||
require.Equal(t, uint64(1), monitor.failedTimes)
|
|
||||||
|
|
||||||
monitor.handleCheckResult(checkErr)
|
|
||||||
require.Equal(t, []string{"normal", "failed"}, events)
|
|
||||||
|
|
||||||
monitor.handleCheckResult(nil)
|
|
||||||
monitor.handleCheckResult(checkErr)
|
|
||||||
require.Equal(t, []string{"normal", "failed", "normal", "failed"}, events)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitorStopCancelsWork(t *testing.T) {
|
|
||||||
t.Run("in-flight check", func(t *testing.T) {
|
|
||||||
requestStarted := make(chan struct{})
|
|
||||||
requestCanceled := make(chan struct{})
|
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) {
|
|
||||||
close(requestStarted)
|
|
||||||
<-r.Context().Done()
|
|
||||||
close(requestCanceled)
|
|
||||||
}))
|
|
||||||
t.Cleanup(server.Close)
|
|
||||||
|
|
||||||
monitor := NewMonitor(
|
|
||||||
context.Background(),
|
|
||||||
v1.HealthCheckConfig{Type: "http", Path: "/health"},
|
|
||||||
strings.TrimPrefix(server.URL, "http://"),
|
|
||||||
func() {},
|
|
||||||
func() {},
|
|
||||||
)
|
|
||||||
t.Cleanup(func() { stopMonitorWorker(t, monitor) })
|
|
||||||
monitor.Start()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-requestStarted:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("health check request did not start")
|
|
||||||
}
|
|
||||||
stopMonitorWorker(t, monitor)
|
|
||||||
select {
|
|
||||||
case <-requestCanceled:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("health check request context was not canceled")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("interval wait", func(t *testing.T) {
|
|
||||||
monitor := NewMonitor(context.Background(), v1.HealthCheckConfig{Type: "tcp"}, "", nil, nil)
|
|
||||||
intervals := observeWorkerIntervals(monitor)
|
|
||||||
t.Cleanup(func() { stopMonitorWorker(t, monitor) })
|
|
||||||
monitor.Start()
|
|
||||||
|
|
||||||
// A real worker has completed its first check and installed a timer
|
|
||||||
// that will never fire. Stop must release the wait and exit the loop.
|
|
||||||
awaitWorkerInterval(t, intervals, workerStatus{})
|
|
||||||
stopMonitorWorker(t, monitor)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
type gatedErrContext struct {
|
|
||||||
context.Context
|
|
||||||
errCalled chan struct{}
|
|
||||||
releaseErr chan struct{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (ctx *gatedErrContext) Err() error {
|
|
||||||
select {
|
|
||||||
case ctx.errCalled <- struct{}{}:
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
<-ctx.releaseErr
|
|
||||||
return ctx.Context.Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitorTimerCancellationAfterTimerFires(t *testing.T) {
|
|
||||||
baseCtx, cancel := context.WithCancel(context.Background())
|
|
||||||
ctx := &gatedErrContext{
|
|
||||||
Context: baseCtx,
|
|
||||||
errCalled: make(chan struct{}, 1),
|
|
||||||
releaseErr: make(chan struct{}),
|
|
||||||
}
|
|
||||||
monitor := NewMonitor(
|
|
||||||
context.Background(),
|
|
||||||
v1.HealthCheckConfig{Type: "tcp"},
|
|
||||||
"",
|
|
||||||
nil,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
monitor.ctx = ctx
|
|
||||||
monitor.cancel = cancel
|
|
||||||
|
|
||||||
timerReady := make(chan time.Time, 1)
|
|
||||||
timerReady <- time.Now()
|
|
||||||
var timerStopped atomic.Bool
|
|
||||||
monitor.timerFactory = func(time.Duration) (<-chan time.Time, func()) {
|
|
||||||
return timerReady, func() { timerStopped.Store(true) }
|
|
||||||
}
|
|
||||||
|
|
||||||
waitResult := make(chan bool, 1)
|
|
||||||
go func() {
|
|
||||||
waitResult <- monitor.waitForNextCheck()
|
|
||||||
}()
|
|
||||||
|
|
||||||
// ctx.Done is not ready when select runs, so receiving timerReady is the
|
|
||||||
// only possible branch. Err signals after that receive and blocks until the
|
|
||||||
// test cancels the context, deterministically exercising the cancellation
|
|
||||||
// re-check in the timer branch.
|
|
||||||
select {
|
|
||||||
case <-ctx.errCalled:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("timer branch did not re-check the monitor context")
|
|
||||||
}
|
|
||||||
cancel()
|
|
||||||
close(ctx.releaseErr)
|
|
||||||
|
|
||||||
select {
|
|
||||||
case shouldContinue := <-waitResult:
|
|
||||||
require.False(t, shouldContinue)
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("timer wait did not observe context cancellation")
|
|
||||||
}
|
|
||||||
require.True(t, timerStopped.Load())
|
|
||||||
}
|
|
||||||
+9
-27
@@ -61,12 +61,11 @@ func NewProxy(
|
|||||||
encryptionKey []byte,
|
encryptionKey []byte,
|
||||||
msgTransporter transport.MessageTransporter,
|
msgTransporter transport.MessageTransporter,
|
||||||
vnetController *vnet.Controller,
|
vnetController *vnet.Controller,
|
||||||
udpPacketCodec string,
|
|
||||||
) (pxy Proxy) {
|
) (pxy Proxy) {
|
||||||
var limiter *rate.Limiter
|
var limiter *rate.Limiter
|
||||||
limitBytes := pxyConf.GetBaseConfig().Transport.BandwidthLimit.Bytes()
|
limitBytes := pxyConf.GetBaseConfig().Transport.BandwidthLimit.Bytes()
|
||||||
if limitBytes > 0 && pxyConf.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeClient {
|
if limitBytes > 0 && pxyConf.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeClient {
|
||||||
limiter = limit.NewBandwidthLimiter(limitBytes)
|
limiter = rate.NewLimiter(rate.Limit(float64(limitBytes)), int(limitBytes))
|
||||||
}
|
}
|
||||||
|
|
||||||
baseProxy := BaseProxy{
|
baseProxy := BaseProxy{
|
||||||
@@ -78,7 +77,6 @@ func NewProxy(
|
|||||||
vnetController: vnetController,
|
vnetController: vnetController,
|
||||||
xl: xlog.FromContextSafe(ctx),
|
xl: xlog.FromContextSafe(ctx),
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
udpPacketCodec: udpPacketCodec,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
factory := proxyFactoryRegistry[reflect.TypeOf(pxyConf)]
|
factory := proxyFactoryRegistry[reflect.TypeOf(pxyConf)]
|
||||||
@@ -100,10 +98,9 @@ type BaseProxy struct {
|
|||||||
proxyPlugin plugin.Plugin
|
proxyPlugin plugin.Plugin
|
||||||
inWorkConnCallback func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) /* continue */ bool
|
inWorkConnCallback func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) /* continue */ bool
|
||||||
|
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
xl *xlog.Logger
|
xl *xlog.Logger
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
udpPacketCodec string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (pxy *BaseProxy) Run() error {
|
func (pxy *BaseProxy) Run() error {
|
||||||
@@ -174,26 +171,6 @@ func (pxy *BaseProxy) HandleTCPWorkConnection(workConn net.Conn, m *msg.StartWor
|
|||||||
xl.Tracef("handle tcp work connection, useEncryption: %t, useCompression: %t",
|
xl.Tracef("handle tcp work connection, useEncryption: %t, useCompression: %t",
|
||||||
baseCfg.Transport.UseEncryption, baseCfg.Transport.UseCompression)
|
baseCfg.Transport.UseEncryption, baseCfg.Transport.UseCompression)
|
||||||
|
|
||||||
var srcAddr, dstAddr *net.TCPAddr
|
|
||||||
if m.SrcAddr != "" && m.SrcPort != 0 {
|
|
||||||
if m.DstAddr == "" {
|
|
||||||
m.DstAddr = "127.0.0.1"
|
|
||||||
}
|
|
||||||
var err error
|
|
||||||
srcAddr, err = net.ResolveTCPAddr("tcp", net.JoinHostPort(m.SrcAddr, strconv.Itoa(int(m.SrcPort))))
|
|
||||||
if err != nil {
|
|
||||||
xl.Warnf("resolve source address [%s] error: %v", m.SrcAddr, err)
|
|
||||||
_ = workConn.Close()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
dstAddr, err = net.ResolveTCPAddr("tcp", net.JoinHostPort(m.DstAddr, strconv.Itoa(int(m.DstPort))))
|
|
||||||
if err != nil {
|
|
||||||
xl.Warnf("resolve destination address [%s] error: %v", m.DstAddr, err)
|
|
||||||
_ = workConn.Close()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
remote, recycleFn, err := pxy.wrapWorkConn(workConn, encKey)
|
remote, recycleFn, err := pxy.wrapWorkConn(workConn, encKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
xl.Errorf("wrap work connection: %v", err)
|
xl.Errorf("wrap work connection: %v", err)
|
||||||
@@ -203,6 +180,11 @@ func (pxy *BaseProxy) HandleTCPWorkConnection(workConn net.Conn, m *msg.StartWor
|
|||||||
// check if we need to send proxy protocol info
|
// check if we need to send proxy protocol info
|
||||||
var connInfo plugin.ConnectionInfo
|
var connInfo plugin.ConnectionInfo
|
||||||
if m.SrcAddr != "" && m.SrcPort != 0 {
|
if m.SrcAddr != "" && m.SrcPort != 0 {
|
||||||
|
if m.DstAddr == "" {
|
||||||
|
m.DstAddr = "127.0.0.1"
|
||||||
|
}
|
||||||
|
srcAddr, _ := net.ResolveTCPAddr("tcp", net.JoinHostPort(m.SrcAddr, strconv.Itoa(int(m.SrcPort))))
|
||||||
|
dstAddr, _ := net.ResolveTCPAddr("tcp", net.JoinHostPort(m.DstAddr, strconv.Itoa(int(m.DstPort))))
|
||||||
connInfo.SrcAddr = srcAddr
|
connInfo.SrcAddr = srcAddr
|
||||||
connInfo.DstAddr = dstAddr
|
connInfo.DstAddr = dstAddr
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -43,8 +43,7 @@ type Manager struct {
|
|||||||
encryptionKey []byte
|
encryptionKey []byte
|
||||||
clientCfg *v1.ClientCommonConfig
|
clientCfg *v1.ClientCommonConfig
|
||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
udpPacketCodec string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewManager(
|
func NewManager(
|
||||||
@@ -53,7 +52,6 @@ func NewManager(
|
|||||||
encryptionKey []byte,
|
encryptionKey []byte,
|
||||||
msgTransporter transport.MessageTransporter,
|
msgTransporter transport.MessageTransporter,
|
||||||
vnetController *vnet.Controller,
|
vnetController *vnet.Controller,
|
||||||
udpPacketCodec string,
|
|
||||||
) *Manager {
|
) *Manager {
|
||||||
return &Manager{
|
return &Manager{
|
||||||
proxies: make(map[string]*Wrapper),
|
proxies: make(map[string]*Wrapper),
|
||||||
@@ -63,7 +61,6 @@ func NewManager(
|
|||||||
encryptionKey: encryptionKey,
|
encryptionKey: encryptionKey,
|
||||||
clientCfg: clientCfg,
|
clientCfg: clientCfg,
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
udpPacketCodec: udpPacketCodec,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -169,7 +166,7 @@ func (pm *Manager) UpdateAll(proxyCfgs []v1.ProxyConfigurer) {
|
|||||||
for _, cfg := range proxyCfgs {
|
for _, cfg := range proxyCfgs {
|
||||||
name := cfg.GetBaseConfig().Name
|
name := cfg.GetBaseConfig().Name
|
||||||
if _, ok := pm.proxies[name]; !ok {
|
if _, ok := pm.proxies[name]; !ok {
|
||||||
pxy := NewWrapper(pm.ctx, cfg, pm.clientCfg, pm.encryptionKey, pm.HandleEvent, pm.msgTransporter, pm.vnetController, pm.udpPacketCodec)
|
pxy := NewWrapper(pm.ctx, cfg, pm.clientCfg, pm.encryptionKey, pm.HandleEvent, pm.msgTransporter, pm.vnetController)
|
||||||
if pm.inWorkConnCallback != nil {
|
if pm.inWorkConnCallback != nil {
|
||||||
pxy.SetInWorkConnCallback(pm.inWorkConnCallback)
|
pxy.SetInWorkConnCallback(pm.inWorkConnCallback)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,47 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
//go:build !frps
|
|
||||||
|
|
||||||
package proxy
|
|
||||||
|
|
||||||
import (
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
|
||||||
"github.com/fatedier/frp/pkg/util/xlog"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestHandleTCPWorkConnectionRejectsInvalidAddress(t *testing.T) {
|
|
||||||
workConn, peerConn := net.Pipe()
|
|
||||||
defer peerConn.Close()
|
|
||||||
|
|
||||||
pxy := &BaseProxy{
|
|
||||||
baseCfg: &v1.ProxyBaseConfig{},
|
|
||||||
xl: xlog.New(),
|
|
||||||
}
|
|
||||||
pxy.HandleTCPWorkConnection(workConn, &msg.StartWorkConn{
|
|
||||||
SrcAddr: "[",
|
|
||||||
SrcPort: 1,
|
|
||||||
}, nil)
|
|
||||||
|
|
||||||
buffer := make([]byte, 1)
|
|
||||||
_, err := peerConn.Read(buffer)
|
|
||||||
require.ErrorIs(t, err, io.EOF)
|
|
||||||
}
|
|
||||||
@@ -99,7 +99,6 @@ func NewWrapper(
|
|||||||
eventHandler event.Handler,
|
eventHandler event.Handler,
|
||||||
msgTransporter transport.MessageTransporter,
|
msgTransporter transport.MessageTransporter,
|
||||||
vnetController *vnet.Controller,
|
vnetController *vnet.Controller,
|
||||||
udpPacketCodec string,
|
|
||||||
) *Wrapper {
|
) *Wrapper {
|
||||||
baseInfo := cfg.GetBaseConfig()
|
baseInfo := cfg.GetBaseConfig()
|
||||||
xl := xlog.FromContextSafe(ctx).Spawn().AppendPrefix(baseInfo.Name)
|
xl := xlog.FromContextSafe(ctx).Spawn().AppendPrefix(baseInfo.Name)
|
||||||
@@ -128,7 +127,7 @@ func NewWrapper(
|
|||||||
xl.Tracef("enable health check monitor")
|
xl.Tracef("enable health check monitor")
|
||||||
}
|
}
|
||||||
|
|
||||||
pw.pxy = NewProxy(pw.ctx, pw.Cfg, clientCfg, encryptionKey, pw.msgTransporter, pw.vnetController, udpPacketCodec)
|
pw.pxy = NewProxy(pw.ctx, pw.Cfg, clientCfg, encryptionKey, pw.msgTransporter, pw.vnetController)
|
||||||
return pw
|
return pw
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -87,13 +87,7 @@ func (pxy *SUDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
workConn := netpkg.WrapReadWriteCloserToConn(remote, conn)
|
workConn := netpkg.WrapReadWriteCloserToConn(remote, conn)
|
||||||
payloadRW, err := msg.NewUDPPacketReadWriter(workConn, pxy.clientCfg.Transport.WireProtocol, pxy.udpPacketCodec)
|
payloadConn := msg.NewConn(workConn, msg.NewReadWriter(workConn, pxy.clientCfg.Transport.WireProtocol))
|
||||||
if err != nil {
|
|
||||||
xl.Errorf("create SUDP packet read writer: %v", err)
|
|
||||||
_ = workConn.Close()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
payloadConn := msg.NewConn(workConn, payloadRW)
|
|
||||||
readCh := make(chan *msg.UDPPacket, 1024)
|
readCh := make(chan *msg.UDPPacket, 1024)
|
||||||
sendCh := make(chan msg.Message, 1024)
|
sendCh := make(chan msg.Message, 1024)
|
||||||
isClose := false
|
isClose := false
|
||||||
|
|||||||
+3
-10
@@ -97,17 +97,10 @@ func (pxy *UDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
workConn := netpkg.WrapReadWriteCloserToConn(remote, conn)
|
|
||||||
// Plain UDP payload follows the configured wire protocol for message framing.
|
|
||||||
payloadRW, err := msg.NewUDPPacketReadWriter(workConn, pxy.clientCfg.Transport.WireProtocol, pxy.udpPacketCodec)
|
|
||||||
if err != nil {
|
|
||||||
xl.Errorf("create UDP packet read writer: %v", err)
|
|
||||||
workConn.Close()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
pxy.mu.Lock()
|
pxy.mu.Lock()
|
||||||
pxy.workConn = workConn
|
pxy.workConn = netpkg.WrapReadWriteCloserToConn(remote, conn)
|
||||||
|
// Plain UDP payload follows the configured wire protocol for message framing.
|
||||||
|
payloadRW := msg.NewReadWriter(pxy.workConn, pxy.clientCfg.Transport.WireProtocol)
|
||||||
pxy.readCh = make(chan *msg.UDPPacket, 1024)
|
pxy.readCh = make(chan *msg.UDPPacket, 1024)
|
||||||
pxy.sendCh = make(chan msg.Message, 1024)
|
pxy.sendCh = make(chan msg.Message, 1024)
|
||||||
pxy.closed = false
|
pxy.closed = false
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ import (
|
|||||||
"reflect"
|
"reflect"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
fmux "github.com/fatedier/yamux"
|
fmux "github.com/hashicorp/yamux"
|
||||||
"github.com/quic-go/quic-go"
|
"github.com/quic-go/quic-go"
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
v1 "github.com/fatedier/frp/pkg/config/v1"
|
||||||
|
|||||||
+4
-16
@@ -22,7 +22,6 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/fatedier/golib/crypto"
|
"github.com/fatedier/golib/crypto"
|
||||||
@@ -33,7 +32,6 @@ import (
|
|||||||
"github.com/fatedier/frp/pkg/config"
|
"github.com/fatedier/frp/pkg/config"
|
||||||
"github.com/fatedier/frp/pkg/config/source"
|
"github.com/fatedier/frp/pkg/config/source"
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
v1 "github.com/fatedier/frp/pkg/config/v1"
|
||||||
"github.com/fatedier/frp/pkg/config/v1/validation"
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
"github.com/fatedier/frp/pkg/msg"
|
||||||
"github.com/fatedier/frp/pkg/policy/security"
|
"github.com/fatedier/frp/pkg/policy/security"
|
||||||
httppkg "github.com/fatedier/frp/pkg/util/http"
|
httppkg "github.com/fatedier/frp/pkg/util/http"
|
||||||
@@ -111,9 +109,6 @@ func setServiceOptionsDefault(options *ServiceOptions) error {
|
|||||||
// Service is the client service that connects to frps and provides proxy services.
|
// Service is the client service that connects to frps and provides proxy services.
|
||||||
type Service struct {
|
type Service struct {
|
||||||
ctlMu sync.RWMutex
|
ctlMu sync.RWMutex
|
||||||
// Stores gracefulShutdownDuration independently from ctlMu, because the
|
|
||||||
// graceful shutdown wait may hold ctlMu for an arbitrary duration.
|
|
||||||
gracefulShutdownDuration atomic.Int64
|
|
||||||
// manager control connection with server
|
// manager control connection with server
|
||||||
ctl *Control
|
ctl *Control
|
||||||
// Uniq id got from frps, it will be attached to loginMsg.
|
// Uniq id got from frps, it will be attached to loginMsg.
|
||||||
@@ -154,7 +149,8 @@ type Service struct {
|
|||||||
// service context
|
// service context
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
// call cancel to stop service
|
// call cancel to stop service
|
||||||
cancel context.CancelCauseFunc
|
cancel context.CancelCauseFunc
|
||||||
|
gracefulShutdownDuration time.Duration
|
||||||
|
|
||||||
connectorCreator func(context.Context, *v1.ClientCommonConfig) Connector
|
connectorCreator func(context.Context, *v1.ClientCommonConfig) Connector
|
||||||
handleWorkConnCb func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) bool
|
handleWorkConnCb func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) bool
|
||||||
@@ -416,7 +412,7 @@ func (svr *Service) Close() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (svr *Service) GracefulClose(d time.Duration) {
|
func (svr *Service) GracefulClose(d time.Duration) {
|
||||||
svr.gracefulShutdownDuration.Store(int64(d))
|
svr.gracefulShutdownDuration = d
|
||||||
svr.cancel(nil)
|
svr.cancel(nil)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -433,8 +429,7 @@ func (svr *Service) stop() {
|
|||||||
svr.ctlMu.Lock()
|
svr.ctlMu.Lock()
|
||||||
defer svr.ctlMu.Unlock()
|
defer svr.ctlMu.Unlock()
|
||||||
if svr.ctl != nil {
|
if svr.ctl != nil {
|
||||||
d := time.Duration(svr.gracefulShutdownDuration.Load())
|
svr.ctl.GracefulClose(svr.gracefulShutdownDuration)
|
||||||
svr.ctl.GracefulClose(d)
|
|
||||||
svr.ctl = nil
|
svr.ctl = nil
|
||||||
}
|
}
|
||||||
if svr.webServer != nil {
|
if svr.webServer != nil {
|
||||||
@@ -511,13 +506,6 @@ func (svr *Service) reloadConfigFromSourcesLocked() error {
|
|||||||
proxies, visitors = config.FilterClientConfigurers(reloadCommon, proxies, visitors)
|
proxies, visitors = config.FilterClientConfigurers(reloadCommon, proxies, visitors)
|
||||||
proxies = config.CompleteProxyConfigurers(proxies)
|
proxies = config.CompleteProxyConfigurers(proxies)
|
||||||
visitors = config.CompleteVisitorConfigurers(visitors)
|
visitors = config.CompleteVisitorConfigurers(visitors)
|
||||||
requirements := validation.GetClientConfigRequirements(reloadCommon, proxies, visitors)
|
|
||||||
if svr.vnetController == nil && requirements.VirtualNet {
|
|
||||||
return errors.New(
|
|
||||||
"VirtualNet-dependent configuration requires a VirtualNet runtime enabled at startup; " +
|
|
||||||
"restart frpc after configuring featureGates.VirtualNet and virtualNet.address",
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Atomically replace the entire configuration
|
// Atomically replace the entire configuration
|
||||||
if err := svr.UpdateAllConfigurer(proxies, visitors); err != nil {
|
if err := svr.UpdateAllConfigurer(proxies, visitors); err != nil {
|
||||||
|
|||||||
@@ -1,95 +0,0 @@
|
|||||||
package client
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"net"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/client/proxy"
|
|
||||||
"github.com/fatedier/frp/client/visitor"
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
|
||||||
)
|
|
||||||
|
|
||||||
type gracefulCloseTestConnector struct {
|
|
||||||
conn net.Conn
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*gracefulCloseTestConnector) Connect() (*msg.Conn, error) { return nil, net.ErrClosed }
|
|
||||||
func (c *gracefulCloseTestConnector) Close() error { return c.conn.Close() }
|
|
||||||
|
|
||||||
func newGracefulCloseTestService() *Service {
|
|
||||||
ctx := context.Background()
|
|
||||||
common := &v1.ClientCommonConfig{}
|
|
||||||
serverConn, clientConn := net.Pipe()
|
|
||||||
ctl := &Control{
|
|
||||||
ctx: ctx,
|
|
||||||
sessionCtx: &SessionContext{
|
|
||||||
Common: common,
|
|
||||||
RunID: "graceful-close-race",
|
|
||||||
Conn: msg.NewConn(clientConn, msg.NewV1ReadWriter(clientConn)),
|
|
||||||
Connector: &gracefulCloseTestConnector{conn: serverConn},
|
|
||||||
},
|
|
||||||
doneCh: make(chan struct{}),
|
|
||||||
}
|
|
||||||
ctl.pm = proxy.NewManager(ctx, common, nil, nil, nil, "")
|
|
||||||
ctl.vm = visitor.NewManager(ctx, "graceful-close-race", common, nil, nil, nil, "")
|
|
||||||
return &Service{ctl: ctl, cancel: context.CancelCauseFunc(func(error) {})}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGracefulCloseAndStopSynchronizeDuration(t *testing.T) {
|
|
||||||
for i := range 10000 {
|
|
||||||
svr := newGracefulCloseTestService()
|
|
||||||
start := make(chan struct{})
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
wg.Add(2)
|
|
||||||
go func() {
|
|
||||||
defer wg.Done()
|
|
||||||
<-start
|
|
||||||
svr.GracefulClose(time.Duration(i))
|
|
||||||
}()
|
|
||||||
go func() {
|
|
||||||
defer wg.Done()
|
|
||||||
<-start
|
|
||||||
svr.stop()
|
|
||||||
}()
|
|
||||||
close(start)
|
|
||||||
wg.Wait()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGracefulCloseDoesNotBlockDuringStop(t *testing.T) {
|
|
||||||
const gracefulDuration = 200 * time.Millisecond
|
|
||||||
|
|
||||||
svr := newGracefulCloseTestService()
|
|
||||||
svr.GracefulClose(gracefulDuration)
|
|
||||||
stopDone := make(chan struct{})
|
|
||||||
go func() {
|
|
||||||
svr.stop()
|
|
||||||
close(stopDone)
|
|
||||||
}()
|
|
||||||
defer func() {
|
|
||||||
select {
|
|
||||||
case <-stopDone:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Error("stop did not finish")
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
deadline := time.Now().Add(time.Second)
|
|
||||||
for svr.ctlMu.TryLock() {
|
|
||||||
svr.ctlMu.Unlock()
|
|
||||||
if time.Now().After(deadline) {
|
|
||||||
t.Fatal("stop did not acquire ctlMu")
|
|
||||||
}
|
|
||||||
time.Sleep(time.Millisecond)
|
|
||||||
}
|
|
||||||
|
|
||||||
start := time.Now()
|
|
||||||
svr.GracefulClose(0)
|
|
||||||
if elapsed := time.Since(start); elapsed >= gracefulDuration/2 {
|
|
||||||
t.Fatalf("GracefulClose blocked for %v while stop was waiting", elapsed)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -113,13 +113,7 @@ func (sv *SUDPVisitor) dispatcher() {
|
|||||||
func (sv *SUDPVisitor) worker(workConn net.Conn, firstPacket *msg.UDPPacket) {
|
func (sv *SUDPVisitor) worker(workConn net.Conn, firstPacket *msg.UDPPacket) {
|
||||||
xl := xlog.FromContextSafe(sv.ctx)
|
xl := xlog.FromContextSafe(sv.ctx)
|
||||||
xl.Debugf("starting sudp proxy worker")
|
xl.Debugf("starting sudp proxy worker")
|
||||||
payloadRW, err := msg.NewUDPPacketReadWriter(workConn, sv.clientCfg.Transport.WireProtocol, udpPacketCodecFromHelper(sv.helper))
|
payloadConn := msg.NewConn(workConn, msg.NewReadWriter(workConn, sv.clientCfg.Transport.WireProtocol))
|
||||||
if err != nil {
|
|
||||||
xl.Errorf("create SUDP packet read writer: %v", err)
|
|
||||||
_ = workConn.Close()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
payloadConn := msg.NewConn(workConn, payloadRW)
|
|
||||||
|
|
||||||
wg := &sync.WaitGroup{}
|
wg := &sync.WaitGroup{}
|
||||||
wg.Add(2)
|
wg.Add(2)
|
||||||
|
|||||||
@@ -50,17 +50,6 @@ type Helper interface {
|
|||||||
RunID() string
|
RunID() string
|
||||||
}
|
}
|
||||||
|
|
||||||
type udpPacketCodecProvider interface {
|
|
||||||
UDPPacketCodec() string
|
|
||||||
}
|
|
||||||
|
|
||||||
func udpPacketCodecFromHelper(helper Helper) string {
|
|
||||||
if provider, ok := helper.(udpPacketCodecProvider); ok {
|
|
||||||
return provider.UDPPacketCodec()
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
// Visitor is used for forward traffics from local port tot remote service.
|
// Visitor is used for forward traffics from local port tot remote service.
|
||||||
type Visitor interface {
|
type Visitor interface {
|
||||||
Run() error
|
Run() error
|
||||||
|
|||||||
@@ -53,12 +53,7 @@ func NewManager(
|
|||||||
connectServer func() (*msg.Conn, error),
|
connectServer func() (*msg.Conn, error),
|
||||||
msgTransporter transport.MessageTransporter,
|
msgTransporter transport.MessageTransporter,
|
||||||
vnetController *vnet.Controller,
|
vnetController *vnet.Controller,
|
||||||
udpPacketCodecs ...string,
|
|
||||||
) *Manager {
|
) *Manager {
|
||||||
udpPacketCodec := ""
|
|
||||||
if len(udpPacketCodecs) > 0 {
|
|
||||||
udpPacketCodec = udpPacketCodecs[0]
|
|
||||||
}
|
|
||||||
m := &Manager{
|
m := &Manager{
|
||||||
clientCfg: clientCfg,
|
clientCfg: clientCfg,
|
||||||
cfgs: make(map[string]v1.VisitorConfigurer),
|
cfgs: make(map[string]v1.VisitorConfigurer),
|
||||||
@@ -73,7 +68,6 @@ func NewManager(
|
|||||||
vnetController: vnetController,
|
vnetController: vnetController,
|
||||||
transferConnFn: m.TransferConn,
|
transferConnFn: m.TransferConn,
|
||||||
runID: runID,
|
runID: runID,
|
||||||
udpPacketCodec: udpPacketCodec,
|
|
||||||
}
|
}
|
||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
@@ -211,7 +205,6 @@ type visitorHelperImpl struct {
|
|||||||
vnetController *vnet.Controller
|
vnetController *vnet.Controller
|
||||||
transferConnFn func(name string, conn net.Conn) error
|
transferConnFn func(name string, conn net.Conn) error
|
||||||
runID string
|
runID string
|
||||||
udpPacketCodec string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *visitorHelperImpl) ConnectServer() (*msg.Conn, error) {
|
func (v *visitorHelperImpl) ConnectServer() (*msg.Conn, error) {
|
||||||
@@ -233,7 +226,3 @@ func (v *visitorHelperImpl) VNetController() *vnet.Controller {
|
|||||||
func (v *visitorHelperImpl) RunID() string {
|
func (v *visitorHelperImpl) RunID() string {
|
||||||
return v.runID
|
return v.runID
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *visitorHelperImpl) UDPPacketCodec() string {
|
|
||||||
return v.udpPacketCodec
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
libio "github.com/fatedier/golib/io"
|
libio "github.com/fatedier/golib/io"
|
||||||
fmux "github.com/fatedier/yamux"
|
fmux "github.com/hashicorp/yamux"
|
||||||
quic "github.com/quic-go/quic-go"
|
quic "github.com/quic-go/quic-go"
|
||||||
"golang.org/x/time/rate"
|
"golang.org/x/time/rate"
|
||||||
|
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ import (
|
|||||||
"github.com/fatedier/frp/pkg/config/source"
|
"github.com/fatedier/frp/pkg/config/source"
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
v1 "github.com/fatedier/frp/pkg/config/v1"
|
||||||
"github.com/fatedier/frp/pkg/config/v1/validation"
|
"github.com/fatedier/frp/pkg/config/v1/validation"
|
||||||
|
"github.com/fatedier/frp/pkg/policy/featuregate"
|
||||||
"github.com/fatedier/frp/pkg/policy/security"
|
"github.com/fatedier/frp/pkg/policy/security"
|
||||||
"github.com/fatedier/frp/pkg/util/log"
|
"github.com/fatedier/frp/pkg/util/log"
|
||||||
"github.com/fatedier/frp/pkg/util/version"
|
"github.com/fatedier/frp/pkg/util/version"
|
||||||
@@ -130,6 +131,12 @@ func runClient(cfgFilePath string, unsafeFeatures *security.UnsafeFeatures) erro
|
|||||||
"please use yaml/json/toml format instead!\n")
|
"please use yaml/json/toml format instead!\n")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if len(result.Common.FeatureGates) > 0 {
|
||||||
|
if err := featuregate.SetFromMap(result.Common.FeatureGates); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return runClientWithAggregator(result, unsafeFeatures, cfgFilePath)
|
return runClientWithAggregator(result, unsafeFeatures, cfgFilePath)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+6
-13
@@ -29,18 +29,6 @@ func init() {
|
|||||||
rootCmd.AddCommand(verifyCmd)
|
rootCmd.AddCommand(verifyCmd)
|
||||||
}
|
}
|
||||||
|
|
||||||
func verifyClientConfig(
|
|
||||||
configFile string,
|
|
||||||
strict bool,
|
|
||||||
unsafeFeatures *security.UnsafeFeatures,
|
|
||||||
) (validation.Warning, error) {
|
|
||||||
cliCfg, proxyCfgs, visitorCfgs, _, err := config.LoadClientConfig(configFile, strict)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return validation.ValidateAllClientConfig(cliCfg, proxyCfgs, visitorCfgs, unsafeFeatures)
|
|
||||||
}
|
|
||||||
|
|
||||||
var verifyCmd = &cobra.Command{
|
var verifyCmd = &cobra.Command{
|
||||||
Use: "verify",
|
Use: "verify",
|
||||||
Short: "Verify that the configures is valid",
|
Short: "Verify that the configures is valid",
|
||||||
@@ -50,8 +38,13 @@ var verifyCmd = &cobra.Command{
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
cliCfg, proxyCfgs, visitorCfgs, _, err := config.LoadClientConfig(cfgFile, strictConfigMode)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Println(err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
unsafeFeatures := security.NewUnsafeFeatures(allowUnsafe)
|
unsafeFeatures := security.NewUnsafeFeatures(allowUnsafe)
|
||||||
warning, err := verifyClientConfig(cfgFile, strictConfigMode, unsafeFeatures)
|
warning, err := validation.ValidateAllClientConfig(cliCfg, proxyCfgs, visitorCfgs, unsafeFeatures)
|
||||||
if warning != nil {
|
if warning != nil {
|
||||||
fmt.Printf("WARNING: %v\n", warning)
|
fmt.Printf("WARNING: %v\n", warning)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,67 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package sub
|
|
||||||
|
|
||||||
import (
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/policy/security"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestVerifyClientConfigFeatureGates(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
content string
|
|
||||||
wantErr string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "VirtualNet enabled",
|
|
||||||
content: `featureGates = { VirtualNet = true }
|
|
||||||
virtualNet.address = "100.86.0.4/24"
|
|
||||||
`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "VirtualNet disabled",
|
|
||||||
content: `featureGates = { VirtualNet = false }
|
|
||||||
virtualNet.address = "100.86.0.4/24"
|
|
||||||
`,
|
|
||||||
wantErr: "VirtualNet feature is not enabled",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "unknown feature gate",
|
|
||||||
content: `featureGates = { UnknownFeature = true }`,
|
|
||||||
wantErr: "unrecognized feature gate: UnknownFeature",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
configFile := filepath.Join(t.TempDir(), "frpc.toml")
|
|
||||||
require.NoError(t, os.WriteFile(configFile, []byte(tc.content), 0o600))
|
|
||||||
|
|
||||||
warning, err := verifyClientConfig(configFile, true, security.NewUnsafeFeatures(nil))
|
|
||||||
require.NoError(t, warning)
|
|
||||||
if tc.wantErr == "" {
|
|
||||||
require.NoError(t, err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
require.ErrorContains(t, err, tc.wantErr)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Binary file not shown.
|
Before Width: | Height: | Size: 49 KiB |
@@ -4,18 +4,19 @@ go 1.25.0
|
|||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5
|
github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5
|
||||||
github.com/coreos/go-oidc/v3 v3.18.0
|
github.com/coreos/go-oidc/v3 v3.14.1
|
||||||
github.com/fatedier/golib v0.8.2
|
github.com/fatedier/golib v0.7.0
|
||||||
github.com/fatedier/yamux v0.2.0
|
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/gorilla/mux v1.8.1
|
github.com/gorilla/mux v1.8.1
|
||||||
github.com/gorilla/websocket v1.5.0
|
github.com/gorilla/websocket v1.5.0
|
||||||
|
github.com/hashicorp/yamux v0.1.1
|
||||||
github.com/onsi/ginkgo/v2 v2.23.4
|
github.com/onsi/ginkgo/v2 v2.23.4
|
||||||
github.com/onsi/gomega v1.36.3
|
github.com/onsi/gomega v1.36.3
|
||||||
github.com/pelletier/go-toml/v2 v2.2.0
|
github.com/pelletier/go-toml/v2 v2.2.0
|
||||||
github.com/pires/go-proxyproto v0.15.0
|
github.com/pion/stun/v3 v3.1.1
|
||||||
|
github.com/pires/go-proxyproto v0.7.0
|
||||||
github.com/prometheus/client_golang v1.19.1
|
github.com/prometheus/client_golang v1.19.1
|
||||||
github.com/quic-go/quic-go v0.60.0
|
github.com/quic-go/quic-go v0.55.0
|
||||||
github.com/rodaine/table v1.2.0
|
github.com/rodaine/table v1.2.0
|
||||||
github.com/samber/lo v1.47.0
|
github.com/samber/lo v1.47.0
|
||||||
github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8
|
github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8
|
||||||
@@ -25,11 +26,11 @@ require (
|
|||||||
github.com/tidwall/gjson v1.17.1
|
github.com/tidwall/gjson v1.17.1
|
||||||
github.com/vishvananda/netlink v1.3.0
|
github.com/vishvananda/netlink v1.3.0
|
||||||
github.com/xtaci/kcp-go/v5 v5.6.13
|
github.com/xtaci/kcp-go/v5 v5.6.13
|
||||||
golang.org/x/crypto v0.54.0
|
golang.org/x/crypto v0.49.0
|
||||||
golang.org/x/net v0.56.0
|
golang.org/x/net v0.52.0
|
||||||
golang.org/x/oauth2 v0.36.0
|
golang.org/x/oauth2 v0.28.0
|
||||||
golang.org/x/sync v0.22.0
|
golang.org/x/sync v0.20.0
|
||||||
golang.org/x/sys v0.47.0
|
golang.org/x/sys v0.42.0
|
||||||
golang.org/x/time v0.10.0
|
golang.org/x/time v0.10.0
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20231211153847-12269c276173
|
golang.zx2c4.com/wireguard v0.0.0-20231211153847-12269c276173
|
||||||
gopkg.in/ini.v1 v1.67.0
|
gopkg.in/ini.v1 v1.67.0
|
||||||
@@ -43,7 +44,7 @@ require (
|
|||||||
github.com/beorn7/perks v1.0.1 // indirect
|
github.com/beorn7/perks v1.0.1 // indirect
|
||||||
github.com/cespare/xxhash/v2 v2.2.0 // indirect
|
github.com/cespare/xxhash/v2 v2.2.0 // indirect
|
||||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||||
github.com/go-jose/go-jose/v4 v4.1.4 // indirect
|
github.com/go-jose/go-jose/v4 v4.0.5 // indirect
|
||||||
github.com/go-logr/logr v1.4.2 // indirect
|
github.com/go-logr/logr v1.4.2 // indirect
|
||||||
github.com/go-task/slim-sprig/v3 v3.0.0 // indirect
|
github.com/go-task/slim-sprig/v3 v3.0.0 // indirect
|
||||||
github.com/golang/snappy v0.0.4 // indirect
|
github.com/golang/snappy v0.0.4 // indirect
|
||||||
@@ -52,6 +53,9 @@ require (
|
|||||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||||
github.com/klauspost/cpuid/v2 v2.2.6 // indirect
|
github.com/klauspost/cpuid/v2 v2.2.6 // indirect
|
||||||
github.com/klauspost/reedsolomon v1.12.0 // indirect
|
github.com/klauspost/reedsolomon v1.12.0 // indirect
|
||||||
|
github.com/pion/dtls/v3 v3.0.10 // indirect
|
||||||
|
github.com/pion/logging v0.2.4 // indirect
|
||||||
|
github.com/pion/transport/v4 v4.0.1 // indirect
|
||||||
github.com/pkg/errors v0.9.1 // indirect
|
github.com/pkg/errors v0.9.1 // indirect
|
||||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||||
github.com/prometheus/client_model v0.5.0 // indirect
|
github.com/prometheus/client_model v0.5.0 // indirect
|
||||||
@@ -63,9 +67,11 @@ require (
|
|||||||
github.com/tidwall/pretty v1.2.0 // indirect
|
github.com/tidwall/pretty v1.2.0 // indirect
|
||||||
github.com/tjfoc/gmsm v1.4.1 // indirect
|
github.com/tjfoc/gmsm v1.4.1 // indirect
|
||||||
github.com/vishvananda/netns v0.0.4 // indirect
|
github.com/vishvananda/netns v0.0.4 // indirect
|
||||||
|
github.com/wlynxg/anet v0.0.5 // indirect
|
||||||
go.uber.org/automaxprocs v1.6.0 // indirect
|
go.uber.org/automaxprocs v1.6.0 // indirect
|
||||||
golang.org/x/text v0.40.0 // indirect
|
golang.org/x/mod v0.33.0 // indirect
|
||||||
golang.org/x/tools v0.47.0 // indirect
|
golang.org/x/text v0.35.0 // indirect
|
||||||
|
golang.org/x/tools v0.42.0 // indirect
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect
|
||||||
google.golang.org/protobuf v1.36.5 // indirect
|
google.golang.org/protobuf v1.36.5 // indirect
|
||||||
gopkg.in/yaml.v2 v2.4.0 // indirect
|
gopkg.in/yaml.v2 v2.4.0 // indirect
|
||||||
@@ -73,3 +79,6 @@ require (
|
|||||||
sigs.k8s.io/json v0.0.0-20221116044647-bc3834ca7abd // indirect
|
sigs.k8s.io/json v0.0.0-20221116044647-bc3834ca7abd // indirect
|
||||||
sigs.k8s.io/yaml v1.3.0 // indirect
|
sigs.k8s.io/yaml v1.3.0 // indirect
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// TODO(fatedier): Temporary use the modified version, update to the official version after merging into the official repository.
|
||||||
|
replace github.com/hashicorp/yamux => github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6
|
||||||
|
|||||||
@@ -11,8 +11,8 @@ github.com/cespare/xxhash/v2 v2.2.0 h1:DC2CZ1Ep5Y4k3ZQ899DldepgrayRUGE6BBZ/cd9Cj
|
|||||||
github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||||
github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw=
|
github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw=
|
||||||
github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc=
|
github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc=
|
||||||
github.com/coreos/go-oidc/v3 v3.18.0 h1:V9orjXynvu5wiC9SemFTWnG4F45v403aIcjWo0d41+A=
|
github.com/coreos/go-oidc/v3 v3.14.1 h1:9ePWwfdwC4QKRlCXsJGou56adA/owXczOzwKdOumLqk=
|
||||||
github.com/coreos/go-oidc/v3 v3.18.0/go.mod h1:DYCf24+ncYi+XkIH97GY1+dqoRlbaSI26KVTCI9SrY4=
|
github.com/coreos/go-oidc/v3 v3.14.1/go.mod h1:HaZ3szPaZ0e4r6ebqvsLWlk2Tn+aejfmrfah6hnSYEU=
|
||||||
github.com/cpuguy83/go-md2man/v2 v2.0.3/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o=
|
github.com/cpuguy83/go-md2man/v2 v2.0.3/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o=
|
||||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||||
@@ -20,12 +20,12 @@ github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSs
|
|||||||
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
|
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
|
||||||
github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
|
github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
|
||||||
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
|
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
|
||||||
github.com/fatedier/golib v0.8.2 h1:02n2Dg7KJ7rR7p7n4/6hBUjaLQf2J7EiHYZQsgGTvww=
|
github.com/fatedier/golib v0.7.0 h1:tMDF9ObcwVt59VUHroJOzHQjVFPLymZVMpGm9WAVwhY=
|
||||||
github.com/fatedier/golib v0.8.2/go.mod h1:ArUGvPg2cOw/py2RAuBt46nNZH2VQ5Z70p109MAZpJw=
|
github.com/fatedier/golib v0.7.0/go.mod h1:ArUGvPg2cOw/py2RAuBt46nNZH2VQ5Z70p109MAZpJw=
|
||||||
github.com/fatedier/yamux v0.2.0 h1:H+2A9iBVh7aJlEOc1Ws1FXWOaecBf2nRv9zpFMPUWg8=
|
github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6 h1:u92UUy6FURPmNsMBUuongRWC0rBqN6gd01Dzu+D21NE=
|
||||||
github.com/fatedier/yamux v0.2.0/go.mod h1:d4FtRDrC9sHvRpiDL6J5EnfjLhqzZZplHe5yToSn2Ac=
|
github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6/go.mod h1:c5/tk6G0dSpXGzJN7Wk1OEie8grdSJAmeawId9Zvd34=
|
||||||
github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA=
|
github.com/go-jose/go-jose/v4 v4.0.5 h1:M6T8+mKZl/+fNNuFHvGIzDz7BTLQPIounk/b9dw3AaE=
|
||||||
github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
|
github.com/go-jose/go-jose/v4 v4.0.5/go.mod h1:s3P1lRrkT8igV8D9OjyL4WRyHvjB6a4JSllnOrmmBOA=
|
||||||
github.com/go-logr/logr v1.4.2 h1:6pFjapn8bFcIbiKo3XT4j/BhANplGihG6tvd+8rYgrY=
|
github.com/go-logr/logr v1.4.2 h1:6pFjapn8bFcIbiKo3XT4j/BhANplGihG6tvd+8rYgrY=
|
||||||
github.com/go-logr/logr v1.4.2/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
|
github.com/go-logr/logr v1.4.2/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
|
||||||
github.com/go-task/slim-sprig/v3 v3.0.0 h1:sUs3vkvUymDpBKi3qH1YSqBQk9+9D/8M2mN1vB6EwHI=
|
github.com/go-task/slim-sprig/v3 v3.0.0 h1:sUs3vkvUymDpBKi3qH1YSqBQk9+9D/8M2mN1vB6EwHI=
|
||||||
@@ -78,8 +78,16 @@ github.com/onsi/gomega v1.36.3 h1:hID7cr8t3Wp26+cYnfcjR6HpJ00fdogN6dqZ1t6IylU=
|
|||||||
github.com/onsi/gomega v1.36.3/go.mod h1:8D9+Txp43QWKhM24yyOBEdpkzN8FvJyAwecBgsU4KU0=
|
github.com/onsi/gomega v1.36.3/go.mod h1:8D9+Txp43QWKhM24yyOBEdpkzN8FvJyAwecBgsU4KU0=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.0 h1:QLgLl2yMN7N+ruc31VynXs1vhMZa7CeHHejIeBAsoHo=
|
github.com/pelletier/go-toml/v2 v2.2.0 h1:QLgLl2yMN7N+ruc31VynXs1vhMZa7CeHHejIeBAsoHo=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.0/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs=
|
github.com/pelletier/go-toml/v2 v2.2.0/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs=
|
||||||
github.com/pires/go-proxyproto v0.15.0 h1:dTshmNbFm/D+0+sbrxUuddPOZ5Y0B7c5NhtsBkm6LqI=
|
github.com/pion/dtls/v3 v3.0.10 h1:k9ekkq1kaZoxnNEbyLKI8DI37j/Nbk1HWmMuywpQJgg=
|
||||||
github.com/pires/go-proxyproto v0.15.0/go.mod h1:OXsCrKwrK2tXS9YrI5tkHx5xaQlO8FH3lFW76orFh24=
|
github.com/pion/dtls/v3 v3.0.10/go.mod h1:YEmmBYIoBsY3jmG56dsziTv/Lca9y4Om83370CXfqJ8=
|
||||||
|
github.com/pion/logging v0.2.4 h1:tTew+7cmQ+Mc1pTBLKH2puKsOvhm32dROumOZ655zB8=
|
||||||
|
github.com/pion/logging v0.2.4/go.mod h1:DffhXTKYdNZU+KtJ5pyQDjvOAh/GsNSyv1lbkFbe3so=
|
||||||
|
github.com/pion/stun/v3 v3.1.1 h1:CkQxveJ4xGQjulGSROXbXq94TAWu8gIX2dT+ePhUkqw=
|
||||||
|
github.com/pion/stun/v3 v3.1.1/go.mod h1:qC1DfmcCTQjl9PBaMa5wSn3x9IPmKxSdcCsxBcDBndM=
|
||||||
|
github.com/pion/transport/v4 v4.0.1 h1:sdROELU6BZ63Ab7FrOLn13M6YdJLY20wldXW2Cu2k8o=
|
||||||
|
github.com/pion/transport/v4 v4.0.1/go.mod h1:nEuEA4AD5lPdcIegQDpVLgNoDGreqM/YqmEx3ovP4jM=
|
||||||
|
github.com/pires/go-proxyproto v0.7.0 h1:IukmRewDQFWC7kfnb66CSomk2q/seBuilHBYFwyq0Hs=
|
||||||
|
github.com/pires/go-proxyproto v0.7.0/go.mod h1:Vz/1JPY/OACxWGQNIRY2BeyDmpoaWmEP40O9LbuiFR4=
|
||||||
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
|
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
|
||||||
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||||
@@ -95,10 +103,8 @@ github.com/prometheus/common v0.48.0 h1:QO8U2CdOzSn1BBsmXJXduaaW+dY/5QLjfB8svtSz
|
|||||||
github.com/prometheus/common v0.48.0/go.mod h1:0/KsvlIEfPQCQ5I2iNSAWKPZziNCvRs5EC6ILDTlAPc=
|
github.com/prometheus/common v0.48.0/go.mod h1:0/KsvlIEfPQCQ5I2iNSAWKPZziNCvRs5EC6ILDTlAPc=
|
||||||
github.com/prometheus/procfs v0.12.0 h1:jluTpSng7V9hY0O2R9DzzJHYb2xULk9VTR1V1R/k6Bo=
|
github.com/prometheus/procfs v0.12.0 h1:jluTpSng7V9hY0O2R9DzzJHYb2xULk9VTR1V1R/k6Bo=
|
||||||
github.com/prometheus/procfs v0.12.0/go.mod h1:pcuDEFsWDnvcgNzo4EEweacyhjeA9Zk3cnaOZAZEfOo=
|
github.com/prometheus/procfs v0.12.0/go.mod h1:pcuDEFsWDnvcgNzo4EEweacyhjeA9Zk3cnaOZAZEfOo=
|
||||||
github.com/quic-go/go-ossfuzz-seeds v0.1.0 h1:APacT+iIaNF6fd8AGEiN3bT/Jtkd2jz4v4TzM7MFjy0=
|
github.com/quic-go/quic-go v0.55.0 h1:zccPQIqYCXDt5NmcEabyYvOnomjs8Tlwl7tISjJh9Mk=
|
||||||
github.com/quic-go/go-ossfuzz-seeds v0.1.0/go.mod h1:3IOHRbJIc+L6YKMwfDtJAM9Vj9k0YY4muhuyUYk5tbk=
|
github.com/quic-go/quic-go v0.55.0/go.mod h1:DR51ilwU1uE164KuWXhinFcKWGlEjzys2l8zUl5Ss1U=
|
||||||
github.com/quic-go/quic-go v0.60.0 h1:xcQioE8OM66UQLeUMHltK1CCcOu3JbVB4JAQdDQSB+0=
|
|
||||||
github.com/quic-go/quic-go v0.60.0/go.mod h1:wpKpjmPpftl30sL6pFh7REVpjbcCVy4zt2vDyK1TuJk=
|
|
||||||
github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY=
|
github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY=
|
||||||
github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc=
|
github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc=
|
||||||
github.com/rodaine/table v1.2.0 h1:38HEnwK4mKSHQJIkavVj+bst1TEY7j9zhLMWu4QJrMA=
|
github.com/rodaine/table v1.2.0 h1:38HEnwK4mKSHQJIkavVj+bst1TEY7j9zhLMWu4QJrMA=
|
||||||
@@ -140,6 +146,8 @@ github.com/vishvananda/netlink v1.3.0 h1:X7l42GfcV4S6E4vHTsw48qbrV+9PVojNfIhZcwQ
|
|||||||
github.com/vishvananda/netlink v1.3.0/go.mod h1:i6NetklAujEcC6fK0JPjT8qSwWyO0HLn4UKG+hGqeJs=
|
github.com/vishvananda/netlink v1.3.0/go.mod h1:i6NetklAujEcC6fK0JPjT8qSwWyO0HLn4UKG+hGqeJs=
|
||||||
github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8=
|
github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8=
|
||||||
github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM=
|
github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM=
|
||||||
|
github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU=
|
||||||
|
github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA=
|
||||||
github.com/xtaci/kcp-go/v5 v5.6.13 h1:FEjtz9+D4p8t2x4WjciGt/jsIuhlWjjgPCCWjrVR4Hk=
|
github.com/xtaci/kcp-go/v5 v5.6.13 h1:FEjtz9+D4p8t2x4WjciGt/jsIuhlWjjgPCCWjrVR4Hk=
|
||||||
github.com/xtaci/kcp-go/v5 v5.6.13/go.mod h1:75S1AKYYzNUSXIv30h+jPKJYZUwqpfvLshu63nCNSOM=
|
github.com/xtaci/kcp-go/v5 v5.6.13/go.mod h1:75S1AKYYzNUSXIv30h+jPKJYZUwqpfvLshu63nCNSOM=
|
||||||
github.com/xtaci/lossyconn v0.0.0-20200209145036-adba10fffc37 h1:EWU6Pktpas0n8lLQwDsRyZfmkPeRbdgPtW609es+/9E=
|
github.com/xtaci/lossyconn v0.0.0-20200209145036-adba10fffc37 h1:EWU6Pktpas0n8lLQwDsRyZfmkPeRbdgPtW609es+/9E=
|
||||||
@@ -151,28 +159,30 @@ go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
|||||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||||
golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||||
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
|
golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4=
|
||||||
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
|
golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA=
|
||||||
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
|
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
|
||||||
golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE=
|
golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE=
|
||||||
golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU=
|
golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU=
|
||||||
golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc=
|
golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc=
|
||||||
|
golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8=
|
||||||
|
golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w=
|
||||||
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
||||||
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
||||||
golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0=
|
||||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
|
||||||
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
|
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
|
||||||
golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
|
golang.org/x/oauth2 v0.28.0 h1:CrgCKl8PPAVtLnU3c+EDw6x11699EWlsDeWNWKdIOkc=
|
||||||
golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q=
|
golang.org/x/oauth2 v0.28.0/go.mod h1:onh5ek6nERTohokkhCD/y2cV4Do3fxFHFuAejCkRWT8=
|
||||||
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||||
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
@@ -180,14 +190,14 @@ golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7w
|
|||||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo=
|
||||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
golang.org/x/term v0.41.0 h1:QCgPso/Q3RTJx2Th4bDLqML4W6iJiaXFq2/ftQF13YU=
|
||||||
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
golang.org/x/term v0.41.0/go.mod h1:3pfBgksrReYfZ5lvYM0kSO0LIkAl4Yl2bXOkKP7Ec2A=
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||||
golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs=
|
golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8=
|
||||||
golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY=
|
golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA=
|
||||||
golang.org/x/time v0.10.0 h1:3usCWA8tQn0L8+hFJQNgzpWbd89begxN66o1Ojdn5L4=
|
golang.org/x/time v0.10.0 h1:3usCWA8tQn0L8+hFJQNgzpWbd89begxN66o1Ojdn5L4=
|
||||||
golang.org/x/time v0.10.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
|
golang.org/x/time v0.10.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
|
||||||
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||||
@@ -195,8 +205,8 @@ golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGm
|
|||||||
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
|
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
|
||||||
golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
|
golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
|
||||||
golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q=
|
golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q=
|
||||||
golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=
|
golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k=
|
||||||
golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA=
|
golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0=
|
||||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
||||||
|
|||||||
@@ -394,10 +394,6 @@ func LoadClientConfigResult(path string, strict bool) (*ClientConfigLoadResult,
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := validateNoDuplicateNames(result.Proxies, result.Visitors); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -421,31 +417,6 @@ func LoadClientConfig(path string, strict bool) (
|
|||||||
return result.Common, proxyCfgs, visitorCfgs, result.IsLegacyFormat, nil
|
return result.Common, proxyCfgs, visitorCfgs, result.IsLegacyFormat, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// validateNoDuplicateNames rejects proxies or visitors that share a name. They are
|
|
||||||
// keyed by name in the config sources, so a duplicate would otherwise be silently
|
|
||||||
// overwritten and never started, with no error or log.
|
|
||||||
func validateNoDuplicateNames(proxies []v1.ProxyConfigurer, visitors []v1.VisitorConfigurer) error {
|
|
||||||
proxyNames := make(map[string]struct{}, len(proxies))
|
|
||||||
for _, p := range proxies {
|
|
||||||
name := p.GetBaseConfig().Name
|
|
||||||
if _, ok := proxyNames[name]; ok {
|
|
||||||
return fmt.Errorf("proxy name [%s] is duplicated", name)
|
|
||||||
}
|
|
||||||
proxyNames[name] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
visitorNames := make(map[string]struct{}, len(visitors))
|
|
||||||
for _, v := range visitors {
|
|
||||||
name := v.GetBaseConfig().Name
|
|
||||||
if _, ok := visitorNames[name]; ok {
|
|
||||||
return fmt.Errorf("visitor name [%s] is duplicated", name)
|
|
||||||
}
|
|
||||||
visitorNames[name] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func CompleteProxyConfigurers(proxies []v1.ProxyConfigurer) []v1.ProxyConfigurer {
|
func CompleteProxyConfigurers(proxies []v1.ProxyConfigurer) []v1.ProxyConfigurer {
|
||||||
proxyCfgs := proxies
|
proxyCfgs := proxies
|
||||||
for _, c := range proxyCfgs {
|
for _, c := range proxyCfgs {
|
||||||
|
|||||||
@@ -17,8 +17,6 @@ package config
|
|||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -464,111 +462,6 @@ func TestFilterClientConfigurers_FilterByStartAndEnabled(t *testing.T) {
|
|||||||
require.Equal("keep", proxies[0].GetBaseConfig().Name)
|
require.Equal("keep", proxies[0].GetBaseConfig().Name)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestLoadClientConfigResult_DuplicateNames(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
content string
|
|
||||||
errSubstr string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "duplicate proxy names",
|
|
||||||
content: `
|
|
||||||
serverAddr = "127.0.0.1"
|
|
||||||
serverPort = 7000
|
|
||||||
|
|
||||||
[[proxies]]
|
|
||||||
name = "dup"
|
|
||||||
type = "tcp"
|
|
||||||
localPort = 22
|
|
||||||
remotePort = 6000
|
|
||||||
|
|
||||||
[[proxies]]
|
|
||||||
name = "dup"
|
|
||||||
type = "tcp"
|
|
||||||
localPort = 3306
|
|
||||||
remotePort = 6001
|
|
||||||
`,
|
|
||||||
errSubstr: "proxy name [dup] is duplicated",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "duplicate visitor names",
|
|
||||||
content: `
|
|
||||||
serverAddr = "127.0.0.1"
|
|
||||||
serverPort = 7000
|
|
||||||
|
|
||||||
[[visitors]]
|
|
||||||
name = "dup"
|
|
||||||
type = "stcp"
|
|
||||||
serverName = "a"
|
|
||||||
secretKey = "secret"
|
|
||||||
bindPort = 9001
|
|
||||||
|
|
||||||
[[visitors]]
|
|
||||||
name = "dup"
|
|
||||||
type = "stcp"
|
|
||||||
serverName = "b"
|
|
||||||
secretKey = "secret"
|
|
||||||
bindPort = 9002
|
|
||||||
`,
|
|
||||||
errSubstr: "visitor name [dup] is duplicated",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "unique names",
|
|
||||||
content: `
|
|
||||||
serverAddr = "127.0.0.1"
|
|
||||||
serverPort = 7000
|
|
||||||
|
|
||||||
[[proxies]]
|
|
||||||
name = "p1"
|
|
||||||
type = "tcp"
|
|
||||||
localPort = 22
|
|
||||||
remotePort = 6000
|
|
||||||
|
|
||||||
[[proxies]]
|
|
||||||
name = "p2"
|
|
||||||
type = "tcp"
|
|
||||||
localPort = 3306
|
|
||||||
remotePort = 6001
|
|
||||||
`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "same name across proxy and visitor",
|
|
||||||
content: `
|
|
||||||
serverAddr = "127.0.0.1"
|
|
||||||
serverPort = 7000
|
|
||||||
|
|
||||||
[[proxies]]
|
|
||||||
name = "same"
|
|
||||||
type = "tcp"
|
|
||||||
localPort = 22
|
|
||||||
remotePort = 6000
|
|
||||||
|
|
||||||
[[visitors]]
|
|
||||||
name = "same"
|
|
||||||
type = "stcp"
|
|
||||||
serverName = "a"
|
|
||||||
secretKey = "secret"
|
|
||||||
bindPort = 9001
|
|
||||||
`,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
require := require.New(t)
|
|
||||||
path := filepath.Join(t.TempDir(), "frpc.toml")
|
|
||||||
require.NoError(os.WriteFile(path, []byte(tc.content), 0o600))
|
|
||||||
|
|
||||||
_, err := LoadClientConfigResult(path, false)
|
|
||||||
if tc.errSubstr == "" {
|
|
||||||
require.NoError(err)
|
|
||||||
} else {
|
|
||||||
require.ErrorContains(err, tc.errSubstr)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestYAMLEdgeCases tests edge cases for YAML parsing, including non-map types
|
// TestYAMLEdgeCases tests edge cases for YAML parsing, including non-map types
|
||||||
func TestYAMLEdgeCases(t *testing.T) {
|
func TestYAMLEdgeCases(t *testing.T) {
|
||||||
require := require.New(t)
|
require := require.New(t)
|
||||||
|
|||||||
@@ -167,7 +167,7 @@ type ServerTransportConfig struct {
|
|||||||
// If negative, keep-alive probes are disabled.
|
// If negative, keep-alive probes are disabled.
|
||||||
TCPKeepAlive int64 `json:"tcpKeepalive,omitempty"`
|
TCPKeepAlive int64 `json:"tcpKeepalive,omitempty"`
|
||||||
// MaxPoolCount specifies the maximum pool size for each proxy. By default,
|
// MaxPoolCount specifies the maximum pool size for each proxy. By default,
|
||||||
// this value is 5. Negative values are invalid.
|
// this value is 5.
|
||||||
MaxPoolCount int64 `json:"maxPoolCount,omitempty"`
|
MaxPoolCount int64 `json:"maxPoolCount,omitempty"`
|
||||||
// HeartBeatTimeout specifies the maximum time to wait for a heartbeat
|
// HeartBeatTimeout specifies the maximum time to wait for a heartbeat
|
||||||
// before terminating the connection. It is not recommended to change this
|
// before terminating the connection. It is not recommended to change this
|
||||||
|
|||||||
@@ -51,51 +51,14 @@ func (v *ConfigValidator) ValidateClientCommonConfig(c *v1.ClientCommonConfig) (
|
|||||||
}
|
}
|
||||||
|
|
||||||
func validateFeatureGates(c *v1.ClientCommonConfig) (Warning, error) {
|
func validateFeatureGates(c *v1.ClientCommonConfig) (Warning, error) {
|
||||||
gates := featuregate.NewFeatureGate()
|
|
||||||
if err := gates.SetFromMap(c.FeatureGates); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.VirtualNet.Address != "" {
|
if c.VirtualNet.Address != "" {
|
||||||
if !gates.Enabled(featuregate.VirtualNet) {
|
if !featuregate.Enabled(featuregate.VirtualNet) {
|
||||||
return nil, fmt.Errorf("VirtualNet feature is not enabled; enable it by setting the appropriate feature gate flag")
|
return nil, fmt.Errorf("VirtualNet feature is not enabled; enable it by setting the appropriate feature gate flag")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ClientConfigRequirements describes runtime capabilities needed by a client configuration.
|
|
||||||
type ClientConfigRequirements struct {
|
|
||||||
VirtualNet bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetClientConfigRequirements returns the runtime capabilities needed by a client configuration.
|
|
||||||
func GetClientConfigRequirements(
|
|
||||||
common *v1.ClientCommonConfig,
|
|
||||||
proxyCfgs []v1.ProxyConfigurer,
|
|
||||||
visitorCfgs []v1.VisitorConfigurer,
|
|
||||||
) ClientConfigRequirements {
|
|
||||||
requirements := ClientConfigRequirements{}
|
|
||||||
if common != nil && common.VirtualNet.Address != "" {
|
|
||||||
requirements.VirtualNet = true
|
|
||||||
}
|
|
||||||
for _, cfg := range proxyCfgs {
|
|
||||||
if cfg.GetBaseConfig().Plugin.Type == v1.PluginVirtualNet {
|
|
||||||
requirements.VirtualNet = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !requirements.VirtualNet {
|
|
||||||
for _, cfg := range visitorCfgs {
|
|
||||||
if cfg.GetBaseConfig().Plugin.Type == v1.VisitorPluginVirtualNet {
|
|
||||||
requirements.VirtualNet = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return requirements
|
|
||||||
}
|
|
||||||
|
|
||||||
func (v *ConfigValidator) validateAuthConfig(c *v1.AuthClientConfig) (Warning, error) {
|
func (v *ConfigValidator) validateAuthConfig(c *v1.AuthClientConfig) (Warning, error) {
|
||||||
var errs error
|
var errs error
|
||||||
if !slices.Contains(SupportedAuthMethods, c.Method) {
|
if !slices.Contains(SupportedAuthMethods, c.Method) {
|
||||||
|
|||||||
@@ -1,140 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package validation
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/policy/featuregate"
|
|
||||||
"github.com/fatedier/frp/pkg/policy/security"
|
|
||||||
)
|
|
||||||
|
|
||||||
func validateClientFeatureGates(t *testing.T, gates map[string]bool, virtualNetAddress string) error {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
cfg := &v1.ClientCommonConfig{
|
|
||||||
FeatureGates: gates,
|
|
||||||
VirtualNet: v1.VirtualNetConfig{
|
|
||||||
Address: virtualNetAddress,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
require.NoError(t, cfg.Complete())
|
|
||||||
|
|
||||||
_, err := NewConfigValidator(security.NewUnsafeFeatures(nil)).ValidateClientCommonConfig(cfg)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestValidateClientFeatureGates(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
featureGates map[string]bool
|
|
||||||
virtualNetAddress string
|
|
||||||
wantErr string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "VirtualNet enabled",
|
|
||||||
featureGates: map[string]bool{"VirtualNet": true},
|
|
||||||
virtualNetAddress: "100.86.0.4/24",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "VirtualNet explicitly disabled",
|
|
||||||
featureGates: map[string]bool{"VirtualNet": false},
|
|
||||||
virtualNetAddress: "100.86.0.4/24",
|
|
||||||
wantErr: "VirtualNet feature is not enabled",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "VirtualNet disabled by default",
|
|
||||||
virtualNetAddress: "100.86.0.4/24",
|
|
||||||
wantErr: "VirtualNet feature is not enabled",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "unknown feature gate",
|
|
||||||
featureGates: map[string]bool{"UnknownFeature": true},
|
|
||||||
wantErr: "unrecognized feature gate: UnknownFeature",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
err := validateClientFeatureGates(t, tc.featureGates, tc.virtualNetAddress)
|
|
||||||
if tc.wantErr == "" {
|
|
||||||
require.NoError(t, err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
require.ErrorContains(t, err, tc.wantErr)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetClientConfigRequirements(t *testing.T) {
|
|
||||||
virtualNetProxy := &v1.STCPProxyConfig{
|
|
||||||
ProxyBaseConfig: v1.ProxyBaseConfig{
|
|
||||||
ProxyBackend: v1.ProxyBackend{
|
|
||||||
Plugin: v1.TypedClientPluginOptions{Type: v1.PluginVirtualNet},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
virtualNetVisitor := &v1.STCPVisitorConfig{
|
|
||||||
VisitorBaseConfig: v1.VisitorBaseConfig{
|
|
||||||
Plugin: v1.TypedVisitorPluginOptions{Type: v1.VisitorPluginVirtualNet},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
common *v1.ClientCommonConfig
|
|
||||||
proxies []v1.ProxyConfigurer
|
|
||||||
visitors []v1.VisitorConfigurer
|
|
||||||
wantVNet bool
|
|
||||||
}{
|
|
||||||
{name: "no requirements"},
|
|
||||||
{
|
|
||||||
name: "common VirtualNet address",
|
|
||||||
common: &v1.ClientCommonConfig{VirtualNet: v1.VirtualNetConfig{Address: "100.86.0.4/24"}},
|
|
||||||
wantVNet: true,
|
|
||||||
},
|
|
||||||
{name: "VirtualNet proxy", proxies: []v1.ProxyConfigurer{virtualNetProxy}, wantVNet: true},
|
|
||||||
{name: "VirtualNet visitor", visitors: []v1.VisitorConfigurer{virtualNetVisitor}, wantVNet: true},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
got := GetClientConfigRequirements(tc.common, tc.proxies, tc.visitors)
|
|
||||||
require.Equal(t, tc.wantVNet, got.VirtualNet)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestValidateClientFeatureGatesAreConfigScoped(t *testing.T) {
|
|
||||||
defaultGatesBefore := featuregate.DefaultFeatureGates.String()
|
|
||||||
|
|
||||||
require.NoError(t, validateClientFeatureGates(
|
|
||||||
t,
|
|
||||||
map[string]bool{"VirtualNet": true},
|
|
||||||
"100.86.0.4/24",
|
|
||||||
))
|
|
||||||
require.Equal(t, defaultGatesBefore, featuregate.DefaultFeatureGates.String())
|
|
||||||
|
|
||||||
err := validateClientFeatureGates(
|
|
||||||
t,
|
|
||||||
map[string]bool{"VirtualNet": false},
|
|
||||||
"100.86.0.4/24",
|
|
||||||
)
|
|
||||||
require.ErrorContains(t, err, "VirtualNet feature is not enabled")
|
|
||||||
require.Equal(t, defaultGatesBefore, featuregate.DefaultFeatureGates.String())
|
|
||||||
}
|
|
||||||
@@ -1,48 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package validation
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"unicode"
|
|
||||||
"unicode/utf8"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
// MaxRunIDLength is the maximum number of bytes accepted for a control run ID.
|
|
||||||
MaxRunIDLength = 64
|
|
||||||
)
|
|
||||||
|
|
||||||
func validateIdentifier(value, kind string, maxLength int) error {
|
|
||||||
if value == "" {
|
|
||||||
return fmt.Errorf("%s cannot be empty", kind)
|
|
||||||
}
|
|
||||||
if len(value) > maxLength {
|
|
||||||
return fmt.Errorf("%s is too long: length %d exceeds maximum %d", kind, len(value), maxLength)
|
|
||||||
}
|
|
||||||
if !utf8.ValidString(value) {
|
|
||||||
return fmt.Errorf("%s must be valid UTF-8", kind)
|
|
||||||
}
|
|
||||||
for _, r := range value {
|
|
||||||
if !unicode.IsPrint(r) {
|
|
||||||
return fmt.Errorf("%s contains non-printable character", kind)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func ValidateRunID(runID string) error {
|
|
||||||
return validateIdentifier(runID, "run id", MaxRunIDLength)
|
|
||||||
}
|
|
||||||
@@ -1,48 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package validation
|
|
||||||
|
|
||||||
import (
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestValidateIdentifiers(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
validate func(string) error
|
|
||||||
value string
|
|
||||||
wantError string
|
|
||||||
}{
|
|
||||||
{name: "run id accepts printable values", validate: ValidateRunID, value: "run-%1000s-中文"},
|
|
||||||
{name: "run id rejects empty", validate: ValidateRunID, wantError: "cannot be empty"},
|
|
||||||
{name: "run id rejects control character", validate: ValidateRunID, value: "run\nforged", wantError: "non-printable"},
|
|
||||||
{name: "run id rejects invalid utf8", validate: ValidateRunID, value: string([]byte{0xff}), wantError: "valid UTF-8"},
|
|
||||||
{name: "run id rejects excessive length", validate: ValidateRunID, value: strings.Repeat("a", MaxRunIDLength+1), wantError: "too long"},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
err := tt.validate(tt.value)
|
|
||||||
if tt.wantError == "" {
|
|
||||||
require.NoError(t, err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
require.ErrorContains(t, err, tt.wantError)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -79,11 +79,9 @@ func validateDomainConfigForClient(c *v1.DomainConfig) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func validateDomainConfigForServer(c *v1.DomainConfig, s *v1.ServerConfig) error {
|
func validateDomainConfigForServer(c *v1.DomainConfig, s *v1.ServerConfig) error {
|
||||||
subDomainHost := strings.ToLower(s.SubDomainHost)
|
|
||||||
for _, domain := range c.CustomDomains {
|
for _, domain := range c.CustomDomains {
|
||||||
canonicalDomain := strings.ToLower(domain)
|
if s.SubDomainHost != "" && len(strings.Split(s.SubDomainHost, ".")) < len(strings.Split(domain, ".")) {
|
||||||
if subDomainHost != "" && len(strings.Split(subDomainHost, ".")) < len(strings.Split(canonicalDomain, ".")) {
|
if strings.HasSuffix(domain, "."+s.SubDomainHost) {
|
||||||
if strings.HasSuffix(canonicalDomain, "."+subDomainHost) {
|
|
||||||
return fmt.Errorf("custom domain [%s] should not belong to subdomain host [%s]", domain, s.SubDomainHost)
|
return fmt.Errorf("custom domain [%s] should not belong to subdomain host [%s]", domain, s.SubDomainHost)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,76 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package validation
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestValidateDomainConfigForServerRejectsSubdomainHostCaseInsensitively(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
subDomainHost string
|
|
||||||
customDomain string
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "lowercase subdomain",
|
|
||||||
subDomainHost: "frp.example.com",
|
|
||||||
customDomain: "victim.frp.example.com",
|
|
||||||
wantErr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "mixed case custom domain",
|
|
||||||
subDomainHost: "frp.example.com",
|
|
||||||
customDomain: "victim.FRP.example.com",
|
|
||||||
wantErr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "mixed case wildcard domain",
|
|
||||||
subDomainHost: "frp.example.com",
|
|
||||||
customDomain: "*.FRP.example.com",
|
|
||||||
wantErr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "mixed case subdomain host",
|
|
||||||
subDomainHost: "FRP.Example.Com",
|
|
||||||
customDomain: "victim.frp.example.com",
|
|
||||||
wantErr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "external domain",
|
|
||||||
subDomainHost: "frp.example.com",
|
|
||||||
customDomain: "victim.example.net",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
err := validateDomainConfigForServer(
|
|
||||||
&v1.DomainConfig{CustomDomains: []string{tt.customDomain}},
|
|
||||||
&v1.ServerConfig{SubDomainHost: tt.subDomainHost},
|
|
||||||
)
|
|
||||||
if tt.wantErr {
|
|
||||||
require.ErrorContains(t, err, "should not belong to subdomain host")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
require.NoError(t, err)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -51,9 +51,6 @@ func (v *ConfigValidator) ValidateServerConfig(c *v1.ServerConfig) (Warning, err
|
|||||||
errs = AppendError(errs, ValidatePort(c.VhostHTTPPort, "vhostHTTPPort"))
|
errs = AppendError(errs, ValidatePort(c.VhostHTTPPort, "vhostHTTPPort"))
|
||||||
errs = AppendError(errs, ValidatePort(c.VhostHTTPSPort, "vhostHTTPSPort"))
|
errs = AppendError(errs, ValidatePort(c.VhostHTTPSPort, "vhostHTTPSPort"))
|
||||||
errs = AppendError(errs, ValidatePort(c.TCPMuxHTTPConnectPort, "tcpMuxHTTPConnectPort"))
|
errs = AppendError(errs, ValidatePort(c.TCPMuxHTTPConnectPort, "tcpMuxHTTPConnectPort"))
|
||||||
if c.Transport.MaxPoolCount < 0 {
|
|
||||||
errs = AppendError(errs, fmt.Errorf("invalid transport.maxPoolCount, must be non-negative"))
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, p := range c.HTTPPlugins {
|
for _, p := range c.HTTPPlugins {
|
||||||
if !lo.Every(SupportedHTTPPluginOps, p.Ops) {
|
if !lo.Every(SupportedHTTPPluginOps, p.Ops) {
|
||||||
|
|||||||
@@ -1,51 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package validation
|
|
||||||
|
|
||||||
import (
|
|
||||||
"math"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestValidateServerConfigMaxPoolCount(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
maxPoolCount int64
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{name: "negative", maxPoolCount: -1, wantErr: true},
|
|
||||||
{name: "zero", maxPoolCount: 0},
|
|
||||||
{name: "positive", maxPoolCount: 5},
|
|
||||||
{name: "maximum int64", maxPoolCount: math.MaxInt64},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
cfg := validServerConfigWithAuth(v1.AuthServerConfig{Method: v1.AuthMethodToken})
|
|
||||||
cfg.Transport.MaxPoolCount = tc.maxPoolCount
|
|
||||||
require.NoError(t, cfg.Complete())
|
|
||||||
|
|
||||||
_, err := NewConfigValidator(nil).ValidateServerConfig(cfg)
|
|
||||||
if tc.wantErr {
|
|
||||||
require.ErrorContains(t, err, "invalid transport.maxPoolCount")
|
|
||||||
require.ErrorContains(t, err, "must be non-negative")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
require.NoError(t, err)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -95,7 +95,9 @@ func (m *serverMetrics) clearUselessInfo(continuousOfflineDuration time.Duration
|
|||||||
defer m.mu.Unlock()
|
defer m.mu.Unlock()
|
||||||
total = len(m.info.ProxyStatistics)
|
total = len(m.info.ProxyStatistics)
|
||||||
for name, data := range m.info.ProxyStatistics {
|
for name, data := range m.info.ProxyStatistics {
|
||||||
if m.shouldClearProxyStats(data, continuousOfflineDuration) {
|
if !data.LastCloseTime.IsZero() &&
|
||||||
|
data.LastStartTime.Before(data.LastCloseTime) &&
|
||||||
|
m.clock.Since(data.LastCloseTime) > continuousOfflineDuration {
|
||||||
delete(m.info.ProxyStatistics, name)
|
delete(m.info.ProxyStatistics, name)
|
||||||
count++
|
count++
|
||||||
log.Tracef("clear proxy [%s]'s statistics data, lastCloseTime: [%s]", name, data.LastCloseTime.String())
|
log.Tracef("clear proxy [%s]'s statistics data, lastCloseTime: [%s]", name, data.LastCloseTime.String())
|
||||||
@@ -104,20 +106,10 @@ func (m *serverMetrics) clearUselessInfo(continuousOfflineDuration time.Duration
|
|||||||
return count, total
|
return count, total
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *serverMetrics) shouldClearProxyStats(data *ProxyStatistics, continuousOfflineDuration time.Duration) bool {
|
|
||||||
return !data.LastCloseTime.IsZero() &&
|
|
||||||
data.LastStartTime.Before(data.LastCloseTime) &&
|
|
||||||
m.clock.Since(data.LastCloseTime) > continuousOfflineDuration
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *serverMetrics) ClearOfflineProxies() (int, int) {
|
func (m *serverMetrics) ClearOfflineProxies() (int, int) {
|
||||||
return m.clearUselessInfo(0)
|
return m.clearUselessInfo(0)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *serverMetrics) PruneOfflineProxies() (int, int) {
|
|
||||||
return m.clearUselessInfo(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *serverMetrics) NewClient() {
|
func (m *serverMetrics) NewClient() {
|
||||||
m.info.ClientCounts.Inc(1)
|
m.info.ClientCounts.Inc(1)
|
||||||
}
|
}
|
||||||
@@ -239,11 +231,9 @@ func toProxyStats(name string, proxyStats *ProxyStatistics) *ProxyStats {
|
|||||||
}
|
}
|
||||||
if !proxyStats.LastStartTime.IsZero() {
|
if !proxyStats.LastStartTime.IsZero() {
|
||||||
ps.LastStartTime = proxyStats.LastStartTime.Format("01-02 15:04:05")
|
ps.LastStartTime = proxyStats.LastStartTime.Format("01-02 15:04:05")
|
||||||
ps.LastStartAt = proxyStats.LastStartTime.Unix()
|
|
||||||
}
|
}
|
||||||
if !proxyStats.LastCloseTime.IsZero() {
|
if !proxyStats.LastCloseTime.IsZero() {
|
||||||
ps.LastCloseTime = proxyStats.LastCloseTime.Format("01-02 15:04:05")
|
ps.LastCloseTime = proxyStats.LastCloseTime.Format("01-02 15:04:05")
|
||||||
ps.LastCloseAt = proxyStats.LastCloseTime.Unix()
|
|
||||||
}
|
}
|
||||||
return ps
|
return ps
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -22,12 +22,6 @@ func TestServerMetricsUsesClockForProxyTimestamps(t *testing.T) {
|
|||||||
clk.SetTime(closedAt)
|
clk.SetTime(closedAt)
|
||||||
metrics.CloseProxy("proxy", "tcp")
|
metrics.CloseProxy("proxy", "tcp")
|
||||||
require.Equal(closedAt, metrics.info.ProxyStatistics["proxy"].LastCloseTime)
|
require.Equal(closedAt, metrics.info.ProxyStatistics["proxy"].LastCloseTime)
|
||||||
|
|
||||||
stats := metrics.GetProxyByName("proxy")
|
|
||||||
require.Equal(start.Format("01-02 15:04:05"), stats.LastStartTime)
|
|
||||||
require.Equal(closedAt.Format("01-02 15:04:05"), stats.LastCloseTime)
|
|
||||||
require.Equal(start.Unix(), stats.LastStartAt)
|
|
||||||
require.Equal(closedAt.Unix(), stats.LastCloseAt)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestServerMetricsClearUselessInfoUsesClock(t *testing.T) {
|
func TestServerMetricsClearUselessInfoUsesClock(t *testing.T) {
|
||||||
@@ -49,70 +43,6 @@ func TestServerMetricsClearUselessInfoUsesClock(t *testing.T) {
|
|||||||
require.Empty(metrics.info.ProxyStatistics)
|
require.Empty(metrics.info.ProxyStatistics)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestServerMetricsClearOfflineProxiesPreservesLegacyTotal(t *testing.T) {
|
|
||||||
require := require.New(t)
|
|
||||||
|
|
||||||
start := time.Date(2026, time.May, 8, 12, 30, 0, 0, time.UTC)
|
|
||||||
clk := clocktesting.NewFakeClock(start.Add(time.Minute))
|
|
||||||
metrics := newServerMetricsWithClock(clk)
|
|
||||||
metrics.info.ProxyStatistics["offline"] = &ProxyStatistics{
|
|
||||||
Name: "offline",
|
|
||||||
LastStartTime: start.Add(-time.Hour),
|
|
||||||
LastCloseTime: start,
|
|
||||||
}
|
|
||||||
metrics.info.ProxyStatistics["online"] = &ProxyStatistics{
|
|
||||||
Name: "online",
|
|
||||||
LastStartTime: start,
|
|
||||||
}
|
|
||||||
|
|
||||||
cleared, total := metrics.ClearOfflineProxies()
|
|
||||||
|
|
||||||
require.Equal(1, cleared)
|
|
||||||
require.Equal(2, total)
|
|
||||||
require.False(metrics.hasProxyStatistics("offline"))
|
|
||||||
require.True(metrics.hasProxyStatistics("online"))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServerMetricsPruneOfflineProxiesReportsTotalStats(t *testing.T) {
|
|
||||||
require := require.New(t)
|
|
||||||
|
|
||||||
start := time.Date(2026, time.May, 8, 12, 30, 0, 0, time.UTC)
|
|
||||||
clk := clocktesting.NewFakeClock(start.Add(time.Minute))
|
|
||||||
metrics := newServerMetricsWithClock(clk)
|
|
||||||
metrics.info.ProxyStatistics["offline"] = &ProxyStatistics{
|
|
||||||
Name: "offline",
|
|
||||||
LastStartTime: start.Add(-time.Hour),
|
|
||||||
LastCloseTime: start,
|
|
||||||
}
|
|
||||||
metrics.info.ProxyStatistics["online"] = &ProxyStatistics{
|
|
||||||
Name: "online",
|
|
||||||
LastStartTime: start,
|
|
||||||
}
|
|
||||||
metrics.info.ProxyStatistics["restarted"] = &ProxyStatistics{
|
|
||||||
Name: "restarted",
|
|
||||||
LastStartTime: start.Add(30 * time.Second),
|
|
||||||
LastCloseTime: start,
|
|
||||||
}
|
|
||||||
metrics.info.ProxyStatistics["same-time"] = &ProxyStatistics{
|
|
||||||
Name: "same-time",
|
|
||||||
LastStartTime: start,
|
|
||||||
LastCloseTime: start,
|
|
||||||
}
|
|
||||||
|
|
||||||
cleared, total := metrics.PruneOfflineProxies()
|
|
||||||
|
|
||||||
require.Equal(1, cleared)
|
|
||||||
require.Equal(4, total)
|
|
||||||
require.False(metrics.hasProxyStatistics("offline"))
|
|
||||||
require.True(metrics.hasProxyStatistics("online"))
|
|
||||||
require.True(metrics.hasProxyStatistics("restarted"))
|
|
||||||
require.True(metrics.hasProxyStatistics("same-time"))
|
|
||||||
|
|
||||||
cleared, total = metrics.PruneOfflineProxies()
|
|
||||||
require.Equal(0, cleared)
|
|
||||||
require.Equal(3, total)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServerMetricsRunUsesClockTicker(t *testing.T) {
|
func TestServerMetricsRunUsesClockTicker(t *testing.T) {
|
||||||
require := require.New(t)
|
require := require.New(t)
|
||||||
|
|
||||||
|
|||||||
@@ -41,8 +41,6 @@ type ProxyStats struct {
|
|||||||
TodayTrafficOut int64
|
TodayTrafficOut int64
|
||||||
LastStartTime string
|
LastStartTime string
|
||||||
LastCloseTime string
|
LastCloseTime string
|
||||||
LastStartAt int64
|
|
||||||
LastCloseAt int64
|
|
||||||
CurConns int64
|
CurConns int64
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -87,5 +85,4 @@ type Collector interface {
|
|||||||
GetProxyByName(proxyName string) *ProxyStats
|
GetProxyByName(proxyName string) *ProxyStats
|
||||||
GetProxyTraffic(name string) *ProxyTrafficInfo
|
GetProxyTraffic(name string) *ProxyTrafficInfo
|
||||||
ClearOfflineProxies() (int, int)
|
ClearOfflineProxies() (int, int)
|
||||||
PruneOfflineProxies() (int, int)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,199 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package msg
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"runtime"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/proto/wire"
|
|
||||||
)
|
|
||||||
|
|
||||||
type udpBenchmarkCase struct {
|
|
||||||
name string
|
|
||||||
packet *UDPPacket
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
|
||||||
udpBenchmarkBytesSink []byte
|
|
||||||
udpBenchmarkMessageSink Message
|
|
||||||
)
|
|
||||||
|
|
||||||
func udpBenchmarkCases(payloadSize int) []udpBenchmarkCase {
|
|
||||||
content := bytes.Repeat([]byte{0x5a}, payloadSize)
|
|
||||||
return []udpBenchmarkCase{
|
|
||||||
{
|
|
||||||
name: "ipv4-remote",
|
|
||||||
packet: &UDPPacket{
|
|
||||||
Content: content,
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 12345},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "ipv4-local-remote",
|
|
||||||
packet: &UDPPacket{
|
|
||||||
Content: content,
|
|
||||||
LocalAddr: &net.UDPAddr{IP: net.ParseIP("192.0.2.2"), Port: 23456},
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 12345},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "ipv6-remote",
|
|
||||||
packet: &UDPPacket{
|
|
||||||
Content: content,
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 12345},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "ipv6-local-remote",
|
|
||||||
packet: &UDPPacket{
|
|
||||||
Content: content,
|
|
||||||
LocalAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::2"), Port: 23456, Zone: "bench0"},
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 12345, Zone: "bench1"},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPPacketV2FrameSizes(t *testing.T) {
|
|
||||||
t.Logf("environment go=%s goos=%s goarch=%s gomaxprocs=%d", runtime.Version(), runtime.GOOS, runtime.GOARCH, runtime.GOMAXPROCS(0))
|
|
||||||
for _, payloadSize := range []int{64, 512, 1200, 1472} {
|
|
||||||
for _, tc := range udpBenchmarkCases(payloadSize) {
|
|
||||||
jsonFrame := udpBenchmarkWireBytes(t, tc.packet, "")
|
|
||||||
binaryFrame := udpBenchmarkWireBytes(t, tc.packet, wire.UDPPacketCodecBinary)
|
|
||||||
saving := 100 * float64(len(jsonFrame)-len(binaryFrame)) / float64(len(jsonFrame))
|
|
||||||
t.Logf("frame payload=%d case=%s json_bytes=%d binary_bytes=%d binary_saving_pct=%.2f", payloadSize, tc.name, len(jsonFrame), len(binaryFrame), saving)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func udpBenchmarkWireBytes(t testing.TB, packet *UDPPacket, codec string) []byte {
|
|
||||||
t.Helper()
|
|
||||||
var buf bytes.Buffer
|
|
||||||
rw, err := NewUDPPacketReadWriter(&buf, wire.ProtocolV2, codec)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("create UDP read writer: %v", err)
|
|
||||||
}
|
|
||||||
if err := rw.WriteMsg(packet); err != nil {
|
|
||||||
t.Fatalf("write UDP packet: %v", err)
|
|
||||||
}
|
|
||||||
return append([]byte(nil), buf.Bytes()...)
|
|
||||||
}
|
|
||||||
|
|
||||||
type udpBenchmarkReadWriter struct {
|
|
||||||
reader bytes.Reader
|
|
||||||
}
|
|
||||||
|
|
||||||
func (rw *udpBenchmarkReadWriter) Read(p []byte) (int, error) {
|
|
||||||
return rw.reader.Read(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (rw *udpBenchmarkReadWriter) Write(p []byte) (int, error) {
|
|
||||||
return len(p), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (rw *udpBenchmarkReadWriter) Reset(p []byte) {
|
|
||||||
rw.reader.Reset(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
func udpBenchmarkValidatePacket(b testing.TB, got, want *UDPPacket) {
|
|
||||||
b.Helper()
|
|
||||||
if !bytes.Equal(got.Content, want.Content) || !udpBenchmarkUDPAddrEqual(got.LocalAddr, want.LocalAddr) ||
|
|
||||||
!udpBenchmarkUDPAddrEqual(got.RemoteAddr, want.RemoteAddr) {
|
|
||||||
b.Fatalf("decoded packet mismatch: got %+v, want %+v", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func udpBenchmarkUDPAddrEqual(got, want *net.UDPAddr) bool {
|
|
||||||
if got == nil || want == nil {
|
|
||||||
return got == want
|
|
||||||
}
|
|
||||||
return got.IP.Equal(want.IP) && got.Port == want.Port && got.Zone == want.Zone
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkUDPPacketV2CodecWrite(b *testing.B) {
|
|
||||||
for _, payloadSize := range []int{64, 512, 1200, 1472} {
|
|
||||||
for _, tc := range udpBenchmarkCases(payloadSize) {
|
|
||||||
for _, codec := range []struct {
|
|
||||||
name string
|
|
||||||
value string
|
|
||||||
}{
|
|
||||||
{name: "json", value: ""},
|
|
||||||
{name: "binary", value: wire.UDPPacketCodecBinary},
|
|
||||||
} {
|
|
||||||
b.Run(fmt.Sprintf("payload-%d/%s/%s", payloadSize, tc.name, codec.name), func(b *testing.B) {
|
|
||||||
var buf bytes.Buffer
|
|
||||||
rw, err := NewUDPPacketReadWriter(&buf, wire.ProtocolV2, codec.value)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
expected := udpBenchmarkWireBytes(b, tc.packet, codec.value)
|
|
||||||
b.SetBytes(int64(len(expected)))
|
|
||||||
for b.Loop() {
|
|
||||||
buf.Reset()
|
|
||||||
if err := rw.WriteMsg(tc.packet); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !bytes.Equal(buf.Bytes(), expected) {
|
|
||||||
b.Fatalf("encoded packet mismatch: got %d bytes, want %d", buf.Len(), len(expected))
|
|
||||||
}
|
|
||||||
udpBenchmarkBytesSink = buf.Bytes()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkUDPPacketV2CodecRead(b *testing.B) {
|
|
||||||
for _, payloadSize := range []int{64, 512, 1200, 1472} {
|
|
||||||
for _, tc := range udpBenchmarkCases(payloadSize) {
|
|
||||||
for _, codec := range []struct {
|
|
||||||
name string
|
|
||||||
value string
|
|
||||||
}{
|
|
||||||
{name: "json", value: ""},
|
|
||||||
{name: "binary", value: wire.UDPPacketCodecBinary},
|
|
||||||
} {
|
|
||||||
b.Run(fmt.Sprintf("payload-%d/%s/%s", payloadSize, tc.name, codec.name), func(b *testing.B) {
|
|
||||||
encoded := udpBenchmarkWireBytes(b, tc.packet, codec.value)
|
|
||||||
stream := &udpBenchmarkReadWriter{}
|
|
||||||
rw, err := NewUDPPacketReadWriter(stream, wire.ProtocolV2, codec.value)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
var decoded Message
|
|
||||||
b.SetBytes(int64(len(encoded)))
|
|
||||||
for b.Loop() {
|
|
||||||
stream.Reset(encoded)
|
|
||||||
decoded, err = rw.ReadMsg()
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
packet, ok := decoded.(*UDPPacket)
|
|
||||||
if !ok {
|
|
||||||
b.Fatalf("decoded message type %T, want *UDPPacket", decoded)
|
|
||||||
}
|
|
||||||
udpBenchmarkValidatePacket(b, packet, tc.packet)
|
|
||||||
udpBenchmarkMessageSink = decoded
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,338 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package msg
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"unicode/utf8"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/proto/wire"
|
|
||||||
)
|
|
||||||
|
|
||||||
const MaxUDPPayloadSize = 65507
|
|
||||||
|
|
||||||
const (
|
|
||||||
udpPacketFlagLocalAddr byte = 1 << 0
|
|
||||||
udpPacketFlagRemoteAddr byte = 1 << 1
|
|
||||||
udpPacketValidFlags = udpPacketFlagLocalAddr | udpPacketFlagRemoteAddr
|
|
||||||
)
|
|
||||||
|
|
||||||
type binaryUDPAddr struct {
|
|
||||||
family byte
|
|
||||||
ip []byte
|
|
||||||
port uint16
|
|
||||||
zone string
|
|
||||||
}
|
|
||||||
|
|
||||||
// EncodeUDPPacketBinary encodes the body of a V2 binary UDP packet message.
|
|
||||||
// RemoteAddr is required by the UDP forwarding path.
|
|
||||||
func EncodeUDPPacketBinary(packet *UDPPacket) ([]byte, error) {
|
|
||||||
if packet == nil {
|
|
||||||
return nil, fmt.Errorf("nil UDP packet")
|
|
||||||
}
|
|
||||||
if packet.RemoteAddr == nil {
|
|
||||||
return nil, fmt.Errorf("UDP packet missing remote address")
|
|
||||||
}
|
|
||||||
if len(packet.Content) > MaxUDPPayloadSize {
|
|
||||||
return nil, fmt.Errorf("UDP payload length %d exceeds limit %d", len(packet.Content), MaxUDPPayloadSize)
|
|
||||||
}
|
|
||||||
|
|
||||||
var flags byte
|
|
||||||
var localAddr, remoteAddr binaryUDPAddr
|
|
||||||
bodyLen := 1 + 2 + len(packet.Content)
|
|
||||||
if packet.LocalAddr != nil {
|
|
||||||
flags |= udpPacketFlagLocalAddr
|
|
||||||
var err error
|
|
||||||
localAddr, err = validateBinaryUDPAddr(packet.LocalAddr)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("local address: %w", err)
|
|
||||||
}
|
|
||||||
bodyLen += binaryUDPAddrLen(localAddr)
|
|
||||||
}
|
|
||||||
flags |= udpPacketFlagRemoteAddr
|
|
||||||
var err error
|
|
||||||
remoteAddr, err = validateBinaryUDPAddr(packet.RemoteAddr)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("remote address: %w", err)
|
|
||||||
}
|
|
||||||
bodyLen += binaryUDPAddrLen(remoteAddr)
|
|
||||||
if 2+bodyLen > wire.DefaultMaxFramePayloadSize {
|
|
||||||
return nil, fmt.Errorf("v2 frame payload length %d exceeds limit %d", 2+bodyLen, wire.DefaultMaxFramePayloadSize)
|
|
||||||
}
|
|
||||||
|
|
||||||
body := make([]byte, bodyLen)
|
|
||||||
body[0] = flags
|
|
||||||
offset := 1
|
|
||||||
if flags&udpPacketFlagLocalAddr != 0 {
|
|
||||||
offset = putBinaryUDPAddr(body, offset, localAddr)
|
|
||||||
}
|
|
||||||
offset = putBinaryUDPAddr(body, offset, remoteAddr)
|
|
||||||
binary.BigEndian.PutUint16(body[offset:offset+2], uint16(len(packet.Content)))
|
|
||||||
offset += 2
|
|
||||||
copy(body[offset:], packet.Content)
|
|
||||||
return body, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// DecodeUDPPacketBinary decodes a V2 binary UDP packet body and returns data
|
|
||||||
// that does not alias the input frame buffer.
|
|
||||||
func DecodeUDPPacketBinary(body []byte) (*UDPPacket, error) {
|
|
||||||
if len(body) < 3 {
|
|
||||||
return nil, fmt.Errorf("UDP packet body too short: %d", len(body))
|
|
||||||
}
|
|
||||||
if 2+len(body) > wire.DefaultMaxFramePayloadSize {
|
|
||||||
return nil, fmt.Errorf("v2 frame payload length %d exceeds limit %d", 2+len(body), wire.DefaultMaxFramePayloadSize)
|
|
||||||
}
|
|
||||||
|
|
||||||
flags := body[0]
|
|
||||||
if flags&^udpPacketValidFlags != 0 {
|
|
||||||
return nil, fmt.Errorf("reserved UDP packet flags set: 0x%02x", flags)
|
|
||||||
}
|
|
||||||
if flags&udpPacketFlagRemoteAddr == 0 {
|
|
||||||
return nil, fmt.Errorf("UDP packet missing remote address")
|
|
||||||
}
|
|
||||||
|
|
||||||
packet := &UDPPacket{}
|
|
||||||
offset := 1
|
|
||||||
var err error
|
|
||||||
if flags&udpPacketFlagLocalAddr != 0 {
|
|
||||||
packet.LocalAddr, offset, err = readBinaryUDPAddr(body, offset)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("local address: %w", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if flags&udpPacketFlagRemoteAddr != 0 {
|
|
||||||
packet.RemoteAddr, offset, err = readBinaryUDPAddr(body, offset)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("remote address: %w", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if len(body)-offset < 2 {
|
|
||||||
return nil, fmt.Errorf("truncated UDP payload length")
|
|
||||||
}
|
|
||||||
payloadLen := int(binary.BigEndian.Uint16(body[offset : offset+2]))
|
|
||||||
offset += 2
|
|
||||||
if payloadLen > MaxUDPPayloadSize {
|
|
||||||
return nil, fmt.Errorf("UDP payload length %d exceeds limit %d", payloadLen, MaxUDPPayloadSize)
|
|
||||||
}
|
|
||||||
remaining := len(body) - offset
|
|
||||||
if remaining < payloadLen {
|
|
||||||
return nil, fmt.Errorf("truncated UDP payload: have %d want %d", remaining, payloadLen)
|
|
||||||
}
|
|
||||||
if remaining > payloadLen {
|
|
||||||
return nil, fmt.Errorf("trailing UDP packet bytes: %d", remaining-payloadLen)
|
|
||||||
}
|
|
||||||
packet.Content = append([]byte(nil), body[offset:offset+payloadLen]...)
|
|
||||||
return packet, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func validateBinaryUDPAddr(addr *net.UDPAddr) (binaryUDPAddr, error) {
|
|
||||||
if addr.Port < 0 || addr.Port > 65535 {
|
|
||||||
return binaryUDPAddr{}, fmt.Errorf("port out of range: %d", addr.Port)
|
|
||||||
}
|
|
||||||
if ip := addr.IP.To4(); ip != nil {
|
|
||||||
if addr.Zone != "" {
|
|
||||||
return binaryUDPAddr{}, fmt.Errorf("IPv4 zone is forbidden")
|
|
||||||
}
|
|
||||||
return binaryUDPAddr{family: 4, ip: ip, port: uint16(addr.Port)}, nil
|
|
||||||
}
|
|
||||||
ip := addr.IP.To16()
|
|
||||||
if ip == nil {
|
|
||||||
return binaryUDPAddr{}, fmt.Errorf("invalid IP")
|
|
||||||
}
|
|
||||||
if len(addr.Zone) > 255 {
|
|
||||||
return binaryUDPAddr{}, fmt.Errorf("zone exceeds 255 bytes")
|
|
||||||
}
|
|
||||||
if !utf8.ValidString(addr.Zone) {
|
|
||||||
return binaryUDPAddr{}, fmt.Errorf("zone is not valid UTF-8")
|
|
||||||
}
|
|
||||||
return binaryUDPAddr{family: 6, ip: ip, port: uint16(addr.Port), zone: addr.Zone}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func binaryUDPAddrLen(addr binaryUDPAddr) int {
|
|
||||||
return 1 + len(addr.ip) + 2 + 1 + len(addr.zone)
|
|
||||||
}
|
|
||||||
|
|
||||||
func putBinaryUDPAddr(body []byte, offset int, addr binaryUDPAddr) int {
|
|
||||||
body[offset] = addr.family
|
|
||||||
offset++
|
|
||||||
copy(body[offset:], addr.ip)
|
|
||||||
offset += len(addr.ip)
|
|
||||||
binary.BigEndian.PutUint16(body[offset:offset+2], addr.port)
|
|
||||||
offset += 2
|
|
||||||
body[offset] = byte(len(addr.zone))
|
|
||||||
offset++
|
|
||||||
copy(body[offset:], addr.zone)
|
|
||||||
return offset + len(addr.zone)
|
|
||||||
}
|
|
||||||
|
|
||||||
func readBinaryUDPAddr(body []byte, offset int) (*net.UDPAddr, int, error) {
|
|
||||||
if offset >= len(body) {
|
|
||||||
return nil, offset, fmt.Errorf("truncated address family")
|
|
||||||
}
|
|
||||||
family := body[offset]
|
|
||||||
offset++
|
|
||||||
var ipLen int
|
|
||||||
switch family {
|
|
||||||
case 4:
|
|
||||||
ipLen = net.IPv4len
|
|
||||||
case 6:
|
|
||||||
ipLen = net.IPv6len
|
|
||||||
default:
|
|
||||||
return nil, offset, fmt.Errorf("unknown address family %d", family)
|
|
||||||
}
|
|
||||||
if len(body)-offset < ipLen+3 {
|
|
||||||
return nil, offset, fmt.Errorf("truncated address")
|
|
||||||
}
|
|
||||||
ip := append(net.IP(nil), body[offset:offset+ipLen]...)
|
|
||||||
offset += ipLen
|
|
||||||
port := binary.BigEndian.Uint16(body[offset : offset+2])
|
|
||||||
offset += 2
|
|
||||||
zoneLen := int(body[offset])
|
|
||||||
offset++
|
|
||||||
if len(body)-offset < zoneLen {
|
|
||||||
return nil, offset, fmt.Errorf("truncated zone")
|
|
||||||
}
|
|
||||||
zoneBytes := body[offset : offset+zoneLen]
|
|
||||||
if family == 4 && zoneLen != 0 {
|
|
||||||
return nil, offset, fmt.Errorf("IPv4 zone is forbidden")
|
|
||||||
}
|
|
||||||
if !utf8.Valid(zoneBytes) {
|
|
||||||
return nil, offset, fmt.Errorf("zone is not valid UTF-8")
|
|
||||||
}
|
|
||||||
offset += zoneLen
|
|
||||||
return &net.UDPAddr{IP: ip, Port: int(port), Zone: string(zoneBytes)}, offset, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2BinaryUDPPacketReadWriter struct {
|
|
||||||
conn *wire.Conn
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewV2BinaryUDPPacketReadWriter(rw io.ReadWriter) *V2BinaryUDPPacketReadWriter {
|
|
||||||
return &V2BinaryUDPPacketReadWriter{conn: wire.NewConn(rw)}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (rw *V2BinaryUDPPacketReadWriter) ReadMsg() (Message, error) {
|
|
||||||
frame, err := rw.conn.ReadFrame()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if isV2MessageType(frame, V2TypeUDPPacketBinary) {
|
|
||||||
return decodeV2BinaryUDPPacketFrame(frame)
|
|
||||||
}
|
|
||||||
if isV2MessageType(frame, V2TypeUDPPacket) {
|
|
||||||
return nil, fmt.Errorf("received JSON UDP packet after binary codec negotiation")
|
|
||||||
}
|
|
||||||
return DecodeV2MessageFrame(frame)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (rw *V2BinaryUDPPacketReadWriter) ReadMsgInto(out Message) error {
|
|
||||||
frame, err := rw.conn.ReadFrame()
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if packetOut, ok := out.(*UDPPacket); ok {
|
|
||||||
if !isV2MessageType(frame, V2TypeUDPPacketBinary) {
|
|
||||||
return unexpectedV2UDPPacketType(frame)
|
|
||||||
}
|
|
||||||
packet, err := decodeV2BinaryUDPPacketFrame(frame)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
*packetOut = *packet
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return DecodeV2MessageFrameInto(frame, out)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (rw *V2BinaryUDPPacketReadWriter) WriteMsg(message Message) error {
|
|
||||||
var packet *UDPPacket
|
|
||||||
switch typed := message.(type) {
|
|
||||||
case *UDPPacket:
|
|
||||||
packet = typed
|
|
||||||
case UDPPacket:
|
|
||||||
packet = &typed
|
|
||||||
default:
|
|
||||||
frame, err := EncodeV2MessageFrame(message)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return rw.conn.WriteFrame(frame)
|
|
||||||
}
|
|
||||||
body, err := EncodeUDPPacketBinary(packet)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
payload := make([]byte, 2+len(body))
|
|
||||||
binary.BigEndian.PutUint16(payload[:2], V2TypeUDPPacketBinary)
|
|
||||||
copy(payload[2:], body)
|
|
||||||
return rw.conn.WriteFrame(&wire.Frame{Type: wire.FrameTypeMessage, Payload: payload})
|
|
||||||
}
|
|
||||||
|
|
||||||
func decodeV2BinaryUDPPacketFrame(frame *wire.Frame) (*UDPPacket, error) {
|
|
||||||
if frame.Type != wire.FrameTypeMessage {
|
|
||||||
return nil, fmt.Errorf("unexpected frame type %d, want %d", frame.Type, wire.FrameTypeMessage)
|
|
||||||
}
|
|
||||||
if len(frame.Payload) < 2 {
|
|
||||||
return nil, fmt.Errorf("message frame payload too short")
|
|
||||||
}
|
|
||||||
if binary.BigEndian.Uint16(frame.Payload[:2]) != V2TypeUDPPacketBinary {
|
|
||||||
return nil, unexpectedV2UDPPacketType(frame)
|
|
||||||
}
|
|
||||||
return DecodeUDPPacketBinary(frame.Payload[2:])
|
|
||||||
}
|
|
||||||
|
|
||||||
func isV2MessageType(frame *wire.Frame, typeID uint16) bool {
|
|
||||||
return frame.Type == wire.FrameTypeMessage && len(frame.Payload) >= 2 && binary.BigEndian.Uint16(frame.Payload[:2]) == typeID
|
|
||||||
}
|
|
||||||
|
|
||||||
func unexpectedV2UDPPacketType(frame *wire.Frame) error {
|
|
||||||
if frame.Type != wire.FrameTypeMessage {
|
|
||||||
return fmt.Errorf("unexpected frame type %d, want %d", frame.Type, wire.FrameTypeMessage)
|
|
||||||
}
|
|
||||||
if len(frame.Payload) < 2 {
|
|
||||||
return fmt.Errorf("message frame payload too short")
|
|
||||||
}
|
|
||||||
typeID := binary.BigEndian.Uint16(frame.Payload[:2])
|
|
||||||
if typeID == V2TypeUDPPacket {
|
|
||||||
return fmt.Errorf("received JSON UDP packet after binary codec negotiation")
|
|
||||||
}
|
|
||||||
return fmt.Errorf("unexpected message type %d, want %d", typeID, V2TypeUDPPacketBinary)
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewUDPPacketReadWriter selects the negotiated packet codec without changing
|
|
||||||
// the framing or codecs used by non-UDP messages on the work connection.
|
|
||||||
func NewUDPPacketReadWriter(rw io.ReadWriter, wireProtocol, udpPacketCodec string) (ReadWriter, error) {
|
|
||||||
switch wireProtocol {
|
|
||||||
case "", wire.ProtocolV1:
|
|
||||||
if udpPacketCodec != "" {
|
|
||||||
return nil, fmt.Errorf("UDP packet codec %q requires wire protocol v2", udpPacketCodec)
|
|
||||||
}
|
|
||||||
return NewV1ReadWriter(rw), nil
|
|
||||||
case wire.ProtocolV2:
|
|
||||||
switch udpPacketCodec {
|
|
||||||
case "":
|
|
||||||
return NewV2ReadWriter(rw), nil
|
|
||||||
case wire.UDPPacketCodecBinary:
|
|
||||||
return NewV2BinaryUDPPacketReadWriter(rw), nil
|
|
||||||
default:
|
|
||||||
return nil, fmt.Errorf("unsupported UDP packet codec %q", udpPacketCodec)
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return nil, fmt.Errorf("unsupported wire protocol %q", wireProtocol)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,248 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
|
|
||||||
package msg
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/binary"
|
|
||||||
"net"
|
|
||||||
"strconv"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/proto/wire"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestUDPPacketBinaryRoundTrip(t *testing.T) {
|
|
||||||
payload := bytes.Repeat([]byte{0xa5}, 1472)
|
|
||||||
in := &UDPPacket{
|
|
||||||
Content: payload,
|
|
||||||
LocalAddr: &net.UDPAddr{
|
|
||||||
IP: net.ParseIP("2001:db8::1"),
|
|
||||||
Port: 1234,
|
|
||||||
Zone: "en0",
|
|
||||||
},
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321},
|
|
||||||
}
|
|
||||||
body, err := EncodeUDPPacketBinary(in)
|
|
||||||
require.NoError(t, err)
|
|
||||||
out, err := DecodeUDPPacketBinary(body)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, in.Content, out.Content)
|
|
||||||
require.Equal(t, in.LocalAddr.String(), out.LocalAddr.String())
|
|
||||||
require.Equal(t, in.RemoteAddr.String(), out.RemoteAddr.String())
|
|
||||||
body[len(body)-1] ^= 0xff
|
|
||||||
body[25] ^= 0xff
|
|
||||||
require.Equal(t, byte(0xa5), out.Content[len(out.Content)-1], "decoded payload must own frame bytes")
|
|
||||||
require.Equal(t, byte(203), out.RemoteAddr.IP.To4()[0], "decoded address must own frame bytes")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPPacketBinarySizesAndOptionalLocalAddress(t *testing.T) {
|
|
||||||
for _, size := range []int{0, 32, 128, 512, 1200, 1472, 4096, 49107, 65507} {
|
|
||||||
t.Run(strconv.Itoa(size), func(t *testing.T) {
|
|
||||||
in := &UDPPacket{
|
|
||||||
Content: bytes.Repeat([]byte{byte(size)}, size),
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321},
|
|
||||||
}
|
|
||||||
body, err := EncodeUDPPacketBinary(in)
|
|
||||||
require.NoError(t, err)
|
|
||||||
out, err := DecodeUDPPacketBinary(body)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, len(in.Content), len(out.Content))
|
|
||||||
if size == 0 {
|
|
||||||
require.Empty(t, out.Content)
|
|
||||||
} else {
|
|
||||||
require.Equal(t, in.Content, out.Content)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPPacketBinaryMalformed(t *testing.T) {
|
|
||||||
valid, err := EncodeUDPPacketBinary(&UDPPacket{
|
|
||||||
Content: []byte("payload"),
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321},
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
tests := [][]byte{
|
|
||||||
{0x80, 0, 0},
|
|
||||||
{0x02, 4, 1, 2},
|
|
||||||
{0x02, 4, 1, 2, 3, 4, 0xd4},
|
|
||||||
append(append([]byte(nil), valid...), 0),
|
|
||||||
}
|
|
||||||
for _, malformed := range tests {
|
|
||||||
_, err := DecodeUDPPacketBinary(malformed)
|
|
||||||
require.Error(t, err)
|
|
||||||
}
|
|
||||||
_, err = DecodeUDPPacketBinary([]byte{0, 0, 0})
|
|
||||||
require.ErrorContains(t, err, "missing remote address")
|
|
||||||
payloadLengthOffset := len(valid) - len("payload") - 2
|
|
||||||
invalidPayloadLength := append([]byte(nil), valid...)
|
|
||||||
binary.BigEndian.PutUint16(invalidPayloadLength[payloadLengthOffset:payloadLengthOffset+2], 0xffff)
|
|
||||||
_, err = DecodeUDPPacketBinary(invalidPayloadLength)
|
|
||||||
require.ErrorContains(t, err, "payload length")
|
|
||||||
truncatedPayload := append([]byte(nil), valid[:payloadLengthOffset+2]...)
|
|
||||||
binary.BigEndian.PutUint16(truncatedPayload[payloadLengthOffset:payloadLengthOffset+2], 1)
|
|
||||||
_, err = DecodeUDPPacketBinary(truncatedPayload)
|
|
||||||
require.ErrorContains(t, err, "truncated UDP payload")
|
|
||||||
_, err = DecodeUDPPacketBinary(make([]byte, wire.DefaultMaxFramePayloadSize))
|
|
||||||
require.ErrorContains(t, err, "frame payload length")
|
|
||||||
|
|
||||||
badIPv4Zone := []byte{2, 4, 203, 0, 113, 9, 0xd4, 0x31, 1, 'z', 0, 0}
|
|
||||||
_, err = DecodeUDPPacketBinary(badIPv4Zone)
|
|
||||||
require.ErrorContains(t, err, "IPv4 zone")
|
|
||||||
badFamily := []byte{2, 9, 0, 0}
|
|
||||||
_, err = DecodeUDPPacketBinary(badFamily)
|
|
||||||
require.ErrorContains(t, err, "unknown address family")
|
|
||||||
badUTF8 := make([]byte, 0, 24)
|
|
||||||
badUTF8 = append(badUTF8, 2, 6)
|
|
||||||
badUTF8 = append(badUTF8, make([]byte, 16)...)
|
|
||||||
badUTF8 = append(badUTF8, 0, 1, 1, 0xff, 0, 0)
|
|
||||||
_, err = DecodeUDPPacketBinary(badUTF8)
|
|
||||||
require.ErrorContains(t, err, "UTF-8")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPPacketBinaryEncodeRejectsInvalidPackets(t *testing.T) {
|
|
||||||
_, err := EncodeUDPPacketBinary(&UDPPacket{})
|
|
||||||
require.ErrorContains(t, err, "missing remote address")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{
|
|
||||||
LocalAddr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 1234},
|
|
||||||
})
|
|
||||||
require.ErrorContains(t, err, "missing remote address")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{
|
|
||||||
Content: make([]byte, MaxUDPPayloadSize+1),
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321},
|
|
||||||
})
|
|
||||||
require.ErrorContains(t, err, "exceeds limit")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 1, Zone: "bad"}})
|
|
||||||
require.ErrorContains(t, err, "IPv4 zone")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 1, Zone: string(bytes.Repeat([]byte{'z'}, 256))}})
|
|
||||||
require.ErrorContains(t, err, "zone exceeds")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 1, Zone: string([]byte{0xff})}})
|
|
||||||
require.ErrorContains(t, err, "UTF-8")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{Port: -1}})
|
|
||||||
require.ErrorContains(t, err, "port out of range")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{Port: 65536}})
|
|
||||||
require.ErrorContains(t, err, "port out of range")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{IP: net.IP{1, 2, 3}}})
|
|
||||||
require.ErrorContains(t, err, "invalid IP")
|
|
||||||
_, err = EncodeUDPPacketBinary(&UDPPacket{
|
|
||||||
Content: make([]byte, MaxUDPPayloadSize),
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Zone: string(bytes.Repeat([]byte{'z'}, 255))},
|
|
||||||
})
|
|
||||||
require.ErrorContains(t, err, "frame payload length")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestV2BinaryUDPPacketReadWriterPreservesOtherMessages(t *testing.T) {
|
|
||||||
var buf bytes.Buffer
|
|
||||||
rw := NewV2BinaryUDPPacketReadWriter(&buf)
|
|
||||||
in := &UDPPacket{Content: []byte("udp"), RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321}}
|
|
||||||
require.NoError(t, rw.WriteMsg(in))
|
|
||||||
require.NoError(t, rw.WriteMsg(&Ping{Timestamp: 7}))
|
|
||||||
frameConn := wire.NewConn(&buf)
|
|
||||||
frame, err := frameConn.ReadFrame()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, V2TypeUDPPacketBinary, binary.BigEndian.Uint16(frame.Payload[:2]))
|
|
||||||
frame, err = frameConn.ReadFrame()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, V2TypePing, binary.BigEndian.Uint16(frame.Payload[:2]))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestV2BinaryUDPPacketReadWriterRoundTripAndCodecInvariant(t *testing.T) {
|
|
||||||
in := &UDPPacket{Content: []byte("udp"), RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321}}
|
|
||||||
var binaryStream bytes.Buffer
|
|
||||||
binaryWriter, err := NewUDPPacketReadWriter(&binaryStream, wire.ProtocolV2, wire.UDPPacketCodecBinary)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, binaryWriter.WriteMsg(in))
|
|
||||||
binaryReader, err := NewUDPPacketReadWriter(&binaryStream, wire.ProtocolV2, wire.UDPPacketCodecBinary)
|
|
||||||
require.NoError(t, err)
|
|
||||||
out, err := binaryReader.ReadMsg()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, in.Content, out.(*UDPPacket).Content)
|
|
||||||
|
|
||||||
for _, read := range []func(ReadWriter) error{
|
|
||||||
func(rw ReadWriter) error {
|
|
||||||
_, err := rw.ReadMsg()
|
|
||||||
return err
|
|
||||||
},
|
|
||||||
func(rw ReadWriter) error {
|
|
||||||
return rw.ReadMsgInto(&UDPPacket{})
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
var jsonStream bytes.Buffer
|
|
||||||
require.NoError(t, NewReadWriter(&jsonStream, wire.ProtocolV2).WriteMsg(in))
|
|
||||||
negotiatedReader, err := NewUDPPacketReadWriter(&jsonStream, wire.ProtocolV2, wire.UDPPacketCodecBinary)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.ErrorContains(t, read(negotiatedReader), "JSON UDP packet after binary codec negotiation")
|
|
||||||
}
|
|
||||||
|
|
||||||
var fallbackStream bytes.Buffer
|
|
||||||
fallbackWriter, err := NewUDPPacketReadWriter(&fallbackStream, wire.ProtocolV2, "")
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, fallbackWriter.WriteMsg(in))
|
|
||||||
frame, err := wire.NewConn(&fallbackStream).ReadFrame()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, V2TypeUDPPacket, binary.BigEndian.Uint16(frame.Payload[:2]))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewUDPPacketReadWriterDefaultProtocolUsesV1(t *testing.T) {
|
|
||||||
var stream bytes.Buffer
|
|
||||||
rw, err := NewUDPPacketReadWriter(&stream, "", "")
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.IsType(t, &V1ReadWriter{}, rw)
|
|
||||||
require.NoError(t, rw.WriteMsg(&UDPPacket{Content: []byte("legacy")}))
|
|
||||||
require.Equal(t, TypeUDPPacket, stream.Bytes()[0])
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewUDPPacketReadWriterRejectsInvalidSelection(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
wireProtocol string
|
|
||||||
udpPacketCodec string
|
|
||||||
errorSubstring string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "binary codec over v1",
|
|
||||||
wireProtocol: wire.ProtocolV1,
|
|
||||||
udpPacketCodec: wire.UDPPacketCodecBinary,
|
|
||||||
errorSubstring: "requires wire protocol v2",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "binary codec over default protocol",
|
|
||||||
udpPacketCodec: wire.UDPPacketCodecBinary,
|
|
||||||
errorSubstring: "requires wire protocol v2",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "unknown v2 codec",
|
|
||||||
wireProtocol: wire.ProtocolV2,
|
|
||||||
udpPacketCodec: "unknown",
|
|
||||||
errorSubstring: "unsupported UDP packet codec",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "unknown wire protocol",
|
|
||||||
wireProtocol: "unknown",
|
|
||||||
errorSubstring: "unsupported wire protocol",
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
rw, err := NewUDPPacketReadWriter(&bytes.Buffer{}, tc.wireProtocol, tc.udpPacketCodec)
|
|
||||||
require.Nil(t, rw)
|
|
||||||
require.ErrorContains(t, err, tc.errorSubstring)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func FuzzDecodeUDPPacketBinary(f *testing.F) {
|
|
||||||
f.Add([]byte{0, 0, 0})
|
|
||||||
f.Add([]byte{2, 4, 203, 0, 113, 9, 0xd4, 0x31, 0, 0, 1})
|
|
||||||
f.Fuzz(func(t *testing.T, body []byte) {
|
|
||||||
_, _ = DecodeUDPPacketBinary(body)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@@ -43,7 +43,6 @@ const (
|
|||||||
V2TypeNatHoleResp uint16 = 16
|
V2TypeNatHoleResp uint16 = 16
|
||||||
V2TypeNatHoleSid uint16 = 17
|
V2TypeNatHoleSid uint16 = 17
|
||||||
V2TypeNatHoleReport uint16 = 18
|
V2TypeNatHoleReport uint16 = 18
|
||||||
V2TypeUDPPacketBinary uint16 = 19
|
|
||||||
)
|
)
|
||||||
|
|
||||||
var v2MsgTypeMap = map[uint16]any{
|
var v2MsgTypeMap = map[uint16]any{
|
||||||
|
|||||||
@@ -84,9 +84,6 @@ func TestV2MessageTypeIDsAreStable(t *testing.T) {
|
|||||||
require.Equal(t, uint16(16), V2TypeNatHoleResp)
|
require.Equal(t, uint16(16), V2TypeNatHoleResp)
|
||||||
require.Equal(t, uint16(17), V2TypeNatHoleSid)
|
require.Equal(t, uint16(17), V2TypeNatHoleSid)
|
||||||
require.Equal(t, uint16(18), V2TypeNatHoleReport)
|
require.Equal(t, uint16(18), V2TypeNatHoleReport)
|
||||||
require.Equal(t, uint16(19), V2TypeUDPPacketBinary)
|
|
||||||
_, registered := v2MsgTypeMap[V2TypeUDPPacketBinary]
|
|
||||||
require.False(t, registered, "binary UDP has a dedicated codec and must not alter generic type registry")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestV2MessageFrameEncoding(t *testing.T) {
|
func TestV2MessageFrameEncoding(t *testing.T) {
|
||||||
|
|||||||
+70
-30
@@ -15,24 +15,31 @@
|
|||||||
package nathole
|
package nathole
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/fatedier/golib/net/stun"
|
"github.com/pion/stun/v3"
|
||||||
)
|
)
|
||||||
|
|
||||||
var responseTimeout = 3 * time.Second
|
var responseTimeout = 3 * time.Second
|
||||||
|
|
||||||
|
type Message struct {
|
||||||
|
Body []byte
|
||||||
|
Addr string
|
||||||
|
}
|
||||||
|
|
||||||
// If the localAddr is empty, it will listen on a random port.
|
// If the localAddr is empty, it will listen on a random port.
|
||||||
func Discover(stunServers []string, localAddr string) ([]string, net.Addr, error) {
|
func Discover(stunServers []string, localAddr string) ([]string, net.Addr, error) {
|
||||||
|
// create a discoverConn and get response from messageChan
|
||||||
discoverConn, err := listen(localAddr)
|
discoverConn, err := listen(localAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, nil, err
|
return nil, nil, err
|
||||||
}
|
}
|
||||||
defer discoverConn.Close()
|
defer discoverConn.Close()
|
||||||
|
|
||||||
|
go discoverConn.readLoop()
|
||||||
|
|
||||||
addresses := make([]string, 0, len(stunServers))
|
addresses := make([]string, 0, len(stunServers))
|
||||||
for _, addr := range stunServers {
|
for _, addr := range stunServers {
|
||||||
// get external address from stun server
|
// get external address from stun server
|
||||||
@@ -51,9 +58,10 @@ type stunResponse struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type discoverConn struct {
|
type discoverConn struct {
|
||||||
conn *net.UDPConn
|
conn *net.UDPConn
|
||||||
client *stun.Client
|
|
||||||
localAddr net.Addr
|
localAddr net.Addr
|
||||||
|
messageChan chan *Message
|
||||||
}
|
}
|
||||||
|
|
||||||
func listen(localAddr string) (*discoverConn, error) {
|
func listen(localAddr string) (*discoverConn, error) {
|
||||||
@@ -69,50 +77,82 @@ func listen(localAddr string) (*discoverConn, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
client, err := stun.NewClient(conn)
|
|
||||||
if err != nil {
|
|
||||||
_ = conn.Close()
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return &discoverConn{
|
return &discoverConn{
|
||||||
conn: conn,
|
conn: conn,
|
||||||
client: client,
|
localAddr: conn.LocalAddr(),
|
||||||
localAddr: conn.LocalAddr(),
|
messageChan: make(chan *Message, 10),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *discoverConn) Close() error {
|
func (c *discoverConn) Close() error {
|
||||||
|
if c.messageChan != nil {
|
||||||
|
close(c.messageChan)
|
||||||
|
c.messageChan = nil
|
||||||
|
}
|
||||||
return c.conn.Close()
|
return c.conn.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *discoverConn) readLoop() {
|
||||||
|
for {
|
||||||
|
buf := make([]byte, 1024)
|
||||||
|
n, addr, err := c.conn.ReadFromUDP(buf)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
buf = buf[:n]
|
||||||
|
|
||||||
|
c.messageChan <- &Message{
|
||||||
|
Body: buf,
|
||||||
|
Addr: addr.String(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (c *discoverConn) doSTUNRequest(addr string) (*stunResponse, error) {
|
func (c *discoverConn) doSTUNRequest(addr string) (*stunResponse, error) {
|
||||||
serverAddr, err := net.ResolveUDPAddr("udp4", addr)
|
serverAddr, err := net.ResolveUDPAddr("udp4", addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
transaction, err := stun.NewBindingTransaction(serverAddr)
|
request, err := stun.Build(stun.TransactionID, stun.BindingRequest)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := c.conn.SetReadDeadline(time.Now().Add(responseTimeout)); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
response, err := c.client.Do(transaction)
|
|
||||||
if err != nil {
|
|
||||||
var netErr net.Error
|
|
||||||
if errors.As(err, &netErr) && netErr.Timeout() {
|
|
||||||
return nil, fmt.Errorf("wait response from stun server timeout")
|
|
||||||
}
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
resp := &stunResponse{}
|
if err = request.NewTransactionID(); err != nil {
|
||||||
if response.MappedAddr != nil {
|
return nil, err
|
||||||
resp.externalAddr = response.MappedAddr.String()
|
|
||||||
}
|
}
|
||||||
if response.OtherAddr != nil {
|
if _, err := c.conn.WriteTo(request.Raw, serverAddr); err != nil {
|
||||||
resp.otherAddr = response.OtherAddr.String()
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var m stun.Message
|
||||||
|
select {
|
||||||
|
case msg := <-c.messageChan:
|
||||||
|
m.Raw = msg.Body
|
||||||
|
if err := m.Decode(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
case <-time.After(responseTimeout):
|
||||||
|
return nil, fmt.Errorf("wait response from stun server timeout")
|
||||||
|
}
|
||||||
|
xorAddrGetter := &stun.XORMappedAddress{}
|
||||||
|
mappedAddrGetter := &stun.MappedAddress{}
|
||||||
|
changedAddrGetter := ChangedAddress{}
|
||||||
|
otherAddrGetter := &stun.OtherAddress{}
|
||||||
|
|
||||||
|
resp := &stunResponse{}
|
||||||
|
if err := mappedAddrGetter.GetFrom(&m); err == nil {
|
||||||
|
resp.externalAddr = mappedAddrGetter.String()
|
||||||
|
}
|
||||||
|
if err := xorAddrGetter.GetFrom(&m); err == nil {
|
||||||
|
resp.externalAddr = xorAddrGetter.String()
|
||||||
|
}
|
||||||
|
if err := changedAddrGetter.GetFrom(&m); err == nil {
|
||||||
|
resp.otherAddr = changedAddrGetter.String()
|
||||||
|
}
|
||||||
|
if err := otherAddrGetter.GetFrom(&m); err == nil {
|
||||||
|
resp.otherAddr = otherAddrGetter.String()
|
||||||
}
|
}
|
||||||
return resp, nil
|
return resp, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,382 +0,0 @@
|
|||||||
package nathole
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/fatedier/golib/net/stun"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
testBindingRequest = 0x0001
|
|
||||||
testBindingSuccess = 0x0101
|
|
||||||
testBindingError = 0x0111
|
|
||||||
testMagicCookie = 0x2112a442
|
|
||||||
testAttrMapped = 0x0001
|
|
||||||
testAttrChanged = 0x0005
|
|
||||||
testAttrErrorCode = 0x0009
|
|
||||||
testAttrXORMapped = 0x0020
|
|
||||||
testAttrOther = 0x802c
|
|
||||||
testSTUNHeaderSize = 20
|
|
||||||
testSTUNServerLimit = time.Second
|
|
||||||
)
|
|
||||||
|
|
||||||
type testSTUNAttribute struct {
|
|
||||||
typ uint16
|
|
||||||
value []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
type testSTUNExchange struct {
|
|
||||||
source *net.UDPAddr
|
|
||||||
err error
|
|
||||||
}
|
|
||||||
|
|
||||||
func listenTestUDP4(t *testing.T) *net.UDPConn {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = conn.Close() })
|
|
||||||
return conn
|
|
||||||
}
|
|
||||||
|
|
||||||
func serveOneSTUNRequest(
|
|
||||||
server *net.UDPConn,
|
|
||||||
buildResponse func([]byte, *net.UDPAddr) ([]byte, error),
|
|
||||||
) <-chan testSTUNExchange {
|
|
||||||
done := make(chan testSTUNExchange, 1)
|
|
||||||
go func() {
|
|
||||||
if err := server.SetDeadline(time.Now().Add(testSTUNServerLimit)); err != nil {
|
|
||||||
done <- testSTUNExchange{err: err}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
buffer := make([]byte, 1024)
|
|
||||||
n, source, err := server.ReadFromUDP(buffer)
|
|
||||||
if err == nil && buildResponse != nil {
|
|
||||||
var response []byte
|
|
||||||
response, err = buildResponse(buffer[:n], source)
|
|
||||||
if err == nil && response != nil {
|
|
||||||
_, err = server.WriteToUDP(response, source)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
done <- testSTUNExchange{source: source, err: err}
|
|
||||||
}()
|
|
||||||
return done
|
|
||||||
}
|
|
||||||
|
|
||||||
func waitSTUNExchange(t *testing.T, done <-chan testSTUNExchange) *net.UDPAddr {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case exchange := <-done:
|
|
||||||
require.NoError(t, exchange.err)
|
|
||||||
return exchange.source
|
|
||||||
case <-time.After(testSTUNServerLimit):
|
|
||||||
t.Fatal("timed out waiting for local STUN server")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func makeTestSTUNResponse(request []byte, typ uint16, attributes ...testSTUNAttribute) ([]byte, error) {
|
|
||||||
if len(request) != testSTUNHeaderSize || binary.BigEndian.Uint16(request[0:2]) != testBindingRequest ||
|
|
||||||
binary.BigEndian.Uint32(request[4:8]) != testMagicCookie {
|
|
||||||
return nil, fmt.Errorf("invalid Binding request")
|
|
||||||
}
|
|
||||||
|
|
||||||
length := 0
|
|
||||||
for _, attribute := range attributes {
|
|
||||||
length += 4 + (len(attribute.value)+3)&^3
|
|
||||||
}
|
|
||||||
response := make([]byte, testSTUNHeaderSize, testSTUNHeaderSize+length)
|
|
||||||
binary.BigEndian.PutUint16(response[0:2], typ)
|
|
||||||
binary.BigEndian.PutUint16(response[2:4], uint16(length))
|
|
||||||
binary.BigEndian.PutUint32(response[4:8], testMagicCookie)
|
|
||||||
copy(response[8:20], request[8:20])
|
|
||||||
|
|
||||||
for _, attribute := range attributes {
|
|
||||||
start := len(response)
|
|
||||||
paddedLength := (len(attribute.value) + 3) &^ 3
|
|
||||||
response = append(response, make([]byte, 4+paddedLength)...)
|
|
||||||
binary.BigEndian.PutUint16(response[start:start+2], attribute.typ)
|
|
||||||
binary.BigEndian.PutUint16(response[start+2:start+4], uint16(len(attribute.value)))
|
|
||||||
copy(response[start+4:], attribute.value)
|
|
||||||
}
|
|
||||||
return response, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func testIPv4AddressValue(ip net.IP, port int, xor bool) []byte {
|
|
||||||
value := make([]byte, 8)
|
|
||||||
value[1] = 0x01
|
|
||||||
binary.BigEndian.PutUint16(value[2:4], uint16(port))
|
|
||||||
copy(value[4:], ip.To4())
|
|
||||||
if xor {
|
|
||||||
binary.BigEndian.PutUint16(value[2:4], binary.BigEndian.Uint16(value[2:4])^uint16(testMagicCookie>>16))
|
|
||||||
for i := range 4 {
|
|
||||||
value[4+i] ^= byte(uint32(testMagicCookie) >> uint(24-8*i))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return value
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDiscoverReusesLocalPortAndPreservesNATClassification(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
secondMapped string
|
|
||||||
secondMappedPort int
|
|
||||||
wantNATType string
|
|
||||||
wantBehavior string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "same mapped address",
|
|
||||||
secondMapped: "198.51.100.10:40000",
|
|
||||||
secondMappedPort: 40000,
|
|
||||||
wantNATType: EasyNAT,
|
|
||||||
wantBehavior: BehaviorNoChange,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "different mapped port",
|
|
||||||
secondMapped: "198.51.100.10:40001",
|
|
||||||
secondMappedPort: 40001,
|
|
||||||
wantNATType: HardNAT,
|
|
||||||
wantBehavior: BehaviorPortChanged,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
primary := listenTestUDP4(t)
|
|
||||||
alternate := listenTestUDP4(t)
|
|
||||||
alternateAddr := alternate.LocalAddr().(*net.UDPAddr)
|
|
||||||
|
|
||||||
primaryDone := serveOneSTUNRequest(primary, func(request []byte, _ *net.UDPAddr) ([]byte, error) {
|
|
||||||
return makeTestSTUNResponse(request, testBindingSuccess,
|
|
||||||
testSTUNAttribute{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.10"), 40000, true)},
|
|
||||||
testSTUNAttribute{typ: testAttrOther, value: testIPv4AddressValue(alternateAddr.IP, alternateAddr.Port, false)},
|
|
||||||
)
|
|
||||||
})
|
|
||||||
alternateDone := serveOneSTUNRequest(alternate, func(request []byte, _ *net.UDPAddr) ([]byte, error) {
|
|
||||||
return makeTestSTUNResponse(request, testBindingSuccess,
|
|
||||||
testSTUNAttribute{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.10"), tt.secondMappedPort, true)},
|
|
||||||
)
|
|
||||||
})
|
|
||||||
|
|
||||||
addresses, localAddr, err := Discover([]string{primary.LocalAddr().String()}, "")
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, []string{"198.51.100.10:40000", tt.secondMapped}, addresses)
|
|
||||||
|
|
||||||
primarySource := waitSTUNExchange(t, primaryDone)
|
|
||||||
alternateSource := waitSTUNExchange(t, alternateDone)
|
|
||||||
require.Equal(t, primarySource.Port, alternateSource.Port)
|
|
||||||
require.Equal(t, localAddr.(*net.UDPAddr).Port, primarySource.Port)
|
|
||||||
|
|
||||||
feature, err := ClassifyNATFeature(addresses, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, tt.wantNATType, feature.NatType)
|
|
||||||
require.Equal(t, tt.wantBehavior, feature.Behavior)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDoSTUNRequestMapsLegacyAndModernAddresses(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
attributes []testSTUNAttribute
|
|
||||||
wantExternal string
|
|
||||||
wantOther string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "legacy",
|
|
||||||
attributes: []testSTUNAttribute{
|
|
||||||
{typ: testAttrMapped, value: testIPv4AddressValue(net.ParseIP("192.0.2.1"), 1000, false)},
|
|
||||||
{typ: testAttrChanged, value: testIPv4AddressValue(net.ParseIP("192.0.2.2"), 2000, false)},
|
|
||||||
},
|
|
||||||
wantExternal: "192.0.2.1:1000",
|
|
||||||
wantOther: "192.0.2.2:2000",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "modern takes precedence",
|
|
||||||
attributes: []testSTUNAttribute{
|
|
||||||
{typ: testAttrMapped, value: testIPv4AddressValue(net.ParseIP("192.0.2.1"), 1000, false)},
|
|
||||||
{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.1"), 3000, true)},
|
|
||||||
{typ: testAttrChanged, value: testIPv4AddressValue(net.ParseIP("192.0.2.2"), 2000, false)},
|
|
||||||
{typ: testAttrOther, value: testIPv4AddressValue(net.ParseIP("198.51.100.2"), 4000, false)},
|
|
||||||
},
|
|
||||||
wantExternal: "198.51.100.1:3000",
|
|
||||||
wantOther: "198.51.100.2:4000",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
server := listenTestUDP4(t)
|
|
||||||
done := serveOneSTUNRequest(server, func(request []byte, _ *net.UDPAddr) ([]byte, error) {
|
|
||||||
return makeTestSTUNResponse(request, testBindingSuccess, tt.attributes...)
|
|
||||||
})
|
|
||||||
conn, err := listen("")
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = conn.Close() })
|
|
||||||
|
|
||||||
response, err := conn.doSTUNRequest(server.LocalAddr().String())
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, tt.wantExternal, response.externalAddr)
|
|
||||||
require.Equal(t, tt.wantOther, response.otherAddr)
|
|
||||||
waitSTUNExchange(t, done)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSTUNResponseErrorsAndMissingAddresses(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
buildResponse func([]byte, *net.UDPAddr) ([]byte, error)
|
|
||||||
request func(*discoverConn, string) error
|
|
||||||
checkError func(*testing.T, error)
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "correlated malformed response",
|
|
||||||
buildResponse: func(request []byte, _ *net.UDPAddr) ([]byte, error) {
|
|
||||||
response, err := makeTestSTUNResponse(request, testBindingSuccess)
|
|
||||||
if err == nil {
|
|
||||||
binary.BigEndian.PutUint16(response[2:4], 4)
|
|
||||||
}
|
|
||||||
return response, err
|
|
||||||
},
|
|
||||||
request: func(conn *discoverConn, server string) error {
|
|
||||||
_, err := conn.doSTUNRequest(server)
|
|
||||||
return err
|
|
||||||
},
|
|
||||||
checkError: func(t *testing.T, err error) {
|
|
||||||
require.ErrorIs(t, err, stun.ErrMalformedResponse)
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "Binding error response",
|
|
||||||
buildResponse: func(request []byte, _ *net.UDPAddr) ([]byte, error) {
|
|
||||||
return makeTestSTUNResponse(request, testBindingError, testSTUNAttribute{
|
|
||||||
typ: testAttrErrorCode,
|
|
||||||
value: []byte{0, 0, 4, 20, 'U', 'n', 'k', 'n', 'o', 'w', 'n'},
|
|
||||||
})
|
|
||||||
},
|
|
||||||
request: func(conn *discoverConn, server string) error {
|
|
||||||
_, err := conn.doSTUNRequest(server)
|
|
||||||
return err
|
|
||||||
},
|
|
||||||
checkError: func(t *testing.T, err error) {
|
|
||||||
var responseErr *stun.ResponseError
|
|
||||||
require.ErrorAs(t, err, &responseErr)
|
|
||||||
require.Equal(t, 420, responseErr.Code)
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "missing mapped address",
|
|
||||||
buildResponse: func(request []byte, _ *net.UDPAddr) ([]byte, error) {
|
|
||||||
return makeTestSTUNResponse(request, testBindingSuccess,
|
|
||||||
testSTUNAttribute{typ: testAttrOther, value: testIPv4AddressValue(net.ParseIP("192.0.2.2"), 2000, false)},
|
|
||||||
)
|
|
||||||
},
|
|
||||||
request: func(conn *discoverConn, server string) error {
|
|
||||||
_, err := conn.discoverFromStunServer(server)
|
|
||||||
return err
|
|
||||||
},
|
|
||||||
checkError: func(t *testing.T, err error) {
|
|
||||||
require.EqualError(t, err, "no external address found")
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
server := listenTestUDP4(t)
|
|
||||||
done := serveOneSTUNRequest(server, tt.buildResponse)
|
|
||||||
conn, err := listen("")
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = conn.Close() })
|
|
||||||
|
|
||||||
err = tt.request(conn, server.LocalAddr().String())
|
|
||||||
tt.checkError(t, err)
|
|
||||||
waitSTUNExchange(t, done)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Run("missing other address", func(t *testing.T) {
|
|
||||||
server := listenTestUDP4(t)
|
|
||||||
done := serveOneSTUNRequest(server, func(request []byte, _ *net.UDPAddr) ([]byte, error) {
|
|
||||||
return makeTestSTUNResponse(request, testBindingSuccess,
|
|
||||||
testSTUNAttribute{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.1"), 3000, true)},
|
|
||||||
)
|
|
||||||
})
|
|
||||||
|
|
||||||
_, err := Prepare([]string{server.LocalAddr().String()}, PrepareOptions{})
|
|
||||||
require.EqualError(t, err, "discover error: not enough addresses")
|
|
||||||
waitSTUNExchange(t, done)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSTUNTimeoutUsesCallerDeadlineWithoutRetry(t *testing.T) {
|
|
||||||
originalTimeout := responseTimeout
|
|
||||||
responseTimeout = 50 * time.Millisecond
|
|
||||||
t.Cleanup(func() { responseTimeout = originalTimeout })
|
|
||||||
|
|
||||||
server := listenTestUDP4(t)
|
|
||||||
done := serveOneSTUNRequest(server, nil)
|
|
||||||
conn, err := listen("")
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = conn.Close() })
|
|
||||||
|
|
||||||
_, err = conn.doSTUNRequest(server.LocalAddr().String())
|
|
||||||
require.EqualError(t, err, "wait response from stun server timeout")
|
|
||||||
waitSTUNExchange(t, done)
|
|
||||||
|
|
||||||
require.NoError(t, server.SetReadDeadline(time.Now().Add(50*time.Millisecond)))
|
|
||||||
_, _, err = server.ReadFromUDP(make([]byte, 1))
|
|
||||||
var netErr net.Error
|
|
||||||
require.ErrorAs(t, err, &netErr)
|
|
||||||
require.True(t, netErr.Timeout())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSTUNClientLeavesSocketAndDeadlineWithCaller(t *testing.T) {
|
|
||||||
originalTimeout := responseTimeout
|
|
||||||
responseTimeout = 100 * time.Millisecond
|
|
||||||
t.Cleanup(func() { responseTimeout = originalTimeout })
|
|
||||||
|
|
||||||
server := listenTestUDP4(t)
|
|
||||||
unrelated := listenTestUDP4(t)
|
|
||||||
done := serveOneSTUNRequest(server, func(request []byte, source *net.UDPAddr) ([]byte, error) {
|
|
||||||
response, err := makeTestSTUNResponse(request, testBindingSuccess,
|
|
||||||
testSTUNAttribute{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.5"), 5000, true)},
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if _, err := unrelated.WriteToUDP(response, source); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return response, nil
|
|
||||||
})
|
|
||||||
conn, err := listen("")
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = conn.Close() })
|
|
||||||
|
|
||||||
response, err := conn.doSTUNRequest(server.LocalAddr().String())
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, "198.51.100.5:5000", response.externalAddr)
|
|
||||||
waitSTUNExchange(t, done)
|
|
||||||
|
|
||||||
_, _, err = conn.conn.ReadFromUDP(make([]byte, 1))
|
|
||||||
var netErr net.Error
|
|
||||||
require.True(t, errors.As(err, &netErr))
|
|
||||||
require.True(t, netErr.Timeout())
|
|
||||||
|
|
||||||
require.NoError(t, conn.conn.SetDeadline(time.Time{}))
|
|
||||||
require.NoError(t, server.SetReadDeadline(time.Now().Add(testSTUNServerLimit)))
|
|
||||||
_, err = conn.conn.WriteToUDP([]byte{1}, server.LocalAddr().(*net.UDPAddr))
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, source, err := server.ReadFromUDP(make([]byte, 1))
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, conn.localAddr.(*net.UDPAddr).Port, source.Port)
|
|
||||||
}
|
|
||||||
@@ -18,8 +18,10 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
|
"strconv"
|
||||||
|
|
||||||
"github.com/fatedier/golib/crypto"
|
"github.com/fatedier/golib/crypto"
|
||||||
|
"github.com/pion/stun/v3"
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
"github.com/fatedier/frp/pkg/msg"
|
||||||
)
|
)
|
||||||
@@ -46,6 +48,20 @@ func DecodeMessageInto(data, key []byte, m msg.Message) error {
|
|||||||
return msg.ReadMsgInto(bytes.NewReader(buf), m)
|
return msg.ReadMsgInto(bytes.NewReader(buf), m)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type ChangedAddress struct {
|
||||||
|
IP net.IP
|
||||||
|
Port int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ChangedAddress) GetFrom(m *stun.Message) error {
|
||||||
|
a := (*stun.MappedAddress)(s)
|
||||||
|
return a.GetFromAs(m, stun.AttrChangedAddress)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ChangedAddress) String() string {
|
||||||
|
return net.JoinHostPort(s.IP.String(), strconv.Itoa(s.Port))
|
||||||
|
}
|
||||||
|
|
||||||
func ListAllLocalIPs() ([]net.IP, error) {
|
func ListAllLocalIPs() ([]net.IP, error) {
|
||||||
addrs, err := net.InterfaceAddrs()
|
addrs, err := net.InterfaceAddrs()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -72,15 +72,6 @@ func (p *TLS2RawPlugin) Handle(ctx context.Context, connInfo *ConnectionInfo) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if connInfo.ProxyProtocolHeader != nil {
|
|
||||||
if _, err := connInfo.ProxyProtocolHeader.WriteTo(rawConn); err != nil {
|
|
||||||
xl.Warnf("tls2raw write proxy protocol header to local conn error: %v", err)
|
|
||||||
rawConn.Close()
|
|
||||||
tlsConn.Close()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
libio.Join(tlsConn, rawConn)
|
libio.Join(tlsConn, rawConn)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"net"
|
"net"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -34,16 +33,10 @@ func init() {
|
|||||||
Register(v1.VisitorPluginVirtualNet, NewVirtualNetPlugin)
|
Register(v1.VisitorPluginVirtualNet, NewVirtualNetPlugin)
|
||||||
}
|
}
|
||||||
|
|
||||||
type clientRouteController interface {
|
|
||||||
RegisterClientRoute(context.Context, string, []net.IPNet, io.ReadWriteCloser)
|
|
||||||
UnregisterClientRoute(string, io.Writer) bool
|
|
||||||
}
|
|
||||||
|
|
||||||
type VirtualNetPlugin struct {
|
type VirtualNetPlugin struct {
|
||||||
pluginCtx PluginContext
|
pluginCtx PluginContext
|
||||||
|
|
||||||
routeController clientRouteController
|
routes []net.IPNet
|
||||||
routes []net.IPNet
|
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
controllerConn net.Conn
|
controllerConn net.Conn
|
||||||
@@ -55,11 +48,6 @@ type VirtualNetPlugin struct {
|
|||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
|
||||||
virtualNetReconnectBaseDelay = 60 * time.Second
|
|
||||||
virtualNetReconnectMaxDelay = 300 * time.Second
|
|
||||||
)
|
|
||||||
|
|
||||||
func NewVirtualNetPlugin(pluginCtx PluginContext, options v1.VisitorPluginOptions) (Plugin, error) {
|
func NewVirtualNetPlugin(pluginCtx PluginContext, options v1.VisitorPluginOptions) (Plugin, error) {
|
||||||
opts := options.(*v1.VirtualNetVisitorPluginOptions)
|
opts := options.(*v1.VirtualNetVisitorPluginOptions)
|
||||||
|
|
||||||
@@ -67,9 +55,6 @@ func NewVirtualNetPlugin(pluginCtx PluginContext, options v1.VisitorPluginOption
|
|||||||
pluginCtx: pluginCtx,
|
pluginCtx: pluginCtx,
|
||||||
routes: make([]net.IPNet, 0),
|
routes: make([]net.IPNet, 0),
|
||||||
}
|
}
|
||||||
if pluginCtx.VnetController != nil {
|
|
||||||
p.routeController = pluginCtx.VnetController
|
|
||||||
}
|
|
||||||
|
|
||||||
p.ctx, p.cancel = context.WithCancel(pluginCtx.Ctx)
|
p.ctx, p.cancel = context.WithCancel(pluginCtx.Ctx)
|
||||||
|
|
||||||
@@ -100,7 +85,7 @@ func (p *VirtualNetPlugin) Name() string {
|
|||||||
|
|
||||||
func (p *VirtualNetPlugin) Start() {
|
func (p *VirtualNetPlugin) Start() {
|
||||||
xl := xlog.FromContextSafe(p.pluginCtx.Ctx)
|
xl := xlog.FromContextSafe(p.pluginCtx.Ctx)
|
||||||
if p.routeController == nil {
|
if p.pluginCtx.VnetController == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -126,17 +111,16 @@ func (p *VirtualNetPlugin) run() {
|
|||||||
select {
|
select {
|
||||||
case <-p.ctx.Done():
|
case <-p.ctx.Done():
|
||||||
xl.Infof("VirtualNetPlugin run loop for visitor [%s] stopping (context cancelled before pipe creation).", p.pluginCtx.Name)
|
xl.Infof("VirtualNetPlugin run loop for visitor [%s] stopping (context cancelled before pipe creation).", p.pluginCtx.Name)
|
||||||
p.cleanupCurrentControllerConn(xl)
|
p.cleanupControllerConn(xl)
|
||||||
return
|
return
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
|
|
||||||
controllerConn, pluginConn := net.Pipe()
|
controllerConn, pluginConn := net.Pipe()
|
||||||
xl.Infof("attempting to register client route for visitor [%s]", p.pluginCtx.Name)
|
|
||||||
if !p.registerControllerConn(controllerConn, pluginConn) {
|
p.mu.Lock()
|
||||||
xl.Infof("VirtualNetPlugin run loop for visitor [%s] stopping (context cancelled before route registration).", p.pluginCtx.Name)
|
p.controllerConn = controllerConn
|
||||||
return
|
p.mu.Unlock()
|
||||||
}
|
|
||||||
|
|
||||||
// Wrap with CloseNotifyConn which supports both close notification and error recording
|
// Wrap with CloseNotifyConn which supports both close notification and error recording
|
||||||
var closeErr error
|
var closeErr error
|
||||||
@@ -145,6 +129,8 @@ func (p *VirtualNetPlugin) run() {
|
|||||||
close(currentCloseSignal) // Signal the run loop on close.
|
close(currentCloseSignal) // Signal the run loop on close.
|
||||||
})
|
})
|
||||||
|
|
||||||
|
xl.Infof("attempting to register client route for visitor [%s]", p.pluginCtx.Name)
|
||||||
|
p.pluginCtx.VnetController.RegisterClientRoute(p.ctx, p.pluginCtx.Name, p.routes, controllerConn)
|
||||||
xl.Infof("successfully registered client route for visitor [%s]. Starting connection handler with CloseNotifyConn.", p.pluginCtx.Name)
|
xl.Infof("successfully registered client route for visitor [%s]. Starting connection handler with CloseNotifyConn.", p.pluginCtx.Name)
|
||||||
|
|
||||||
// Pass the CloseNotifyConn to the visitor for handling.
|
// Pass the CloseNotifyConn to the visitor for handling.
|
||||||
@@ -155,7 +141,7 @@ func (p *VirtualNetPlugin) run() {
|
|||||||
select {
|
select {
|
||||||
case <-p.ctx.Done():
|
case <-p.ctx.Done():
|
||||||
xl.Infof("VirtualNetPlugin run loop stopping for visitor [%s] (context cancelled while waiting).", p.pluginCtx.Name)
|
xl.Infof("VirtualNetPlugin run loop stopping for visitor [%s] (context cancelled while waiting).", p.pluginCtx.Name)
|
||||||
p.cleanupControllerConn(xl, controllerConn)
|
p.cleanupControllerConn(xl)
|
||||||
return
|
return
|
||||||
case <-currentCloseSignal:
|
case <-currentCloseSignal:
|
||||||
// Determine reconnect delay based on error with exponential backoff
|
// Determine reconnect delay based on error with exponential backoff
|
||||||
@@ -166,7 +152,8 @@ func (p *VirtualNetPlugin) run() {
|
|||||||
p.pluginCtx.Name, p.consecutiveErrors, closeErr)
|
p.pluginCtx.Name, p.consecutiveErrors, closeErr)
|
||||||
|
|
||||||
// Exponential backoff: 60s, 120s, 240s, 300s (capped)
|
// Exponential backoff: 60s, 120s, 240s, 300s (capped)
|
||||||
reconnectDelay = virtualNetReconnectDelay(p.consecutiveErrors)
|
baseDelay := 60 * time.Second
|
||||||
|
reconnectDelay = min(baseDelay*time.Duration(1<<uint(p.consecutiveErrors-1)), 300*time.Second)
|
||||||
} else {
|
} else {
|
||||||
// Reset consecutive errors on successful connection
|
// Reset consecutive errors on successful connection
|
||||||
if p.consecutiveErrors > 0 {
|
if p.consecutiveErrors > 0 {
|
||||||
@@ -180,7 +167,7 @@ func (p *VirtualNetPlugin) run() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// The visitor closed the plugin side. Close the controller side.
|
// The visitor closed the plugin side. Close the controller side.
|
||||||
p.cleanupControllerConn(xl, controllerConn)
|
p.cleanupControllerConn(xl)
|
||||||
|
|
||||||
xl.Infof("waiting %v before attempting reconnection for visitor [%s]...", reconnectDelay, p.pluginCtx.Name)
|
xl.Infof("waiting %v before attempting reconnection for visitor [%s]...", reconnectDelay, p.pluginCtx.Name)
|
||||||
select {
|
select {
|
||||||
@@ -195,66 +182,16 @@ func (p *VirtualNetPlugin) run() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// registerControllerConn publishes and registers controllerConn atomically with
|
// cleanupControllerConn closes the current controllerConn (if it exists) under lock.
|
||||||
// respect to Close. A canceled plugin cannot register a new route.
|
func (p *VirtualNetPlugin) cleanupControllerConn(xl *xlog.Logger) {
|
||||||
func (p *VirtualNetPlugin) registerControllerConn(controllerConn, pluginConn net.Conn) bool {
|
|
||||||
p.mu.Lock()
|
|
||||||
if p.ctx.Err() != nil || p.routeController == nil {
|
|
||||||
p.mu.Unlock()
|
|
||||||
_ = controllerConn.Close()
|
|
||||||
_ = pluginConn.Close()
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
p.controllerConn = controllerConn
|
|
||||||
p.routeController.RegisterClientRoute(p.ctx, p.pluginCtx.Name, p.routes, controllerConn)
|
|
||||||
p.mu.Unlock()
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// virtualNetReconnectDelay returns a bounded reconnect delay without allowing
|
|
||||||
// the exponential shift to overflow for large consecutive error counts.
|
|
||||||
func virtualNetReconnectDelay(consecutiveErrors int) time.Duration {
|
|
||||||
if consecutiveErrors <= 1 {
|
|
||||||
return virtualNetReconnectBaseDelay
|
|
||||||
}
|
|
||||||
if consecutiveErrors >= 4 {
|
|
||||||
return virtualNetReconnectMaxDelay
|
|
||||||
}
|
|
||||||
return virtualNetReconnectBaseDelay * time.Duration(1<<uint(consecutiveErrors-1))
|
|
||||||
}
|
|
||||||
|
|
||||||
// cleanupControllerConn unregisters and closes one connection round without
|
|
||||||
// affecting a replacement route owned by another connection.
|
|
||||||
func (p *VirtualNetPlugin) cleanupControllerConn(xl *xlog.Logger, controllerConn net.Conn) {
|
|
||||||
p.mu.Lock()
|
p.mu.Lock()
|
||||||
defer p.mu.Unlock()
|
defer p.mu.Unlock()
|
||||||
p.cleanupControllerConnLocked(xl, controllerConn)
|
if p.controllerConn != nil {
|
||||||
}
|
xl.Debugf("cleaning up controllerConn for visitor [%s]", p.pluginCtx.Name)
|
||||||
|
p.controllerConn.Close()
|
||||||
func (p *VirtualNetPlugin) cleanupCurrentControllerConn(xl *xlog.Logger) {
|
|
||||||
p.mu.Lock()
|
|
||||||
defer p.mu.Unlock()
|
|
||||||
p.cleanupControllerConnLocked(xl, p.controllerConn)
|
|
||||||
}
|
|
||||||
|
|
||||||
// cleanupControllerConnLocked must be called with p.mu held.
|
|
||||||
func (p *VirtualNetPlugin) cleanupControllerConnLocked(xl *xlog.Logger, controllerConn net.Conn) {
|
|
||||||
if controllerConn == nil {
|
|
||||||
p.closeSignal = nil
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if p.routeController != nil &&
|
|
||||||
p.routeController.UnregisterClientRoute(p.pluginCtx.Name, controllerConn) {
|
|
||||||
xl.Infof("unregistered client route for visitor [%s]", p.pluginCtx.Name)
|
|
||||||
}
|
|
||||||
xl.Debugf("cleaning up controllerConn for visitor [%s]", p.pluginCtx.Name)
|
|
||||||
_ = controllerConn.Close()
|
|
||||||
if p.controllerConn == controllerConn {
|
|
||||||
p.controllerConn = nil
|
p.controllerConn = nil
|
||||||
p.closeSignal = nil
|
|
||||||
}
|
}
|
||||||
|
p.closeSignal = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close initiates the plugin shutdown.
|
// Close initiates the plugin shutdown.
|
||||||
@@ -265,9 +202,15 @@ func (p *VirtualNetPlugin) Close() error {
|
|||||||
// Signal the run loop goroutine to stop.
|
// Signal the run loop goroutine to stop.
|
||||||
p.cancel()
|
p.cancel()
|
||||||
|
|
||||||
// Unregister and close the current connection while holding the same lock
|
// Unregister the route from the controller.
|
||||||
// used to check cancellation and register a route in run.
|
if p.pluginCtx.VnetController != nil {
|
||||||
p.cleanupCurrentControllerConn(xl)
|
p.pluginCtx.VnetController.UnregisterClientRoute(p.pluginCtx.Name)
|
||||||
|
xl.Infof("unregistered client route for visitor [%s]", p.pluginCtx.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Explicitly close the controller side of the pipe.
|
||||||
|
// This ensures the pipe is broken even if the run loop is stuck or the visitor hasn't closed its end.
|
||||||
|
p.cleanupControllerConn(xl)
|
||||||
xl.Infof("finished cleaning up connections during close for visitor [%s]", p.pluginCtx.Name)
|
xl.Infof("finished cleaning up connections during close for visitor [%s]", p.pluginCtx.Name)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -1,245 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
//go:build !frps
|
|
||||||
|
|
||||||
package visitor
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/util/xlog"
|
|
||||||
)
|
|
||||||
|
|
||||||
const testVirtualNetVisitorName = "vnet-visitor"
|
|
||||||
|
|
||||||
type fakeClientRouteController struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
|
|
||||||
routes map[string]io.Writer
|
|
||||||
|
|
||||||
beforeRegister func()
|
|
||||||
registerCalls int
|
|
||||||
unregisterCalls int
|
|
||||||
}
|
|
||||||
|
|
||||||
func newFakeClientRouteController() *fakeClientRouteController {
|
|
||||||
return &fakeClientRouteController{
|
|
||||||
routes: make(map[string]io.Writer),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *fakeClientRouteController) RegisterClientRoute(
|
|
||||||
_ context.Context,
|
|
||||||
name string,
|
|
||||||
_ []net.IPNet,
|
|
||||||
conn io.ReadWriteCloser,
|
|
||||||
) {
|
|
||||||
if c.beforeRegister != nil {
|
|
||||||
c.beforeRegister()
|
|
||||||
}
|
|
||||||
|
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
c.registerCalls++
|
|
||||||
c.routes[name] = conn
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *fakeClientRouteController) UnregisterClientRoute(name string, conn io.Writer) bool {
|
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
c.unregisterCalls++
|
|
||||||
owner, ok := c.routes[name]
|
|
||||||
if !ok || owner != conn {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
delete(c.routes, name)
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *fakeClientRouteController) owner() io.Writer {
|
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
return c.routes[testVirtualNetVisitorName]
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *fakeClientRouteController) callCounts() (register, unregister int) {
|
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
return c.registerCalls, c.unregisterCalls
|
|
||||||
}
|
|
||||||
|
|
||||||
type trackedConn struct {
|
|
||||||
net.Conn
|
|
||||||
closed atomic.Bool
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *trackedConn) Close() error {
|
|
||||||
c.closed.Store(true)
|
|
||||||
return c.Conn.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTrackedPipe(t *testing.T) (*trackedConn, *trackedConn) {
|
|
||||||
t.Helper()
|
|
||||||
left, right := net.Pipe()
|
|
||||||
trackedLeft := &trackedConn{Conn: left}
|
|
||||||
trackedRight := &trackedConn{Conn: right}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = trackedLeft.Close()
|
|
||||||
_ = trackedRight.Close()
|
|
||||||
})
|
|
||||||
return trackedLeft, trackedRight
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTestVirtualNetPlugin(t *testing.T, controller *fakeClientRouteController) *VirtualNetPlugin {
|
|
||||||
t.Helper()
|
|
||||||
pluginCtx := context.Background()
|
|
||||||
ctx, cancel := context.WithCancel(pluginCtx)
|
|
||||||
p := &VirtualNetPlugin{
|
|
||||||
pluginCtx: PluginContext{
|
|
||||||
Name: testVirtualNetVisitorName,
|
|
||||||
Ctx: pluginCtx,
|
|
||||||
},
|
|
||||||
routeController: controller,
|
|
||||||
routes: []net.IPNet{{
|
|
||||||
IP: net.ParseIP("10.1.0.1"),
|
|
||||||
Mask: net.CIDRMask(32, 32),
|
|
||||||
}},
|
|
||||||
ctx: ctx,
|
|
||||||
cancel: cancel,
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = p.Close()
|
|
||||||
})
|
|
||||||
return p
|
|
||||||
}
|
|
||||||
|
|
||||||
func waitResult[T any](t *testing.T, ch <-chan T) T {
|
|
||||||
t.Helper()
|
|
||||||
select {
|
|
||||||
case result := <-ch:
|
|
||||||
return result
|
|
||||||
case <-time.After(5 * time.Second):
|
|
||||||
t.Fatal("timed out waiting for concurrent operation")
|
|
||||||
var zero T
|
|
||||||
return zero
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestVirtualNetReconnectDelay verifies the documented exponential backoff and
|
|
||||||
// ensures large error counts remain capped instead of overflowing to zero.
|
|
||||||
func TestVirtualNetReconnectDelay(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
consecutiveErrors int
|
|
||||||
want time.Duration
|
|
||||||
}{
|
|
||||||
{name: "first error", consecutiveErrors: 1, want: 60 * time.Second},
|
|
||||||
{name: "second error", consecutiveErrors: 2, want: 120 * time.Second},
|
|
||||||
{name: "third error", consecutiveErrors: 3, want: 240 * time.Second},
|
|
||||||
{name: "fourth error", consecutiveErrors: 4, want: 300 * time.Second},
|
|
||||||
{name: "shift width boundary", consecutiveErrors: 64, want: 300 * time.Second},
|
|
||||||
{name: "observed retry storm", consecutiveErrors: 329769, want: 300 * time.Second},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
require.Equal(t, tt.want, virtualNetReconnectDelay(tt.consecutiveErrors))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestVirtualNetPluginCloseBeforeRegisterDoesNotReplaceNewRoute(t *testing.T) {
|
|
||||||
controller := newFakeClientRouteController()
|
|
||||||
oldPlugin := newTestVirtualNetPlugin(t, controller)
|
|
||||||
newPlugin := newTestVirtualNetPlugin(t, controller)
|
|
||||||
oldControllerConn, oldPluginConn := newTrackedPipe(t)
|
|
||||||
newControllerConn, newPluginConn := newTrackedPipe(t)
|
|
||||||
|
|
||||||
allowOldRegister := make(chan struct{})
|
|
||||||
oldRegisterResult := make(chan bool, 1)
|
|
||||||
go func() {
|
|
||||||
<-allowOldRegister
|
|
||||||
oldRegisterResult <- oldPlugin.registerControllerConn(oldControllerConn, oldPluginConn)
|
|
||||||
}()
|
|
||||||
|
|
||||||
require.NoError(t, oldPlugin.Close())
|
|
||||||
require.True(t, newPlugin.registerControllerConn(newControllerConn, newPluginConn))
|
|
||||||
close(allowOldRegister)
|
|
||||||
|
|
||||||
require.False(t, waitResult(t, oldRegisterResult))
|
|
||||||
require.Same(t, newControllerConn, controller.owner())
|
|
||||||
registerCalls, _ := controller.callCounts()
|
|
||||||
require.Equal(t, 1, registerCalls)
|
|
||||||
require.True(t, oldControllerConn.closed.Load())
|
|
||||||
require.True(t, oldPluginConn.closed.Load())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestVirtualNetPluginRegisterBeforeCloseIsCleanedUp(t *testing.T) {
|
|
||||||
controller := newFakeClientRouteController()
|
|
||||||
p := newTestVirtualNetPlugin(t, controller)
|
|
||||||
controllerConn, pluginConn := newTrackedPipe(t)
|
|
||||||
registerEntered := make(chan struct{})
|
|
||||||
var registerEnteredOnce sync.Once
|
|
||||||
controller.beforeRegister = func() {
|
|
||||||
registerEnteredOnce.Do(func() {
|
|
||||||
close(registerEntered)
|
|
||||||
})
|
|
||||||
<-p.ctx.Done()
|
|
||||||
}
|
|
||||||
|
|
||||||
registerResult := make(chan bool, 1)
|
|
||||||
go func() {
|
|
||||||
registerResult <- p.registerControllerConn(controllerConn, pluginConn)
|
|
||||||
}()
|
|
||||||
waitResult(t, registerEntered)
|
|
||||||
|
|
||||||
closeResult := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
closeResult <- p.Close()
|
|
||||||
}()
|
|
||||||
|
|
||||||
require.NoError(t, waitResult(t, closeResult))
|
|
||||||
require.True(t, waitResult(t, registerResult))
|
|
||||||
require.Nil(t, controller.owner())
|
|
||||||
registerCalls, unregisterCalls := controller.callCounts()
|
|
||||||
require.Equal(t, 1, registerCalls)
|
|
||||||
require.Equal(t, 1, unregisterCalls)
|
|
||||||
require.True(t, controllerConn.closed.Load())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestVirtualNetPluginOldConnectionCleanupKeepsReplacementRoute(t *testing.T) {
|
|
||||||
controller := newFakeClientRouteController()
|
|
||||||
oldPlugin := newTestVirtualNetPlugin(t, controller)
|
|
||||||
newPlugin := newTestVirtualNetPlugin(t, controller)
|
|
||||||
oldControllerConn, oldPluginConn := newTrackedPipe(t)
|
|
||||||
newControllerConn, newPluginConn := newTrackedPipe(t)
|
|
||||||
|
|
||||||
require.True(t, oldPlugin.registerControllerConn(oldControllerConn, oldPluginConn))
|
|
||||||
require.True(t, newPlugin.registerControllerConn(newControllerConn, newPluginConn))
|
|
||||||
require.Same(t, newControllerConn, controller.owner())
|
|
||||||
|
|
||||||
oldPlugin.cleanupControllerConn(xlog.FromContextSafe(oldPlugin.ctx), oldControllerConn)
|
|
||||||
|
|
||||||
require.Same(t, newControllerConn, controller.owner())
|
|
||||||
require.True(t, oldControllerConn.closed.Load())
|
|
||||||
}
|
|
||||||
@@ -63,13 +63,9 @@ func ForwardUserConn(udpConn *net.UDPConn, readCh <-chan *msg.UDPPacket, sendCh
|
|||||||
// NewUDPPacket copies buf[:n], so the read buffer can be reused
|
// NewUDPPacket copies buf[:n], so the read buffer can be reused
|
||||||
udpMsg := NewUDPPacket(buf[:n], nil, remoteAddr)
|
udpMsg := NewUDPPacket(buf[:n], nil, remoteAddr)
|
||||||
|
|
||||||
if err = errors.PanicToError(func() {
|
select {
|
||||||
select {
|
case sendCh <- udpMsg:
|
||||||
case sendCh <- udpMsg:
|
default:
|
||||||
default:
|
|
||||||
}
|
|
||||||
}); err != nil {
|
|
||||||
return
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,13 +1,9 @@
|
|||||||
package udp
|
package udp
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestUdpPacket(t *testing.T) {
|
func TestUdpPacket(t *testing.T) {
|
||||||
@@ -20,33 +16,3 @@ func TestUdpPacket(t *testing.T) {
|
|||||||
require.NoError(err)
|
require.NoError(err)
|
||||||
require.EqualValues(buf, newBuf)
|
require.EqualValues(buf, newBuf)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestForwardUserConnReturnsWhenSendChannelIsClosed(t *testing.T) {
|
|
||||||
listener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = listener.Close() })
|
|
||||||
|
|
||||||
readCh := make(chan *msg.UDPPacket)
|
|
||||||
sendCh := make(chan *msg.UDPPacket)
|
|
||||||
close(sendCh)
|
|
||||||
t.Cleanup(func() { close(readCh) })
|
|
||||||
|
|
||||||
done := make(chan struct{})
|
|
||||||
go func() {
|
|
||||||
ForwardUserConn(listener, readCh, sendCh, 1500)
|
|
||||||
close(done)
|
|
||||||
}()
|
|
||||||
|
|
||||||
sender, err := net.DialUDP("udp4", nil, listener.LocalAddr().(*net.UDPAddr))
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = sender.Close() })
|
|
||||||
|
|
||||||
_, err = sender.Write([]byte("trigger"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("ForwardUserConn did not return after sending to a closed channel")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -68,8 +68,7 @@ func NewServerHello(clientHello ClientHello) (ServerHello, error) {
|
|||||||
return ServerHello{
|
return ServerHello{
|
||||||
Selected: ServerSelection{
|
Selected: ServerSelection{
|
||||||
Message: MessageSelection{
|
Message: MessageSelection{
|
||||||
Codec: MessageCodecJSON,
|
Codec: MessageCodecJSON,
|
||||||
UDPPacketCodec: selectUDPPacketCodec(clientHello.Capabilities.Message.UDPPacketCodecs),
|
|
||||||
},
|
},
|
||||||
Crypto: CryptoSelection{
|
Crypto: CryptoSelection{
|
||||||
Algorithm: algorithm,
|
Algorithm: algorithm,
|
||||||
@@ -93,15 +92,6 @@ func ValidateServerHelloForClient(clientHello ClientHello, serverHello ServerHel
|
|||||||
if serverHello.Selected.Message.Codec != MessageCodecJSON {
|
if serverHello.Selected.Message.Codec != MessageCodecJSON {
|
||||||
return fmt.Errorf("unsupported selected message codec: %s", serverHello.Selected.Message.Codec)
|
return fmt.Errorf("unsupported selected message codec: %s", serverHello.Selected.Message.Codec)
|
||||||
}
|
}
|
||||||
udpPacketCodec := serverHello.Selected.Message.UDPPacketCodec
|
|
||||||
if udpPacketCodec != "" {
|
|
||||||
if udpPacketCodec != UDPPacketCodecBinary {
|
|
||||||
return fmt.Errorf("unsupported selected UDP packet codec: %s", udpPacketCodec)
|
|
||||||
}
|
|
||||||
if !Supports(clientHello.Capabilities.Message.UDPPacketCodecs, udpPacketCodec) {
|
|
||||||
return fmt.Errorf("selected UDP packet codec was not advertised by client: %s", udpPacketCodec)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
cryptoSelection := serverHello.Selected.Crypto
|
cryptoSelection := serverHello.Selected.Crypto
|
||||||
if !IsSupportedAEADAlgorithm(cryptoSelection.Algorithm) {
|
if !IsSupportedAEADAlgorithm(cryptoSelection.Algorithm) {
|
||||||
return fmt.Errorf("unknown selected crypto algorithm: %s", cryptoSelection.Algorithm)
|
return fmt.Errorf("unknown selected crypto algorithm: %s", cryptoSelection.Algorithm)
|
||||||
@@ -115,13 +105,6 @@ func ValidateServerHelloForClient(clientHello ClientHello, serverHello ServerHel
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func selectUDPPacketCodec(codecs []string) string {
|
|
||||||
if Supports(codecs, UDPPacketCodecBinary) {
|
|
||||||
return UDPPacketCodecBinary
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewCryptoContext(algorithm string, clientHelloPayload, serverHelloPayload []byte) *CryptoContext {
|
func NewCryptoContext(algorithm string, clientHelloPayload, serverHelloPayload []byte) *CryptoContext {
|
||||||
return &CryptoContext{
|
return &CryptoContext{
|
||||||
Algorithm: algorithm,
|
Algorithm: algorithm,
|
||||||
|
|||||||
@@ -36,7 +36,6 @@ const (
|
|||||||
FrameTypeMessage uint16 = 16
|
FrameTypeMessage uint16 = 16
|
||||||
|
|
||||||
MessageCodecJSON = "json"
|
MessageCodecJSON = "json"
|
||||||
UDPPacketCodecBinary = "binary-v1"
|
|
||||||
DefaultMaxFramePayloadSize = 64 * 1024
|
DefaultMaxFramePayloadSize = 64 * 1024
|
||||||
|
|
||||||
MagicV2 = "FRP\x00\x02\r\n"
|
MagicV2 = "FRP\x00\x02\r\n"
|
||||||
@@ -183,8 +182,7 @@ type ClientCapabilities struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type MessageCapabilities struct {
|
type MessageCapabilities struct {
|
||||||
Codecs []string `json:"codecs,omitempty"`
|
Codecs []string `json:"codecs,omitempty"`
|
||||||
UDPPacketCodecs []string `json:"udpPacketCodecs,omitempty"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type CryptoCapabilities struct {
|
type CryptoCapabilities struct {
|
||||||
@@ -203,8 +201,7 @@ type ServerSelection struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type MessageSelection struct {
|
type MessageSelection struct {
|
||||||
Codec string `json:"codec,omitempty"`
|
Codec string `json:"codec,omitempty"`
|
||||||
UDPPacketCodec string `json:"udpPacketCodec,omitempty"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type CryptoSelection struct {
|
type CryptoSelection struct {
|
||||||
@@ -217,8 +214,7 @@ func clientHelloWithCryptoRandom(bootstrap BootstrapInfo, clientRandom []byte) C
|
|||||||
Bootstrap: bootstrap,
|
Bootstrap: bootstrap,
|
||||||
Capabilities: ClientCapabilities{
|
Capabilities: ClientCapabilities{
|
||||||
Message: MessageCapabilities{
|
Message: MessageCapabilities{
|
||||||
Codecs: []string{MessageCodecJSON},
|
Codecs: []string{MessageCodecJSON},
|
||||||
UDPPacketCodecs: []string{UDPPacketCodecBinary},
|
|
||||||
},
|
},
|
||||||
Crypto: CryptoCapabilities{
|
Crypto: CryptoCapabilities{
|
||||||
Algorithms: PreferredAEADAlgorithms(),
|
Algorithms: PreferredAEADAlgorithms(),
|
||||||
|
|||||||
@@ -148,40 +148,10 @@ func TestNewServerHelloSelectsFirstSupportedAEADAlgorithm(t *testing.T) {
|
|||||||
serverHello, err := NewServerHello(hello)
|
serverHello, err := NewServerHello(hello)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, MessageCodecJSON, serverHello.Selected.Message.Codec)
|
require.Equal(t, MessageCodecJSON, serverHello.Selected.Message.Codec)
|
||||||
require.Equal(t, UDPPacketCodecBinary, serverHello.Selected.Message.UDPPacketCodec)
|
|
||||||
require.Equal(t, AEADAlgorithmXChaCha20Poly1305, serverHello.Selected.Crypto.Algorithm)
|
require.Equal(t, AEADAlgorithmXChaCha20Poly1305, serverHello.Selected.Crypto.Algorithm)
|
||||||
require.Len(t, serverHello.Selected.Crypto.ServerRandom, CryptoRandomSize)
|
require.Len(t, serverHello.Selected.Crypto.ServerRandom, CryptoRandomSize)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUDPPacketCodecNegotiationFallbackAndValidation(t *testing.T) {
|
|
||||||
hello := mustClientHello(t, BootstrapInfo{})
|
|
||||||
serverHello, err := NewServerHello(hello)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, UDPPacketCodecBinary, serverHello.Selected.Message.UDPPacketCodec)
|
|
||||||
require.NoError(t, ValidateServerHelloForClient(hello, serverHello))
|
|
||||||
|
|
||||||
legacyHello := hello
|
|
||||||
legacyHello.Capabilities.Message.UDPPacketCodecs = nil
|
|
||||||
legacyServerHello, err := NewServerHello(legacyHello)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Empty(t, legacyServerHello.Selected.Message.UDPPacketCodec)
|
|
||||||
require.NoError(t, ValidateServerHelloForClient(legacyHello, legacyServerHello))
|
|
||||||
|
|
||||||
unknownOffer := hello
|
|
||||||
unknownOffer.Capabilities.Message.UDPPacketCodecs = []string{"unknown"}
|
|
||||||
unknownServerHello, err := NewServerHello(unknownOffer)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Empty(t, unknownServerHello.Selected.Message.UDPPacketCodec)
|
|
||||||
|
|
||||||
rejected := serverHello
|
|
||||||
rejected.Selected.Message.UDPPacketCodec = "unknown"
|
|
||||||
require.ErrorContains(t, ValidateServerHelloForClient(hello, rejected), "unsupported selected UDP packet codec")
|
|
||||||
|
|
||||||
unadvertised := serverHello
|
|
||||||
unadvertised.Selected.Message.UDPPacketCodec = UDPPacketCodecBinary
|
|
||||||
require.ErrorContains(t, ValidateServerHelloForClient(legacyHello, unadvertised), "was not advertised")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewClientCryptoContextValidatesServerHello(t *testing.T) {
|
func TestNewClientCryptoContextValidatesServerHello(t *testing.T) {
|
||||||
hello := mustClientHello(t, BootstrapInfo{})
|
hello := mustClientHello(t, BootstrapInfo{})
|
||||||
serverHello, err := NewServerHello(hello)
|
serverHello, err := NewServerHello(hello)
|
||||||
|
|||||||
+5
-21
@@ -16,6 +16,7 @@ package ssh
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/binary"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
@@ -51,11 +52,6 @@ type tcpipForward struct {
|
|||||||
Port uint32
|
Port uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
// https://datatracker.ietf.org/doc/html/rfc4254#section-6.5
|
|
||||||
type execPayload struct {
|
|
||||||
Command string
|
|
||||||
}
|
|
||||||
|
|
||||||
// https://datatracker.ietf.org/doc/html/rfc4254#page-16
|
// https://datatracker.ietf.org/doc/html/rfc4254#page-16
|
||||||
type forwardedTCPPayload struct {
|
type forwardedTCPPayload struct {
|
||||||
Addr string
|
Addr string
|
||||||
@@ -70,7 +66,6 @@ type TunnelServer struct {
|
|||||||
sshConn *ssh.ServerConn
|
sshConn *ssh.ServerConn
|
||||||
sc *ssh.ServerConfig
|
sc *ssh.ServerConfig
|
||||||
firstChannel ssh.Channel
|
firstChannel ssh.Channel
|
||||||
firstChannelMu sync.Mutex
|
|
||||||
|
|
||||||
vc *virtual.Client
|
vc *virtual.Client
|
||||||
peerServerListener *netpkg.InternalListener
|
peerServerListener *netpkg.InternalListener
|
||||||
@@ -192,8 +187,6 @@ func (s *TunnelServer) Run() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *TunnelServer) writeToClient(data string) {
|
func (s *TunnelServer) writeToClient(data string) {
|
||||||
s.firstChannelMu.Lock()
|
|
||||||
defer s.firstChannelMu.Unlock()
|
|
||||||
if s.firstChannel == nil {
|
if s.firstChannel == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -307,24 +300,23 @@ func (s *TunnelServer) handleNewChannel(channel ssh.NewChannel, extraPayloadCh c
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
s.firstChannelMu.Lock()
|
|
||||||
if s.firstChannel == nil {
|
if s.firstChannel == nil {
|
||||||
s.firstChannel = ch
|
s.firstChannel = ch
|
||||||
}
|
}
|
||||||
s.firstChannelMu.Unlock()
|
|
||||||
go s.keepAlive(ch)
|
go s.keepAlive(ch)
|
||||||
|
|
||||||
for req := range reqs {
|
for req := range reqs {
|
||||||
if req.WantReply {
|
if req.WantReply {
|
||||||
_ = req.Reply(true, nil)
|
_ = req.Reply(true, nil)
|
||||||
}
|
}
|
||||||
if req.Type != "exec" {
|
if req.Type != "exec" || len(req.Payload) <= 4 {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
extraPayload, ok := parseExecPayload(req.Payload)
|
end := 4 + binary.BigEndian.Uint32(req.Payload[:4])
|
||||||
if !ok {
|
if len(req.Payload) < int(end) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
extraPayload := string(req.Payload[4:end])
|
||||||
select {
|
select {
|
||||||
case extraPayloadCh <- extraPayload:
|
case extraPayloadCh <- extraPayload:
|
||||||
default:
|
default:
|
||||||
@@ -332,14 +324,6 @@ func (s *TunnelServer) handleNewChannel(channel ssh.NewChannel, extraPayloadCh c
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseExecPayload(payload []byte) (string, bool) {
|
|
||||||
var msg execPayload
|
|
||||||
if err := ssh.Unmarshal(payload, &msg); err != nil {
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
return msg.Command, true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *TunnelServer) keepAlive(ch ssh.Channel) {
|
func (s *TunnelServer) keepAlive(ch ssh.Channel) {
|
||||||
tk := time.NewTicker(time.Second * 30)
|
tk := time.NewTicker(time.Second * 30)
|
||||||
defer tk.Stop()
|
defer tk.Stop()
|
||||||
|
|||||||
@@ -1,115 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package ssh
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"io"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
cryptossh "golang.org/x/crypto/ssh"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestParseExecPayload(t *testing.T) {
|
|
||||||
payload := cryptossh.Marshal(&execPayload{Command: "tcp --remote_port 6000"})
|
|
||||||
|
|
||||||
got, ok := parseExecPayload(payload)
|
|
||||||
|
|
||||||
require.True(t, ok)
|
|
||||||
require.Equal(t, "tcp --remote_port 6000", got)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseExecPayloadRejectsMalformedPayloads(t *testing.T) {
|
|
||||||
overflowLength := make([]byte, 5)
|
|
||||||
binary.BigEndian.PutUint32(overflowLength[:4], ^uint32(0))
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
payload []byte
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "empty",
|
|
||||||
payload: nil,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "short length prefix",
|
|
||||||
payload: []byte{0, 0, 0},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "declared length exceeds remaining payload",
|
|
||||||
payload: []byte{0, 0, 0, 2, 'x'},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "overflow length",
|
|
||||||
payload: overflowLength,
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
var (
|
|
||||||
got string
|
|
||||||
ok bool
|
|
||||||
)
|
|
||||||
require.NotPanics(t, func() {
|
|
||||||
got, ok = parseExecPayload(tc.payload)
|
|
||||||
})
|
|
||||||
require.False(t, ok)
|
|
||||||
require.Empty(t, got)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type trackingChannel struct {
|
|
||||||
active atomic.Int32
|
|
||||||
concurrent atomic.Bool
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *trackingChannel) Read([]byte) (int, error) { return 0, io.EOF }
|
|
||||||
|
|
||||||
func (c *trackingChannel) Write(p []byte) (int, error) {
|
|
||||||
if c.active.Add(1) != 1 {
|
|
||||||
c.concurrent.Store(true)
|
|
||||||
}
|
|
||||||
time.Sleep(time.Millisecond)
|
|
||||||
c.active.Add(-1)
|
|
||||||
return len(p), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *trackingChannel) Close() error { return nil }
|
|
||||||
func (c *trackingChannel) CloseWrite() error { return nil }
|
|
||||||
func (c *trackingChannel) SendRequest(string, bool, []byte) (bool, error) { return false, nil }
|
|
||||||
func (c *trackingChannel) Stderr() io.ReadWriter { return nil }
|
|
||||||
|
|
||||||
func TestWriteToClientSerializesChannelWrites(t *testing.T) {
|
|
||||||
channel := &trackingChannel{}
|
|
||||||
s := &TunnelServer{firstChannel: channel}
|
|
||||||
start := make(chan struct{})
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
for range 8 {
|
|
||||||
wg.Go(func() {
|
|
||||||
<-start
|
|
||||||
s.writeToClient("message")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
close(start)
|
|
||||||
wg.Wait()
|
|
||||||
|
|
||||||
if channel.concurrent.Load() {
|
|
||||||
t.Fatal("channel writes were concurrent")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -26,12 +26,6 @@ type GeneralResponse struct {
|
|||||||
Msg string
|
Msg string
|
||||||
}
|
}
|
||||||
|
|
||||||
type V2Response struct {
|
|
||||||
Code int `json:"code"`
|
|
||||||
Msg string `json:"msg"`
|
|
||||||
Data any `json:"data"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// APIHandler is a handler function that returns a response object or an error.
|
// APIHandler is a handler function that returns a response object or an error.
|
||||||
type APIHandler func(ctx *Context) (any, error)
|
type APIHandler func(ctx *Context) (any, error)
|
||||||
|
|
||||||
@@ -70,27 +64,3 @@ func MakeHTTPHandlerFunc(handler APIHandler) http.HandlerFunc {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// MakeHTTPHandlerFuncV2 wraps a handler response in the dashboard API v2 envelope.
|
|
||||||
func MakeHTTPHandlerFuncV2(handler APIHandler) http.HandlerFunc {
|
|
||||||
return func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
ctx := NewContext(w, r)
|
|
||||||
res, err := handler(ctx)
|
|
||||||
if err != nil {
|
|
||||||
log.Warnf("http response [%s]: error: %v", r.URL.Path, err)
|
|
||||||
code := http.StatusInternalServerError
|
|
||||||
if e, ok := err.(*Error); ok {
|
|
||||||
code = e.Code
|
|
||||||
}
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
w.WriteHeader(code)
|
|
||||||
_ = json.NewEncoder(w).Encode(V2Response{Code: code, Msg: err.Error(), Data: nil})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
w.WriteHeader(http.StatusOK)
|
|
||||||
_ = json.NewEncoder(w).Encode(V2Response{Code: http.StatusOK, Msg: "success", Data: res})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,37 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package limit
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
"golang.org/x/time/rate"
|
|
||||||
)
|
|
||||||
|
|
||||||
// NewBandwidthLimiter creates a limiter whose rate preserves the configured
|
|
||||||
// byte limit while keeping the burst representable as an int on all targets.
|
|
||||||
func NewBandwidthLimiter(bytes int64) *rate.Limiter {
|
|
||||||
if bytes <= 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
maxInt := int64(^uint(0) >> 1)
|
|
||||||
burst := min(bytes, maxInt)
|
|
||||||
return rate.NewLimiter(rate.Limit(float64(bytes)), int(burst))
|
|
||||||
}
|
|
||||||
|
|
||||||
func invalidBurstError(burst int) error {
|
|
||||||
return fmt.Errorf("invalid limiter burst: %d", burst)
|
|
||||||
}
|
|
||||||
@@ -1,65 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package limit
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"golang.org/x/time/rate"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewBandwidthLimiterClampsBurstToTargetInt(t *testing.T) {
|
|
||||||
const bytesPerSecond = int64(1 << 31)
|
|
||||||
|
|
||||||
limiter := NewBandwidthLimiter(bytesPerSecond)
|
|
||||||
require.NotNil(t, limiter)
|
|
||||||
|
|
||||||
wantBurst := bytesPerSecond
|
|
||||||
maxInt := int64(^uint(0) >> 1)
|
|
||||||
if wantBurst > maxInt {
|
|
||||||
wantBurst = maxInt
|
|
||||||
}
|
|
||||||
require.Equal(t, int(wantBurst), limiter.Burst())
|
|
||||||
require.Equal(t, rate.Limit(float64(bytesPerSecond)), limiter.Limit())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewBandwidthLimiterDisablesNonPositiveLimit(t *testing.T) {
|
|
||||||
require.Nil(t, NewBandwidthLimiter(0))
|
|
||||||
require.Nil(t, NewBandwidthLimiter(-1))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReaderAndWriterRejectInvalidBurst(t *testing.T) {
|
|
||||||
for _, burst := range []int{0, -1} {
|
|
||||||
t.Run("reader/"+strconv.Itoa(burst), func(t *testing.T) {
|
|
||||||
reader := NewReader(strings.NewReader("payload"), rate.NewLimiter(rate.Limit(1), burst))
|
|
||||||
n, err := reader.Read(make([]byte, 1))
|
|
||||||
require.Zero(t, n)
|
|
||||||
require.EqualError(t, err, "invalid limiter burst: "+strconv.Itoa(burst))
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("writer/"+strconv.Itoa(burst), func(t *testing.T) {
|
|
||||||
var dst bytes.Buffer
|
|
||||||
writer := NewWriter(&dst, rate.NewLimiter(rate.Limit(1), burst))
|
|
||||||
n, err := writer.Write([]byte("payload"))
|
|
||||||
require.Zero(t, n)
|
|
||||||
require.EqualError(t, err, "invalid limiter burst: "+strconv.Itoa(burst))
|
|
||||||
require.Empty(t, dst.Bytes())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -35,12 +35,6 @@ func NewReader(r io.Reader, limiter *rate.Limiter) *Reader {
|
|||||||
|
|
||||||
func (r *Reader) Read(p []byte) (n int, err error) {
|
func (r *Reader) Read(p []byte) (n int, err error) {
|
||||||
b := r.limiter.Burst()
|
b := r.limiter.Burst()
|
||||||
if b <= 0 {
|
|
||||||
if len(p) == 0 {
|
|
||||||
return 0, nil
|
|
||||||
}
|
|
||||||
return 0, invalidBurstError(b)
|
|
||||||
}
|
|
||||||
if b < len(p) {
|
if b < len(p) {
|
||||||
p = p[:b]
|
p = p[:b]
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -34,15 +34,8 @@ func NewWriter(w io.Writer, limiter *rate.Limiter) *Writer {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (w *Writer) Write(p []byte) (n int, err error) {
|
func (w *Writer) Write(p []byte) (n int, err error) {
|
||||||
if len(p) == 0 {
|
|
||||||
return 0, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var nn int
|
var nn int
|
||||||
b := w.limiter.Burst()
|
b := w.limiter.Burst()
|
||||||
if b <= 0 {
|
|
||||||
return 0, invalidBurstError(b)
|
|
||||||
}
|
|
||||||
for {
|
for {
|
||||||
end := len(p)
|
end := len(p)
|
||||||
if end == 0 {
|
if end == 0 {
|
||||||
|
|||||||
@@ -16,7 +16,6 @@ package net
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/hkdf"
|
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
@@ -26,6 +25,7 @@ import (
|
|||||||
|
|
||||||
libcrypto "github.com/fatedier/golib/crypto"
|
libcrypto "github.com/fatedier/golib/crypto"
|
||||||
quic "github.com/quic-go/quic-go"
|
quic "github.com/quic-go/quic-go"
|
||||||
|
"golang.org/x/crypto/hkdf"
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/util/xlog"
|
"github.com/fatedier/frp/pkg/util/xlog"
|
||||||
)
|
)
|
||||||
@@ -335,6 +335,11 @@ func deriveAEADControlKeys(key []byte, algorithm string, transcriptHash []byte)
|
|||||||
}
|
}
|
||||||
|
|
||||||
func deriveAEADControlKey(key []byte, algorithm string, transcriptHash []byte, direction string) ([]byte, error) {
|
func deriveAEADControlKey(key []byte, algorithm string, transcriptHash []byte, direction string) ([]byte, error) {
|
||||||
info := aeadControlHKDFInfoPrefix + " " + algorithm + " " + direction
|
info := []byte(aeadControlHKDFInfoPrefix + " " + algorithm + " " + direction)
|
||||||
return hkdf.Key(sha256.New, key, transcriptHash, info, libcrypto.AEADKeySize)
|
reader := hkdf.New(sha256.New, key, transcriptHash, info)
|
||||||
|
out := make([]byte, libcrypto.AEADKeySize)
|
||||||
|
if _, err := io.ReadFull(reader, out); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -114,11 +114,5 @@ func TestDeriveAEADControlKeysUsesDistinctDirections(t *testing.T) {
|
|||||||
bytes.Repeat([]byte{0x44}, 32),
|
bytes.Repeat([]byte{0x44}, 32),
|
||||||
)
|
)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, []byte{
|
|
||||||
0xa0, 0x58, 0xcd, 0x02, 0x5d, 0x96, 0x98, 0x5f,
|
|
||||||
0xeb, 0xeb, 0xff, 0x79, 0xa1, 0x9f, 0x62, 0xb7,
|
|
||||||
0x15, 0xe0, 0x53, 0x91, 0x3d, 0xfc, 0x74, 0x77,
|
|
||||||
0x05, 0x91, 0x4c, 0x62, 0x4b, 0xf3, 0xd4, 0x95,
|
|
||||||
}, clientToServerKey)
|
|
||||||
require.NotEqual(t, clientToServerKey, serverToClientKey)
|
require.NotEqual(t, clientToServerKey, serverToClientKey)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -45,11 +45,6 @@ func DialHookWebsocket(protocol string, host string) libnet.AfterHookFunc {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, nil, err
|
return nil, nil, err
|
||||||
}
|
}
|
||||||
// The tunnel payload is a raw byte stream (yamux), not UTF-8 text.
|
|
||||||
// Send it as binary frames; otherwise RFC 6455-compliant intermediaries
|
|
||||||
// (e.g. API gateways/reverse proxies) UTF-8-validate the default text
|
|
||||||
// frames and close the connection on invalid bytes.
|
|
||||||
conn.PayloadType = websocket.BinaryFrame
|
|
||||||
return ctx, conn, nil
|
return ctx, conn, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -32,11 +32,6 @@ func NewWebsocketListener(ln net.Listener) (wl *WebsocketListener) {
|
|||||||
|
|
||||||
muxer := http.NewServeMux()
|
muxer := http.NewServeMux()
|
||||||
muxer.Handle(FrpWebsocketPath, websocket.Handler(func(c *websocket.Conn) {
|
muxer.Handle(FrpWebsocketPath, websocket.Handler(func(c *websocket.Conn) {
|
||||||
// The tunnel payload is a raw byte stream (yamux), not UTF-8 text.
|
|
||||||
// Send it as binary frames; otherwise RFC 6455-compliant intermediaries
|
|
||||||
// (e.g. API gateways/reverse proxies) UTF-8-validate the default text
|
|
||||||
// frames and close the connection on invalid bytes.
|
|
||||||
c.PayloadType = websocket.BinaryFrame
|
|
||||||
notifyCh := make(chan struct{})
|
notifyCh := make(chan struct{})
|
||||||
conn := WrapCloseNotifyConn(c, func(_ error) {
|
conn := WrapCloseNotifyConn(c, func(_ error) {
|
||||||
close(notifyCh)
|
close(notifyCh)
|
||||||
|
|||||||
@@ -14,7 +14,7 @@
|
|||||||
|
|
||||||
package version
|
package version
|
||||||
|
|
||||||
var version = "0.71.0"
|
var version = "0.69.1"
|
||||||
|
|
||||||
func Full() string {
|
func Full() string {
|
||||||
return version
|
return version
|
||||||
|
|||||||
@@ -28,6 +28,8 @@ import (
|
|||||||
|
|
||||||
libio "github.com/fatedier/golib/io"
|
libio "github.com/fatedier/golib/io"
|
||||||
"github.com/fatedier/golib/pool"
|
"github.com/fatedier/golib/pool"
|
||||||
|
"golang.org/x/net/http2"
|
||||||
|
"golang.org/x/net/http2/h2c"
|
||||||
|
|
||||||
httppkg "github.com/fatedier/frp/pkg/util/http"
|
httppkg "github.com/fatedier/frp/pkg/util/http"
|
||||||
"github.com/fatedier/frp/pkg/util/log"
|
"github.com/fatedier/frp/pkg/util/log"
|
||||||
@@ -137,7 +139,7 @@ func NewHTTPReverseProxy(option HTTPReverseProxyOptions, vhostRouter *Routers) *
|
|||||||
_, _ = rw.Write(getNotFoundPageContent())
|
_, _ = rw.Write(getNotFoundPageContent())
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
rp.proxy = proxy
|
rp.proxy = h2c.NewHandler(proxy, &http2.Server{})
|
||||||
return rp
|
return rp
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,94 +1,14 @@
|
|||||||
package vhost
|
package vhost
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bufio"
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
httppkg "github.com/fatedier/frp/pkg/util/http"
|
httppkg "github.com/fatedier/frp/pkg/util/http"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestHTTPServerProtocols(t *testing.T) {
|
|
||||||
rp := NewHTTPReverseProxy(HTTPReverseProxyOptions{}, NewRouters())
|
|
||||||
protocols := new(http.Protocols)
|
|
||||||
protocols.SetHTTP1(true)
|
|
||||||
protocols.SetUnencryptedHTTP2(true)
|
|
||||||
server := &http.Server{
|
|
||||||
Handler: rp,
|
|
||||||
ReadHeaderTimeout: time.Second,
|
|
||||||
Protocols: protocols,
|
|
||||||
}
|
|
||||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
serveErr := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
serveErr <- server.Serve(listener)
|
|
||||||
}()
|
|
||||||
defer func() {
|
|
||||||
require.NoError(t, server.Close())
|
|
||||||
require.ErrorIs(t, <-serveErr, http.ErrServerClosed)
|
|
||||||
}()
|
|
||||||
|
|
||||||
require.True(t, server.Protocols.HTTP1())
|
|
||||||
require.True(t, server.Protocols.UnencryptedHTTP2())
|
|
||||||
|
|
||||||
t.Run("HTTP/1.1", func(t *testing.T) {
|
|
||||||
transport := &http.Transport{Protocols: httpProtocols(true, false)}
|
|
||||||
defer transport.CloseIdleConnections()
|
|
||||||
client := &http.Client{Transport: transport}
|
|
||||||
response, err := client.Get("http://" + listener.Addr().String() + "/")
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer response.Body.Close()
|
|
||||||
|
|
||||||
require.Equal(t, "HTTP/1.1", response.Proto)
|
|
||||||
require.Equal(t, http.StatusNotFound, response.StatusCode)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("HTTP/2 prior knowledge", func(t *testing.T) {
|
|
||||||
transport := &http.Transport{Protocols: httpProtocols(false, true)}
|
|
||||||
defer transport.CloseIdleConnections()
|
|
||||||
client := &http.Client{Transport: transport}
|
|
||||||
response, err := client.Get("http://" + listener.Addr().String() + "/")
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer response.Body.Close()
|
|
||||||
|
|
||||||
require.Equal(t, "HTTP/2.0", response.Proto)
|
|
||||||
require.Equal(t, http.StatusNotFound, response.StatusCode)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("HTTP/1.1 Upgrade h2c", func(t *testing.T) {
|
|
||||||
conn, err := net.Dial("tcp", listener.Addr().String())
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer conn.Close()
|
|
||||||
|
|
||||||
_, err = fmt.Fprintf(conn,
|
|
||||||
"GET / HTTP/1.1\r\nHost: %s\r\n"+
|
|
||||||
"Connection: Upgrade, HTTP2-Settings\r\nUpgrade: h2c\r\n"+
|
|
||||||
"HTTP2-Settings: AAMAAABkAAQCAAAAAAIAAAAA\r\n\r\n",
|
|
||||||
listener.Addr())
|
|
||||||
require.NoError(t, err)
|
|
||||||
response, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer response.Body.Close()
|
|
||||||
|
|
||||||
require.NotEqual(t, http.StatusSwitchingProtocols, response.StatusCode)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func httpProtocols(http1, unencryptedHTTP2 bool) *http.Protocols {
|
|
||||||
protocols := new(http.Protocols)
|
|
||||||
protocols.SetHTTP1(http1)
|
|
||||||
protocols.SetUnencryptedHTTP2(unencryptedHTTP2)
|
|
||||||
return protocols
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCheckRouteAuthByRequest(t *testing.T) {
|
func TestCheckRouteAuthByRequest(t *testing.T) {
|
||||||
rc := &RouteConfig{
|
rc := &RouteConfig{
|
||||||
Username: "alice",
|
Username: "alice",
|
||||||
|
|||||||
@@ -96,21 +96,21 @@ func (l *Logger) Spawn() *Logger {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (l *Logger) Errorf(format string, v ...any) {
|
func (l *Logger) Errorf(format string, v ...any) {
|
||||||
log.Logger.WithPrefix(l.prefixString).Errorf(format, v...)
|
log.Logger.Errorf(l.prefixString+format, v...)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (l *Logger) Warnf(format string, v ...any) {
|
func (l *Logger) Warnf(format string, v ...any) {
|
||||||
log.Logger.WithPrefix(l.prefixString).Warnf(format, v...)
|
log.Logger.Warnf(l.prefixString+format, v...)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (l *Logger) Infof(format string, v ...any) {
|
func (l *Logger) Infof(format string, v ...any) {
|
||||||
log.Logger.WithPrefix(l.prefixString).Infof(format, v...)
|
log.Logger.Infof(l.prefixString+format, v...)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (l *Logger) Debugf(format string, v ...any) {
|
func (l *Logger) Debugf(format string, v ...any) {
|
||||||
log.Logger.WithPrefix(l.prefixString).Debugf(format, v...)
|
log.Logger.Debugf(l.prefixString+format, v...)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (l *Logger) Tracef(format string, v ...any) {
|
func (l *Logger) Tracef(format string, v ...any) {
|
||||||
log.Logger.WithPrefix(l.prefixString).Tracef(format, v...)
|
log.Logger.Tracef(l.prefixString+format, v...)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,76 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package xlog
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
goliblog "github.com/fatedier/golib/log"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
frplog "github.com/fatedier/frp/pkg/util/log"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestPrefixIsNotPartOfFormatString(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
log func(*Logger)
|
|
||||||
}{
|
|
||||||
{name: "error", log: func(xl *Logger) { xl.Errorf("value [%s]", "ok") }},
|
|
||||||
{name: "warn", log: func(xl *Logger) { xl.Warnf("value [%s]", "ok") }},
|
|
||||||
{name: "info", log: func(xl *Logger) { xl.Infof("value [%s]", "ok") }},
|
|
||||||
{name: "debug", log: func(xl *Logger) { xl.Debugf("value [%s]", "ok") }},
|
|
||||||
{name: "trace", log: func(xl *Logger) { xl.Tracef("value [%s]", "ok") }},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
output := captureLogs(t)
|
|
||||||
|
|
||||||
tt.log(New().AppendPrefix("%1000000s"))
|
|
||||||
|
|
||||||
require.Contains(t, output.String(), "[%1000000s] value [ok]")
|
|
||||||
require.Less(t, output.Len(), 1024)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFormattingSemanticsArePreserved(t *testing.T) {
|
|
||||||
output := captureLogs(t)
|
|
||||||
xl := New().AppendPrefix("run")
|
|
||||||
|
|
||||||
xl.Infof("%[2]s %[1]s", "first", "second")
|
|
||||||
xl.Infof("100% complete")
|
|
||||||
|
|
||||||
require.Contains(t, output.String(), "[run] second first")
|
|
||||||
require.Contains(t, output.String(), "[run] 100% complete")
|
|
||||||
}
|
|
||||||
|
|
||||||
func captureLogs(t *testing.T) *bytes.Buffer {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
output := bytes.NewBuffer(nil)
|
|
||||||
oldLogger := frplog.Logger
|
|
||||||
frplog.Logger = goliblog.New(
|
|
||||||
goliblog.WithOutput(output),
|
|
||||||
goliblog.WithLevel(goliblog.TraceLevel),
|
|
||||||
goliblog.WithCaller(false),
|
|
||||||
)
|
|
||||||
t.Cleanup(func() {
|
|
||||||
frplog.Logger = oldLogger
|
|
||||||
})
|
|
||||||
return output
|
|
||||||
}
|
|
||||||
@@ -246,9 +246,9 @@ func (c *Controller) RegisterClientRoute(ctx context.Context, name string, route
|
|||||||
go c.readLoopClient(ctx, conn)
|
go c.readLoopClient(ctx, conn)
|
||||||
}
|
}
|
||||||
|
|
||||||
// UnregisterClientRoute removes a client route only when it is still owned by conn.
|
// UnregisterClientRoute Remove client route from routing table
|
||||||
func (c *Controller) UnregisterClientRoute(name string, conn io.Writer) bool {
|
func (c *Controller) UnregisterClientRoute(name string) {
|
||||||
return c.clientRouter.delRoute(name, conn)
|
c.clientRouter.delRoute(name)
|
||||||
}
|
}
|
||||||
|
|
||||||
// StartServerConnReadLoop starts the read loop for a server connection
|
// StartServerConnReadLoop starts the read loop for a server connection
|
||||||
@@ -304,15 +304,10 @@ func (r *clientRouter) findConn(dst net.IP) (io.Writer, error) {
|
|||||||
return nil, fmt.Errorf("no route found for destination %s", dst)
|
return nil, fmt.Errorf("no route found for destination %s", dst)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *clientRouter) delRoute(name string, conn io.Writer) bool {
|
func (r *clientRouter) delRoute(name string) {
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
re, ok := r.routes[name]
|
|
||||||
if !ok || re.conn != conn {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
delete(r.routes, name)
|
delete(r.routes, name)
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *clientRouter) removeConnRoute(conn io.Writer) {
|
func (r *clientRouter) removeConnRoute(conn io.Writer) {
|
||||||
|
|||||||
@@ -1,64 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package vnet
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestClientRouterDeleteRouteRequiresMatchingConnection verifies that a stale
|
|
||||||
// visitor cannot remove a replacement route registered under the same name.
|
|
||||||
func TestClientRouterDeleteRouteRequiresMatchingConnection(t *testing.T) {
|
|
||||||
require := require.New(t)
|
|
||||||
controller := NewController(v1.VirtualNetConfig{})
|
|
||||||
|
|
||||||
_, route, err := net.ParseCIDR("10.1.0.1/32")
|
|
||||||
require.NoError(err)
|
|
||||||
|
|
||||||
oldConn, oldPeer := net.Pipe()
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = oldConn.Close()
|
|
||||||
_ = oldPeer.Close()
|
|
||||||
})
|
|
||||||
replacementConn, replacementPeer := net.Pipe()
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = replacementConn.Close()
|
|
||||||
_ = replacementPeer.Close()
|
|
||||||
})
|
|
||||||
|
|
||||||
controller.clientRouter.addRoute("vnet-visitor", []net.IPNet{*route}, oldConn)
|
|
||||||
controller.clientRouter.addRoute("vnet-visitor", []net.IPNet{*route}, replacementConn)
|
|
||||||
|
|
||||||
require.False(controller.UnregisterClientRoute("vnet-visitor", oldConn))
|
|
||||||
got, err := controller.clientRouter.findConn(net.ParseIP("10.1.0.1"))
|
|
||||||
require.NoError(err)
|
|
||||||
require.Same(replacementConn, got)
|
|
||||||
|
|
||||||
// The read loop for an old connection can exit after a replacement route
|
|
||||||
// has already been registered. Its deferred cleanup must keep the new owner.
|
|
||||||
controller.clientRouter.removeConnRoute(oldConn)
|
|
||||||
got, err = controller.clientRouter.findConn(net.ParseIP("10.1.0.1"))
|
|
||||||
require.NoError(err)
|
|
||||||
require.Same(replacementConn, got)
|
|
||||||
|
|
||||||
require.True(controller.UnregisterClientRoute("vnet-visitor", replacementConn))
|
|
||||||
_, err = controller.clientRouter.findConn(net.ParseIP("10.1.0.1"))
|
|
||||||
require.Error(err)
|
|
||||||
}
|
|
||||||
@@ -48,17 +48,6 @@ func (svr *Service) registerRouteHandlers(helper *httppkg.RouterRegisterHelper)
|
|||||||
subRouter.HandleFunc("/api/clients/{key}", httppkg.MakeHTTPHandlerFunc(apiController.APIClientDetail)).Methods("GET")
|
subRouter.HandleFunc("/api/clients/{key}", httppkg.MakeHTTPHandlerFunc(apiController.APIClientDetail)).Methods("GET")
|
||||||
subRouter.HandleFunc("/api/proxies", httppkg.MakeHTTPHandlerFunc(apiController.DeleteProxies)).Methods("DELETE")
|
subRouter.HandleFunc("/api/proxies", httppkg.MakeHTTPHandlerFunc(apiController.DeleteProxies)).Methods("DELETE")
|
||||||
|
|
||||||
subRouter.HandleFunc("/api/v2/users", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2UserList)).Methods("GET")
|
|
||||||
subRouter.HandleFunc("/api/v2/system/info", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2SystemInfo)).Methods("GET")
|
|
||||||
subRouter.HandleFunc("/api/v2/system/prune", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2SystemPrune)).Methods("POST")
|
|
||||||
subRouter.HandleFunc("/api/v2/clients", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ClientList)).Methods("GET")
|
|
||||||
v2EncodedPathRouter := subRouter.NewRoute().Subrouter()
|
|
||||||
v2EncodedPathRouter.UseEncodedPath()
|
|
||||||
v2EncodedPathRouter.HandleFunc("/api/v2/clients/{key}", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ClientDetail)).Methods("GET")
|
|
||||||
subRouter.HandleFunc("/api/v2/proxies", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyList)).Methods("GET")
|
|
||||||
v2EncodedPathRouter.HandleFunc("/api/v2/proxies/{name}/traffic", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyTraffic)).Methods("GET")
|
|
||||||
v2EncodedPathRouter.HandleFunc("/api/v2/proxies/{name}", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyDetail)).Methods("GET")
|
|
||||||
|
|
||||||
// view
|
// view
|
||||||
subRouter.Handle("/favicon.ico", http.FileServer(helper.AssetsFS)).Methods("GET")
|
subRouter.Handle("/favicon.ico", http.FileServer(helper.AssetsFS)).Methods("GET")
|
||||||
subRouter.PathPrefix("/static/").Handler(
|
subRouter.PathPrefix("/static/").Handler(
|
||||||
|
|||||||
+78
-466
@@ -17,8 +17,6 @@ package server
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
|
||||||
"net"
|
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -42,313 +40,55 @@ import (
|
|||||||
"github.com/fatedier/frp/server/registry"
|
"github.com/fatedier/frp/server/registry"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ControlID uint64
|
|
||||||
|
|
||||||
var nextControlID atomic.Uint64
|
|
||||||
|
|
||||||
const workConnPoolCapacityOffset = 10
|
|
||||||
|
|
||||||
type controlEntry struct {
|
|
||||||
ctl *Control
|
|
||||||
id ControlID
|
|
||||||
// runMu serializes lifecycle and routing decisions for one run ID.
|
|
||||||
// Replacements inherit it; removing the entry releases the manager's reference.
|
|
||||||
runMu *sync.Mutex
|
|
||||||
|
|
||||||
registryOnline bool
|
|
||||||
registryControlID ControlID
|
|
||||||
}
|
|
||||||
|
|
||||||
type ControlManager struct {
|
type ControlManager struct {
|
||||||
// controls indexed by run id
|
// controls indexed by run id
|
||||||
ctlsByRunID map[string]*controlEntry
|
ctlsByRunID map[string]*Control
|
||||||
registry *registry.ClientRegistry
|
|
||||||
closed bool
|
|
||||||
|
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewControlManager(clientRegistry *registry.ClientRegistry) *ControlManager {
|
func NewControlManager() *ControlManager {
|
||||||
return &ControlManager{
|
return &ControlManager{
|
||||||
ctlsByRunID: make(map[string]*controlEntry),
|
ctlsByRunID: make(map[string]*Control),
|
||||||
registry: clientRegistry,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// lockCurrentRun returns the current entry with its run gate held. It never
|
func (cm *ControlManager) Add(runID string, ctl *Control) (old *Control) {
|
||||||
// waits for the gate while holding cm.mu and revalidates the gate after waiting.
|
|
||||||
// The global order is runMu, cm.mu, ctl.lifecycleMu, then registry locks.
|
|
||||||
func (cm *ControlManager) lockCurrentRun(runID string, allowClosed bool) (*controlEntry, bool) {
|
|
||||||
cm.mu.RLock()
|
|
||||||
entry, ok := cm.ctlsByRunID[runID]
|
|
||||||
if cm.closed && !allowClosed {
|
|
||||||
ok = false
|
|
||||||
}
|
|
||||||
cm.mu.RUnlock()
|
|
||||||
if !ok {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
|
|
||||||
runMu := entry.runMu
|
|
||||||
runMu.Lock()
|
|
||||||
cm.mu.RLock()
|
|
||||||
entry, ok = cm.ctlsByRunID[runID]
|
|
||||||
if (cm.closed && !allowClosed) || !ok || entry.runMu != runMu {
|
|
||||||
ok = false
|
|
||||||
}
|
|
||||||
cm.mu.RUnlock()
|
|
||||||
if !ok {
|
|
||||||
runMu.Unlock()
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
return entry, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add makes ctl the pending current generation and records the predecessor
|
|
||||||
// finalization barrier it must wait for before activation.
|
|
||||||
func (cm *ControlManager) Add(ctl *Control) error {
|
|
||||||
for {
|
|
||||||
// Never wait for a run gate while holding cm.mu.
|
|
||||||
cm.mu.RLock()
|
|
||||||
old := cm.ctlsByRunID[ctl.runID]
|
|
||||||
cm.mu.RUnlock()
|
|
||||||
if old != nil {
|
|
||||||
old.runMu.Lock()
|
|
||||||
}
|
|
||||||
|
|
||||||
cm.mu.Lock()
|
|
||||||
if cm.closed {
|
|
||||||
cm.mu.Unlock()
|
|
||||||
if old != nil {
|
|
||||||
old.runMu.Unlock()
|
|
||||||
}
|
|
||||||
return fmt.Errorf("control manager is closed")
|
|
||||||
}
|
|
||||||
if cm.ctlsByRunID[ctl.runID] != old {
|
|
||||||
cm.mu.Unlock()
|
|
||||||
if old != nil {
|
|
||||||
old.runMu.Unlock()
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
id := ControlID(nextControlID.Add(1))
|
|
||||||
if err := ctl.admit(cm, id); err != nil {
|
|
||||||
cm.mu.Unlock()
|
|
||||||
if old != nil {
|
|
||||||
old.runMu.Unlock()
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
runMu := &sync.Mutex{}
|
|
||||||
if old != nil {
|
|
||||||
runMu = old.runMu
|
|
||||||
}
|
|
||||||
entry := &controlEntry{ctl: ctl, id: id, runMu: runMu}
|
|
||||||
var (
|
|
||||||
oldCtl *Control
|
|
||||||
barrier <-chan struct{}
|
|
||||||
)
|
|
||||||
if old != nil {
|
|
||||||
oldCtl = old.ctl
|
|
||||||
barrier = oldCtl.markReplaced()
|
|
||||||
ctl.setHandoffBarrier(barrier)
|
|
||||||
entry.registryOnline = old.registryOnline
|
|
||||||
entry.registryControlID = old.registryControlID
|
|
||||||
}
|
|
||||||
cm.ctlsByRunID[ctl.runID] = entry
|
|
||||||
cm.mu.Unlock()
|
|
||||||
if old != nil {
|
|
||||||
old.runMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
if oldCtl != nil {
|
|
||||||
oldCtl.Replaced(ctl)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Activate registers ctl as online only if it is still the pending current
|
|
||||||
// generation.
|
|
||||||
func (cm *ControlManager) Activate(ctl *Control) (bool, error) {
|
|
||||||
entry, ok := cm.lockCurrentRun(ctl.runID, false)
|
|
||||||
if !ok {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
defer entry.runMu.Unlock()
|
|
||||||
cm.mu.Lock()
|
cm.mu.Lock()
|
||||||
defer cm.mu.Unlock()
|
defer cm.mu.Unlock()
|
||||||
|
|
||||||
if cm.closed || cm.ctlsByRunID[ctl.runID] != entry || entry.ctl != ctl || entry.id != ctl.controlID {
|
var ok bool
|
||||||
return false, nil
|
old, ok = cm.ctlsByRunID[runID]
|
||||||
|
if ok {
|
||||||
|
old.Replaced(ctl)
|
||||||
}
|
}
|
||||||
|
cm.ctlsByRunID[runID] = ctl
|
||||||
ctl.lifecycleMu.Lock()
|
return
|
||||||
defer ctl.lifecycleMu.Unlock()
|
|
||||||
if ctl.state != controlStatePending {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
if ctl.activated {
|
|
||||||
return true, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
loginMsg := ctl.sessionCtx.LoginMsg
|
|
||||||
remoteAddr := ctl.sessionCtx.Conn.RemoteAddr().String()
|
|
||||||
if host, _, err := net.SplitHostPort(remoteAddr); err == nil {
|
|
||||||
remoteAddr = host
|
|
||||||
}
|
|
||||||
_, conflict := cm.registry.RegisterWithControlID(
|
|
||||||
loginMsg.User,
|
|
||||||
loginMsg.ClientID,
|
|
||||||
ctl.runID,
|
|
||||||
loginMsg.Hostname,
|
|
||||||
loginMsg.Version,
|
|
||||||
remoteAddr,
|
|
||||||
ctl.sessionCtx.WireProtocol,
|
|
||||||
uint64(entry.id),
|
|
||||||
)
|
|
||||||
if conflict {
|
|
||||||
return true, fmt.Errorf("client_id [%s] for user [%s] is already online", loginMsg.ClientID, loginMsg.User)
|
|
||||||
}
|
|
||||||
|
|
||||||
entry.registryOnline = true
|
|
||||||
entry.registryControlID = entry.id
|
|
||||||
ctl.activated = true
|
|
||||||
return true, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// completeLogin reserves ctl's current ownership with its run gate while the
|
// we should make sure if it's the same control to prevent delete a new one
|
||||||
// bounded successful LoginResp write runs, then transitions it to running.
|
func (cm *ControlManager) Del(runID string, ctl *Control) {
|
||||||
// The callback must only perform that bounded write; it must not call back into
|
|
||||||
// the control manager or the same control lifecycle.
|
|
||||||
func (cm *ControlManager) completeLogin(ctl *Control, writeSuccess func() error) (bool, error) {
|
|
||||||
entry, ok := cm.lockCurrentRun(ctl.runID, false)
|
|
||||||
if !ok {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
defer entry.runMu.Unlock()
|
|
||||||
if entry.ctl != ctl || entry.id != ctl.controlID {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
ctl.lifecycleMu.Lock()
|
|
||||||
defer ctl.lifecycleMu.Unlock()
|
|
||||||
if ctl.state != controlStatePending || !ctl.activated {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
if err := writeSuccess(); err != nil {
|
|
||||||
return false, err
|
|
||||||
}
|
|
||||||
if !ctl.startLocked() {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
return true, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Remove deletes and offlines ctl only if it is still the current generation.
|
|
||||||
func (cm *ControlManager) Remove(ctl *Control) bool {
|
|
||||||
entry, ok := cm.lockCurrentRun(ctl.runID, true)
|
|
||||||
if !ok {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
defer entry.runMu.Unlock()
|
|
||||||
cm.mu.Lock()
|
cm.mu.Lock()
|
||||||
defer cm.mu.Unlock()
|
defer cm.mu.Unlock()
|
||||||
|
if c, ok := cm.ctlsByRunID[runID]; ok && c == ctl {
|
||||||
if cm.ctlsByRunID[ctl.runID] != entry || entry.ctl != ctl || entry.id != ctl.controlID {
|
delete(cm.ctlsByRunID, runID)
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
delete(cm.ctlsByRunID, ctl.runID)
|
|
||||||
if entry.registryOnline {
|
|
||||||
cm.registry.MarkOfflineByRunIDAndControlID(ctl.runID, uint64(entry.registryControlID))
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *ControlManager) GetByID(runID string) (ctl *Control, ok bool) {
|
func (cm *ControlManager) GetByID(runID string) (ctl *Control, ok bool) {
|
||||||
entry, ok := cm.lockCurrentRun(runID, false)
|
cm.mu.RLock()
|
||||||
if !ok {
|
defer cm.mu.RUnlock()
|
||||||
return nil, false
|
ctl, ok = cm.ctlsByRunID[runID]
|
||||||
}
|
return
|
||||||
defer entry.runMu.Unlock()
|
|
||||||
ctl = entry.ctl
|
|
||||||
|
|
||||||
ctl.lifecycleMu.Lock()
|
|
||||||
defer ctl.lifecycleMu.Unlock()
|
|
||||||
if ctl.state != controlStateRunning {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
return ctl, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// admitVisitorByRunID commits a visitor admission against the current running
|
|
||||||
// control while its run and lifecycle ownership are held. The callback must
|
|
||||||
// only perform the in-memory, buffered visitor admission.
|
|
||||||
func (cm *ControlManager) admitVisitorByRunID(runID string, admit func(user, wireProtocol, udpPacketCodec string) error) (bool, error) {
|
|
||||||
entry, ok := cm.lockCurrentRun(runID, false)
|
|
||||||
if !ok {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
defer entry.runMu.Unlock()
|
|
||||||
ctl := entry.ctl
|
|
||||||
|
|
||||||
ctl.lifecycleMu.Lock()
|
|
||||||
defer ctl.lifecycleMu.Unlock()
|
|
||||||
if ctl.state != controlStateRunning {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
return true, admit(ctl.sessionCtx.LoginMsg.User, ctl.sessionCtx.WireProtocol, ctl.sessionCtx.UDPPacketCodec)
|
|
||||||
}
|
|
||||||
|
|
||||||
// RegisterWorkConn transfers conn to ctl only if ctl is still the current
|
|
||||||
// running generation. On error, ownership remains with the caller.
|
|
||||||
func (cm *ControlManager) RegisterWorkConn(ctl *Control, conn *proxy.WorkConn) error {
|
|
||||||
entry, ok := cm.lockCurrentRun(ctl.runID, false)
|
|
||||||
if !ok {
|
|
||||||
cm.mu.RLock()
|
|
||||||
closed := cm.closed
|
|
||||||
cm.mu.RUnlock()
|
|
||||||
if closed {
|
|
||||||
return fmt.Errorf("control manager is closed")
|
|
||||||
}
|
|
||||||
return fmt.Errorf("client control for run id [%s] is no longer current", ctl.runID)
|
|
||||||
}
|
|
||||||
defer entry.runMu.Unlock()
|
|
||||||
if entry.ctl != ctl || entry.id != ctl.controlID {
|
|
||||||
return fmt.Errorf("client control for run id [%s] is no longer current", ctl.runID)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctl.lifecycleMu.Lock()
|
|
||||||
defer ctl.lifecycleMu.Unlock()
|
|
||||||
if ctl.state != controlStateRunning {
|
|
||||||
return fmt.Errorf("client control for run id [%s] is not running", ctl.runID)
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case ctl.workConnCh <- conn:
|
|
||||||
ctl.xl.Debugf("new work connection registered")
|
|
||||||
return nil
|
|
||||||
default:
|
|
||||||
ctl.xl.Debugf("work connection pool is full, discarding")
|
|
||||||
return fmt.Errorf("work connection pool is full, discarding")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *ControlManager) Close() error {
|
func (cm *ControlManager) Close() error {
|
||||||
cm.mu.Lock()
|
cm.mu.Lock()
|
||||||
cm.closed = true
|
defer cm.mu.Unlock()
|
||||||
ctls := make([]*Control, 0, len(cm.ctlsByRunID))
|
for _, ctl := range cm.ctlsByRunID {
|
||||||
for _, entry := range cm.ctlsByRunID {
|
ctl.Close()
|
||||||
ctls = append(ctls, entry.ctl)
|
|
||||||
}
|
|
||||||
cm.mu.Unlock()
|
|
||||||
|
|
||||||
for _, ctl := range ctls {
|
|
||||||
cm.Remove(ctl)
|
|
||||||
_ = ctl.Close()
|
|
||||||
}
|
}
|
||||||
|
cm.ctlsByRunID = make(map[string]*Control)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -370,21 +110,12 @@ type SessionContext struct {
|
|||||||
LoginMsg *msg.Login
|
LoginMsg *msg.Login
|
||||||
// server configuration
|
// server configuration
|
||||||
ServerCfg *v1.ServerConfig
|
ServerCfg *v1.ServerConfig
|
||||||
|
// client registry
|
||||||
|
ClientRegistry *registry.ClientRegistry
|
||||||
// negotiated wire protocol for this client session
|
// negotiated wire protocol for this client session
|
||||||
WireProtocol string
|
WireProtocol string
|
||||||
UDPPacketCodec string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type controlState uint8
|
|
||||||
|
|
||||||
const (
|
|
||||||
controlStateCreated controlState = iota
|
|
||||||
controlStatePending
|
|
||||||
controlStateRunning
|
|
||||||
controlStateClosing
|
|
||||||
controlStateClosed
|
|
||||||
)
|
|
||||||
|
|
||||||
type Control struct {
|
type Control struct {
|
||||||
// session context
|
// session context
|
||||||
sessionCtx *SessionContext
|
sessionCtx *SessionContext
|
||||||
@@ -411,59 +142,30 @@ type Control struct {
|
|||||||
// last time got the Ping message
|
// last time got the Ping message
|
||||||
lastPing atomic.Value
|
lastPing atomic.Value
|
||||||
|
|
||||||
// runID never changes during the lifetime of a control. controlID is assigned
|
// A new run id will be generated when a new client login.
|
||||||
// once by ControlManager and distinguishes same-runID generations.
|
// If run id got from login message has same run id, it means it's the same client, so we can
|
||||||
runID string
|
// replace old controller instantly.
|
||||||
controlID ControlID
|
runID string
|
||||||
manager *ControlManager
|
|
||||||
|
|
||||||
lifecycleMu sync.Mutex
|
|
||||||
state controlState
|
|
||||||
activated bool
|
|
||||||
handoffBarrier <-chan struct{}
|
|
||||||
|
|
||||||
interruptOnce sync.Once
|
|
||||||
interruptErr error
|
|
||||||
|
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
|
||||||
xl *xlog.Logger
|
xl *xlog.Logger
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
doneCh chan struct{}
|
doneCh chan struct{}
|
||||||
serverMetrics metrics.ServerMetrics
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, error) {
|
func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, error) {
|
||||||
if sessionCtx.LoginMsg.PoolCount < 0 {
|
poolCount := min(sessionCtx.LoginMsg.PoolCount, int(sessionCtx.ServerCfg.Transport.MaxPoolCount))
|
||||||
return nil, fmt.Errorf("invalid pool count %d, must be non-negative", sessionCtx.LoginMsg.PoolCount)
|
|
||||||
}
|
|
||||||
if sessionCtx.ServerCfg.Transport.MaxPoolCount < 0 {
|
|
||||||
return nil, fmt.Errorf(
|
|
||||||
"invalid max pool count %d, must be non-negative",
|
|
||||||
sessionCtx.ServerCfg.Transport.MaxPoolCount,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
effectivePoolCount := min(int64(sessionCtx.LoginMsg.PoolCount), sessionCtx.ServerCfg.Transport.MaxPoolCount)
|
|
||||||
maxPoolCountForChannel := int64(math.MaxInt) - int64(workConnPoolCapacityOffset)
|
|
||||||
if effectivePoolCount > maxPoolCountForChannel {
|
|
||||||
return nil, fmt.Errorf(
|
|
||||||
"invalid effective pool count %d, cannot safely add %d for work connection pool capacity",
|
|
||||||
effectivePoolCount, workConnPoolCapacityOffset,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
poolCount := int(effectivePoolCount)
|
|
||||||
ctl := &Control{
|
ctl := &Control{
|
||||||
sessionCtx: sessionCtx,
|
sessionCtx: sessionCtx,
|
||||||
workConnCh: make(chan *proxy.WorkConn, poolCount+workConnPoolCapacityOffset),
|
workConnCh: make(chan *proxy.WorkConn, poolCount+10),
|
||||||
proxies: make(map[string]proxy.Proxy),
|
proxies: make(map[string]proxy.Proxy),
|
||||||
poolCount: poolCount,
|
poolCount: poolCount,
|
||||||
portsUsedNum: 0,
|
portsUsedNum: 0,
|
||||||
runID: sessionCtx.LoginMsg.RunID,
|
runID: sessionCtx.LoginMsg.RunID,
|
||||||
state: controlStateCreated,
|
xl: xlog.FromContextSafe(ctx),
|
||||||
xl: xlog.FromContextSafe(ctx),
|
ctx: ctx,
|
||||||
ctx: ctx,
|
doneCh: make(chan struct{}),
|
||||||
doneCh: make(chan struct{}),
|
|
||||||
serverMetrics: metrics.Server,
|
|
||||||
}
|
}
|
||||||
ctl.lastPing.Store(time.Now())
|
ctl.lastPing.Store(time.Now())
|
||||||
|
|
||||||
@@ -473,121 +175,48 @@ func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, erro
|
|||||||
return ctl, nil
|
return ctl, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ctl *Control) RunID() string {
|
|
||||||
return ctl.runID
|
|
||||||
}
|
|
||||||
|
|
||||||
func (ctl *Control) ID() ControlID {
|
|
||||||
ctl.lifecycleMu.Lock()
|
|
||||||
defer ctl.lifecycleMu.Unlock()
|
|
||||||
return ctl.controlID
|
|
||||||
}
|
|
||||||
|
|
||||||
func (ctl *Control) admit(manager *ControlManager, id ControlID) error {
|
|
||||||
ctl.lifecycleMu.Lock()
|
|
||||||
defer ctl.lifecycleMu.Unlock()
|
|
||||||
if ctl.state != controlStateCreated {
|
|
||||||
return fmt.Errorf("control [%s] is not in created state", ctl.runID)
|
|
||||||
}
|
|
||||||
ctl.manager = manager
|
|
||||||
ctl.controlID = id
|
|
||||||
ctl.state = controlStatePending
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (ctl *Control) setHandoffBarrier(barrier <-chan struct{}) {
|
|
||||||
ctl.lifecycleMu.Lock()
|
|
||||||
ctl.handoffBarrier = barrier
|
|
||||||
ctl.lifecycleMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (ctl *Control) WaitForHandoff() {
|
|
||||||
ctl.lifecycleMu.Lock()
|
|
||||||
barrier := ctl.handoffBarrier
|
|
||||||
ctl.lifecycleMu.Unlock()
|
|
||||||
if barrier != nil {
|
|
||||||
<-barrier
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Start starts the control session workers after login succeeds.
|
// Start starts the control session workers after login succeeds.
|
||||||
func (ctl *Control) Start() bool {
|
func (ctl *Control) Start() {
|
||||||
ctl.lifecycleMu.Lock()
|
go func() {
|
||||||
defer ctl.lifecycleMu.Unlock()
|
for i := 0; i < ctl.poolCount; i++ {
|
||||||
return ctl.startLocked()
|
// ignore error here, that means that this control is closed
|
||||||
}
|
_ = ctl.msgDispatcher.Send(&msg.ReqWorkConn{})
|
||||||
|
}
|
||||||
func (ctl *Control) startLocked() bool {
|
}()
|
||||||
if ctl.state != controlStatePending || !ctl.activated {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
ctl.state = controlStateRunning
|
|
||||||
go ctl.worker()
|
go ctl.worker()
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ctl *Control) Close() error {
|
func (ctl *Control) Close() error {
|
||||||
ctl.lifecycleMu.Lock()
|
ctl.sessionCtx.Conn.Close()
|
||||||
switch ctl.state {
|
return nil
|
||||||
case controlStateCreated, controlStatePending:
|
|
||||||
ctl.state = controlStateClosing
|
|
||||||
ctl.finishLocked()
|
|
||||||
case controlStateRunning:
|
|
||||||
ctl.state = controlStateClosing
|
|
||||||
}
|
|
||||||
ctl.lifecycleMu.Unlock()
|
|
||||||
return ctl.interruptReadAndClose()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ctl *Control) Replaced(newCtl *Control) {
|
func (ctl *Control) Replaced(newCtl *Control) {
|
||||||
ctl.markReplaced()
|
xl := ctl.xl
|
||||||
ctl.xl.Infof("replaced by client [%s] (control ID %d)", newCtl.runID, newCtl.ID())
|
xl.Infof("replaced by client [%s]", newCtl.runID)
|
||||||
_ = ctl.interruptReadAndClose()
|
ctl.runID = ""
|
||||||
|
ctl.sessionCtx.Conn.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
// markReplaced returns the transitive predecessor barrier. A pending control
|
func (ctl *Control) RegisterWorkConn(conn *proxy.WorkConn) error {
|
||||||
// has no worker, so it finishes immediately and passes its inherited barrier
|
xl := ctl.xl
|
||||||
// to the replacement. A running control is finished only by its worker.
|
defer func() {
|
||||||
func (ctl *Control) markReplaced() <-chan struct{} {
|
if err := recover(); err != nil {
|
||||||
ctl.lifecycleMu.Lock()
|
xl.Errorf("panic error: %v", err)
|
||||||
defer ctl.lifecycleMu.Unlock()
|
xl.Errorf(string(debug.Stack()))
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
switch ctl.state {
|
select {
|
||||||
case controlStateCreated:
|
case ctl.workConnCh <- conn:
|
||||||
ctl.state = controlStateClosing
|
xl.Debugf("new work connection registered")
|
||||||
ctl.finishLocked()
|
|
||||||
return nil
|
return nil
|
||||||
case controlStatePending:
|
|
||||||
barrier := ctl.handoffBarrier
|
|
||||||
ctl.state = controlStateClosing
|
|
||||||
ctl.finishLocked()
|
|
||||||
return barrier
|
|
||||||
case controlStateRunning:
|
|
||||||
ctl.state = controlStateClosing
|
|
||||||
return ctl.doneCh
|
|
||||||
case controlStateClosing, controlStateClosed:
|
|
||||||
return ctl.doneCh
|
|
||||||
default:
|
default:
|
||||||
return ctl.doneCh
|
xl.Debugf("work connection pool is full, discarding")
|
||||||
|
return fmt.Errorf("work connection pool is full, discarding")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ctl *Control) interruptReadAndClose() error {
|
|
||||||
ctl.interruptOnce.Do(func() {
|
|
||||||
_ = ctl.sessionCtx.Conn.SetReadDeadline(time.Now())
|
|
||||||
ctl.interruptErr = ctl.sessionCtx.Conn.Close()
|
|
||||||
})
|
|
||||||
return ctl.interruptErr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (ctl *Control) finishLocked() {
|
|
||||||
if ctl.state == controlStateClosed {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
ctl.state = controlStateClosed
|
|
||||||
close(ctl.doneCh)
|
|
||||||
}
|
|
||||||
|
|
||||||
// When frps get one user connection, we get one work connection from the pool and return it.
|
// When frps get one user connection, we get one work connection from the pool and return it.
|
||||||
// If no workConn available in the pool, send message to frpc to get one or more
|
// If no workConn available in the pool, send message to frpc to get one or more
|
||||||
// and wait until it is available.
|
// and wait until it is available.
|
||||||
@@ -642,10 +271,10 @@ func (ctl *Control) heartbeatWorker() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
xl := ctl.xl
|
xl := ctl.xl
|
||||||
wait.Until(func() {
|
go wait.Until(func() {
|
||||||
if time.Since(ctl.lastPing.Load().(time.Time)) > time.Duration(ctl.sessionCtx.ServerCfg.Transport.HeartbeatTimeout)*time.Second {
|
if time.Since(ctl.lastPing.Load().(time.Time)) > time.Duration(ctl.sessionCtx.ServerCfg.Transport.HeartbeatTimeout)*time.Second {
|
||||||
xl.Warnf("heartbeat timeout")
|
xl.Warnf("heartbeat timeout")
|
||||||
_ = ctl.Close()
|
ctl.sessionCtx.Conn.Close()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}, time.Second, ctl.doneCh)
|
}, time.Second, ctl.doneCh)
|
||||||
@@ -660,14 +289,14 @@ func (ctl *Control) loginUserInfo() plugin.UserInfo {
|
|||||||
return plugin.UserInfo{
|
return plugin.UserInfo{
|
||||||
User: ctl.sessionCtx.LoginMsg.User,
|
User: ctl.sessionCtx.LoginMsg.User,
|
||||||
Metas: ctl.sessionCtx.LoginMsg.Metas,
|
Metas: ctl.sessionCtx.LoginMsg.Metas,
|
||||||
RunID: ctl.runID,
|
RunID: ctl.sessionCtx.LoginMsg.RunID,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ctl *Control) closeProxy(pxy proxy.Proxy) {
|
func (ctl *Control) closeProxy(pxy proxy.Proxy) {
|
||||||
pxy.Close()
|
pxy.Close()
|
||||||
ctl.sessionCtx.PxyManager.Del(pxy.GetName())
|
ctl.sessionCtx.PxyManager.Del(pxy.GetName())
|
||||||
ctl.serverMetrics.CloseProxy(pxy.GetName(), pxy.GetConfigurer().GetBaseConfig().Type)
|
metrics.Server.CloseProxy(pxy.GetName(), pxy.GetConfigurer().GetBaseConfig().Type)
|
||||||
|
|
||||||
notifyContent := &plugin.CloseProxyContent{
|
notifyContent := &plugin.CloseProxyContent{
|
||||||
User: ctl.loginUserInfo(),
|
User: ctl.loginUserInfo(),
|
||||||
@@ -682,24 +311,12 @@ func (ctl *Control) closeProxy(pxy proxy.Proxy) {
|
|||||||
|
|
||||||
func (ctl *Control) worker() {
|
func (ctl *Control) worker() {
|
||||||
xl := ctl.xl
|
xl := ctl.xl
|
||||||
ctl.serverMetrics.NewClient()
|
|
||||||
|
|
||||||
go ctl.heartbeatWorker()
|
go ctl.heartbeatWorker()
|
||||||
go ctl.msgDispatcher.Run()
|
go ctl.msgDispatcher.Run()
|
||||||
go func() {
|
|
||||||
for i := 0; i < ctl.poolCount; i++ {
|
|
||||||
// Ignore the error: it means this control is already closing.
|
|
||||||
_ = ctl.msgDispatcher.Send(&msg.ReqWorkConn{})
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
<-ctl.msgDispatcher.Done()
|
<-ctl.msgDispatcher.Done()
|
||||||
ctl.lifecycleMu.Lock()
|
ctl.sessionCtx.Conn.Close()
|
||||||
if ctl.state == controlStateRunning {
|
|
||||||
ctl.state = controlStateClosing
|
|
||||||
}
|
|
||||||
ctl.lifecycleMu.Unlock()
|
|
||||||
_ = ctl.interruptReadAndClose()
|
|
||||||
|
|
||||||
ctl.mu.Lock()
|
ctl.mu.Lock()
|
||||||
close(ctl.workConnCh)
|
close(ctl.workConnCh)
|
||||||
@@ -714,14 +331,10 @@ func (ctl *Control) worker() {
|
|||||||
ctl.closeProxy(pxy)
|
ctl.closeProxy(pxy)
|
||||||
}
|
}
|
||||||
|
|
||||||
ctl.serverMetrics.CloseClient()
|
metrics.Server.CloseClient()
|
||||||
if ctl.manager != nil {
|
ctl.sessionCtx.ClientRegistry.MarkOfflineByRunID(ctl.runID)
|
||||||
ctl.manager.Remove(ctl)
|
|
||||||
}
|
|
||||||
xl.Infof("client exit success")
|
xl.Infof("client exit success")
|
||||||
ctl.lifecycleMu.Lock()
|
close(ctl.doneCh)
|
||||||
ctl.finishLocked()
|
|
||||||
ctl.lifecycleMu.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ctl *Control) registerMsgHandlers() {
|
func (ctl *Control) registerMsgHandlers() {
|
||||||
@@ -761,9 +374,9 @@ func (ctl *Control) handleNewProxy(m msg.Message) {
|
|||||||
xl.Infof("new proxy [%s] type [%s] success", inMsg.ProxyName, inMsg.ProxyType)
|
xl.Infof("new proxy [%s] type [%s] success", inMsg.ProxyName, inMsg.ProxyType)
|
||||||
clientID := ctl.sessionCtx.LoginMsg.ClientID
|
clientID := ctl.sessionCtx.LoginMsg.ClientID
|
||||||
if clientID == "" {
|
if clientID == "" {
|
||||||
clientID = ctl.runID
|
clientID = ctl.sessionCtx.LoginMsg.RunID
|
||||||
}
|
}
|
||||||
ctl.serverMetrics.NewProxy(inMsg.ProxyName, inMsg.ProxyType, ctl.sessionCtx.LoginMsg.User, clientID)
|
metrics.Server.NewProxy(inMsg.ProxyName, inMsg.ProxyType, ctl.sessionCtx.LoginMsg.User, clientID)
|
||||||
}
|
}
|
||||||
_ = ctl.msgDispatcher.Send(resp)
|
_ = ctl.msgDispatcher.Send(resp)
|
||||||
}
|
}
|
||||||
@@ -842,7 +455,6 @@ func (ctl *Control) RegisterProxy(pxyMsg *msg.NewProxy) (remoteAddr string, err
|
|||||||
ServerCfg: ctl.sessionCtx.ServerCfg,
|
ServerCfg: ctl.sessionCtx.ServerCfg,
|
||||||
EncryptionKey: ctl.sessionCtx.EncryptionKey,
|
EncryptionKey: ctl.sessionCtx.EncryptionKey,
|
||||||
WireProtocol: ctl.sessionCtx.WireProtocol,
|
WireProtocol: ctl.sessionCtx.WireProtocol,
|
||||||
UDPPacketCodec: ctl.sessionCtx.UDPPacketCodec,
|
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return remoteAddr, err
|
return remoteAddr, err
|
||||||
|
|||||||
@@ -1,595 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package server
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"math"
|
|
||||||
"net"
|
|
||||||
"os"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/auth"
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
|
||||||
plugin "github.com/fatedier/frp/pkg/plugin/server"
|
|
||||||
"github.com/fatedier/frp/server/controller"
|
|
||||||
"github.com/fatedier/frp/server/proxy"
|
|
||||||
"github.com/fatedier/frp/server/registry"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestControlPendingReplacementFinishesWithoutStarting(t *testing.T) {
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
manager := NewControlManager(clientRegistry)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
oldCtl, oldConn := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
newCtl, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
|
|
||||||
mustAddAndActivate(t, manager, oldCtl)
|
|
||||||
|
|
||||||
err := manager.Add(newCtl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
waitForControlDone(t, oldCtl)
|
|
||||||
require.False(t, oldCtl.Start())
|
|
||||||
require.Equal(t, []string{"deadline", "close"}, oldConn.eventsSnapshot())
|
|
||||||
require.Equal(t, int64(0), metrics.newClients())
|
|
||||||
require.Equal(t, int64(0), metrics.closedClients())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewControlPoolCountBoundaries(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
poolCount int
|
|
||||||
maxPoolCount int64
|
|
||||||
wantErr string
|
|
||||||
wantPoolCount int
|
|
||||||
wantCapacity int
|
|
||||||
}{
|
|
||||||
{name: "negative pool count below offset", poolCount: -11, maxPoolCount: 5, wantErr: "invalid pool count"},
|
|
||||||
{name: "negative pool count at offset", poolCount: -10, maxPoolCount: 5, wantErr: "invalid pool count"},
|
|
||||||
{name: "negative pool count", poolCount: -1, maxPoolCount: 5, wantErr: "invalid pool count"},
|
|
||||||
{name: "zero pool count", poolCount: 0, maxPoolCount: 5, wantPoolCount: 0, wantCapacity: 10},
|
|
||||||
{name: "pool count capped", poolCount: 10, maxPoolCount: 5, wantPoolCount: 5, wantCapacity: 15},
|
|
||||||
{name: "maximum int pool count capped", poolCount: math.MaxInt, maxPoolCount: 5, wantPoolCount: 5, wantCapacity: 15},
|
|
||||||
{name: "negative maximum", poolCount: 1, maxPoolCount: -1, wantErr: "invalid max pool count"},
|
|
||||||
{name: "maximum int64 with small client pool", poolCount: 1, maxPoolCount: math.MaxInt64, wantPoolCount: 1, wantCapacity: 11},
|
|
||||||
{name: "maximum int client and server overflow", poolCount: math.MaxInt, maxPoolCount: math.MaxInt64, wantErr: "cannot safely add"},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
conn := newDeadlineReadConn()
|
|
||||||
msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn))
|
|
||||||
cfg := &v1.ServerConfig{}
|
|
||||||
cfg.Transport.MaxPoolCount = tc.maxPoolCount
|
|
||||||
|
|
||||||
ctl, err := NewControl(context.Background(), &SessionContext{
|
|
||||||
RC: &controller.ResourceController{},
|
|
||||||
PxyManager: proxy.NewManager(),
|
|
||||||
PluginManager: plugin.NewManager(),
|
|
||||||
AuthVerifier: auth.AlwaysPassVerifier,
|
|
||||||
Conn: msgConn,
|
|
||||||
LoginMsg: &msg.Login{
|
|
||||||
RunID: "pool-count-run",
|
|
||||||
PoolCount: tc.poolCount,
|
|
||||||
},
|
|
||||||
ServerCfg: cfg,
|
|
||||||
})
|
|
||||||
if tc.wantErr != "" {
|
|
||||||
require.Nil(t, ctl)
|
|
||||||
require.ErrorContains(t, err, tc.wantErr)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, tc.wantPoolCount, ctl.poolCount)
|
|
||||||
require.Equal(t, tc.wantCapacity, cap(ctl.workConnCh))
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlRunningReplacementFinishesInWorker(t *testing.T) {
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
manager := NewControlManager(clientRegistry)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
oldCtl, oldConn := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
newCtl, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
|
|
||||||
mustAddAndActivate(t, manager, oldCtl)
|
|
||||||
require.True(t, oldCtl.Start())
|
|
||||||
waitForSignal(t, oldConn.readStarted, "control reader to start")
|
|
||||||
|
|
||||||
err := manager.Add(newCtl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
waitForControlDone(t, oldCtl)
|
|
||||||
require.Equal(t, []string{"deadline", "close"}, oldConn.eventsSnapshot())
|
|
||||||
require.Equal(t, int64(1), metrics.newClients())
|
|
||||||
require.Equal(t, int64(1), metrics.closedClients())
|
|
||||||
|
|
||||||
_, ok := manager.GetByID("same-run")
|
|
||||||
require.False(t, ok)
|
|
||||||
require.Same(t, newCtl, currentControlForTest(manager, "same-run"))
|
|
||||||
info, ok := clientRegistry.GetByKey("client")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.True(t, info.Online)
|
|
||||||
require.Equal(t, uint64(oldCtl.ID()), info.ControlID)
|
|
||||||
|
|
||||||
active, err := manager.Activate(newCtl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.True(t, active)
|
|
||||||
_, ok = manager.GetByID("same-run")
|
|
||||||
require.False(t, ok)
|
|
||||||
info, ok = clientRegistry.GetByKey("client")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.Equal(t, uint64(newCtl.ID()), info.ControlID)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlClosePendingAndRunning(t *testing.T) {
|
|
||||||
t.Run("pending", func(t *testing.T) {
|
|
||||||
manager := NewControlManager(registry.NewClientRegistry())
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
ctl, conn := newLifecycleTestControl(t, "pending", "pending", metrics)
|
|
||||||
err := manager.Add(ctl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
require.Equal(t, []string{"deadline", "close"}, conn.eventsSnapshot())
|
|
||||||
require.Equal(t, int64(0), metrics.newClients())
|
|
||||||
require.Equal(t, int64(0), metrics.closedClients())
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("running", func(t *testing.T) {
|
|
||||||
manager := NewControlManager(registry.NewClientRegistry())
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
ctl, conn := newLifecycleTestControl(t, "running", "running", metrics)
|
|
||||||
mustAddAndActivate(t, manager, ctl)
|
|
||||||
require.True(t, ctl.Start())
|
|
||||||
waitForSignal(t, conn.readStarted, "control reader to start")
|
|
||||||
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
require.Equal(t, []string{"deadline", "close"}, conn.eventsSnapshot())
|
|
||||||
require.Equal(t, int64(1), metrics.newClients())
|
|
||||||
require.Equal(t, int64(1), metrics.closedClients())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlCloseAndReplacedAreIdempotent(t *testing.T) {
|
|
||||||
manager := NewControlManager(registry.NewClientRegistry())
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
ctl, conn := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
replacement, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
|
|
||||||
err := manager.Add(ctl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
err = manager.Add(replacement)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
ctl.Replaced(replacement)
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
|
|
||||||
require.Equal(t, []string{"deadline", "close"}, conn.eventsSnapshot())
|
|
||||||
require.Equal(t, int64(0), metrics.newClients())
|
|
||||||
require.Equal(t, int64(0), metrics.closedClients())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlHeartbeatTimeoutInterruptsRead(t *testing.T) {
|
|
||||||
manager := NewControlManager(registry.NewClientRegistry())
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
ctl, conn := newLifecycleTestControl(t, "heartbeat", "heartbeat", metrics)
|
|
||||||
ctl.sessionCtx.ServerCfg.Transport.HeartbeatTimeout = 1
|
|
||||||
ctl.lastPing.Store(time.Now().Add(-2 * time.Second))
|
|
||||||
|
|
||||||
mustAddAndActivate(t, manager, ctl)
|
|
||||||
require.True(t, ctl.Start())
|
|
||||||
waitForSignal(t, conn.readStarted, "control reader to start")
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
|
|
||||||
require.Equal(t, []string{"deadline", "close"}, conn.eventsSnapshot())
|
|
||||||
require.Equal(t, int64(1), metrics.newClients())
|
|
||||||
require.Equal(t, int64(1), metrics.closedClients())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlStartReplacementRacePairsMetrics(t *testing.T) {
|
|
||||||
for range 100 {
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
manager := NewControlManager(clientRegistry)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
ctl, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
replacement, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
|
|
||||||
mustAddAndActivate(t, manager, ctl)
|
|
||||||
|
|
||||||
startGate := make(chan struct{})
|
|
||||||
startedCh := make(chan bool, 1)
|
|
||||||
addErrCh := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
<-startGate
|
|
||||||
startedCh <- ctl.Start()
|
|
||||||
}()
|
|
||||||
go func() {
|
|
||||||
<-startGate
|
|
||||||
addErr := manager.Add(replacement)
|
|
||||||
addErrCh <- addErr
|
|
||||||
}()
|
|
||||||
close(startGate)
|
|
||||||
|
|
||||||
started := <-startedCh
|
|
||||||
require.NoError(t, <-addErrCh)
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
if started {
|
|
||||||
require.Equal(t, int64(1), metrics.newClients())
|
|
||||||
require.Equal(t, int64(1), metrics.closedClients())
|
|
||||||
} else {
|
|
||||||
require.Equal(t, int64(0), metrics.newClients())
|
|
||||||
require.Equal(t, int64(0), metrics.closedClients())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlManagerRejectsStaleActivateAndRemove(t *testing.T) {
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
manager := NewControlManager(clientRegistry)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
oldCtl, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
newCtl, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
|
|
||||||
mustAddAndActivate(t, manager, oldCtl)
|
|
||||||
err := manager.Add(newCtl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Greater(t, uint64(newCtl.ID()), uint64(oldCtl.ID()))
|
|
||||||
|
|
||||||
active, err := manager.Activate(oldCtl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.False(t, active)
|
|
||||||
require.False(t, manager.Remove(oldCtl))
|
|
||||||
|
|
||||||
_, ok := manager.GetByID("same-run")
|
|
||||||
require.False(t, ok)
|
|
||||||
require.Same(t, newCtl, currentControlForTest(manager, "same-run"))
|
|
||||||
info, ok := clientRegistry.GetByKey("client")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.True(t, info.Online)
|
|
||||||
require.Equal(t, uint64(oldCtl.ID()), info.ControlID)
|
|
||||||
|
|
||||||
active, err = manager.Activate(newCtl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.True(t, active)
|
|
||||||
info, ok = clientRegistry.GetByKey("client")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.True(t, info.Online)
|
|
||||||
require.Equal(t, uint64(newCtl.ID()), info.ControlID)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlManagerPreservesClientIDConflict(t *testing.T) {
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
manager := NewControlManager(clientRegistry)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
first, _ := newLifecycleTestControl(t, "run-one", "shared-client", metrics)
|
|
||||||
conflicting, _ := newLifecycleTestControl(t, "run-two", "shared-client", metrics)
|
|
||||||
|
|
||||||
mustAddAndActivate(t, manager, first)
|
|
||||||
err := manager.Add(conflicting)
|
|
||||||
require.NoError(t, err)
|
|
||||||
active, err := manager.Activate(conflicting)
|
|
||||||
require.True(t, active)
|
|
||||||
require.ErrorContains(t, err, "already online")
|
|
||||||
|
|
||||||
require.True(t, manager.Remove(conflicting))
|
|
||||||
info, ok := clientRegistry.GetByKey("shared-client")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.True(t, info.Online)
|
|
||||||
require.Equal(t, "run-one", info.RunID)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlManagerFailedLoginWriteReleasesRunWithoutStarting(t *testing.T) {
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
manager := NewControlManager(clientRegistry)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
ctl, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
replacement, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
|
|
||||||
mustAddAndActivate(t, manager, ctl)
|
|
||||||
|
|
||||||
writeErr := errors.New("write failed")
|
|
||||||
committed, err := manager.completeLogin(ctl, func() error { return writeErr })
|
|
||||||
require.ErrorIs(t, err, writeErr)
|
|
||||||
require.False(t, committed)
|
|
||||||
|
|
||||||
err = manager.Add(replacement)
|
|
||||||
require.NoError(t, err)
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
require.Same(t, replacement, currentControlForTest(manager, "same-run"))
|
|
||||||
require.Equal(t, int64(0), metrics.newClients())
|
|
||||||
require.Equal(t, int64(0), metrics.closedClients())
|
|
||||||
require.True(t, manager.Remove(replacement))
|
|
||||||
info, ok := clientRegistry.GetByKey("client")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.False(t, info.Online)
|
|
||||||
require.Empty(t, info.RunID)
|
|
||||||
require.Zero(t, info.ControlID)
|
|
||||||
require.False(t, info.DisconnectedAt.IsZero())
|
|
||||||
require.NoError(t, replacement.Close())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControlManagerCloseWaitsForInFlightLoginRun(t *testing.T) {
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
manager := NewControlManager(clientRegistry)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
ctl, _ := newLifecycleTestControl(t, "same-run", "client", metrics)
|
|
||||||
|
|
||||||
mustAddAndActivate(t, manager, ctl)
|
|
||||||
|
|
||||||
writeEntered := make(chan struct{})
|
|
||||||
resumeWrite := make(chan struct{})
|
|
||||||
loginDone := make(chan struct {
|
|
||||||
committed bool
|
|
||||||
err error
|
|
||||||
}, 1)
|
|
||||||
go func() {
|
|
||||||
committed, loginErr := manager.completeLogin(ctl, func() error {
|
|
||||||
close(writeEntered)
|
|
||||||
<-resumeWrite
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
loginDone <- struct {
|
|
||||||
committed bool
|
|
||||||
err error
|
|
||||||
}{committed: committed, err: loginErr}
|
|
||||||
}()
|
|
||||||
waitForSignal(t, writeEntered, "LoginResp write")
|
|
||||||
|
|
||||||
closeDone := make(chan error, 1)
|
|
||||||
go func() { closeDone <- manager.Close() }()
|
|
||||||
waitForManagerClosed(t, manager)
|
|
||||||
select {
|
|
||||||
case err := <-closeDone:
|
|
||||||
t.Fatalf("manager close completed during LoginResp write: %v", err)
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
|
|
||||||
close(resumeWrite)
|
|
||||||
result := <-loginDone
|
|
||||||
require.NoError(t, result.err)
|
|
||||||
require.True(t, result.committed)
|
|
||||||
require.NoError(t, <-closeDone)
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
require.Nil(t, currentControlForTest(manager, "same-run"))
|
|
||||||
require.Equal(t, int64(1), metrics.newClients())
|
|
||||||
require.Equal(t, int64(1), metrics.closedClients())
|
|
||||||
info, ok := clientRegistry.GetByKey("client")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.False(t, info.Online)
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLifecycleTestControl(
|
|
||||||
t *testing.T,
|
|
||||||
runID string,
|
|
||||||
clientID string,
|
|
||||||
serverMetrics *countingServerMetrics,
|
|
||||||
) (*Control, *deadlineReadConn) {
|
|
||||||
t.Helper()
|
|
||||||
conn := newDeadlineReadConn()
|
|
||||||
msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn))
|
|
||||||
ctl, err := NewControl(context.Background(), &SessionContext{
|
|
||||||
RC: &controller.ResourceController{},
|
|
||||||
PxyManager: proxy.NewManager(),
|
|
||||||
PluginManager: plugin.NewManager(),
|
|
||||||
AuthVerifier: auth.AlwaysPassVerifier,
|
|
||||||
Conn: msgConn,
|
|
||||||
LoginMsg: &msg.Login{
|
|
||||||
RunID: runID,
|
|
||||||
ClientID: clientID,
|
|
||||||
},
|
|
||||||
ServerCfg: &v1.ServerConfig{},
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
ctl.serverMetrics = serverMetrics
|
|
||||||
t.Cleanup(func() { _ = ctl.Close() })
|
|
||||||
return ctl, conn
|
|
||||||
}
|
|
||||||
|
|
||||||
func mustAddAndActivate(t *testing.T, manager *ControlManager, ctl *Control) {
|
|
||||||
t.Helper()
|
|
||||||
require.NoError(t, manager.Add(ctl))
|
|
||||||
active, err := manager.Activate(ctl)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.True(t, active)
|
|
||||||
}
|
|
||||||
|
|
||||||
func waitForControlDone(t *testing.T, ctl *Control) {
|
|
||||||
t.Helper()
|
|
||||||
done := make(chan struct{})
|
|
||||||
go func() {
|
|
||||||
ctl.WaitClosed()
|
|
||||||
close(done)
|
|
||||||
}()
|
|
||||||
waitForSignal(t, done, "control to finish")
|
|
||||||
}
|
|
||||||
|
|
||||||
func currentControlForTest(manager *ControlManager, runID string) *Control {
|
|
||||||
manager.mu.RLock()
|
|
||||||
defer manager.mu.RUnlock()
|
|
||||||
entry := manager.ctlsByRunID[runID]
|
|
||||||
if entry == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return entry.ctl
|
|
||||||
}
|
|
||||||
|
|
||||||
func currentRunGateForTest(manager *ControlManager, runID string) *sync.Mutex {
|
|
||||||
manager.mu.RLock()
|
|
||||||
defer manager.mu.RUnlock()
|
|
||||||
entry := manager.ctlsByRunID[runID]
|
|
||||||
if entry == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return entry.runMu
|
|
||||||
}
|
|
||||||
|
|
||||||
func waitForManagerClosed(t *testing.T, manager *ControlManager) {
|
|
||||||
t.Helper()
|
|
||||||
deadline := time.Now().Add(3 * time.Second)
|
|
||||||
for time.Now().Before(deadline) {
|
|
||||||
manager.mu.RLock()
|
|
||||||
closed := manager.closed
|
|
||||||
manager.mu.RUnlock()
|
|
||||||
if closed {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
t.Fatal("timed out waiting for control manager to close")
|
|
||||||
}
|
|
||||||
|
|
||||||
func waitForSignal(t *testing.T, ch <-chan struct{}, description string) {
|
|
||||||
t.Helper()
|
|
||||||
select {
|
|
||||||
case <-ch:
|
|
||||||
case <-time.After(3 * time.Second):
|
|
||||||
t.Fatalf("timed out waiting for %s", description)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type deadlineReadConn struct {
|
|
||||||
readStarted chan struct{}
|
|
||||||
unblockRead chan struct{}
|
|
||||||
|
|
||||||
readOnce sync.Once
|
|
||||||
unblockOnce sync.Once
|
|
||||||
deadlineOnce sync.Once
|
|
||||||
closeOnce sync.Once
|
|
||||||
|
|
||||||
eventsMu sync.Mutex
|
|
||||||
events []string
|
|
||||||
}
|
|
||||||
|
|
||||||
func newDeadlineReadConn() *deadlineReadConn {
|
|
||||||
return &deadlineReadConn{
|
|
||||||
readStarted: make(chan struct{}),
|
|
||||||
unblockRead: make(chan struct{}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *deadlineReadConn) Read([]byte) (int, error) {
|
|
||||||
c.readOnce.Do(func() { close(c.readStarted) })
|
|
||||||
<-c.unblockRead
|
|
||||||
return 0, os.ErrDeadlineExceeded
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*deadlineReadConn) Write(p []byte) (int, error) { return len(p), nil }
|
|
||||||
|
|
||||||
func (c *deadlineReadConn) Close() error {
|
|
||||||
c.closeOnce.Do(func() {
|
|
||||||
c.recordEvent("close")
|
|
||||||
c.unblockOnce.Do(func() { close(c.unblockRead) })
|
|
||||||
})
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*deadlineReadConn) LocalAddr() net.Addr { return lifecycleTestAddr("local") }
|
|
||||||
func (*deadlineReadConn) RemoteAddr() net.Addr { return lifecycleTestAddr("remote") }
|
|
||||||
|
|
||||||
func (c *deadlineReadConn) SetDeadline(deadline time.Time) error {
|
|
||||||
if err := c.SetReadDeadline(deadline); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return c.SetWriteDeadline(deadline)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *deadlineReadConn) SetReadDeadline(deadline time.Time) error {
|
|
||||||
if deadline.IsZero() {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
c.deadlineOnce.Do(func() {
|
|
||||||
c.recordEvent("deadline")
|
|
||||||
c.unblockOnce.Do(func() { close(c.unblockRead) })
|
|
||||||
})
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*deadlineReadConn) SetWriteDeadline(time.Time) error { return nil }
|
|
||||||
|
|
||||||
func (c *deadlineReadConn) recordEvent(event string) {
|
|
||||||
c.eventsMu.Lock()
|
|
||||||
c.events = append(c.events, event)
|
|
||||||
c.eventsMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *deadlineReadConn) eventsSnapshot() []string {
|
|
||||||
c.eventsMu.Lock()
|
|
||||||
defer c.eventsMu.Unlock()
|
|
||||||
return append([]string(nil), c.events...)
|
|
||||||
}
|
|
||||||
|
|
||||||
type lifecycleTestAddr string
|
|
||||||
|
|
||||||
func (a lifecycleTestAddr) Network() string { return string(a) }
|
|
||||||
func (a lifecycleTestAddr) String() string { return string(a) }
|
|
||||||
|
|
||||||
type countingServerMetrics struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
newCount int64
|
|
||||||
closeCount int64
|
|
||||||
closeEnter chan struct{}
|
|
||||||
closeResume chan struct{}
|
|
||||||
closeOnce sync.Once
|
|
||||||
}
|
|
||||||
|
|
||||||
func newCountingServerMetrics() *countingServerMetrics {
|
|
||||||
return &countingServerMetrics{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *countingServerMetrics) NewClient() {
|
|
||||||
m.mu.Lock()
|
|
||||||
m.newCount++
|
|
||||||
m.mu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *countingServerMetrics) CloseClient() {
|
|
||||||
m.mu.Lock()
|
|
||||||
m.closeCount++
|
|
||||||
closeEnter := m.closeEnter
|
|
||||||
closeResume := m.closeResume
|
|
||||||
m.mu.Unlock()
|
|
||||||
if closeEnter != nil {
|
|
||||||
m.closeOnce.Do(func() { close(closeEnter) })
|
|
||||||
<-closeResume
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*countingServerMetrics) NewProxy(string, string, string, string) {}
|
|
||||||
func (*countingServerMetrics) CloseProxy(string, string) {}
|
|
||||||
func (*countingServerMetrics) OpenConnection(string, string) {}
|
|
||||||
func (*countingServerMetrics) CloseConnection(string, string) {}
|
|
||||||
func (*countingServerMetrics) AddTrafficIn(string, string, int64) {}
|
|
||||||
func (*countingServerMetrics) AddTrafficOut(string, string, int64) {}
|
|
||||||
|
|
||||||
func (m *countingServerMetrics) newClients() int64 {
|
|
||||||
m.mu.Lock()
|
|
||||||
defer m.mu.Unlock()
|
|
||||||
return m.newCount
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *countingServerMetrics) closedClients() int64 {
|
|
||||||
m.mu.Lock()
|
|
||||||
defer m.mu.Unlock()
|
|
||||||
return m.closeCount
|
|
||||||
}
|
|
||||||
@@ -58,12 +58,8 @@ func NewController(
|
|||||||
|
|
||||||
// /api/serverinfo
|
// /api/serverinfo
|
||||||
func (c *Controller) APIServerInfo(ctx *httppkg.Context) (any, error) {
|
func (c *Controller) APIServerInfo(ctx *httppkg.Context) (any, error) {
|
||||||
return c.buildServerInfoResp(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Controller) buildServerInfoResp() model.ServerInfoResp {
|
|
||||||
serverStats := mem.StatsCollector.GetServer()
|
serverStats := mem.StatsCollector.GetServer()
|
||||||
return model.ServerInfoResp{
|
svrResp := model.ServerInfoResp{
|
||||||
Version: version.Full(),
|
Version: version.Full(),
|
||||||
BindPort: c.serverCfg.BindPort,
|
BindPort: c.serverCfg.BindPort,
|
||||||
VhostHTTPPort: c.serverCfg.VhostHTTPPort,
|
VhostHTTPPort: c.serverCfg.VhostHTTPPort,
|
||||||
@@ -84,6 +80,8 @@ func (c *Controller) buildServerInfoResp() model.ServerInfoResp {
|
|||||||
ClientCounts: serverStats.ClientCounts,
|
ClientCounts: serverStats.ClientCounts,
|
||||||
ProxyTypeCounts: serverStats.ProxyTypeCounts,
|
ProxyTypeCounts: serverStats.ProxyTypeCounts,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return svrResp, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// /api/clients
|
// /api/clients
|
||||||
|
|||||||
@@ -1,647 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package http
|
|
||||||
|
|
||||||
import (
|
|
||||||
"cmp"
|
|
||||||
"fmt"
|
|
||||||
"maps"
|
|
||||||
"math"
|
|
||||||
"net/http"
|
|
||||||
"net/url"
|
|
||||||
"slices"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/metrics/mem"
|
|
||||||
httppkg "github.com/fatedier/frp/pkg/util/http"
|
|
||||||
"github.com/fatedier/frp/server/http/model"
|
|
||||||
"github.com/fatedier/frp/server/registry"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
defaultV2Page = 1
|
|
||||||
defaultV2PageSize = 50
|
|
||||||
maxV2PageSize = 200
|
|
||||||
|
|
||||||
v2SystemPruneTypeOfflineProxies = "offline_proxies"
|
|
||||||
v2ProxyTrafficDefaultDays = 7
|
|
||||||
v2ProxyTrafficUnit = "bytes"
|
|
||||||
v2ProxyTrafficGranularity = "day"
|
|
||||||
)
|
|
||||||
|
|
||||||
var apiV2ProxyTypes = []string{
|
|
||||||
string(v1.ProxyTypeTCP),
|
|
||||||
string(v1.ProxyTypeUDP),
|
|
||||||
string(v1.ProxyTypeHTTP),
|
|
||||||
string(v1.ProxyTypeHTTPS),
|
|
||||||
string(v1.ProxyTypeTCPMUX),
|
|
||||||
string(v1.ProxyTypeSTCP),
|
|
||||||
string(v1.ProxyTypeXTCP),
|
|
||||||
string(v1.ProxyTypeSUDP),
|
|
||||||
}
|
|
||||||
|
|
||||||
// /api/v2/system/info
|
|
||||||
func (c *Controller) APIV2SystemInfo(ctx *httppkg.Context) (any, error) {
|
|
||||||
info := c.buildServerInfoResp()
|
|
||||||
proxyTypeCounts := info.ProxyTypeCounts
|
|
||||||
if proxyTypeCounts == nil {
|
|
||||||
proxyTypeCounts = map[string]int64{}
|
|
||||||
}
|
|
||||||
|
|
||||||
return model.V2SystemInfoResp{
|
|
||||||
Version: info.Version,
|
|
||||||
Config: model.V2SystemInfoConfigResp{
|
|
||||||
BindPort: info.BindPort,
|
|
||||||
VhostHTTPPort: info.VhostHTTPPort,
|
|
||||||
VhostHTTPSPort: info.VhostHTTPSPort,
|
|
||||||
TCPMuxHTTPConnectPort: info.TCPMuxHTTPConnectPort,
|
|
||||||
KCPBindPort: info.KCPBindPort,
|
|
||||||
QUICBindPort: info.QUICBindPort,
|
|
||||||
SubdomainHost: info.SubdomainHost,
|
|
||||||
MaxPoolCount: info.MaxPoolCount,
|
|
||||||
MaxPortsPerClient: info.MaxPortsPerClient,
|
|
||||||
HeartbeatTimeout: info.HeartBeatTimeout,
|
|
||||||
AllowPortsStr: info.AllowPortsStr,
|
|
||||||
TLSForce: info.TLSForce,
|
|
||||||
},
|
|
||||||
Status: model.V2SystemInfoStatusResp{
|
|
||||||
TotalTrafficIn: info.TotalTrafficIn,
|
|
||||||
TotalTrafficOut: info.TotalTrafficOut,
|
|
||||||
CurConns: info.CurConns,
|
|
||||||
ClientCounts: info.ClientCounts,
|
|
||||||
ProxyTypeCounts: proxyTypeCounts,
|
|
||||||
},
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// /api/v2/system/prune
|
|
||||||
func (c *Controller) APIV2SystemPrune(ctx *httppkg.Context) (any, error) {
|
|
||||||
pruneType, err := parseV2SystemPruneType(ctx.Query("type"))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
cleared, total := mem.StatsCollector.PruneOfflineProxies()
|
|
||||||
return model.V2SystemPruneResp{
|
|
||||||
Type: pruneType,
|
|
||||||
Cleared: cleared,
|
|
||||||
Total: total,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// /api/v2/users
|
|
||||||
func (c *Controller) APIV2UserList(ctx *httppkg.Context) (any, error) {
|
|
||||||
page, pageSize, err := parseV2PageParams(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if c.clientRegistry == nil {
|
|
||||||
return nil, fmt.Errorf("client registry unavailable")
|
|
||||||
}
|
|
||||||
|
|
||||||
userStats := make(map[string]*model.V2UserResp)
|
|
||||||
for _, info := range c.clientRegistry.List() {
|
|
||||||
item := getOrCreateV2User(userStats, info.User)
|
|
||||||
item.ClientCount++
|
|
||||||
}
|
|
||||||
for _, proxyInfo := range c.listV2ProxyStats("") {
|
|
||||||
item := getOrCreateV2User(userStats, proxyInfo.User)
|
|
||||||
item.ProxyCount++
|
|
||||||
}
|
|
||||||
|
|
||||||
q := strings.ToLower(ctx.Query("q"))
|
|
||||||
items := make([]model.V2UserResp, 0, len(userStats))
|
|
||||||
for _, item := range userStats {
|
|
||||||
if q != "" && !strings.Contains(strings.ToLower(item.User), q) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
items = append(items, *item)
|
|
||||||
}
|
|
||||||
slices.SortFunc(items, func(a, b model.V2UserResp) int {
|
|
||||||
return cmp.Compare(a.User, b.User)
|
|
||||||
})
|
|
||||||
|
|
||||||
return buildV2PageResp(items, page, pageSize), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// /api/v2/clients
|
|
||||||
func (c *Controller) APIV2ClientList(ctx *httppkg.Context) (any, error) {
|
|
||||||
page, pageSize, err := parseV2PageParams(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if c.clientRegistry == nil {
|
|
||||||
return nil, fmt.Errorf("client registry unavailable")
|
|
||||||
}
|
|
||||||
statusFilter, err := parseV2StatusFilter(ctx.Query("status"))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
userFilter, filterByUser := queryValue(ctx, "user")
|
|
||||||
clientIDFilter := ctx.Query("clientID")
|
|
||||||
runIDFilter := ctx.Query("runID")
|
|
||||||
q := strings.ToLower(ctx.Query("q"))
|
|
||||||
|
|
||||||
records := c.clientRegistry.List()
|
|
||||||
items := make([]model.ClientInfoResp, 0, len(records))
|
|
||||||
for _, info := range records {
|
|
||||||
if filterByUser && info.User != userFilter {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if clientIDFilter != "" && info.ClientID() != clientIDFilter {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if runIDFilter != "" && info.RunID != runIDFilter {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if !matchV2StatusFilter(info.Online, statusFilter) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
resp := buildClientInfoResp(info)
|
|
||||||
if q != "" && !matchV2ClientQuery(resp, q) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
items = append(items, resp)
|
|
||||||
}
|
|
||||||
|
|
||||||
slices.SortFunc(items, func(a, b model.ClientInfoResp) int {
|
|
||||||
if v := cmp.Compare(a.User, b.User); v != 0 {
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
if v := cmp.Compare(a.ClientID, b.ClientID); v != 0 {
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
return cmp.Compare(a.Key, b.Key)
|
|
||||||
})
|
|
||||||
|
|
||||||
return buildV2PageResp(items, page, pageSize), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// /api/v2/clients/{key}
|
|
||||||
func (c *Controller) APIV2ClientDetail(ctx *httppkg.Context) (any, error) {
|
|
||||||
key, err := decodeV2PathParam(ctx, "key", "client key")
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.clientRegistry == nil {
|
|
||||||
return nil, fmt.Errorf("client registry unavailable")
|
|
||||||
}
|
|
||||||
|
|
||||||
info, ok := c.clientRegistry.GetByKey(key)
|
|
||||||
if !ok {
|
|
||||||
return nil, httppkg.NewError(http.StatusNotFound, fmt.Sprintf("client %s not found", key))
|
|
||||||
}
|
|
||||||
|
|
||||||
resp := buildClientInfoResp(info)
|
|
||||||
status := c.buildV2ClientStatus(info)
|
|
||||||
return model.V2ClientDetailResp{
|
|
||||||
ClientInfoResp: resp,
|
|
||||||
Status: status,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// /api/v2/proxies
|
|
||||||
func (c *Controller) APIV2ProxyList(ctx *httppkg.Context) (any, error) {
|
|
||||||
page, pageSize, err := parseV2PageParams(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
statusFilter, err := parseV2StatusFilter(ctx.Query("status"))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
proxyType, err := parseV2ProxyTypeFilter(ctx.Query("type"))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
userFilter, filterByUser := queryValue(ctx, "user")
|
|
||||||
clientIDFilter := ctx.Query("clientID")
|
|
||||||
q := strings.ToLower(ctx.Query("q"))
|
|
||||||
|
|
||||||
stats := c.listV2ProxyStats(proxyType)
|
|
||||||
items := make([]model.V2ProxyResp, 0, len(stats))
|
|
||||||
for _, ps := range stats {
|
|
||||||
resp := c.buildV2ProxyResp(ps)
|
|
||||||
if filterByUser && resp.User != userFilter {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if clientIDFilter != "" && resp.ClientID != clientIDFilter {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if !matchV2StatusFilter(resp.Status.State == "online", statusFilter) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if q != "" && !matchV2ProxyQuery(resp, q) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
items = append(items, resp)
|
|
||||||
}
|
|
||||||
|
|
||||||
slices.SortFunc(items, func(a, b model.V2ProxyResp) int {
|
|
||||||
if v := cmp.Compare(a.Spec.Type, b.Spec.Type); v != 0 {
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
return cmp.Compare(a.Name, b.Name)
|
|
||||||
})
|
|
||||||
|
|
||||||
return buildV2PageResp(items, page, pageSize), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// /api/v2/proxies/{name}
|
|
||||||
func (c *Controller) APIV2ProxyDetail(ctx *httppkg.Context) (any, error) {
|
|
||||||
name, err := decodeV2PathParam(ctx, "name", "proxy name")
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
ps := mem.StatsCollector.GetProxyByName(name)
|
|
||||||
if ps == nil {
|
|
||||||
return nil, httppkg.NewError(http.StatusNotFound, "no proxy info found")
|
|
||||||
}
|
|
||||||
return c.buildV2ProxyResp(ps), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// /api/v2/proxies/{name}/traffic
|
|
||||||
func (c *Controller) APIV2ProxyTraffic(ctx *httppkg.Context) (any, error) {
|
|
||||||
name, err := decodeV2PathParam(ctx, "name", "proxy name")
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
proxyTrafficInfo := mem.StatsCollector.GetProxyTraffic(name)
|
|
||||||
if proxyTrafficInfo == nil {
|
|
||||||
return nil, httppkg.NewError(http.StatusNotFound, "no proxy info found")
|
|
||||||
}
|
|
||||||
|
|
||||||
return buildV2ProxyTrafficResp(name, proxyTrafficInfo, time.Now()), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func decodeV2PathParam(ctx *httppkg.Context, key string, label string) (string, error) {
|
|
||||||
raw := ctx.Param(key)
|
|
||||||
if raw == "" {
|
|
||||||
return "", fmt.Errorf("missing %s", label)
|
|
||||||
}
|
|
||||||
decoded, err := url.PathUnescape(raw)
|
|
||||||
if err != nil {
|
|
||||||
return "", httppkg.NewError(http.StatusBadRequest, fmt.Sprintf("invalid %s", label))
|
|
||||||
}
|
|
||||||
return decoded, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func getOrCreateV2User(items map[string]*model.V2UserResp, user string) *model.V2UserResp {
|
|
||||||
item, ok := items[user]
|
|
||||||
if !ok {
|
|
||||||
item = &model.V2UserResp{User: user}
|
|
||||||
items[user] = item
|
|
||||||
}
|
|
||||||
return item
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseV2PageParams(ctx *httppkg.Context) (int, int, error) {
|
|
||||||
page, err := parseV2PositiveInt(ctx.Query("page"), defaultV2Page, "page")
|
|
||||||
if err != nil {
|
|
||||||
return 0, 0, err
|
|
||||||
}
|
|
||||||
pageSize, err := parseV2PositiveInt(ctx.Query("pageSize"), defaultV2PageSize, "pageSize")
|
|
||||||
if err != nil {
|
|
||||||
return 0, 0, err
|
|
||||||
}
|
|
||||||
if pageSize > maxV2PageSize {
|
|
||||||
return 0, 0, httppkg.NewError(http.StatusBadRequest, fmt.Sprintf("pageSize must be between 1 and %d", maxV2PageSize))
|
|
||||||
}
|
|
||||||
if page > math.MaxInt/pageSize {
|
|
||||||
return 0, 0, httppkg.NewError(http.StatusBadRequest, "page is too large")
|
|
||||||
}
|
|
||||||
return page, pageSize, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseV2PositiveInt(raw string, defaultValue int, name string) (int, error) {
|
|
||||||
if raw == "" {
|
|
||||||
return defaultValue, nil
|
|
||||||
}
|
|
||||||
value, err := strconv.Atoi(raw)
|
|
||||||
if err != nil || value < 1 {
|
|
||||||
return 0, httppkg.NewError(http.StatusBadRequest, fmt.Sprintf("%s must be a positive integer", name))
|
|
||||||
}
|
|
||||||
return value, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseV2StatusFilter(raw string) (string, error) {
|
|
||||||
status := strings.ToLower(raw)
|
|
||||||
switch status {
|
|
||||||
case "", "all", "online", "offline":
|
|
||||||
return status, nil
|
|
||||||
default:
|
|
||||||
return "", httppkg.NewError(http.StatusBadRequest, "status must be one of all, online, offline")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseV2ProxyTypeFilter(raw string) (string, error) {
|
|
||||||
proxyType := strings.ToLower(raw)
|
|
||||||
if proxyType == "" {
|
|
||||||
return "", nil
|
|
||||||
}
|
|
||||||
if slices.Contains(apiV2ProxyTypes, proxyType) {
|
|
||||||
return proxyType, nil
|
|
||||||
}
|
|
||||||
return "", httppkg.NewError(http.StatusBadRequest, "type must be one of tcp, udp, http, https, tcpmux, stcp, xtcp, sudp")
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseV2SystemPruneType(raw string) (string, error) {
|
|
||||||
pruneType := strings.ToLower(raw)
|
|
||||||
switch pruneType {
|
|
||||||
case "":
|
|
||||||
return "", httppkg.NewError(http.StatusBadRequest, "type is required")
|
|
||||||
case v2SystemPruneTypeOfflineProxies:
|
|
||||||
return pruneType, nil
|
|
||||||
default:
|
|
||||||
return "", httppkg.NewError(http.StatusBadRequest, "type must be one of offline_proxies")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func matchV2StatusFilter(online bool, filter string) bool {
|
|
||||||
switch filter {
|
|
||||||
case "", "all":
|
|
||||||
return true
|
|
||||||
case "online":
|
|
||||||
return online
|
|
||||||
case "offline":
|
|
||||||
return !online
|
|
||||||
default:
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func buildV2PageResp[T any](items []T, page, pageSize int) model.V2PageResp[T] {
|
|
||||||
total := len(items)
|
|
||||||
return model.V2PageResp[T]{
|
|
||||||
Total: total,
|
|
||||||
Page: page,
|
|
||||||
PageSize: pageSize,
|
|
||||||
Items: paginateV2Items(items, page, pageSize),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func paginateV2Items[T any](items []T, page, pageSize int) []T {
|
|
||||||
start := (page - 1) * pageSize
|
|
||||||
if start >= len(items) {
|
|
||||||
return []T{}
|
|
||||||
}
|
|
||||||
end := min(start+pageSize, len(items))
|
|
||||||
return items[start:end]
|
|
||||||
}
|
|
||||||
|
|
||||||
func queryValue(ctx *httppkg.Context, key string) (string, bool) {
|
|
||||||
values, ok := ctx.Req.URL.Query()[key]
|
|
||||||
if !ok {
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
if len(values) == 0 {
|
|
||||||
return "", true
|
|
||||||
}
|
|
||||||
return values[0], true
|
|
||||||
}
|
|
||||||
|
|
||||||
func matchV2ClientQuery(item model.ClientInfoResp, q string) bool {
|
|
||||||
return containsV2Query(q,
|
|
||||||
item.Key,
|
|
||||||
item.User,
|
|
||||||
item.ClientID,
|
|
||||||
item.RunID,
|
|
||||||
item.Version,
|
|
||||||
item.WireProtocol,
|
|
||||||
item.Hostname,
|
|
||||||
item.ClientIP,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
func matchV2ProxyQuery(item model.V2ProxyResp, q string) bool {
|
|
||||||
values := []string{
|
|
||||||
item.Name,
|
|
||||||
item.Spec.Type,
|
|
||||||
item.User,
|
|
||||||
item.ClientID,
|
|
||||||
item.Status.State,
|
|
||||||
}
|
|
||||||
|
|
||||||
switch item.Spec.Type {
|
|
||||||
case string(v1.ProxyTypeTCP):
|
|
||||||
if item.Spec.TCP != nil && item.Spec.TCP.RemotePort != nil {
|
|
||||||
values = append(values, strconv.Itoa(*item.Spec.TCP.RemotePort))
|
|
||||||
}
|
|
||||||
case string(v1.ProxyTypeUDP):
|
|
||||||
if item.Spec.UDP != nil && item.Spec.UDP.RemotePort != nil {
|
|
||||||
values = append(values, strconv.Itoa(*item.Spec.UDP.RemotePort))
|
|
||||||
}
|
|
||||||
case string(v1.ProxyTypeHTTP):
|
|
||||||
if item.Spec.HTTP != nil {
|
|
||||||
values = append(values, item.Spec.HTTP.CustomDomains...)
|
|
||||||
values = append(values, item.Spec.HTTP.Subdomain)
|
|
||||||
}
|
|
||||||
case string(v1.ProxyTypeHTTPS):
|
|
||||||
if item.Spec.HTTPS != nil {
|
|
||||||
values = append(values, item.Spec.HTTPS.CustomDomains...)
|
|
||||||
values = append(values, item.Spec.HTTPS.Subdomain)
|
|
||||||
}
|
|
||||||
case string(v1.ProxyTypeTCPMUX):
|
|
||||||
if item.Spec.TCPMux != nil {
|
|
||||||
values = append(values, item.Spec.TCPMux.CustomDomains...)
|
|
||||||
values = append(values, item.Spec.TCPMux.Subdomain)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return containsV2Query(q, values...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func containsV2Query(q string, values ...string) bool {
|
|
||||||
for _, value := range values {
|
|
||||||
if strings.Contains(strings.ToLower(value), q) {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Controller) listV2ProxyStats(proxyType string) []*mem.ProxyStats {
|
|
||||||
if proxyType != "" {
|
|
||||||
return mem.StatsCollector.GetProxiesByType(proxyType)
|
|
||||||
}
|
|
||||||
|
|
||||||
items := make([]*mem.ProxyStats, 0)
|
|
||||||
for _, t := range apiV2ProxyTypes {
|
|
||||||
items = append(items, mem.StatsCollector.GetProxiesByType(t)...)
|
|
||||||
}
|
|
||||||
return items
|
|
||||||
}
|
|
||||||
|
|
||||||
func buildV2ProxyTrafficResp(name string, traffic *mem.ProxyTrafficInfo, now time.Time) model.V2ProxyTrafficResp {
|
|
||||||
history := make([]model.V2ProxyTrafficPointResp, 0, v2ProxyTrafficDefaultDays)
|
|
||||||
for age := v2ProxyTrafficDefaultDays - 1; age >= 0; age-- {
|
|
||||||
history = append(history, model.V2ProxyTrafficPointResp{
|
|
||||||
Date: now.AddDate(0, 0, -age).Format(time.DateOnly),
|
|
||||||
TrafficIn: v2TrafficValueAt(traffic.TrafficIn, age),
|
|
||||||
TrafficOut: v2TrafficValueAt(traffic.TrafficOut, age),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
return model.V2ProxyTrafficResp{
|
|
||||||
Name: name,
|
|
||||||
Unit: v2ProxyTrafficUnit,
|
|
||||||
Granularity: v2ProxyTrafficGranularity,
|
|
||||||
History: history,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func v2TrafficValueAt(values []int64, todayFirstIndex int) int64 {
|
|
||||||
if todayFirstIndex >= len(values) {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
return values[todayFirstIndex]
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Controller) buildV2ClientStatus(info registry.ClientInfo) model.V2ClientStatusResp {
|
|
||||||
status := model.V2ClientStatusResp{State: "offline"}
|
|
||||||
if info.Online {
|
|
||||||
status.State = "online"
|
|
||||||
}
|
|
||||||
|
|
||||||
user := info.User
|
|
||||||
clientID := info.ClientID()
|
|
||||||
for _, ps := range c.listV2ProxyStats("") {
|
|
||||||
if ps.User != user || ps.ClientID != clientID {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
status.CurConns += ps.CurConns
|
|
||||||
status.ProxyCount++
|
|
||||||
}
|
|
||||||
return status
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Controller) buildV2ProxyResp(ps *mem.ProxyStats) model.V2ProxyResp {
|
|
||||||
state := "offline"
|
|
||||||
var cfg v1.ProxyConfigurer
|
|
||||||
if c.pxyManager != nil {
|
|
||||||
if pxy, ok := c.pxyManager.GetByName(ps.Name); ok {
|
|
||||||
state = "online"
|
|
||||||
cfg = pxy.GetConfigurer()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return model.V2ProxyResp{
|
|
||||||
Name: ps.Name,
|
|
||||||
User: ps.User,
|
|
||||||
ClientID: ps.ClientID,
|
|
||||||
Spec: buildV2ProxySpec(ps.Type, cfg),
|
|
||||||
Status: model.V2ProxyStatusResp{
|
|
||||||
State: state,
|
|
||||||
TodayTrafficIn: ps.TodayTrafficIn,
|
|
||||||
TodayTrafficOut: ps.TodayTrafficOut,
|
|
||||||
CurConns: ps.CurConns,
|
|
||||||
LastStartAt: ps.LastStartAt,
|
|
||||||
LastCloseAt: ps.LastCloseAt,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func buildV2ProxySpec(proxyType string, cfg v1.ProxyConfigurer) model.V2ProxySpec {
|
|
||||||
spec := model.V2ProxySpec{Type: proxyType}
|
|
||||||
|
|
||||||
switch proxyType {
|
|
||||||
case string(v1.ProxyTypeTCP):
|
|
||||||
block := &model.V2TCPProxySpec{}
|
|
||||||
if c, ok := cfg.(*v1.TCPProxyConfig); ok {
|
|
||||||
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
|
|
||||||
block.RemotePort = &c.RemotePort
|
|
||||||
}
|
|
||||||
spec.TCP = block
|
|
||||||
case string(v1.ProxyTypeUDP):
|
|
||||||
block := &model.V2UDPProxySpec{}
|
|
||||||
if c, ok := cfg.(*v1.UDPProxyConfig); ok {
|
|
||||||
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
|
|
||||||
block.RemotePort = &c.RemotePort
|
|
||||||
}
|
|
||||||
spec.UDP = block
|
|
||||||
case string(v1.ProxyTypeHTTP):
|
|
||||||
block := &model.V2HTTPProxySpec{}
|
|
||||||
if c, ok := cfg.(*v1.HTTPProxyConfig); ok {
|
|
||||||
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
|
|
||||||
block.CustomDomains = slices.Clone(c.CustomDomains)
|
|
||||||
block.Subdomain = c.SubDomain
|
|
||||||
block.Locations = slices.Clone(c.Locations)
|
|
||||||
block.HostHeaderRewrite = c.HostHeaderRewrite
|
|
||||||
}
|
|
||||||
spec.HTTP = block
|
|
||||||
case string(v1.ProxyTypeHTTPS):
|
|
||||||
block := &model.V2HTTPSProxySpec{}
|
|
||||||
if c, ok := cfg.(*v1.HTTPSProxyConfig); ok {
|
|
||||||
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
|
|
||||||
block.CustomDomains = slices.Clone(c.CustomDomains)
|
|
||||||
block.Subdomain = c.SubDomain
|
|
||||||
}
|
|
||||||
spec.HTTPS = block
|
|
||||||
case string(v1.ProxyTypeTCPMUX):
|
|
||||||
block := &model.V2TCPMuxProxySpec{}
|
|
||||||
if c, ok := cfg.(*v1.TCPMuxProxyConfig); ok {
|
|
||||||
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
|
|
||||||
block.CustomDomains = slices.Clone(c.CustomDomains)
|
|
||||||
block.Subdomain = c.SubDomain
|
|
||||||
block.Multiplexer = c.Multiplexer
|
|
||||||
block.RouteByHTTPUser = c.RouteByHTTPUser
|
|
||||||
}
|
|
||||||
spec.TCPMux = block
|
|
||||||
case string(v1.ProxyTypeSTCP):
|
|
||||||
block := &model.V2STCPProxySpec{}
|
|
||||||
if c, ok := cfg.(*v1.STCPProxyConfig); ok {
|
|
||||||
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
|
|
||||||
}
|
|
||||||
spec.STCP = block
|
|
||||||
case string(v1.ProxyTypeSUDP):
|
|
||||||
block := &model.V2SUDPProxySpec{}
|
|
||||||
if c, ok := cfg.(*v1.SUDPProxyConfig); ok {
|
|
||||||
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
|
|
||||||
}
|
|
||||||
spec.SUDP = block
|
|
||||||
case string(v1.ProxyTypeXTCP):
|
|
||||||
block := &model.V2XTCPProxySpec{}
|
|
||||||
if c, ok := cfg.(*v1.XTCPProxyConfig); ok {
|
|
||||||
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
|
|
||||||
}
|
|
||||||
spec.XTCP = block
|
|
||||||
}
|
|
||||||
|
|
||||||
return spec
|
|
||||||
}
|
|
||||||
|
|
||||||
func buildV2ProxyBaseSpec(base *v1.ProxyBaseConfig) model.V2ProxyBaseSpec {
|
|
||||||
return model.V2ProxyBaseSpec{
|
|
||||||
Annotations: maps.Clone(base.Annotations),
|
|
||||||
Metadatas: maps.Clone(base.Metadatas),
|
|
||||||
Transport: &model.V2ProxyTransportSpec{
|
|
||||||
UseEncryption: base.Transport.UseEncryption,
|
|
||||||
UseCompression: base.Transport.UseCompression,
|
|
||||||
BandwidthLimit: base.Transport.BandwidthLimit.String(),
|
|
||||||
BandwidthLimitMode: base.Transport.BandwidthLimitMode,
|
|
||||||
},
|
|
||||||
LoadBalancer: &model.V2ProxyLoadBalancerSpec{
|
|
||||||
Group: base.LoadBalancer.Group,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,393 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package http
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
configtypes "github.com/fatedier/frp/pkg/config/types"
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/metrics/mem"
|
|
||||||
"github.com/fatedier/frp/server/http/model"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestBuildV2ProxySpecAllTypesAndRedaction(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
proxyType string
|
|
||||||
cfg v1.ProxyConfigurer
|
|
||||||
blockKeys []string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
proxyType: "tcp",
|
|
||||||
cfg: &v1.TCPProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "tcp"),
|
|
||||||
RemotePort: 6000,
|
|
||||||
},
|
|
||||||
blockKeys: []string{"annotations", "loadBalancer", "metadatas", "remotePort", "transport"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
proxyType: "udp",
|
|
||||||
cfg: &v1.UDPProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "udp"),
|
|
||||||
RemotePort: 7000,
|
|
||||||
},
|
|
||||||
blockKeys: []string{"annotations", "loadBalancer", "metadatas", "remotePort", "transport"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
proxyType: "http",
|
|
||||||
cfg: &v1.HTTPProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "http"),
|
|
||||||
DomainConfig: v1.DomainConfig{CustomDomains: []string{"app.example.com"}, SubDomain: "app"},
|
|
||||||
Locations: []string{"/api"},
|
|
||||||
HTTPUser: "secret-http-user",
|
|
||||||
HTTPPassword: "secret-http-password",
|
|
||||||
HostHeaderRewrite: "backend.example.com",
|
|
||||||
RequestHeaders: v1.HeaderOperations{Set: map[string]string{"X-Secret": "secret-request-header"}},
|
|
||||||
ResponseHeaders: v1.HeaderOperations{Set: map[string]string{"X-Secret": "secret-response-header"}},
|
|
||||||
RouteByHTTPUser: "secret-http-route-user",
|
|
||||||
},
|
|
||||||
blockKeys: []string{"annotations", "customDomains", "hostHeaderRewrite", "loadBalancer", "locations", "metadatas", "subdomain", "transport"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
proxyType: "https",
|
|
||||||
cfg: &v1.HTTPSProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "https"),
|
|
||||||
DomainConfig: v1.DomainConfig{CustomDomains: []string{"secure.example.com"}, SubDomain: "secure"},
|
|
||||||
},
|
|
||||||
blockKeys: []string{"annotations", "customDomains", "loadBalancer", "metadatas", "subdomain", "transport"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
proxyType: "tcpmux",
|
|
||||||
cfg: &v1.TCPMuxProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "tcpmux"),
|
|
||||||
DomainConfig: v1.DomainConfig{CustomDomains: []string{"mux.example.com"}, SubDomain: "mux"},
|
|
||||||
HTTPUser: strings.Join([]string{"secret", "mux-http-user"}, "-"),
|
|
||||||
HTTPPassword: strings.Join([]string{"secret", "mux-http-password"}, "-"),
|
|
||||||
RouteByHTTPUser: "displayed-mux-user",
|
|
||||||
Multiplexer: "httpconnect",
|
|
||||||
},
|
|
||||||
blockKeys: []string{"annotations", "customDomains", "loadBalancer", "metadatas", "multiplexer", "routeByHTTPUser", "subdomain", "transport"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
proxyType: "stcp",
|
|
||||||
cfg: &v1.STCPProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "stcp"),
|
|
||||||
Secretkey: strings.Join([]string{"secret", "stcp-key"}, "-"),
|
|
||||||
AllowUsers: []string{strings.Join([]string{"secret", "stcp-user"}, "-")},
|
|
||||||
},
|
|
||||||
blockKeys: []string{"annotations", "loadBalancer", "metadatas", "transport"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
proxyType: "sudp",
|
|
||||||
cfg: &v1.SUDPProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "sudp"),
|
|
||||||
Secretkey: strings.Join([]string{"secret", "sudp-key"}, "-"),
|
|
||||||
AllowUsers: []string{strings.Join([]string{"secret", "sudp-user"}, "-")},
|
|
||||||
},
|
|
||||||
blockKeys: []string{"annotations", "loadBalancer", "metadatas", "transport"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
proxyType: "xtcp",
|
|
||||||
cfg: &v1.XTCPProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "xtcp"),
|
|
||||||
Secretkey: strings.Join([]string{"secret", "xtcp-key"}, "-"),
|
|
||||||
AllowUsers: []string{strings.Join([]string{"secret", "xtcp-user"}, "-")},
|
|
||||||
},
|
|
||||||
blockKeys: []string{"annotations", "loadBalancer", "metadatas", "transport"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.proxyType, func(t *testing.T) {
|
|
||||||
spec := buildV2ProxySpec(tt.proxyType, tt.cfg)
|
|
||||||
raw := mustMarshalJSON(t, spec)
|
|
||||||
|
|
||||||
var specObject map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(raw, &specObject); err != nil {
|
|
||||||
t.Fatalf("unmarshal spec failed: %v", err)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, specObject, tt.proxyType, "type")
|
|
||||||
|
|
||||||
var gotType string
|
|
||||||
if err := json.Unmarshal(specObject["type"], &gotType); err != nil {
|
|
||||||
t.Fatalf("unmarshal spec type failed: %v", err)
|
|
||||||
}
|
|
||||||
if gotType != tt.proxyType {
|
|
||||||
t.Fatalf("spec type mismatch, want %q got %q", tt.proxyType, gotType)
|
|
||||||
}
|
|
||||||
|
|
||||||
var block map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(specObject[tt.proxyType], &block); err != nil {
|
|
||||||
t.Fatalf("unmarshal active block failed: %v", err)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, block, tt.blockKeys...)
|
|
||||||
assertV2ProxyCommonSpec(t, block)
|
|
||||||
assertV2ProxyTypeFields(t, tt.proxyType, specObject[tt.proxyType])
|
|
||||||
assertNoV2ProxySensitiveFields(t, block)
|
|
||||||
|
|
||||||
content := string(raw)
|
|
||||||
for _, secret := range []string{
|
|
||||||
"secret-proxy-name",
|
|
||||||
"secret-group-key",
|
|
||||||
"secret-local-host",
|
|
||||||
"secret-plugin-user",
|
|
||||||
"secret-plugin-password",
|
|
||||||
"secret-health-path",
|
|
||||||
"secret-http-user",
|
|
||||||
"secret-http-password",
|
|
||||||
"secret-request-header",
|
|
||||||
"secret-response-header",
|
|
||||||
"secret-http-route-user",
|
|
||||||
"secret-mux-http-user",
|
|
||||||
"secret-mux-http-password",
|
|
||||||
"secret-stcp-key",
|
|
||||||
"secret-stcp-user",
|
|
||||||
"secret-sudp-key",
|
|
||||||
"secret-sudp-user",
|
|
||||||
"secret-xtcp-key",
|
|
||||||
"secret-xtcp-user",
|
|
||||||
} {
|
|
||||||
if strings.Contains(content, secret) {
|
|
||||||
t.Fatalf("sensitive value %q leaked in spec: %s", secret, content)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func assertV2ProxyTypeFields(t *testing.T, proxyType string, raw json.RawMessage) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
switch proxyType {
|
|
||||||
case "tcp":
|
|
||||||
var block model.V2TCPProxySpec
|
|
||||||
if err := json.Unmarshal(raw, &block); err != nil {
|
|
||||||
t.Fatalf("unmarshal tcp block failed: %v", err)
|
|
||||||
}
|
|
||||||
if block.RemotePort == nil || *block.RemotePort != 6000 {
|
|
||||||
t.Fatalf("tcp remote port mismatch: %#v", block.RemotePort)
|
|
||||||
}
|
|
||||||
case "udp":
|
|
||||||
var block model.V2UDPProxySpec
|
|
||||||
if err := json.Unmarshal(raw, &block); err != nil {
|
|
||||||
t.Fatalf("unmarshal udp block failed: %v", err)
|
|
||||||
}
|
|
||||||
if block.RemotePort == nil || *block.RemotePort != 7000 {
|
|
||||||
t.Fatalf("udp remote port mismatch: %#v", block.RemotePort)
|
|
||||||
}
|
|
||||||
case "http":
|
|
||||||
var block model.V2HTTPProxySpec
|
|
||||||
if err := json.Unmarshal(raw, &block); err != nil {
|
|
||||||
t.Fatalf("unmarshal http block failed: %v", err)
|
|
||||||
}
|
|
||||||
if len(block.CustomDomains) != 1 || block.CustomDomains[0] != "app.example.com" ||
|
|
||||||
block.Subdomain != "app" || len(block.Locations) != 1 || block.Locations[0] != "/api" ||
|
|
||||||
block.HostHeaderRewrite != "backend.example.com" {
|
|
||||||
t.Fatalf("http fields mismatch: %#v", block)
|
|
||||||
}
|
|
||||||
case "https":
|
|
||||||
var block model.V2HTTPSProxySpec
|
|
||||||
if err := json.Unmarshal(raw, &block); err != nil {
|
|
||||||
t.Fatalf("unmarshal https block failed: %v", err)
|
|
||||||
}
|
|
||||||
if len(block.CustomDomains) != 1 || block.CustomDomains[0] != "secure.example.com" || block.Subdomain != "secure" {
|
|
||||||
t.Fatalf("https fields mismatch: %#v", block)
|
|
||||||
}
|
|
||||||
case "tcpmux":
|
|
||||||
var block model.V2TCPMuxProxySpec
|
|
||||||
if err := json.Unmarshal(raw, &block); err != nil {
|
|
||||||
t.Fatalf("unmarshal tcpmux block failed: %v", err)
|
|
||||||
}
|
|
||||||
if len(block.CustomDomains) != 1 || block.CustomDomains[0] != "mux.example.com" ||
|
|
||||||
block.Subdomain != "mux" || block.Multiplexer != "httpconnect" || block.RouteByHTTPUser != "displayed-mux-user" {
|
|
||||||
t.Fatalf("tcpmux fields mismatch: %#v", block)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildV2ProxyRespOfflineTypedShells(t *testing.T) {
|
|
||||||
for _, proxyType := range apiV2ProxyTypes {
|
|
||||||
t.Run(proxyType, func(t *testing.T) {
|
|
||||||
resp := (&Controller{}).buildV2ProxyResp(&mem.ProxyStats{
|
|
||||||
Name: "offline-" + proxyType,
|
|
||||||
Type: proxyType,
|
|
||||||
})
|
|
||||||
if resp.Status.State != "offline" {
|
|
||||||
t.Fatalf("offline phase mismatch: %#v", resp.Status)
|
|
||||||
}
|
|
||||||
|
|
||||||
var specObject map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(mustMarshalJSON(t, resp.Spec), &specObject); err != nil {
|
|
||||||
t.Fatalf("unmarshal offline spec failed: %v", err)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, specObject, proxyType, "type")
|
|
||||||
assertRawJSONKeysFromMessage(t, specObject[proxyType])
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildV2ProxySpecDoesNotPopulateMismatchedBlock(t *testing.T) {
|
|
||||||
spec := buildV2ProxySpec("tcp", &v1.UDPProxyConfig{
|
|
||||||
ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "udp"),
|
|
||||||
RemotePort: 7000,
|
|
||||||
})
|
|
||||||
|
|
||||||
var specObject map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(mustMarshalJSON(t, spec), &specObject); err != nil {
|
|
||||||
t.Fatalf("unmarshal mismatched spec failed: %v", err)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, specObject, "tcp", "type")
|
|
||||||
assertRawJSONKeysFromMessage(t, specObject["tcp"])
|
|
||||||
}
|
|
||||||
|
|
||||||
func newV2ProxyTestBaseConfig(t *testing.T, proxyType string) v1.ProxyBaseConfig {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
bandwidthLimit, err := configtypes.NewBandwidthQuantity("10MB")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("create bandwidth limit failed: %v", err)
|
|
||||||
}
|
|
||||||
enabled := false
|
|
||||||
return v1.ProxyBaseConfig{
|
|
||||||
Name: "secret-proxy-name",
|
|
||||||
Type: proxyType,
|
|
||||||
Enabled: &enabled,
|
|
||||||
Annotations: map[string]string{"annotation-key": "annotation-value"},
|
|
||||||
Metadatas: map[string]string{"metadata-key": "metadata-value"},
|
|
||||||
Transport: v1.ProxyTransport{
|
|
||||||
UseEncryption: true,
|
|
||||||
UseCompression: true,
|
|
||||||
BandwidthLimit: bandwidthLimit,
|
|
||||||
BandwidthLimitMode: configtypes.BandwidthLimitModeServer,
|
|
||||||
ProxyProtocolVersion: "v2",
|
|
||||||
},
|
|
||||||
LoadBalancer: v1.LoadBalancerConfig{
|
|
||||||
Group: "public-group",
|
|
||||||
GroupKey: "secret-group-key",
|
|
||||||
},
|
|
||||||
HealthCheck: v1.HealthCheckConfig{
|
|
||||||
Type: "http",
|
|
||||||
Path: "secret-health-path",
|
|
||||||
},
|
|
||||||
ProxyBackend: v1.ProxyBackend{
|
|
||||||
LocalIP: "secret-local-host",
|
|
||||||
LocalPort: 8080,
|
|
||||||
Plugin: v1.TypedClientPluginOptions{
|
|
||||||
Type: v1.PluginHTTPProxy,
|
|
||||||
ClientPluginOptions: &v1.HTTPProxyPluginOptions{
|
|
||||||
Type: v1.PluginHTTPProxy,
|
|
||||||
HTTPUser: "secret-plugin-user",
|
|
||||||
HTTPPassword: "secret-plugin-password",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func assertV2ProxyCommonSpec(t *testing.T, block map[string]json.RawMessage) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var annotations map[string]string
|
|
||||||
if err := json.Unmarshal(block["annotations"], &annotations); err != nil {
|
|
||||||
t.Fatalf("unmarshal annotations failed: %v", err)
|
|
||||||
}
|
|
||||||
if annotations["annotation-key"] != "annotation-value" {
|
|
||||||
t.Fatalf("annotations mismatch: %#v", annotations)
|
|
||||||
}
|
|
||||||
|
|
||||||
var metadatas map[string]string
|
|
||||||
if err := json.Unmarshal(block["metadatas"], &metadatas); err != nil {
|
|
||||||
t.Fatalf("unmarshal metadatas failed: %v", err)
|
|
||||||
}
|
|
||||||
if metadatas["metadata-key"] != "metadata-value" {
|
|
||||||
t.Fatalf("metadatas mismatch: %#v", metadatas)
|
|
||||||
}
|
|
||||||
|
|
||||||
assertRawJSONKeysFromMessage(t, block["transport"],
|
|
||||||
"bandwidthLimit",
|
|
||||||
"bandwidthLimitMode",
|
|
||||||
"useCompression",
|
|
||||||
"useEncryption",
|
|
||||||
)
|
|
||||||
var transport model.V2ProxyTransportSpec
|
|
||||||
if err := json.Unmarshal(block["transport"], &transport); err != nil {
|
|
||||||
t.Fatalf("unmarshal transport failed: %v", err)
|
|
||||||
}
|
|
||||||
if !transport.UseEncryption || !transport.UseCompression ||
|
|
||||||
transport.BandwidthLimit != "10MB" || transport.BandwidthLimitMode != "server" {
|
|
||||||
t.Fatalf("transport mismatch: %#v", transport)
|
|
||||||
}
|
|
||||||
|
|
||||||
assertRawJSONKeysFromMessage(t, block["loadBalancer"], "group")
|
|
||||||
var loadBalancer model.V2ProxyLoadBalancerSpec
|
|
||||||
if err := json.Unmarshal(block["loadBalancer"], &loadBalancer); err != nil {
|
|
||||||
t.Fatalf("unmarshal load balancer failed: %v", err)
|
|
||||||
}
|
|
||||||
if loadBalancer.Group != "public-group" {
|
|
||||||
t.Fatalf("load balancer mismatch: %#v", loadBalancer)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func assertNoV2ProxySensitiveFields(t *testing.T, value any) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
forbidden := map[string]struct{}{
|
|
||||||
"allowUsers": {},
|
|
||||||
"enabled": {},
|
|
||||||
"groupKey": {},
|
|
||||||
"healthCheck": {},
|
|
||||||
"httpPassword": {},
|
|
||||||
"httpUser": {},
|
|
||||||
"localIP": {},
|
|
||||||
"localPort": {},
|
|
||||||
"name": {},
|
|
||||||
"natTraversal": {},
|
|
||||||
"plugin": {},
|
|
||||||
"proxyProtocolVersion": {},
|
|
||||||
"requestHeaders": {},
|
|
||||||
"responseHeaders": {},
|
|
||||||
"secretKey": {},
|
|
||||||
"type": {},
|
|
||||||
}
|
|
||||||
|
|
||||||
var walk func(any)
|
|
||||||
walk = func(current any) {
|
|
||||||
switch current := current.(type) {
|
|
||||||
case map[string]any:
|
|
||||||
for key, nested := range current {
|
|
||||||
if _, ok := forbidden[key]; ok {
|
|
||||||
t.Fatalf("sensitive field %q leaked in active block", key)
|
|
||||||
}
|
|
||||||
walk(nested)
|
|
||||||
}
|
|
||||||
case []any:
|
|
||||||
for _, nested := range current {
|
|
||||||
walk(nested)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
raw, err := json.Marshal(value)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("marshal active block failed: %v", err)
|
|
||||||
}
|
|
||||||
var decoded any
|
|
||||||
if err := json.Unmarshal(raw, &decoded); err != nil {
|
|
||||||
t.Fatalf("decode active block failed: %v", err)
|
|
||||||
}
|
|
||||||
walk(decoded)
|
|
||||||
}
|
|
||||||
@@ -1,908 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package http
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"math"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"net/url"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/gorilla/mux"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/config/types"
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/metrics/mem"
|
|
||||||
httppkg "github.com/fatedier/frp/pkg/util/http"
|
|
||||||
"github.com/fatedier/frp/server/http/model"
|
|
||||||
serverproxy "github.com/fatedier/frp/server/proxy"
|
|
||||||
"github.com/fatedier/frp/server/registry"
|
|
||||||
)
|
|
||||||
|
|
||||||
type v2EnvelopeForTest[T any] struct {
|
|
||||||
Code int `json:"code"`
|
|
||||||
Msg string `json:"msg"`
|
|
||||||
Data T `json:"data"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type fakeStatsCollector struct {
|
|
||||||
server *mem.ServerStats
|
|
||||||
proxies map[string]*mem.ProxyStats
|
|
||||||
traffic map[string]*mem.ProxyTrafficInfo
|
|
||||||
pruneable map[string]bool
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *fakeStatsCollector) GetServer() *mem.ServerStats {
|
|
||||||
if f.server != nil {
|
|
||||||
return f.server
|
|
||||||
}
|
|
||||||
return &mem.ServerStats{ProxyTypeCounts: map[string]int64{}}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *fakeStatsCollector) GetProxiesByType(proxyType string) []*mem.ProxyStats {
|
|
||||||
items := make([]*mem.ProxyStats, 0)
|
|
||||||
for _, ps := range f.proxies {
|
|
||||||
if ps.Type == proxyType {
|
|
||||||
items = append(items, ps)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return items
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *fakeStatsCollector) GetProxiesByTypeAndName(proxyType string, proxyName string) *mem.ProxyStats {
|
|
||||||
ps := f.proxies[proxyName]
|
|
||||||
if ps != nil && ps.Type == proxyType {
|
|
||||||
return ps
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *fakeStatsCollector) GetProxyByName(proxyName string) *mem.ProxyStats {
|
|
||||||
return f.proxies[proxyName]
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *fakeStatsCollector) GetProxyTraffic(name string) *mem.ProxyTrafficInfo {
|
|
||||||
return f.traffic[name]
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *fakeStatsCollector) ClearOfflineProxies() (int, int) {
|
|
||||||
return 0, len(f.proxies)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *fakeStatsCollector) PruneOfflineProxies() (int, int) {
|
|
||||||
total := len(f.proxies)
|
|
||||||
cleared := 0
|
|
||||||
for name := range f.pruneable {
|
|
||||||
if _, ok := f.proxies[name]; ok {
|
|
||||||
delete(f.proxies, name)
|
|
||||||
cleared++
|
|
||||||
}
|
|
||||||
}
|
|
||||||
f.pruneable = map[string]bool{}
|
|
||||||
return cleared, total
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2SystemInfoEnvelope(t *testing.T) {
|
|
||||||
oldStatsCollector := mem.StatsCollector
|
|
||||||
mem.StatsCollector = &fakeStatsCollector{
|
|
||||||
server: &mem.ServerStats{
|
|
||||||
TotalTrafficIn: 1024,
|
|
||||||
TotalTrafficOut: 2048,
|
|
||||||
CurConns: 3,
|
|
||||||
ClientCounts: 4,
|
|
||||||
ProxyTypeCounts: map[string]int64{
|
|
||||||
"tcp": 2,
|
|
||||||
"http": 1,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
proxies: map[string]*mem.ProxyStats{},
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
mem.StatsCollector = oldStatsCollector
|
|
||||||
})
|
|
||||||
|
|
||||||
controller := NewController(&v1.ServerConfig{
|
|
||||||
BindPort: 7000,
|
|
||||||
VhostHTTPPort: 8080,
|
|
||||||
VhostHTTPSPort: 8443,
|
|
||||||
TCPMuxHTTPConnectPort: 9000,
|
|
||||||
KCPBindPort: 7001,
|
|
||||||
QUICBindPort: 7002,
|
|
||||||
SubDomainHost: "example.com",
|
|
||||||
MaxPortsPerClient: 8,
|
|
||||||
AllowPorts: []types.PortsRange{
|
|
||||||
{Start: 1000, End: 1002},
|
|
||||||
{Single: 2000},
|
|
||||||
},
|
|
||||||
Transport: v1.ServerTransportConfig{
|
|
||||||
MaxPoolCount: 5,
|
|
||||||
HeartbeatTimeout: 90,
|
|
||||||
TLS: v1.TLSServerConfig{
|
|
||||||
Force: true,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}, registry.NewClientRegistry(), serverproxy.NewManager())
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/v2/system/info")
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d", http.StatusOK, resp.Code)
|
|
||||||
}
|
|
||||||
|
|
||||||
rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp)
|
|
||||||
if rawResp.Code != http.StatusOK || rawResp.Msg != "success" {
|
|
||||||
t.Fatalf("envelope mismatch: %#v", rawResp)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, rawResp.Data, "config", "status", "version")
|
|
||||||
assertRawJSONKeysFromMessage(t, rawResp.Data["config"],
|
|
||||||
"allowPortsStr",
|
|
||||||
"bindPort",
|
|
||||||
"heartbeatTimeout",
|
|
||||||
"kcpBindPort",
|
|
||||||
"maxPoolCount",
|
|
||||||
"maxPortsPerClient",
|
|
||||||
"quicBindPort",
|
|
||||||
"subdomainHost",
|
|
||||||
"tcpmuxHTTPConnectPort",
|
|
||||||
"tlsForce",
|
|
||||||
"vhostHTTPPort",
|
|
||||||
"vhostHTTPSPort",
|
|
||||||
)
|
|
||||||
assertRawJSONKeysFromMessage(t, rawResp.Data["status"],
|
|
||||||
"clientCounts",
|
|
||||||
"curConns",
|
|
||||||
"proxyTypeCount",
|
|
||||||
"totalTrafficIn",
|
|
||||||
"totalTrafficOut",
|
|
||||||
)
|
|
||||||
|
|
||||||
systemResp := decodeResponse[v2EnvelopeForTest[model.V2SystemInfoResp]](t, resp)
|
|
||||||
if systemResp.Data.Version == "" {
|
|
||||||
t.Fatal("version should be set at top level")
|
|
||||||
}
|
|
||||||
if systemResp.Data.Config.BindPort != 7000 ||
|
|
||||||
systemResp.Data.Config.VhostHTTPPort != 8080 ||
|
|
||||||
systemResp.Data.Config.VhostHTTPSPort != 8443 ||
|
|
||||||
systemResp.Data.Config.TCPMuxHTTPConnectPort != 9000 ||
|
|
||||||
systemResp.Data.Config.KCPBindPort != 7001 ||
|
|
||||||
systemResp.Data.Config.QUICBindPort != 7002 ||
|
|
||||||
systemResp.Data.Config.SubdomainHost != "example.com" ||
|
|
||||||
systemResp.Data.Config.MaxPoolCount != 5 ||
|
|
||||||
systemResp.Data.Config.MaxPortsPerClient != 8 ||
|
|
||||||
systemResp.Data.Config.HeartbeatTimeout != 90 ||
|
|
||||||
systemResp.Data.Config.AllowPortsStr != "1000-1002,2000" ||
|
|
||||||
!systemResp.Data.Config.TLSForce {
|
|
||||||
t.Fatalf("config mismatch: %#v", systemResp.Data.Config)
|
|
||||||
}
|
|
||||||
if systemResp.Data.Status.TotalTrafficIn != 1024 ||
|
|
||||||
systemResp.Data.Status.TotalTrafficOut != 2048 ||
|
|
||||||
systemResp.Data.Status.CurConns != 3 ||
|
|
||||||
systemResp.Data.Status.ClientCounts != 4 ||
|
|
||||||
systemResp.Data.Status.ProxyTypeCounts["tcp"] != 2 ||
|
|
||||||
systemResp.Data.Status.ProxyTypeCounts["http"] != 1 {
|
|
||||||
t.Fatalf("status mismatch: %#v", systemResp.Data.Status)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2SystemPruneOfflineProxies(t *testing.T) {
|
|
||||||
oldStatsCollector := mem.StatsCollector
|
|
||||||
collector := &fakeStatsCollector{
|
|
||||||
proxies: map[string]*mem.ProxyStats{
|
|
||||||
"tcp-offline": {Name: "tcp-offline", Type: "tcp"},
|
|
||||||
"http-offline": {Name: "http-offline", Type: "http"},
|
|
||||||
"udp-offline": {Name: "udp-offline", Type: "udp"},
|
|
||||||
"tcp-online": {Name: "tcp-online", Type: "tcp"},
|
|
||||||
"http-online": {Name: "http-online", Type: "http"},
|
|
||||||
"udp-online": {Name: "udp-online", Type: "udp"},
|
|
||||||
"stcp-restarted": {Name: "stcp-restarted", Type: "stcp"},
|
|
||||||
"xtcp-restarted": {Name: "xtcp-restarted", Type: "xtcp"},
|
|
||||||
"sudp-same-time": {Name: "sudp-same-time", Type: "sudp"},
|
|
||||||
"tcpmux-running": {Name: "tcpmux-running", Type: "tcpmux"},
|
|
||||||
},
|
|
||||||
pruneable: map[string]bool{
|
|
||||||
"tcp-offline": true,
|
|
||||||
"http-offline": true,
|
|
||||||
"udp-offline": true,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
mem.StatsCollector = collector
|
|
||||||
t.Cleanup(func() {
|
|
||||||
mem.StatsCollector = oldStatsCollector
|
|
||||||
})
|
|
||||||
|
|
||||||
controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager())
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=offline_proxies")
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
|
|
||||||
}
|
|
||||||
rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp)
|
|
||||||
if rawResp.Code != http.StatusOK || rawResp.Msg != "success" {
|
|
||||||
t.Fatalf("envelope mismatch: %#v", rawResp)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, rawResp.Data, "cleared", "total", "type")
|
|
||||||
pruneResp := decodeResponse[v2EnvelopeForTest[model.V2SystemPruneResp]](t, resp)
|
|
||||||
if pruneResp.Data.Type != "offline_proxies" || pruneResp.Data.Cleared != 3 || pruneResp.Data.Total != 10 {
|
|
||||||
t.Fatalf("prune response mismatch: %#v", pruneResp.Data)
|
|
||||||
}
|
|
||||||
if _, ok := collector.proxies["tcp-offline"]; ok {
|
|
||||||
t.Fatal("pruned proxy statistics should be removed")
|
|
||||||
}
|
|
||||||
if _, ok := collector.proxies["tcp-online"]; !ok {
|
|
||||||
t.Fatal("online proxy statistics should remain")
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=offline_proxies")
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("second prune status mismatch, want %d got %d", http.StatusOK, resp.Code)
|
|
||||||
}
|
|
||||||
pruneResp = decodeResponse[v2EnvelopeForTest[model.V2SystemPruneResp]](t, resp)
|
|
||||||
if pruneResp.Data.Cleared != 0 || pruneResp.Data.Total != 7 {
|
|
||||||
t.Fatalf("second prune response mismatch: %#v", pruneResp.Data)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2SystemPruneTypeErrorsUseEnvelope(t *testing.T) {
|
|
||||||
controller := newV2TestController(t)
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune")
|
|
||||||
if resp.Code != http.StatusBadRequest {
|
|
||||||
t.Fatalf("missing type status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
|
|
||||||
}
|
|
||||||
errResp := decodeResponse[httppkg.V2Response](t, resp)
|
|
||||||
if errResp.Code != http.StatusBadRequest || errResp.Msg != "type is required" || errResp.Data != nil {
|
|
||||||
t.Fatalf("missing type error envelope mismatch: %#v", errResp)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=clients")
|
|
||||||
if resp.Code != http.StatusBadRequest {
|
|
||||||
t.Fatalf("invalid type status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
|
|
||||||
}
|
|
||||||
errResp = decodeResponse[httppkg.V2Response](t, resp)
|
|
||||||
if errResp.Code != http.StatusBadRequest || errResp.Msg != "type must be one of offline_proxies" || errResp.Data != nil {
|
|
||||||
t.Fatalf("invalid type error envelope mismatch: %#v", errResp)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2ClientListEnvelopePaginationAndFilters(t *testing.T) {
|
|
||||||
controller := newV2TestController(t)
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/v2/clients?page=1&pageSize=1")
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d", http.StatusOK, resp.Code)
|
|
||||||
}
|
|
||||||
pageResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.ClientInfoResp]]](t, resp)
|
|
||||||
if pageResp.Code != http.StatusOK || pageResp.Msg != "success" {
|
|
||||||
t.Fatalf("envelope mismatch: %#v", pageResp)
|
|
||||||
}
|
|
||||||
if pageResp.Data.Total != 3 || pageResp.Data.Page != 1 || pageResp.Data.PageSize != 1 || len(pageResp.Data.Items) != 1 {
|
|
||||||
t.Fatalf("page data mismatch: %#v", pageResp.Data)
|
|
||||||
}
|
|
||||||
if got := pageResp.Data.Items[0].User; got != "" {
|
|
||||||
t.Fatalf("first sorted user mismatch, want empty got %q", got)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/v2/clients?user=&page=1&pageSize=50")
|
|
||||||
emptyUserResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.ClientInfoResp]]](t, resp)
|
|
||||||
if emptyUserResp.Data.Total != 1 || emptyUserResp.Data.Items[0].User != "" {
|
|
||||||
t.Fatalf("empty user filter mismatch: %#v", emptyUserResp.Data)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/v2/clients?user=alice&status=online&q=alice-host")
|
|
||||||
aliceResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.ClientInfoResp]]](t, resp)
|
|
||||||
if aliceResp.Data.Total != 1 || aliceResp.Data.Items[0].User != "alice" {
|
|
||||||
t.Fatalf("alice filter mismatch: %#v", aliceResp.Data)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/v2/clients?status=offline")
|
|
||||||
offlineResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.ClientInfoResp]]](t, resp)
|
|
||||||
if offlineResp.Data.Total != 1 || offlineResp.Data.Items[0].User != "bob" {
|
|
||||||
t.Fatalf("offline filter mismatch: %#v", offlineResp.Data)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2PageParamErrorsUseEnvelope(t *testing.T) {
|
|
||||||
controller := newV2TestController(t)
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/v2/clients?page=0")
|
|
||||||
if resp.Code != http.StatusBadRequest {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
|
|
||||||
}
|
|
||||||
errResp := decodeResponse[httppkg.V2Response](t, resp)
|
|
||||||
if errResp.Code != http.StatusBadRequest || errResp.Data != nil {
|
|
||||||
t.Fatalf("error envelope mismatch: %#v", errResp)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/v2/clients?pageSize=201")
|
|
||||||
if resp.Code != http.StatusBadRequest {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, fmt.Sprintf("/api/v2/clients?page=%d&pageSize=2", math.MaxInt))
|
|
||||||
if resp.Code != http.StatusBadRequest {
|
|
||||||
t.Fatalf("status mismatch for overflowing page offset, want %d got %d", http.StatusBadRequest, resp.Code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2ClientDetailEnvelope(t *testing.T) {
|
|
||||||
controller := newV2TestController(t)
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/v2/clients/alice.client-a")
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d", http.StatusOK, resp.Code)
|
|
||||||
}
|
|
||||||
detailResp := decodeResponse[v2EnvelopeForTest[model.V2ClientDetailResp]](t, resp)
|
|
||||||
if detailResp.Data.User != "alice" || detailResp.Data.ClientID != "client-a" {
|
|
||||||
t.Fatalf("client detail mismatch: %#v", detailResp.Data)
|
|
||||||
}
|
|
||||||
if detailResp.Data.Status.State != "online" || detailResp.Data.Status.CurConns != 5 || detailResp.Data.Status.ProxyCount != 2 {
|
|
||||||
t.Fatalf("client detail status mismatch: %#v", detailResp.Data.Status)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2ClientDetailEncodedKey(t *testing.T) {
|
|
||||||
oldStatsCollector := mem.StatsCollector
|
|
||||||
mem.StatsCollector = &fakeStatsCollector{
|
|
||||||
proxies: map[string]*mem.ProxyStats{
|
|
||||||
"tcp-url": {
|
|
||||||
Name: "tcp-url",
|
|
||||||
Type: "tcp",
|
|
||||||
User: "url",
|
|
||||||
ClientID: "client/a?b#c",
|
|
||||||
CurConns: 7,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
mem.StatsCollector = oldStatsCollector
|
|
||||||
})
|
|
||||||
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
clientRegistry.Register("url", "client/a?b#c", "run-url", "url-host", "1.0.0", "127.0.0.4", "v2")
|
|
||||||
controller := NewController(&v1.ServerConfig{}, clientRegistry, serverproxy.NewManager())
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
encodedKey := url.PathEscape("url.client/a?b#c")
|
|
||||||
resp := performRequest(router, "/api/v2/clients/"+encodedKey)
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("encoded client key status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
|
|
||||||
}
|
|
||||||
encodedResp := decodeResponse[v2EnvelopeForTest[model.V2ClientDetailResp]](t, resp)
|
|
||||||
if encodedResp.Data.User != "url" || encodedResp.Data.ClientID != "client/a?b#c" {
|
|
||||||
t.Fatalf("encoded client detail mismatch: %#v", encodedResp.Data)
|
|
||||||
}
|
|
||||||
if encodedResp.Data.Status.CurConns != 7 || encodedResp.Data.Status.ProxyCount != 1 {
|
|
||||||
t.Fatalf("encoded client detail status mismatch: %#v", encodedResp.Data.Status)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2ProxyListDetailAndUsers(t *testing.T) {
|
|
||||||
controller := newV2TestController(t)
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/v2/proxies?type=invalid")
|
|
||||||
if resp.Code != http.StatusBadRequest {
|
|
||||||
t.Fatalf("invalid proxy type status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
|
|
||||||
}
|
|
||||||
errResp := decodeResponse[httppkg.V2Response](t, resp)
|
|
||||||
if errResp.Code != http.StatusBadRequest || errResp.Data != nil {
|
|
||||||
t.Fatalf("invalid proxy type error envelope mismatch: %#v", errResp)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/v2/proxies?type=tcp&user=&page=1&pageSize=50")
|
|
||||||
proxyResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.V2ProxyResp]]](t, resp)
|
|
||||||
if proxyResp.Data.Total != 1 {
|
|
||||||
t.Fatalf("proxy filter total mismatch: %#v", proxyResp.Data)
|
|
||||||
}
|
|
||||||
proxyItem := proxyResp.Data.Items[0]
|
|
||||||
if proxyItem.Name != "tcp-empty" || proxyItem.Spec.Type != "tcp" || proxyItem.User != "" || proxyItem.Status.State != "offline" {
|
|
||||||
t.Fatalf("proxy item mismatch: %#v", proxyItem)
|
|
||||||
}
|
|
||||||
rawProxyResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[map[string]json.RawMessage]]](t, resp)
|
|
||||||
assertRawJSONKeys(t, rawProxyResp.Data.Items[0], "clientID", "name", "spec", "status", "user")
|
|
||||||
var rawListSpec map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(rawProxyResp.Data.Items[0]["spec"], &rawListSpec); err != nil {
|
|
||||||
t.Fatalf("unmarshal list proxy spec failed: %v", err)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, rawListSpec, "tcp", "type")
|
|
||||||
assertRawJSONKeysFromMessage(t, rawListSpec["tcp"])
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/v2/proxies/tcp-alice")
|
|
||||||
rawProxyDetailResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp)
|
|
||||||
assertRawJSONKeysFromMessage(t, rawProxyDetailResp.Data["status"],
|
|
||||||
"curConns",
|
|
||||||
"lastCloseAt",
|
|
||||||
"lastStartAt",
|
|
||||||
"phase",
|
|
||||||
"todayTrafficIn",
|
|
||||||
"todayTrafficOut",
|
|
||||||
)
|
|
||||||
proxyDetailResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyResp]](t, resp)
|
|
||||||
if proxyDetailResp.Data.Name != "tcp-alice" || proxyDetailResp.Data.User != "alice" {
|
|
||||||
t.Fatalf("proxy detail mismatch: %#v", proxyDetailResp.Data)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, rawProxyDetailResp.Data, "clientID", "name", "spec", "status", "user")
|
|
||||||
var rawDetailSpec map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(rawProxyDetailResp.Data["spec"], &rawDetailSpec); err != nil {
|
|
||||||
t.Fatalf("unmarshal detail proxy spec failed: %v", err)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, rawDetailSpec, "tcp", "type")
|
|
||||||
assertRawJSONKeysFromMessage(t, rawDetailSpec["tcp"])
|
|
||||||
if proxyDetailResp.Data.Status.LastStartAt != 1783504200 || proxyDetailResp.Data.Status.LastCloseAt != 1783504300 {
|
|
||||||
t.Fatalf("proxy detail timestamp mismatch: %#v", proxyDetailResp.Data.Status)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/v2/users?page=1&pageSize=50")
|
|
||||||
userResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.V2UserResp]]](t, resp)
|
|
||||||
if userResp.Data.Total != 3 {
|
|
||||||
t.Fatalf("user total mismatch: %#v", userResp.Data)
|
|
||||||
}
|
|
||||||
expectedProxyCounts := map[string]int{
|
|
||||||
"": 1,
|
|
||||||
"alice": 2,
|
|
||||||
"bob": 1,
|
|
||||||
}
|
|
||||||
for _, item := range userResp.Data.Items {
|
|
||||||
if item.ClientCount != 1 || item.ProxyCount != expectedProxyCounts[item.User] {
|
|
||||||
t.Fatalf("user counts mismatch: %#v", item)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2ProxyTrafficEnvelopeSchemaAndHistory(t *testing.T) {
|
|
||||||
oldStatsCollector := mem.StatsCollector
|
|
||||||
mem.StatsCollector = &fakeStatsCollector{
|
|
||||||
proxies: map[string]*mem.ProxyStats{
|
|
||||||
"ssh": {Name: "ssh", Type: "tcp"},
|
|
||||||
},
|
|
||||||
traffic: map[string]*mem.ProxyTrafficInfo{
|
|
||||||
"ssh": {
|
|
||||||
Name: "ssh",
|
|
||||||
TrafficIn: []int64{70, 60, 50, 40, 30, 20, 10},
|
|
||||||
TrafficOut: []int64{700, 600, 500, 400, 300, 200, 100},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
mem.StatsCollector = oldStatsCollector
|
|
||||||
})
|
|
||||||
|
|
||||||
controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager())
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/v2/proxies/ssh/traffic")
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
|
|
||||||
}
|
|
||||||
rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp)
|
|
||||||
if rawResp.Code != http.StatusOK || rawResp.Msg != "success" {
|
|
||||||
t.Fatalf("envelope mismatch: %#v", rawResp)
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, rawResp.Data, "granularity", "history", "name", "unit")
|
|
||||||
|
|
||||||
trafficResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyTrafficResp]](t, resp)
|
|
||||||
if trafficResp.Data.Name != "ssh" || trafficResp.Data.Unit != "bytes" || trafficResp.Data.Granularity != "day" {
|
|
||||||
t.Fatalf("traffic metadata mismatch: %#v", trafficResp.Data)
|
|
||||||
}
|
|
||||||
if len(trafficResp.Data.History) != 7 {
|
|
||||||
t.Fatalf("history length mismatch, want 7 got %d: %#v", len(trafficResp.Data.History), trafficResp.Data.History)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantIn := []int64{10, 20, 30, 40, 50, 60, 70}
|
|
||||||
wantOut := []int64{100, 200, 300, 400, 500, 600, 700}
|
|
||||||
var prevDate time.Time
|
|
||||||
for i, point := range trafficResp.Data.History {
|
|
||||||
assertRawJSONKeysFromMessage(t, mustMarshalJSON(t, point), "date", "trafficIn", "trafficOut")
|
|
||||||
if point.TrafficIn != wantIn[i] || point.TrafficOut != wantOut[i] {
|
|
||||||
t.Fatalf("history[%d] traffic mismatch: %#v", i, point)
|
|
||||||
}
|
|
||||||
parsedDate, err := time.Parse(time.DateOnly, point.Date)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("history[%d] date should be yyyy-mm-dd, got %q: %v", i, point.Date, err)
|
|
||||||
}
|
|
||||||
if i > 0 && !parsedDate.Equal(prevDate.AddDate(0, 0, 1)) {
|
|
||||||
t.Fatalf("history dates should be oldest to newest, got %s after %s", point.Date, prevDate.Format(time.DateOnly))
|
|
||||||
}
|
|
||||||
prevDate = parsedDate
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2ProxyTrafficNotFoundEnvelope(t *testing.T) {
|
|
||||||
controller := newV2TestController(t)
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/v2/proxies/missing/traffic")
|
|
||||||
if resp.Code != http.StatusNotFound {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusNotFound, resp.Code, resp.Body.String())
|
|
||||||
}
|
|
||||||
errResp := decodeResponse[httppkg.V2Response](t, resp)
|
|
||||||
if errResp.Code != http.StatusNotFound || errResp.Msg != "no proxy info found" || errResp.Data != nil {
|
|
||||||
t.Fatalf("not found envelope mismatch: %#v", errResp)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2ProxyDetailAndTrafficEncodedName(t *testing.T) {
|
|
||||||
name := "folder/ssh?x#y"
|
|
||||||
oldStatsCollector := mem.StatsCollector
|
|
||||||
mem.StatsCollector = &fakeStatsCollector{
|
|
||||||
proxies: map[string]*mem.ProxyStats{
|
|
||||||
name: {Name: name, Type: "tcp", User: "encoded"},
|
|
||||||
},
|
|
||||||
traffic: map[string]*mem.ProxyTrafficInfo{
|
|
||||||
name: {
|
|
||||||
Name: name,
|
|
||||||
TrafficIn: []int64{1},
|
|
||||||
TrafficOut: []int64{2},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
mem.StatsCollector = oldStatsCollector
|
|
||||||
})
|
|
||||||
|
|
||||||
controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager())
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
encodedName := url.PathEscape(name)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/v2/proxies/"+encodedName)
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("encoded proxy detail status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
|
|
||||||
}
|
|
||||||
detailResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyResp]](t, resp)
|
|
||||||
if detailResp.Data.Name != name || detailResp.Data.User != "encoded" {
|
|
||||||
t.Fatalf("encoded proxy detail mismatch: %#v", detailResp.Data)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/v2/proxies/"+encodedName+"/traffic")
|
|
||||||
if resp.Code != http.StatusOK {
|
|
||||||
t.Fatalf("encoded traffic status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
|
|
||||||
}
|
|
||||||
trafficResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyTrafficResp]](t, resp)
|
|
||||||
if trafficResp.Data.Name != name {
|
|
||||||
t.Fatalf("encoded traffic name mismatch: %#v", trafficResp.Data)
|
|
||||||
}
|
|
||||||
if got := trafficResp.Data.History[len(trafficResp.Data.History)-1]; got.TrafficIn != 1 || got.TrafficOut != 2 {
|
|
||||||
t.Fatalf("encoded traffic latest point mismatch: %#v", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAPIV2ProxyTrafficInvalidEncodedNameUses400Envelope(t *testing.T) {
|
|
||||||
controller := newV2TestController(t)
|
|
||||||
handler := httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyTraffic)
|
|
||||||
req := httptest.NewRequest(http.MethodGet, "/api/v2/proxies/%25ZZ/traffic", nil)
|
|
||||||
req = mux.SetURLVars(req, map[string]string{"name": "%ZZ"})
|
|
||||||
resp := httptest.NewRecorder()
|
|
||||||
handler.ServeHTTP(resp, req)
|
|
||||||
|
|
||||||
if resp.Code != http.StatusBadRequest {
|
|
||||||
t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusBadRequest, resp.Code, resp.Body.String())
|
|
||||||
}
|
|
||||||
errResp := decodeResponse[httppkg.V2Response](t, resp)
|
|
||||||
if errResp.Code != http.StatusBadRequest || errResp.Msg != "invalid proxy name" || errResp.Data != nil {
|
|
||||||
t.Fatalf("invalid encoded name envelope mismatch: %#v", errResp)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMatchV2ProxyQueryMatchesSpecFields(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
item model.V2ProxyResp
|
|
||||||
q string
|
|
||||||
want bool
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "tcp remote port",
|
|
||||||
item: model.V2ProxyResp{Name: "tcp-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "tcp",
|
|
||||||
TCP: &model.V2TCPProxySpec{RemotePort: v2TestIntPtr(6000)},
|
|
||||||
}},
|
|
||||||
q: "6000",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "udp remote port",
|
|
||||||
item: model.V2ProxyResp{Name: "udp-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "udp",
|
|
||||||
UDP: &model.V2UDPProxySpec{RemotePort: v2TestIntPtr(7000)},
|
|
||||||
}},
|
|
||||||
q: "7000",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "remote port does not match colon form",
|
|
||||||
item: model.V2ProxyResp{Name: "tcp-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "tcp",
|
|
||||||
TCP: &model.V2TCPProxySpec{RemotePort: v2TestIntPtr(6000)},
|
|
||||||
}},
|
|
||||||
q: ":6000",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "http custom domain",
|
|
||||||
item: model.V2ProxyResp{Name: "http-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "http",
|
|
||||||
HTTP: &model.V2HTTPProxySpec{CustomDomains: []string{"app.example.com"}},
|
|
||||||
}},
|
|
||||||
q: "app.example.com",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "https subdomain",
|
|
||||||
item: model.V2ProxyResp{Name: "https-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "https",
|
|
||||||
HTTPS: &model.V2HTTPSProxySpec{Subdomain: "portal"},
|
|
||||||
}},
|
|
||||||
q: "portal",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "subdomain does not match expanded host",
|
|
||||||
item: model.V2ProxyResp{Name: "https-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "https",
|
|
||||||
HTTPS: &model.V2HTTPSProxySpec{Subdomain: "portal"},
|
|
||||||
}},
|
|
||||||
q: "portal.example.com",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "tcpmux custom domain",
|
|
||||||
item: model.V2ProxyResp{Name: "tcpmux-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "tcpmux",
|
|
||||||
TCPMux: &model.V2TCPMuxProxySpec{CustomDomains: []string{"mux.example.com"}},
|
|
||||||
}},
|
|
||||||
q: "mux.example.com",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "offline shell does not match online spec fields",
|
|
||||||
item: model.V2ProxyResp{Name: "offline-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "tcp",
|
|
||||||
TCP: &model.V2TCPProxySpec{},
|
|
||||||
}},
|
|
||||||
q: "6000",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "offline shell does not contribute zero remote port",
|
|
||||||
item: model.V2ProxyResp{Name: "offline-proxy", Spec: model.V2ProxySpec{
|
|
||||||
Type: "tcp",
|
|
||||||
TCP: &model.V2TCPProxySpec{},
|
|
||||||
}},
|
|
||||||
q: "0",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
if got := matchV2ProxyQuery(tt.item, tt.q); got != tt.want {
|
|
||||||
t.Fatalf("matchV2ProxyQuery() = %v, want %v", got, tt.want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLegacyAPIResponsesRemainBare(t *testing.T) {
|
|
||||||
controller := newV2TestController(t)
|
|
||||||
router := newV2TestRouter(controller)
|
|
||||||
|
|
||||||
resp := performRequest(router, "/api/serverinfo")
|
|
||||||
var serverInfo model.ServerInfoResp
|
|
||||||
if err := json.Unmarshal(resp.Body.Bytes(), &serverInfo); err != nil {
|
|
||||||
t.Fatalf("legacy serverinfo should be a bare object: %v, body: %s", err, resp.Body.String())
|
|
||||||
}
|
|
||||||
if serverInfo.Version == "" {
|
|
||||||
t.Fatal("legacy serverinfo version should be set")
|
|
||||||
}
|
|
||||||
var serverInfoRaw map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(resp.Body.Bytes(), &serverInfoRaw); err != nil {
|
|
||||||
t.Fatalf("unmarshal legacy serverinfo object failed: %v", err)
|
|
||||||
}
|
|
||||||
if _, ok := serverInfoRaw["data"]; ok {
|
|
||||||
t.Fatalf("legacy serverinfo should not use v2 envelope: %s", resp.Body.String())
|
|
||||||
}
|
|
||||||
if _, ok := serverInfoRaw["config"]; ok {
|
|
||||||
t.Fatalf("legacy serverinfo should stay flat, got config in: %s", resp.Body.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/clients")
|
|
||||||
var clients []model.ClientInfoResp
|
|
||||||
if err := json.Unmarshal(resp.Body.Bytes(), &clients); err != nil {
|
|
||||||
t.Fatalf("legacy clients should be a bare array: %v, body: %s", err, resp.Body.String())
|
|
||||||
}
|
|
||||||
if len(clients) != 3 {
|
|
||||||
t.Fatalf("legacy clients total mismatch, want 3 got %d", len(clients))
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/proxy/tcp")
|
|
||||||
var proxies model.GetProxyInfoResp
|
|
||||||
if err := json.Unmarshal(resp.Body.Bytes(), &proxies); err != nil {
|
|
||||||
t.Fatalf("legacy proxy response should be {proxies}: %v, body: %s", err, resp.Body.String())
|
|
||||||
}
|
|
||||||
if len(proxies.Proxies) != 2 {
|
|
||||||
t.Fatalf("legacy tcp proxy total mismatch, want 2 got %d", len(proxies.Proxies))
|
|
||||||
}
|
|
||||||
var envelope httppkg.V2Response
|
|
||||||
if err := json.Unmarshal(resp.Body.Bytes(), &envelope); err == nil && envelope.Code != 0 {
|
|
||||||
t.Fatalf("legacy proxy response should not use v2 envelope: %#v", envelope)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp = performRequest(router, "/api/traffic/tcp-alice")
|
|
||||||
var traffic model.GetProxyTrafficResp
|
|
||||||
if err := json.Unmarshal(resp.Body.Bytes(), &traffic); err != nil {
|
|
||||||
t.Fatalf("legacy traffic should be a bare object: %v, body: %s", err, resp.Body.String())
|
|
||||||
}
|
|
||||||
if traffic.Name != "tcp-alice" ||
|
|
||||||
len(traffic.TrafficIn) != 2 || traffic.TrafficIn[0] != 7 || traffic.TrafficIn[1] != 6 ||
|
|
||||||
len(traffic.TrafficOut) != 2 || traffic.TrafficOut[0] != 70 || traffic.TrafficOut[1] != 60 {
|
|
||||||
t.Fatalf("legacy traffic should preserve today-first arrays, got: %#v", traffic)
|
|
||||||
}
|
|
||||||
var trafficRaw map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(resp.Body.Bytes(), &trafficRaw); err != nil {
|
|
||||||
t.Fatalf("unmarshal legacy traffic object failed: %v", err)
|
|
||||||
}
|
|
||||||
if _, ok := trafficRaw["data"]; ok {
|
|
||||||
t.Fatalf("legacy traffic should not use v2 envelope: %s", resp.Body.String())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func v2TestIntPtr(value int) *int {
|
|
||||||
return &value
|
|
||||||
}
|
|
||||||
|
|
||||||
func newV2TestController(t *testing.T) *Controller {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
oldStatsCollector := mem.StatsCollector
|
|
||||||
mem.StatsCollector = &fakeStatsCollector{
|
|
||||||
proxies: map[string]*mem.ProxyStats{
|
|
||||||
"tcp-empty": {
|
|
||||||
Name: "tcp-empty",
|
|
||||||
Type: "tcp",
|
|
||||||
User: "",
|
|
||||||
ClientID: "legacy-client",
|
|
||||||
TodayTrafficIn: 10,
|
|
||||||
TodayTrafficOut: 20,
|
|
||||||
CurConns: 1,
|
|
||||||
},
|
|
||||||
"tcp-alice": {
|
|
||||||
Name: "tcp-alice",
|
|
||||||
Type: "tcp",
|
|
||||||
User: "alice",
|
|
||||||
ClientID: "client-a",
|
|
||||||
TodayTrafficIn: 30,
|
|
||||||
TodayTrafficOut: 40,
|
|
||||||
CurConns: 2,
|
|
||||||
LastStartTime: "07-08 12:30:00",
|
|
||||||
LastCloseTime: "07-08 12:31:40",
|
|
||||||
LastStartAt: 1783504200,
|
|
||||||
LastCloseAt: 1783504300,
|
|
||||||
},
|
|
||||||
"http-alice": {
|
|
||||||
Name: "http-alice",
|
|
||||||
Type: "http",
|
|
||||||
User: "alice",
|
|
||||||
ClientID: "client-a",
|
|
||||||
CurConns: 3,
|
|
||||||
},
|
|
||||||
"udp-bob": {
|
|
||||||
Name: "udp-bob",
|
|
||||||
Type: "udp",
|
|
||||||
User: "bob",
|
|
||||||
ClientID: "client-b",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
traffic: map[string]*mem.ProxyTrafficInfo{
|
|
||||||
"tcp-alice": {
|
|
||||||
Name: "tcp-alice",
|
|
||||||
TrafficIn: []int64{7, 6},
|
|
||||||
TrafficOut: []int64{70, 60},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
mem.StatsCollector = oldStatsCollector
|
|
||||||
})
|
|
||||||
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
clientRegistry.Register("", "legacy-client", "run-empty", "empty-host", "1.0.0", "127.0.0.1", "v1")
|
|
||||||
clientRegistry.Register("alice", "client-a", "run-a", "alice-host", "1.0.0", "127.0.0.2", "v2")
|
|
||||||
clientRegistry.Register("bob", "client-b", "run-b", "bob-host", "1.0.0", "127.0.0.3", "v1")
|
|
||||||
clientRegistry.MarkOfflineByRunID("run-b")
|
|
||||||
|
|
||||||
return NewController(&v1.ServerConfig{}, clientRegistry, serverproxy.NewManager())
|
|
||||||
}
|
|
||||||
|
|
||||||
func newV2TestRouter(controller *Controller) *mux.Router {
|
|
||||||
router := mux.NewRouter()
|
|
||||||
router.HandleFunc("/api/v2/users", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2UserList)).Methods(http.MethodGet)
|
|
||||||
router.HandleFunc("/api/v2/system/info", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2SystemInfo)).Methods(http.MethodGet)
|
|
||||||
router.HandleFunc("/api/v2/system/prune", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2SystemPrune)).Methods(http.MethodPost)
|
|
||||||
router.HandleFunc("/api/v2/clients", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ClientList)).Methods(http.MethodGet)
|
|
||||||
encodedPathRouter := router.NewRoute().Subrouter()
|
|
||||||
encodedPathRouter.UseEncodedPath()
|
|
||||||
encodedPathRouter.HandleFunc("/api/v2/clients/{key}", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ClientDetail)).Methods(http.MethodGet)
|
|
||||||
router.HandleFunc("/api/v2/proxies", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyList)).Methods(http.MethodGet)
|
|
||||||
encodedPathRouter.HandleFunc("/api/v2/proxies/{name}/traffic", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyTraffic)).Methods(http.MethodGet)
|
|
||||||
encodedPathRouter.HandleFunc("/api/v2/proxies/{name}", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyDetail)).Methods(http.MethodGet)
|
|
||||||
router.HandleFunc("/api/serverinfo", httppkg.MakeHTTPHandlerFunc(controller.APIServerInfo)).Methods(http.MethodGet)
|
|
||||||
router.HandleFunc("/api/clients", httppkg.MakeHTTPHandlerFunc(controller.APIClientList)).Methods(http.MethodGet)
|
|
||||||
router.HandleFunc("/api/proxy/{type}", httppkg.MakeHTTPHandlerFunc(controller.APIProxyByType)).Methods(http.MethodGet)
|
|
||||||
router.HandleFunc("/api/traffic/{name}", httppkg.MakeHTTPHandlerFunc(controller.APIProxyTraffic)).Methods(http.MethodGet)
|
|
||||||
return router
|
|
||||||
}
|
|
||||||
|
|
||||||
func performRequest(handler http.Handler, target string) *httptest.ResponseRecorder {
|
|
||||||
return performRequestWithMethod(handler, http.MethodGet, target)
|
|
||||||
}
|
|
||||||
|
|
||||||
func performRequestWithMethod(handler http.Handler, method, target string) *httptest.ResponseRecorder {
|
|
||||||
req := httptest.NewRequest(method, target, nil)
|
|
||||||
resp := httptest.NewRecorder()
|
|
||||||
handler.ServeHTTP(resp, req)
|
|
||||||
return resp
|
|
||||||
}
|
|
||||||
|
|
||||||
func decodeResponse[T any](t *testing.T, resp *httptest.ResponseRecorder) T {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var out T
|
|
||||||
if err := json.Unmarshal(resp.Body.Bytes(), &out); err != nil {
|
|
||||||
t.Fatalf("unmarshal response failed: %v, body: %s", err, resp.Body.String())
|
|
||||||
}
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func assertRawJSONKeys(t *testing.T, raw map[string]json.RawMessage, want ...string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
if len(raw) != len(want) {
|
|
||||||
t.Fatalf("json keys mismatch, want %v got %v", want, raw)
|
|
||||||
}
|
|
||||||
for _, key := range want {
|
|
||||||
if _, ok := raw[key]; !ok {
|
|
||||||
t.Fatalf("json key %q missing from %v", key, raw)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func assertRawJSONKeysFromMessage(t *testing.T, raw json.RawMessage, want ...string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var out map[string]json.RawMessage
|
|
||||||
if err := json.Unmarshal(raw, &out); err != nil {
|
|
||||||
t.Fatalf("unmarshal raw json object failed: %v, body: %s", err, string(raw))
|
|
||||||
}
|
|
||||||
assertRawJSONKeys(t, out, want...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func mustMarshalJSON(t *testing.T, value any) json.RawMessage {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
out, err := json.Marshal(value)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("marshal json failed: %v", err)
|
|
||||||
}
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
@@ -1,179 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package model
|
|
||||||
|
|
||||||
type V2PageResp[T any] struct {
|
|
||||||
Total int `json:"total"`
|
|
||||||
Page int `json:"page"`
|
|
||||||
PageSize int `json:"pageSize"`
|
|
||||||
Items []T `json:"items"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2SystemInfoResp struct {
|
|
||||||
Version string `json:"version"`
|
|
||||||
Config V2SystemInfoConfigResp `json:"config"`
|
|
||||||
Status V2SystemInfoStatusResp `json:"status"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2SystemInfoConfigResp struct {
|
|
||||||
BindPort int `json:"bindPort"`
|
|
||||||
VhostHTTPPort int `json:"vhostHTTPPort"`
|
|
||||||
VhostHTTPSPort int `json:"vhostHTTPSPort"`
|
|
||||||
TCPMuxHTTPConnectPort int `json:"tcpmuxHTTPConnectPort"`
|
|
||||||
KCPBindPort int `json:"kcpBindPort"`
|
|
||||||
QUICBindPort int `json:"quicBindPort"`
|
|
||||||
SubdomainHost string `json:"subdomainHost"`
|
|
||||||
MaxPoolCount int64 `json:"maxPoolCount"`
|
|
||||||
MaxPortsPerClient int64 `json:"maxPortsPerClient"`
|
|
||||||
HeartbeatTimeout int64 `json:"heartbeatTimeout"`
|
|
||||||
AllowPortsStr string `json:"allowPortsStr"`
|
|
||||||
TLSForce bool `json:"tlsForce"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2SystemInfoStatusResp struct {
|
|
||||||
TotalTrafficIn int64 `json:"totalTrafficIn"`
|
|
||||||
TotalTrafficOut int64 `json:"totalTrafficOut"`
|
|
||||||
CurConns int64 `json:"curConns"`
|
|
||||||
ClientCounts int64 `json:"clientCounts"`
|
|
||||||
ProxyTypeCounts map[string]int64 `json:"proxyTypeCount"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2SystemPruneResp struct {
|
|
||||||
Type string `json:"type"`
|
|
||||||
Cleared int `json:"cleared"`
|
|
||||||
Total int `json:"total"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2UserResp struct {
|
|
||||||
User string `json:"user"`
|
|
||||||
ClientCount int `json:"clientCount"`
|
|
||||||
ProxyCount int `json:"proxyCount"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ClientDetailResp struct {
|
|
||||||
ClientInfoResp
|
|
||||||
Status V2ClientStatusResp `json:"status"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ClientStatusResp struct {
|
|
||||||
State string `json:"phase"`
|
|
||||||
CurConns int64 `json:"curConns"`
|
|
||||||
ProxyCount int64 `json:"proxyCount"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ProxyResp struct {
|
|
||||||
Name string `json:"name"`
|
|
||||||
User string `json:"user"`
|
|
||||||
ClientID string `json:"clientID"`
|
|
||||||
Spec V2ProxySpec `json:"spec"`
|
|
||||||
Status V2ProxyStatusResp `json:"status"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ProxySpec struct {
|
|
||||||
Type string `json:"type"`
|
|
||||||
|
|
||||||
TCP *V2TCPProxySpec `json:"tcp,omitempty"`
|
|
||||||
UDP *V2UDPProxySpec `json:"udp,omitempty"`
|
|
||||||
HTTP *V2HTTPProxySpec `json:"http,omitempty"`
|
|
||||||
HTTPS *V2HTTPSProxySpec `json:"https,omitempty"`
|
|
||||||
TCPMux *V2TCPMuxProxySpec `json:"tcpmux,omitempty"`
|
|
||||||
STCP *V2STCPProxySpec `json:"stcp,omitempty"`
|
|
||||||
SUDP *V2SUDPProxySpec `json:"sudp,omitempty"`
|
|
||||||
XTCP *V2XTCPProxySpec `json:"xtcp,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ProxyBaseSpec struct {
|
|
||||||
Annotations map[string]string `json:"annotations,omitempty"`
|
|
||||||
Metadatas map[string]string `json:"metadatas,omitempty"`
|
|
||||||
Transport *V2ProxyTransportSpec `json:"transport,omitempty"`
|
|
||||||
LoadBalancer *V2ProxyLoadBalancerSpec `json:"loadBalancer,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ProxyTransportSpec struct {
|
|
||||||
UseEncryption bool `json:"useEncryption"`
|
|
||||||
UseCompression bool `json:"useCompression"`
|
|
||||||
BandwidthLimit string `json:"bandwidthLimit"`
|
|
||||||
BandwidthLimitMode string `json:"bandwidthLimitMode"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ProxyLoadBalancerSpec struct {
|
|
||||||
Group string `json:"group"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2TCPProxySpec struct {
|
|
||||||
V2ProxyBaseSpec
|
|
||||||
RemotePort *int `json:"remotePort,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2UDPProxySpec struct {
|
|
||||||
V2ProxyBaseSpec
|
|
||||||
RemotePort *int `json:"remotePort,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2HTTPProxySpec struct {
|
|
||||||
V2ProxyBaseSpec
|
|
||||||
CustomDomains []string `json:"customDomains,omitempty"`
|
|
||||||
Subdomain string `json:"subdomain,omitempty"`
|
|
||||||
Locations []string `json:"locations,omitempty"`
|
|
||||||
HostHeaderRewrite string `json:"hostHeaderRewrite,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2HTTPSProxySpec struct {
|
|
||||||
V2ProxyBaseSpec
|
|
||||||
CustomDomains []string `json:"customDomains,omitempty"`
|
|
||||||
Subdomain string `json:"subdomain,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2TCPMuxProxySpec struct {
|
|
||||||
V2ProxyBaseSpec
|
|
||||||
CustomDomains []string `json:"customDomains,omitempty"`
|
|
||||||
Subdomain string `json:"subdomain,omitempty"`
|
|
||||||
Multiplexer string `json:"multiplexer,omitempty"`
|
|
||||||
RouteByHTTPUser string `json:"routeByHTTPUser,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2STCPProxySpec struct {
|
|
||||||
V2ProxyBaseSpec
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2SUDPProxySpec struct {
|
|
||||||
V2ProxyBaseSpec
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2XTCPProxySpec struct {
|
|
||||||
V2ProxyBaseSpec
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ProxyStatusResp struct {
|
|
||||||
State string `json:"phase"`
|
|
||||||
TodayTrafficIn int64 `json:"todayTrafficIn"`
|
|
||||||
TodayTrafficOut int64 `json:"todayTrafficOut"`
|
|
||||||
CurConns int64 `json:"curConns"`
|
|
||||||
LastStartAt int64 `json:"lastStartAt,omitempty"`
|
|
||||||
LastCloseAt int64 `json:"lastCloseAt,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ProxyTrafficResp struct {
|
|
||||||
Name string `json:"name"`
|
|
||||||
Unit string `json:"unit"`
|
|
||||||
Granularity string `json:"granularity"`
|
|
||||||
History []V2ProxyTrafficPointResp `json:"history"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type V2ProxyTrafficPointResp struct {
|
|
||||||
Date string `json:"date"`
|
|
||||||
TrafficIn int64 `json:"trafficIn"`
|
|
||||||
TrafficOut int64 `json:"trafficOut"`
|
|
||||||
}
|
|
||||||
+34
-93
@@ -82,20 +82,19 @@ type Proxy interface {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type BaseProxy struct {
|
type BaseProxy struct {
|
||||||
name string
|
name string
|
||||||
rc *controller.ResourceController
|
rc *controller.ResourceController
|
||||||
listeners []net.Listener
|
listeners []net.Listener
|
||||||
usedPortsNum int
|
usedPortsNum int
|
||||||
poolCount int
|
poolCount int
|
||||||
getWorkConnFn GetWorkConnFn
|
getWorkConnFn GetWorkConnFn
|
||||||
serverCfg *v1.ServerConfig
|
serverCfg *v1.ServerConfig
|
||||||
encryptionKey []byte
|
encryptionKey []byte
|
||||||
limiter *rate.Limiter
|
limiter *rate.Limiter
|
||||||
userInfo plugin.UserInfo
|
userInfo plugin.UserInfo
|
||||||
loginMsg *msg.Login
|
loginMsg *msg.Login
|
||||||
configurer v1.ProxyConfigurer
|
configurer v1.ProxyConfigurer
|
||||||
wireProtocol string
|
wireProtocol string
|
||||||
udpPacketCodec string
|
|
||||||
|
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
xl *xlog.Logger
|
xl *xlog.Logger
|
||||||
@@ -328,18 +327,10 @@ func (pxy *BaseProxy) handleUserTCPConnection(userConn net.Conn) {
|
|||||||
|
|
||||||
func (pxy *BaseProxy) joinUserConnection(local io.ReadWriteCloser, userConn net.Conn, proxyType string, xl *xlog.Logger) (int64, int64, []error) {
|
func (pxy *BaseProxy) joinUserConnection(local io.ReadWriteCloser, userConn net.Conn, proxyType string, xl *xlog.Logger) (int64, int64, []error) {
|
||||||
visitorWireProtocol := wireProtocolFromConn(userConn)
|
visitorWireProtocol := wireProtocolFromConn(userConn)
|
||||||
visitorUDPPacketCodec := udpPacketCodecFromConn(userConn)
|
if proxyType == string(v1.ProxyTypeSUDP) && isMixedWireProtocol(pxy.wireProtocol, visitorWireProtocol) {
|
||||||
if proxyType == string(v1.ProxyTypeSUDP) {
|
xl.Infof("bridge mixed SUDP payload codecs, proxy wireProtocol [%s], visitor wireProtocol [%s]",
|
||||||
mixed, err := isMixedSUDPPacketEncoding(pxy.wireProtocol, pxy.udpPacketCodec, visitorWireProtocol, visitorUDPPacketCodec)
|
normalizeWireProtocol(pxy.wireProtocol), normalizeWireProtocol(visitorWireProtocol))
|
||||||
if err != nil {
|
return joinSUDPMessageBridge(local, userConn, pxy.wireProtocol, visitorWireProtocol, xl)
|
||||||
return 0, 0, []error{err}
|
|
||||||
}
|
|
||||||
if mixed {
|
|
||||||
xl.Infof("bridge mixed SUDP payload codecs, proxy [%s/%s], visitor [%s/%s]",
|
|
||||||
normalizeWireProtocol(pxy.wireProtocol), pxy.udpPacketCodec,
|
|
||||||
normalizeWireProtocol(visitorWireProtocol), visitorUDPPacketCodec)
|
|
||||||
return joinSUDPMessageBridge(local, userConn, pxy.wireProtocol, pxy.udpPacketCodec, visitorWireProtocol, visitorUDPPacketCodec, xl)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return libio.Join(local, userConn)
|
return libio.Join(local, userConn)
|
||||||
}
|
}
|
||||||
@@ -348,10 +339,6 @@ type wireProtocolGetter interface {
|
|||||||
WireProtocol() string
|
WireProtocol() string
|
||||||
}
|
}
|
||||||
|
|
||||||
type udpPacketCodecGetter interface {
|
|
||||||
UDPPacketCodec() string
|
|
||||||
}
|
|
||||||
|
|
||||||
func wireProtocolFromConn(conn net.Conn) string {
|
func wireProtocolFromConn(conn net.Conn) string {
|
||||||
if getter, ok := conn.(wireProtocolGetter); ok {
|
if getter, ok := conn.(wireProtocolGetter); ok {
|
||||||
return getter.WireProtocol()
|
return getter.WireProtocol()
|
||||||
@@ -359,46 +346,10 @@ func wireProtocolFromConn(conn net.Conn) string {
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
func udpPacketCodecFromConn(conn net.Conn) string {
|
|
||||||
if getter, ok := conn.(udpPacketCodecGetter); ok {
|
|
||||||
return getter.UDPPacketCodec()
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func isMixedWireProtocol(left, right string) bool {
|
func isMixedWireProtocol(left, right string) bool {
|
||||||
return normalizeWireProtocol(left) != normalizeWireProtocol(right)
|
return normalizeWireProtocol(left) != normalizeWireProtocol(right)
|
||||||
}
|
}
|
||||||
|
|
||||||
func isMixedSUDPPacketEncoding(leftWire, leftCodec, rightWire, rightCodec string) (bool, error) {
|
|
||||||
leftCodec, err := normalizeUDPPacketCodec(leftWire, leftCodec)
|
|
||||||
if err != nil {
|
|
||||||
return false, fmt.Errorf("invalid left SUDP packet encoding: %w", err)
|
|
||||||
}
|
|
||||||
rightCodec, err = normalizeUDPPacketCodec(rightWire, rightCodec)
|
|
||||||
if err != nil {
|
|
||||||
return false, fmt.Errorf("invalid right SUDP packet encoding: %w", err)
|
|
||||||
}
|
|
||||||
return normalizeWireProtocol(leftWire) != normalizeWireProtocol(rightWire) || leftCodec != rightCodec, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func normalizeUDPPacketCodec(wireProtocol, codec string) (string, error) {
|
|
||||||
switch wireProtocol {
|
|
||||||
case "", wire.ProtocolV1:
|
|
||||||
if codec != "" {
|
|
||||||
return "", fmt.Errorf("UDP packet codec %q requires wire protocol v2", codec)
|
|
||||||
}
|
|
||||||
return "", nil
|
|
||||||
case wire.ProtocolV2:
|
|
||||||
if codec == "" || codec == wire.UDPPacketCodecBinary {
|
|
||||||
return codec, nil
|
|
||||||
}
|
|
||||||
return "", fmt.Errorf("unsupported UDP packet codec %q", codec)
|
|
||||||
default:
|
|
||||||
return "", fmt.Errorf("unsupported wire protocol %q", wireProtocol)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func normalizeWireProtocol(wireProtocol string) string {
|
func normalizeWireProtocol(wireProtocol string) string {
|
||||||
if wireProtocol == wire.ProtocolV2 {
|
if wireProtocol == wire.ProtocolV2 {
|
||||||
return wire.ProtocolV2
|
return wire.ProtocolV2
|
||||||
@@ -410,21 +361,13 @@ func joinSUDPMessageBridge(
|
|||||||
proxyConn io.ReadWriteCloser,
|
proxyConn io.ReadWriteCloser,
|
||||||
visitorConn io.ReadWriteCloser,
|
visitorConn io.ReadWriteCloser,
|
||||||
proxyWireProtocol string,
|
proxyWireProtocol string,
|
||||||
proxyUDPPacketCodec string,
|
|
||||||
visitorWireProtocol string,
|
visitorWireProtocol string,
|
||||||
visitorUDPPacketCodec string,
|
|
||||||
xl *xlog.Logger,
|
xl *xlog.Logger,
|
||||||
) (inCount int64, outCount int64, errs []error) {
|
) (inCount int64, outCount int64, errs []error) {
|
||||||
// The mixed bridge decodes and re-encodes messages, so raw framed byte counts
|
// The mixed bridge decodes and re-encodes messages, so raw framed byte counts
|
||||||
// are not available. Count UDP payload bytes and ignore heartbeat traffic.
|
// are not available. Count UDP payload bytes and ignore heartbeat traffic.
|
||||||
proxyRW, err := msg.NewUDPPacketReadWriter(proxyConn, proxyWireProtocol, proxyUDPPacketCodec)
|
proxyRW := msg.NewReadWriter(proxyConn, proxyWireProtocol)
|
||||||
if err != nil {
|
visitorRW := msg.NewReadWriter(visitorConn, visitorWireProtocol)
|
||||||
return 0, 0, []error{err}
|
|
||||||
}
|
|
||||||
visitorRW, err := msg.NewUDPPacketReadWriter(visitorConn, visitorWireProtocol, visitorUDPPacketCodec)
|
|
||||||
if err != nil {
|
|
||||||
return 0, 0, []error{err}
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
var (
|
||||||
once sync.Once
|
once sync.Once
|
||||||
@@ -526,7 +469,6 @@ type Options struct {
|
|||||||
ServerCfg *v1.ServerConfig
|
ServerCfg *v1.ServerConfig
|
||||||
EncryptionKey []byte
|
EncryptionKey []byte
|
||||||
WireProtocol string
|
WireProtocol string
|
||||||
UDPPacketCodec string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) {
|
func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) {
|
||||||
@@ -536,25 +478,24 @@ func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) {
|
|||||||
var limiter *rate.Limiter
|
var limiter *rate.Limiter
|
||||||
limitBytes := configurer.GetBaseConfig().Transport.BandwidthLimit.Bytes()
|
limitBytes := configurer.GetBaseConfig().Transport.BandwidthLimit.Bytes()
|
||||||
if limitBytes > 0 && configurer.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeServer {
|
if limitBytes > 0 && configurer.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeServer {
|
||||||
limiter = limit.NewBandwidthLimiter(limitBytes)
|
limiter = rate.NewLimiter(rate.Limit(float64(limitBytes)), int(limitBytes))
|
||||||
}
|
}
|
||||||
|
|
||||||
basePxy := BaseProxy{
|
basePxy := BaseProxy{
|
||||||
name: configurer.GetBaseConfig().Name,
|
name: configurer.GetBaseConfig().Name,
|
||||||
rc: options.ResourceController,
|
rc: options.ResourceController,
|
||||||
listeners: make([]net.Listener, 0),
|
listeners: make([]net.Listener, 0),
|
||||||
poolCount: options.PoolCount,
|
poolCount: options.PoolCount,
|
||||||
getWorkConnFn: options.GetWorkConnFn,
|
getWorkConnFn: options.GetWorkConnFn,
|
||||||
serverCfg: options.ServerCfg,
|
serverCfg: options.ServerCfg,
|
||||||
encryptionKey: options.EncryptionKey,
|
encryptionKey: options.EncryptionKey,
|
||||||
limiter: limiter,
|
limiter: limiter,
|
||||||
xl: xl,
|
xl: xl,
|
||||||
ctx: xlog.NewContext(ctx, xl),
|
ctx: xlog.NewContext(ctx, xl),
|
||||||
userInfo: options.UserInfo,
|
userInfo: options.UserInfo,
|
||||||
loginMsg: options.LoginMsg,
|
loginMsg: options.LoginMsg,
|
||||||
configurer: configurer,
|
configurer: configurer,
|
||||||
wireProtocol: options.WireProtocol,
|
wireProtocol: options.WireProtocol,
|
||||||
udpPacketCodec: options.UDPPacketCodec,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
factory := proxyFactoryRegistry[reflect.TypeOf(configurer)]
|
factory := proxyFactoryRegistry[reflect.TypeOf(configurer)]
|
||||||
|
|||||||
@@ -1,232 +0,0 @@
|
|||||||
// Copyright 2026 The frp Authors
|
|
||||||
//
|
|
||||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
// you may not use this file except in compliance with the License.
|
|
||||||
// You may obtain a copy of the License at
|
|
||||||
//
|
|
||||||
// http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
//
|
|
||||||
// Unless required by applicable law or agreed to in writing, software
|
|
||||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
// See the License for the specific language governing permissions and
|
|
||||||
// limitations under the License.
|
|
||||||
|
|
||||||
package proxy
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
|
||||||
"github.com/fatedier/frp/pkg/proto/wire"
|
|
||||||
)
|
|
||||||
|
|
||||||
type sudpPathBenchmarkCase struct {
|
|
||||||
name string
|
|
||||||
packet *msg.UDPPacket
|
|
||||||
}
|
|
||||||
|
|
||||||
var sudpPathBenchmarkBytesSink []byte
|
|
||||||
|
|
||||||
func sudpPathBenchmarkCases(payloadSize int) []sudpPathBenchmarkCase {
|
|
||||||
content := bytes.Repeat([]byte{0x5a}, payloadSize)
|
|
||||||
return []sudpPathBenchmarkCase{
|
|
||||||
{
|
|
||||||
name: "ipv4-remote",
|
|
||||||
packet: &msg.UDPPacket{
|
|
||||||
Content: content,
|
|
||||||
RemoteAddr: &net.UDPAddr{
|
|
||||||
IP: net.ParseIP("192.0.2.1"), Port: 12345,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "ipv4-local-remote",
|
|
||||||
packet: &msg.UDPPacket{
|
|
||||||
Content: content,
|
|
||||||
LocalAddr: &net.UDPAddr{
|
|
||||||
IP: net.ParseIP("192.0.2.2"), Port: 23456,
|
|
||||||
},
|
|
||||||
RemoteAddr: &net.UDPAddr{
|
|
||||||
IP: net.ParseIP("192.0.2.1"), Port: 12345,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "ipv6-remote",
|
|
||||||
packet: &msg.UDPPacket{
|
|
||||||
Content: content,
|
|
||||||
RemoteAddr: &net.UDPAddr{
|
|
||||||
IP: net.ParseIP("2001:db8::1"), Port: 12345,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "ipv6-local-remote",
|
|
||||||
packet: &msg.UDPPacket{
|
|
||||||
Content: content,
|
|
||||||
LocalAddr: &net.UDPAddr{
|
|
||||||
IP: net.ParseIP("2001:db8::2"), Port: 23456, Zone: "bench0",
|
|
||||||
},
|
|
||||||
RemoteAddr: &net.UDPAddr{
|
|
||||||
IP: net.ParseIP("2001:db8::1"), Port: 12345, Zone: "bench1",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type sudpPathReadWriter struct {
|
|
||||||
reader bytes.Reader
|
|
||||||
}
|
|
||||||
|
|
||||||
func (rw *sudpPathReadWriter) Read(p []byte) (int, error) { return rw.reader.Read(p) }
|
|
||||||
func (rw *sudpPathReadWriter) Write(p []byte) (int, error) { return len(p), nil }
|
|
||||||
func (rw *sudpPathReadWriter) Reset(p []byte) { rw.reader.Reset(p) }
|
|
||||||
|
|
||||||
func sudpPathWireBytes(b testing.TB, packet *msg.UDPPacket, codec string) []byte {
|
|
||||||
b.Helper()
|
|
||||||
var buf bytes.Buffer
|
|
||||||
rw, err := msg.NewUDPPacketReadWriter(&buf, wire.ProtocolV2, codec)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := rw.WriteMsg(packet); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
return append([]byte(nil), buf.Bytes()...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func sudpPathCopyFrame(dst, src []byte) []byte {
|
|
||||||
copy(dst, src)
|
|
||||||
return dst
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkSUDPInMemoryFrameCopy(b *testing.B) {
|
|
||||||
// This is an in-memory copy of an already encoded frame. It is a proxy for
|
|
||||||
// frame-size-dependent copy work, not a benchmark of libio.Join or sockets.
|
|
||||||
for _, payloadSize := range []int{64, 512, 1200, 1472} {
|
|
||||||
for _, tc := range sudpPathBenchmarkCases(payloadSize) {
|
|
||||||
for _, codec := range []struct {
|
|
||||||
name string
|
|
||||||
value string
|
|
||||||
}{
|
|
||||||
{name: "json", value: ""},
|
|
||||||
{name: "binary", value: wire.UDPPacketCodecBinary},
|
|
||||||
} {
|
|
||||||
b.Run(fmt.Sprintf("payload-%d/%s/%s", payloadSize, tc.name, codec.name), func(b *testing.B) {
|
|
||||||
encoded := sudpPathWireBytes(b, tc.packet, codec.value)
|
|
||||||
dst := make([]byte, len(encoded))
|
|
||||||
b.SetBytes(int64(len(encoded)))
|
|
||||||
for b.Loop() {
|
|
||||||
dst = sudpPathCopyFrame(dst, encoded)
|
|
||||||
}
|
|
||||||
if !bytes.Equal(dst, encoded) {
|
|
||||||
b.Fatal("copied frame does not match source")
|
|
||||||
}
|
|
||||||
sudpPathBenchmarkBytesSink = dst
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkSUDPEndpointCodecPair(b *testing.B) {
|
|
||||||
// This measures an in-memory decode and re-encode with the same codec. It
|
|
||||||
// does not include the live SUDP server path, sockets, goroutines, or I/O.
|
|
||||||
for _, payloadSize := range []int{64, 512, 1200, 1472} {
|
|
||||||
for _, tc := range sudpPathBenchmarkCases(payloadSize) {
|
|
||||||
for _, codec := range []struct {
|
|
||||||
name string
|
|
||||||
value string
|
|
||||||
}{
|
|
||||||
{name: "json", value: ""},
|
|
||||||
{name: "binary", value: wire.UDPPacketCodecBinary},
|
|
||||||
} {
|
|
||||||
b.Run(fmt.Sprintf("payload-%d/%s/%s", payloadSize, tc.name, codec.name), func(b *testing.B) {
|
|
||||||
encoded := sudpPathWireBytes(b, tc.packet, codec.value)
|
|
||||||
from := &sudpPathReadWriter{}
|
|
||||||
to := &bytes.Buffer{}
|
|
||||||
fromRW, err := msg.NewUDPPacketReadWriter(from, wire.ProtocolV2, codec.value)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
toRW, err := msg.NewUDPPacketReadWriter(to, wire.ProtocolV2, codec.value)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
b.SetBytes(int64(len(encoded)))
|
|
||||||
for b.Loop() {
|
|
||||||
from.Reset(encoded)
|
|
||||||
to.Reset()
|
|
||||||
m, err := fromRW.ReadMsg()
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := toRW.WriteMsg(m); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !bytes.Equal(to.Bytes(), encoded) {
|
|
||||||
b.Fatalf("re-encoded packet mismatch: got %d bytes, want %d", to.Len(), len(encoded))
|
|
||||||
}
|
|
||||||
sudpPathBenchmarkBytesSink = to.Bytes()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkSUDPMixedCodecTranscodeModel(b *testing.B) {
|
|
||||||
// This exercises the real codec decode/re-encode pair used by the mixed
|
|
||||||
// bridge, excluding sockets, goroutines, crypto, compression, and framing I/O.
|
|
||||||
for _, payloadSize := range []int{64, 512, 1200, 1472} {
|
|
||||||
for _, tc := range sudpPathBenchmarkCases(payloadSize) {
|
|
||||||
for _, direction := range []struct {
|
|
||||||
name string
|
|
||||||
from string
|
|
||||||
to string
|
|
||||||
}{
|
|
||||||
{name: "json-to-binary", from: "", to: wire.UDPPacketCodecBinary},
|
|
||||||
{name: "binary-to-json", from: wire.UDPPacketCodecBinary, to: ""},
|
|
||||||
} {
|
|
||||||
b.Run(fmt.Sprintf("payload-%d/%s/%s", payloadSize, tc.name, direction.name), func(b *testing.B) {
|
|
||||||
encoded := sudpPathWireBytes(b, tc.packet, direction.from)
|
|
||||||
expected := sudpPathWireBytes(b, tc.packet, direction.to)
|
|
||||||
from := &sudpPathReadWriter{}
|
|
||||||
to := &bytes.Buffer{}
|
|
||||||
fromRW, err := msg.NewUDPPacketReadWriter(from, wire.ProtocolV2, direction.from)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
toRW, err := msg.NewUDPPacketReadWriter(to, wire.ProtocolV2, direction.to)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
b.SetBytes(int64(len(encoded)))
|
|
||||||
for b.Loop() {
|
|
||||||
from.Reset(encoded)
|
|
||||||
to.Reset()
|
|
||||||
m, err := fromRW.ReadMsg()
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := toRW.WriteMsg(m); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !bytes.Equal(to.Bytes(), expected) {
|
|
||||||
b.Fatalf("transcoded packet mismatch: got %d bytes, want %d", to.Len(), len(expected))
|
|
||||||
}
|
|
||||||
sudpPathBenchmarkBytesSink = to.Bytes()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var _ io.ReadWriter = (*sudpPathReadWriter)(nil)
|
|
||||||
+18
-235
@@ -18,27 +18,22 @@ import (
|
|||||||
"bufio"
|
"bufio"
|
||||||
"bytes"
|
"bytes"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
"github.com/fatedier/frp/pkg/msg"
|
||||||
"github.com/fatedier/frp/pkg/proto/wire"
|
"github.com/fatedier/frp/pkg/proto/wire"
|
||||||
"github.com/fatedier/frp/pkg/util/xlog"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestSUDPBridgeTranscodesProxyV1ToVisitorV2(t *testing.T) {
|
func TestSUDPBridgeTranscodesProxyV1ToVisitorV2(t *testing.T) {
|
||||||
var in, out bytes.Buffer
|
var in, out bytes.Buffer
|
||||||
writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, "", &msg.UDPPacket{Content: []byte("proxy-to-visitor")})
|
writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, &msg.UDPPacket{Content: []byte("proxy-to-visitor")})
|
||||||
|
|
||||||
var count int64
|
var count int64
|
||||||
err := bridgeSUDPProxyToVisitor(
|
err := bridgeSUDPProxyToVisitor(
|
||||||
newSUDPBridgeRW(t, &in, wire.ProtocolV1, ""),
|
msg.NewReadWriter(&in, wire.ProtocolV1),
|
||||||
newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""),
|
msg.NewReadWriter(&out, wire.ProtocolV2),
|
||||||
&count,
|
&count,
|
||||||
nil,
|
nil,
|
||||||
)
|
)
|
||||||
@@ -58,12 +53,12 @@ func TestSUDPBridgeTranscodesProxyV1ToVisitorV2(t *testing.T) {
|
|||||||
|
|
||||||
func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) {
|
func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) {
|
||||||
var in, out bytes.Buffer
|
var in, out bytes.Buffer
|
||||||
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", &msg.UDPPacket{Content: []byte("visitor-to-proxy")})
|
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, &msg.UDPPacket{Content: []byte("visitor-to-proxy")})
|
||||||
|
|
||||||
var count int64
|
var count int64
|
||||||
err := bridgeSUDPVisitorToProxy(
|
err := bridgeSUDPVisitorToProxy(
|
||||||
newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
|
msg.NewReadWriter(&in, wire.ProtocolV2),
|
||||||
newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
|
msg.NewReadWriter(&out, wire.ProtocolV1),
|
||||||
&count,
|
&count,
|
||||||
nil,
|
nil,
|
||||||
)
|
)
|
||||||
@@ -81,67 +76,33 @@ func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) {
|
|||||||
require.Equal(t, []byte("visitor-to-proxy"), got.Content)
|
require.Equal(t, []byte("visitor-to-proxy"), got.Content)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSUDPBridgeTranscodesProxyV2BinaryToVisitorV2JSON(t *testing.T) {
|
|
||||||
var in, out bytes.Buffer
|
|
||||||
packet := newSUDPBridgeUDPPacket("proxy-binary-to-json")
|
|
||||||
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, wire.UDPPacketCodecBinary, packet)
|
|
||||||
|
|
||||||
var count int64
|
|
||||||
err := bridgeSUDPProxyToVisitor(
|
|
||||||
newSUDPBridgeRW(t, &in, wire.ProtocolV2, wire.UDPPacketCodecBinary),
|
|
||||||
newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""),
|
|
||||||
&count,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, int64(len(packet.Content)), count)
|
|
||||||
requireV2UDPPacketFrame(t, &out, msg.V2TypeUDPPacket, packet)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSUDPBridgeTranscodesVisitorV2JSONToProxyV2Binary(t *testing.T) {
|
|
||||||
var in, out bytes.Buffer
|
|
||||||
packet := newSUDPBridgeUDPPacket("visitor-json-to-binary")
|
|
||||||
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", packet)
|
|
||||||
|
|
||||||
var count int64
|
|
||||||
err := bridgeSUDPVisitorToProxy(
|
|
||||||
newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
|
|
||||||
newSUDPBridgeRW(t, &out, wire.ProtocolV2, wire.UDPPacketCodecBinary),
|
|
||||||
&count,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, int64(len(packet.Content)), count)
|
|
||||||
requireV2UDPPacketFrame(t, &out, msg.V2TypeUDPPacketBinary, packet)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSUDPBridgeForwardsProxyPing(t *testing.T) {
|
func TestSUDPBridgeForwardsProxyPing(t *testing.T) {
|
||||||
var in, out bytes.Buffer
|
var in, out bytes.Buffer
|
||||||
writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, "", &msg.Ping{})
|
writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, &msg.Ping{})
|
||||||
|
|
||||||
var count int64
|
var count int64
|
||||||
err := bridgeSUDPProxyToVisitor(
|
err := bridgeSUDPProxyToVisitor(
|
||||||
newSUDPBridgeRW(t, &in, wire.ProtocolV1, ""),
|
msg.NewReadWriter(&in, wire.ProtocolV1),
|
||||||
newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""),
|
msg.NewReadWriter(&out, wire.ProtocolV2),
|
||||||
&count,
|
&count,
|
||||||
nil,
|
nil,
|
||||||
)
|
)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Zero(t, count)
|
require.Zero(t, count)
|
||||||
|
|
||||||
rawMsg, err := newSUDPBridgeRW(t, &out, wire.ProtocolV2, "").ReadMsg()
|
rawMsg, err := msg.NewReadWriter(&out, wire.ProtocolV2).ReadMsg()
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.IsType(t, &msg.Ping{}, rawMsg)
|
require.IsType(t, &msg.Ping{}, rawMsg)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSUDPBridgeDropsVisitorPing(t *testing.T) {
|
func TestSUDPBridgeDropsVisitorPing(t *testing.T) {
|
||||||
var in, out bytes.Buffer
|
var in, out bytes.Buffer
|
||||||
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", &msg.Ping{})
|
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, &msg.Ping{})
|
||||||
|
|
||||||
var count int64
|
var count int64
|
||||||
err := bridgeSUDPVisitorToProxy(
|
err := bridgeSUDPVisitorToProxy(
|
||||||
newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
|
msg.NewReadWriter(&in, wire.ProtocolV2),
|
||||||
newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
|
msg.NewReadWriter(&out, wire.ProtocolV1),
|
||||||
&count,
|
&count,
|
||||||
nil,
|
nil,
|
||||||
)
|
)
|
||||||
@@ -152,12 +113,12 @@ func TestSUDPBridgeDropsVisitorPing(t *testing.T) {
|
|||||||
|
|
||||||
func TestSUDPBridgeRejectsUnknownVisitorMessage(t *testing.T) {
|
func TestSUDPBridgeRejectsUnknownVisitorMessage(t *testing.T) {
|
||||||
var in, out bytes.Buffer
|
var in, out bytes.Buffer
|
||||||
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", &msg.Pong{})
|
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, &msg.Pong{})
|
||||||
|
|
||||||
var count int64
|
var count int64
|
||||||
err := bridgeSUDPVisitorToProxy(
|
err := bridgeSUDPVisitorToProxy(
|
||||||
newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
|
msg.NewReadWriter(&in, wire.ProtocolV2),
|
||||||
newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
|
msg.NewReadWriter(&out, wire.ProtocolV1),
|
||||||
&count,
|
&count,
|
||||||
nil,
|
nil,
|
||||||
)
|
)
|
||||||
@@ -166,22 +127,6 @@ func TestSUDPBridgeRejectsUnknownVisitorMessage(t *testing.T) {
|
|||||||
require.Empty(t, out.Bytes())
|
require.Empty(t, out.Bytes())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSUDPBridgeRejectsMismatchedPacketCodecOnStream(t *testing.T) {
|
|
||||||
var in, out bytes.Buffer
|
|
||||||
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", newSUDPBridgeUDPPacket("json-on-binary-stream"))
|
|
||||||
|
|
||||||
var count int64
|
|
||||||
err := bridgeSUDPProxyToVisitor(
|
|
||||||
newSUDPBridgeRW(t, &in, wire.ProtocolV2, wire.UDPPacketCodecBinary),
|
|
||||||
newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""),
|
|
||||||
&count,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
require.ErrorContains(t, err, "received JSON UDP packet after binary codec negotiation")
|
|
||||||
require.Zero(t, count)
|
|
||||||
require.Empty(t, out.Bytes())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSUDPBridgeDetectsMixedWireProtocol(t *testing.T) {
|
func TestSUDPBridgeDetectsMixedWireProtocol(t *testing.T) {
|
||||||
require.False(t, isMixedWireProtocol("", wire.ProtocolV1))
|
require.False(t, isMixedWireProtocol("", wire.ProtocolV1))
|
||||||
require.False(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV2))
|
require.False(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV2))
|
||||||
@@ -189,170 +134,8 @@ func TestSUDPBridgeDetectsMixedWireProtocol(t *testing.T) {
|
|||||||
require.True(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV1))
|
require.True(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV1))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSUDPBridgeDetectsMixedPacketEncoding(t *testing.T) {
|
func writeSUDPBridgeMsg(t *testing.T, buf *bytes.Buffer, wireProtocol string, m msg.Message) {
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
leftWire string
|
|
||||||
leftCodec string
|
|
||||||
rightWire string
|
|
||||||
rightCodec string
|
|
||||||
mixed bool
|
|
||||||
}{
|
|
||||||
{name: "legacy v1 aliases explicit v1", leftWire: "", rightWire: wire.ProtocolV1},
|
|
||||||
{name: "v2 json matches v2 json", leftWire: wire.ProtocolV2, rightWire: wire.ProtocolV2},
|
|
||||||
{
|
|
||||||
name: "v2 binary matches v2 binary",
|
|
||||||
leftWire: wire.ProtocolV2,
|
|
||||||
leftCodec: wire.UDPPacketCodecBinary,
|
|
||||||
rightWire: wire.ProtocolV2,
|
|
||||||
rightCodec: wire.UDPPacketCodecBinary,
|
|
||||||
},
|
|
||||||
{name: "v1 json differs from v2 json", leftWire: wire.ProtocolV1, rightWire: wire.ProtocolV2, mixed: true},
|
|
||||||
{
|
|
||||||
name: "v2 json differs from v2 binary",
|
|
||||||
leftWire: wire.ProtocolV2,
|
|
||||||
rightWire: wire.ProtocolV2,
|
|
||||||
rightCodec: wire.UDPPacketCodecBinary,
|
|
||||||
mixed: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "v2 binary differs from v1 json",
|
|
||||||
leftWire: wire.ProtocolV2,
|
|
||||||
leftCodec: wire.UDPPacketCodecBinary,
|
|
||||||
rightWire: wire.ProtocolV1,
|
|
||||||
mixed: true,
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
mixed, err := isMixedSUDPPacketEncoding(tc.leftWire, tc.leftCodec, tc.rightWire, tc.rightCodec)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, tc.mixed, mixed)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSUDPBridgeRejectsInvalidEncodingMetadata(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
leftWire string
|
|
||||||
leftCodec string
|
|
||||||
rightWire string
|
|
||||||
rightCodec string
|
|
||||||
wantErr string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "left v1 binary",
|
|
||||||
leftWire: wire.ProtocolV1,
|
|
||||||
leftCodec: wire.UDPPacketCodecBinary,
|
|
||||||
rightWire: wire.ProtocolV1,
|
|
||||||
wantErr: "invalid left SUDP packet encoding",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "right unknown v2 codec",
|
|
||||||
leftWire: wire.ProtocolV2,
|
|
||||||
rightWire: wire.ProtocolV2,
|
|
||||||
rightCodec: "snappy",
|
|
||||||
wantErr: "invalid right SUDP packet encoding",
|
|
||||||
},
|
|
||||||
{name: "left unknown wire", leftWire: "v3", rightWire: wire.ProtocolV2, wantErr: "unsupported wire protocol"},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
mixed, err := isMixedSUDPPacketEncoding(tc.leftWire, tc.leftCodec, tc.rightWire, tc.rightCodec)
|
|
||||||
require.False(t, mixed)
|
|
||||||
require.ErrorContains(t, err, tc.wantErr)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSUDPJoinUsesRawPathForSameEncodingState(t *testing.T) {
|
|
||||||
proxyClient, proxyServer := net.Pipe()
|
|
||||||
visitorClient, visitorServer := net.Pipe()
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = proxyClient.Close()
|
|
||||||
_ = proxyServer.Close()
|
|
||||||
_ = visitorClient.Close()
|
|
||||||
_ = visitorServer.Close()
|
|
||||||
})
|
|
||||||
deadline := time.Now().Add(3 * time.Second)
|
|
||||||
require.NoError(t, proxyClient.SetDeadline(deadline))
|
|
||||||
require.NoError(t, proxyServer.SetDeadline(deadline))
|
|
||||||
require.NoError(t, visitorClient.SetDeadline(deadline))
|
|
||||||
require.NoError(t, visitorServer.SetDeadline(deadline))
|
|
||||||
|
|
||||||
pxy := &BaseProxy{
|
|
||||||
configurer: &v1.SUDPProxyConfig{},
|
|
||||||
wireProtocol: wire.ProtocolV2,
|
|
||||||
udpPacketCodec: wire.UDPPacketCodecBinary,
|
|
||||||
}
|
|
||||||
visitorConn := &metadataConn{Conn: visitorServer, wireProtocol: wire.ProtocolV2, udpPacketCodec: wire.UDPPacketCodecBinary}
|
|
||||||
joinDone := make(chan []error, 1)
|
|
||||||
go func() {
|
|
||||||
_, _, errs := pxy.joinUserConnection(proxyServer, visitorConn, string(v1.ProxyTypeSUDP), xlog.New())
|
|
||||||
joinDone <- errs
|
|
||||||
}()
|
|
||||||
|
|
||||||
raw := []byte{0, 16, 0, 0, 0, 4, 0xde, 0xad, 0xbe, 0xef}
|
|
||||||
_, err := proxyClient.Write(raw)
|
|
||||||
require.NoError(t, err)
|
|
||||||
got := make([]byte, len(raw))
|
|
||||||
_, err = io.ReadFull(visitorClient, got)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, raw, got)
|
|
||||||
|
|
||||||
_ = proxyClient.Close()
|
|
||||||
_ = visitorClient.Close()
|
|
||||||
<-joinDone
|
|
||||||
}
|
|
||||||
|
|
||||||
func newSUDPBridgeRW(t *testing.T, buf *bytes.Buffer, wireProtocol, udpPacketCodec string) msg.ReadWriter {
|
|
||||||
t.Helper()
|
t.Helper()
|
||||||
rw, err := msg.NewUDPPacketReadWriter(buf, wireProtocol, udpPacketCodec)
|
|
||||||
require.NoError(t, err)
|
|
||||||
return rw
|
|
||||||
}
|
|
||||||
|
|
||||||
func writeSUDPBridgeMsg(t *testing.T, buf *bytes.Buffer, wireProtocol, udpPacketCodec string, m msg.Message) {
|
require.NoError(t, msg.NewReadWriter(buf, wireProtocol).WriteMsg(m))
|
||||||
t.Helper()
|
|
||||||
require.NoError(t, newSUDPBridgeRW(t, buf, wireProtocol, udpPacketCodec).WriteMsg(m))
|
|
||||||
}
|
|
||||||
|
|
||||||
func newSUDPBridgeUDPPacket(content string) *msg.UDPPacket {
|
|
||||||
return &msg.UDPPacket{
|
|
||||||
Content: []byte(content),
|
|
||||||
RemoteAddr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 12345},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func requireV2UDPPacketFrame(t *testing.T, buf *bytes.Buffer, wantType uint16, want *msg.UDPPacket) {
|
|
||||||
t.Helper()
|
|
||||||
frame, err := wire.NewConn(buf).ReadFrame()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, wire.FrameTypeMessage, frame.Type)
|
|
||||||
require.GreaterOrEqual(t, len(frame.Payload), 2)
|
|
||||||
require.Equal(t, wantType, binary.BigEndian.Uint16(frame.Payload[:2]))
|
|
||||||
var got *msg.UDPPacket
|
|
||||||
if wantType == msg.V2TypeUDPPacketBinary {
|
|
||||||
got, err = msg.DecodeUDPPacketBinary(frame.Payload[2:])
|
|
||||||
} else {
|
|
||||||
var decoded msg.UDPPacket
|
|
||||||
err = msg.DecodeV2MessageFrameInto(frame, &decoded)
|
|
||||||
got = &decoded
|
|
||||||
}
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, want.Content, got.Content)
|
|
||||||
require.Equal(t, want.RemoteAddr.String(), got.RemoteAddr.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
type metadataConn struct {
|
|
||||||
net.Conn
|
|
||||||
wireProtocol string
|
|
||||||
udpPacketCodec string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *metadataConn) WireProtocol() string {
|
|
||||||
return c.wireProtocol
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *metadataConn) UDPPacketCodec() string {
|
|
||||||
return c.udpPacketCodec
|
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-7
@@ -224,13 +224,7 @@ func (pxy *UDPProxy) Run() (remoteAddr string, err error) {
|
|||||||
|
|
||||||
pxy.workConn = netpkg.WrapReadWriteCloserToConn(rwc, workConn)
|
pxy.workConn = netpkg.WrapReadWriteCloserToConn(rwc, workConn)
|
||||||
// Plain UDP payload follows the negotiated wire protocol for message framing.
|
// Plain UDP payload follows the negotiated wire protocol for message framing.
|
||||||
payloadRW, err := msg.NewUDPPacketReadWriter(pxy.workConn, pxy.wireProtocol, pxy.udpPacketCodec)
|
payloadConn := msg.NewConn(pxy.workConn, msg.NewReadWriter(pxy.workConn, pxy.wireProtocol))
|
||||||
if err != nil {
|
|
||||||
xl.Errorf("create UDP packet read writer: %v", err)
|
|
||||||
pxy.workConn.Close()
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
payloadConn := msg.NewConn(pxy.workConn, payloadRW)
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
go workConnReaderFn(payloadConn)
|
go workConnReaderFn(payloadConn)
|
||||||
go workConnSenderFn(payloadConn, ctx)
|
go workConnSenderFn(payloadConn, ctx)
|
||||||
|
|||||||
@@ -28,7 +28,6 @@ type ClientInfo struct {
|
|||||||
User string
|
User string
|
||||||
RawClientID string
|
RawClientID string
|
||||||
RunID string
|
RunID string
|
||||||
ControlID uint64
|
|
||||||
Hostname string
|
Hostname string
|
||||||
IP string
|
IP string
|
||||||
Version string
|
Version string
|
||||||
@@ -65,16 +64,6 @@ func newClientRegistryWithClock(clk clock.PassiveClock) *ClientRegistry {
|
|||||||
|
|
||||||
// Register stores/updates metadata for a client and returns the registry key plus whether it conflicts with an online client.
|
// Register stores/updates metadata for a client and returns the registry key plus whether it conflicts with an online client.
|
||||||
func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol string) (key string, conflict bool) {
|
func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol string) (key string, conflict bool) {
|
||||||
return cr.RegisterWithControlID(user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// RegisterWithControlID is the generation-aware form used by ControlManager.
|
|
||||||
// A control ID is process-local and prevents an older control generation from
|
|
||||||
// changing the registry entry now owned by a newer generation with the same run ID.
|
|
||||||
func (cr *ClientRegistry) RegisterWithControlID(
|
|
||||||
user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol string,
|
|
||||||
controlID uint64,
|
|
||||||
) (key string, conflict bool) {
|
|
||||||
if runID == "" {
|
if runID == "" {
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
@@ -94,16 +83,6 @@ func (cr *ClientRegistry) RegisterWithControlID(
|
|||||||
if enforceUnique && exists && info.Online && info.RunID != "" && info.RunID != runID {
|
if enforceUnique && exists && info.Online && info.RunID != "" && info.RunID != runID {
|
||||||
return key, true
|
return key, true
|
||||||
}
|
}
|
||||||
if previousKey, ok := cr.runIndex[runID]; ok && previousKey != key {
|
|
||||||
if previous, ok := cr.clients[previousKey]; ok && previous.RunID == runID {
|
|
||||||
if previous.RawClientID == "" {
|
|
||||||
delete(cr.clients, previousKey)
|
|
||||||
} else {
|
|
||||||
setClientOffline(previous, now)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
delete(cr.runIndex, runID)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !exists {
|
if !exists {
|
||||||
info = &ClientInfo{
|
info = &ClientInfo{
|
||||||
@@ -118,7 +97,6 @@ func (cr *ClientRegistry) RegisterWithControlID(
|
|||||||
|
|
||||||
info.RawClientID = rawClientID
|
info.RawClientID = rawClientID
|
||||||
info.RunID = runID
|
info.RunID = runID
|
||||||
info.ControlID = controlID
|
|
||||||
info.Hostname = hostname
|
info.Hostname = hostname
|
||||||
info.IP = remoteAddr
|
info.IP = remoteAddr
|
||||||
info.Version = version
|
info.Version = version
|
||||||
@@ -136,16 +114,6 @@ func (cr *ClientRegistry) RegisterWithControlID(
|
|||||||
|
|
||||||
// MarkOfflineByRunID marks the client as offline when the corresponding control disconnects.
|
// MarkOfflineByRunID marks the client as offline when the corresponding control disconnects.
|
||||||
func (cr *ClientRegistry) MarkOfflineByRunID(runID string) {
|
func (cr *ClientRegistry) MarkOfflineByRunID(runID string) {
|
||||||
cr.markOfflineByRunID(runID, 0, false)
|
|
||||||
}
|
|
||||||
|
|
||||||
// MarkOfflineByRunIDAndControlID marks a client offline only when the registry
|
|
||||||
// entry still belongs to the supplied control generation.
|
|
||||||
func (cr *ClientRegistry) MarkOfflineByRunIDAndControlID(runID string, controlID uint64) {
|
|
||||||
cr.markOfflineByRunID(runID, controlID, true)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (cr *ClientRegistry) markOfflineByRunID(runID string, controlID uint64, matchControlID bool) {
|
|
||||||
cr.mu.Lock()
|
cr.mu.Lock()
|
||||||
defer cr.mu.Unlock()
|
defer cr.mu.Unlock()
|
||||||
|
|
||||||
@@ -153,23 +121,17 @@ func (cr *ClientRegistry) markOfflineByRunID(runID string, controlID uint64, mat
|
|||||||
if !ok {
|
if !ok {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if info, ok := cr.clients[key]; ok && info.RunID == runID && (!matchControlID || info.ControlID == controlID) {
|
if info, ok := cr.clients[key]; ok && info.RunID == runID {
|
||||||
if info.RawClientID == "" {
|
if info.RawClientID == "" {
|
||||||
delete(cr.clients, key)
|
delete(cr.clients, key)
|
||||||
} else {
|
} else {
|
||||||
setClientOffline(info, cr.clock.Now())
|
info.RunID = ""
|
||||||
|
info.Online = false
|
||||||
|
now := cr.clock.Now()
|
||||||
|
info.DisconnectedAt = now
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if info, ok := cr.clients[key]; !ok || info.RunID != runID {
|
delete(cr.runIndex, runID)
|
||||||
delete(cr.runIndex, runID)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func setClientOffline(info *ClientInfo, now time.Time) {
|
|
||||||
info.RunID = ""
|
|
||||||
info.ControlID = 0
|
|
||||||
info.Online = false
|
|
||||||
info.DisconnectedAt = now
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// List returns a snapshot of all known clients.
|
// List returns a snapshot of all known clients.
|
||||||
|
|||||||
@@ -72,89 +72,3 @@ func TestClientRegistryUsesClockForTimestamps(t *testing.T) {
|
|||||||
t.Fatalf("disconnected time mismatch, want %s got %s", disconnectedAt, info.DisconnectedAt)
|
t.Fatalf("disconnected time mismatch, want %s got %s", disconnectedAt, info.DisconnectedAt)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestClientRegistryControlIDPreventsStaleOffline(t *testing.T) {
|
|
||||||
registry := NewClientRegistry()
|
|
||||||
key, conflict := registry.RegisterWithControlID(
|
|
||||||
"user", "client-id", "run-id", "old-host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1,
|
|
||||||
)
|
|
||||||
if conflict {
|
|
||||||
t.Fatal("unexpected client conflict")
|
|
||||||
}
|
|
||||||
_, conflict = registry.RegisterWithControlID(
|
|
||||||
"user", "client-id", "run-id", "new-host", "1.0.1", "127.0.0.2", wire.ProtocolV2, 2,
|
|
||||||
)
|
|
||||||
if conflict {
|
|
||||||
t.Fatal("same run ID replacement should not conflict")
|
|
||||||
}
|
|
||||||
|
|
||||||
registry.MarkOfflineByRunIDAndControlID("run-id", 1)
|
|
||||||
info, ok := registry.GetByKey(key)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("client %q not found", key)
|
|
||||||
}
|
|
||||||
if !info.Online || info.ControlID != 2 || info.Hostname != "new-host" {
|
|
||||||
t.Fatalf("stale offline changed current generation: %+v", info)
|
|
||||||
}
|
|
||||||
|
|
||||||
registry.MarkOfflineByRunIDAndControlID("run-id", 2)
|
|
||||||
info, ok = registry.GetByKey(key)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("client %q not found after disconnect", key)
|
|
||||||
}
|
|
||||||
if info.Online || info.ControlID != 0 || info.RunID != "" {
|
|
||||||
t.Fatalf("current generation was not marked offline: %+v", info)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClientRegistryClientIDConflictSemantics(t *testing.T) {
|
|
||||||
registry := NewClientRegistry()
|
|
||||||
_, conflict := registry.RegisterWithControlID(
|
|
||||||
"user", "client-id", "run-one", "host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1,
|
|
||||||
)
|
|
||||||
if conflict {
|
|
||||||
t.Fatal("unexpected initial client conflict")
|
|
||||||
}
|
|
||||||
_, conflict = registry.RegisterWithControlID(
|
|
||||||
"user", "client-id", "run-two", "host", "1.0.0", "127.0.0.2", wire.ProtocolV1, 2,
|
|
||||||
)
|
|
||||||
if !conflict {
|
|
||||||
t.Fatal("different online run IDs with the same explicit client ID must conflict")
|
|
||||||
}
|
|
||||||
|
|
||||||
registry.MarkOfflineByRunIDAndControlID("run-one", 1)
|
|
||||||
_, conflict = registry.RegisterWithControlID(
|
|
||||||
"user", "client-id", "run-two", "host", "1.0.0", "127.0.0.2", wire.ProtocolV1, 2,
|
|
||||||
)
|
|
||||||
if conflict {
|
|
||||||
t.Fatal("offline explicit client ID should be reusable")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClientRegistrySameRunIDMovesBetweenClientKeys(t *testing.T) {
|
|
||||||
registry := NewClientRegistry()
|
|
||||||
oldKey, conflict := registry.RegisterWithControlID(
|
|
||||||
"user", "old-client", "run-id", "old-host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1,
|
|
||||||
)
|
|
||||||
if conflict {
|
|
||||||
t.Fatal("unexpected initial client conflict")
|
|
||||||
}
|
|
||||||
newKey, conflict := registry.RegisterWithControlID(
|
|
||||||
"user", "new-client", "run-id", "new-host", "1.0.1", "127.0.0.2", wire.ProtocolV2, 2,
|
|
||||||
)
|
|
||||||
if conflict {
|
|
||||||
t.Fatal("same run ID moving to a new client key should not conflict")
|
|
||||||
}
|
|
||||||
|
|
||||||
oldInfo, ok := registry.GetByKey(oldKey)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("old explicit client %q should remain as offline history", oldKey)
|
|
||||||
}
|
|
||||||
if oldInfo.Online || oldInfo.RunID != "" || oldInfo.ControlID != 0 {
|
|
||||||
t.Fatalf("old client key remained online: %+v", oldInfo)
|
|
||||||
}
|
|
||||||
newInfo, ok := registry.GetByKey(newKey)
|
|
||||||
if !ok || !newInfo.Online || newInfo.RunID != "run-id" || newInfo.ControlID != 2 {
|
|
||||||
t.Fatalf("new client key was not registered: %+v", newInfo)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+45
-113
@@ -18,7 +18,6 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"crypto/tls"
|
"crypto/tls"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
@@ -29,13 +28,12 @@ import (
|
|||||||
|
|
||||||
"github.com/fatedier/golib/crypto"
|
"github.com/fatedier/golib/crypto"
|
||||||
"github.com/fatedier/golib/net/mux"
|
"github.com/fatedier/golib/net/mux"
|
||||||
fmux "github.com/fatedier/yamux"
|
fmux "github.com/hashicorp/yamux"
|
||||||
quic "github.com/quic-go/quic-go"
|
quic "github.com/quic-go/quic-go"
|
||||||
"github.com/samber/lo"
|
"github.com/samber/lo"
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/auth"
|
"github.com/fatedier/frp/pkg/auth"
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
v1 "github.com/fatedier/frp/pkg/config/v1"
|
||||||
"github.com/fatedier/frp/pkg/config/v1/validation"
|
|
||||||
modelmetrics "github.com/fatedier/frp/pkg/metrics"
|
modelmetrics "github.com/fatedier/frp/pkg/metrics"
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
"github.com/fatedier/frp/pkg/msg"
|
||||||
"github.com/fatedier/frp/pkg/nathole"
|
"github.com/fatedier/frp/pkg/nathole"
|
||||||
@@ -53,6 +51,7 @@ import (
|
|||||||
"github.com/fatedier/frp/pkg/util/xlog"
|
"github.com/fatedier/frp/pkg/util/xlog"
|
||||||
"github.com/fatedier/frp/server/controller"
|
"github.com/fatedier/frp/server/controller"
|
||||||
"github.com/fatedier/frp/server/group"
|
"github.com/fatedier/frp/server/group"
|
||||||
|
"github.com/fatedier/frp/server/metrics"
|
||||||
"github.com/fatedier/frp/server/ports"
|
"github.com/fatedier/frp/server/ports"
|
||||||
"github.com/fatedier/frp/server/proxy"
|
"github.com/fatedier/frp/server/proxy"
|
||||||
"github.com/fatedier/frp/server/registry"
|
"github.com/fatedier/frp/server/registry"
|
||||||
@@ -65,8 +64,6 @@ const (
|
|||||||
vhostReadWriteTimeout time.Duration = 30 * time.Second
|
vhostReadWriteTimeout time.Duration = 30 * time.Second
|
||||||
)
|
)
|
||||||
|
|
||||||
var errControlReplaced = errors.New("control was replaced during login")
|
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
crypto.DefaultSalt = "frp"
|
crypto.DefaultSalt = "frp"
|
||||||
// Disable quic-go's receive buffer warning.
|
// Disable quic-go's receive buffer warning.
|
||||||
@@ -164,10 +161,9 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
svr := &Service{
|
svr := &Service{
|
||||||
ctlManager: NewControlManager(clientRegistry),
|
ctlManager: NewControlManager(),
|
||||||
clientRegistry: clientRegistry,
|
clientRegistry: registry.NewClientRegistry(),
|
||||||
pxyManager: proxy.NewManager(),
|
pxyManager: proxy.NewManager(),
|
||||||
pluginManager: plugin.NewManager(),
|
pluginManager: plugin.NewManager(),
|
||||||
rc: &controller.ResourceController{
|
rc: &controller.ResourceController{
|
||||||
@@ -301,14 +297,10 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) {
|
|||||||
svr.rc.HTTPReverseProxy = rp
|
svr.rc.HTTPReverseProxy = rp
|
||||||
|
|
||||||
address := net.JoinHostPort(cfg.ProxyBindAddr, strconv.Itoa(cfg.VhostHTTPPort))
|
address := net.JoinHostPort(cfg.ProxyBindAddr, strconv.Itoa(cfg.VhostHTTPPort))
|
||||||
protocols := new(http.Protocols)
|
|
||||||
protocols.SetHTTP1(true)
|
|
||||||
protocols.SetUnencryptedHTTP2(true)
|
|
||||||
server := &http.Server{
|
server := &http.Server{
|
||||||
Addr: address,
|
Addr: address,
|
||||||
Handler: rp,
|
Handler: rp,
|
||||||
ReadHeaderTimeout: 60 * time.Second,
|
ReadHeaderTimeout: 60 * time.Second,
|
||||||
Protocols: protocols,
|
|
||||||
}
|
}
|
||||||
var l net.Listener
|
var l net.Listener
|
||||||
if httpMuxOn {
|
if httpMuxOn {
|
||||||
@@ -471,15 +463,12 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err == nil {
|
if err == nil {
|
||||||
ctl, err = svr.RegisterControl(controlConn, m, internal, acceptedConn.wireProtocol, acceptedConn.udpPacketCodec)
|
ctl, err = svr.RegisterControl(controlConn, m, internal, acceptedConn.wireProtocol)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
xl.Warnf("register control error: %v", err)
|
xl.Warnf("register control error: %v", err)
|
||||||
if ctl != nil {
|
|
||||||
svr.ctlManager.Remove(ctl)
|
|
||||||
}
|
|
||||||
if writeErr := writeWithDeadline(conn, connWriteTimeout, func() error {
|
if writeErr := writeWithDeadline(conn, connWriteTimeout, func() error {
|
||||||
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
|
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
|
||||||
Version: version.Full(),
|
Version: version.Full(),
|
||||||
@@ -488,34 +477,31 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
|
|||||||
}); writeErr != nil {
|
}); writeErr != nil {
|
||||||
xl.Warnf("write login error response error: %v", writeErr)
|
xl.Warnf("write login error response error: %v", writeErr)
|
||||||
}
|
}
|
||||||
if ctl != nil {
|
conn.Close()
|
||||||
_ = ctl.Close()
|
|
||||||
} else {
|
|
||||||
conn.Close()
|
|
||||||
}
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err = svr.completeControlLogin(ctl, func() error {
|
if err = writeWithDeadline(conn, connWriteTimeout, func() error {
|
||||||
return writeWithDeadline(conn, connWriteTimeout, func() error {
|
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
|
||||||
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
|
Version: version.Full(),
|
||||||
Version: version.Full(),
|
RunID: ctl.runID,
|
||||||
RunID: ctl.runID,
|
Error: "",
|
||||||
Error: "",
|
|
||||||
})
|
|
||||||
})
|
})
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
xl.Warnf("complete control login error: %v", err)
|
xl.Warnf("write login response error: %v", err)
|
||||||
svr.ctlManager.Remove(ctl)
|
svr.ctlManager.Del(m.RunID, ctl)
|
||||||
_ = ctl.Close()
|
svr.clientRegistry.MarkOfflineByRunID(m.RunID)
|
||||||
|
conn.Close()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
ctl.Start()
|
||||||
|
metrics.Server.NewClient()
|
||||||
|
go func() {
|
||||||
|
// block until control closed
|
||||||
|
ctl.WaitClosed()
|
||||||
|
svr.ctlManager.Del(m.RunID, ctl)
|
||||||
|
}()
|
||||||
case *msg.NewWorkConn:
|
case *msg.NewWorkConn:
|
||||||
if err := svr.RegisterWorkConn(
|
if err := svr.RegisterWorkConn(acceptedConn.conn, m); err != nil {
|
||||||
acceptedConn.conn,
|
|
||||||
m,
|
|
||||||
acceptedConn.wireProtocol,
|
|
||||||
acceptedConn.clientHelloPresent,
|
|
||||||
); err != nil {
|
|
||||||
_ = acceptedConn.conn.WriteMsg(&msg.StartWorkConn{
|
_ = acceptedConn.conn.WriteMsg(&msg.StartWorkConn{
|
||||||
Error: util.GenerateResponseErrorString("invalid NewWorkConn", err, lo.FromPtr(svr.cfg.DetailedErrorsToClient)),
|
Error: util.GenerateResponseErrorString("invalid NewWorkConn", err, lo.FromPtr(svr.cfg.DetailedErrorsToClient)),
|
||||||
})
|
})
|
||||||
@@ -541,24 +527,11 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (svr *Service) completeControlLogin(ctl *Control, writeSuccess func() error) error {
|
|
||||||
committed, err := svr.ctlManager.completeLogin(ctl, writeSuccess)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if !committed {
|
|
||||||
return errControlReplaced
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type acceptedConnection struct {
|
type acceptedConnection struct {
|
||||||
conn *msg.Conn
|
conn *msg.Conn
|
||||||
wireProtocol string
|
wireProtocol string
|
||||||
clientHelloPresent bool
|
cryptoContext *wire.CryptoContext
|
||||||
udpPacketCodec string
|
firstMsg msg.Message
|
||||||
cryptoContext *wire.CryptoContext
|
|
||||||
firstMsg msg.Message
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (svr *Service) acceptConnection(ctx context.Context, conn net.Conn) (*acceptedConnection, error) {
|
func (svr *Service) acceptConnection(ctx context.Context, conn net.Conn) (*acceptedConnection, error) {
|
||||||
@@ -626,7 +599,6 @@ func (ac *acceptedConnection) readFirstV2Msg(conn net.Conn, wireConn *wire.Conn)
|
|||||||
return nil, fmt.Errorf("read v2 frame: %w", err)
|
return nil, fmt.Errorf("read v2 frame: %w", err)
|
||||||
}
|
}
|
||||||
if frame.Type == wire.FrameTypeClientHello {
|
if frame.Type == wire.FrameTypeClientHello {
|
||||||
ac.clientHelloPresent = true
|
|
||||||
if err := ac.handleClientHello(conn, wireConn, frame); err != nil {
|
if err := ac.handleClientHello(conn, wireConn, frame); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -675,7 +647,6 @@ func (ac *acceptedConnection) handleClientHello(conn net.Conn, wireConn *wire.Co
|
|||||||
return fmt.Errorf("write ServerHello: %w", err)
|
return fmt.Errorf("write ServerHello: %w", err)
|
||||||
}
|
}
|
||||||
ac.cryptoContext = cryptoContext
|
ac.cryptoContext = cryptoContext
|
||||||
ac.udpPacketCodec = serverHello.Selected.Message.UDPPacketCodec
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -769,20 +740,7 @@ func (svr *Service) RegisterControl(
|
|||||||
loginMsg *msg.Login,
|
loginMsg *msg.Login,
|
||||||
internal bool,
|
internal bool,
|
||||||
wireProtocol string,
|
wireProtocol string,
|
||||||
udpPacketCodec string,
|
|
||||||
) (*Control, error) {
|
) (*Control, error) {
|
||||||
switch wireProtocol {
|
|
||||||
case wire.ProtocolV1:
|
|
||||||
if udpPacketCodec != "" {
|
|
||||||
return nil, fmt.Errorf("UDP packet codec %q requires wire protocol v2", udpPacketCodec)
|
|
||||||
}
|
|
||||||
case wire.ProtocolV2:
|
|
||||||
if udpPacketCodec != "" && udpPacketCodec != wire.UDPPacketCodecBinary {
|
|
||||||
return nil, fmt.Errorf("unsupported UDP packet codec selection: %s", udpPacketCodec)
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return nil, fmt.Errorf("unsupported wire protocol: %s", wireProtocol)
|
|
||||||
}
|
|
||||||
// If client's RunID is empty, it's a new client, we just create a new controller.
|
// If client's RunID is empty, it's a new client, we just create a new controller.
|
||||||
// Otherwise, we check if there is one controller has the same run id. If so, we release previous controller and start new one.
|
// Otherwise, we check if there is one controller has the same run id. If so, we release previous controller and start new one.
|
||||||
var err error
|
var err error
|
||||||
@@ -792,9 +750,6 @@ func (svr *Service) RegisterControl(
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := validation.ValidateRunID(loginMsg.RunID); err != nil {
|
|
||||||
return nil, fmt.Errorf("invalid run id: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx := netpkg.NewContextFromConn(ctlConn)
|
ctx := netpkg.NewContextFromConn(ctlConn)
|
||||||
xl := xlog.FromContextSafe(ctx)
|
xl := xlog.FromContextSafe(ctx)
|
||||||
@@ -821,8 +776,8 @@ func (svr *Service) RegisterControl(
|
|||||||
Conn: ctlConn,
|
Conn: ctlConn,
|
||||||
LoginMsg: loginMsg,
|
LoginMsg: loginMsg,
|
||||||
ServerCfg: svr.cfg,
|
ServerCfg: svr.cfg,
|
||||||
|
ClientRegistry: svr.clientRegistry,
|
||||||
WireProtocol: wireProtocol,
|
WireProtocol: wireProtocol,
|
||||||
UDPPacketCodec: udpPacketCodec,
|
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
xl.Warnf("create new controller error: %v", err)
|
xl.Warnf("create new controller error: %v", err)
|
||||||
@@ -830,41 +785,31 @@ func (svr *Service) RegisterControl(
|
|||||||
return nil, fmt.Errorf("unexpected error when creating new controller")
|
return nil, fmt.Errorf("unexpected error when creating new controller")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := svr.ctlManager.Add(ctl); err != nil {
|
if oldCtl := svr.ctlManager.Add(loginMsg.RunID, ctl); oldCtl != nil {
|
||||||
return ctl, err
|
oldCtl.WaitClosed()
|
||||||
}
|
}
|
||||||
ctl.WaitForHandoff()
|
|
||||||
|
|
||||||
active, err := svr.ctlManager.Activate(ctl)
|
remoteAddr := ctlConn.RemoteAddr().String()
|
||||||
if err != nil {
|
if host, _, err := net.SplitHostPort(remoteAddr); err == nil {
|
||||||
return ctl, err
|
remoteAddr = host
|
||||||
}
|
}
|
||||||
if !active {
|
_, conflict := svr.clientRegistry.Register(loginMsg.User, loginMsg.ClientID, loginMsg.RunID, loginMsg.Hostname, loginMsg.Version, remoteAddr, wireProtocol)
|
||||||
return ctl, errControlReplaced
|
if conflict {
|
||||||
|
svr.ctlManager.Del(loginMsg.RunID, ctl)
|
||||||
|
return nil, fmt.Errorf("client_id [%s] for user [%s] is already online", loginMsg.ClientID, loginMsg.User)
|
||||||
}
|
}
|
||||||
|
|
||||||
return ctl, nil
|
return ctl, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// RegisterWorkConn register a new work connection to control and proxies need it.
|
// RegisterWorkConn register a new work connection to control and proxies need it.
|
||||||
func (svr *Service) RegisterWorkConn(
|
func (svr *Service) RegisterWorkConn(workConn *msg.Conn, newMsg *msg.NewWorkConn) error {
|
||||||
workConn *msg.Conn,
|
|
||||||
newMsg *msg.NewWorkConn,
|
|
||||||
workWireProtocol string,
|
|
||||||
workClientHelloPresent bool,
|
|
||||||
) error {
|
|
||||||
if workClientHelloPresent {
|
|
||||||
return fmt.Errorf("ClientHello is not allowed on work connections")
|
|
||||||
}
|
|
||||||
xl := netpkg.NewLogFromConn(workConn)
|
xl := netpkg.NewLogFromConn(workConn)
|
||||||
ctl, exist := svr.ctlManager.GetByID(newMsg.RunID)
|
ctl, exist := svr.ctlManager.GetByID(newMsg.RunID)
|
||||||
if !exist {
|
if !exist {
|
||||||
xl.Warnf("no client control found for run id [%s]", newMsg.RunID)
|
xl.Warnf("no client control found for run id [%s]", newMsg.RunID)
|
||||||
return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID)
|
return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID)
|
||||||
}
|
}
|
||||||
if workWireProtocol != ctl.sessionCtx.WireProtocol {
|
|
||||||
return fmt.Errorf("work connection wire protocol mismatch: got %s want %s", workWireProtocol, ctl.sessionCtx.WireProtocol)
|
|
||||||
}
|
|
||||||
|
|
||||||
// server plugin hook
|
// server plugin hook
|
||||||
content := &plugin.NewWorkConnContent{
|
content := &plugin.NewWorkConnContent{
|
||||||
@@ -885,33 +830,20 @@ func (svr *Service) RegisterWorkConn(
|
|||||||
xl.Warnf("invalid NewWorkConn with run id [%s]", newMsg.RunID)
|
xl.Warnf("invalid NewWorkConn with run id [%s]", newMsg.RunID)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return svr.ctlManager.RegisterWorkConn(ctl, proxy.NewWorkConn(workConn))
|
return ctl.RegisterWorkConn(proxy.NewWorkConn(workConn))
|
||||||
}
|
}
|
||||||
|
|
||||||
func (svr *Service) RegisterVisitorConn(visitorConn net.Conn, newMsg *msg.NewVisitorConn, wireProtocol string) error {
|
func (svr *Service) RegisterVisitorConn(visitorConn net.Conn, newMsg *msg.NewVisitorConn, wireProtocol string) error {
|
||||||
admit := func(visitorUser, visitorWireProtocol, visitorUDPPacketCodec string) error {
|
visitorUser := ""
|
||||||
if visitorWireProtocol == "" {
|
|
||||||
visitorWireProtocol = wireProtocol
|
|
||||||
}
|
|
||||||
return svr.rc.VisitorManager.NewConn(newMsg.ProxyName, visitorConn, newMsg.Timestamp, newMsg.SignKey,
|
|
||||||
newMsg.UseEncryption, newMsg.UseCompression, visitorUser, visitorWireProtocol, visitorUDPPacketCodec)
|
|
||||||
}
|
|
||||||
// TODO(deprecation): Compatible with old versions, can be without runID, user is empty. In later versions, it will be mandatory to include runID.
|
// TODO(deprecation): Compatible with old versions, can be without runID, user is empty. In later versions, it will be mandatory to include runID.
|
||||||
// If runID is required, it is not compatible with versions prior to v0.50.0.
|
// If runID is required, it is not compatible with versions prior to v0.50.0.
|
||||||
if newMsg.RunID != "" {
|
if newMsg.RunID != "" {
|
||||||
admitted, err := svr.ctlManager.admitVisitorByRunID(newMsg.RunID, func(visitorUser, controlWireProtocol, controlUDPPacketCodec string) error {
|
ctl, exist := svr.ctlManager.GetByID(newMsg.RunID)
|
||||||
if wireProtocol != controlWireProtocol {
|
if !exist {
|
||||||
return fmt.Errorf("visitor connection wire protocol mismatch: got %s want %s", wireProtocol, controlWireProtocol)
|
|
||||||
}
|
|
||||||
return admit(visitorUser, controlWireProtocol, controlUDPPacketCodec)
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if !admitted {
|
|
||||||
return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID)
|
return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID)
|
||||||
}
|
}
|
||||||
return nil
|
visitorUser = ctl.sessionCtx.LoginMsg.User
|
||||||
}
|
}
|
||||||
return admit("", wireProtocol, "")
|
return svr.rc.VisitorManager.NewConn(newMsg.ProxyName, visitorConn, newMsg.Timestamp, newMsg.SignKey,
|
||||||
|
newMsg.UseEncryption, newMsg.UseCompression, visitorUser, wireProtocol)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -15,33 +15,12 @@
|
|||||||
package server
|
package server
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
|
||||||
"math"
|
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
|
||||||
"runtime"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/fatedier/golib/net/mux"
|
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
"github.com/fatedier/frp/pkg/auth"
|
|
||||||
v1 "github.com/fatedier/frp/pkg/config/v1"
|
|
||||||
"github.com/fatedier/frp/pkg/config/v1/validation"
|
|
||||||
"github.com/fatedier/frp/pkg/msg"
|
|
||||||
plugin "github.com/fatedier/frp/pkg/plugin/server"
|
|
||||||
"github.com/fatedier/frp/pkg/proto/wire"
|
|
||||||
"github.com/fatedier/frp/pkg/util/util"
|
|
||||||
"github.com/fatedier/frp/server/controller"
|
|
||||||
"github.com/fatedier/frp/server/proxy"
|
|
||||||
"github.com/fatedier/frp/server/registry"
|
|
||||||
"github.com/fatedier/frp/server/visitor"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) {
|
func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) {
|
||||||
@@ -82,895 +61,3 @@ func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) {
|
|||||||
t.Fatal("timed out waiting for write after deadline reset")
|
t.Fatal("timed out waiting for write after deadline reset")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestServiceAcceptConnectionTracksClientHelloPresence(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
clientHelloPresent bool
|
|
||||||
offeredCodecs []string
|
|
||||||
expectedCodec string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "absent Hello",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "present Hello with JSON fallback",
|
|
||||||
clientHelloPresent: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "present Hello with binary codec",
|
|
||||||
clientHelloPresent: true,
|
|
||||||
offeredCodecs: []string{wire.UDPPacketCodecBinary},
|
|
||||||
expectedCodec: wire.UDPPacketCodecBinary,
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
serverConn, clientConn := net.Pipe()
|
|
||||||
defer serverConn.Close()
|
|
||||||
defer clientConn.Close()
|
|
||||||
|
|
||||||
clientErrCh := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
if err := wire.WriteMagic(clientConn); err != nil {
|
|
||||||
clientErrCh <- err
|
|
||||||
return
|
|
||||||
}
|
|
||||||
wireConn := wire.NewConn(clientConn)
|
|
||||||
if tc.clientHelloPresent {
|
|
||||||
hello, err := wire.NewClientHello(wire.BootstrapInfo{})
|
|
||||||
if err != nil {
|
|
||||||
clientErrCh <- err
|
|
||||||
return
|
|
||||||
}
|
|
||||||
hello.Capabilities.Message.UDPPacketCodecs = tc.offeredCodecs
|
|
||||||
if err := wireConn.WriteJSONFrame(wire.FrameTypeClientHello, hello); err != nil {
|
|
||||||
clientErrCh <- err
|
|
||||||
return
|
|
||||||
}
|
|
||||||
var serverHello wire.ServerHello
|
|
||||||
if err := wireConn.ReadJSONFrame(wire.FrameTypeServerHello, &serverHello); err != nil {
|
|
||||||
clientErrCh <- err
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
clientErrCh <- msg.NewV2ReadWriterWithConn(wireConn).WriteMsg(&msg.NewWorkConn{RunID: "shared-run"})
|
|
||||||
}()
|
|
||||||
|
|
||||||
acceptedConn, err := (&Service{}).acceptConnection(t.Context(), serverConn)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, <-clientErrCh)
|
|
||||||
require.Equal(t, tc.clientHelloPresent, acceptedConn.clientHelloPresent)
|
|
||||||
require.Equal(t, tc.expectedCodec, acceptedConn.udpPacketCodec)
|
|
||||||
require.IsType(t, &msg.NewWorkConn{}, acceptedConn.firstMsg)
|
|
||||||
require.NoError(t, acceptedConn.conn.Close())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSharedPortHTTPListenerProtocols(t *testing.T) {
|
|
||||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
sharedMux := mux.NewMux(listener)
|
|
||||||
httpListener := sharedMux.ListenHTTP(1)
|
|
||||||
muxServeErr := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
muxServeErr <- sharedMux.Serve()
|
|
||||||
}()
|
|
||||||
|
|
||||||
newProtocols := func(http1, unencryptedHTTP2 bool) *http.Protocols {
|
|
||||||
protocols := new(http.Protocols)
|
|
||||||
protocols.SetHTTP1(http1)
|
|
||||||
protocols.SetUnencryptedHTTP2(unencryptedHTTP2)
|
|
||||||
return protocols
|
|
||||||
}
|
|
||||||
|
|
||||||
const handlerProtocolHeader = "X-Test-Handler-Protocol"
|
|
||||||
httpServer := &http.Server{
|
|
||||||
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
w.Header().Set(handlerProtocolHeader, r.Proto)
|
|
||||||
w.WriteHeader(http.StatusNoContent)
|
|
||||||
}),
|
|
||||||
ReadHeaderTimeout: time.Second,
|
|
||||||
Protocols: newProtocols(true, true),
|
|
||||||
}
|
|
||||||
httpServeErr := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
httpServeErr <- httpServer.Serve(httpListener)
|
|
||||||
}()
|
|
||||||
t.Cleanup(func() {
|
|
||||||
require.NoError(t, httpServer.Close())
|
|
||||||
require.ErrorIs(t, waitForResult(t, httpServeErr, "shared HTTP server to stop"), http.ErrServerClosed)
|
|
||||||
require.NoError(t, sharedMux.Close())
|
|
||||||
require.ErrorIs(t, waitForResult(t, muxServeErr, "shared mux to stop"), net.ErrClosed)
|
|
||||||
})
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
http1 bool
|
|
||||||
unencryptedHTTP2 bool
|
|
||||||
expectedProtocol string
|
|
||||||
}{
|
|
||||||
{name: "HTTP/1.1", http1: true, expectedProtocol: "HTTP/1.1"},
|
|
||||||
{name: "HTTP/2 prior knowledge", unencryptedHTTP2: true, expectedProtocol: "HTTP/2.0"},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
transport := &http.Transport{
|
|
||||||
Protocols: newProtocols(tc.http1, tc.unencryptedHTTP2),
|
|
||||||
}
|
|
||||||
defer transport.CloseIdleConnections()
|
|
||||||
client := &http.Client{Transport: transport, Timeout: 3 * time.Second}
|
|
||||||
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "http://"+listener.Addr().String()+"/", nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
response, err := client.Do(request)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, http.StatusNoContent, response.StatusCode)
|
|
||||||
require.Equal(t, tc.expectedProtocol, response.Proto)
|
|
||||||
require.Equal(t, tc.expectedProtocol, response.Header.Get(handlerProtocolHeader))
|
|
||||||
require.NoError(t, response.Body.Close())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceControlHandoffSkipsStalePendingGeneration(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
metrics.closeEnter = make(chan struct{})
|
|
||||||
metrics.closeResume = make(chan struct{})
|
|
||||||
|
|
||||||
ctlA, connA, err := registerLifecycleTestControl(svr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
ctlA.serverMetrics = metrics
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctlA, func() error { return nil }))
|
|
||||||
waitForSignal(t, connA.readStarted, "A reader to start")
|
|
||||||
|
|
||||||
require.NoError(t, ctlA.Close())
|
|
||||||
waitForSignal(t, metrics.closeEnter, "A finalization barrier")
|
|
||||||
|
|
||||||
type registerResult struct {
|
|
||||||
ctl *Control
|
|
||||||
conn *deadlineReadConn
|
|
||||||
err error
|
|
||||||
}
|
|
||||||
resultB := make(chan registerResult, 1)
|
|
||||||
go func() {
|
|
||||||
ctl, conn, registerErr := registerLifecycleTestControl(svr)
|
|
||||||
resultB <- registerResult{ctl: ctl, conn: conn, err: registerErr}
|
|
||||||
}()
|
|
||||||
ctlB := waitForDifferentCurrentControl(t, svr.ctlManager, "shared-run", ctlA)
|
|
||||||
ctlB.serverMetrics = metrics
|
|
||||||
|
|
||||||
resultC := make(chan registerResult, 1)
|
|
||||||
go func() {
|
|
||||||
ctl, conn, registerErr := registerLifecycleTestControl(svr)
|
|
||||||
resultC <- registerResult{ctl: ctl, conn: conn, err: registerErr}
|
|
||||||
}()
|
|
||||||
ctlC := waitForDifferentCurrentControl(t, svr.ctlManager, "shared-run", ctlB)
|
|
||||||
ctlC.serverMetrics = metrics
|
|
||||||
waitForControlDone(t, ctlB)
|
|
||||||
|
|
||||||
select {
|
|
||||||
case result := <-resultB:
|
|
||||||
t.Fatalf("B returned before A finalized: %v", result.err)
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
select {
|
|
||||||
case result := <-resultC:
|
|
||||||
t.Fatalf("C returned before A finalized: %v", result.err)
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
|
|
||||||
close(metrics.closeResume)
|
|
||||||
waitForControlDone(t, ctlA)
|
|
||||||
|
|
||||||
b := <-resultB
|
|
||||||
require.Same(t, ctlB, b.ctl)
|
|
||||||
require.ErrorIs(t, b.err, errControlReplaced)
|
|
||||||
require.False(t, svr.ctlManager.Remove(ctlB))
|
|
||||||
require.NoError(t, ctlB.Close())
|
|
||||||
|
|
||||||
c := <-resultC
|
|
||||||
require.NoError(t, c.err)
|
|
||||||
require.Same(t, ctlC, c.ctl)
|
|
||||||
_, ok := svr.ctlManager.GetByID("shared-run")
|
|
||||||
require.False(t, ok)
|
|
||||||
require.Same(t, ctlC, currentControlForTest(svr.ctlManager, "shared-run"))
|
|
||||||
|
|
||||||
info, ok := svr.clientRegistry.GetByKey("client")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.True(t, info.Online)
|
|
||||||
require.Equal(t, uint64(ctlC.ID()), info.ControlID)
|
|
||||||
|
|
||||||
var staleWrites atomic.Int64
|
|
||||||
err = svr.completeControlLogin(ctlB, func() error {
|
|
||||||
staleWrites.Add(1)
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
require.ErrorIs(t, err, errControlReplaced)
|
|
||||||
require.Equal(t, int64(0), staleWrites.Load())
|
|
||||||
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctlC, func() error { return nil }))
|
|
||||||
waitForSignal(t, c.conn.readStarted, "C reader to start")
|
|
||||||
current, ok := svr.ctlManager.GetByID("shared-run")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.Same(t, ctlC, current)
|
|
||||||
require.Equal(t, int64(2), metrics.newClients())
|
|
||||||
require.Equal(t, int64(1), metrics.closedClients())
|
|
||||||
|
|
||||||
require.NoError(t, ctlC.Close())
|
|
||||||
waitForControlDone(t, ctlC)
|
|
||||||
require.Equal(t, int64(2), metrics.newClients())
|
|
||||||
require.Equal(t, int64(2), metrics.closedClients())
|
|
||||||
_, ok = svr.ctlManager.GetByID("shared-run")
|
|
||||||
require.False(t, ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceLoginResponseSynchronizationIsScopedToRun(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
metrics := newCountingServerMetrics()
|
|
||||||
ctlA, connA, err := registerLifecycleTestControl(svr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
ctlA.serverMetrics = metrics
|
|
||||||
|
|
||||||
writeEntered := make(chan struct{})
|
|
||||||
resumeWrite := make(chan struct{})
|
|
||||||
var resumeWriteOnce sync.Once
|
|
||||||
resume := func() {
|
|
||||||
resumeWriteOnce.Do(func() { close(resumeWrite) })
|
|
||||||
}
|
|
||||||
t.Cleanup(resume)
|
|
||||||
writeCount := atomic.Int64{}
|
|
||||||
loginDone := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
loginDone <- svr.completeControlLogin(ctlA, func() error {
|
|
||||||
close(writeEntered)
|
|
||||||
<-resumeWrite
|
|
||||||
writeCount.Add(1)
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
}()
|
|
||||||
waitForSignal(t, writeEntered, "A LoginResp write")
|
|
||||||
|
|
||||||
runMu := currentRunGateForTest(svr.ctlManager, "shared-run")
|
|
||||||
require.NotNil(t, runMu)
|
|
||||||
if !svr.ctlManager.mu.TryLock() {
|
|
||||||
t.Fatal("ControlManager mutex was held while LoginResp write was in progress")
|
|
||||||
}
|
|
||||||
svr.ctlManager.mu.Unlock()
|
|
||||||
|
|
||||||
ctlB, connB := newLifecycleTestControl(t, "shared-run", "client", metrics)
|
|
||||||
gateAvailable := make(chan bool)
|
|
||||||
addDone := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
if runMu.TryLock() {
|
|
||||||
runMu.Unlock()
|
|
||||||
gateAvailable <- true
|
|
||||||
} else {
|
|
||||||
gateAvailable <- false
|
|
||||||
}
|
|
||||||
addErr := svr.ctlManager.Add(ctlB)
|
|
||||||
addDone <- addErr
|
|
||||||
}()
|
|
||||||
available := waitForResult(t, gateAvailable, "same-run replacement gate probe")
|
|
||||||
require.False(t, available, "same-run gate was available to replacement during LoginResp write")
|
|
||||||
select {
|
|
||||||
case addErr := <-addDone:
|
|
||||||
t.Fatalf("same-run replacement completed during LoginResp write: %v", addErr)
|
|
||||||
case <-time.After(20 * time.Millisecond):
|
|
||||||
}
|
|
||||||
require.Same(t, ctlA, currentControlForTest(svr.ctlManager, "shared-run"))
|
|
||||||
|
|
||||||
otherMetrics := newCountingServerMetrics()
|
|
||||||
otherCtl, otherConn := newLifecycleTestControl(t, "other-run", "other-client", otherMetrics)
|
|
||||||
type unrelatedResult struct {
|
|
||||||
addErr error
|
|
||||||
active bool
|
|
||||||
activateErr error
|
|
||||||
loginErr error
|
|
||||||
current *Control
|
|
||||||
found bool
|
|
||||||
}
|
|
||||||
unrelatedDone := make(chan unrelatedResult, 1)
|
|
||||||
go func() {
|
|
||||||
result := unrelatedResult{}
|
|
||||||
result.addErr = svr.ctlManager.Add(otherCtl)
|
|
||||||
if result.addErr == nil {
|
|
||||||
result.active, result.activateErr = svr.ctlManager.Activate(otherCtl)
|
|
||||||
}
|
|
||||||
if result.activateErr == nil && result.active {
|
|
||||||
result.loginErr = svr.completeControlLogin(otherCtl, func() error { return nil })
|
|
||||||
}
|
|
||||||
result.current, result.found = svr.ctlManager.GetByID("other-run")
|
|
||||||
unrelatedDone <- result
|
|
||||||
}()
|
|
||||||
result := waitForResult(t, unrelatedDone, "unrelated run lifecycle")
|
|
||||||
require.NoError(t, result.addErr)
|
|
||||||
require.NoError(t, result.activateErr)
|
|
||||||
require.True(t, result.active)
|
|
||||||
require.NoError(t, result.loginErr)
|
|
||||||
require.True(t, result.found)
|
|
||||||
require.Same(t, otherCtl, result.current)
|
|
||||||
waitForSignal(t, otherConn.readStarted, "unrelated control reader to start")
|
|
||||||
require.Equal(t, int64(1), otherMetrics.newClients())
|
|
||||||
|
|
||||||
resume()
|
|
||||||
require.NoError(t, waitForResult(t, loginDone, "LoginResp completion"))
|
|
||||||
require.NoError(t, waitForResult(t, addDone, "replacement"))
|
|
||||||
waitForControlDone(t, ctlA)
|
|
||||||
require.Same(t, ctlB, currentControlForTest(svr.ctlManager, "shared-run"))
|
|
||||||
require.Equal(t, int64(1), writeCount.Load())
|
|
||||||
require.Equal(t, int64(1), metrics.newClients())
|
|
||||||
require.Equal(t, int64(1), metrics.closedClients())
|
|
||||||
require.Equal(t, []string{"deadline", "close"}, connA.eventsSnapshot())
|
|
||||||
|
|
||||||
require.False(t, svr.ctlManager.Remove(ctlA))
|
|
||||||
require.NoError(t, ctlA.Close())
|
|
||||||
require.True(t, svr.ctlManager.Remove(ctlB))
|
|
||||||
require.NoError(t, ctlB.Close())
|
|
||||||
require.Equal(t, []string{"deadline", "close"}, connB.eventsSnapshot())
|
|
||||||
|
|
||||||
require.NoError(t, otherCtl.Close())
|
|
||||||
waitForControlDone(t, otherCtl)
|
|
||||||
require.Equal(t, int64(1), otherMetrics.newClients())
|
|
||||||
require.Equal(t, int64(1), otherMetrics.closedClients())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceVisitorAdmissionSerializesReplacement(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
ctlA, controlConn, err := registerLifecycleTestControl(svr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
ctlA.sessionCtx.LoginMsg.User = "old-user"
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctlA, func() error { return nil }))
|
|
||||||
waitForSignal(t, controlConn.readStarted, "A reader to start")
|
|
||||||
|
|
||||||
admissionEntered := make(chan struct{})
|
|
||||||
resumeAdmission := make(chan struct{})
|
|
||||||
var resumeOnce sync.Once
|
|
||||||
resume := func() {
|
|
||||||
resumeOnce.Do(func() { close(resumeAdmission) })
|
|
||||||
}
|
|
||||||
t.Cleanup(resume)
|
|
||||||
type admissionResult struct {
|
|
||||||
admitted bool
|
|
||||||
user string
|
|
||||||
wireProtocol string
|
|
||||||
udpPacketCodec string
|
|
||||||
err error
|
|
||||||
}
|
|
||||||
admissionDone := make(chan admissionResult, 1)
|
|
||||||
go func() {
|
|
||||||
result := admissionResult{}
|
|
||||||
result.admitted, result.err = svr.ctlManager.admitVisitorByRunID("shared-run", func(user, wireProtocol, udpPacketCodec string) error {
|
|
||||||
result.user = user
|
|
||||||
result.wireProtocol = wireProtocol
|
|
||||||
result.udpPacketCodec = udpPacketCodec
|
|
||||||
close(admissionEntered)
|
|
||||||
<-resumeAdmission
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
admissionDone <- result
|
|
||||||
}()
|
|
||||||
waitForSignal(t, admissionEntered, "visitor admission callback")
|
|
||||||
runMu := currentRunGateForTest(svr.ctlManager, "shared-run")
|
|
||||||
require.NotNil(t, runMu)
|
|
||||||
|
|
||||||
type registerResult struct {
|
|
||||||
ctl *Control
|
|
||||||
err error
|
|
||||||
}
|
|
||||||
gateAvailable := make(chan bool)
|
|
||||||
replacementDone := make(chan registerResult, 1)
|
|
||||||
go func() {
|
|
||||||
if runMu.TryLock() {
|
|
||||||
runMu.Unlock()
|
|
||||||
gateAvailable <- true
|
|
||||||
} else {
|
|
||||||
gateAvailable <- false
|
|
||||||
}
|
|
||||||
ctl, _, registerErr := registerLifecycleTestControl(svr)
|
|
||||||
replacementDone <- registerResult{ctl: ctl, err: registerErr}
|
|
||||||
}()
|
|
||||||
available := waitForResult(t, gateAvailable, "visitor replacement gate probe")
|
|
||||||
require.False(t, available, "same-run gate was available during visitor admission")
|
|
||||||
select {
|
|
||||||
case result := <-replacementDone:
|
|
||||||
t.Fatalf("replacement completed during visitor admission: %v", result.err)
|
|
||||||
case <-time.After(20 * time.Millisecond):
|
|
||||||
}
|
|
||||||
require.Same(t, ctlA, currentControlForTest(svr.ctlManager, "shared-run"))
|
|
||||||
|
|
||||||
resume()
|
|
||||||
admission := waitForResult(t, admissionDone, "visitor admission")
|
|
||||||
require.NoError(t, admission.err)
|
|
||||||
require.True(t, admission.admitted)
|
|
||||||
require.Equal(t, "old-user", admission.user)
|
|
||||||
require.Equal(t, wire.ProtocolV1, admission.wireProtocol)
|
|
||||||
require.Empty(t, admission.udpPacketCodec)
|
|
||||||
replacement := waitForResult(t, replacementDone, "replacement")
|
|
||||||
require.NoError(t, replacement.err)
|
|
||||||
ctlB := replacement.ctl
|
|
||||||
require.Same(t, ctlB, currentControlForTest(svr.ctlManager, "shared-run"))
|
|
||||||
waitForControlDone(t, ctlA)
|
|
||||||
require.True(t, svr.ctlManager.Remove(ctlB))
|
|
||||||
require.NoError(t, ctlB.Close())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
ctl, controlConn, err := registerLifecycleTestControl(svr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
pendingConn := newCountingCloseConn()
|
|
||||||
pendingMsgConn := msg.NewConn(pendingConn, msg.NewV1ReadWriter(pendingConn))
|
|
||||||
err = registerWorkConnAsCaller(svr, pendingMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false)
|
|
||||||
require.Error(t, err)
|
|
||||||
require.Equal(t, int64(1), pendingConn.closeCount.Load())
|
|
||||||
require.Len(t, ctl.workConnCh, 0)
|
|
||||||
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
|
|
||||||
waitForSignal(t, controlConn.readStarted, "control reader to start")
|
|
||||||
current, ok := svr.ctlManager.GetByID("shared-run")
|
|
||||||
require.True(t, ok)
|
|
||||||
require.Same(t, ctl, current)
|
|
||||||
require.Len(t, ctl.workConnCh, 0)
|
|
||||||
|
|
||||||
runningConn := newCountingCloseConn()
|
|
||||||
runningMsgConn := msg.NewConn(runningConn, msg.NewV1ReadWriter(runningConn))
|
|
||||||
require.NoError(t, svr.RegisterWorkConn(runningMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false))
|
|
||||||
require.Len(t, ctl.workConnCh, 1)
|
|
||||||
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
require.Equal(t, int64(1), runningConn.closeCount.Load())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceWorkConnRoutingRejectsWireProtocolMismatch(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
ctl, controlConn, err := registerLifecycleTestControl(svr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
|
|
||||||
waitForSignal(t, controlConn.readStarted, "control reader to start")
|
|
||||||
|
|
||||||
workConn := newCountingCloseConn()
|
|
||||||
workMsgConn := msg.NewConn(workConn, msg.NewV2ReadWriter(workConn))
|
|
||||||
err = svr.RegisterWorkConn(workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV2, false)
|
|
||||||
require.ErrorContains(t, err, "wire protocol mismatch")
|
|
||||||
require.Len(t, ctl.workConnCh, 0)
|
|
||||||
_ = workMsgConn.Close()
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceWorkConnRoutingClientHelloPolicy(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
controlUDPPacketCodec string
|
|
||||||
workClientHelloPresent bool
|
|
||||||
errorSubstring string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "JSON control allows work connection without Hello",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "binary control allows work connection without Hello",
|
|
||||||
controlUDPPacketCodec: wire.UDPPacketCodecBinary,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "JSON control rejects work connection with Hello",
|
|
||||||
workClientHelloPresent: true,
|
|
||||||
errorSubstring: "ClientHello is not allowed",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "binary control rejects work connection with Hello",
|
|
||||||
controlUDPPacketCodec: wire.UDPPacketCodecBinary,
|
|
||||||
workClientHelloPresent: true,
|
|
||||||
errorSubstring: "ClientHello is not allowed",
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
controlConn := newDeadlineReadConn()
|
|
||||||
controlMsgConn := msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn))
|
|
||||||
ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{
|
|
||||||
RunID: "shared-run",
|
|
||||||
ClientID: "client",
|
|
||||||
ClientSpec: msg.ClientSpec{
|
|
||||||
AlwaysAuthPass: true,
|
|
||||||
},
|
|
||||||
}, true, wire.ProtocolV2, tc.controlUDPPacketCodec)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
|
|
||||||
waitForSignal(t, controlConn.readStarted, "control reader to start")
|
|
||||||
|
|
||||||
workConn := newCountingCloseConn()
|
|
||||||
workMsgConn := msg.NewConn(workConn, msg.NewV2ReadWriter(workConn))
|
|
||||||
err = svr.RegisterWorkConn(
|
|
||||||
workMsgConn,
|
|
||||||
&msg.NewWorkConn{RunID: "shared-run"},
|
|
||||||
wire.ProtocolV2,
|
|
||||||
tc.workClientHelloPresent,
|
|
||||||
)
|
|
||||||
if tc.errorSubstring != "" {
|
|
||||||
require.ErrorContains(t, err, tc.errorSubstring)
|
|
||||||
require.Len(t, ctl.workConnCh, 0)
|
|
||||||
require.NoError(t, workMsgConn.Close())
|
|
||||||
} else {
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Len(t, ctl.workConnCh, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
require.Equal(t, int64(1), workConn.closeCount.Load())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceRegisterControlRejectsInvalidCodecSelection(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
wireProtocol string
|
|
||||||
udpPacketCodec string
|
|
||||||
errorSubstring string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "binary codec over v1",
|
|
||||||
wireProtocol: wire.ProtocolV1,
|
|
||||||
udpPacketCodec: wire.UDPPacketCodecBinary,
|
|
||||||
errorSubstring: "requires wire protocol v2",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "unknown v2 codec",
|
|
||||||
wireProtocol: wire.ProtocolV2,
|
|
||||||
udpPacketCodec: "unknown",
|
|
||||||
errorSubstring: "unsupported UDP packet codec",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "unknown wire protocol",
|
|
||||||
wireProtocol: "unknown",
|
|
||||||
errorSubstring: "unsupported wire protocol",
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
conn := newDeadlineReadConn()
|
|
||||||
msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn))
|
|
||||||
ctl, err := svr.RegisterControl(msgConn, &msg.Login{}, true, tc.wireProtocol, tc.udpPacketCodec)
|
|
||||||
require.Nil(t, ctl)
|
|
||||||
require.ErrorContains(t, err, tc.errorSubstring)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceRegisterControlRejectsInvalidRunID(t *testing.T) {
|
|
||||||
for _, runID := range []string{
|
|
||||||
"run\nforged",
|
|
||||||
strings.Repeat("a", validation.MaxRunIDLength+1),
|
|
||||||
} {
|
|
||||||
t.Run(fmt.Sprintf("run_id_%d", len(runID)), func(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
conn := newDeadlineReadConn()
|
|
||||||
msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn))
|
|
||||||
ctl, err := svr.RegisterControl(msgConn, &msg.Login{RunID: runID}, true, wire.ProtocolV1, "")
|
|
||||||
require.Nil(t, ctl)
|
|
||||||
require.ErrorContains(t, err, "invalid run id")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceRegisterControlPoolCountBoundaries(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
poolCount int
|
|
||||||
wantErr bool
|
|
||||||
wantPool int
|
|
||||||
}{
|
|
||||||
{name: "less than channel offset", poolCount: -11, wantErr: true},
|
|
||||||
{name: "channel offset", poolCount: -10, wantErr: true},
|
|
||||||
{name: "negative", poolCount: -1, wantErr: true},
|
|
||||||
{name: "zero", poolCount: 0, wantPool: 0},
|
|
||||||
{name: "capped by server maximum", poolCount: 10, wantPool: 5},
|
|
||||||
{name: "maximum int capped", poolCount: math.MaxInt, wantPool: 5},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
svr.cfg.Transport.MaxPoolCount = 5
|
|
||||||
conn := newDeadlineReadConn()
|
|
||||||
msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn))
|
|
||||||
const timestamp = int64(1)
|
|
||||||
|
|
||||||
ctl, err := svr.RegisterControl(msgConn, &msg.Login{
|
|
||||||
RunID: "pool-count-run",
|
|
||||||
ClientID: "client",
|
|
||||||
Timestamp: timestamp,
|
|
||||||
PrivilegeKey: util.GetAuthKey("", timestamp),
|
|
||||||
PoolCount: tc.poolCount,
|
|
||||||
}, false, wire.ProtocolV1, "")
|
|
||||||
if tc.wantErr {
|
|
||||||
require.Nil(t, ctl)
|
|
||||||
require.ErrorContains(t, err, "unexpected error when creating new controller")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, tc.wantPool, ctl.poolCount)
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceWorkConnRoutingRejectsLostGeneration(t *testing.T) {
|
|
||||||
for _, action := range []string{"replace", "close"} {
|
|
||||||
t.Run(action, func(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
ctl, controlConn, err := registerLifecycleTestControl(svr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
|
|
||||||
waitForSignal(t, controlConn.readStarted, "control reader to start")
|
|
||||||
|
|
||||||
barrier := newWorkConnBarrierPlugin()
|
|
||||||
svr.pluginManager.Register(barrier)
|
|
||||||
workConn := newCountingCloseConn()
|
|
||||||
workMsgConn := msg.NewConn(workConn, msg.NewV1ReadWriter(workConn))
|
|
||||||
routeDone := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
routeDone <- registerWorkConnAsCaller(svr, workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false)
|
|
||||||
}()
|
|
||||||
waitForSignal(t, barrier.entered, "work connection plugin barrier")
|
|
||||||
|
|
||||||
var replacement *Control
|
|
||||||
switch action {
|
|
||||||
case "replace":
|
|
||||||
replacement, _, err = registerLifecycleTestControl(svr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
case "close":
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
}
|
|
||||||
|
|
||||||
close(barrier.resume)
|
|
||||||
require.Error(t, waitForResult(t, routeDone, "work connection route to finish"))
|
|
||||||
require.Equal(t, int64(1), workConn.closeCount.Load())
|
|
||||||
require.Len(t, ctl.workConnCh, 0)
|
|
||||||
|
|
||||||
if replacement != nil {
|
|
||||||
require.Len(t, replacement.workConnCh, 0)
|
|
||||||
require.True(t, svr.ctlManager.Remove(replacement))
|
|
||||||
require.NoError(t, replacement.Close())
|
|
||||||
waitForControlDone(t, replacement)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceVisitorRoutingExcludesPendingUser(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{"pending-user"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = listener.Close() })
|
|
||||||
|
|
||||||
controlConn := newDeadlineReadConn()
|
|
||||||
controlMsgConn := msg.NewConn(controlConn, msg.NewV1ReadWriter(controlConn))
|
|
||||||
ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{
|
|
||||||
RunID: "visitor-run",
|
|
||||||
User: "pending-user",
|
|
||||||
ClientID: "visitor-client",
|
|
||||||
ClientSpec: msg.ClientSpec{
|
|
||||||
AlwaysAuthPass: true,
|
|
||||||
},
|
|
||||||
}, true, wire.ProtocolV1, "")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
timestamp := time.Now().Unix()
|
|
||||||
visitorMsg := &msg.NewVisitorConn{
|
|
||||||
RunID: "visitor-run",
|
|
||||||
ProxyName: "visitor",
|
|
||||||
Timestamp: timestamp,
|
|
||||||
SignKey: util.GetAuthKey("secret", timestamp),
|
|
||||||
}
|
|
||||||
pendingConn := newCountingCloseConn()
|
|
||||||
err = svr.RegisterVisitorConn(pendingConn, visitorMsg, wire.ProtocolV1)
|
|
||||||
require.ErrorContains(t, err, "no client control found")
|
|
||||||
require.NoError(t, pendingConn.Close())
|
|
||||||
require.Equal(t, int64(1), pendingConn.closeCount.Load())
|
|
||||||
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
|
|
||||||
waitForSignal(t, controlConn.readStarted, "control reader to start")
|
|
||||||
runningConn := newCountingCloseConn()
|
|
||||||
require.NoError(t, svr.RegisterVisitorConn(runningConn, visitorMsg, wire.ProtocolV1))
|
|
||||||
accepted, err := listener.Accept()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, accepted.Close())
|
|
||||||
require.Equal(t, int64(1), runningConn.closeCount.Load())
|
|
||||||
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceVisitorRoutingCarriesControlPacketCodec(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{"visitor-user"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = listener.Close() })
|
|
||||||
|
|
||||||
controlConn := newDeadlineReadConn()
|
|
||||||
controlMsgConn := msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn))
|
|
||||||
ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{
|
|
||||||
RunID: "visitor-binary-run",
|
|
||||||
User: "visitor-user",
|
|
||||||
ClientID: "visitor-client",
|
|
||||||
ClientSpec: msg.ClientSpec{
|
|
||||||
AlwaysAuthPass: true,
|
|
||||||
},
|
|
||||||
}, true, wire.ProtocolV2, wire.UDPPacketCodecBinary)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
timestamp := time.Now().Unix()
|
|
||||||
visitorMsg := &msg.NewVisitorConn{
|
|
||||||
RunID: "visitor-binary-run",
|
|
||||||
ProxyName: "visitor",
|
|
||||||
Timestamp: timestamp,
|
|
||||||
SignKey: util.GetAuthKey("secret", timestamp),
|
|
||||||
}
|
|
||||||
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
|
|
||||||
waitForSignal(t, controlConn.readStarted, "binary visitor control reader to start")
|
|
||||||
|
|
||||||
runningConn := newCountingCloseConn()
|
|
||||||
require.NoError(t, svr.RegisterVisitorConn(runningConn, visitorMsg, wire.ProtocolV2))
|
|
||||||
accepted, err := listener.Accept()
|
|
||||||
require.NoError(t, err)
|
|
||||||
metadata, ok := accepted.(interface {
|
|
||||||
WireProtocol() string
|
|
||||||
UDPPacketCodec() string
|
|
||||||
})
|
|
||||||
require.True(t, ok)
|
|
||||||
require.Equal(t, wire.ProtocolV2, metadata.WireProtocol())
|
|
||||||
require.Equal(t, wire.UDPPacketCodecBinary, metadata.UDPPacketCodec())
|
|
||||||
require.NoError(t, accepted.Close())
|
|
||||||
require.Equal(t, int64(1), runningConn.closeCount.Load())
|
|
||||||
|
|
||||||
mismatchConn := newCountingCloseConn()
|
|
||||||
err = svr.RegisterVisitorConn(mismatchConn, visitorMsg, wire.ProtocolV1)
|
|
||||||
require.ErrorContains(t, err, "visitor connection wire protocol mismatch")
|
|
||||||
require.NoError(t, mismatchConn.Close())
|
|
||||||
require.Equal(t, int64(1), mismatchConn.closeCount.Load())
|
|
||||||
|
|
||||||
require.NoError(t, ctl.Close())
|
|
||||||
waitForControlDone(t, ctl)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceVisitorRoutingLegacyFallsBackToJSONPacketCodec(t *testing.T) {
|
|
||||||
svr := newControlTestService(t)
|
|
||||||
listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{""})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = listener.Close() })
|
|
||||||
|
|
||||||
timestamp := time.Now().Unix()
|
|
||||||
visitorMsg := &msg.NewVisitorConn{
|
|
||||||
ProxyName: "visitor",
|
|
||||||
Timestamp: timestamp,
|
|
||||||
SignKey: util.GetAuthKey("secret", timestamp),
|
|
||||||
}
|
|
||||||
visitorConn := newCountingCloseConn()
|
|
||||||
require.NoError(t, svr.RegisterVisitorConn(visitorConn, visitorMsg, wire.ProtocolV2))
|
|
||||||
accepted, err := listener.Accept()
|
|
||||||
require.NoError(t, err)
|
|
||||||
metadata, ok := accepted.(interface{ UDPPacketCodec() string })
|
|
||||||
require.True(t, ok)
|
|
||||||
require.Empty(t, metadata.UDPPacketCodec())
|
|
||||||
require.NoError(t, accepted.Close())
|
|
||||||
require.Equal(t, int64(1), visitorConn.closeCount.Load())
|
|
||||||
}
|
|
||||||
|
|
||||||
func newControlTestService(t *testing.T) *Service {
|
|
||||||
t.Helper()
|
|
||||||
cfg := &v1.ServerConfig{}
|
|
||||||
cfg.Auth.Method = v1.AuthMethodToken
|
|
||||||
authRuntime, err := auth.BuildServerAuth(&cfg.Auth)
|
|
||||||
require.NoError(t, err)
|
|
||||||
clientRegistry := registry.NewClientRegistry()
|
|
||||||
return &Service{
|
|
||||||
ctlManager: NewControlManager(clientRegistry),
|
|
||||||
clientRegistry: clientRegistry,
|
|
||||||
pxyManager: proxy.NewManager(),
|
|
||||||
pluginManager: plugin.NewManager(),
|
|
||||||
rc: &controller.ResourceController{
|
|
||||||
VisitorManager: visitor.NewManager(),
|
|
||||||
},
|
|
||||||
auth: authRuntime,
|
|
||||||
cfg: cfg,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func registerLifecycleTestControl(svr *Service) (*Control, *deadlineReadConn, error) {
|
|
||||||
conn := newDeadlineReadConn()
|
|
||||||
msgConn := msg.NewConn(conn, msg.NewReadWriter(conn, wire.ProtocolV1))
|
|
||||||
ctl, err := svr.RegisterControl(msgConn, &msg.Login{
|
|
||||||
RunID: "shared-run",
|
|
||||||
ClientID: "client",
|
|
||||||
ClientSpec: msg.ClientSpec{
|
|
||||||
AlwaysAuthPass: true,
|
|
||||||
},
|
|
||||||
}, true, wire.ProtocolV1, "")
|
|
||||||
return ctl, conn, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func waitForDifferentCurrentControl(t *testing.T, manager *ControlManager, runID string, old *Control) *Control {
|
|
||||||
t.Helper()
|
|
||||||
deadline := time.Now().Add(3 * time.Second)
|
|
||||||
for time.Now().Before(deadline) {
|
|
||||||
if ctl := currentControlForTest(manager, runID); ctl != nil && ctl != old {
|
|
||||||
return ctl
|
|
||||||
}
|
|
||||||
runtime.Gosched()
|
|
||||||
}
|
|
||||||
t.Fatalf("timed out waiting for a new current control after ID %d", old.ID())
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func registerWorkConnAsCaller(
|
|
||||||
svr *Service,
|
|
||||||
workConn *msg.Conn,
|
|
||||||
newMsg *msg.NewWorkConn,
|
|
||||||
wireProtocol string,
|
|
||||||
clientHelloPresent bool,
|
|
||||||
) error {
|
|
||||||
err := svr.RegisterWorkConn(workConn, newMsg, wireProtocol, clientHelloPresent)
|
|
||||||
if err != nil {
|
|
||||||
_ = workConn.Close()
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func waitForResult[T any](t *testing.T, ch <-chan T, description string) T {
|
|
||||||
t.Helper()
|
|
||||||
select {
|
|
||||||
case result := <-ch:
|
|
||||||
return result
|
|
||||||
case <-time.After(3 * time.Second):
|
|
||||||
t.Fatalf("timed out waiting for %s", description)
|
|
||||||
var zero T
|
|
||||||
return zero
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type workConnBarrierPlugin struct {
|
|
||||||
entered chan struct{}
|
|
||||||
resume chan struct{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newWorkConnBarrierPlugin() *workConnBarrierPlugin {
|
|
||||||
return &workConnBarrierPlugin{
|
|
||||||
entered: make(chan struct{}),
|
|
||||||
resume: make(chan struct{}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*workConnBarrierPlugin) Name() string { return "work-conn-barrier" }
|
|
||||||
|
|
||||||
func (*workConnBarrierPlugin) IsSupport(op string) bool { return op == plugin.OpNewWorkConn }
|
|
||||||
|
|
||||||
func (p *workConnBarrierPlugin) Handle(
|
|
||||||
context.Context,
|
|
||||||
string,
|
|
||||||
any,
|
|
||||||
) (*plugin.Response, any, error) {
|
|
||||||
close(p.entered)
|
|
||||||
<-p.resume
|
|
||||||
return &plugin.Response{Unchange: true}, nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type countingCloseConn struct {
|
|
||||||
closeCount atomic.Int64
|
|
||||||
}
|
|
||||||
|
|
||||||
func newCountingCloseConn() *countingCloseConn { return &countingCloseConn{} }
|
|
||||||
|
|
||||||
func (*countingCloseConn) Read([]byte) (int, error) { return 0, net.ErrClosed }
|
|
||||||
func (*countingCloseConn) Write(p []byte) (int, error) { return len(p), nil }
|
|
||||||
func (c *countingCloseConn) Close() error { c.closeCount.Add(1); return nil }
|
|
||||||
func (*countingCloseConn) LocalAddr() net.Addr { return lifecycleTestAddr("local") }
|
|
||||||
func (*countingCloseConn) RemoteAddr() net.Addr { return lifecycleTestAddr("remote") }
|
|
||||||
func (*countingCloseConn) SetDeadline(time.Time) error { return nil }
|
|
||||||
func (*countingCloseConn) SetReadDeadline(time.Time) error { return nil }
|
|
||||||
func (*countingCloseConn) SetWriteDeadline(time.Time) error { return nil }
|
|
||||||
|
|||||||
@@ -65,12 +65,8 @@ func (vm *Manager) Listen(name string, sk string, allowUsers []string) (*netpkg.
|
|||||||
|
|
||||||
func (vm *Manager) NewConn(name string, conn net.Conn, timestamp int64, signKey string,
|
func (vm *Manager) NewConn(name string, conn net.Conn, timestamp int64, signKey string,
|
||||||
useEncryption bool, useCompression bool, visitorUser string,
|
useEncryption bool, useCompression bool, visitorUser string,
|
||||||
wireProtocol string, udpPacketCodecs ...string,
|
wireProtocol string,
|
||||||
) (err error) {
|
) (err error) {
|
||||||
udpPacketCodec := ""
|
|
||||||
if len(udpPacketCodecs) > 0 {
|
|
||||||
udpPacketCodec = udpPacketCodecs[0]
|
|
||||||
}
|
|
||||||
vm.mu.RLock()
|
vm.mu.RLock()
|
||||||
defer vm.mu.RUnlock()
|
defer vm.mu.RUnlock()
|
||||||
|
|
||||||
@@ -97,9 +93,8 @@ func (vm *Manager) NewConn(name string, conn net.Conn, timestamp int64, signKey
|
|||||||
}
|
}
|
||||||
visitorConn := netpkg.WrapReadWriteCloserToConn(rwc, conn)
|
visitorConn := netpkg.WrapReadWriteCloserToConn(rwc, conn)
|
||||||
err = l.l.PutConn(&wireProtocolConn{
|
err = l.l.PutConn(&wireProtocolConn{
|
||||||
Conn: visitorConn,
|
Conn: visitorConn,
|
||||||
wireProtocol: wireProtocol,
|
wireProtocol: wireProtocol,
|
||||||
udpPacketCodec: udpPacketCodec,
|
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
err = fmt.Errorf("custom listener for [%s] doesn't exist", name)
|
err = fmt.Errorf("custom listener for [%s] doesn't exist", name)
|
||||||
@@ -110,18 +105,13 @@ func (vm *Manager) NewConn(name string, conn net.Conn, timestamp int64, signKey
|
|||||||
|
|
||||||
type wireProtocolConn struct {
|
type wireProtocolConn struct {
|
||||||
net.Conn
|
net.Conn
|
||||||
wireProtocol string
|
wireProtocol string
|
||||||
udpPacketCodec string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *wireProtocolConn) WireProtocol() string {
|
func (c *wireProtocolConn) WireProtocol() string {
|
||||||
return c.wireProtocol
|
return c.wireProtocol
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *wireProtocolConn) UDPPacketCodec() string {
|
|
||||||
return c.udpPacketCodec
|
|
||||||
}
|
|
||||||
|
|
||||||
func (vm *Manager) CloseListener(name string) {
|
func (vm *Manager) CloseListener(name string) {
|
||||||
vm.mu.Lock()
|
vm.mu.Lock()
|
||||||
defer vm.mu.Unlock()
|
defer vm.mu.Unlock()
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user