Compare commits

..
11 Commits
Author SHA1 Message Date
fatedierandGitHub 7b6e01f04f Merge pull request #5411 from fatedier/dev
Release v0.70.0
2026-07-11 18:50:35 +08:00
fatedierandGitHub 8dd26c6961 Merge pull request #5350 from fatedier/dev
Release v0.69.1
2026-06-01 18:02:22 +08:00
fatedierandGitHub c8c1e5116c Merge pull request #5323 from fatedier/dev
Release v0.69.0
2026-05-22 00:55:23 +08:00
fatedierandGitHub 4ec8de973f Merge pull request #5287 from fatedier/dev
bump version to v0.68.1
2026-04-14 01:28:33 +08:00
fatedierandGitHub 5bfcea3d0c merge dev to master (#5254)
* ci: bump github actions to latest major versions (#5251)

* docker: copy shared web directory for npm workspace builds
2026-03-20 15:54:26 +08:00
fatedierandGitHub 0a1b4ab21f Merge pull request #5249 from fatedier/dev
bump version
2026-03-20 13:56:28 +08:00
fatedierandGitHub 5f575b8442 Merge pull request #5147 from fatedier/dev
bump version
2026-01-31 14:01:40 +08:00
fatedierandGitHub a1348cdf00 bump version (#5112) 2026-01-04 14:54:13 +08:00
fatedierandGitHub 2f5e1f7945 Merge pull request #4999 from fatedier/dev
bump version
2025-09-25 20:23:42 +08:00
fatedierandGitHub 22ae8166d3 Merge pull request #4925 from fatedier/dev
bump version
2025-08-10 23:26:32 +08:00
fatedierandGitHub af6bc6369d Merge pull request #4849 from fatedier/dev
bump version
2025-06-25 11:51:19 +08:00
118 changed files with 1608 additions and 9665 deletions
+7 -8
View File
@@ -7,15 +7,14 @@ jobs:
steps:
- checkout
- run:
name: Test and build web assets
command: make web-ci
name: Build web assets (frps)
command: make install build
working_directory: web/frps
- run:
name: Check Go formatting and build binaries
command: |
set -e
make env fmt
git diff --exit-code
make build
name: Build web assets (frpc)
command: make install build
working_directory: web/frpc
- run: make
- run: make alltest
workflows:
+2 -2
View File
@@ -65,7 +65,7 @@ jobs:
with:
context: .
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
tags: |
${{ env.TAG_FRPC }}
@@ -76,7 +76,7 @@ jobs:
with:
context: .
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
tags: |
${{ env.TAG_FRPS }}
+7 -3
View File
@@ -22,10 +22,14 @@ jobs:
- uses: actions/setup-node@v6
with:
node-version: '22'
- name: Test and build web assets
run: make web-ci
- name: Build web assets (frps)
run: make build
working-directory: web/frps
- name: Build web assets (frpc)
run: make build
working-directory: web/frpc
- name: golangci-lint
uses: golangci/golangci-lint-action@v9
with:
# 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
+1 -4
View File
@@ -5,7 +5,7 @@ NOWEB_TAG = $(shell [ ! -d web/frps/dist ] || [ ! -d web/frpc/dist ] && echo ',n
FRP_COMPAT_BASELINE_COUNT ?= 8
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
@@ -16,9 +16,6 @@ env:
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:
$(MAKE) -C web/frps build
+1 -1
View File
@@ -9,7 +9,7 @@ all: build
build: app
app:
@set -e; $(foreach n, $(os-archs), \
@$(foreach n, $(os-archs), \
os=$(shell echo "$(n)" | cut -d : -f 1); \
arch=$(shell echo "$(n)" | cut -d : -f 2); \
extra=$(shell echo "$(n)" | cut -d : -f 3); \
+10 -11
View File
@@ -12,6 +12,16 @@ frp is an open source project with its ongoing development made possible entirel
<h3 align="center">Gold Sponsors</h3>
<!--gold sponsors start-->
<div align="center">
## Recall.ai - API for meeting recordings
If you're looking for a meeting recording API, consider checking out [Recall.ai](https://www.recall.ai/?utm_source=github&utm_medium=sponsorship&utm_campaign=fatedier-frp),
an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
</div>
<p align="center">
<a href="https://jb.gg/frp" target="_blank">
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
@@ -29,17 +39,6 @@ frp is an open source project with its ongoing development made possible entirel
<sub>An open source, self-hosted alternative to public clouds, built for data ownership and privacy</sub>
</a>
</p>
<div align="center">
## Recall.ai - API for meeting recordings
If you're looking for a meeting recording API, consider checking out [Recall.ai](https://www.recall.ai/?utm_source=github&utm_medium=sponsorship&utm_campaign=fatedier-frp),
an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
</div>
<!--gold sponsors end-->
## What is frp?
+11 -11
View File
@@ -2,6 +2,7 @@
[![Build Status](https://circleci.com/gh/fatedier/frp.svg?style=shield)](https://circleci.com/gh/fatedier/frp)
[![GitHub release](https://img.shields.io/github/tag/fatedier/frp.svg?label=release)](https://github.com/fatedier/frp/releases)
[![Go Report Card](https://goreportcard.com/badge/github.com/fatedier/frp)](https://goreportcard.com/report/github.com/fatedier/frp)
[![GitHub Releases Stats](https://img.shields.io/github/downloads/fatedier/frp/total.svg?logo=github)](https://somsubhra.github.io/github-release-stats/?username=fatedier&repository=frp)
[README](README.md) | [中文文档](README_zh.md)
@@ -14,6 +15,16 @@ frp 是一个完全开源的项目,我们的开发工作完全依靠赞助者
<h3 align="center">Gold Sponsors</h3>
<!--gold sponsors start-->
<div align="center">
## Recall.ai - API for meeting recordings
If you're looking for a meeting recording API, consider checking out [Recall.ai](https://www.recall.ai/?utm_source=github&utm_medium=sponsorship&utm_campaign=fatedier-frp),
an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
</div>
<p align="center">
<a href="https://jb.gg/frp" target="_blank">
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
@@ -31,17 +42,6 @@ frp 是一个完全开源的项目,我们的开发工作完全依靠赞助者
<sub>An open source, self-hosted alternative to public clouds, built for data ownership and privacy</sub>
</a>
</p>
<div align="center">
## Recall.ai - API for meeting recordings
If you're looking for a meeting recording API, consider checking out [Recall.ai](https://www.recall.ai/?utm_source=github&utm_medium=sponsorship&utm_campaign=fatedier-frp),
an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
</div>
<!--gold sponsors end-->
## 为什么使用 frp
+7 -1
View File
@@ -1,3 +1,9 @@
## Features
* Expanded the frps dashboard API v2 migration across Clients, Proxies, Server Overview, Client Detail, and Proxy Detail, covering paginated users/clients/proxies, detail data, proxy traffic history, server system info, offline proxy statistics pruning, server-side pagination, search, and proxy type filtering.
## Fixes
* Fixed VirtualNet route lifecycle issues during reconnect and shutdown, including stale route cleanup, shutdown races, and reconnect backoff overflow.
* WebSocket and WSS tunnel payloads are now sent as binary frames, avoiding disconnects through RFC-compliant intermediaries that validate text frames as UTF-8.
* The `tls2raw` client plugin now writes the proxy protocol header to the local raw connection when proxy protocol is enabled.
* frpc now rejects duplicate proxy and visitor names in config files instead of silently overwriting earlier entries.
-254
View File
@@ -2,16 +2,12 @@ package client
import (
"errors"
"os"
"path/filepath"
"strings"
"testing"
"github.com/fatedier/frp/client/configmgmt"
"github.com/fatedier/frp/pkg/config/source"
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 {
@@ -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) {
storeSource, err := source.NewStoreSource(source.StoreSourceConfig{
Path: filepath.Join(t.TempDir(), "store.json"),
+1 -1
View File
@@ -24,7 +24,7 @@ import (
"time"
libnet "github.com/fatedier/golib/net"
fmux "github.com/fatedier/yamux"
fmux "github.com/hashicorp/yamux"
quic "github.com/quic-go/quic-go"
"github.com/samber/lo"
+2 -11
View File
@@ -47,8 +47,6 @@ type SessionContext struct {
Connector MessageConnector
// Virtual net controller
VnetController *vnet.Controller
// UDPPacketCodec is immutable for the lifetime of this negotiated session.
UDPPacketCodec string
}
type Control struct {
@@ -94,16 +92,9 @@ func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, erro
ctl.registerMsgHandlers()
ctl.msgTransporter = transport.NewMessageTransporter(ctl.msgDispatcher)
ctl.pm = proxy.NewManager(
ctl.ctx,
sessionCtx.Common,
sessionCtx.Auth.EncryptionKey(),
ctl.msgTransporter,
sessionCtx.VnetController,
sessionCtx.UDPPacketCodec,
)
ctl.pm = proxy.NewManager(ctl.ctx, sessionCtx.Common, sessionCtx.Auth.EncryptionKey(), ctl.msgTransporter, sessionCtx.VnetController)
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
}
+4 -9
View File
@@ -99,7 +99,6 @@ func (d *controlSessionDialer) Dial(previousRunID string) (*SessionContext, erro
Auth: d.auth,
Connector: newMessageConnector(connector, d.common.Transport.WireProtocol),
VnetController: d.vnetController,
UDPPacketCodec: loginResult.udpPacketCodec,
}, nil
}
@@ -128,9 +127,8 @@ func (d *controlSessionDialer) buildLoginMsg(previousRunID string) (*msg.Login,
}
type loginExchangeResult struct {
resp *msg.LoginResp
crypto *wire.CryptoContext
udpPacketCodec string
resp *msg.LoginResp
crypto *wire.CryptoContext
}
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 udpPacketCodec string
if wireConn != nil {
serverHelloFrame, err := wireConn.ReadFrame()
if err != nil {
@@ -194,7 +191,6 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login)
if err != nil {
return nil, err
}
udpPacketCodec = serverHello.Selected.Message.UDPPacketCodec
}
var loginRespMsg msg.LoginResp
@@ -202,9 +198,8 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login)
return nil, err
}
return &loginExchangeResult{
resp: &loginRespMsg,
crypto: cryptoContext,
udpPacketCodec: udpPacketCodec,
resp: &loginRespMsg,
crypto: cryptoContext,
}, nil
}
-2
View File
@@ -117,7 +117,6 @@ func TestControlSessionDialerDialV1(t *testing.T) {
defer sessionCtx.Connector.Close()
require.Equal(t, "run-v1", sessionCtx.RunID)
require.Empty(t, sessionCtx.UDPPacketCodec)
require.NotNil(t, sessionCtx.Conn)
require.NotNil(t, sessionCtx.Connector)
require.False(t, connector.closed.Load())
@@ -226,7 +225,6 @@ func TestControlSessionDialerDialV2(t *testing.T) {
defer sessionCtx.Connector.Close()
require.Equal(t, "run-v2", sessionCtx.RunID)
require.Equal(t, wire.UDPPacketCodecBinary, sessionCtx.UDPPacketCodec)
require.NotNil(t, sessionCtx.Conn)
require.NotNil(t, sessionCtx.Connector)
require.False(t, connector.closed.Load())
-125
View File
@@ -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)
}
-1
View File
@@ -119,7 +119,6 @@ func (monitor *Monitor) checkWorker() {
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
-65
View File
@@ -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 health
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/require"
v1 "github.com/fatedier/frp/pkg/config/v1"
)
func TestMonitorResetsFailedTimesAfterSuccess(t *testing.T) {
var checkCount atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
count := checkCount.Add(1)
if count == 1 || count == 2 || count == 4 {
w.WriteHeader(http.StatusServiceUnavailable)
return
}
w.WriteHeader(http.StatusOK)
}))
defer server.Close()
var failedCount atomic.Int32
monitor := NewMonitor(
context.Background(),
v1.HealthCheckConfig{
Type: "http",
Path: "/health",
TimeoutSeconds: 1,
IntervalSeconds: 1,
MaxFailed: 3,
},
strings.TrimPrefix(server.URL, "http://"),
func() {},
func() { failedCount.Add(1) },
)
monitor.interval = 10 * time.Millisecond
monitor.Start()
defer monitor.Stop()
require.Eventually(t, func() bool {
return checkCount.Load() >= 5
}, time.Second, 10*time.Millisecond)
require.Equal(t, int32(0), failedCount.Load())
}
+9 -27
View File
@@ -61,12 +61,11 @@ func NewProxy(
encryptionKey []byte,
msgTransporter transport.MessageTransporter,
vnetController *vnet.Controller,
udpPacketCodec string,
) (pxy Proxy) {
var limiter *rate.Limiter
limitBytes := pxyConf.GetBaseConfig().Transport.BandwidthLimit.Bytes()
if limitBytes > 0 && pxyConf.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeClient {
limiter = limit.NewBandwidthLimiter(limitBytes)
limiter = rate.NewLimiter(rate.Limit(float64(limitBytes)), int(limitBytes))
}
baseProxy := BaseProxy{
@@ -78,7 +77,6 @@ func NewProxy(
vnetController: vnetController,
xl: xlog.FromContextSafe(ctx),
ctx: ctx,
udpPacketCodec: udpPacketCodec,
}
factory := proxyFactoryRegistry[reflect.TypeOf(pxyConf)]
@@ -100,10 +98,9 @@ type BaseProxy struct {
proxyPlugin plugin.Plugin
inWorkConnCallback func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) /* continue */ bool
mu sync.RWMutex
xl *xlog.Logger
ctx context.Context
udpPacketCodec string
mu sync.RWMutex
xl *xlog.Logger
ctx context.Context
}
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",
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)
if err != nil {
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
var connInfo plugin.ConnectionInfo
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.DstAddr = dstAddr
}
+2 -5
View File
@@ -43,8 +43,7 @@ type Manager struct {
encryptionKey []byte
clientCfg *v1.ClientCommonConfig
ctx context.Context
udpPacketCodec string
ctx context.Context
}
func NewManager(
@@ -53,7 +52,6 @@ func NewManager(
encryptionKey []byte,
msgTransporter transport.MessageTransporter,
vnetController *vnet.Controller,
udpPacketCodec string,
) *Manager {
return &Manager{
proxies: make(map[string]*Wrapper),
@@ -63,7 +61,6 @@ func NewManager(
encryptionKey: encryptionKey,
clientCfg: clientCfg,
ctx: ctx,
udpPacketCodec: udpPacketCodec,
}
}
@@ -169,7 +166,7 @@ func (pm *Manager) UpdateAll(proxyCfgs []v1.ProxyConfigurer) {
for _, cfg := range proxyCfgs {
name := cfg.GetBaseConfig().Name
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 {
pxy.SetInWorkConnCallback(pm.inWorkConnCallback)
}
-47
View File
@@ -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)
}
+1 -2
View File
@@ -99,7 +99,6 @@ func NewWrapper(
eventHandler event.Handler,
msgTransporter transport.MessageTransporter,
vnetController *vnet.Controller,
udpPacketCodec string,
) *Wrapper {
baseInfo := cfg.GetBaseConfig()
xl := xlog.FromContextSafe(ctx).Spawn().AppendPrefix(baseInfo.Name)
@@ -128,7 +127,7 @@ func NewWrapper(
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
}
+1 -7
View File
@@ -87,13 +87,7 @@ func (pxy *SUDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) {
}
workConn := netpkg.WrapReadWriteCloserToConn(remote, conn)
payloadRW, err := msg.NewUDPPacketReadWriter(workConn, pxy.clientCfg.Transport.WireProtocol, pxy.udpPacketCodec)
if err != nil {
xl.Errorf("create SUDP packet read writer: %v", err)
_ = workConn.Close()
return
}
payloadConn := msg.NewConn(workConn, payloadRW)
payloadConn := msg.NewConn(workConn, msg.NewReadWriter(workConn, pxy.clientCfg.Transport.WireProtocol))
readCh := make(chan *msg.UDPPacket, 1024)
sendCh := make(chan msg.Message, 1024)
isClose := false
+3 -10
View File
@@ -97,17 +97,10 @@ func (pxy *UDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) {
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.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.sendCh = make(chan msg.Message, 1024)
pxy.closed = false
+1 -1
View File
@@ -22,7 +22,7 @@ import (
"reflect"
"time"
fmux "github.com/fatedier/yamux"
fmux "github.com/hashicorp/yamux"
"github.com/quic-go/quic-go"
v1 "github.com/fatedier/frp/pkg/config/v1"
+4 -16
View File
@@ -22,7 +22,6 @@ import (
"net/http"
"os"
"sync"
"sync/atomic"
"time"
"github.com/fatedier/golib/crypto"
@@ -33,7 +32,6 @@ import (
"github.com/fatedier/frp/pkg/config"
"github.com/fatedier/frp/pkg/config/source"
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/policy/security"
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.
type Service struct {
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
ctl *Control
// Uniq id got from frps, it will be attached to loginMsg.
@@ -154,7 +149,8 @@ type Service struct {
// service context
ctx context.Context
// call cancel to stop service
cancel context.CancelCauseFunc
cancel context.CancelCauseFunc
gracefulShutdownDuration time.Duration
connectorCreator func(context.Context, *v1.ClientCommonConfig) Connector
handleWorkConnCb func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) bool
@@ -416,7 +412,7 @@ func (svr *Service) Close() {
}
func (svr *Service) GracefulClose(d time.Duration) {
svr.gracefulShutdownDuration.Store(int64(d))
svr.gracefulShutdownDuration = d
svr.cancel(nil)
}
@@ -433,8 +429,7 @@ func (svr *Service) stop() {
svr.ctlMu.Lock()
defer svr.ctlMu.Unlock()
if svr.ctl != nil {
d := time.Duration(svr.gracefulShutdownDuration.Load())
svr.ctl.GracefulClose(d)
svr.ctl.GracefulClose(svr.gracefulShutdownDuration)
svr.ctl = nil
}
if svr.webServer != nil {
@@ -511,13 +506,6 @@ func (svr *Service) reloadConfigFromSourcesLocked() error {
proxies, visitors = config.FilterClientConfigurers(reloadCommon, proxies, visitors)
proxies = config.CompleteProxyConfigurers(proxies)
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
if err := svr.UpdateAllConfigurer(proxies, visitors); err != nil {
-95
View File
@@ -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)
}
}
+1 -7
View File
@@ -113,13 +113,7 @@ func (sv *SUDPVisitor) dispatcher() {
func (sv *SUDPVisitor) worker(workConn net.Conn, firstPacket *msg.UDPPacket) {
xl := xlog.FromContextSafe(sv.ctx)
xl.Debugf("starting sudp proxy worker")
payloadRW, err := msg.NewUDPPacketReadWriter(workConn, sv.clientCfg.Transport.WireProtocol, udpPacketCodecFromHelper(sv.helper))
if err != nil {
xl.Errorf("create SUDP packet read writer: %v", err)
_ = workConn.Close()
return
}
payloadConn := msg.NewConn(workConn, payloadRW)
payloadConn := msg.NewConn(workConn, msg.NewReadWriter(workConn, sv.clientCfg.Transport.WireProtocol))
wg := &sync.WaitGroup{}
wg.Add(2)
-11
View File
@@ -50,17 +50,6 @@ type Helper interface {
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.
type Visitor interface {
Run() error
-11
View File
@@ -53,12 +53,7 @@ func NewManager(
connectServer func() (*msg.Conn, error),
msgTransporter transport.MessageTransporter,
vnetController *vnet.Controller,
udpPacketCodecs ...string,
) *Manager {
udpPacketCodec := ""
if len(udpPacketCodecs) > 0 {
udpPacketCodec = udpPacketCodecs[0]
}
m := &Manager{
clientCfg: clientCfg,
cfgs: make(map[string]v1.VisitorConfigurer),
@@ -73,7 +68,6 @@ func NewManager(
vnetController: vnetController,
transferConnFn: m.TransferConn,
runID: runID,
udpPacketCodec: udpPacketCodec,
}
return m
}
@@ -211,7 +205,6 @@ type visitorHelperImpl struct {
vnetController *vnet.Controller
transferConnFn func(name string, conn net.Conn) error
runID string
udpPacketCodec string
}
func (v *visitorHelperImpl) ConnectServer() (*msg.Conn, error) {
@@ -233,7 +226,3 @@ func (v *visitorHelperImpl) VNetController() *vnet.Controller {
func (v *visitorHelperImpl) RunID() string {
return v.runID
}
func (v *visitorHelperImpl) UDPPacketCodec() string {
return v.udpPacketCodec
}
+1 -1
View File
@@ -25,7 +25,7 @@ import (
"time"
libio "github.com/fatedier/golib/io"
fmux "github.com/fatedier/yamux"
fmux "github.com/hashicorp/yamux"
quic "github.com/quic-go/quic-go"
"golang.org/x/time/rate"
+7
View File
@@ -33,6 +33,7 @@ import (
"github.com/fatedier/frp/pkg/config/source"
v1 "github.com/fatedier/frp/pkg/config/v1"
"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/util/log"
"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")
}
if len(result.Common.FeatureGates) > 0 {
if err := featuregate.SetFromMap(result.Common.FeatureGates); err != nil {
return err
}
}
return runClientWithAggregator(result, unsafeFeatures, cfgFilePath)
}
+6 -13
View File
@@ -29,18 +29,6 @@ func init() {
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{
Use: "verify",
Short: "Verify that the configures is valid",
@@ -50,8 +38,13 @@ var verifyCmd = &cobra.Command{
return nil
}
cliCfg, proxyCfgs, visitorCfgs, _, err := config.LoadClientConfig(cfgFile, strictConfigMode)
if err != nil {
fmt.Println(err)
os.Exit(1)
}
unsafeFeatures := security.NewUnsafeFeatures(allowUnsafe)
warning, err := verifyClientConfig(cfgFile, strictConfigMode, unsafeFeatures)
warning, err := validation.ValidateAllClientConfig(cliCfg, proxyCfgs, visitorCfgs, unsafeFeatures)
if warning != nil {
fmt.Printf("WARNING: %v\n", warning)
}
-67
View File
@@ -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)
})
}
}
+22 -13
View File
@@ -4,18 +4,19 @@ go 1.25.0
require (
github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5
github.com/coreos/go-oidc/v3 v3.18.0
github.com/fatedier/golib v0.8.2
github.com/fatedier/yamux v0.2.0
github.com/coreos/go-oidc/v3 v3.14.1
github.com/fatedier/golib v0.7.0
github.com/google/uuid v1.6.0
github.com/gorilla/mux v1.8.1
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/gomega v1.36.3
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/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/samber/lo v1.47.0
github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8
@@ -25,11 +26,11 @@ require (
github.com/tidwall/gjson v1.17.1
github.com/vishvananda/netlink v1.3.0
github.com/xtaci/kcp-go/v5 v5.6.13
golang.org/x/crypto v0.54.0
golang.org/x/net v0.56.0
golang.org/x/oauth2 v0.36.0
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
golang.org/x/crypto v0.49.0
golang.org/x/net v0.52.0
golang.org/x/oauth2 v0.28.0
golang.org/x/sync v0.20.0
golang.org/x/sys v0.42.0
golang.org/x/time v0.10.0
golang.zx2c4.com/wireguard v0.0.0-20231211153847-12269c276173
gopkg.in/ini.v1 v1.67.0
@@ -43,7 +44,7 @@ require (
github.com/beorn7/perks v1.0.1 // indirect
github.com/cespare/xxhash/v2 v2.2.0 // 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-task/slim-sprig/v3 v3.0.0 // indirect
github.com/golang/snappy v0.0.4 // indirect
@@ -52,6 +53,9 @@ require (
github.com/inconshreveable/mousetrap v1.1.0 // indirect
github.com/klauspost/cpuid/v2 v2.2.6 // 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/pmezard/go-difflib v1.0.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/tjfoc/gmsm v1.4.1 // 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
golang.org/x/text v0.40.0 // indirect
golang.org/x/tools v0.47.0 // indirect
golang.org/x/mod v0.33.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
google.golang.org/protobuf v1.36.5 // 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/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
+40 -30
View File
@@ -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/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/coreos/go-oidc/v3 v3.18.0 h1:V9orjXynvu5wiC9SemFTWnG4F45v403aIcjWo0d41+A=
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 h1:9ePWwfdwC4QKRlCXsJGou56adA/owXczOzwKdOumLqk=
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/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=
@@ -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.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
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.8.2/go.mod h1:ArUGvPg2cOw/py2RAuBt46nNZH2VQ5Z70p109MAZpJw=
github.com/fatedier/yamux v0.2.0 h1:H+2A9iBVh7aJlEOc1Ws1FXWOaecBf2nRv9zpFMPUWg8=
github.com/fatedier/yamux v0.2.0/go.mod h1:d4FtRDrC9sHvRpiDL6J5EnfjLhqzZZplHe5yToSn2Ac=
github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA=
github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
github.com/fatedier/golib v0.7.0 h1:tMDF9ObcwVt59VUHroJOzHQjVFPLymZVMpGm9WAVwhY=
github.com/fatedier/golib v0.7.0/go.mod h1:ArUGvPg2cOw/py2RAuBt46nNZH2VQ5Z70p109MAZpJw=
github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6 h1:u92UUy6FURPmNsMBUuongRWC0rBqN6gd01Dzu+D21NE=
github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6/go.mod h1:c5/tk6G0dSpXGzJN7Wk1OEie8grdSJAmeawId9Zvd34=
github.com/go-jose/go-jose/v4 v4.0.5 h1:M6T8+mKZl/+fNNuFHvGIzDz7BTLQPIounk/b9dw3AaE=
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/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
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/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/pires/go-proxyproto v0.15.0 h1:dTshmNbFm/D+0+sbrxUuddPOZ5Y0B7c5NhtsBkm6LqI=
github.com/pires/go-proxyproto v0.15.0/go.mod h1:OXsCrKwrK2tXS9YrI5tkHx5xaQlO8FH3lFW76orFh24=
github.com/pion/dtls/v3 v3.0.10 h1:k9ekkq1kaZoxnNEbyLKI8DI37j/Nbk1HWmMuywpQJgg=
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/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
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/procfs v0.12.0 h1:jluTpSng7V9hY0O2R9DzzJHYb2xULk9VTR1V1R/k6Bo=
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/go-ossfuzz-seeds v0.1.0/go.mod h1:3IOHRbJIc+L6YKMwfDtJAM9Vj9k0YY4muhuyUYk5tbk=
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/quic-go/quic-go v0.55.0 h1:zccPQIqYCXDt5NmcEabyYvOnomjs8Tlwl7tISjJh9Mk=
github.com/quic-go/quic-go v0.55.0/go.mod h1:DR51ilwU1uE164KuWXhinFcKWGlEjzys2l8zUl5Ss1U=
github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY=
github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc=
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/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8=
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/go.mod h1:75S1AKYYzNUSXIv30h+jPKJYZUwqpfvLshu63nCNSOM=
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-20200622213623-75b288015ac9/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.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4=
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/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-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-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-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-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0=
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.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q=
golang.org/x/oauth2 v0.28.0 h1:CrgCKl8PPAVtLnU3c+EDw6x11699EWlsDeWNWKdIOkc=
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-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.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
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-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
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.5.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.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo=
golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.41.0 h1:QCgPso/Q3RTJx2Th4bDLqML4W6iJiaXFq2/ftQF13YU=
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.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs=
golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY=
golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8=
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/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
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-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
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.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA=
golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k=
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.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=
+1 -1
View File
@@ -167,7 +167,7 @@ type ServerTransportConfig struct {
// If negative, keep-alive probes are disabled.
TCPKeepAlive int64 `json:"tcpKeepalive,omitempty"`
// 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"`
// HeartBeatTimeout specifies the maximum time to wait for a heartbeat
// before terminating the connection. It is not recommended to change this
+1 -38
View File
@@ -51,51 +51,14 @@ func (v *ConfigValidator) ValidateClientCommonConfig(c *v1.ClientCommonConfig) (
}
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 !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, 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) {
var errs error
if !slices.Contains(SupportedAuthMethods, c.Method) {
-140
View File
@@ -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())
}
-48
View File
@@ -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)
}
-48
View File
@@ -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)
})
}
}
+2 -4
View File
@@ -79,11 +79,9 @@ func validateDomainConfigForClient(c *v1.DomainConfig) error {
}
func validateDomainConfigForServer(c *v1.DomainConfig, s *v1.ServerConfig) error {
subDomainHost := strings.ToLower(s.SubDomainHost)
for _, domain := range c.CustomDomains {
canonicalDomain := strings.ToLower(domain)
if subDomainHost != "" && len(strings.Split(subDomainHost, ".")) < len(strings.Split(canonicalDomain, ".")) {
if strings.HasSuffix(canonicalDomain, "."+subDomainHost) {
if s.SubDomainHost != "" && len(strings.Split(s.SubDomainHost, ".")) < len(strings.Split(domain, ".")) {
if strings.HasSuffix(domain, "."+s.SubDomainHost) {
return fmt.Errorf("custom domain [%s] should not belong to subdomain host [%s]", domain, s.SubDomainHost)
}
}
-76
View File
@@ -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)
})
}
}
-3
View File
@@ -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.VhostHTTPSPort, "vhostHTTPSPort"))
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 {
if !lo.Every(SupportedHTTPPluginOps, p.Ops) {
-51
View File
@@ -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)
})
}
}
-199
View File
@@ -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
})
}
}
}
}
-338
View File
@@ -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)
}
}
-248
View File
@@ -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)
})
}
-1
View File
@@ -43,7 +43,6 @@ const (
V2TypeNatHoleResp uint16 = 16
V2TypeNatHoleSid uint16 = 17
V2TypeNatHoleReport uint16 = 18
V2TypeUDPPacketBinary uint16 = 19
)
var v2MsgTypeMap = map[uint16]any{
-3
View File
@@ -84,9 +84,6 @@ func TestV2MessageTypeIDsAreStable(t *testing.T) {
require.Equal(t, uint16(16), V2TypeNatHoleResp)
require.Equal(t, uint16(17), V2TypeNatHoleSid)
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) {
+70 -30
View File
@@ -15,24 +15,31 @@
package nathole
import (
"errors"
"fmt"
"net"
"time"
"github.com/fatedier/golib/net/stun"
"github.com/pion/stun/v3"
)
var responseTimeout = 3 * time.Second
type Message struct {
Body []byte
Addr string
}
// If the localAddr is empty, it will listen on a random port.
func Discover(stunServers []string, localAddr string) ([]string, net.Addr, error) {
// create a discoverConn and get response from messageChan
discoverConn, err := listen(localAddr)
if err != nil {
return nil, nil, err
}
defer discoverConn.Close()
go discoverConn.readLoop()
addresses := make([]string, 0, len(stunServers))
for _, addr := range stunServers {
// get external address from stun server
@@ -51,9 +58,10 @@ type stunResponse struct {
}
type discoverConn struct {
conn *net.UDPConn
client *stun.Client
localAddr net.Addr
conn *net.UDPConn
localAddr net.Addr
messageChan chan *Message
}
func listen(localAddr string) (*discoverConn, error) {
@@ -69,50 +77,82 @@ func listen(localAddr string) (*discoverConn, error) {
if err != nil {
return nil, err
}
client, err := stun.NewClient(conn)
if err != nil {
_ = conn.Close()
return nil, err
}
return &discoverConn{
conn: conn,
client: client,
localAddr: conn.LocalAddr(),
conn: conn,
localAddr: conn.LocalAddr(),
messageChan: make(chan *Message, 10),
}, nil
}
func (c *discoverConn) Close() error {
if c.messageChan != nil {
close(c.messageChan)
c.messageChan = nil
}
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) {
serverAddr, err := net.ResolveUDPAddr("udp4", addr)
if err != nil {
return nil, err
}
transaction, err := stun.NewBindingTransaction(serverAddr)
request, err := stun.Build(stun.TransactionID, stun.BindingRequest)
if err != nil {
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 response.MappedAddr != nil {
resp.externalAddr = response.MappedAddr.String()
if err = request.NewTransactionID(); err != nil {
return nil, err
}
if response.OtherAddr != nil {
resp.otherAddr = response.OtherAddr.String()
if _, err := c.conn.WriteTo(request.Raw, serverAddr); err != nil {
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
}
-382
View File
@@ -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)
}
+16
View File
@@ -18,8 +18,10 @@ import (
"bytes"
"fmt"
"net"
"strconv"
"github.com/fatedier/golib/crypto"
"github.com/pion/stun/v3"
"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)
}
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) {
addrs, err := net.InterfaceAddrs()
if err != nil {
+28 -85
View File
@@ -20,7 +20,6 @@ import (
"context"
"errors"
"fmt"
"io"
"net"
"sync"
"time"
@@ -34,16 +33,10 @@ func init() {
Register(v1.VisitorPluginVirtualNet, NewVirtualNetPlugin)
}
type clientRouteController interface {
RegisterClientRoute(context.Context, string, []net.IPNet, io.ReadWriteCloser)
UnregisterClientRoute(string, io.Writer) bool
}
type VirtualNetPlugin struct {
pluginCtx PluginContext
routeController clientRouteController
routes []net.IPNet
routes []net.IPNet
mu sync.Mutex
controllerConn net.Conn
@@ -55,11 +48,6 @@ type VirtualNetPlugin struct {
cancel context.CancelFunc
}
const (
virtualNetReconnectBaseDelay = 60 * time.Second
virtualNetReconnectMaxDelay = 300 * time.Second
)
func NewVirtualNetPlugin(pluginCtx PluginContext, options v1.VisitorPluginOptions) (Plugin, error) {
opts := options.(*v1.VirtualNetVisitorPluginOptions)
@@ -67,9 +55,6 @@ func NewVirtualNetPlugin(pluginCtx PluginContext, options v1.VisitorPluginOption
pluginCtx: pluginCtx,
routes: make([]net.IPNet, 0),
}
if pluginCtx.VnetController != nil {
p.routeController = pluginCtx.VnetController
}
p.ctx, p.cancel = context.WithCancel(pluginCtx.Ctx)
@@ -100,7 +85,7 @@ func (p *VirtualNetPlugin) Name() string {
func (p *VirtualNetPlugin) Start() {
xl := xlog.FromContextSafe(p.pluginCtx.Ctx)
if p.routeController == nil {
if p.pluginCtx.VnetController == nil {
return
}
@@ -126,17 +111,16 @@ func (p *VirtualNetPlugin) run() {
select {
case <-p.ctx.Done():
xl.Infof("VirtualNetPlugin run loop for visitor [%s] stopping (context cancelled before pipe creation).", p.pluginCtx.Name)
p.cleanupCurrentControllerConn(xl)
p.cleanupControllerConn(xl)
return
default:
}
controllerConn, pluginConn := net.Pipe()
xl.Infof("attempting to register client route for visitor [%s]", p.pluginCtx.Name)
if !p.registerControllerConn(controllerConn, pluginConn) {
xl.Infof("VirtualNetPlugin run loop for visitor [%s] stopping (context cancelled before route registration).", p.pluginCtx.Name)
return
}
p.mu.Lock()
p.controllerConn = controllerConn
p.mu.Unlock()
// Wrap with CloseNotifyConn which supports both close notification and error recording
var closeErr error
@@ -145,6 +129,8 @@ func (p *VirtualNetPlugin) run() {
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)
// Pass the CloseNotifyConn to the visitor for handling.
@@ -155,7 +141,7 @@ func (p *VirtualNetPlugin) run() {
select {
case <-p.ctx.Done():
xl.Infof("VirtualNetPlugin run loop stopping for visitor [%s] (context cancelled while waiting).", p.pluginCtx.Name)
p.cleanupControllerConn(xl, controllerConn)
p.cleanupControllerConn(xl)
return
case <-currentCloseSignal:
// Determine reconnect delay based on error with exponential backoff
@@ -166,7 +152,8 @@ func (p *VirtualNetPlugin) run() {
p.pluginCtx.Name, p.consecutiveErrors, closeErr)
// 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 {
// Reset consecutive errors on successful connection
if p.consecutiveErrors > 0 {
@@ -180,7 +167,7 @@ func (p *VirtualNetPlugin) run() {
}
// 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)
select {
@@ -195,66 +182,16 @@ func (p *VirtualNetPlugin) run() {
}
}
// registerControllerConn publishes and registers controllerConn atomically with
// respect to Close. A canceled plugin cannot register a new route.
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) {
// cleanupControllerConn closes the current controllerConn (if it exists) under lock.
func (p *VirtualNetPlugin) cleanupControllerConn(xl *xlog.Logger) {
p.mu.Lock()
defer p.mu.Unlock()
p.cleanupControllerConnLocked(xl, controllerConn)
}
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 {
if p.controllerConn != nil {
xl.Debugf("cleaning up controllerConn for visitor [%s]", p.pluginCtx.Name)
p.controllerConn.Close()
p.controllerConn = nil
p.closeSignal = nil
}
p.closeSignal = nil
}
// Close initiates the plugin shutdown.
@@ -265,9 +202,15 @@ func (p *VirtualNetPlugin) Close() error {
// Signal the run loop goroutine to stop.
p.cancel()
// Unregister and close the current connection while holding the same lock
// used to check cancellation and register a route in run.
p.cleanupCurrentControllerConn(xl)
// Unregister the route from the controller.
if p.pluginCtx.VnetController != nil {
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)
return nil
-245
View File
@@ -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())
}
+3 -7
View File
@@ -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
udpMsg := NewUDPPacket(buf[:n], nil, remoteAddr)
if err = errors.PanicToError(func() {
select {
case sendCh <- udpMsg:
default:
}
}); err != nil {
return
select {
case sendCh <- udpMsg:
default:
}
}
}
-34
View File
@@ -1,13 +1,9 @@
package udp
import (
"net"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/fatedier/frp/pkg/msg"
)
func TestUdpPacket(t *testing.T) {
@@ -20,33 +16,3 @@ func TestUdpPacket(t *testing.T) {
require.NoError(err)
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")
}
}
+1 -18
View File
@@ -68,8 +68,7 @@ func NewServerHello(clientHello ClientHello) (ServerHello, error) {
return ServerHello{
Selected: ServerSelection{
Message: MessageSelection{
Codec: MessageCodecJSON,
UDPPacketCodec: selectUDPPacketCodec(clientHello.Capabilities.Message.UDPPacketCodecs),
Codec: MessageCodecJSON,
},
Crypto: CryptoSelection{
Algorithm: algorithm,
@@ -93,15 +92,6 @@ func ValidateServerHelloForClient(clientHello ClientHello, serverHello ServerHel
if serverHello.Selected.Message.Codec != MessageCodecJSON {
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
if !IsSupportedAEADAlgorithm(cryptoSelection.Algorithm) {
return fmt.Errorf("unknown selected crypto algorithm: %s", cryptoSelection.Algorithm)
@@ -115,13 +105,6 @@ func ValidateServerHelloForClient(clientHello ClientHello, serverHello ServerHel
return nil
}
func selectUDPPacketCodec(codecs []string) string {
if Supports(codecs, UDPPacketCodecBinary) {
return UDPPacketCodecBinary
}
return ""
}
func NewCryptoContext(algorithm string, clientHelloPayload, serverHelloPayload []byte) *CryptoContext {
return &CryptoContext{
Algorithm: algorithm,
+3 -7
View File
@@ -36,7 +36,6 @@ const (
FrameTypeMessage uint16 = 16
MessageCodecJSON = "json"
UDPPacketCodecBinary = "binary-v1"
DefaultMaxFramePayloadSize = 64 * 1024
MagicV2 = "FRP\x00\x02\r\n"
@@ -183,8 +182,7 @@ type ClientCapabilities struct {
}
type MessageCapabilities struct {
Codecs []string `json:"codecs,omitempty"`
UDPPacketCodecs []string `json:"udpPacketCodecs,omitempty"`
Codecs []string `json:"codecs,omitempty"`
}
type CryptoCapabilities struct {
@@ -203,8 +201,7 @@ type ServerSelection struct {
}
type MessageSelection struct {
Codec string `json:"codec,omitempty"`
UDPPacketCodec string `json:"udpPacketCodec,omitempty"`
Codec string `json:"codec,omitempty"`
}
type CryptoSelection struct {
@@ -217,8 +214,7 @@ func clientHelloWithCryptoRandom(bootstrap BootstrapInfo, clientRandom []byte) C
Bootstrap: bootstrap,
Capabilities: ClientCapabilities{
Message: MessageCapabilities{
Codecs: []string{MessageCodecJSON},
UDPPacketCodecs: []string{UDPPacketCodecBinary},
Codecs: []string{MessageCodecJSON},
},
Crypto: CryptoCapabilities{
Algorithms: PreferredAEADAlgorithms(),
-30
View File
@@ -148,40 +148,10 @@ func TestNewServerHelloSelectsFirstSupportedAEADAlgorithm(t *testing.T) {
serverHello, err := NewServerHello(hello)
require.NoError(t, err)
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.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) {
hello := mustClientHello(t, BootstrapInfo{})
serverHello, err := NewServerHello(hello)
+5 -21
View File
@@ -16,6 +16,7 @@ package ssh
import (
"context"
"encoding/binary"
"errors"
"fmt"
"net"
@@ -51,11 +52,6 @@ type tcpipForward struct {
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
type forwardedTCPPayload struct {
Addr string
@@ -70,7 +66,6 @@ type TunnelServer struct {
sshConn *ssh.ServerConn
sc *ssh.ServerConfig
firstChannel ssh.Channel
firstChannelMu sync.Mutex
vc *virtual.Client
peerServerListener *netpkg.InternalListener
@@ -192,8 +187,6 @@ func (s *TunnelServer) Run() error {
}
func (s *TunnelServer) writeToClient(data string) {
s.firstChannelMu.Lock()
defer s.firstChannelMu.Unlock()
if s.firstChannel == nil {
return
}
@@ -307,24 +300,23 @@ func (s *TunnelServer) handleNewChannel(channel ssh.NewChannel, extraPayloadCh c
if err != nil {
return
}
s.firstChannelMu.Lock()
if s.firstChannel == nil {
s.firstChannel = ch
}
s.firstChannelMu.Unlock()
go s.keepAlive(ch)
for req := range reqs {
if req.WantReply {
_ = req.Reply(true, nil)
}
if req.Type != "exec" {
if req.Type != "exec" || len(req.Payload) <= 4 {
continue
}
extraPayload, ok := parseExecPayload(req.Payload)
if !ok {
end := 4 + binary.BigEndian.Uint32(req.Payload[:4])
if len(req.Payload) < int(end) {
continue
}
extraPayload := string(req.Payload[4:end])
select {
case extraPayloadCh <- extraPayload:
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) {
tk := time.NewTicker(time.Second * 30)
defer tk.Stop()
-115
View File
@@ -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")
}
}
-37
View File
@@ -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)
}
-65
View File
@@ -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())
})
}
}
-6
View File
@@ -35,12 +35,6 @@ func NewReader(r io.Reader, limiter *rate.Limiter) *Reader {
func (r *Reader) Read(p []byte) (n int, err error) {
b := r.limiter.Burst()
if b <= 0 {
if len(p) == 0 {
return 0, nil
}
return 0, invalidBurstError(b)
}
if b < len(p) {
p = p[:b]
}
-7
View File
@@ -34,15 +34,8 @@ func NewWriter(w io.Writer, limiter *rate.Limiter) *Writer {
}
func (w *Writer) Write(p []byte) (n int, err error) {
if len(p) == 0 {
return 0, nil
}
var nn int
b := w.limiter.Burst()
if b <= 0 {
return 0, invalidBurstError(b)
}
for {
end := len(p)
if end == 0 {
+8 -3
View File
@@ -16,7 +16,6 @@ package net
import (
"context"
"crypto/hkdf"
"crypto/sha256"
"errors"
"io"
@@ -26,6 +25,7 @@ import (
libcrypto "github.com/fatedier/golib/crypto"
quic "github.com/quic-go/quic-go"
"golang.org/x/crypto/hkdf"
"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) {
info := aeadControlHKDFInfoPrefix + " " + algorithm + " " + direction
return hkdf.Key(sha256.New, key, transcriptHash, info, libcrypto.AEADKeySize)
info := []byte(aeadControlHKDFInfoPrefix + " " + algorithm + " " + direction)
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
}
-6
View File
@@ -114,11 +114,5 @@ func TestDeriveAEADControlKeysUsesDistinctDirections(t *testing.T) {
bytes.Repeat([]byte{0x44}, 32),
)
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)
}
+1 -1
View File
@@ -14,7 +14,7 @@
package version
var version = "0.71.0"
var version = "0.70.0"
func Full() string {
return version
+3 -1
View File
@@ -28,6 +28,8 @@ import (
libio "github.com/fatedier/golib/io"
"github.com/fatedier/golib/pool"
"golang.org/x/net/http2"
"golang.org/x/net/http2/h2c"
httppkg "github.com/fatedier/frp/pkg/util/http"
"github.com/fatedier/frp/pkg/util/log"
@@ -137,7 +139,7 @@ func NewHTTPReverseProxy(option HTTPReverseProxyOptions, vhostRouter *Routers) *
_, _ = rw.Write(getNotFoundPageContent())
},
}
rp.proxy = proxy
rp.proxy = h2c.NewHandler(proxy, &http2.Server{})
return rp
}
-80
View File
@@ -1,94 +1,14 @@
package vhost
import (
"bufio"
"fmt"
"net"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/stretchr/testify/require"
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) {
rc := &RouteConfig{
Username: "alice",
+5 -5
View File
@@ -96,21 +96,21 @@ func (l *Logger) Spawn() *Logger {
}
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) {
log.Logger.WithPrefix(l.prefixString).Warnf(format, v...)
log.Logger.Warnf(l.prefixString+format, v...)
}
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) {
log.Logger.WithPrefix(l.prefixString).Debugf(format, v...)
log.Logger.Debugf(l.prefixString+format, v...)
}
func (l *Logger) Tracef(format string, v ...any) {
log.Logger.WithPrefix(l.prefixString).Tracef(format, v...)
log.Logger.Tracef(l.prefixString+format, v...)
}
-76
View File
@@ -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
}
+4 -9
View File
@@ -246,9 +246,9 @@ func (c *Controller) RegisterClientRoute(ctx context.Context, name string, route
go c.readLoopClient(ctx, conn)
}
// UnregisterClientRoute removes a client route only when it is still owned by conn.
func (c *Controller) UnregisterClientRoute(name string, conn io.Writer) bool {
return c.clientRouter.delRoute(name, conn)
// UnregisterClientRoute Remove client route from routing table
func (c *Controller) UnregisterClientRoute(name string) {
c.clientRouter.delRoute(name)
}
// 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)
}
func (r *clientRouter) delRoute(name string, conn io.Writer) bool {
func (r *clientRouter) delRoute(name string) {
r.mu.Lock()
defer r.mu.Unlock()
re, ok := r.routes[name]
if !ok || re.conn != conn {
return false
}
delete(r.routes, name)
return true
}
func (r *clientRouter) removeConnRoute(conn io.Writer) {
-64
View File
@@ -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)
}
+78 -466
View File
@@ -17,8 +17,6 @@ package server
import (
"context"
"fmt"
"math"
"net"
"runtime/debug"
"sync"
"sync/atomic"
@@ -42,313 +40,55 @@ import (
"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 {
// controls indexed by run id
ctlsByRunID map[string]*controlEntry
registry *registry.ClientRegistry
closed bool
ctlsByRunID map[string]*Control
mu sync.RWMutex
}
func NewControlManager(clientRegistry *registry.ClientRegistry) *ControlManager {
func NewControlManager() *ControlManager {
return &ControlManager{
ctlsByRunID: make(map[string]*controlEntry),
registry: clientRegistry,
ctlsByRunID: make(map[string]*Control),
}
}
// lockCurrentRun returns the current entry with its run gate held. It never
// 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()
func (cm *ControlManager) Add(runID string, ctl *Control) (old *Control) {
cm.mu.Lock()
defer cm.mu.Unlock()
if cm.closed || cm.ctlsByRunID[ctl.runID] != entry || entry.ctl != ctl || entry.id != ctl.controlID {
return false, nil
var ok bool
old, ok = cm.ctlsByRunID[runID]
if ok {
old.Replaced(ctl)
}
ctl.lifecycleMu.Lock()
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
cm.ctlsByRunID[runID] = ctl
return
}
// completeLogin reserves ctl's current ownership with its run gate while the
// bounded successful LoginResp write runs, then transitions it to running.
// 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()
// we should make sure if it's the same control to prevent delete a new one
func (cm *ControlManager) Del(runID string, ctl *Control) {
cm.mu.Lock()
defer cm.mu.Unlock()
if cm.ctlsByRunID[ctl.runID] != entry || entry.ctl != ctl || entry.id != ctl.controlID {
return false
if c, ok := cm.ctlsByRunID[runID]; ok && c == ctl {
delete(cm.ctlsByRunID, runID)
}
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) {
entry, ok := cm.lockCurrentRun(runID, false)
if !ok {
return nil, false
}
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")
}
cm.mu.RLock()
defer cm.mu.RUnlock()
ctl, ok = cm.ctlsByRunID[runID]
return
}
func (cm *ControlManager) Close() error {
cm.mu.Lock()
cm.closed = true
ctls := make([]*Control, 0, len(cm.ctlsByRunID))
for _, entry := range cm.ctlsByRunID {
ctls = append(ctls, entry.ctl)
}
cm.mu.Unlock()
for _, ctl := range ctls {
cm.Remove(ctl)
_ = ctl.Close()
defer cm.mu.Unlock()
for _, ctl := range cm.ctlsByRunID {
ctl.Close()
}
cm.ctlsByRunID = make(map[string]*Control)
return nil
}
@@ -370,21 +110,12 @@ type SessionContext struct {
LoginMsg *msg.Login
// server configuration
ServerCfg *v1.ServerConfig
// client registry
ClientRegistry *registry.ClientRegistry
// negotiated wire protocol for this client session
WireProtocol string
UDPPacketCodec string
WireProtocol string
}
type controlState uint8
const (
controlStateCreated controlState = iota
controlStatePending
controlStateRunning
controlStateClosing
controlStateClosed
)
type Control struct {
// session context
sessionCtx *SessionContext
@@ -411,59 +142,30 @@ type Control struct {
// last time got the Ping message
lastPing atomic.Value
// runID never changes during the lifetime of a control. controlID is assigned
// once by ControlManager and distinguishes same-runID generations.
runID string
controlID ControlID
manager *ControlManager
lifecycleMu sync.Mutex
state controlState
activated bool
handoffBarrier <-chan struct{}
interruptOnce sync.Once
interruptErr error
// A new run id will be generated when a new client login.
// If run id got from login message has same run id, it means it's the same client, so we can
// replace old controller instantly.
runID string
mu sync.RWMutex
xl *xlog.Logger
ctx context.Context
doneCh chan struct{}
serverMetrics metrics.ServerMetrics
xl *xlog.Logger
ctx context.Context
doneCh chan struct{}
}
func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, error) {
if sessionCtx.LoginMsg.PoolCount < 0 {
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)
poolCount := min(sessionCtx.LoginMsg.PoolCount, int(sessionCtx.ServerCfg.Transport.MaxPoolCount))
ctl := &Control{
sessionCtx: sessionCtx,
workConnCh: make(chan *proxy.WorkConn, poolCount+workConnPoolCapacityOffset),
proxies: make(map[string]proxy.Proxy),
poolCount: poolCount,
portsUsedNum: 0,
runID: sessionCtx.LoginMsg.RunID,
state: controlStateCreated,
xl: xlog.FromContextSafe(ctx),
ctx: ctx,
doneCh: make(chan struct{}),
serverMetrics: metrics.Server,
sessionCtx: sessionCtx,
workConnCh: make(chan *proxy.WorkConn, poolCount+10),
proxies: make(map[string]proxy.Proxy),
poolCount: poolCount,
portsUsedNum: 0,
runID: sessionCtx.LoginMsg.RunID,
xl: xlog.FromContextSafe(ctx),
ctx: ctx,
doneCh: make(chan struct{}),
}
ctl.lastPing.Store(time.Now())
@@ -473,121 +175,48 @@ func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, erro
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.
func (ctl *Control) Start() bool {
ctl.lifecycleMu.Lock()
defer ctl.lifecycleMu.Unlock()
return ctl.startLocked()
}
func (ctl *Control) startLocked() bool {
if ctl.state != controlStatePending || !ctl.activated {
return false
}
ctl.state = controlStateRunning
func (ctl *Control) Start() {
go func() {
for i := 0; i < ctl.poolCount; i++ {
// ignore error here, that means that this control is closed
_ = ctl.msgDispatcher.Send(&msg.ReqWorkConn{})
}
}()
go ctl.worker()
return true
}
func (ctl *Control) Close() error {
ctl.lifecycleMu.Lock()
switch ctl.state {
case controlStateCreated, controlStatePending:
ctl.state = controlStateClosing
ctl.finishLocked()
case controlStateRunning:
ctl.state = controlStateClosing
}
ctl.lifecycleMu.Unlock()
return ctl.interruptReadAndClose()
ctl.sessionCtx.Conn.Close()
return nil
}
func (ctl *Control) Replaced(newCtl *Control) {
ctl.markReplaced()
ctl.xl.Infof("replaced by client [%s] (control ID %d)", newCtl.runID, newCtl.ID())
_ = ctl.interruptReadAndClose()
xl := ctl.xl
xl.Infof("replaced by client [%s]", newCtl.runID)
ctl.runID = ""
ctl.sessionCtx.Conn.Close()
}
// markReplaced returns the transitive predecessor barrier. A pending control
// has no worker, so it finishes immediately and passes its inherited barrier
// to the replacement. A running control is finished only by its worker.
func (ctl *Control) markReplaced() <-chan struct{} {
ctl.lifecycleMu.Lock()
defer ctl.lifecycleMu.Unlock()
func (ctl *Control) RegisterWorkConn(conn *proxy.WorkConn) error {
xl := ctl.xl
defer func() {
if err := recover(); err != nil {
xl.Errorf("panic error: %v", err)
xl.Errorf(string(debug.Stack()))
}
}()
switch ctl.state {
case controlStateCreated:
ctl.state = controlStateClosing
ctl.finishLocked()
select {
case ctl.workConnCh <- conn:
xl.Debugf("new work connection registered")
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:
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.
// If no workConn available in the pool, send message to frpc to get one or more
// and wait until it is available.
@@ -642,10 +271,10 @@ func (ctl *Control) heartbeatWorker() {
}
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 {
xl.Warnf("heartbeat timeout")
_ = ctl.Close()
ctl.sessionCtx.Conn.Close()
return
}
}, time.Second, ctl.doneCh)
@@ -660,14 +289,14 @@ func (ctl *Control) loginUserInfo() plugin.UserInfo {
return plugin.UserInfo{
User: ctl.sessionCtx.LoginMsg.User,
Metas: ctl.sessionCtx.LoginMsg.Metas,
RunID: ctl.runID,
RunID: ctl.sessionCtx.LoginMsg.RunID,
}
}
func (ctl *Control) closeProxy(pxy proxy.Proxy) {
pxy.Close()
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{
User: ctl.loginUserInfo(),
@@ -682,24 +311,12 @@ func (ctl *Control) closeProxy(pxy proxy.Proxy) {
func (ctl *Control) worker() {
xl := ctl.xl
ctl.serverMetrics.NewClient()
go ctl.heartbeatWorker()
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.lifecycleMu.Lock()
if ctl.state == controlStateRunning {
ctl.state = controlStateClosing
}
ctl.lifecycleMu.Unlock()
_ = ctl.interruptReadAndClose()
ctl.sessionCtx.Conn.Close()
ctl.mu.Lock()
close(ctl.workConnCh)
@@ -714,14 +331,10 @@ func (ctl *Control) worker() {
ctl.closeProxy(pxy)
}
ctl.serverMetrics.CloseClient()
if ctl.manager != nil {
ctl.manager.Remove(ctl)
}
metrics.Server.CloseClient()
ctl.sessionCtx.ClientRegistry.MarkOfflineByRunID(ctl.runID)
xl.Infof("client exit success")
ctl.lifecycleMu.Lock()
ctl.finishLocked()
ctl.lifecycleMu.Unlock()
close(ctl.doneCh)
}
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)
clientID := ctl.sessionCtx.LoginMsg.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)
}
@@ -842,7 +455,6 @@ func (ctl *Control) RegisterProxy(pxyMsg *msg.NewProxy) (remoteAddr string, err
ServerCfg: ctl.sessionCtx.ServerCfg,
EncryptionKey: ctl.sessionCtx.EncryptionKey,
WireProtocol: ctl.sessionCtx.WireProtocol,
UDPPacketCodec: ctl.sessionCtx.UDPPacketCodec,
})
if err != nil {
return remoteAddr, err
-595
View File
@@ -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
}
+34 -93
View File
@@ -82,20 +82,19 @@ type Proxy interface {
}
type BaseProxy struct {
name string
rc *controller.ResourceController
listeners []net.Listener
usedPortsNum int
poolCount int
getWorkConnFn GetWorkConnFn
serverCfg *v1.ServerConfig
encryptionKey []byte
limiter *rate.Limiter
userInfo plugin.UserInfo
loginMsg *msg.Login
configurer v1.ProxyConfigurer
wireProtocol string
udpPacketCodec string
name string
rc *controller.ResourceController
listeners []net.Listener
usedPortsNum int
poolCount int
getWorkConnFn GetWorkConnFn
serverCfg *v1.ServerConfig
encryptionKey []byte
limiter *rate.Limiter
userInfo plugin.UserInfo
loginMsg *msg.Login
configurer v1.ProxyConfigurer
wireProtocol string
mu sync.RWMutex
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) {
visitorWireProtocol := wireProtocolFromConn(userConn)
visitorUDPPacketCodec := udpPacketCodecFromConn(userConn)
if proxyType == string(v1.ProxyTypeSUDP) {
mixed, err := isMixedSUDPPacketEncoding(pxy.wireProtocol, pxy.udpPacketCodec, visitorWireProtocol, visitorUDPPacketCodec)
if err != nil {
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)
}
if proxyType == string(v1.ProxyTypeSUDP) && isMixedWireProtocol(pxy.wireProtocol, visitorWireProtocol) {
xl.Infof("bridge mixed SUDP payload codecs, proxy wireProtocol [%s], visitor wireProtocol [%s]",
normalizeWireProtocol(pxy.wireProtocol), normalizeWireProtocol(visitorWireProtocol))
return joinSUDPMessageBridge(local, userConn, pxy.wireProtocol, visitorWireProtocol, xl)
}
return libio.Join(local, userConn)
}
@@ -348,10 +339,6 @@ type wireProtocolGetter interface {
WireProtocol() string
}
type udpPacketCodecGetter interface {
UDPPacketCodec() string
}
func wireProtocolFromConn(conn net.Conn) string {
if getter, ok := conn.(wireProtocolGetter); ok {
return getter.WireProtocol()
@@ -359,46 +346,10 @@ func wireProtocolFromConn(conn net.Conn) string {
return ""
}
func udpPacketCodecFromConn(conn net.Conn) string {
if getter, ok := conn.(udpPacketCodecGetter); ok {
return getter.UDPPacketCodec()
}
return ""
}
func isMixedWireProtocol(left, right string) bool {
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 {
if wireProtocol == wire.ProtocolV2 {
return wire.ProtocolV2
@@ -410,21 +361,13 @@ func joinSUDPMessageBridge(
proxyConn io.ReadWriteCloser,
visitorConn io.ReadWriteCloser,
proxyWireProtocol string,
proxyUDPPacketCodec string,
visitorWireProtocol string,
visitorUDPPacketCodec string,
xl *xlog.Logger,
) (inCount int64, outCount int64, errs []error) {
// The mixed bridge decodes and re-encodes messages, so raw framed byte counts
// are not available. Count UDP payload bytes and ignore heartbeat traffic.
proxyRW, err := msg.NewUDPPacketReadWriter(proxyConn, proxyWireProtocol, proxyUDPPacketCodec)
if err != nil {
return 0, 0, []error{err}
}
visitorRW, err := msg.NewUDPPacketReadWriter(visitorConn, visitorWireProtocol, visitorUDPPacketCodec)
if err != nil {
return 0, 0, []error{err}
}
proxyRW := msg.NewReadWriter(proxyConn, proxyWireProtocol)
visitorRW := msg.NewReadWriter(visitorConn, visitorWireProtocol)
var (
once sync.Once
@@ -526,7 +469,6 @@ type Options struct {
ServerCfg *v1.ServerConfig
EncryptionKey []byte
WireProtocol string
UDPPacketCodec string
}
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
limitBytes := configurer.GetBaseConfig().Transport.BandwidthLimit.Bytes()
if limitBytes > 0 && configurer.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeServer {
limiter = limit.NewBandwidthLimiter(limitBytes)
limiter = rate.NewLimiter(rate.Limit(float64(limitBytes)), int(limitBytes))
}
basePxy := BaseProxy{
name: configurer.GetBaseConfig().Name,
rc: options.ResourceController,
listeners: make([]net.Listener, 0),
poolCount: options.PoolCount,
getWorkConnFn: options.GetWorkConnFn,
serverCfg: options.ServerCfg,
encryptionKey: options.EncryptionKey,
limiter: limiter,
xl: xl,
ctx: xlog.NewContext(ctx, xl),
userInfo: options.UserInfo,
loginMsg: options.LoginMsg,
configurer: configurer,
wireProtocol: options.WireProtocol,
udpPacketCodec: options.UDPPacketCodec,
name: configurer.GetBaseConfig().Name,
rc: options.ResourceController,
listeners: make([]net.Listener, 0),
poolCount: options.PoolCount,
getWorkConnFn: options.GetWorkConnFn,
serverCfg: options.ServerCfg,
encryptionKey: options.EncryptionKey,
limiter: limiter,
xl: xl,
ctx: xlog.NewContext(ctx, xl),
userInfo: options.UserInfo,
loginMsg: options.LoginMsg,
configurer: configurer,
wireProtocol: options.WireProtocol,
}
factory := proxyFactoryRegistry[reflect.TypeOf(configurer)]
-232
View File
@@ -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
View File
@@ -18,27 +18,22 @@ import (
"bufio"
"bytes"
"encoding/binary"
"io"
"net"
"testing"
"time"
"github.com/stretchr/testify/require"
v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/msg"
"github.com/fatedier/frp/pkg/proto/wire"
"github.com/fatedier/frp/pkg/util/xlog"
)
func TestSUDPBridgeTranscodesProxyV1ToVisitorV2(t *testing.T) {
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
err := bridgeSUDPProxyToVisitor(
newSUDPBridgeRW(t, &in, wire.ProtocolV1, ""),
newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""),
msg.NewReadWriter(&in, wire.ProtocolV1),
msg.NewReadWriter(&out, wire.ProtocolV2),
&count,
nil,
)
@@ -58,12 +53,12 @@ func TestSUDPBridgeTranscodesProxyV1ToVisitorV2(t *testing.T) {
func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) {
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
err := bridgeSUDPVisitorToProxy(
newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
msg.NewReadWriter(&in, wire.ProtocolV2),
msg.NewReadWriter(&out, wire.ProtocolV1),
&count,
nil,
)
@@ -81,67 +76,33 @@ func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) {
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) {
var in, out bytes.Buffer
writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, "", &msg.Ping{})
writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, &msg.Ping{})
var count int64
err := bridgeSUDPProxyToVisitor(
newSUDPBridgeRW(t, &in, wire.ProtocolV1, ""),
newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""),
msg.NewReadWriter(&in, wire.ProtocolV1),
msg.NewReadWriter(&out, wire.ProtocolV2),
&count,
nil,
)
require.NoError(t, err)
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.IsType(t, &msg.Ping{}, rawMsg)
}
func TestSUDPBridgeDropsVisitorPing(t *testing.T) {
var in, out bytes.Buffer
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", &msg.Ping{})
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, &msg.Ping{})
var count int64
err := bridgeSUDPVisitorToProxy(
newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
msg.NewReadWriter(&in, wire.ProtocolV2),
msg.NewReadWriter(&out, wire.ProtocolV1),
&count,
nil,
)
@@ -152,12 +113,12 @@ func TestSUDPBridgeDropsVisitorPing(t *testing.T) {
func TestSUDPBridgeRejectsUnknownVisitorMessage(t *testing.T) {
var in, out bytes.Buffer
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", &msg.Pong{})
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, &msg.Pong{})
var count int64
err := bridgeSUDPVisitorToProxy(
newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
msg.NewReadWriter(&in, wire.ProtocolV2),
msg.NewReadWriter(&out, wire.ProtocolV1),
&count,
nil,
)
@@ -166,22 +127,6 @@ func TestSUDPBridgeRejectsUnknownVisitorMessage(t *testing.T) {
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) {
require.False(t, isMixedWireProtocol("", wire.ProtocolV1))
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))
}
func TestSUDPBridgeDetectsMixedPacketEncoding(t *testing.T) {
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 {
func writeSUDPBridgeMsg(t *testing.T, buf *bytes.Buffer, wireProtocol string, m msg.Message) {
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) {
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
require.NoError(t, msg.NewReadWriter(buf, wireProtocol).WriteMsg(m))
}
+1 -7
View File
@@ -224,13 +224,7 @@ func (pxy *UDPProxy) Run() (remoteAddr string, err error) {
pxy.workConn = netpkg.WrapReadWriteCloserToConn(rwc, workConn)
// Plain UDP payload follows the negotiated wire protocol for message framing.
payloadRW, err := msg.NewUDPPacketReadWriter(pxy.workConn, pxy.wireProtocol, pxy.udpPacketCodec)
if err != nil {
xl.Errorf("create UDP packet read writer: %v", err)
pxy.workConn.Close()
continue
}
payloadConn := msg.NewConn(pxy.workConn, payloadRW)
payloadConn := msg.NewConn(pxy.workConn, msg.NewReadWriter(pxy.workConn, pxy.wireProtocol))
ctx, cancel := context.WithCancel(context.Background())
go workConnReaderFn(payloadConn)
go workConnSenderFn(payloadConn, ctx)
+6 -44
View File
@@ -28,7 +28,6 @@ type ClientInfo struct {
User string
RawClientID string
RunID string
ControlID uint64
Hostname string
IP 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.
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 == "" {
return "", false
}
@@ -94,16 +83,6 @@ func (cr *ClientRegistry) RegisterWithControlID(
if enforceUnique && exists && info.Online && info.RunID != "" && info.RunID != runID {
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 {
info = &ClientInfo{
@@ -118,7 +97,6 @@ func (cr *ClientRegistry) RegisterWithControlID(
info.RawClientID = rawClientID
info.RunID = runID
info.ControlID = controlID
info.Hostname = hostname
info.IP = remoteAddr
info.Version = version
@@ -136,16 +114,6 @@ func (cr *ClientRegistry) RegisterWithControlID(
// MarkOfflineByRunID marks the client as offline when the corresponding control disconnects.
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()
defer cr.mu.Unlock()
@@ -153,23 +121,17 @@ func (cr *ClientRegistry) markOfflineByRunID(runID string, controlID uint64, mat
if !ok {
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 == "" {
delete(cr.clients, key)
} 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)
}
}
func setClientOffline(info *ClientInfo, now time.Time) {
info.RunID = ""
info.ControlID = 0
info.Online = false
info.DisconnectedAt = now
delete(cr.runIndex, runID)
}
// List returns a snapshot of all known clients.
-86
View File
@@ -72,89 +72,3 @@ func TestClientRegistryUsesClockForTimestamps(t *testing.T) {
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
View File
@@ -18,7 +18,6 @@ import (
"bytes"
"context"
"crypto/tls"
"errors"
"fmt"
"io"
"net"
@@ -29,13 +28,12 @@ import (
"github.com/fatedier/golib/crypto"
"github.com/fatedier/golib/net/mux"
fmux "github.com/fatedier/yamux"
fmux "github.com/hashicorp/yamux"
quic "github.com/quic-go/quic-go"
"github.com/samber/lo"
"github.com/fatedier/frp/pkg/auth"
v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/config/v1/validation"
modelmetrics "github.com/fatedier/frp/pkg/metrics"
"github.com/fatedier/frp/pkg/msg"
"github.com/fatedier/frp/pkg/nathole"
@@ -53,6 +51,7 @@ import (
"github.com/fatedier/frp/pkg/util/xlog"
"github.com/fatedier/frp/server/controller"
"github.com/fatedier/frp/server/group"
"github.com/fatedier/frp/server/metrics"
"github.com/fatedier/frp/server/ports"
"github.com/fatedier/frp/server/proxy"
"github.com/fatedier/frp/server/registry"
@@ -65,8 +64,6 @@ const (
vhostReadWriteTimeout time.Duration = 30 * time.Second
)
var errControlReplaced = errors.New("control was replaced during login")
func init() {
crypto.DefaultSalt = "frp"
// Disable quic-go's receive buffer warning.
@@ -164,10 +161,9 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) {
return nil, err
}
clientRegistry := registry.NewClientRegistry()
svr := &Service{
ctlManager: NewControlManager(clientRegistry),
clientRegistry: clientRegistry,
ctlManager: NewControlManager(),
clientRegistry: registry.NewClientRegistry(),
pxyManager: proxy.NewManager(),
pluginManager: plugin.NewManager(),
rc: &controller.ResourceController{
@@ -301,14 +297,10 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) {
svr.rc.HTTPReverseProxy = rp
address := net.JoinHostPort(cfg.ProxyBindAddr, strconv.Itoa(cfg.VhostHTTPPort))
protocols := new(http.Protocols)
protocols.SetHTTP1(true)
protocols.SetUnencryptedHTTP2(true)
server := &http.Server{
Addr: address,
Handler: rp,
ReadHeaderTimeout: 60 * time.Second,
Protocols: protocols,
}
var l net.Listener
if httpMuxOn {
@@ -471,15 +463,12 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
}
}
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 {
xl.Warnf("register control error: %v", err)
if ctl != nil {
svr.ctlManager.Remove(ctl)
}
if writeErr := writeWithDeadline(conn, connWriteTimeout, func() error {
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
Version: version.Full(),
@@ -488,34 +477,31 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
}); writeErr != nil {
xl.Warnf("write login error response error: %v", writeErr)
}
if ctl != nil {
_ = ctl.Close()
} else {
conn.Close()
}
conn.Close()
return
}
if err = svr.completeControlLogin(ctl, func() error {
return writeWithDeadline(conn, connWriteTimeout, func() error {
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
Version: version.Full(),
RunID: ctl.runID,
Error: "",
})
if err = writeWithDeadline(conn, connWriteTimeout, func() error {
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
Version: version.Full(),
RunID: ctl.runID,
Error: "",
})
}); err != nil {
xl.Warnf("complete control login error: %v", err)
svr.ctlManager.Remove(ctl)
_ = ctl.Close()
xl.Warnf("write login response error: %v", err)
svr.ctlManager.Del(m.RunID, ctl)
svr.clientRegistry.MarkOfflineByRunID(m.RunID)
conn.Close()
return
}
ctl.Start()
metrics.Server.NewClient()
go func() {
// block until control closed
ctl.WaitClosed()
svr.ctlManager.Del(m.RunID, ctl)
}()
case *msg.NewWorkConn:
if err := svr.RegisterWorkConn(
acceptedConn.conn,
m,
acceptedConn.wireProtocol,
acceptedConn.clientHelloPresent,
); err != nil {
if err := svr.RegisterWorkConn(acceptedConn.conn, m); err != nil {
_ = acceptedConn.conn.WriteMsg(&msg.StartWorkConn{
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 {
conn *msg.Conn
wireProtocol string
clientHelloPresent bool
udpPacketCodec string
cryptoContext *wire.CryptoContext
firstMsg msg.Message
conn *msg.Conn
wireProtocol string
cryptoContext *wire.CryptoContext
firstMsg msg.Message
}
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)
}
if frame.Type == wire.FrameTypeClientHello {
ac.clientHelloPresent = true
if err := ac.handleClientHello(conn, wireConn, frame); err != nil {
return nil, err
}
@@ -675,7 +647,6 @@ func (ac *acceptedConnection) handleClientHello(conn net.Conn, wireConn *wire.Co
return fmt.Errorf("write ServerHello: %w", err)
}
ac.cryptoContext = cryptoContext
ac.udpPacketCodec = serverHello.Selected.Message.UDPPacketCodec
return nil
}
@@ -769,20 +740,7 @@ func (svr *Service) RegisterControl(
loginMsg *msg.Login,
internal bool,
wireProtocol string,
udpPacketCodec string,
) (*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.
// 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
@@ -792,9 +750,6 @@ func (svr *Service) RegisterControl(
return nil, err
}
}
if err := validation.ValidateRunID(loginMsg.RunID); err != nil {
return nil, fmt.Errorf("invalid run id: %w", err)
}
ctx := netpkg.NewContextFromConn(ctlConn)
xl := xlog.FromContextSafe(ctx)
@@ -821,8 +776,8 @@ func (svr *Service) RegisterControl(
Conn: ctlConn,
LoginMsg: loginMsg,
ServerCfg: svr.cfg,
ClientRegistry: svr.clientRegistry,
WireProtocol: wireProtocol,
UDPPacketCodec: udpPacketCodec,
})
if err != nil {
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")
}
if err := svr.ctlManager.Add(ctl); err != nil {
return ctl, err
if oldCtl := svr.ctlManager.Add(loginMsg.RunID, ctl); oldCtl != nil {
oldCtl.WaitClosed()
}
ctl.WaitForHandoff()
active, err := svr.ctlManager.Activate(ctl)
if err != nil {
return ctl, err
remoteAddr := ctlConn.RemoteAddr().String()
if host, _, err := net.SplitHostPort(remoteAddr); err == nil {
remoteAddr = host
}
if !active {
return ctl, errControlReplaced
_, conflict := svr.clientRegistry.Register(loginMsg.User, loginMsg.ClientID, loginMsg.RunID, loginMsg.Hostname, loginMsg.Version, remoteAddr, wireProtocol)
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
}
// RegisterWorkConn register a new work connection to control and proxies need it.
func (svr *Service) RegisterWorkConn(
workConn *msg.Conn,
newMsg *msg.NewWorkConn,
workWireProtocol string,
workClientHelloPresent bool,
) error {
if workClientHelloPresent {
return fmt.Errorf("ClientHello is not allowed on work connections")
}
func (svr *Service) RegisterWorkConn(workConn *msg.Conn, newMsg *msg.NewWorkConn) error {
xl := netpkg.NewLogFromConn(workConn)
ctl, exist := svr.ctlManager.GetByID(newMsg.RunID)
if !exist {
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)
}
if workWireProtocol != ctl.sessionCtx.WireProtocol {
return fmt.Errorf("work connection wire protocol mismatch: got %s want %s", workWireProtocol, ctl.sessionCtx.WireProtocol)
}
// server plugin hook
content := &plugin.NewWorkConnContent{
@@ -885,33 +830,20 @@ func (svr *Service) RegisterWorkConn(
xl.Warnf("invalid NewWorkConn with run id [%s]", newMsg.RunID)
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 {
admit := func(visitorUser, visitorWireProtocol, visitorUDPPacketCodec string) error {
if visitorWireProtocol == "" {
visitorWireProtocol = wireProtocol
}
return svr.rc.VisitorManager.NewConn(newMsg.ProxyName, visitorConn, newMsg.Timestamp, newMsg.SignKey,
newMsg.UseEncryption, newMsg.UseCompression, visitorUser, visitorWireProtocol, visitorUDPPacketCodec)
}
visitorUser := ""
// 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 newMsg.RunID != "" {
admitted, err := svr.ctlManager.admitVisitorByRunID(newMsg.RunID, func(visitorUser, controlWireProtocol, controlUDPPacketCodec string) error {
if wireProtocol != controlWireProtocol {
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 {
ctl, exist := svr.ctlManager.GetByID(newMsg.RunID)
if !exist {
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)
}
-913
View File
@@ -15,33 +15,12 @@
package server
import (
"context"
"errors"
"fmt"
"math"
"net"
"net/http"
"runtime"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/fatedier/golib/net/mux"
"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) {
@@ -82,895 +61,3 @@ func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) {
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 }
+4 -14
View File
@@ -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,
useEncryption bool, useCompression bool, visitorUser string,
wireProtocol string, udpPacketCodecs ...string,
wireProtocol string,
) (err error) {
udpPacketCodec := ""
if len(udpPacketCodecs) > 0 {
udpPacketCodec = udpPacketCodecs[0]
}
vm.mu.RLock()
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)
err = l.l.PutConn(&wireProtocolConn{
Conn: visitorConn,
wireProtocol: wireProtocol,
udpPacketCodec: udpPacketCodec,
Conn: visitorConn,
wireProtocol: wireProtocol,
})
} else {
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 {
net.Conn
wireProtocol string
udpPacketCodec string
wireProtocol string
}
func (c *wireProtocolConn) WireProtocol() string {
return c.wireProtocol
}
func (c *wireProtocolConn) UDPPacketCodec() string {
return c.udpPacketCodec
}
func (vm *Manager) CloseListener(name string) {
vm.mu.Lock()
defer vm.mu.Unlock()
+3 -8
View File
@@ -25,7 +25,7 @@ import (
"github.com/fatedier/frp/pkg/util/util"
)
func TestManagerNewConnCarriesWireProtocolAndUDPPacketCodec(t *testing.T) {
func TestManagerNewConnCarriesWireProtocol(t *testing.T) {
vm := NewManager()
listener, err := vm.Listen("sudp", "secret", []string{"*"})
require.NoError(t, err)
@@ -47,7 +47,6 @@ func TestManagerNewConnCarriesWireProtocolAndUDPPacketCodec(t *testing.T) {
false,
"user",
wire.ProtocolV2,
wire.UDPPacketCodecBinary,
)
}()
@@ -55,12 +54,8 @@ func TestManagerNewConnCarriesWireProtocolAndUDPPacketCodec(t *testing.T) {
require.NoError(t, err)
defer acceptedConn.Close()
metadata, ok := acceptedConn.(interface {
WireProtocol() string
UDPPacketCodec() string
})
getter, ok := acceptedConn.(interface{ WireProtocol() string })
require.True(t, ok)
require.Equal(t, wire.ProtocolV2, metadata.WireProtocol())
require.Equal(t, wire.UDPPacketCodecBinary, metadata.UDPPacketCodec())
require.Equal(t, wire.ProtocolV2, getter.WireProtocol())
require.NoError(t, <-errCh)
}
@@ -192,41 +192,6 @@ transport.wireProtocol = "v2"
})
})
var _ = ginkgo.Describe("[Compatibility: BinaryUDPPacket]", func() {
f := framework.NewDefaultFramework()
ginkgo.BeforeEach(func() {
supportsV2, knownVersion := baselineSupportsControlWireProtocolV2(compatCtx.BaselineVersion)
if !knownVersion || !supportsV2 {
ginkgo.Skip(fmt.Sprintf("baseline version %q does not have known wire protocol v2 support", compatCtx.BaselineVersion))
}
})
ginkgo.It("current frps falls back to JSON for baseline frpc", func() {
portName := port.GenName("CompatBinaryUDPBaselineFRPC")
clientConf := udpClientConfig("udp", portName, `transport.wireProtocol = "v2"`)
f.RunProcessesWithBinaries(
compatCtx.CurrentFRPSPath,
compatCtx.BaselineFRPCPath,
consts.DefaultServerConfig,
[]string{clientConf},
)
framework.NewRequestExpect(f).Protocol("udp").PortName(portName).Ensure()
})
ginkgo.It("current frpc falls back to JSON for baseline frps", func() {
portName := port.GenName("CompatBinaryUDPBaselineFRPS")
clientConf := udpClientConfig("udp", portName, `transport.wireProtocol = "v2"`)
f.RunProcessesWithBinaries(
compatCtx.BaselineFRPSPath,
compatCtx.CurrentFRPCPath,
consts.DefaultServerConfig,
[]string{clientConf},
)
framework.NewRequestExpect(f).Protocol("udp").PortName(portName).Ensure()
})
})
func tcpClientConfig(proxyName string, remotePortName string, extra string) string {
return fmt.Sprintf(`
serverAddr = "127.0.0.1"
@@ -243,22 +208,6 @@ remotePort = {{ .%s }}
`, consts.PortServerName, extra, proxyName, framework.TCPEchoServerPort, remotePortName)
}
func udpClientConfig(proxyName string, remotePortName string, extra string) string {
return fmt.Sprintf(`
serverAddr = "127.0.0.1"
serverPort = {{ .%s }}
loginFailExit = true
log.level = "trace"
%s
[[proxies]]
name = "%s"
type = "udp"
localPort = {{ .%s }}
remotePort = {{ .%s }}
`, consts.PortServerName, extra, proxyName, framework.UDPEchoServerPort, remotePortName)
}
func expectProcessExit(p *process.Process, timeout time.Duration) {
select {
case <-p.Done():
+8 -28
View File
@@ -94,9 +94,11 @@ func (f *Framework) RunProcessesWithBinaries(
}
func (f *Framework) RunFrps(args ...string) (*process.Process, string, error) {
p, output, err := f.StartFrps(args...)
p := process.NewWithEnvs(TestContext.FRPServerPath, args, f.osEnvs)
f.serverProcesses = append(f.serverProcesses, p)
err := p.Start()
if err != nil {
return p, output, err
return p, p.Output(), err
}
select {
case <-p.Done():
@@ -105,39 +107,17 @@ func (f *Framework) RunFrps(args ...string) (*process.Process, string, error) {
return p, p.Output(), nil
}
// StartFrps starts frps without an implicit sleep so tests can wait on an
// explicit readiness event.
func (f *Framework) StartFrps(args ...string) (*process.Process, string, error) {
p := process.NewWithEnvs(TestContext.FRPServerPath, args, f.osEnvs)
f.serverProcesses = append(f.serverProcesses, p)
err := p.Start()
if err != nil {
return p, p.Output(), err
}
return p, p.Output(), nil
}
func (f *Framework) RunFrpc(args ...string) (*process.Process, string, error) {
p, output, err := f.StartFrpc(args...)
if err != nil {
return p, output, err
}
select {
case <-p.Done():
case <-time.After(1500 * time.Millisecond):
}
return p, p.Output(), nil
}
// StartFrpc starts frpc without an implicit sleep so tests can wait on an
// explicit login or proxy-readiness event.
func (f *Framework) StartFrpc(args ...string) (*process.Process, string, error) {
p := process.NewWithEnvs(TestContext.FRPClientPath, args, f.osEnvs)
f.clientProcesses = append(f.clientProcesses, p)
err := p.Start()
if err != nil {
return p, p.Output(), err
}
select {
case <-p.Done():
case <-time.After(1500 * time.Millisecond):
}
return p, p.Output(), nil
}
-208
View File
@@ -1,208 +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 relay
import (
"fmt"
"io"
"net"
"strconv"
"sync"
"time"
)
// HalfOpen forwards TCP connections until Blackhole is called. A blackholed
// pair stops forwarding but deliberately retains the upstream socket so the
// peer sees a real half-open connection until the relay is closed.
type HalfOpen struct {
bindAddr string
bindPort int
upstreamAddr string
listener net.Listener
done chan struct{}
accepted chan struct{}
mu sync.Mutex
pairs []*connectionPair
wg sync.WaitGroup
closeOnce sync.Once
}
type connectionPair struct {
downstream net.Conn
upstream net.Conn
mu sync.Mutex
blackholed bool
}
func New(upstreamAddr string) *HalfOpen {
return &HalfOpen{
bindAddr: "127.0.0.1",
upstreamAddr: upstreamAddr,
done: make(chan struct{}),
accepted: make(chan struct{}, 1),
}
}
func (r *HalfOpen) Run() error {
listener, err := net.Listen("tcp", net.JoinHostPort(r.bindAddr, strconv.Itoa(r.bindPort)))
if err != nil {
return err
}
r.listener = listener
r.bindPort = listener.Addr().(*net.TCPAddr).Port
r.wg.Add(1)
go r.acceptLoop()
return nil
}
func (r *HalfOpen) acceptLoop() {
defer r.wg.Done()
for {
downstream, err := r.listener.Accept()
if err != nil {
return
}
upstream, err := net.DialTimeout("tcp", r.upstreamAddr, 3*time.Second)
if err != nil {
_ = downstream.Close()
continue
}
pair := &connectionPair{downstream: downstream, upstream: upstream}
r.mu.Lock()
r.pairs = append(r.pairs, pair)
r.mu.Unlock()
select {
case r.accepted <- struct{}{}:
default:
}
r.wg.Add(1)
go r.servePair(pair)
}
}
func (r *HalfOpen) servePair(pair *connectionPair) {
defer r.wg.Done()
copyDone := make(chan struct{}, 2)
go func() {
_, _ = io.Copy(pair.upstream, pair.downstream)
copyDone <- struct{}{}
}()
go func() {
_, _ = io.Copy(pair.downstream, pair.upstream)
copyDone <- struct{}{}
}()
completed := 0
select {
case <-copyDone:
completed = 1
if !pair.isBlackholed() {
_ = pair.downstream.Close()
_ = pair.upstream.Close()
}
case <-r.done:
_ = pair.downstream.Close()
_ = pair.upstream.Close()
}
for completed < 2 {
<-copyDone
completed++
}
if pair.isBlackholed() {
<-r.done
_ = pair.downstream.Close()
_ = pair.upstream.Close()
}
}
func (r *HalfOpen) WaitForConnections(count int, timeout time.Duration) error {
timer := time.NewTimer(timeout)
defer timer.Stop()
for {
r.mu.Lock()
accepted := len(r.pairs)
r.mu.Unlock()
if accepted >= count {
return nil
}
select {
case <-r.accepted:
case <-r.done:
return fmt.Errorf("relay closed after accepting %d of %d connections", accepted, count)
case <-timer.C:
return fmt.Errorf("timed out after accepting %d of %d connections", accepted, count)
}
}
}
// Blackhole uses a one-based connection index in accept order.
func (r *HalfOpen) Blackhole(index int) error {
r.mu.Lock()
if index <= 0 || index > len(r.pairs) {
accepted := len(r.pairs)
r.mu.Unlock()
return fmt.Errorf("connection %d is unavailable; accepted %d", index, accepted)
}
pair := r.pairs[index-1]
r.mu.Unlock()
pair.mu.Lock()
if pair.blackholed {
pair.mu.Unlock()
return nil
}
pair.blackholed = true
pair.mu.Unlock()
now := time.Now()
_ = pair.downstream.SetDeadline(now)
_ = pair.upstream.SetDeadline(now)
return nil
}
func (r *HalfOpen) Close() error {
r.closeOnce.Do(func() {
close(r.done)
if r.listener != nil {
_ = r.listener.Close()
}
r.mu.Lock()
pairs := append([]*connectionPair(nil), r.pairs...)
r.mu.Unlock()
for _, pair := range pairs {
_ = pair.downstream.Close()
_ = pair.upstream.Close()
}
r.wg.Wait()
})
return nil
}
func (r *HalfOpen) BindAddr() string { return r.bindAddr }
func (r *HalfOpen) BindPort() int { return r.bindPort }
func (p *connectionPair) isBlackholed() bool {
p.mu.Lock()
defer p.mu.Unlock()
return p.blackholed
}
+3 -85
View File
@@ -5,7 +5,6 @@ import (
"fmt"
"io"
"net/http"
"time"
"github.com/onsi/ginkgo/v2"
@@ -110,12 +109,12 @@ var _ = ginkgo.Describe("[Feature: WireProtocol]", func() {
name: "default sudp visitor",
},
{
name: "v2 binary raw sudp visitor",
name: "v2 sudp visitor",
proxyWireConfig: `transport.wireProtocol = "v2"`,
visitorWireConfig: `transport.wireProtocol = "v2"`,
},
{
name: "v1 JSON proxy -> v2 Binary visitor transcode",
name: "mixed sudp proxy v1 visitor v2",
proxyWireConfig: `transport.wireProtocol = "v1"`,
visitorWireConfig: `transport.wireProtocol = "v2"`,
extraProxyConfig: `
@@ -128,7 +127,7 @@ var _ = ginkgo.Describe("[Feature: WireProtocol]", func() {
`,
},
{
name: "v2 Binary proxy -> v1 JSON visitor transcode",
name: "mixed sudp proxy v2 visitor v1",
proxyWireConfig: `transport.wireProtocol = "v2"`,
visitorWireConfig: `transport.wireProtocol = "v1"`,
},
@@ -205,87 +204,6 @@ var _ = ginkgo.Describe("[Feature: WireProtocol]", func() {
})
})
var _ = ginkgo.Describe("[Feature: BinaryUDPPacket]", func() {
f := framework.NewDefaultFramework()
for _, tc := range []struct {
name string
protocol string
extraServer string
extraTransport string
}{
{name: "tcp mux on", protocol: "tcp", extraTransport: "transport.tcpMux = true"},
{name: "tcp mux off", protocol: "tcp", extraServer: "transport.tcpMux = false", extraTransport: "transport.tcpMux = false"},
{name: "kcp", protocol: "kcp"},
{name: "quic stream", protocol: "quic"},
{name: "websocket", protocol: "websocket"},
} {
ginkgo.It(tc.name, func() {
runClientServerTest(f, &generalTestConfigures{
server: renderBindPortConfig(tc.protocol) + "\n" + tc.extraServer,
client: fmt.Sprintf(`
transport.wireProtocol = "v2"
transport.protocol = %q
%s
`, tc.protocol, tc.extraTransport),
})
})
}
ginkgo.It("wss", func() {
wssPort := f.AllocPort()
runClientServerTest(f, &generalTestConfigures{
clientPrefix: fmt.Sprintf(`
serverAddr = "127.0.0.1"
serverPort = %d
loginFailExit = false
transport.protocol = "wss"
transport.wireProtocol = "v2"
log.level = "trace"
`, wssPort),
client2: fmt.Sprintf(`
[[proxies]]
name = "wss2ws"
type = "tcp"
remotePort = %d
[proxies.plugin]
type = "https2http"
localAddr = "127.0.0.1:{{ .%s }}"
`, wssPort, consts.PortServerName),
testDelay: 10 * time.Second,
})
})
for _, tc := range []struct {
name string
transport string
}{
{name: "plain"},
{name: "aes-cfb", transport: "transport.useEncryption = true"},
{name: "snappy", transport: "transport.useCompression = true"},
{name: "snappy and aes-cfb", transport: "transport.useEncryption = true\ntransport.useCompression = true"},
{name: "limiter", transport: "transport.bandwidthLimit = \"1MB\""},
} {
ginkgo.It(tc.name, func() {
serverConf := consts.DefaultServerConfig
udpPortName := port.GenName("BinaryUDPPacket")
clientConf := consts.DefaultClientConfig + fmt.Sprintf(`
transport.wireProtocol = "v2"
[[proxies]]
name = "udp"
type = "udp"
localPort = {{ .%s }}
remotePort = {{ .%s }}
%s
`, framework.UDPEchoServerPort, udpPortName, tc.transport)
f.RunProcesses(serverConf, []string{clientConf})
framework.NewRequestExpect(f).Protocol("udp").PortName(udpPortName).Ensure()
})
}
})
type wireClientInfo struct {
ClientID string `json:"clientID"`
WireProtocol string `json:"wireProtocol"`
-300
View File
@@ -1,300 +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 features
import (
"encoding/json"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"
"github.com/onsi/ginkgo/v2"
"github.com/fatedier/frp/test/e2e/framework"
"github.com/fatedier/frp/test/e2e/framework/consts"
"github.com/fatedier/frp/test/e2e/pkg/relay"
"github.com/fatedier/frp/test/e2e/pkg/request"
)
var _ = ginkgo.Describe("[Feature: ControlReplacement]", func() {
f := framework.NewDefaultFramework()
for _, wireProtocol := range []string{"v1", "v2"} {
for _, tcpMux := range []bool{true, false} {
ginkgo.It(fmt.Sprintf("recovers a %s control through a half-open relay with tcpMux=%t", wireProtocol, tcpMux), func() {
runHalfOpenControlReplacement(f, wireProtocol, tcpMux)
})
}
}
})
func runHalfOpenControlReplacement(f *framework.Framework, wireProtocol string, tcpMux bool) {
serverPort := f.AllocPort()
dashboardPort := f.AllocPort()
remotePort := f.AllocPort()
heartbeatTimeout := int64(-1)
if !tcpMux {
heartbeatTimeout = 3
}
serverConfig := fmt.Sprintf(`
bindAddr = "127.0.0.1"
bindPort = %d
log.level = "trace"
transport.tcpMux = %t
transport.tcpMuxKeepaliveInterval = 30
transport.heartbeatTimeout = %d
webServer.addr = "127.0.0.1"
webServer.port = %d
webServer.pprofEnable = true
enablePrometheus = true
`, serverPort, tcpMux, heartbeatTimeout, dashboardPort)
serverConfigPath := f.WriteTempFile("issue-5391-frps.toml", serverConfig)
serverProcess, _, err := f.StartFrps("-c", serverConfigPath)
framework.ExpectNoError(err)
framework.ExpectNoError(framework.WaitForTCPReady(fmt.Sprintf("127.0.0.1:%d", serverPort), 5*time.Second))
halfOpenRelay := relay.New(fmt.Sprintf("127.0.0.1:%d", serverPort))
f.RunServer("", halfOpenRelay)
heartbeatInterval := int64(-1)
clientHeartbeatTimeout := int64(-1)
if !tcpMux {
heartbeatInterval = 1
clientHeartbeatTimeout = 3
}
clientConfig := fmt.Sprintf(`
serverAddr = "127.0.0.1"
serverPort = %d
clientID = "issue-5391"
loginFailExit = false
log.level = "trace"
transport.wireProtocol = %q
transport.tcpMux = %t
transport.tcpMuxKeepaliveInterval = 1
transport.heartbeatInterval = %d
transport.heartbeatTimeout = %d
transport.tls.enable = false
[[proxies]]
name = "issue-5391-tcp"
type = "tcp"
localPort = %d
remotePort = %d
`, halfOpenRelay.BindPort(), wireProtocol, tcpMux, heartbeatInterval, clientHeartbeatTimeout,
f.PortByName(framework.TCPEchoServerPort), remotePort)
clientConfigPath := f.WriteTempFile("issue-5391-frpc.toml", clientConfig)
clientProcess, _, err := f.StartFrpc("-c", clientConfigPath)
framework.ExpectNoError(err)
framework.ExpectNoError(halfOpenRelay.WaitForConnections(1, 5*time.Second))
framework.ExpectNoError(clientProcess.WaitForOutput("login to server success", 1, 10*time.Second))
framework.ExpectNoError(clientProcess.WaitForOutput("[issue-5391-tcp] start proxy success", 1, 10*time.Second))
framework.ExpectNoError(waitForReplacementState(dashboardPort, remotePort, 10*time.Second))
replacementCount := 1
if tcpMux {
replacementCount = 3
}
for i := 0; i < replacementCount; i++ {
connectionIndex := 1
if tcpMux {
connectionIndex = i + 1
}
framework.ExpectNoError(halfOpenRelay.Blackhole(connectionIndex))
if tcpMux {
framework.ExpectNoError(halfOpenRelay.WaitForConnections(connectionIndex+1, 15*time.Second))
}
framework.ExpectNoError(clientProcess.WaitForOutput("login to server success", i+2, 15*time.Second))
framework.ExpectNoError(clientProcess.WaitForOutput("[issue-5391-tcp] start proxy success", i+2, 15*time.Second))
framework.ExpectNoError(waitForReplacementState(dashboardPort, remotePort, 10*time.Second))
}
_ = clientProcess.Stop()
select {
case <-clientProcess.Done():
case <-time.After(5 * time.Second):
framework.Failf("frpc did not exit")
}
framework.ExpectNoError(waitForReplacementShutdown(dashboardPort, 10*time.Second))
framework.ExpectNoError(halfOpenRelay.Close())
framework.ExpectNoError(waitForNoHandoffWaiters(dashboardPort, 5*time.Second))
_ = serverProcess.Stop()
select {
case <-serverProcess.Done():
case <-time.After(5 * time.Second):
framework.Failf("frps did not exit")
}
}
func waitForReplacementState(dashboardPort, remotePort int, timeout time.Duration) error {
return waitForLifecycleCondition(timeout, func() error {
metricsBody, err := getLifecycleEndpoint(dashboardPort, "/metrics")
if err != nil {
return err
}
if err := expectMetricValue(metricsBody, "frp_server_client_counts", "", 1); err != nil {
return err
}
if err := expectMetricValue(metricsBody, "frp_server_proxy_counts", `type="tcp"`, 1); err != nil {
return err
}
clients, err := getOnlineLifecycleClients(dashboardPort)
if err != nil {
return err
}
if len(clients) != 1 || clients[0].ClientID != "issue-5391" {
return fmt.Errorf("expected one online client, got %+v", clients)
}
profile, err := getLifecycleEndpoint(dashboardPort, "/debug/pprof/goroutine?debug=2")
if err != nil {
return err
}
if err := expectNoHandoffWaiter(profile, "after replacement"); err != nil {
return err
}
resp, err := request.New().
TCP().
Port(remotePort).
Timeout(time.Second).
Body([]byte(consts.TestString)).
Do()
if err != nil {
return err
}
if string(resp.Content) != consts.TestString {
return fmt.Errorf("unexpected proxy response %q", resp.Content)
}
return nil
})
}
func waitForReplacementShutdown(dashboardPort int, timeout time.Duration) error {
return waitForLifecycleCondition(timeout, func() error {
metricsBody, err := getLifecycleEndpoint(dashboardPort, "/metrics")
if err != nil {
return err
}
if err := expectMetricValue(metricsBody, "frp_server_client_counts", "", 0); err != nil {
return err
}
if err := expectMetricValue(metricsBody, "frp_server_proxy_counts", `type="tcp"`, 0); err != nil {
return err
}
clients, err := getOnlineLifecycleClients(dashboardPort)
if err != nil {
return err
}
if len(clients) != 0 {
return fmt.Errorf("expected no online clients, got %+v", clients)
}
return nil
})
}
func waitForNoHandoffWaiters(dashboardPort int, timeout time.Duration) error {
return waitForLifecycleCondition(timeout, func() error {
profile, err := getLifecycleEndpoint(dashboardPort, "/debug/pprof/goroutine?debug=2")
if err != nil {
return err
}
return expectNoHandoffWaiter(profile, "after relay shutdown")
})
}
type lifecycleClient struct {
ClientID string `json:"clientID"`
}
func expectNoHandoffWaiter(profile, phase string) error {
if strings.Contains(profile, "(*Control).WaitForHandoff") {
return fmt.Errorf("control handoff waiter remained %s", phase)
}
return nil
}
func getOnlineLifecycleClients(dashboardPort int) ([]lifecycleClient, error) {
body, err := getLifecycleEndpoint(dashboardPort, "/api/clients?status=online")
if err != nil {
return nil, err
}
var clients []lifecycleClient
if err := json.Unmarshal([]byte(body), &clients); err != nil {
return nil, err
}
return clients, nil
}
func getLifecycleEndpoint(port int, path string) (string, error) {
client := &http.Client{Timeout: time.Second}
resp, err := client.Get(fmt.Sprintf("http://127.0.0.1:%d%s", port, path))
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("GET %s returned %s", path, resp.Status)
}
body, err := io.ReadAll(resp.Body)
if err != nil {
return "", err
}
return string(body), nil
}
func expectMetricValue(body, name, labels string, want float64) error {
prefix := name
if labels != "" {
prefix += "{" + labels + "}"
}
for line := range strings.SplitSeq(body, "\n") {
fields := strings.Fields(line)
if len(fields) != 2 || fields[0] != prefix {
continue
}
got, err := strconv.ParseFloat(fields[1], 64)
if err != nil {
return err
}
if got != want {
return fmt.Errorf("metric %s = %v, want %v", prefix, got, want)
}
return nil
}
return fmt.Errorf("metric %s not found", prefix)
}
func waitForLifecycleCondition(timeout time.Duration, condition func() error) error {
timer := time.NewTimer(timeout)
defer timer.Stop()
ticker := time.NewTicker(25 * time.Millisecond)
defer ticker.Stop()
var lastErr error
for {
err := condition()
if err == nil {
return nil
}
lastErr = err
select {
case <-ticker.C:
case <-timer.C:
return fmt.Errorf("condition was not met: %w", lastErr)
}
}
}
-1
View File
@@ -3,7 +3,6 @@
// @ts-nocheck
// noinspection JSUnusedGlobalSymbols
// Generated by unplugin-auto-import
// biome-ignore lint: disable
export {}
declare global {
+2 -7
View File
@@ -1,14 +1,10 @@
/* eslint-disable */
/* prettier-ignore */
// @ts-nocheck
// biome-ignore lint: disable
// oxlint-disable
// ------
// Generated by unplugin-vue-components
// Read more: https://github.com/vuejs/core/pull/3399
export {}
/* prettier-ignore */
declare module 'vue' {
export interface GlobalComponents {
ConfigField: typeof import('./src/components/ConfigField.vue')['default']
@@ -42,11 +38,10 @@ declare module 'vue' {
VisitorBaseSection: typeof import('./src/components/visitor-form/VisitorBaseSection.vue')['default']
VisitorConnectionSection: typeof import('./src/components/visitor-form/VisitorConnectionSection.vue')['default']
VisitorFormLayout: typeof import('./src/components/visitor-form/VisitorFormLayout.vue')['default']
VisitorPluginSection: typeof import('./src/components/visitor-form/VisitorPluginSection.vue')['default']
VisitorTransportSection: typeof import('./src/components/visitor-form/VisitorTransportSection.vue')['default']
VisitorXtcpSection: typeof import('./src/components/visitor-form/VisitorXtcpSection.vue')['default']
}
export interface GlobalDirectives {
export interface ComponentCustomProperties {
vLoading: typeof import('element-plus/es')['ElLoadingDirective']
}
}
View File
+13 -14
View File
@@ -9,14 +9,13 @@
"preview": "vite preview",
"build-only": "vite build",
"type-check": "vue-tsc --noEmit",
"lint": "eslint . --fix",
"lint:check": "eslint ."
"lint": "eslint --fix"
},
"dependencies": {
"element-plus": "^2.14.3",
"element-plus": "^2.13.0",
"pinia": "^3.0.4",
"vue": "^3.5.40",
"vue-router": "^5.2.0"
"vue": "^3.5.26",
"vue-router": "^4.6.4"
},
"devDependencies": {
"@types/node": "24",
@@ -24,19 +23,19 @@
"@vue/eslint-config-prettier": "^10.2.0",
"@vue/eslint-config-typescript": "^14.7.0",
"@vue/tsconfig": "^0.8.1",
"@vueuse/core": "^14.3.0",
"eslint": "^10.8.0",
"eslint-plugin-vue": "^10.10.0",
"@vueuse/core": "^14.1.0",
"eslint": "^9.39.0",
"eslint-plugin-vue": "^9.33.0",
"npm-run-all": "^4.1.5",
"prettier": "^3.9.6",
"sass": "^1.102.0",
"terser": "^5.49.0",
"prettier": "^3.7.4",
"sass": "^1.97.2",
"terser": "^5.44.1",
"typescript": "^5.9.3",
"unplugin-auto-import": "^21.0.0",
"unplugin-auto-import": "^0.17.5",
"unplugin-element-plus": "^0.11.2",
"unplugin-vue-components": "^32.1.0",
"unplugin-vue-components": "^0.26.0",
"vite": "^7.3.0",
"vite-svg-loader": "^5.1.0",
"vue-tsc": "^3.3.8"
"vue-tsc": "^3.2.2"
}
}
+2 -28
View File
@@ -12,7 +12,7 @@
<!-- number -->
<el-input
v-else-if="type === 'number'"
:model-value="numberDraft"
:model-value="modelValue != null ? String(modelValue) : ''"
:placeholder="placeholder"
:disabled="disabled"
@update:model-value="handleNumberInput($event)"
@@ -112,7 +112,7 @@
</template>
<script setup lang="ts">
import { computed, ref, watch } from 'vue'
import { computed } from 'vue'
import KeyValueEditor from './KeyValueEditor.vue'
import StringListEditor from './StringListEditor.vue'
import PopoverMenu from '@shared/components/PopoverMenu.vue'
@@ -154,42 +154,16 @@ const emit = defineEmits<{
'update:modelValue': [value: any]
}>()
const numberDraft = ref(
props.modelValue != null ? String(props.modelValue) : '',
)
const preserveNumberDraft = ref(false)
watch(
() => props.modelValue,
(value) => {
if (preserveNumberDraft.value && value == null) {
preserveNumberDraft.value = false
return
}
preserveNumberDraft.value = false
numberDraft.value = value != null ? String(value) : ''
},
)
const handleNumberInput = (val: string) => {
numberDraft.value = val
if (val === '') {
emit('update:modelValue', undefined)
return
}
// Keep a leading sign as an editing state so negative values can be typed
// naturally. The form value is updated once the input becomes a number.
if (val === '-' || val === '+') {
preserveNumberDraft.value = true
emit('update:modelValue', undefined)
return
}
const num = Number(val)
if (!isNaN(num)) {
let clamped = num
if (props.min != null && clamped < props.min) clamped = props.min
if (props.max != null && clamped > props.max) clamped = props.max
numberDraft.value = String(clamped)
emit('update:modelValue', clamped)
}
}
@@ -12,8 +12,7 @@
<ConfigField label="Bind Address" type="text" v-model="form.bindAddr"
placeholder="127.0.0.1" :readonly="readonly" />
<ConfigField label="Bind Port" type="number" v-model="form.bindPort"
:min="bindPortMin" :max="65535" prop="bindPort" :readonly="readonly"
:tip="bindPortTip" />
:min="bindPortMin" :max="65535" prop="bindPort" :readonly="readonly" />
</div>
</ConfigSection>
</template>
@@ -37,9 +36,6 @@ const form = computed({
})
const bindPortMin = computed(() => (form.value.type === 'sudp' ? 1 : undefined))
const bindPortTip = computed(() => form.value.type === 'sudp'
? ''
: 'Use -1 to skip the local listener when connections come from another visitor or plugin.')
</script>
<style scoped lang="scss">
@@ -4,11 +4,6 @@
<VisitorBaseSection v-model="form" :readonly="readonly" :editing="editing" />
</ConfigSection>
<VisitorConnectionSection v-model="form" :readonly="readonly" />
<VisitorPluginSection
v-if="form.type === 'stcp' || form.type === 'xtcp'"
v-model="form"
:readonly="readonly"
/>
<VisitorTransportSection v-model="form" :readonly="readonly" />
<VisitorXtcpSection v-if="form.type === 'xtcp'" v-model="form" :readonly="readonly" />
</div>
@@ -20,7 +15,6 @@ import type { VisitorFormData } from '../../types'
import ConfigSection from '../ConfigSection.vue'
import VisitorBaseSection from './VisitorBaseSection.vue'
import VisitorConnectionSection from './VisitorConnectionSection.vue'
import VisitorPluginSection from './VisitorPluginSection.vue'
import VisitorTransportSection from './VisitorTransportSection.vue'
import VisitorXtcpSection from './VisitorXtcpSection.vue'
@@ -1,52 +0,0 @@
<template>
<ConfigSection title="Plugin" :readonly="readonly">
<div class="field-row two-col">
<ConfigField
label="Plugin Type"
type="select"
v-model="form.pluginType"
:options="pluginOptions"
placeholder="None"
:readonly="readonly"
/>
<ConfigField
v-if="form.pluginType === 'virtual_net'"
label="Destination IP"
type="text"
v-model="form.pluginDestinationIP"
prop="pluginDestinationIP"
placeholder="10.10.10.10"
tip="Destination address in the frp virtual network."
:readonly="readonly"
/>
</div>
</ConfigSection>
</template>
<script setup lang="ts">
import { computed } from 'vue'
import type { VisitorFormData } from '../../types'
import ConfigField from '../ConfigField.vue'
import ConfigSection from '../ConfigSection.vue'
const pluginOptions = [
{ label: 'None', value: '' },
{ label: 'virtual_net', value: 'virtual_net' },
]
const props = withDefaults(defineProps<{
modelValue: VisitorFormData
readonly?: boolean
}>(), { readonly: false })
const emit = defineEmits<{ 'update:modelValue': [value: VisitorFormData] }>()
const form = computed({
get: () => props.modelValue,
set: (val) => emit('update:modelValue', val),
})
</script>
<style scoped lang="scss">
@use '@/assets/css/form-layout';
</style>
-19
View File
@@ -190,16 +190,6 @@ export function formToStoreVisitor(form: VisitorFormData): VisitorDefinition {
block.bindPort = form.bindPort
}
if (
form.pluginType === 'virtual_net' &&
(form.type === 'stcp' || form.type === 'xtcp')
) {
block.plugin = {
type: 'virtual_net',
destinationIP: form.pluginDestinationIP.trim(),
}
}
if (form.type === 'xtcp') {
if (form.protocol && form.protocol !== 'quic') {
block.protocol = form.protocol
@@ -458,15 +448,6 @@ export function storeVisitorToForm(
form.bindAddr = c.bindAddr || '127.0.0.1'
form.bindPort = c.bindPort
// Visitor plugin (only supported by stcp/xtcp; sudp runtime never uses it)
if (
c.plugin?.type === 'virtual_net' &&
(form.type === 'stcp' || form.type === 'xtcp')
) {
form.pluginType = 'virtual_net'
form.pluginDestinationIP = c.plugin.destinationIP || ''
}
// XTCP specific
form.protocol = c.protocol || 'quic'
form.keepTunnelOpen = c.keepTunnelOpen || false
-7
View File
@@ -79,10 +79,6 @@ export interface VisitorFormData {
bindAddr: string
bindPort: number | undefined
// Visitor plugin
pluginType: '' | 'virtual_net'
pluginDestinationIP: string
// XTCP specific (XTCPVisitorConfig)
protocol: string
keepTunnelOpen: boolean
@@ -160,9 +156,6 @@ export function createDefaultVisitorForm(): VisitorFormData {
bindAddr: '127.0.0.1',
bindPort: undefined,
pluginType: '',
pluginDestinationIP: '',
protocol: 'quic',
keepTunnelOpen: false,
maxRetriesAnHour: undefined,

Some files were not shown because too many files have changed in this diff Show More