diff --git a/adapter/outbound/vmess.go b/adapter/outbound/vmess.go index e071aa5d..ecac585c 100644 --- a/adapter/outbound/vmess.go +++ b/adapter/outbound/vmess.go @@ -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) diff --git a/docs/config.yaml b/docs/config.yaml index be5f0a4f..5e038280 100644 --- a/docs/config.yaml +++ b/docs/config.yaml @@ -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 diff --git a/listener/config/mekya.go b/listener/config/mekya.go new file mode 100644 index 00000000..613edd10 --- /dev/null +++ b/listener/config/mekya.go @@ -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, + } +} diff --git a/listener/config/vmess.go b/listener/config/vmess.go index f786e2e0..13a02eef 100644 --- a/listener/config/vmess.go +++ b/listener/config/vmess.go @@ -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"` } diff --git a/listener/inbound/mekya.go b/listener/inbound/mekya.go new file mode 100644 index 00000000..bef1f816 --- /dev/null +++ b/listener/inbound/mekya.go @@ -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(), + } +} diff --git a/listener/inbound/vmess.go b/listener/inbound/vmess.go index 8337b4c6..5eb52e0d 100644 --- a/listener/inbound/vmess.go +++ b/listener/inbound/vmess.go @@ -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(), }, diff --git a/listener/inbound/vmess_mekya_interop_test.go b/listener/inbound/vmess_mekya_interop_test.go new file mode 100644 index 00000000..59bc63dc --- /dev/null +++ b/listener/inbound/vmess_mekya_interop_test.go @@ -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, + } +} diff --git a/listener/inbound/vmess_test.go b/listener/inbound/vmess_test.go index c847bea3..eee41095 100644 --- a/listener/inbound/vmess_test.go +++ b/listener/inbound/vmess_test.go @@ -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{ diff --git a/listener/sing_vmess/server.go b/listener/sing_vmess/server.go index 54b146c4..19859bd1 100644 --- a/listener/sing_vmess/server.go +++ b/listener/sing_vmess/server.go @@ -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() { diff --git a/transport/mekya/bundle.go b/transport/mekya/bundle.go new file mode 100644 index 00000000..44814a99 --- /dev/null +++ b/transport/mekya/bundle.go @@ -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 +} diff --git a/transport/mekya/client.go b/transport/mekya/client.go new file mode 100644 index 00000000..530ce1cc --- /dev/null +++ b/transport/mekya/client.go @@ -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 +} diff --git a/transport/mekya/config.go b/transport/mekya/config.go new file mode 100644 index 00000000..8cc58a3f --- /dev/null +++ b/transport/mekya/config.go @@ -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 +} diff --git a/transport/mekya/deadline.go b/transport/mekya/deadline.go new file mode 100644 index 00000000..f97385ad --- /dev/null +++ b/transport/mekya/deadline.go @@ -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 +} diff --git a/transport/mekya/mekya_test.go b/transport/mekya/mekya_test.go new file mode 100644 index 00000000..fa9d4cec --- /dev/null +++ b/transport/mekya/mekya_test.go @@ -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 +} diff --git a/transport/mekya/packetconn.go b/transport/mekya/packetconn.go new file mode 100644 index 00000000..3f8fce2b --- /dev/null +++ b/transport/mekya/packetconn.go @@ -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) diff --git a/transport/mekya/server.go b/transport/mekya/server.go new file mode 100644 index 00000000..c1cddcd4 --- /dev/null +++ b/transport/mekya/server.go @@ -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)