Compare commits

..
12 Commits
Author SHA1 Message Date
fatedierandGitHub fa3bcca2b0 Merge pull request #5444 from fatedier/dev
Release v0.70.1
2026-07-23 15:34:31 +08:00
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
102 changed files with 1353 additions and 6672 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:
+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
+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
+3 -1
View File
@@ -1,3 +1,5 @@
## Fixes
* Fixed VirtualNet route lifecycle issues during reconnect and shutdown, including stale route cleanup, shutdown races, and reconnect backoff overflow.
* HTTP vhost servers no longer support HTTP/1.1 `Upgrade: h2c` requests. Cleartext HTTP/2 prior-knowledge remains supported.
* Fixed control-session replacement leaks when frpc reconnects through a half-open TCP multiplexed connection.
* Fixed an SSH tunnel gateway panic when handling malformed exec requests.
-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"
-8
View File
@@ -33,7 +33,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"
@@ -511,13 +510,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 {
+2 -2
View File
@@ -34,8 +34,8 @@ func newGracefulCloseTestService() *Service {
},
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, "")
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) {})}
}
+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)
})
}
}
+6 -3
View File
@@ -5,15 +5,15 @@ 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/fatedier/golib v0.8.1
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/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/rodaine/table v1.2.0
@@ -73,3 +73,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
+6 -6
View File
@@ -20,10 +20,10 @@ 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/fatedier/golib v0.8.1 h1:pHcIu0zAcZ6VTkO1dW/meelCGN5nem52DKCBY7cUvyA=
github.com/fatedier/golib v0.8.1/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.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA=
github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
github.com/go-logr/logr v1.4.2 h1:6pFjapn8bFcIbiKo3XT4j/BhANplGihG6tvd+8rYgrY=
@@ -78,8 +78,8 @@ 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/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=
+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) {
+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
View File
@@ -70,7 +70,6 @@ type TunnelServer struct {
sshConn *ssh.ServerConn
sc *ssh.ServerConfig
firstChannel ssh.Channel
firstChannelMu sync.Mutex
vc *virtual.Client
peerServerListener *netpkg.InternalListener
@@ -192,8 +191,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,11 +304,9 @@ 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 {
-44
View File
@@ -16,11 +16,7 @@ package ssh
import (
"encoding/binary"
"io"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/require"
cryptossh "golang.org/x/crypto/ssh"
@@ -73,43 +69,3 @@ func TestParseExecPayloadRejectsMalformedPayloads(t *testing.T) {
})
}
}
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 {
+1 -1
View File
@@ -14,7 +14,7 @@
package version
var version = "0.71.0"
var version = "0.70.1"
func Full() string {
return version
+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)
}
+5 -27
View File
@@ -17,7 +17,6 @@ package server
import (
"context"
"fmt"
"math"
"net"
"runtime/debug"
"sync"
@@ -46,8 +45,6 @@ type ControlID uint64
var nextControlID atomic.Uint64
const workConnPoolCapacityOffset = 10
type controlEntry struct {
ctl *Control
id ControlID
@@ -286,7 +283,7 @@ func (cm *ControlManager) GetByID(runID string) (ctl *Control, ok bool) {
// 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) {
func (cm *ControlManager) admitVisitorByRunID(runID string, admit func(user string) error) (bool, error) {
entry, ok := cm.lockCurrentRun(runID, false)
if !ok {
return false, nil
@@ -299,7 +296,7 @@ func (cm *ControlManager) admitVisitorByRunID(runID string, admit func(user, wir
if ctl.state != controlStateRunning {
return false, nil
}
return true, admit(ctl.sessionCtx.LoginMsg.User, ctl.sessionCtx.WireProtocol, ctl.sessionCtx.UDPPacketCodec)
return true, admit(ctl.sessionCtx.LoginMsg.User)
}
// RegisterWorkConn transfers conn to ctl only if ctl is still the current
@@ -371,8 +368,7 @@ type SessionContext struct {
// server configuration
ServerCfg *v1.ServerConfig
// negotiated wire protocol for this client session
WireProtocol string
UDPPacketCodec string
WireProtocol string
}
type controlState uint8
@@ -434,27 +430,10 @@ type Control 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),
workConnCh: make(chan *proxy.WorkConn, poolCount+10),
proxies: make(map[string]proxy.Proxy),
poolCount: poolCount,
portsUsedNum: 0,
@@ -842,7 +821,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
-52
View File
@@ -17,7 +17,6 @@ package server
import (
"context"
"errors"
"math"
"net"
"os"
"sync"
@@ -53,57 +52,6 @@ func TestControlPendingReplacementFinishesWithoutStarting(t *testing.T) {
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)
+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)
+21 -67
View File
@@ -29,13 +29,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"
@@ -471,7 +470,7 @@ 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)
}
}
@@ -510,12 +509,7 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
return
}
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)),
})
@@ -553,12 +547,10 @@ func (svr *Service) completeControlLogin(ctl *Control, writeSuccess func() error
}
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 +618,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 +666,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 +759,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 +769,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)
@@ -813,16 +787,15 @@ func (svr *Service) RegisterControl(
}
ctl, err := NewControl(ctx, &SessionContext{
RC: svr.rc,
PxyManager: svr.pxyManager,
PluginManager: svr.pluginManager,
AuthVerifier: authVerifier,
EncryptionKey: svr.auth.EncryptionKey(),
Conn: ctlConn,
LoginMsg: loginMsg,
ServerCfg: svr.cfg,
WireProtocol: wireProtocol,
UDPPacketCodec: udpPacketCodec,
RC: svr.rc,
PxyManager: svr.pxyManager,
PluginManager: svr.pluginManager,
AuthVerifier: authVerifier,
EncryptionKey: svr.auth.EncryptionKey(),
Conn: ctlConn,
LoginMsg: loginMsg,
ServerCfg: svr.cfg,
WireProtocol: wireProtocol,
})
if err != nil {
xl.Warnf("create new controller error: %v", err)
@@ -847,24 +820,13 @@ func (svr *Service) RegisterControl(
}
// 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{
@@ -889,22 +851,14 @@ func (svr *Service) RegisterWorkConn(
}
func (svr *Service) RegisterVisitorConn(visitorConn net.Conn, newMsg *msg.NewVisitorConn, wireProtocol string) error {
admit := func(visitorUser, visitorWireProtocol, visitorUDPPacketCodec string) error {
if visitorWireProtocol == "" {
visitorWireProtocol = wireProtocol
}
admit := func(visitorUser string) error {
return svr.rc.VisitorManager.NewConn(newMsg.ProxyName, visitorConn, newMsg.Timestamp, newMsg.SignKey,
newMsg.UseEncryption, newMsg.UseCompression, visitorUser, visitorWireProtocol, visitorUDPPacketCodec)
newMsg.UseEncryption, newMsg.UseCompression, visitorUser, wireProtocol)
}
// 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)
})
admitted, err := svr.ctlManager.admitVisitorByRunID(newMsg.RunID, admit)
if err != nil {
return err
}
@@ -913,5 +867,5 @@ func (svr *Service) RegisterVisitorConn(visitorConn net.Conn, newMsg *msg.NewVis
}
return nil
}
return admit("", wireProtocol, "")
return admit("")
}
+14 -344
View File
@@ -17,12 +17,9 @@ package server
import (
"context"
"errors"
"fmt"
"math"
"net"
"net/http"
"runtime"
"strings"
"sync"
"sync/atomic"
"testing"
@@ -33,7 +30,6 @@ import (
"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"
@@ -83,70 +79,6 @@ func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) {
}
}
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)
@@ -431,24 +363,20 @@ func TestServiceVisitorAdmissionSerializesReplacement(t *testing.T) {
}
t.Cleanup(resume)
type admissionResult struct {
admitted bool
user string
wireProtocol string
udpPacketCodec string
err error
admitted bool
user 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
var admittedUser string
admitted, admitErr := svr.ctlManager.admitVisitorByRunID("shared-run", func(user string) error {
admittedUser = user
close(admissionEntered)
<-resumeAdmission
return nil
})
admissionDone <- result
admissionDone <- admissionResult{admitted: admitted, user: admittedUser, err: admitErr}
}()
waitForSignal(t, admissionEntered, "visitor admission callback")
runMu := currentRunGateForTest(svr.ctlManager, "shared-run")
@@ -484,8 +412,6 @@ func TestServiceVisitorAdmissionSerializesReplacement(t *testing.T) {
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
@@ -502,7 +428,7 @@ func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) {
pendingConn := newCountingCloseConn()
pendingMsgConn := msg.NewConn(pendingConn, msg.NewV1ReadWriter(pendingConn))
err = registerWorkConnAsCaller(svr, pendingMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false)
err = registerWorkConnAsCaller(svr, pendingMsgConn, &msg.NewWorkConn{RunID: "shared-run"})
require.Error(t, err)
require.Equal(t, int64(1), pendingConn.closeCount.Load())
require.Len(t, ctl.workConnCh, 0)
@@ -516,7 +442,7 @@ func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) {
runningConn := newCountingCloseConn()
runningMsgConn := msg.NewConn(runningConn, msg.NewV1ReadWriter(runningConn))
require.NoError(t, svr.RegisterWorkConn(runningMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false))
require.NoError(t, svr.RegisterWorkConn(runningMsgConn, &msg.NewWorkConn{RunID: "shared-run"}))
require.Len(t, ctl.workConnCh, 1)
require.NoError(t, ctl.Close())
@@ -524,181 +450,6 @@ func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) {
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) {
@@ -714,7 +465,7 @@ func TestServiceWorkConnRoutingRejectsLostGeneration(t *testing.T) {
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)
routeDone <- registerWorkConnAsCaller(svr, workMsgConn, &msg.NewWorkConn{RunID: "shared-run"})
}()
waitForSignal(t, barrier.entered, "work connection plugin barrier")
@@ -758,7 +509,7 @@ func TestServiceVisitorRoutingExcludesPendingUser(t *testing.T) {
ClientSpec: msg.ClientSpec{
AlwaysAuthPass: true,
},
}, true, wire.ProtocolV1, "")
}, true, wire.ProtocolV1)
require.NoError(t, err)
timestamp := time.Now().Unix()
@@ -787,81 +538,6 @@ func TestServiceVisitorRoutingExcludesPendingUser(t *testing.T) {
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{}
@@ -891,7 +567,7 @@ func registerLifecycleTestControl(svr *Service) (*Control, *deadlineReadConn, er
ClientSpec: msg.ClientSpec{
AlwaysAuthPass: true,
},
}, true, wire.ProtocolV1, "")
}, true, wire.ProtocolV1)
return ctl, conn, err
}
@@ -908,14 +584,8 @@ func waitForDifferentCurrentControl(t *testing.T, manager *ControlManager, runID
return nil
}
func registerWorkConnAsCaller(
svr *Service,
workConn *msg.Conn,
newMsg *msg.NewWorkConn,
wireProtocol string,
clientHelloPresent bool,
) error {
err := svr.RegisterWorkConn(workConn, newMsg, wireProtocol, clientHelloPresent)
func registerWorkConnAsCaller(svr *Service, workConn *msg.Conn, newMsg *msg.NewWorkConn) error {
err := svr.RegisterWorkConn(workConn, newMsg)
if err != nil {
_ = workConn.Close()
}
+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():
+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"`
-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']
}
}
+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,
-16
View File
@@ -107,22 +107,6 @@ const formRules: FormRules = {
trigger: 'blur',
},
],
pluginDestinationIP: [
{
validator: (_rule, value, callback) => {
if (
form.value.pluginType === 'virtual_net' &&
(form.value.type === 'stcp' || form.value.type === 'xtcp') &&
!value?.trim()
) {
callback(new Error('Destination IP is required for virtual_net'))
return
}
callback()
},
trigger: 'blur',
},
],
}
const goBack = () => {
-73
View File
@@ -1,73 +0,0 @@
import { defineComponent } from 'vue'
import { describe, expect, it } from 'vitest'
import { mount } from '@vue/test-utils'
import ConfigField from '../src/components/ConfigField.vue'
const InputStub = defineComponent({
props: ['modelValue', 'disabled', 'placeholder', 'type'],
emits: ['update:modelValue'],
template:
'<input :value="modelValue ?? \'\'" :disabled="disabled" :placeholder="placeholder" :type="type === \'password\' ? \'password\' : \'text\'" @input="$emit(\'update:modelValue\', $event.target.value)" />',
})
const FormItemStub = defineComponent({
template: '<div class="form-item-stub"><slot /></div>',
})
const stubs = {
'el-form-item': FormItemStub,
'el-input': InputStub,
'el-switch': defineComponent({
props: ['modelValue', 'disabled'],
emits: ['update:modelValue'],
template: '<button :disabled="disabled" @click="$emit(\'update:modelValue\', !modelValue)">switch</button>',
}),
KeyValueEditor: defineComponent({ template: '<div class="key-value-stub" />' }),
StringListEditor: defineComponent({ template: '<div class="string-list-stub" />' }),
PopoverMenu: defineComponent({ template: '<div class="popover-stub"><slot /></div>' }),
PopoverMenuItem: defineComponent({ template: '<div><slot /></div>' }),
}
describe('ConfigField', () => {
it('clamps numeric input and emits the public model update', async () => {
const wrapper = mount(ConfigField, {
props: { label: 'Port', type: 'number', modelValue: 100, min: 1, max: 65535 },
global: { stubs },
})
const input = wrapper.find('input')
await input.setValue('70000')
expect(wrapper.emitted('update:modelValue')).toContainEqual([65535])
expect((input.element as HTMLInputElement).value).toBe('65535')
})
it('keeps a leading sign as an editable draft until it becomes numeric', async () => {
const wrapper = mount(ConfigField, {
props: { label: 'Port', type: 'number', modelValue: 100 },
global: { stubs },
})
const input = wrapper.find('input')
await wrapper.setProps({ modelValue: 200 })
expect((input.element as HTMLInputElement).value).toBe('200')
await input.setValue('-')
expect(wrapper.emitted('update:modelValue')).toContainEqual([undefined])
await wrapper.setProps({ modelValue: undefined })
expect((input.element as HTMLInputElement).value).toBe('-')
})
it('renders an empty readonly field as a disabled placeholder', () => {
const wrapper = mount(ConfigField, {
props: { label: 'Secret', readonly: true, modelValue: '' },
global: { stubs },
})
const input = wrapper.find('input').element as HTMLInputElement
expect(wrapper.find('.config-field-label').text()).toBe('Secret')
expect(input.disabled).toBe(true)
expect(input.value).toBe('—')
})
})
-96
View File
@@ -1,96 +0,0 @@
import { describe, expect, it } from 'vitest'
import {
createDefaultProxyForm,
createDefaultVisitorForm,
} from '../src/types/proxy-form'
import {
formToStoreProxy,
formToStoreVisitor,
storeProxyToForm,
} from '../src/types/proxy-converters'
describe('proxy converters', () => {
it('serializes only configured proxy fields into the typed store block', () => {
const form = createDefaultProxyForm()
Object.assign(form, {
name: 'web',
type: 'http',
enabled: false,
localIP: '10.0.0.8',
localPort: 8080,
useEncryption: true,
customDomains: ['example.com', ''],
locations: ['/api', ''],
metadatas: [{ key: 'team', value: 'edge' }],
})
expect(formToStoreProxy(form)).toEqual({
name: 'web',
type: 'http',
http: {
enabled: false,
localIP: '10.0.0.8',
localPort: 8080,
transport: { useEncryption: true },
metadatas: { team: 'edge' },
customDomains: ['example.com'],
locations: ['/api'],
},
})
})
it('round-trips store values while applying form defaults', () => {
const form = storeProxyToForm({
name: 'tcp-proxy',
type: 'tcp',
tcp: {
localPort: 9000,
customDomains: 'unused-for-tcp',
enabled: false,
metadatas: { owner: 'platform' },
},
})
expect(form).toMatchObject({
name: 'tcp-proxy',
type: 'tcp',
enabled: false,
localIP: '127.0.0.1',
localPort: 9000,
metadatas: [{ key: 'owner', value: 'platform' }],
multiplexer: 'httpconnect',
})
})
it.each(['stcp', 'xtcp'] as const)(
'serializes virtual_net for %s visitors',
(type) => {
const form = createDefaultVisitorForm()
Object.assign(form, {
name: 'visitor',
type,
pluginType: 'virtual_net',
pluginDestinationIP: ' 192.0.2.10 ',
bindPort: 6000,
})
expect(formToStoreVisitor(form)[type]).toMatchObject({
plugin: { type: 'virtual_net', destinationIP: '192.0.2.10' },
bindPort: 6000,
})
},
)
it('does not serialize virtual_net for SUDP visitors', () => {
const form = createDefaultVisitorForm()
Object.assign(form, {
name: 'visitor',
type: 'sudp',
pluginType: 'virtual_net',
pluginDestinationIP: '192.0.2.10',
bindPort: 6000,
})
expect(formToStoreVisitor(form).sudp).toEqual({ bindPort: 6000 })
})
})
+1 -1
View File
@@ -6,5 +6,5 @@
"moduleResolution": "bundler",
"allowSyntheticDefaultImports": true
},
"include": ["vite.config.mts"]
"include": ["vite.config.ts"]
}
+5
View File
@@ -28,10 +28,15 @@ export default defineConfig({
'@shared': fileURLToPath(new URL('../shared', import.meta.url)),
},
dedupe: ['vue', 'element-plus', '@element-plus/icons-vue'],
modules: [
fileURLToPath(new URL('../node_modules', import.meta.url)),
'node_modules',
],
},
css: {
preprocessorOptions: {
scss: {
api: 'modern',
additionalData: `@use "@shared/css/_index.scss" as *;`,
},
},
-1
View File
@@ -3,7 +3,6 @@
// @ts-nocheck
// noinspection JSUnusedGlobalSymbols
// Generated by unplugin-auto-import
// biome-ignore lint: disable
export {}
declare global {
+2 -6
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 {
ClientCard: typeof import('./src/components/ClientCard.vue')['default']
@@ -29,7 +25,7 @@ declare module 'vue' {
StatCard: typeof import('./src/components/StatCard.vue')['default']
Traffic: typeof import('./src/components/Traffic.vue')['default']
}
export interface GlobalDirectives {
export interface ComponentCustomProperties {
vLoading: typeof import('element-plus/es')['ElLoadingDirective']
}
}
+13 -14
View File
@@ -9,13 +9,12 @@
"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",
"vue": "^3.5.40",
"vue-router": "^5.2.0"
"element-plus": "^2.13.0",
"vue": "^3.5.26",
"vue-router": "^4.6.4"
},
"devDependencies": {
"@types/node": "24",
@@ -23,19 +22,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"
}
}
-48
View File
@@ -1,48 +0,0 @@
import { defineComponent } from 'vue'
import { describe, expect, it, vi } from 'vitest'
import { mount } from '@vue/test-utils'
const router = vi.hoisted(() => ({ push: vi.fn() }))
vi.mock('vue-router', () => ({ useRouter: () => router }))
import StatCard from '../src/components/StatCard.vue'
const stubs = {
'el-card': defineComponent({
template: '<div class="el-card-stub" @click="$emit(\'click\', $event)"><slot /></div>',
emits: ['click'],
}),
'el-icon': defineComponent({ template: '<span class="el-icon-stub"><slot /></span>' }),
}
describe('StatCard', () => {
it('shows its content and does not navigate without a destination', async () => {
router.push.mockClear()
const wrapper = mount(StatCard, {
props: { label: 'Clients', value: 3, subtitle: 'active' },
global: { stubs },
})
expect(wrapper.find('.stat-value').text()).toBe('3')
expect(wrapper.find('.stat-label').text()).toBe('Clients')
expect(wrapper.find('.stat-subtitle').text()).toBe('active')
expect(wrapper.find('.arrow-icon').exists()).toBe(false)
await wrapper.find('.stat-card').trigger('click')
expect(router.push).not.toHaveBeenCalled()
})
it('navigates only when a destination is provided', async () => {
router.push.mockClear()
const wrapper = mount(StatCard, {
props: { label: 'Proxies', value: '12', type: 'proxies', to: '/proxies' },
global: { stubs },
})
await wrapper.find('.stat-card').trigger('click')
expect(router.push).toHaveBeenCalledWith('/proxies')
expect(wrapper.find('.arrow-icon').exists()).toBe(true)
})
})
+1 -1
View File
@@ -6,5 +6,5 @@
"moduleResolution": "bundler",
"allowSyntheticDefaultImports": true
},
"include": ["vite.config.mts"]
"include": ["vite.config.ts"]
}
+5
View File
@@ -28,10 +28,15 @@ export default defineConfig({
'@shared': fileURLToPath(new URL('../shared', import.meta.url)),
},
dedupe: ['vue', 'element-plus', '@element-plus/icons-vue'],
modules: [
fileURLToPath(new URL('../node_modules', import.meta.url)),
'node_modules',
],
},
css: {
preprocessorOptions: {
scss: {
api: 'modern',
additionalData: `@use "@shared/css/_index.scss" as *;`,
},
},
+1063 -2255
View File
File diff suppressed because it is too large Load Diff
+1 -11
View File
@@ -1,15 +1,5 @@
{
"name": "frp-web",
"private": true,
"workspaces": ["shared", "frpc", "frps"],
"scripts": {
"test:unit": "vitest run",
"test:unit:watch": "vitest"
},
"devDependencies": {
"@vue/test-utils": "2.4.11",
"jsdom": "29.1.1",
"vitest": "4.1.10",
"vue": "^3.5.40"
}
"workspaces": ["shared", "frpc", "frps"]
}
-29
View File
@@ -1,29 +0,0 @@
import { describe, expect, it } from 'vitest'
import { mount } from '@vue/test-utils'
import ActionButton from '../components/ActionButton.vue'
describe('ActionButton', () => {
it('emits click events for an enabled button', async () => {
const wrapper = mount(ActionButton, { slots: { default: 'Save' } })
await wrapper.trigger('click')
expect(wrapper.emitted('click')).toHaveLength(1)
expect(wrapper.text()).toBe('Save')
})
it('disables itself and shows loading text while loading', async () => {
const wrapper = mount(ActionButton, {
props: { loading: true, loadingText: 'Saving...' },
slots: { default: 'Save' },
})
await wrapper.trigger('click')
expect((wrapper.element as HTMLButtonElement).disabled).toBe(true)
expect(wrapper.classes()).toContain('is-loading')
expect(wrapper.find('.spinner').exists()).toBe(true)
expect(wrapper.text()).toBe('Saving...')
expect(wrapper.emitted('click')).toBeUndefined()
})
})
-25
View File
@@ -1,25 +0,0 @@
import { describe, expect, it, vi } from 'vitest'
const capturedFetch = globalThis.fetch
describe('unit-test environment', () => {
it('blocks real network requests before test modules are evaluated', async () => {
await expect(capturedFetch('https://example.invalid/')).rejects.toThrow(
'Real network requests are disabled in unit tests',
)
})
it('allows a test to replace fetch explicitly', async () => {
const replacement = vi.fn(async () => new Response('ok'))
vi.stubGlobal('fetch', replacement)
await expect(fetch('/test-endpoint')).resolves.toBeInstanceOf(Response)
expect(replacement).toHaveBeenCalledWith('/test-endpoint')
})
it('restores the network guard after a test replacement', async () => {
await expect(fetch('/must-remain-isolated')).rejects.toThrow(
'Real network requests are disabled in unit tests',
)
})
})

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