diff --git a/.circleci/config.yml b/.circleci/config.yml index 21f4159d..7a61cbb8 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -7,14 +7,15 @@ jobs: steps: - checkout - run: - name: Build web assets (frps) - command: make install build - working_directory: web/frps + name: Test and build web assets + command: make web-ci - run: - name: Build web assets (frpc) - command: make install build - working_directory: web/frpc - - run: make + name: Check Go formatting and build binaries + command: | + set -e + make env fmt + git diff --exit-code + make build - run: make alltest workflows: diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index 3217c234..36e8d50b 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -22,12 +22,8 @@ jobs: - uses: actions/setup-node@v6 with: node-version: '22' - - 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: Test and build web assets + run: make web-ci - name: golangci-lint uses: golangci/golangci-lint-action@v9 with: diff --git a/Makefile b/Makefile index 7d949d60..2654fd84 100644 --- a/Makefile +++ b/Makefile @@ -7,7 +7,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 frps-web frpc-web frps frpc e2e-compatibility-smoke e2e-compatibility e2e-compatibility-floor +.PHONY: web web-ci frps-web frpc-web frps frpc e2e-compatibility-smoke e2e-compatibility e2e-compatibility-floor all: env fmt web build @@ -18,6 +18,9 @@ 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 diff --git a/Makefile.cross-compiles b/Makefile.cross-compiles index d084bbef..eb74d647 100644 --- a/Makefile.cross-compiles +++ b/Makefile.cross-compiles @@ -9,7 +9,7 @@ all: build build: app app: - @$(foreach n, $(os-archs), \ + @set -e; $(foreach n, $(os-archs), \ os=$(shell echo "$(n)" | cut -d : -f 1); \ arch=$(shell echo "$(n)" | cut -d : -f 2); \ extra=$(shell echo "$(n)" | cut -d : -f 3); \ diff --git a/README.md b/README.md index d850c854..c3bc2f60 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,6 @@ [![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) @@ -13,14 +12,6 @@ frp is an open source project with its ongoing development made possible entirel

Gold Sponsors

-

- - -
- The complete IDE crafted for professional Go developers -
-

-

@@ -40,6 +31,14 @@ If you're looking for a meeting recording API, consider checking out [Recall.ai] an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more. + +

+ + +
+ The complete IDE crafted for professional Go developers +
+

## What is frp? diff --git a/README_zh.md b/README_zh.md index a08d4401..600cd4fd 100644 --- a/README_zh.md +++ b/README_zh.md @@ -2,7 +2,6 @@ [![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) @@ -15,14 +14,6 @@ frp 是一个完全开源的项目,我们的开发工作完全依靠赞助者

Gold Sponsors

-

- - -
- The complete IDE crafted for professional Go developers -
-

-

@@ -42,6 +33,14 @@ If you're looking for a meeting recording API, consider checking out [Recall.ai] an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more. + +

+ + +
+ The complete IDE crafted for professional Go developers +
+

## 为什么使用 frp ? diff --git a/Release.md b/Release.md index 2549a736..33fc0aea 100644 --- a/Release.md +++ b/Release.md @@ -1,9 +1,9 @@ ## Features -* `transport.wireProtocol = "v2"` now also applies to UDP-based proxy payloads, including ordinary UDP and SUDP, so their payload framing is consistent with the selected wire protocol. -* Improved SUDP compatibility during mixed `transport.wireProtocol` deployments, allowing frps to bridge payloads between v1/default and v2 SUDP clients. -* XTCP work connection `NatHoleSid` messages now follow the selected `transport.wireProtocol`. +* UDP packet payloads for ordinary UDP proxies and SUDP now use a dedicated binary codec when frpc and frps successfully negotiate the capability under wire protocol v2, using a more compact wire representation. Wire protocol v1 remains JSON; wire protocol v2 falls back to JSON `UDPPacket` when the peer does not support or did not negotiate the capability. -## Compatibility Notes +## Fixes -* When enabling `transport.wireProtocol = "v2"` for SUDP, upgrade both the proxy and visitor frpc instances first, or keep them on `v1` until both sides are upgraded. +* Fixed a server panic and remote denial of service caused by a client sending a negative `pool_count`. Negative values are now rejected before work-connection pool resources are allocated. +* Fixed `frpc verify` ignoring configured `featureGates`, which caused VirtualNet configurations to be rejected even when the feature was enabled. +* Fixed a case-insensitive validation bypass that allowed `customDomains` under the configured `subDomainHost` to be registered using mixed-case domain names. diff --git a/client/config_manager_test.go b/client/config_manager_test.go index 07ae3297..2758b2d7 100644 --- a/client/config_manager_test.go +++ b/client/config_manager_test.go @@ -2,12 +2,16 @@ 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 { @@ -22,6 +26,256 @@ 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"), diff --git a/client/control.go b/client/control.go index bf1f6e0b..53c0b234 100644 --- a/client/control.go +++ b/client/control.go @@ -49,6 +49,8 @@ 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,9 +96,16 @@ 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) + ctl.pm = proxy.NewManager( + ctl.ctx, + sessionCtx.Common, + sessionCtx.Auth.EncryptionKey(), + ctl.msgTransporter, + sessionCtx.VnetController, + sessionCtx.UDPPacketCodec, + ) ctl.vm = visitor.NewManager(ctl.ctx, sessionCtx.RunID, sessionCtx.Common, - ctl.connectServer, ctl.msgTransporter, sessionCtx.VnetController) + ctl.connectServer, ctl.msgTransporter, sessionCtx.VnetController, sessionCtx.UDPPacketCodec) return ctl, nil } diff --git a/client/control_session.go b/client/control_session.go index d533ba2d..e3ed27fc 100644 --- a/client/control_session.go +++ b/client/control_session.go @@ -99,6 +99,7 @@ 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 } @@ -127,8 +128,9 @@ func (d *controlSessionDialer) buildLoginMsg(previousRunID string) (*msg.Login, } type loginExchangeResult struct { - resp *msg.LoginResp - crypto *wire.CryptoContext + resp *msg.LoginResp + crypto *wire.CryptoContext + udpPacketCodec string } func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) (*loginExchangeResult, error) { @@ -172,6 +174,7 @@ 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 { @@ -191,6 +194,7 @@ 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 @@ -198,8 +202,9 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) return nil, err } return &loginExchangeResult{ - resp: &loginRespMsg, - crypto: cryptoContext, + resp: &loginRespMsg, + crypto: cryptoContext, + udpPacketCodec: udpPacketCodec, }, nil } diff --git a/client/control_session_test.go b/client/control_session_test.go index a0778fba..4bbcd453 100644 --- a/client/control_session_test.go +++ b/client/control_session_test.go @@ -117,6 +117,7 @@ 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()) @@ -225,6 +226,7 @@ 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()) diff --git a/client/control_udp_test.go b/client/control_udp_test.go new file mode 100644 index 00000000..f3c91232 --- /dev/null +++ b/client/control_udp_test.go @@ -0,0 +1,125 @@ +// 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) +} diff --git a/client/proxy/proxy.go b/client/proxy/proxy.go index 4a0f4cef..ca3373a1 100644 --- a/client/proxy/proxy.go +++ b/client/proxy/proxy.go @@ -63,11 +63,12 @@ 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 = rate.NewLimiter(rate.Limit(float64(limitBytes)), int(limitBytes)) + limiter = limit.NewBandwidthLimiter(limitBytes) } baseProxy := BaseProxy{ @@ -80,6 +81,7 @@ func NewProxy( vnetController: vnetController, xl: xlog.FromContextSafe(ctx), ctx: ctx, + udpPacketCodec: udpPacketCodec, } factory := proxyFactoryRegistry[reflect.TypeOf(pxyConf)] @@ -102,9 +104,10 @@ 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 + mu sync.RWMutex + xl *xlog.Logger + ctx context.Context + udpPacketCodec string } func (pxy *BaseProxy) Run() error { @@ -209,6 +212,26 @@ 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) @@ -218,11 +241,6 @@ 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 } diff --git a/client/proxy/proxy_manager.go b/client/proxy/proxy_manager.go index e1372417..f14a8911 100644 --- a/client/proxy/proxy_manager.go +++ b/client/proxy/proxy_manager.go @@ -43,7 +43,8 @@ type Manager struct { encryptionKey []byte clientCfg *v1.ClientCommonConfig - ctx context.Context + ctx context.Context + udpPacketCodec string } func NewManager( @@ -52,6 +53,7 @@ func NewManager( encryptionKey []byte, msgTransporter transport.MessageTransporter, vnetController *vnet.Controller, + udpPacketCodec string, ) *Manager { return &Manager{ proxies: make(map[string]*Wrapper), @@ -61,6 +63,7 @@ func NewManager( encryptionKey: encryptionKey, clientCfg: clientCfg, ctx: ctx, + udpPacketCodec: udpPacketCodec, } } @@ -166,7 +169,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) + pxy := NewWrapper(pm.ctx, cfg, pm.clientCfg, pm.encryptionKey, pm.HandleEvent, pm.msgTransporter, pm.vnetController, pm.udpPacketCodec) if pm.inWorkConnCallback != nil { pxy.SetInWorkConnCallback(pm.inWorkConnCallback) } diff --git a/client/proxy/proxy_test.go b/client/proxy/proxy_test.go new file mode 100644 index 00000000..c51c6997 --- /dev/null +++ b/client/proxy/proxy_test.go @@ -0,0 +1,47 @@ +// 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) +} diff --git a/client/proxy/proxy_wrapper.go b/client/proxy/proxy_wrapper.go index 718c02e6..37a71277 100644 --- a/client/proxy/proxy_wrapper.go +++ b/client/proxy/proxy_wrapper.go @@ -99,6 +99,7 @@ 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) @@ -127,7 +128,7 @@ func NewWrapper( xl.Tracef("enable health check monitor") } - pw.pxy = NewProxy(pw.ctx, pw.Cfg, clientCfg, encryptionKey, pw.msgTransporter, pw.vnetController) + pw.pxy = NewProxy(pw.ctx, pw.Cfg, clientCfg, encryptionKey, pw.msgTransporter, pw.vnetController, udpPacketCodec) return pw } diff --git a/client/proxy/sudp.go b/client/proxy/sudp.go index 7db6d991..b5613856 100644 --- a/client/proxy/sudp.go +++ b/client/proxy/sudp.go @@ -87,7 +87,13 @@ func (pxy *SUDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) { } workConn := netpkg.WrapReadWriteCloserToConn(remote, conn) - payloadConn := msg.NewConn(workConn, msg.NewReadWriter(workConn, pxy.clientCfg.Transport.WireProtocol)) + 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) readCh := make(chan *msg.UDPPacket, 1024) sendCh := make(chan msg.Message, 1024) isClose := false diff --git a/client/proxy/udp.go b/client/proxy/udp.go index 60071560..aefe0e10 100644 --- a/client/proxy/udp.go +++ b/client/proxy/udp.go @@ -97,10 +97,17 @@ func (pxy *UDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) { return } - pxy.mu.Lock() - pxy.workConn = netpkg.WrapReadWriteCloserToConn(remote, conn) + 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) + 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.readCh = make(chan *msg.UDPPacket, 1024) pxy.sendCh = make(chan msg.Message, 1024) pxy.closed = false diff --git a/client/service.go b/client/service.go index 96a8192b..a8e1051b 100644 --- a/client/service.go +++ b/client/service.go @@ -22,6 +22,7 @@ import ( "net/http" "os" "sync" + "sync/atomic" "time" "github.com/fatedier/golib/crypto" @@ -32,6 +33,7 @@ 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" @@ -109,6 +111,9 @@ func setServiceOptionsDefault(options *ServiceOptions) error { // Service is the client service that connects to frps and provides proxy services. type Service struct { ctlMu sync.RWMutex + // Stores gracefulShutdownDuration independently from ctlMu, because the + // graceful shutdown wait may hold ctlMu for an arbitrary duration. + gracefulShutdownDuration atomic.Int64 // manager control connection with server ctl *Control // Uniq id got from frps, it will be attached to loginMsg. @@ -149,8 +154,7 @@ type Service struct { // service context ctx context.Context // call cancel to stop service - cancel context.CancelCauseFunc - gracefulShutdownDuration time.Duration + cancel context.CancelCauseFunc connectorCreator func(context.Context, *v1.ClientCommonConfig) Connector handleWorkConnCb func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) bool @@ -412,7 +416,7 @@ func (svr *Service) Close() { } func (svr *Service) GracefulClose(d time.Duration) { - svr.gracefulShutdownDuration = d + svr.gracefulShutdownDuration.Store(int64(d)) svr.cancel(nil) } @@ -429,7 +433,8 @@ func (svr *Service) stop() { svr.ctlMu.Lock() defer svr.ctlMu.Unlock() if svr.ctl != nil { - svr.ctl.GracefulClose(svr.gracefulShutdownDuration) + d := time.Duration(svr.gracefulShutdownDuration.Load()) + svr.ctl.GracefulClose(d) svr.ctl = nil } if svr.webServer != nil { @@ -506,6 +511,13 @@ 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 { diff --git a/client/service_shutdown_test.go b/client/service_shutdown_test.go new file mode 100644 index 00000000..73c7bd3d --- /dev/null +++ b/client/service_shutdown_test.go @@ -0,0 +1,95 @@ +package client + +import ( + "context" + "net" + "sync" + "testing" + "time" + + "github.com/fatedier/frp/client/proxy" + "github.com/fatedier/frp/client/visitor" + v1 "github.com/fatedier/frp/pkg/config/v1" + "github.com/fatedier/frp/pkg/msg" +) + +type gracefulCloseTestConnector struct { + conn net.Conn +} + +func (*gracefulCloseTestConnector) Connect() (*msg.Conn, error) { return nil, net.ErrClosed } +func (c *gracefulCloseTestConnector) Close() error { return c.conn.Close() } + +func newGracefulCloseTestService() *Service { + ctx := context.Background() + common := &v1.ClientCommonConfig{} + serverConn, clientConn := net.Pipe() + ctl := &Control{ + ctx: ctx, + sessionCtx: &SessionContext{ + Common: common, + RunID: "graceful-close-race", + Conn: msg.NewConn(clientConn, msg.NewV1ReadWriter(clientConn)), + Connector: &gracefulCloseTestConnector{conn: serverConn}, + }, + doneCh: make(chan struct{}), + } + ctl.pm = proxy.NewManager(ctx, common, nil, nil, nil, "") + ctl.vm = visitor.NewManager(ctx, "graceful-close-race", common, nil, nil, nil, "") + return &Service{ctl: ctl, cancel: context.CancelCauseFunc(func(error) {})} +} + +func TestGracefulCloseAndStopSynchronizeDuration(t *testing.T) { + for i := range 10000 { + svr := newGracefulCloseTestService() + start := make(chan struct{}) + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + <-start + svr.GracefulClose(time.Duration(i)) + }() + go func() { + defer wg.Done() + <-start + svr.stop() + }() + close(start) + wg.Wait() + } +} + +func TestGracefulCloseDoesNotBlockDuringStop(t *testing.T) { + const gracefulDuration = 200 * time.Millisecond + + svr := newGracefulCloseTestService() + svr.GracefulClose(gracefulDuration) + stopDone := make(chan struct{}) + go func() { + svr.stop() + close(stopDone) + }() + defer func() { + select { + case <-stopDone: + case <-time.After(time.Second): + t.Error("stop did not finish") + } + }() + + deadline := time.Now().Add(time.Second) + for svr.ctlMu.TryLock() { + svr.ctlMu.Unlock() + if time.Now().After(deadline) { + t.Fatal("stop did not acquire ctlMu") + } + time.Sleep(time.Millisecond) + } + + start := time.Now() + svr.GracefulClose(0) + if elapsed := time.Since(start); elapsed >= gracefulDuration/2 { + t.Fatalf("GracefulClose blocked for %v while stop was waiting", elapsed) + } +} diff --git a/client/visitor/sudp.go b/client/visitor/sudp.go index 91d57da6..f1939bda 100644 --- a/client/visitor/sudp.go +++ b/client/visitor/sudp.go @@ -113,7 +113,13 @@ 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") - payloadConn := msg.NewConn(workConn, msg.NewReadWriter(workConn, sv.clientCfg.Transport.WireProtocol)) + 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) wg := &sync.WaitGroup{} wg.Add(2) diff --git a/client/visitor/visitor.go b/client/visitor/visitor.go index dff0bb94..eccad013 100644 --- a/client/visitor/visitor.go +++ b/client/visitor/visitor.go @@ -50,6 +50,17 @@ 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 diff --git a/client/visitor/visitor_manager.go b/client/visitor/visitor_manager.go index 1ca194bd..3e614f7a 100644 --- a/client/visitor/visitor_manager.go +++ b/client/visitor/visitor_manager.go @@ -53,7 +53,12 @@ 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), @@ -68,6 +73,7 @@ func NewManager( vnetController: vnetController, transferConnFn: m.TransferConn, runID: runID, + udpPacketCodec: udpPacketCodec, } return m } @@ -205,6 +211,7 @@ 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) { @@ -226,3 +233,7 @@ func (v *visitorHelperImpl) VNetController() *vnet.Controller { func (v *visitorHelperImpl) RunID() string { return v.runID } + +func (v *visitorHelperImpl) UDPPacketCodec() string { + return v.udpPacketCodec +} diff --git a/cmd/frpc/sub/root.go b/cmd/frpc/sub/root.go index 0428d412..1b31be9d 100644 --- a/cmd/frpc/sub/root.go +++ b/cmd/frpc/sub/root.go @@ -36,7 +36,6 @@ 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/banner" "github.com/fatedier/frp/pkg/util/log" @@ -176,12 +175,6 @@ 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) } @@ -445,12 +438,6 @@ func runClientWithConfig(configBytes []byte, unsafeFeatures *security.UnsafeFeat proxyCfgs = config.CompleteProxyConfigurers(proxyCfgs) visitorCfgs = config.CompleteVisitorConfigurers(visitorCfgs) - if len(cfg.FeatureGates) > 0 { - if err := featuregate.SetFromMap(cfg.FeatureGates); err != nil { - return err - } - } - warning, err := validation.ValidateAllClientConfig(cfg, proxyCfgs, visitorCfgs, unsafeFeatures) if warning != nil { fmt.Printf("WARNING: %v\n", warning) diff --git a/cmd/frpc/sub/verify.go b/cmd/frpc/sub/verify.go index 9f8ddccf..e3f90210 100644 --- a/cmd/frpc/sub/verify.go +++ b/cmd/frpc/sub/verify.go @@ -29,6 +29,18 @@ 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", @@ -39,13 +51,8 @@ var verifyCmd = &cobra.Command{ } cfgFile := cfgFiles[0] - cliCfg, proxyCfgs, visitorCfgs, _, err := config.LoadClientConfig(cfgFile, strictConfigMode) - if err != nil { - fmt.Println(err) - os.Exit(1) - } unsafeFeatures := security.NewUnsafeFeatures(allowUnsafe) - warning, err := validation.ValidateAllClientConfig(cliCfg, proxyCfgs, visitorCfgs, unsafeFeatures) + warning, err := verifyClientConfig(cfgFile, strictConfigMode, unsafeFeatures) if warning != nil { fmt.Printf("WARNING: %v\n", warning) } diff --git a/cmd/frpc/sub/verify_test.go b/cmd/frpc/sub/verify_test.go new file mode 100644 index 00000000..b93db23b --- /dev/null +++ b/cmd/frpc/sub/verify_test.go @@ -0,0 +1,67 @@ +// 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) + }) + } +} diff --git a/go.mod b/go.mod index 68ea1333..f265091c 100644 --- a/go.mod +++ b/go.mod @@ -6,8 +6,8 @@ require ( github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5 github.com/charmbracelet/lipgloss v1.1.0 github.com/charmbracelet/log v0.4.2 - github.com/coreos/go-oidc/v3 v3.14.1 - github.com/fatedier/golib v0.7.0 + github.com/coreos/go-oidc/v3 v3.18.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 @@ -15,10 +15,9 @@ require ( 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/pion/stun/v3 v3.1.1 - github.com/pires/go-proxyproto v0.7.0 + github.com/pires/go-proxyproto v0.15.0 github.com/prometheus/client_golang v1.19.1 - github.com/quic-go/quic-go v0.55.0 + github.com/quic-go/quic-go v0.60.0 github.com/rodaine/table v1.2.0 github.com/samber/lo v1.47.0 github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8 @@ -28,11 +27,11 @@ require ( github.com/tidwall/gjson v1.17.1 github.com/vishvananda/netlink v1.3.0 github.com/xtaci/kcp-go/v5 v5.6.13 - golang.org/x/crypto v0.49.0 - golang.org/x/net v0.52.0 - golang.org/x/oauth2 v0.28.0 - golang.org/x/sync v0.20.0 - golang.org/x/sys v0.42.0 + golang.org/x/crypto v0.54.0 + golang.org/x/net v0.56.0 + golang.org/x/oauth2 v0.36.0 + golang.org/x/sync v0.22.0 + golang.org/x/sys v0.47.0 golang.org/x/time v0.10.0 golang.zx2c4.com/wireguard v0.0.0-20231211153847-12269c276173 gopkg.in/ini.v1 v1.67.0 @@ -51,7 +50,7 @@ require ( github.com/charmbracelet/x/cellbuf v0.0.13-0.20250311204145-2c3ea96c31dd // indirect github.com/charmbracelet/x/term v0.2.1 // indirect github.com/davecgh/go-spew v1.1.1 // indirect - github.com/go-jose/go-jose/v4 v4.0.5 // indirect + github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/go-logfmt/logfmt v0.6.0 // indirect github.com/go-logr/logr v1.4.2 // indirect github.com/go-task/slim-sprig/v3 v3.0.0 // indirect @@ -65,9 +64,6 @@ require ( github.com/mattn/go-isatty v0.0.20 // indirect github.com/mattn/go-runewidth v0.0.16 // indirect github.com/muesli/termenv v0.16.0 // indirect - github.com/pion/dtls/v3 v3.0.10 // indirect - github.com/pion/logging v0.2.4 // indirect - github.com/pion/transport/v4 v4.0.1 // indirect github.com/pkg/errors v0.9.1 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect github.com/prometheus/client_model v0.5.0 // indirect @@ -80,13 +76,11 @@ require ( github.com/tidwall/pretty v1.2.0 // indirect github.com/tjfoc/gmsm v1.4.1 // indirect github.com/vishvananda/netns v0.0.4 // indirect - github.com/wlynxg/anet v0.0.5 // indirect github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect go.uber.org/automaxprocs v1.6.0 // indirect golang.org/x/exp v0.0.0-20231006140011-7918f672742d // indirect - golang.org/x/mod v0.33.0 // indirect - golang.org/x/text v0.35.0 // indirect - golang.org/x/tools v0.42.0 // indirect + golang.org/x/text v0.40.0 // indirect + golang.org/x/tools v0.47.0 // indirect golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect google.golang.org/protobuf v1.36.5 // indirect gopkg.in/yaml.v2 v2.4.0 // indirect @@ -99,4 +93,4 @@ require ( replace github.com/hashicorp/yamux => github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6 // Use the Lolia-FRP fork of golib: io.Join relays with adaptively sized buffers. -replace github.com/fatedier/golib => github.com/Lolia-FRP/golib v0.0.0-20260704205217-7f676961e707 +replace github.com/fatedier/golib => github.com/Lolia-FRP/golib v0.0.0-20260810035216-6e2252e33564 diff --git a/go.sum b/go.sum index 3c07a076..c1dd13d6 100644 --- a/go.sum +++ b/go.sum @@ -2,8 +2,8 @@ cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMT github.com/Azure/go-ntlmssp v0.1.0 h1:DjFo6YtWzNqNvQdrwEyr/e4nhU3vRiwenz5QX7sFz+A= github.com/Azure/go-ntlmssp v0.1.0/go.mod h1:NYqdhxd/8aAct/s4qSYZEerdPuH1liG2/X9DiVTbhpk= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= -github.com/Lolia-FRP/golib v0.0.0-20260704205217-7f676961e707 h1:ROQVjgk+RvaN8J8uR8ekUu4TKgJLXYObVeQC/0KADn4= -github.com/Lolia-FRP/golib v0.0.0-20260704205217-7f676961e707/go.mod h1:ArUGvPg2cOw/py2RAuBt46nNZH2VQ5Z70p109MAZpJw= +github.com/Lolia-FRP/golib v0.0.0-20260810035216-6e2252e33564 h1:L9XoKV/oAoTJNe/h93oNPlhMf6EsM0V54/tpucjbz1A= +github.com/Lolia-FRP/golib v0.0.0-20260810035216-6e2252e33564/go.mod h1:ArUGvPg2cOw/py2RAuBt46nNZH2VQ5Z70p109MAZpJw= github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5 h1:0CwZNZbxp69SHPdPJAN/hZIm0C4OItdklCFmMRWYpio= github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5/go.mod h1:wHh0iHkYZB8zMSxRWpUBQtwG5a7fFgvEO+odwuTv2gs= github.com/aymanbagabas/go-osc52/v2 v2.0.1 h1:HwpRHbFMcZLEVr42D4p7XBqjyuxQH5SMiErDT4WkJ2k= @@ -27,8 +27,8 @@ github.com/charmbracelet/x/term v0.2.1 h1:AQeHeLZ1OqSXhrAWpYUtZyX1T3zVxfpZuEQMIQ github.com/charmbracelet/x/term v0.2.1/go.mod h1:oQ4enTYFV7QN4m0i9mzHrViD7TQKvNEEkHUMCmsxdUg= github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw= github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc= -github.com/coreos/go-oidc/v3 v3.14.1 h1:9ePWwfdwC4QKRlCXsJGou56adA/owXczOzwKdOumLqk= -github.com/coreos/go-oidc/v3 v3.14.1/go.mod h1:HaZ3szPaZ0e4r6ebqvsLWlk2Tn+aejfmrfah6hnSYEU= +github.com/coreos/go-oidc/v3 v3.18.0 h1:V9orjXynvu5wiC9SemFTWnG4F45v403aIcjWo0d41+A= +github.com/coreos/go-oidc/v3 v3.18.0/go.mod h1:DYCf24+ncYi+XkIH97GY1+dqoRlbaSI26KVTCI9SrY4= github.com/cpuguy83/go-md2man/v2 v2.0.3/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= @@ -38,8 +38,8 @@ github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1m github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c= github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6 h1:u92UUy6FURPmNsMBUuongRWC0rBqN6gd01Dzu+D21NE= github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6/go.mod h1:c5/tk6G0dSpXGzJN7Wk1OEie8grdSJAmeawId9Zvd34= -github.com/go-jose/go-jose/v4 v4.0.5 h1:M6T8+mKZl/+fNNuFHvGIzDz7BTLQPIounk/b9dw3AaE= -github.com/go-jose/go-jose/v4 v4.0.5/go.mod h1:s3P1lRrkT8igV8D9OjyL4WRyHvjB6a4JSllnOrmmBOA= +github.com/go-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-logfmt/logfmt v0.6.0 h1:wGYYu3uicYdqXVgoYbvnkrPVXkuLM1p1ifugDMEdRi4= github.com/go-logfmt/logfmt v0.6.0/go.mod h1:WYhtIu8zTZfxdn5+rREduYbwxfcBr/Vr6KEVveWlfTs= github.com/go-logr/logr v1.4.2 h1:6pFjapn8bFcIbiKo3XT4j/BhANplGihG6tvd+8rYgrY= @@ -101,16 +101,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/pion/dtls/v3 v3.0.10 h1:k9ekkq1kaZoxnNEbyLKI8DI37j/Nbk1HWmMuywpQJgg= -github.com/pion/dtls/v3 v3.0.10/go.mod h1:YEmmBYIoBsY3jmG56dsziTv/Lca9y4Om83370CXfqJ8= -github.com/pion/logging v0.2.4 h1:tTew+7cmQ+Mc1pTBLKH2puKsOvhm32dROumOZ655zB8= -github.com/pion/logging v0.2.4/go.mod h1:DffhXTKYdNZU+KtJ5pyQDjvOAh/GsNSyv1lbkFbe3so= -github.com/pion/stun/v3 v3.1.1 h1:CkQxveJ4xGQjulGSROXbXq94TAWu8gIX2dT+ePhUkqw= -github.com/pion/stun/v3 v3.1.1/go.mod h1:qC1DfmcCTQjl9PBaMa5wSn3x9IPmKxSdcCsxBcDBndM= -github.com/pion/transport/v4 v4.0.1 h1:sdROELU6BZ63Ab7FrOLn13M6YdJLY20wldXW2Cu2k8o= -github.com/pion/transport/v4 v4.0.1/go.mod h1:nEuEA4AD5lPdcIegQDpVLgNoDGreqM/YqmEx3ovP4jM= -github.com/pires/go-proxyproto v0.7.0 h1:IukmRewDQFWC7kfnb66CSomk2q/seBuilHBYFwyq0Hs= -github.com/pires/go-proxyproto v0.7.0/go.mod h1:Vz/1JPY/OACxWGQNIRY2BeyDmpoaWmEP40O9LbuiFR4= +github.com/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/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= @@ -126,8 +118,10 @@ github.com/prometheus/common v0.48.0 h1:QO8U2CdOzSn1BBsmXJXduaaW+dY/5QLjfB8svtSz github.com/prometheus/common v0.48.0/go.mod h1:0/KsvlIEfPQCQ5I2iNSAWKPZziNCvRs5EC6ILDTlAPc= github.com/prometheus/procfs v0.12.0 h1:jluTpSng7V9hY0O2R9DzzJHYb2xULk9VTR1V1R/k6Bo= github.com/prometheus/procfs v0.12.0/go.mod h1:pcuDEFsWDnvcgNzo4EEweacyhjeA9Zk3cnaOZAZEfOo= -github.com/quic-go/quic-go v0.55.0 h1:zccPQIqYCXDt5NmcEabyYvOnomjs8Tlwl7tISjJh9Mk= -github.com/quic-go/quic-go v0.55.0/go.mod h1:DR51ilwU1uE164KuWXhinFcKWGlEjzys2l8zUl5Ss1U= +github.com/quic-go/go-ossfuzz-seeds v0.1.0 h1:APacT+iIaNF6fd8AGEiN3bT/Jtkd2jz4v4TzM7MFjy0= +github.com/quic-go/go-ossfuzz-seeds v0.1.0/go.mod h1:3IOHRbJIc+L6YKMwfDtJAM9Vj9k0YY4muhuyUYk5tbk= +github.com/quic-go/quic-go v0.60.0 h1:xcQioE8OM66UQLeUMHltK1CCcOu3JbVB4JAQdDQSB+0= +github.com/quic-go/quic-go v0.60.0/go.mod h1:wpKpjmPpftl30sL6pFh7REVpjbcCVy4zt2vDyK1TuJk= github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= @@ -170,8 +164,6 @@ github.com/vishvananda/netlink v1.3.0 h1:X7l42GfcV4S6E4vHTsw48qbrV+9PVojNfIhZcwQ github.com/vishvananda/netlink v1.3.0/go.mod h1:i6NetklAujEcC6fK0JPjT8qSwWyO0HLn4UKG+hGqeJs= github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8= github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM= -github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU= -github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA= github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e h1:JVG44RsyaB9T2KIHavMF/ppJZNG9ZpyihvCd0w101no= github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e/go.mod h1:RbqR21r5mrJuqunuUZ/Dhy/avygyECGrLceyNeo4LiM= github.com/xtaci/kcp-go/v5 v5.6.13 h1:FEjtz9+D4p8t2x4WjciGt/jsIuhlWjjgPCCWjrVR4Hk= @@ -185,32 +177,30 @@ go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= -golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4= -golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA= +golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= +golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= golang.org/x/exp v0.0.0-20231006140011-7918f672742d h1:jtJma62tbqLibJ5sFQz8bKtEM8rJBtfilJ2qTU199MI= golang.org/x/exp v0.0.0-20231006140011-7918f672742d/go.mod h1:ldy0pHrwJyGW56pPQzzkH36rKxoZW1tw7ZJpeKx+hdo= golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU= golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= -golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8= -golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= -golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0= -golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw= +golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o= +golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec= golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= -golang.org/x/oauth2 v0.28.0 h1:CrgCKl8PPAVtLnU3c+EDw6x11699EWlsDeWNWKdIOkc= -golang.org/x/oauth2 v0.28.0/go.mod h1:onh5ek6nERTohokkhCD/y2cV4Do3fxFHFuAejCkRWT8= +golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs= +golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= -golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= @@ -219,14 +209,14 @@ golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo= -golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -golang.org/x/term v0.41.0 h1:QCgPso/Q3RTJx2Th4bDLqML4W6iJiaXFq2/ftQF13YU= -golang.org/x/term v0.41.0/go.mod h1:3pfBgksrReYfZ5lvYM0kSO0LIkAl4Yl2bXOkKP7Ec2A= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0= +golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= -golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8= -golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA= +golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= +golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= golang.org/x/time v0.10.0 h1:3usCWA8tQn0L8+hFJQNgzpWbd89begxN66o1Ojdn5L4= golang.org/x/time v0.10.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= @@ -234,8 +224,8 @@ golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGm golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY= golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q= -golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k= -golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0= +golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= +golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI= diff --git a/pkg/config/load.go b/pkg/config/load.go index 38634fc5..2aaf1457 100644 --- a/pkg/config/load.go +++ b/pkg/config/load.go @@ -394,6 +394,10 @@ func LoadClientConfigResult(path string, strict bool) (*ClientConfigLoadResult, } } + if err := validateNoDuplicateNames(result.Proxies, result.Visitors); err != nil { + return nil, err + } + return result, nil } @@ -417,6 +421,31 @@ func LoadClientConfig(path string, strict bool) ( return result.Common, proxyCfgs, visitorCfgs, result.IsLegacyFormat, nil } +// validateNoDuplicateNames rejects proxies or visitors that share a name. They are +// keyed by name in the config sources, so a duplicate would otherwise be silently +// overwritten and never started, with no error or log. +func validateNoDuplicateNames(proxies []v1.ProxyConfigurer, visitors []v1.VisitorConfigurer) error { + proxyNames := make(map[string]struct{}, len(proxies)) + for _, p := range proxies { + name := p.GetBaseConfig().Name + if _, ok := proxyNames[name]; ok { + return fmt.Errorf("proxy name [%s] is duplicated", name) + } + proxyNames[name] = struct{}{} + } + + visitorNames := make(map[string]struct{}, len(visitors)) + for _, v := range visitors { + name := v.GetBaseConfig().Name + if _, ok := visitorNames[name]; ok { + return fmt.Errorf("visitor name [%s] is duplicated", name) + } + visitorNames[name] = struct{}{} + } + + return nil +} + func CompleteProxyConfigurers(proxies []v1.ProxyConfigurer) []v1.ProxyConfigurer { proxyCfgs := proxies for _, c := range proxyCfgs { diff --git a/pkg/config/load_test.go b/pkg/config/load_test.go index b711a5c1..a1e29247 100644 --- a/pkg/config/load_test.go +++ b/pkg/config/load_test.go @@ -17,6 +17,8 @@ package config import ( "encoding/json" "fmt" + "os" + "path/filepath" "strings" "testing" @@ -462,6 +464,111 @@ func TestFilterClientConfigurers_FilterByStartAndEnabled(t *testing.T) { require.Equal("keep", proxies[0].GetBaseConfig().Name) } +func TestLoadClientConfigResult_DuplicateNames(t *testing.T) { + tests := []struct { + name string + content string + errSubstr string + }{ + { + name: "duplicate proxy names", + content: ` +serverAddr = "127.0.0.1" +serverPort = 7000 + +[[proxies]] +name = "dup" +type = "tcp" +localPort = 22 +remotePort = 6000 + +[[proxies]] +name = "dup" +type = "tcp" +localPort = 3306 +remotePort = 6001 +`, + errSubstr: "proxy name [dup] is duplicated", + }, + { + name: "duplicate visitor names", + content: ` +serverAddr = "127.0.0.1" +serverPort = 7000 + +[[visitors]] +name = "dup" +type = "stcp" +serverName = "a" +secretKey = "secret" +bindPort = 9001 + +[[visitors]] +name = "dup" +type = "stcp" +serverName = "b" +secretKey = "secret" +bindPort = 9002 +`, + errSubstr: "visitor name [dup] is duplicated", + }, + { + name: "unique names", + content: ` +serverAddr = "127.0.0.1" +serverPort = 7000 + +[[proxies]] +name = "p1" +type = "tcp" +localPort = 22 +remotePort = 6000 + +[[proxies]] +name = "p2" +type = "tcp" +localPort = 3306 +remotePort = 6001 +`, + }, + { + name: "same name across proxy and visitor", + content: ` +serverAddr = "127.0.0.1" +serverPort = 7000 + +[[proxies]] +name = "same" +type = "tcp" +localPort = 22 +remotePort = 6000 + +[[visitors]] +name = "same" +type = "stcp" +serverName = "a" +secretKey = "secret" +bindPort = 9001 +`, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + require := require.New(t) + path := filepath.Join(t.TempDir(), "frpc.toml") + require.NoError(os.WriteFile(path, []byte(tc.content), 0o600)) + + _, err := LoadClientConfigResult(path, false) + if tc.errSubstr == "" { + require.NoError(err) + } else { + require.ErrorContains(err, tc.errSubstr) + } + }) + } +} + // TestYAMLEdgeCases tests edge cases for YAML parsing, including non-map types func TestYAMLEdgeCases(t *testing.T) { require := require.New(t) diff --git a/pkg/config/v1/server.go b/pkg/config/v1/server.go index b0b2baa4..da5d5677 100644 --- a/pkg/config/v1/server.go +++ b/pkg/config/v1/server.go @@ -174,7 +174,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. + // this value is 5. Negative values are invalid. 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 diff --git a/pkg/config/v1/validation/client.go b/pkg/config/v1/validation/client.go index 77004317..28409c4b 100644 --- a/pkg/config/v1/validation/client.go +++ b/pkg/config/v1/validation/client.go @@ -51,14 +51,51 @@ 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 !featuregate.Enabled(featuregate.VirtualNet) { + if !gates.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) { diff --git a/pkg/config/v1/validation/client_test.go b/pkg/config/v1/validation/client_test.go new file mode 100644 index 00000000..1370e890 --- /dev/null +++ b/pkg/config/v1/validation/client_test.go @@ -0,0 +1,140 @@ +// 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()) +} diff --git a/pkg/config/v1/validation/proxy.go b/pkg/config/v1/validation/proxy.go index 744620f0..5d430e6b 100644 --- a/pkg/config/v1/validation/proxy.go +++ b/pkg/config/v1/validation/proxy.go @@ -79,9 +79,11 @@ 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 { - if s.SubDomainHost != "" && len(strings.Split(s.SubDomainHost, ".")) < len(strings.Split(domain, ".")) { - if strings.HasSuffix(domain, "."+s.SubDomainHost) { + canonicalDomain := strings.ToLower(domain) + if subDomainHost != "" && len(strings.Split(subDomainHost, ".")) < len(strings.Split(canonicalDomain, ".")) { + if strings.HasSuffix(canonicalDomain, "."+subDomainHost) { return fmt.Errorf("custom domain [%s] should not belong to subdomain host [%s]", domain, s.SubDomainHost) } } diff --git a/pkg/config/v1/validation/proxy_test.go b/pkg/config/v1/validation/proxy_test.go new file mode 100644 index 00000000..b9411dc8 --- /dev/null +++ b/pkg/config/v1/validation/proxy_test.go @@ -0,0 +1,76 @@ +// 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) + }) + } +} diff --git a/pkg/config/v1/validation/server.go b/pkg/config/v1/validation/server.go index 8be740aa..54bd8348 100644 --- a/pkg/config/v1/validation/server.go +++ b/pkg/config/v1/validation/server.go @@ -51,6 +51,9 @@ 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) { diff --git a/pkg/config/v1/validation/server_test.go b/pkg/config/v1/validation/server_test.go new file mode 100644 index 00000000..7c0c57ad --- /dev/null +++ b/pkg/config/v1/validation/server_test.go @@ -0,0 +1,51 @@ +// 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) + }) + } +} diff --git a/pkg/metrics/mem/server.go b/pkg/metrics/mem/server.go index 999719fd..3ff9c566 100644 --- a/pkg/metrics/mem/server.go +++ b/pkg/metrics/mem/server.go @@ -95,9 +95,7 @@ func (m *serverMetrics) clearUselessInfo(continuousOfflineDuration time.Duration defer m.mu.Unlock() total = len(m.info.ProxyStatistics) for name, data := range m.info.ProxyStatistics { - if !data.LastCloseTime.IsZero() && - data.LastStartTime.Before(data.LastCloseTime) && - m.clock.Since(data.LastCloseTime) > continuousOfflineDuration { + if m.shouldClearProxyStats(data, continuousOfflineDuration) { delete(m.info.ProxyStatistics, name) count++ log.Tracef("clear proxy [%s]'s statistics data, lastCloseTime: [%s]", name, data.LastCloseTime.String()) @@ -106,10 +104,20 @@ func (m *serverMetrics) clearUselessInfo(continuousOfflineDuration time.Duration return count, total } +func (m *serverMetrics) shouldClearProxyStats(data *ProxyStatistics, continuousOfflineDuration time.Duration) bool { + return !data.LastCloseTime.IsZero() && + data.LastStartTime.Before(data.LastCloseTime) && + m.clock.Since(data.LastCloseTime) > continuousOfflineDuration +} + func (m *serverMetrics) ClearOfflineProxies() (int, int) { return m.clearUselessInfo(0) } +func (m *serverMetrics) PruneOfflineProxies() (int, int) { + return m.clearUselessInfo(0) +} + func (m *serverMetrics) NewClient() { m.info.ClientCounts.Inc(1) } @@ -231,9 +239,11 @@ func toProxyStats(name string, proxyStats *ProxyStatistics) *ProxyStats { } if !proxyStats.LastStartTime.IsZero() { ps.LastStartTime = proxyStats.LastStartTime.Format("01-02 15:04:05") + ps.LastStartAt = proxyStats.LastStartTime.Unix() } if !proxyStats.LastCloseTime.IsZero() { ps.LastCloseTime = proxyStats.LastCloseTime.Format("01-02 15:04:05") + ps.LastCloseAt = proxyStats.LastCloseTime.Unix() } return ps } diff --git a/pkg/metrics/mem/server_test.go b/pkg/metrics/mem/server_test.go index fe9f9984..12040c9f 100644 --- a/pkg/metrics/mem/server_test.go +++ b/pkg/metrics/mem/server_test.go @@ -22,6 +22,12 @@ func TestServerMetricsUsesClockForProxyTimestamps(t *testing.T) { clk.SetTime(closedAt) metrics.CloseProxy("proxy", "tcp") require.Equal(closedAt, metrics.info.ProxyStatistics["proxy"].LastCloseTime) + + stats := metrics.GetProxyByName("proxy") + require.Equal(start.Format("01-02 15:04:05"), stats.LastStartTime) + require.Equal(closedAt.Format("01-02 15:04:05"), stats.LastCloseTime) + require.Equal(start.Unix(), stats.LastStartAt) + require.Equal(closedAt.Unix(), stats.LastCloseAt) } func TestServerMetricsClearUselessInfoUsesClock(t *testing.T) { @@ -43,6 +49,70 @@ func TestServerMetricsClearUselessInfoUsesClock(t *testing.T) { require.Empty(metrics.info.ProxyStatistics) } +func TestServerMetricsClearOfflineProxiesPreservesLegacyTotal(t *testing.T) { + require := require.New(t) + + start := time.Date(2026, time.May, 8, 12, 30, 0, 0, time.UTC) + clk := clocktesting.NewFakeClock(start.Add(time.Minute)) + metrics := newServerMetricsWithClock(clk) + metrics.info.ProxyStatistics["offline"] = &ProxyStatistics{ + Name: "offline", + LastStartTime: start.Add(-time.Hour), + LastCloseTime: start, + } + metrics.info.ProxyStatistics["online"] = &ProxyStatistics{ + Name: "online", + LastStartTime: start, + } + + cleared, total := metrics.ClearOfflineProxies() + + require.Equal(1, cleared) + require.Equal(2, total) + require.False(metrics.hasProxyStatistics("offline")) + require.True(metrics.hasProxyStatistics("online")) +} + +func TestServerMetricsPruneOfflineProxiesReportsTotalStats(t *testing.T) { + require := require.New(t) + + start := time.Date(2026, time.May, 8, 12, 30, 0, 0, time.UTC) + clk := clocktesting.NewFakeClock(start.Add(time.Minute)) + metrics := newServerMetricsWithClock(clk) + metrics.info.ProxyStatistics["offline"] = &ProxyStatistics{ + Name: "offline", + LastStartTime: start.Add(-time.Hour), + LastCloseTime: start, + } + metrics.info.ProxyStatistics["online"] = &ProxyStatistics{ + Name: "online", + LastStartTime: start, + } + metrics.info.ProxyStatistics["restarted"] = &ProxyStatistics{ + Name: "restarted", + LastStartTime: start.Add(30 * time.Second), + LastCloseTime: start, + } + metrics.info.ProxyStatistics["same-time"] = &ProxyStatistics{ + Name: "same-time", + LastStartTime: start, + LastCloseTime: start, + } + + cleared, total := metrics.PruneOfflineProxies() + + require.Equal(1, cleared) + require.Equal(4, total) + require.False(metrics.hasProxyStatistics("offline")) + require.True(metrics.hasProxyStatistics("online")) + require.True(metrics.hasProxyStatistics("restarted")) + require.True(metrics.hasProxyStatistics("same-time")) + + cleared, total = metrics.PruneOfflineProxies() + require.Equal(0, cleared) + require.Equal(3, total) +} + func TestServerMetricsRunUsesClockTicker(t *testing.T) { require := require.New(t) diff --git a/pkg/metrics/mem/types.go b/pkg/metrics/mem/types.go index b7661ba8..6361693b 100644 --- a/pkg/metrics/mem/types.go +++ b/pkg/metrics/mem/types.go @@ -41,6 +41,8 @@ type ProxyStats struct { TodayTrafficOut int64 LastStartTime string LastCloseTime string + LastStartAt int64 + LastCloseAt int64 CurConns int64 } @@ -85,4 +87,5 @@ type Collector interface { GetProxyByName(proxyName string) *ProxyStats GetProxyTraffic(name string) *ProxyTrafficInfo ClearOfflineProxies() (int, int) + PruneOfflineProxies() (int, int) } diff --git a/pkg/msg/udp_benchmark_test.go b/pkg/msg/udp_benchmark_test.go new file mode 100644 index 00000000..575938ec --- /dev/null +++ b/pkg/msg/udp_benchmark_test.go @@ -0,0 +1,199 @@ +// 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 + }) + } + } + } +} diff --git a/pkg/msg/udp_binary.go b/pkg/msg/udp_binary.go new file mode 100644 index 00000000..7493c70f --- /dev/null +++ b/pkg/msg/udp_binary.go @@ -0,0 +1,338 @@ +// 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) + } +} diff --git a/pkg/msg/udp_binary_test.go b/pkg/msg/udp_binary_test.go new file mode 100644 index 00000000..7dac31fb --- /dev/null +++ b/pkg/msg/udp_binary_test.go @@ -0,0 +1,248 @@ +// 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) + }) +} diff --git a/pkg/msg/wire_v2.go b/pkg/msg/wire_v2.go index 8d2cd88d..f3da2330 100644 --- a/pkg/msg/wire_v2.go +++ b/pkg/msg/wire_v2.go @@ -43,6 +43,7 @@ const ( V2TypeNatHoleResp uint16 = 16 V2TypeNatHoleSid uint16 = 17 V2TypeNatHoleReport uint16 = 18 + V2TypeUDPPacketBinary uint16 = 19 ) var v2MsgTypeMap = map[uint16]any{ diff --git a/pkg/msg/wire_v2_test.go b/pkg/msg/wire_v2_test.go index f6f25e55..ea6e8d5d 100644 --- a/pkg/msg/wire_v2_test.go +++ b/pkg/msg/wire_v2_test.go @@ -84,6 +84,9 @@ 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) { diff --git a/pkg/nathole/discovery.go b/pkg/nathole/discovery.go index 6fa0140f..241b0f6c 100644 --- a/pkg/nathole/discovery.go +++ b/pkg/nathole/discovery.go @@ -15,31 +15,24 @@ package nathole import ( + "errors" "fmt" "net" "time" - "github.com/pion/stun/v3" + "github.com/fatedier/golib/net/stun" ) var responseTimeout = 3 * time.Second -type Message struct { - Body []byte - Addr string -} - // If the localAddr is empty, it will listen on a random port. func Discover(stunServers []string, localAddr string) ([]string, net.Addr, error) { - // create a discoverConn and get response from messageChan discoverConn, err := listen(localAddr) if err != nil { return nil, nil, err } defer discoverConn.Close() - go discoverConn.readLoop() - addresses := make([]string, 0, len(stunServers)) for _, addr := range stunServers { // get external address from stun server @@ -58,10 +51,9 @@ type stunResponse struct { } type discoverConn struct { - conn *net.UDPConn - - localAddr net.Addr - messageChan chan *Message + conn *net.UDPConn + client *stun.Client + localAddr net.Addr } func listen(localAddr string) (*discoverConn, error) { @@ -77,82 +69,50 @@ func listen(localAddr string) (*discoverConn, error) { if err != nil { return nil, err } + client, err := stun.NewClient(conn) + if err != nil { + _ = conn.Close() + return nil, err + } return &discoverConn{ - conn: conn, - localAddr: conn.LocalAddr(), - messageChan: make(chan *Message, 10), + conn: conn, + client: client, + localAddr: conn.LocalAddr(), }, nil } func (c *discoverConn) Close() error { - if c.messageChan != nil { - close(c.messageChan) - c.messageChan = nil - } return c.conn.Close() } -func (c *discoverConn) readLoop() { - for { - buf := make([]byte, 1024) - n, addr, err := c.conn.ReadFromUDP(buf) - if err != nil { - return - } - buf = buf[:n] - - c.messageChan <- &Message{ - Body: buf, - Addr: addr.String(), - } - } -} - func (c *discoverConn) doSTUNRequest(addr string) (*stunResponse, error) { serverAddr, err := net.ResolveUDPAddr("udp4", addr) if err != nil { return nil, err } - request, err := stun.Build(stun.TransactionID, stun.BindingRequest) + transaction, err := stun.NewBindingTransaction(serverAddr) if err != nil { return nil, err } - - if err = request.NewTransactionID(); err != nil { + if err := c.conn.SetReadDeadline(time.Now().Add(responseTimeout)); err != nil { return nil, err } - if _, err := c.conn.WriteTo(request.Raw, serverAddr); err != nil { - return nil, err - } - - var m stun.Message - select { - case msg := <-c.messageChan: - m.Raw = msg.Body - if err := m.Decode(); err != nil { - return nil, err + response, err := c.client.Do(transaction) + if err != nil { + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return nil, fmt.Errorf("wait response from stun server timeout") } - case <-time.After(responseTimeout): - return nil, fmt.Errorf("wait response from stun server timeout") + return nil, err } - xorAddrGetter := &stun.XORMappedAddress{} - mappedAddrGetter := &stun.MappedAddress{} - changedAddrGetter := ChangedAddress{} - otherAddrGetter := &stun.OtherAddress{} resp := &stunResponse{} - if err := mappedAddrGetter.GetFrom(&m); err == nil { - resp.externalAddr = mappedAddrGetter.String() + if response.MappedAddr != nil { + resp.externalAddr = response.MappedAddr.String() } - if err := xorAddrGetter.GetFrom(&m); err == nil { - resp.externalAddr = xorAddrGetter.String() - } - if err := changedAddrGetter.GetFrom(&m); err == nil { - resp.otherAddr = changedAddrGetter.String() - } - if err := otherAddrGetter.GetFrom(&m); err == nil { - resp.otherAddr = otherAddrGetter.String() + if response.OtherAddr != nil { + resp.otherAddr = response.OtherAddr.String() } return resp, nil } diff --git a/pkg/nathole/discovery_test.go b/pkg/nathole/discovery_test.go new file mode 100644 index 00000000..6596314d --- /dev/null +++ b/pkg/nathole/discovery_test.go @@ -0,0 +1,382 @@ +package nathole + +import ( + "encoding/binary" + "errors" + "fmt" + "net" + "testing" + "time" + + "github.com/fatedier/golib/net/stun" + "github.com/stretchr/testify/require" +) + +const ( + testBindingRequest = 0x0001 + testBindingSuccess = 0x0101 + testBindingError = 0x0111 + testMagicCookie = 0x2112a442 + testAttrMapped = 0x0001 + testAttrChanged = 0x0005 + testAttrErrorCode = 0x0009 + testAttrXORMapped = 0x0020 + testAttrOther = 0x802c + testSTUNHeaderSize = 20 + testSTUNServerLimit = time.Second +) + +type testSTUNAttribute struct { + typ uint16 + value []byte +} + +type testSTUNExchange struct { + source *net.UDPAddr + err error +} + +func listenTestUDP4(t *testing.T) *net.UDPConn { + t.Helper() + + conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + return conn +} + +func serveOneSTUNRequest( + server *net.UDPConn, + buildResponse func([]byte, *net.UDPAddr) ([]byte, error), +) <-chan testSTUNExchange { + done := make(chan testSTUNExchange, 1) + go func() { + if err := server.SetDeadline(time.Now().Add(testSTUNServerLimit)); err != nil { + done <- testSTUNExchange{err: err} + return + } + buffer := make([]byte, 1024) + n, source, err := server.ReadFromUDP(buffer) + if err == nil && buildResponse != nil { + var response []byte + response, err = buildResponse(buffer[:n], source) + if err == nil && response != nil { + _, err = server.WriteToUDP(response, source) + } + } + done <- testSTUNExchange{source: source, err: err} + }() + return done +} + +func waitSTUNExchange(t *testing.T, done <-chan testSTUNExchange) *net.UDPAddr { + t.Helper() + + select { + case exchange := <-done: + require.NoError(t, exchange.err) + return exchange.source + case <-time.After(testSTUNServerLimit): + t.Fatal("timed out waiting for local STUN server") + return nil + } +} + +func makeTestSTUNResponse(request []byte, typ uint16, attributes ...testSTUNAttribute) ([]byte, error) { + if len(request) != testSTUNHeaderSize || binary.BigEndian.Uint16(request[0:2]) != testBindingRequest || + binary.BigEndian.Uint32(request[4:8]) != testMagicCookie { + return nil, fmt.Errorf("invalid Binding request") + } + + length := 0 + for _, attribute := range attributes { + length += 4 + (len(attribute.value)+3)&^3 + } + response := make([]byte, testSTUNHeaderSize, testSTUNHeaderSize+length) + binary.BigEndian.PutUint16(response[0:2], typ) + binary.BigEndian.PutUint16(response[2:4], uint16(length)) + binary.BigEndian.PutUint32(response[4:8], testMagicCookie) + copy(response[8:20], request[8:20]) + + for _, attribute := range attributes { + start := len(response) + paddedLength := (len(attribute.value) + 3) &^ 3 + response = append(response, make([]byte, 4+paddedLength)...) + binary.BigEndian.PutUint16(response[start:start+2], attribute.typ) + binary.BigEndian.PutUint16(response[start+2:start+4], uint16(len(attribute.value))) + copy(response[start+4:], attribute.value) + } + return response, nil +} + +func testIPv4AddressValue(ip net.IP, port int, xor bool) []byte { + value := make([]byte, 8) + value[1] = 0x01 + binary.BigEndian.PutUint16(value[2:4], uint16(port)) + copy(value[4:], ip.To4()) + if xor { + binary.BigEndian.PutUint16(value[2:4], binary.BigEndian.Uint16(value[2:4])^uint16(testMagicCookie>>16)) + for i := range 4 { + value[4+i] ^= byte(uint32(testMagicCookie) >> uint(24-8*i)) + } + } + return value +} + +func TestDiscoverReusesLocalPortAndPreservesNATClassification(t *testing.T) { + tests := []struct { + name string + secondMapped string + secondMappedPort int + wantNATType string + wantBehavior string + }{ + { + name: "same mapped address", + secondMapped: "198.51.100.10:40000", + secondMappedPort: 40000, + wantNATType: EasyNAT, + wantBehavior: BehaviorNoChange, + }, + { + name: "different mapped port", + secondMapped: "198.51.100.10:40001", + secondMappedPort: 40001, + wantNATType: HardNAT, + wantBehavior: BehaviorPortChanged, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + primary := listenTestUDP4(t) + alternate := listenTestUDP4(t) + alternateAddr := alternate.LocalAddr().(*net.UDPAddr) + + primaryDone := serveOneSTUNRequest(primary, func(request []byte, _ *net.UDPAddr) ([]byte, error) { + return makeTestSTUNResponse(request, testBindingSuccess, + testSTUNAttribute{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.10"), 40000, true)}, + testSTUNAttribute{typ: testAttrOther, value: testIPv4AddressValue(alternateAddr.IP, alternateAddr.Port, false)}, + ) + }) + alternateDone := serveOneSTUNRequest(alternate, func(request []byte, _ *net.UDPAddr) ([]byte, error) { + return makeTestSTUNResponse(request, testBindingSuccess, + testSTUNAttribute{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.10"), tt.secondMappedPort, true)}, + ) + }) + + addresses, localAddr, err := Discover([]string{primary.LocalAddr().String()}, "") + require.NoError(t, err) + require.Equal(t, []string{"198.51.100.10:40000", tt.secondMapped}, addresses) + + primarySource := waitSTUNExchange(t, primaryDone) + alternateSource := waitSTUNExchange(t, alternateDone) + require.Equal(t, primarySource.Port, alternateSource.Port) + require.Equal(t, localAddr.(*net.UDPAddr).Port, primarySource.Port) + + feature, err := ClassifyNATFeature(addresses, nil) + require.NoError(t, err) + require.Equal(t, tt.wantNATType, feature.NatType) + require.Equal(t, tt.wantBehavior, feature.Behavior) + }) + } +} + +func TestDoSTUNRequestMapsLegacyAndModernAddresses(t *testing.T) { + tests := []struct { + name string + attributes []testSTUNAttribute + wantExternal string + wantOther string + }{ + { + name: "legacy", + attributes: []testSTUNAttribute{ + {typ: testAttrMapped, value: testIPv4AddressValue(net.ParseIP("192.0.2.1"), 1000, false)}, + {typ: testAttrChanged, value: testIPv4AddressValue(net.ParseIP("192.0.2.2"), 2000, false)}, + }, + wantExternal: "192.0.2.1:1000", + wantOther: "192.0.2.2:2000", + }, + { + name: "modern takes precedence", + attributes: []testSTUNAttribute{ + {typ: testAttrMapped, value: testIPv4AddressValue(net.ParseIP("192.0.2.1"), 1000, false)}, + {typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.1"), 3000, true)}, + {typ: testAttrChanged, value: testIPv4AddressValue(net.ParseIP("192.0.2.2"), 2000, false)}, + {typ: testAttrOther, value: testIPv4AddressValue(net.ParseIP("198.51.100.2"), 4000, false)}, + }, + wantExternal: "198.51.100.1:3000", + wantOther: "198.51.100.2:4000", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := listenTestUDP4(t) + done := serveOneSTUNRequest(server, func(request []byte, _ *net.UDPAddr) ([]byte, error) { + return makeTestSTUNResponse(request, testBindingSuccess, tt.attributes...) + }) + conn, err := listen("") + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + + response, err := conn.doSTUNRequest(server.LocalAddr().String()) + require.NoError(t, err) + require.Equal(t, tt.wantExternal, response.externalAddr) + require.Equal(t, tt.wantOther, response.otherAddr) + waitSTUNExchange(t, done) + }) + } +} + +func TestSTUNResponseErrorsAndMissingAddresses(t *testing.T) { + tests := []struct { + name string + buildResponse func([]byte, *net.UDPAddr) ([]byte, error) + request func(*discoverConn, string) error + checkError func(*testing.T, error) + }{ + { + name: "correlated malformed response", + buildResponse: func(request []byte, _ *net.UDPAddr) ([]byte, error) { + response, err := makeTestSTUNResponse(request, testBindingSuccess) + if err == nil { + binary.BigEndian.PutUint16(response[2:4], 4) + } + return response, err + }, + request: func(conn *discoverConn, server string) error { + _, err := conn.doSTUNRequest(server) + return err + }, + checkError: func(t *testing.T, err error) { + require.ErrorIs(t, err, stun.ErrMalformedResponse) + }, + }, + { + name: "Binding error response", + buildResponse: func(request []byte, _ *net.UDPAddr) ([]byte, error) { + return makeTestSTUNResponse(request, testBindingError, testSTUNAttribute{ + typ: testAttrErrorCode, + value: []byte{0, 0, 4, 20, 'U', 'n', 'k', 'n', 'o', 'w', 'n'}, + }) + }, + request: func(conn *discoverConn, server string) error { + _, err := conn.doSTUNRequest(server) + return err + }, + checkError: func(t *testing.T, err error) { + var responseErr *stun.ResponseError + require.ErrorAs(t, err, &responseErr) + require.Equal(t, 420, responseErr.Code) + }, + }, + { + name: "missing mapped address", + buildResponse: func(request []byte, _ *net.UDPAddr) ([]byte, error) { + return makeTestSTUNResponse(request, testBindingSuccess, + testSTUNAttribute{typ: testAttrOther, value: testIPv4AddressValue(net.ParseIP("192.0.2.2"), 2000, false)}, + ) + }, + request: func(conn *discoverConn, server string) error { + _, err := conn.discoverFromStunServer(server) + return err + }, + checkError: func(t *testing.T, err error) { + require.EqualError(t, err, "no external address found") + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := listenTestUDP4(t) + done := serveOneSTUNRequest(server, tt.buildResponse) + conn, err := listen("") + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + + err = tt.request(conn, server.LocalAddr().String()) + tt.checkError(t, err) + waitSTUNExchange(t, done) + }) + } + + t.Run("missing other address", func(t *testing.T) { + server := listenTestUDP4(t) + done := serveOneSTUNRequest(server, func(request []byte, _ *net.UDPAddr) ([]byte, error) { + return makeTestSTUNResponse(request, testBindingSuccess, + testSTUNAttribute{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.1"), 3000, true)}, + ) + }) + + _, err := Prepare([]string{server.LocalAddr().String()}, PrepareOptions{}) + require.EqualError(t, err, "discover error: not enough addresses") + waitSTUNExchange(t, done) + }) +} + +func TestSTUNTimeoutUsesCallerDeadlineWithoutRetry(t *testing.T) { + originalTimeout := responseTimeout + responseTimeout = 50 * time.Millisecond + t.Cleanup(func() { responseTimeout = originalTimeout }) + + server := listenTestUDP4(t) + done := serveOneSTUNRequest(server, nil) + conn, err := listen("") + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + + _, err = conn.doSTUNRequest(server.LocalAddr().String()) + require.EqualError(t, err, "wait response from stun server timeout") + waitSTUNExchange(t, done) + + require.NoError(t, server.SetReadDeadline(time.Now().Add(50*time.Millisecond))) + _, _, err = server.ReadFromUDP(make([]byte, 1)) + var netErr net.Error + require.ErrorAs(t, err, &netErr) + require.True(t, netErr.Timeout()) +} + +func TestSTUNClientLeavesSocketAndDeadlineWithCaller(t *testing.T) { + originalTimeout := responseTimeout + responseTimeout = 100 * time.Millisecond + t.Cleanup(func() { responseTimeout = originalTimeout }) + + server := listenTestUDP4(t) + unrelated := listenTestUDP4(t) + done := serveOneSTUNRequest(server, func(request []byte, source *net.UDPAddr) ([]byte, error) { + response, err := makeTestSTUNResponse(request, testBindingSuccess, + testSTUNAttribute{typ: testAttrXORMapped, value: testIPv4AddressValue(net.ParseIP("198.51.100.5"), 5000, true)}, + ) + if err != nil { + return nil, err + } + if _, err := unrelated.WriteToUDP(response, source); err != nil { + return nil, err + } + return response, nil + }) + conn, err := listen("") + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + + response, err := conn.doSTUNRequest(server.LocalAddr().String()) + require.NoError(t, err) + require.Equal(t, "198.51.100.5:5000", response.externalAddr) + waitSTUNExchange(t, done) + + _, _, err = conn.conn.ReadFromUDP(make([]byte, 1)) + var netErr net.Error + require.True(t, errors.As(err, &netErr)) + require.True(t, netErr.Timeout()) + + require.NoError(t, conn.conn.SetDeadline(time.Time{})) + require.NoError(t, server.SetReadDeadline(time.Now().Add(testSTUNServerLimit))) + _, err = conn.conn.WriteToUDP([]byte{1}, server.LocalAddr().(*net.UDPAddr)) + require.NoError(t, err) + _, source, err := server.ReadFromUDP(make([]byte, 1)) + require.NoError(t, err) + require.Equal(t, conn.localAddr.(*net.UDPAddr).Port, source.Port) +} diff --git a/pkg/nathole/utils.go b/pkg/nathole/utils.go index 2a882830..0d71692f 100644 --- a/pkg/nathole/utils.go +++ b/pkg/nathole/utils.go @@ -18,10 +18,8 @@ import ( "bytes" "fmt" "net" - "strconv" "github.com/fatedier/golib/crypto" - "github.com/pion/stun/v3" "github.com/fatedier/frp/pkg/msg" ) @@ -48,20 +46,6 @@ func DecodeMessageInto(data, key []byte, m msg.Message) error { return msg.ReadMsgInto(bytes.NewReader(buf), m) } -type ChangedAddress struct { - IP net.IP - Port int -} - -func (s *ChangedAddress) GetFrom(m *stun.Message) error { - a := (*stun.MappedAddress)(s) - return a.GetFromAs(m, stun.AttrChangedAddress) -} - -func (s *ChangedAddress) String() string { - return net.JoinHostPort(s.IP.String(), strconv.Itoa(s.Port)) -} - func ListAllLocalIPs() ([]net.IP, error) { addrs, err := net.InterfaceAddrs() if err != nil { diff --git a/pkg/plugin/client/tls2raw.go b/pkg/plugin/client/tls2raw.go index 29176238..7562383b 100644 --- a/pkg/plugin/client/tls2raw.go +++ b/pkg/plugin/client/tls2raw.go @@ -81,6 +81,15 @@ func (p *TLS2RawPlugin) Handle(ctx context.Context, connInfo *ConnectionInfo) { return } + if connInfo.ProxyProtocolHeader != nil { + if _, err := connInfo.ProxyProtocolHeader.WriteTo(rawConn); err != nil { + xl.Warnf("tls2raw write proxy protocol header to local conn error: %v", err) + rawConn.Close() + tlsConn.Close() + return + } + } + libio.Join(tlsConn, rawConn) } diff --git a/pkg/proto/udp/udp.go b/pkg/proto/udp/udp.go index 70b51900..ec3bd94f 100644 --- a/pkg/proto/udp/udp.go +++ b/pkg/proto/udp/udp.go @@ -63,9 +63,13 @@ 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) - select { - case sendCh <- udpMsg: - default: + if err = errors.PanicToError(func() { + select { + case sendCh <- udpMsg: + default: + } + }); err != nil { + return } } } diff --git a/pkg/proto/udp/udp_test.go b/pkg/proto/udp/udp_test.go index 1a7f0091..b959a26a 100644 --- a/pkg/proto/udp/udp_test.go +++ b/pkg/proto/udp/udp_test.go @@ -1,9 +1,13 @@ package udp import ( + "net" "testing" + "time" "github.com/stretchr/testify/require" + + "github.com/fatedier/frp/pkg/msg" ) func TestUdpPacket(t *testing.T) { @@ -16,3 +20,33 @@ 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") + } +} diff --git a/pkg/proto/wire/crypto.go b/pkg/proto/wire/crypto.go index 69aedc69..38916a28 100644 --- a/pkg/proto/wire/crypto.go +++ b/pkg/proto/wire/crypto.go @@ -68,7 +68,8 @@ func NewServerHello(clientHello ClientHello) (ServerHello, error) { return ServerHello{ Selected: ServerSelection{ Message: MessageSelection{ - Codec: MessageCodecJSON, + Codec: MessageCodecJSON, + UDPPacketCodec: selectUDPPacketCodec(clientHello.Capabilities.Message.UDPPacketCodecs), }, Crypto: CryptoSelection{ Algorithm: algorithm, @@ -92,6 +93,15 @@ 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) @@ -105,6 +115,13 @@ 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, diff --git a/pkg/proto/wire/wire.go b/pkg/proto/wire/wire.go index 47bf5984..5a109d67 100644 --- a/pkg/proto/wire/wire.go +++ b/pkg/proto/wire/wire.go @@ -36,6 +36,7 @@ const ( FrameTypeMessage uint16 = 16 MessageCodecJSON = "json" + UDPPacketCodecBinary = "binary-v1" DefaultMaxFramePayloadSize = 64 * 1024 MagicV2 = "FRP\x00\x02\r\n" @@ -182,7 +183,8 @@ type ClientCapabilities struct { } type MessageCapabilities struct { - Codecs []string `json:"codecs,omitempty"` + Codecs []string `json:"codecs,omitempty"` + UDPPacketCodecs []string `json:"udpPacketCodecs,omitempty"` } type CryptoCapabilities struct { @@ -201,7 +203,8 @@ type ServerSelection struct { } type MessageSelection struct { - Codec string `json:"codec,omitempty"` + Codec string `json:"codec,omitempty"` + UDPPacketCodec string `json:"udpPacketCodec,omitempty"` } type CryptoSelection struct { @@ -214,7 +217,8 @@ func clientHelloWithCryptoRandom(bootstrap BootstrapInfo, clientRandom []byte) C Bootstrap: bootstrap, Capabilities: ClientCapabilities{ Message: MessageCapabilities{ - Codecs: []string{MessageCodecJSON}, + Codecs: []string{MessageCodecJSON}, + UDPPacketCodecs: []string{UDPPacketCodecBinary}, }, Crypto: CryptoCapabilities{ Algorithms: PreferredAEADAlgorithms(), diff --git a/pkg/proto/wire/wire_test.go b/pkg/proto/wire/wire_test.go index b564f712..b1c1b062 100644 --- a/pkg/proto/wire/wire_test.go +++ b/pkg/proto/wire/wire_test.go @@ -148,10 +148,40 @@ 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) diff --git a/pkg/ssh/server.go b/pkg/ssh/server.go index bffc40bc..8a69fc30 100644 --- a/pkg/ssh/server.go +++ b/pkg/ssh/server.go @@ -16,7 +16,6 @@ package ssh import ( "context" - "encoding/binary" "errors" "fmt" "net" @@ -52,6 +51,11 @@ type tcpipForward struct { Port uint32 } +// https://datatracker.ietf.org/doc/html/rfc4254#section-6.5 +type execPayload struct { + Command string +} + // https://datatracker.ietf.org/doc/html/rfc4254#page-16 type forwardedTCPPayload struct { Addr string @@ -66,6 +70,7 @@ type TunnelServer struct { sshConn *ssh.ServerConn sc *ssh.ServerConfig firstChannel ssh.Channel + firstChannelMu sync.Mutex vc *virtual.Client peerServerListener *netpkg.InternalListener @@ -187,6 +192,8 @@ func (s *TunnelServer) Run() error { } func (s *TunnelServer) writeToClient(data string) { + s.firstChannelMu.Lock() + defer s.firstChannelMu.Unlock() if s.firstChannel == nil { return } @@ -300,23 +307,24 @@ func (s *TunnelServer) handleNewChannel(channel ssh.NewChannel, extraPayloadCh c if err != nil { return } + s.firstChannelMu.Lock() if s.firstChannel == nil { s.firstChannel = ch } + s.firstChannelMu.Unlock() go s.keepAlive(ch) for req := range reqs { if req.WantReply { _ = req.Reply(true, nil) } - if req.Type != "exec" || len(req.Payload) <= 4 { + if req.Type != "exec" { continue } - end := 4 + binary.BigEndian.Uint32(req.Payload[:4]) - if len(req.Payload) < int(end) { + extraPayload, ok := parseExecPayload(req.Payload) + if !ok { continue } - extraPayload := string(req.Payload[4:end]) select { case extraPayloadCh <- extraPayload: default: @@ -324,6 +332,14 @@ func (s *TunnelServer) handleNewChannel(channel ssh.NewChannel, extraPayloadCh c } } +func parseExecPayload(payload []byte) (string, bool) { + var msg execPayload + if err := ssh.Unmarshal(payload, &msg); err != nil { + return "", false + } + return msg.Command, true +} + func (s *TunnelServer) keepAlive(ch ssh.Channel) { tk := time.NewTicker(time.Second * 30) defer tk.Stop() diff --git a/pkg/ssh/server_test.go b/pkg/ssh/server_test.go new file mode 100644 index 00000000..56b1ba5b --- /dev/null +++ b/pkg/ssh/server_test.go @@ -0,0 +1,115 @@ +// Copyright 2026 The frp Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package ssh + +import ( + "encoding/binary" + "io" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" + cryptossh "golang.org/x/crypto/ssh" +) + +func TestParseExecPayload(t *testing.T) { + payload := cryptossh.Marshal(&execPayload{Command: "tcp --remote_port 6000"}) + + got, ok := parseExecPayload(payload) + + require.True(t, ok) + require.Equal(t, "tcp --remote_port 6000", got) +} + +func TestParseExecPayloadRejectsMalformedPayloads(t *testing.T) { + overflowLength := make([]byte, 5) + binary.BigEndian.PutUint32(overflowLength[:4], ^uint32(0)) + + for _, tc := range []struct { + name string + payload []byte + }{ + { + name: "empty", + payload: nil, + }, + { + name: "short length prefix", + payload: []byte{0, 0, 0}, + }, + { + name: "declared length exceeds remaining payload", + payload: []byte{0, 0, 0, 2, 'x'}, + }, + { + name: "overflow length", + payload: overflowLength, + }, + } { + t.Run(tc.name, func(t *testing.T) { + var ( + got string + ok bool + ) + require.NotPanics(t, func() { + got, ok = parseExecPayload(tc.payload) + }) + require.False(t, ok) + require.Empty(t, got) + }) + } +} + +type trackingChannel struct { + active atomic.Int32 + concurrent atomic.Bool +} + +func (c *trackingChannel) Read([]byte) (int, error) { return 0, io.EOF } + +func (c *trackingChannel) Write(p []byte) (int, error) { + if c.active.Add(1) != 1 { + c.concurrent.Store(true) + } + time.Sleep(time.Millisecond) + c.active.Add(-1) + return len(p), nil +} + +func (c *trackingChannel) Close() error { return nil } +func (c *trackingChannel) CloseWrite() error { return nil } +func (c *trackingChannel) SendRequest(string, bool, []byte) (bool, error) { return false, nil } +func (c *trackingChannel) Stderr() io.ReadWriter { return nil } + +func TestWriteToClientSerializesChannelWrites(t *testing.T) { + channel := &trackingChannel{} + s := &TunnelServer{firstChannel: channel} + start := make(chan struct{}) + var wg sync.WaitGroup + for range 8 { + wg.Go(func() { + <-start + s.writeToClient("message") + }) + } + close(start) + wg.Wait() + + if channel.concurrent.Load() { + t.Fatal("channel writes were concurrent") + } +} diff --git a/pkg/util/limit/limiter.go b/pkg/util/limit/limiter.go new file mode 100644 index 00000000..43ba7434 --- /dev/null +++ b/pkg/util/limit/limiter.go @@ -0,0 +1,37 @@ +// 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) +} diff --git a/pkg/util/limit/limiter_test.go b/pkg/util/limit/limiter_test.go new file mode 100644 index 00000000..f245bcfc --- /dev/null +++ b/pkg/util/limit/limiter_test.go @@ -0,0 +1,65 @@ +// 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()) + }) + } +} diff --git a/pkg/util/limit/reader.go b/pkg/util/limit/reader.go index efa828f4..eccca5a5 100644 --- a/pkg/util/limit/reader.go +++ b/pkg/util/limit/reader.go @@ -35,6 +35,12 @@ 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] } diff --git a/pkg/util/limit/writer.go b/pkg/util/limit/writer.go index 5256d1e2..56357228 100644 --- a/pkg/util/limit/writer.go +++ b/pkg/util/limit/writer.go @@ -34,8 +34,15 @@ 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 { diff --git a/pkg/util/net/conn.go b/pkg/util/net/conn.go index 5dd605a5..fd71162e 100644 --- a/pkg/util/net/conn.go +++ b/pkg/util/net/conn.go @@ -16,6 +16,7 @@ package net import ( "context" + "crypto/hkdf" "crypto/sha256" "errors" "io" @@ -25,7 +26,6 @@ import ( libcrypto "github.com/fatedier/golib/crypto" quic "github.com/quic-go/quic-go" - "golang.org/x/crypto/hkdf" "github.com/fatedier/frp/pkg/util/xlog" ) @@ -335,11 +335,6 @@ func deriveAEADControlKeys(key []byte, algorithm string, transcriptHash []byte) } func deriveAEADControlKey(key []byte, algorithm string, transcriptHash []byte, direction string) ([]byte, error) { - info := []byte(aeadControlHKDFInfoPrefix + " " + algorithm + " " + direction) - reader := hkdf.New(sha256.New, key, transcriptHash, info) - out := make([]byte, libcrypto.AEADKeySize) - if _, err := io.ReadFull(reader, out); err != nil { - return nil, err - } - return out, nil + info := aeadControlHKDFInfoPrefix + " " + algorithm + " " + direction + return hkdf.Key(sha256.New, key, transcriptHash, info, libcrypto.AEADKeySize) } diff --git a/pkg/util/net/conn_test.go b/pkg/util/net/conn_test.go index 42d3c06b..93387384 100644 --- a/pkg/util/net/conn_test.go +++ b/pkg/util/net/conn_test.go @@ -114,5 +114,11 @@ func TestDeriveAEADControlKeysUsesDistinctDirections(t *testing.T) { bytes.Repeat([]byte{0x44}, 32), ) require.NoError(t, err) + require.Equal(t, []byte{ + 0xa0, 0x58, 0xcd, 0x02, 0x5d, 0x96, 0x98, 0x5f, + 0xeb, 0xeb, 0xff, 0x79, 0xa1, 0x9f, 0x62, 0xb7, + 0x15, 0xe0, 0x53, 0x91, 0x3d, 0xfc, 0x74, 0x77, + 0x05, 0x91, 0x4c, 0x62, 0x4b, 0xf3, 0xd4, 0x95, + }, clientToServerKey) require.NotEqual(t, clientToServerKey, serverToClientKey) } diff --git a/pkg/util/net/dial.go b/pkg/util/net/dial.go index 1a3859ed..b0a5e437 100644 --- a/pkg/util/net/dial.go +++ b/pkg/util/net/dial.go @@ -45,6 +45,11 @@ func DialHookWebsocket(protocol string, host string) libnet.AfterHookFunc { if err != nil { return nil, nil, err } + // The tunnel payload is a raw byte stream (yamux), not UTF-8 text. + // Send it as binary frames; otherwise RFC 6455-compliant intermediaries + // (e.g. API gateways/reverse proxies) UTF-8-validate the default text + // frames and close the connection on invalid bytes. + conn.PayloadType = websocket.BinaryFrame return ctx, conn, nil } } diff --git a/pkg/util/net/websocket.go b/pkg/util/net/websocket.go index 6c2f39c4..57641b31 100644 --- a/pkg/util/net/websocket.go +++ b/pkg/util/net/websocket.go @@ -32,6 +32,11 @@ func NewWebsocketListener(ln net.Listener) (wl *WebsocketListener) { muxer := http.NewServeMux() muxer.Handle(FrpWebsocketPath, websocket.Handler(func(c *websocket.Conn) { + // The tunnel payload is a raw byte stream (yamux), not UTF-8 text. + // Send it as binary frames; otherwise RFC 6455-compliant intermediaries + // (e.g. API gateways/reverse proxies) UTF-8-validate the default text + // frames and close the connection on invalid bytes. + c.PayloadType = websocket.BinaryFrame notifyCh := make(chan struct{}) conn := WrapCloseNotifyConn(c, func(_ error) { close(notifyCh) diff --git a/pkg/util/vhost/http.go b/pkg/util/vhost/http.go index 725662d8..d10c99ad 100644 --- a/pkg/util/vhost/http.go +++ b/pkg/util/vhost/http.go @@ -28,8 +28,6 @@ import ( libio "github.com/fatedier/golib/io" "github.com/fatedier/golib/pool" - "golang.org/x/net/http2" - "golang.org/x/net/http2/h2c" httppkg "github.com/fatedier/frp/pkg/util/http" "github.com/fatedier/frp/pkg/util/log" @@ -144,7 +142,7 @@ func NewHTTPReverseProxy(option HTTPReverseProxyOptions, vhostRouter *Routers) * _, _ = rw.Write(getNotFoundPageContent()) }, } - rp.proxy = h2c.NewHandler(proxy, &http2.Server{}) + rp.proxy = proxy return rp } diff --git a/pkg/util/vhost/http_test.go b/pkg/util/vhost/http_test.go index 237ae903..a6493b08 100644 --- a/pkg/util/vhost/http_test.go +++ b/pkg/util/vhost/http_test.go @@ -1,14 +1,94 @@ package vhost import ( + "bufio" + "fmt" + "net" + "net/http" "net/http/httptest" "testing" + "time" "github.com/stretchr/testify/require" httppkg "github.com/fatedier/frp/pkg/util/http" ) +func TestHTTPServerProtocols(t *testing.T) { + rp := NewHTTPReverseProxy(HTTPReverseProxyOptions{}, NewRouters()) + protocols := new(http.Protocols) + protocols.SetHTTP1(true) + protocols.SetUnencryptedHTTP2(true) + server := &http.Server{ + Handler: rp, + ReadHeaderTimeout: time.Second, + Protocols: protocols, + } + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + serveErr := make(chan error, 1) + go func() { + serveErr <- server.Serve(listener) + }() + defer func() { + require.NoError(t, server.Close()) + require.ErrorIs(t, <-serveErr, http.ErrServerClosed) + }() + + require.True(t, server.Protocols.HTTP1()) + require.True(t, server.Protocols.UnencryptedHTTP2()) + + t.Run("HTTP/1.1", func(t *testing.T) { + transport := &http.Transport{Protocols: httpProtocols(true, false)} + defer transport.CloseIdleConnections() + client := &http.Client{Transport: transport} + response, err := client.Get("http://" + listener.Addr().String() + "/") + require.NoError(t, err) + defer response.Body.Close() + + require.Equal(t, "HTTP/1.1", response.Proto) + require.Equal(t, http.StatusNotFound, response.StatusCode) + }) + + t.Run("HTTP/2 prior knowledge", func(t *testing.T) { + transport := &http.Transport{Protocols: httpProtocols(false, true)} + defer transport.CloseIdleConnections() + client := &http.Client{Transport: transport} + response, err := client.Get("http://" + listener.Addr().String() + "/") + require.NoError(t, err) + defer response.Body.Close() + + require.Equal(t, "HTTP/2.0", response.Proto) + require.Equal(t, http.StatusNotFound, response.StatusCode) + }) + + t.Run("HTTP/1.1 Upgrade h2c", func(t *testing.T) { + conn, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer conn.Close() + + _, err = fmt.Fprintf(conn, + "GET / HTTP/1.1\r\nHost: %s\r\n"+ + "Connection: Upgrade, HTTP2-Settings\r\nUpgrade: h2c\r\n"+ + "HTTP2-Settings: AAMAAABkAAQCAAAAAAIAAAAA\r\n\r\n", + listener.Addr()) + require.NoError(t, err) + response, err := http.ReadResponse(bufio.NewReader(conn), nil) + require.NoError(t, err) + defer response.Body.Close() + + require.NotEqual(t, http.StatusSwitchingProtocols, response.StatusCode) + }) +} + +func httpProtocols(http1, unencryptedHTTP2 bool) *http.Protocols { + protocols := new(http.Protocols) + protocols.SetHTTP1(http1) + protocols.SetUnencryptedHTTP2(unencryptedHTTP2) + return protocols +} + func TestCheckRouteAuthByRequest(t *testing.T) { rc := &RouteConfig{ Username: "alice", diff --git a/server/api_router.go b/server/api_router.go index e0ac31ed..43b7eb92 100644 --- a/server/api_router.go +++ b/server/api_router.go @@ -50,10 +50,15 @@ func (svr *Service) registerRouteHandlers(helper *httppkg.RouterRegisterHelper) subRouter.HandleFunc("/api/proxies", httppkg.MakeHTTPHandlerFunc(apiController.DeleteProxies)).Methods("DELETE") subRouter.HandleFunc("/api/v2/users", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2UserList)).Methods("GET") + subRouter.HandleFunc("/api/v2/system/info", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2SystemInfo)).Methods("GET") + subRouter.HandleFunc("/api/v2/system/prune", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2SystemPrune)).Methods("POST") subRouter.HandleFunc("/api/v2/clients", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ClientList)).Methods("GET") - subRouter.HandleFunc("/api/v2/clients/{key}", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ClientDetail)).Methods("GET") + v2EncodedPathRouter := subRouter.NewRoute().Subrouter() + v2EncodedPathRouter.UseEncodedPath() + v2EncodedPathRouter.HandleFunc("/api/v2/clients/{key}", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ClientDetail)).Methods("GET") subRouter.HandleFunc("/api/v2/proxies", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyList)).Methods("GET") - subRouter.HandleFunc("/api/v2/proxies/{name}", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyDetail)).Methods("GET") + v2EncodedPathRouter.HandleFunc("/api/v2/proxies/{name}/traffic", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyTraffic)).Methods("GET") + v2EncodedPathRouter.HandleFunc("/api/v2/proxies/{name}", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyDetail)).Methods("GET") // view subRouter.Handle("/favicon.ico", http.FileServer(helper.AssetsFS)).Methods("GET") diff --git a/server/control.go b/server/control.go index eb496938..890182b2 100644 --- a/server/control.go +++ b/server/control.go @@ -17,6 +17,8 @@ package server import ( "context" "fmt" + "math" + "net" "runtime/debug" "sync" "sync/atomic" @@ -40,55 +42,313 @@ import ( "github.com/fatedier/frp/server/registry" ) +type ControlID uint64 + +var nextControlID atomic.Uint64 + +const workConnPoolCapacityOffset = 10 + +type controlEntry struct { + ctl *Control + id ControlID + // runMu serializes lifecycle and routing decisions for one run ID. + // Replacements inherit it; removing the entry releases the manager's reference. + runMu *sync.Mutex + + registryOnline bool + registryControlID ControlID +} + type ControlManager struct { // controls indexed by run id - ctlsByRunID map[string]*Control + ctlsByRunID map[string]*controlEntry + registry *registry.ClientRegistry + closed bool mu sync.RWMutex } -func NewControlManager() *ControlManager { +func NewControlManager(clientRegistry *registry.ClientRegistry) *ControlManager { return &ControlManager{ - ctlsByRunID: make(map[string]*Control), + ctlsByRunID: make(map[string]*controlEntry), + registry: clientRegistry, } } -func (cm *ControlManager) Add(runID string, ctl *Control) (old *Control) { - cm.mu.Lock() - defer cm.mu.Unlock() - - var ok bool - old, ok = cm.ctlsByRunID[runID] - if ok { - old.Replaced(ctl) +// lockCurrentRun returns the current entry with its run gate held. It never +// waits for the gate while holding cm.mu and revalidates the gate after waiting. +// The global order is runMu, cm.mu, ctl.lifecycleMu, then registry locks. +func (cm *ControlManager) lockCurrentRun(runID string, allowClosed bool) (*controlEntry, bool) { + cm.mu.RLock() + entry, ok := cm.ctlsByRunID[runID] + if cm.closed && !allowClosed { + ok = false } - cm.ctlsByRunID[runID] = ctl - return + cm.mu.RUnlock() + if !ok { + return nil, false + } + + runMu := entry.runMu + runMu.Lock() + cm.mu.RLock() + entry, ok = cm.ctlsByRunID[runID] + if (cm.closed && !allowClosed) || !ok || entry.runMu != runMu { + ok = false + } + cm.mu.RUnlock() + if !ok { + runMu.Unlock() + return nil, false + } + return entry, true } -// we should make sure if it's the same control to prevent delete a new one -func (cm *ControlManager) Del(runID string, ctl *Control) { +// Add makes ctl the pending current generation and records the predecessor +// finalization barrier it must wait for before activation. +func (cm *ControlManager) Add(ctl *Control) error { + for { + // Never wait for a run gate while holding cm.mu. + cm.mu.RLock() + old := cm.ctlsByRunID[ctl.runID] + cm.mu.RUnlock() + if old != nil { + old.runMu.Lock() + } + + cm.mu.Lock() + if cm.closed { + cm.mu.Unlock() + if old != nil { + old.runMu.Unlock() + } + return fmt.Errorf("control manager is closed") + } + if cm.ctlsByRunID[ctl.runID] != old { + cm.mu.Unlock() + if old != nil { + old.runMu.Unlock() + } + continue + } + + id := ControlID(nextControlID.Add(1)) + if err := ctl.admit(cm, id); err != nil { + cm.mu.Unlock() + if old != nil { + old.runMu.Unlock() + } + return err + } + + runMu := &sync.Mutex{} + if old != nil { + runMu = old.runMu + } + entry := &controlEntry{ctl: ctl, id: id, runMu: runMu} + var ( + oldCtl *Control + barrier <-chan struct{} + ) + if old != nil { + oldCtl = old.ctl + barrier = oldCtl.markReplaced() + ctl.setHandoffBarrier(barrier) + entry.registryOnline = old.registryOnline + entry.registryControlID = old.registryControlID + } + cm.ctlsByRunID[ctl.runID] = entry + cm.mu.Unlock() + if old != nil { + old.runMu.Unlock() + } + + if oldCtl != nil { + oldCtl.Replaced(ctl) + } + return nil + } +} + +// Activate registers ctl as online only if it is still the pending current +// generation. +func (cm *ControlManager) Activate(ctl *Control) (bool, error) { + entry, ok := cm.lockCurrentRun(ctl.runID, false) + if !ok { + return false, nil + } + defer entry.runMu.Unlock() cm.mu.Lock() defer cm.mu.Unlock() - if c, ok := cm.ctlsByRunID[runID]; ok && c == ctl { - delete(cm.ctlsByRunID, runID) + + if cm.closed || cm.ctlsByRunID[ctl.runID] != entry || entry.ctl != ctl || entry.id != ctl.controlID { + return false, nil } + + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + if ctl.state != controlStatePending { + return false, nil + } + if ctl.activated { + return true, nil + } + + loginMsg := ctl.sessionCtx.LoginMsg + remoteAddr := ctl.sessionCtx.Conn.RemoteAddr().String() + if host, _, err := net.SplitHostPort(remoteAddr); err == nil { + remoteAddr = host + } + _, conflict := cm.registry.RegisterWithControlID( + loginMsg.User, + loginMsg.ClientID, + ctl.runID, + loginMsg.Hostname, + loginMsg.Version, + remoteAddr, + ctl.sessionCtx.WireProtocol, + uint64(entry.id), + ) + if conflict { + return true, fmt.Errorf("client_id [%s] for user [%s] is already online", loginMsg.ClientID, loginMsg.User) + } + + entry.registryOnline = true + entry.registryControlID = entry.id + ctl.activated = true + return true, nil +} + +// completeLogin reserves ctl's current ownership with its run gate while the +// bounded successful LoginResp write runs, then transitions it to running. +// The callback must only perform that bounded write; it must not call back into +// the control manager or the same control lifecycle. +func (cm *ControlManager) completeLogin(ctl *Control, writeSuccess func() error) (bool, error) { + entry, ok := cm.lockCurrentRun(ctl.runID, false) + if !ok { + return false, nil + } + defer entry.runMu.Unlock() + if entry.ctl != ctl || entry.id != ctl.controlID { + return false, nil + } + + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + if ctl.state != controlStatePending || !ctl.activated { + return false, nil + } + if err := writeSuccess(); err != nil { + return false, err + } + if !ctl.startLocked() { + return false, nil + } + return true, nil +} + +// Remove deletes and offlines ctl only if it is still the current generation. +func (cm *ControlManager) Remove(ctl *Control) bool { + entry, ok := cm.lockCurrentRun(ctl.runID, true) + if !ok { + return false + } + defer entry.runMu.Unlock() + cm.mu.Lock() + defer cm.mu.Unlock() + + if cm.ctlsByRunID[ctl.runID] != entry || entry.ctl != ctl || entry.id != ctl.controlID { + return false + } + delete(cm.ctlsByRunID, ctl.runID) + if entry.registryOnline { + cm.registry.MarkOfflineByRunIDAndControlID(ctl.runID, uint64(entry.registryControlID)) + } + return true } func (cm *ControlManager) GetByID(runID string) (ctl *Control, ok bool) { - cm.mu.RLock() - defer cm.mu.RUnlock() - ctl, ok = cm.ctlsByRunID[runID] - return + entry, ok := cm.lockCurrentRun(runID, false) + if !ok { + return nil, false + } + defer entry.runMu.Unlock() + ctl = entry.ctl + + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + if ctl.state != controlStateRunning { + return nil, false + } + return ctl, true +} + +// admitVisitorByRunID commits a visitor admission against the current running +// control while its run and lifecycle ownership are held. The callback must +// only perform the in-memory, buffered visitor admission. +func (cm *ControlManager) admitVisitorByRunID(runID string, admit func(user, wireProtocol, udpPacketCodec string) error) (bool, error) { + entry, ok := cm.lockCurrentRun(runID, false) + if !ok { + return false, nil + } + defer entry.runMu.Unlock() + ctl := entry.ctl + + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + if ctl.state != controlStateRunning { + return false, nil + } + return true, admit(ctl.sessionCtx.LoginMsg.User, ctl.sessionCtx.WireProtocol, ctl.sessionCtx.UDPPacketCodec) +} + +// RegisterWorkConn transfers conn to ctl only if ctl is still the current +// running generation. On error, ownership remains with the caller. +func (cm *ControlManager) RegisterWorkConn(ctl *Control, conn *proxy.WorkConn) error { + entry, ok := cm.lockCurrentRun(ctl.runID, false) + if !ok { + cm.mu.RLock() + closed := cm.closed + cm.mu.RUnlock() + if closed { + return fmt.Errorf("control manager is closed") + } + return fmt.Errorf("client control for run id [%s] is no longer current", ctl.runID) + } + defer entry.runMu.Unlock() + if entry.ctl != ctl || entry.id != ctl.controlID { + return fmt.Errorf("client control for run id [%s] is no longer current", ctl.runID) + } + + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + if ctl.state != controlStateRunning { + return fmt.Errorf("client control for run id [%s] is not running", ctl.runID) + } + + select { + case ctl.workConnCh <- conn: + ctl.xl.Debugf("new work connection registered") + return nil + default: + ctl.xl.Debugf("work connection pool is full, discarding") + return fmt.Errorf("work connection pool is full, discarding") + } } func (cm *ControlManager) Close() error { cm.mu.Lock() - defer cm.mu.Unlock() - for _, ctl := range cm.ctlsByRunID { - ctl.Close() + cm.closed = true + ctls := make([]*Control, 0, len(cm.ctlsByRunID)) + for _, entry := range cm.ctlsByRunID { + ctls = append(ctls, entry.ctl) + } + cm.mu.Unlock() + + for _, ctl := range ctls { + cm.Remove(ctl) + _ = ctl.Close() } - cm.ctlsByRunID = make(map[string]*Control) return nil } @@ -96,7 +356,8 @@ func (cm *ControlManager) Close() error { func (cm *ControlManager) CloseAllProxyByName(proxyName string) error { cm.mu.RLock() var target *Control - for _, ctl := range cm.ctlsByRunID { + for _, entry := range cm.ctlsByRunID { + ctl := entry.ctl ctl.mu.RLock() _, ok := ctl.proxies[proxyName] ctl.mu.RUnlock() @@ -117,7 +378,8 @@ func (cm *ControlManager) CloseAllProxyByName(proxyName string) error { func (cm *ControlManager) KickByProxyName(proxyName string) error { cm.mu.RLock() var target *Control - for _, ctl := range cm.ctlsByRunID { + for _, entry := range cm.ctlsByRunID { + ctl := entry.ctl ctl.mu.RLock() _, ok := ctl.proxies[proxyName] ctl.mu.RUnlock() @@ -155,12 +417,21 @@ type SessionContext struct { LoginMsg *msg.Login // server configuration ServerCfg *v1.ServerConfig - // client registry - ClientRegistry *registry.ClientRegistry // negotiated wire protocol for this client session - WireProtocol string + WireProtocol string + UDPPacketCodec string } +type controlState uint8 + +const ( + controlStateCreated controlState = iota + controlStatePending + controlStateRunning + controlStateClosing + controlStateClosed +) + type Control struct { // session context sessionCtx *SessionContext @@ -187,30 +458,59 @@ type Control struct { // last time got the Ping message lastPing atomic.Value - // A new run id will be generated when a new client login. - // If run id got from login message has same run id, it means it's the same client, so we can - // replace old controller instantly. - runID string + // runID never changes during the lifetime of a control. controlID is assigned + // once by ControlManager and distinguishes same-runID generations. + runID string + controlID ControlID + manager *ControlManager + + lifecycleMu sync.Mutex + state controlState + activated bool + handoffBarrier <-chan struct{} + + interruptOnce sync.Once + interruptErr error mu sync.RWMutex - xl *xlog.Logger - ctx context.Context - doneCh chan struct{} + xl *xlog.Logger + ctx context.Context + doneCh chan struct{} + serverMetrics metrics.ServerMetrics } func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, error) { - poolCount := min(sessionCtx.LoginMsg.PoolCount, int(sessionCtx.ServerCfg.Transport.MaxPoolCount)) + 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) ctl := &Control{ - sessionCtx: sessionCtx, - workConnCh: make(chan *proxy.WorkConn, poolCount+10), - proxies: make(map[string]proxy.Proxy), - poolCount: poolCount, - portsUsedNum: 0, - runID: sessionCtx.LoginMsg.RunID, - xl: xlog.FromContextSafe(ctx), - ctx: ctx, - doneCh: make(chan struct{}), + sessionCtx: sessionCtx, + workConnCh: make(chan *proxy.WorkConn, poolCount+workConnPoolCapacityOffset), + proxies: make(map[string]proxy.Proxy), + poolCount: poolCount, + portsUsedNum: 0, + runID: sessionCtx.LoginMsg.RunID, + state: controlStateCreated, + xl: xlog.FromContextSafe(ctx), + ctx: ctx, + doneCh: make(chan struct{}), + serverMetrics: metrics.Server, } ctl.lastPing.Store(time.Now()) @@ -220,48 +520,121 @@ func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, erro return ctl, nil } -// Start starts the control session workers after login succeeds. -func (ctl *Control) Start() { - go func() { - for i := 0; i < ctl.poolCount; i++ { - // ignore error here, that means that this control is closed - _ = ctl.msgDispatcher.Send(&msg.ReqWorkConn{}) - } - }() - go ctl.worker() +func (ctl *Control) RunID() string { + return ctl.runID } -func (ctl *Control) Close() error { - ctl.sessionCtx.Conn.Close() +func (ctl *Control) ID() ControlID { + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + return ctl.controlID +} + +func (ctl *Control) admit(manager *ControlManager, id ControlID) error { + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + if ctl.state != controlStateCreated { + return fmt.Errorf("control [%s] is not in created state", ctl.runID) + } + ctl.manager = manager + ctl.controlID = id + ctl.state = controlStatePending return nil } -func (ctl *Control) Replaced(newCtl *Control) { - xl := ctl.xl - xl.Infof("replaced by client [%s]", newCtl.runID) - ctl.runID = "" - ctl.sessionCtx.Conn.Close() +func (ctl *Control) setHandoffBarrier(barrier <-chan struct{}) { + ctl.lifecycleMu.Lock() + ctl.handoffBarrier = barrier + ctl.lifecycleMu.Unlock() } -func (ctl *Control) RegisterWorkConn(conn *proxy.WorkConn) error { - xl := ctl.xl - defer func() { - if err := recover(); err != nil { - xl.Errorf("panic error: %v", err) - xl.Errorf(string(debug.Stack())) - } - }() - - select { - case ctl.workConnCh <- conn: - xl.Debugf("new work connection registered") - return nil - default: - xl.Debugf("work connection pool is full, discarding") - return fmt.Errorf("work connection pool is full, discarding") +func (ctl *Control) WaitForHandoff() { + ctl.lifecycleMu.Lock() + barrier := ctl.handoffBarrier + ctl.lifecycleMu.Unlock() + if barrier != nil { + <-barrier } } +// Start starts the control session workers after login succeeds. +func (ctl *Control) Start() bool { + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + return ctl.startLocked() +} + +func (ctl *Control) startLocked() bool { + if ctl.state != controlStatePending || !ctl.activated { + return false + } + ctl.state = controlStateRunning + go ctl.worker() + return true +} + +func (ctl *Control) Close() error { + ctl.lifecycleMu.Lock() + switch ctl.state { + case controlStateCreated, controlStatePending: + ctl.state = controlStateClosing + ctl.finishLocked() + case controlStateRunning: + ctl.state = controlStateClosing + } + ctl.lifecycleMu.Unlock() + return ctl.interruptReadAndClose() +} + +func (ctl *Control) Replaced(newCtl *Control) { + ctl.markReplaced() + ctl.xl.Infof("replaced by client [%s] (control ID %d)", newCtl.runID, newCtl.ID()) + _ = ctl.interruptReadAndClose() +} + +// markReplaced returns the transitive predecessor barrier. A pending control +// has no worker, so it finishes immediately and passes its inherited barrier +// to the replacement. A running control is finished only by its worker. +func (ctl *Control) markReplaced() <-chan struct{} { + ctl.lifecycleMu.Lock() + defer ctl.lifecycleMu.Unlock() + + switch ctl.state { + case controlStateCreated: + ctl.state = controlStateClosing + ctl.finishLocked() + return nil + case controlStatePending: + barrier := ctl.handoffBarrier + ctl.state = controlStateClosing + ctl.finishLocked() + return barrier + case controlStateRunning: + ctl.state = controlStateClosing + return ctl.doneCh + case controlStateClosing, controlStateClosed: + return ctl.doneCh + default: + return ctl.doneCh + } +} + +func (ctl *Control) interruptReadAndClose() error { + ctl.interruptOnce.Do(func() { + _ = ctl.sessionCtx.Conn.SetReadDeadline(time.Now()) + ctl.interruptErr = ctl.sessionCtx.Conn.Close() + }) + return ctl.interruptErr +} + +func (ctl *Control) finishLocked() { + if ctl.state == controlStateClosed { + return + } + ctl.state = controlStateClosed + close(ctl.doneCh) +} + // When frps get one user connection, we get one work connection from the pool and return it. // If no workConn available in the pool, send message to frpc to get one or more // and wait until it is available. @@ -316,10 +689,10 @@ func (ctl *Control) heartbeatWorker() { } xl := ctl.xl - go wait.Until(func() { + wait.Until(func() { if time.Since(ctl.lastPing.Load().(time.Time)) > time.Duration(ctl.sessionCtx.ServerCfg.Transport.HeartbeatTimeout)*time.Second { xl.Warnf("heartbeat timeout") - ctl.sessionCtx.Conn.Close() + _ = ctl.Close() return } }, time.Second, ctl.doneCh) @@ -334,14 +707,14 @@ func (ctl *Control) loginUserInfo() plugin.UserInfo { return plugin.UserInfo{ User: ctl.sessionCtx.LoginMsg.User, Metas: ctl.sessionCtx.LoginMsg.Metas, - RunID: ctl.sessionCtx.LoginMsg.RunID, + RunID: ctl.runID, } } func (ctl *Control) closeProxy(pxy proxy.Proxy) { pxy.Close() ctl.sessionCtx.PxyManager.Del(pxy.GetName()) - metrics.Server.CloseProxy(pxy.GetName(), pxy.GetConfigurer().GetBaseConfig().Type) + ctl.serverMetrics.CloseProxy(pxy.GetName(), pxy.GetConfigurer().GetBaseConfig().Type) notifyContent := &plugin.CloseProxyContent{ User: ctl.loginUserInfo(), @@ -356,12 +729,24 @@ func (ctl *Control) closeProxy(pxy proxy.Proxy) { func (ctl *Control) worker() { xl := ctl.xl + ctl.serverMetrics.NewClient() go ctl.heartbeatWorker() go ctl.msgDispatcher.Run() + go func() { + for i := 0; i < ctl.poolCount; i++ { + // Ignore the error: it means this control is already closing. + _ = ctl.msgDispatcher.Send(&msg.ReqWorkConn{}) + } + }() <-ctl.msgDispatcher.Done() - ctl.sessionCtx.Conn.Close() + ctl.lifecycleMu.Lock() + if ctl.state == controlStateRunning { + ctl.state = controlStateClosing + } + ctl.lifecycleMu.Unlock() + _ = ctl.interruptReadAndClose() ctl.mu.Lock() close(ctl.workConnCh) @@ -376,10 +761,14 @@ func (ctl *Control) worker() { ctl.closeProxy(pxy) } - metrics.Server.CloseClient() - ctl.sessionCtx.ClientRegistry.MarkOfflineByRunID(ctl.runID) + ctl.serverMetrics.CloseClient() + if ctl.manager != nil { + ctl.manager.Remove(ctl) + } xl.Infof("client exit success") - close(ctl.doneCh) + ctl.lifecycleMu.Lock() + ctl.finishLocked() + ctl.lifecycleMu.Unlock() } func (ctl *Control) registerMsgHandlers() { @@ -419,9 +808,9 @@ func (ctl *Control) handleNewProxy(m msg.Message) { xl.Infof("new proxy [%s] type [%s] success", inMsg.ProxyName, inMsg.ProxyType) clientID := ctl.sessionCtx.LoginMsg.ClientID if clientID == "" { - clientID = ctl.sessionCtx.LoginMsg.RunID + clientID = ctl.runID } - metrics.Server.NewProxy(inMsg.ProxyName, inMsg.ProxyType, ctl.sessionCtx.LoginMsg.User, clientID) + ctl.serverMetrics.NewProxy(inMsg.ProxyName, inMsg.ProxyType, ctl.sessionCtx.LoginMsg.User, clientID) } _ = ctl.msgDispatcher.Send(resp) } @@ -500,6 +889,7 @@ 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 diff --git a/server/control_test.go b/server/control_test.go new file mode 100644 index 00000000..965ea5ea --- /dev/null +++ b/server/control_test.go @@ -0,0 +1,595 @@ +// Copyright 2026 The frp Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package server + +import ( + "context" + "errors" + "math" + "net" + "os" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/fatedier/frp/pkg/auth" + v1 "github.com/fatedier/frp/pkg/config/v1" + "github.com/fatedier/frp/pkg/msg" + plugin "github.com/fatedier/frp/pkg/plugin/server" + "github.com/fatedier/frp/server/controller" + "github.com/fatedier/frp/server/proxy" + "github.com/fatedier/frp/server/registry" +) + +func TestControlPendingReplacementFinishesWithoutStarting(t *testing.T) { + clientRegistry := registry.NewClientRegistry() + manager := NewControlManager(clientRegistry) + metrics := newCountingServerMetrics() + oldCtl, oldConn := newLifecycleTestControl(t, "same-run", "client", metrics) + newCtl, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + + mustAddAndActivate(t, manager, oldCtl) + + err := manager.Add(newCtl) + require.NoError(t, err) + waitForControlDone(t, oldCtl) + require.False(t, oldCtl.Start()) + require.Equal(t, []string{"deadline", "close"}, oldConn.eventsSnapshot()) + require.Equal(t, int64(0), metrics.newClients()) + require.Equal(t, int64(0), metrics.closedClients()) +} + +func TestNewControlPoolCountBoundaries(t *testing.T) { + for _, tc := range []struct { + name string + poolCount int + maxPoolCount int64 + wantErr string + wantPoolCount int + wantCapacity int + }{ + {name: "negative pool count below offset", poolCount: -11, maxPoolCount: 5, wantErr: "invalid pool count"}, + {name: "negative pool count at offset", poolCount: -10, maxPoolCount: 5, wantErr: "invalid pool count"}, + {name: "negative pool count", poolCount: -1, maxPoolCount: 5, wantErr: "invalid pool count"}, + {name: "zero pool count", poolCount: 0, maxPoolCount: 5, wantPoolCount: 0, wantCapacity: 10}, + {name: "pool count capped", poolCount: 10, maxPoolCount: 5, wantPoolCount: 5, wantCapacity: 15}, + {name: "maximum int pool count capped", poolCount: math.MaxInt, maxPoolCount: 5, wantPoolCount: 5, wantCapacity: 15}, + {name: "negative maximum", poolCount: 1, maxPoolCount: -1, wantErr: "invalid max pool count"}, + {name: "maximum int64 with small client pool", poolCount: 1, maxPoolCount: math.MaxInt64, wantPoolCount: 1, wantCapacity: 11}, + {name: "maximum int client and server overflow", poolCount: math.MaxInt, maxPoolCount: math.MaxInt64, wantErr: "cannot safely add"}, + } { + t.Run(tc.name, func(t *testing.T) { + conn := newDeadlineReadConn() + msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn)) + cfg := &v1.ServerConfig{} + cfg.Transport.MaxPoolCount = tc.maxPoolCount + + ctl, err := NewControl(context.Background(), &SessionContext{ + RC: &controller.ResourceController{}, + PxyManager: proxy.NewManager(), + PluginManager: plugin.NewManager(), + AuthVerifier: auth.AlwaysPassVerifier, + Conn: msgConn, + LoginMsg: &msg.Login{ + RunID: "pool-count-run", + PoolCount: tc.poolCount, + }, + ServerCfg: cfg, + }) + if tc.wantErr != "" { + require.Nil(t, ctl) + require.ErrorContains(t, err, tc.wantErr) + return + } + + require.NoError(t, err) + require.Equal(t, tc.wantPoolCount, ctl.poolCount) + require.Equal(t, tc.wantCapacity, cap(ctl.workConnCh)) + require.NoError(t, ctl.Close()) + }) + } +} + +func TestControlRunningReplacementFinishesInWorker(t *testing.T) { + clientRegistry := registry.NewClientRegistry() + manager := NewControlManager(clientRegistry) + metrics := newCountingServerMetrics() + oldCtl, oldConn := newLifecycleTestControl(t, "same-run", "client", metrics) + newCtl, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + + mustAddAndActivate(t, manager, oldCtl) + require.True(t, oldCtl.Start()) + waitForSignal(t, oldConn.readStarted, "control reader to start") + + err := manager.Add(newCtl) + require.NoError(t, err) + waitForControlDone(t, oldCtl) + require.Equal(t, []string{"deadline", "close"}, oldConn.eventsSnapshot()) + require.Equal(t, int64(1), metrics.newClients()) + require.Equal(t, int64(1), metrics.closedClients()) + + _, ok := manager.GetByID("same-run") + require.False(t, ok) + require.Same(t, newCtl, currentControlForTest(manager, "same-run")) + info, ok := clientRegistry.GetByKey("client") + require.True(t, ok) + require.True(t, info.Online) + require.Equal(t, uint64(oldCtl.ID()), info.ControlID) + + active, err := manager.Activate(newCtl) + require.NoError(t, err) + require.True(t, active) + _, ok = manager.GetByID("same-run") + require.False(t, ok) + info, ok = clientRegistry.GetByKey("client") + require.True(t, ok) + require.Equal(t, uint64(newCtl.ID()), info.ControlID) +} + +func TestControlClosePendingAndRunning(t *testing.T) { + t.Run("pending", func(t *testing.T) { + manager := NewControlManager(registry.NewClientRegistry()) + metrics := newCountingServerMetrics() + ctl, conn := newLifecycleTestControl(t, "pending", "pending", metrics) + err := manager.Add(ctl) + require.NoError(t, err) + + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) + require.Equal(t, []string{"deadline", "close"}, conn.eventsSnapshot()) + require.Equal(t, int64(0), metrics.newClients()) + require.Equal(t, int64(0), metrics.closedClients()) + }) + + t.Run("running", func(t *testing.T) { + manager := NewControlManager(registry.NewClientRegistry()) + metrics := newCountingServerMetrics() + ctl, conn := newLifecycleTestControl(t, "running", "running", metrics) + mustAddAndActivate(t, manager, ctl) + require.True(t, ctl.Start()) + waitForSignal(t, conn.readStarted, "control reader to start") + + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) + require.Equal(t, []string{"deadline", "close"}, conn.eventsSnapshot()) + require.Equal(t, int64(1), metrics.newClients()) + require.Equal(t, int64(1), metrics.closedClients()) + }) +} + +func TestControlCloseAndReplacedAreIdempotent(t *testing.T) { + manager := NewControlManager(registry.NewClientRegistry()) + metrics := newCountingServerMetrics() + ctl, conn := newLifecycleTestControl(t, "same-run", "client", metrics) + replacement, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + + err := manager.Add(ctl) + require.NoError(t, err) + err = manager.Add(replacement) + require.NoError(t, err) + require.NoError(t, ctl.Close()) + ctl.Replaced(replacement) + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) + + require.Equal(t, []string{"deadline", "close"}, conn.eventsSnapshot()) + require.Equal(t, int64(0), metrics.newClients()) + require.Equal(t, int64(0), metrics.closedClients()) +} + +func TestControlHeartbeatTimeoutInterruptsRead(t *testing.T) { + manager := NewControlManager(registry.NewClientRegistry()) + metrics := newCountingServerMetrics() + ctl, conn := newLifecycleTestControl(t, "heartbeat", "heartbeat", metrics) + ctl.sessionCtx.ServerCfg.Transport.HeartbeatTimeout = 1 + ctl.lastPing.Store(time.Now().Add(-2 * time.Second)) + + mustAddAndActivate(t, manager, ctl) + require.True(t, ctl.Start()) + waitForSignal(t, conn.readStarted, "control reader to start") + waitForControlDone(t, ctl) + + require.Equal(t, []string{"deadline", "close"}, conn.eventsSnapshot()) + require.Equal(t, int64(1), metrics.newClients()) + require.Equal(t, int64(1), metrics.closedClients()) +} + +func TestControlStartReplacementRacePairsMetrics(t *testing.T) { + for range 100 { + clientRegistry := registry.NewClientRegistry() + manager := NewControlManager(clientRegistry) + metrics := newCountingServerMetrics() + ctl, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + replacement, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + + mustAddAndActivate(t, manager, ctl) + + startGate := make(chan struct{}) + startedCh := make(chan bool, 1) + addErrCh := make(chan error, 1) + go func() { + <-startGate + startedCh <- ctl.Start() + }() + go func() { + <-startGate + addErr := manager.Add(replacement) + addErrCh <- addErr + }() + close(startGate) + + started := <-startedCh + require.NoError(t, <-addErrCh) + waitForControlDone(t, ctl) + if started { + require.Equal(t, int64(1), metrics.newClients()) + require.Equal(t, int64(1), metrics.closedClients()) + } else { + require.Equal(t, int64(0), metrics.newClients()) + require.Equal(t, int64(0), metrics.closedClients()) + } + } +} + +func TestControlManagerRejectsStaleActivateAndRemove(t *testing.T) { + clientRegistry := registry.NewClientRegistry() + manager := NewControlManager(clientRegistry) + metrics := newCountingServerMetrics() + oldCtl, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + newCtl, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + + mustAddAndActivate(t, manager, oldCtl) + err := manager.Add(newCtl) + require.NoError(t, err) + require.Greater(t, uint64(newCtl.ID()), uint64(oldCtl.ID())) + + active, err := manager.Activate(oldCtl) + require.NoError(t, err) + require.False(t, active) + require.False(t, manager.Remove(oldCtl)) + + _, ok := manager.GetByID("same-run") + require.False(t, ok) + require.Same(t, newCtl, currentControlForTest(manager, "same-run")) + info, ok := clientRegistry.GetByKey("client") + require.True(t, ok) + require.True(t, info.Online) + require.Equal(t, uint64(oldCtl.ID()), info.ControlID) + + active, err = manager.Activate(newCtl) + require.NoError(t, err) + require.True(t, active) + info, ok = clientRegistry.GetByKey("client") + require.True(t, ok) + require.True(t, info.Online) + require.Equal(t, uint64(newCtl.ID()), info.ControlID) +} + +func TestControlManagerPreservesClientIDConflict(t *testing.T) { + clientRegistry := registry.NewClientRegistry() + manager := NewControlManager(clientRegistry) + metrics := newCountingServerMetrics() + first, _ := newLifecycleTestControl(t, "run-one", "shared-client", metrics) + conflicting, _ := newLifecycleTestControl(t, "run-two", "shared-client", metrics) + + mustAddAndActivate(t, manager, first) + err := manager.Add(conflicting) + require.NoError(t, err) + active, err := manager.Activate(conflicting) + require.True(t, active) + require.ErrorContains(t, err, "already online") + + require.True(t, manager.Remove(conflicting)) + info, ok := clientRegistry.GetByKey("shared-client") + require.True(t, ok) + require.True(t, info.Online) + require.Equal(t, "run-one", info.RunID) +} + +func TestControlManagerFailedLoginWriteReleasesRunWithoutStarting(t *testing.T) { + clientRegistry := registry.NewClientRegistry() + manager := NewControlManager(clientRegistry) + metrics := newCountingServerMetrics() + ctl, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + replacement, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + + mustAddAndActivate(t, manager, ctl) + + writeErr := errors.New("write failed") + committed, err := manager.completeLogin(ctl, func() error { return writeErr }) + require.ErrorIs(t, err, writeErr) + require.False(t, committed) + + err = manager.Add(replacement) + require.NoError(t, err) + waitForControlDone(t, ctl) + require.Same(t, replacement, currentControlForTest(manager, "same-run")) + require.Equal(t, int64(0), metrics.newClients()) + require.Equal(t, int64(0), metrics.closedClients()) + require.True(t, manager.Remove(replacement)) + info, ok := clientRegistry.GetByKey("client") + require.True(t, ok) + require.False(t, info.Online) + require.Empty(t, info.RunID) + require.Zero(t, info.ControlID) + require.False(t, info.DisconnectedAt.IsZero()) + require.NoError(t, replacement.Close()) +} + +func TestControlManagerCloseWaitsForInFlightLoginRun(t *testing.T) { + clientRegistry := registry.NewClientRegistry() + manager := NewControlManager(clientRegistry) + metrics := newCountingServerMetrics() + ctl, _ := newLifecycleTestControl(t, "same-run", "client", metrics) + + mustAddAndActivate(t, manager, ctl) + + writeEntered := make(chan struct{}) + resumeWrite := make(chan struct{}) + loginDone := make(chan struct { + committed bool + err error + }, 1) + go func() { + committed, loginErr := manager.completeLogin(ctl, func() error { + close(writeEntered) + <-resumeWrite + return nil + }) + loginDone <- struct { + committed bool + err error + }{committed: committed, err: loginErr} + }() + waitForSignal(t, writeEntered, "LoginResp write") + + closeDone := make(chan error, 1) + go func() { closeDone <- manager.Close() }() + waitForManagerClosed(t, manager) + select { + case err := <-closeDone: + t.Fatalf("manager close completed during LoginResp write: %v", err) + default: + } + + close(resumeWrite) + result := <-loginDone + require.NoError(t, result.err) + require.True(t, result.committed) + require.NoError(t, <-closeDone) + waitForControlDone(t, ctl) + require.Nil(t, currentControlForTest(manager, "same-run")) + require.Equal(t, int64(1), metrics.newClients()) + require.Equal(t, int64(1), metrics.closedClients()) + info, ok := clientRegistry.GetByKey("client") + require.True(t, ok) + require.False(t, info.Online) +} + +func newLifecycleTestControl( + t *testing.T, + runID string, + clientID string, + serverMetrics *countingServerMetrics, +) (*Control, *deadlineReadConn) { + t.Helper() + conn := newDeadlineReadConn() + msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn)) + ctl, err := NewControl(context.Background(), &SessionContext{ + RC: &controller.ResourceController{}, + PxyManager: proxy.NewManager(), + PluginManager: plugin.NewManager(), + AuthVerifier: auth.AlwaysPassVerifier, + Conn: msgConn, + LoginMsg: &msg.Login{ + RunID: runID, + ClientID: clientID, + }, + ServerCfg: &v1.ServerConfig{}, + }) + require.NoError(t, err) + ctl.serverMetrics = serverMetrics + t.Cleanup(func() { _ = ctl.Close() }) + return ctl, conn +} + +func mustAddAndActivate(t *testing.T, manager *ControlManager, ctl *Control) { + t.Helper() + require.NoError(t, manager.Add(ctl)) + active, err := manager.Activate(ctl) + require.NoError(t, err) + require.True(t, active) +} + +func waitForControlDone(t *testing.T, ctl *Control) { + t.Helper() + done := make(chan struct{}) + go func() { + ctl.WaitClosed() + close(done) + }() + waitForSignal(t, done, "control to finish") +} + +func currentControlForTest(manager *ControlManager, runID string) *Control { + manager.mu.RLock() + defer manager.mu.RUnlock() + entry := manager.ctlsByRunID[runID] + if entry == nil { + return nil + } + return entry.ctl +} + +func currentRunGateForTest(manager *ControlManager, runID string) *sync.Mutex { + manager.mu.RLock() + defer manager.mu.RUnlock() + entry := manager.ctlsByRunID[runID] + if entry == nil { + return nil + } + return entry.runMu +} + +func waitForManagerClosed(t *testing.T, manager *ControlManager) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + manager.mu.RLock() + closed := manager.closed + manager.mu.RUnlock() + if closed { + return + } + } + t.Fatal("timed out waiting for control manager to close") +} + +func waitForSignal(t *testing.T, ch <-chan struct{}, description string) { + t.Helper() + select { + case <-ch: + case <-time.After(3 * time.Second): + t.Fatalf("timed out waiting for %s", description) + } +} + +type deadlineReadConn struct { + readStarted chan struct{} + unblockRead chan struct{} + + readOnce sync.Once + unblockOnce sync.Once + deadlineOnce sync.Once + closeOnce sync.Once + + eventsMu sync.Mutex + events []string +} + +func newDeadlineReadConn() *deadlineReadConn { + return &deadlineReadConn{ + readStarted: make(chan struct{}), + unblockRead: make(chan struct{}), + } +} + +func (c *deadlineReadConn) Read([]byte) (int, error) { + c.readOnce.Do(func() { close(c.readStarted) }) + <-c.unblockRead + return 0, os.ErrDeadlineExceeded +} + +func (*deadlineReadConn) Write(p []byte) (int, error) { return len(p), nil } + +func (c *deadlineReadConn) Close() error { + c.closeOnce.Do(func() { + c.recordEvent("close") + c.unblockOnce.Do(func() { close(c.unblockRead) }) + }) + return nil +} + +func (*deadlineReadConn) LocalAddr() net.Addr { return lifecycleTestAddr("local") } +func (*deadlineReadConn) RemoteAddr() net.Addr { return lifecycleTestAddr("remote") } + +func (c *deadlineReadConn) SetDeadline(deadline time.Time) error { + if err := c.SetReadDeadline(deadline); err != nil { + return err + } + return c.SetWriteDeadline(deadline) +} + +func (c *deadlineReadConn) SetReadDeadline(deadline time.Time) error { + if deadline.IsZero() { + return nil + } + c.deadlineOnce.Do(func() { + c.recordEvent("deadline") + c.unblockOnce.Do(func() { close(c.unblockRead) }) + }) + return nil +} + +func (*deadlineReadConn) SetWriteDeadline(time.Time) error { return nil } + +func (c *deadlineReadConn) recordEvent(event string) { + c.eventsMu.Lock() + c.events = append(c.events, event) + c.eventsMu.Unlock() +} + +func (c *deadlineReadConn) eventsSnapshot() []string { + c.eventsMu.Lock() + defer c.eventsMu.Unlock() + return append([]string(nil), c.events...) +} + +type lifecycleTestAddr string + +func (a lifecycleTestAddr) Network() string { return string(a) } +func (a lifecycleTestAddr) String() string { return string(a) } + +type countingServerMetrics struct { + mu sync.Mutex + newCount int64 + closeCount int64 + closeEnter chan struct{} + closeResume chan struct{} + closeOnce sync.Once +} + +func newCountingServerMetrics() *countingServerMetrics { + return &countingServerMetrics{} +} + +func (m *countingServerMetrics) NewClient() { + m.mu.Lock() + m.newCount++ + m.mu.Unlock() +} + +func (m *countingServerMetrics) CloseClient() { + m.mu.Lock() + m.closeCount++ + closeEnter := m.closeEnter + closeResume := m.closeResume + m.mu.Unlock() + if closeEnter != nil { + m.closeOnce.Do(func() { close(closeEnter) }) + <-closeResume + } +} + +func (*countingServerMetrics) NewProxy(string, string, string, string) {} +func (*countingServerMetrics) CloseProxy(string, string) {} +func (*countingServerMetrics) OpenConnection(string, string) {} +func (*countingServerMetrics) CloseConnection(string, string) {} +func (*countingServerMetrics) AddTrafficIn(string, string, int64) {} +func (*countingServerMetrics) AddTrafficOut(string, string, int64) {} + +func (m *countingServerMetrics) newClients() int64 { + m.mu.Lock() + defer m.mu.Unlock() + return m.newCount +} + +func (m *countingServerMetrics) closedClients() int64 { + m.mu.Lock() + defer m.mu.Unlock() + return m.closeCount +} diff --git a/server/http/controller.go b/server/http/controller.go index 1a8db7dc..caf43755 100644 --- a/server/http/controller.go +++ b/server/http/controller.go @@ -65,8 +65,12 @@ func NewController( // /api/serverinfo func (c *Controller) APIServerInfo(ctx *httppkg.Context) (any, error) { + return c.buildServerInfoResp(), nil +} + +func (c *Controller) buildServerInfoResp() model.ServerInfoResp { serverStats := mem.StatsCollector.GetServer() - svrResp := model.ServerInfoResp{ + return model.ServerInfoResp{ Version: version.Full(), BindPort: c.serverCfg.BindPort, VhostHTTPPort: c.serverCfg.VhostHTTPPort, @@ -87,8 +91,6 @@ func (c *Controller) APIServerInfo(ctx *httppkg.Context) (any, error) { ClientCounts: serverStats.ClientCounts, ProxyTypeCounts: serverStats.ProxyTypeCounts, } - - return svrResp, nil } // /api/clients diff --git a/server/http/controller_v2.go b/server/http/controller_v2.go index 8a50a54f..dfa7b495 100644 --- a/server/http/controller_v2.go +++ b/server/http/controller_v2.go @@ -17,22 +17,31 @@ package http import ( "cmp" "fmt" + "maps" "math" "net/http" + "net/url" "slices" "strconv" "strings" + "time" v1 "github.com/fatedier/frp/pkg/config/v1" "github.com/fatedier/frp/pkg/metrics/mem" httppkg "github.com/fatedier/frp/pkg/util/http" "github.com/fatedier/frp/server/http/model" + "github.com/fatedier/frp/server/registry" ) const ( defaultV2Page = 1 defaultV2PageSize = 50 maxV2PageSize = 200 + + v2SystemPruneTypeOfflineProxies = "offline_proxies" + v2ProxyTrafficDefaultDays = 7 + v2ProxyTrafficUnit = "bytes" + v2ProxyTrafficGranularity = "day" ) var apiV2ProxyTypes = []string{ @@ -46,6 +55,55 @@ var apiV2ProxyTypes = []string{ string(v1.ProxyTypeSUDP), } +// /api/v2/system/info +func (c *Controller) APIV2SystemInfo(ctx *httppkg.Context) (any, error) { + info := c.buildServerInfoResp() + proxyTypeCounts := info.ProxyTypeCounts + if proxyTypeCounts == nil { + proxyTypeCounts = map[string]int64{} + } + + return model.V2SystemInfoResp{ + Version: info.Version, + Config: model.V2SystemInfoConfigResp{ + BindPort: info.BindPort, + VhostHTTPPort: info.VhostHTTPPort, + VhostHTTPSPort: info.VhostHTTPSPort, + TCPMuxHTTPConnectPort: info.TCPMuxHTTPConnectPort, + KCPBindPort: info.KCPBindPort, + QUICBindPort: info.QUICBindPort, + SubdomainHost: info.SubdomainHost, + MaxPoolCount: info.MaxPoolCount, + MaxPortsPerClient: info.MaxPortsPerClient, + HeartbeatTimeout: info.HeartBeatTimeout, + AllowPortsStr: info.AllowPortsStr, + TLSForce: info.TLSForce, + }, + Status: model.V2SystemInfoStatusResp{ + TotalTrafficIn: info.TotalTrafficIn, + TotalTrafficOut: info.TotalTrafficOut, + CurConns: info.CurConns, + ClientCounts: info.ClientCounts, + ProxyTypeCounts: proxyTypeCounts, + }, + }, nil +} + +// /api/v2/system/prune +func (c *Controller) APIV2SystemPrune(ctx *httppkg.Context) (any, error) { + pruneType, err := parseV2SystemPruneType(ctx.Query("type")) + if err != nil { + return nil, err + } + + cleared, total := mem.StatsCollector.PruneOfflineProxies() + return model.V2SystemPruneResp{ + Type: pruneType, + Cleared: cleared, + Total: total, + }, nil +} + // /api/v2/users func (c *Controller) APIV2UserList(ctx *httppkg.Context) (any, error) { page, pageSize, err := parseV2PageParams(ctx) @@ -137,7 +195,26 @@ func (c *Controller) APIV2ClientList(ctx *httppkg.Context) (any, error) { // /api/v2/clients/{key} func (c *Controller) APIV2ClientDetail(ctx *httppkg.Context) (any, error) { - return c.APIClientDetail(ctx) + key, err := decodeV2PathParam(ctx, "key", "client key") + if err != nil { + return nil, err + } + + if c.clientRegistry == nil { + return nil, fmt.Errorf("client registry unavailable") + } + + info, ok := c.clientRegistry.GetByKey(key) + if !ok { + return nil, httppkg.NewError(http.StatusNotFound, fmt.Sprintf("client %s not found", key)) + } + + resp := buildClientInfoResp(info) + status := c.buildV2ClientStatus(info) + return model.V2ClientDetailResp{ + ClientInfoResp: resp, + Status: status, + }, nil } // /api/v2/proxies @@ -179,7 +256,7 @@ func (c *Controller) APIV2ProxyList(ctx *httppkg.Context) (any, error) { } slices.SortFunc(items, func(a, b model.V2ProxyResp) int { - if v := cmp.Compare(a.Type, b.Type); v != 0 { + if v := cmp.Compare(a.Spec.Type, b.Spec.Type); v != 0 { return v } return cmp.Compare(a.Name, b.Name) @@ -190,9 +267,9 @@ func (c *Controller) APIV2ProxyList(ctx *httppkg.Context) (any, error) { // /api/v2/proxies/{name} func (c *Controller) APIV2ProxyDetail(ctx *httppkg.Context) (any, error) { - name := ctx.Param("name") - if name == "" { - return nil, fmt.Errorf("missing proxy name") + name, err := decodeV2PathParam(ctx, "name", "proxy name") + if err != nil { + return nil, err } ps := mem.StatsCollector.GetProxyByName(name) @@ -202,6 +279,33 @@ func (c *Controller) APIV2ProxyDetail(ctx *httppkg.Context) (any, error) { return c.buildV2ProxyResp(ps), nil } +// /api/v2/proxies/{name}/traffic +func (c *Controller) APIV2ProxyTraffic(ctx *httppkg.Context) (any, error) { + name, err := decodeV2PathParam(ctx, "name", "proxy name") + if err != nil { + return nil, err + } + + proxyTrafficInfo := mem.StatsCollector.GetProxyTraffic(name) + if proxyTrafficInfo == nil { + return nil, httppkg.NewError(http.StatusNotFound, "no proxy info found") + } + + return buildV2ProxyTrafficResp(name, proxyTrafficInfo, time.Now()), nil +} + +func decodeV2PathParam(ctx *httppkg.Context, key string, label string) (string, error) { + raw := ctx.Param(key) + if raw == "" { + return "", fmt.Errorf("missing %s", label) + } + decoded, err := url.PathUnescape(raw) + if err != nil { + return "", httppkg.NewError(http.StatusBadRequest, fmt.Sprintf("invalid %s", label)) + } + return decoded, nil +} + func getOrCreateV2User(items map[string]*model.V2UserResp, user string) *model.V2UserResp { item, ok := items[user] if !ok { @@ -261,6 +365,18 @@ func parseV2ProxyTypeFilter(raw string) (string, error) { return "", httppkg.NewError(http.StatusBadRequest, "type must be one of tcp, udp, http, https, tcpmux, stcp, xtcp, sudp") } +func parseV2SystemPruneType(raw string) (string, error) { + pruneType := strings.ToLower(raw) + switch pruneType { + case "": + return "", httppkg.NewError(http.StatusBadRequest, "type is required") + case v2SystemPruneTypeOfflineProxies: + return pruneType, nil + default: + return "", httppkg.NewError(http.StatusBadRequest, "type must be one of offline_proxies") + } +} + func matchV2StatusFilter(online bool, filter string) bool { switch filter { case "", "all": @@ -320,26 +436,36 @@ func matchV2ClientQuery(item model.ClientInfoResp, q string) bool { func matchV2ProxyQuery(item model.V2ProxyResp, q string) bool { values := []string{ item.Name, - item.Type, + item.Spec.Type, item.User, item.ClientID, item.Status.State, } - switch spec := item.Spec.(type) { - case *model.TCPOutConf: - values = append(values, strconv.Itoa(spec.RemotePort)) - case *model.UDPOutConf: - values = append(values, strconv.Itoa(spec.RemotePort)) - case *model.HTTPOutConf: - values = append(values, spec.CustomDomains...) - values = append(values, spec.SubDomain) - case *model.HTTPSOutConf: - values = append(values, spec.CustomDomains...) - values = append(values, spec.SubDomain) - case *model.TCPMuxOutConf: - values = append(values, spec.CustomDomains...) - values = append(values, spec.SubDomain) + switch item.Spec.Type { + case string(v1.ProxyTypeTCP): + if item.Spec.TCP != nil && item.Spec.TCP.RemotePort != nil { + values = append(values, strconv.Itoa(*item.Spec.TCP.RemotePort)) + } + case string(v1.ProxyTypeUDP): + if item.Spec.UDP != nil && item.Spec.UDP.RemotePort != nil { + values = append(values, strconv.Itoa(*item.Spec.UDP.RemotePort)) + } + case string(v1.ProxyTypeHTTP): + if item.Spec.HTTP != nil { + values = append(values, item.Spec.HTTP.CustomDomains...) + values = append(values, item.Spec.HTTP.Subdomain) + } + case string(v1.ProxyTypeHTTPS): + if item.Spec.HTTPS != nil { + values = append(values, item.Spec.HTTPS.CustomDomains...) + values = append(values, item.Spec.HTTPS.Subdomain) + } + case string(v1.ProxyTypeTCPMUX): + if item.Spec.TCPMux != nil { + values = append(values, item.Spec.TCPMux.CustomDomains...) + values = append(values, item.Spec.TCPMux.Subdomain) + } } return containsV2Query(q, values...) @@ -366,29 +492,156 @@ func (c *Controller) listV2ProxyStats(proxyType string) []*mem.ProxyStats { return items } +func buildV2ProxyTrafficResp(name string, traffic *mem.ProxyTrafficInfo, now time.Time) model.V2ProxyTrafficResp { + history := make([]model.V2ProxyTrafficPointResp, 0, v2ProxyTrafficDefaultDays) + for age := v2ProxyTrafficDefaultDays - 1; age >= 0; age-- { + history = append(history, model.V2ProxyTrafficPointResp{ + Date: now.AddDate(0, 0, -age).Format(time.DateOnly), + TrafficIn: v2TrafficValueAt(traffic.TrafficIn, age), + TrafficOut: v2TrafficValueAt(traffic.TrafficOut, age), + }) + } + + return model.V2ProxyTrafficResp{ + Name: name, + Unit: v2ProxyTrafficUnit, + Granularity: v2ProxyTrafficGranularity, + History: history, + } +} + +func v2TrafficValueAt(values []int64, todayFirstIndex int) int64 { + if todayFirstIndex >= len(values) { + return 0 + } + return values[todayFirstIndex] +} + +func (c *Controller) buildV2ClientStatus(info registry.ClientInfo) model.V2ClientStatusResp { + status := model.V2ClientStatusResp{State: "offline"} + if info.Online { + status.State = "online" + } + + user := info.User + clientID := info.ClientID() + for _, ps := range c.listV2ProxyStats("") { + if ps.User != user || ps.ClientID != clientID { + continue + } + status.CurConns += ps.CurConns + status.ProxyCount++ + } + return status +} + func (c *Controller) buildV2ProxyResp(ps *mem.ProxyStats) model.V2ProxyResp { state := "offline" - var spec any + var cfg v1.ProxyConfigurer if c.pxyManager != nil { if pxy, ok := c.pxyManager.GetByName(ps.Name); ok { state = "online" - spec = getConfFromConfigurer(pxy.GetConfigurer()) + cfg = pxy.GetConfigurer() } } return model.V2ProxyResp{ Name: ps.Name, - Type: ps.Type, User: ps.User, ClientID: ps.ClientID, - Spec: spec, + Spec: buildV2ProxySpec(ps.Type, cfg), Status: model.V2ProxyStatusResp{ State: state, TodayTrafficIn: ps.TodayTrafficIn, TodayTrafficOut: ps.TodayTrafficOut, CurConns: ps.CurConns, - LastStartTime: ps.LastStartTime, - LastCloseTime: ps.LastCloseTime, + LastStartAt: ps.LastStartAt, + LastCloseAt: ps.LastCloseAt, + }, + } +} + +func buildV2ProxySpec(proxyType string, cfg v1.ProxyConfigurer) model.V2ProxySpec { + spec := model.V2ProxySpec{Type: proxyType} + + switch proxyType { + case string(v1.ProxyTypeTCP): + block := &model.V2TCPProxySpec{} + if c, ok := cfg.(*v1.TCPProxyConfig); ok { + block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig()) + block.RemotePort = &c.RemotePort + } + spec.TCP = block + case string(v1.ProxyTypeUDP): + block := &model.V2UDPProxySpec{} + if c, ok := cfg.(*v1.UDPProxyConfig); ok { + block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig()) + block.RemotePort = &c.RemotePort + } + spec.UDP = block + case string(v1.ProxyTypeHTTP): + block := &model.V2HTTPProxySpec{} + if c, ok := cfg.(*v1.HTTPProxyConfig); ok { + block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig()) + block.CustomDomains = slices.Clone(c.CustomDomains) + block.Subdomain = c.SubDomain + block.Locations = slices.Clone(c.Locations) + block.HostHeaderRewrite = c.HostHeaderRewrite + } + spec.HTTP = block + case string(v1.ProxyTypeHTTPS): + block := &model.V2HTTPSProxySpec{} + if c, ok := cfg.(*v1.HTTPSProxyConfig); ok { + block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig()) + block.CustomDomains = slices.Clone(c.CustomDomains) + block.Subdomain = c.SubDomain + } + spec.HTTPS = block + case string(v1.ProxyTypeTCPMUX): + block := &model.V2TCPMuxProxySpec{} + if c, ok := cfg.(*v1.TCPMuxProxyConfig); ok { + block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig()) + block.CustomDomains = slices.Clone(c.CustomDomains) + block.Subdomain = c.SubDomain + block.Multiplexer = c.Multiplexer + block.RouteByHTTPUser = c.RouteByHTTPUser + } + spec.TCPMux = block + case string(v1.ProxyTypeSTCP): + block := &model.V2STCPProxySpec{} + if c, ok := cfg.(*v1.STCPProxyConfig); ok { + block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig()) + } + spec.STCP = block + case string(v1.ProxyTypeSUDP): + block := &model.V2SUDPProxySpec{} + if c, ok := cfg.(*v1.SUDPProxyConfig); ok { + block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig()) + } + spec.SUDP = block + case string(v1.ProxyTypeXTCP): + block := &model.V2XTCPProxySpec{} + if c, ok := cfg.(*v1.XTCPProxyConfig); ok { + block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig()) + } + spec.XTCP = block + } + + return spec +} + +func buildV2ProxyBaseSpec(base *v1.ProxyBaseConfig) model.V2ProxyBaseSpec { + return model.V2ProxyBaseSpec{ + Annotations: maps.Clone(base.Annotations), + Metadatas: maps.Clone(base.Metadatas), + Transport: &model.V2ProxyTransportSpec{ + UseEncryption: base.Transport.UseEncryption, + UseCompression: base.Transport.UseCompression, + BandwidthLimit: base.Transport.BandwidthLimit.String(), + BandwidthLimitMode: base.Transport.BandwidthLimitMode, + }, + LoadBalancer: &model.V2ProxyLoadBalancerSpec{ + Group: base.LoadBalancer.Group, }, } } diff --git a/server/http/controller_v2_proxy_spec_test.go b/server/http/controller_v2_proxy_spec_test.go new file mode 100644 index 00000000..82229bd7 --- /dev/null +++ b/server/http/controller_v2_proxy_spec_test.go @@ -0,0 +1,393 @@ +// Copyright 2026 The frp Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package http + +import ( + "encoding/json" + "strings" + "testing" + + configtypes "github.com/fatedier/frp/pkg/config/types" + v1 "github.com/fatedier/frp/pkg/config/v1" + "github.com/fatedier/frp/pkg/metrics/mem" + "github.com/fatedier/frp/server/http/model" +) + +func TestBuildV2ProxySpecAllTypesAndRedaction(t *testing.T) { + tests := []struct { + proxyType string + cfg v1.ProxyConfigurer + blockKeys []string + }{ + { + proxyType: "tcp", + cfg: &v1.TCPProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "tcp"), + RemotePort: 6000, + }, + blockKeys: []string{"annotations", "loadBalancer", "metadatas", "remotePort", "transport"}, + }, + { + proxyType: "udp", + cfg: &v1.UDPProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "udp"), + RemotePort: 7000, + }, + blockKeys: []string{"annotations", "loadBalancer", "metadatas", "remotePort", "transport"}, + }, + { + proxyType: "http", + cfg: &v1.HTTPProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "http"), + DomainConfig: v1.DomainConfig{CustomDomains: []string{"app.example.com"}, SubDomain: "app"}, + Locations: []string{"/api"}, + HTTPUser: "secret-http-user", + HTTPPassword: "secret-http-password", + HostHeaderRewrite: "backend.example.com", + RequestHeaders: v1.HeaderOperations{Set: map[string]string{"X-Secret": "secret-request-header"}}, + ResponseHeaders: v1.HeaderOperations{Set: map[string]string{"X-Secret": "secret-response-header"}}, + RouteByHTTPUser: "secret-http-route-user", + }, + blockKeys: []string{"annotations", "customDomains", "hostHeaderRewrite", "loadBalancer", "locations", "metadatas", "subdomain", "transport"}, + }, + { + proxyType: "https", + cfg: &v1.HTTPSProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "https"), + DomainConfig: v1.DomainConfig{CustomDomains: []string{"secure.example.com"}, SubDomain: "secure"}, + }, + blockKeys: []string{"annotations", "customDomains", "loadBalancer", "metadatas", "subdomain", "transport"}, + }, + { + proxyType: "tcpmux", + cfg: &v1.TCPMuxProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "tcpmux"), + DomainConfig: v1.DomainConfig{CustomDomains: []string{"mux.example.com"}, SubDomain: "mux"}, + HTTPUser: strings.Join([]string{"secret", "mux-http-user"}, "-"), + HTTPPassword: strings.Join([]string{"secret", "mux-http-password"}, "-"), + RouteByHTTPUser: "displayed-mux-user", + Multiplexer: "httpconnect", + }, + blockKeys: []string{"annotations", "customDomains", "loadBalancer", "metadatas", "multiplexer", "routeByHTTPUser", "subdomain", "transport"}, + }, + { + proxyType: "stcp", + cfg: &v1.STCPProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "stcp"), + Secretkey: strings.Join([]string{"secret", "stcp-key"}, "-"), + AllowUsers: []string{strings.Join([]string{"secret", "stcp-user"}, "-")}, + }, + blockKeys: []string{"annotations", "loadBalancer", "metadatas", "transport"}, + }, + { + proxyType: "sudp", + cfg: &v1.SUDPProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "sudp"), + Secretkey: strings.Join([]string{"secret", "sudp-key"}, "-"), + AllowUsers: []string{strings.Join([]string{"secret", "sudp-user"}, "-")}, + }, + blockKeys: []string{"annotations", "loadBalancer", "metadatas", "transport"}, + }, + { + proxyType: "xtcp", + cfg: &v1.XTCPProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "xtcp"), + Secretkey: strings.Join([]string{"secret", "xtcp-key"}, "-"), + AllowUsers: []string{strings.Join([]string{"secret", "xtcp-user"}, "-")}, + }, + blockKeys: []string{"annotations", "loadBalancer", "metadatas", "transport"}, + }, + } + + for _, tt := range tests { + t.Run(tt.proxyType, func(t *testing.T) { + spec := buildV2ProxySpec(tt.proxyType, tt.cfg) + raw := mustMarshalJSON(t, spec) + + var specObject map[string]json.RawMessage + if err := json.Unmarshal(raw, &specObject); err != nil { + t.Fatalf("unmarshal spec failed: %v", err) + } + assertRawJSONKeys(t, specObject, tt.proxyType, "type") + + var gotType string + if err := json.Unmarshal(specObject["type"], &gotType); err != nil { + t.Fatalf("unmarshal spec type failed: %v", err) + } + if gotType != tt.proxyType { + t.Fatalf("spec type mismatch, want %q got %q", tt.proxyType, gotType) + } + + var block map[string]json.RawMessage + if err := json.Unmarshal(specObject[tt.proxyType], &block); err != nil { + t.Fatalf("unmarshal active block failed: %v", err) + } + assertRawJSONKeys(t, block, tt.blockKeys...) + assertV2ProxyCommonSpec(t, block) + assertV2ProxyTypeFields(t, tt.proxyType, specObject[tt.proxyType]) + assertNoV2ProxySensitiveFields(t, block) + + content := string(raw) + for _, secret := range []string{ + "secret-proxy-name", + "secret-group-key", + "secret-local-host", + "secret-plugin-user", + "secret-plugin-password", + "secret-health-path", + "secret-http-user", + "secret-http-password", + "secret-request-header", + "secret-response-header", + "secret-http-route-user", + "secret-mux-http-user", + "secret-mux-http-password", + "secret-stcp-key", + "secret-stcp-user", + "secret-sudp-key", + "secret-sudp-user", + "secret-xtcp-key", + "secret-xtcp-user", + } { + if strings.Contains(content, secret) { + t.Fatalf("sensitive value %q leaked in spec: %s", secret, content) + } + } + }) + } +} + +func assertV2ProxyTypeFields(t *testing.T, proxyType string, raw json.RawMessage) { + t.Helper() + + switch proxyType { + case "tcp": + var block model.V2TCPProxySpec + if err := json.Unmarshal(raw, &block); err != nil { + t.Fatalf("unmarshal tcp block failed: %v", err) + } + if block.RemotePort == nil || *block.RemotePort != 6000 { + t.Fatalf("tcp remote port mismatch: %#v", block.RemotePort) + } + case "udp": + var block model.V2UDPProxySpec + if err := json.Unmarshal(raw, &block); err != nil { + t.Fatalf("unmarshal udp block failed: %v", err) + } + if block.RemotePort == nil || *block.RemotePort != 7000 { + t.Fatalf("udp remote port mismatch: %#v", block.RemotePort) + } + case "http": + var block model.V2HTTPProxySpec + if err := json.Unmarshal(raw, &block); err != nil { + t.Fatalf("unmarshal http block failed: %v", err) + } + if len(block.CustomDomains) != 1 || block.CustomDomains[0] != "app.example.com" || + block.Subdomain != "app" || len(block.Locations) != 1 || block.Locations[0] != "/api" || + block.HostHeaderRewrite != "backend.example.com" { + t.Fatalf("http fields mismatch: %#v", block) + } + case "https": + var block model.V2HTTPSProxySpec + if err := json.Unmarshal(raw, &block); err != nil { + t.Fatalf("unmarshal https block failed: %v", err) + } + if len(block.CustomDomains) != 1 || block.CustomDomains[0] != "secure.example.com" || block.Subdomain != "secure" { + t.Fatalf("https fields mismatch: %#v", block) + } + case "tcpmux": + var block model.V2TCPMuxProxySpec + if err := json.Unmarshal(raw, &block); err != nil { + t.Fatalf("unmarshal tcpmux block failed: %v", err) + } + if len(block.CustomDomains) != 1 || block.CustomDomains[0] != "mux.example.com" || + block.Subdomain != "mux" || block.Multiplexer != "httpconnect" || block.RouteByHTTPUser != "displayed-mux-user" { + t.Fatalf("tcpmux fields mismatch: %#v", block) + } + } +} + +func TestBuildV2ProxyRespOfflineTypedShells(t *testing.T) { + for _, proxyType := range apiV2ProxyTypes { + t.Run(proxyType, func(t *testing.T) { + resp := (&Controller{}).buildV2ProxyResp(&mem.ProxyStats{ + Name: "offline-" + proxyType, + Type: proxyType, + }) + if resp.Status.State != "offline" { + t.Fatalf("offline phase mismatch: %#v", resp.Status) + } + + var specObject map[string]json.RawMessage + if err := json.Unmarshal(mustMarshalJSON(t, resp.Spec), &specObject); err != nil { + t.Fatalf("unmarshal offline spec failed: %v", err) + } + assertRawJSONKeys(t, specObject, proxyType, "type") + assertRawJSONKeysFromMessage(t, specObject[proxyType]) + }) + } +} + +func TestBuildV2ProxySpecDoesNotPopulateMismatchedBlock(t *testing.T) { + spec := buildV2ProxySpec("tcp", &v1.UDPProxyConfig{ + ProxyBaseConfig: newV2ProxyTestBaseConfig(t, "udp"), + RemotePort: 7000, + }) + + var specObject map[string]json.RawMessage + if err := json.Unmarshal(mustMarshalJSON(t, spec), &specObject); err != nil { + t.Fatalf("unmarshal mismatched spec failed: %v", err) + } + assertRawJSONKeys(t, specObject, "tcp", "type") + assertRawJSONKeysFromMessage(t, specObject["tcp"]) +} + +func newV2ProxyTestBaseConfig(t *testing.T, proxyType string) v1.ProxyBaseConfig { + t.Helper() + + bandwidthLimit, err := configtypes.NewBandwidthQuantity("10MB") + if err != nil { + t.Fatalf("create bandwidth limit failed: %v", err) + } + enabled := false + return v1.ProxyBaseConfig{ + Name: "secret-proxy-name", + Type: proxyType, + Enabled: &enabled, + Annotations: map[string]string{"annotation-key": "annotation-value"}, + Metadatas: map[string]string{"metadata-key": "metadata-value"}, + Transport: v1.ProxyTransport{ + UseEncryption: true, + UseCompression: true, + BandwidthLimit: bandwidthLimit, + BandwidthLimitMode: configtypes.BandwidthLimitModeServer, + ProxyProtocolVersion: "v2", + }, + LoadBalancer: v1.LoadBalancerConfig{ + Group: "public-group", + GroupKey: "secret-group-key", + }, + HealthCheck: v1.HealthCheckConfig{ + Type: "http", + Path: "secret-health-path", + }, + ProxyBackend: v1.ProxyBackend{ + LocalIP: "secret-local-host", + LocalPort: 8080, + Plugin: v1.TypedClientPluginOptions{ + Type: v1.PluginHTTPProxy, + ClientPluginOptions: &v1.HTTPProxyPluginOptions{ + Type: v1.PluginHTTPProxy, + HTTPUser: "secret-plugin-user", + HTTPPassword: "secret-plugin-password", + }, + }, + }, + } +} + +func assertV2ProxyCommonSpec(t *testing.T, block map[string]json.RawMessage) { + t.Helper() + + var annotations map[string]string + if err := json.Unmarshal(block["annotations"], &annotations); err != nil { + t.Fatalf("unmarshal annotations failed: %v", err) + } + if annotations["annotation-key"] != "annotation-value" { + t.Fatalf("annotations mismatch: %#v", annotations) + } + + var metadatas map[string]string + if err := json.Unmarshal(block["metadatas"], &metadatas); err != nil { + t.Fatalf("unmarshal metadatas failed: %v", err) + } + if metadatas["metadata-key"] != "metadata-value" { + t.Fatalf("metadatas mismatch: %#v", metadatas) + } + + assertRawJSONKeysFromMessage(t, block["transport"], + "bandwidthLimit", + "bandwidthLimitMode", + "useCompression", + "useEncryption", + ) + var transport model.V2ProxyTransportSpec + if err := json.Unmarshal(block["transport"], &transport); err != nil { + t.Fatalf("unmarshal transport failed: %v", err) + } + if !transport.UseEncryption || !transport.UseCompression || + transport.BandwidthLimit != "10MB" || transport.BandwidthLimitMode != "server" { + t.Fatalf("transport mismatch: %#v", transport) + } + + assertRawJSONKeysFromMessage(t, block["loadBalancer"], "group") + var loadBalancer model.V2ProxyLoadBalancerSpec + if err := json.Unmarshal(block["loadBalancer"], &loadBalancer); err != nil { + t.Fatalf("unmarshal load balancer failed: %v", err) + } + if loadBalancer.Group != "public-group" { + t.Fatalf("load balancer mismatch: %#v", loadBalancer) + } +} + +func assertNoV2ProxySensitiveFields(t *testing.T, value any) { + t.Helper() + + forbidden := map[string]struct{}{ + "allowUsers": {}, + "enabled": {}, + "groupKey": {}, + "healthCheck": {}, + "httpPassword": {}, + "httpUser": {}, + "localIP": {}, + "localPort": {}, + "name": {}, + "natTraversal": {}, + "plugin": {}, + "proxyProtocolVersion": {}, + "requestHeaders": {}, + "responseHeaders": {}, + "secretKey": {}, + "type": {}, + } + + var walk func(any) + walk = func(current any) { + switch current := current.(type) { + case map[string]any: + for key, nested := range current { + if _, ok := forbidden[key]; ok { + t.Fatalf("sensitive field %q leaked in active block", key) + } + walk(nested) + } + case []any: + for _, nested := range current { + walk(nested) + } + } + } + + raw, err := json.Marshal(value) + if err != nil { + t.Fatalf("marshal active block failed: %v", err) + } + var decoded any + if err := json.Unmarshal(raw, &decoded); err != nil { + t.Fatalf("decode active block failed: %v", err) + } + walk(decoded) +} diff --git a/server/http/controller_v2_test.go b/server/http/controller_v2_test.go index 008012d8..1510fcef 100644 --- a/server/http/controller_v2_test.go +++ b/server/http/controller_v2_test.go @@ -20,10 +20,13 @@ import ( "math" "net/http" "net/http/httptest" + "net/url" "testing" + "time" "github.com/gorilla/mux" + "github.com/fatedier/frp/pkg/config/types" v1 "github.com/fatedier/frp/pkg/config/v1" "github.com/fatedier/frp/pkg/metrics/mem" httppkg "github.com/fatedier/frp/pkg/util/http" @@ -32,6 +35,10 @@ import ( "github.com/fatedier/frp/server/registry" ) +type stubControlManager struct{} + +func (stubControlManager) CloseAllProxyByName(string) error { return nil } + type v2EnvelopeForTest[T any] struct { Code int `json:"code"` Msg string `json:"msg"` @@ -39,10 +46,16 @@ type v2EnvelopeForTest[T any] struct { } type fakeStatsCollector struct { - proxies map[string]*mem.ProxyStats + server *mem.ServerStats + proxies map[string]*mem.ProxyStats + traffic map[string]*mem.ProxyTrafficInfo + pruneable map[string]bool } func (f *fakeStatsCollector) GetServer() *mem.ServerStats { + if f.server != nil { + return f.server + } return &mem.ServerStats{ProxyTypeCounts: map[string]int64{}} } @@ -69,13 +82,210 @@ func (f *fakeStatsCollector) GetProxyByName(proxyName string) *mem.ProxyStats { } func (f *fakeStatsCollector) GetProxyTraffic(name string) *mem.ProxyTrafficInfo { - return nil + return f.traffic[name] } func (f *fakeStatsCollector) ClearOfflineProxies() (int, int) { return 0, len(f.proxies) } +func (f *fakeStatsCollector) PruneOfflineProxies() (int, int) { + total := len(f.proxies) + cleared := 0 + for name := range f.pruneable { + if _, ok := f.proxies[name]; ok { + delete(f.proxies, name) + cleared++ + } + } + f.pruneable = map[string]bool{} + return cleared, total +} + +func TestAPIV2SystemInfoEnvelope(t *testing.T) { + oldStatsCollector := mem.StatsCollector + mem.StatsCollector = &fakeStatsCollector{ + server: &mem.ServerStats{ + TotalTrafficIn: 1024, + TotalTrafficOut: 2048, + CurConns: 3, + ClientCounts: 4, + ProxyTypeCounts: map[string]int64{ + "tcp": 2, + "http": 1, + }, + }, + proxies: map[string]*mem.ProxyStats{}, + } + t.Cleanup(func() { + mem.StatsCollector = oldStatsCollector + }) + + controller := NewController(&v1.ServerConfig{ + BindPort: 7000, + VhostHTTPPort: 8080, + VhostHTTPSPort: 8443, + TCPMuxHTTPConnectPort: 9000, + KCPBindPort: 7001, + QUICBindPort: 7002, + SubDomainHost: "example.com", + MaxPortsPerClient: 8, + AllowPorts: []types.PortsRange{ + {Start: 1000, End: 1002}, + {Single: 2000}, + }, + Transport: v1.ServerTransportConfig{ + MaxPoolCount: 5, + HeartbeatTimeout: 90, + TLS: v1.TLSServerConfig{ + Force: true, + }, + }, + }, registry.NewClientRegistry(), serverproxy.NewManager(), stubControlManager{}) + router := newV2TestRouter(controller) + + resp := performRequest(router, "/api/v2/system/info") + if resp.Code != http.StatusOK { + t.Fatalf("status mismatch, want %d got %d", http.StatusOK, resp.Code) + } + + rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp) + if rawResp.Code != http.StatusOK || rawResp.Msg != "success" { + t.Fatalf("envelope mismatch: %#v", rawResp) + } + assertRawJSONKeys(t, rawResp.Data, "config", "status", "version") + assertRawJSONKeysFromMessage(t, rawResp.Data["config"], + "allowPortsStr", + "bindPort", + "heartbeatTimeout", + "kcpBindPort", + "maxPoolCount", + "maxPortsPerClient", + "quicBindPort", + "subdomainHost", + "tcpmuxHTTPConnectPort", + "tlsForce", + "vhostHTTPPort", + "vhostHTTPSPort", + ) + assertRawJSONKeysFromMessage(t, rawResp.Data["status"], + "clientCounts", + "curConns", + "proxyTypeCount", + "totalTrafficIn", + "totalTrafficOut", + ) + + systemResp := decodeResponse[v2EnvelopeForTest[model.V2SystemInfoResp]](t, resp) + if systemResp.Data.Version == "" { + t.Fatal("version should be set at top level") + } + if systemResp.Data.Config.BindPort != 7000 || + systemResp.Data.Config.VhostHTTPPort != 8080 || + systemResp.Data.Config.VhostHTTPSPort != 8443 || + systemResp.Data.Config.TCPMuxHTTPConnectPort != 9000 || + systemResp.Data.Config.KCPBindPort != 7001 || + systemResp.Data.Config.QUICBindPort != 7002 || + systemResp.Data.Config.SubdomainHost != "example.com" || + systemResp.Data.Config.MaxPoolCount != 5 || + systemResp.Data.Config.MaxPortsPerClient != 8 || + systemResp.Data.Config.HeartbeatTimeout != 90 || + systemResp.Data.Config.AllowPortsStr != "1000-1002,2000" || + !systemResp.Data.Config.TLSForce { + t.Fatalf("config mismatch: %#v", systemResp.Data.Config) + } + if systemResp.Data.Status.TotalTrafficIn != 1024 || + systemResp.Data.Status.TotalTrafficOut != 2048 || + systemResp.Data.Status.CurConns != 3 || + systemResp.Data.Status.ClientCounts != 4 || + systemResp.Data.Status.ProxyTypeCounts["tcp"] != 2 || + systemResp.Data.Status.ProxyTypeCounts["http"] != 1 { + t.Fatalf("status mismatch: %#v", systemResp.Data.Status) + } +} + +func TestAPIV2SystemPruneOfflineProxies(t *testing.T) { + oldStatsCollector := mem.StatsCollector + collector := &fakeStatsCollector{ + proxies: map[string]*mem.ProxyStats{ + "tcp-offline": {Name: "tcp-offline", Type: "tcp"}, + "http-offline": {Name: "http-offline", Type: "http"}, + "udp-offline": {Name: "udp-offline", Type: "udp"}, + "tcp-online": {Name: "tcp-online", Type: "tcp"}, + "http-online": {Name: "http-online", Type: "http"}, + "udp-online": {Name: "udp-online", Type: "udp"}, + "stcp-restarted": {Name: "stcp-restarted", Type: "stcp"}, + "xtcp-restarted": {Name: "xtcp-restarted", Type: "xtcp"}, + "sudp-same-time": {Name: "sudp-same-time", Type: "sudp"}, + "tcpmux-running": {Name: "tcpmux-running", Type: "tcpmux"}, + }, + pruneable: map[string]bool{ + "tcp-offline": true, + "http-offline": true, + "udp-offline": true, + }, + } + mem.StatsCollector = collector + t.Cleanup(func() { + mem.StatsCollector = oldStatsCollector + }) + + controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager(), stubControlManager{}) + router := newV2TestRouter(controller) + + resp := performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=offline_proxies") + if resp.Code != http.StatusOK { + t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String()) + } + rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp) + if rawResp.Code != http.StatusOK || rawResp.Msg != "success" { + t.Fatalf("envelope mismatch: %#v", rawResp) + } + assertRawJSONKeys(t, rawResp.Data, "cleared", "total", "type") + pruneResp := decodeResponse[v2EnvelopeForTest[model.V2SystemPruneResp]](t, resp) + if pruneResp.Data.Type != "offline_proxies" || pruneResp.Data.Cleared != 3 || pruneResp.Data.Total != 10 { + t.Fatalf("prune response mismatch: %#v", pruneResp.Data) + } + if _, ok := collector.proxies["tcp-offline"]; ok { + t.Fatal("pruned proxy statistics should be removed") + } + if _, ok := collector.proxies["tcp-online"]; !ok { + t.Fatal("online proxy statistics should remain") + } + + resp = performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=offline_proxies") + if resp.Code != http.StatusOK { + t.Fatalf("second prune status mismatch, want %d got %d", http.StatusOK, resp.Code) + } + pruneResp = decodeResponse[v2EnvelopeForTest[model.V2SystemPruneResp]](t, resp) + if pruneResp.Data.Cleared != 0 || pruneResp.Data.Total != 7 { + t.Fatalf("second prune response mismatch: %#v", pruneResp.Data) + } +} + +func TestAPIV2SystemPruneTypeErrorsUseEnvelope(t *testing.T) { + controller := newV2TestController(t) + router := newV2TestRouter(controller) + + resp := performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune") + if resp.Code != http.StatusBadRequest { + t.Fatalf("missing type status mismatch, want %d got %d", http.StatusBadRequest, resp.Code) + } + errResp := decodeResponse[httppkg.V2Response](t, resp) + if errResp.Code != http.StatusBadRequest || errResp.Msg != "type is required" || errResp.Data != nil { + t.Fatalf("missing type error envelope mismatch: %#v", errResp) + } + + resp = performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=clients") + if resp.Code != http.StatusBadRequest { + t.Fatalf("invalid type status mismatch, want %d got %d", http.StatusBadRequest, resp.Code) + } + errResp = decodeResponse[httppkg.V2Response](t, resp) + if errResp.Code != http.StatusBadRequest || errResp.Msg != "type must be one of offline_proxies" || errResp.Data != nil { + t.Fatalf("invalid type error envelope mismatch: %#v", errResp) + } +} + func TestAPIV2ClientListEnvelopePaginationAndFilters(t *testing.T) { controller := newV2TestController(t) router := newV2TestRouter(controller) @@ -146,10 +356,49 @@ func TestAPIV2ClientDetailEnvelope(t *testing.T) { if resp.Code != http.StatusOK { t.Fatalf("status mismatch, want %d got %d", http.StatusOK, resp.Code) } - detailResp := decodeResponse[v2EnvelopeForTest[model.ClientInfoResp]](t, resp) + detailResp := decodeResponse[v2EnvelopeForTest[model.V2ClientDetailResp]](t, resp) if detailResp.Data.User != "alice" || detailResp.Data.ClientID != "client-a" { t.Fatalf("client detail mismatch: %#v", detailResp.Data) } + if detailResp.Data.Status.State != "online" || detailResp.Data.Status.CurConns != 5 || detailResp.Data.Status.ProxyCount != 2 { + t.Fatalf("client detail status mismatch: %#v", detailResp.Data.Status) + } +} + +func TestAPIV2ClientDetailEncodedKey(t *testing.T) { + oldStatsCollector := mem.StatsCollector + mem.StatsCollector = &fakeStatsCollector{ + proxies: map[string]*mem.ProxyStats{ + "tcp-url": { + Name: "tcp-url", + Type: "tcp", + User: "url", + ClientID: "client/a?b#c", + CurConns: 7, + }, + }, + } + t.Cleanup(func() { + mem.StatsCollector = oldStatsCollector + }) + + clientRegistry := registry.NewClientRegistry() + clientRegistry.Register("url", "client/a?b#c", "run-url", "url-host", "1.0.0", "127.0.0.4", "v2") + controller := NewController(&v1.ServerConfig{}, clientRegistry, serverproxy.NewManager(), stubControlManager{}) + router := newV2TestRouter(controller) + + encodedKey := url.PathEscape("url.client/a?b#c") + resp := performRequest(router, "/api/v2/clients/"+encodedKey) + if resp.Code != http.StatusOK { + t.Fatalf("encoded client key status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String()) + } + encodedResp := decodeResponse[v2EnvelopeForTest[model.V2ClientDetailResp]](t, resp) + if encodedResp.Data.User != "url" || encodedResp.Data.ClientID != "client/a?b#c" { + t.Fatalf("encoded client detail mismatch: %#v", encodedResp.Data) + } + if encodedResp.Data.Status.CurConns != 7 || encodedResp.Data.Status.ProxyCount != 1 { + t.Fatalf("encoded client detail status mismatch: %#v", encodedResp.Data.Status) + } } func TestAPIV2ProxyListDetailAndUsers(t *testing.T) { @@ -171,28 +420,194 @@ func TestAPIV2ProxyListDetailAndUsers(t *testing.T) { t.Fatalf("proxy filter total mismatch: %#v", proxyResp.Data) } proxyItem := proxyResp.Data.Items[0] - if proxyItem.Name != "tcp-empty" || proxyItem.Type != "tcp" || proxyItem.User != "" || proxyItem.Status.State != "offline" { + if proxyItem.Name != "tcp-empty" || proxyItem.Spec.Type != "tcp" || proxyItem.User != "" || proxyItem.Status.State != "offline" { t.Fatalf("proxy item mismatch: %#v", proxyItem) } + rawProxyResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[map[string]json.RawMessage]]](t, resp) + assertRawJSONKeys(t, rawProxyResp.Data.Items[0], "clientID", "name", "spec", "status", "user") + var rawListSpec map[string]json.RawMessage + if err := json.Unmarshal(rawProxyResp.Data.Items[0]["spec"], &rawListSpec); err != nil { + t.Fatalf("unmarshal list proxy spec failed: %v", err) + } + assertRawJSONKeys(t, rawListSpec, "tcp", "type") + assertRawJSONKeysFromMessage(t, rawListSpec["tcp"]) resp = performRequest(router, "/api/v2/proxies/tcp-alice") + rawProxyDetailResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp) + assertRawJSONKeysFromMessage(t, rawProxyDetailResp.Data["status"], + "curConns", + "lastCloseAt", + "lastStartAt", + "phase", + "todayTrafficIn", + "todayTrafficOut", + ) proxyDetailResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyResp]](t, resp) if proxyDetailResp.Data.Name != "tcp-alice" || proxyDetailResp.Data.User != "alice" { t.Fatalf("proxy detail mismatch: %#v", proxyDetailResp.Data) } + assertRawJSONKeys(t, rawProxyDetailResp.Data, "clientID", "name", "spec", "status", "user") + var rawDetailSpec map[string]json.RawMessage + if err := json.Unmarshal(rawProxyDetailResp.Data["spec"], &rawDetailSpec); err != nil { + t.Fatalf("unmarshal detail proxy spec failed: %v", err) + } + assertRawJSONKeys(t, rawDetailSpec, "tcp", "type") + assertRawJSONKeysFromMessage(t, rawDetailSpec["tcp"]) + if proxyDetailResp.Data.Status.LastStartAt != 1783504200 || proxyDetailResp.Data.Status.LastCloseAt != 1783504300 { + t.Fatalf("proxy detail timestamp mismatch: %#v", proxyDetailResp.Data.Status) + } resp = performRequest(router, "/api/v2/users?page=1&pageSize=50") userResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.V2UserResp]]](t, resp) if userResp.Data.Total != 3 { t.Fatalf("user total mismatch: %#v", userResp.Data) } + expectedProxyCounts := map[string]int{ + "": 1, + "alice": 2, + "bob": 1, + } for _, item := range userResp.Data.Items { - if item.ClientCount != 1 || item.ProxyCount != 1 { + if item.ClientCount != 1 || item.ProxyCount != expectedProxyCounts[item.User] { t.Fatalf("user counts mismatch: %#v", item) } } } +func TestAPIV2ProxyTrafficEnvelopeSchemaAndHistory(t *testing.T) { + oldStatsCollector := mem.StatsCollector + mem.StatsCollector = &fakeStatsCollector{ + proxies: map[string]*mem.ProxyStats{ + "ssh": {Name: "ssh", Type: "tcp"}, + }, + traffic: map[string]*mem.ProxyTrafficInfo{ + "ssh": { + Name: "ssh", + TrafficIn: []int64{70, 60, 50, 40, 30, 20, 10}, + TrafficOut: []int64{700, 600, 500, 400, 300, 200, 100}, + }, + }, + } + t.Cleanup(func() { + mem.StatsCollector = oldStatsCollector + }) + + controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager(), stubControlManager{}) + router := newV2TestRouter(controller) + + resp := performRequest(router, "/api/v2/proxies/ssh/traffic") + if resp.Code != http.StatusOK { + t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String()) + } + rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp) + if rawResp.Code != http.StatusOK || rawResp.Msg != "success" { + t.Fatalf("envelope mismatch: %#v", rawResp) + } + assertRawJSONKeys(t, rawResp.Data, "granularity", "history", "name", "unit") + + trafficResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyTrafficResp]](t, resp) + if trafficResp.Data.Name != "ssh" || trafficResp.Data.Unit != "bytes" || trafficResp.Data.Granularity != "day" { + t.Fatalf("traffic metadata mismatch: %#v", trafficResp.Data) + } + if len(trafficResp.Data.History) != 7 { + t.Fatalf("history length mismatch, want 7 got %d: %#v", len(trafficResp.Data.History), trafficResp.Data.History) + } + + wantIn := []int64{10, 20, 30, 40, 50, 60, 70} + wantOut := []int64{100, 200, 300, 400, 500, 600, 700} + var prevDate time.Time + for i, point := range trafficResp.Data.History { + assertRawJSONKeysFromMessage(t, mustMarshalJSON(t, point), "date", "trafficIn", "trafficOut") + if point.TrafficIn != wantIn[i] || point.TrafficOut != wantOut[i] { + t.Fatalf("history[%d] traffic mismatch: %#v", i, point) + } + parsedDate, err := time.Parse(time.DateOnly, point.Date) + if err != nil { + t.Fatalf("history[%d] date should be yyyy-mm-dd, got %q: %v", i, point.Date, err) + } + if i > 0 && !parsedDate.Equal(prevDate.AddDate(0, 0, 1)) { + t.Fatalf("history dates should be oldest to newest, got %s after %s", point.Date, prevDate.Format(time.DateOnly)) + } + prevDate = parsedDate + } +} + +func TestAPIV2ProxyTrafficNotFoundEnvelope(t *testing.T) { + controller := newV2TestController(t) + router := newV2TestRouter(controller) + + resp := performRequest(router, "/api/v2/proxies/missing/traffic") + if resp.Code != http.StatusNotFound { + t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusNotFound, resp.Code, resp.Body.String()) + } + errResp := decodeResponse[httppkg.V2Response](t, resp) + if errResp.Code != http.StatusNotFound || errResp.Msg != "no proxy info found" || errResp.Data != nil { + t.Fatalf("not found envelope mismatch: %#v", errResp) + } +} + +func TestAPIV2ProxyDetailAndTrafficEncodedName(t *testing.T) { + name := "folder/ssh?x#y" + oldStatsCollector := mem.StatsCollector + mem.StatsCollector = &fakeStatsCollector{ + proxies: map[string]*mem.ProxyStats{ + name: {Name: name, Type: "tcp", User: "encoded"}, + }, + traffic: map[string]*mem.ProxyTrafficInfo{ + name: { + Name: name, + TrafficIn: []int64{1}, + TrafficOut: []int64{2}, + }, + }, + } + t.Cleanup(func() { + mem.StatsCollector = oldStatsCollector + }) + + controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager(), stubControlManager{}) + router := newV2TestRouter(controller) + encodedName := url.PathEscape(name) + + resp := performRequest(router, "/api/v2/proxies/"+encodedName) + if resp.Code != http.StatusOK { + t.Fatalf("encoded proxy detail status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String()) + } + detailResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyResp]](t, resp) + if detailResp.Data.Name != name || detailResp.Data.User != "encoded" { + t.Fatalf("encoded proxy detail mismatch: %#v", detailResp.Data) + } + + resp = performRequest(router, "/api/v2/proxies/"+encodedName+"/traffic") + if resp.Code != http.StatusOK { + t.Fatalf("encoded traffic status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String()) + } + trafficResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyTrafficResp]](t, resp) + if trafficResp.Data.Name != name { + t.Fatalf("encoded traffic name mismatch: %#v", trafficResp.Data) + } + if got := trafficResp.Data.History[len(trafficResp.Data.History)-1]; got.TrafficIn != 1 || got.TrafficOut != 2 { + t.Fatalf("encoded traffic latest point mismatch: %#v", got) + } +} + +func TestAPIV2ProxyTrafficInvalidEncodedNameUses400Envelope(t *testing.T) { + controller := newV2TestController(t) + handler := httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyTraffic) + req := httptest.NewRequest(http.MethodGet, "/api/v2/proxies/%25ZZ/traffic", nil) + req = mux.SetURLVars(req, map[string]string{"name": "%ZZ"}) + resp := httptest.NewRecorder() + handler.ServeHTTP(resp, req) + + if resp.Code != http.StatusBadRequest { + t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusBadRequest, resp.Code, resp.Body.String()) + } + errResp := decodeResponse[httppkg.V2Response](t, resp) + if errResp.Code != http.StatusBadRequest || errResp.Msg != "invalid proxy name" || errResp.Data != nil { + t.Fatalf("invalid encoded name envelope mismatch: %#v", errResp) + } +} + func TestMatchV2ProxyQueryMatchesSpecFields(t *testing.T) { tests := []struct { name string @@ -202,66 +617,85 @@ func TestMatchV2ProxyQueryMatchesSpecFields(t *testing.T) { }{ { name: "tcp remote port", - item: model.V2ProxyResp{Name: "tcp-proxy", Type: "tcp", Spec: &model.TCPOutConf{ - RemotePort: 6000, + item: model.V2ProxyResp{Name: "tcp-proxy", Spec: model.V2ProxySpec{ + Type: "tcp", + TCP: &model.V2TCPProxySpec{RemotePort: v2TestIntPtr(6000)}, }}, q: "6000", want: true, }, { name: "udp remote port", - item: model.V2ProxyResp{Name: "udp-proxy", Type: "udp", Spec: &model.UDPOutConf{ - RemotePort: 7000, + item: model.V2ProxyResp{Name: "udp-proxy", Spec: model.V2ProxySpec{ + Type: "udp", + UDP: &model.V2UDPProxySpec{RemotePort: v2TestIntPtr(7000)}, }}, q: "7000", want: true, }, { name: "remote port does not match colon form", - item: model.V2ProxyResp{Name: "tcp-proxy", Type: "tcp", Spec: &model.TCPOutConf{ - RemotePort: 6000, + item: model.V2ProxyResp{Name: "tcp-proxy", Spec: model.V2ProxySpec{ + Type: "tcp", + TCP: &model.V2TCPProxySpec{RemotePort: v2TestIntPtr(6000)}, }}, q: ":6000", want: false, }, { name: "http custom domain", - item: model.V2ProxyResp{Name: "http-proxy", Type: "http", Spec: &model.HTTPOutConf{ - DomainConfig: v1.DomainConfig{CustomDomains: []string{"app.example.com"}}, + item: model.V2ProxyResp{Name: "http-proxy", Spec: model.V2ProxySpec{ + Type: "http", + HTTP: &model.V2HTTPProxySpec{CustomDomains: []string{"app.example.com"}}, }}, q: "app.example.com", want: true, }, { name: "https subdomain", - item: model.V2ProxyResp{Name: "https-proxy", Type: "https", Spec: &model.HTTPSOutConf{ - DomainConfig: v1.DomainConfig{SubDomain: "portal"}, + item: model.V2ProxyResp{Name: "https-proxy", Spec: model.V2ProxySpec{ + Type: "https", + HTTPS: &model.V2HTTPSProxySpec{Subdomain: "portal"}, }}, q: "portal", want: true, }, { name: "subdomain does not match expanded host", - item: model.V2ProxyResp{Name: "https-proxy", Type: "https", Spec: &model.HTTPSOutConf{ - DomainConfig: v1.DomainConfig{SubDomain: "portal"}, + item: model.V2ProxyResp{Name: "https-proxy", Spec: model.V2ProxySpec{ + Type: "https", + HTTPS: &model.V2HTTPSProxySpec{Subdomain: "portal"}, }}, q: "portal.example.com", want: false, }, { name: "tcpmux custom domain", - item: model.V2ProxyResp{Name: "tcpmux-proxy", Type: "tcpmux", Spec: &model.TCPMuxOutConf{ - DomainConfig: v1.DomainConfig{CustomDomains: []string{"mux.example.com"}}, + item: model.V2ProxyResp{Name: "tcpmux-proxy", Spec: model.V2ProxySpec{ + Type: "tcpmux", + TCPMux: &model.V2TCPMuxProxySpec{CustomDomains: []string{"mux.example.com"}}, }}, q: "mux.example.com", want: true, }, { - name: "nil spec does not match spec fields", - item: model.V2ProxyResp{Name: "offline-proxy", Type: "tcp", Spec: nil}, + name: "offline shell does not match online spec fields", + item: model.V2ProxyResp{Name: "offline-proxy", Spec: model.V2ProxySpec{ + Type: "tcp", + TCP: &model.V2TCPProxySpec{}, + }}, q: "6000", want: false, }, + { + name: "offline shell does not contribute zero remote port", + item: model.V2ProxyResp{Name: "offline-proxy", Spec: model.V2ProxySpec{ + Type: "tcp", + TCP: &model.V2TCPProxySpec{}, + }}, + q: "0", + want: false, + }, } for _, tt := range tests { @@ -277,7 +711,26 @@ func TestLegacyAPIResponsesRemainBare(t *testing.T) { controller := newV2TestController(t) router := newV2TestRouter(controller) - resp := performRequest(router, "/api/clients") + resp := performRequest(router, "/api/serverinfo") + var serverInfo model.ServerInfoResp + if err := json.Unmarshal(resp.Body.Bytes(), &serverInfo); err != nil { + t.Fatalf("legacy serverinfo should be a bare object: %v, body: %s", err, resp.Body.String()) + } + if serverInfo.Version == "" { + t.Fatal("legacy serverinfo version should be set") + } + var serverInfoRaw map[string]json.RawMessage + if err := json.Unmarshal(resp.Body.Bytes(), &serverInfoRaw); err != nil { + t.Fatalf("unmarshal legacy serverinfo object failed: %v", err) + } + if _, ok := serverInfoRaw["data"]; ok { + t.Fatalf("legacy serverinfo should not use v2 envelope: %s", resp.Body.String()) + } + if _, ok := serverInfoRaw["config"]; ok { + t.Fatalf("legacy serverinfo should stay flat, got config in: %s", resp.Body.String()) + } + + resp = performRequest(router, "/api/clients") var clients []model.ClientInfoResp if err := json.Unmarshal(resp.Body.Bytes(), &clients); err != nil { t.Fatalf("legacy clients should be a bare array: %v, body: %s", err, resp.Body.String()) @@ -298,6 +751,28 @@ func TestLegacyAPIResponsesRemainBare(t *testing.T) { if err := json.Unmarshal(resp.Body.Bytes(), &envelope); err == nil && envelope.Code != 0 { t.Fatalf("legacy proxy response should not use v2 envelope: %#v", envelope) } + + resp = performRequest(router, "/api/traffic/tcp-alice") + var traffic model.GetProxyTrafficResp + if err := json.Unmarshal(resp.Body.Bytes(), &traffic); err != nil { + t.Fatalf("legacy traffic should be a bare object: %v, body: %s", err, resp.Body.String()) + } + if traffic.Name != "tcp-alice" || + len(traffic.TrafficIn) != 2 || traffic.TrafficIn[0] != 7 || traffic.TrafficIn[1] != 6 || + len(traffic.TrafficOut) != 2 || traffic.TrafficOut[0] != 70 || traffic.TrafficOut[1] != 60 { + t.Fatalf("legacy traffic should preserve today-first arrays, got: %#v", traffic) + } + var trafficRaw map[string]json.RawMessage + if err := json.Unmarshal(resp.Body.Bytes(), &trafficRaw); err != nil { + t.Fatalf("unmarshal legacy traffic object failed: %v", err) + } + if _, ok := trafficRaw["data"]; ok { + t.Fatalf("legacy traffic should not use v2 envelope: %s", resp.Body.String()) + } +} + +func v2TestIntPtr(value int) *int { + return &value } func newV2TestController(t *testing.T) *Controller { @@ -322,6 +797,18 @@ func newV2TestController(t *testing.T) *Controller { ClientID: "client-a", TodayTrafficIn: 30, TodayTrafficOut: 40, + CurConns: 2, + LastStartTime: "07-08 12:30:00", + LastCloseTime: "07-08 12:31:40", + LastStartAt: 1783504200, + LastCloseAt: 1783504300, + }, + "http-alice": { + Name: "http-alice", + Type: "http", + User: "alice", + ClientID: "client-a", + CurConns: 3, }, "udp-bob": { Name: "udp-bob", @@ -330,6 +817,13 @@ func newV2TestController(t *testing.T) *Controller { ClientID: "client-b", }, }, + traffic: map[string]*mem.ProxyTrafficInfo{ + "tcp-alice": { + Name: "tcp-alice", + TrafficIn: []int64{7, 6}, + TrafficOut: []int64{70, 60}, + }, + }, } t.Cleanup(func() { mem.StatsCollector = oldStatsCollector @@ -347,17 +841,28 @@ func newV2TestController(t *testing.T) *Controller { func newV2TestRouter(controller *Controller) *mux.Router { router := mux.NewRouter() router.HandleFunc("/api/v2/users", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2UserList)).Methods(http.MethodGet) + router.HandleFunc("/api/v2/system/info", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2SystemInfo)).Methods(http.MethodGet) + router.HandleFunc("/api/v2/system/prune", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2SystemPrune)).Methods(http.MethodPost) router.HandleFunc("/api/v2/clients", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ClientList)).Methods(http.MethodGet) - router.HandleFunc("/api/v2/clients/{key}", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ClientDetail)).Methods(http.MethodGet) + encodedPathRouter := router.NewRoute().Subrouter() + encodedPathRouter.UseEncodedPath() + encodedPathRouter.HandleFunc("/api/v2/clients/{key}", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ClientDetail)).Methods(http.MethodGet) router.HandleFunc("/api/v2/proxies", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyList)).Methods(http.MethodGet) - router.HandleFunc("/api/v2/proxies/{name}", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyDetail)).Methods(http.MethodGet) + encodedPathRouter.HandleFunc("/api/v2/proxies/{name}/traffic", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyTraffic)).Methods(http.MethodGet) + encodedPathRouter.HandleFunc("/api/v2/proxies/{name}", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyDetail)).Methods(http.MethodGet) + router.HandleFunc("/api/serverinfo", httppkg.MakeHTTPHandlerFunc(controller.APIServerInfo)).Methods(http.MethodGet) router.HandleFunc("/api/clients", httppkg.MakeHTTPHandlerFunc(controller.APIClientList)).Methods(http.MethodGet) router.HandleFunc("/api/proxy/{type}", httppkg.MakeHTTPHandlerFunc(controller.APIProxyByType)).Methods(http.MethodGet) + router.HandleFunc("/api/traffic/{name}", httppkg.MakeHTTPHandlerFunc(controller.APIProxyTraffic)).Methods(http.MethodGet) return router } func performRequest(handler http.Handler, target string) *httptest.ResponseRecorder { - req := httptest.NewRequest(http.MethodGet, target, nil) + return performRequestWithMethod(handler, http.MethodGet, target) +} + +func performRequestWithMethod(handler http.Handler, method, target string) *httptest.ResponseRecorder { + req := httptest.NewRequest(method, target, nil) resp := httptest.NewRecorder() handler.ServeHTTP(resp, req) return resp @@ -372,3 +877,36 @@ func decodeResponse[T any](t *testing.T, resp *httptest.ResponseRecorder) T { } return out } + +func assertRawJSONKeys(t *testing.T, raw map[string]json.RawMessage, want ...string) { + t.Helper() + + if len(raw) != len(want) { + t.Fatalf("json keys mismatch, want %v got %v", want, raw) + } + for _, key := range want { + if _, ok := raw[key]; !ok { + t.Fatalf("json key %q missing from %v", key, raw) + } + } +} + +func assertRawJSONKeysFromMessage(t *testing.T, raw json.RawMessage, want ...string) { + t.Helper() + + var out map[string]json.RawMessage + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("unmarshal raw json object failed: %v, body: %s", err, string(raw)) + } + assertRawJSONKeys(t, out, want...) +} + +func mustMarshalJSON(t *testing.T, value any) json.RawMessage { + t.Helper() + + out, err := json.Marshal(value) + if err != nil { + t.Fatalf("marshal json failed: %v", err) + } + return out +} diff --git a/server/http/model/v2.go b/server/http/model/v2.go index edd3212d..383041b1 100644 --- a/server/http/model/v2.go +++ b/server/http/model/v2.go @@ -21,26 +21,159 @@ type V2PageResp[T any] struct { Items []T `json:"items"` } +type V2SystemInfoResp struct { + Version string `json:"version"` + Config V2SystemInfoConfigResp `json:"config"` + Status V2SystemInfoStatusResp `json:"status"` +} + +type V2SystemInfoConfigResp struct { + BindPort int `json:"bindPort"` + VhostHTTPPort int `json:"vhostHTTPPort"` + VhostHTTPSPort int `json:"vhostHTTPSPort"` + TCPMuxHTTPConnectPort int `json:"tcpmuxHTTPConnectPort"` + KCPBindPort int `json:"kcpBindPort"` + QUICBindPort int `json:"quicBindPort"` + SubdomainHost string `json:"subdomainHost"` + MaxPoolCount int64 `json:"maxPoolCount"` + MaxPortsPerClient int64 `json:"maxPortsPerClient"` + HeartbeatTimeout int64 `json:"heartbeatTimeout"` + AllowPortsStr string `json:"allowPortsStr"` + TLSForce bool `json:"tlsForce"` +} + +type V2SystemInfoStatusResp struct { + TotalTrafficIn int64 `json:"totalTrafficIn"` + TotalTrafficOut int64 `json:"totalTrafficOut"` + CurConns int64 `json:"curConns"` + ClientCounts int64 `json:"clientCounts"` + ProxyTypeCounts map[string]int64 `json:"proxyTypeCount"` +} + +type V2SystemPruneResp struct { + Type string `json:"type"` + Cleared int `json:"cleared"` + Total int `json:"total"` +} + type V2UserResp struct { User string `json:"user"` ClientCount int `json:"clientCount"` ProxyCount int `json:"proxyCount"` } +type V2ClientDetailResp struct { + ClientInfoResp + Status V2ClientStatusResp `json:"status"` +} + +type V2ClientStatusResp struct { + State string `json:"phase"` + CurConns int64 `json:"curConns"` + ProxyCount int64 `json:"proxyCount"` +} + type V2ProxyResp struct { Name string `json:"name"` - Type string `json:"type"` User string `json:"user"` ClientID string `json:"clientID"` - Spec any `json:"spec"` + Spec V2ProxySpec `json:"spec"` Status V2ProxyStatusResp `json:"status"` } +type V2ProxySpec struct { + Type string `json:"type"` + + TCP *V2TCPProxySpec `json:"tcp,omitempty"` + UDP *V2UDPProxySpec `json:"udp,omitempty"` + HTTP *V2HTTPProxySpec `json:"http,omitempty"` + HTTPS *V2HTTPSProxySpec `json:"https,omitempty"` + TCPMux *V2TCPMuxProxySpec `json:"tcpmux,omitempty"` + STCP *V2STCPProxySpec `json:"stcp,omitempty"` + SUDP *V2SUDPProxySpec `json:"sudp,omitempty"` + XTCP *V2XTCPProxySpec `json:"xtcp,omitempty"` +} + +type V2ProxyBaseSpec struct { + Annotations map[string]string `json:"annotations,omitempty"` + Metadatas map[string]string `json:"metadatas,omitempty"` + Transport *V2ProxyTransportSpec `json:"transport,omitempty"` + LoadBalancer *V2ProxyLoadBalancerSpec `json:"loadBalancer,omitempty"` +} + +type V2ProxyTransportSpec struct { + UseEncryption bool `json:"useEncryption"` + UseCompression bool `json:"useCompression"` + BandwidthLimit string `json:"bandwidthLimit"` + BandwidthLimitMode string `json:"bandwidthLimitMode"` +} + +type V2ProxyLoadBalancerSpec struct { + Group string `json:"group"` +} + +type V2TCPProxySpec struct { + V2ProxyBaseSpec + RemotePort *int `json:"remotePort,omitempty"` +} + +type V2UDPProxySpec struct { + V2ProxyBaseSpec + RemotePort *int `json:"remotePort,omitempty"` +} + +type V2HTTPProxySpec struct { + V2ProxyBaseSpec + CustomDomains []string `json:"customDomains,omitempty"` + Subdomain string `json:"subdomain,omitempty"` + Locations []string `json:"locations,omitempty"` + HostHeaderRewrite string `json:"hostHeaderRewrite,omitempty"` +} + +type V2HTTPSProxySpec struct { + V2ProxyBaseSpec + CustomDomains []string `json:"customDomains,omitempty"` + Subdomain string `json:"subdomain,omitempty"` +} + +type V2TCPMuxProxySpec struct { + V2ProxyBaseSpec + CustomDomains []string `json:"customDomains,omitempty"` + Subdomain string `json:"subdomain,omitempty"` + Multiplexer string `json:"multiplexer,omitempty"` + RouteByHTTPUser string `json:"routeByHTTPUser,omitempty"` +} + +type V2STCPProxySpec struct { + V2ProxyBaseSpec +} + +type V2SUDPProxySpec struct { + V2ProxyBaseSpec +} + +type V2XTCPProxySpec struct { + V2ProxyBaseSpec +} + type V2ProxyStatusResp struct { State string `json:"phase"` TodayTrafficIn int64 `json:"todayTrafficIn"` TodayTrafficOut int64 `json:"todayTrafficOut"` CurConns int64 `json:"curConns"` - LastStartTime string `json:"lastStartTime"` - LastCloseTime string `json:"lastCloseTime"` + LastStartAt int64 `json:"lastStartAt,omitempty"` + LastCloseAt int64 `json:"lastCloseAt,omitempty"` +} + +type V2ProxyTrafficResp struct { + Name string `json:"name"` + Unit string `json:"unit"` + Granularity string `json:"granularity"` + History []V2ProxyTrafficPointResp `json:"history"` +} + +type V2ProxyTrafficPointResp struct { + Date string `json:"date"` + TrafficIn int64 `json:"trafficIn"` + TrafficOut int64 `json:"trafficOut"` } diff --git a/server/proxy/proxy.go b/server/proxy/proxy.go index 002b5f8a..292fb410 100644 --- a/server/proxy/proxy.go +++ b/server/proxy/proxy.go @@ -82,19 +82,20 @@ 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 + 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 mu sync.RWMutex xl *xlog.Logger @@ -354,10 +355,18 @@ 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) - 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) + 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) + } } return libio.Join(local, userConn) } @@ -366,6 +375,10 @@ 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() @@ -373,10 +386,46 @@ 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 @@ -388,13 +437,21 @@ 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 := msg.NewReadWriter(proxyConn, proxyWireProtocol) - visitorRW := msg.NewReadWriter(visitorConn, visitorWireProtocol) + 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} + } var ( once sync.Once @@ -496,6 +553,7 @@ type Options struct { ServerCfg *v1.ServerConfig EncryptionKey []byte WireProtocol string + UDPPacketCodec string } func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) { @@ -505,24 +563,25 @@ 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 = rate.NewLimiter(rate.Limit(float64(limitBytes)), int(limitBytes)) + limiter = limit.NewBandwidthLimiter(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, + 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, } factory := proxyFactoryRegistry[reflect.TypeOf(configurer)] diff --git a/server/proxy/sudp_benchmark_test.go b/server/proxy/sudp_benchmark_test.go new file mode 100644 index 00000000..a2e337d9 --- /dev/null +++ b/server/proxy/sudp_benchmark_test.go @@ -0,0 +1,232 @@ +// 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) diff --git a/server/proxy/sudp_test.go b/server/proxy/sudp_test.go index e2986eed..b7ea9bc6 100644 --- a/server/proxy/sudp_test.go +++ b/server/proxy/sudp_test.go @@ -18,22 +18,27 @@ 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( - msg.NewReadWriter(&in, wire.ProtocolV1), - msg.NewReadWriter(&out, wire.ProtocolV2), + newSUDPBridgeRW(t, &in, wire.ProtocolV1, ""), + newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""), &count, nil, ) @@ -53,12 +58,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( - msg.NewReadWriter(&in, wire.ProtocolV2), - msg.NewReadWriter(&out, wire.ProtocolV1), + newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""), + newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""), &count, nil, ) @@ -76,33 +81,67 @@ func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) { require.Equal(t, []byte("visitor-to-proxy"), got.Content) } -func TestSUDPBridgeForwardsProxyPing(t *testing.T) { +func TestSUDPBridgeTranscodesProxyV2BinaryToVisitorV2JSON(t *testing.T) { var in, out bytes.Buffer - writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, &msg.Ping{}) + packet := newSUDPBridgeUDPPacket("proxy-binary-to-json") + writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, wire.UDPPacketCodecBinary, packet) var count int64 err := bridgeSUDPProxyToVisitor( - msg.NewReadWriter(&in, wire.ProtocolV1), - msg.NewReadWriter(&out, wire.ProtocolV2), + 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{}) + + var count int64 + err := bridgeSUDPProxyToVisitor( + newSUDPBridgeRW(t, &in, wire.ProtocolV1, ""), + newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""), &count, nil, ) require.NoError(t, err) require.Zero(t, count) - rawMsg, err := msg.NewReadWriter(&out, wire.ProtocolV2).ReadMsg() + rawMsg, err := newSUDPBridgeRW(t, &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( - msg.NewReadWriter(&in, wire.ProtocolV2), - msg.NewReadWriter(&out, wire.ProtocolV1), + newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""), + newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""), &count, nil, ) @@ -113,12 +152,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( - msg.NewReadWriter(&in, wire.ProtocolV2), - msg.NewReadWriter(&out, wire.ProtocolV1), + newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""), + newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""), &count, nil, ) @@ -127,6 +166,22 @@ 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)) @@ -134,8 +189,170 @@ func TestSUDPBridgeDetectsMixedWireProtocol(t *testing.T) { require.True(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV1)) } -func writeSUDPBridgeMsg(t *testing.T, buf *bytes.Buffer, wireProtocol string, m msg.Message) { - t.Helper() - - require.NoError(t, msg.NewReadWriter(buf, wireProtocol).WriteMsg(m)) +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 { + 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 } diff --git a/server/proxy/udp.go b/server/proxy/udp.go index 609d9c42..47bdfc8c 100644 --- a/server/proxy/udp.go +++ b/server/proxy/udp.go @@ -224,7 +224,13 @@ 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. - payloadConn := msg.NewConn(pxy.workConn, msg.NewReadWriter(pxy.workConn, pxy.wireProtocol)) + 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) ctx, cancel := context.WithCancel(context.Background()) go workConnReaderFn(payloadConn) go workConnSenderFn(payloadConn, ctx) diff --git a/server/registry/registry.go b/server/registry/registry.go index 21771bce..a3632f5d 100644 --- a/server/registry/registry.go +++ b/server/registry/registry.go @@ -28,6 +28,7 @@ type ClientInfo struct { User string RawClientID string RunID string + ControlID uint64 Hostname string IP string Version string @@ -64,6 +65,16 @@ func newClientRegistryWithClock(clk clock.PassiveClock) *ClientRegistry { // Register stores/updates metadata for a client and returns the registry key plus whether it conflicts with an online client. func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol string) (key string, conflict bool) { + return cr.RegisterWithControlID(user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol, 0) +} + +// RegisterWithControlID is the generation-aware form used by ControlManager. +// A control ID is process-local and prevents an older control generation from +// changing the registry entry now owned by a newer generation with the same run ID. +func (cr *ClientRegistry) RegisterWithControlID( + user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol string, + controlID uint64, +) (key string, conflict bool) { if runID == "" { return "", false } @@ -83,6 +94,16 @@ func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version, if enforceUnique && exists && info.Online && info.RunID != "" && info.RunID != runID { return key, true } + if previousKey, ok := cr.runIndex[runID]; ok && previousKey != key { + if previous, ok := cr.clients[previousKey]; ok && previous.RunID == runID { + if previous.RawClientID == "" { + delete(cr.clients, previousKey) + } else { + setClientOffline(previous, now) + } + } + delete(cr.runIndex, runID) + } if !exists { info = &ClientInfo{ @@ -97,6 +118,7 @@ func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version, info.RawClientID = rawClientID info.RunID = runID + info.ControlID = controlID info.Hostname = hostname info.IP = remoteAddr info.Version = version @@ -114,6 +136,16 @@ func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version, // MarkOfflineByRunID marks the client as offline when the corresponding control disconnects. func (cr *ClientRegistry) MarkOfflineByRunID(runID string) { + cr.markOfflineByRunID(runID, 0, false) +} + +// MarkOfflineByRunIDAndControlID marks a client offline only when the registry +// entry still belongs to the supplied control generation. +func (cr *ClientRegistry) MarkOfflineByRunIDAndControlID(runID string, controlID uint64) { + cr.markOfflineByRunID(runID, controlID, true) +} + +func (cr *ClientRegistry) markOfflineByRunID(runID string, controlID uint64, matchControlID bool) { cr.mu.Lock() defer cr.mu.Unlock() @@ -121,17 +153,23 @@ func (cr *ClientRegistry) MarkOfflineByRunID(runID string) { if !ok { return } - if info, ok := cr.clients[key]; ok && info.RunID == runID { + if info, ok := cr.clients[key]; ok && info.RunID == runID && (!matchControlID || info.ControlID == controlID) { if info.RawClientID == "" { delete(cr.clients, key) } else { - info.RunID = "" - info.Online = false - now := cr.clock.Now() - info.DisconnectedAt = now + setClientOffline(info, cr.clock.Now()) } } - delete(cr.runIndex, runID) + if info, ok := cr.clients[key]; !ok || info.RunID != runID { + delete(cr.runIndex, runID) + } +} + +func setClientOffline(info *ClientInfo, now time.Time) { + info.RunID = "" + info.ControlID = 0 + info.Online = false + info.DisconnectedAt = now } // List returns a snapshot of all known clients. diff --git a/server/registry/registry_test.go b/server/registry/registry_test.go index 0ff083b8..bacac949 100644 --- a/server/registry/registry_test.go +++ b/server/registry/registry_test.go @@ -72,3 +72,89 @@ func TestClientRegistryUsesClockForTimestamps(t *testing.T) { t.Fatalf("disconnected time mismatch, want %s got %s", disconnectedAt, info.DisconnectedAt) } } + +func TestClientRegistryControlIDPreventsStaleOffline(t *testing.T) { + registry := NewClientRegistry() + key, conflict := registry.RegisterWithControlID( + "user", "client-id", "run-id", "old-host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1, + ) + if conflict { + t.Fatal("unexpected client conflict") + } + _, conflict = registry.RegisterWithControlID( + "user", "client-id", "run-id", "new-host", "1.0.1", "127.0.0.2", wire.ProtocolV2, 2, + ) + if conflict { + t.Fatal("same run ID replacement should not conflict") + } + + registry.MarkOfflineByRunIDAndControlID("run-id", 1) + info, ok := registry.GetByKey(key) + if !ok { + t.Fatalf("client %q not found", key) + } + if !info.Online || info.ControlID != 2 || info.Hostname != "new-host" { + t.Fatalf("stale offline changed current generation: %+v", info) + } + + registry.MarkOfflineByRunIDAndControlID("run-id", 2) + info, ok = registry.GetByKey(key) + if !ok { + t.Fatalf("client %q not found after disconnect", key) + } + if info.Online || info.ControlID != 0 || info.RunID != "" { + t.Fatalf("current generation was not marked offline: %+v", info) + } +} + +func TestClientRegistryClientIDConflictSemantics(t *testing.T) { + registry := NewClientRegistry() + _, conflict := registry.RegisterWithControlID( + "user", "client-id", "run-one", "host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1, + ) + if conflict { + t.Fatal("unexpected initial client conflict") + } + _, conflict = registry.RegisterWithControlID( + "user", "client-id", "run-two", "host", "1.0.0", "127.0.0.2", wire.ProtocolV1, 2, + ) + if !conflict { + t.Fatal("different online run IDs with the same explicit client ID must conflict") + } + + registry.MarkOfflineByRunIDAndControlID("run-one", 1) + _, conflict = registry.RegisterWithControlID( + "user", "client-id", "run-two", "host", "1.0.0", "127.0.0.2", wire.ProtocolV1, 2, + ) + if conflict { + t.Fatal("offline explicit client ID should be reusable") + } +} + +func TestClientRegistrySameRunIDMovesBetweenClientKeys(t *testing.T) { + registry := NewClientRegistry() + oldKey, conflict := registry.RegisterWithControlID( + "user", "old-client", "run-id", "old-host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1, + ) + if conflict { + t.Fatal("unexpected initial client conflict") + } + newKey, conflict := registry.RegisterWithControlID( + "user", "new-client", "run-id", "new-host", "1.0.1", "127.0.0.2", wire.ProtocolV2, 2, + ) + if conflict { + t.Fatal("same run ID moving to a new client key should not conflict") + } + + oldInfo, ok := registry.GetByKey(oldKey) + if !ok { + t.Fatalf("old explicit client %q should remain as offline history", oldKey) + } + if oldInfo.Online || oldInfo.RunID != "" || oldInfo.ControlID != 0 { + t.Fatalf("old client key remained online: %+v", oldInfo) + } + newInfo, ok := registry.GetByKey(newKey) + if !ok || !newInfo.Online || newInfo.RunID != "run-id" || newInfo.ControlID != 2 { + t.Fatalf("new client key was not registered: %+v", newInfo) + } +} diff --git a/server/service.go b/server/service.go index 05f24a87..3f88b1c4 100644 --- a/server/service.go +++ b/server/service.go @@ -18,6 +18,7 @@ import ( "bytes" "context" "crypto/tls" + "errors" "fmt" "io" "net" @@ -51,7 +52,6 @@ import ( "github.com/fatedier/frp/pkg/util/xlog" "github.com/fatedier/frp/server/controller" "github.com/fatedier/frp/server/group" - "github.com/fatedier/frp/server/metrics" "github.com/fatedier/frp/server/ports" "github.com/fatedier/frp/server/proxy" "github.com/fatedier/frp/server/registry" @@ -64,6 +64,8 @@ const ( vhostReadWriteTimeout time.Duration = 30 * time.Second ) +var errControlReplaced = errors.New("control was replaced during login") + func init() { crypto.DefaultSalt = "frp" // Disable quic-go's receive buffer warning. @@ -161,9 +163,10 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) { return nil, err } + clientRegistry := registry.NewClientRegistry() svr := &Service{ - ctlManager: NewControlManager(), - clientRegistry: registry.NewClientRegistry(), + ctlManager: NewControlManager(clientRegistry), + clientRegistry: clientRegistry, pxyManager: proxy.NewManager(), pluginManager: plugin.NewManager(), rc: &controller.ResourceController{ @@ -303,10 +306,14 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) { svr.rc.HTTPReverseProxy = rp address := net.JoinHostPort(cfg.ProxyBindAddr, strconv.Itoa(cfg.VhostHTTPPort)) + protocols := new(http.Protocols) + protocols.SetHTTP1(true) + protocols.SetUnencryptedHTTP2(true) server := &http.Server{ Addr: address, Handler: rp, ReadHeaderTimeout: 60 * time.Second, + Protocols: protocols, } var l net.Listener if httpMuxOn { @@ -469,12 +476,15 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna } } if err == nil { - ctl, err = svr.RegisterControl(controlConn, m, internal, acceptedConn.wireProtocol) + ctl, err = svr.RegisterControl(controlConn, m, internal, acceptedConn.wireProtocol, acceptedConn.udpPacketCodec) } } if err != nil { xl.Warnf("register control error: %v", err) + if ctl != nil { + svr.ctlManager.Remove(ctl) + } if writeErr := writeWithDeadline(conn, connWriteTimeout, func() error { return acceptedConn.conn.WriteMsg(&msg.LoginResp{ Version: version.Full(), @@ -483,31 +493,34 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna }); writeErr != nil { xl.Warnf("write login error response error: %v", writeErr) } - conn.Close() + if ctl != nil { + _ = ctl.Close() + } else { + conn.Close() + } return } - if err = writeWithDeadline(conn, connWriteTimeout, func() error { - return acceptedConn.conn.WriteMsg(&msg.LoginResp{ - Version: version.Full(), - RunID: ctl.runID, - Error: "", + if err = svr.completeControlLogin(ctl, func() error { + return writeWithDeadline(conn, connWriteTimeout, func() error { + return acceptedConn.conn.WriteMsg(&msg.LoginResp{ + Version: version.Full(), + RunID: ctl.runID, + Error: "", + }) }) }); err != nil { - xl.Warnf("write login response error: %v", err) - svr.ctlManager.Del(m.RunID, ctl) - svr.clientRegistry.MarkOfflineByRunID(m.RunID) - conn.Close() + xl.Warnf("complete control login error: %v", err) + svr.ctlManager.Remove(ctl) + _ = ctl.Close() return } - ctl.Start() - metrics.Server.NewClient() - go func() { - // block until control closed - ctl.WaitClosed() - svr.ctlManager.Del(m.RunID, ctl) - }() case *msg.NewWorkConn: - if err := svr.RegisterWorkConn(acceptedConn.conn, m); err != nil { + if err := svr.RegisterWorkConn( + acceptedConn.conn, + m, + acceptedConn.wireProtocol, + acceptedConn.clientHelloPresent, + ); err != nil { _ = acceptedConn.conn.WriteMsg(&msg.StartWorkConn{ Error: util.GenerateResponseErrorString("invalid NewWorkConn", err, lo.FromPtr(svr.cfg.DetailedErrorsToClient)), }) @@ -533,11 +546,24 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna } } +func (svr *Service) completeControlLogin(ctl *Control, writeSuccess func() error) error { + committed, err := svr.ctlManager.completeLogin(ctl, writeSuccess) + if err != nil { + return err + } + if !committed { + return errControlReplaced + } + return nil +} + type acceptedConnection struct { - conn *msg.Conn - wireProtocol string - cryptoContext *wire.CryptoContext - firstMsg msg.Message + conn *msg.Conn + wireProtocol string + clientHelloPresent bool + udpPacketCodec string + cryptoContext *wire.CryptoContext + firstMsg msg.Message } func (svr *Service) acceptConnection(ctx context.Context, conn net.Conn) (*acceptedConnection, error) { @@ -605,6 +631,7 @@ 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 } @@ -653,6 +680,7 @@ 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 } @@ -746,7 +774,20 @@ 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 @@ -782,8 +823,8 @@ func (svr *Service) RegisterControl( Conn: ctlConn, LoginMsg: loginMsg, ServerCfg: svr.cfg, - ClientRegistry: svr.clientRegistry, WireProtocol: wireProtocol, + UDPPacketCodec: udpPacketCodec, }) if err != nil { xl.Warnf("create new controller error: %v", err) @@ -791,31 +832,41 @@ func (svr *Service) RegisterControl( return nil, fmt.Errorf("unexpected error when creating new controller") } - if oldCtl := svr.ctlManager.Add(loginMsg.RunID, ctl); oldCtl != nil { - oldCtl.WaitClosed() + if err := svr.ctlManager.Add(ctl); err != nil { + return ctl, err } + ctl.WaitForHandoff() - remoteAddr := ctlConn.RemoteAddr().String() - if host, _, err := net.SplitHostPort(remoteAddr); err == nil { - remoteAddr = host + active, err := svr.ctlManager.Activate(ctl) + if err != nil { + return ctl, err } - _, conflict := svr.clientRegistry.Register(loginMsg.User, loginMsg.ClientID, loginMsg.RunID, loginMsg.Hostname, loginMsg.Version, remoteAddr, wireProtocol) - if conflict { - svr.ctlManager.Del(loginMsg.RunID, ctl) - return nil, fmt.Errorf("client_id [%s] for user [%s] is already online", loginMsg.ClientID, loginMsg.User) + if !active { + return ctl, errControlReplaced } return ctl, nil } // RegisterWorkConn register a new work connection to control and proxies need it. -func (svr *Service) RegisterWorkConn(workConn *msg.Conn, newMsg *msg.NewWorkConn) error { +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") + } 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{ @@ -836,20 +887,33 @@ func (svr *Service) RegisterWorkConn(workConn *msg.Conn, newMsg *msg.NewWorkConn xl.Warnf("invalid NewWorkConn with run id [%s]", newMsg.RunID) return err } - return ctl.RegisterWorkConn(proxy.NewWorkConn(workConn)) + return svr.ctlManager.RegisterWorkConn(ctl, proxy.NewWorkConn(workConn)) } func (svr *Service) RegisterVisitorConn(visitorConn net.Conn, newMsg *msg.NewVisitorConn, wireProtocol string) error { - visitorUser := "" + admit := func(visitorUser, visitorWireProtocol, visitorUDPPacketCodec string) error { + if visitorWireProtocol == "" { + visitorWireProtocol = wireProtocol + } + return svr.rc.VisitorManager.NewConn(newMsg.ProxyName, visitorConn, newMsg.Timestamp, newMsg.SignKey, + newMsg.UseEncryption, newMsg.UseCompression, visitorUser, visitorWireProtocol, visitorUDPPacketCodec) + } // 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 != "" { - ctl, exist := svr.ctlManager.GetByID(newMsg.RunID) - if !exist { + admitted, err := svr.ctlManager.admitVisitorByRunID(newMsg.RunID, func(visitorUser, controlWireProtocol, controlUDPPacketCodec string) error { + if wireProtocol != controlWireProtocol { + return fmt.Errorf("visitor connection wire protocol mismatch: got %s want %s", wireProtocol, controlWireProtocol) + } + return admit(visitorUser, controlWireProtocol, controlUDPPacketCodec) + }) + if err != nil { + return err + } + if !admitted { return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID) } - visitorUser = ctl.sessionCtx.LoginMsg.User + return nil } - return svr.rc.VisitorManager.NewConn(newMsg.ProxyName, visitorConn, newMsg.Timestamp, newMsg.SignKey, - newMsg.UseEncryption, newMsg.UseCompression, visitorUser, wireProtocol) + return admit("", wireProtocol, "") } diff --git a/server/service_test.go b/server/service_test.go index cf5eec5e..e35f11b9 100644 --- a/server/service_test.go +++ b/server/service_test.go @@ -15,12 +15,30 @@ package server import ( + "context" "errors" + "math" "net" + "net/http" + "runtime" + "sync" + "sync/atomic" "testing" "time" + "github.com/fatedier/golib/net/mux" "github.com/stretchr/testify/require" + + "github.com/fatedier/frp/pkg/auth" + v1 "github.com/fatedier/frp/pkg/config/v1" + "github.com/fatedier/frp/pkg/msg" + plugin "github.com/fatedier/frp/pkg/plugin/server" + "github.com/fatedier/frp/pkg/proto/wire" + "github.com/fatedier/frp/pkg/util/util" + "github.com/fatedier/frp/server/controller" + "github.com/fatedier/frp/server/proxy" + "github.com/fatedier/frp/server/registry" + "github.com/fatedier/frp/server/visitor" ) func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) { @@ -61,3 +79,879 @@ func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) { t.Fatal("timed out waiting for write after deadline reset") } } + +func TestServiceAcceptConnectionTracksClientHelloPresence(t *testing.T) { + for _, tc := range []struct { + name string + clientHelloPresent bool + offeredCodecs []string + expectedCodec string + }{ + { + name: "absent Hello", + }, + { + name: "present Hello with JSON fallback", + clientHelloPresent: true, + }, + { + name: "present Hello with binary codec", + clientHelloPresent: true, + offeredCodecs: []string{wire.UDPPacketCodecBinary}, + expectedCodec: wire.UDPPacketCodecBinary, + }, + } { + t.Run(tc.name, func(t *testing.T) { + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + clientErrCh := make(chan error, 1) + go func() { + if err := wire.WriteMagic(clientConn); err != nil { + clientErrCh <- err + return + } + wireConn := wire.NewConn(clientConn) + if tc.clientHelloPresent { + hello, err := wire.NewClientHello(wire.BootstrapInfo{}) + if err != nil { + clientErrCh <- err + return + } + hello.Capabilities.Message.UDPPacketCodecs = tc.offeredCodecs + if err := wireConn.WriteJSONFrame(wire.FrameTypeClientHello, hello); err != nil { + clientErrCh <- err + return + } + var serverHello wire.ServerHello + if err := wireConn.ReadJSONFrame(wire.FrameTypeServerHello, &serverHello); err != nil { + clientErrCh <- err + return + } + } + clientErrCh <- msg.NewV2ReadWriterWithConn(wireConn).WriteMsg(&msg.NewWorkConn{RunID: "shared-run"}) + }() + + acceptedConn, err := (&Service{}).acceptConnection(t.Context(), serverConn) + require.NoError(t, err) + require.NoError(t, <-clientErrCh) + require.Equal(t, tc.clientHelloPresent, acceptedConn.clientHelloPresent) + require.Equal(t, tc.expectedCodec, acceptedConn.udpPacketCodec) + require.IsType(t, &msg.NewWorkConn{}, acceptedConn.firstMsg) + require.NoError(t, acceptedConn.conn.Close()) + }) + } +} + +func TestSharedPortHTTPListenerProtocols(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + sharedMux := mux.NewMux(listener) + httpListener := sharedMux.ListenHTTP(1) + muxServeErr := make(chan error, 1) + go func() { + muxServeErr <- sharedMux.Serve() + }() + + newProtocols := func(http1, unencryptedHTTP2 bool) *http.Protocols { + protocols := new(http.Protocols) + protocols.SetHTTP1(http1) + protocols.SetUnencryptedHTTP2(unencryptedHTTP2) + return protocols + } + + const handlerProtocolHeader = "X-Test-Handler-Protocol" + httpServer := &http.Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set(handlerProtocolHeader, r.Proto) + w.WriteHeader(http.StatusNoContent) + }), + ReadHeaderTimeout: time.Second, + Protocols: newProtocols(true, true), + } + httpServeErr := make(chan error, 1) + go func() { + httpServeErr <- httpServer.Serve(httpListener) + }() + t.Cleanup(func() { + require.NoError(t, httpServer.Close()) + require.ErrorIs(t, waitForResult(t, httpServeErr, "shared HTTP server to stop"), http.ErrServerClosed) + require.NoError(t, sharedMux.Close()) + require.ErrorIs(t, waitForResult(t, muxServeErr, "shared mux to stop"), net.ErrClosed) + }) + + for _, tc := range []struct { + name string + http1 bool + unencryptedHTTP2 bool + expectedProtocol string + }{ + {name: "HTTP/1.1", http1: true, expectedProtocol: "HTTP/1.1"}, + {name: "HTTP/2 prior knowledge", unencryptedHTTP2: true, expectedProtocol: "HTTP/2.0"}, + } { + t.Run(tc.name, func(t *testing.T) { + transport := &http.Transport{ + Protocols: newProtocols(tc.http1, tc.unencryptedHTTP2), + } + defer transport.CloseIdleConnections() + client := &http.Client{Transport: transport, Timeout: 3 * time.Second} + request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "http://"+listener.Addr().String()+"/", nil) + require.NoError(t, err) + response, err := client.Do(request) + require.NoError(t, err) + require.Equal(t, http.StatusNoContent, response.StatusCode) + require.Equal(t, tc.expectedProtocol, response.Proto) + require.Equal(t, tc.expectedProtocol, response.Header.Get(handlerProtocolHeader)) + require.NoError(t, response.Body.Close()) + }) + } +} + +func TestServiceControlHandoffSkipsStalePendingGeneration(t *testing.T) { + svr := newControlTestService(t) + metrics := newCountingServerMetrics() + metrics.closeEnter = make(chan struct{}) + metrics.closeResume = make(chan struct{}) + + ctlA, connA, err := registerLifecycleTestControl(svr) + require.NoError(t, err) + ctlA.serverMetrics = metrics + require.NoError(t, svr.completeControlLogin(ctlA, func() error { return nil })) + waitForSignal(t, connA.readStarted, "A reader to start") + + require.NoError(t, ctlA.Close()) + waitForSignal(t, metrics.closeEnter, "A finalization barrier") + + type registerResult struct { + ctl *Control + conn *deadlineReadConn + err error + } + resultB := make(chan registerResult, 1) + go func() { + ctl, conn, registerErr := registerLifecycleTestControl(svr) + resultB <- registerResult{ctl: ctl, conn: conn, err: registerErr} + }() + ctlB := waitForDifferentCurrentControl(t, svr.ctlManager, "shared-run", ctlA) + ctlB.serverMetrics = metrics + + resultC := make(chan registerResult, 1) + go func() { + ctl, conn, registerErr := registerLifecycleTestControl(svr) + resultC <- registerResult{ctl: ctl, conn: conn, err: registerErr} + }() + ctlC := waitForDifferentCurrentControl(t, svr.ctlManager, "shared-run", ctlB) + ctlC.serverMetrics = metrics + waitForControlDone(t, ctlB) + + select { + case result := <-resultB: + t.Fatalf("B returned before A finalized: %v", result.err) + default: + } + select { + case result := <-resultC: + t.Fatalf("C returned before A finalized: %v", result.err) + default: + } + + close(metrics.closeResume) + waitForControlDone(t, ctlA) + + b := <-resultB + require.Same(t, ctlB, b.ctl) + require.ErrorIs(t, b.err, errControlReplaced) + require.False(t, svr.ctlManager.Remove(ctlB)) + require.NoError(t, ctlB.Close()) + + c := <-resultC + require.NoError(t, c.err) + require.Same(t, ctlC, c.ctl) + _, ok := svr.ctlManager.GetByID("shared-run") + require.False(t, ok) + require.Same(t, ctlC, currentControlForTest(svr.ctlManager, "shared-run")) + + info, ok := svr.clientRegistry.GetByKey("client") + require.True(t, ok) + require.True(t, info.Online) + require.Equal(t, uint64(ctlC.ID()), info.ControlID) + + var staleWrites atomic.Int64 + err = svr.completeControlLogin(ctlB, func() error { + staleWrites.Add(1) + return nil + }) + require.ErrorIs(t, err, errControlReplaced) + require.Equal(t, int64(0), staleWrites.Load()) + + require.NoError(t, svr.completeControlLogin(ctlC, func() error { return nil })) + waitForSignal(t, c.conn.readStarted, "C reader to start") + current, ok := svr.ctlManager.GetByID("shared-run") + require.True(t, ok) + require.Same(t, ctlC, current) + require.Equal(t, int64(2), metrics.newClients()) + require.Equal(t, int64(1), metrics.closedClients()) + + require.NoError(t, ctlC.Close()) + waitForControlDone(t, ctlC) + require.Equal(t, int64(2), metrics.newClients()) + require.Equal(t, int64(2), metrics.closedClients()) + _, ok = svr.ctlManager.GetByID("shared-run") + require.False(t, ok) +} + +func TestServiceLoginResponseSynchronizationIsScopedToRun(t *testing.T) { + svr := newControlTestService(t) + metrics := newCountingServerMetrics() + ctlA, connA, err := registerLifecycleTestControl(svr) + require.NoError(t, err) + ctlA.serverMetrics = metrics + + writeEntered := make(chan struct{}) + resumeWrite := make(chan struct{}) + var resumeWriteOnce sync.Once + resume := func() { + resumeWriteOnce.Do(func() { close(resumeWrite) }) + } + t.Cleanup(resume) + writeCount := atomic.Int64{} + loginDone := make(chan error, 1) + go func() { + loginDone <- svr.completeControlLogin(ctlA, func() error { + close(writeEntered) + <-resumeWrite + writeCount.Add(1) + return nil + }) + }() + waitForSignal(t, writeEntered, "A LoginResp write") + + runMu := currentRunGateForTest(svr.ctlManager, "shared-run") + require.NotNil(t, runMu) + if !svr.ctlManager.mu.TryLock() { + t.Fatal("ControlManager mutex was held while LoginResp write was in progress") + } + svr.ctlManager.mu.Unlock() + + ctlB, connB := newLifecycleTestControl(t, "shared-run", "client", metrics) + gateAvailable := make(chan bool) + addDone := make(chan error, 1) + go func() { + if runMu.TryLock() { + runMu.Unlock() + gateAvailable <- true + } else { + gateAvailable <- false + } + addErr := svr.ctlManager.Add(ctlB) + addDone <- addErr + }() + available := waitForResult(t, gateAvailable, "same-run replacement gate probe") + require.False(t, available, "same-run gate was available to replacement during LoginResp write") + select { + case addErr := <-addDone: + t.Fatalf("same-run replacement completed during LoginResp write: %v", addErr) + case <-time.After(20 * time.Millisecond): + } + require.Same(t, ctlA, currentControlForTest(svr.ctlManager, "shared-run")) + + otherMetrics := newCountingServerMetrics() + otherCtl, otherConn := newLifecycleTestControl(t, "other-run", "other-client", otherMetrics) + type unrelatedResult struct { + addErr error + active bool + activateErr error + loginErr error + current *Control + found bool + } + unrelatedDone := make(chan unrelatedResult, 1) + go func() { + result := unrelatedResult{} + result.addErr = svr.ctlManager.Add(otherCtl) + if result.addErr == nil { + result.active, result.activateErr = svr.ctlManager.Activate(otherCtl) + } + if result.activateErr == nil && result.active { + result.loginErr = svr.completeControlLogin(otherCtl, func() error { return nil }) + } + result.current, result.found = svr.ctlManager.GetByID("other-run") + unrelatedDone <- result + }() + result := waitForResult(t, unrelatedDone, "unrelated run lifecycle") + require.NoError(t, result.addErr) + require.NoError(t, result.activateErr) + require.True(t, result.active) + require.NoError(t, result.loginErr) + require.True(t, result.found) + require.Same(t, otherCtl, result.current) + waitForSignal(t, otherConn.readStarted, "unrelated control reader to start") + require.Equal(t, int64(1), otherMetrics.newClients()) + + resume() + require.NoError(t, waitForResult(t, loginDone, "LoginResp completion")) + require.NoError(t, waitForResult(t, addDone, "replacement")) + waitForControlDone(t, ctlA) + require.Same(t, ctlB, currentControlForTest(svr.ctlManager, "shared-run")) + require.Equal(t, int64(1), writeCount.Load()) + require.Equal(t, int64(1), metrics.newClients()) + require.Equal(t, int64(1), metrics.closedClients()) + require.Equal(t, []string{"deadline", "close"}, connA.eventsSnapshot()) + + require.False(t, svr.ctlManager.Remove(ctlA)) + require.NoError(t, ctlA.Close()) + require.True(t, svr.ctlManager.Remove(ctlB)) + require.NoError(t, ctlB.Close()) + require.Equal(t, []string{"deadline", "close"}, connB.eventsSnapshot()) + + require.NoError(t, otherCtl.Close()) + waitForControlDone(t, otherCtl) + require.Equal(t, int64(1), otherMetrics.newClients()) + require.Equal(t, int64(1), otherMetrics.closedClients()) +} + +func TestServiceVisitorAdmissionSerializesReplacement(t *testing.T) { + svr := newControlTestService(t) + ctlA, controlConn, err := registerLifecycleTestControl(svr) + require.NoError(t, err) + ctlA.sessionCtx.LoginMsg.User = "old-user" + require.NoError(t, svr.completeControlLogin(ctlA, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "A reader to start") + + admissionEntered := make(chan struct{}) + resumeAdmission := make(chan struct{}) + var resumeOnce sync.Once + resume := func() { + resumeOnce.Do(func() { close(resumeAdmission) }) + } + t.Cleanup(resume) + type admissionResult struct { + admitted bool + user string + wireProtocol string + udpPacketCodec string + err error + } + admissionDone := make(chan admissionResult, 1) + go func() { + result := admissionResult{} + result.admitted, result.err = svr.ctlManager.admitVisitorByRunID("shared-run", func(user, wireProtocol, udpPacketCodec string) error { + result.user = user + result.wireProtocol = wireProtocol + result.udpPacketCodec = udpPacketCodec + close(admissionEntered) + <-resumeAdmission + return nil + }) + admissionDone <- result + }() + waitForSignal(t, admissionEntered, "visitor admission callback") + runMu := currentRunGateForTest(svr.ctlManager, "shared-run") + require.NotNil(t, runMu) + + type registerResult struct { + ctl *Control + err error + } + gateAvailable := make(chan bool) + replacementDone := make(chan registerResult, 1) + go func() { + if runMu.TryLock() { + runMu.Unlock() + gateAvailable <- true + } else { + gateAvailable <- false + } + ctl, _, registerErr := registerLifecycleTestControl(svr) + replacementDone <- registerResult{ctl: ctl, err: registerErr} + }() + available := waitForResult(t, gateAvailable, "visitor replacement gate probe") + require.False(t, available, "same-run gate was available during visitor admission") + select { + case result := <-replacementDone: + t.Fatalf("replacement completed during visitor admission: %v", result.err) + case <-time.After(20 * time.Millisecond): + } + require.Same(t, ctlA, currentControlForTest(svr.ctlManager, "shared-run")) + + resume() + admission := waitForResult(t, admissionDone, "visitor admission") + require.NoError(t, admission.err) + require.True(t, admission.admitted) + require.Equal(t, "old-user", admission.user) + require.Equal(t, wire.ProtocolV1, admission.wireProtocol) + require.Empty(t, admission.udpPacketCodec) + replacement := waitForResult(t, replacementDone, "replacement") + require.NoError(t, replacement.err) + ctlB := replacement.ctl + require.Same(t, ctlB, currentControlForTest(svr.ctlManager, "shared-run")) + waitForControlDone(t, ctlA) + require.True(t, svr.ctlManager.Remove(ctlB)) + require.NoError(t, ctlB.Close()) +} + +func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) { + svr := newControlTestService(t) + ctl, controlConn, err := registerLifecycleTestControl(svr) + require.NoError(t, err) + + pendingConn := newCountingCloseConn() + pendingMsgConn := msg.NewConn(pendingConn, msg.NewV1ReadWriter(pendingConn)) + err = registerWorkConnAsCaller(svr, pendingMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false) + require.Error(t, err) + require.Equal(t, int64(1), pendingConn.closeCount.Load()) + require.Len(t, ctl.workConnCh, 0) + + require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "control reader to start") + current, ok := svr.ctlManager.GetByID("shared-run") + require.True(t, ok) + require.Same(t, ctl, current) + require.Len(t, ctl.workConnCh, 0) + + runningConn := newCountingCloseConn() + runningMsgConn := msg.NewConn(runningConn, msg.NewV1ReadWriter(runningConn)) + require.NoError(t, svr.RegisterWorkConn(runningMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false)) + require.Len(t, ctl.workConnCh, 1) + + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) + require.Equal(t, int64(1), runningConn.closeCount.Load()) +} + +func TestServiceWorkConnRoutingRejectsWireProtocolMismatch(t *testing.T) { + svr := newControlTestService(t) + ctl, controlConn, err := registerLifecycleTestControl(svr) + require.NoError(t, err) + require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "control reader to start") + + workConn := newCountingCloseConn() + workMsgConn := msg.NewConn(workConn, msg.NewV2ReadWriter(workConn)) + err = svr.RegisterWorkConn(workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV2, false) + require.ErrorContains(t, err, "wire protocol mismatch") + require.Len(t, ctl.workConnCh, 0) + _ = workMsgConn.Close() + require.NoError(t, ctl.Close()) +} + +func TestServiceWorkConnRoutingClientHelloPolicy(t *testing.T) { + for _, tc := range []struct { + name string + controlUDPPacketCodec string + workClientHelloPresent bool + errorSubstring string + }{ + { + name: "JSON control allows work connection without Hello", + }, + { + name: "binary control allows work connection without Hello", + controlUDPPacketCodec: wire.UDPPacketCodecBinary, + }, + { + name: "JSON control rejects work connection with Hello", + workClientHelloPresent: true, + errorSubstring: "ClientHello is not allowed", + }, + { + name: "binary control rejects work connection with Hello", + controlUDPPacketCodec: wire.UDPPacketCodecBinary, + workClientHelloPresent: true, + errorSubstring: "ClientHello is not allowed", + }, + } { + t.Run(tc.name, func(t *testing.T) { + svr := newControlTestService(t) + controlConn := newDeadlineReadConn() + controlMsgConn := msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn)) + ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{ + RunID: "shared-run", + ClientID: "client", + ClientSpec: msg.ClientSpec{ + AlwaysAuthPass: true, + }, + }, true, wire.ProtocolV2, tc.controlUDPPacketCodec) + require.NoError(t, err) + require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "control reader to start") + + workConn := newCountingCloseConn() + workMsgConn := msg.NewConn(workConn, msg.NewV2ReadWriter(workConn)) + err = svr.RegisterWorkConn( + workMsgConn, + &msg.NewWorkConn{RunID: "shared-run"}, + wire.ProtocolV2, + tc.workClientHelloPresent, + ) + if tc.errorSubstring != "" { + require.ErrorContains(t, err, tc.errorSubstring) + require.Len(t, ctl.workConnCh, 0) + require.NoError(t, workMsgConn.Close()) + } else { + require.NoError(t, err) + require.Len(t, ctl.workConnCh, 1) + } + + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) + require.Equal(t, int64(1), workConn.closeCount.Load()) + }) + } +} + +func TestServiceRegisterControlRejectsInvalidCodecSelection(t *testing.T) { + for _, tc := range []struct { + name string + wireProtocol string + udpPacketCodec string + errorSubstring string + }{ + { + name: "binary codec over v1", + wireProtocol: wire.ProtocolV1, + udpPacketCodec: wire.UDPPacketCodecBinary, + errorSubstring: "requires wire protocol v2", + }, + { + name: "unknown v2 codec", + wireProtocol: wire.ProtocolV2, + udpPacketCodec: "unknown", + errorSubstring: "unsupported UDP packet codec", + }, + { + name: "unknown wire protocol", + wireProtocol: "unknown", + errorSubstring: "unsupported wire protocol", + }, + } { + t.Run(tc.name, func(t *testing.T) { + svr := newControlTestService(t) + conn := newDeadlineReadConn() + msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn)) + ctl, err := svr.RegisterControl(msgConn, &msg.Login{}, true, tc.wireProtocol, tc.udpPacketCodec) + require.Nil(t, ctl) + require.ErrorContains(t, err, tc.errorSubstring) + }) + } +} + +func TestServiceRegisterControlPoolCountBoundaries(t *testing.T) { + for _, tc := range []struct { + name string + poolCount int + wantErr bool + wantPool int + }{ + {name: "less than channel offset", poolCount: -11, wantErr: true}, + {name: "channel offset", poolCount: -10, wantErr: true}, + {name: "negative", poolCount: -1, wantErr: true}, + {name: "zero", poolCount: 0, wantPool: 0}, + {name: "capped by server maximum", poolCount: 10, wantPool: 5}, + {name: "maximum int capped", poolCount: math.MaxInt, wantPool: 5}, + } { + t.Run(tc.name, func(t *testing.T) { + svr := newControlTestService(t) + svr.cfg.Transport.MaxPoolCount = 5 + conn := newDeadlineReadConn() + msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn)) + const timestamp = int64(1) + + ctl, err := svr.RegisterControl(msgConn, &msg.Login{ + RunID: "pool-count-run", + ClientID: "client", + Timestamp: timestamp, + PrivilegeKey: util.GetAuthKey("", timestamp), + PoolCount: tc.poolCount, + }, false, wire.ProtocolV1, "") + if tc.wantErr { + require.Nil(t, ctl) + require.ErrorContains(t, err, "unexpected error when creating new controller") + return + } + + require.NoError(t, err) + require.Equal(t, tc.wantPool, ctl.poolCount) + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) + }) + } +} + +func TestServiceWorkConnRoutingRejectsLostGeneration(t *testing.T) { + for _, action := range []string{"replace", "close"} { + t.Run(action, func(t *testing.T) { + svr := newControlTestService(t) + ctl, controlConn, err := registerLifecycleTestControl(svr) + require.NoError(t, err) + require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "control reader to start") + + barrier := newWorkConnBarrierPlugin() + svr.pluginManager.Register(barrier) + workConn := newCountingCloseConn() + workMsgConn := msg.NewConn(workConn, msg.NewV1ReadWriter(workConn)) + routeDone := make(chan error, 1) + go func() { + routeDone <- registerWorkConnAsCaller(svr, workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false) + }() + waitForSignal(t, barrier.entered, "work connection plugin barrier") + + var replacement *Control + switch action { + case "replace": + replacement, _, err = registerLifecycleTestControl(svr) + require.NoError(t, err) + case "close": + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) + } + + close(barrier.resume) + require.Error(t, waitForResult(t, routeDone, "work connection route to finish")) + require.Equal(t, int64(1), workConn.closeCount.Load()) + require.Len(t, ctl.workConnCh, 0) + + if replacement != nil { + require.Len(t, replacement.workConnCh, 0) + require.True(t, svr.ctlManager.Remove(replacement)) + require.NoError(t, replacement.Close()) + waitForControlDone(t, replacement) + } + }) + } +} + +func TestServiceVisitorRoutingExcludesPendingUser(t *testing.T) { + svr := newControlTestService(t) + listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{"pending-user"}) + require.NoError(t, err) + t.Cleanup(func() { _ = listener.Close() }) + + controlConn := newDeadlineReadConn() + controlMsgConn := msg.NewConn(controlConn, msg.NewV1ReadWriter(controlConn)) + ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{ + RunID: "visitor-run", + User: "pending-user", + ClientID: "visitor-client", + ClientSpec: msg.ClientSpec{ + AlwaysAuthPass: true, + }, + }, true, wire.ProtocolV1, "") + require.NoError(t, err) + + timestamp := time.Now().Unix() + visitorMsg := &msg.NewVisitorConn{ + RunID: "visitor-run", + ProxyName: "visitor", + Timestamp: timestamp, + SignKey: util.GetAuthKey("secret", timestamp), + } + pendingConn := newCountingCloseConn() + err = svr.RegisterVisitorConn(pendingConn, visitorMsg, wire.ProtocolV1) + require.ErrorContains(t, err, "no client control found") + require.NoError(t, pendingConn.Close()) + require.Equal(t, int64(1), pendingConn.closeCount.Load()) + + require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "control reader to start") + runningConn := newCountingCloseConn() + require.NoError(t, svr.RegisterVisitorConn(runningConn, visitorMsg, wire.ProtocolV1)) + accepted, err := listener.Accept() + require.NoError(t, err) + require.NoError(t, accepted.Close()) + require.Equal(t, int64(1), runningConn.closeCount.Load()) + + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) +} + +func TestServiceVisitorRoutingCarriesControlPacketCodec(t *testing.T) { + svr := newControlTestService(t) + listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{"visitor-user"}) + require.NoError(t, err) + t.Cleanup(func() { _ = listener.Close() }) + + controlConn := newDeadlineReadConn() + controlMsgConn := msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn)) + ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{ + RunID: "visitor-binary-run", + User: "visitor-user", + ClientID: "visitor-client", + ClientSpec: msg.ClientSpec{ + AlwaysAuthPass: true, + }, + }, true, wire.ProtocolV2, wire.UDPPacketCodecBinary) + require.NoError(t, err) + + timestamp := time.Now().Unix() + visitorMsg := &msg.NewVisitorConn{ + RunID: "visitor-binary-run", + ProxyName: "visitor", + Timestamp: timestamp, + SignKey: util.GetAuthKey("secret", timestamp), + } + require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "binary visitor control reader to start") + + runningConn := newCountingCloseConn() + require.NoError(t, svr.RegisterVisitorConn(runningConn, visitorMsg, wire.ProtocolV2)) + accepted, err := listener.Accept() + require.NoError(t, err) + metadata, ok := accepted.(interface { + WireProtocol() string + UDPPacketCodec() string + }) + require.True(t, ok) + require.Equal(t, wire.ProtocolV2, metadata.WireProtocol()) + require.Equal(t, wire.UDPPacketCodecBinary, metadata.UDPPacketCodec()) + require.NoError(t, accepted.Close()) + require.Equal(t, int64(1), runningConn.closeCount.Load()) + + mismatchConn := newCountingCloseConn() + err = svr.RegisterVisitorConn(mismatchConn, visitorMsg, wire.ProtocolV1) + require.ErrorContains(t, err, "visitor connection wire protocol mismatch") + require.NoError(t, mismatchConn.Close()) + require.Equal(t, int64(1), mismatchConn.closeCount.Load()) + + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) +} + +func TestServiceVisitorRoutingLegacyFallsBackToJSONPacketCodec(t *testing.T) { + svr := newControlTestService(t) + listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{""}) + require.NoError(t, err) + t.Cleanup(func() { _ = listener.Close() }) + + timestamp := time.Now().Unix() + visitorMsg := &msg.NewVisitorConn{ + ProxyName: "visitor", + Timestamp: timestamp, + SignKey: util.GetAuthKey("secret", timestamp), + } + visitorConn := newCountingCloseConn() + require.NoError(t, svr.RegisterVisitorConn(visitorConn, visitorMsg, wire.ProtocolV2)) + accepted, err := listener.Accept() + require.NoError(t, err) + metadata, ok := accepted.(interface{ UDPPacketCodec() string }) + require.True(t, ok) + require.Empty(t, metadata.UDPPacketCodec()) + require.NoError(t, accepted.Close()) + require.Equal(t, int64(1), visitorConn.closeCount.Load()) +} + +func newControlTestService(t *testing.T) *Service { + t.Helper() + cfg := &v1.ServerConfig{} + cfg.Auth.Method = v1.AuthMethodToken + authRuntime, err := auth.BuildServerAuth(&cfg.Auth) + require.NoError(t, err) + clientRegistry := registry.NewClientRegistry() + return &Service{ + ctlManager: NewControlManager(clientRegistry), + clientRegistry: clientRegistry, + pxyManager: proxy.NewManager(), + pluginManager: plugin.NewManager(), + rc: &controller.ResourceController{ + VisitorManager: visitor.NewManager(), + }, + auth: authRuntime, + cfg: cfg, + } +} + +func registerLifecycleTestControl(svr *Service) (*Control, *deadlineReadConn, error) { + conn := newDeadlineReadConn() + msgConn := msg.NewConn(conn, msg.NewReadWriter(conn, wire.ProtocolV1)) + ctl, err := svr.RegisterControl(msgConn, &msg.Login{ + RunID: "shared-run", + ClientID: "client", + ClientSpec: msg.ClientSpec{ + AlwaysAuthPass: true, + }, + }, true, wire.ProtocolV1, "") + return ctl, conn, err +} + +func waitForDifferentCurrentControl(t *testing.T, manager *ControlManager, runID string, old *Control) *Control { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + if ctl := currentControlForTest(manager, runID); ctl != nil && ctl != old { + return ctl + } + runtime.Gosched() + } + t.Fatalf("timed out waiting for a new current control after ID %d", old.ID()) + return nil +} + +func registerWorkConnAsCaller( + svr *Service, + workConn *msg.Conn, + newMsg *msg.NewWorkConn, + wireProtocol string, + clientHelloPresent bool, +) error { + err := svr.RegisterWorkConn(workConn, newMsg, wireProtocol, clientHelloPresent) + if err != nil { + _ = workConn.Close() + } + return err +} + +func waitForResult[T any](t *testing.T, ch <-chan T, description string) T { + t.Helper() + select { + case result := <-ch: + return result + case <-time.After(3 * time.Second): + t.Fatalf("timed out waiting for %s", description) + var zero T + return zero + } +} + +type workConnBarrierPlugin struct { + entered chan struct{} + resume chan struct{} +} + +func newWorkConnBarrierPlugin() *workConnBarrierPlugin { + return &workConnBarrierPlugin{ + entered: make(chan struct{}), + resume: make(chan struct{}), + } +} + +func (*workConnBarrierPlugin) Name() string { return "work-conn-barrier" } + +func (*workConnBarrierPlugin) IsSupport(op string) bool { return op == plugin.OpNewWorkConn } + +func (p *workConnBarrierPlugin) Handle( + context.Context, + string, + any, +) (*plugin.Response, any, error) { + close(p.entered) + <-p.resume + return &plugin.Response{Unchange: true}, nil, nil +} + +type countingCloseConn struct { + closeCount atomic.Int64 +} + +func newCountingCloseConn() *countingCloseConn { return &countingCloseConn{} } + +func (*countingCloseConn) Read([]byte) (int, error) { return 0, net.ErrClosed } +func (*countingCloseConn) Write(p []byte) (int, error) { return len(p), nil } +func (c *countingCloseConn) Close() error { c.closeCount.Add(1); return nil } +func (*countingCloseConn) LocalAddr() net.Addr { return lifecycleTestAddr("local") } +func (*countingCloseConn) RemoteAddr() net.Addr { return lifecycleTestAddr("remote") } +func (*countingCloseConn) SetDeadline(time.Time) error { return nil } +func (*countingCloseConn) SetReadDeadline(time.Time) error { return nil } +func (*countingCloseConn) SetWriteDeadline(time.Time) error { return nil } diff --git a/server/visitor/visitor.go b/server/visitor/visitor.go index 11a44367..7776aad8 100644 --- a/server/visitor/visitor.go +++ b/server/visitor/visitor.go @@ -65,8 +65,12 @@ 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, + wireProtocol string, udpPacketCodecs ...string, ) (err error) { + udpPacketCodec := "" + if len(udpPacketCodecs) > 0 { + udpPacketCodec = udpPacketCodecs[0] + } vm.mu.RLock() defer vm.mu.RUnlock() @@ -93,8 +97,9 @@ 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, + Conn: visitorConn, + wireProtocol: wireProtocol, + udpPacketCodec: udpPacketCodec, }) } else { err = fmt.Errorf("custom listener for [%s] doesn't exist", name) @@ -105,13 +110,18 @@ func (vm *Manager) NewConn(name string, conn net.Conn, timestamp int64, signKey type wireProtocolConn struct { net.Conn - wireProtocol string + wireProtocol string + udpPacketCodec 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() diff --git a/server/visitor/visitor_test.go b/server/visitor/visitor_test.go index 67b907fd..04c50ff0 100644 --- a/server/visitor/visitor_test.go +++ b/server/visitor/visitor_test.go @@ -25,7 +25,7 @@ import ( "github.com/fatedier/frp/pkg/util/util" ) -func TestManagerNewConnCarriesWireProtocol(t *testing.T) { +func TestManagerNewConnCarriesWireProtocolAndUDPPacketCodec(t *testing.T) { vm := NewManager() listener, err := vm.Listen("sudp", "secret", []string{"*"}) require.NoError(t, err) @@ -47,6 +47,7 @@ func TestManagerNewConnCarriesWireProtocol(t *testing.T) { false, "user", wire.ProtocolV2, + wire.UDPPacketCodecBinary, ) }() @@ -54,8 +55,12 @@ func TestManagerNewConnCarriesWireProtocol(t *testing.T) { require.NoError(t, err) defer acceptedConn.Close() - getter, ok := acceptedConn.(interface{ WireProtocol() string }) + metadata, ok := acceptedConn.(interface { + WireProtocol() string + UDPPacketCodec() string + }) require.True(t, ok) - require.Equal(t, wire.ProtocolV2, getter.WireProtocol()) + require.Equal(t, wire.ProtocolV2, metadata.WireProtocol()) + require.Equal(t, wire.UDPPacketCodecBinary, metadata.UDPPacketCodec()) require.NoError(t, <-errCh) } diff --git a/test/e2e/compatibility/compatibility_test.go b/test/e2e/compatibility/compatibility_test.go index 246bcac1..ad505510 100644 --- a/test/e2e/compatibility/compatibility_test.go +++ b/test/e2e/compatibility/compatibility_test.go @@ -192,6 +192,41 @@ 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" @@ -208,6 +243,22 @@ 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(): diff --git a/test/e2e/framework/process.go b/test/e2e/framework/process.go index 8d27e887..ffbec33b 100644 --- a/test/e2e/framework/process.go +++ b/test/e2e/framework/process.go @@ -94,11 +94,9 @@ func (f *Framework) RunProcessesWithBinaries( } func (f *Framework) RunFrps(args ...string) (*process.Process, string, error) { - p := process.NewWithEnvs(TestContext.FRPServerPath, args, f.osEnvs) - f.serverProcesses = append(f.serverProcesses, p) - err := p.Start() + p, output, err := f.StartFrps(args...) if err != nil { - return p, p.Output(), err + return p, output, err } select { case <-p.Done(): @@ -107,17 +105,39 @@ func (f *Framework) RunFrps(args ...string) (*process.Process, string, error) { return p, p.Output(), nil } +// StartFrps starts frps without an implicit sleep so tests can wait on an +// explicit readiness event. +func (f *Framework) StartFrps(args ...string) (*process.Process, string, error) { + p := process.NewWithEnvs(TestContext.FRPServerPath, args, f.osEnvs) + f.serverProcesses = append(f.serverProcesses, p) + err := p.Start() + if err != nil { + return p, p.Output(), err + } + return p, p.Output(), nil +} + func (f *Framework) RunFrpc(args ...string) (*process.Process, string, error) { + p, output, err := f.StartFrpc(args...) + if err != nil { + return p, output, err + } + select { + case <-p.Done(): + case <-time.After(1500 * time.Millisecond): + } + return p, p.Output(), nil +} + +// StartFrpc starts frpc without an implicit sleep so tests can wait on an +// explicit login or proxy-readiness event. +func (f *Framework) StartFrpc(args ...string) (*process.Process, string, error) { p := process.NewWithEnvs(TestContext.FRPClientPath, args, f.osEnvs) f.clientProcesses = append(f.clientProcesses, p) err := p.Start() if err != nil { return p, p.Output(), err } - select { - case <-p.Done(): - case <-time.After(1500 * time.Millisecond): - } return p, p.Output(), nil } diff --git a/test/e2e/pkg/relay/halfopen.go b/test/e2e/pkg/relay/halfopen.go new file mode 100644 index 00000000..34389e53 --- /dev/null +++ b/test/e2e/pkg/relay/halfopen.go @@ -0,0 +1,208 @@ +// Copyright 2026 The frp Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package relay + +import ( + "fmt" + "io" + "net" + "strconv" + "sync" + "time" +) + +// HalfOpen forwards TCP connections until Blackhole is called. A blackholed +// pair stops forwarding but deliberately retains the upstream socket so the +// peer sees a real half-open connection until the relay is closed. +type HalfOpen struct { + bindAddr string + bindPort int + upstreamAddr string + + listener net.Listener + done chan struct{} + accepted chan struct{} + + mu sync.Mutex + pairs []*connectionPair + + wg sync.WaitGroup + closeOnce sync.Once +} + +type connectionPair struct { + downstream net.Conn + upstream net.Conn + + mu sync.Mutex + blackholed bool +} + +func New(upstreamAddr string) *HalfOpen { + return &HalfOpen{ + bindAddr: "127.0.0.1", + upstreamAddr: upstreamAddr, + done: make(chan struct{}), + accepted: make(chan struct{}, 1), + } +} + +func (r *HalfOpen) Run() error { + listener, err := net.Listen("tcp", net.JoinHostPort(r.bindAddr, strconv.Itoa(r.bindPort))) + if err != nil { + return err + } + r.listener = listener + r.bindPort = listener.Addr().(*net.TCPAddr).Port + + r.wg.Add(1) + go r.acceptLoop() + return nil +} + +func (r *HalfOpen) acceptLoop() { + defer r.wg.Done() + for { + downstream, err := r.listener.Accept() + if err != nil { + return + } + upstream, err := net.DialTimeout("tcp", r.upstreamAddr, 3*time.Second) + if err != nil { + _ = downstream.Close() + continue + } + + pair := &connectionPair{downstream: downstream, upstream: upstream} + r.mu.Lock() + r.pairs = append(r.pairs, pair) + r.mu.Unlock() + select { + case r.accepted <- struct{}{}: + default: + } + + r.wg.Add(1) + go r.servePair(pair) + } +} + +func (r *HalfOpen) servePair(pair *connectionPair) { + defer r.wg.Done() + copyDone := make(chan struct{}, 2) + go func() { + _, _ = io.Copy(pair.upstream, pair.downstream) + copyDone <- struct{}{} + }() + go func() { + _, _ = io.Copy(pair.downstream, pair.upstream) + copyDone <- struct{}{} + }() + + completed := 0 + select { + case <-copyDone: + completed = 1 + if !pair.isBlackholed() { + _ = pair.downstream.Close() + _ = pair.upstream.Close() + } + case <-r.done: + _ = pair.downstream.Close() + _ = pair.upstream.Close() + } + for completed < 2 { + <-copyDone + completed++ + } + + if pair.isBlackholed() { + <-r.done + _ = pair.downstream.Close() + _ = pair.upstream.Close() + } +} + +func (r *HalfOpen) WaitForConnections(count int, timeout time.Duration) error { + timer := time.NewTimer(timeout) + defer timer.Stop() + for { + r.mu.Lock() + accepted := len(r.pairs) + r.mu.Unlock() + if accepted >= count { + return nil + } + select { + case <-r.accepted: + case <-r.done: + return fmt.Errorf("relay closed after accepting %d of %d connections", accepted, count) + case <-timer.C: + return fmt.Errorf("timed out after accepting %d of %d connections", accepted, count) + } + } +} + +// Blackhole uses a one-based connection index in accept order. +func (r *HalfOpen) Blackhole(index int) error { + r.mu.Lock() + if index <= 0 || index > len(r.pairs) { + accepted := len(r.pairs) + r.mu.Unlock() + return fmt.Errorf("connection %d is unavailable; accepted %d", index, accepted) + } + pair := r.pairs[index-1] + r.mu.Unlock() + + pair.mu.Lock() + if pair.blackholed { + pair.mu.Unlock() + return nil + } + pair.blackholed = true + pair.mu.Unlock() + + now := time.Now() + _ = pair.downstream.SetDeadline(now) + _ = pair.upstream.SetDeadline(now) + return nil +} + +func (r *HalfOpen) Close() error { + r.closeOnce.Do(func() { + close(r.done) + if r.listener != nil { + _ = r.listener.Close() + } + r.mu.Lock() + pairs := append([]*connectionPair(nil), r.pairs...) + r.mu.Unlock() + for _, pair := range pairs { + _ = pair.downstream.Close() + _ = pair.upstream.Close() + } + r.wg.Wait() + }) + return nil +} + +func (r *HalfOpen) BindAddr() string { return r.bindAddr } +func (r *HalfOpen) BindPort() int { return r.bindPort } + +func (p *connectionPair) isBlackholed() bool { + p.mu.Lock() + defer p.mu.Unlock() + return p.blackholed +} diff --git a/test/e2e/v1/basic/wire.go b/test/e2e/v1/basic/wire.go index 051e6a6c..bf475287 100644 --- a/test/e2e/v1/basic/wire.go +++ b/test/e2e/v1/basic/wire.go @@ -5,6 +5,7 @@ import ( "fmt" "io" "net/http" + "time" "github.com/onsi/ginkgo/v2" @@ -109,12 +110,12 @@ var _ = ginkgo.Describe("[Feature: WireProtocol]", func() { name: "default sudp visitor", }, { - name: "v2 sudp visitor", + name: "v2 binary raw sudp visitor", proxyWireConfig: `transport.wireProtocol = "v2"`, visitorWireConfig: `transport.wireProtocol = "v2"`, }, { - name: "mixed sudp proxy v1 visitor v2", + name: "v1 JSON proxy -> v2 Binary visitor transcode", proxyWireConfig: `transport.wireProtocol = "v1"`, visitorWireConfig: `transport.wireProtocol = "v2"`, extraProxyConfig: ` @@ -127,7 +128,7 @@ var _ = ginkgo.Describe("[Feature: WireProtocol]", func() { `, }, { - name: "mixed sudp proxy v2 visitor v1", + name: "v2 Binary proxy -> v1 JSON visitor transcode", proxyWireConfig: `transport.wireProtocol = "v2"`, visitorWireConfig: `transport.wireProtocol = "v1"`, }, @@ -204,6 +205,87 @@ 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"` diff --git a/test/e2e/v1/features/control_replacement.go b/test/e2e/v1/features/control_replacement.go new file mode 100644 index 00000000..dbb95f41 --- /dev/null +++ b/test/e2e/v1/features/control_replacement.go @@ -0,0 +1,300 @@ +// Copyright 2026 The frp Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package features + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "strconv" + "strings" + "time" + + "github.com/onsi/ginkgo/v2" + + "github.com/fatedier/frp/test/e2e/framework" + "github.com/fatedier/frp/test/e2e/framework/consts" + "github.com/fatedier/frp/test/e2e/pkg/relay" + "github.com/fatedier/frp/test/e2e/pkg/request" +) + +var _ = ginkgo.Describe("[Feature: ControlReplacement]", func() { + f := framework.NewDefaultFramework() + + for _, wireProtocol := range []string{"v1", "v2"} { + for _, tcpMux := range []bool{true, false} { + ginkgo.It(fmt.Sprintf("recovers a %s control through a half-open relay with tcpMux=%t", wireProtocol, tcpMux), func() { + runHalfOpenControlReplacement(f, wireProtocol, tcpMux) + }) + } + } +}) + +func runHalfOpenControlReplacement(f *framework.Framework, wireProtocol string, tcpMux bool) { + serverPort := f.AllocPort() + dashboardPort := f.AllocPort() + remotePort := f.AllocPort() + heartbeatTimeout := int64(-1) + if !tcpMux { + heartbeatTimeout = 3 + } + + serverConfig := fmt.Sprintf(` +bindAddr = "127.0.0.1" +bindPort = %d +log.level = "trace" +transport.tcpMux = %t +transport.tcpMuxKeepaliveInterval = 30 +transport.heartbeatTimeout = %d +webServer.addr = "127.0.0.1" +webServer.port = %d +webServer.pprofEnable = true +enablePrometheus = true +`, serverPort, tcpMux, heartbeatTimeout, dashboardPort) + serverConfigPath := f.WriteTempFile("issue-5391-frps.toml", serverConfig) + serverProcess, _, err := f.StartFrps("-c", serverConfigPath) + framework.ExpectNoError(err) + framework.ExpectNoError(framework.WaitForTCPReady(fmt.Sprintf("127.0.0.1:%d", serverPort), 5*time.Second)) + + halfOpenRelay := relay.New(fmt.Sprintf("127.0.0.1:%d", serverPort)) + f.RunServer("", halfOpenRelay) + + heartbeatInterval := int64(-1) + clientHeartbeatTimeout := int64(-1) + if !tcpMux { + heartbeatInterval = 1 + clientHeartbeatTimeout = 3 + } + clientConfig := fmt.Sprintf(` +serverAddr = "127.0.0.1" +serverPort = %d +clientID = "issue-5391" +loginFailExit = false +log.level = "trace" +transport.wireProtocol = %q +transport.tcpMux = %t +transport.tcpMuxKeepaliveInterval = 1 +transport.heartbeatInterval = %d +transport.heartbeatTimeout = %d +transport.tls.enable = false + +[[proxies]] +name = "issue-5391-tcp" +type = "tcp" +localPort = %d +remotePort = %d +`, halfOpenRelay.BindPort(), wireProtocol, tcpMux, heartbeatInterval, clientHeartbeatTimeout, + f.PortByName(framework.TCPEchoServerPort), remotePort) + clientConfigPath := f.WriteTempFile("issue-5391-frpc.toml", clientConfig) + clientProcess, _, err := f.StartFrpc("-c", clientConfigPath) + framework.ExpectNoError(err) + framework.ExpectNoError(halfOpenRelay.WaitForConnections(1, 5*time.Second)) + framework.ExpectNoError(clientProcess.WaitForOutput("login to server success", 1, 10*time.Second)) + framework.ExpectNoError(clientProcess.WaitForOutput("[issue-5391-tcp] start proxy success", 1, 10*time.Second)) + framework.ExpectNoError(waitForReplacementState(dashboardPort, remotePort, 10*time.Second)) + + replacementCount := 1 + if tcpMux { + replacementCount = 3 + } + for i := 0; i < replacementCount; i++ { + connectionIndex := 1 + if tcpMux { + connectionIndex = i + 1 + } + framework.ExpectNoError(halfOpenRelay.Blackhole(connectionIndex)) + if tcpMux { + framework.ExpectNoError(halfOpenRelay.WaitForConnections(connectionIndex+1, 15*time.Second)) + } + framework.ExpectNoError(clientProcess.WaitForOutput("login to server success", i+2, 15*time.Second)) + framework.ExpectNoError(clientProcess.WaitForOutput("[issue-5391-tcp] start proxy success", i+2, 15*time.Second)) + framework.ExpectNoError(waitForReplacementState(dashboardPort, remotePort, 10*time.Second)) + } + + _ = clientProcess.Stop() + select { + case <-clientProcess.Done(): + case <-time.After(5 * time.Second): + framework.Failf("frpc did not exit") + } + framework.ExpectNoError(waitForReplacementShutdown(dashboardPort, 10*time.Second)) + framework.ExpectNoError(halfOpenRelay.Close()) + framework.ExpectNoError(waitForNoHandoffWaiters(dashboardPort, 5*time.Second)) + + _ = serverProcess.Stop() + select { + case <-serverProcess.Done(): + case <-time.After(5 * time.Second): + framework.Failf("frps did not exit") + } +} + +func waitForReplacementState(dashboardPort, remotePort int, timeout time.Duration) error { + return waitForLifecycleCondition(timeout, func() error { + metricsBody, err := getLifecycleEndpoint(dashboardPort, "/metrics") + if err != nil { + return err + } + if err := expectMetricValue(metricsBody, "frp_server_client_counts", "", 1); err != nil { + return err + } + if err := expectMetricValue(metricsBody, "frp_server_proxy_counts", `type="tcp"`, 1); err != nil { + return err + } + clients, err := getOnlineLifecycleClients(dashboardPort) + if err != nil { + return err + } + if len(clients) != 1 || clients[0].ClientID != "issue-5391" { + return fmt.Errorf("expected one online client, got %+v", clients) + } + profile, err := getLifecycleEndpoint(dashboardPort, "/debug/pprof/goroutine?debug=2") + if err != nil { + return err + } + if err := expectNoHandoffWaiter(profile, "after replacement"); err != nil { + return err + } + resp, err := request.New(). + TCP(). + Port(remotePort). + Timeout(time.Second). + Body([]byte(consts.TestString)). + Do() + if err != nil { + return err + } + if string(resp.Content) != consts.TestString { + return fmt.Errorf("unexpected proxy response %q", resp.Content) + } + return nil + }) +} + +func waitForReplacementShutdown(dashboardPort int, timeout time.Duration) error { + return waitForLifecycleCondition(timeout, func() error { + metricsBody, err := getLifecycleEndpoint(dashboardPort, "/metrics") + if err != nil { + return err + } + if err := expectMetricValue(metricsBody, "frp_server_client_counts", "", 0); err != nil { + return err + } + if err := expectMetricValue(metricsBody, "frp_server_proxy_counts", `type="tcp"`, 0); err != nil { + return err + } + clients, err := getOnlineLifecycleClients(dashboardPort) + if err != nil { + return err + } + if len(clients) != 0 { + return fmt.Errorf("expected no online clients, got %+v", clients) + } + return nil + }) +} + +func waitForNoHandoffWaiters(dashboardPort int, timeout time.Duration) error { + return waitForLifecycleCondition(timeout, func() error { + profile, err := getLifecycleEndpoint(dashboardPort, "/debug/pprof/goroutine?debug=2") + if err != nil { + return err + } + return expectNoHandoffWaiter(profile, "after relay shutdown") + }) +} + +type lifecycleClient struct { + ClientID string `json:"clientID"` +} + +func expectNoHandoffWaiter(profile, phase string) error { + if strings.Contains(profile, "(*Control).WaitForHandoff") { + return fmt.Errorf("control handoff waiter remained %s", phase) + } + return nil +} + +func getOnlineLifecycleClients(dashboardPort int) ([]lifecycleClient, error) { + body, err := getLifecycleEndpoint(dashboardPort, "/api/clients?status=online") + if err != nil { + return nil, err + } + var clients []lifecycleClient + if err := json.Unmarshal([]byte(body), &clients); err != nil { + return nil, err + } + return clients, nil +} + +func getLifecycleEndpoint(port int, path string) (string, error) { + client := &http.Client{Timeout: time.Second} + resp, err := client.Get(fmt.Sprintf("http://127.0.0.1:%d%s", port, path)) + if err != nil { + return "", err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("GET %s returned %s", path, resp.Status) + } + body, err := io.ReadAll(resp.Body) + if err != nil { + return "", err + } + return string(body), nil +} + +func expectMetricValue(body, name, labels string, want float64) error { + prefix := name + if labels != "" { + prefix += "{" + labels + "}" + } + for line := range strings.SplitSeq(body, "\n") { + fields := strings.Fields(line) + if len(fields) != 2 || fields[0] != prefix { + continue + } + got, err := strconv.ParseFloat(fields[1], 64) + if err != nil { + return err + } + if got != want { + return fmt.Errorf("metric %s = %v, want %v", prefix, got, want) + } + return nil + } + return fmt.Errorf("metric %s not found", prefix) +} + +func waitForLifecycleCondition(timeout time.Duration, condition func() error) error { + timer := time.NewTimer(timeout) + defer timer.Stop() + ticker := time.NewTicker(25 * time.Millisecond) + defer ticker.Stop() + var lastErr error + for { + err := condition() + if err == nil { + return nil + } + lastErr = err + select { + case <-ticker.C: + case <-timer.C: + return fmt.Errorf("condition was not met: %w", lastErr) + } + } +} diff --git a/test/e2e/v1/plugin/client.go b/test/e2e/v1/plugin/client.go index 1ea25a6e..f2afad41 100644 --- a/test/e2e/v1/plugin/client.go +++ b/test/e2e/v1/plugin/client.go @@ -1,18 +1,24 @@ package plugin import ( + "bufio" "crypto/tls" "fmt" + "io" + "net" "net/http" "strconv" "strings" "github.com/onsi/ginkgo/v2" + pp "github.com/pires/go-proxyproto" "github.com/fatedier/frp/pkg/transport" + "github.com/fatedier/frp/pkg/util/log" "github.com/fatedier/frp/test/e2e/framework" "github.com/fatedier/frp/test/e2e/framework/consts" "github.com/fatedier/frp/test/e2e/mock/server/httpserver" + "github.com/fatedier/frp/test/e2e/mock/server/streamserver" "github.com/fatedier/frp/test/e2e/pkg/cert" "github.com/fatedier/frp/test/e2e/pkg/port" "github.com/fatedier/frp/test/e2e/pkg/request" @@ -450,4 +456,85 @@ var _ = ginkgo.Describe("[Feature: Client-Plugins]", func() { ExpectResp([]byte("test")). Ensure() }) + + ginkgo.It("tls2raw with proxy protocol v2", func() { + generator := &cert.SelfSignedCertGenerator{} + artifacts, err := generator.Generate("example.com") + framework.ExpectNoError(err) + crtPath := f.WriteTempFile("tls2raw_proxy_protocol_server.crt", string(artifacts.Cert)) + keyPath := f.WriteTempFile("tls2raw_proxy_protocol_server.key", string(artifacts.Key)) + + serverConf := consts.DefaultServerConfig + vhostHTTPSPort := f.AllocPort() + serverConf += fmt.Sprintf(` + vhostHTTPSPort = %d + `, vhostHTTPSPort) + + localPort := f.AllocPort() + clientConf := consts.DefaultClientConfig + fmt.Sprintf(` + [[proxies]] + name = "tls2raw-proxy-protocol-test" + type = "https" + customDomains = ["example.com"] + transport.proxyProtocolVersion = "v2" + [proxies.plugin] + type = "tls2raw" + localAddr = "127.0.0.1:%d" + crtPath = "%s" + keyPath = "%s" + `, localPort, crtPath, keyPath) + + f.RunProcesses(serverConf, []string{clientConf}) + + localServer := streamserver.New(streamserver.TCP, streamserver.WithBindPort(localPort), + streamserver.WithCustomHandler(func(c net.Conn) { + defer c.Close() + + writeResp := func(body string) { + _, _ = fmt.Fprintf(c, "HTTP/1.1 200 OK\r\nContent-Length: %d\r\nConnection: close\r\n\r\n%s", len(body), body) + } + + rd := bufio.NewReader(c) + ppHeader, err := pp.Read(rd) + if err != nil { + log.Errorf("read proxy protocol error: %v", err) + writeResp("missing proxy protocol") + return + } + if ppHeader.Version != 2 { + log.Errorf("unexpected proxy protocol version: %d", ppHeader.Version) + writeResp("unexpected proxy protocol version") + return + } + srcAddr, ok := ppHeader.SourceAddr.(*net.TCPAddr) + if !ok || srcAddr.IP.String() != "127.0.0.1" { + log.Errorf("unexpected proxy protocol source address: %v", ppHeader.SourceAddr) + writeResp("unexpected proxy protocol source address") + return + } + + req, err := http.ReadRequest(rd) + if err != nil { + log.Errorf("read http request after proxy protocol error: %v", err) + writeResp("missing http request") + return + } + _, _ = io.Copy(io.Discard, req.Body) + _ = req.Body.Close() + + writeResp("test") + })) + f.RunServer("", localServer) + + framework.NewRequestExpect(f). + Port(vhostHTTPSPort). + RequestModify(func(r *request.Request) { + r.HTTPS().HTTPHost("example.com").TLSConfig(&tls.Config{ + ServerName: "example.com", + InsecureSkipVerify: true, + }) + }). + ExpectResp([]byte("test")). + Ensure() + }) }) diff --git a/web/frpc/auto-imports.d.ts b/web/frpc/auto-imports.d.ts index 1d89ee8c..9d240079 100644 --- a/web/frpc/auto-imports.d.ts +++ b/web/frpc/auto-imports.d.ts @@ -3,6 +3,7 @@ // @ts-nocheck // noinspection JSUnusedGlobalSymbols // Generated by unplugin-auto-import +// biome-ignore lint: disable export {} declare global { diff --git a/web/frpc/components.d.ts b/web/frpc/components.d.ts index c58d0bca..262d5831 100644 --- a/web/frpc/components.d.ts +++ b/web/frpc/components.d.ts @@ -1,10 +1,14 @@ /* 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'] @@ -38,10 +42,11 @@ 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 ComponentCustomProperties { + export interface GlobalDirectives { vLoading: typeof import('element-plus/es')['ElLoadingDirective'] } } diff --git a/web/frpc/package.json b/web/frpc/package.json index 6e20d893..17e8c594 100644 --- a/web/frpc/package.json +++ b/web/frpc/package.json @@ -9,13 +9,14 @@ "preview": "vite preview", "build-only": "vite build", "type-check": "vue-tsc --noEmit", - "lint": "eslint --fix" + "lint": "eslint . --fix", + "lint:check": "eslint ." }, "dependencies": { - "element-plus": "^2.13.0", + "element-plus": "^2.14.3", "pinia": "^3.0.4", - "vue": "^3.5.26", - "vue-router": "^4.6.4" + "vue": "^3.5.40", + "vue-router": "^5.2.0" }, "devDependencies": { "@types/node": "24", @@ -23,19 +24,19 @@ "@vue/eslint-config-prettier": "^10.2.0", "@vue/eslint-config-typescript": "^14.7.0", "@vue/tsconfig": "^0.8.1", - "@vueuse/core": "^14.1.0", - "eslint": "^9.39.0", - "eslint-plugin-vue": "^9.33.0", + "@vueuse/core": "^14.3.0", + "eslint": "^10.8.0", + "eslint-plugin-vue": "^10.10.0", "npm-run-all": "^4.1.5", - "prettier": "^3.7.4", - "sass": "^1.97.2", - "terser": "^5.44.1", + "prettier": "^3.9.6", + "sass": "^1.102.0", + "terser": "^5.49.0", "typescript": "^5.9.3", - "unplugin-auto-import": "^0.17.5", + "unplugin-auto-import": "^21.0.0", "unplugin-element-plus": "^0.11.2", - "unplugin-vue-components": "^0.26.0", + "unplugin-vue-components": "^32.1.0", "vite": "^7.3.0", "vite-svg-loader": "^5.1.0", - "vue-tsc": "^3.2.2" + "vue-tsc": "^3.3.8" } } diff --git a/web/frpc/src/components/ConfigField.vue b/web/frpc/src/components/ConfigField.vue index 27719847..67be4ba0 100644 --- a/web/frpc/src/components/ConfigField.vue +++ b/web/frpc/src/components/ConfigField.vue @@ -12,7 +12,7 @@ diff --git a/web/frpc/env.d.ts b/web/frpc/src/env.d.ts similarity index 100% rename from web/frpc/env.d.ts rename to web/frpc/src/env.d.ts diff --git a/web/frpc/src/types/proxy-converters.ts b/web/frpc/src/types/proxy-converters.ts index 6587986e..eb4c4c17 100644 --- a/web/frpc/src/types/proxy-converters.ts +++ b/web/frpc/src/types/proxy-converters.ts @@ -190,6 +190,16 @@ 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 @@ -448,6 +458,15 @@ 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 diff --git a/web/frpc/src/types/proxy-form.ts b/web/frpc/src/types/proxy-form.ts index 8f6bd614..793fe516 100644 --- a/web/frpc/src/types/proxy-form.ts +++ b/web/frpc/src/types/proxy-form.ts @@ -79,6 +79,10 @@ export interface VisitorFormData { bindAddr: string bindPort: number | undefined + // Visitor plugin + pluginType: '' | 'virtual_net' + pluginDestinationIP: string + // XTCP specific (XTCPVisitorConfig) protocol: string keepTunnelOpen: boolean @@ -156,6 +160,9 @@ export function createDefaultVisitorForm(): VisitorFormData { bindAddr: '127.0.0.1', bindPort: undefined, + pluginType: '', + pluginDestinationIP: '', + protocol: 'quic', keepTunnelOpen: false, maxRetriesAnHour: undefined, diff --git a/web/frpc/src/views/VisitorEdit.vue b/web/frpc/src/views/VisitorEdit.vue index 615e97e9..f095103d 100644 --- a/web/frpc/src/views/VisitorEdit.vue +++ b/web/frpc/src/views/VisitorEdit.vue @@ -107,6 +107,22 @@ 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 = () => { diff --git a/web/frpc/test/config-field.test.ts b/web/frpc/test/config-field.test.ts new file mode 100644 index 00000000..64986274 --- /dev/null +++ b/web/frpc/test/config-field.test.ts @@ -0,0 +1,73 @@ +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: + '', +}) + +const FormItemStub = defineComponent({ + template: '
', +}) + +const stubs = { + 'el-form-item': FormItemStub, + 'el-input': InputStub, + 'el-switch': defineComponent({ + props: ['modelValue', 'disabled'], + emits: ['update:modelValue'], + template: '', + }), + KeyValueEditor: defineComponent({ template: '
' }), + StringListEditor: defineComponent({ template: '
' }), + PopoverMenu: defineComponent({ template: '
' }), + PopoverMenuItem: defineComponent({ template: '
' }), +} + +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('—') + }) +}) diff --git a/web/frpc/test/proxy-converters.test.ts b/web/frpc/test/proxy-converters.test.ts new file mode 100644 index 00000000..69e3410b --- /dev/null +++ b/web/frpc/test/proxy-converters.test.ts @@ -0,0 +1,96 @@ +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 }) + }) +}) diff --git a/web/frpc/tsconfig.node.json b/web/frpc/tsconfig.node.json index 42872c59..a8583534 100644 --- a/web/frpc/tsconfig.node.json +++ b/web/frpc/tsconfig.node.json @@ -6,5 +6,5 @@ "moduleResolution": "bundler", "allowSyntheticDefaultImports": true }, - "include": ["vite.config.ts"] + "include": ["vite.config.mts"] } diff --git a/web/frpc/vite.config.mts b/web/frpc/vite.config.mts index 6a9205d0..209a91fd 100644 --- a/web/frpc/vite.config.mts +++ b/web/frpc/vite.config.mts @@ -28,15 +28,10 @@ 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 *;`, }, }, diff --git a/web/frps/auto-imports.d.ts b/web/frps/auto-imports.d.ts index 1d89ee8c..9d240079 100644 --- a/web/frps/auto-imports.d.ts +++ b/web/frps/auto-imports.d.ts @@ -3,6 +3,7 @@ // @ts-nocheck // noinspection JSUnusedGlobalSymbols // Generated by unplugin-auto-import +// biome-ignore lint: disable export {} declare global { diff --git a/web/frps/components.d.ts b/web/frps/components.d.ts index b5a18d79..e177245b 100644 --- a/web/frps/components.d.ts +++ b/web/frps/components.d.ts @@ -1,10 +1,14 @@ /* 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'] @@ -15,7 +19,6 @@ declare module 'vue' { ElEmpty: typeof import('element-plus/es')['ElEmpty'] ElIcon: typeof import('element-plus/es')['ElIcon'] ElInput: typeof import('element-plus/es')['ElInput'] - ElPopover: typeof import('element-plus/es')['ElPopover'] ElRow: typeof import('element-plus/es')['ElRow'] ElSwitch: typeof import('element-plus/es')['ElSwitch'] ElTag: typeof import('element-plus/es')['ElTag'] @@ -26,7 +29,7 @@ declare module 'vue' { StatCard: typeof import('./src/components/StatCard.vue')['default'] Traffic: typeof import('./src/components/Traffic.vue')['default'] } - export interface ComponentCustomProperties { + export interface GlobalDirectives { vLoading: typeof import('element-plus/es')['ElLoadingDirective'] } } diff --git a/web/frps/package.json b/web/frps/package.json index c44c0dd4..8cc42d50 100644 --- a/web/frps/package.json +++ b/web/frps/package.json @@ -9,12 +9,13 @@ "preview": "vite preview", "build-only": "vite build", "type-check": "vue-tsc --noEmit", - "lint": "eslint --fix" + "lint": "eslint . --fix", + "lint:check": "eslint ." }, "dependencies": { - "element-plus": "^2.13.0", - "vue": "^3.5.26", - "vue-router": "^4.6.4" + "element-plus": "^2.14.3", + "vue": "^3.5.40", + "vue-router": "^5.2.0" }, "devDependencies": { "@types/node": "24", @@ -22,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.1.0", - "eslint": "^9.39.0", - "eslint-plugin-vue": "^9.33.0", + "@vueuse/core": "^14.3.0", + "eslint": "^10.8.0", + "eslint-plugin-vue": "^10.10.0", "npm-run-all": "^4.1.5", - "prettier": "^3.7.4", - "sass": "^1.97.2", - "terser": "^5.44.1", + "prettier": "^3.9.6", + "sass": "^1.102.0", + "terser": "^5.49.0", "typescript": "^5.9.3", - "unplugin-auto-import": "^0.17.5", + "unplugin-auto-import": "^21.0.0", "unplugin-element-plus": "^0.11.2", - "unplugin-vue-components": "^0.26.0", + "unplugin-vue-components": "^32.1.0", "vite": "^7.3.0", "vite-svg-loader": "^5.1.0", - "vue-tsc": "^3.2.2" + "vue-tsc": "^3.3.8" } } diff --git a/web/frps/src/api/client.ts b/web/frps/src/api/client.ts index 2be7d239..7a3982b7 100644 --- a/web/frps/src/api/client.ts +++ b/web/frps/src/api/client.ts @@ -24,3 +24,7 @@ export const getClientsV2 = (params: ClientListV2Params = {}) => { export const getClient = (key: string) => { return http.get(`../api/clients/${key}`) } + +export const getClientV2 = (key: string) => { + return http.getV2(`../api/v2/clients/${encodeURIComponent(key)}`) +} diff --git a/web/frps/src/api/http.ts b/web/frps/src/api/http.ts index 44100b07..2f59c1b7 100644 --- a/web/frps/src/api/http.ts +++ b/web/frps/src/api/http.ts @@ -98,6 +98,13 @@ export const http = { request(url, { ...options, method: 'GET' }), getV2: (url: string, options?: RequestInit) => requestV2(url, { ...options, method: 'GET' }), + postV2: (url: string, body?: any, options?: RequestInit) => + requestV2(url, { + ...options, + method: 'POST', + headers: { 'Content-Type': 'application/json', ...options?.headers }, + body: JSON.stringify(body), + }), post: (url: string, body?: any, options?: RequestInit) => request(url, { ...options, diff --git a/web/frps/src/api/proxy.ts b/web/frps/src/api/proxy.ts index 20125842..de07cac7 100644 --- a/web/frps/src/api/proxy.ts +++ b/web/frps/src/api/proxy.ts @@ -1,13 +1,23 @@ import { buildQueryString, http } from './http' +import { formatUnixSeconds } from '../utils/format' import type { V2Page } from './http' import type { GetProxyResponse, ProxyListV2Params, ProxyStatsInfo, ProxyV2Info, + ProxyV2Spec, + ProxyV2SpecBlocks, + ProxyV2Type, TrafficResponse, } from '../types/proxy' +export interface SystemPruneResponse { + type: 'offline_proxies' + cleared: number + total: number +} + export const getProxiesByType = (type: string) => { return http.get(`../api/proxy/${type}`) } @@ -32,32 +42,77 @@ export const getProxiesV2 = async (params: ProxyListV2Params = {}) => { } } -const toLegacyProxyStats = (proxy: ProxyV2Info): ProxyStatsInfo => ({ - name: proxy.name, - type: proxy.type, - conf: proxy.spec, - user: proxy.user, - clientID: proxy.clientID, - todayTrafficIn: proxy.status.todayTrafficIn, - todayTrafficOut: proxy.status.todayTrafficOut, - curConns: proxy.status.curConns, - lastStartTime: proxy.status.lastStartTime, - lastCloseTime: proxy.status.lastCloseTime, - status: proxy.status.phase, -}) +const getActiveProxySpec = ( + spec: ProxyV2Spec, +): ProxyV2SpecBlocks[ProxyV2Type] => { + switch (spec.type) { + case 'tcp': + return spec.tcp + case 'udp': + return spec.udp + case 'http': + return spec.http + case 'https': + return spec.https + case 'tcpmux': + return spec.tcpmux + case 'stcp': + return spec.stcp + case 'sudp': + return spec.sudp + case 'xtcp': + return spec.xtcp + default: + return assertNever(spec) + } +} + +const assertNever = (value: never): never => { + throw new Error(`Unsupported proxy spec: ${JSON.stringify(value)}`) +} + +export const toLegacyProxyStats = (proxy: ProxyV2Info): ProxyStatsInfo => { + const type = proxy.spec.type + const activeSpec = getActiveProxySpec(proxy.spec) + + return { + name: proxy.name, + type, + conf: proxy.status.phase === 'offline' ? null : activeSpec, + user: proxy.user, + clientID: proxy.clientID, + todayTrafficIn: proxy.status.todayTrafficIn, + todayTrafficOut: proxy.status.todayTrafficOut, + curConns: proxy.status.curConns, + lastStartTime: formatUnixSeconds(proxy.status.lastStartAt), + lastCloseTime: formatUnixSeconds(proxy.status.lastCloseAt), + status: proxy.status.phase, + } +} export const getProxy = (type: string, name: string) => { return http.get(`../api/proxy/${type}/${name}`) } +export const getProxyByNameV2 = async (name: string) => { + const proxy = await http.getV2( + `../api/v2/proxies/${encodeURIComponent(name)}`, + ) + return toLegacyProxyStats(proxy) +} + export const getProxyByName = (name: string) => { return http.get(`../api/proxies/${name}`) } export const getProxyTraffic = (name: string) => { - return http.get(`../api/traffic/${name}`) + return http.getV2( + `../api/v2/proxies/${encodeURIComponent(name)}/traffic`, + ) } export const clearOfflineProxies = () => { - return http.delete('../api/proxies?status=offline') + return http.postV2( + '../api/v2/system/prune?type=offline_proxies', + ) } diff --git a/web/frps/src/api/server.ts b/web/frps/src/api/server.ts index f46f21d3..f6621ca8 100644 --- a/web/frps/src/api/server.ts +++ b/web/frps/src/api/server.ts @@ -2,5 +2,5 @@ import { http } from './http' import type { ServerInfo } from '../types/server' export const getServerInfo = () => { - return http.get('../api/serverinfo') + return http.getV2('../api/v2/system/info') } diff --git a/web/frps/src/components/Traffic.vue b/web/frps/src/components/Traffic.vue index 9142f3e0..9f4d422c 100644 --- a/web/frps/src/components/Traffic.vue +++ b/web/frps/src/components/Traffic.vue @@ -54,6 +54,7 @@ import { ref, onMounted } from 'vue' import { ElMessage } from 'element-plus' import { formatFileSize } from '../utils/format' import { getProxyTraffic } from '../api/proxy' +import type { TrafficResponse } from '../types/proxy' const props = defineProps<{ proxyName: string @@ -71,41 +72,24 @@ const chartData = ref< >([]) const maxVal = ref(0) -const processData = (trafficIn: number[], trafficOut: number[]) => { - // Ensure we have arrays and reverse them (server returns newest first) - const inArr = [...(trafficIn || [])].reverse() - const outArr = [...(trafficOut || [])].reverse() +const formatDateLabel = (date: string) => { + const parts = date.split('-') + if (parts.length !== 3) return date + return `${Number(parts[1])}-${Number(parts[2])}` +} - // Pad with zeros if less than 7 days - while (inArr.length < 7) inArr.unshift(0) - while (outArr.length < 7) outArr.unshift(0) - - // Slice to last 7 entries just in case - const finalIn = inArr.slice(-7) - const finalOut = outArr.slice(-7) - - // Calculate dates (last 7 days ending today) - const dates: string[] = [] - const d = new Date() - d.setDate(d.getDate() - 6) - - for (let i = 0; i < 7; i++) { - dates.push(`${d.getMonth() + 1}-${d.getDate()}`) - d.setDate(d.getDate() + 1) - } - - // Find max value for scaling - const maxIn = Math.max(...finalIn) - const maxOut = Math.max(...finalOut) +const processData = (history: TrafficResponse['history'] = []) => { + const points = history || [] + const maxIn = Math.max(0, ...points.map((item) => item.trafficIn)) + const maxOut = Math.max(0, ...points.map((item) => item.trafficOut)) maxVal.value = Math.max(maxIn, maxOut, 100) // Minimum scale 100 bytes - // Build chart data - chartData.value = dates.map((date, i) => ({ - date, - in: finalIn[i], - out: finalOut[i], - inPercent: (finalIn[i] / maxVal.value) * 100, - outPercent: (finalOut[i] / maxVal.value) * 100, + chartData.value = points.map((item) => ({ + date: formatDateLabel(item.date), + in: item.trafficIn, + out: item.trafficOut, + inPercent: (item.trafficIn / maxVal.value) * 100, + outPercent: (item.trafficOut / maxVal.value) * 100, })) } @@ -113,7 +97,7 @@ const fetchData = () => { loading.value = true getProxyTraffic(props.proxyName) .then((json) => { - processData(json.trafficIn, json.trafficOut) + processData(json.history) }) .catch((err) => { ElMessage({ diff --git a/web/frps/env.d.ts b/web/frps/src/env.d.ts similarity index 100% rename from web/frps/env.d.ts rename to web/frps/src/env.d.ts diff --git a/web/frps/src/types/client.ts b/web/frps/src/types/client.ts index 96700fe9..94d59098 100644 --- a/web/frps/src/types/client.ts +++ b/web/frps/src/types/client.ts @@ -7,11 +7,17 @@ export interface ClientInfoData { wireProtocol?: string hostname: string clientIP?: string - metas?: Record firstConnectedAt: number lastConnectedAt: number disconnectedAt?: number online: boolean + status?: ClientStatus +} + +export interface ClientStatus { + phase: 'online' | 'offline' + curConns: number + proxyCount: number } export interface ClientListV2Params { diff --git a/web/frps/src/types/proxy.ts b/web/frps/src/types/proxy.ts index 53645d48..27606ef5 100644 --- a/web/frps/src/types/proxy.ts +++ b/web/frps/src/types/proxy.ts @@ -28,24 +28,102 @@ export interface ProxyListV2Params { export interface ProxyV2Info { name: string - type: string user: string clientID: string - spec: any + spec: ProxyV2Spec status: ProxyV2Status } +export interface ProxyV2BaseSpec { + annotations?: Record + metadatas?: Record + transport?: { + useEncryption: boolean + useCompression: boolean + bandwidthLimit: string + bandwidthLimitMode: string + } + loadBalancer?: { + group: string + } +} + +export interface ProxyV2TCPBlock extends ProxyV2BaseSpec { + remotePort?: number +} + +export interface ProxyV2UDPBlock extends ProxyV2BaseSpec { + remotePort?: number +} + +export interface ProxyV2HTTPBlock extends ProxyV2BaseSpec { + customDomains?: string[] + subdomain?: string + locations?: string[] + hostHeaderRewrite?: string +} + +export interface ProxyV2HTTPSBlock extends ProxyV2BaseSpec { + customDomains?: string[] + subdomain?: string +} + +export interface ProxyV2TCPMuxBlock extends ProxyV2BaseSpec { + customDomains?: string[] + subdomain?: string + multiplexer?: string + routeByHTTPUser?: string +} + +export type ProxyV2STCPBlock = ProxyV2BaseSpec + +export type ProxyV2SUDPBlock = ProxyV2BaseSpec + +export type ProxyV2XTCPBlock = ProxyV2BaseSpec + +export interface ProxyV2SpecBlocks { + tcp: ProxyV2TCPBlock + udp: ProxyV2UDPBlock + http: ProxyV2HTTPBlock + https: ProxyV2HTTPSBlock + tcpmux: ProxyV2TCPMuxBlock + stcp: ProxyV2STCPBlock + sudp: ProxyV2SUDPBlock + xtcp: ProxyV2XTCPBlock +} + +export type ProxyV2Type = keyof ProxyV2SpecBlocks + +type ProxyV2SpecFor = { + type: T +} & { + [K in T]: ProxyV2SpecBlocks[K] +} & { + [K in Exclude]?: never +} + +export type ProxyV2Spec = { + [T in ProxyV2Type]: ProxyV2SpecFor +}[ProxyV2Type] + export interface ProxyV2Status { - phase: string + phase: 'online' | 'offline' todayTrafficIn: number todayTrafficOut: number curConns: number - lastStartTime: string - lastCloseTime: string + lastStartAt?: number + lastCloseAt?: number } export interface TrafficResponse { name: string - trafficIn: number[] - trafficOut: number[] + unit: 'bytes' + granularity: 'day' + history: TrafficPoint[] +} + +export interface TrafficPoint { + date: string + trafficIn: number + trafficOut: number } diff --git a/web/frps/src/types/server.ts b/web/frps/src/types/server.ts index ada31cc5..b4796e90 100644 --- a/web/frps/src/types/server.ts +++ b/web/frps/src/types/server.ts @@ -1,5 +1,10 @@ export interface ServerInfo { version: string + config: ServerInfoConfig + status: ServerInfoStatus +} + +export interface ServerInfoConfig { bindPort: number vhostHTTPPort: number vhostHTTPSPort: number @@ -12,8 +17,9 @@ export interface ServerInfo { heartbeatTimeout: number allowPortsStr: string tlsForce: boolean +} - // Stats +export interface ServerInfoStatus { totalTrafficIn: number totalTrafficOut: number curConns: number diff --git a/web/frps/src/utils/client.ts b/web/frps/src/utils/client.ts index 8d26eeb8..0cb9e5a0 100644 --- a/web/frps/src/utils/client.ts +++ b/web/frps/src/utils/client.ts @@ -1,5 +1,5 @@ import { formatDistanceToNow } from './format' -import type { ClientInfoData } from '../types/client' +import type { ClientInfoData, ClientStatus } from '../types/client' export class Client { key: string @@ -10,11 +10,11 @@ export class Client { wireProtocol: string hostname: string ip: string - metas: Map firstConnectedAt: Date lastConnectedAt: Date disconnectedAt?: Date online: boolean + status: ClientStatus constructor(data: ClientInfoData) { this.key = data.key @@ -25,18 +25,17 @@ export class Client { this.wireProtocol = data.wireProtocol || '' this.hostname = data.hostname this.ip = data.clientIP || '' - this.metas = new Map() - if (data.metas) { - for (const [key, value] of Object.entries(data.metas)) { - this.metas.set(key, value) - } - } this.firstConnectedAt = new Date(data.firstConnectedAt * 1000) this.lastConnectedAt = new Date(data.lastConnectedAt * 1000) if (data.disconnectedAt && data.disconnectedAt > 0) { this.disconnectedAt = new Date(data.disconnectedAt * 1000) } this.online = data.online + this.status = data.status || { + phase: this.online ? 'online' : 'offline', + curConns: 0, + proxyCount: 0, + } } get displayName(): string { @@ -46,10 +45,6 @@ export class Client { return this.runID } - get shortRunId(): string { - return this.runID.substring(0, 8) - } - get wireProtocolLabel(): string { if (!this.wireProtocol) return '' return `Protocol ${this.wireProtocol}` @@ -67,28 +62,4 @@ export class Client { if (!this.disconnectedAt) return '' return formatDistanceToNow(this.disconnectedAt) } - - get statusColor(): string { - return this.online ? 'success' : 'danger' - } - - get metasArray(): Array<{ key: string; value: string }> { - const arr: Array<{ key: string; value: string }> = [] - this.metas.forEach((value, key) => { - arr.push({ key, value }) - }) - return arr - } - - matchesFilter(searchText: string): boolean { - const search = searchText.toLowerCase() - return ( - this.key.toLowerCase().includes(search) || - this.user.toLowerCase().includes(search) || - this.clientID.toLowerCase().includes(search) || - this.runID.toLowerCase().includes(search) || - this.wireProtocol.toLowerCase().includes(search) || - this.hostname.toLowerCase().includes(search) - ) - } } diff --git a/web/frps/src/utils/format.ts b/web/frps/src/utils/format.ts index 11cd398f..cca9e31a 100644 --- a/web/frps/src/utils/format.ts +++ b/web/frps/src/utils/format.ts @@ -19,6 +19,14 @@ export function formatDistanceToNow(date: Date): string { return Math.floor(seconds) + ' seconds ago' } +export function formatUnixSeconds(seconds?: number): string { + if (seconds == null || !Number.isFinite(seconds) || seconds <= 0) return '' + + const date = new Date(seconds * 1000) + const pad = (value: number) => value.toString().padStart(2, '0') + return `${pad(date.getMonth() + 1)}-${pad(date.getDate())} ${pad(date.getHours())}:${pad(date.getMinutes())}:${pad(date.getSeconds())}` +} + export function formatFileSize(bytes: number): string { if (!Number.isFinite(bytes) || bytes < 0) return '0 B' if (bytes === 0) return '0 B' diff --git a/web/frps/src/views/ClientDetail.vue b/web/frps/src/views/ClientDetail.vue index 42c9a25e..17a50468 100644 --- a/web/frps/src/views/ClientDetail.vue +++ b/web/frps/src/views/ClientDetail.vue @@ -55,7 +55,7 @@
Connections - {{ totalConnections }} + {{ client.status.curConns }}
Run ID @@ -85,7 +85,7 @@

Proxies

- {{ filteredProxies.length }} + {{ total }}
Loading...
-
+
-
+

No proxies match "{{ proxySearch }}"

No proxies found

+
+ +
@@ -130,13 +141,13 @@ @@ -483,6 +568,12 @@ html.dark .status-badge.online { padding: 16px; } +.pagination-section { + display: flex; + justify-content: center; + padding: 0 20px 20px; +} + .proxies-list { display: flex; flex-direction: column; diff --git a/web/frps/src/views/Proxies.vue b/web/frps/src/views/Proxies.vue index 80548bb8..734c85a5 100644 --- a/web/frps/src/views/Proxies.vue +++ b/web/frps/src/views/Proxies.vue @@ -27,36 +27,6 @@ clearable class="main-search" /> - - - -
@@ -111,7 +81,7 @@