add support for websocket
This commit is contained in:
6
.github/workflows/build.yml
vendored
6
.github/workflows/build.yml
vendored
@@ -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
|
||||
|
||||
20
README.md
20
README.md
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
15
cmd/gateway/log_notunix.go
Normal file
15
cmd/gateway/log_notunix.go
Normal 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
24
cmd/gateway/log_unix.go
Normal 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")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
58
cmd/gateway/tcp.go
Normal 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
97
cmd/gateway/websocket.go
Normal 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
1
go.mod
@@ -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
2
go.sum
@@ -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=
|
||||
|
||||
Reference in New Issue
Block a user