add support for websocket

This commit is contained in:
2025-07-06 21:07:54 +08:00
parent 55c2f0a888
commit 9e80512948
12 changed files with 296 additions and 103 deletions

View File

@@ -153,6 +153,12 @@ jobs:
- goos: solaris
goarch: amd64
name: solaris-amd64
- goos: js
goarch: wasm
name: js-wasm
- goos: wasip1
goarch: wasm
name: wasip1-wasm
steps:
- uses: actions/checkout@v4

View File

@@ -58,14 +58,16 @@ hosts 使用期望的 host 做 key转发的目的地址为 value。参考`con
}
```
### TCP
### tcp
| 配置 | 类型 | 备注 |
| ------ | ---- | -------- |
| enable | bool | 是否启用 |
| port | int | 端口 |
### KCP
> 默认端口为 25565与 Minecraft 服务端保持一致
### kcp
| 配置 | 类型 | 备注 |
| ------------- | ---- | -------- |
@@ -74,7 +76,7 @@ hosts 使用期望的 host 做 key转发的目的地址为 value。参考`con
| data_shards | int | 数据分片 |
| parity_Shards | int | 校验分片 |
### QUIC
### quic
| 配置 | 类型 | 备注 |
| --------------------- | -------- | ------------ |
@@ -84,3 +86,15 @@ hosts 使用期望的 host 做 key转发的目的地址为 value。参考`con
> application_protocols 只要客户端与服务端有一个能够对应上就可以成功连接
> 默认值为 ["minecraft", "quic", "raw", "h3"]
### websocket
| 配置 | 类型 | 备注 |
| ------ | ---- | -------- |
| enable | bool | 是否启用 |
| port | int | 端口 |
| path | str | 接口路径 |
> path 默认为 "/",会对所有路径的请求进行处理
>
> port 默认为 25566

View File

