Compare commits

..
77 Commits
Author SHA1 Message Date
4a23aa181c Release v0.71.0 (#5495)
* web/frpc: support virtual net visitor plugin (#5414)

* web/frpc: fix visitor plugin validation for sudp (#5448)

* web: patch vulnerable transitive dependencies (#5449)

* web: fix Vite config tsconfig includes (#5450)

* test(web): add frontend unit test baseline (#5451)

* build(web): upgrade ESLint to v10 (#5452)

* build(web): upgrade auto import plugins (#5453)

* chore(web): upgrade frontend dependencies (#5454)

* feat: add negotiated binary UDP packets (#5456)

* docs: add release note for #5456 (#5458)

* server: validate control pool counts (#5459)

* feat: use binary codec for SUDP packets (#5461)

* update version to 0.71.0 (#5464)

* fix(frpc): respect feature gates in verify (#5465)

* remove retired Go Report Card badge (#5467)

* udp: handle closed forwarding channel (#5470)

* limit: clamp bandwidth limiter burst (#5471)

* client: reject invalid work connection addresses (#5472)

* ssh: serialize tunnel channel writes (#5473)

* config: reject case-insensitive subdomain domains (#5474)

* docs: add release note for domain validation fix (#5475)

* log: improve prefix handling (#5489)

* web: patch vulnerable transitive dependencies (#5490)

---------

Co-authored-by: barkure <43804451+barkure@users.noreply.github.com>
Co-authored-by: Yarden Shoham <git@yardenshoham.com>
2026-08-14 14:13:20 +08:00
fatedierandGitHub 6c8a8d0a97 web: patch vulnerable transitive dependencies (#5490) 2026-08-13 01:12:34 +08:00
fatedierandGitHub da04e1e07a log: improve prefix handling (#5489) 2026-08-12 23:56:00 +08:00
fatedierandGitHub 71a2bf30f9 docs: add release note for domain validation fix (#5475) 2026-08-09 23:04:09 +08:00
fatedierandGitHub a6a782bed4 config: reject case-insensitive subdomain domains (#5474) 2026-08-09 22:52:30 +08:00
fatedierandGitHub 223b44336c ssh: serialize tunnel channel writes (#5473) 2026-08-09 19:12:12 +08:00
fatedierandGitHub f6688e2a0d client: reject invalid work connection addresses (#5472) 2026-08-09 18:39:21 +08:00
fatedierandGitHub f666d97b64 limit: clamp bandwidth limiter burst (#5471) 2026-08-09 17:30:39 +08:00
fatedierandGitHub 5b68148f11 udp: handle closed forwarding channel (#5470) 2026-08-09 15:31:11 +08:00
Yarden ShohamandGitHub d1928f9689 remove retired Go Report Card badge (#5467) 2026-08-05 10:52:16 +08:00
fatedierandGitHub 1a3a872bd2 fix(frpc): respect feature gates in verify (#5465) 2026-08-04 12:53:10 +08:00
fatedierandGitHub 290017cb53 update version to 0.71.0 (#5464) 2026-08-03 23:37:36 +08:00
fatedierandGitHub 2291e8835f feat: use binary codec for SUDP packets (#5461) 2026-07-31 16:53:16 +08:00
fatedierandGitHub 1ab59e763c server: validate control pool counts (#5459) 2026-07-30 23:30:16 +08:00
fatedierandGitHub 9d45a55720 docs: add release note for #5456 (#5458) 2026-07-30 14:40:54 +08:00
fatedierandGitHub effa496859 feat: add negotiated binary UDP packets (#5456) 2026-07-30 13:19:55 +08:00
fatedierandGitHub 5c6d761c12 chore(web): upgrade frontend dependencies (#5454) 2026-07-28 19:36:44 +08:00
fatedierandGitHub 3d89ab7ff1 build(web): upgrade auto import plugins (#5453) 2026-07-28 17:09:46 +08:00
fatedierandGitHub f3828cbfdb build(web): upgrade ESLint to v10 (#5452) 2026-07-28 16:10:29 +08:00
fatedierandGitHub 00c24e39bb test(web): add frontend unit test baseline (#5451) 2026-07-28 01:27:35 +08:00
fatedierandGitHub 2175557d48 web: fix Vite config tsconfig includes (#5450) 2026-07-27 19:21:28 +08:00
fatedierandGitHub de4b483710 web: patch vulnerable transitive dependencies (#5449) 2026-07-27 14:57:58 +08:00
fatedierandGitHub 2d63f6b1f9 web/frpc: fix visitor plugin validation for sudp (#5448) 2026-07-27 12:04:00 +08:00
barkureandGitHub 40274d8c92 web/frpc: support virtual net visitor plugin (#5414) 2026-07-23 22:02:49 +08:00
fatedierandGitHub fa3bcca2b0 Merge pull request #5444 from fatedier/dev
Release v0.70.1
2026-07-23 15:34:31 +08:00
fatedierandGitHub 18eef83b69 fix(client): synchronize graceful shutdown duration (#5442) 2026-07-23 13:55:00 +08:00
fatedierandGitHub 466a94f5d1 update quic-go dependency to v0.60.0 (#5443) 2026-07-23 12:28:10 +08:00
barkureandGitHub 898bcf73d5 web: restore Vite client declarations for both dashboards (#5415)
The 2024 bundler-mode migration changed the dashboard tsconfig include
lists to src/**, leaving the root env.d.ts files outside the TypeScript
programs. Move the Vite client declarations into src/ for both the frpc
and frps dashboards so they are part of the programs again.

Drop the *.vue wildcard module shim: the current vue-tsc toolchain
resolves SFC types without it, and the shim would only mask real
component prop types under plain tsc.
2026-07-23 11:14:56 +08:00
fatedierandGitHub 0e53833a2c refactor: use standard library HKDF (#5440) 2026-07-23 01:17:37 +08:00
fatedierandGitHub c23934bb67 web: update vulnerable transitive dependencies (#5441) 2026-07-23 01:16:40 +08:00
fatedierandGitHub 23512f577e fix: upgrade go-oidc to v3.18.0 (#5439) 2026-07-23 00:00:47 +08:00
fatedierandGitHub 84b907a198 deps: bump golib to v0.8.1 (#5437) 2026-07-22 21:47:36 +08:00
fatedierandGitHub f8fc3c6b1b server: drop HTTP/1.1 h2c upgrade handling (#5436)
Use net/http Server.Protocols and the merged golib listener path instead of the deprecated h2c handler. Keep HTTP/1.1 and cleartext HTTP/2 prior-knowledge support, while intentionally dropping HTTP/1.1 Upgrade: h2c.
2026-07-22 19:55:42 +08:00
fatedierandGitHub 7dc7be930e ssh: fix malformed exec payload panic (#5428) 2026-07-21 23:04:56 +08:00
fatedierandGitHub d486018885 fix(server): prevent control replacement lifecycle leaks (#5424) 2026-07-21 18:09:09 +08:00
fatedierandGitHub fe79598ee4 fix: fail fast on cross-compile errors (#5420) 2026-07-16 16:40:48 +08:00
fatedierandGitHub 269a26c5d5 fix(nathole): migrate STUN client to golib (#5419) 2026-07-16 13:11:45 +08:00
fatedierandGitHub 2886393f5b ci: remove unsupported s390x image target (#5412) 2026-07-11 22:44:45 +08:00
fatedierandGitHub 7b6e01f04f Merge pull request #5411 from fatedier/dev
Release v0.70.0
2026-07-11 18:50:35 +08:00
fatedierandGitHub 5396947c7c remove retired Go Report Card badge (#5409) 2026-07-11 15:59:51 +08:00
fatedierandGitHub fe61093bc2 bump version to 0.70.0 (#5406) 2026-07-11 14:28:19 +08:00
fatedierandGitHub 54aeb2a7b0 feat(server): add typed frps v2 proxy specs (#5405) 2026-07-11 10:34:04 +08:00
fatedierandGitHub 84be1938e4 api: expose v2 proxy timestamps as unix seconds (#5402) 2026-07-09 14:47:19 +08:00
fatedierandGitHub becee40715 docs: update API v2 release notes (#5400) 2026-07-08 15:54:16 +08:00
fatedierandGitHub 17e788d43b refactor: clean up frps v2 frontend models (#5399) 2026-07-08 13:10:55 +08:00
fatedierandGitHub 68509f5d44 Add frps proxy traffic API v2 (#5398) 2026-07-08 02:05:28 +08:00
fatedierandGitHub 5cd722b177 feat: add system prune API v2 (#5395) 2026-07-07 13:00:53 +08:00
fatedierandGitHub 5876beceac feat: add system info API v2 (#5394) 2026-07-07 02:00:44 +08:00
fatedierandGitHub 7fe152e3aa web/frps: use API v2 for client and proxy details (#5386) 2026-06-30 01:33:57 +08:00
fatedierandGitHub 7c343fc6e7 web: bump esbuild to 0.28.1 and remove unused ElPopover type (#5385) 2026-06-29 23:21:22 +08:00
fatedierandGitHub a3b3b35b69 feat(ui): default proxies view to all tab (#5384) 2026-06-29 22:49:08 +08:00
fatedierandGitHub ae1c0504ec feat(dashboard): add v2 client detail status (#5381) 2026-06-26 21:15:30 +08:00
fatedierandGitHub 393a533744 docs: add release note for duplicate config names (#5379) 2026-06-24 15:03:22 +08:00
MAAZIZ Adel AyoubandGitHub 035889c360 fix(client): reject duplicate proxy and visitor names (#5378) 2026-06-24 13:37:29 +08:00
fatedierandGitHub 940bde5c46 docs: update release notes (#5377) 2026-06-23 17:06:16 +08:00
fatedierandGitHub 14628df63c test: cover tls2raw proxy protocol header (#5376) 2026-06-23 00:04:38 +08:00
Shani PathakandGitHub ba7adcab8f fix(websocket): send tunnel payload as binary frames (#5363)
The ws/wss transport carries a raw byte stream (yamux), but the
golang.org/x/net/websocket Conn defaults to text frames (PayloadType
TextFrame). Per RFC 6455 §5.6 a text frame must contain valid UTF-8, so
RFC-compliant intermediaries (API gateways / reverse proxies) validate
the payload and close the connection when the binary tunnel data is not
valid UTF-8.

This goes unnoticed peer-to-peer because x/net/websocket does not
validate UTF-8 on read, but it breaks the connection through a compliant
validating proxy. Set PayloadType to BinaryFrame on both the server
listener and the client dialer so the tunnel is framed as binary.
2026-06-22 23:34:02 +08:00
4cc826e236 fix(client): write proxy protocol header in tls2raw plugin (#5362)
Co-authored-by: futrobo <futrobo@163.com>
2026-06-22 23:08:50 +08:00
fatedierandGitHub 54c6ccdfec feat: remove proxies client filter (#5375) 2026-06-22 22:53:57 +08:00
fatedierandGitHub 9bde0b07de feat: paginate dashboard clients and proxies via API v2 (#5354)
Move the frps dashboard Clients and Proxies views to the paginated
/api/v2/clients and /api/v2/proxies endpoints instead of fetching all
data at once, and extend server-side proxy search so the search box
keeps working under pagination.

Frontend:
- Add V2Envelope/V2Page types and getV2 HTTP helper to api/http.ts
- Add v2 paginated fetch functions to api/client.ts and api/proxy.ts
- Add ClientV2Info and ProxyV2Info types for v2 API responses
- Rewrite Clients.vue with server-side pagination, status/user search
  filtering, and ElPagination component
- Rewrite Proxies.vue with server-side pagination, type tabs, client
  dropdown filter, and a search box that passes q to the API
- Default page size 10, selectable sizes [10, 20, 50, 100]

Backend:
- Extend /api/v2/proxies q matching to also cover online proxy spec
  fields: TCP/UDP remotePort and HTTP/HTTPS/TCPMux customDomains and
  subdomain, so dashboard search no longer needs to scan every page
- Add controller_v2 tests for the new spec-field matching
2026-06-03 14:08:45 +08:00
fatedierandGitHub c6c545289c fix: normalize web package lockfile (#5353) 2026-06-02 13:39:37 +08:00
fatedier 503afe78b7 feat: add dashboard API v2 pagination endpoints (#5351) 2026-06-01 20:09:25 +08:00
fatedierandGitHub 8dd26c6961 Merge pull request #5350 from fatedier/dev
Release v0.69.1
2026-06-01 18:02:22 +08:00
fatedierandGitHub 9ea1d86f03 test: handle wire v2 compatibility baselines (#5349) 2026-06-01 17:52:57 +08:00
fatedier ac3e82db4e Release v0.69.1 (#5348) 2026-06-01 16:22:34 +08:00
fatedier 0773938d70 feat: bridge mixed wire protocol SUDP payloads (#5347)
SUDP payload codec follows transport wireProtocol; same-protocol v1/v1 and v2/v2 keep raw join; only mixed proxy/visitor protocols use message-aware bridge; no new capability/selection field.
2026-06-01 16:22:34 +08:00
fatedier 9bacce22a2 feat: use wire v2 framing for XTCP NatHoleSid (#5343) 2026-06-01 16:22:34 +08:00
fatedier 7f8d68b666 feat: use wire v2 framing for UDP workConn payload (#5340) 2026-06-01 16:22:34 +08:00
fatedierandGitHub c8c1e5116c Merge pull request #5323 from fatedier/dev
Release v0.69.0
2026-05-22 00:55:23 +08:00
fatedierandGitHub 4ec8de973f Merge pull request #5287 from fatedier/dev
bump version to v0.68.1
2026-04-14 01:28:33 +08:00
fatedierandGitHub 5bfcea3d0c merge dev to master (#5254)
* ci: bump github actions to latest major versions (#5251)

* docker: copy shared web directory for npm workspace builds
2026-03-20 15:54:26 +08:00
fatedierandGitHub 0a1b4ab21f Merge pull request #5249 from fatedier/dev
bump version
2026-03-20 13:56:28 +08:00
fatedierandGitHub 5f575b8442 Merge pull request #5147 from fatedier/dev
bump version
2026-01-31 14:01:40 +08:00
fatedierandGitHub a1348cdf00 bump version (#5112) 2026-01-04 14:54:13 +08:00
fatedierandGitHub 2f5e1f7945 Merge pull request #4999 from fatedier/dev
bump version
2025-09-25 20:23:42 +08:00
fatedierandGitHub 22ae8166d3 Merge pull request #4925 from fatedier/dev
bump version
2025-08-10 23:26:32 +08:00
fatedierandGitHub af6bc6369d Merge pull request #4849 from fatedier/dev
bump version
2025-06-25 11:51:19 +08:00
140 changed files with 12772 additions and 2354 deletions
+8 -7
View File
@@ -7,14 +7,15 @@ jobs:
steps: steps:
- checkout - checkout
- run: - run:
name: Build web assets (frps) name: Test and build web assets
command: make install build command: make web-ci
working_directory: web/frps
- run: - run:
name: Build web assets (frpc) name: Check Go formatting and build binaries
command: make install build command: |
working_directory: web/frpc set -e
- run: make make env fmt
git diff --exit-code
make build
- run: make alltest - run: make alltest
workflows: workflows:
+2 -2
View File
@@ -65,7 +65,7 @@ jobs:
with: with:
context: . context: .
file: ./dockerfiles/Dockerfile-for-frpc file: ./dockerfiles/Dockerfile-for-frpc
platforms: linux/amd64,linux/arm/v7,linux/arm64,linux/ppc64le,linux/s390x platforms: linux/amd64,linux/arm/v7,linux/arm64,linux/ppc64le
push: true push: true
tags: | tags: |
${{ env.TAG_FRPC }} ${{ env.TAG_FRPC }}
@@ -76,7 +76,7 @@ jobs:
with: with:
context: . context: .
file: ./dockerfiles/Dockerfile-for-frps file: ./dockerfiles/Dockerfile-for-frps
platforms: linux/amd64,linux/arm/v7,linux/arm64,linux/ppc64le,linux/s390x platforms: linux/amd64,linux/arm/v7,linux/arm64,linux/ppc64le
push: true push: true
tags: | tags: |
${{ env.TAG_FRPS }} ${{ env.TAG_FRPS }}
+2 -6
View File
@@ -22,12 +22,8 @@ jobs:
- uses: actions/setup-node@v6 - uses: actions/setup-node@v6
with: with:
node-version: '22' node-version: '22'
- name: Build web assets (frps) - name: Test and build web assets
run: make build run: make web-ci
working-directory: web/frps
- name: Build web assets (frpc)
run: make build
working-directory: web/frpc
- name: golangci-lint - name: golangci-lint
uses: golangci/golangci-lint-action@v9 uses: golangci/golangci-lint-action@v9
with: with:
+4 -1
View File
@@ -5,7 +5,7 @@ NOWEB_TAG = $(shell [ ! -d web/frps/dist ] || [ ! -d web/frpc/dist ] && echo ',n
FRP_COMPAT_BASELINE_COUNT ?= 8 FRP_COMPAT_BASELINE_COUNT ?= 8
FRP_COMPAT_FLOOR_VERSION ?= 0.61.0 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 all: env fmt web build
@@ -16,6 +16,9 @@ env:
web: frps-web frpc-web web: frps-web frpc-web
web-ci:
cd web && npm ci && npm run lint:check --workspace frps && npm run lint:check --workspace frpc && npm run test:unit && npm run build --workspace frps && npm run build --workspace frpc
frps-web: frps-web:
$(MAKE) -C web/frps build $(MAKE) -C web/frps build
+1 -1
View File
@@ -9,7 +9,7 @@ all: build
build: app build: app
app: app:
@$(foreach n, $(os-archs), \ @set -e; $(foreach n, $(os-archs), \
os=$(shell echo "$(n)" | cut -d : -f 1); \ os=$(shell echo "$(n)" | cut -d : -f 1); \
arch=$(shell echo "$(n)" | cut -d : -f 2); \ arch=$(shell echo "$(n)" | cut -d : -f 2); \
extra=$(shell echo "$(n)" | cut -d : -f 3); \ extra=$(shell echo "$(n)" | cut -d : -f 3); \
+8 -9
View File
@@ -2,7 +2,6 @@
[![Build Status](https://circleci.com/gh/fatedier/frp.svg?style=shield)](https://circleci.com/gh/fatedier/frp) [![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) [![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) [![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) [README](README.md) | [中文文档](README_zh.md)
@@ -13,14 +12,6 @@ frp is an open source project with its ongoing development made possible entirel
<h3 align="center">Gold Sponsors</h3> <h3 align="center">Gold Sponsors</h3>
<!--gold sponsors start--> <!--gold sponsors start-->
<p align="center">
<a href="https://jb.gg/frp" target="_blank">
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
<br>
<b>The complete IDE crafted for professional Go developers</b>
</a>
</p>
<p align="center"> <p align="center">
<a href="https://github.com/beclab/Olares" target="_blank"> <a href="https://github.com/beclab/Olares" target="_blank">
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_olares.jpeg"> <img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_olares.jpeg">
@@ -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. an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
</div> </div>
<p align="center">
<a href="https://jb.gg/frp" target="_blank">
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
<br>
<b>The complete IDE crafted for professional Go developers</b>
</a>
</p>
<!--gold sponsors end--> <!--gold sponsors end-->
## What is frp? ## What is frp?
+8 -9
View File
@@ -2,7 +2,6 @@
[![Build Status](https://circleci.com/gh/fatedier/frp.svg?style=shield)](https://circleci.com/gh/fatedier/frp) [![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) [![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) [![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) [README](README.md) | [中文文档](README_zh.md)
@@ -15,14 +14,6 @@ frp 是一个完全开源的项目,我们的开发工作完全依靠赞助者
<h3 align="center">Gold Sponsors</h3> <h3 align="center">Gold Sponsors</h3>
<!--gold sponsors start--> <!--gold sponsors start-->
<p align="center">
<a href="https://jb.gg/frp" target="_blank">
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
<br>
<b>The complete IDE crafted for professional Go developers</b>
</a>
</p>
<p align="center"> <p align="center">
<a href="https://github.com/beclab/Olares" target="_blank"> <a href="https://github.com/beclab/Olares" target="_blank">
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_olares.jpeg"> <img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_olares.jpeg">
@@ -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. an API that records Zoom, Google Meet, Microsoft Teams, in-person meetings, and more.
</div> </div>
<p align="center">
<a href="https://jb.gg/frp" target="_blank">
<img width="420px" src="https://raw.githubusercontent.com/fatedier/frp/dev/doc/pic/sponsor_jetbrains.jpg">
<br>
<b>The complete IDE crafted for professional Go developers</b>
</a>
</p>
<!--gold sponsors end--> <!--gold sponsors end-->
## 为什么使用 frp ? ## 为什么使用 frp ?
+5 -7
View File
@@ -1,11 +1,9 @@
## Features ## Features
* When `transport.wireProtocol = "v2"` is enabled, ordinary UDP proxy work connection payloads now use wire protocol v2 message framing. This keeps UDP message payloads aligned with the negotiated frpc/frps wire protocol. * 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.
* SUDP proxy payloads now also follow the connection wire protocol. SUDP v2 endpoints use wire protocol v2 message framing, while v1/default endpoints continue to use the legacy message codec. When the SUDP proxy frpc and visitor frpc use mixed v1/v2 wire protocols, frps bridges UDPPacket messages between the two codecs.
## Compatibility Notes ## Fixes
* The default/empty `transport.wireProtocol` and `transport.wireProtocol = "v1"` continue to use the legacy message codec for ordinary UDP and SUDP proxy payloads. * 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.
* Raw stream proxy paths such as TCP, HTTP, and STCP remain unframed and are not affected by the UDP/SUDP payload framing change. * Fixed `frpc verify` ignoring configured `featureGates`, which caused VirtualNet configurations to be rejected even when the feature was enabled.
* Direct NAT hole UDP sid probing packets are not changed by this release. * Fixed a case-insensitive validation bypass that allowed `customDomains` under the configured `subDomainHost` to be registered using mixed-case domain names.
* `transport.wireProtocol = "v2"` requires peers to use versions that support the same wire v2 payload semantics. Mixing a newer peer that sends v2-framed UDP or SUDP payloads with an older v2-capable peer that still expects the legacy payload codec can break that proxy traffic. During rolling upgrades, upgrade both SUDP proxy and visitor frpc instances before enabling `transport.wireProtocol = "v2"` for SUDP, or keep those clients on `transport.wireProtocol = "v1"` until both sides are upgraded.
+254
View File
@@ -2,12 +2,16 @@ package client
import ( import (
"errors" "errors"
"os"
"path/filepath" "path/filepath"
"strings"
"testing" "testing"
"github.com/fatedier/frp/client/configmgmt" "github.com/fatedier/frp/client/configmgmt"
"github.com/fatedier/frp/pkg/config/source" "github.com/fatedier/frp/pkg/config/source"
v1 "github.com/fatedier/frp/pkg/config/v1" v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/policy/security"
"github.com/fatedier/frp/pkg/vnet"
) )
func newTestRawTCPProxyConfig(name string) *v1.TCPProxyConfig { func newTestRawTCPProxyConfig(name string) *v1.TCPProxyConfig {
@@ -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) { func TestServiceConfigManagerCreateStoreProxyConflict(t *testing.T) {
storeSource, err := source.NewStoreSource(source.StoreSourceConfig{ storeSource, err := source.NewStoreSource(source.StoreSourceConfig{
Path: filepath.Join(t.TempDir(), "store.json"), Path: filepath.Join(t.TempDir(), "store.json"),
+11 -2
View File
@@ -47,6 +47,8 @@ type SessionContext struct {
Connector MessageConnector Connector MessageConnector
// Virtual net controller // Virtual net controller
VnetController *vnet.Controller VnetController *vnet.Controller
// UDPPacketCodec is immutable for the lifetime of this negotiated session.
UDPPacketCodec string
} }
type Control struct { type Control struct {
@@ -92,9 +94,16 @@ func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, erro
ctl.registerMsgHandlers() ctl.registerMsgHandlers()
ctl.msgTransporter = transport.NewMessageTransporter(ctl.msgDispatcher) ctl.msgTransporter = transport.NewMessageTransporter(ctl.msgDispatcher)
ctl.pm = proxy.NewManager(ctl.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.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 return ctl, nil
} }
+9 -4
View File
@@ -99,6 +99,7 @@ func (d *controlSessionDialer) Dial(previousRunID string) (*SessionContext, erro
Auth: d.auth, Auth: d.auth,
Connector: newMessageConnector(connector, d.common.Transport.WireProtocol), Connector: newMessageConnector(connector, d.common.Transport.WireProtocol),
VnetController: d.vnetController, VnetController: d.vnetController,
UDPPacketCodec: loginResult.udpPacketCodec,
}, nil }, nil
} }
@@ -127,8 +128,9 @@ func (d *controlSessionDialer) buildLoginMsg(previousRunID string) (*msg.Login,
} }
type loginExchangeResult struct { type loginExchangeResult struct {
resp *msg.LoginResp resp *msg.LoginResp
crypto *wire.CryptoContext crypto *wire.CryptoContext
udpPacketCodec string
} }
func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) (*loginExchangeResult, error) { func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) (*loginExchangeResult, error) {
@@ -172,6 +174,7 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login)
}() }()
var cryptoContext *wire.CryptoContext var cryptoContext *wire.CryptoContext
var udpPacketCodec string
if wireConn != nil { if wireConn != nil {
serverHelloFrame, err := wireConn.ReadFrame() serverHelloFrame, err := wireConn.ReadFrame()
if err != nil { if err != nil {
@@ -191,6 +194,7 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login)
if err != nil { if err != nil {
return nil, err return nil, err
} }
udpPacketCodec = serverHello.Selected.Message.UDPPacketCodec
} }
var loginRespMsg msg.LoginResp var loginRespMsg msg.LoginResp
@@ -198,8 +202,9 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login)
return nil, err return nil, err
} }
return &loginExchangeResult{ return &loginExchangeResult{
resp: &loginRespMsg, resp: &loginRespMsg,
crypto: cryptoContext, crypto: cryptoContext,
udpPacketCodec: udpPacketCodec,
}, nil }, nil
} }
+2
View File
@@ -117,6 +117,7 @@ func TestControlSessionDialerDialV1(t *testing.T) {
defer sessionCtx.Connector.Close() defer sessionCtx.Connector.Close()
require.Equal(t, "run-v1", sessionCtx.RunID) require.Equal(t, "run-v1", sessionCtx.RunID)
require.Empty(t, sessionCtx.UDPPacketCodec)
require.NotNil(t, sessionCtx.Conn) require.NotNil(t, sessionCtx.Conn)
require.NotNil(t, sessionCtx.Connector) require.NotNil(t, sessionCtx.Connector)
require.False(t, connector.closed.Load()) require.False(t, connector.closed.Load())
@@ -225,6 +226,7 @@ func TestControlSessionDialerDialV2(t *testing.T) {
defer sessionCtx.Connector.Close() defer sessionCtx.Connector.Close()
require.Equal(t, "run-v2", sessionCtx.RunID) require.Equal(t, "run-v2", sessionCtx.RunID)
require.Equal(t, wire.UDPPacketCodecBinary, sessionCtx.UDPPacketCodec)
require.NotNil(t, sessionCtx.Conn) require.NotNil(t, sessionCtx.Conn)
require.NotNil(t, sessionCtx.Connector) require.NotNil(t, sessionCtx.Connector)
require.False(t, connector.closed.Load()) require.False(t, connector.closed.Load())
+125
View File
@@ -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)
}
+27 -9
View File
@@ -61,11 +61,12 @@ func NewProxy(
encryptionKey []byte, encryptionKey []byte,
msgTransporter transport.MessageTransporter, msgTransporter transport.MessageTransporter,
vnetController *vnet.Controller, vnetController *vnet.Controller,
udpPacketCodec string,
) (pxy Proxy) { ) (pxy Proxy) {
var limiter *rate.Limiter var limiter *rate.Limiter
limitBytes := pxyConf.GetBaseConfig().Transport.BandwidthLimit.Bytes() limitBytes := pxyConf.GetBaseConfig().Transport.BandwidthLimit.Bytes()
if limitBytes > 0 && pxyConf.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeClient { if limitBytes > 0 && pxyConf.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeClient {
limiter = rate.NewLimiter(rate.Limit(float64(limitBytes)), int(limitBytes)) limiter = limit.NewBandwidthLimiter(limitBytes)
} }
baseProxy := BaseProxy{ baseProxy := BaseProxy{
@@ -77,6 +78,7 @@ func NewProxy(
vnetController: vnetController, vnetController: vnetController,
xl: xlog.FromContextSafe(ctx), xl: xlog.FromContextSafe(ctx),
ctx: ctx, ctx: ctx,
udpPacketCodec: udpPacketCodec,
} }
factory := proxyFactoryRegistry[reflect.TypeOf(pxyConf)] factory := proxyFactoryRegistry[reflect.TypeOf(pxyConf)]
@@ -98,9 +100,10 @@ type BaseProxy struct {
proxyPlugin plugin.Plugin proxyPlugin plugin.Plugin
inWorkConnCallback func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) /* continue */ bool inWorkConnCallback func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) /* continue */ bool
mu sync.RWMutex mu sync.RWMutex
xl *xlog.Logger xl *xlog.Logger
ctx context.Context ctx context.Context
udpPacketCodec string
} }
func (pxy *BaseProxy) Run() error { func (pxy *BaseProxy) Run() error {
@@ -171,6 +174,26 @@ func (pxy *BaseProxy) HandleTCPWorkConnection(workConn net.Conn, m *msg.StartWor
xl.Tracef("handle tcp work connection, useEncryption: %t, useCompression: %t", xl.Tracef("handle tcp work connection, useEncryption: %t, useCompression: %t",
baseCfg.Transport.UseEncryption, baseCfg.Transport.UseCompression) baseCfg.Transport.UseEncryption, baseCfg.Transport.UseCompression)
var srcAddr, dstAddr *net.TCPAddr
if m.SrcAddr != "" && m.SrcPort != 0 {
if m.DstAddr == "" {
m.DstAddr = "127.0.0.1"
}
var err error
srcAddr, err = net.ResolveTCPAddr("tcp", net.JoinHostPort(m.SrcAddr, strconv.Itoa(int(m.SrcPort))))
if err != nil {
xl.Warnf("resolve source address [%s] error: %v", m.SrcAddr, err)
_ = workConn.Close()
return
}
dstAddr, err = net.ResolveTCPAddr("tcp", net.JoinHostPort(m.DstAddr, strconv.Itoa(int(m.DstPort))))
if err != nil {
xl.Warnf("resolve destination address [%s] error: %v", m.DstAddr, err)
_ = workConn.Close()
return
}
}
remote, recycleFn, err := pxy.wrapWorkConn(workConn, encKey) remote, recycleFn, err := pxy.wrapWorkConn(workConn, encKey)
if err != nil { if err != nil {
xl.Errorf("wrap work connection: %v", err) xl.Errorf("wrap work connection: %v", err)
@@ -180,11 +203,6 @@ func (pxy *BaseProxy) HandleTCPWorkConnection(workConn net.Conn, m *msg.StartWor
// check if we need to send proxy protocol info // check if we need to send proxy protocol info
var connInfo plugin.ConnectionInfo var connInfo plugin.ConnectionInfo
if m.SrcAddr != "" && m.SrcPort != 0 { if m.SrcAddr != "" && m.SrcPort != 0 {
if m.DstAddr == "" {
m.DstAddr = "127.0.0.1"
}
srcAddr, _ := net.ResolveTCPAddr("tcp", net.JoinHostPort(m.SrcAddr, strconv.Itoa(int(m.SrcPort))))
dstAddr, _ := net.ResolveTCPAddr("tcp", net.JoinHostPort(m.DstAddr, strconv.Itoa(int(m.DstPort))))
connInfo.SrcAddr = srcAddr connInfo.SrcAddr = srcAddr
connInfo.DstAddr = dstAddr connInfo.DstAddr = dstAddr
} }
+5 -2
View File
@@ -43,7 +43,8 @@ type Manager struct {
encryptionKey []byte encryptionKey []byte
clientCfg *v1.ClientCommonConfig clientCfg *v1.ClientCommonConfig
ctx context.Context ctx context.Context
udpPacketCodec string
} }
func NewManager( func NewManager(
@@ -52,6 +53,7 @@ func NewManager(
encryptionKey []byte, encryptionKey []byte,
msgTransporter transport.MessageTransporter, msgTransporter transport.MessageTransporter,
vnetController *vnet.Controller, vnetController *vnet.Controller,
udpPacketCodec string,
) *Manager { ) *Manager {
return &Manager{ return &Manager{
proxies: make(map[string]*Wrapper), proxies: make(map[string]*Wrapper),
@@ -61,6 +63,7 @@ func NewManager(
encryptionKey: encryptionKey, encryptionKey: encryptionKey,
clientCfg: clientCfg, clientCfg: clientCfg,
ctx: ctx, ctx: ctx,
udpPacketCodec: udpPacketCodec,
} }
} }
@@ -166,7 +169,7 @@ func (pm *Manager) UpdateAll(proxyCfgs []v1.ProxyConfigurer) {
for _, cfg := range proxyCfgs { for _, cfg := range proxyCfgs {
name := cfg.GetBaseConfig().Name name := cfg.GetBaseConfig().Name
if _, ok := pm.proxies[name]; !ok { if _, ok := pm.proxies[name]; !ok {
pxy := NewWrapper(pm.ctx, cfg, pm.clientCfg, pm.encryptionKey, pm.HandleEvent, pm.msgTransporter, pm.vnetController) pxy := NewWrapper(pm.ctx, cfg, pm.clientCfg, pm.encryptionKey, pm.HandleEvent, pm.msgTransporter, pm.vnetController, pm.udpPacketCodec)
if pm.inWorkConnCallback != nil { if pm.inWorkConnCallback != nil {
pxy.SetInWorkConnCallback(pm.inWorkConnCallback) pxy.SetInWorkConnCallback(pm.inWorkConnCallback)
} }
+47
View File
@@ -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)
}
+2 -1
View File
@@ -99,6 +99,7 @@ func NewWrapper(
eventHandler event.Handler, eventHandler event.Handler,
msgTransporter transport.MessageTransporter, msgTransporter transport.MessageTransporter,
vnetController *vnet.Controller, vnetController *vnet.Controller,
udpPacketCodec string,
) *Wrapper { ) *Wrapper {
baseInfo := cfg.GetBaseConfig() baseInfo := cfg.GetBaseConfig()
xl := xlog.FromContextSafe(ctx).Spawn().AppendPrefix(baseInfo.Name) xl := xlog.FromContextSafe(ctx).Spawn().AppendPrefix(baseInfo.Name)
@@ -127,7 +128,7 @@ func NewWrapper(
xl.Tracef("enable health check monitor") xl.Tracef("enable health check monitor")
} }
pw.pxy = NewProxy(pw.ctx, pw.Cfg, clientCfg, encryptionKey, pw.msgTransporter, pw.vnetController) pw.pxy = NewProxy(pw.ctx, pw.Cfg, clientCfg, encryptionKey, pw.msgTransporter, pw.vnetController, udpPacketCodec)
return pw return pw
} }
+7 -1
View File
@@ -87,7 +87,13 @@ func (pxy *SUDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) {
} }
workConn := netpkg.WrapReadWriteCloserToConn(remote, conn) 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) readCh := make(chan *msg.UDPPacket, 1024)
sendCh := make(chan msg.Message, 1024) sendCh := make(chan msg.Message, 1024)
isClose := false isClose := false
+10 -3
View File
@@ -97,10 +97,17 @@ func (pxy *UDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) {
return return
} }
pxy.mu.Lock() workConn := netpkg.WrapReadWriteCloserToConn(remote, conn)
pxy.workConn = netpkg.WrapReadWriteCloserToConn(remote, conn)
// Plain UDP payload follows the configured wire protocol for message framing. // 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.readCh = make(chan *msg.UDPPacket, 1024)
pxy.sendCh = make(chan msg.Message, 1024) pxy.sendCh = make(chan msg.Message, 1024)
pxy.closed = false pxy.closed = false
+16 -4
View File
@@ -22,6 +22,7 @@ import (
"net/http" "net/http"
"os" "os"
"sync" "sync"
"sync/atomic"
"time" "time"
"github.com/fatedier/golib/crypto" "github.com/fatedier/golib/crypto"
@@ -32,6 +33,7 @@ import (
"github.com/fatedier/frp/pkg/config" "github.com/fatedier/frp/pkg/config"
"github.com/fatedier/frp/pkg/config/source" "github.com/fatedier/frp/pkg/config/source"
v1 "github.com/fatedier/frp/pkg/config/v1" v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/config/v1/validation"
"github.com/fatedier/frp/pkg/msg" "github.com/fatedier/frp/pkg/msg"
"github.com/fatedier/frp/pkg/policy/security" "github.com/fatedier/frp/pkg/policy/security"
httppkg "github.com/fatedier/frp/pkg/util/http" httppkg "github.com/fatedier/frp/pkg/util/http"
@@ -109,6 +111,9 @@ func setServiceOptionsDefault(options *ServiceOptions) error {
// Service is the client service that connects to frps and provides proxy services. // Service is the client service that connects to frps and provides proxy services.
type Service struct { type Service struct {
ctlMu sync.RWMutex ctlMu sync.RWMutex
// Stores gracefulShutdownDuration independently from ctlMu, because the
// graceful shutdown wait may hold ctlMu for an arbitrary duration.
gracefulShutdownDuration atomic.Int64
// manager control connection with server // manager control connection with server
ctl *Control ctl *Control
// Uniq id got from frps, it will be attached to loginMsg. // Uniq id got from frps, it will be attached to loginMsg.
@@ -149,8 +154,7 @@ type Service struct {
// service context // service context
ctx context.Context ctx context.Context
// call cancel to stop service // call cancel to stop service
cancel context.CancelCauseFunc cancel context.CancelCauseFunc
gracefulShutdownDuration time.Duration
connectorCreator func(context.Context, *v1.ClientCommonConfig) Connector connectorCreator func(context.Context, *v1.ClientCommonConfig) Connector
handleWorkConnCb func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) bool handleWorkConnCb func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) bool
@@ -412,7 +416,7 @@ func (svr *Service) Close() {
} }
func (svr *Service) GracefulClose(d time.Duration) { func (svr *Service) GracefulClose(d time.Duration) {
svr.gracefulShutdownDuration = d svr.gracefulShutdownDuration.Store(int64(d))
svr.cancel(nil) svr.cancel(nil)
} }
@@ -429,7 +433,8 @@ func (svr *Service) stop() {
svr.ctlMu.Lock() svr.ctlMu.Lock()
defer svr.ctlMu.Unlock() defer svr.ctlMu.Unlock()
if svr.ctl != nil { if svr.ctl != nil {
svr.ctl.GracefulClose(svr.gracefulShutdownDuration) d := time.Duration(svr.gracefulShutdownDuration.Load())
svr.ctl.GracefulClose(d)
svr.ctl = nil svr.ctl = nil
} }
if svr.webServer != nil { if svr.webServer != nil {
@@ -506,6 +511,13 @@ func (svr *Service) reloadConfigFromSourcesLocked() error {
proxies, visitors = config.FilterClientConfigurers(reloadCommon, proxies, visitors) proxies, visitors = config.FilterClientConfigurers(reloadCommon, proxies, visitors)
proxies = config.CompleteProxyConfigurers(proxies) proxies = config.CompleteProxyConfigurers(proxies)
visitors = config.CompleteVisitorConfigurers(visitors) visitors = config.CompleteVisitorConfigurers(visitors)
requirements := validation.GetClientConfigRequirements(reloadCommon, proxies, visitors)
if svr.vnetController == nil && requirements.VirtualNet {
return errors.New(
"VirtualNet-dependent configuration requires a VirtualNet runtime enabled at startup; " +
"restart frpc after configuring featureGates.VirtualNet and virtualNet.address",
)
}
// Atomically replace the entire configuration // Atomically replace the entire configuration
if err := svr.UpdateAllConfigurer(proxies, visitors); err != nil { if err := svr.UpdateAllConfigurer(proxies, visitors); err != nil {
+95
View File
@@ -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)
}
}
+7 -1
View File
@@ -113,7 +113,13 @@ func (sv *SUDPVisitor) dispatcher() {
func (sv *SUDPVisitor) worker(workConn net.Conn, firstPacket *msg.UDPPacket) { func (sv *SUDPVisitor) worker(workConn net.Conn, firstPacket *msg.UDPPacket) {
xl := xlog.FromContextSafe(sv.ctx) xl := xlog.FromContextSafe(sv.ctx)
xl.Debugf("starting sudp proxy worker") xl.Debugf("starting sudp proxy worker")
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 := &sync.WaitGroup{}
wg.Add(2) wg.Add(2)
+11
View File
@@ -50,6 +50,17 @@ type Helper interface {
RunID() string RunID() string
} }
type udpPacketCodecProvider interface {
UDPPacketCodec() string
}
func udpPacketCodecFromHelper(helper Helper) string {
if provider, ok := helper.(udpPacketCodecProvider); ok {
return provider.UDPPacketCodec()
}
return ""
}
// Visitor is used for forward traffics from local port tot remote service. // Visitor is used for forward traffics from local port tot remote service.
type Visitor interface { type Visitor interface {
Run() error Run() error
+11
View File
@@ -53,7 +53,12 @@ func NewManager(
connectServer func() (*msg.Conn, error), connectServer func() (*msg.Conn, error),
msgTransporter transport.MessageTransporter, msgTransporter transport.MessageTransporter,
vnetController *vnet.Controller, vnetController *vnet.Controller,
udpPacketCodecs ...string,
) *Manager { ) *Manager {
udpPacketCodec := ""
if len(udpPacketCodecs) > 0 {
udpPacketCodec = udpPacketCodecs[0]
}
m := &Manager{ m := &Manager{
clientCfg: clientCfg, clientCfg: clientCfg,
cfgs: make(map[string]v1.VisitorConfigurer), cfgs: make(map[string]v1.VisitorConfigurer),
@@ -68,6 +73,7 @@ func NewManager(
vnetController: vnetController, vnetController: vnetController,
transferConnFn: m.TransferConn, transferConnFn: m.TransferConn,
runID: runID, runID: runID,
udpPacketCodec: udpPacketCodec,
} }
return m return m
} }
@@ -205,6 +211,7 @@ type visitorHelperImpl struct {
vnetController *vnet.Controller vnetController *vnet.Controller
transferConnFn func(name string, conn net.Conn) error transferConnFn func(name string, conn net.Conn) error
runID string runID string
udpPacketCodec string
} }
func (v *visitorHelperImpl) ConnectServer() (*msg.Conn, error) { func (v *visitorHelperImpl) ConnectServer() (*msg.Conn, error) {
@@ -226,3 +233,7 @@ func (v *visitorHelperImpl) VNetController() *vnet.Controller {
func (v *visitorHelperImpl) RunID() string { func (v *visitorHelperImpl) RunID() string {
return v.runID return v.runID
} }
func (v *visitorHelperImpl) UDPPacketCodec() string {
return v.udpPacketCodec
}
-7
View File
@@ -33,7 +33,6 @@ import (
"github.com/fatedier/frp/pkg/config/source" "github.com/fatedier/frp/pkg/config/source"
v1 "github.com/fatedier/frp/pkg/config/v1" v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/config/v1/validation" "github.com/fatedier/frp/pkg/config/v1/validation"
"github.com/fatedier/frp/pkg/policy/featuregate"
"github.com/fatedier/frp/pkg/policy/security" "github.com/fatedier/frp/pkg/policy/security"
"github.com/fatedier/frp/pkg/util/log" "github.com/fatedier/frp/pkg/util/log"
"github.com/fatedier/frp/pkg/util/version" "github.com/fatedier/frp/pkg/util/version"
@@ -131,12 +130,6 @@ func runClient(cfgFilePath string, unsafeFeatures *security.UnsafeFeatures) erro
"please use yaml/json/toml format instead!\n") "please use yaml/json/toml format instead!\n")
} }
if len(result.Common.FeatureGates) > 0 {
if err := featuregate.SetFromMap(result.Common.FeatureGates); err != nil {
return err
}
}
return runClientWithAggregator(result, unsafeFeatures, cfgFilePath) return runClientWithAggregator(result, unsafeFeatures, cfgFilePath)
} }
+13 -6
View File
@@ -29,6 +29,18 @@ func init() {
rootCmd.AddCommand(verifyCmd) rootCmd.AddCommand(verifyCmd)
} }
func verifyClientConfig(
configFile string,
strict bool,
unsafeFeatures *security.UnsafeFeatures,
) (validation.Warning, error) {
cliCfg, proxyCfgs, visitorCfgs, _, err := config.LoadClientConfig(configFile, strict)
if err != nil {
return nil, err
}
return validation.ValidateAllClientConfig(cliCfg, proxyCfgs, visitorCfgs, unsafeFeatures)
}
var verifyCmd = &cobra.Command{ var verifyCmd = &cobra.Command{
Use: "verify", Use: "verify",
Short: "Verify that the configures is valid", Short: "Verify that the configures is valid",
@@ -38,13 +50,8 @@ var verifyCmd = &cobra.Command{
return nil return nil
} }
cliCfg, proxyCfgs, visitorCfgs, _, err := config.LoadClientConfig(cfgFile, strictConfigMode)
if err != nil {
fmt.Println(err)
os.Exit(1)
}
unsafeFeatures := security.NewUnsafeFeatures(allowUnsafe) unsafeFeatures := security.NewUnsafeFeatures(allowUnsafe)
warning, err := validation.ValidateAllClientConfig(cliCfg, proxyCfgs, visitorCfgs, unsafeFeatures) warning, err := verifyClientConfig(cfgFile, strictConfigMode, unsafeFeatures)
if warning != nil { if warning != nil {
fmt.Printf("WARNING: %v\n", warning) fmt.Printf("WARNING: %v\n", warning)
} }
+67
View File
@@ -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)
})
}
}
+12 -18
View File
@@ -4,8 +4,8 @@ go 1.25.0
require ( require (
github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5 github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5
github.com/coreos/go-oidc/v3 v3.14.1 github.com/coreos/go-oidc/v3 v3.18.0
github.com/fatedier/golib v0.7.0 github.com/fatedier/golib v0.8.2
github.com/google/uuid v1.6.0 github.com/google/uuid v1.6.0
github.com/gorilla/mux v1.8.1 github.com/gorilla/mux v1.8.1
github.com/gorilla/websocket v1.5.0 github.com/gorilla/websocket v1.5.0
@@ -13,10 +13,9 @@ require (
github.com/onsi/ginkgo/v2 v2.23.4 github.com/onsi/ginkgo/v2 v2.23.4
github.com/onsi/gomega v1.36.3 github.com/onsi/gomega v1.36.3
github.com/pelletier/go-toml/v2 v2.2.0 github.com/pelletier/go-toml/v2 v2.2.0
github.com/pion/stun/v3 v3.1.1 github.com/pires/go-proxyproto v0.15.0
github.com/pires/go-proxyproto v0.7.0
github.com/prometheus/client_golang v1.19.1 github.com/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/rodaine/table v1.2.0
github.com/samber/lo v1.47.0 github.com/samber/lo v1.47.0
github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8 github.com/songgao/water v0.0.0-20200317203138-2b4b6d7c09d8
@@ -26,11 +25,11 @@ require (
github.com/tidwall/gjson v1.17.1 github.com/tidwall/gjson v1.17.1
github.com/vishvananda/netlink v1.3.0 github.com/vishvananda/netlink v1.3.0
github.com/xtaci/kcp-go/v5 v5.6.13 github.com/xtaci/kcp-go/v5 v5.6.13
golang.org/x/crypto v0.49.0 golang.org/x/crypto v0.54.0
golang.org/x/net v0.52.0 golang.org/x/net v0.56.0
golang.org/x/oauth2 v0.28.0 golang.org/x/oauth2 v0.36.0
golang.org/x/sync v0.20.0 golang.org/x/sync v0.22.0
golang.org/x/sys v0.42.0 golang.org/x/sys v0.47.0
golang.org/x/time v0.10.0 golang.org/x/time v0.10.0
golang.zx2c4.com/wireguard v0.0.0-20231211153847-12269c276173 golang.zx2c4.com/wireguard v0.0.0-20231211153847-12269c276173
gopkg.in/ini.v1 v1.67.0 gopkg.in/ini.v1 v1.67.0
@@ -44,7 +43,7 @@ require (
github.com/beorn7/perks v1.0.1 // indirect github.com/beorn7/perks v1.0.1 // indirect
github.com/cespare/xxhash/v2 v2.2.0 // indirect github.com/cespare/xxhash/v2 v2.2.0 // indirect
github.com/davecgh/go-spew v1.1.1 // indirect github.com/davecgh/go-spew v1.1.1 // indirect
github.com/go-jose/go-jose/v4 v4.0.5 // indirect github.com/go-jose/go-jose/v4 v4.1.4 // indirect
github.com/go-logr/logr v1.4.2 // indirect github.com/go-logr/logr v1.4.2 // indirect
github.com/go-task/slim-sprig/v3 v3.0.0 // indirect github.com/go-task/slim-sprig/v3 v3.0.0 // indirect
github.com/golang/snappy v0.0.4 // indirect github.com/golang/snappy v0.0.4 // indirect
@@ -53,9 +52,6 @@ require (
github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect
github.com/klauspost/cpuid/v2 v2.2.6 // indirect github.com/klauspost/cpuid/v2 v2.2.6 // indirect
github.com/klauspost/reedsolomon v1.12.0 // indirect github.com/klauspost/reedsolomon v1.12.0 // indirect
github.com/pion/dtls/v3 v3.0.10 // indirect
github.com/pion/logging v0.2.4 // indirect
github.com/pion/transport/v4 v4.0.1 // indirect
github.com/pkg/errors v0.9.1 // indirect github.com/pkg/errors v0.9.1 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/prometheus/client_model v0.5.0 // indirect github.com/prometheus/client_model v0.5.0 // indirect
@@ -67,11 +63,9 @@ require (
github.com/tidwall/pretty v1.2.0 // indirect github.com/tidwall/pretty v1.2.0 // indirect
github.com/tjfoc/gmsm v1.4.1 // indirect github.com/tjfoc/gmsm v1.4.1 // indirect
github.com/vishvananda/netns v0.0.4 // indirect github.com/vishvananda/netns v0.0.4 // indirect
github.com/wlynxg/anet v0.0.5 // indirect
go.uber.org/automaxprocs v1.6.0 // indirect go.uber.org/automaxprocs v1.6.0 // indirect
golang.org/x/mod v0.33.0 // indirect golang.org/x/text v0.40.0 // indirect
golang.org/x/text v0.35.0 // indirect golang.org/x/tools v0.47.0 // indirect
golang.org/x/tools v0.42.0 // indirect
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect
google.golang.org/protobuf v1.36.5 // indirect google.golang.org/protobuf v1.36.5 // indirect
gopkg.in/yaml.v2 v2.4.0 // indirect gopkg.in/yaml.v2 v2.4.0 // indirect
+28 -38
View File
@@ -11,8 +11,8 @@ github.com/cespare/xxhash/v2 v2.2.0 h1:DC2CZ1Ep5Y4k3ZQ899DldepgrayRUGE6BBZ/cd9Cj
github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw= github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw=
github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc= github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc=
github.com/coreos/go-oidc/v3 v3.14.1 h1:9ePWwfdwC4QKRlCXsJGou56adA/owXczOzwKdOumLqk= github.com/coreos/go-oidc/v3 v3.18.0 h1:V9orjXynvu5wiC9SemFTWnG4F45v403aIcjWo0d41+A=
github.com/coreos/go-oidc/v3 v3.14.1/go.mod h1:HaZ3szPaZ0e4r6ebqvsLWlk2Tn+aejfmrfah6hnSYEU= 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/cpuguy83/go-md2man/v2 v2.0.3/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
@@ -20,12 +20,12 @@ github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSs
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98= github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c= github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
github.com/fatedier/golib v0.7.0 h1:tMDF9ObcwVt59VUHroJOzHQjVFPLymZVMpGm9WAVwhY= github.com/fatedier/golib v0.8.2 h1:02n2Dg7KJ7rR7p7n4/6hBUjaLQf2J7EiHYZQsgGTvww=
github.com/fatedier/golib v0.7.0/go.mod h1:ArUGvPg2cOw/py2RAuBt46nNZH2VQ5Z70p109MAZpJw= github.com/fatedier/golib v0.8.2/go.mod h1:ArUGvPg2cOw/py2RAuBt46nNZH2VQ5Z70p109MAZpJw=
github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6 h1:u92UUy6FURPmNsMBUuongRWC0rBqN6gd01Dzu+D21NE= github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6 h1:u92UUy6FURPmNsMBUuongRWC0rBqN6gd01Dzu+D21NE=
github.com/fatedier/yamux v0.0.0-20250825093530-d0154be01cd6/go.mod h1:c5/tk6G0dSpXGzJN7Wk1OEie8grdSJAmeawId9Zvd34= 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.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA=
github.com/go-jose/go-jose/v4 v4.0.5/go.mod h1:s3P1lRrkT8igV8D9OjyL4WRyHvjB6a4JSllnOrmmBOA= github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
github.com/go-logr/logr v1.4.2 h1:6pFjapn8bFcIbiKo3XT4j/BhANplGihG6tvd+8rYgrY= github.com/go-logr/logr v1.4.2 h1:6pFjapn8bFcIbiKo3XT4j/BhANplGihG6tvd+8rYgrY=
github.com/go-logr/logr v1.4.2/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= github.com/go-logr/logr v1.4.2/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
github.com/go-task/slim-sprig/v3 v3.0.0 h1:sUs3vkvUymDpBKi3qH1YSqBQk9+9D/8M2mN1vB6EwHI= github.com/go-task/slim-sprig/v3 v3.0.0 h1:sUs3vkvUymDpBKi3qH1YSqBQk9+9D/8M2mN1vB6EwHI=
@@ -78,16 +78,8 @@ github.com/onsi/gomega v1.36.3 h1:hID7cr8t3Wp26+cYnfcjR6HpJ00fdogN6dqZ1t6IylU=
github.com/onsi/gomega v1.36.3/go.mod h1:8D9+Txp43QWKhM24yyOBEdpkzN8FvJyAwecBgsU4KU0= github.com/onsi/gomega v1.36.3/go.mod h1:8D9+Txp43QWKhM24yyOBEdpkzN8FvJyAwecBgsU4KU0=
github.com/pelletier/go-toml/v2 v2.2.0 h1:QLgLl2yMN7N+ruc31VynXs1vhMZa7CeHHejIeBAsoHo= github.com/pelletier/go-toml/v2 v2.2.0 h1:QLgLl2yMN7N+ruc31VynXs1vhMZa7CeHHejIeBAsoHo=
github.com/pelletier/go-toml/v2 v2.2.0/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs= github.com/pelletier/go-toml/v2 v2.2.0/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs=
github.com/pion/dtls/v3 v3.0.10 h1:k9ekkq1kaZoxnNEbyLKI8DI37j/Nbk1HWmMuywpQJgg= github.com/pires/go-proxyproto v0.15.0 h1:dTshmNbFm/D+0+sbrxUuddPOZ5Y0B7c5NhtsBkm6LqI=
github.com/pion/dtls/v3 v3.0.10/go.mod h1:YEmmBYIoBsY3jmG56dsziTv/Lca9y4Om83370CXfqJ8= github.com/pires/go-proxyproto v0.15.0/go.mod h1:OXsCrKwrK2tXS9YrI5tkHx5xaQlO8FH3lFW76orFh24=
github.com/pion/logging v0.2.4 h1:tTew+7cmQ+Mc1pTBLKH2puKsOvhm32dROumOZ655zB8=
github.com/pion/logging v0.2.4/go.mod h1:DffhXTKYdNZU+KtJ5pyQDjvOAh/GsNSyv1lbkFbe3so=
github.com/pion/stun/v3 v3.1.1 h1:CkQxveJ4xGQjulGSROXbXq94TAWu8gIX2dT+ePhUkqw=
github.com/pion/stun/v3 v3.1.1/go.mod h1:qC1DfmcCTQjl9PBaMa5wSn3x9IPmKxSdcCsxBcDBndM=
github.com/pion/transport/v4 v4.0.1 h1:sdROELU6BZ63Ab7FrOLn13M6YdJLY20wldXW2Cu2k8o=
github.com/pion/transport/v4 v4.0.1/go.mod h1:nEuEA4AD5lPdcIegQDpVLgNoDGreqM/YqmEx3ovP4jM=
github.com/pires/go-proxyproto v0.7.0 h1:IukmRewDQFWC7kfnb66CSomk2q/seBuilHBYFwyq0Hs=
github.com/pires/go-proxyproto v0.7.0/go.mod h1:Vz/1JPY/OACxWGQNIRY2BeyDmpoaWmEP40O9LbuiFR4=
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
@@ -103,8 +95,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/common v0.48.0/go.mod h1:0/KsvlIEfPQCQ5I2iNSAWKPZziNCvRs5EC6ILDTlAPc=
github.com/prometheus/procfs v0.12.0 h1:jluTpSng7V9hY0O2R9DzzJHYb2xULk9VTR1V1R/k6Bo= github.com/prometheus/procfs v0.12.0 h1:jluTpSng7V9hY0O2R9DzzJHYb2xULk9VTR1V1R/k6Bo=
github.com/prometheus/procfs v0.12.0/go.mod h1:pcuDEFsWDnvcgNzo4EEweacyhjeA9Zk3cnaOZAZEfOo= github.com/prometheus/procfs v0.12.0/go.mod h1:pcuDEFsWDnvcgNzo4EEweacyhjeA9Zk3cnaOZAZEfOo=
github.com/quic-go/quic-go v0.55.0 h1:zccPQIqYCXDt5NmcEabyYvOnomjs8Tlwl7tISjJh9Mk= github.com/quic-go/go-ossfuzz-seeds v0.1.0 h1:APacT+iIaNF6fd8AGEiN3bT/Jtkd2jz4v4TzM7MFjy0=
github.com/quic-go/quic-go v0.55.0/go.mod h1:DR51ilwU1uE164KuWXhinFcKWGlEjzys2l8zUl5Ss1U= 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 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY= github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY=
github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc=
github.com/rodaine/table v1.2.0 h1:38HEnwK4mKSHQJIkavVj+bst1TEY7j9zhLMWu4QJrMA= github.com/rodaine/table v1.2.0 h1:38HEnwK4mKSHQJIkavVj+bst1TEY7j9zhLMWu4QJrMA=
@@ -146,8 +140,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/netlink v1.3.0/go.mod h1:i6NetklAujEcC6fK0JPjT8qSwWyO0HLn4UKG+hGqeJs=
github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8= github.com/vishvananda/netns v0.0.4 h1:Oeaw1EM2JMxD51g9uhtC0D7erkIjgmj8+JZc26m1YX8=
github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM= github.com/vishvananda/netns v0.0.4/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM=
github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU=
github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA=
github.com/xtaci/kcp-go/v5 v5.6.13 h1:FEjtz9+D4p8t2x4WjciGt/jsIuhlWjjgPCCWjrVR4Hk= github.com/xtaci/kcp-go/v5 v5.6.13 h1:FEjtz9+D4p8t2x4WjciGt/jsIuhlWjjgPCCWjrVR4Hk=
github.com/xtaci/kcp-go/v5 v5.6.13/go.mod h1:75S1AKYYzNUSXIv30h+jPKJYZUwqpfvLshu63nCNSOM= github.com/xtaci/kcp-go/v5 v5.6.13/go.mod h1:75S1AKYYzNUSXIv30h+jPKJYZUwqpfvLshu63nCNSOM=
github.com/xtaci/lossyconn v0.0.0-20200209145036-adba10fffc37 h1:EWU6Pktpas0n8lLQwDsRyZfmkPeRbdgPtW609es+/9E= github.com/xtaci/lossyconn v0.0.0-20200209145036-adba10fffc37 h1:EWU6Pktpas0n8lLQwDsRyZfmkPeRbdgPtW609es+/9E=
@@ -159,30 +151,28 @@ go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4= golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA= 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-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE=
golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU= golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU=
golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc=
golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8=
golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w=
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0= golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw= 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.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.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
golang.org/x/oauth2 v0.28.0/go.mod h1:onh5ek6nERTohokkhCD/y2cV4Do3fxFHFuAejCkRWT8= 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-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= 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-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -190,14 +180,14 @@ golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.41.0 h1:QCgPso/Q3RTJx2Th4bDLqML4W6iJiaXFq2/ftQF13YU= golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
golang.org/x/term v0.41.0/go.mod h1:3pfBgksrReYfZ5lvYM0kSO0LIkAl4Yl2bXOkKP7Ec2A= 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.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8= golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs=
golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA= 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 h1:3usCWA8tQn0L8+hFJQNgzpWbd89begxN66o1Ojdn5L4=
golang.org/x/time v0.10.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM= golang.org/x/time v0.10.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
@@ -205,8 +195,8 @@ golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGm
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY= golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q= golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q=
golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k= golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=
golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0= 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.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI= golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
+29
View File
@@ -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 return result, nil
} }
@@ -417,6 +421,31 @@ func LoadClientConfig(path string, strict bool) (
return result.Common, proxyCfgs, visitorCfgs, result.IsLegacyFormat, nil return result.Common, proxyCfgs, visitorCfgs, result.IsLegacyFormat, nil
} }
// validateNoDuplicateNames rejects proxies or visitors that share a name. They are
// keyed by name in the config sources, so a duplicate would otherwise be silently
// overwritten and never started, with no error or log.
func validateNoDuplicateNames(proxies []v1.ProxyConfigurer, visitors []v1.VisitorConfigurer) error {
proxyNames := make(map[string]struct{}, len(proxies))
for _, p := range proxies {
name := p.GetBaseConfig().Name
if _, ok := proxyNames[name]; ok {
return fmt.Errorf("proxy name [%s] is duplicated", name)
}
proxyNames[name] = struct{}{}
}
visitorNames := make(map[string]struct{}, len(visitors))
for _, v := range visitors {
name := v.GetBaseConfig().Name
if _, ok := visitorNames[name]; ok {
return fmt.Errorf("visitor name [%s] is duplicated", name)
}
visitorNames[name] = struct{}{}
}
return nil
}
func CompleteProxyConfigurers(proxies []v1.ProxyConfigurer) []v1.ProxyConfigurer { func CompleteProxyConfigurers(proxies []v1.ProxyConfigurer) []v1.ProxyConfigurer {
proxyCfgs := proxies proxyCfgs := proxies
for _, c := range proxyCfgs { for _, c := range proxyCfgs {
+107
View File
@@ -17,6 +17,8 @@ package config
import ( import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"os"
"path/filepath"
"strings" "strings"
"testing" "testing"
@@ -462,6 +464,111 @@ func TestFilterClientConfigurers_FilterByStartAndEnabled(t *testing.T) {
require.Equal("keep", proxies[0].GetBaseConfig().Name) require.Equal("keep", proxies[0].GetBaseConfig().Name)
} }
func TestLoadClientConfigResult_DuplicateNames(t *testing.T) {
tests := []struct {
name string
content string
errSubstr string
}{
{
name: "duplicate proxy names",
content: `
serverAddr = "127.0.0.1"
serverPort = 7000
[[proxies]]
name = "dup"
type = "tcp"
localPort = 22
remotePort = 6000
[[proxies]]
name = "dup"
type = "tcp"
localPort = 3306
remotePort = 6001
`,
errSubstr: "proxy name [dup] is duplicated",
},
{
name: "duplicate visitor names",
content: `
serverAddr = "127.0.0.1"
serverPort = 7000
[[visitors]]
name = "dup"
type = "stcp"
serverName = "a"
secretKey = "secret"
bindPort = 9001
[[visitors]]
name = "dup"
type = "stcp"
serverName = "b"
secretKey = "secret"
bindPort = 9002
`,
errSubstr: "visitor name [dup] is duplicated",
},
{
name: "unique names",
content: `
serverAddr = "127.0.0.1"
serverPort = 7000
[[proxies]]
name = "p1"
type = "tcp"
localPort = 22
remotePort = 6000
[[proxies]]
name = "p2"
type = "tcp"
localPort = 3306
remotePort = 6001
`,
},
{
name: "same name across proxy and visitor",
content: `
serverAddr = "127.0.0.1"
serverPort = 7000
[[proxies]]
name = "same"
type = "tcp"
localPort = 22
remotePort = 6000
[[visitors]]
name = "same"
type = "stcp"
serverName = "a"
secretKey = "secret"
bindPort = 9001
`,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
require := require.New(t)
path := filepath.Join(t.TempDir(), "frpc.toml")
require.NoError(os.WriteFile(path, []byte(tc.content), 0o600))
_, err := LoadClientConfigResult(path, false)
if tc.errSubstr == "" {
require.NoError(err)
} else {
require.ErrorContains(err, tc.errSubstr)
}
})
}
}
// TestYAMLEdgeCases tests edge cases for YAML parsing, including non-map types // TestYAMLEdgeCases tests edge cases for YAML parsing, including non-map types
func TestYAMLEdgeCases(t *testing.T) { func TestYAMLEdgeCases(t *testing.T) {
require := require.New(t) require := require.New(t)
+1 -1
View File
@@ -167,7 +167,7 @@ type ServerTransportConfig struct {
// If negative, keep-alive probes are disabled. // If negative, keep-alive probes are disabled.
TCPKeepAlive int64 `json:"tcpKeepalive,omitempty"` TCPKeepAlive int64 `json:"tcpKeepalive,omitempty"`
// MaxPoolCount specifies the maximum pool size for each proxy. By default, // MaxPoolCount specifies the maximum pool size for each proxy. By default,
// this value is 5. // this value is 5. Negative values are invalid.
MaxPoolCount int64 `json:"maxPoolCount,omitempty"` MaxPoolCount int64 `json:"maxPoolCount,omitempty"`
// HeartBeatTimeout specifies the maximum time to wait for a heartbeat // HeartBeatTimeout specifies the maximum time to wait for a heartbeat
// before terminating the connection. It is not recommended to change this // before terminating the connection. It is not recommended to change this
+38 -1
View File
@@ -51,14 +51,51 @@ func (v *ConfigValidator) ValidateClientCommonConfig(c *v1.ClientCommonConfig) (
} }
func validateFeatureGates(c *v1.ClientCommonConfig) (Warning, error) { func validateFeatureGates(c *v1.ClientCommonConfig) (Warning, error) {
gates := featuregate.NewFeatureGate()
if err := gates.SetFromMap(c.FeatureGates); err != nil {
return nil, err
}
if c.VirtualNet.Address != "" { if c.VirtualNet.Address != "" {
if !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, fmt.Errorf("VirtualNet feature is not enabled; enable it by setting the appropriate feature gate flag")
} }
} }
return nil, nil return nil, nil
} }
// ClientConfigRequirements describes runtime capabilities needed by a client configuration.
type ClientConfigRequirements struct {
VirtualNet bool
}
// GetClientConfigRequirements returns the runtime capabilities needed by a client configuration.
func GetClientConfigRequirements(
common *v1.ClientCommonConfig,
proxyCfgs []v1.ProxyConfigurer,
visitorCfgs []v1.VisitorConfigurer,
) ClientConfigRequirements {
requirements := ClientConfigRequirements{}
if common != nil && common.VirtualNet.Address != "" {
requirements.VirtualNet = true
}
for _, cfg := range proxyCfgs {
if cfg.GetBaseConfig().Plugin.Type == v1.PluginVirtualNet {
requirements.VirtualNet = true
break
}
}
if !requirements.VirtualNet {
for _, cfg := range visitorCfgs {
if cfg.GetBaseConfig().Plugin.Type == v1.VisitorPluginVirtualNet {
requirements.VirtualNet = true
break
}
}
}
return requirements
}
func (v *ConfigValidator) validateAuthConfig(c *v1.AuthClientConfig) (Warning, error) { func (v *ConfigValidator) validateAuthConfig(c *v1.AuthClientConfig) (Warning, error) {
var errs error var errs error
if !slices.Contains(SupportedAuthMethods, c.Method) { if !slices.Contains(SupportedAuthMethods, c.Method) {
+140
View File
@@ -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())
}
+48
View File
@@ -0,0 +1,48 @@
// Copyright 2026 The frp Authors
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package validation
import (
"fmt"
"unicode"
"unicode/utf8"
)
const (
// MaxRunIDLength is the maximum number of bytes accepted for a control run ID.
MaxRunIDLength = 64
)
func validateIdentifier(value, kind string, maxLength int) error {
if value == "" {
return fmt.Errorf("%s cannot be empty", kind)
}
if len(value) > maxLength {
return fmt.Errorf("%s is too long: length %d exceeds maximum %d", kind, len(value), maxLength)
}
if !utf8.ValidString(value) {
return fmt.Errorf("%s must be valid UTF-8", kind)
}
for _, r := range value {
if !unicode.IsPrint(r) {
return fmt.Errorf("%s contains non-printable character", kind)
}
}
return nil
}
func ValidateRunID(runID string) error {
return validateIdentifier(runID, "run id", MaxRunIDLength)
}
+48
View File
@@ -0,0 +1,48 @@
// Copyright 2026 The frp Authors
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package validation
import (
"strings"
"testing"
"github.com/stretchr/testify/require"
)
func TestValidateIdentifiers(t *testing.T) {
tests := []struct {
name string
validate func(string) error
value string
wantError string
}{
{name: "run id accepts printable values", validate: ValidateRunID, value: "run-%1000s-中文"},
{name: "run id rejects empty", validate: ValidateRunID, wantError: "cannot be empty"},
{name: "run id rejects control character", validate: ValidateRunID, value: "run\nforged", wantError: "non-printable"},
{name: "run id rejects invalid utf8", validate: ValidateRunID, value: string([]byte{0xff}), wantError: "valid UTF-8"},
{name: "run id rejects excessive length", validate: ValidateRunID, value: strings.Repeat("a", MaxRunIDLength+1), wantError: "too long"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := tt.validate(tt.value)
if tt.wantError == "" {
require.NoError(t, err)
return
}
require.ErrorContains(t, err, tt.wantError)
})
}
}
+4 -2
View File
@@ -79,9 +79,11 @@ func validateDomainConfigForClient(c *v1.DomainConfig) error {
} }
func validateDomainConfigForServer(c *v1.DomainConfig, s *v1.ServerConfig) error { func validateDomainConfigForServer(c *v1.DomainConfig, s *v1.ServerConfig) error {
subDomainHost := strings.ToLower(s.SubDomainHost)
for _, domain := range c.CustomDomains { for _, domain := range c.CustomDomains {
if s.SubDomainHost != "" && len(strings.Split(s.SubDomainHost, ".")) < len(strings.Split(domain, ".")) { canonicalDomain := strings.ToLower(domain)
if strings.HasSuffix(domain, "."+s.SubDomainHost) { 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) return fmt.Errorf("custom domain [%s] should not belong to subdomain host [%s]", domain, s.SubDomainHost)
} }
} }
+76
View File
@@ -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)
})
}
}
+3
View File
@@ -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.VhostHTTPPort, "vhostHTTPPort"))
errs = AppendError(errs, ValidatePort(c.VhostHTTPSPort, "vhostHTTPSPort")) errs = AppendError(errs, ValidatePort(c.VhostHTTPSPort, "vhostHTTPSPort"))
errs = AppendError(errs, ValidatePort(c.TCPMuxHTTPConnectPort, "tcpMuxHTTPConnectPort")) errs = AppendError(errs, ValidatePort(c.TCPMuxHTTPConnectPort, "tcpMuxHTTPConnectPort"))
if c.Transport.MaxPoolCount < 0 {
errs = AppendError(errs, fmt.Errorf("invalid transport.maxPoolCount, must be non-negative"))
}
for _, p := range c.HTTPPlugins { for _, p := range c.HTTPPlugins {
if !lo.Every(SupportedHTTPPluginOps, p.Ops) { if !lo.Every(SupportedHTTPPluginOps, p.Ops) {
+51
View File
@@ -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)
})
}
}
+13 -3
View File
@@ -95,9 +95,7 @@ func (m *serverMetrics) clearUselessInfo(continuousOfflineDuration time.Duration
defer m.mu.Unlock() defer m.mu.Unlock()
total = len(m.info.ProxyStatistics) total = len(m.info.ProxyStatistics)
for name, data := range m.info.ProxyStatistics { for name, data := range m.info.ProxyStatistics {
if !data.LastCloseTime.IsZero() && if m.shouldClearProxyStats(data, continuousOfflineDuration) {
data.LastStartTime.Before(data.LastCloseTime) &&
m.clock.Since(data.LastCloseTime) > continuousOfflineDuration {
delete(m.info.ProxyStatistics, name) delete(m.info.ProxyStatistics, name)
count++ count++
log.Tracef("clear proxy [%s]'s statistics data, lastCloseTime: [%s]", name, data.LastCloseTime.String()) log.Tracef("clear proxy [%s]'s statistics data, lastCloseTime: [%s]", name, data.LastCloseTime.String())
@@ -106,10 +104,20 @@ func (m *serverMetrics) clearUselessInfo(continuousOfflineDuration time.Duration
return count, total return count, total
} }
func (m *serverMetrics) shouldClearProxyStats(data *ProxyStatistics, continuousOfflineDuration time.Duration) bool {
return !data.LastCloseTime.IsZero() &&
data.LastStartTime.Before(data.LastCloseTime) &&
m.clock.Since(data.LastCloseTime) > continuousOfflineDuration
}
func (m *serverMetrics) ClearOfflineProxies() (int, int) { func (m *serverMetrics) ClearOfflineProxies() (int, int) {
return m.clearUselessInfo(0) return m.clearUselessInfo(0)
} }
func (m *serverMetrics) PruneOfflineProxies() (int, int) {
return m.clearUselessInfo(0)
}
func (m *serverMetrics) NewClient() { func (m *serverMetrics) NewClient() {
m.info.ClientCounts.Inc(1) m.info.ClientCounts.Inc(1)
} }
@@ -231,9 +239,11 @@ func toProxyStats(name string, proxyStats *ProxyStatistics) *ProxyStats {
} }
if !proxyStats.LastStartTime.IsZero() { if !proxyStats.LastStartTime.IsZero() {
ps.LastStartTime = proxyStats.LastStartTime.Format("01-02 15:04:05") ps.LastStartTime = proxyStats.LastStartTime.Format("01-02 15:04:05")
ps.LastStartAt = proxyStats.LastStartTime.Unix()
} }
if !proxyStats.LastCloseTime.IsZero() { if !proxyStats.LastCloseTime.IsZero() {
ps.LastCloseTime = proxyStats.LastCloseTime.Format("01-02 15:04:05") ps.LastCloseTime = proxyStats.LastCloseTime.Format("01-02 15:04:05")
ps.LastCloseAt = proxyStats.LastCloseTime.Unix()
} }
return ps return ps
} }
+70
View File
@@ -22,6 +22,12 @@ func TestServerMetricsUsesClockForProxyTimestamps(t *testing.T) {
clk.SetTime(closedAt) clk.SetTime(closedAt)
metrics.CloseProxy("proxy", "tcp") metrics.CloseProxy("proxy", "tcp")
require.Equal(closedAt, metrics.info.ProxyStatistics["proxy"].LastCloseTime) require.Equal(closedAt, metrics.info.ProxyStatistics["proxy"].LastCloseTime)
stats := metrics.GetProxyByName("proxy")
require.Equal(start.Format("01-02 15:04:05"), stats.LastStartTime)
require.Equal(closedAt.Format("01-02 15:04:05"), stats.LastCloseTime)
require.Equal(start.Unix(), stats.LastStartAt)
require.Equal(closedAt.Unix(), stats.LastCloseAt)
} }
func TestServerMetricsClearUselessInfoUsesClock(t *testing.T) { func TestServerMetricsClearUselessInfoUsesClock(t *testing.T) {
@@ -43,6 +49,70 @@ func TestServerMetricsClearUselessInfoUsesClock(t *testing.T) {
require.Empty(metrics.info.ProxyStatistics) require.Empty(metrics.info.ProxyStatistics)
} }
func TestServerMetricsClearOfflineProxiesPreservesLegacyTotal(t *testing.T) {
require := require.New(t)
start := time.Date(2026, time.May, 8, 12, 30, 0, 0, time.UTC)
clk := clocktesting.NewFakeClock(start.Add(time.Minute))
metrics := newServerMetricsWithClock(clk)
metrics.info.ProxyStatistics["offline"] = &ProxyStatistics{
Name: "offline",
LastStartTime: start.Add(-time.Hour),
LastCloseTime: start,
}
metrics.info.ProxyStatistics["online"] = &ProxyStatistics{
Name: "online",
LastStartTime: start,
}
cleared, total := metrics.ClearOfflineProxies()
require.Equal(1, cleared)
require.Equal(2, total)
require.False(metrics.hasProxyStatistics("offline"))
require.True(metrics.hasProxyStatistics("online"))
}
func TestServerMetricsPruneOfflineProxiesReportsTotalStats(t *testing.T) {
require := require.New(t)
start := time.Date(2026, time.May, 8, 12, 30, 0, 0, time.UTC)
clk := clocktesting.NewFakeClock(start.Add(time.Minute))
metrics := newServerMetricsWithClock(clk)
metrics.info.ProxyStatistics["offline"] = &ProxyStatistics{
Name: "offline",
LastStartTime: start.Add(-time.Hour),
LastCloseTime: start,
}
metrics.info.ProxyStatistics["online"] = &ProxyStatistics{
Name: "online",
LastStartTime: start,
}
metrics.info.ProxyStatistics["restarted"] = &ProxyStatistics{
Name: "restarted",
LastStartTime: start.Add(30 * time.Second),
LastCloseTime: start,
}
metrics.info.ProxyStatistics["same-time"] = &ProxyStatistics{
Name: "same-time",
LastStartTime: start,
LastCloseTime: start,
}
cleared, total := metrics.PruneOfflineProxies()
require.Equal(1, cleared)
require.Equal(4, total)
require.False(metrics.hasProxyStatistics("offline"))
require.True(metrics.hasProxyStatistics("online"))
require.True(metrics.hasProxyStatistics("restarted"))
require.True(metrics.hasProxyStatistics("same-time"))
cleared, total = metrics.PruneOfflineProxies()
require.Equal(0, cleared)
require.Equal(3, total)
}
func TestServerMetricsRunUsesClockTicker(t *testing.T) { func TestServerMetricsRunUsesClockTicker(t *testing.T) {
require := require.New(t) require := require.New(t)
+3
View File
@@ -41,6 +41,8 @@ type ProxyStats struct {
TodayTrafficOut int64 TodayTrafficOut int64
LastStartTime string LastStartTime string
LastCloseTime string LastCloseTime string
LastStartAt int64
LastCloseAt int64
CurConns int64 CurConns int64
} }
@@ -85,4 +87,5 @@ type Collector interface {
GetProxyByName(proxyName string) *ProxyStats GetProxyByName(proxyName string) *ProxyStats
GetProxyTraffic(name string) *ProxyTrafficInfo GetProxyTraffic(name string) *ProxyTrafficInfo
ClearOfflineProxies() (int, int) ClearOfflineProxies() (int, int)
PruneOfflineProxies() (int, int)
} }
+199
View File
@@ -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
})
}
}
}
}
+338
View File
@@ -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)
}
}
+248
View File
@@ -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)
})
}
+1
View File
@@ -43,6 +43,7 @@ const (
V2TypeNatHoleResp uint16 = 16 V2TypeNatHoleResp uint16 = 16
V2TypeNatHoleSid uint16 = 17 V2TypeNatHoleSid uint16 = 17
V2TypeNatHoleReport uint16 = 18 V2TypeNatHoleReport uint16 = 18
V2TypeUDPPacketBinary uint16 = 19
) )
var v2MsgTypeMap = map[uint16]any{ var v2MsgTypeMap = map[uint16]any{
+3
View File
@@ -84,6 +84,9 @@ func TestV2MessageTypeIDsAreStable(t *testing.T) {
require.Equal(t, uint16(16), V2TypeNatHoleResp) require.Equal(t, uint16(16), V2TypeNatHoleResp)
require.Equal(t, uint16(17), V2TypeNatHoleSid) require.Equal(t, uint16(17), V2TypeNatHoleSid)
require.Equal(t, uint16(18), V2TypeNatHoleReport) require.Equal(t, uint16(18), V2TypeNatHoleReport)
require.Equal(t, uint16(19), V2TypeUDPPacketBinary)
_, registered := v2MsgTypeMap[V2TypeUDPPacketBinary]
require.False(t, registered, "binary UDP has a dedicated codec and must not alter generic type registry")
} }
func TestV2MessageFrameEncoding(t *testing.T) { func TestV2MessageFrameEncoding(t *testing.T) {
+25 -65
View File
@@ -15,31 +15,24 @@
package nathole package nathole
import ( import (
"errors"
"fmt" "fmt"
"net" "net"
"time" "time"
"github.com/pion/stun/v3" "github.com/fatedier/golib/net/stun"
) )
var responseTimeout = 3 * time.Second var responseTimeout = 3 * time.Second
type Message struct {
Body []byte
Addr string
}
// If the localAddr is empty, it will listen on a random port. // If the localAddr is empty, it will listen on a random port.
func Discover(stunServers []string, localAddr string) ([]string, net.Addr, error) { func Discover(stunServers []string, localAddr string) ([]string, net.Addr, error) {
// create a discoverConn and get response from messageChan
discoverConn, err := listen(localAddr) discoverConn, err := listen(localAddr)
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
defer discoverConn.Close() defer discoverConn.Close()
go discoverConn.readLoop()
addresses := make([]string, 0, len(stunServers)) addresses := make([]string, 0, len(stunServers))
for _, addr := range stunServers { for _, addr := range stunServers {
// get external address from stun server // get external address from stun server
@@ -58,10 +51,9 @@ type stunResponse struct {
} }
type discoverConn struct { type discoverConn struct {
conn *net.UDPConn conn *net.UDPConn
client *stun.Client
localAddr net.Addr localAddr net.Addr
messageChan chan *Message
} }
func listen(localAddr string) (*discoverConn, error) { func listen(localAddr string) (*discoverConn, error) {
@@ -77,82 +69,50 @@ func listen(localAddr string) (*discoverConn, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
client, err := stun.NewClient(conn)
if err != nil {
_ = conn.Close()
return nil, err
}
return &discoverConn{ return &discoverConn{
conn: conn, conn: conn,
localAddr: conn.LocalAddr(), client: client,
messageChan: make(chan *Message, 10), localAddr: conn.LocalAddr(),
}, nil }, nil
} }
func (c *discoverConn) Close() error { func (c *discoverConn) Close() error {
if c.messageChan != nil {
close(c.messageChan)
c.messageChan = nil
}
return c.conn.Close() return c.conn.Close()
} }
func (c *discoverConn) readLoop() {
for {
buf := make([]byte, 1024)
n, addr, err := c.conn.ReadFromUDP(buf)
if err != nil {
return
}
buf = buf[:n]
c.messageChan <- &Message{
Body: buf,
Addr: addr.String(),
}
}
}
func (c *discoverConn) doSTUNRequest(addr string) (*stunResponse, error) { func (c *discoverConn) doSTUNRequest(addr string) (*stunResponse, error) {
serverAddr, err := net.ResolveUDPAddr("udp4", addr) serverAddr, err := net.ResolveUDPAddr("udp4", addr)
if err != nil { if err != nil {
return nil, err return nil, err
} }
request, err := stun.Build(stun.TransactionID, stun.BindingRequest) transaction, err := stun.NewBindingTransaction(serverAddr)
if err != nil { if err != nil {
return nil, err return nil, err
} }
if err := c.conn.SetReadDeadline(time.Now().Add(responseTimeout)); err != nil {
if err = request.NewTransactionID(); err != nil {
return nil, err return nil, err
} }
if _, err := c.conn.WriteTo(request.Raw, serverAddr); err != nil { response, err := c.client.Do(transaction)
return nil, err if err != nil {
} var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
var m stun.Message return nil, fmt.Errorf("wait response from stun server timeout")
select {
case msg := <-c.messageChan:
m.Raw = msg.Body
if err := m.Decode(); err != nil {
return nil, err
} }
case <-time.After(responseTimeout): return nil, err
return nil, fmt.Errorf("wait response from stun server timeout")
} }
xorAddrGetter := &stun.XORMappedAddress{}
mappedAddrGetter := &stun.MappedAddress{}
changedAddrGetter := ChangedAddress{}
otherAddrGetter := &stun.OtherAddress{}
resp := &stunResponse{} resp := &stunResponse{}
if err := mappedAddrGetter.GetFrom(&m); err == nil { if response.MappedAddr != nil {
resp.externalAddr = mappedAddrGetter.String() resp.externalAddr = response.MappedAddr.String()
} }
if err := xorAddrGetter.GetFrom(&m); err == nil { if response.OtherAddr != nil {
resp.externalAddr = xorAddrGetter.String() resp.otherAddr = response.OtherAddr.String()
}
if err := changedAddrGetter.GetFrom(&m); err == nil {
resp.otherAddr = changedAddrGetter.String()
}
if err := otherAddrGetter.GetFrom(&m); err == nil {
resp.otherAddr = otherAddrGetter.String()
} }
return resp, nil return resp, nil
} }
+382
View File
@@ -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)
}
-16
View File
@@ -18,10 +18,8 @@ import (
"bytes" "bytes"
"fmt" "fmt"
"net" "net"
"strconv"
"github.com/fatedier/golib/crypto" "github.com/fatedier/golib/crypto"
"github.com/pion/stun/v3"
"github.com/fatedier/frp/pkg/msg" "github.com/fatedier/frp/pkg/msg"
) )
@@ -48,20 +46,6 @@ func DecodeMessageInto(data, key []byte, m msg.Message) error {
return msg.ReadMsgInto(bytes.NewReader(buf), m) return msg.ReadMsgInto(bytes.NewReader(buf), m)
} }
type ChangedAddress struct {
IP net.IP
Port int
}
func (s *ChangedAddress) GetFrom(m *stun.Message) error {
a := (*stun.MappedAddress)(s)
return a.GetFromAs(m, stun.AttrChangedAddress)
}
func (s *ChangedAddress) String() string {
return net.JoinHostPort(s.IP.String(), strconv.Itoa(s.Port))
}
func ListAllLocalIPs() ([]net.IP, error) { func ListAllLocalIPs() ([]net.IP, error) {
addrs, err := net.InterfaceAddrs() addrs, err := net.InterfaceAddrs()
if err != nil { if err != nil {
+9
View File
@@ -72,6 +72,15 @@ func (p *TLS2RawPlugin) Handle(ctx context.Context, connInfo *ConnectionInfo) {
return return
} }
if connInfo.ProxyProtocolHeader != nil {
if _, err := connInfo.ProxyProtocolHeader.WriteTo(rawConn); err != nil {
xl.Warnf("tls2raw write proxy protocol header to local conn error: %v", err)
rawConn.Close()
tlsConn.Close()
return
}
}
libio.Join(tlsConn, rawConn) libio.Join(tlsConn, rawConn)
} }
+7 -3
View File
@@ -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 // NewUDPPacket copies buf[:n], so the read buffer can be reused
udpMsg := NewUDPPacket(buf[:n], nil, remoteAddr) udpMsg := NewUDPPacket(buf[:n], nil, remoteAddr)
select { if err = errors.PanicToError(func() {
case sendCh <- udpMsg: select {
default: case sendCh <- udpMsg:
default:
}
}); err != nil {
return
} }
} }
} }
+34
View File
@@ -1,9 +1,13 @@
package udp package udp
import ( import (
"net"
"testing" "testing"
"time"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"github.com/fatedier/frp/pkg/msg"
) )
func TestUdpPacket(t *testing.T) { func TestUdpPacket(t *testing.T) {
@@ -16,3 +20,33 @@ func TestUdpPacket(t *testing.T) {
require.NoError(err) require.NoError(err)
require.EqualValues(buf, newBuf) require.EqualValues(buf, newBuf)
} }
func TestForwardUserConnReturnsWhenSendChannelIsClosed(t *testing.T) {
listener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
require.NoError(t, err)
t.Cleanup(func() { _ = listener.Close() })
readCh := make(chan *msg.UDPPacket)
sendCh := make(chan *msg.UDPPacket)
close(sendCh)
t.Cleanup(func() { close(readCh) })
done := make(chan struct{})
go func() {
ForwardUserConn(listener, readCh, sendCh, 1500)
close(done)
}()
sender, err := net.DialUDP("udp4", nil, listener.LocalAddr().(*net.UDPAddr))
require.NoError(t, err)
t.Cleanup(func() { _ = sender.Close() })
_, err = sender.Write([]byte("trigger"))
require.NoError(t, err)
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("ForwardUserConn did not return after sending to a closed channel")
}
}
+18 -1
View File
@@ -68,7 +68,8 @@ func NewServerHello(clientHello ClientHello) (ServerHello, error) {
return ServerHello{ return ServerHello{
Selected: ServerSelection{ Selected: ServerSelection{
Message: MessageSelection{ Message: MessageSelection{
Codec: MessageCodecJSON, Codec: MessageCodecJSON,
UDPPacketCodec: selectUDPPacketCodec(clientHello.Capabilities.Message.UDPPacketCodecs),
}, },
Crypto: CryptoSelection{ Crypto: CryptoSelection{
Algorithm: algorithm, Algorithm: algorithm,
@@ -92,6 +93,15 @@ func ValidateServerHelloForClient(clientHello ClientHello, serverHello ServerHel
if serverHello.Selected.Message.Codec != MessageCodecJSON { if serverHello.Selected.Message.Codec != MessageCodecJSON {
return fmt.Errorf("unsupported selected message codec: %s", serverHello.Selected.Message.Codec) return fmt.Errorf("unsupported selected message codec: %s", serverHello.Selected.Message.Codec)
} }
udpPacketCodec := serverHello.Selected.Message.UDPPacketCodec
if udpPacketCodec != "" {
if udpPacketCodec != UDPPacketCodecBinary {
return fmt.Errorf("unsupported selected UDP packet codec: %s", udpPacketCodec)
}
if !Supports(clientHello.Capabilities.Message.UDPPacketCodecs, udpPacketCodec) {
return fmt.Errorf("selected UDP packet codec was not advertised by client: %s", udpPacketCodec)
}
}
cryptoSelection := serverHello.Selected.Crypto cryptoSelection := serverHello.Selected.Crypto
if !IsSupportedAEADAlgorithm(cryptoSelection.Algorithm) { if !IsSupportedAEADAlgorithm(cryptoSelection.Algorithm) {
return fmt.Errorf("unknown selected crypto algorithm: %s", cryptoSelection.Algorithm) return fmt.Errorf("unknown selected crypto algorithm: %s", cryptoSelection.Algorithm)
@@ -105,6 +115,13 @@ func ValidateServerHelloForClient(clientHello ClientHello, serverHello ServerHel
return nil return nil
} }
func selectUDPPacketCodec(codecs []string) string {
if Supports(codecs, UDPPacketCodecBinary) {
return UDPPacketCodecBinary
}
return ""
}
func NewCryptoContext(algorithm string, clientHelloPayload, serverHelloPayload []byte) *CryptoContext { func NewCryptoContext(algorithm string, clientHelloPayload, serverHelloPayload []byte) *CryptoContext {
return &CryptoContext{ return &CryptoContext{
Algorithm: algorithm, Algorithm: algorithm,
+7 -3
View File
@@ -36,6 +36,7 @@ const (
FrameTypeMessage uint16 = 16 FrameTypeMessage uint16 = 16
MessageCodecJSON = "json" MessageCodecJSON = "json"
UDPPacketCodecBinary = "binary-v1"
DefaultMaxFramePayloadSize = 64 * 1024 DefaultMaxFramePayloadSize = 64 * 1024
MagicV2 = "FRP\x00\x02\r\n" MagicV2 = "FRP\x00\x02\r\n"
@@ -182,7 +183,8 @@ type ClientCapabilities struct {
} }
type MessageCapabilities struct { type MessageCapabilities struct {
Codecs []string `json:"codecs,omitempty"` Codecs []string `json:"codecs,omitempty"`
UDPPacketCodecs []string `json:"udpPacketCodecs,omitempty"`
} }
type CryptoCapabilities struct { type CryptoCapabilities struct {
@@ -201,7 +203,8 @@ type ServerSelection struct {
} }
type MessageSelection struct { type MessageSelection struct {
Codec string `json:"codec,omitempty"` Codec string `json:"codec,omitempty"`
UDPPacketCodec string `json:"udpPacketCodec,omitempty"`
} }
type CryptoSelection struct { type CryptoSelection struct {
@@ -214,7 +217,8 @@ func clientHelloWithCryptoRandom(bootstrap BootstrapInfo, clientRandom []byte) C
Bootstrap: bootstrap, Bootstrap: bootstrap,
Capabilities: ClientCapabilities{ Capabilities: ClientCapabilities{
Message: MessageCapabilities{ Message: MessageCapabilities{
Codecs: []string{MessageCodecJSON}, Codecs: []string{MessageCodecJSON},
UDPPacketCodecs: []string{UDPPacketCodecBinary},
}, },
Crypto: CryptoCapabilities{ Crypto: CryptoCapabilities{
Algorithms: PreferredAEADAlgorithms(), Algorithms: PreferredAEADAlgorithms(),
+30
View File
@@ -148,10 +148,40 @@ func TestNewServerHelloSelectsFirstSupportedAEADAlgorithm(t *testing.T) {
serverHello, err := NewServerHello(hello) serverHello, err := NewServerHello(hello)
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, MessageCodecJSON, serverHello.Selected.Message.Codec) require.Equal(t, MessageCodecJSON, serverHello.Selected.Message.Codec)
require.Equal(t, UDPPacketCodecBinary, serverHello.Selected.Message.UDPPacketCodec)
require.Equal(t, AEADAlgorithmXChaCha20Poly1305, serverHello.Selected.Crypto.Algorithm) require.Equal(t, AEADAlgorithmXChaCha20Poly1305, serverHello.Selected.Crypto.Algorithm)
require.Len(t, serverHello.Selected.Crypto.ServerRandom, CryptoRandomSize) require.Len(t, serverHello.Selected.Crypto.ServerRandom, CryptoRandomSize)
} }
func TestUDPPacketCodecNegotiationFallbackAndValidation(t *testing.T) {
hello := mustClientHello(t, BootstrapInfo{})
serverHello, err := NewServerHello(hello)
require.NoError(t, err)
require.Equal(t, UDPPacketCodecBinary, serverHello.Selected.Message.UDPPacketCodec)
require.NoError(t, ValidateServerHelloForClient(hello, serverHello))
legacyHello := hello
legacyHello.Capabilities.Message.UDPPacketCodecs = nil
legacyServerHello, err := NewServerHello(legacyHello)
require.NoError(t, err)
require.Empty(t, legacyServerHello.Selected.Message.UDPPacketCodec)
require.NoError(t, ValidateServerHelloForClient(legacyHello, legacyServerHello))
unknownOffer := hello
unknownOffer.Capabilities.Message.UDPPacketCodecs = []string{"unknown"}
unknownServerHello, err := NewServerHello(unknownOffer)
require.NoError(t, err)
require.Empty(t, unknownServerHello.Selected.Message.UDPPacketCodec)
rejected := serverHello
rejected.Selected.Message.UDPPacketCodec = "unknown"
require.ErrorContains(t, ValidateServerHelloForClient(hello, rejected), "unsupported selected UDP packet codec")
unadvertised := serverHello
unadvertised.Selected.Message.UDPPacketCodec = UDPPacketCodecBinary
require.ErrorContains(t, ValidateServerHelloForClient(legacyHello, unadvertised), "was not advertised")
}
func TestNewClientCryptoContextValidatesServerHello(t *testing.T) { func TestNewClientCryptoContextValidatesServerHello(t *testing.T) {
hello := mustClientHello(t, BootstrapInfo{}) hello := mustClientHello(t, BootstrapInfo{})
serverHello, err := NewServerHello(hello) serverHello, err := NewServerHello(hello)
+21 -5
View File
@@ -16,7 +16,6 @@ package ssh
import ( import (
"context" "context"
"encoding/binary"
"errors" "errors"
"fmt" "fmt"
"net" "net"
@@ -52,6 +51,11 @@ type tcpipForward struct {
Port uint32 Port uint32
} }
// https://datatracker.ietf.org/doc/html/rfc4254#section-6.5
type execPayload struct {
Command string
}
// https://datatracker.ietf.org/doc/html/rfc4254#page-16 // https://datatracker.ietf.org/doc/html/rfc4254#page-16
type forwardedTCPPayload struct { type forwardedTCPPayload struct {
Addr string Addr string
@@ -66,6 +70,7 @@ type TunnelServer struct {
sshConn *ssh.ServerConn sshConn *ssh.ServerConn
sc *ssh.ServerConfig sc *ssh.ServerConfig
firstChannel ssh.Channel firstChannel ssh.Channel
firstChannelMu sync.Mutex
vc *virtual.Client vc *virtual.Client
peerServerListener *netpkg.InternalListener peerServerListener *netpkg.InternalListener
@@ -187,6 +192,8 @@ func (s *TunnelServer) Run() error {
} }
func (s *TunnelServer) writeToClient(data string) { func (s *TunnelServer) writeToClient(data string) {
s.firstChannelMu.Lock()
defer s.firstChannelMu.Unlock()
if s.firstChannel == nil { if s.firstChannel == nil {
return return
} }
@@ -300,23 +307,24 @@ func (s *TunnelServer) handleNewChannel(channel ssh.NewChannel, extraPayloadCh c
if err != nil { if err != nil {
return return
} }
s.firstChannelMu.Lock()
if s.firstChannel == nil { if s.firstChannel == nil {
s.firstChannel = ch s.firstChannel = ch
} }
s.firstChannelMu.Unlock()
go s.keepAlive(ch) go s.keepAlive(ch)
for req := range reqs { for req := range reqs {
if req.WantReply { if req.WantReply {
_ = req.Reply(true, nil) _ = req.Reply(true, nil)
} }
if req.Type != "exec" || len(req.Payload) <= 4 { if req.Type != "exec" {
continue continue
} }
end := 4 + binary.BigEndian.Uint32(req.Payload[:4]) extraPayload, ok := parseExecPayload(req.Payload)
if len(req.Payload) < int(end) { if !ok {
continue continue
} }
extraPayload := string(req.Payload[4:end])
select { select {
case extraPayloadCh <- extraPayload: case extraPayloadCh <- extraPayload:
default: 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) { func (s *TunnelServer) keepAlive(ch ssh.Channel) {
tk := time.NewTicker(time.Second * 30) tk := time.NewTicker(time.Second * 30)
defer tk.Stop() defer tk.Stop()
+115
View File
@@ -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")
}
}
+30
View File
@@ -26,6 +26,12 @@ type GeneralResponse struct {
Msg string Msg string
} }
type V2Response struct {
Code int `json:"code"`
Msg string `json:"msg"`
Data any `json:"data"`
}
// APIHandler is a handler function that returns a response object or an error. // APIHandler is a handler function that returns a response object or an error.
type APIHandler func(ctx *Context) (any, error) type APIHandler func(ctx *Context) (any, error)
@@ -64,3 +70,27 @@ func MakeHTTPHandlerFunc(handler APIHandler) http.HandlerFunc {
} }
} }
} }
// MakeHTTPHandlerFuncV2 wraps a handler response in the dashboard API v2 envelope.
func MakeHTTPHandlerFuncV2(handler APIHandler) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := NewContext(w, r)
res, err := handler(ctx)
if err != nil {
log.Warnf("http response [%s]: error: %v", r.URL.Path, err)
code := http.StatusInternalServerError
if e, ok := err.(*Error); ok {
code = e.Code
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(code)
_ = json.NewEncoder(w).Encode(V2Response{Code: code, Msg: err.Error(), Data: nil})
return
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_ = json.NewEncoder(w).Encode(V2Response{Code: http.StatusOK, Msg: "success", Data: res})
}
}
+37
View File
@@ -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)
}
+65
View File
@@ -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())
})
}
}
+6
View File
@@ -35,6 +35,12 @@ func NewReader(r io.Reader, limiter *rate.Limiter) *Reader {
func (r *Reader) Read(p []byte) (n int, err error) { func (r *Reader) Read(p []byte) (n int, err error) {
b := r.limiter.Burst() b := r.limiter.Burst()
if b <= 0 {
if len(p) == 0 {
return 0, nil
}
return 0, invalidBurstError(b)
}
if b < len(p) { if b < len(p) {
p = p[:b] p = p[:b]
} }
+7
View File
@@ -34,8 +34,15 @@ func NewWriter(w io.Writer, limiter *rate.Limiter) *Writer {
} }
func (w *Writer) Write(p []byte) (n int, err error) { func (w *Writer) Write(p []byte) (n int, err error) {
if len(p) == 0 {
return 0, nil
}
var nn int var nn int
b := w.limiter.Burst() b := w.limiter.Burst()
if b <= 0 {
return 0, invalidBurstError(b)
}
for { for {
end := len(p) end := len(p)
if end == 0 { if end == 0 {
+3 -8
View File
@@ -16,6 +16,7 @@ package net
import ( import (
"context" "context"
"crypto/hkdf"
"crypto/sha256" "crypto/sha256"
"errors" "errors"
"io" "io"
@@ -25,7 +26,6 @@ import (
libcrypto "github.com/fatedier/golib/crypto" libcrypto "github.com/fatedier/golib/crypto"
quic "github.com/quic-go/quic-go" quic "github.com/quic-go/quic-go"
"golang.org/x/crypto/hkdf"
"github.com/fatedier/frp/pkg/util/xlog" "github.com/fatedier/frp/pkg/util/xlog"
) )
@@ -335,11 +335,6 @@ func deriveAEADControlKeys(key []byte, algorithm string, transcriptHash []byte)
} }
func deriveAEADControlKey(key []byte, algorithm string, transcriptHash []byte, direction string) ([]byte, error) { func deriveAEADControlKey(key []byte, algorithm string, transcriptHash []byte, direction string) ([]byte, error) {
info := []byte(aeadControlHKDFInfoPrefix + " " + algorithm + " " + direction) info := aeadControlHKDFInfoPrefix + " " + algorithm + " " + direction
reader := hkdf.New(sha256.New, key, transcriptHash, info) return hkdf.Key(sha256.New, key, transcriptHash, info, libcrypto.AEADKeySize)
out := make([]byte, libcrypto.AEADKeySize)
if _, err := io.ReadFull(reader, out); err != nil {
return nil, err
}
return out, nil
} }
+6
View File
@@ -114,5 +114,11 @@ func TestDeriveAEADControlKeysUsesDistinctDirections(t *testing.T) {
bytes.Repeat([]byte{0x44}, 32), bytes.Repeat([]byte{0x44}, 32),
) )
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, []byte{
0xa0, 0x58, 0xcd, 0x02, 0x5d, 0x96, 0x98, 0x5f,
0xeb, 0xeb, 0xff, 0x79, 0xa1, 0x9f, 0x62, 0xb7,
0x15, 0xe0, 0x53, 0x91, 0x3d, 0xfc, 0x74, 0x77,
0x05, 0x91, 0x4c, 0x62, 0x4b, 0xf3, 0xd4, 0x95,
}, clientToServerKey)
require.NotEqual(t, clientToServerKey, serverToClientKey) require.NotEqual(t, clientToServerKey, serverToClientKey)
} }
+5
View File
@@ -45,6 +45,11 @@ func DialHookWebsocket(protocol string, host string) libnet.AfterHookFunc {
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
// The tunnel payload is a raw byte stream (yamux), not UTF-8 text.
// Send it as binary frames; otherwise RFC 6455-compliant intermediaries
// (e.g. API gateways/reverse proxies) UTF-8-validate the default text
// frames and close the connection on invalid bytes.
conn.PayloadType = websocket.BinaryFrame
return ctx, conn, nil return ctx, conn, nil
} }
} }
+5
View File
@@ -32,6 +32,11 @@ func NewWebsocketListener(ln net.Listener) (wl *WebsocketListener) {
muxer := http.NewServeMux() muxer := http.NewServeMux()
muxer.Handle(FrpWebsocketPath, websocket.Handler(func(c *websocket.Conn) { muxer.Handle(FrpWebsocketPath, websocket.Handler(func(c *websocket.Conn) {
// The tunnel payload is a raw byte stream (yamux), not UTF-8 text.
// Send it as binary frames; otherwise RFC 6455-compliant intermediaries
// (e.g. API gateways/reverse proxies) UTF-8-validate the default text
// frames and close the connection on invalid bytes.
c.PayloadType = websocket.BinaryFrame
notifyCh := make(chan struct{}) notifyCh := make(chan struct{})
conn := WrapCloseNotifyConn(c, func(_ error) { conn := WrapCloseNotifyConn(c, func(_ error) {
close(notifyCh) close(notifyCh)
+1 -1
View File
@@ -14,7 +14,7 @@
package version package version
var version = "0.69.0" var version = "0.71.0"
func Full() string { func Full() string {
return version return version
+1 -3
View File
@@ -28,8 +28,6 @@ import (
libio "github.com/fatedier/golib/io" libio "github.com/fatedier/golib/io"
"github.com/fatedier/golib/pool" "github.com/fatedier/golib/pool"
"golang.org/x/net/http2"
"golang.org/x/net/http2/h2c"
httppkg "github.com/fatedier/frp/pkg/util/http" httppkg "github.com/fatedier/frp/pkg/util/http"
"github.com/fatedier/frp/pkg/util/log" "github.com/fatedier/frp/pkg/util/log"
@@ -139,7 +137,7 @@ func NewHTTPReverseProxy(option HTTPReverseProxyOptions, vhostRouter *Routers) *
_, _ = rw.Write(getNotFoundPageContent()) _, _ = rw.Write(getNotFoundPageContent())
}, },
} }
rp.proxy = h2c.NewHandler(proxy, &http2.Server{}) rp.proxy = proxy
return rp return rp
} }
+80
View File
@@ -1,14 +1,94 @@
package vhost package vhost
import ( import (
"bufio"
"fmt"
"net"
"net/http"
"net/http/httptest" "net/http/httptest"
"testing" "testing"
"time"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
httppkg "github.com/fatedier/frp/pkg/util/http" httppkg "github.com/fatedier/frp/pkg/util/http"
) )
func TestHTTPServerProtocols(t *testing.T) {
rp := NewHTTPReverseProxy(HTTPReverseProxyOptions{}, NewRouters())
protocols := new(http.Protocols)
protocols.SetHTTP1(true)
protocols.SetUnencryptedHTTP2(true)
server := &http.Server{
Handler: rp,
ReadHeaderTimeout: time.Second,
Protocols: protocols,
}
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
serveErr := make(chan error, 1)
go func() {
serveErr <- server.Serve(listener)
}()
defer func() {
require.NoError(t, server.Close())
require.ErrorIs(t, <-serveErr, http.ErrServerClosed)
}()
require.True(t, server.Protocols.HTTP1())
require.True(t, server.Protocols.UnencryptedHTTP2())
t.Run("HTTP/1.1", func(t *testing.T) {
transport := &http.Transport{Protocols: httpProtocols(true, false)}
defer transport.CloseIdleConnections()
client := &http.Client{Transport: transport}
response, err := client.Get("http://" + listener.Addr().String() + "/")
require.NoError(t, err)
defer response.Body.Close()
require.Equal(t, "HTTP/1.1", response.Proto)
require.Equal(t, http.StatusNotFound, response.StatusCode)
})
t.Run("HTTP/2 prior knowledge", func(t *testing.T) {
transport := &http.Transport{Protocols: httpProtocols(false, true)}
defer transport.CloseIdleConnections()
client := &http.Client{Transport: transport}
response, err := client.Get("http://" + listener.Addr().String() + "/")
require.NoError(t, err)
defer response.Body.Close()
require.Equal(t, "HTTP/2.0", response.Proto)
require.Equal(t, http.StatusNotFound, response.StatusCode)
})
t.Run("HTTP/1.1 Upgrade h2c", func(t *testing.T) {
conn, err := net.Dial("tcp", listener.Addr().String())
require.NoError(t, err)
defer conn.Close()
_, err = fmt.Fprintf(conn,
"GET / HTTP/1.1\r\nHost: %s\r\n"+
"Connection: Upgrade, HTTP2-Settings\r\nUpgrade: h2c\r\n"+
"HTTP2-Settings: AAMAAABkAAQCAAAAAAIAAAAA\r\n\r\n",
listener.Addr())
require.NoError(t, err)
response, err := http.ReadResponse(bufio.NewReader(conn), nil)
require.NoError(t, err)
defer response.Body.Close()
require.NotEqual(t, http.StatusSwitchingProtocols, response.StatusCode)
})
}
func httpProtocols(http1, unencryptedHTTP2 bool) *http.Protocols {
protocols := new(http.Protocols)
protocols.SetHTTP1(http1)
protocols.SetUnencryptedHTTP2(unencryptedHTTP2)
return protocols
}
func TestCheckRouteAuthByRequest(t *testing.T) { func TestCheckRouteAuthByRequest(t *testing.T) {
rc := &RouteConfig{ rc := &RouteConfig{
Username: "alice", Username: "alice",
+5 -5
View File
@@ -96,21 +96,21 @@ func (l *Logger) Spawn() *Logger {
} }
func (l *Logger) Errorf(format string, v ...any) { func (l *Logger) Errorf(format string, v ...any) {
log.Logger.Errorf(l.prefixString+format, v...) log.Logger.WithPrefix(l.prefixString).Errorf(format, v...)
} }
func (l *Logger) Warnf(format string, v ...any) { func (l *Logger) Warnf(format string, v ...any) {
log.Logger.Warnf(l.prefixString+format, v...) log.Logger.WithPrefix(l.prefixString).Warnf(format, v...)
} }
func (l *Logger) Infof(format string, v ...any) { func (l *Logger) Infof(format string, v ...any) {
log.Logger.Infof(l.prefixString+format, v...) log.Logger.WithPrefix(l.prefixString).Infof(format, v...)
} }
func (l *Logger) Debugf(format string, v ...any) { func (l *Logger) Debugf(format string, v ...any) {
log.Logger.Debugf(l.prefixString+format, v...) log.Logger.WithPrefix(l.prefixString).Debugf(format, v...)
} }
func (l *Logger) Tracef(format string, v ...any) { func (l *Logger) Tracef(format string, v ...any) {
log.Logger.Tracef(l.prefixString+format, v...) log.Logger.WithPrefix(l.prefixString).Tracef(format, v...)
} }
+76
View File
@@ -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 xlog
import (
"bytes"
"testing"
goliblog "github.com/fatedier/golib/log"
"github.com/stretchr/testify/require"
frplog "github.com/fatedier/frp/pkg/util/log"
)
func TestPrefixIsNotPartOfFormatString(t *testing.T) {
tests := []struct {
name string
log func(*Logger)
}{
{name: "error", log: func(xl *Logger) { xl.Errorf("value [%s]", "ok") }},
{name: "warn", log: func(xl *Logger) { xl.Warnf("value [%s]", "ok") }},
{name: "info", log: func(xl *Logger) { xl.Infof("value [%s]", "ok") }},
{name: "debug", log: func(xl *Logger) { xl.Debugf("value [%s]", "ok") }},
{name: "trace", log: func(xl *Logger) { xl.Tracef("value [%s]", "ok") }},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
output := captureLogs(t)
tt.log(New().AppendPrefix("%1000000s"))
require.Contains(t, output.String(), "[%1000000s] value [ok]")
require.Less(t, output.Len(), 1024)
})
}
}
func TestFormattingSemanticsArePreserved(t *testing.T) {
output := captureLogs(t)
xl := New().AppendPrefix("run")
xl.Infof("%[2]s %[1]s", "first", "second")
xl.Infof("100% complete")
require.Contains(t, output.String(), "[run] second first")
require.Contains(t, output.String(), "[run] 100% complete")
}
func captureLogs(t *testing.T) *bytes.Buffer {
t.Helper()
output := bytes.NewBuffer(nil)
oldLogger := frplog.Logger
frplog.Logger = goliblog.New(
goliblog.WithOutput(output),
goliblog.WithLevel(goliblog.TraceLevel),
goliblog.WithCaller(false),
)
t.Cleanup(func() {
frplog.Logger = oldLogger
})
return output
}
+11
View File
@@ -48,6 +48,17 @@ func (svr *Service) registerRouteHandlers(helper *httppkg.RouterRegisterHelper)
subRouter.HandleFunc("/api/clients/{key}", httppkg.MakeHTTPHandlerFunc(apiController.APIClientDetail)).Methods("GET") subRouter.HandleFunc("/api/clients/{key}", httppkg.MakeHTTPHandlerFunc(apiController.APIClientDetail)).Methods("GET")
subRouter.HandleFunc("/api/proxies", httppkg.MakeHTTPHandlerFunc(apiController.DeleteProxies)).Methods("DELETE") subRouter.HandleFunc("/api/proxies", httppkg.MakeHTTPHandlerFunc(apiController.DeleteProxies)).Methods("DELETE")
subRouter.HandleFunc("/api/v2/users", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2UserList)).Methods("GET")
subRouter.HandleFunc("/api/v2/system/info", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2SystemInfo)).Methods("GET")
subRouter.HandleFunc("/api/v2/system/prune", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2SystemPrune)).Methods("POST")
subRouter.HandleFunc("/api/v2/clients", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ClientList)).Methods("GET")
v2EncodedPathRouter := subRouter.NewRoute().Subrouter()
v2EncodedPathRouter.UseEncodedPath()
v2EncodedPathRouter.HandleFunc("/api/v2/clients/{key}", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ClientDetail)).Methods("GET")
subRouter.HandleFunc("/api/v2/proxies", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyList)).Methods("GET")
v2EncodedPathRouter.HandleFunc("/api/v2/proxies/{name}/traffic", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyTraffic)).Methods("GET")
v2EncodedPathRouter.HandleFunc("/api/v2/proxies/{name}", httppkg.MakeHTTPHandlerFuncV2(apiController.APIV2ProxyDetail)).Methods("GET")
// view // view
subRouter.Handle("/favicon.ico", http.FileServer(helper.AssetsFS)).Methods("GET") subRouter.Handle("/favicon.ico", http.FileServer(helper.AssetsFS)).Methods("GET")
subRouter.PathPrefix("/static/").Handler( subRouter.PathPrefix("/static/").Handler(
+475 -87
View File
@@ -17,6 +17,8 @@ package server
import ( import (
"context" "context"
"fmt" "fmt"
"math"
"net"
"runtime/debug" "runtime/debug"
"sync" "sync"
"sync/atomic" "sync/atomic"
@@ -40,55 +42,313 @@ import (
"github.com/fatedier/frp/server/registry" "github.com/fatedier/frp/server/registry"
) )
type ControlID uint64
var nextControlID atomic.Uint64
const workConnPoolCapacityOffset = 10
type controlEntry struct {
ctl *Control
id ControlID
// runMu serializes lifecycle and routing decisions for one run ID.
// Replacements inherit it; removing the entry releases the manager's reference.
runMu *sync.Mutex
registryOnline bool
registryControlID ControlID
}
type ControlManager struct { type ControlManager struct {
// controls indexed by run id // controls indexed by run id
ctlsByRunID map[string]*Control ctlsByRunID map[string]*controlEntry
registry *registry.ClientRegistry
closed bool
mu sync.RWMutex mu sync.RWMutex
} }
func NewControlManager() *ControlManager { func NewControlManager(clientRegistry *registry.ClientRegistry) *ControlManager {
return &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) { // lockCurrentRun returns the current entry with its run gate held. It never
cm.mu.Lock() // waits for the gate while holding cm.mu and revalidates the gate after waiting.
defer cm.mu.Unlock() // The global order is runMu, cm.mu, ctl.lifecycleMu, then registry locks.
func (cm *ControlManager) lockCurrentRun(runID string, allowClosed bool) (*controlEntry, bool) {
var ok bool cm.mu.RLock()
old, ok = cm.ctlsByRunID[runID] entry, ok := cm.ctlsByRunID[runID]
if ok { if cm.closed && !allowClosed {
old.Replaced(ctl) ok = false
} }
cm.ctlsByRunID[runID] = ctl cm.mu.RUnlock()
return 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 // Add makes ctl the pending current generation and records the predecessor
func (cm *ControlManager) Del(runID string, ctl *Control) { // finalization barrier it must wait for before activation.
func (cm *ControlManager) Add(ctl *Control) error {
for {
// Never wait for a run gate while holding cm.mu.
cm.mu.RLock()
old := cm.ctlsByRunID[ctl.runID]
cm.mu.RUnlock()
if old != nil {
old.runMu.Lock()
}
cm.mu.Lock()
if cm.closed {
cm.mu.Unlock()
if old != nil {
old.runMu.Unlock()
}
return fmt.Errorf("control manager is closed")
}
if cm.ctlsByRunID[ctl.runID] != old {
cm.mu.Unlock()
if old != nil {
old.runMu.Unlock()
}
continue
}
id := ControlID(nextControlID.Add(1))
if err := ctl.admit(cm, id); err != nil {
cm.mu.Unlock()
if old != nil {
old.runMu.Unlock()
}
return err
}
runMu := &sync.Mutex{}
if old != nil {
runMu = old.runMu
}
entry := &controlEntry{ctl: ctl, id: id, runMu: runMu}
var (
oldCtl *Control
barrier <-chan struct{}
)
if old != nil {
oldCtl = old.ctl
barrier = oldCtl.markReplaced()
ctl.setHandoffBarrier(barrier)
entry.registryOnline = old.registryOnline
entry.registryControlID = old.registryControlID
}
cm.ctlsByRunID[ctl.runID] = entry
cm.mu.Unlock()
if old != nil {
old.runMu.Unlock()
}
if oldCtl != nil {
oldCtl.Replaced(ctl)
}
return nil
}
}
// Activate registers ctl as online only if it is still the pending current
// generation.
func (cm *ControlManager) Activate(ctl *Control) (bool, error) {
entry, ok := cm.lockCurrentRun(ctl.runID, false)
if !ok {
return false, nil
}
defer entry.runMu.Unlock()
cm.mu.Lock() cm.mu.Lock()
defer cm.mu.Unlock() defer cm.mu.Unlock()
if 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) { func (cm *ControlManager) GetByID(runID string) (ctl *Control, ok bool) {
cm.mu.RLock() entry, ok := cm.lockCurrentRun(runID, false)
defer cm.mu.RUnlock() if !ok {
ctl, ok = cm.ctlsByRunID[runID] return nil, false
return }
defer entry.runMu.Unlock()
ctl = entry.ctl
ctl.lifecycleMu.Lock()
defer ctl.lifecycleMu.Unlock()
if ctl.state != controlStateRunning {
return nil, false
}
return ctl, true
}
// admitVisitorByRunID commits a visitor admission against the current running
// control while its run and lifecycle ownership are held. The callback must
// only perform the in-memory, buffered visitor admission.
func (cm *ControlManager) admitVisitorByRunID(runID string, admit func(user, wireProtocol, udpPacketCodec string) error) (bool, error) {
entry, ok := cm.lockCurrentRun(runID, false)
if !ok {
return false, nil
}
defer entry.runMu.Unlock()
ctl := entry.ctl
ctl.lifecycleMu.Lock()
defer ctl.lifecycleMu.Unlock()
if ctl.state != controlStateRunning {
return false, nil
}
return true, admit(ctl.sessionCtx.LoginMsg.User, ctl.sessionCtx.WireProtocol, ctl.sessionCtx.UDPPacketCodec)
}
// RegisterWorkConn transfers conn to ctl only if ctl is still the current
// running generation. On error, ownership remains with the caller.
func (cm *ControlManager) RegisterWorkConn(ctl *Control, conn *proxy.WorkConn) error {
entry, ok := cm.lockCurrentRun(ctl.runID, false)
if !ok {
cm.mu.RLock()
closed := cm.closed
cm.mu.RUnlock()
if closed {
return fmt.Errorf("control manager is closed")
}
return fmt.Errorf("client control for run id [%s] is no longer current", ctl.runID)
}
defer entry.runMu.Unlock()
if entry.ctl != ctl || entry.id != ctl.controlID {
return fmt.Errorf("client control for run id [%s] is no longer current", ctl.runID)
}
ctl.lifecycleMu.Lock()
defer ctl.lifecycleMu.Unlock()
if ctl.state != controlStateRunning {
return fmt.Errorf("client control for run id [%s] is not running", ctl.runID)
}
select {
case ctl.workConnCh <- conn:
ctl.xl.Debugf("new work connection registered")
return nil
default:
ctl.xl.Debugf("work connection pool is full, discarding")
return fmt.Errorf("work connection pool is full, discarding")
}
} }
func (cm *ControlManager) Close() error { func (cm *ControlManager) Close() error {
cm.mu.Lock() cm.mu.Lock()
defer cm.mu.Unlock() cm.closed = true
for _, ctl := range cm.ctlsByRunID { ctls := make([]*Control, 0, len(cm.ctlsByRunID))
ctl.Close() 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 return nil
} }
@@ -110,12 +370,21 @@ type SessionContext struct {
LoginMsg *msg.Login LoginMsg *msg.Login
// server configuration // server configuration
ServerCfg *v1.ServerConfig ServerCfg *v1.ServerConfig
// client registry
ClientRegistry *registry.ClientRegistry
// negotiated wire protocol for this client session // negotiated wire protocol for this client session
WireProtocol string WireProtocol string
UDPPacketCodec string
} }
type controlState uint8
const (
controlStateCreated controlState = iota
controlStatePending
controlStateRunning
controlStateClosing
controlStateClosed
)
type Control struct { type Control struct {
// session context // session context
sessionCtx *SessionContext sessionCtx *SessionContext
@@ -142,30 +411,59 @@ type Control struct {
// last time got the Ping message // last time got the Ping message
lastPing atomic.Value lastPing atomic.Value
// A new run id will be generated when a new client login. // runID never changes during the lifetime of a control. controlID is assigned
// If run id got from login message has same run id, it means it's the same client, so we can // once by ControlManager and distinguishes same-runID generations.
// replace old controller instantly. runID string
runID string controlID ControlID
manager *ControlManager
lifecycleMu sync.Mutex
state controlState
activated bool
handoffBarrier <-chan struct{}
interruptOnce sync.Once
interruptErr error
mu sync.RWMutex mu sync.RWMutex
xl *xlog.Logger xl *xlog.Logger
ctx context.Context ctx context.Context
doneCh chan struct{} doneCh chan struct{}
serverMetrics metrics.ServerMetrics
} }
func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, error) { func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, error) {
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{ ctl := &Control{
sessionCtx: sessionCtx, sessionCtx: sessionCtx,
workConnCh: make(chan *proxy.WorkConn, poolCount+10), workConnCh: make(chan *proxy.WorkConn, poolCount+workConnPoolCapacityOffset),
proxies: make(map[string]proxy.Proxy), proxies: make(map[string]proxy.Proxy),
poolCount: poolCount, poolCount: poolCount,
portsUsedNum: 0, portsUsedNum: 0,
runID: sessionCtx.LoginMsg.RunID, runID: sessionCtx.LoginMsg.RunID,
xl: xlog.FromContextSafe(ctx), state: controlStateCreated,
ctx: ctx, xl: xlog.FromContextSafe(ctx),
doneCh: make(chan struct{}), ctx: ctx,
doneCh: make(chan struct{}),
serverMetrics: metrics.Server,
} }
ctl.lastPing.Store(time.Now()) ctl.lastPing.Store(time.Now())
@@ -175,48 +473,121 @@ func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, erro
return ctl, nil return ctl, nil
} }
// Start starts the control session workers after login succeeds. func (ctl *Control) RunID() string {
func (ctl *Control) Start() { return ctl.runID
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) Close() error { func (ctl *Control) ID() ControlID {
ctl.sessionCtx.Conn.Close() 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 return nil
} }
func (ctl *Control) Replaced(newCtl *Control) { func (ctl *Control) setHandoffBarrier(barrier <-chan struct{}) {
xl := ctl.xl ctl.lifecycleMu.Lock()
xl.Infof("replaced by client [%s]", newCtl.runID) ctl.handoffBarrier = barrier
ctl.runID = "" ctl.lifecycleMu.Unlock()
ctl.sessionCtx.Conn.Close()
} }
func (ctl *Control) RegisterWorkConn(conn *proxy.WorkConn) error { func (ctl *Control) WaitForHandoff() {
xl := ctl.xl ctl.lifecycleMu.Lock()
defer func() { barrier := ctl.handoffBarrier
if err := recover(); err != nil { ctl.lifecycleMu.Unlock()
xl.Errorf("panic error: %v", err) if barrier != nil {
xl.Errorf(string(debug.Stack())) <-barrier
}
}()
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")
} }
} }
// 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. // When frps get one user connection, we get one work connection from the pool and return it.
// If no workConn available in the pool, send message to frpc to get one or more // If no workConn available in the pool, send message to frpc to get one or more
// and wait until it is available. // and wait until it is available.
@@ -271,10 +642,10 @@ func (ctl *Control) heartbeatWorker() {
} }
xl := ctl.xl 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 { if time.Since(ctl.lastPing.Load().(time.Time)) > time.Duration(ctl.sessionCtx.ServerCfg.Transport.HeartbeatTimeout)*time.Second {
xl.Warnf("heartbeat timeout") xl.Warnf("heartbeat timeout")
ctl.sessionCtx.Conn.Close() _ = ctl.Close()
return return
} }
}, time.Second, ctl.doneCh) }, time.Second, ctl.doneCh)
@@ -289,14 +660,14 @@ func (ctl *Control) loginUserInfo() plugin.UserInfo {
return plugin.UserInfo{ return plugin.UserInfo{
User: ctl.sessionCtx.LoginMsg.User, User: ctl.sessionCtx.LoginMsg.User,
Metas: ctl.sessionCtx.LoginMsg.Metas, Metas: ctl.sessionCtx.LoginMsg.Metas,
RunID: ctl.sessionCtx.LoginMsg.RunID, RunID: ctl.runID,
} }
} }
func (ctl *Control) closeProxy(pxy proxy.Proxy) { func (ctl *Control) closeProxy(pxy proxy.Proxy) {
pxy.Close() pxy.Close()
ctl.sessionCtx.PxyManager.Del(pxy.GetName()) ctl.sessionCtx.PxyManager.Del(pxy.GetName())
metrics.Server.CloseProxy(pxy.GetName(), pxy.GetConfigurer().GetBaseConfig().Type) ctl.serverMetrics.CloseProxy(pxy.GetName(), pxy.GetConfigurer().GetBaseConfig().Type)
notifyContent := &plugin.CloseProxyContent{ notifyContent := &plugin.CloseProxyContent{
User: ctl.loginUserInfo(), User: ctl.loginUserInfo(),
@@ -311,12 +682,24 @@ func (ctl *Control) closeProxy(pxy proxy.Proxy) {
func (ctl *Control) worker() { func (ctl *Control) worker() {
xl := ctl.xl xl := ctl.xl
ctl.serverMetrics.NewClient()
go ctl.heartbeatWorker() go ctl.heartbeatWorker()
go ctl.msgDispatcher.Run() go ctl.msgDispatcher.Run()
go func() {
for i := 0; i < ctl.poolCount; i++ {
// Ignore the error: it means this control is already closing.
_ = ctl.msgDispatcher.Send(&msg.ReqWorkConn{})
}
}()
<-ctl.msgDispatcher.Done() <-ctl.msgDispatcher.Done()
ctl.sessionCtx.Conn.Close() ctl.lifecycleMu.Lock()
if ctl.state == controlStateRunning {
ctl.state = controlStateClosing
}
ctl.lifecycleMu.Unlock()
_ = ctl.interruptReadAndClose()
ctl.mu.Lock() ctl.mu.Lock()
close(ctl.workConnCh) close(ctl.workConnCh)
@@ -331,10 +714,14 @@ func (ctl *Control) worker() {
ctl.closeProxy(pxy) ctl.closeProxy(pxy)
} }
metrics.Server.CloseClient() ctl.serverMetrics.CloseClient()
ctl.sessionCtx.ClientRegistry.MarkOfflineByRunID(ctl.runID) if ctl.manager != nil {
ctl.manager.Remove(ctl)
}
xl.Infof("client exit success") xl.Infof("client exit success")
close(ctl.doneCh) ctl.lifecycleMu.Lock()
ctl.finishLocked()
ctl.lifecycleMu.Unlock()
} }
func (ctl *Control) registerMsgHandlers() { func (ctl *Control) registerMsgHandlers() {
@@ -374,9 +761,9 @@ func (ctl *Control) handleNewProxy(m msg.Message) {
xl.Infof("new proxy [%s] type [%s] success", inMsg.ProxyName, inMsg.ProxyType) xl.Infof("new proxy [%s] type [%s] success", inMsg.ProxyName, inMsg.ProxyType)
clientID := ctl.sessionCtx.LoginMsg.ClientID clientID := ctl.sessionCtx.LoginMsg.ClientID
if clientID == "" { if clientID == "" {
clientID = ctl.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) _ = ctl.msgDispatcher.Send(resp)
} }
@@ -455,6 +842,7 @@ func (ctl *Control) RegisterProxy(pxyMsg *msg.NewProxy) (remoteAddr string, err
ServerCfg: ctl.sessionCtx.ServerCfg, ServerCfg: ctl.sessionCtx.ServerCfg,
EncryptionKey: ctl.sessionCtx.EncryptionKey, EncryptionKey: ctl.sessionCtx.EncryptionKey,
WireProtocol: ctl.sessionCtx.WireProtocol, WireProtocol: ctl.sessionCtx.WireProtocol,
UDPPacketCodec: ctl.sessionCtx.UDPPacketCodec,
}) })
if err != nil { if err != nil {
return remoteAddr, err return remoteAddr, err
+595
View File
@@ -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
}
+5 -3
View File
@@ -58,8 +58,12 @@ func NewController(
// /api/serverinfo // /api/serverinfo
func (c *Controller) APIServerInfo(ctx *httppkg.Context) (any, error) { func (c *Controller) APIServerInfo(ctx *httppkg.Context) (any, error) {
return c.buildServerInfoResp(), nil
}
func (c *Controller) buildServerInfoResp() model.ServerInfoResp {
serverStats := mem.StatsCollector.GetServer() serverStats := mem.StatsCollector.GetServer()
svrResp := model.ServerInfoResp{ return model.ServerInfoResp{
Version: version.Full(), Version: version.Full(),
BindPort: c.serverCfg.BindPort, BindPort: c.serverCfg.BindPort,
VhostHTTPPort: c.serverCfg.VhostHTTPPort, VhostHTTPPort: c.serverCfg.VhostHTTPPort,
@@ -80,8 +84,6 @@ func (c *Controller) APIServerInfo(ctx *httppkg.Context) (any, error) {
ClientCounts: serverStats.ClientCounts, ClientCounts: serverStats.ClientCounts,
ProxyTypeCounts: serverStats.ProxyTypeCounts, ProxyTypeCounts: serverStats.ProxyTypeCounts,
} }
return svrResp, nil
} }
// /api/clients // /api/clients
+647
View File
@@ -0,0 +1,647 @@
// Copyright 2026 The frp Authors
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package http
import (
"cmp"
"fmt"
"maps"
"math"
"net/http"
"net/url"
"slices"
"strconv"
"strings"
"time"
v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/metrics/mem"
httppkg "github.com/fatedier/frp/pkg/util/http"
"github.com/fatedier/frp/server/http/model"
"github.com/fatedier/frp/server/registry"
)
const (
defaultV2Page = 1
defaultV2PageSize = 50
maxV2PageSize = 200
v2SystemPruneTypeOfflineProxies = "offline_proxies"
v2ProxyTrafficDefaultDays = 7
v2ProxyTrafficUnit = "bytes"
v2ProxyTrafficGranularity = "day"
)
var apiV2ProxyTypes = []string{
string(v1.ProxyTypeTCP),
string(v1.ProxyTypeUDP),
string(v1.ProxyTypeHTTP),
string(v1.ProxyTypeHTTPS),
string(v1.ProxyTypeTCPMUX),
string(v1.ProxyTypeSTCP),
string(v1.ProxyTypeXTCP),
string(v1.ProxyTypeSUDP),
}
// /api/v2/system/info
func (c *Controller) APIV2SystemInfo(ctx *httppkg.Context) (any, error) {
info := c.buildServerInfoResp()
proxyTypeCounts := info.ProxyTypeCounts
if proxyTypeCounts == nil {
proxyTypeCounts = map[string]int64{}
}
return model.V2SystemInfoResp{
Version: info.Version,
Config: model.V2SystemInfoConfigResp{
BindPort: info.BindPort,
VhostHTTPPort: info.VhostHTTPPort,
VhostHTTPSPort: info.VhostHTTPSPort,
TCPMuxHTTPConnectPort: info.TCPMuxHTTPConnectPort,
KCPBindPort: info.KCPBindPort,
QUICBindPort: info.QUICBindPort,
SubdomainHost: info.SubdomainHost,
MaxPoolCount: info.MaxPoolCount,
MaxPortsPerClient: info.MaxPortsPerClient,
HeartbeatTimeout: info.HeartBeatTimeout,
AllowPortsStr: info.AllowPortsStr,
TLSForce: info.TLSForce,
},
Status: model.V2SystemInfoStatusResp{
TotalTrafficIn: info.TotalTrafficIn,
TotalTrafficOut: info.TotalTrafficOut,
CurConns: info.CurConns,
ClientCounts: info.ClientCounts,
ProxyTypeCounts: proxyTypeCounts,
},
}, nil
}
// /api/v2/system/prune
func (c *Controller) APIV2SystemPrune(ctx *httppkg.Context) (any, error) {
pruneType, err := parseV2SystemPruneType(ctx.Query("type"))
if err != nil {
return nil, err
}
cleared, total := mem.StatsCollector.PruneOfflineProxies()
return model.V2SystemPruneResp{
Type: pruneType,
Cleared: cleared,
Total: total,
}, nil
}
// /api/v2/users
func (c *Controller) APIV2UserList(ctx *httppkg.Context) (any, error) {
page, pageSize, err := parseV2PageParams(ctx)
if err != nil {
return nil, err
}
if c.clientRegistry == nil {
return nil, fmt.Errorf("client registry unavailable")
}
userStats := make(map[string]*model.V2UserResp)
for _, info := range c.clientRegistry.List() {
item := getOrCreateV2User(userStats, info.User)
item.ClientCount++
}
for _, proxyInfo := range c.listV2ProxyStats("") {
item := getOrCreateV2User(userStats, proxyInfo.User)
item.ProxyCount++
}
q := strings.ToLower(ctx.Query("q"))
items := make([]model.V2UserResp, 0, len(userStats))
for _, item := range userStats {
if q != "" && !strings.Contains(strings.ToLower(item.User), q) {
continue
}
items = append(items, *item)
}
slices.SortFunc(items, func(a, b model.V2UserResp) int {
return cmp.Compare(a.User, b.User)
})
return buildV2PageResp(items, page, pageSize), nil
}
// /api/v2/clients
func (c *Controller) APIV2ClientList(ctx *httppkg.Context) (any, error) {
page, pageSize, err := parseV2PageParams(ctx)
if err != nil {
return nil, err
}
if c.clientRegistry == nil {
return nil, fmt.Errorf("client registry unavailable")
}
statusFilter, err := parseV2StatusFilter(ctx.Query("status"))
if err != nil {
return nil, err
}
userFilter, filterByUser := queryValue(ctx, "user")
clientIDFilter := ctx.Query("clientID")
runIDFilter := ctx.Query("runID")
q := strings.ToLower(ctx.Query("q"))
records := c.clientRegistry.List()
items := make([]model.ClientInfoResp, 0, len(records))
for _, info := range records {
if filterByUser && info.User != userFilter {
continue
}
if clientIDFilter != "" && info.ClientID() != clientIDFilter {
continue
}
if runIDFilter != "" && info.RunID != runIDFilter {
continue
}
if !matchV2StatusFilter(info.Online, statusFilter) {
continue
}
resp := buildClientInfoResp(info)
if q != "" && !matchV2ClientQuery(resp, q) {
continue
}
items = append(items, resp)
}
slices.SortFunc(items, func(a, b model.ClientInfoResp) int {
if v := cmp.Compare(a.User, b.User); v != 0 {
return v
}
if v := cmp.Compare(a.ClientID, b.ClientID); v != 0 {
return v
}
return cmp.Compare(a.Key, b.Key)
})
return buildV2PageResp(items, page, pageSize), nil
}
// /api/v2/clients/{key}
func (c *Controller) APIV2ClientDetail(ctx *httppkg.Context) (any, error) {
key, err := decodeV2PathParam(ctx, "key", "client key")
if err != nil {
return nil, err
}
if c.clientRegistry == nil {
return nil, fmt.Errorf("client registry unavailable")
}
info, ok := c.clientRegistry.GetByKey(key)
if !ok {
return nil, httppkg.NewError(http.StatusNotFound, fmt.Sprintf("client %s not found", key))
}
resp := buildClientInfoResp(info)
status := c.buildV2ClientStatus(info)
return model.V2ClientDetailResp{
ClientInfoResp: resp,
Status: status,
}, nil
}
// /api/v2/proxies
func (c *Controller) APIV2ProxyList(ctx *httppkg.Context) (any, error) {
page, pageSize, err := parseV2PageParams(ctx)
if err != nil {
return nil, err
}
statusFilter, err := parseV2StatusFilter(ctx.Query("status"))
if err != nil {
return nil, err
}
proxyType, err := parseV2ProxyTypeFilter(ctx.Query("type"))
if err != nil {
return nil, err
}
userFilter, filterByUser := queryValue(ctx, "user")
clientIDFilter := ctx.Query("clientID")
q := strings.ToLower(ctx.Query("q"))
stats := c.listV2ProxyStats(proxyType)
items := make([]model.V2ProxyResp, 0, len(stats))
for _, ps := range stats {
resp := c.buildV2ProxyResp(ps)
if filterByUser && resp.User != userFilter {
continue
}
if clientIDFilter != "" && resp.ClientID != clientIDFilter {
continue
}
if !matchV2StatusFilter(resp.Status.State == "online", statusFilter) {
continue
}
if q != "" && !matchV2ProxyQuery(resp, q) {
continue
}
items = append(items, resp)
}
slices.SortFunc(items, func(a, b model.V2ProxyResp) int {
if v := cmp.Compare(a.Spec.Type, b.Spec.Type); v != 0 {
return v
}
return cmp.Compare(a.Name, b.Name)
})
return buildV2PageResp(items, page, pageSize), nil
}
// /api/v2/proxies/{name}
func (c *Controller) APIV2ProxyDetail(ctx *httppkg.Context) (any, error) {
name, err := decodeV2PathParam(ctx, "name", "proxy name")
if err != nil {
return nil, err
}
ps := mem.StatsCollector.GetProxyByName(name)
if ps == nil {
return nil, httppkg.NewError(http.StatusNotFound, "no proxy info found")
}
return c.buildV2ProxyResp(ps), nil
}
// /api/v2/proxies/{name}/traffic
func (c *Controller) APIV2ProxyTraffic(ctx *httppkg.Context) (any, error) {
name, err := decodeV2PathParam(ctx, "name", "proxy name")
if err != nil {
return nil, err
}
proxyTrafficInfo := mem.StatsCollector.GetProxyTraffic(name)
if proxyTrafficInfo == nil {
return nil, httppkg.NewError(http.StatusNotFound, "no proxy info found")
}
return buildV2ProxyTrafficResp(name, proxyTrafficInfo, time.Now()), nil
}
func decodeV2PathParam(ctx *httppkg.Context, key string, label string) (string, error) {
raw := ctx.Param(key)
if raw == "" {
return "", fmt.Errorf("missing %s", label)
}
decoded, err := url.PathUnescape(raw)
if err != nil {
return "", httppkg.NewError(http.StatusBadRequest, fmt.Sprintf("invalid %s", label))
}
return decoded, nil
}
func getOrCreateV2User(items map[string]*model.V2UserResp, user string) *model.V2UserResp {
item, ok := items[user]
if !ok {
item = &model.V2UserResp{User: user}
items[user] = item
}
return item
}
func parseV2PageParams(ctx *httppkg.Context) (int, int, error) {
page, err := parseV2PositiveInt(ctx.Query("page"), defaultV2Page, "page")
if err != nil {
return 0, 0, err
}
pageSize, err := parseV2PositiveInt(ctx.Query("pageSize"), defaultV2PageSize, "pageSize")
if err != nil {
return 0, 0, err
}
if pageSize > maxV2PageSize {
return 0, 0, httppkg.NewError(http.StatusBadRequest, fmt.Sprintf("pageSize must be between 1 and %d", maxV2PageSize))
}
if page > math.MaxInt/pageSize {
return 0, 0, httppkg.NewError(http.StatusBadRequest, "page is too large")
}
return page, pageSize, nil
}
func parseV2PositiveInt(raw string, defaultValue int, name string) (int, error) {
if raw == "" {
return defaultValue, nil
}
value, err := strconv.Atoi(raw)
if err != nil || value < 1 {
return 0, httppkg.NewError(http.StatusBadRequest, fmt.Sprintf("%s must be a positive integer", name))
}
return value, nil
}
func parseV2StatusFilter(raw string) (string, error) {
status := strings.ToLower(raw)
switch status {
case "", "all", "online", "offline":
return status, nil
default:
return "", httppkg.NewError(http.StatusBadRequest, "status must be one of all, online, offline")
}
}
func parseV2ProxyTypeFilter(raw string) (string, error) {
proxyType := strings.ToLower(raw)
if proxyType == "" {
return "", nil
}
if slices.Contains(apiV2ProxyTypes, proxyType) {
return proxyType, nil
}
return "", httppkg.NewError(http.StatusBadRequest, "type must be one of tcp, udp, http, https, tcpmux, stcp, xtcp, sudp")
}
func parseV2SystemPruneType(raw string) (string, error) {
pruneType := strings.ToLower(raw)
switch pruneType {
case "":
return "", httppkg.NewError(http.StatusBadRequest, "type is required")
case v2SystemPruneTypeOfflineProxies:
return pruneType, nil
default:
return "", httppkg.NewError(http.StatusBadRequest, "type must be one of offline_proxies")
}
}
func matchV2StatusFilter(online bool, filter string) bool {
switch filter {
case "", "all":
return true
case "online":
return online
case "offline":
return !online
default:
return true
}
}
func buildV2PageResp[T any](items []T, page, pageSize int) model.V2PageResp[T] {
total := len(items)
return model.V2PageResp[T]{
Total: total,
Page: page,
PageSize: pageSize,
Items: paginateV2Items(items, page, pageSize),
}
}
func paginateV2Items[T any](items []T, page, pageSize int) []T {
start := (page - 1) * pageSize
if start >= len(items) {
return []T{}
}
end := min(start+pageSize, len(items))
return items[start:end]
}
func queryValue(ctx *httppkg.Context, key string) (string, bool) {
values, ok := ctx.Req.URL.Query()[key]
if !ok {
return "", false
}
if len(values) == 0 {
return "", true
}
return values[0], true
}
func matchV2ClientQuery(item model.ClientInfoResp, q string) bool {
return containsV2Query(q,
item.Key,
item.User,
item.ClientID,
item.RunID,
item.Version,
item.WireProtocol,
item.Hostname,
item.ClientIP,
)
}
func matchV2ProxyQuery(item model.V2ProxyResp, q string) bool {
values := []string{
item.Name,
item.Spec.Type,
item.User,
item.ClientID,
item.Status.State,
}
switch item.Spec.Type {
case string(v1.ProxyTypeTCP):
if item.Spec.TCP != nil && item.Spec.TCP.RemotePort != nil {
values = append(values, strconv.Itoa(*item.Spec.TCP.RemotePort))
}
case string(v1.ProxyTypeUDP):
if item.Spec.UDP != nil && item.Spec.UDP.RemotePort != nil {
values = append(values, strconv.Itoa(*item.Spec.UDP.RemotePort))
}
case string(v1.ProxyTypeHTTP):
if item.Spec.HTTP != nil {
values = append(values, item.Spec.HTTP.CustomDomains...)
values = append(values, item.Spec.HTTP.Subdomain)
}
case string(v1.ProxyTypeHTTPS):
if item.Spec.HTTPS != nil {
values = append(values, item.Spec.HTTPS.CustomDomains...)
values = append(values, item.Spec.HTTPS.Subdomain)
}
case string(v1.ProxyTypeTCPMUX):
if item.Spec.TCPMux != nil {
values = append(values, item.Spec.TCPMux.CustomDomains...)
values = append(values, item.Spec.TCPMux.Subdomain)
}
}
return containsV2Query(q, values...)
}
func containsV2Query(q string, values ...string) bool {
for _, value := range values {
if strings.Contains(strings.ToLower(value), q) {
return true
}
}
return false
}
func (c *Controller) listV2ProxyStats(proxyType string) []*mem.ProxyStats {
if proxyType != "" {
return mem.StatsCollector.GetProxiesByType(proxyType)
}
items := make([]*mem.ProxyStats, 0)
for _, t := range apiV2ProxyTypes {
items = append(items, mem.StatsCollector.GetProxiesByType(t)...)
}
return items
}
func buildV2ProxyTrafficResp(name string, traffic *mem.ProxyTrafficInfo, now time.Time) model.V2ProxyTrafficResp {
history := make([]model.V2ProxyTrafficPointResp, 0, v2ProxyTrafficDefaultDays)
for age := v2ProxyTrafficDefaultDays - 1; age >= 0; age-- {
history = append(history, model.V2ProxyTrafficPointResp{
Date: now.AddDate(0, 0, -age).Format(time.DateOnly),
TrafficIn: v2TrafficValueAt(traffic.TrafficIn, age),
TrafficOut: v2TrafficValueAt(traffic.TrafficOut, age),
})
}
return model.V2ProxyTrafficResp{
Name: name,
Unit: v2ProxyTrafficUnit,
Granularity: v2ProxyTrafficGranularity,
History: history,
}
}
func v2TrafficValueAt(values []int64, todayFirstIndex int) int64 {
if todayFirstIndex >= len(values) {
return 0
}
return values[todayFirstIndex]
}
func (c *Controller) buildV2ClientStatus(info registry.ClientInfo) model.V2ClientStatusResp {
status := model.V2ClientStatusResp{State: "offline"}
if info.Online {
status.State = "online"
}
user := info.User
clientID := info.ClientID()
for _, ps := range c.listV2ProxyStats("") {
if ps.User != user || ps.ClientID != clientID {
continue
}
status.CurConns += ps.CurConns
status.ProxyCount++
}
return status
}
func (c *Controller) buildV2ProxyResp(ps *mem.ProxyStats) model.V2ProxyResp {
state := "offline"
var cfg v1.ProxyConfigurer
if c.pxyManager != nil {
if pxy, ok := c.pxyManager.GetByName(ps.Name); ok {
state = "online"
cfg = pxy.GetConfigurer()
}
}
return model.V2ProxyResp{
Name: ps.Name,
User: ps.User,
ClientID: ps.ClientID,
Spec: buildV2ProxySpec(ps.Type, cfg),
Status: model.V2ProxyStatusResp{
State: state,
TodayTrafficIn: ps.TodayTrafficIn,
TodayTrafficOut: ps.TodayTrafficOut,
CurConns: ps.CurConns,
LastStartAt: ps.LastStartAt,
LastCloseAt: ps.LastCloseAt,
},
}
}
func buildV2ProxySpec(proxyType string, cfg v1.ProxyConfigurer) model.V2ProxySpec {
spec := model.V2ProxySpec{Type: proxyType}
switch proxyType {
case string(v1.ProxyTypeTCP):
block := &model.V2TCPProxySpec{}
if c, ok := cfg.(*v1.TCPProxyConfig); ok {
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
block.RemotePort = &c.RemotePort
}
spec.TCP = block
case string(v1.ProxyTypeUDP):
block := &model.V2UDPProxySpec{}
if c, ok := cfg.(*v1.UDPProxyConfig); ok {
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
block.RemotePort = &c.RemotePort
}
spec.UDP = block
case string(v1.ProxyTypeHTTP):
block := &model.V2HTTPProxySpec{}
if c, ok := cfg.(*v1.HTTPProxyConfig); ok {
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
block.CustomDomains = slices.Clone(c.CustomDomains)
block.Subdomain = c.SubDomain
block.Locations = slices.Clone(c.Locations)
block.HostHeaderRewrite = c.HostHeaderRewrite
}
spec.HTTP = block
case string(v1.ProxyTypeHTTPS):
block := &model.V2HTTPSProxySpec{}
if c, ok := cfg.(*v1.HTTPSProxyConfig); ok {
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
block.CustomDomains = slices.Clone(c.CustomDomains)
block.Subdomain = c.SubDomain
}
spec.HTTPS = block
case string(v1.ProxyTypeTCPMUX):
block := &model.V2TCPMuxProxySpec{}
if c, ok := cfg.(*v1.TCPMuxProxyConfig); ok {
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
block.CustomDomains = slices.Clone(c.CustomDomains)
block.Subdomain = c.SubDomain
block.Multiplexer = c.Multiplexer
block.RouteByHTTPUser = c.RouteByHTTPUser
}
spec.TCPMux = block
case string(v1.ProxyTypeSTCP):
block := &model.V2STCPProxySpec{}
if c, ok := cfg.(*v1.STCPProxyConfig); ok {
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
}
spec.STCP = block
case string(v1.ProxyTypeSUDP):
block := &model.V2SUDPProxySpec{}
if c, ok := cfg.(*v1.SUDPProxyConfig); ok {
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
}
spec.SUDP = block
case string(v1.ProxyTypeXTCP):
block := &model.V2XTCPProxySpec{}
if c, ok := cfg.(*v1.XTCPProxyConfig); ok {
block.V2ProxyBaseSpec = buildV2ProxyBaseSpec(c.GetBaseConfig())
}
spec.XTCP = block
}
return spec
}
func buildV2ProxyBaseSpec(base *v1.ProxyBaseConfig) model.V2ProxyBaseSpec {
return model.V2ProxyBaseSpec{
Annotations: maps.Clone(base.Annotations),
Metadatas: maps.Clone(base.Metadatas),
Transport: &model.V2ProxyTransportSpec{
UseEncryption: base.Transport.UseEncryption,
UseCompression: base.Transport.UseCompression,
BandwidthLimit: base.Transport.BandwidthLimit.String(),
BandwidthLimitMode: base.Transport.BandwidthLimitMode,
},
LoadBalancer: &model.V2ProxyLoadBalancerSpec{
Group: base.LoadBalancer.Group,
},
}
}
@@ -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)
}
+908
View File
@@ -0,0 +1,908 @@
// Copyright 2026 The frp Authors
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package http
import (
"encoding/json"
"fmt"
"math"
"net/http"
"net/http/httptest"
"net/url"
"testing"
"time"
"github.com/gorilla/mux"
"github.com/fatedier/frp/pkg/config/types"
v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/metrics/mem"
httppkg "github.com/fatedier/frp/pkg/util/http"
"github.com/fatedier/frp/server/http/model"
serverproxy "github.com/fatedier/frp/server/proxy"
"github.com/fatedier/frp/server/registry"
)
type v2EnvelopeForTest[T any] struct {
Code int `json:"code"`
Msg string `json:"msg"`
Data T `json:"data"`
}
type fakeStatsCollector struct {
server *mem.ServerStats
proxies map[string]*mem.ProxyStats
traffic map[string]*mem.ProxyTrafficInfo
pruneable map[string]bool
}
func (f *fakeStatsCollector) GetServer() *mem.ServerStats {
if f.server != nil {
return f.server
}
return &mem.ServerStats{ProxyTypeCounts: map[string]int64{}}
}
func (f *fakeStatsCollector) GetProxiesByType(proxyType string) []*mem.ProxyStats {
items := make([]*mem.ProxyStats, 0)
for _, ps := range f.proxies {
if ps.Type == proxyType {
items = append(items, ps)
}
}
return items
}
func (f *fakeStatsCollector) GetProxiesByTypeAndName(proxyType string, proxyName string) *mem.ProxyStats {
ps := f.proxies[proxyName]
if ps != nil && ps.Type == proxyType {
return ps
}
return nil
}
func (f *fakeStatsCollector) GetProxyByName(proxyName string) *mem.ProxyStats {
return f.proxies[proxyName]
}
func (f *fakeStatsCollector) GetProxyTraffic(name string) *mem.ProxyTrafficInfo {
return f.traffic[name]
}
func (f *fakeStatsCollector) ClearOfflineProxies() (int, int) {
return 0, len(f.proxies)
}
func (f *fakeStatsCollector) PruneOfflineProxies() (int, int) {
total := len(f.proxies)
cleared := 0
for name := range f.pruneable {
if _, ok := f.proxies[name]; ok {
delete(f.proxies, name)
cleared++
}
}
f.pruneable = map[string]bool{}
return cleared, total
}
func TestAPIV2SystemInfoEnvelope(t *testing.T) {
oldStatsCollector := mem.StatsCollector
mem.StatsCollector = &fakeStatsCollector{
server: &mem.ServerStats{
TotalTrafficIn: 1024,
TotalTrafficOut: 2048,
CurConns: 3,
ClientCounts: 4,
ProxyTypeCounts: map[string]int64{
"tcp": 2,
"http": 1,
},
},
proxies: map[string]*mem.ProxyStats{},
}
t.Cleanup(func() {
mem.StatsCollector = oldStatsCollector
})
controller := NewController(&v1.ServerConfig{
BindPort: 7000,
VhostHTTPPort: 8080,
VhostHTTPSPort: 8443,
TCPMuxHTTPConnectPort: 9000,
KCPBindPort: 7001,
QUICBindPort: 7002,
SubDomainHost: "example.com",
MaxPortsPerClient: 8,
AllowPorts: []types.PortsRange{
{Start: 1000, End: 1002},
{Single: 2000},
},
Transport: v1.ServerTransportConfig{
MaxPoolCount: 5,
HeartbeatTimeout: 90,
TLS: v1.TLSServerConfig{
Force: true,
},
},
}, registry.NewClientRegistry(), serverproxy.NewManager())
router := newV2TestRouter(controller)
resp := performRequest(router, "/api/v2/system/info")
if resp.Code != http.StatusOK {
t.Fatalf("status mismatch, want %d got %d", http.StatusOK, resp.Code)
}
rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp)
if rawResp.Code != http.StatusOK || rawResp.Msg != "success" {
t.Fatalf("envelope mismatch: %#v", rawResp)
}
assertRawJSONKeys(t, rawResp.Data, "config", "status", "version")
assertRawJSONKeysFromMessage(t, rawResp.Data["config"],
"allowPortsStr",
"bindPort",
"heartbeatTimeout",
"kcpBindPort",
"maxPoolCount",
"maxPortsPerClient",
"quicBindPort",
"subdomainHost",
"tcpmuxHTTPConnectPort",
"tlsForce",
"vhostHTTPPort",
"vhostHTTPSPort",
)
assertRawJSONKeysFromMessage(t, rawResp.Data["status"],
"clientCounts",
"curConns",
"proxyTypeCount",
"totalTrafficIn",
"totalTrafficOut",
)
systemResp := decodeResponse[v2EnvelopeForTest[model.V2SystemInfoResp]](t, resp)
if systemResp.Data.Version == "" {
t.Fatal("version should be set at top level")
}
if systemResp.Data.Config.BindPort != 7000 ||
systemResp.Data.Config.VhostHTTPPort != 8080 ||
systemResp.Data.Config.VhostHTTPSPort != 8443 ||
systemResp.Data.Config.TCPMuxHTTPConnectPort != 9000 ||
systemResp.Data.Config.KCPBindPort != 7001 ||
systemResp.Data.Config.QUICBindPort != 7002 ||
systemResp.Data.Config.SubdomainHost != "example.com" ||
systemResp.Data.Config.MaxPoolCount != 5 ||
systemResp.Data.Config.MaxPortsPerClient != 8 ||
systemResp.Data.Config.HeartbeatTimeout != 90 ||
systemResp.Data.Config.AllowPortsStr != "1000-1002,2000" ||
!systemResp.Data.Config.TLSForce {
t.Fatalf("config mismatch: %#v", systemResp.Data.Config)
}
if systemResp.Data.Status.TotalTrafficIn != 1024 ||
systemResp.Data.Status.TotalTrafficOut != 2048 ||
systemResp.Data.Status.CurConns != 3 ||
systemResp.Data.Status.ClientCounts != 4 ||
systemResp.Data.Status.ProxyTypeCounts["tcp"] != 2 ||
systemResp.Data.Status.ProxyTypeCounts["http"] != 1 {
t.Fatalf("status mismatch: %#v", systemResp.Data.Status)
}
}
func TestAPIV2SystemPruneOfflineProxies(t *testing.T) {
oldStatsCollector := mem.StatsCollector
collector := &fakeStatsCollector{
proxies: map[string]*mem.ProxyStats{
"tcp-offline": {Name: "tcp-offline", Type: "tcp"},
"http-offline": {Name: "http-offline", Type: "http"},
"udp-offline": {Name: "udp-offline", Type: "udp"},
"tcp-online": {Name: "tcp-online", Type: "tcp"},
"http-online": {Name: "http-online", Type: "http"},
"udp-online": {Name: "udp-online", Type: "udp"},
"stcp-restarted": {Name: "stcp-restarted", Type: "stcp"},
"xtcp-restarted": {Name: "xtcp-restarted", Type: "xtcp"},
"sudp-same-time": {Name: "sudp-same-time", Type: "sudp"},
"tcpmux-running": {Name: "tcpmux-running", Type: "tcpmux"},
},
pruneable: map[string]bool{
"tcp-offline": true,
"http-offline": true,
"udp-offline": true,
},
}
mem.StatsCollector = collector
t.Cleanup(func() {
mem.StatsCollector = oldStatsCollector
})
controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager())
router := newV2TestRouter(controller)
resp := performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=offline_proxies")
if resp.Code != http.StatusOK {
t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
}
rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp)
if rawResp.Code != http.StatusOK || rawResp.Msg != "success" {
t.Fatalf("envelope mismatch: %#v", rawResp)
}
assertRawJSONKeys(t, rawResp.Data, "cleared", "total", "type")
pruneResp := decodeResponse[v2EnvelopeForTest[model.V2SystemPruneResp]](t, resp)
if pruneResp.Data.Type != "offline_proxies" || pruneResp.Data.Cleared != 3 || pruneResp.Data.Total != 10 {
t.Fatalf("prune response mismatch: %#v", pruneResp.Data)
}
if _, ok := collector.proxies["tcp-offline"]; ok {
t.Fatal("pruned proxy statistics should be removed")
}
if _, ok := collector.proxies["tcp-online"]; !ok {
t.Fatal("online proxy statistics should remain")
}
resp = performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=offline_proxies")
if resp.Code != http.StatusOK {
t.Fatalf("second prune status mismatch, want %d got %d", http.StatusOK, resp.Code)
}
pruneResp = decodeResponse[v2EnvelopeForTest[model.V2SystemPruneResp]](t, resp)
if pruneResp.Data.Cleared != 0 || pruneResp.Data.Total != 7 {
t.Fatalf("second prune response mismatch: %#v", pruneResp.Data)
}
}
func TestAPIV2SystemPruneTypeErrorsUseEnvelope(t *testing.T) {
controller := newV2TestController(t)
router := newV2TestRouter(controller)
resp := performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune")
if resp.Code != http.StatusBadRequest {
t.Fatalf("missing type status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
}
errResp := decodeResponse[httppkg.V2Response](t, resp)
if errResp.Code != http.StatusBadRequest || errResp.Msg != "type is required" || errResp.Data != nil {
t.Fatalf("missing type error envelope mismatch: %#v", errResp)
}
resp = performRequestWithMethod(router, http.MethodPost, "/api/v2/system/prune?type=clients")
if resp.Code != http.StatusBadRequest {
t.Fatalf("invalid type status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
}
errResp = decodeResponse[httppkg.V2Response](t, resp)
if errResp.Code != http.StatusBadRequest || errResp.Msg != "type must be one of offline_proxies" || errResp.Data != nil {
t.Fatalf("invalid type error envelope mismatch: %#v", errResp)
}
}
func TestAPIV2ClientListEnvelopePaginationAndFilters(t *testing.T) {
controller := newV2TestController(t)
router := newV2TestRouter(controller)
resp := performRequest(router, "/api/v2/clients?page=1&pageSize=1")
if resp.Code != http.StatusOK {
t.Fatalf("status mismatch, want %d got %d", http.StatusOK, resp.Code)
}
pageResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.ClientInfoResp]]](t, resp)
if pageResp.Code != http.StatusOK || pageResp.Msg != "success" {
t.Fatalf("envelope mismatch: %#v", pageResp)
}
if pageResp.Data.Total != 3 || pageResp.Data.Page != 1 || pageResp.Data.PageSize != 1 || len(pageResp.Data.Items) != 1 {
t.Fatalf("page data mismatch: %#v", pageResp.Data)
}
if got := pageResp.Data.Items[0].User; got != "" {
t.Fatalf("first sorted user mismatch, want empty got %q", got)
}
resp = performRequest(router, "/api/v2/clients?user=&page=1&pageSize=50")
emptyUserResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.ClientInfoResp]]](t, resp)
if emptyUserResp.Data.Total != 1 || emptyUserResp.Data.Items[0].User != "" {
t.Fatalf("empty user filter mismatch: %#v", emptyUserResp.Data)
}
resp = performRequest(router, "/api/v2/clients?user=alice&status=online&q=alice-host")
aliceResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.ClientInfoResp]]](t, resp)
if aliceResp.Data.Total != 1 || aliceResp.Data.Items[0].User != "alice" {
t.Fatalf("alice filter mismatch: %#v", aliceResp.Data)
}
resp = performRequest(router, "/api/v2/clients?status=offline")
offlineResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.ClientInfoResp]]](t, resp)
if offlineResp.Data.Total != 1 || offlineResp.Data.Items[0].User != "bob" {
t.Fatalf("offline filter mismatch: %#v", offlineResp.Data)
}
}
func TestAPIV2PageParamErrorsUseEnvelope(t *testing.T) {
controller := newV2TestController(t)
router := newV2TestRouter(controller)
resp := performRequest(router, "/api/v2/clients?page=0")
if resp.Code != http.StatusBadRequest {
t.Fatalf("status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
}
errResp := decodeResponse[httppkg.V2Response](t, resp)
if errResp.Code != http.StatusBadRequest || errResp.Data != nil {
t.Fatalf("error envelope mismatch: %#v", errResp)
}
resp = performRequest(router, "/api/v2/clients?pageSize=201")
if resp.Code != http.StatusBadRequest {
t.Fatalf("status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
}
resp = performRequest(router, fmt.Sprintf("/api/v2/clients?page=%d&pageSize=2", math.MaxInt))
if resp.Code != http.StatusBadRequest {
t.Fatalf("status mismatch for overflowing page offset, want %d got %d", http.StatusBadRequest, resp.Code)
}
}
func TestAPIV2ClientDetailEnvelope(t *testing.T) {
controller := newV2TestController(t)
router := newV2TestRouter(controller)
resp := performRequest(router, "/api/v2/clients/alice.client-a")
if resp.Code != http.StatusOK {
t.Fatalf("status mismatch, want %d got %d", http.StatusOK, resp.Code)
}
detailResp := decodeResponse[v2EnvelopeForTest[model.V2ClientDetailResp]](t, resp)
if detailResp.Data.User != "alice" || detailResp.Data.ClientID != "client-a" {
t.Fatalf("client detail mismatch: %#v", detailResp.Data)
}
if detailResp.Data.Status.State != "online" || detailResp.Data.Status.CurConns != 5 || detailResp.Data.Status.ProxyCount != 2 {
t.Fatalf("client detail status mismatch: %#v", detailResp.Data.Status)
}
}
func TestAPIV2ClientDetailEncodedKey(t *testing.T) {
oldStatsCollector := mem.StatsCollector
mem.StatsCollector = &fakeStatsCollector{
proxies: map[string]*mem.ProxyStats{
"tcp-url": {
Name: "tcp-url",
Type: "tcp",
User: "url",
ClientID: "client/a?b#c",
CurConns: 7,
},
},
}
t.Cleanup(func() {
mem.StatsCollector = oldStatsCollector
})
clientRegistry := registry.NewClientRegistry()
clientRegistry.Register("url", "client/a?b#c", "run-url", "url-host", "1.0.0", "127.0.0.4", "v2")
controller := NewController(&v1.ServerConfig{}, clientRegistry, serverproxy.NewManager())
router := newV2TestRouter(controller)
encodedKey := url.PathEscape("url.client/a?b#c")
resp := performRequest(router, "/api/v2/clients/"+encodedKey)
if resp.Code != http.StatusOK {
t.Fatalf("encoded client key status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
}
encodedResp := decodeResponse[v2EnvelopeForTest[model.V2ClientDetailResp]](t, resp)
if encodedResp.Data.User != "url" || encodedResp.Data.ClientID != "client/a?b#c" {
t.Fatalf("encoded client detail mismatch: %#v", encodedResp.Data)
}
if encodedResp.Data.Status.CurConns != 7 || encodedResp.Data.Status.ProxyCount != 1 {
t.Fatalf("encoded client detail status mismatch: %#v", encodedResp.Data.Status)
}
}
func TestAPIV2ProxyListDetailAndUsers(t *testing.T) {
controller := newV2TestController(t)
router := newV2TestRouter(controller)
resp := performRequest(router, "/api/v2/proxies?type=invalid")
if resp.Code != http.StatusBadRequest {
t.Fatalf("invalid proxy type status mismatch, want %d got %d", http.StatusBadRequest, resp.Code)
}
errResp := decodeResponse[httppkg.V2Response](t, resp)
if errResp.Code != http.StatusBadRequest || errResp.Data != nil {
t.Fatalf("invalid proxy type error envelope mismatch: %#v", errResp)
}
resp = performRequest(router, "/api/v2/proxies?type=tcp&user=&page=1&pageSize=50")
proxyResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.V2ProxyResp]]](t, resp)
if proxyResp.Data.Total != 1 {
t.Fatalf("proxy filter total mismatch: %#v", proxyResp.Data)
}
proxyItem := proxyResp.Data.Items[0]
if proxyItem.Name != "tcp-empty" || proxyItem.Spec.Type != "tcp" || proxyItem.User != "" || proxyItem.Status.State != "offline" {
t.Fatalf("proxy item mismatch: %#v", proxyItem)
}
rawProxyResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[map[string]json.RawMessage]]](t, resp)
assertRawJSONKeys(t, rawProxyResp.Data.Items[0], "clientID", "name", "spec", "status", "user")
var rawListSpec map[string]json.RawMessage
if err := json.Unmarshal(rawProxyResp.Data.Items[0]["spec"], &rawListSpec); err != nil {
t.Fatalf("unmarshal list proxy spec failed: %v", err)
}
assertRawJSONKeys(t, rawListSpec, "tcp", "type")
assertRawJSONKeysFromMessage(t, rawListSpec["tcp"])
resp = performRequest(router, "/api/v2/proxies/tcp-alice")
rawProxyDetailResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp)
assertRawJSONKeysFromMessage(t, rawProxyDetailResp.Data["status"],
"curConns",
"lastCloseAt",
"lastStartAt",
"phase",
"todayTrafficIn",
"todayTrafficOut",
)
proxyDetailResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyResp]](t, resp)
if proxyDetailResp.Data.Name != "tcp-alice" || proxyDetailResp.Data.User != "alice" {
t.Fatalf("proxy detail mismatch: %#v", proxyDetailResp.Data)
}
assertRawJSONKeys(t, rawProxyDetailResp.Data, "clientID", "name", "spec", "status", "user")
var rawDetailSpec map[string]json.RawMessage
if err := json.Unmarshal(rawProxyDetailResp.Data["spec"], &rawDetailSpec); err != nil {
t.Fatalf("unmarshal detail proxy spec failed: %v", err)
}
assertRawJSONKeys(t, rawDetailSpec, "tcp", "type")
assertRawJSONKeysFromMessage(t, rawDetailSpec["tcp"])
if proxyDetailResp.Data.Status.LastStartAt != 1783504200 || proxyDetailResp.Data.Status.LastCloseAt != 1783504300 {
t.Fatalf("proxy detail timestamp mismatch: %#v", proxyDetailResp.Data.Status)
}
resp = performRequest(router, "/api/v2/users?page=1&pageSize=50")
userResp := decodeResponse[v2EnvelopeForTest[model.V2PageResp[model.V2UserResp]]](t, resp)
if userResp.Data.Total != 3 {
t.Fatalf("user total mismatch: %#v", userResp.Data)
}
expectedProxyCounts := map[string]int{
"": 1,
"alice": 2,
"bob": 1,
}
for _, item := range userResp.Data.Items {
if item.ClientCount != 1 || item.ProxyCount != expectedProxyCounts[item.User] {
t.Fatalf("user counts mismatch: %#v", item)
}
}
}
func TestAPIV2ProxyTrafficEnvelopeSchemaAndHistory(t *testing.T) {
oldStatsCollector := mem.StatsCollector
mem.StatsCollector = &fakeStatsCollector{
proxies: map[string]*mem.ProxyStats{
"ssh": {Name: "ssh", Type: "tcp"},
},
traffic: map[string]*mem.ProxyTrafficInfo{
"ssh": {
Name: "ssh",
TrafficIn: []int64{70, 60, 50, 40, 30, 20, 10},
TrafficOut: []int64{700, 600, 500, 400, 300, 200, 100},
},
},
}
t.Cleanup(func() {
mem.StatsCollector = oldStatsCollector
})
controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager())
router := newV2TestRouter(controller)
resp := performRequest(router, "/api/v2/proxies/ssh/traffic")
if resp.Code != http.StatusOK {
t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
}
rawResp := decodeResponse[v2EnvelopeForTest[map[string]json.RawMessage]](t, resp)
if rawResp.Code != http.StatusOK || rawResp.Msg != "success" {
t.Fatalf("envelope mismatch: %#v", rawResp)
}
assertRawJSONKeys(t, rawResp.Data, "granularity", "history", "name", "unit")
trafficResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyTrafficResp]](t, resp)
if trafficResp.Data.Name != "ssh" || trafficResp.Data.Unit != "bytes" || trafficResp.Data.Granularity != "day" {
t.Fatalf("traffic metadata mismatch: %#v", trafficResp.Data)
}
if len(trafficResp.Data.History) != 7 {
t.Fatalf("history length mismatch, want 7 got %d: %#v", len(trafficResp.Data.History), trafficResp.Data.History)
}
wantIn := []int64{10, 20, 30, 40, 50, 60, 70}
wantOut := []int64{100, 200, 300, 400, 500, 600, 700}
var prevDate time.Time
for i, point := range trafficResp.Data.History {
assertRawJSONKeysFromMessage(t, mustMarshalJSON(t, point), "date", "trafficIn", "trafficOut")
if point.TrafficIn != wantIn[i] || point.TrafficOut != wantOut[i] {
t.Fatalf("history[%d] traffic mismatch: %#v", i, point)
}
parsedDate, err := time.Parse(time.DateOnly, point.Date)
if err != nil {
t.Fatalf("history[%d] date should be yyyy-mm-dd, got %q: %v", i, point.Date, err)
}
if i > 0 && !parsedDate.Equal(prevDate.AddDate(0, 0, 1)) {
t.Fatalf("history dates should be oldest to newest, got %s after %s", point.Date, prevDate.Format(time.DateOnly))
}
prevDate = parsedDate
}
}
func TestAPIV2ProxyTrafficNotFoundEnvelope(t *testing.T) {
controller := newV2TestController(t)
router := newV2TestRouter(controller)
resp := performRequest(router, "/api/v2/proxies/missing/traffic")
if resp.Code != http.StatusNotFound {
t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusNotFound, resp.Code, resp.Body.String())
}
errResp := decodeResponse[httppkg.V2Response](t, resp)
if errResp.Code != http.StatusNotFound || errResp.Msg != "no proxy info found" || errResp.Data != nil {
t.Fatalf("not found envelope mismatch: %#v", errResp)
}
}
func TestAPIV2ProxyDetailAndTrafficEncodedName(t *testing.T) {
name := "folder/ssh?x#y"
oldStatsCollector := mem.StatsCollector
mem.StatsCollector = &fakeStatsCollector{
proxies: map[string]*mem.ProxyStats{
name: {Name: name, Type: "tcp", User: "encoded"},
},
traffic: map[string]*mem.ProxyTrafficInfo{
name: {
Name: name,
TrafficIn: []int64{1},
TrafficOut: []int64{2},
},
},
}
t.Cleanup(func() {
mem.StatsCollector = oldStatsCollector
})
controller := NewController(&v1.ServerConfig{}, registry.NewClientRegistry(), serverproxy.NewManager())
router := newV2TestRouter(controller)
encodedName := url.PathEscape(name)
resp := performRequest(router, "/api/v2/proxies/"+encodedName)
if resp.Code != http.StatusOK {
t.Fatalf("encoded proxy detail status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
}
detailResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyResp]](t, resp)
if detailResp.Data.Name != name || detailResp.Data.User != "encoded" {
t.Fatalf("encoded proxy detail mismatch: %#v", detailResp.Data)
}
resp = performRequest(router, "/api/v2/proxies/"+encodedName+"/traffic")
if resp.Code != http.StatusOK {
t.Fatalf("encoded traffic status mismatch, want %d got %d, body: %s", http.StatusOK, resp.Code, resp.Body.String())
}
trafficResp := decodeResponse[v2EnvelopeForTest[model.V2ProxyTrafficResp]](t, resp)
if trafficResp.Data.Name != name {
t.Fatalf("encoded traffic name mismatch: %#v", trafficResp.Data)
}
if got := trafficResp.Data.History[len(trafficResp.Data.History)-1]; got.TrafficIn != 1 || got.TrafficOut != 2 {
t.Fatalf("encoded traffic latest point mismatch: %#v", got)
}
}
func TestAPIV2ProxyTrafficInvalidEncodedNameUses400Envelope(t *testing.T) {
controller := newV2TestController(t)
handler := httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyTraffic)
req := httptest.NewRequest(http.MethodGet, "/api/v2/proxies/%25ZZ/traffic", nil)
req = mux.SetURLVars(req, map[string]string{"name": "%ZZ"})
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, req)
if resp.Code != http.StatusBadRequest {
t.Fatalf("status mismatch, want %d got %d, body: %s", http.StatusBadRequest, resp.Code, resp.Body.String())
}
errResp := decodeResponse[httppkg.V2Response](t, resp)
if errResp.Code != http.StatusBadRequest || errResp.Msg != "invalid proxy name" || errResp.Data != nil {
t.Fatalf("invalid encoded name envelope mismatch: %#v", errResp)
}
}
func TestMatchV2ProxyQueryMatchesSpecFields(t *testing.T) {
tests := []struct {
name string
item model.V2ProxyResp
q string
want bool
}{
{
name: "tcp remote port",
item: model.V2ProxyResp{Name: "tcp-proxy", Spec: model.V2ProxySpec{
Type: "tcp",
TCP: &model.V2TCPProxySpec{RemotePort: v2TestIntPtr(6000)},
}},
q: "6000",
want: true,
},
{
name: "udp remote port",
item: model.V2ProxyResp{Name: "udp-proxy", Spec: model.V2ProxySpec{
Type: "udp",
UDP: &model.V2UDPProxySpec{RemotePort: v2TestIntPtr(7000)},
}},
q: "7000",
want: true,
},
{
name: "remote port does not match colon form",
item: model.V2ProxyResp{Name: "tcp-proxy", Spec: model.V2ProxySpec{
Type: "tcp",
TCP: &model.V2TCPProxySpec{RemotePort: v2TestIntPtr(6000)},
}},
q: ":6000",
want: false,
},
{
name: "http custom domain",
item: model.V2ProxyResp{Name: "http-proxy", Spec: model.V2ProxySpec{
Type: "http",
HTTP: &model.V2HTTPProxySpec{CustomDomains: []string{"app.example.com"}},
}},
q: "app.example.com",
want: true,
},
{
name: "https subdomain",
item: model.V2ProxyResp{Name: "https-proxy", Spec: model.V2ProxySpec{
Type: "https",
HTTPS: &model.V2HTTPSProxySpec{Subdomain: "portal"},
}},
q: "portal",
want: true,
},
{
name: "subdomain does not match expanded host",
item: model.V2ProxyResp{Name: "https-proxy", Spec: model.V2ProxySpec{
Type: "https",
HTTPS: &model.V2HTTPSProxySpec{Subdomain: "portal"},
}},
q: "portal.example.com",
want: false,
},
{
name: "tcpmux custom domain",
item: model.V2ProxyResp{Name: "tcpmux-proxy", Spec: model.V2ProxySpec{
Type: "tcpmux",
TCPMux: &model.V2TCPMuxProxySpec{CustomDomains: []string{"mux.example.com"}},
}},
q: "mux.example.com",
want: true,
},
{
name: "offline shell does not match online spec fields",
item: model.V2ProxyResp{Name: "offline-proxy", Spec: model.V2ProxySpec{
Type: "tcp",
TCP: &model.V2TCPProxySpec{},
}},
q: "6000",
want: false,
},
{
name: "offline shell does not contribute zero remote port",
item: model.V2ProxyResp{Name: "offline-proxy", Spec: model.V2ProxySpec{
Type: "tcp",
TCP: &model.V2TCPProxySpec{},
}},
q: "0",
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := matchV2ProxyQuery(tt.item, tt.q); got != tt.want {
t.Fatalf("matchV2ProxyQuery() = %v, want %v", got, tt.want)
}
})
}
}
func TestLegacyAPIResponsesRemainBare(t *testing.T) {
controller := newV2TestController(t)
router := newV2TestRouter(controller)
resp := performRequest(router, "/api/serverinfo")
var serverInfo model.ServerInfoResp
if err := json.Unmarshal(resp.Body.Bytes(), &serverInfo); err != nil {
t.Fatalf("legacy serverinfo should be a bare object: %v, body: %s", err, resp.Body.String())
}
if serverInfo.Version == "" {
t.Fatal("legacy serverinfo version should be set")
}
var serverInfoRaw map[string]json.RawMessage
if err := json.Unmarshal(resp.Body.Bytes(), &serverInfoRaw); err != nil {
t.Fatalf("unmarshal legacy serverinfo object failed: %v", err)
}
if _, ok := serverInfoRaw["data"]; ok {
t.Fatalf("legacy serverinfo should not use v2 envelope: %s", resp.Body.String())
}
if _, ok := serverInfoRaw["config"]; ok {
t.Fatalf("legacy serverinfo should stay flat, got config in: %s", resp.Body.String())
}
resp = performRequest(router, "/api/clients")
var clients []model.ClientInfoResp
if err := json.Unmarshal(resp.Body.Bytes(), &clients); err != nil {
t.Fatalf("legacy clients should be a bare array: %v, body: %s", err, resp.Body.String())
}
if len(clients) != 3 {
t.Fatalf("legacy clients total mismatch, want 3 got %d", len(clients))
}
resp = performRequest(router, "/api/proxy/tcp")
var proxies model.GetProxyInfoResp
if err := json.Unmarshal(resp.Body.Bytes(), &proxies); err != nil {
t.Fatalf("legacy proxy response should be {proxies}: %v, body: %s", err, resp.Body.String())
}
if len(proxies.Proxies) != 2 {
t.Fatalf("legacy tcp proxy total mismatch, want 2 got %d", len(proxies.Proxies))
}
var envelope httppkg.V2Response
if err := json.Unmarshal(resp.Body.Bytes(), &envelope); err == nil && envelope.Code != 0 {
t.Fatalf("legacy proxy response should not use v2 envelope: %#v", envelope)
}
resp = performRequest(router, "/api/traffic/tcp-alice")
var traffic model.GetProxyTrafficResp
if err := json.Unmarshal(resp.Body.Bytes(), &traffic); err != nil {
t.Fatalf("legacy traffic should be a bare object: %v, body: %s", err, resp.Body.String())
}
if traffic.Name != "tcp-alice" ||
len(traffic.TrafficIn) != 2 || traffic.TrafficIn[0] != 7 || traffic.TrafficIn[1] != 6 ||
len(traffic.TrafficOut) != 2 || traffic.TrafficOut[0] != 70 || traffic.TrafficOut[1] != 60 {
t.Fatalf("legacy traffic should preserve today-first arrays, got: %#v", traffic)
}
var trafficRaw map[string]json.RawMessage
if err := json.Unmarshal(resp.Body.Bytes(), &trafficRaw); err != nil {
t.Fatalf("unmarshal legacy traffic object failed: %v", err)
}
if _, ok := trafficRaw["data"]; ok {
t.Fatalf("legacy traffic should not use v2 envelope: %s", resp.Body.String())
}
}
func v2TestIntPtr(value int) *int {
return &value
}
func newV2TestController(t *testing.T) *Controller {
t.Helper()
oldStatsCollector := mem.StatsCollector
mem.StatsCollector = &fakeStatsCollector{
proxies: map[string]*mem.ProxyStats{
"tcp-empty": {
Name: "tcp-empty",
Type: "tcp",
User: "",
ClientID: "legacy-client",
TodayTrafficIn: 10,
TodayTrafficOut: 20,
CurConns: 1,
},
"tcp-alice": {
Name: "tcp-alice",
Type: "tcp",
User: "alice",
ClientID: "client-a",
TodayTrafficIn: 30,
TodayTrafficOut: 40,
CurConns: 2,
LastStartTime: "07-08 12:30:00",
LastCloseTime: "07-08 12:31:40",
LastStartAt: 1783504200,
LastCloseAt: 1783504300,
},
"http-alice": {
Name: "http-alice",
Type: "http",
User: "alice",
ClientID: "client-a",
CurConns: 3,
},
"udp-bob": {
Name: "udp-bob",
Type: "udp",
User: "bob",
ClientID: "client-b",
},
},
traffic: map[string]*mem.ProxyTrafficInfo{
"tcp-alice": {
Name: "tcp-alice",
TrafficIn: []int64{7, 6},
TrafficOut: []int64{70, 60},
},
},
}
t.Cleanup(func() {
mem.StatsCollector = oldStatsCollector
})
clientRegistry := registry.NewClientRegistry()
clientRegistry.Register("", "legacy-client", "run-empty", "empty-host", "1.0.0", "127.0.0.1", "v1")
clientRegistry.Register("alice", "client-a", "run-a", "alice-host", "1.0.0", "127.0.0.2", "v2")
clientRegistry.Register("bob", "client-b", "run-b", "bob-host", "1.0.0", "127.0.0.3", "v1")
clientRegistry.MarkOfflineByRunID("run-b")
return NewController(&v1.ServerConfig{}, clientRegistry, serverproxy.NewManager())
}
func newV2TestRouter(controller *Controller) *mux.Router {
router := mux.NewRouter()
router.HandleFunc("/api/v2/users", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2UserList)).Methods(http.MethodGet)
router.HandleFunc("/api/v2/system/info", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2SystemInfo)).Methods(http.MethodGet)
router.HandleFunc("/api/v2/system/prune", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2SystemPrune)).Methods(http.MethodPost)
router.HandleFunc("/api/v2/clients", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ClientList)).Methods(http.MethodGet)
encodedPathRouter := router.NewRoute().Subrouter()
encodedPathRouter.UseEncodedPath()
encodedPathRouter.HandleFunc("/api/v2/clients/{key}", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ClientDetail)).Methods(http.MethodGet)
router.HandleFunc("/api/v2/proxies", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyList)).Methods(http.MethodGet)
encodedPathRouter.HandleFunc("/api/v2/proxies/{name}/traffic", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyTraffic)).Methods(http.MethodGet)
encodedPathRouter.HandleFunc("/api/v2/proxies/{name}", httppkg.MakeHTTPHandlerFuncV2(controller.APIV2ProxyDetail)).Methods(http.MethodGet)
router.HandleFunc("/api/serverinfo", httppkg.MakeHTTPHandlerFunc(controller.APIServerInfo)).Methods(http.MethodGet)
router.HandleFunc("/api/clients", httppkg.MakeHTTPHandlerFunc(controller.APIClientList)).Methods(http.MethodGet)
router.HandleFunc("/api/proxy/{type}", httppkg.MakeHTTPHandlerFunc(controller.APIProxyByType)).Methods(http.MethodGet)
router.HandleFunc("/api/traffic/{name}", httppkg.MakeHTTPHandlerFunc(controller.APIProxyTraffic)).Methods(http.MethodGet)
return router
}
func performRequest(handler http.Handler, target string) *httptest.ResponseRecorder {
return performRequestWithMethod(handler, http.MethodGet, target)
}
func performRequestWithMethod(handler http.Handler, method, target string) *httptest.ResponseRecorder {
req := httptest.NewRequest(method, target, nil)
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, req)
return resp
}
func decodeResponse[T any](t *testing.T, resp *httptest.ResponseRecorder) T {
t.Helper()
var out T
if err := json.Unmarshal(resp.Body.Bytes(), &out); err != nil {
t.Fatalf("unmarshal response failed: %v, body: %s", err, resp.Body.String())
}
return out
}
func assertRawJSONKeys(t *testing.T, raw map[string]json.RawMessage, want ...string) {
t.Helper()
if len(raw) != len(want) {
t.Fatalf("json keys mismatch, want %v got %v", want, raw)
}
for _, key := range want {
if _, ok := raw[key]; !ok {
t.Fatalf("json key %q missing from %v", key, raw)
}
}
}
func assertRawJSONKeysFromMessage(t *testing.T, raw json.RawMessage, want ...string) {
t.Helper()
var out map[string]json.RawMessage
if err := json.Unmarshal(raw, &out); err != nil {
t.Fatalf("unmarshal raw json object failed: %v, body: %s", err, string(raw))
}
assertRawJSONKeys(t, out, want...)
}
func mustMarshalJSON(t *testing.T, value any) json.RawMessage {
t.Helper()
out, err := json.Marshal(value)
if err != nil {
t.Fatalf("marshal json failed: %v", err)
}
return out
}
+179
View File
@@ -0,0 +1,179 @@
// Copyright 2026 The frp Authors
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package model
type V2PageResp[T any] struct {
Total int `json:"total"`
Page int `json:"page"`
PageSize int `json:"pageSize"`
Items []T `json:"items"`
}
type V2SystemInfoResp struct {
Version string `json:"version"`
Config V2SystemInfoConfigResp `json:"config"`
Status V2SystemInfoStatusResp `json:"status"`
}
type V2SystemInfoConfigResp struct {
BindPort int `json:"bindPort"`
VhostHTTPPort int `json:"vhostHTTPPort"`
VhostHTTPSPort int `json:"vhostHTTPSPort"`
TCPMuxHTTPConnectPort int `json:"tcpmuxHTTPConnectPort"`
KCPBindPort int `json:"kcpBindPort"`
QUICBindPort int `json:"quicBindPort"`
SubdomainHost string `json:"subdomainHost"`
MaxPoolCount int64 `json:"maxPoolCount"`
MaxPortsPerClient int64 `json:"maxPortsPerClient"`
HeartbeatTimeout int64 `json:"heartbeatTimeout"`
AllowPortsStr string `json:"allowPortsStr"`
TLSForce bool `json:"tlsForce"`
}
type V2SystemInfoStatusResp struct {
TotalTrafficIn int64 `json:"totalTrafficIn"`
TotalTrafficOut int64 `json:"totalTrafficOut"`
CurConns int64 `json:"curConns"`
ClientCounts int64 `json:"clientCounts"`
ProxyTypeCounts map[string]int64 `json:"proxyTypeCount"`
}
type V2SystemPruneResp struct {
Type string `json:"type"`
Cleared int `json:"cleared"`
Total int `json:"total"`
}
type V2UserResp struct {
User string `json:"user"`
ClientCount int `json:"clientCount"`
ProxyCount int `json:"proxyCount"`
}
type V2ClientDetailResp struct {
ClientInfoResp
Status V2ClientStatusResp `json:"status"`
}
type V2ClientStatusResp struct {
State string `json:"phase"`
CurConns int64 `json:"curConns"`
ProxyCount int64 `json:"proxyCount"`
}
type V2ProxyResp struct {
Name string `json:"name"`
User string `json:"user"`
ClientID string `json:"clientID"`
Spec V2ProxySpec `json:"spec"`
Status V2ProxyStatusResp `json:"status"`
}
type V2ProxySpec struct {
Type string `json:"type"`
TCP *V2TCPProxySpec `json:"tcp,omitempty"`
UDP *V2UDPProxySpec `json:"udp,omitempty"`
HTTP *V2HTTPProxySpec `json:"http,omitempty"`
HTTPS *V2HTTPSProxySpec `json:"https,omitempty"`
TCPMux *V2TCPMuxProxySpec `json:"tcpmux,omitempty"`
STCP *V2STCPProxySpec `json:"stcp,omitempty"`
SUDP *V2SUDPProxySpec `json:"sudp,omitempty"`
XTCP *V2XTCPProxySpec `json:"xtcp,omitempty"`
}
type V2ProxyBaseSpec struct {
Annotations map[string]string `json:"annotations,omitempty"`
Metadatas map[string]string `json:"metadatas,omitempty"`
Transport *V2ProxyTransportSpec `json:"transport,omitempty"`
LoadBalancer *V2ProxyLoadBalancerSpec `json:"loadBalancer,omitempty"`
}
type V2ProxyTransportSpec struct {
UseEncryption bool `json:"useEncryption"`
UseCompression bool `json:"useCompression"`
BandwidthLimit string `json:"bandwidthLimit"`
BandwidthLimitMode string `json:"bandwidthLimitMode"`
}
type V2ProxyLoadBalancerSpec struct {
Group string `json:"group"`
}
type V2TCPProxySpec struct {
V2ProxyBaseSpec
RemotePort *int `json:"remotePort,omitempty"`
}
type V2UDPProxySpec struct {
V2ProxyBaseSpec
RemotePort *int `json:"remotePort,omitempty"`
}
type V2HTTPProxySpec struct {
V2ProxyBaseSpec
CustomDomains []string `json:"customDomains,omitempty"`
Subdomain string `json:"subdomain,omitempty"`
Locations []string `json:"locations,omitempty"`
HostHeaderRewrite string `json:"hostHeaderRewrite,omitempty"`
}
type V2HTTPSProxySpec struct {
V2ProxyBaseSpec
CustomDomains []string `json:"customDomains,omitempty"`
Subdomain string `json:"subdomain,omitempty"`
}
type V2TCPMuxProxySpec struct {
V2ProxyBaseSpec
CustomDomains []string `json:"customDomains,omitempty"`
Subdomain string `json:"subdomain,omitempty"`
Multiplexer string `json:"multiplexer,omitempty"`
RouteByHTTPUser string `json:"routeByHTTPUser,omitempty"`
}
type V2STCPProxySpec struct {
V2ProxyBaseSpec
}
type V2SUDPProxySpec struct {
V2ProxyBaseSpec
}
type V2XTCPProxySpec struct {
V2ProxyBaseSpec
}
type V2ProxyStatusResp struct {
State string `json:"phase"`
TodayTrafficIn int64 `json:"todayTrafficIn"`
TodayTrafficOut int64 `json:"todayTrafficOut"`
CurConns int64 `json:"curConns"`
LastStartAt int64 `json:"lastStartAt,omitempty"`
LastCloseAt int64 `json:"lastCloseAt,omitempty"`
}
type V2ProxyTrafficResp struct {
Name string `json:"name"`
Unit string `json:"unit"`
Granularity string `json:"granularity"`
History []V2ProxyTrafficPointResp `json:"history"`
}
type V2ProxyTrafficPointResp struct {
Date string `json:"date"`
TrafficIn int64 `json:"trafficIn"`
TrafficOut int64 `json:"trafficOut"`
}
+93 -34
View File
@@ -82,19 +82,20 @@ type Proxy interface {
} }
type BaseProxy struct { type BaseProxy struct {
name string name string
rc *controller.ResourceController rc *controller.ResourceController
listeners []net.Listener listeners []net.Listener
usedPortsNum int usedPortsNum int
poolCount int poolCount int
getWorkConnFn GetWorkConnFn getWorkConnFn GetWorkConnFn
serverCfg *v1.ServerConfig serverCfg *v1.ServerConfig
encryptionKey []byte encryptionKey []byte
limiter *rate.Limiter limiter *rate.Limiter
userInfo plugin.UserInfo userInfo plugin.UserInfo
loginMsg *msg.Login loginMsg *msg.Login
configurer v1.ProxyConfigurer configurer v1.ProxyConfigurer
wireProtocol string wireProtocol string
udpPacketCodec string
mu sync.RWMutex mu sync.RWMutex
xl *xlog.Logger xl *xlog.Logger
@@ -327,10 +328,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) { func (pxy *BaseProxy) joinUserConnection(local io.ReadWriteCloser, userConn net.Conn, proxyType string, xl *xlog.Logger) (int64, int64, []error) {
visitorWireProtocol := wireProtocolFromConn(userConn) visitorWireProtocol := wireProtocolFromConn(userConn)
if proxyType == string(v1.ProxyTypeSUDP) && isMixedWireProtocol(pxy.wireProtocol, visitorWireProtocol) { visitorUDPPacketCodec := udpPacketCodecFromConn(userConn)
xl.Infof("bridge mixed SUDP payload codecs, proxy wireProtocol [%s], visitor wireProtocol [%s]", if proxyType == string(v1.ProxyTypeSUDP) {
normalizeWireProtocol(pxy.wireProtocol), normalizeWireProtocol(visitorWireProtocol)) mixed, err := isMixedSUDPPacketEncoding(pxy.wireProtocol, pxy.udpPacketCodec, visitorWireProtocol, visitorUDPPacketCodec)
return joinSUDPMessageBridge(local, userConn, pxy.wireProtocol, visitorWireProtocol, xl) 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) return libio.Join(local, userConn)
} }
@@ -339,6 +348,10 @@ type wireProtocolGetter interface {
WireProtocol() string WireProtocol() string
} }
type udpPacketCodecGetter interface {
UDPPacketCodec() string
}
func wireProtocolFromConn(conn net.Conn) string { func wireProtocolFromConn(conn net.Conn) string {
if getter, ok := conn.(wireProtocolGetter); ok { if getter, ok := conn.(wireProtocolGetter); ok {
return getter.WireProtocol() return getter.WireProtocol()
@@ -346,10 +359,46 @@ func wireProtocolFromConn(conn net.Conn) string {
return "" return ""
} }
func udpPacketCodecFromConn(conn net.Conn) string {
if getter, ok := conn.(udpPacketCodecGetter); ok {
return getter.UDPPacketCodec()
}
return ""
}
func isMixedWireProtocol(left, right string) bool { func isMixedWireProtocol(left, right string) bool {
return normalizeWireProtocol(left) != normalizeWireProtocol(right) return normalizeWireProtocol(left) != normalizeWireProtocol(right)
} }
func isMixedSUDPPacketEncoding(leftWire, leftCodec, rightWire, rightCodec string) (bool, error) {
leftCodec, err := normalizeUDPPacketCodec(leftWire, leftCodec)
if err != nil {
return false, fmt.Errorf("invalid left SUDP packet encoding: %w", err)
}
rightCodec, err = normalizeUDPPacketCodec(rightWire, rightCodec)
if err != nil {
return false, fmt.Errorf("invalid right SUDP packet encoding: %w", err)
}
return normalizeWireProtocol(leftWire) != normalizeWireProtocol(rightWire) || leftCodec != rightCodec, nil
}
func normalizeUDPPacketCodec(wireProtocol, codec string) (string, error) {
switch wireProtocol {
case "", wire.ProtocolV1:
if codec != "" {
return "", fmt.Errorf("UDP packet codec %q requires wire protocol v2", codec)
}
return "", nil
case wire.ProtocolV2:
if codec == "" || codec == wire.UDPPacketCodecBinary {
return codec, nil
}
return "", fmt.Errorf("unsupported UDP packet codec %q", codec)
default:
return "", fmt.Errorf("unsupported wire protocol %q", wireProtocol)
}
}
func normalizeWireProtocol(wireProtocol string) string { func normalizeWireProtocol(wireProtocol string) string {
if wireProtocol == wire.ProtocolV2 { if wireProtocol == wire.ProtocolV2 {
return wire.ProtocolV2 return wire.ProtocolV2
@@ -361,13 +410,21 @@ func joinSUDPMessageBridge(
proxyConn io.ReadWriteCloser, proxyConn io.ReadWriteCloser,
visitorConn io.ReadWriteCloser, visitorConn io.ReadWriteCloser,
proxyWireProtocol string, proxyWireProtocol string,
proxyUDPPacketCodec string,
visitorWireProtocol string, visitorWireProtocol string,
visitorUDPPacketCodec string,
xl *xlog.Logger, xl *xlog.Logger,
) (inCount int64, outCount int64, errs []error) { ) (inCount int64, outCount int64, errs []error) {
// The mixed bridge decodes and re-encodes messages, so raw framed byte counts // The mixed bridge decodes and re-encodes messages, so raw framed byte counts
// are not available. Count UDP payload bytes and ignore heartbeat traffic. // are not available. Count UDP payload bytes and ignore heartbeat traffic.
proxyRW := msg.NewReadWriter(proxyConn, proxyWireProtocol) proxyRW, err := msg.NewUDPPacketReadWriter(proxyConn, proxyWireProtocol, proxyUDPPacketCodec)
visitorRW := msg.NewReadWriter(visitorConn, visitorWireProtocol) if err != nil {
return 0, 0, []error{err}
}
visitorRW, err := msg.NewUDPPacketReadWriter(visitorConn, visitorWireProtocol, visitorUDPPacketCodec)
if err != nil {
return 0, 0, []error{err}
}
var ( var (
once sync.Once once sync.Once
@@ -469,6 +526,7 @@ type Options struct {
ServerCfg *v1.ServerConfig ServerCfg *v1.ServerConfig
EncryptionKey []byte EncryptionKey []byte
WireProtocol string WireProtocol string
UDPPacketCodec string
} }
func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) { func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) {
@@ -478,24 +536,25 @@ func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) {
var limiter *rate.Limiter var limiter *rate.Limiter
limitBytes := configurer.GetBaseConfig().Transport.BandwidthLimit.Bytes() limitBytes := configurer.GetBaseConfig().Transport.BandwidthLimit.Bytes()
if limitBytes > 0 && configurer.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeServer { if limitBytes > 0 && configurer.GetBaseConfig().Transport.BandwidthLimitMode == types.BandwidthLimitModeServer {
limiter = rate.NewLimiter(rate.Limit(float64(limitBytes)), int(limitBytes)) limiter = limit.NewBandwidthLimiter(limitBytes)
} }
basePxy := BaseProxy{ basePxy := BaseProxy{
name: configurer.GetBaseConfig().Name, name: configurer.GetBaseConfig().Name,
rc: options.ResourceController, rc: options.ResourceController,
listeners: make([]net.Listener, 0), listeners: make([]net.Listener, 0),
poolCount: options.PoolCount, poolCount: options.PoolCount,
getWorkConnFn: options.GetWorkConnFn, getWorkConnFn: options.GetWorkConnFn,
serverCfg: options.ServerCfg, serverCfg: options.ServerCfg,
encryptionKey: options.EncryptionKey, encryptionKey: options.EncryptionKey,
limiter: limiter, limiter: limiter,
xl: xl, xl: xl,
ctx: xlog.NewContext(ctx, xl), ctx: xlog.NewContext(ctx, xl),
userInfo: options.UserInfo, userInfo: options.UserInfo,
loginMsg: options.LoginMsg, loginMsg: options.LoginMsg,
configurer: configurer, configurer: configurer,
wireProtocol: options.WireProtocol, wireProtocol: options.WireProtocol,
udpPacketCodec: options.UDPPacketCodec,
} }
factory := proxyFactoryRegistry[reflect.TypeOf(configurer)] factory := proxyFactoryRegistry[reflect.TypeOf(configurer)]
+232
View File
@@ -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)
+238 -21
View File
@@ -18,22 +18,27 @@ import (
"bufio" "bufio"
"bytes" "bytes"
"encoding/binary" "encoding/binary"
"io"
"net"
"testing" "testing"
"time"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/msg" "github.com/fatedier/frp/pkg/msg"
"github.com/fatedier/frp/pkg/proto/wire" "github.com/fatedier/frp/pkg/proto/wire"
"github.com/fatedier/frp/pkg/util/xlog"
) )
func TestSUDPBridgeTranscodesProxyV1ToVisitorV2(t *testing.T) { func TestSUDPBridgeTranscodesProxyV1ToVisitorV2(t *testing.T) {
var in, out bytes.Buffer var in, out bytes.Buffer
writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, &msg.UDPPacket{Content: []byte("proxy-to-visitor")}) writeSUDPBridgeMsg(t, &in, wire.ProtocolV1, "", &msg.UDPPacket{Content: []byte("proxy-to-visitor")})
var count int64 var count int64
err := bridgeSUDPProxyToVisitor( err := bridgeSUDPProxyToVisitor(
msg.NewReadWriter(&in, wire.ProtocolV1), newSUDPBridgeRW(t, &in, wire.ProtocolV1, ""),
msg.NewReadWriter(&out, wire.ProtocolV2), newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""),
&count, &count,
nil, nil,
) )
@@ -53,12 +58,12 @@ func TestSUDPBridgeTranscodesProxyV1ToVisitorV2(t *testing.T) {
func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) { func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) {
var in, out bytes.Buffer var in, out bytes.Buffer
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, &msg.UDPPacket{Content: []byte("visitor-to-proxy")}) writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", &msg.UDPPacket{Content: []byte("visitor-to-proxy")})
var count int64 var count int64
err := bridgeSUDPVisitorToProxy( err := bridgeSUDPVisitorToProxy(
msg.NewReadWriter(&in, wire.ProtocolV2), newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
msg.NewReadWriter(&out, wire.ProtocolV1), newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
&count, &count,
nil, nil,
) )
@@ -76,33 +81,67 @@ func TestSUDPBridgeTranscodesVisitorV2ToProxyV1(t *testing.T) {
require.Equal(t, []byte("visitor-to-proxy"), got.Content) require.Equal(t, []byte("visitor-to-proxy"), got.Content)
} }
func TestSUDPBridgeForwardsProxyPing(t *testing.T) { func TestSUDPBridgeTranscodesProxyV2BinaryToVisitorV2JSON(t *testing.T) {
var in, out bytes.Buffer 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 var count int64
err := bridgeSUDPProxyToVisitor( err := bridgeSUDPProxyToVisitor(
msg.NewReadWriter(&in, wire.ProtocolV1), newSUDPBridgeRW(t, &in, wire.ProtocolV2, wire.UDPPacketCodecBinary),
msg.NewReadWriter(&out, wire.ProtocolV2), 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, &count,
nil, nil,
) )
require.NoError(t, err) require.NoError(t, err)
require.Zero(t, count) 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.NoError(t, err)
require.IsType(t, &msg.Ping{}, rawMsg) require.IsType(t, &msg.Ping{}, rawMsg)
} }
func TestSUDPBridgeDropsVisitorPing(t *testing.T) { func TestSUDPBridgeDropsVisitorPing(t *testing.T) {
var in, out bytes.Buffer var in, out bytes.Buffer
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, &msg.Ping{}) writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", &msg.Ping{})
var count int64 var count int64
err := bridgeSUDPVisitorToProxy( err := bridgeSUDPVisitorToProxy(
msg.NewReadWriter(&in, wire.ProtocolV2), newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
msg.NewReadWriter(&out, wire.ProtocolV1), newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
&count, &count,
nil, nil,
) )
@@ -113,12 +152,12 @@ func TestSUDPBridgeDropsVisitorPing(t *testing.T) {
func TestSUDPBridgeRejectsUnknownVisitorMessage(t *testing.T) { func TestSUDPBridgeRejectsUnknownVisitorMessage(t *testing.T) {
var in, out bytes.Buffer var in, out bytes.Buffer
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, &msg.Pong{}) writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", &msg.Pong{})
var count int64 var count int64
err := bridgeSUDPVisitorToProxy( err := bridgeSUDPVisitorToProxy(
msg.NewReadWriter(&in, wire.ProtocolV2), newSUDPBridgeRW(t, &in, wire.ProtocolV2, ""),
msg.NewReadWriter(&out, wire.ProtocolV1), newSUDPBridgeRW(t, &out, wire.ProtocolV1, ""),
&count, &count,
nil, nil,
) )
@@ -127,6 +166,22 @@ func TestSUDPBridgeRejectsUnknownVisitorMessage(t *testing.T) {
require.Empty(t, out.Bytes()) require.Empty(t, out.Bytes())
} }
func TestSUDPBridgeRejectsMismatchedPacketCodecOnStream(t *testing.T) {
var in, out bytes.Buffer
writeSUDPBridgeMsg(t, &in, wire.ProtocolV2, "", newSUDPBridgeUDPPacket("json-on-binary-stream"))
var count int64
err := bridgeSUDPProxyToVisitor(
newSUDPBridgeRW(t, &in, wire.ProtocolV2, wire.UDPPacketCodecBinary),
newSUDPBridgeRW(t, &out, wire.ProtocolV2, ""),
&count,
nil,
)
require.ErrorContains(t, err, "received JSON UDP packet after binary codec negotiation")
require.Zero(t, count)
require.Empty(t, out.Bytes())
}
func TestSUDPBridgeDetectsMixedWireProtocol(t *testing.T) { func TestSUDPBridgeDetectsMixedWireProtocol(t *testing.T) {
require.False(t, isMixedWireProtocol("", wire.ProtocolV1)) require.False(t, isMixedWireProtocol("", wire.ProtocolV1))
require.False(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV2)) require.False(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV2))
@@ -134,8 +189,170 @@ func TestSUDPBridgeDetectsMixedWireProtocol(t *testing.T) {
require.True(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV1)) require.True(t, isMixedWireProtocol(wire.ProtocolV2, wire.ProtocolV1))
} }
func writeSUDPBridgeMsg(t *testing.T, buf *bytes.Buffer, wireProtocol string, m msg.Message) { func TestSUDPBridgeDetectsMixedPacketEncoding(t *testing.T) {
t.Helper() for _, tc := range []struct {
name string
require.NoError(t, msg.NewReadWriter(buf, wireProtocol).WriteMsg(m)) 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
} }
+7 -1
View File
@@ -224,7 +224,13 @@ func (pxy *UDPProxy) Run() (remoteAddr string, err error) {
pxy.workConn = netpkg.WrapReadWriteCloserToConn(rwc, workConn) pxy.workConn = netpkg.WrapReadWriteCloserToConn(rwc, workConn)
// Plain UDP payload follows the negotiated wire protocol for message framing. // Plain UDP payload follows the negotiated wire protocol for message framing.
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()) ctx, cancel := context.WithCancel(context.Background())
go workConnReaderFn(payloadConn) go workConnReaderFn(payloadConn)
go workConnSenderFn(payloadConn, ctx) go workConnSenderFn(payloadConn, ctx)
+44 -6
View File
@@ -28,6 +28,7 @@ type ClientInfo struct {
User string User string
RawClientID string RawClientID string
RunID string RunID string
ControlID uint64
Hostname string Hostname string
IP string IP string
Version string Version string
@@ -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. // Register stores/updates metadata for a client and returns the registry key plus whether it conflicts with an online client.
func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol string) (key string, conflict bool) { func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol string) (key string, conflict bool) {
return cr.RegisterWithControlID(user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol, 0)
}
// RegisterWithControlID is the generation-aware form used by ControlManager.
// A control ID is process-local and prevents an older control generation from
// changing the registry entry now owned by a newer generation with the same run ID.
func (cr *ClientRegistry) RegisterWithControlID(
user, rawClientID, runID, hostname, version, remoteAddr, wireProtocol string,
controlID uint64,
) (key string, conflict bool) {
if runID == "" { if runID == "" {
return "", false return "", false
} }
@@ -83,6 +94,16 @@ func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version,
if enforceUnique && exists && info.Online && info.RunID != "" && info.RunID != runID { if enforceUnique && exists && info.Online && info.RunID != "" && info.RunID != runID {
return key, true return key, true
} }
if previousKey, ok := cr.runIndex[runID]; ok && previousKey != key {
if previous, ok := cr.clients[previousKey]; ok && previous.RunID == runID {
if previous.RawClientID == "" {
delete(cr.clients, previousKey)
} else {
setClientOffline(previous, now)
}
}
delete(cr.runIndex, runID)
}
if !exists { if !exists {
info = &ClientInfo{ info = &ClientInfo{
@@ -97,6 +118,7 @@ func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version,
info.RawClientID = rawClientID info.RawClientID = rawClientID
info.RunID = runID info.RunID = runID
info.ControlID = controlID
info.Hostname = hostname info.Hostname = hostname
info.IP = remoteAddr info.IP = remoteAddr
info.Version = version info.Version = version
@@ -114,6 +136,16 @@ func (cr *ClientRegistry) Register(user, rawClientID, runID, hostname, version,
// MarkOfflineByRunID marks the client as offline when the corresponding control disconnects. // MarkOfflineByRunID marks the client as offline when the corresponding control disconnects.
func (cr *ClientRegistry) MarkOfflineByRunID(runID string) { func (cr *ClientRegistry) MarkOfflineByRunID(runID string) {
cr.markOfflineByRunID(runID, 0, false)
}
// MarkOfflineByRunIDAndControlID marks a client offline only when the registry
// entry still belongs to the supplied control generation.
func (cr *ClientRegistry) MarkOfflineByRunIDAndControlID(runID string, controlID uint64) {
cr.markOfflineByRunID(runID, controlID, true)
}
func (cr *ClientRegistry) markOfflineByRunID(runID string, controlID uint64, matchControlID bool) {
cr.mu.Lock() cr.mu.Lock()
defer cr.mu.Unlock() defer cr.mu.Unlock()
@@ -121,17 +153,23 @@ func (cr *ClientRegistry) MarkOfflineByRunID(runID string) {
if !ok { if !ok {
return 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 == "" { if info.RawClientID == "" {
delete(cr.clients, key) delete(cr.clients, key)
} else { } else {
info.RunID = "" setClientOffline(info, cr.clock.Now())
info.Online = false
now := cr.clock.Now()
info.DisconnectedAt = 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. // List returns a snapshot of all known clients.
+86
View File
@@ -72,3 +72,89 @@ func TestClientRegistryUsesClockForTimestamps(t *testing.T) {
t.Fatalf("disconnected time mismatch, want %s got %s", disconnectedAt, info.DisconnectedAt) t.Fatalf("disconnected time mismatch, want %s got %s", disconnectedAt, info.DisconnectedAt)
} }
} }
func TestClientRegistryControlIDPreventsStaleOffline(t *testing.T) {
registry := NewClientRegistry()
key, conflict := registry.RegisterWithControlID(
"user", "client-id", "run-id", "old-host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1,
)
if conflict {
t.Fatal("unexpected client conflict")
}
_, conflict = registry.RegisterWithControlID(
"user", "client-id", "run-id", "new-host", "1.0.1", "127.0.0.2", wire.ProtocolV2, 2,
)
if conflict {
t.Fatal("same run ID replacement should not conflict")
}
registry.MarkOfflineByRunIDAndControlID("run-id", 1)
info, ok := registry.GetByKey(key)
if !ok {
t.Fatalf("client %q not found", key)
}
if !info.Online || info.ControlID != 2 || info.Hostname != "new-host" {
t.Fatalf("stale offline changed current generation: %+v", info)
}
registry.MarkOfflineByRunIDAndControlID("run-id", 2)
info, ok = registry.GetByKey(key)
if !ok {
t.Fatalf("client %q not found after disconnect", key)
}
if info.Online || info.ControlID != 0 || info.RunID != "" {
t.Fatalf("current generation was not marked offline: %+v", info)
}
}
func TestClientRegistryClientIDConflictSemantics(t *testing.T) {
registry := NewClientRegistry()
_, conflict := registry.RegisterWithControlID(
"user", "client-id", "run-one", "host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1,
)
if conflict {
t.Fatal("unexpected initial client conflict")
}
_, conflict = registry.RegisterWithControlID(
"user", "client-id", "run-two", "host", "1.0.0", "127.0.0.2", wire.ProtocolV1, 2,
)
if !conflict {
t.Fatal("different online run IDs with the same explicit client ID must conflict")
}
registry.MarkOfflineByRunIDAndControlID("run-one", 1)
_, conflict = registry.RegisterWithControlID(
"user", "client-id", "run-two", "host", "1.0.0", "127.0.0.2", wire.ProtocolV1, 2,
)
if conflict {
t.Fatal("offline explicit client ID should be reusable")
}
}
func TestClientRegistrySameRunIDMovesBetweenClientKeys(t *testing.T) {
registry := NewClientRegistry()
oldKey, conflict := registry.RegisterWithControlID(
"user", "old-client", "run-id", "old-host", "1.0.0", "127.0.0.1", wire.ProtocolV1, 1,
)
if conflict {
t.Fatal("unexpected initial client conflict")
}
newKey, conflict := registry.RegisterWithControlID(
"user", "new-client", "run-id", "new-host", "1.0.1", "127.0.0.2", wire.ProtocolV2, 2,
)
if conflict {
t.Fatal("same run ID moving to a new client key should not conflict")
}
oldInfo, ok := registry.GetByKey(oldKey)
if !ok {
t.Fatalf("old explicit client %q should remain as offline history", oldKey)
}
if oldInfo.Online || oldInfo.RunID != "" || oldInfo.ControlID != 0 {
t.Fatalf("old client key remained online: %+v", oldInfo)
}
newInfo, ok := registry.GetByKey(newKey)
if !ok || !newInfo.Online || newInfo.RunID != "run-id" || newInfo.ControlID != 2 {
t.Fatalf("new client key was not registered: %+v", newInfo)
}
}
+112 -44
View File
@@ -18,6 +18,7 @@ import (
"bytes" "bytes"
"context" "context"
"crypto/tls" "crypto/tls"
"errors"
"fmt" "fmt"
"io" "io"
"net" "net"
@@ -34,6 +35,7 @@ import (
"github.com/fatedier/frp/pkg/auth" "github.com/fatedier/frp/pkg/auth"
v1 "github.com/fatedier/frp/pkg/config/v1" v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/config/v1/validation"
modelmetrics "github.com/fatedier/frp/pkg/metrics" modelmetrics "github.com/fatedier/frp/pkg/metrics"
"github.com/fatedier/frp/pkg/msg" "github.com/fatedier/frp/pkg/msg"
"github.com/fatedier/frp/pkg/nathole" "github.com/fatedier/frp/pkg/nathole"
@@ -51,7 +53,6 @@ import (
"github.com/fatedier/frp/pkg/util/xlog" "github.com/fatedier/frp/pkg/util/xlog"
"github.com/fatedier/frp/server/controller" "github.com/fatedier/frp/server/controller"
"github.com/fatedier/frp/server/group" "github.com/fatedier/frp/server/group"
"github.com/fatedier/frp/server/metrics"
"github.com/fatedier/frp/server/ports" "github.com/fatedier/frp/server/ports"
"github.com/fatedier/frp/server/proxy" "github.com/fatedier/frp/server/proxy"
"github.com/fatedier/frp/server/registry" "github.com/fatedier/frp/server/registry"
@@ -64,6 +65,8 @@ const (
vhostReadWriteTimeout time.Duration = 30 * time.Second vhostReadWriteTimeout time.Duration = 30 * time.Second
) )
var errControlReplaced = errors.New("control was replaced during login")
func init() { func init() {
crypto.DefaultSalt = "frp" crypto.DefaultSalt = "frp"
// Disable quic-go's receive buffer warning. // Disable quic-go's receive buffer warning.
@@ -161,9 +164,10 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) {
return nil, err return nil, err
} }
clientRegistry := registry.NewClientRegistry()
svr := &Service{ svr := &Service{
ctlManager: NewControlManager(), ctlManager: NewControlManager(clientRegistry),
clientRegistry: registry.NewClientRegistry(), clientRegistry: clientRegistry,
pxyManager: proxy.NewManager(), pxyManager: proxy.NewManager(),
pluginManager: plugin.NewManager(), pluginManager: plugin.NewManager(),
rc: &controller.ResourceController{ rc: &controller.ResourceController{
@@ -297,10 +301,14 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) {
svr.rc.HTTPReverseProxy = rp svr.rc.HTTPReverseProxy = rp
address := net.JoinHostPort(cfg.ProxyBindAddr, strconv.Itoa(cfg.VhostHTTPPort)) address := net.JoinHostPort(cfg.ProxyBindAddr, strconv.Itoa(cfg.VhostHTTPPort))
protocols := new(http.Protocols)
protocols.SetHTTP1(true)
protocols.SetUnencryptedHTTP2(true)
server := &http.Server{ server := &http.Server{
Addr: address, Addr: address,
Handler: rp, Handler: rp,
ReadHeaderTimeout: 60 * time.Second, ReadHeaderTimeout: 60 * time.Second,
Protocols: protocols,
} }
var l net.Listener var l net.Listener
if httpMuxOn { if httpMuxOn {
@@ -463,12 +471,15 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
} }
} }
if err == nil { if err == nil {
ctl, err = svr.RegisterControl(controlConn, m, internal, acceptedConn.wireProtocol) ctl, err = svr.RegisterControl(controlConn, m, internal, acceptedConn.wireProtocol, acceptedConn.udpPacketCodec)
} }
} }
if err != nil { if err != nil {
xl.Warnf("register control error: %v", err) xl.Warnf("register control error: %v", err)
if ctl != nil {
svr.ctlManager.Remove(ctl)
}
if writeErr := writeWithDeadline(conn, connWriteTimeout, func() error { if writeErr := writeWithDeadline(conn, connWriteTimeout, func() error {
return acceptedConn.conn.WriteMsg(&msg.LoginResp{ return acceptedConn.conn.WriteMsg(&msg.LoginResp{
Version: version.Full(), Version: version.Full(),
@@ -477,31 +488,34 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
}); writeErr != nil { }); writeErr != nil {
xl.Warnf("write login error response error: %v", writeErr) xl.Warnf("write login error response error: %v", writeErr)
} }
conn.Close() if ctl != nil {
_ = ctl.Close()
} else {
conn.Close()
}
return return
} }
if err = writeWithDeadline(conn, connWriteTimeout, func() error { if err = svr.completeControlLogin(ctl, func() error {
return acceptedConn.conn.WriteMsg(&msg.LoginResp{ return writeWithDeadline(conn, connWriteTimeout, func() error {
Version: version.Full(), return acceptedConn.conn.WriteMsg(&msg.LoginResp{
RunID: ctl.runID, Version: version.Full(),
Error: "", RunID: ctl.runID,
Error: "",
})
}) })
}); err != nil { }); err != nil {
xl.Warnf("write login response error: %v", err) xl.Warnf("complete control login error: %v", err)
svr.ctlManager.Del(m.RunID, ctl) svr.ctlManager.Remove(ctl)
svr.clientRegistry.MarkOfflineByRunID(m.RunID) _ = ctl.Close()
conn.Close()
return return
} }
ctl.Start()
metrics.Server.NewClient()
go func() {
// block until control closed
ctl.WaitClosed()
svr.ctlManager.Del(m.RunID, ctl)
}()
case *msg.NewWorkConn: case *msg.NewWorkConn:
if err := svr.RegisterWorkConn(acceptedConn.conn, m); err != nil { if err := svr.RegisterWorkConn(
acceptedConn.conn,
m,
acceptedConn.wireProtocol,
acceptedConn.clientHelloPresent,
); err != nil {
_ = acceptedConn.conn.WriteMsg(&msg.StartWorkConn{ _ = acceptedConn.conn.WriteMsg(&msg.StartWorkConn{
Error: util.GenerateResponseErrorString("invalid NewWorkConn", err, lo.FromPtr(svr.cfg.DetailedErrorsToClient)), Error: util.GenerateResponseErrorString("invalid NewWorkConn", err, lo.FromPtr(svr.cfg.DetailedErrorsToClient)),
}) })
@@ -527,11 +541,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 { type acceptedConnection struct {
conn *msg.Conn conn *msg.Conn
wireProtocol string wireProtocol string
cryptoContext *wire.CryptoContext clientHelloPresent bool
firstMsg msg.Message udpPacketCodec string
cryptoContext *wire.CryptoContext
firstMsg msg.Message
} }
func (svr *Service) acceptConnection(ctx context.Context, conn net.Conn) (*acceptedConnection, error) { func (svr *Service) acceptConnection(ctx context.Context, conn net.Conn) (*acceptedConnection, error) {
@@ -599,6 +626,7 @@ func (ac *acceptedConnection) readFirstV2Msg(conn net.Conn, wireConn *wire.Conn)
return nil, fmt.Errorf("read v2 frame: %w", err) return nil, fmt.Errorf("read v2 frame: %w", err)
} }
if frame.Type == wire.FrameTypeClientHello { if frame.Type == wire.FrameTypeClientHello {
ac.clientHelloPresent = true
if err := ac.handleClientHello(conn, wireConn, frame); err != nil { if err := ac.handleClientHello(conn, wireConn, frame); err != nil {
return nil, err return nil, err
} }
@@ -647,6 +675,7 @@ func (ac *acceptedConnection) handleClientHello(conn net.Conn, wireConn *wire.Co
return fmt.Errorf("write ServerHello: %w", err) return fmt.Errorf("write ServerHello: %w", err)
} }
ac.cryptoContext = cryptoContext ac.cryptoContext = cryptoContext
ac.udpPacketCodec = serverHello.Selected.Message.UDPPacketCodec
return nil return nil
} }
@@ -740,7 +769,20 @@ func (svr *Service) RegisterControl(
loginMsg *msg.Login, loginMsg *msg.Login,
internal bool, internal bool,
wireProtocol string, wireProtocol string,
udpPacketCodec string,
) (*Control, error) { ) (*Control, error) {
switch wireProtocol {
case wire.ProtocolV1:
if udpPacketCodec != "" {
return nil, fmt.Errorf("UDP packet codec %q requires wire protocol v2", udpPacketCodec)
}
case wire.ProtocolV2:
if udpPacketCodec != "" && udpPacketCodec != wire.UDPPacketCodecBinary {
return nil, fmt.Errorf("unsupported UDP packet codec selection: %s", udpPacketCodec)
}
default:
return nil, fmt.Errorf("unsupported wire protocol: %s", wireProtocol)
}
// If client's RunID is empty, it's a new client, we just create a new controller. // If client's RunID is empty, it's a new client, we just create a new controller.
// Otherwise, we check if there is one controller has the same run id. If so, we release previous controller and start new one. // Otherwise, we check if there is one controller has the same run id. If so, we release previous controller and start new one.
var err error var err error
@@ -750,6 +792,9 @@ func (svr *Service) RegisterControl(
return nil, err return nil, err
} }
} }
if err := validation.ValidateRunID(loginMsg.RunID); err != nil {
return nil, fmt.Errorf("invalid run id: %w", err)
}
ctx := netpkg.NewContextFromConn(ctlConn) ctx := netpkg.NewContextFromConn(ctlConn)
xl := xlog.FromContextSafe(ctx) xl := xlog.FromContextSafe(ctx)
@@ -776,8 +821,8 @@ func (svr *Service) RegisterControl(
Conn: ctlConn, Conn: ctlConn,
LoginMsg: loginMsg, LoginMsg: loginMsg,
ServerCfg: svr.cfg, ServerCfg: svr.cfg,
ClientRegistry: svr.clientRegistry,
WireProtocol: wireProtocol, WireProtocol: wireProtocol,
UDPPacketCodec: udpPacketCodec,
}) })
if err != nil { if err != nil {
xl.Warnf("create new controller error: %v", err) xl.Warnf("create new controller error: %v", err)
@@ -785,31 +830,41 @@ func (svr *Service) RegisterControl(
return nil, fmt.Errorf("unexpected error when creating new controller") return nil, fmt.Errorf("unexpected error when creating new controller")
} }
if oldCtl := svr.ctlManager.Add(loginMsg.RunID, ctl); oldCtl != nil { if err := svr.ctlManager.Add(ctl); err != nil {
oldCtl.WaitClosed() return ctl, err
} }
ctl.WaitForHandoff()
remoteAddr := ctlConn.RemoteAddr().String() active, err := svr.ctlManager.Activate(ctl)
if host, _, err := net.SplitHostPort(remoteAddr); err == nil { if err != nil {
remoteAddr = host return ctl, err
} }
_, conflict := svr.clientRegistry.Register(loginMsg.User, loginMsg.ClientID, loginMsg.RunID, loginMsg.Hostname, loginMsg.Version, remoteAddr, wireProtocol) if !active {
if conflict { return ctl, errControlReplaced
svr.ctlManager.Del(loginMsg.RunID, ctl)
return nil, fmt.Errorf("client_id [%s] for user [%s] is already online", loginMsg.ClientID, loginMsg.User)
} }
return ctl, nil return ctl, nil
} }
// RegisterWorkConn register a new work connection to control and proxies need it. // RegisterWorkConn register a new work connection to control and proxies need it.
func (svr *Service) RegisterWorkConn(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) xl := netpkg.NewLogFromConn(workConn)
ctl, exist := svr.ctlManager.GetByID(newMsg.RunID) ctl, exist := svr.ctlManager.GetByID(newMsg.RunID)
if !exist { if !exist {
xl.Warnf("no client control found for run id [%s]", newMsg.RunID) xl.Warnf("no client control found for run id [%s]", newMsg.RunID)
return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID) return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID)
} }
if workWireProtocol != ctl.sessionCtx.WireProtocol {
return fmt.Errorf("work connection wire protocol mismatch: got %s want %s", workWireProtocol, ctl.sessionCtx.WireProtocol)
}
// server plugin hook // server plugin hook
content := &plugin.NewWorkConnContent{ content := &plugin.NewWorkConnContent{
@@ -830,20 +885,33 @@ func (svr *Service) RegisterWorkConn(workConn *msg.Conn, newMsg *msg.NewWorkConn
xl.Warnf("invalid NewWorkConn with run id [%s]", newMsg.RunID) xl.Warnf("invalid NewWorkConn with run id [%s]", newMsg.RunID)
return err return err
} }
return 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 { 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. // TODO(deprecation): Compatible with old versions, can be without runID, user is empty. In later versions, it will be mandatory to include runID.
// If runID is required, it is not compatible with versions prior to v0.50.0. // If runID is required, it is not compatible with versions prior to v0.50.0.
if newMsg.RunID != "" { if newMsg.RunID != "" {
ctl, exist := svr.ctlManager.GetByID(newMsg.RunID) admitted, err := svr.ctlManager.admitVisitorByRunID(newMsg.RunID, func(visitorUser, controlWireProtocol, controlUDPPacketCodec string) error {
if !exist { 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) 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, return admit("", wireProtocol, "")
newMsg.UseEncryption, newMsg.UseCompression, visitorUser, wireProtocol)
} }
+913
View File
@@ -15,12 +15,33 @@
package server package server
import ( import (
"context"
"errors" "errors"
"fmt"
"math"
"net" "net"
"net/http"
"runtime"
"strings"
"sync"
"sync/atomic"
"testing" "testing"
"time" "time"
"github.com/fatedier/golib/net/mux"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"github.com/fatedier/frp/pkg/auth"
v1 "github.com/fatedier/frp/pkg/config/v1"
"github.com/fatedier/frp/pkg/config/v1/validation"
"github.com/fatedier/frp/pkg/msg"
plugin "github.com/fatedier/frp/pkg/plugin/server"
"github.com/fatedier/frp/pkg/proto/wire"
"github.com/fatedier/frp/pkg/util/util"
"github.com/fatedier/frp/server/controller"
"github.com/fatedier/frp/server/proxy"
"github.com/fatedier/frp/server/registry"
"github.com/fatedier/frp/server/visitor"
) )
func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) { func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) {
@@ -61,3 +82,895 @@ func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) {
t.Fatal("timed out waiting for write after deadline reset") t.Fatal("timed out waiting for write after deadline reset")
} }
} }
func TestServiceAcceptConnectionTracksClientHelloPresence(t *testing.T) {
for _, tc := range []struct {
name string
clientHelloPresent bool
offeredCodecs []string
expectedCodec string
}{
{
name: "absent Hello",
},
{
name: "present Hello with JSON fallback",
clientHelloPresent: true,
},
{
name: "present Hello with binary codec",
clientHelloPresent: true,
offeredCodecs: []string{wire.UDPPacketCodecBinary},
expectedCodec: wire.UDPPacketCodecBinary,
},
} {
t.Run(tc.name, func(t *testing.T) {
serverConn, clientConn := net.Pipe()
defer serverConn.Close()
defer clientConn.Close()
clientErrCh := make(chan error, 1)
go func() {
if err := wire.WriteMagic(clientConn); err != nil {
clientErrCh <- err
return
}
wireConn := wire.NewConn(clientConn)
if tc.clientHelloPresent {
hello, err := wire.NewClientHello(wire.BootstrapInfo{})
if err != nil {
clientErrCh <- err
return
}
hello.Capabilities.Message.UDPPacketCodecs = tc.offeredCodecs
if err := wireConn.WriteJSONFrame(wire.FrameTypeClientHello, hello); err != nil {
clientErrCh <- err
return
}
var serverHello wire.ServerHello
if err := wireConn.ReadJSONFrame(wire.FrameTypeServerHello, &serverHello); err != nil {
clientErrCh <- err
return
}
}
clientErrCh <- msg.NewV2ReadWriterWithConn(wireConn).WriteMsg(&msg.NewWorkConn{RunID: "shared-run"})
}()
acceptedConn, err := (&Service{}).acceptConnection(t.Context(), serverConn)
require.NoError(t, err)
require.NoError(t, <-clientErrCh)
require.Equal(t, tc.clientHelloPresent, acceptedConn.clientHelloPresent)
require.Equal(t, tc.expectedCodec, acceptedConn.udpPacketCodec)
require.IsType(t, &msg.NewWorkConn{}, acceptedConn.firstMsg)
require.NoError(t, acceptedConn.conn.Close())
})
}
}
func TestSharedPortHTTPListenerProtocols(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
sharedMux := mux.NewMux(listener)
httpListener := sharedMux.ListenHTTP(1)
muxServeErr := make(chan error, 1)
go func() {
muxServeErr <- sharedMux.Serve()
}()
newProtocols := func(http1, unencryptedHTTP2 bool) *http.Protocols {
protocols := new(http.Protocols)
protocols.SetHTTP1(http1)
protocols.SetUnencryptedHTTP2(unencryptedHTTP2)
return protocols
}
const handlerProtocolHeader = "X-Test-Handler-Protocol"
httpServer := &http.Server{
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set(handlerProtocolHeader, r.Proto)
w.WriteHeader(http.StatusNoContent)
}),
ReadHeaderTimeout: time.Second,
Protocols: newProtocols(true, true),
}
httpServeErr := make(chan error, 1)
go func() {
httpServeErr <- httpServer.Serve(httpListener)
}()
t.Cleanup(func() {
require.NoError(t, httpServer.Close())
require.ErrorIs(t, waitForResult(t, httpServeErr, "shared HTTP server to stop"), http.ErrServerClosed)
require.NoError(t, sharedMux.Close())
require.ErrorIs(t, waitForResult(t, muxServeErr, "shared mux to stop"), net.ErrClosed)
})
for _, tc := range []struct {
name string
http1 bool
unencryptedHTTP2 bool
expectedProtocol string
}{
{name: "HTTP/1.1", http1: true, expectedProtocol: "HTTP/1.1"},
{name: "HTTP/2 prior knowledge", unencryptedHTTP2: true, expectedProtocol: "HTTP/2.0"},
} {
t.Run(tc.name, func(t *testing.T) {
transport := &http.Transport{
Protocols: newProtocols(tc.http1, tc.unencryptedHTTP2),
}
defer transport.CloseIdleConnections()
client := &http.Client{Transport: transport, Timeout: 3 * time.Second}
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "http://"+listener.Addr().String()+"/", nil)
require.NoError(t, err)
response, err := client.Do(request)
require.NoError(t, err)
require.Equal(t, http.StatusNoContent, response.StatusCode)
require.Equal(t, tc.expectedProtocol, response.Proto)
require.Equal(t, tc.expectedProtocol, response.Header.Get(handlerProtocolHeader))
require.NoError(t, response.Body.Close())
})
}
}
func TestServiceControlHandoffSkipsStalePendingGeneration(t *testing.T) {
svr := newControlTestService(t)
metrics := newCountingServerMetrics()
metrics.closeEnter = make(chan struct{})
metrics.closeResume = make(chan struct{})
ctlA, connA, err := registerLifecycleTestControl(svr)
require.NoError(t, err)
ctlA.serverMetrics = metrics
require.NoError(t, svr.completeControlLogin(ctlA, func() error { return nil }))
waitForSignal(t, connA.readStarted, "A reader to start")
require.NoError(t, ctlA.Close())
waitForSignal(t, metrics.closeEnter, "A finalization barrier")
type registerResult struct {
ctl *Control
conn *deadlineReadConn
err error
}
resultB := make(chan registerResult, 1)
go func() {
ctl, conn, registerErr := registerLifecycleTestControl(svr)
resultB <- registerResult{ctl: ctl, conn: conn, err: registerErr}
}()
ctlB := waitForDifferentCurrentControl(t, svr.ctlManager, "shared-run", ctlA)
ctlB.serverMetrics = metrics
resultC := make(chan registerResult, 1)
go func() {
ctl, conn, registerErr := registerLifecycleTestControl(svr)
resultC <- registerResult{ctl: ctl, conn: conn, err: registerErr}
}()
ctlC := waitForDifferentCurrentControl(t, svr.ctlManager, "shared-run", ctlB)
ctlC.serverMetrics = metrics
waitForControlDone(t, ctlB)
select {
case result := <-resultB:
t.Fatalf("B returned before A finalized: %v", result.err)
default:
}
select {
case result := <-resultC:
t.Fatalf("C returned before A finalized: %v", result.err)
default:
}
close(metrics.closeResume)
waitForControlDone(t, ctlA)
b := <-resultB
require.Same(t, ctlB, b.ctl)
require.ErrorIs(t, b.err, errControlReplaced)
require.False(t, svr.ctlManager.Remove(ctlB))
require.NoError(t, ctlB.Close())
c := <-resultC
require.NoError(t, c.err)
require.Same(t, ctlC, c.ctl)
_, ok := svr.ctlManager.GetByID("shared-run")
require.False(t, ok)
require.Same(t, ctlC, currentControlForTest(svr.ctlManager, "shared-run"))
info, ok := svr.clientRegistry.GetByKey("client")
require.True(t, ok)
require.True(t, info.Online)
require.Equal(t, uint64(ctlC.ID()), info.ControlID)
var staleWrites atomic.Int64
err = svr.completeControlLogin(ctlB, func() error {
staleWrites.Add(1)
return nil
})
require.ErrorIs(t, err, errControlReplaced)
require.Equal(t, int64(0), staleWrites.Load())
require.NoError(t, svr.completeControlLogin(ctlC, func() error { return nil }))
waitForSignal(t, c.conn.readStarted, "C reader to start")
current, ok := svr.ctlManager.GetByID("shared-run")
require.True(t, ok)
require.Same(t, ctlC, current)
require.Equal(t, int64(2), metrics.newClients())
require.Equal(t, int64(1), metrics.closedClients())
require.NoError(t, ctlC.Close())
waitForControlDone(t, ctlC)
require.Equal(t, int64(2), metrics.newClients())
require.Equal(t, int64(2), metrics.closedClients())
_, ok = svr.ctlManager.GetByID("shared-run")
require.False(t, ok)
}
func TestServiceLoginResponseSynchronizationIsScopedToRun(t *testing.T) {
svr := newControlTestService(t)
metrics := newCountingServerMetrics()
ctlA, connA, err := registerLifecycleTestControl(svr)
require.NoError(t, err)
ctlA.serverMetrics = metrics
writeEntered := make(chan struct{})
resumeWrite := make(chan struct{})
var resumeWriteOnce sync.Once
resume := func() {
resumeWriteOnce.Do(func() { close(resumeWrite) })
}
t.Cleanup(resume)
writeCount := atomic.Int64{}
loginDone := make(chan error, 1)
go func() {
loginDone <- svr.completeControlLogin(ctlA, func() error {
close(writeEntered)
<-resumeWrite
writeCount.Add(1)
return nil
})
}()
waitForSignal(t, writeEntered, "A LoginResp write")
runMu := currentRunGateForTest(svr.ctlManager, "shared-run")
require.NotNil(t, runMu)
if !svr.ctlManager.mu.TryLock() {
t.Fatal("ControlManager mutex was held while LoginResp write was in progress")
}
svr.ctlManager.mu.Unlock()
ctlB, connB := newLifecycleTestControl(t, "shared-run", "client", metrics)
gateAvailable := make(chan bool)
addDone := make(chan error, 1)
go func() {
if runMu.TryLock() {
runMu.Unlock()
gateAvailable <- true
} else {
gateAvailable <- false
}
addErr := svr.ctlManager.Add(ctlB)
addDone <- addErr
}()
available := waitForResult(t, gateAvailable, "same-run replacement gate probe")
require.False(t, available, "same-run gate was available to replacement during LoginResp write")
select {
case addErr := <-addDone:
t.Fatalf("same-run replacement completed during LoginResp write: %v", addErr)
case <-time.After(20 * time.Millisecond):
}
require.Same(t, ctlA, currentControlForTest(svr.ctlManager, "shared-run"))
otherMetrics := newCountingServerMetrics()
otherCtl, otherConn := newLifecycleTestControl(t, "other-run", "other-client", otherMetrics)
type unrelatedResult struct {
addErr error
active bool
activateErr error
loginErr error
current *Control
found bool
}
unrelatedDone := make(chan unrelatedResult, 1)
go func() {
result := unrelatedResult{}
result.addErr = svr.ctlManager.Add(otherCtl)
if result.addErr == nil {
result.active, result.activateErr = svr.ctlManager.Activate(otherCtl)
}
if result.activateErr == nil && result.active {
result.loginErr = svr.completeControlLogin(otherCtl, func() error { return nil })
}
result.current, result.found = svr.ctlManager.GetByID("other-run")
unrelatedDone <- result
}()
result := waitForResult(t, unrelatedDone, "unrelated run lifecycle")
require.NoError(t, result.addErr)
require.NoError(t, result.activateErr)
require.True(t, result.active)
require.NoError(t, result.loginErr)
require.True(t, result.found)
require.Same(t, otherCtl, result.current)
waitForSignal(t, otherConn.readStarted, "unrelated control reader to start")
require.Equal(t, int64(1), otherMetrics.newClients())
resume()
require.NoError(t, waitForResult(t, loginDone, "LoginResp completion"))
require.NoError(t, waitForResult(t, addDone, "replacement"))
waitForControlDone(t, ctlA)
require.Same(t, ctlB, currentControlForTest(svr.ctlManager, "shared-run"))
require.Equal(t, int64(1), writeCount.Load())
require.Equal(t, int64(1), metrics.newClients())
require.Equal(t, int64(1), metrics.closedClients())
require.Equal(t, []string{"deadline", "close"}, connA.eventsSnapshot())
require.False(t, svr.ctlManager.Remove(ctlA))
require.NoError(t, ctlA.Close())
require.True(t, svr.ctlManager.Remove(ctlB))
require.NoError(t, ctlB.Close())
require.Equal(t, []string{"deadline", "close"}, connB.eventsSnapshot())
require.NoError(t, otherCtl.Close())
waitForControlDone(t, otherCtl)
require.Equal(t, int64(1), otherMetrics.newClients())
require.Equal(t, int64(1), otherMetrics.closedClients())
}
func TestServiceVisitorAdmissionSerializesReplacement(t *testing.T) {
svr := newControlTestService(t)
ctlA, controlConn, err := registerLifecycleTestControl(svr)
require.NoError(t, err)
ctlA.sessionCtx.LoginMsg.User = "old-user"
require.NoError(t, svr.completeControlLogin(ctlA, func() error { return nil }))
waitForSignal(t, controlConn.readStarted, "A reader to start")
admissionEntered := make(chan struct{})
resumeAdmission := make(chan struct{})
var resumeOnce sync.Once
resume := func() {
resumeOnce.Do(func() { close(resumeAdmission) })
}
t.Cleanup(resume)
type admissionResult struct {
admitted bool
user string
wireProtocol string
udpPacketCodec string
err error
}
admissionDone := make(chan admissionResult, 1)
go func() {
result := admissionResult{}
result.admitted, result.err = svr.ctlManager.admitVisitorByRunID("shared-run", func(user, wireProtocol, udpPacketCodec string) error {
result.user = user
result.wireProtocol = wireProtocol
result.udpPacketCodec = udpPacketCodec
close(admissionEntered)
<-resumeAdmission
return nil
})
admissionDone <- result
}()
waitForSignal(t, admissionEntered, "visitor admission callback")
runMu := currentRunGateForTest(svr.ctlManager, "shared-run")
require.NotNil(t, runMu)
type registerResult struct {
ctl *Control
err error
}
gateAvailable := make(chan bool)
replacementDone := make(chan registerResult, 1)
go func() {
if runMu.TryLock() {
runMu.Unlock()
gateAvailable <- true
} else {
gateAvailable <- false
}
ctl, _, registerErr := registerLifecycleTestControl(svr)
replacementDone <- registerResult{ctl: ctl, err: registerErr}
}()
available := waitForResult(t, gateAvailable, "visitor replacement gate probe")
require.False(t, available, "same-run gate was available during visitor admission")
select {
case result := <-replacementDone:
t.Fatalf("replacement completed during visitor admission: %v", result.err)
case <-time.After(20 * time.Millisecond):
}
require.Same(t, ctlA, currentControlForTest(svr.ctlManager, "shared-run"))
resume()
admission := waitForResult(t, admissionDone, "visitor admission")
require.NoError(t, admission.err)
require.True(t, admission.admitted)
require.Equal(t, "old-user", admission.user)
require.Equal(t, wire.ProtocolV1, admission.wireProtocol)
require.Empty(t, admission.udpPacketCodec)
replacement := waitForResult(t, replacementDone, "replacement")
require.NoError(t, replacement.err)
ctlB := replacement.ctl
require.Same(t, ctlB, currentControlForTest(svr.ctlManager, "shared-run"))
waitForControlDone(t, ctlA)
require.True(t, svr.ctlManager.Remove(ctlB))
require.NoError(t, ctlB.Close())
}
func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) {
svr := newControlTestService(t)
ctl, controlConn, err := registerLifecycleTestControl(svr)
require.NoError(t, err)
pendingConn := newCountingCloseConn()
pendingMsgConn := msg.NewConn(pendingConn, msg.NewV1ReadWriter(pendingConn))
err = registerWorkConnAsCaller(svr, pendingMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false)
require.Error(t, err)
require.Equal(t, int64(1), pendingConn.closeCount.Load())
require.Len(t, ctl.workConnCh, 0)
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
waitForSignal(t, controlConn.readStarted, "control reader to start")
current, ok := svr.ctlManager.GetByID("shared-run")
require.True(t, ok)
require.Same(t, ctl, current)
require.Len(t, ctl.workConnCh, 0)
runningConn := newCountingCloseConn()
runningMsgConn := msg.NewConn(runningConn, msg.NewV1ReadWriter(runningConn))
require.NoError(t, svr.RegisterWorkConn(runningMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false))
require.Len(t, ctl.workConnCh, 1)
require.NoError(t, ctl.Close())
waitForControlDone(t, ctl)
require.Equal(t, int64(1), runningConn.closeCount.Load())
}
func TestServiceWorkConnRoutingRejectsWireProtocolMismatch(t *testing.T) {
svr := newControlTestService(t)
ctl, controlConn, err := registerLifecycleTestControl(svr)
require.NoError(t, err)
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
waitForSignal(t, controlConn.readStarted, "control reader to start")
workConn := newCountingCloseConn()
workMsgConn := msg.NewConn(workConn, msg.NewV2ReadWriter(workConn))
err = svr.RegisterWorkConn(workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV2, false)
require.ErrorContains(t, err, "wire protocol mismatch")
require.Len(t, ctl.workConnCh, 0)
_ = workMsgConn.Close()
require.NoError(t, ctl.Close())
}
func TestServiceWorkConnRoutingClientHelloPolicy(t *testing.T) {
for _, tc := range []struct {
name string
controlUDPPacketCodec string
workClientHelloPresent bool
errorSubstring string
}{
{
name: "JSON control allows work connection without Hello",
},
{
name: "binary control allows work connection without Hello",
controlUDPPacketCodec: wire.UDPPacketCodecBinary,
},
{
name: "JSON control rejects work connection with Hello",
workClientHelloPresent: true,
errorSubstring: "ClientHello is not allowed",
},
{
name: "binary control rejects work connection with Hello",
controlUDPPacketCodec: wire.UDPPacketCodecBinary,
workClientHelloPresent: true,
errorSubstring: "ClientHello is not allowed",
},
} {
t.Run(tc.name, func(t *testing.T) {
svr := newControlTestService(t)
controlConn := newDeadlineReadConn()
controlMsgConn := msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn))
ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{
RunID: "shared-run",
ClientID: "client",
ClientSpec: msg.ClientSpec{
AlwaysAuthPass: true,
},
}, true, wire.ProtocolV2, tc.controlUDPPacketCodec)
require.NoError(t, err)
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
waitForSignal(t, controlConn.readStarted, "control reader to start")
workConn := newCountingCloseConn()
workMsgConn := msg.NewConn(workConn, msg.NewV2ReadWriter(workConn))
err = svr.RegisterWorkConn(
workMsgConn,
&msg.NewWorkConn{RunID: "shared-run"},
wire.ProtocolV2,
tc.workClientHelloPresent,
)
if tc.errorSubstring != "" {
require.ErrorContains(t, err, tc.errorSubstring)
require.Len(t, ctl.workConnCh, 0)
require.NoError(t, workMsgConn.Close())
} else {
require.NoError(t, err)
require.Len(t, ctl.workConnCh, 1)
}
require.NoError(t, ctl.Close())
waitForControlDone(t, ctl)
require.Equal(t, int64(1), workConn.closeCount.Load())
})
}
}
func TestServiceRegisterControlRejectsInvalidCodecSelection(t *testing.T) {
for _, tc := range []struct {
name string
wireProtocol string
udpPacketCodec string
errorSubstring string
}{
{
name: "binary codec over v1",
wireProtocol: wire.ProtocolV1,
udpPacketCodec: wire.UDPPacketCodecBinary,
errorSubstring: "requires wire protocol v2",
},
{
name: "unknown v2 codec",
wireProtocol: wire.ProtocolV2,
udpPacketCodec: "unknown",
errorSubstring: "unsupported UDP packet codec",
},
{
name: "unknown wire protocol",
wireProtocol: "unknown",
errorSubstring: "unsupported wire protocol",
},
} {
t.Run(tc.name, func(t *testing.T) {
svr := newControlTestService(t)
conn := newDeadlineReadConn()
msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn))
ctl, err := svr.RegisterControl(msgConn, &msg.Login{}, true, tc.wireProtocol, tc.udpPacketCodec)
require.Nil(t, ctl)
require.ErrorContains(t, err, tc.errorSubstring)
})
}
}
func TestServiceRegisterControlRejectsInvalidRunID(t *testing.T) {
for _, runID := range []string{
"run\nforged",
strings.Repeat("a", validation.MaxRunIDLength+1),
} {
t.Run(fmt.Sprintf("run_id_%d", len(runID)), func(t *testing.T) {
svr := newControlTestService(t)
conn := newDeadlineReadConn()
msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn))
ctl, err := svr.RegisterControl(msgConn, &msg.Login{RunID: runID}, true, wire.ProtocolV1, "")
require.Nil(t, ctl)
require.ErrorContains(t, err, "invalid run id")
})
}
}
func TestServiceRegisterControlPoolCountBoundaries(t *testing.T) {
for _, tc := range []struct {
name string
poolCount int
wantErr bool
wantPool int
}{
{name: "less than channel offset", poolCount: -11, wantErr: true},
{name: "channel offset", poolCount: -10, wantErr: true},
{name: "negative", poolCount: -1, wantErr: true},
{name: "zero", poolCount: 0, wantPool: 0},
{name: "capped by server maximum", poolCount: 10, wantPool: 5},
{name: "maximum int capped", poolCount: math.MaxInt, wantPool: 5},
} {
t.Run(tc.name, func(t *testing.T) {
svr := newControlTestService(t)
svr.cfg.Transport.MaxPoolCount = 5
conn := newDeadlineReadConn()
msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn))
const timestamp = int64(1)
ctl, err := svr.RegisterControl(msgConn, &msg.Login{
RunID: "pool-count-run",
ClientID: "client",
Timestamp: timestamp,
PrivilegeKey: util.GetAuthKey("", timestamp),
PoolCount: tc.poolCount,
}, false, wire.ProtocolV1, "")
if tc.wantErr {
require.Nil(t, ctl)
require.ErrorContains(t, err, "unexpected error when creating new controller")
return
}
require.NoError(t, err)
require.Equal(t, tc.wantPool, ctl.poolCount)
require.NoError(t, ctl.Close())
waitForControlDone(t, ctl)
})
}
}
func TestServiceWorkConnRoutingRejectsLostGeneration(t *testing.T) {
for _, action := range []string{"replace", "close"} {
t.Run(action, func(t *testing.T) {
svr := newControlTestService(t)
ctl, controlConn, err := registerLifecycleTestControl(svr)
require.NoError(t, err)
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
waitForSignal(t, controlConn.readStarted, "control reader to start")
barrier := newWorkConnBarrierPlugin()
svr.pluginManager.Register(barrier)
workConn := newCountingCloseConn()
workMsgConn := msg.NewConn(workConn, msg.NewV1ReadWriter(workConn))
routeDone := make(chan error, 1)
go func() {
routeDone <- registerWorkConnAsCaller(svr, workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false)
}()
waitForSignal(t, barrier.entered, "work connection plugin barrier")
var replacement *Control
switch action {
case "replace":
replacement, _, err = registerLifecycleTestControl(svr)
require.NoError(t, err)
case "close":
require.NoError(t, ctl.Close())
waitForControlDone(t, ctl)
}
close(barrier.resume)
require.Error(t, waitForResult(t, routeDone, "work connection route to finish"))
require.Equal(t, int64(1), workConn.closeCount.Load())
require.Len(t, ctl.workConnCh, 0)
if replacement != nil {
require.Len(t, replacement.workConnCh, 0)
require.True(t, svr.ctlManager.Remove(replacement))
require.NoError(t, replacement.Close())
waitForControlDone(t, replacement)
}
})
}
}
func TestServiceVisitorRoutingExcludesPendingUser(t *testing.T) {
svr := newControlTestService(t)
listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{"pending-user"})
require.NoError(t, err)
t.Cleanup(func() { _ = listener.Close() })
controlConn := newDeadlineReadConn()
controlMsgConn := msg.NewConn(controlConn, msg.NewV1ReadWriter(controlConn))
ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{
RunID: "visitor-run",
User: "pending-user",
ClientID: "visitor-client",
ClientSpec: msg.ClientSpec{
AlwaysAuthPass: true,
},
}, true, wire.ProtocolV1, "")
require.NoError(t, err)
timestamp := time.Now().Unix()
visitorMsg := &msg.NewVisitorConn{
RunID: "visitor-run",
ProxyName: "visitor",
Timestamp: timestamp,
SignKey: util.GetAuthKey("secret", timestamp),
}
pendingConn := newCountingCloseConn()
err = svr.RegisterVisitorConn(pendingConn, visitorMsg, wire.ProtocolV1)
require.ErrorContains(t, err, "no client control found")
require.NoError(t, pendingConn.Close())
require.Equal(t, int64(1), pendingConn.closeCount.Load())
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
waitForSignal(t, controlConn.readStarted, "control reader to start")
runningConn := newCountingCloseConn()
require.NoError(t, svr.RegisterVisitorConn(runningConn, visitorMsg, wire.ProtocolV1))
accepted, err := listener.Accept()
require.NoError(t, err)
require.NoError(t, accepted.Close())
require.Equal(t, int64(1), runningConn.closeCount.Load())
require.NoError(t, ctl.Close())
waitForControlDone(t, ctl)
}
func TestServiceVisitorRoutingCarriesControlPacketCodec(t *testing.T) {
svr := newControlTestService(t)
listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{"visitor-user"})
require.NoError(t, err)
t.Cleanup(func() { _ = listener.Close() })
controlConn := newDeadlineReadConn()
controlMsgConn := msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn))
ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{
RunID: "visitor-binary-run",
User: "visitor-user",
ClientID: "visitor-client",
ClientSpec: msg.ClientSpec{
AlwaysAuthPass: true,
},
}, true, wire.ProtocolV2, wire.UDPPacketCodecBinary)
require.NoError(t, err)
timestamp := time.Now().Unix()
visitorMsg := &msg.NewVisitorConn{
RunID: "visitor-binary-run",
ProxyName: "visitor",
Timestamp: timestamp,
SignKey: util.GetAuthKey("secret", timestamp),
}
require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil }))
waitForSignal(t, controlConn.readStarted, "binary visitor control reader to start")
runningConn := newCountingCloseConn()
require.NoError(t, svr.RegisterVisitorConn(runningConn, visitorMsg, wire.ProtocolV2))
accepted, err := listener.Accept()
require.NoError(t, err)
metadata, ok := accepted.(interface {
WireProtocol() string
UDPPacketCodec() string
})
require.True(t, ok)
require.Equal(t, wire.ProtocolV2, metadata.WireProtocol())
require.Equal(t, wire.UDPPacketCodecBinary, metadata.UDPPacketCodec())
require.NoError(t, accepted.Close())
require.Equal(t, int64(1), runningConn.closeCount.Load())
mismatchConn := newCountingCloseConn()
err = svr.RegisterVisitorConn(mismatchConn, visitorMsg, wire.ProtocolV1)
require.ErrorContains(t, err, "visitor connection wire protocol mismatch")
require.NoError(t, mismatchConn.Close())
require.Equal(t, int64(1), mismatchConn.closeCount.Load())
require.NoError(t, ctl.Close())
waitForControlDone(t, ctl)
}
func TestServiceVisitorRoutingLegacyFallsBackToJSONPacketCodec(t *testing.T) {
svr := newControlTestService(t)
listener, err := svr.rc.VisitorManager.Listen("visitor", "secret", []string{""})
require.NoError(t, err)
t.Cleanup(func() { _ = listener.Close() })
timestamp := time.Now().Unix()
visitorMsg := &msg.NewVisitorConn{
ProxyName: "visitor",
Timestamp: timestamp,
SignKey: util.GetAuthKey("secret", timestamp),
}
visitorConn := newCountingCloseConn()
require.NoError(t, svr.RegisterVisitorConn(visitorConn, visitorMsg, wire.ProtocolV2))
accepted, err := listener.Accept()
require.NoError(t, err)
metadata, ok := accepted.(interface{ UDPPacketCodec() string })
require.True(t, ok)
require.Empty(t, metadata.UDPPacketCodec())
require.NoError(t, accepted.Close())
require.Equal(t, int64(1), visitorConn.closeCount.Load())
}
func newControlTestService(t *testing.T) *Service {
t.Helper()
cfg := &v1.ServerConfig{}
cfg.Auth.Method = v1.AuthMethodToken
authRuntime, err := auth.BuildServerAuth(&cfg.Auth)
require.NoError(t, err)
clientRegistry := registry.NewClientRegistry()
return &Service{
ctlManager: NewControlManager(clientRegistry),
clientRegistry: clientRegistry,
pxyManager: proxy.NewManager(),
pluginManager: plugin.NewManager(),
rc: &controller.ResourceController{
VisitorManager: visitor.NewManager(),
},
auth: authRuntime,
cfg: cfg,
}
}
func registerLifecycleTestControl(svr *Service) (*Control, *deadlineReadConn, error) {
conn := newDeadlineReadConn()
msgConn := msg.NewConn(conn, msg.NewReadWriter(conn, wire.ProtocolV1))
ctl, err := svr.RegisterControl(msgConn, &msg.Login{
RunID: "shared-run",
ClientID: "client",
ClientSpec: msg.ClientSpec{
AlwaysAuthPass: true,
},
}, true, wire.ProtocolV1, "")
return ctl, conn, err
}
func waitForDifferentCurrentControl(t *testing.T, manager *ControlManager, runID string, old *Control) *Control {
t.Helper()
deadline := time.Now().Add(3 * time.Second)
for time.Now().Before(deadline) {
if ctl := currentControlForTest(manager, runID); ctl != nil && ctl != old {
return ctl
}
runtime.Gosched()
}
t.Fatalf("timed out waiting for a new current control after ID %d", old.ID())
return nil
}
func registerWorkConnAsCaller(
svr *Service,
workConn *msg.Conn,
newMsg *msg.NewWorkConn,
wireProtocol string,
clientHelloPresent bool,
) error {
err := svr.RegisterWorkConn(workConn, newMsg, wireProtocol, clientHelloPresent)
if err != nil {
_ = workConn.Close()
}
return err
}
func waitForResult[T any](t *testing.T, ch <-chan T, description string) T {
t.Helper()
select {
case result := <-ch:
return result
case <-time.After(3 * time.Second):
t.Fatalf("timed out waiting for %s", description)
var zero T
return zero
}
}
type workConnBarrierPlugin struct {
entered chan struct{}
resume chan struct{}
}
func newWorkConnBarrierPlugin() *workConnBarrierPlugin {
return &workConnBarrierPlugin{
entered: make(chan struct{}),
resume: make(chan struct{}),
}
}
func (*workConnBarrierPlugin) Name() string { return "work-conn-barrier" }
func (*workConnBarrierPlugin) IsSupport(op string) bool { return op == plugin.OpNewWorkConn }
func (p *workConnBarrierPlugin) Handle(
context.Context,
string,
any,
) (*plugin.Response, any, error) {
close(p.entered)
<-p.resume
return &plugin.Response{Unchange: true}, nil, nil
}
type countingCloseConn struct {
closeCount atomic.Int64
}
func newCountingCloseConn() *countingCloseConn { return &countingCloseConn{} }
func (*countingCloseConn) Read([]byte) (int, error) { return 0, net.ErrClosed }
func (*countingCloseConn) Write(p []byte) (int, error) { return len(p), nil }
func (c *countingCloseConn) Close() error { c.closeCount.Add(1); return nil }
func (*countingCloseConn) LocalAddr() net.Addr { return lifecycleTestAddr("local") }
func (*countingCloseConn) RemoteAddr() net.Addr { return lifecycleTestAddr("remote") }
func (*countingCloseConn) SetDeadline(time.Time) error { return nil }
func (*countingCloseConn) SetReadDeadline(time.Time) error { return nil }
func (*countingCloseConn) SetWriteDeadline(time.Time) error { return nil }
+14 -4
View File
@@ -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, func (vm *Manager) NewConn(name string, conn net.Conn, timestamp int64, signKey string,
useEncryption bool, useCompression bool, visitorUser string, useEncryption bool, useCompression bool, visitorUser string,
wireProtocol string, wireProtocol string, udpPacketCodecs ...string,
) (err error) { ) (err error) {
udpPacketCodec := ""
if len(udpPacketCodecs) > 0 {
udpPacketCodec = udpPacketCodecs[0]
}
vm.mu.RLock() vm.mu.RLock()
defer vm.mu.RUnlock() defer vm.mu.RUnlock()
@@ -93,8 +97,9 @@ func (vm *Manager) NewConn(name string, conn net.Conn, timestamp int64, signKey
} }
visitorConn := netpkg.WrapReadWriteCloserToConn(rwc, conn) visitorConn := netpkg.WrapReadWriteCloserToConn(rwc, conn)
err = l.l.PutConn(&wireProtocolConn{ err = l.l.PutConn(&wireProtocolConn{
Conn: visitorConn, Conn: visitorConn,
wireProtocol: wireProtocol, wireProtocol: wireProtocol,
udpPacketCodec: udpPacketCodec,
}) })
} else { } else {
err = fmt.Errorf("custom listener for [%s] doesn't exist", name) err = fmt.Errorf("custom listener for [%s] doesn't exist", name)
@@ -105,13 +110,18 @@ func (vm *Manager) NewConn(name string, conn net.Conn, timestamp int64, signKey
type wireProtocolConn struct { type wireProtocolConn struct {
net.Conn net.Conn
wireProtocol string wireProtocol string
udpPacketCodec string
} }
func (c *wireProtocolConn) WireProtocol() string { func (c *wireProtocolConn) WireProtocol() string {
return c.wireProtocol return c.wireProtocol
} }
func (c *wireProtocolConn) UDPPacketCodec() string {
return c.udpPacketCodec
}
func (vm *Manager) CloseListener(name string) { func (vm *Manager) CloseListener(name string) {
vm.mu.Lock() vm.mu.Lock()
defer vm.mu.Unlock() defer vm.mu.Unlock()
+8 -3
View File
@@ -25,7 +25,7 @@ import (
"github.com/fatedier/frp/pkg/util/util" "github.com/fatedier/frp/pkg/util/util"
) )
func TestManagerNewConnCarriesWireProtocol(t *testing.T) { func TestManagerNewConnCarriesWireProtocolAndUDPPacketCodec(t *testing.T) {
vm := NewManager() vm := NewManager()
listener, err := vm.Listen("sudp", "secret", []string{"*"}) listener, err := vm.Listen("sudp", "secret", []string{"*"})
require.NoError(t, err) require.NoError(t, err)
@@ -47,6 +47,7 @@ func TestManagerNewConnCarriesWireProtocol(t *testing.T) {
false, false,
"user", "user",
wire.ProtocolV2, wire.ProtocolV2,
wire.UDPPacketCodecBinary,
) )
}() }()
@@ -54,8 +55,12 @@ func TestManagerNewConnCarriesWireProtocol(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
defer acceptedConn.Close() defer acceptedConn.Close()
getter, ok := acceptedConn.(interface{ WireProtocol() string }) metadata, ok := acceptedConn.(interface {
WireProtocol() string
UDPPacketCodec() string
})
require.True(t, ok) 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) require.NoError(t, <-errCh)
} }
+102 -1
View File
@@ -7,6 +7,8 @@ import (
"io" "io"
"net/http" "net/http"
"os" "os"
"strconv"
"strings"
"testing" "testing"
"time" "time"
@@ -158,7 +160,7 @@ webServer.port = %d
framework.NewRequestExpect(f).PortName(portName).Ensure() framework.NewRequestExpect(f).PortName(portName).Ensure()
}) })
ginkgo.It("baseline frps rejects current frpc forced to v2", func() { ginkgo.It("baseline frps handles current frpc forced to v2 according to baseline support", func() {
portName := port.GenName("CompatBaselineFRPSForcedV2") portName := port.GenName("CompatBaselineFRPSForcedV2")
clientConf := tcpClientConfig("tcp", portName, ` clientConf := tcpClientConfig("tcp", portName, `
transport.wireProtocol = "v2" transport.wireProtocol = "v2"
@@ -170,11 +172,61 @@ transport.wireProtocol = "v2"
consts.DefaultServerConfig, consts.DefaultServerConfig,
[]string{clientConf}, []string{clientConf},
) )
// frp v0.69.0 added control connection wireProtocol v2 support, so
// baseline frps v0.69.0 and newer should accept a current frpc forced
// to v2. Older known baselines still must reject the unsupported protocol.
// For custom baselines, the version is unknown and the binary may be either
// side of the support boundary, so this versioned expectation is skipped.
supportsV2, knownVersion := baselineSupportsControlWireProtocolV2(compatCtx.BaselineVersion)
if !knownVersion {
ginkgo.Skip(fmt.Sprintf("baseline version %q is not semver; skip versioned forced-v2 expectation", compatCtx.BaselineVersion))
}
if supportsV2 {
framework.NewRequestExpect(f).PortName(portName).Ensure()
return
}
expectProcessExit(clientProcesses[0], 5*time.Second) expectProcessExit(clientProcesses[0], 5*time.Second)
framework.NewRequestExpect(f).PortName(portName).ExpectError(true).Ensure() framework.NewRequestExpect(f).PortName(portName).ExpectError(true).Ensure()
}) })
}) })
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 { func tcpClientConfig(proxyName string, remotePortName string, extra string) string {
return fmt.Sprintf(` return fmt.Sprintf(`
serverAddr = "127.0.0.1" serverAddr = "127.0.0.1"
@@ -191,6 +243,22 @@ remotePort = {{ .%s }}
`, consts.PortServerName, extra, proxyName, framework.TCPEchoServerPort, remotePortName) `, 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) { func expectProcessExit(p *process.Process, timeout time.Duration) {
select { select {
case <-p.Done(): case <-p.Done():
@@ -199,6 +267,39 @@ func expectProcessExit(p *process.Process, timeout time.Duration) {
} }
} }
func baselineSupportsControlWireProtocolV2(version string) (supports bool, known bool) {
version = strings.TrimPrefix(version, "v")
parts := strings.Split(version, ".")
if len(parts) != 3 {
return false, false
}
major, err := strconv.Atoi(parts[0])
if err != nil {
return false, false
}
minor, err := strconv.Atoi(parts[1])
if err != nil {
return false, false
}
patch, err := strconv.Atoi(parts[2])
if err != nil {
return false, false
}
return compareSemanticVersion(major, minor, patch, 0, 69, 0) >= 0, true
}
func compareSemanticVersion(major, minor, patch int, baseMajor, baseMinor, basePatch int) int {
if major != baseMajor {
return major - baseMajor
}
if minor != baseMinor {
return minor - baseMinor
}
return patch - basePatch
}
type wireClientInfo struct { type wireClientInfo struct {
ClientID string `json:"clientID"` ClientID string `json:"clientID"`
WireProtocol string `json:"wireProtocol"` WireProtocol string `json:"wireProtocol"`
+28 -8
View File
@@ -94,11 +94,9 @@ func (f *Framework) RunProcessesWithBinaries(
} }
func (f *Framework) RunFrps(args ...string) (*process.Process, string, error) { func (f *Framework) RunFrps(args ...string) (*process.Process, string, error) {
p := process.NewWithEnvs(TestContext.FRPServerPath, args, f.osEnvs) p, output, err := f.StartFrps(args...)
f.serverProcesses = append(f.serverProcesses, p)
err := p.Start()
if err != nil { if err != nil {
return p, p.Output(), err return p, output, err
} }
select { select {
case <-p.Done(): case <-p.Done():
@@ -107,17 +105,39 @@ func (f *Framework) RunFrps(args ...string) (*process.Process, string, error) {
return p, p.Output(), nil 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) { 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) p := process.NewWithEnvs(TestContext.FRPClientPath, args, f.osEnvs)
f.clientProcesses = append(f.clientProcesses, p) f.clientProcesses = append(f.clientProcesses, p)
err := p.Start() err := p.Start()
if err != nil { if err != nil {
return p, p.Output(), err return p, p.Output(), err
} }
select {
case <-p.Done():
case <-time.After(1500 * time.Millisecond):
}
return p, p.Output(), nil return p, p.Output(), nil
} }
+208
View File
@@ -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
}
+85 -3
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
"time"
"github.com/onsi/ginkgo/v2" "github.com/onsi/ginkgo/v2"
@@ -109,12 +110,12 @@ var _ = ginkgo.Describe("[Feature: WireProtocol]", func() {
name: "default sudp visitor", name: "default sudp visitor",
}, },
{ {
name: "v2 sudp visitor", name: "v2 binary raw sudp visitor",
proxyWireConfig: `transport.wireProtocol = "v2"`, proxyWireConfig: `transport.wireProtocol = "v2"`,
visitorWireConfig: `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"`, proxyWireConfig: `transport.wireProtocol = "v1"`,
visitorWireConfig: `transport.wireProtocol = "v2"`, visitorWireConfig: `transport.wireProtocol = "v2"`,
extraProxyConfig: ` 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"`, proxyWireConfig: `transport.wireProtocol = "v2"`,
visitorWireConfig: `transport.wireProtocol = "v1"`, 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 { type wireClientInfo struct {
ClientID string `json:"clientID"` ClientID string `json:"clientID"`
WireProtocol string `json:"wireProtocol"` WireProtocol string `json:"wireProtocol"`
+300
View File
@@ -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)
}
}
}
+87
View File
@@ -1,18 +1,24 @@
package plugin package plugin
import ( import (
"bufio"
"crypto/tls" "crypto/tls"
"fmt" "fmt"
"io"
"net"
"net/http" "net/http"
"strconv" "strconv"
"strings" "strings"
"github.com/onsi/ginkgo/v2" "github.com/onsi/ginkgo/v2"
pp "github.com/pires/go-proxyproto"
"github.com/fatedier/frp/pkg/transport" "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"
"github.com/fatedier/frp/test/e2e/framework/consts" "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/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/cert"
"github.com/fatedier/frp/test/e2e/pkg/port" "github.com/fatedier/frp/test/e2e/pkg/port"
"github.com/fatedier/frp/test/e2e/pkg/request" "github.com/fatedier/frp/test/e2e/pkg/request"
@@ -450,4 +456,85 @@ var _ = ginkgo.Describe("[Feature: Client-Plugins]", func() {
ExpectResp([]byte("test")). ExpectResp([]byte("test")).
Ensure() 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()
})
}) })
+1
View File
@@ -3,6 +3,7 @@
// @ts-nocheck // @ts-nocheck
// noinspection JSUnusedGlobalSymbols // noinspection JSUnusedGlobalSymbols
// Generated by unplugin-auto-import // Generated by unplugin-auto-import
// biome-ignore lint: disable
export {} export {}
declare global { declare global {
+7 -2
View File
@@ -1,10 +1,14 @@
/* eslint-disable */ /* eslint-disable */
/* prettier-ignore */
// @ts-nocheck // @ts-nocheck
// biome-ignore lint: disable
// oxlint-disable
// ------
// Generated by unplugin-vue-components // Generated by unplugin-vue-components
// Read more: https://github.com/vuejs/core/pull/3399 // Read more: https://github.com/vuejs/core/pull/3399
export {} export {}
/* prettier-ignore */
declare module 'vue' { declare module 'vue' {
export interface GlobalComponents { export interface GlobalComponents {
ConfigField: typeof import('./src/components/ConfigField.vue')['default'] 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'] VisitorBaseSection: typeof import('./src/components/visitor-form/VisitorBaseSection.vue')['default']
VisitorConnectionSection: typeof import('./src/components/visitor-form/VisitorConnectionSection.vue')['default'] VisitorConnectionSection: typeof import('./src/components/visitor-form/VisitorConnectionSection.vue')['default']
VisitorFormLayout: typeof import('./src/components/visitor-form/VisitorFormLayout.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'] VisitorTransportSection: typeof import('./src/components/visitor-form/VisitorTransportSection.vue')['default']
VisitorXtcpSection: typeof import('./src/components/visitor-form/VisitorXtcpSection.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'] vLoading: typeof import('element-plus/es')['ElLoadingDirective']
} }
} }
+14 -13
View File
@@ -9,13 +9,14 @@
"preview": "vite preview", "preview": "vite preview",
"build-only": "vite build", "build-only": "vite build",
"type-check": "vue-tsc --noEmit", "type-check": "vue-tsc --noEmit",
"lint": "eslint --fix" "lint": "eslint . --fix",
"lint:check": "eslint ."
}, },
"dependencies": { "dependencies": {
"element-plus": "^2.13.0", "element-plus": "^2.14.3",
"pinia": "^3.0.4", "pinia": "^3.0.4",
"vue": "^3.5.26", "vue": "^3.5.40",
"vue-router": "^4.6.4" "vue-router": "^5.2.0"
}, },
"devDependencies": { "devDependencies": {
"@types/node": "24", "@types/node": "24",
@@ -23,19 +24,19 @@
"@vue/eslint-config-prettier": "^10.2.0", "@vue/eslint-config-prettier": "^10.2.0",
"@vue/eslint-config-typescript": "^14.7.0", "@vue/eslint-config-typescript": "^14.7.0",
"@vue/tsconfig": "^0.8.1", "@vue/tsconfig": "^0.8.1",
"@vueuse/core": "^14.1.0", "@vueuse/core": "^14.3.0",
"eslint": "^9.39.0", "eslint": "^10.8.0",
"eslint-plugin-vue": "^9.33.0", "eslint-plugin-vue": "^10.10.0",
"npm-run-all": "^4.1.5", "npm-run-all": "^4.1.5",
"prettier": "^3.7.4", "prettier": "^3.9.6",
"sass": "^1.97.2", "sass": "^1.102.0",
"terser": "^5.44.1", "terser": "^5.49.0",
"typescript": "^5.9.3", "typescript": "^5.9.3",
"unplugin-auto-import": "^0.17.5", "unplugin-auto-import": "^21.0.0",
"unplugin-element-plus": "^0.11.2", "unplugin-element-plus": "^0.11.2",
"unplugin-vue-components": "^0.26.0", "unplugin-vue-components": "^32.1.0",
"vite": "^7.3.0", "vite": "^7.3.0",
"vite-svg-loader": "^5.1.0", "vite-svg-loader": "^5.1.0",
"vue-tsc": "^3.2.2" "vue-tsc": "^3.3.8"
} }
} }

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