1
0
mirror of https://github.com/MetaCubeX/mihomo.git synced 2026-10-10 04:03:11 +08:00

feat: support mekya for vmess

This commit is contained in:
wwqgtxx
2026-07-07 22:51:06 +08:00
parent de51c4b840
commit e10a1a707f
16 changed files with 1614 additions and 1 deletions
+64 -1
View File
@@ -18,6 +18,7 @@ import (
C "github.com/metacubex/mihomo/constant"
"github.com/metacubex/mihomo/ntp"
"github.com/metacubex/mihomo/transport/gun"
"github.com/metacubex/mihomo/transport/mekya"
"github.com/metacubex/mihomo/transport/mkcp"
mihomoVMess "github.com/metacubex/mihomo/transport/vmess"
@@ -36,7 +37,8 @@ type Vmess struct {
option *VmessOption
// for gun mux
gunClient *gun.Client
gunClient *gun.Client
mekyaClient *mekya.Client
realityConfig *tlsC.RealityConfig
echConfig *ech.Config
@@ -62,6 +64,7 @@ type VmessOption struct {
ECHOpts ECHOptions `proxy:"ech-opts,omitempty"`
RealityOpts RealityOptions `proxy:"reality-opts,omitempty"`
TLSMirrorOpts TLSMirrorOptions `proxy:"tlsmirror-opts,omitempty"`
MekyaOpts MekyaOptions `proxy:"mekya-opts,omitempty"`
MKCPOpts MKCPOptions `proxy:"mkcp-opts,omitempty"`
HTTPOpts HTTPOptions `proxy:"http-opts,omitempty"`
HTTP2Opts HTTP2Options `proxy:"h2-opts,omitempty"`
@@ -101,6 +104,34 @@ func (o MKCPOptions) Build() mkcp.Config {
}
}
type MekyaOptions struct {
URL string `proxy:"url,omitempty"`
H2PoolSize int `proxy:"h2-pool-size,omitempty"`
MaxWriteDelay int `proxy:"max-write-delay,omitempty"`
MaxRequestSize int `proxy:"max-request-size,omitempty"`
PollingIntervalInitial int `proxy:"polling-interval-initial,omitempty"`
MaxWriteSize int `proxy:"max-write-size,omitempty"`
MaxWriteDurationMs int `proxy:"max-write-duration-ms,omitempty"`
MaxSimultaneousWriteConnection int `proxy:"max-simultaneous-write-connection,omitempty"`
PacketWritingBuffer int `proxy:"packet-writing-buffer,omitempty"`
KCP MKCPOptions `proxy:"kcp,omitempty"`
}
func (o MekyaOptions) Build() mekya.Config {
return mekya.Config{
KCP: o.KCP.Build(),
URL: o.URL,
H2PoolSize: o.H2PoolSize,
MaxWriteDelay: o.MaxWriteDelay,
MaxRequestSize: o.MaxRequestSize,
PollingIntervalInitial: o.PollingIntervalInitial,
MaxWriteSize: o.MaxWriteSize,
MaxWriteDurationMs: o.MaxWriteDurationMs,
MaxSimultaneousWriteConnection: o.MaxSimultaneousWriteConnection,
PacketWritingBuffer: o.PacketWritingBuffer,
}
}
type HTTPOptions struct {
Method string `proxy:"method,omitempty"`
Path []string `proxy:"path,omitempty"`
@@ -213,6 +244,8 @@ func (v *Vmess) StreamConnContext(ctx context.Context, c net.Conn, metadata *C.M
c, err = mihomoVMess.StreamH2Conn(ctx, c, h2Opts)
case "grpc":
break // already handle in dialContext
case "mekya":
break // already handle in dialContext
default:
// default tcp network
// handle TLS
@@ -317,6 +350,8 @@ func (v *Vmess) dialContext(ctx context.Context) (c net.Conn, err error) {
switch v.option.Network {
case "grpc": // gun transport
return v.gunClient.Dial()
case "mekya":
return v.mekyaClient.Dial(ctx)
case "mkcp", "kcp":
rawConn, err := v.dialer.DialContext(ctx, "udp", v.addr)
if err != nil {
@@ -385,6 +420,11 @@ func (v *Vmess) Close() error {
errs = append(errs, err)
}
}
if v.mekyaClient != nil {
if err := v.mekyaClient.Close(); err != nil {
errs = append(errs, err)
}
}
return errors.Join(errs...)
}
@@ -452,6 +492,29 @@ func NewVmess(option VmessOption) (*Vmess, error) {
if len(option.HTTP2Opts.Host) == 0 {
option.HTTP2Opts.Host = append(option.HTTP2Opts.Host, "www.example.com")
}
case "mekya":
if len(v.option.ALPN) == 0 {
v.option.ALPN = []string{"h2", "http/1.1"}
}
cfg := option.MekyaOpts.Build()
if cfg.URL == "" {
cfg.URL = "https://" + v.addr
}
v.mekyaClient, err = mekya.NewClient(context.Background(), func(ctx context.Context) (net.Conn, error) {
rawConn, err := v.dialer.DialContext(ctx, "tcp", v.addr)
if err != nil {
return nil, err
}
conn, err := v.streamTLSConn(ctx, rawConn, false)
if err != nil {
_ = rawConn.Close()
return nil, err
}
return conn, nil
}, cfg)
if err != nil {
return nil, err
}
case "grpc":
dialFn := func(ctx context.Context, network, addr string) (net.Conn, error) {
c, err := v.dialer.DialContext(ctx, "tcp", v.addr)
+43
View File
@@ -737,6 +737,32 @@ proxies: # socks5
seed: "" # 启用 AES-GCM 认证时使用的种子,留空使用默认认证
header: "" # 伪装包头,可选:none/srtp/utp/wechat-video/dtls/wireguard
- name: "vmess-mekya"
type: vmess
server: server
port: 443
uuid: uuid
alterId: 32
cipher: auto
network: mekya
tls: true
mekya-opts:
url: https://server:443/mekya
max-write-delay: 80 # 首包后的最大聚合等待时间,单位毫秒
max-request-size: 96000 # 单次 HTTP 请求的最大负载大小,单位字节
polling-interval-initial: 200 # 空轮询间隔,单位毫秒
h2-pool-size: 8 # HTTP/2 连接池大小
kcp:
mtu: 1350 # 最大传输单元
tti: 15 # 传输时间间隔,单位毫秒
uplink-capacity: 40 # 上行容量,单位 MB/s
downlink-capacity: 2000 # 下行容量,单位 MB/s
congestion: false # 是否启用拥塞控制
write-buffer: 67108864 # 写缓冲区大小,单位字节
read-buffer: 67108864 # 读缓冲区大小,单位字节
seed: "" # 启用 AES-GCM 认证时使用的种子,留空使用默认认证
header: "" # 伪装包头,可选:none/srtp/utp/wechat-video/dtls/wireguard
- name: "vmess-h2"
type: vmess
server: server
@@ -1879,6 +1905,23 @@ listeners:
alterId: 1
# ws-path: "/" # 如果不为空则开启 websocket 传输层
# grpc-service-name: "GunService" # 如果不为空则开启 grpc 传输层
# 如果填写 mekya-config 并设置 enable: true 则启用 v2ray 兼容的 Mekya 入站监听(不可与 mkcp/ws/grpc 同时使用)
# mekya-config:
# enable: true
# max-write-size: 10485760 # 单个响应写回的最大负载大小,单位字节
# max-write-duration-ms: 5000 # 单个响应写回的最大持续时间,单位毫秒
# max-simultaneous-write-connection: 128 # 同一会话允许同时等待写回的请求数
# packet-writing-buffer: 65536 # 写包缓冲区大小
# kcp:
# mtu: 1350 # 最大传输单元
# tti: 15 # 传输时间间隔,单位毫秒
# uplink-capacity: 40 # 上行容量,单位 MB/s
# downlink-capacity: 2000 # 下行容量,单位 MB/s
# congestion: false # 是否启用拥塞控制
# write-buffer: 67108864 # 写缓冲区大小,单位字节
# read-buffer: 67108864 # 读缓冲区大小,单位字节
# seed: "" # 启用 AES-GCM 认证时使用的种子,留空使用默认认证
# header: "" # 伪装包头,可选:none/srtp/utp/wechat-video/dtls/wireguard
# 如果填写 mkcp-config 并设置 enable: true 则启用 v2ray 兼容的 mKCP 入站监听
# mkcp-config:
# enable: true
+32
View File
@@ -0,0 +1,32 @@
package config
import "github.com/metacubex/mihomo/transport/mekya"
type MekyaConfig struct {
Enable bool `yaml:"enable" json:"enable,omitempty"`
URL string `yaml:"url" json:"url,omitempty"`
H2PoolSize int `yaml:"h2-pool-size" json:"h2-pool-size,omitempty"`
MaxWriteDelay int `yaml:"max-write-delay" json:"max-write-delay,omitempty"`
MaxRequestSize int `yaml:"max-request-size" json:"max-request-size,omitempty"`
PollingIntervalInitial int `yaml:"polling-interval-initial" json:"polling-interval-initial,omitempty"`
MaxWriteSize int `yaml:"max-write-size" json:"max-write-size,omitempty"`
MaxWriteDurationMs int `yaml:"max-write-duration-ms" json:"max-write-duration-ms,omitempty"`
MaxSimultaneousWriteConnection int `yaml:"max-simultaneous-write-connection" json:"max-simultaneous-write-connection,omitempty"`
PacketWritingBuffer int `yaml:"packet-writing-buffer" json:"packet-writing-buffer,omitempty"`
KCP MKCPConfig `yaml:"kcp" json:"kcp,omitempty"`
}
func (c MekyaConfig) Build() mekya.Config {
return mekya.Config{
KCP: c.KCP.Build(),
URL: c.URL,
H2PoolSize: c.H2PoolSize,
MaxWriteDelay: c.MaxWriteDelay,
MaxRequestSize: c.MaxRequestSize,
PollingIntervalInitial: c.PollingIntervalInitial,
MaxWriteSize: c.MaxWriteSize,
MaxWriteDurationMs: c.MaxWriteDurationMs,
MaxSimultaneousWriteConnection: c.MaxSimultaneousWriteConnection,
PacketWritingBuffer: c.PacketWritingBuffer,
}
}
+1
View File
@@ -27,6 +27,7 @@ type VmessServer struct {
EchKey string
RealityConfig reality.Config
TLSMirrorConfig TLSMirrorConfig `yaml:"tlsmirror-config" json:"tlsmirror-config,omitempty"`
MekyaConfig MekyaConfig `yaml:"mekya-config" json:"mekya-config,omitempty"`
MKCPConfig MKCPConfig `yaml:"mkcp-config" json:"mkcp-config,omitempty"`
MuxOption sing.MuxOption `yaml:"mux-option" json:"mux-option,omitempty"`
}
+33
View File
@@ -0,0 +1,33 @@
package inbound
import LC "github.com/metacubex/mihomo/listener/config"
type MekyaConfig struct {
Enable bool `inbound:"enable,omitempty"`
URL string `inbound:"url,omitempty"`
H2PoolSize int `inbound:"h2-pool-size,omitempty"`
MaxWriteDelay int `inbound:"max-write-delay,omitempty"`
MaxRequestSize int `inbound:"max-request-size,omitempty"`
PollingIntervalInitial int `inbound:"polling-interval-initial,omitempty"`
MaxWriteSize int `inbound:"max-write-size,omitempty"`
MaxWriteDurationMs int `inbound:"max-write-duration-ms,omitempty"`
MaxSimultaneousWriteConnection int `inbound:"max-simultaneous-write-connection,omitempty"`
PacketWritingBuffer int `inbound:"packet-writing-buffer,omitempty"`
KCP MKCPConfig `inbound:"kcp,omitempty"`
}
func (c MekyaConfig) Build() LC.MekyaConfig {
return LC.MekyaConfig{
Enable: c.Enable,
URL: c.URL,
H2PoolSize: c.H2PoolSize,
MaxWriteDelay: c.MaxWriteDelay,
MaxRequestSize: c.MaxRequestSize,
PollingIntervalInitial: c.PollingIntervalInitial,
MaxWriteSize: c.MaxWriteSize,
MaxWriteDurationMs: c.MaxWriteDurationMs,
MaxSimultaneousWriteConnection: c.MaxSimultaneousWriteConnection,
PacketWritingBuffer: c.PacketWritingBuffer,
KCP: c.KCP.Build(),
}
}
+2
View File
@@ -21,6 +21,7 @@ type VmessOption struct {
EchKey string `inbound:"ech-key,omitempty"`
RealityConfig RealityConfig `inbound:"reality-config,omitempty"`
TLSMirrorConfig TLSMirrorConfig `inbound:"tlsmirror-config,omitempty"`
MekyaConfig MekyaConfig `inbound:"mekya-config,omitempty"`
MKCPConfig MKCPConfig `inbound:"mkcp-config,omitempty"`
MuxOption MuxOption `inbound:"mux-option,omitempty"`
}
@@ -71,6 +72,7 @@ func NewVmess(options *VmessOption) (*Vmess, error) {
EchKey: options.EchKey,
RealityConfig: options.RealityConfig.Build(),
TLSMirrorConfig: options.TLSMirrorConfig.Build(),
MekyaConfig: options.MekyaConfig.Build(),
MKCPConfig: options.MKCPConfig.Build(),
MuxOption: options.MuxOption.Build(),
},
@@ -0,0 +1,224 @@
package inbound_test
import (
"context"
"fmt"
"net"
"os"
"path/filepath"
"testing"
"github.com/metacubex/mihomo/adapter/outbound"
"github.com/metacubex/mihomo/listener/inbound"
"github.com/stretchr/testify/require"
)
func TestInboundVMess_Mekya_V2RayInterop(t *testing.T) {
vmessInteropSkip(t)
v2rayBin := vmessInteropV2RayBinary(t)
t.Run("mihomo client to v2ray server", func(t *testing.T) {
echoAddr := startVMessInteropEcho(t)
v2rayPort := vmessInteropReserveTCPPort(t)
certFile, keyFile := mekyaInteropCertificateFiles(t)
config := mekyaInteropServerConfig(t, v2rayPort.Port(), userUUID, certFile, keyFile)
startVMessInteropV2Ray(t, v2rayBin, config, v2rayPort.Release, net.JoinHostPort("127.0.0.1", fmt.Sprint(v2rayPort.Port())))
out, err := outbound.NewVmess(outbound.VmessOption{
Name: "vmess_mekya_v2ray_server",
Server: "127.0.0.1",
Port: v2rayPort.Port(),
UUID: userUUID,
Cipher: "auto",
Network: "mekya",
TLS: true,
Fingerprint: tlsFingerprint,
MekyaOpts: mekyaInteropOutboundOptions(net.JoinHostPort("127.0.0.1", fmt.Sprint(v2rayPort.Port()))),
})
require.NoError(t, err)
t.Cleanup(func() { _ = out.Close() })
vmessInteropRoundTripWithRetry(t, func() (net.Conn, error) {
return out.DialContext(context.Background(), vmessInteropMetadata(t, echoAddr))
}, 64*1024)
})
t.Run("v2ray client to mihomo server", func(t *testing.T) {
echoAddr := startVMessInteropEcho(t)
v2rayPort := vmessInteropReserveTCPPort(t)
in, err := inbound.NewVmess(&inbound.VmessOption{
BaseOption: inbound.BaseOption{
NameStr: "vmess_mekya_v2ray_client",
Listen: "127.0.0.1",
Port: "0",
},
Users: []inbound.VmessUser{
{Username: "test", UUID: userUUID},
},
Certificate: tlsCertificate,
PrivateKey: tlsPrivateKey,
MekyaConfig: mekyaInteropInboundConfig(),
})
require.NoError(t, err)
tunnel := vmessInteropDirectTunnel(t)
require.NoError(t, in.Listen(tunnel))
t.Cleanup(func() { _ = in.Close() })
inboundPort := vmessInteropParsePort(t, vmessInteropPort(in.Address()))
config := mekyaInteropClientConfig(t, v2rayPort.Port(), inboundPort, vmessInteropPort(echoAddr), userUUID)
startVMessInteropV2Ray(t, v2rayBin, config, v2rayPort.Release, net.JoinHostPort("127.0.0.1", fmt.Sprint(v2rayPort.Port())))
vmessInteropRoundTripWithRetry(t, func() (net.Conn, error) {
return net.Dial("tcp", net.JoinHostPort("127.0.0.1", fmt.Sprint(v2rayPort.Port())))
}, 64*1024)
})
}
func mekyaInteropCertificateFiles(t *testing.T) (string, string) {
t.Helper()
dir := t.TempDir()
certFile := filepath.Join(dir, "cert.pem")
keyFile := filepath.Join(dir, "key.pem")
require.NoError(t, os.WriteFile(certFile, []byte(tlsCertificate), 0o600))
require.NoError(t, os.WriteFile(keyFile, []byte(tlsPrivateKey), 0o600))
return certFile, keyFile
}
func mekyaInteropInboundConfig() inbound.MekyaConfig {
return inbound.MekyaConfig{
Enable: true,
MaxWriteSize: 10 * 1024 * 1024,
MaxWriteDurationMs: 500,
MaxSimultaneousWriteConnection: 128,
PacketWritingBuffer: 65536,
KCP: inbound.MKCPConfig{
MTU: 1350,
TTI: 15,
UplinkCapacity: 40,
DownlinkCapacity: 2000,
WriteBuffer: 64 * 1024 * 1024,
ReadBuffer: 64 * 1024 * 1024,
},
}
}
func mekyaInteropOutboundOptions(addr string) outbound.MekyaOptions {
return outbound.MekyaOptions{
URL: "https://" + addr + "/mekya",
MaxWriteDelay: 80,
MaxRequestSize: 96000,
PollingIntervalInitial: 200,
H2PoolSize: 8,
KCP: outbound.MKCPOptions{
MTU: 1350,
TTI: 15,
UplinkCapacity: 40,
DownlinkCapacity: 2000,
WriteBuffer: 64 * 1024 * 1024,
ReadBuffer: 64 * 1024 * 1024,
},
}
}
func mekyaInteropServerConfig(t *testing.T, listenPort int, userID, certFile, keyFile string) []byte {
t.Helper()
config := vmessInteropBaseConfig()
config["inbounds"] = []any{map[string]any{
"protocol": "vmess",
"listen": "127.0.0.1",
"port": listenPort,
"settings": map[string]any{
"users": []string{userID},
},
"streamSettings": mekyaInteropStreamConfig(
"http://127.0.0.1:"+fmt.Sprint(listenPort),
mekyaInteropServerSecuritySettings(certFile, keyFile),
),
}}
config["outbounds"] = []any{vmessInteropDirectOutbound()}
return vmessInteropMarshalJSONConfig(t, config)
}
func mekyaInteropClientConfig(t *testing.T, listenPort, serverPort int, targetPort string, userID string) []byte {
t.Helper()
targetPortValue := vmessInteropParsePort(t, targetPort)
config := vmessInteropBaseConfig()
config["inbounds"] = []any{map[string]any{
"protocol": "dokodemo-door",
"listen": "127.0.0.1",
"port": listenPort,
"settings": map[string]any{
"address": "127.0.0.1",
"port": targetPortValue,
"networks": "tcp",
},
}}
config["outbounds"] = []any{
map[string]any{
"protocol": "vmess",
"streamSettings": mekyaInteropStreamConfig(
"https://127.0.0.1:"+fmt.Sprint(serverPort)+"/mekya",
mekyaInteropClientSecuritySettings(),
),
"settings": map[string]any{
"address": "127.0.0.1",
"port": serverPort,
"uuid": userID,
},
},
}
return vmessInteropMarshalJSONConfig(t, config)
}
func mekyaInteropStreamConfig(url string, securitySettings map[string]any) map[string]any {
return map[string]any{
"transport": "mekya",
"transportSettings": mekyaInteropTransportSettings(url),
"security": "tls",
"securitySettings": securitySettings,
}
}
func mekyaInteropTransportSettings(url string) map[string]any {
return map[string]any{
"url": url,
"maxWriteDelay": 80,
"maxRequestSize": 96000,
"pollingIntervalInitial": 200,
"h2_pool_size": 8,
"maxWriteSize": 10 * 1024 * 1024,
"maxWriteDurationMs": 500,
"maxSimultaneousWriteConnection": 128,
"packetWritingBuffer": 65536,
"kcp": map[string]any{
"mtu": map[string]any{"value": 1350},
"tti": map[string]any{"value": 15},
"uplink_capacity": map[string]any{"value": 40},
"downlink_capacity": map[string]any{"value": 2000},
"congestion": false,
"write_buffer": map[string]any{"size": 64 * 1024 * 1024},
"read_buffer": map[string]any{"size": 64 * 1024 * 1024},
},
}
}
func mekyaInteropServerSecuritySettings(certFile, keyFile string) map[string]any {
return map[string]any{
"certificate": []any{map[string]any{
"usage": "ENCIPHERMENT",
"certificateFile": certFile,
"keyFile": keyFile,
}},
}
}
func mekyaInteropClientSecuritySettings() map[string]any {
return map[string]any{
"pinnedPeerCertificateChainSha256": []string{tlsMirrorInteropCertChainHash([]byte(tlsCertificate))},
"allowInsecureIfPinnedPeerCertificate": true,
}
}
+30
View File
@@ -64,6 +64,9 @@ func testInboundVMess(t *testing.T, inboundOptions inbound.VmessOption, outbound
if outboundOptions.Network == "mkcp" { // don't test sing-mux over mkcp
return
}
if outboundOptions.Network == "mekya" { // don't test sing-mux over mekya
return
}
if outboundOptions.TLSMirrorOpts.PrimaryKey != "" { // don't test sing-mux over tlsmirror
return
}
@@ -188,6 +191,33 @@ func TestInboundVMess_MKCP(t *testing.T) {
}
}
func TestInboundVMess_Mekya(t *testing.T) {
inboundOptions := inbound.VmessOption{
MekyaConfig: inbound.MekyaConfig{
Enable: true,
MaxWriteSize: 1 << 20,
MaxWriteDurationMs: 100,
MaxSimultaneousWriteConnection: 16,
PacketWritingBuffer: 1024,
KCP: inbound.MKCPConfig{
TTI: 15,
},
},
}
outboundOptions := outbound.VmessOption{
Network: "mekya",
MekyaOpts: outbound.MekyaOptions{
MaxWriteDelay: 20,
MaxRequestSize: 96000,
PollingIntervalInitial: 20,
KCP: outbound.MKCPOptions{
TTI: 15,
},
},
}
testInboundVMess(t, inboundOptions, outboundOptions)
}
func TestInboundVMess_TLSMirror(t *testing.T) {
inboundOptions := inbound.VmessOption{
TLSMirrorConfig: inbound.TLSMirrorConfig{
+24
View File
@@ -17,6 +17,7 @@ import (
"github.com/metacubex/mihomo/listener/tlsmirror"
"github.com/metacubex/mihomo/ntp"
"github.com/metacubex/mihomo/transport/gun"
"github.com/metacubex/mihomo/transport/mekya"
"github.com/metacubex/mihomo/transport/mkcp"
mihomoVMess "github.com/metacubex/mihomo/transport/vmess"
@@ -26,6 +27,7 @@ import (
"github.com/metacubex/sing/common"
"github.com/metacubex/sing/common/metadata"
"github.com/metacubex/tls"
"golang.org/x/exp/slices"
)
type Listener struct {
@@ -56,6 +58,14 @@ func New(config LC.VmessServer, lc C.InboundListenConfig, tunnel C.Tunnel, addit
if err != nil {
return nil, err
}
if config.MekyaConfig.Enable {
if config.MKCPConfig.Enable {
return nil, errors.New("mkcp-config is unavailable in mekya")
}
if config.WsPath != "" || config.GrpcServiceName != "" {
return nil, errors.New("ws and grpc are unavailable in mekya")
}
}
service := vmess.NewService[string](h, vmess.ServiceWithDisableHeaderProtection(), vmess.ServiceWithTimeFunc(ntp.Now))
err = service.UpdateUsers(
@@ -174,6 +184,14 @@ func New(config LC.VmessServer, lc C.InboundListenConfig, tunnel C.Tunnel, addit
httpServer.Protocols.SetUnencryptedHTTP2(true)
tlsConfig.NextProtos = append([]string{"h2"}, tlsConfig.NextProtos...) // h2 must before http/1.1
}
if config.MekyaConfig.Enable {
if !slices.Contains(tlsConfig.NextProtos, "http/1.1") {
tlsConfig.NextProtos = append([]string{"http/1.1"}, tlsConfig.NextProtos...)
}
if !slices.Contains(tlsConfig.NextProtos, "h2") {
tlsConfig.NextProtos = append([]string{"h2"}, tlsConfig.NextProtos...)
}
}
for _, addr := range strings.Split(config.Listen, ",") {
addr := addr
@@ -203,6 +221,12 @@ func New(config LC.VmessServer, lc C.InboundListenConfig, tunnel C.Tunnel, addit
} else if tlsConfig.GetCertificate != nil {
l = tls.NewListener(l, tlsConfig)
}
if config.MekyaConfig.Enable {
l, err = mekya.Listen(context.Background(), l, config.MekyaConfig.Build())
if err != nil {
return nil, err
}
}
sl.listeners = append(sl.listeners, l)
go func() {
+34
View File
@@ -0,0 +1,34 @@
package mekya
import (
"encoding/binary"
"io"
)
const packetBundleOverhead = 2
func writePacketBundle(w io.Writer, packet []byte) error {
if len(packet) > 0xffff {
return io.ErrShortBuffer
}
var header [packetBundleOverhead]byte
binary.BigEndian.PutUint16(header[:], uint16(len(packet)))
if _, err := w.Write(header[:]); err != nil {
return err
}
_, err := w.Write(packet)
return err
}
func readPacketBundle(r io.Reader) ([]byte, error) {
var header [packetBundleOverhead]byte
if _, err := io.ReadFull(r, header[:]); err != nil {
return nil, err
}
length := binary.BigEndian.Uint16(header[:])
packet := make([]byte, length)
if _, err := io.ReadFull(r, packet); err != nil {
return nil, err
}
return packet, nil
}
+540
View File
@@ -0,0 +1,540 @@
package mekya
import (
"bytes"
"context"
"crypto/rand"
"encoding/base64"
"errors"
"fmt"
"net"
"net/url"
"os"
"sync"
"time"
"github.com/metacubex/mihomo/common/httputils"
tlsC "github.com/metacubex/mihomo/component/tls"
"github.com/metacubex/mihomo/transport/mkcp"
"github.com/metacubex/http"
"github.com/metacubex/http/http2"
"github.com/metacubex/http/httptrace"
"github.com/metacubex/tls"
)
type DialFunc func(ctx context.Context) (net.Conn, error)
type Client struct {
ctx context.Context
cancel context.CancelFunc
cfg Config
url string
rt http.RoundTripper
once sync.Once
}
func NewClient(ctx context.Context, dial DialFunc, cfg Config) (*Client, error) {
ctx, cancel := context.WithCancel(ctx)
roundTripURL, err := normalizeURL(cfg.URL)
if err != nil {
cancel()
return nil, err
}
c := &Client{
ctx: ctx,
cancel: cancel,
cfg: cfg,
url: roundTripURL,
rt: newRoundTripper(dial, cfg.H2PoolSize),
}
return c, nil
}
func (c *Client) Dial(ctx context.Context) (net.Conn, error) {
if err := c.ctx.Err(); err != nil {
return nil, err
}
if err := ctx.Err(); err != nil {
return nil, err
}
raw, err := c.newSession()
if err != nil {
return nil, err
}
conn, err := mkcp.Dial(ctx, raw, c.cfg.KCP)
if err != nil {
_ = raw.Close()
return nil, err
}
return &clientConn{Conn: conn, raw: raw}, nil
}
func (c *Client) Close() error {
c.once.Do(func() {
c.cancel()
httputils.CloseTransport(c.rt)
})
return nil
}
func newRoundTripper(dial DialFunc, h2PoolSize int) http.RoundTripper {
rt := &alpnAwareRoundTripper{
dial: dial,
connectWithH1: make(map[string]bool),
pendingConn: make(map[pendingConnKey]*pendingConn),
}
rt.h1 = &http.Transport{
DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return rt.dialOrGetTLSWithExpectedALPN(ctx, addr, false)
},
DisableCompression: true,
}
rt.h2 = newH2RoundTripper(h2PoolSize, func(ctx context.Context, network, addr string, cfg *tls.Config) (net.Conn, error) {
return rt.dialOrGetTLSWithExpectedALPN(ctx, addr, true)
})
return rt
}
func newH2RoundTripper(h2PoolSize int, dialTLSContext func(context.Context, string, string, *tls.Config) (net.Conn, error)) http.RoundTripper {
newTransport := func() http.RoundTripper {
return &http2.Transport{
DialTLSContext: dialTLSContext,
DisableCompression: true,
}
}
if h2PoolSize >= 2 {
pool := &roundTripperPool{roundTrippers: make([]http.RoundTripper, h2PoolSize)}
for i := range pool.roundTrippers {
pool.roundTrippers[i] = newTransport()
}
return pool
}
return newTransport()
}
var (
errUnexpectedALPN = errors.New("mekya: incorrect ALPN negotiated, try again")
errUnexpectedALPNTooMany = errors.New("mekya: incorrect ALPN negotiated")
)
type alpnAwareRoundTripper struct {
mu sync.Mutex
connectWithH1 map[string]bool
pendingConn map[pendingConnKey]*pendingConn
dial DialFunc
h1 http.RoundTripper
h2 http.RoundTripper
}
type pendingConnKey struct {
addr string
h2 bool
}
func (r *alpnAwareRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
if req.URL.Scheme != "https" {
return nil, fmt.Errorf("mekya: unsupported url scheme %q", req.URL.Scheme)
}
addr := roundTripAddr(req)
for retry := 0; retry < 5; retry++ {
var rt http.RoundTripper
if r.shouldConnectWithH1(addr) {
rt = r.h1
} else {
rt = r.h2
}
resp, err := rt.RoundTrip(req)
if errors.Is(err, errUnexpectedALPN) {
continue
}
return resp, err
}
return nil, errUnexpectedALPNTooMany
}
func (r *alpnAwareRoundTripper) Close() error {
r.mu.Lock()
pending := r.pendingConn
r.pendingConn = nil
r.mu.Unlock()
for _, conn := range pending {
conn.close()
}
httputils.CloseTransport(r.h1)
httputils.CloseTransport(r.h2)
return nil
}
func (r *alpnAwareRoundTripper) shouldConnectWithH1(addr string) bool {
r.mu.Lock()
defer r.mu.Unlock()
return r.connectWithH1[addr]
}
func (r *alpnAwareRoundTripper) dialOrGetTLSWithExpectedALPN(ctx context.Context, addr string, expectedH2 bool) (net.Conn, error) {
r.mu.Lock()
defer r.mu.Unlock()
if r.connectWithH1[addr] == expectedH2 {
return nil, errUnexpectedALPN
}
if conn := r.getPendingConnLocked(addr, expectedH2); conn != nil {
return conn, nil
}
conn, err := r.dial(ctx)
if err != nil {
return nil, err
}
protocolIsH2 := tlsC.GetTLSConnectionState(conn).NegotiatedProtocol == http2.NextProtoTLS
if protocolIsH2 == expectedH2 {
return conn, nil
}
r.putPendingConnLocked(addr, protocolIsH2, conn)
r.connectWithH1[addr] = !protocolIsH2
return nil, errUnexpectedALPN
}
func (r *alpnAwareRoundTripper) getPendingConnLocked(addr string, h2 bool) net.Conn {
if r.pendingConn == nil {
return nil
}
key := pendingConnKey{addr: addr, h2: h2}
pending := r.pendingConn[key]
if pending == nil {
return nil
}
delete(r.pendingConn, key)
return pending.claim()
}
func (r *alpnAwareRoundTripper) putPendingConnLocked(addr string, h2 bool, conn net.Conn) {
if r.pendingConn == nil {
_ = conn.Close()
return
}
key := pendingConnKey{addr: addr, h2: h2}
if old := r.pendingConn[key]; old != nil {
old.close()
}
r.pendingConn[key] = newPendingConn(conn)
}
func roundTripAddr(req *http.Request) string {
port := req.URL.Port()
if port == "" {
port = "443"
}
return net.JoinHostPort(req.URL.Hostname(), port)
}
type pendingConn struct {
conn net.Conn
timer *time.Timer
mu sync.Mutex
claimed bool
}
func newPendingConn(conn net.Conn) *pendingConn {
p := &pendingConn{conn: conn}
p.timer = time.AfterFunc(time.Minute, p.close)
return p
}
func (p *pendingConn) claim() net.Conn {
p.mu.Lock()
defer p.mu.Unlock()
if p.claimed {
return nil
}
p.claimed = true
if p.timer != nil {
p.timer.Stop()
}
return p.conn
}
func (p *pendingConn) close() {
p.mu.Lock()
defer p.mu.Unlock()
if p.claimed {
return
}
p.claimed = true
if p.timer != nil {
p.timer.Stop()
}
_ = p.conn.Close()
}
type roundTripperPool struct {
mu sync.Mutex
next int
roundTrippers []http.RoundTripper
}
func (p *roundTripperPool) RoundTrip(req *http.Request) (*http.Response, error) {
p.mu.Lock()
rt := p.roundTrippers[p.next]
p.next = (p.next + 1) % len(p.roundTrippers)
p.mu.Unlock()
return rt.RoundTrip(req)
}
func (p *roundTripperPool) Close() error {
for _, rt := range p.roundTrippers {
httputils.CloseTransport(rt)
}
return nil
}
func (c *Client) newSession() (*requestClientSession, error) {
sessionID := make([]byte, 16)
if _, err := rand.Read(sessionID); err != nil {
return nil, err
}
ctx, cancel := context.WithCancel(c.ctx)
session := &requestClientSession{
ctx: ctx,
cancel: cancel,
sessionID: sessionID,
url: c.url,
currentPollingInterval: c.cfg.PollingIntervalInitial,
maxRequestSize: c.cfg.MaxRequestSize,
maxWriteDelay: c.cfg.MaxWriteDelay,
writerChan: make(chan []byte, 256),
readerChan: make(chan []byte, 256),
deadlines: newPipeDeadlines(),
rt: c.rt,
}
go session.keepRunning()
return session, nil
}
type clientConn struct {
net.Conn
once sync.Once
raw *requestClientSession
}
func (c *clientConn) Close() error {
c.once.Do(func() {
_ = c.raw.Close()
})
return c.Conn.Close()
}
func (c *clientConn) LocalAddr() net.Addr {
if addr := c.raw.LocalAddr(); addr != nil {
return addr
}
return c.Conn.LocalAddr()
}
func (c *clientConn) RemoteAddr() net.Addr {
if addr := c.raw.RemoteAddr(); addr != nil {
return addr
}
return c.Conn.RemoteAddr()
}
type requestClientSession struct {
ctx context.Context
cancel context.CancelFunc
sessionID []byte
rt http.RoundTripper
url string
currentPollingInterval int
maxRequestSize int
maxWriteDelay int
writerChan chan []byte
readerChan chan []byte
nextWrite []byte
deadlines pipeDeadlines
addrMu sync.RWMutex
localAddr net.Addr
remoteAddr net.Addr
}
func (s *requestClientSession) keepRunning() {
for s.ctx.Err() == nil {
s.runOnce()
}
}
func (s *requestClientSession) runOnce() {
requestBody := bytes.NewBuffer(nil)
waitTimer := time.NewTimer(time.Duration(s.currentPollingInterval) * time.Millisecond)
defer waitTimer.Stop()
seenPacket := false
if s.nextWrite != nil {
seenPacket = true
if !waitTimer.Stop() {
select {
case <-waitTimer.C:
default:
}
}
waitTimer.Reset(time.Duration(s.maxWriteDelay) * time.Millisecond)
if !s.writePacket(requestBody, s.nextWrite) {
return
}
s.nextWrite = nil
}
copyFromChan:
for {
select {
case <-s.ctx.Done():
return
case <-waitTimer.C:
break copyFromChan
case packet := <-s.writerChan:
if !seenPacket {
seenPacket = true
if !waitTimer.Stop() {
select {
case <-waitTimer.C:
default:
}
}
waitTimer.Reset(time.Duration(s.maxWriteDelay) * time.Millisecond)
}
if !s.writePacket(requestBody, packet) {
break copyFromChan
}
}
}
go s.roundTrip(requestBody.Bytes())
}
func (s *requestClientSession) writePacket(requestBody *bytes.Buffer, packet []byte) bool {
sizeOffset := packetBundleOverhead + len(packet)
if s.maxRequestSize > 0 && requestBody.Len()+sizeOffset > s.maxRequestSize {
s.nextWrite = packet
return false
}
if err := writePacketBundle(requestBody, packet); err != nil {
return false
}
return true
}
func (s *requestClientSession) roundTrip(body []byte) {
trace := &httptrace.ClientTrace{
GotConn: func(info httptrace.GotConnInfo) {
if info.Conn != nil {
s.setUnderlyingAddr(info.Conn)
}
},
}
ctx := httptrace.WithClientTrace(s.ctx, trace)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, s.url, bytes.NewReader(body))
if err != nil {
return
}
req.Header.Set("X-Session-ID", base64.RawURLEncoding.EncodeToString(s.sessionID))
resp, err := s.rt.RoundTrip(req)
if err != nil {
return
}
defer resp.Body.Close()
for {
packet, err := readPacketBundle(resp.Body)
if err != nil {
return
}
select {
case <-s.ctx.Done():
return
case s.readerChan <- packet:
}
}
}
func (s *requestClientSession) Read(p []byte) (int, error) {
select {
case <-s.ctx.Done():
return 0, s.ctx.Err()
case <-s.deadlines.read.Wait():
return 0, os.ErrDeadlineExceeded
case packet := <-s.readerChan:
return copy(p, packet), nil
}
}
func (s *requestClientSession) Write(p []byte) (int, error) {
packet := append([]byte(nil), p...)
select {
case <-s.ctx.Done():
return 0, s.ctx.Err()
case <-s.deadlines.write.Wait():
return 0, os.ErrDeadlineExceeded
case s.writerChan <- packet:
return len(p), nil
}
}
func (s *requestClientSession) Close() error {
s.cancel()
return nil
}
func (s *requestClientSession) setUnderlyingAddr(conn net.Conn) {
s.addrMu.Lock()
s.localAddr = conn.LocalAddr()
s.remoteAddr = conn.RemoteAddr()
s.addrMu.Unlock()
}
func (s *requestClientSession) LocalAddr() net.Addr {
s.addrMu.RLock()
defer s.addrMu.RUnlock()
return s.localAddr
}
func (s *requestClientSession) RemoteAddr() net.Addr {
s.addrMu.RLock()
defer s.addrMu.RUnlock()
return s.remoteAddr
}
func (s *requestClientSession) SetDeadline(t time.Time) error {
return s.deadlines.SetDeadline(t)
}
func (s *requestClientSession) SetReadDeadline(t time.Time) error {
return s.deadlines.SetReadDeadline(t)
}
func (s *requestClientSession) SetWriteDeadline(t time.Time) error {
return s.deadlines.SetWriteDeadline(t)
}
var _ net.Conn = (*requestClientSession)(nil)
func normalizeURL(raw string) (string, error) {
if raw == "" {
return "", fmt.Errorf("mekya: empty url")
}
u, err := url.Parse(raw)
if err != nil {
return "", err
}
if u.Scheme == "" {
u.Scheme = "https"
}
if u.Scheme != "https" {
return "", fmt.Errorf("mekya: unsupported url scheme %q", u.Scheme)
}
if u.Host == "" {
return "", fmt.Errorf("mekya: empty url host")
}
return u.String(), nil
}
+16
View File
@@ -0,0 +1,16 @@
package mekya
import "github.com/metacubex/mihomo/transport/mkcp"
type Config struct {
KCP mkcp.Config
URL string
H2PoolSize int
MaxWriteDelay int
MaxRequestSize int
PollingIntervalInitial int
MaxWriteSize int
MaxWriteDurationMs int
MaxSimultaneousWriteConnection int
PacketWritingBuffer int
}
+35
View File
@@ -0,0 +1,35 @@
package mekya
import (
"time"
"github.com/metacubex/mihomo/common/net/deadline"
)
type pipeDeadlines struct {
read deadline.PipeDeadline
write deadline.PipeDeadline
}
func newPipeDeadlines() pipeDeadlines {
return pipeDeadlines{
read: deadline.MakePipeDeadline(),
write: deadline.MakePipeDeadline(),
}
}
func (d *pipeDeadlines) SetDeadline(t time.Time) error {
d.read.Set(t)
d.write.Set(t)
return nil
}
func (d *pipeDeadlines) SetReadDeadline(t time.Time) error {
d.read.Set(t)
return nil
}
func (d *pipeDeadlines) SetWriteDeadline(t time.Time) error {
d.write.Set(t)
return nil
}
+82
View File
@@ -0,0 +1,82 @@
package mekya
import (
"bytes"
"context"
"io"
"net"
"testing"
"time"
"github.com/metacubex/mihomo/transport/mkcp"
"github.com/stretchr/testify/require"
)
func TestRoundTrip(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
cfg := testConfig()
server, err := Listen(ctx, ln, cfg)
require.NoError(t, err)
defer server.Close()
serverErr := make(chan error, 1)
go func() {
conn, err := server.Accept()
if err != nil {
serverErr <- err
return
}
defer conn.Close()
serverErr <- echo(conn)
}()
client, err := NewClient(ctx, func(ctx context.Context) (net.Conn, error) {
var d net.Dialer
return d.DialContext(ctx, "tcp", server.Addr().String())
}, cfg)
require.NoError(t, err)
defer client.Close()
conn, err := client.Dial(ctx)
require.NoError(t, err)
defer conn.Close()
require.NoError(t, conn.SetDeadline(time.Now().Add(5*time.Second)))
payload := bytes.Repeat([]byte("m"), 64*1024)
_, err = conn.Write(payload)
require.NoError(t, err)
got := make([]byte, len(payload))
_, err = io.ReadFull(conn, got)
require.NoError(t, err)
require.Equal(t, payload, got)
require.NotNil(t, conn.LocalAddr())
require.NotNil(t, conn.RemoteAddr())
require.NoError(t, conn.Close())
}
func testConfig() Config {
return Config{
KCP: mkcp.Config{
TTI: 15,
},
URL: "https://example.invalid/mekya",
H2PoolSize: 2,
MaxWriteDelay: 20,
MaxRequestSize: 96000,
PollingIntervalInitial: 20,
MaxWriteSize: 1 << 20,
MaxWriteDurationMs: 100,
MaxSimultaneousWriteConnection: 16,
PacketWritingBuffer: 1024,
}
}
func echo(conn net.Conn) error {
_, err := io.Copy(conn, conn)
return err
}
+139
View File
@@ -0,0 +1,139 @@
package mekya
import (
"context"
"io"
"net"
"os"
"sync"
"time"
)
type packet struct {
addr net.Addr
data []byte
}
type wrappedPacketConn struct {
ctx context.Context
cancel context.CancelFunc
mu sync.Mutex
sessions map[string]*serverSession
readChan chan packet
local net.Addr
deadlines pipeDeadlines
}
func newWrappedPacketConn(ctx context.Context, local net.Addr) *wrappedPacketConn {
ctx, cancel := context.WithCancel(ctx)
return &wrappedPacketConn{
ctx: ctx,
cancel: cancel,
sessions: make(map[string]*serverSession),
readChan: make(chan packet, 16),
local: local,
deadlines: newPipeDeadlines(),
}
}
func (c *wrappedPacketConn) addSession(session *serverSession) error {
select {
case <-c.ctx.Done():
return net.ErrClosed
default:
}
c.mu.Lock()
if c.sessions == nil {
c.mu.Unlock()
return net.ErrClosed
}
c.sessions[string(session.sessionID)] = session
c.mu.Unlock()
go c.readSession(session)
return nil
}
func (c *wrappedPacketConn) readSession(session *serverSession) {
buf := make([]byte, 2000)
for {
n, err := session.Read(buf)
if err != nil || n > len(buf) {
return
}
payload := append([]byte(nil), buf[:n]...)
select {
case <-c.ctx.Done():
return
case c.readChan <- packet{addr: session, data: payload}:
}
}
}
func (c *wrappedPacketConn) ReadFrom(p []byte) (int, net.Addr, error) {
select {
case <-c.ctx.Done():
return 0, nil, c.ctx.Err()
case <-c.deadlines.read.Wait():
return 0, nil, os.ErrDeadlineExceeded
case packet := <-c.readChan:
n := copy(p, packet.data)
if n < len(packet.data) {
return n, packet.addr, io.ErrShortBuffer
}
return n, packet.addr, nil
}
}
func (c *wrappedPacketConn) WriteTo(p []byte, addr net.Addr) (int, error) {
select {
case <-c.ctx.Done():
return 0, c.ctx.Err()
case <-c.deadlines.write.Wait():
return 0, os.ErrDeadlineExceeded
default:
}
session, ok := addr.(*serverSession)
if !ok {
return 0, net.ErrClosed
}
c.mu.Lock()
session = c.sessions[string(session.sessionID)]
c.mu.Unlock()
if session == nil {
return 0, net.ErrClosed
}
return session.Write(p)
}
func (c *wrappedPacketConn) Close() error {
c.cancel()
c.mu.Lock()
sessions := make([]*serverSession, 0, len(c.sessions))
for _, session := range c.sessions {
sessions = append(sessions, session)
}
c.sessions = nil
c.mu.Unlock()
for _, session := range sessions {
_ = session.Close()
}
return nil
}
func (c *wrappedPacketConn) LocalAddr() net.Addr {
return c.local
}
func (c *wrappedPacketConn) SetDeadline(t time.Time) error {
return c.deadlines.SetDeadline(t)
}
func (c *wrappedPacketConn) SetReadDeadline(t time.Time) error {
return c.deadlines.SetReadDeadline(t)
}
func (c *wrappedPacketConn) SetWriteDeadline(t time.Time) error {
return c.deadlines.SetWriteDeadline(t)
}
var _ net.PacketConn = (*wrappedPacketConn)(nil)
+315
View File
@@ -0,0 +1,315 @@
package mekya
import (
"bytes"
"context"
"encoding/base64"
"errors"
"io"
"net"
"sync"
"time"
"github.com/metacubex/mihomo/transport/mkcp"
"github.com/metacubex/http"
)
type Listener struct {
outer net.Listener
packetConn *wrappedPacketConn
mkcp *mkcp.Listener
server *http.Server
done chan struct{}
once sync.Once
}
func Listen(ctx context.Context, ln net.Listener, cfg Config) (*Listener, error) {
packetConn := newWrappedPacketConn(ctx, ln.Addr())
handler := newServer(ctx, cfg, packetConn)
mkcpListener, err := mkcp.Listen(ctx, packetConn, cfg.KCP)
if err != nil {
return nil, err
}
protocols := new(http.Protocols)
protocols.SetHTTP1(true)
protocols.SetHTTP2(true)
protocols.SetUnencryptedHTTP2(true)
server := &http.Server{
Handler: handler,
Protocols: protocols,
ReadHeaderTimeout: 240 * time.Second,
ReadTimeout: 240 * time.Second,
WriteTimeout: 240 * time.Second,
IdleTimeout: 240 * time.Second,
}
l := &Listener{
outer: ln,
packetConn: packetConn,
mkcp: mkcpListener,
server: server,
done: make(chan struct{}),
}
go func() {
defer close(l.done)
err := server.Serve(ln)
if err != nil && !errors.Is(err, http.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
_ = mkcpListener.Close()
_ = packetConn.Close()
}
}()
go func() {
select {
case <-ctx.Done():
_ = l.Close()
case <-l.done:
}
}()
return l, nil
}
func (l *Listener) Accept() (net.Conn, error) {
return l.mkcp.Accept()
}
func (l *Listener) Close() error {
var err error
l.once.Do(func() {
err = errors.Join(l.server.Close(), l.mkcp.Close(), l.packetConn.Close(), l.outer.Close())
<-l.done
})
return err
}
func (l *Listener) Addr() net.Addr {
return l.outer.Addr()
}
type server struct {
ctx context.Context
cfg Config
packetConn *wrappedPacketConn
sessions sync.Map
}
func newServer(ctx context.Context, cfg Config, packetConn *wrappedPacketConn) *server {
return &server{ctx: ctx, cfg: cfg, packetConn: packetConn}
}
func (s *server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
sessionID, err := base64.RawURLEncoding.DecodeString(r.Header.Get("X-Session-ID"))
if err != nil {
http.Error(w, "invalid session id", http.StatusBadRequest)
return
}
body, err := io.ReadAll(r.Body)
_ = r.Body.Close()
if err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
session, err := s.getSession(r.Context(), sessionID, parseRemoteAddr(r.RemoteAddr))
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
if err := session.ingest(body); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
w.WriteHeader(http.StatusOK)
session.writeResponse(r.Context(), w)
}
func (s *server) getSession(_ context.Context, sessionID []byte, remoteAddr *net.TCPAddr) (*serverSession, error) {
key := string(sessionID)
if session, ok := s.sessions.Load(key); ok {
return session.(*serverSession), nil
}
sessionCtx, cancel := context.WithCancel(s.ctx)
session := &serverSession{
ctx: sessionCtx,
cancel: cancel,
sessionID: append([]byte(nil), sessionID...),
remoteAddr: remoteAddr,
server: s,
writerChan: make(chan []byte, s.cfg.PacketWritingBuffer),
readerChan: make(chan []byte, 256),
maxWriteSize: s.cfg.MaxWriteSize,
maxWriteDuration: time.Duration(s.cfg.MaxWriteDurationMs) * time.Millisecond,
maxSimultaneousWriteConnection: s.cfg.MaxSimultaneousWriteConnection,
}
actual, loaded := s.sessions.LoadOrStore(key, session)
if loaded {
cancel()
return actual.(*serverSession), nil
}
if err := s.packetConn.addSession(session); err != nil {
cancel()
s.sessions.Delete(key)
return nil, err
}
return session, nil
}
func (s *server) removeSession(sessionID []byte) {
s.sessions.Delete(string(sessionID))
}
func parseRemoteAddr(addr string) *net.TCPAddr {
if tcpAddr, err := net.ResolveTCPAddr("tcp", addr); err == nil {
return tcpAddr
}
return &net.TCPAddr{}
}
type serverSession struct {
ctx context.Context
cancel context.CancelFunc
sessionID []byte
remoteAddr *net.TCPAddr
server *server
writerChan chan []byte
readerChan chan []byte
maxWriteSize int
maxWriteDuration time.Duration
maxSimultaneousWriteConnection int
writingMu sync.Mutex
writingConns []*writingConnection
}
func (s *serverSession) ingest(body []byte) error {
reader := bytes.NewReader(body)
for reader.Len() > 0 {
packet, err := readPacketBundle(reader)
if err != nil {
return err
}
select {
case <-s.ctx.Done():
return s.ctx.Err()
case s.readerChan <- packet:
}
}
return nil
}
type writingConnection struct {
ctx context.Context
cancel context.CancelFunc
}
func (s *serverSession) beginWritingConnection(ctx context.Context) (*writingConnection, context.Context) {
writeCtx, cancel := context.WithCancel(ctx)
conn := &writingConnection{ctx: writeCtx, cancel: cancel}
var stale []*writingConnection
s.writingMu.Lock()
s.writingConns = append(s.writingConns, conn)
if s.maxSimultaneousWriteConnection > 0 {
for len(s.writingConns) > s.maxSimultaneousWriteConnection {
old := s.writingConns[0]
s.writingConns[0] = nil
s.writingConns = s.writingConns[1:]
stale = append(stale, old)
}
}
s.writingMu.Unlock()
for _, old := range stale {
old.cancel()
}
return conn, writeCtx
}
func (s *serverSession) finishWritingConnection(conn *writingConnection) {
conn.cancel()
s.writingMu.Lock()
for i, item := range s.writingConns {
if item == conn {
copy(s.writingConns[i:], s.writingConns[i+1:])
s.writingConns[len(s.writingConns)-1] = nil
s.writingConns = s.writingConns[:len(s.writingConns)-1]
break
}
}
s.writingMu.Unlock()
}
func (s *serverSession) writeResponse(ctx context.Context, w http.ResponseWriter) {
writeConn, writeCtx := s.beginWritingConnection(ctx)
defer s.finishWritingConnection(writeConn)
flusher, _ := w.(http.Flusher)
timer := time.NewTimer(s.maxWriteDuration)
defer timer.Stop()
bytesSent := 0
for {
select {
case <-writeCtx.Done():
return
case <-s.ctx.Done():
return
case packet := <-s.writerChan:
if err := writePacketBundle(w, packet); err != nil {
return
}
if flusher != nil {
flusher.Flush()
}
bytesSent += packetBundleOverhead + len(packet)
if s.maxWriteSize > 0 && bytesSent >= s.maxWriteSize {
return
}
case <-timer.C:
return
}
}
}
func (s *serverSession) Read(p []byte) (int, error) {
select {
case <-s.ctx.Done():
return 0, s.ctx.Err()
case packet := <-s.readerChan:
return copy(p, packet), nil
}
}
func (s *serverSession) Write(p []byte) (int, error) {
packet := append([]byte(nil), p...)
select {
case <-s.ctx.Done():
return 0, s.ctx.Err()
case s.writerChan <- packet:
return len(p), nil
default:
return len(p), nil
}
}
func (s *serverSession) Close() error {
s.server.removeSession(s.sessionID)
s.cancel()
return nil
}
func (s *serverSession) Network() string {
if s.remoteAddr == nil {
return ""
}
return s.remoteAddr.Network()
}
func (s *serverSession) String() string {
if s.remoteAddr == nil {
return ""
}
return s.remoteAddr.String()
}
var _ net.Listener = (*Listener)(nil)
var _ net.Addr = (*serverSession)(nil)
var _ io.ReadWriteCloser = (*serverSession)(nil)