@@ -22,12 +22,13 @@ var (
type (
Config struct {
Tcp ProtocolConfig `toml:"tcp"`
Quic QuicConfig `toml:"quic"`
Kcp KcpConfig `toml:"kcp"`
Hosts map[string]string `toml:"hosts"`
Log LogConfig `toml:"log"`
PidFile string `toml:"pid_file"`
Tcp ProtocolConfig `toml:"tcp"`
Quic QuicConfig `toml:"quic"`
Kcp KcpConfig `toml:"kcp"`
WebSocket WebSocketConfig `toml:"websocket"`
Hosts map[string]string `toml:"hosts"`
Log LogConfig `toml:"log"`
PidFile string `toml:"pid_file"`
}
ProtocolConfig struct {
@@ -48,12 +49,51 @@ type (
ApplicationProtocols []string `toml:"application_protocols"`
}
WebSocketConfig struct {
Enable bool `toml:"enable"`
Port int `toml:"port"`
Path string `toml:"path"`
}
// LogConfig 定义了日志的配置,包括日志级别和日志文件路径
// 日志级别可以是 "trade", "debug", "info", "warn", "error", "fatal", "disabled"
// 默认日志级别为 "info"
// 日志文件路径指定了日志输出的位置
// 如果日志文件路径为空,则日志将只输出到标准输出
LogConfig struct {
Level string `toml:"level"`
File string `toml:"file"`
}
// serviceConfig 定义了一个服务的配置,包括是否启用和运行函数
// 运行函数接收一个 WaitGroup用于在服务运行时进行同步
// 这样可以确保所有服务在主函数退出前都能正确关闭
// 运行函数通常会在 goroutine 中执行,以便并发处理多个服务
serviceConfig struct {
enable *bool
run func(wg *sync.WaitGroup)
}
)
var services = []serviceConfig{
{
enable: &config.Tcp.Enable,
run: runTcp,
},
{
enable: &config.Kcp.Enable,
run: runKcp,
},
{
enable: &config.Quic.Enable,
run: runQuic,
},
{
enable: &config.WebSocket.Enable,
run: runWebSocket,
},
}
func loadConfig() error {
configLoadLock.Lock()
defer configLoadLock.Unlock()

View File

@@ -3,15 +3,17 @@ package main
import (
"io"
"os"
"os/signal"
"path/filepath"
"syscall"
"github.com/rs/zerolog"
"github.com/rs/zerolog/log"
)
func loadLogger() error {
if len(config.Log.Level) == 0 {
config.Log.Level = "info"
}
level, err := zerolog.ParseLevel(config.Log.Level)
if err != nil {
return err
@@ -37,18 +39,6 @@ func loadLogger() error {
return nil
}
func handleLogRotate() {
signalChan := make(chan os.Signal, 1)
signal.Notify(signalChan, syscall.SIGHUP)
for range signalChan {
log.Info().Msg("Received SIGHUP, reopening log file")
if err := reopenLogFile(); err != nil {
log.Error().Err(err).Msg("Failed to reopen log file")
}
}
}
func reopenLogFile() error {
configLoadLock.Lock()
defer configLoadLock.Unlock()

View File

@@ -0,0 +1,15 @@
// pid_unix.go
//go:build !unix && !plan9
package main
import (
"github.com/rs/zerolog/log"
)
func handleLogRotate() {
// No-op for non-unix platforms
// Log rotation is not supported on this platform
// This function can be left empty or removed if not needed
log.Info().Msg("Log rotation is not supported on this platform")
}

24
cmd/gateway/log_unix.go Normal file
View File

@@ -0,0 +1,24 @@
// pid_unix.go
//go:build unix || plan9
package main
import (
"os"
"os/signal"
"syscall"
"github.com/rs/zerolog/log"
)
func handleLogRotate() {
signalChan := make(chan os.Signal, 1)
signal.Notify(signalChan, syscall.SIGHUP)
for range signalChan {
log.Info().Msg("Received SIGHUP, reopening log file")
if err := reopenLogFile(); err != nil {
log.Error().Err(err).Msg("Failed to reopen log file")
}
}
}

View File

@@ -1,13 +1,12 @@
package main
import (
"fmt"
"io"
"net"
"strings"
"sync"
"time"
"github.com/rs/zerolog"
"github.com/rs/zerolog/log"
"github.com/tursom/mc-gateway/protocol"
)
@@ -30,50 +29,13 @@ func main() {
var wg sync.WaitGroup
defer wg.Wait()
// 启动QUIC服务
if config.Quic.Enable {
wg.Add(1)
go runQuic(&wg)
}
if config.Kcp.Enable {
wg.Add(1)
go runKcp(&wg)
}
// 监听TCP端口
if config.Tcp.Enable {
wg.Add(1)
go runTcp(&wg)
}
}
func runTcp(wg *sync.WaitGroup) {
if wg != nil {
defer wg.Done()
}
listener, err := net.Listen("tcp", fmt.Sprintf(":%d", config.Tcp.Port))
if err != nil {
log.Fatal().Err(err).
Int("port", config.Tcp.Port).
Msg("Failed to listen on port")
}
defer listener.Close()
log.Info().
Int("port", config.Tcp.Port).
Msg("Listening for TCP connections")
for {
// 接受传入的连接
conn, err := listener.Accept()
if err != nil {
log.Err(err).Msg("Error accepting")
for _, service := range services {
if !*service.enable {
continue
}
setSocketOptions(conn)
// 处理连接
go handleRequest(conn)
wg.Add(1)
go service.run(&wg)
}
}
@@ -105,10 +67,10 @@ func handleRequest(conn net.Conn) {
defer client.Close()
var wg sync.WaitGroup
wg.Add(1)
go handleRead(client, conn, &wg)
handleWrite(client, conn, nil)
wg.Add(1)
go copyData(client, conn, &wg)
copyData(conn, client, nil)
// 等待所有读写操作完成
// 不放在 defer 中,以防报错时无法关闭连接
@@ -175,43 +137,27 @@ func mapToHost(conn net.Conn) net.Conn {
return client
}
func upstreamTcp(host string) net.Conn {
conn, err := net.Dial("tcp", host)
if err != nil {
log.Err(err).Str("host", host).Msg("Error dialing upstream")
return nil
}
setSocketOptions(conn)
return conn
func copyData(dst io.Writer, src io.Reader, wg *sync.WaitGroup) {
defer func() {
if r := recover(); r != nil {
var event *zerolog.Event
if err, ok := r.(error); ok {
event = log.Err(err)
} else if str, ok := r.(string); ok {
event = log.Error().Str("panic", str)
} else {
event = log.Error().Any("panic", r)
}
event.Msg("Panic in copyData")
}
}()
}
func handleRead(srv, cli net.Conn, wg *sync.WaitGroup) {
if wg != nil {
defer wg.Done()
}
_, err := io.Copy(srv, cli)
_, err := io.Copy(dst, src)
if err != nil && err != io.EOF {
log.Err(err).Msg("Error copying data")
}
}
func handleWrite(srv, cli net.Conn, wg *sync.WaitGroup) {
if wg != nil {
defer wg.Done()
}
_, err := io.Copy(cli, srv)
if err != nil && err != io.EOF {
log.Err(err).Msg("Error copying data")
}
}
func setSocketOptions(conn net.Conn) {
if tcpConn, ok := conn.(*net.TCPConn); ok {
tcpConn.SetNoDelay(true) // 禁用 Nagle 算法
tcpConn.SetKeepAlive(true)
tcpConn.SetKeepAlivePeriod(30 * time.Second)
}
}

View File

@@ -9,20 +9,20 @@ var currentPidFile string
func writePIDFile() error {
newPidFile := getPidFileFromConfig()
if newPidFile == currentPidFile {
if newPidFile == "" || newPidFile == currentPidFile {
return nil
}
if currentPidFile != "" {
removePIDFile()
}
currentPidFile = newPidFile
pid := os.Getpid()
if err := os.WriteFile(newPidFile, []byte(fmt.Sprintf("%d\n", pid)), 0644); err != nil {
return err
}
currentPidFile = newPidFile
return nil
}

58
cmd/gateway/tcp.go Normal file
View File

@@ -0,0 +1,58 @@
package main
import (
"fmt"
"net"
"sync"
"time"
"github.com/rs/zerolog/log"
)
func runTcp(wg *sync.WaitGroup) {
if wg != nil {
defer wg.Done()
}
listener, err := net.Listen("tcp", fmt.Sprintf(":%d", config.Tcp.Port))
if err != nil {
log.Fatal().Err(err).
Int("port", config.Tcp.Port).
Msg("Failed to listen on port")
}
defer listener.Close()
log.Info().
Int("port", config.Tcp.Port).
Msg("Listening for TCP connections")
for {
// 接受传入的连接
conn, err := listener.Accept()
if err != nil {
log.Err(err).Msg("Error accepting")
continue
}
setSocketOptions(conn)
// 处理连接
go handleRequest(conn)
}
}
func upstreamTcp(host string) net.Conn {
conn, err := net.Dial("tcp", host)
if err != nil {
log.Err(err).Str("host", host).Msg("Error dialing upstream")
return nil
}
setSocketOptions(conn)
return conn
}
func setSocketOptions(conn net.Conn) {
if tcpConn, ok := conn.(*net.TCPConn); ok {
tcpConn.SetNoDelay(true) // 禁用 Nagle 算法
tcpConn.SetKeepAlive(true)
tcpConn.SetKeepAlivePeriod(30 * time.Second)
}
}

97
cmd/gateway/websocket.go Normal file
View File

@@ -0,0 +1,97 @@
package main
import (
"fmt"
"net/http"
"sync"
"time"
"github.com/gorilla/websocket"
"github.com/rs/zerolog/log"
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
// 允许所有来源的连接(生产环境中应该更严格)
return true
},
}
type (
// WebSocket 连接适配器,实现 net.Conn 接口
webSocketConn struct {
*websocket.Conn
messageRemain []byte // 用于存储未处理的消息
}
)
// WebSocket 处理函数
func handleWebSocket(w http.ResponseWriter, r *http.Request) {
// 升级 HTTP 连接为 WebSocket
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
log.Err(err).Msg("Failed to upgrade WebSocket connection")
return
}
defer conn.Close()
handleRequest(&webSocketConn{Conn: conn})
}
func (w *webSocketConn) Read(b []byte) (n int, err error) {
// 如果有未处理的消息,直接从 messageRemain 中读取
if len(w.messageRemain) > 0 {
copied := copy(b, w.messageRemain)
if copied < len(w.messageRemain) {
w.messageRemain = w.messageRemain[copied:]
} else {
w.messageRemain = nil // 清空已处理的消息
}
return copied, nil
}
// 读取新消息
_, message, err := w.ReadMessage()
if err != nil {
return 0, err
}
copied := copy(b, message)
if copied < len(message) {
w.messageRemain = message[copied:]
}
return copied, nil
}
func (w *webSocketConn) Write(b []byte) (n int, err error) {
if err = w.WriteMessage(websocket.BinaryMessage, b); err != nil {
return 0, err
}
return len(b), nil
}
func (w *webSocketConn) SetDeadline(t time.Time) error {
return w.SetReadDeadline(t)
}
// 启动 WebSocket 服务器
func runWebSocket(wg *sync.WaitGroup) {
if wg != nil {
defer wg.Done()
}
path := config.WebSocket.Path
if path == "" {
path = "/" // 默认路径,全部处理
}
http.HandleFunc(path, handleWebSocket)
port := config.WebSocket.Port
if port == 0 {
port = 25566 // 默认端口
}
log.Info().Int("port", port).Str("path", path).Msg("Starting WebSocket server")
if err := http.ListenAndServe(fmt.Sprintf(":%d", port), nil); err != nil {
log.Fatal().Err(err).Msg("Failed to start WebSocket server")
}
}

1
go.mod
View File

@@ -5,6 +5,7 @@ go 1.23.1
require (
github.com/BurntSushi/toml v1.5.0
github.com/fsnotify/fsnotify v1.7.0
github.com/gorilla/websocket v1.5.3
github.com/quic-go/quic-go v0.52.0
github.com/rs/zerolog v1.33.0
github.com/xtaci/kcp-go v5.4.20+incompatible

2
go.sum
View File

@@ -43,6 +43,8 @@ github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/google/pprof v0.0.0-20210407192527-94a9f03dee38 h1:yAJXTCF9TqKcTiHJAE8dj7HMvPfh66eeA2JYW7eFpSE=
github.com/google/pprof v0.0.0-20210407192527-94a9f03dee38/go.mod h1:kpwsk12EmLew5upagYY7GY0pfYCcupk39gWOCRROcvE=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/ianlancetaylor/demangle v0.0.0-20200824232613-28f6c0f3b639/go.mod h1:aSSvb/t6k1mPoxDqO4vJh6VOCGPwU4O0C2/Eqndh1Sc=
github.com/klauspost/cpuid/v2 v2.2.8 h1:+StwCXwm9PdpiEkPyzBXIy+M9KUb4ODm0Zarf1kS5BM=
github.com/klauspost/cpuid/v2 v2.2.8/go.mod h1:Lcz8mBdAVJIBVzewtcLocK12l3Y+JytZYpaMropDUws=