netns: fix network namespaces for listeners
This commit is contained in:
@ -2,6 +2,8 @@ package dns
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"io"
|
||||
@ -9,12 +11,12 @@ import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
admission "github.com/go-gost/x/admission/wrapper"
|
||||
limiter "github.com/go-gost/x/limiter/traffic/wrapper"
|
||||
|
||||
"github.com/go-gost/core/listener"
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
admission "github.com/go-gost/x/admission/wrapper"
|
||||
xnet "github.com/go-gost/x/internal/net"
|
||||
limiter "github.com/go-gost/x/limiter/traffic/wrapper"
|
||||
metrics "github.com/go-gost/x/metrics/wrapper"
|
||||
stats "github.com/go-gost/x/observer/stats/wrapper"
|
||||
"github.com/go-gost/x/registry"
|
||||
@ -51,48 +53,144 @@ func (l *dnsListener) Init(md md.Metadata) (err error) {
|
||||
return
|
||||
}
|
||||
|
||||
l.addr, err = net.ResolveTCPAddr("tcp", l.options.Addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
switch strings.ToLower(l.md.mode) {
|
||||
case "tcp":
|
||||
l.server = &dns.Server{
|
||||
Net: "tcp",
|
||||
Addr: l.options.Addr,
|
||||
Handler: l,
|
||||
ReadTimeout: l.md.readTimeout,
|
||||
WriteTimeout: l.md.writeTimeout,
|
||||
l.addr, err = net.ResolveTCPAddr("tcp", l.options.Addr)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
network := "tcp"
|
||||
if xnet.IsIPv4(l.options.Addr) {
|
||||
network = "tcp4"
|
||||
}
|
||||
|
||||
lc := net.ListenConfig{}
|
||||
if l.md.mptcp {
|
||||
lc.SetMultipathTCP(true)
|
||||
l.logger.Debugf("mptcp enabled: %v", lc.MultipathTCP())
|
||||
}
|
||||
|
||||
var ln net.Listener
|
||||
ln, err = lc.Listen(context.Background(), network, l.options.Addr)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
l.server = &dnsServer{
|
||||
server: &dns.Server{
|
||||
Net: "tcp",
|
||||
Addr: l.options.Addr,
|
||||
Listener: ln,
|
||||
Handler: l,
|
||||
ReadTimeout: l.md.readTimeout,
|
||||
WriteTimeout: l.md.writeTimeout,
|
||||
},
|
||||
}
|
||||
|
||||
case "tls":
|
||||
l.server = &dns.Server{
|
||||
Net: "tcp-tls",
|
||||
Addr: l.options.Addr,
|
||||
Handler: l,
|
||||
TLSConfig: l.options.TLSConfig,
|
||||
ReadTimeout: l.md.readTimeout,
|
||||
WriteTimeout: l.md.writeTimeout,
|
||||
l.addr, err = net.ResolveTCPAddr("tcp", l.options.Addr)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
network := "tcp"
|
||||
if xnet.IsIPv4(l.options.Addr) {
|
||||
network = "tcp4"
|
||||
}
|
||||
|
||||
lc := net.ListenConfig{}
|
||||
if l.md.mptcp {
|
||||
lc.SetMultipathTCP(true)
|
||||
l.logger.Debugf("mptcp enabled: %v", lc.MultipathTCP())
|
||||
}
|
||||
|
||||
var ln net.Listener
|
||||
ln, err = lc.Listen(context.Background(), network, l.options.Addr)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
ln = tls.NewListener(ln, l.options.TLSConfig)
|
||||
|
||||
l.server = &dnsServer{
|
||||
server: &dns.Server{
|
||||
Net: "tcp-tls",
|
||||
Addr: l.options.Addr,
|
||||
Listener: ln,
|
||||
Handler: l,
|
||||
TLSConfig: l.options.TLSConfig,
|
||||
ReadTimeout: l.md.readTimeout,
|
||||
WriteTimeout: l.md.writeTimeout,
|
||||
},
|
||||
}
|
||||
|
||||
case "https":
|
||||
l.addr, err = net.ResolveTCPAddr("tcp", l.options.Addr)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
network := "tcp"
|
||||
if xnet.IsIPv4(l.options.Addr) {
|
||||
network = "tcp4"
|
||||
}
|
||||
|
||||
lc := net.ListenConfig{}
|
||||
if l.md.mptcp {
|
||||
lc.SetMultipathTCP(true)
|
||||
l.logger.Debugf("mptcp enabled: %v", lc.MultipathTCP())
|
||||
}
|
||||
|
||||
var ln net.Listener
|
||||
ln, err = lc.Listen(context.Background(), network, l.options.Addr)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
ln = tls.NewListener(ln, l.options.TLSConfig)
|
||||
|
||||
l.server = &dohServer{
|
||||
addr: l.options.Addr,
|
||||
tlsConfig: l.options.TLSConfig,
|
||||
listener: ln,
|
||||
server: &http.Server{
|
||||
Handler: l,
|
||||
ReadTimeout: l.md.readTimeout,
|
||||
WriteTimeout: l.md.writeTimeout,
|
||||
},
|
||||
}
|
||||
|
||||
default:
|
||||
l.addr, err = net.ResolveUDPAddr("udp", l.options.Addr)
|
||||
l.server = &dns.Server{
|
||||
Net: "udp",
|
||||
Addr: l.options.Addr,
|
||||
Handler: l,
|
||||
UDPSize: l.md.readBufferSize,
|
||||
ReadTimeout: l.md.readTimeout,
|
||||
WriteTimeout: l.md.writeTimeout,
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
network := "udp"
|
||||
if xnet.IsIPv4(l.options.Addr) {
|
||||
network = "udp4"
|
||||
}
|
||||
|
||||
lc := net.ListenConfig{}
|
||||
if l.md.mptcp {
|
||||
lc.SetMultipathTCP(true)
|
||||
l.logger.Debugf("mptcp enabled: %v", lc.MultipathTCP())
|
||||
}
|
||||
|
||||
var pc net.PacketConn
|
||||
pc, err = lc.ListenPacket(context.Background(), network, l.options.Addr)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
l.server = &dnsServer{
|
||||
server: &dns.Server{
|
||||
Net: "udp",
|
||||
Addr: l.options.Addr,
|
||||
PacketConn: pc,
|
||||
Handler: l,
|
||||
UDPSize: l.md.readBufferSize,
|
||||
ReadTimeout: l.md.readTimeout,
|
||||
WriteTimeout: l.md.writeTimeout,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@ -104,7 +202,7 @@ func (l *dnsListener) Init(md md.Metadata) (err error) {
|
||||
l.errChan = make(chan error, 1)
|
||||
|
||||
go func() {
|
||||
err := l.server.ListenAndServe()
|
||||
err := l.server.Serve()
|
||||
if err != nil {
|
||||
l.errChan <- err
|
||||
}
|
||||
|
@ -17,6 +17,7 @@ type metadata struct {
|
||||
readTimeout time.Duration
|
||||
writeTimeout time.Duration
|
||||
backlog int
|
||||
mptcp bool
|
||||
}
|
||||
|
||||
func (l *dnsListener) parseMetadata(md mdata.Metadata) (err error) {
|
||||
@ -37,6 +38,7 @@ func (l *dnsListener) parseMetadata(md mdata.Metadata) (err error) {
|
||||
if l.md.backlog <= 0 {
|
||||
l.md.backlog = defaultBacklog
|
||||
}
|
||||
l.md.mptcp = mdutil.GetBool(md, "mptcp")
|
||||
|
||||
return
|
||||
}
|
||||
|
@ -10,29 +10,47 @@ import (
|
||||
"time"
|
||||
|
||||
xnet "github.com/go-gost/x/internal/net"
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
type Server interface {
|
||||
ListenAndServe() error
|
||||
Serve() error
|
||||
Shutdown() error
|
||||
}
|
||||
|
||||
type dnsServer struct {
|
||||
server *dns.Server
|
||||
}
|
||||
|
||||
func (s *dnsServer) Serve() error {
|
||||
return s.server.ActivateAndServe()
|
||||
}
|
||||
|
||||
func (s *dnsServer) Shutdown() error {
|
||||
return s.server.Shutdown()
|
||||
}
|
||||
|
||||
type dohServer struct {
|
||||
addr string
|
||||
tlsConfig *tls.Config
|
||||
listener net.Listener
|
||||
server *http.Server
|
||||
}
|
||||
|
||||
func (s *dohServer) ListenAndServe() error {
|
||||
network := "tcp"
|
||||
if xnet.IsIPv4(s.addr) {
|
||||
network = "tcp4"
|
||||
func (s *dohServer) Serve() error {
|
||||
var err error
|
||||
ln := s.listener
|
||||
if ln == nil {
|
||||
network := "tcp"
|
||||
if xnet.IsIPv4(s.addr) {
|
||||
network = "tcp4"
|
||||
}
|
||||
ln, err = net.Listen(network, s.addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ln = tls.NewListener(ln, s.tlsConfig)
|
||||
}
|
||||
ln, err := net.Listen(network, s.addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ln = tls.NewListener(ln, s.tlsConfig)
|
||||
return s.server.Serve(ln)
|
||||
}
|
||||
|
||||
|
Reference in New Issue
Block a user